use crate::{ databases::postgres::chat::ConversationState, dto::api::BaseLLMRequest, middlewares::auth::middleware::Auth, providers::ollama::client::OllamaProvider, }; use axum::{ Json, extract::{Extension, State}, response::{ IntoResponse, Response, 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, 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; use crate::state::app_state::AppState; #[utoipa::path( post, path = "/completions", tag = "chat", request_body( content = api::CompletionRequest, description = "Text completion request", content_type = "application/json" ), responses( ( status = 200, description = "Text completion response. If stream=true, response is SSE stream of chunks ending in [DONE].", body = api::CompletionResponse, content_type = "application/json" ), ( status = 400, description = "Invalid request: missing prompt, model, or invalid format", body = api::ErrorResponse, example = json!({ "error": "prompt is required and cannot be empty" }) ), ( status = 422, description = "Model not found or not available locally", body = api::ErrorResponse, example = json!({ "error": "model 'llama3' is not available — run `ollama pull llama3` first" }) ), ( status = 500, description = "Internal server error (Ollama or network failure)", body = api::ErrorResponse, example = json!({ "error": "connection refused" }) ) ) )] pub async fn completions( State(state): State, Json(body): Json, ) -> Result)> { tracing::debug!("Received /completion with body {:?}", body); if body.base.stream { let stream = state.ollama.completions_stream(&body).await.map_err(|e| { let (code, msg) = into_http_response(e); (code, Json(api::ErrorResponse::new(msg))) })?; Ok(Sse::new(stream) .keep_alive(KeepAlive::default()) .into_response()) } else { let response = state.ollama.completions(&body).await.map_err(|e| { let (code, msg) = into_http_response(e); (code, Json(api::ErrorResponse::new(msg))) })?; Ok(Json(response).into_response()) } } #[utoipa::path( post, path = "/chat/completions", tag = "chat", request_body( content = api::ChatRequest, description = "Chat completion request with message history", content_type = "application/json" ), responses( ( status = 200, description = "Chat completion response. If stream=false returns JSON. If stream=true returns SSE stream of chunks ending with [DONE].", body = api::ChatCompletionResponse, content_type = "application/json" ), ( status = 400, description = "Invalid request", body = api::ErrorResponse, example = json!({ "error": "messages array with at least one user message is required" }) ), ( status = 401, description = "Unauthorized", body = api::ErrorResponse, example = json!({ "error": "missing or invalid token" }) ), ( status = 422, description = "Model not found or unavailable", body = api::ErrorResponse, example = json!({ "error": "model 'llama3' is not available — run `ollama pull llama3` first" }) ), ( status = 500, description = "Internal server error", body = api::ErrorResponse, example = json!({ "error": "connection refused" }) ) ) )] pub async fn chat_completions( State(state): State, Extension(auth): Extension, Json(body): Json, ) -> 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: {:?}", body.conversation_id.is_some() ); let conversation_state = get_or_create_conversation(&state.postgres, body.conversation_id, auth.user_id()) .await .map_err(crate::databases::postgres::errors::into_http_response)?; let id = match &conversation_state { ConversationState::Existing(uuid) => *uuid, ConversationState::Created(uuid) => { let conversation_id = *uuid; let title = generate_conversation_title(&state.ollama, body.messages[0].content.as_str()) .await; if let Err(e) = set_conversation_title(&state.postgres, conversation_id, auth.user_id(), &title) .await { tracing::warn!( conversation_id = %conversation_id, error = %e, "Failed to set conversation title, continuing with no title" ); } tracing::debug!(conversation_id = %conversation_id, %title, "generated conversation title"); conversation_id } }; tracing::debug!("Using conversation_id: {:?}", id); Some(id) } else { None }; if body.base.stream { let stream = state .ollama .chat_completions_stream(&body) .await .map_err(into_http_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()) } } async fn generate_conversation_title(ollama: &OllamaProvider, first_message: &str) -> String { let request = CompletionRequest { base: BaseLLMRequest { model: "llama3:latest".to_string(), ..Default::default() }, prompt: format!( "Generate a short title (max 6 words) for the following chat conversation: {}", first_message ), }; ollama .completions(&request) .await .ok() .and_then(|r| r.choices.first().map(|c| c.text.trim().to_string())) .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 }