feat: create conversation
This commit is contained in:
+22
@@ -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"
|
||||||
|
}
|
||||||
+22
@@ -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"
|
||||||
|
}
|
||||||
@@ -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")]
|
#[error("database timeout")]
|
||||||
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 {
|
match e {
|
||||||
DbError::Connection(_) => (
|
DbError::Connection(_) => (
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
StatusCode::INTERNAL_SERVER_ERROR,
|
||||||
"database connection error".to_string(),
|
"database connection error".to_string(),
|
||||||
),
|
),
|
||||||
DbError::Timeout => (StatusCode::REQUEST_TIMEOUT, "database timeout".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 api_key;
|
||||||
|
pub mod chat;
|
||||||
pub mod errors;
|
pub mod errors;
|
||||||
pub mod pool;
|
pub mod pool;
|
||||||
pub mod user;
|
pub mod user;
|
||||||
|
|||||||
+20
-2
@@ -1,12 +1,14 @@
|
|||||||
|
use crate::middlewares::auth::middleware::Auth;
|
||||||
use axum::{
|
use axum::{
|
||||||
Json,
|
Json,
|
||||||
extract::State,
|
extract::{Extension, State},
|
||||||
response::{
|
response::{
|
||||||
IntoResponse, Response,
|
IntoResponse, Response,
|
||||||
sse::{KeepAlive, Sse},
|
sse::{KeepAlive, Sse},
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
use crate::databases::postgres::chat::get_or_create_conversation;
|
||||||
use crate::dto::api;
|
use crate::dto::api;
|
||||||
use crate::providers::ollama::errors::into_http_response;
|
use crate::providers::ollama::errors::into_http_response;
|
||||||
use crate::state::app_state::AppState;
|
use crate::state::app_state::AppState;
|
||||||
@@ -115,9 +117,25 @@ pub async fn completions(
|
|||||||
)]
|
)]
|
||||||
pub async fn chat_completions(
|
pub async fn chat_completions(
|
||||||
State(state): State<AppState>,
|
State(state): State<AppState>,
|
||||||
|
Extension(auth): Extension<Auth>,
|
||||||
Json(body): Json<api::ChatRequest>,
|
Json(body): Json<api::ChatRequest>,
|
||||||
) -> Result<Response, (axum::http::StatusCode, String)> {
|
) -> 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 {
|
if body.base.stream {
|
||||||
let stream = state
|
let stream = state
|
||||||
.ollama
|
.ollama
|
||||||
|
|||||||
Reference in New Issue
Block a user