diff --git a/.sqlx/query-93e7661faaa9c5a28468b6b5a4bbb0e40510c081e2d4d5e99bdbf5c153112da2.json b/.sqlx/query-93e7661faaa9c5a28468b6b5a4bbb0e40510c081e2d4d5e99bdbf5c153112da2.json new file mode 100644 index 0000000..8da031b --- /dev/null +++ b/.sqlx/query-93e7661faaa9c5a28468b6b5a4bbb0e40510c081e2d4d5e99bdbf5c153112da2.json @@ -0,0 +1,22 @@ +{ + "db_name": "PostgreSQL", + "query": "INSERT INTO chat.conversation (user_id) VALUES ($1) RETURNING id", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + } + ], + "parameters": { + "Left": [ + "Uuid" + ] + }, + "nullable": [ + false + ] + }, + "hash": "93e7661faaa9c5a28468b6b5a4bbb0e40510c081e2d4d5e99bdbf5c153112da2" +} diff --git a/.sqlx/query-db88cb8d5b447e51fa1bd5ef4feed427e828be9029f4496ec63e704f44982f9e.json b/.sqlx/query-db88cb8d5b447e51fa1bd5ef4feed427e828be9029f4496ec63e704f44982f9e.json new file mode 100644 index 0000000..252033f --- /dev/null +++ b/.sqlx/query-db88cb8d5b447e51fa1bd5ef4feed427e828be9029f4496ec63e704f44982f9e.json @@ -0,0 +1,22 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT id FROM chat.conversation WHERE id = $1", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + } + ], + "parameters": { + "Left": [ + "Uuid" + ] + }, + "nullable": [ + false + ] + }, + "hash": "db88cb8d5b447e51fa1bd5ef4feed427e828be9029f4496ec63e704f44982f9e" +} diff --git a/src/databases/postgres/chat.rs b/src/databases/postgres/chat.rs index e69de29..70e7c56 100644 --- a/src/databases/postgres/chat.rs +++ b/src/databases/postgres/chat.rs @@ -0,0 +1,67 @@ +use super::errors::DbError; +use sqlx::{Acquire, PgPool}; +use uuid::Uuid; + +async fn validate_conversation<'e, E>(executor: E, conversation_id: Uuid) -> Result<(), DbError> +where + E: sqlx::Executor<'e, Database = sqlx::Postgres>, +{ + tracing::debug!("Validating Conversation"); + let rec = sqlx::query!( + r#"SELECT id FROM chat.conversation WHERE id = $1"#, + conversation_id + ) + .fetch_optional(executor) + .await?; + + match rec { + Some(_) => Ok(()), + None => Err(DbError::NotFound), + } +} + +async fn create_conversation<'e, E>(executor: E, user_id: Uuid) -> Result +where + E: sqlx::Executor<'e, Database = sqlx::Postgres>, +{ + tracing::debug!("Creating Conversation"); + + let rec = sqlx::query!( + r#"INSERT INTO chat.conversation (user_id) VALUES ($1) RETURNING id"#, + user_id + ) + .fetch_one(executor) + .await?; + + Ok(rec.id) +} + +pub async fn get_or_create_conversation( + pool: &PgPool, + conversation_id: Option, + user_id: Uuid, +) -> Result { + tracing::debug!("Testing conversation"); + + let mut tx = pool.begin().await?; + + let result = { + 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?; + + match conversation_id { + Some(id) => { + validate_conversation(&mut *conn, id).await?; + id + } + None => create_conversation(&mut *conn, user_id).await?, + } + }; + + tx.commit().await?; + Ok(result) +} diff --git a/src/databases/postgres/errors.rs b/src/databases/postgres/errors.rs index b974181..32e6128 100644 --- a/src/databases/postgres/errors.rs +++ b/src/databases/postgres/errors.rs @@ -8,14 +8,18 @@ pub enum DbError { #[error("database timeout")] Timeout, + + #[error("not found")] + NotFound, } -pub fn _into_http_response(e: DbError) -> (StatusCode, String) { +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 ef62612..e70c180 100644 --- a/src/databases/postgres/mod.rs +++ b/src/databases/postgres/mod.rs @@ -1,4 +1,5 @@ pub mod api_key; +pub mod chat; pub mod errors; pub mod pool; pub mod user; diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index 0bb00f5..581edc9 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -1,12 +1,14 @@ +use crate::middlewares::auth::middleware::Auth; use axum::{ Json, - extract::State, + extract::{Extension, State}, response::{ IntoResponse, Response, sse::{KeepAlive, Sse}, }, }; +use crate::databases::postgres::chat::get_or_create_conversation; use crate::dto::api; use crate::providers::ollama::errors::into_http_response; use crate::state::app_state::AppState; @@ -115,9 +117,25 @@ pub async fn completions( )] pub async fn chat_completions( State(state): State, + Extension(auth): Extension, Json(body): Json, ) -> Result { - tracing::debug!("Received /completion with body {:?}", body); + tracing::debug!("Received /chat/completion with body {:?}", body); + + if matches!(&auth, Auth::Jwt(_)) { + tracing::debug!( + "Is conversation_id existing: {:?}", + body.conversation_id.is_some() + ); + + let conversation_id = + get_or_create_conversation(&state.postgres, body.conversation_id, auth.user_id()) + .await + .map_err(crate::databases::postgres::errors::into_http_response)?; + + tracing::debug!("Using conversation_id: {:?}", conversation_id); + } + if body.base.stream { let stream = state .ollama