refactor: postgres files

This commit is contained in:
2026-05-21 22:27:03 +02:00
parent 0d31ade1f6
commit 739539309a
17 changed files with 143 additions and 97 deletions
+66
View File
@@ -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<DbError> for ApiError {
fn from(e: DbError) -> Self {
ApiError::Db(e)
}
}
impl From<OllamaError> 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()
}
}
+1
View File
@@ -0,0 +1 @@
pub mod errors;
+1
View File
@@ -0,0 +1 @@
pub mod queries;
@@ -1,4 +1,5 @@
use super::errors::DbError;
use crate::databases::postgres::errors::DbError;
use sqlx::PgPool;
use uuid::Uuid;
+2
View File
@@ -0,0 +1,2 @@
pub mod queries;
pub mod types;
@@ -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<Uuid>,
user_id: Uuid,
) -> Result<ConversationState, DbError> {
) -> Result<types::ConversationState, DbError> {
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<Uuid>,
role: MessageRole,
role: types::MessageRole,
content: &str,
tokens: Option<u32>,
) -> Result<Uuid, DbError> {
@@ -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<chrono::DateTime<chrono::Utc>>, // cursor
) -> Result<Vec<postgres::ConversationSummary>, DbError> {
before: Option<chrono::DateTime<chrono::Utc>>,
) -> Result<Vec<types::ConversationSummary>, 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<chrono::DateTime<chrono::Utc>>,
) -> Result<Vec<postgres::MessageSummary>, DbError> {
) -> Result<Vec<types::MessageSummary>, 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
@@ -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,
-12
View File
@@ -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()),
}
}
+3 -2
View File
@@ -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;
+1
View File
@@ -0,0 +1 @@
pub mod queries;
@@ -1,4 +1,5 @@
use super::errors::DbError;
use crate::databases::postgres::errors::DbError;
use sqlx::PgPool;
use uuid::Uuid;
+3 -3
View File
@@ -234,8 +234,8 @@ pub struct ConversationQuery {
#[derive(Debug, Serialize)]
pub struct ConversationListResponse {
pub conversations: Vec<super::postgres::ConversationSummary>,
pub has_more: bool, // client knows if there are more pages
pub conversations: Vec<crate::databases::postgres::chat::types::ConversationSummary>,
pub has_more: bool,
}
#[derive(Debug, Deserialize)]
@@ -246,6 +246,6 @@ pub struct MessageQuery {
#[derive(Debug, Serialize)]
pub struct MessageListResponse {
pub messages: Vec<super::postgres::MessageSummary>,
pub messages: Vec<crate::databases::postgres::chat::types::MessageSummary>,
pub has_more: bool,
}
-1
View File
@@ -1,3 +1,2 @@
pub mod api;
pub mod ollama;
pub mod postgres;
+1
View File
@@ -1,3 +1,4 @@
pub mod api;
pub mod databases;
pub mod dto;
pub mod middlewares;
+1
View File
@@ -1,3 +1,4 @@
mod api;
mod databases;
mod docs;
mod dto;
+3 -1
View File
@@ -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;
+27 -49
View File
@@ -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<Uuid>,
first_message: &str,
) -> Result<Uuid, (axum::http::StatusCode, String)> {
) -> Result<Uuid, ApiError> {
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<Uuid>,
) -> Result<Response, (axum::http::StatusCode, String)> {
) -> Result<Response, ApiError> {
// 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<Event, crate::providers::ollama::errors::OllamaError> {
@@ -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<Uuid>,
) -> Result<Response, (axum::http::StatusCode, String)> {
let mut response = state
.ollama
.chat_completions(&body)
.await
.map_err(into_http_response)?;
) -> Result<Response, ApiError> {
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<AppState>,
Extension(auth): Extension<Auth>,
Json(mut body): Json<api::ChatRequest>,
) -> Result<Response, (axum::http::StatusCode, String)> {
) -> Result<Response, ApiError> {
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<api::Message> = 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<u32>,
) -> Result<Uuid, errors::DbError> {
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<u32>,
) -> Result<Uuid, errors::DbError> {
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<AppState>,
Extension(auth): Extension<Auth>,
Query(params): Query<api::ConversationQuery>,
) -> Result<Json<api::ConversationListResponse>, (axum::http::StatusCode, Json<api::ErrorResponse>)>
{
) -> Result<Json<api::ConversationListResponse>, 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<Auth>,
Path(conversation_id): Path<Uuid>,
Query(params): Query<api::MessageQuery>,
) -> Result<Json<api::MessageListResponse>, (axum::http::StatusCode, Json<api::ErrorResponse>)> {
) -> Result<Json<api::MessageListResponse>, 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;