diff --git a/readme.md b/readme.md index acf46d0..f705e86 100644 --- a/readme.md +++ b/readme.md @@ -275,6 +275,10 @@ This project turns Ollama into: 👉 A controllable model runtime 👉 A foundation for a full LLM gateway +# Bugs +## Ollama +- When model answer onoly with 1 tokens, answer is empty + # TODO - open api doc for bearer token - Unify check before sending to ollama payload diff --git a/src/api/docs.rs b/src/api/docs.rs index d91865d..2cd135a 100644 --- a/src/api/docs.rs +++ b/src/api/docs.rs @@ -24,24 +24,24 @@ use crate::api::routes; components( schemas( api::errors::ErrorResponse, - api::types::ModelsResponse, - api::types::ModelInfo, - api::types::LoadModelResponse, - api::types::LoadModelRequest, - api::types::UnloadModelResponse, - api::types::LLMOptions, - api::types::CompletionRequest, - api::types::CompletionObject, - api::types::FinishReason, - api::types::CompletionResponse, + api::types::ApiModelsResponse, + api::types::ApiModelInfo, + api::types::ApiLoadModelResponse, + api::types::ApiLoadModelRequest, + api::types::ApiUnloadModelResponse, + api::types::ApiLlmOptions, + api::types::ApiCompletionRequest, + api::types::ApiCompletionObject, + api::types::ApiFinishReason, + api::types::ApiCompletionResponse, api::types::Choice, api::types::Usage, api::types::CompletionChunk, - api::types::ChatRequest, - api::types::Message, - api::types::Role, - api::types::ChatResponse, - api::types::ChatChoice, + api::types::ApiChatRequest, + api::types::ApiMessage, + api::types::ApiRole, + api::types::ApiChatResponse, + api::types::ApiChatChoice, api::types::ChatCompletionChunk, api::types::ChatChunkChoice, api::types::Delta, diff --git a/src/api/errors.rs b/src/api/errors.rs index 470f048..308cdf6 100644 --- a/src/api/errors.rs +++ b/src/api/errors.rs @@ -17,6 +17,8 @@ pub struct ErrorResponse { pub code: String, } +#[derive(Debug, thiserror::Error)] +#[error("{message}")] pub struct ApiError { pub status: StatusCode, pub code: &'static str, diff --git a/src/api/middlewares/auth.rs b/src/api/middlewares/auth.rs index 28438ff..36ec3c8 100644 --- a/src/api/middlewares/auth.rs +++ b/src/api/middlewares/auth.rs @@ -58,8 +58,14 @@ pub async fn auth_middleware( }; match handle_auth(&state, req, next, auth).await { - Ok(response) => response, - Err(err) => err.into_response(), + Ok(response) => { + tracing::debug!("User authentified"); + response + } + Err(err) => { + tracing::debug!("Error during authentification {:?}", err); + err.into_response() + } } } diff --git a/src/api/routes/v1/chat.rs b/src/api/routes/v1/chat.rs index f0c76e5..e3a5a5e 100644 --- a/src/api/routes/v1/chat.rs +++ b/src/api/routes/v1/chat.rs @@ -20,7 +20,7 @@ use uuid::Uuid; path = "/completions", tag = "chat", request_body( - content = api::types::CompletionRequest, + content = api::types::ApiCompletionRequest, description = "Text completion request", content_type = "application/json" ), @@ -28,7 +28,7 @@ use uuid::Uuid; ( status = 200, description = "Text completion response. If stream=true, response is SSE stream of chunks ending in [DONE].", - body = api::types::CompletionResponse, + body = api::types::ApiCompletionResponse, content_type = "application/json" ), ( @@ -53,17 +53,16 @@ use uuid::Uuid; )] pub async fn completions( State(state): State, - Json(body): Json, + Json(body): Json, ) -> Result { tracing::debug!("Received /completion with body {:?}", body); let response = state.chat_service.complete(body.into()).await?; match response { - CompletionResult::NoStream(res) => Ok(Json::< - crate::core::llm::completions::CompletionResultNoStream, - >(res) - .into_response()), + CompletionResult::NoStream(res) => { + Ok(Json::(res.into()).into_response()) + } CompletionResult::Stream(stream) => { let sse_stream = stream.map(|item| match item { Ok(event) => { @@ -85,7 +84,7 @@ pub async fn completions( path = "/chat/completions", tag = "chat", request_body( - content = api::types::ChatRequest, + content = api::types::ApiChatRequest, description = "Chat completion request with message history", content_type = "application/json" ), @@ -93,7 +92,7 @@ pub async fn completions( ( status = 200, description = "Chat completion response. If stream=false returns JSON. If stream=true returns SSE stream of chunks ending with [DONE].", - body = api::types::ChatResponse, + body = api::types::ApiChatResponse, content_type = "application/json" ), ( @@ -125,7 +124,7 @@ pub async fn completions( pub async fn chat_completions( State(state): State, Extension(auth): Extension, - Json(body): Json, + Json(body): Json, ) -> Result { tracing::debug!("Received /chat/completion with body {:?}", body); @@ -133,15 +132,68 @@ pub async fn chat_completions( match response { crate::core::llm::chat::ChatCompletionResult::NoStream(res) => { - Ok(Json::(res).into_response()) + Ok(Json::(res.into()).into_response()) } crate::core::llm::chat::ChatCompletionResult::Stream(stream) => { - let sse_stream = stream.map(|item| match item { - Ok(event) => { - let data = serde_json::to_string(&event).unwrap_or_default(); - Ok(Event::default().data(data)) + let sse_stream = stream.map(|item| -> Result { + tracing::debug!("{:?}", item); + match item { + Err(e) => { + tracing::debug!("Error in chat_completions: {:?}", e); + Err(ApiError::from(e)) + } + + Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Start { + conversation_id, + message_id, + created_at, + }) => { + let payload = serde_json::to_string(&api::types::StreamEvent::Start( + api::types::StartEventData { + conversation_id, + created: created_at, + id: message_id, + }, + )) + .unwrap_or_default(); + Ok(Event::default().data(payload)) + } + + Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Token(tok)) => { + let chunk = api::types::ChatCompletionChunk { + id: String::new(), + object: "chat.completion.chunk".to_string(), + choices: vec![api::types::ChatChunkChoice { + index: 0, + delta: api::types::Delta { + content: Some(tok), + role: None, + }, + finish_reason: None, + }], + usage: None, + }; + let payload = serde_json::to_string(&api::types::StreamEvent::Delta(chunk)) + .unwrap_or_default(); + Ok(Event::default().data(payload)) + } + + Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Final(res)) => { + let payload = serde_json::to_string(&api::types::StreamEvent::End( + api::types::EndEventData { + created: res.created_at.parse().unwrap_or(0), + id: res.id, + usage: api::types::Usage { + prompt_tokens: res.prompt_tokens, + completion_tokens: res.completion_tokens, + total_tokens: res.prompt_tokens + res.completion_tokens, + }, + }, + )) + .unwrap_or_default(); + Ok(Event::default().data(payload)) + } } - Err(e) => Err(e), }); Ok(Sse::new(sse_stream) @@ -192,146 +244,3 @@ pub async fn get_messages( Ok(Json(messages.into())) } - -// async fn handle_stream( -// state: AppState, -// auth: Auth, -// body: api::ChatRequest, -// conversation_id: Option, -// ) -> Result { -// // Handle anonymous (API key) path early — no DB logging -// let Some(conv_id) = conversation_id else { -// let stream = state.ollama.chat_completions_stream(&body).await?; - -// let plain_stream = stream.map( -// |item| -> Result { -// match item { -// Ok(chunk) => Ok( -// Event::default().data(serde_json::to_string(&chunk).unwrap_or_default()) -// ), -// Err(e) => Err(e), -// } -// }, -// ); -// return Ok(Sse::new(plain_stream) -// .keep_alive(KeepAlive::default()) -// .into_response()); -// }; - -// // From here conv_id is a plain Uuid — all variables stay in scope -// let user_msg_id = log_user_message( -// &state.postgres, -// auth.user_id(), -// conv_id, -// body.parent_id, -// body.messages -// .last() -// .map(|m| m.content.as_str()) -// .unwrap_or(""), -// None, -// ) -// .await?; - -// let start_event = api::StreamEvent::Start(api::StartEventData { -// conversation_id: conv_id, -// created: chrono::Utc::now().timestamp() as u64, -// id: user_msg_id, -// }); - -// let (tx, rx) = tokio::sync::mpsc::channel::< -// Result, -// >(32); - -// // Send start event immediately, before Ollama is contacted -// let _ = tx -// .send(Ok(Event::default() -// .event("metadata") -// .data(serde_json::to_string(&start_event).unwrap()))) -// .await; - -// let pool = state.postgres.clone(); -// let user_id = auth.user_id(); - -// tokio::spawn(async move { -// // Ollama called inside spawn — start event already queued -// let stream = match state.ollama.chat_completions_stream(&body).await { -// Ok(s) => s, -// Err(e) => { -// let _ = tx.send(Err(e)).await; -// return; -// } -// }; - -// let mut stream = stream; -// let mut accumulated = String::new(); - -// while let Some(item) = futures::StreamExt::next(&mut stream).await { -// match item { -// Ok(chunk) => { -// let is_done = chunk.choices[0].finish_reason == Some(api::FinishReason::Stop); - -// if let Some(content) = chunk.choices[0].delta.content.as_ref() { -// accumulated.push_str(content); -// } - -// if is_done { -// let prompt_tokens = chunk.usage.as_ref().map(|u| u.prompt_tokens); -// let completion_tokens = chunk.usage.as_ref().map(|u| u.completion_tokens); - -// if let Some(pt) = prompt_tokens { -// let _ = update_message_tokens(&pool, user_id, user_msg_id, pt).await; -// } - -// let assistant_msg_id = log_assistant_message( -// &pool, -// user_id, -// conv_id, -// user_msg_id, -// &accumulated, -// completion_tokens, -// ) -// .await; - -// if let Ok(msg_id) = assistant_msg_id { -// let end_event = api::StreamEvent::End(api::EndEventData { -// usage: api::Usage { -// prompt_tokens: prompt_tokens.unwrap_or(0), -// completion_tokens: completion_tokens.unwrap_or(0), -// total_tokens: chunk -// .usage -// .as_ref() -// .map(|u| u.total_tokens) -// .unwrap_or(0), -// }, -// id: msg_id, -// created: chrono::Utc::now().timestamp() as u64, -// }); - -// let _ = tx -// .send(Ok(Event::default() -// .event("metadata") -// .data(serde_json::to_string(&end_event).unwrap()))) -// .await; -// } - -// break; -// } - -// let data = api::StreamEvent::Delta(chunk); -// let json = serde_json::to_string(&data).unwrap(); -// if tx.send(Ok(Event::default().data(json))).await.is_err() { -// break; -// } -// } -// Err(e) => { -// let _ = tx.send(Err(e)).await; -// break; -// } -// } -// } -// }); - -// Ok(Sse::new(tokio_stream::wrappers::ReceiverStream::new(rx)) -// .keep_alive(KeepAlive::default()) -// .into_response()) -// } diff --git a/src/api/routes/v1/mod.rs b/src/api/routes/v1/mod.rs index 3cd966e..dce662a 100644 --- a/src/api/routes/v1/mod.rs +++ b/src/api/routes/v1/mod.rs @@ -43,9 +43,11 @@ pub fn protected_router() -> Router { pub fn router(state: SharedState) -> Router { Router::new() .merge(public_router()) - .merge(protected_router().layer(middleware::from_fn_with_state( - state.clone(), - auth_middleware, - ))) + .merge( + protected_router().route_layer(middleware::from_fn_with_state( + state.clone(), + auth_middleware, + )), + ) .with_state(state) } diff --git a/src/api/routes/v1/models.rs b/src/api/routes/v1/models.rs index 11cd32a..63e4347 100644 --- a/src/api/routes/v1/models.rs +++ b/src/api/routes/v1/models.rs @@ -13,7 +13,7 @@ use axum::{ ( status = 200, description = "List of locally available Ollama models", - body = api::types::ModelsResponse, + body = api::types::ApiModelsResponse, content_type = "application/json", ), ( @@ -27,7 +27,7 @@ use axum::{ // #[axum::debug_handler] pub async fn list_models( State(state): State, -) -> Result, api::errors::ApiError> { +) -> Result, api::errors::ApiError> { let models = state.chat_service.list_models().await?; Ok(Json(models.into())) @@ -41,7 +41,7 @@ pub async fn list_models( ("model" = String, Path, description = "Name of the model to load into memory (e.g. 'llama3')") ), request_body( - content = api::types::LoadModelRequest, + content = api::types::ApiLoadModelRequest, description = "Load model request", content_type = "application/json", example = json!({ "keep_alive": "10m" }) @@ -50,7 +50,7 @@ pub async fn list_models( ( status = 200, description = "Model successfully loaded into memory", - body = api::types::LoadModelResponse, + body = api::types::ApiLoadModelResponse, content_type = "application/json", ), ( @@ -79,8 +79,8 @@ pub async fn list_models( pub async fn load_model( State(state): State, Path(model): Path, - Json(body): Json, -) -> Result, api::errors::ApiError> { + Json(body): Json, +) -> Result, api::errors::ApiError> { let response = state .chat_service .load_model(crate::core::llm::models::LoadModelRequest { @@ -89,7 +89,7 @@ pub async fn load_model( }) .await?; - Ok(Json(api::types::LoadModelResponse { + Ok(Json(api::types::ApiLoadModelResponse { model: response.model, keep_alive: body.keep_alive, status: "loaded".to_string(), diff --git a/src/api/types.rs b/src/api/types.rs index 39de1be..5c88bc6 100644 --- a/src/api/types.rs +++ b/src/api/types.rs @@ -6,43 +6,43 @@ use uuid::Uuid; // ------ Models ------ #[derive(Debug, Serialize, ToSchema)] -pub struct ModelInfo { +pub struct ApiModelInfo { pub name: String, pub family: Option, pub parameter_size: Option, - pub metadata: ModelMetadata, + pub metadata: ApiModelMetadata, } #[derive(Debug, Clone, Default, Serialize, ToSchema)] -pub struct ModelMetadata { +pub struct ApiModelMetadata { pub extra: HashMap, } // ------ Endpoint: /models ------ #[derive(Debug, Serialize, ToSchema)] -pub struct ModelsResponse { - pub models: Vec, +pub struct ApiModelsResponse { + pub models: Vec, } // ------ Endpoint: /models/{model}/load ------ #[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct LoadModelResponse { +pub struct ApiLoadModelResponse { pub model: String, pub status: String, pub keep_alive: String, } #[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct LoadModelRequest { +pub struct ApiLoadModelRequest { pub keep_alive: String, } // ------ Endpoint: /models/{model}/unload ------ #[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct UnloadModelResponse { +pub struct ApiUnloadModelResponse { pub model: String, pub status: String, } @@ -50,7 +50,7 @@ pub struct UnloadModelResponse { // ------ Completions ------ #[derive(Debug, Deserialize, Serialize, ToSchema, Default)] -pub struct LLMOptions { +pub struct ApiLlmOptions { #[serde(default)] pub stream: bool, @@ -74,18 +74,18 @@ pub struct LLMOptions { // ------ Endpoint: /completions ------ #[derive(Debug, Deserialize, Serialize, ToSchema)] -pub struct CompletionRequest { +pub struct ApiCompletionRequest { #[serde(flatten)] - pub options: LLMOptions, + pub options: ApiLlmOptions, pub model: String, pub prompt: String, } #[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct CompletionResponse { +pub struct ApiCompletionResponse { pub id: Uuid, - pub object: CompletionObject, + pub object: ApiCompletionObject, pub created: String, pub model: String, pub choices: Vec, @@ -94,7 +94,7 @@ pub struct CompletionResponse { #[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)] #[serde(rename_all = "snake_case")] -pub enum CompletionObject { +pub enum ApiCompletionObject { TextCompletion, } @@ -102,7 +102,7 @@ pub enum CompletionObject { #[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)] #[serde(rename_all = "snake_case")] -pub enum FinishReason { +pub enum ApiFinishReason { Stop, Length, ContentFilter, @@ -114,7 +114,7 @@ pub enum FinishReason { pub struct Choice { pub text: String, pub index: u32, - pub finish_reason: FinishReason, + pub finish_reason: ApiFinishReason, } #[derive(Debug, Serialize, Deserialize, ToSchema, Clone, Copy)] @@ -132,12 +132,12 @@ pub struct CompletionChunk { } #[derive(Debug, Deserialize, Serialize, ToSchema)] -pub struct ChatRequest { +pub struct ApiChatRequest { #[serde(flatten)] - pub base: LLMOptions, + pub base: ApiLlmOptions, pub model: String, - pub message: Message, + pub message: ApiMessage, // Non standard Open AI pub conversation_id: Option, @@ -145,36 +145,36 @@ pub struct ChatRequest { } #[derive(Debug, Deserialize, Serialize, ToSchema)] -pub struct Message { - pub role: Role, +pub struct ApiMessage { + pub role: ApiRole, pub content: String, } #[derive(Debug, Deserialize, Serialize, ToSchema, PartialEq, Eq)] #[serde(rename_all = "lowercase")] -pub enum Role { +pub enum ApiRole { System, User, Assistant, } #[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct ChatResponse { - pub id: String, - pub object: String, - pub created: u64, +pub struct ApiChatResponse { + pub id: Uuid, + pub object: ApiCompletionObject, + pub created: String, pub model: String, - pub choices: Vec, + pub choices: Vec, pub usage: Usage, pub conversation_id: Uuid, } #[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct ChatChoice { +pub struct ApiChatChoice { pub index: u32, - pub message: Message, - pub finish_reason: FinishReason, + pub message: ApiMessage, + pub finish_reason: ApiFinishReason, } #[derive(Debug, Serialize, Deserialize, ToSchema)] @@ -189,39 +189,39 @@ pub struct ChatCompletionChunk { pub struct ChatChunkChoice { pub index: u32, pub delta: Delta, - pub finish_reason: Option, + pub finish_reason: Option, } #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct Delta { pub content: Option, - pub role: Option, + pub role: Option, } -// #[derive(Debug, Serialize, Deserialize, ToSchema)] -// pub struct StartEventData { -// pub conversation_id: Uuid, -// pub created: u64, -// pub id: Uuid, -// } +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct StartEventData { + pub conversation_id: Uuid, + pub created: u64, + pub id: Uuid, +} -// #[derive(Debug, Serialize, Deserialize, ToSchema)] -// pub struct EndEventData { -// pub created: u64, -// pub id: Uuid, -// pub usage: Usage, -// } +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct EndEventData { + pub created: u64, + pub id: Uuid, + pub usage: Usage, +} -// #[derive(Debug, Serialize, Deserialize, ToSchema)] -// #[serde(tag = "type", content = "data")] -// pub enum StreamEvent { -// #[serde(rename = "start")] -// Start(StartEventData), -// #[serde(rename = "end")] -// End(EndEventData), -// #[serde(rename = "delta")] -// Delta(ChatCompletionChunk), -// } +#[derive(Debug, Serialize, Deserialize, ToSchema)] +#[serde(tag = "type", content = "data")] +pub enum StreamEvent { + #[serde(rename = "start")] + Start(StartEventData), + #[serde(rename = "end")] + End(EndEventData), + #[serde(rename = "delta")] + Delta(ChatCompletionChunk), +} // ------ Api Key ------ diff --git a/src/core/llm/chat.rs b/src/core/llm/chat.rs index f4e37ab..3cf389d 100644 --- a/src/core/llm/chat.rs +++ b/src/core/llm/chat.rs @@ -1,10 +1,11 @@ use super::ChatRole; -use crate::providers::ollama::errors::LlmError; +use crate::services::errors::ServiceError; use futures::Stream; use serde::Serialize; use std::pin::Pin; +use uuid::Uuid; #[derive(Debug, Clone, Default)] pub struct ChatCompletionOptions { @@ -45,24 +46,30 @@ pub struct ChatCompletionResultNoStream { pub created_at: String, pub id: uuid::Uuid, + pub conversation_id: Uuid, pub prompt_tokens: u32, pub completion_tokens: u32, - pub done_reason: Option, + pub done_reason: String, - pub total_duration: Option, - pub load_duration: Option, + pub total_duration: u64, + pub load_duration: u64, } -#[derive(Serialize)] +#[derive(Debug, Serialize)] pub enum ChatCompletionStreamEvent { + Start { + conversation_id: Uuid, + message_id: Uuid, + created_at: u64, + }, Token(String), Final(ChatCompletionResultNoStream), } pub type ChatCompletionStream = - Pin> + Send>>; + Pin> + Send>>; pub enum ChatCompletionResult { Stream(ChatCompletionStream), diff --git a/src/mappers/api_to_core.rs b/src/mappers/api_to_core.rs index b589f7d..eeaf891 100644 --- a/src/mappers/api_to_core.rs +++ b/src/mappers/api_to_core.rs @@ -1,7 +1,7 @@ use crate::{api, core}; -impl From for core::llm::completions::CompletionRequest { - fn from(m: api::types::CompletionRequest) -> Self { +impl From for core::llm::completions::CompletionRequest { + fn from(m: api::types::ApiCompletionRequest) -> Self { Self { model: m.model, prompt: m.prompt, @@ -20,8 +20,8 @@ impl From for core::llm::completions::CompletionR } } -impl From for core::llm::chat::ChatCompletionRequest { - fn from(m: api::types::ChatRequest) -> Self { +impl From for core::llm::chat::ChatCompletionRequest { + fn from(m: api::types::ApiChatRequest) -> Self { Self { model: m.model, message: m.message.into(), @@ -43,13 +43,13 @@ impl From for core::llm::chat::ChatCompletionRequest { } } -impl From for core::llm::chat::Message { - fn from(m: api::types::Message) -> Self { +impl From for core::llm::chat::Message { + fn from(m: api::types::ApiMessage) -> Self { Self { role: match m.role { - api::types::Role::System => core::llm::ChatRole::System, - api::types::Role::User => core::llm::ChatRole::User, - api::types::Role::Assistant => core::llm::ChatRole::Assistant, + api::types::ApiRole::System => core::llm::ChatRole::System, + api::types::ApiRole::User => core::llm::ChatRole::User, + api::types::ApiRole::Assistant => core::llm::ChatRole::Assistant, }, content: m.content, } diff --git a/src/mappers/core_to_api.rs b/src/mappers/core_to_api.rs index 2886675..8c73ea0 100644 --- a/src/mappers/core_to_api.rs +++ b/src/mappers/core_to_api.rs @@ -1,6 +1,6 @@ use crate::{api, core}; -impl From for api::types::ModelsResponse { +impl From for api::types::ApiModelsResponse { fn from(m: core::llm::models::Models) -> Self { Self { models: m.models.into_iter().map(Into::into).collect(), @@ -8,7 +8,7 @@ impl From for api::types::ModelsResponse { } } -impl From for api::types::ModelMetadata { +impl From for api::types::ApiModelMetadata { fn from(m: core::llm::models::ModelMetadata) -> Self { Self { extra: m @@ -20,7 +20,7 @@ impl From for api::types::ModelMetadata { } } -impl From for api::types::ModelInfo { +impl From for api::types::ApiModelInfo { fn from(m: core::llm::models::Model) -> Self { Self { name: m.name, @@ -79,15 +79,28 @@ impl From for api::types::MessageLi } } -impl From for api::types::CompletionResponse { +impl From for api::types::ApiMessage { + fn from(c: core::llm::chat::Message) -> Self { + Self { + role: match c.role { + core::llm::ChatRole::User => api::types::ApiRole::User, + core::llm::ChatRole::Assistant => api::types::ApiRole::Assistant, + core::llm::ChatRole::System => api::types::ApiRole::System, + }, + content: c.content, + } + } +} + +impl From for api::types::ApiCompletionResponse { fn from(m: core::llm::completions::CompletionResultNoStream) -> Self { Self { id: m.id, - object: api::types::CompletionObject::TextCompletion, + object: api::types::ApiCompletionObject::TextCompletion, created: m.created_at, model: m.model, choices: vec![api::types::Choice { - finish_reason: api::types::FinishReason::Stop, + finish_reason: api::types::ApiFinishReason::Stop, text: m.text, index: 0, }], @@ -99,3 +112,25 @@ impl From for api::types::Comp } } } + +impl From for api::types::ApiChatResponse { + fn from(m: core::llm::chat::ChatCompletionResultNoStream) -> Self { + Self { + id: m.id, + conversation_id: m.conversation_id, + object: api::types::ApiCompletionObject::TextCompletion, + created: m.created_at, + model: m.model, + choices: vec![api::types::ApiChatChoice { + finish_reason: api::types::ApiFinishReason::Stop, + message: m.message.into(), + index: 0, + }], + usage: api::types::Usage { + completion_tokens: m.completion_tokens, + prompt_tokens: m.prompt_tokens, + total_tokens: m.completion_tokens + m.prompt_tokens, + }, + } + } +} diff --git a/src/providers/ollama/client.rs b/src/providers/ollama/client.rs index 95f9995..5c319a6 100644 --- a/src/providers/ollama/client.rs +++ b/src/providers/ollama/client.rs @@ -285,41 +285,62 @@ impl OllamaProvider { .await? .bytes_stream(); - let stream = byte_stream.flat_map(|chunk_result| { - let mut out: Vec> = Vec::new(); - - let chunk = match chunk_result { - Ok(b) => b, - Err(e) => { - out.push(Err(LlmError::Http(e))); - return futures::stream::iter(out); - } - }; - - for line in chunk.split(|&b| b == b'\n') { - if line.is_empty() { - continue; - } - - let parsed: ollama::types::OllamaChatResponse = match serde_json::from_slice(line) { - Ok(v) => v, - Err(_) => continue, + let stream = byte_stream + .flat_map(|chunk_result| { + let mut out: Vec> = + Vec::new(); + let chunk = match chunk_result { + Ok(b) => b, + Err(e) => { + tracing::debug!("Error: {:?}", e); + out.push(Err(LlmError::Http(e))); + return futures::stream::iter(out); + } }; - if !parsed.message.content.is_empty() && !parsed.done { - out.push(Ok(super::types::OllamaChatStreamEvent::Token( - parsed.message.content.clone(), - ))); - } + tracing::debug!("Raw line: {:?}", std::str::from_utf8(&chunk)); - if parsed.done { - out.push(Ok(super::types::OllamaChatStreamEvent::Final(parsed))); - return futures::stream::iter(out); - } - } + for line in chunk.split(|&b| b == b'\n') { + if line.is_empty() { + continue; + } - futures::stream::iter(out) - }); + let parsed: ollama::types::OllamaChatResponse = + match serde_json::from_slice(line) { + Ok(v) => v, + Err(_) => continue, + }; + + tracing::debug!("Parsed: {:?}", parsed); + + if !parsed.message.content.is_empty() { + out.push(Ok(super::types::OllamaChatStreamEvent::Token( + parsed.message.content.clone(), + ))); + } + + if parsed.done { + out.push(Ok(super::types::OllamaChatStreamEvent::Final(parsed))); + return futures::stream::iter(out); + } + } + futures::stream::iter(out) + }) + .scan(String::new(), |acc, event| { + let result = match event { + Ok(ollama::types::OllamaChatStreamEvent::Token(ref tok)) => { + acc.push_str(tok); + Some(event) + } + Ok(ollama::types::OllamaChatStreamEvent::Final(mut resp)) => { + tracing::debug!("Acc: {:?}", acc); + resp.message.content = std::mem::take(acc); + Some(Ok(ollama::types::OllamaChatStreamEvent::Final(resp))) + } + Err(_) => Some(event), + }; + futures::future::ready(result) + }); Ok(Box::pin(stream)) } diff --git a/src/providers/ollama/types.rs b/src/providers/ollama/types.rs index b991917..4d4511d 100644 --- a/src/providers/ollama/types.rs +++ b/src/providers/ollama/types.rs @@ -117,13 +117,13 @@ pub struct OllamaChatResponse { pub message: OllamaMessage, pub done: bool, - pub done_reason: Option, + pub done_reason: String, - pub total_duration: Option, - pub load_duration: Option, + pub total_duration: u64, + pub load_duration: u64, - pub prompt_eval_count: Option, - pub eval_count: Option, + pub prompt_eval_count: u32, + pub eval_count: u32, } #[derive(Debug)] diff --git a/src/services/chat_service.rs b/src/services/chat_service.rs index 845fad8..45ca5ab 100644 --- a/src/services/chat_service.rs +++ b/src/services/chat_service.rs @@ -109,7 +109,14 @@ impl ChatService { .resolve_conversation_with_title(auth.user_id(), &body) .await?; + tracing::debug!( + "Received conversation_id={:?}, parent_id={:?}", + conversation_id, + body.parent_id + ); + let user_msg_id = self + .conversation .log_user_message( auth.user_id(), conversation_id, @@ -122,72 +129,100 @@ impl ChatService { let stream = body.options.stream; let history = self - .build_chat_history( - auth.user_id(), - conversation_id, - body.message.clone(), - body.options.context_depth, - ) + .build_chat_history(auth.user_id(), conversation_id, body.options.context_depth) .await?; let mut request: crate::providers::ollama::types::OllamaChatRequest = body.into(); request.messages = history.into_iter().map(Into::into).collect(); if stream { - let ollama_stream = self.ollama.chat_completions_stream(&request).await?; + let created_at = 3; - let mapped = ollama_stream.map(|item| { - item.map(|event| match event { - OllamaChatStreamEvent::Token(tok) => { - core::llm::chat::ChatCompletionStreamEvent::Token(tok) - } - OllamaChatStreamEvent::Final(resp) => { - core::llm::chat::ChatCompletionStreamEvent::Final( - core::llm::chat::ChatCompletionResultNoStream { - id: Uuid::new_v4(), - created_at: resp.created_at, - model: resp.model, - message: resp.message.into(), - prompt_tokens: resp.prompt_eval_count.unwrap_or(0), - completion_tokens: resp.eval_count.unwrap_or(0), - done_reason: resp.done_reason, - total_duration: resp.total_duration, - load_duration: resp.load_duration, - }, - ) - } + let start_event = futures::stream::once(async move { + Ok(core::llm::chat::ChatCompletionStreamEvent::Start { + conversation_id, + message_id: user_msg_id, + created_at, }) }); + let ollama_stream = self.ollama.chat_completions_stream(&request).await?; + + let conversation_svc = self.conversation.clone(); + let user_id = auth.user_id(); + + let mapped = ollama_stream.then(move |item| { + let conversation_svc = conversation_svc.clone(); + async move { + match item { + Err(e) => Err(e.into()), + Ok(OllamaChatStreamEvent::Token(tok)) => { + Ok(core::llm::chat::ChatCompletionStreamEvent::Token(tok)) + } + Ok(OllamaChatStreamEvent::Final(resp)) => { + tracing::debug!("Inserting {:?}", resp.message.content); + + let assistant_message_id = conversation_svc + .log_assistant_message( + user_id, + conversation_id, + user_msg_id, + &resp.message.content, + resp.eval_count, + ) + .await?; + + conversation_svc + .update_message_tokens(user_id, user_msg_id, resp.prompt_eval_count) + .await?; + + Ok(core::llm::chat::ChatCompletionStreamEvent::Final( + core::llm::chat::ChatCompletionResultNoStream { + id: assistant_message_id, + conversation_id, + created_at: resp.created_at, + model: resp.model, + message: resp.message.into(), + prompt_tokens: resp.prompt_eval_count, + completion_tokens: resp.eval_count, + done_reason: resp.done_reason, + total_duration: resp.total_duration, + load_duration: resp.load_duration, + }, + )) + } + } + } + }); + Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin( - mapped, + start_event.chain(mapped), ))) } else { let response = self.ollama.chat_completions(&request).await?; - self.log_assistant_message( - auth.user_id(), - conversation_id, - user_msg_id, - &response.message.content, - response.eval_count, - ) - .await?; - self.conversation - .update_message_tokens( + let assistant_message_id = self + .conversation + .log_assistant_message( auth.user_id(), + conversation_id, user_msg_id, - response.prompt_eval_count.unwrap_or_default(), + &response.message.content, + response.eval_count, ) .await?; + self.conversation + .update_message_tokens(auth.user_id(), user_msg_id, response.prompt_eval_count) + .await?; let enriched = core::llm::chat::ChatCompletionResultNoStream { - id: Uuid::new_v4(), + id: assistant_message_id, + conversation_id, created_at: response.created_at, model: response.model, message: response.message.into(), - prompt_tokens: response.prompt_eval_count.unwrap_or(0), - completion_tokens: response.eval_count.unwrap_or(0), + prompt_tokens: response.prompt_eval_count, + completion_tokens: response.eval_count, done_reason: response.done_reason, total_duration: response.total_duration, load_duration: response.load_duration, @@ -247,7 +282,6 @@ impl ChatService { &self, auth_user_id: Uuid, conversation_id: Uuid, - body_messages: crate::core::llm::chat::Message, context_depth: u32, ) -> Result, ServiceError> { let messages = self @@ -270,48 +304,6 @@ impl ChatService { history.reverse(); - history.push(body_messages); - Ok(history) } - - async fn log_user_message( - &self, - user_id: Uuid, - conversation_id: Uuid, - parent_id: Option, - content: &str, - tokens: Option, - ) -> Result { - self.conversation - .log_message( - user_id, - conversation_id, - parent_id, - llm::ChatRole::User, - content, - tokens, - ) - .await - } - - async fn log_assistant_message( - &self, - user_id: Uuid, - conversation_id: Uuid, - parent_id: Uuid, - content: &str, - tokens: Option, - ) -> Result { - self.conversation - .log_message( - user_id, - conversation_id, - Some(parent_id), - llm::ChatRole::Assistant, - content, - tokens, - ) - .await - } } diff --git a/src/services/conversation_service.rs b/src/services/conversation_service.rs index 5e378dd..c4aea3d 100644 --- a/src/services/conversation_service.rs +++ b/src/services/conversation_service.rs @@ -46,17 +46,23 @@ impl ConversationService { ) -> Result { let limit = pointer.limit; - let messages = postgres::chat::queries::get_conversation_messages( + let mut messages = postgres::chat::queries::get_conversation_messages( &self.postgres, user_id, conversation_id, - limit, + limit + 1, pointer.before, ) .await?; let has_more = messages.len() == limit as usize; + if has_more { + messages.pop(); + } + + messages.reverse(); + Ok(crate::core::databases::conversations::MessageList { messages: messages.into_iter().map(Into::into).collect(), has_more, @@ -118,6 +124,44 @@ impl ConversationService { Ok(id) } + pub async fn log_user_message( + &self, + user_id: Uuid, + conversation_id: Uuid, + parent_id: Option, + content: &str, + tokens: Option, + ) -> Result { + self.log_message( + user_id, + conversation_id, + parent_id, + crate::core::llm::ChatRole::User, + content, + tokens, + ) + .await + } + + pub async fn log_assistant_message( + &self, + user_id: Uuid, + conversation_id: Uuid, + parent_id: Uuid, + content: &str, + tokens: u32, + ) -> Result { + self.log_message( + user_id, + conversation_id, + Some(parent_id), + crate::core::llm::ChatRole::Assistant, + content, + Some(tokens), + ) + .await + } + pub async fn update_message_tokens( &self, user_id: Uuid, diff --git a/src/services/errors.rs b/src/services/errors.rs index 0b61500..c3bac93 100644 --- a/src/services/errors.rs +++ b/src/services/errors.rs @@ -2,6 +2,7 @@ use crate::databases::errors::DbError; use crate::providers::keycloak::errors::AuthError; use crate::providers::ollama::errors::LlmError; +#[derive(Debug)] pub enum ServiceError { Db(DbError), Llm(LlmError),