feat: create conversation
This commit is contained in:
@@ -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<Uuid, DbError>
|
||||
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<Uuid>,
|
||||
user_id: Uuid,
|
||||
) -> Result<Uuid, DbError> {
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -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()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
pub mod api_key;
|
||||
pub mod chat;
|
||||
pub mod errors;
|
||||
pub mod pool;
|
||||
pub mod user;
|
||||
|
||||
+20
-2
@@ -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<AppState>,
|
||||
Extension(auth): Extension<Auth>,
|
||||
Json(body): Json<api::ChatRequest>,
|
||||
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user