feat: add streaming
CI / Rust CI (push) Successful in 4m29s

This commit is contained in:
2026-04-10 14:15:27 +02:00
parent 686f9ff747
commit 1a99490e22
4 changed files with 329 additions and 100 deletions
Generated
+29
View File
@@ -168,6 +168,7 @@ version = "0.1.0"
dependencies = [ dependencies = [
"axum", "axum",
"dotenvy", "dotenvy",
"futures",
"jsonwebtoken", "jsonwebtoken",
"once_cell", "once_cell",
"reqwest", "reqwest",
@@ -175,6 +176,7 @@ dependencies = [
"serde_json", "serde_json",
"thiserror 2.0.18", "thiserror 2.0.18",
"tokio", "tokio",
"tokio-stream",
"wiremock", "wiremock",
] ]
@@ -1156,6 +1158,7 @@ dependencies = [
"bytes", "bytes",
"encoding_rs", "encoding_rs",
"futures-core", "futures-core",
"futures-util",
"h2", "h2",
"http", "http",
"http-body", "http-body",
@@ -1177,12 +1180,14 @@ dependencies = [
"sync_wrapper", "sync_wrapper",
"tokio", "tokio",
"tokio-rustls", "tokio-rustls",
"tokio-util",
"tower", "tower",
"tower-http", "tower-http",
"tower-service", "tower-service",
"url", "url",
"wasm-bindgen", "wasm-bindgen",
"wasm-bindgen-futures", "wasm-bindgen-futures",
"wasm-streams",
"web-sys", "web-sys",
] ]
@@ -1663,6 +1668,17 @@ dependencies = [
"tokio", "tokio",
] ]
[[package]]
name = "tokio-stream"
version = "0.1.18"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32da49809aab5c3bc678af03902d4ccddea2a87d028d86392a4b1560c6906c70"
dependencies = [
"futures-core",
"pin-project-lite",
"tokio",
]
[[package]] [[package]]
name = "tokio-util" name = "tokio-util"
version = "0.7.18" version = "0.7.18"
@@ -1874,6 +1890,19 @@ dependencies = [
"unicode-ident", "unicode-ident",
] ]
[[package]]
name = "wasm-streams"
version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9d1ec4f6517c9e11ae630e200b2b65d193279042e28edd4a2cda233e46670bbb"
dependencies = [
"futures-util",
"js-sys",
"wasm-bindgen",
"wasm-bindgen-futures",
"web-sys",
]
[[package]] [[package]]
name = "web-sys" name = "web-sys"
version = "0.3.94" version = "0.3.94"
+3 -1
View File
@@ -13,7 +13,9 @@ tokio = { version = "1", features = ["full"] }
serde = { version = "1", features = ["derive"] } serde = { version = "1", features = ["derive"] }
serde_json = "1" serde_json = "1"
jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] } jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] }
reqwest = { version = "0.13.2", features = ["json"] } reqwest = { version = "0.13.2", features = ["json", "stream"] }
once_cell = "1" once_cell = "1"
dotenvy = "0.15" dotenvy = "0.15"
thiserror = "2.0.18" thiserror = "2.0.18"
tokio-stream = "0.1"
futures = "0.3"
+219 -67
View File
@@ -1,6 +1,9 @@
use crate::errors::OllamaError; use crate::errors::OllamaError;
use axum::response::sse::Event;
use futures::StreamExt;
use reqwest::Client; use reqwest::Client;
use serde_json::{Value, json}; use serde_json::{Value, json};
use tokio_stream::wrappers::ReceiverStream;
#[derive(Clone)] #[derive(Clone)]
pub struct OllamaProvider { pub struct OllamaProvider {
@@ -66,6 +69,99 @@ impl OllamaProvider {
} }
} }
fn extract_completion_params<'a>(
&self,
body: &'a Value,
) -> Result<(&'a str, &'a str), OllamaError> {
let prompt = body
.get("prompt")
.and_then(|v| v.as_str())
.filter(|s| !s.trim().is_empty())
.ok_or(OllamaError::MissingPrompt)?;
let model = body
.get("model")
.and_then(|v| v.as_str())
.unwrap_or("llama3");
Ok((prompt, model))
}
fn format_completion_response(&self, res: &Value) -> Value {
json!({
"id": "cmpl-ollama",
"object": "text_completion",
"model": res.get("model"),
"choices": [{
"text": res.get("response"),
"index": 0,
"finish_reason": if res.get("done").and_then(|v| v.as_bool()).unwrap_or(false) {
"stop"
} else {
"length"
},
}],
"usage": {
"prompt_tokens": res.get("prompt_eval_count"),
"completion_tokens": res.get("eval_count"),
"total_tokens": null,
}
})
}
fn extract_chat_params<'a>(
&self,
body: &'a Value,
) -> Result<(&'a str, &'a Vec<Value>), OllamaError> {
let model = body
.get("model")
.and_then(|v| v.as_str())
.unwrap_or("llama3");
let messages = body
.get("messages")
.and_then(|v| v.as_array())
.filter(|arr| !arr.is_empty())
.ok_or(OllamaError::MissingMessages)?;
let has_user_msg = messages
.iter()
.any(|m| m.get("role").and_then(|r| r.as_str()) == Some("user"));
if !has_user_msg {
return Err(OllamaError::MissingMessages);
}
Ok((model, messages))
}
fn format_chat_response(&self, res: &Value) -> Value {
json!({
"id": "chatcmpl-ollama",
"object": "chat.completion",
"model": res.get("model"),
"choices": [{
"index": 0,
"message": {
"role": res.get("message").and_then(|m| m.get("role")),
"content": res.get("message").and_then(|m| m.get("content")),
},
"finish_reason": res
.get("done_reason")
.and_then(|v| v.as_str())
.unwrap_or("stop"),
}],
"usage": {
"prompt_tokens": res.get("prompt_eval_count"),
"completion_tokens": res.get("eval_count"),
"total_tokens": res.get("prompt_eval_count")
.and_then(|p| p.as_u64())
.zip(res.get("eval_count").and_then(|e| e.as_u64()))
.map(|(p, e)| p + e),
}
})
}
// ── public endpoints ───────────────────────────────────────────────────── // ── public endpoints ─────────────────────────────────────────────────────
pub async fn list_models(&self) -> Result<Value, OllamaError> { pub async fn list_models(&self) -> Result<Value, OllamaError> {
@@ -133,19 +229,10 @@ impl OllamaProvider {
} }
pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> { pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
let prompt = body let (prompt, model) = self.extract_completion_params(&body)?;
.get("prompt")
.and_then(|v| v.as_str())
.filter(|s| !s.trim().is_empty())
.ok_or(OllamaError::MissingPrompt)?;
let model = body
.get("model")
.and_then(|v| v.as_str())
.unwrap_or("llama3");
self.validate_model(model).await?; self.validate_model(model).await?;
let ollama_payload = json!({ let payload = json!({
"model": model, "model": model,
"prompt": prompt, "prompt": prompt,
"stream": false, "stream": false,
@@ -155,56 +242,80 @@ impl OllamaProvider {
let res = self let res = self
.client .client
.post(format!("{}/api/generate", self.base_url)) .post(format!("{}/api/generate", self.base_url))
.json(&ollama_payload) .json(&payload)
.send() .send()
.await? .await?
.json::<Value>() .json::<Value>()
.await?; .await?;
Ok(json!({ Ok(self.format_completion_response(&res))
}
pub async fn completions_stream(
&self,
body: Value,
) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
let (prompt, model) = self.extract_completion_params(&body)?;
self.validate_model(model).await?;
let payload = json!({
"model": model,
"prompt": prompt,
"stream": true,
"options": Self::build_options(&body),
});
let mut byte_stream = self
.client
.post(format!("{}/api/generate", self.base_url))
.json(&payload)
.send()
.await?
.bytes_stream();
let (tx, rx) = tokio::sync::mpsc::channel(32);
tokio::spawn(async move {
while let Some(chunk) = byte_stream.next().await {
let chunk = match chunk {
Ok(b) => b,
Err(e) => {
let _ = tx.send(Err(OllamaError::Http(e))).await;
break;
}
};
if let Ok(json) = serde_json::from_slice::<Value>(&chunk) {
let token = json.get("response").and_then(|v| v.as_str()).unwrap_or("");
let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
// OpenAI-compatible SSE chunk
let event_data = serde_json::to_string(&json!({
"id": "cmpl-ollama", "id": "cmpl-ollama",
"object": "text_completion", "object": "text_completion",
"model": res.get("model"), "choices": [{ "text": token, "index": 0, "finish_reason": null }],
"choices": [{
"text": res.get("response"),
"index": 0,
"finish_reason": if res.get("done").and_then(|v| v.as_bool()).unwrap_or(false) {
"stop"
} else {
"length"
},
}],
"usage": {
"prompt_tokens": res.get("prompt_eval_count"),
"completion_tokens": res.get("eval_count"),
"total_tokens": null,
}
})) }))
.unwrap_or_default();
let _ = tx.send(Ok(Event::default().data(event_data))).await;
if done {
// Final [DONE] sentinel — matches OpenAI streaming protocol
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
break;
}
}
}
});
Ok(ReceiverStream::new(rx))
} }
pub async fn chat_completions(&self, body: Value) -> Result<Value, OllamaError> { pub async fn chat_completions(&self, body: Value) -> Result<Value, OllamaError> {
let model = body let (model, messages) = self.extract_chat_params(&body)?;
.get("model")
.and_then(|v| v.as_str())
.unwrap_or("llama3");
let messages = body
.get("messages")
.and_then(|v| v.as_array())
.filter(|arr| !arr.is_empty())
.ok_or(OllamaError::MissingMessages)?;
let has_user_msg = messages
.iter()
.any(|m| m.get("role").and_then(|r| r.as_str()) == Some("user"));
if !has_user_msg {
return Err(OllamaError::MissingMessages);
}
self.validate_model(model).await?; self.validate_model(model).await?;
let ollama_payload = json!({ let payload = json!({
"model": model, "model": model,
"messages": messages, "messages": messages,
"stream": false, "stream": false,
@@ -214,35 +325,76 @@ impl OllamaProvider {
let res = self let res = self
.client .client
.post(format!("{}/api/chat", self.base_url)) .post(format!("{}/api/chat", self.base_url))
.json(&ollama_payload) .json(&payload)
.send() .send()
.await? .await?
.json::<Value>() .json::<Value>()
.await?; .await?;
Ok(json!({ Ok(self.format_chat_response(&res))
}
pub async fn chat_completions_stream(
&self,
body: Value,
) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
let (model, messages) = self.extract_chat_params(&body)?;
self.validate_model(model).await?;
let payload = json!({
"model": model,
"messages": messages,
"stream": true,
"options": Self::build_options(&body),
});
let mut byte_stream = self
.client
.post(format!("{}/api/chat", self.base_url))
.json(&payload)
.send()
.await?
.bytes_stream();
let (tx, rx) = tokio::sync::mpsc::channel(32);
tokio::spawn(async move {
while let Some(chunk) = byte_stream.next().await {
let chunk = match chunk {
Ok(b) => b,
Err(e) => {
let _ = tx.send(Err(OllamaError::Http(e))).await;
break;
}
};
if let Ok(json) = serde_json::from_slice::<Value>(&chunk) {
let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
let event_data = serde_json::to_string(&json!({
"id": "chatcmpl-ollama", "id": "chatcmpl-ollama",
"object": "chat.completion", "object": "chat.completion.chunk",
"model": res.get("model"),
"choices": [{ "choices": [{
"index": 0, "index": 0,
"message": { "delta": {
"role": res.get("message").and_then(|m| m.get("role")), "role": json.get("message").and_then(|m| m.get("role")),
"content": res.get("message").and_then(|m| m.get("content")), "content": json.get("message").and_then(|m| m.get("content")),
}, },
"finish_reason": res "finish_reason": if done { json!("stop") } else { json!(null) },
.get("done_reason")
.and_then(|v| v.as_str())
.unwrap_or("stop"),
}], }],
"usage": {
"prompt_tokens": res.get("prompt_eval_count"),
"completion_tokens": res.get("eval_count"),
"total_tokens": res.get("prompt_eval_count")
.and_then(|p| p.as_u64())
.zip(res.get("eval_count").and_then(|e| e.as_u64()))
.map(|(p, e)| p + e),
}
})) }))
.unwrap_or_default();
let _ = tx.send(Ok(Event::default().data(event_data))).await;
if done {
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
break;
}
}
}
});
Ok(ReceiverStream::new(rx))
} }
} }
+71 -25
View File
@@ -1,4 +1,11 @@
use axum::{Json, extract::State}; use axum::{
Json,
extract::State,
response::{
IntoResponse, Response,
sse::{KeepAlive, Sse},
},
};
use serde_json::Value; use serde_json::Value;
use crate::errors::OllamaError; use crate::errors::OllamaError;
@@ -7,38 +14,77 @@ use crate::state::app_state::AppState;
pub async fn completions( pub async fn completions(
State(state): State<AppState>, State(state): State<AppState>,
Json(body): Json<Value>, Json(body): Json<Value>,
) -> Result<Json<Value>, (axum::http::StatusCode, String)> { ) -> Result<Response, (axum::http::StatusCode, String)> {
match state.ollama.completions(body).await { let wants_stream = body
Ok(response) => Ok(Json(response)), .get("stream")
Err(OllamaError::MissingPrompt) => Err(( .and_then(|v| v.as_bool())
axum::http::StatusCode::BAD_REQUEST, .unwrap_or(false);
"prompt is required and cannot be empty".to_string(),
)), if wants_stream {
Err(OllamaError::ModelNotFound(m)) => Err(( let stream = state
axum::http::StatusCode::UNPROCESSABLE_ENTITY, .ollama
format!("model '{m}' is not available — run `ollama pull {m}` first"), .completions_stream(body)
)), .await
Err(e) => Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())), .map_err(ollama_err)?;
Ok(Sse::new(stream)
.keep_alive(KeepAlive::default())
.into_response())
} else {
let response = state.ollama.completions(body).await.map_err(ollama_err)?;
Ok(Json(response).into_response())
} }
} }
pub async fn chat_completions( pub async fn chat_completions(
State(state): State<AppState>, State(state): State<AppState>,
Json(body): Json<Value>, Json(body): Json<Value>,
) -> Result<Json<Value>, (axum::http::StatusCode, String)> { ) -> Result<Response, (axum::http::StatusCode, String)> {
match state.ollama.chat_completions(body).await { let wants_stream = body
Ok(response) => Ok(Json(response)), .get("stream")
Err(OllamaError::MissingMessages) => Err(( .and_then(|v| v.as_bool())
.unwrap_or(false);
if wants_stream {
let stream = state
.ollama
.chat_completions_stream(body)
.await
.map_err(ollama_err)?;
Ok(Sse::new(stream)
.keep_alive(KeepAlive::default())
.into_response())
} else {
let response = state
.ollama
.chat_completions(body)
.await
.map_err(ollama_err)?;
Ok(Json(response).into_response())
}
}
fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
match e {
OllamaError::MissingPrompt => (
axum::http::StatusCode::BAD_REQUEST, axum::http::StatusCode::BAD_REQUEST,
"messages must be a non-empty array with at least one user message".to_string(), "prompt is required and cannot be empty".to_string(),
)), ),
Err(OllamaError::ModelNotFound(m)) => Err(( OllamaError::ModelNotFound(m) => (
axum::http::StatusCode::UNPROCESSABLE_ENTITY, axum::http::StatusCode::UNPROCESSABLE_ENTITY,
format!("model '{m}' is not available — run `ollama pull {m}` first"), format!("model '{m}' is not available — run `ollama pull {m}` first"),
)), ),
Err(OllamaError::Http(e)) => { OllamaError::InvalidKeepAlive(v) => (
Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())) axum::http::StatusCode::BAD_REQUEST,
} format!("invalid keep_alive '{v}'"),
Err(e) => Err((axum::http::StatusCode::BAD_REQUEST, e.to_string())), ),
OllamaError::MissingMessages => (
axum::http::StatusCode::BAD_REQUEST,
"messages array with at least one user message is required".to_string(),
),
OllamaError::Http(e) => (axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string()),
} }
} }