diff --git a/src/api/errors.rs b/src/api/errors.rs new file mode 100644 index 0000000..a48025f --- /dev/null +++ b/src/api/errors.rs @@ -0,0 +1,66 @@ +use crate::databases::postgres::errors::DbError; +use crate::providers::ollama::errors::OllamaError; + +use axum::http::StatusCode; +use axum::response::{IntoResponse, Response}; + +pub enum ApiError { + Db(DbError), + Ollama(OllamaError), +} + +impl From for ApiError { + fn from(e: DbError) -> Self { + ApiError::Db(e) + } +} + +impl From for ApiError { + fn from(e: OllamaError) -> Self { + ApiError::Ollama(e) + } +} + +impl IntoResponse for ApiError { + fn into_response(self) -> Response { + let (status, msg) = match self { + ApiError::Db(db_err) => match db_err { + DbError::Connection(_) => ( + StatusCode::INTERNAL_SERVER_ERROR, + "database connection error".to_string(), + ), + DbError::Timeout => (StatusCode::REQUEST_TIMEOUT, "database timeout".to_string()), + DbError::NotFound => (StatusCode::NOT_FOUND, "not found".to_string()), + }, + + ApiError::Ollama(ollama_err) => match ollama_err { + OllamaError::MissingPrompt => ( + StatusCode::BAD_REQUEST, + "prompt is required and cannot be empty".to_string(), + ), + OllamaError::MissingModel => ( + StatusCode::BAD_REQUEST, + "model is required and cannot be empty".to_string(), + ), + OllamaError::ModelNotFound(m) => ( + StatusCode::UNPROCESSABLE_ENTITY, + format!("model '{m}' is not available — run `ollama pull {m}` first"), + ), + OllamaError::MissingKeepAlive => ( + StatusCode::BAD_REQUEST, + "keep alive is required and cannot be empty".to_string(), + ), + OllamaError::InvalidKeepAlive(v) => { + (StatusCode::BAD_REQUEST, format!("invalid keep_alive '{v}'")) + } + OllamaError::MissingMessages => ( + StatusCode::BAD_REQUEST, + "messages array with at least one user message is required".to_string(), + ), + OllamaError::Http(e) => (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), + }, + }; + + (status, msg).into_response() + } +} diff --git a/src/api/mod.rs b/src/api/mod.rs new file mode 100644 index 0000000..629e98f --- /dev/null +++ b/src/api/mod.rs @@ -0,0 +1 @@ +pub mod errors; diff --git a/src/databases/postgres/api_key/mod.rs b/src/databases/postgres/api_key/mod.rs new file mode 100644 index 0000000..84c032e --- /dev/null +++ b/src/databases/postgres/api_key/mod.rs @@ -0,0 +1 @@ +pub mod queries; diff --git a/src/databases/postgres/api_key.rs b/src/databases/postgres/api_key/queries.rs similarity index 86% rename from src/databases/postgres/api_key.rs rename to src/databases/postgres/api_key/queries.rs index a087c50..d38b7d7 100644 --- a/src/databases/postgres/api_key.rs +++ b/src/databases/postgres/api_key/queries.rs @@ -1,4 +1,5 @@ -use super::errors::DbError; +use crate::databases::postgres::errors::DbError; + use sqlx::PgPool; use uuid::Uuid; diff --git a/src/databases/postgres/chat/mod.rs b/src/databases/postgres/chat/mod.rs new file mode 100644 index 0000000..0333ab5 --- /dev/null +++ b/src/databases/postgres/chat/mod.rs @@ -0,0 +1,2 @@ +pub mod queries; +pub mod types; diff --git a/src/databases/postgres/chat.rs b/src/databases/postgres/chat/queries.rs similarity index 87% rename from src/databases/postgres/chat.rs rename to src/databases/postgres/chat/queries.rs index c7fdc5d..facc050 100644 --- a/src/databases/postgres/chat.rs +++ b/src/databases/postgres/chat/queries.rs @@ -1,5 +1,6 @@ -use super::errors::DbError; -use crate::dto::postgres; +use crate::databases::postgres::chat::types; +use crate::databases::postgres::errors::DbError; + use sqlx::{Acquire, PgPool}; use uuid::Uuid; @@ -39,17 +40,11 @@ where Ok(rec.id) } -#[derive(Debug)] -pub enum ConversationState { - Existing(Uuid), - Created(Uuid), -} - pub async fn get_or_create_conversation( pool: &PgPool, conversation_id: Option, user_id: Uuid, -) -> Result { +) -> Result { tracing::debug!("Testing conversation"); let mut tx = pool.begin().await?; @@ -65,9 +60,11 @@ pub async fn get_or_create_conversation( match conversation_id { Some(id) => { validate_conversation(&mut *conn, id).await?; - ConversationState::Existing(id) + types::ConversationState::Existing(id) + } + None => { + types::ConversationState::Created(create_conversation(&mut *conn, user_id).await?) } - None => ConversationState::Created(create_conversation(&mut *conn, user_id).await?), } }; @@ -101,21 +98,12 @@ pub async fn set_conversation_title( 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, + role: types::MessageRole, content: &str, tokens: Option, ) -> Result { @@ -135,7 +123,7 @@ pub async fn insert_message( "#, conversation_id, parent_id, - role as MessageRole, + role as types::MessageRole, content, tokens.unwrap_or(0) as i32 ) @@ -178,8 +166,8 @@ pub async fn get_conversations_entries( pool: &PgPool, user_id: Uuid, limit: i64, - before: Option>, // cursor -) -> Result, DbError> { + before: Option>, +) -> Result, DbError> { let mut tx = pool.begin().await?; sqlx::query("SELECT set_config('app.current_user_id', $1::text, true)") @@ -188,7 +176,7 @@ pub async fn get_conversations_entries( .await?; let rows = sqlx::query_as!( - postgres::ConversationSummary, + types::ConversationSummary, r#" SELECT id, title, created_at, updated_at FROM chat.conversation @@ -214,7 +202,7 @@ pub async fn get_conversation_messages( conversation_id: Uuid, limit: i64, before: Option>, -) -> Result, DbError> { +) -> Result, DbError> { let mut conn = pool.acquire().await?; sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#) @@ -223,7 +211,7 @@ pub async fn get_conversation_messages( .await?; let rows = sqlx::query_as!( - postgres::MessageSummary, + types::MessageSummary, r#" SELECT id, parent_id, role, content, created_at, tokens FROM chat.message diff --git a/src/dto/postgres.rs b/src/databases/postgres/chat/types.rs similarity index 67% rename from src/dto/postgres.rs rename to src/databases/postgres/chat/types.rs index 17fb63f..c2f0583 100644 --- a/src/dto/postgres.rs +++ b/src/databases/postgres/chat/types.rs @@ -1,6 +1,21 @@ use serde::Serialize; use uuid::Uuid; +#[derive(Debug, Clone, sqlx::Type)] +#[sqlx(type_name = "text")] +#[sqlx(rename_all = "lowercase")] +pub enum MessageRole { + User, + Assistant, + System, +} + +#[derive(Debug)] +pub enum ConversationState { + Existing(Uuid), + Created(Uuid), +} + #[derive(Debug, sqlx::FromRow, Serialize)] pub struct ConversationSummary { pub id: Uuid, diff --git a/src/databases/postgres/errors.rs b/src/databases/postgres/errors.rs index 32e6128..6534304 100644 --- a/src/databases/postgres/errors.rs +++ b/src/databases/postgres/errors.rs @@ -1,4 +1,3 @@ -use axum::http::StatusCode; use thiserror::Error; #[derive(Debug, Error)] @@ -12,14 +11,3 @@ pub enum DbError { #[error("not found")] NotFound, } - -pub fn into_http_response(e: DbError) -> (StatusCode, String) { - match e { - DbError::Connection(_) => ( - StatusCode::INTERNAL_SERVER_ERROR, - "database connection error".to_string(), - ), - DbError::Timeout => (StatusCode::REQUEST_TIMEOUT, "database timeout".to_string()), - DbError::NotFound => (StatusCode::NOT_FOUND, "not found".to_string()), - } -} diff --git a/src/databases/postgres/mod.rs b/src/databases/postgres/mod.rs index e70c180..4db98f0 100644 --- a/src/databases/postgres/mod.rs +++ b/src/databases/postgres/mod.rs @@ -1,5 +1,6 @@ -pub mod api_key; -pub mod chat; pub mod errors; pub mod pool; + +pub mod api_key; +pub mod chat; pub mod user; diff --git a/src/databases/postgres/user/mod.rs b/src/databases/postgres/user/mod.rs new file mode 100644 index 0000000..84c032e --- /dev/null +++ b/src/databases/postgres/user/mod.rs @@ -0,0 +1 @@ +pub mod queries; diff --git a/src/databases/postgres/user.rs b/src/databases/postgres/user/queries.rs similarity index 87% rename from src/databases/postgres/user.rs rename to src/databases/postgres/user/queries.rs index 8e0a535..8e6a0ef 100644 --- a/src/databases/postgres/user.rs +++ b/src/databases/postgres/user/queries.rs @@ -1,4 +1,5 @@ -use super::errors::DbError; +use crate::databases::postgres::errors::DbError; + use sqlx::PgPool; use uuid::Uuid; diff --git a/src/dto/api.rs b/src/dto/api.rs index 1b934a7..a2c8306 100644 --- a/src/dto/api.rs +++ b/src/dto/api.rs @@ -234,8 +234,8 @@ pub struct ConversationQuery { #[derive(Debug, Serialize)] pub struct ConversationListResponse { - pub conversations: Vec, - pub has_more: bool, // client knows if there are more pages + pub conversations: Vec, + pub has_more: bool, } #[derive(Debug, Deserialize)] @@ -246,6 +246,6 @@ pub struct MessageQuery { #[derive(Debug, Serialize)] pub struct MessageListResponse { - pub messages: Vec, + pub messages: Vec, pub has_more: bool, } diff --git a/src/dto/mod.rs b/src/dto/mod.rs index 46e927b..b6fe8b2 100644 --- a/src/dto/mod.rs +++ b/src/dto/mod.rs @@ -1,3 +1,2 @@ pub mod api; pub mod ollama; -pub mod postgres; diff --git a/src/lib.rs b/src/lib.rs index fe7cb5c..ea8d223 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,4 @@ +pub mod api; pub mod databases; pub mod dto; pub mod middlewares; diff --git a/src/main.rs b/src/main.rs index 2ada6cb..b5f7acb 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,3 +1,4 @@ +mod api; mod databases; mod docs; mod dto; diff --git a/src/middlewares/auth/middleware.rs b/src/middlewares/auth/middleware.rs index e972dfb..a437cc4 100644 --- a/src/middlewares/auth/middleware.rs +++ b/src/middlewares/auth/middleware.rs @@ -5,7 +5,9 @@ use axum::{ response::{IntoResponse, Response}, }; -use crate::databases::postgres::{api_key::update_last_access, user::ensure_user_exists}; +use crate::databases::postgres::{ + api_key::queries::update_last_access, user::queries::ensure_user_exists, +}; use crate::middlewares::auth::apikey::{ApiKeyClaims, ApiKeyClaimsRoles}; use crate::middlewares::auth::keycloak::{KeycloakClaims, get_jwks, refresh_jwks, validate_token}; use crate::state::app_state::AppState; diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index 2e1d518..83491bb 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -1,6 +1,6 @@ use crate::{ - databases::postgres::chat::ConversationState, dto::api::BaseLLMRequest, - middlewares::auth::middleware::Auth, providers::ollama::client::OllamaProvider, + dto::api::BaseLLMRequest, middlewares::auth::middleware::Auth, + providers::ollama::client::OllamaProvider, }; use axum::{ Json, @@ -14,13 +14,14 @@ use sqlx::PgPool; use tokio_stream::StreamExt; use uuid::Uuid; -use crate::databases::postgres::chat::{ +use crate::api::errors::ApiError; +use crate::databases::postgres::chat::queries::{ get_conversation_messages, get_conversations_entries, get_or_create_conversation, - set_conversation_title, update_message_tokens, + insert_message, set_conversation_title, update_message_tokens, }; -use crate::databases::postgres::{chat, errors}; +use crate::databases::postgres::chat::types::{ConversationState, MessageRole}; +use crate::databases::postgres::errors; use crate::dto::api; -use crate::dto::api::CompletionRequest; use crate::providers::ollama::errors::into_http_response; use crate::state::app_state::AppState; @@ -89,11 +90,9 @@ async fn ensure_conversation( auth: &Auth, conversation_id: Option, first_message: &str, -) -> Result { +) -> Result { let conversation_state = - get_or_create_conversation(&state.postgres, conversation_id, auth.user_id()) - .await - .map_err(crate::databases::postgres::errors::into_http_response)?; + get_or_create_conversation(&state.postgres, conversation_id, auth.user_id()).await?; let id = match conversation_state { ConversationState::Existing(uuid) => uuid, @@ -127,14 +126,10 @@ async fn handle_stream( auth: Auth, body: api::ChatRequest, conversation_id: Option, -) -> Result { +) -> 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 - .map_err(into_http_response)?; + let stream = state.ollama.chat_completions_stream(&body).await?; let plain_stream = stream.map( |item| -> Result { @@ -163,8 +158,7 @@ async fn handle_stream( .unwrap_or(""), None, ) - .await - .map_err(errors::into_http_response)?; + .await?; let start_event = api::StreamEvent::Start(api::StartEventData { conversation_id: conv_id, @@ -275,12 +269,8 @@ async fn handle_non_stream( auth: Auth, body: api::ChatRequest, conversation_id: Option, -) -> Result { - let mut response = state - .ollama - .chat_completions(&body) - .await - .map_err(into_http_response)?; +) -> Result { + let mut response = state.ollama.chat_completions(&body).await?; response.conversation_id = conversation_id; @@ -296,8 +286,7 @@ async fn handle_non_stream( .unwrap_or(""), response.usage.map(|u| u.prompt_tokens), ) - .await - .map_err(errors::into_http_response)?; + .await?; log_assistant_message( &state.postgres, @@ -307,8 +296,7 @@ async fn handle_non_stream( &response.choices[0].message.content, response.usage.map(|u| u.completion_tokens), ) - .await - .map_err(errors::into_http_response)?; + .await?; } Ok(Json(response).into_response()) @@ -360,7 +348,7 @@ pub async fn chat_completions( State(state): State, Extension(auth): Extension, Json(mut body): Json, -) -> Result { +) -> Result { tracing::debug!("Received /chat/completion with body {:?}", body); let conversation_id = if matches!(&auth, Auth::Jwt(_)) { @@ -378,8 +366,7 @@ pub async fn chat_completions( depth as i64, None, // no cursor — fetch the most recent N messages ) - .await - .map_err(errors::into_http_response)?; + .await?; // Map MessageSummary → api::Message and prepend to the outgoing request let history_messages: Vec = history @@ -410,7 +397,7 @@ pub async fn chat_completions( } async fn generate_conversation_title(ollama: &OllamaProvider, first_message: &str) -> String { - let request = CompletionRequest { + let request = api::CompletionRequest { base: BaseLLMRequest { model: "llama3:latest".to_string(), ..Default::default() @@ -438,12 +425,12 @@ async fn log_user_message( content: &str, tokens: Option, ) -> Result { - chat::insert_message( + insert_message( pool, user_id, conversation_id, parent_id, - chat::MessageRole::User, + MessageRole::User, content, tokens, ) @@ -458,12 +445,12 @@ async fn log_assistant_message( content: &str, tokens: Option, ) -> Result { - chat::insert_message( + insert_message( pool, user_id, conversation_id, Some(parent_id), - chat::MessageRole::Assistant, + MessageRole::Assistant, content, tokens, ) @@ -475,8 +462,7 @@ pub async fn get_conversations( State(state): State, Extension(auth): Extension, Query(params): Query, -) -> Result, (axum::http::StatusCode, Json)> -{ +) -> Result, ApiError> { tracing::debug!("Conversation hit: {:?}", auth); let conversations = get_conversations_entries( @@ -485,11 +471,7 @@ pub async fn get_conversations( params.limit.unwrap_or(20), params.before, ) - .await - .map_err(|e| { - let (code, msg) = errors::into_http_response(e); - (code, Json(api::ErrorResponse::new(msg))) - })?; + .await?; let has_more = conversations.len() == params.limit.unwrap_or(20) as usize; @@ -504,7 +486,7 @@ pub async fn get_messages( Extension(auth): Extension, Path(conversation_id): Path, Query(params): Query, -) -> Result, (axum::http::StatusCode, Json)> { +) -> Result, ApiError> { tracing::debug!("Messages hit: {:?}", auth); let messages = get_conversation_messages( @@ -514,11 +496,7 @@ pub async fn get_messages( params.limit.unwrap_or(50), params.before, ) - .await - .map_err(|e| { - let (code, msg) = errors::into_http_response(e); - (code, Json(api::ErrorResponse::new(msg))) - })?; + .await?; let has_more = messages.len() == params.limit.unwrap_or(50) as usize;