diff --git a/.sqlx/query-426a96fdb055d248f1f73d52dd10014340c6174b1491c1da36a5ba63fa554a9f.json b/.sqlx/query-426a96fdb055d248f1f73d52dd10014340c6174b1491c1da36a5ba63fa554a9f.json new file mode 100644 index 0000000..3b77ab0 --- /dev/null +++ b/.sqlx/query-426a96fdb055d248f1f73d52dd10014340c6174b1491c1da36a5ba63fa554a9f.json @@ -0,0 +1,15 @@ +{ + "db_name": "PostgreSQL", + "query": "UPDATE chat.message SET tokens = $1 WHERE id = $2", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Int4", + "Uuid" + ] + }, + "nullable": [] + }, + "hash": "426a96fdb055d248f1f73d52dd10014340c6174b1491c1da36a5ba63fa554a9f" +} diff --git a/.sqlx/query-c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859.json b/.sqlx/query-c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859.json new file mode 100644 index 0000000..7ae1020 --- /dev/null +++ b/.sqlx/query-c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859.json @@ -0,0 +1,26 @@ +{ + "db_name": "PostgreSQL", + "query": "\n INSERT INTO chat.message (conversation_id, parent_id, role, content, tokens)\n VALUES ($1, $2, $3, $4, $5)\n RETURNING id\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + } + ], + "parameters": { + "Left": [ + "Uuid", + "Uuid", + "Text", + "Text", + "Int4" + ] + }, + "nullable": [ + false + ] + }, + "hash": "c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859" +} diff --git a/src/databases/postgres/chat.rs b/src/databases/postgres/chat.rs index 2ce598a..9fd2feb 100644 --- a/src/databases/postgres/chat.rs +++ b/src/databases/postgres/chat.rs @@ -97,3 +97,74 @@ pub async fn set_conversation_title( tx.commit().await?; Ok(()) } + +#[derive(Debug, Clone, sqlx::Type)] +#[sqlx(type_name = "text")] +#[sqlx(rename_all = "lowercase")] +pub enum MessageRole { + User, + Assistant, + System, +} + +pub async fn insert_message( + pool: &PgPool, + user_id: Uuid, + conversation_id: Uuid, + parent_id: Option, + role: MessageRole, + content: &str, + tokens: Option, +) -> Result { + let mut tx = pool.begin().await?; + let conn = tx.acquire().await?; + + sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#) + .bind(user_id.to_string()) + .execute(&mut *conn) + .await?; + + let row = sqlx::query!( + r#" + INSERT INTO chat.message (conversation_id, parent_id, role, content, tokens) + VALUES ($1, $2, $3, $4, $5) + RETURNING id + "#, + conversation_id, + parent_id, + role as MessageRole, + content, + tokens + ) + .fetch_one(&mut *conn) + .await?; + + tx.commit().await?; + Ok(row.id) +} + +pub async fn update_message_tokens( + pool: &PgPool, + user_id: Uuid, + message_id: Uuid, + tokens: i32, +) -> Result<(), DbError> { + let mut tx = pool.begin().await?; + let conn = tx.acquire().await?; + + sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#) + .bind(user_id.to_string()) + .execute(&mut *conn) + .await?; + + sqlx::query!( + r#"UPDATE chat.message SET tokens = $1 WHERE id = $2"#, + tokens, + message_id + ) + .execute(&mut *conn) + .await?; + + tx.commit().await?; + Ok(()) +} diff --git a/src/dto/api.rs b/src/dto/api.rs index f841fd5..8764fe9 100644 --- a/src/dto/api.rs +++ b/src/dto/api.rs @@ -108,7 +108,7 @@ pub struct Choice { pub finish_reason: FinishReason, } -#[derive(Debug, Serialize, Deserialize, ToSchema)] +#[derive(Debug, Serialize, Deserialize, ToSchema, Clone, Copy)] pub struct Usage { pub prompt_tokens: u32, pub completion_tokens: u32, @@ -131,6 +131,7 @@ pub struct ChatRequest { // Non standard Open AI pub conversation_id: Option, + pub parent_id: Option, } #[derive(Debug, Deserialize, Serialize, ToSchema)] @@ -171,6 +172,7 @@ pub struct ChatCompletionChunk { pub id: String, pub object: String, pub choices: Vec, + pub usage: Option, } #[derive(Debug, Serialize, ToSchema)] diff --git a/src/providers/ollama/client.rs b/src/providers/ollama/client.rs index 50a63be..9d5fc91 100644 --- a/src/providers/ollama/client.rs +++ b/src/providers/ollama/client.rs @@ -345,7 +345,7 @@ impl OllamaProvider { pub async fn chat_completions_stream( &self, body: &api::ChatRequest, - ) -> Result>, OllamaError> { + ) -> Result>, OllamaError> { let url = format!("{}/api/chat", self.base_url); let (messages, model) = self.extract_chat_params(body)?; @@ -399,6 +399,18 @@ impl OllamaProvider { Err(_) => continue, }; + let usage = if parsed.done { + Some(api::Usage { + prompt_tokens: parsed.prompt_eval_count.unwrap_or(0) as u32, + completion_tokens: parsed.eval_count.unwrap_or(0) as u32, + total_tokens: (parsed.prompt_eval_count.unwrap_or(0) + + parsed.eval_count.unwrap_or(0)) + as u32, + }) + } else { + None + }; + let event = api::ChatCompletionChunk { id: stream_id.clone(), object: "chat.completion.chunk".to_string(), @@ -414,14 +426,12 @@ impl OllamaProvider { None }, }], + usage, }; - let event_data = serde_json::to_string(&event).unwrap_or_default(); - - let _ = tx.send(Ok(Event::default().data(event_data))).await; + let _ = tx.send(Ok(event)).await; if parsed.done { - let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; break; } } diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index ca568c8..60d7af1 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -7,11 +7,17 @@ use axum::{ extract::{Extension, State}, response::{ IntoResponse, Response, - sse::{KeepAlive, Sse}, + sse::{Event, KeepAlive, Sse}, }, }; +use sqlx::PgPool; +use tokio_stream::StreamExt; +use uuid::Uuid; -use crate::databases::postgres::chat::{get_or_create_conversation, set_conversation_title}; +use crate::databases::postgres::chat::{ + get_or_create_conversation, set_conversation_title, update_message_tokens, +}; +use crate::databases::postgres::{chat, errors}; use crate::dto::api; use crate::dto::api::CompletionRequest; use crate::providers::ollama::errors::into_http_response; @@ -126,6 +132,7 @@ pub async fn chat_completions( ) -> Result { tracing::debug!("Received /chat/completion with body {:?}", body); + // Get current conversation or create it let conversation_id = if matches!(&auth, Auth::Jwt(_)) { tracing::debug!( "Is conversation_id existing: {:?}", @@ -173,16 +180,157 @@ pub async fn chat_completions( .chat_completions_stream(&body) .await .map_err(into_http_response)?; - Ok(Sse::new(stream) - .keep_alive(KeepAlive::default()) - .into_response()) + + if let Some(conversation_id) = conversation_id { + let user_msg_id = log_user_message( + &state.postgres, + auth.user_id(), + conversation_id, + body.parent_id, + body.messages + .last() + .map(|m| m.content.as_str()) + .unwrap_or(""), + None, // not known at this stage + ) + .await + .map_err(errors::into_http_response)?; + + let start_event = serde_json::json!({ + "type": "start", + "model": body.base.model, + "conversation_id": conversation_id, + "created": chrono::Utc::now().timestamp() as u64, + }); + + let start_stream = futures::stream::once(async move { + Ok::( + Event::default() + .event("metadata") + .data(start_event.to_string()), + ) + }); + + let mut accumulated = String::new(); + + let wrapped_stream = stream.map(move |item| 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 content = accumulated.clone(); + let pool = state.postgres.clone(); + let user_id = auth.user_id(); + let prompt_tokens = chunk.usage.as_ref().map(|u| u.prompt_tokens as i32); + let completion_tokens = + chunk.usage.as_ref().map(|u| u.completion_tokens as i32); + + let end_event = serde_json::json!({ + "type": "end", + "usage": { + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "total_tokens": chunk.usage.as_ref().map(|u| u.total_tokens), + } + }); + + tokio::spawn(async move { + if let Some(prompt_tokens) = prompt_tokens { + let _ = update_message_tokens( + &pool, + user_id, + user_msg_id, + prompt_tokens, + ) + .await; + } + + let _ = log_assistant_message( + &pool, + user_id, + conversation_id, + user_msg_id, + &content, + completion_tokens, + ) + .await; + }); + + return Ok(Event::default() + .event("metadata") + .data(end_event.to_string())); + } + + let data = serde_json::to_string(&chunk).unwrap_or_default(); + Ok(Event::default().data(data)) + } + Err(e) => Err(e), + }); + + let full_stream = start_stream.chain(wrapped_stream); + + Ok(Sse::new(full_stream) + .keep_alive(KeepAlive::default()) + .into_response()) + } else { + let plain_stream = stream.map( + |item| -> Result { + match item { + Ok(chunk) => { + let data = serde_json::to_string(&chunk).unwrap_or_default(); + Ok(Event::default().data(data)) + } + Err(e) => Err(e), + } + }, + ); + + Ok(Sse::new(plain_stream) + .keep_alive(KeepAlive::default()) + .into_response()) + } } else { let mut response = state .ollama .chat_completions(&body) .await .map_err(into_http_response)?; + response.conversation_id = conversation_id; + + if let Some(conversation_id) = conversation_id { + // log user message, get its id as parent for assistant + let user_msg_id = log_user_message( + &state.postgres, + auth.user_id(), + conversation_id, + body.parent_id, + body.messages + .last() + .map(|m| m.content.as_str()) + .unwrap_or(""), + response.usage.map(|u| u.prompt_tokens as i32), + ) + .await + .map_err(errors::into_http_response)?; + + // log assistant response + log_assistant_message( + &state.postgres, + auth.user_id(), + conversation_id, + user_msg_id, + &response.choices[0].message.content, + response.usage.map(|u| u.completion_tokens as i32), + ) + .await + .map_err(errors::into_http_response)?; + } + Ok(Json(response).into_response()) } } @@ -207,3 +355,43 @@ async fn generate_conversation_title(ollama: &OllamaProvider, first_message: &st .filter(|t| !t.is_empty()) .unwrap_or_else(|| "New Conversation".to_string()) // ← default on any error } + +async fn log_user_message( + pool: &PgPool, + user_id: Uuid, + conversation_id: Uuid, + parent_id: Option, + content: &str, + tokens: Option, +) -> Result { + chat::insert_message( + pool, + user_id, + conversation_id, + parent_id, + chat::MessageRole::User, + content, + tokens, + ) + .await +} + +async fn log_assistant_message( + pool: &PgPool, + user_id: Uuid, + conversation_id: Uuid, + parent_id: Uuid, + content: &str, + tokens: Option, +) -> Result { + chat::insert_message( + pool, + user_id, + conversation_id, + Some(parent_id), + chat::MessageRole::Assistant, + content, + tokens, + ) + .await +} diff --git a/tests/ollama_provider.rs b/tests/ollama_provider.rs index ecc557c..a5f3827 100644 --- a/tests/ollama_provider.rs +++ b/tests/ollama_provider.rs @@ -184,6 +184,7 @@ async fn test_chat_completions_ok() { content: "What is 2+2?".to_string(), }], conversation_id: None, + parent_id: None, }; let res = provider.chat_completions(&req).await.unwrap(); @@ -219,6 +220,7 @@ async fn test_chat_completions_missing_messages() { }, messages: vec![], conversation_id: None, + parent_id: None, }; let err = provider.chat_completions(&req).await.unwrap_err(); @@ -246,6 +248,7 @@ async fn test_chat_completions_no_user_message() { content: "be helpful".to_string(), }], conversation_id: None, + parent_id: None, }; let err = provider.chat_completions(&req).await.unwrap_err(); @@ -273,6 +276,7 @@ async fn test_chat_completions_model_not_found() { content: "hi".to_string(), }], conversation_id: None, + parent_id: None, }; let err = provider.chat_completions(&req).await.unwrap_err();