diff --git a/Cargo.lock b/Cargo.lock index 831ae8b..0cf0890 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -168,6 +168,7 @@ version = "0.1.0" dependencies = [ "axum", "dotenvy", + "futures", "jsonwebtoken", "once_cell", "reqwest", @@ -175,6 +176,7 @@ dependencies = [ "serde_json", "thiserror 2.0.18", "tokio", + "tokio-stream", "wiremock", ] @@ -1156,6 +1158,7 @@ dependencies = [ "bytes", "encoding_rs", "futures-core", + "futures-util", "h2", "http", "http-body", @@ -1177,12 +1180,14 @@ dependencies = [ "sync_wrapper", "tokio", "tokio-rustls", + "tokio-util", "tower", "tower-http", "tower-service", "url", "wasm-bindgen", "wasm-bindgen-futures", + "wasm-streams", "web-sys", ] @@ -1663,6 +1668,17 @@ dependencies = [ "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]] name = "tokio-util" version = "0.7.18" @@ -1874,6 +1890,19 @@ dependencies = [ "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]] name = "web-sys" version = "0.3.94" diff --git a/Cargo.toml b/Cargo.toml index 8cc7e0f..e301768 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,7 +13,9 @@ tokio = { version = "1", features = ["full"] } serde = { version = "1", features = ["derive"] } serde_json = "1" 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" dotenvy = "0.15" thiserror = "2.0.18" +tokio-stream = "0.1" +futures = "0.3" \ No newline at end of file diff --git a/src/providers/ollama.rs b/src/providers/ollama.rs index 5c7cffe..cbd5f44 100644 --- a/src/providers/ollama.rs +++ b/src/providers/ollama.rs @@ -1,6 +1,9 @@ use crate::errors::OllamaError; +use axum::response::sse::Event; +use futures::StreamExt; use reqwest::Client; use serde_json::{Value, json}; +use tokio_stream::wrappers::ReceiverStream; #[derive(Clone)] 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), 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 ───────────────────────────────────────────────────── pub async fn list_models(&self) -> Result { @@ -133,19 +229,10 @@ impl OllamaProvider { } pub async fn completions(&self, body: Value) -> Result { - 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"); + let (prompt, model) = self.extract_completion_params(&body)?; self.validate_model(model).await?; - let ollama_payload = json!({ + let payload = json!({ "model": model, "prompt": prompt, "stream": false, @@ -155,56 +242,80 @@ impl OllamaProvider { let res = self .client .post(format!("{}/api/generate", self.base_url)) - .json(&ollama_payload) + .json(&payload) .send() .await? .json::() .await?; - Ok(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, + Ok(self.format_completion_response(&res)) + } + + pub async fn completions_stream( + &self, + body: Value, + ) -> Result>, 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::(&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", + "object": "text_completion", + "choices": [{ "text": token, "index": 0, "finish_reason": 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 { - 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); - } - + let (model, messages) = self.extract_chat_params(&body)?; self.validate_model(model).await?; - let ollama_payload = json!({ + let payload = json!({ "model": model, "messages": messages, "stream": false, @@ -214,35 +325,76 @@ impl OllamaProvider { let res = self .client .post(format!("{}/api/chat", self.base_url)) - .json(&ollama_payload) + .json(&payload) .send() .await? .json::() .await?; - Ok(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), + Ok(self.format_chat_response(&res)) + } + + pub async fn chat_completions_stream( + &self, + body: Value, + ) -> Result>, 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::(&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", + "object": "chat.completion.chunk", + "choices": [{ + "index": 0, + "delta": { + "role": json.get("message").and_then(|m| m.get("role")), + "content": json.get("message").and_then(|m| m.get("content")), + }, + "finish_reason": if done { json!("stop") } else { json!(null) }, + }], + })) + .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)) } } diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index ac3a292..c3108af 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -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 crate::errors::OllamaError; @@ -7,38 +14,77 @@ use crate::state::app_state::AppState; pub async fn completions( State(state): State, Json(body): Json, -) -> Result, (axum::http::StatusCode, String)> { - match state.ollama.completions(body).await { - Ok(response) => Ok(Json(response)), - Err(OllamaError::MissingPrompt) => Err(( - axum::http::StatusCode::BAD_REQUEST, - "prompt is required and cannot be empty".to_string(), - )), - Err(OllamaError::ModelNotFound(m)) => Err(( - axum::http::StatusCode::UNPROCESSABLE_ENTITY, - format!("model '{m}' is not available — run `ollama pull {m}` first"), - )), - Err(e) => Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())), +) -> Result { + let wants_stream = body + .get("stream") + .and_then(|v| v.as_bool()) + .unwrap_or(false); + + if wants_stream { + let stream = state + .ollama + .completions_stream(body) + .await + .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( State(state): State, Json(body): Json, -) -> Result, (axum::http::StatusCode, String)> { - match state.ollama.chat_completions(body).await { - Ok(response) => Ok(Json(response)), - Err(OllamaError::MissingMessages) => Err(( +) -> Result { + let wants_stream = body + .get("stream") + .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, - "messages must be a non-empty array with at least one user message".to_string(), - )), - Err(OllamaError::ModelNotFound(m)) => Err(( + "prompt is required and cannot be empty".to_string(), + ), + OllamaError::ModelNotFound(m) => ( axum::http::StatusCode::UNPROCESSABLE_ENTITY, format!("model '{m}' is not available — run `ollama pull {m}` first"), - )), - Err(OllamaError::Http(e)) => { - Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())) - } - Err(e) => Err((axum::http::StatusCode::BAD_REQUEST, e.to_string())), + ), + OllamaError::InvalidKeepAlive(v) => ( + axum::http::StatusCode::BAD_REQUEST, + format!("invalid keep_alive '{v}'"), + ), + 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()), } }