From e31cdf131f5dd504ebcff0188b23e6c2043b3a5e Mon Sep 17 00:00:00 2001 From: LucasX Ubuntu Date: Fri, 10 Apr 2026 21:33:52 +0200 Subject: [PATCH] feat: chat typing --- src/dto/api.rs | 80 +++++++--- src/dto/ollama.rs | 21 ++- src/providers/ollama/client.rs | 281 +++++++++++++++------------------ src/providers/ollama/mapper.rs | 58 +++++-- src/routes/v1/chat.rs | 54 +++---- src/routes/v1/mod.rs | 4 +- 6 files changed, 277 insertions(+), 221 deletions(-) diff --git a/src/dto/api.rs b/src/dto/api.rs index fdd8ca4..06b0bde 100644 --- a/src/dto/api.rs +++ b/src/dto/api.rs @@ -3,27 +3,6 @@ use utoipa::ToSchema; use crate::errors::OllamaError; -#[derive(Debug, Serialize, Deserialize, ToSchema)] -#[serde(rename_all = "lowercase")] -pub enum Role { - System, - User, - Assistant, -} - -#[derive(Serialize, Deserialize, ToSchema)] -pub struct ChatResponse { - pub response: String, -} - -#[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct Message { - pub role: Role, - pub content: String, -} - -// --------------------------- - #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct ModelsResponse { pub models: Vec, @@ -155,3 +134,62 @@ pub struct CompletionChunk { pub object: String, pub choices: Vec, } + +#[derive(Debug, Deserialize, Serialize, ToSchema)] +pub struct ChatRequest { + #[serde(flatten)] + pub base: BaseLLMRequest, + + pub messages: Vec, +} + +#[derive(Debug, Deserialize, Serialize, ToSchema)] +pub struct Message { + pub role: Role, + pub content: String, +} + +#[derive(Debug, Deserialize, Serialize, ToSchema)] +#[serde(rename_all = "lowercase")] +pub enum Role { + System, + User, + Assistant, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct ChatCompletionResponse { + pub id: String, + pub object: String, + pub created: u64, + pub model: String, + pub choices: Vec, + pub usage: Option, // optional (Ollama may not always provide) +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct ChatChoice { + pub index: u32, + pub message: Message, + pub finish_reason: FinishReason, +} + +#[derive(Debug, Serialize)] +pub struct ChatCompletionChunk { + pub id: String, + pub object: String, + pub choices: Vec, +} + +#[derive(Debug, Serialize)] +pub struct ChatChunkChoice { + pub index: u32, + pub delta: ChatDelta, + pub finish_reason: Option, +} + +#[derive(Debug, Serialize)] +pub struct ChatDelta { + pub role: Option, + pub content: Option, +} diff --git a/src/dto/ollama.rs b/src/dto/ollama.rs index ce870b5..6625748 100644 --- a/src/dto/ollama.rs +++ b/src/dto/ollama.rs @@ -1,5 +1,6 @@ use serde::{Deserialize, Serialize}; -use utoipa::ToSchema; + +use crate::dto::api; #[derive(Debug, Serialize, Deserialize)] pub struct OllamaModels { @@ -59,3 +60,21 @@ pub struct OllamaGenerateResponse { pub prompt_eval_count: Option, pub eval_count: Option, } + +#[derive(Debug, Serialize)] +pub struct OllamaChatRequest<'a> { + pub model: &'a str, + pub messages: &'a [api::Message], + pub stream: bool, + pub options: OllamaOptions, +} + +#[derive(Debug, Deserialize)] +pub struct OllamaChatResponse { + pub model: String, + pub message: api::Message, + pub done: bool, + + pub prompt_eval_count: Option, + pub eval_count: Option, +} diff --git a/src/providers/ollama/client.rs b/src/providers/ollama/client.rs index 6e0684c..a6db618 100644 --- a/src/providers/ollama/client.rs +++ b/src/providers/ollama/client.rs @@ -1,10 +1,9 @@ use crate::dto::{api, ollama}; use crate::errors::OllamaError; -use axum::Json; use axum::response::sse::Event; use futures::StreamExt; use reqwest::Client; -use serde_json::{Value, json}; +use serde_json::json; use tokio_stream::wrappers::ReceiverStream; #[derive(Clone)] @@ -48,80 +47,14 @@ impl OllamaProvider { 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 api::ChatRequest, + ) -> Result<(&'a [api::Message], &'a str), OllamaError> { + let model = body.base.model.as_str(); - // 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), - // } - // }) - // } + Ok((&body.messages, model)) + } // // ── public endpoints ───────────────────────────────────────────────────── @@ -166,7 +99,7 @@ impl OllamaProvider { .json(&payload) .send() .await? - .json::() + .json::() .await?; Ok(api::LoadModelResponse { @@ -197,7 +130,7 @@ impl OllamaProvider { .json(&payload) .send() .await? - .json::() + .json::() .await?; Ok(api::UnloadModelResponse { @@ -227,7 +160,7 @@ impl OllamaProvider { return Err(OllamaError::ModelNotFound(model.to_string())); } - let options = ollama::OllamaOptions::from(body); + let options = ollama::OllamaOptions::from(&body.base); let payload = ollama::OllamaGenerateRequest { model, @@ -269,7 +202,7 @@ impl OllamaProvider { return Err(OllamaError::ModelNotFound(model.to_string())); } - let options = ollama::OllamaOptions::from(body); + let options = ollama::OllamaOptions::from(&body.base); let payload = ollama::OllamaGenerateRequest { model, @@ -332,90 +265,134 @@ impl OllamaProvider { Ok(ReceiverStream::new(rx)) } - // pub async fn chat_completions(&self, body: Value) -> Result { - // let (model, messages) = self.extract_chat_params(&body)?; - // self.validate_model(model).await?; + pub async fn chat_completions( + &self, + body: &api::ChatRequest, + ) -> Result { + let url = format!("{}/api/chat", self.base_url); - // let payload = json!({ - // "model": model, - // "messages": messages, - // "stream": false, - // "options": Self::build_options(&body), - // }); + let (messages, model) = self.extract_chat_params(body)?; - // let res = self - // .client - // .post(format!("{}/api/chat", self.base_url)) - // .json(&payload) - // .send() - // .await? - // .json::() - // .await?; + if body.messages.is_empty() { + return Err(OllamaError::MissingMessages); + } - // Ok(self.format_chat_response(&res)) - // } + if model.is_empty() { + return Err(OllamaError::MissingModel); + } - // 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 exists = self.model_exists(model).await?; + if !exists { + return Err(OllamaError::ModelNotFound(model.to_string())); + } - // let payload = json!({ - // "model": model, - // "messages": messages, - // "stream": true, - // "options": Self::build_options(&body), - // }); + // let options = ollama::OllamaOptions::from(body); + let options = ollama::OllamaOptions::from(&body.base); - // let mut byte_stream = self - // .client - // .post(format!("{}/api/chat", self.base_url)) - // .json(&payload) - // .send() - // .await? - // .bytes_stream(); + let payload = ollama::OllamaChatRequest { + model, + messages, + stream: false, + options, + }; - // let (tx, rx) = tokio::sync::mpsc::channel(32); + let res = self + .client + .post(url) + .json(&payload) + .send() + .await? + .json::() + .await?; - // 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; - // } - // }; + Ok(api::ChatCompletionResponse::from(res)) + } - // if let Ok(json) = serde_json::from_slice::(&chunk) { - // let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false); + pub async fn chat_completions_stream( + &self, + body: &api::ChatRequest, + ) -> Result>, OllamaError> { + let url = format!("{}/api/chat", self.base_url); - // 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 (messages, model) = self.extract_chat_params(body)?; - // let _ = tx.send(Ok(Event::default().data(event_data))).await; + if messages.is_empty() { + return Err(OllamaError::MissingMessages); + } - // if done { - // let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; - // break; - // } - // } - // } - // }); + if model.is_empty() { + return Err(OllamaError::MissingModel); + } - // Ok(ReceiverStream::new(rx)) - // } + let exists = self.model_exists(model).await?; + if !exists { + return Err(OllamaError::ModelNotFound(model.to_string())); + } + + let options = ollama::OllamaOptions::from(&body.base); + + let payload = ollama::OllamaChatRequest { + model, + messages, + stream: true, + options, + }; + + let mut byte_stream = self + .client + .post(url) + .json(&payload) + .send() + .await? + .bytes_stream(); + + let (tx, rx) = tokio::sync::mpsc::channel(32); + + tokio::spawn(async move { + let stream_id = format!("chatcmpl-{}", uuid::Uuid::new_v4()); + + 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; + } + }; + + let parsed: ollama::OllamaChatResponse = match serde_json::from_slice(&chunk) { + Ok(v) => v, + Err(_) => continue, + }; + + let event = api::ChatCompletionChunk { + id: stream_id.clone(), + object: "chat.completion.chunk".to_string(), + choices: vec![api::ChatChunkChoice { + index: 0, + delta: api::ChatDelta { + role: Some(parsed.message.role), + content: Some(parsed.message.content), + }, + finish_reason: if parsed.done { + Some(api::FinishReason::Stop) + } else { + None + }, + }], + }; + + let event_data = serde_json::to_string(&event).unwrap_or_default(); + + let _ = tx.send(Ok(Event::default().data(event_data))).await; + + if parsed.done { + let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; + break; + } + } + }); + + Ok(ReceiverStream::new(rx)) + } } diff --git a/src/providers/ollama/mapper.rs b/src/providers/ollama/mapper.rs index 87ccb5a..b3863df 100644 --- a/src/providers/ollama/mapper.rs +++ b/src/providers/ollama/mapper.rs @@ -18,20 +18,6 @@ impl From for api::ModelInfo { } } -impl From<&api::CompletionRequest> for ollama::OllamaOptions { - fn from(req: &api::CompletionRequest) -> Self { - Self { - temperature: req.base.temperature, - top_p: req.base.top_p, - top_k: req.base.top_k, - repeat_penalty: req.base.repeat_penalty, - seed: req.base.seed, - num_ctx: req.base.num_ctx, - num_predict: req.base.num_predict, - } - } -} - impl From for api::CompletionResponse { fn from(res: ollama::OllamaGenerateResponse) -> Self { Self { @@ -54,3 +40,47 @@ impl From for api::CompletionResponse { } } } + +impl From<&api::BaseLLMRequest> for ollama::OllamaOptions { + fn from(base: &api::BaseLLMRequest) -> Self { + Self { + temperature: base.temperature, + top_p: base.top_p, + top_k: base.top_k, + repeat_penalty: base.repeat_penalty, + seed: base.seed, + num_ctx: base.num_ctx, + num_predict: base.num_predict, + } + } +} + +impl From for api::ChatCompletionResponse { + fn from(res: ollama::OllamaChatResponse) -> Self { + let prompt_tokens = res.prompt_eval_count.unwrap_or(0); + let completion_tokens = res.eval_count.unwrap_or(0); + + Self { + id: Uuid::new_v4().to_string(), + object: "chat.completion".to_string(), + created: Utc::now().timestamp() as u64, + model: res.model, + + choices: vec![api::ChatChoice { + index: 0, + message: res.message, + finish_reason: if res.done { + api::FinishReason::Stop + } else { + api::FinishReason::Length + }, + }], + + usage: Some(api::Usage { + prompt_tokens, + completion_tokens, + total_tokens: prompt_tokens + completion_tokens, + }), + } + } +} diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index 646546a..9c7ba44 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -6,7 +6,6 @@ use axum::{ sse::{KeepAlive, Sse}, }, }; -use serde_json::Value; use crate::dto::api; use crate::errors::OllamaError; @@ -23,9 +22,7 @@ pub async fn completions( State(state): State, Json(body): Json, ) -> Result { - let wants_stream = body.base.stream; - - if wants_stream { + if body.base.stream { let stream = state .ollama .completions_stream(&body) @@ -42,35 +39,30 @@ pub async fn completions( } } -// pub async fn chat_completions( -// State(state): State, -// Json(body): Json, -// ) -> Result { -// let wants_stream = body -// .get("stream") -// .and_then(|v| v.as_bool()) -// .unwrap_or(false); +pub async fn chat_completions( + State(state): State, + Json(body): Json, +) -> Result { + if body.base.stream { + let stream = state + .ollama + .chat_completions_stream(&body) + .await + .map_err(ollama_err)?; -// 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(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()) -// } -// } + Ok(Json(response).into_response()) + } +} fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) { match e { diff --git a/src/routes/v1/mod.rs b/src/routes/v1/mod.rs index a1aaf47..0d7f6a0 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -14,7 +14,7 @@ pub fn protected_router() -> Router { Router::new() .route("/models", get(models::list_models)) .route("/completions", post(chat::completions)) - // .route("/chat/completions", post(chat::chat_completions)) + .route("/chat/completions", post(chat::chat_completions)) .route("/models/{model}/load", post(models::load_model)) .route("/models/{model}/unload", post(models::unload_model)) } @@ -22,5 +22,5 @@ pub fn protected_router() -> Router { pub fn router() -> Router { Router::new() // .merge(public_router()) - .merge(protected_router()) //.layer(middleware::from_fn(auth_middleware))) + .merge(protected_router().layer(middleware::from_fn(auth_middleware))) }