feat: conversation retrieveing
This commit is contained in:
+42
@@ -0,0 +1,42 @@
|
|||||||
|
{
|
||||||
|
"db_name": "PostgreSQL",
|
||||||
|
"query": "\n SELECT id, title, created_at, updated_at\n FROM chat.conversation\n WHERE user_id = $1\n AND ($2::timestamptz IS NULL OR updated_at < $2)\n ORDER BY updated_at DESC\n LIMIT $3\n ",
|
||||||
|
"describe": {
|
||||||
|
"columns": [
|
||||||
|
{
|
||||||
|
"ordinal": 0,
|
||||||
|
"name": "id",
|
||||||
|
"type_info": "Uuid"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ordinal": 1,
|
||||||
|
"name": "title",
|
||||||
|
"type_info": "Text"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ordinal": 2,
|
||||||
|
"name": "created_at",
|
||||||
|
"type_info": "Timestamptz"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ordinal": 3,
|
||||||
|
"name": "updated_at",
|
||||||
|
"type_info": "Timestamptz"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"parameters": {
|
||||||
|
"Left": [
|
||||||
|
"Uuid",
|
||||||
|
"Timestamptz",
|
||||||
|
"Int8"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"nullable": [
|
||||||
|
false,
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
false
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"hash": "6ea72b05912c3075a17f36dc71691fa5dde6934c29e96e9d98f4d65db14c1d88"
|
||||||
|
}
|
||||||
+54
@@ -0,0 +1,54 @@
|
|||||||
|
{
|
||||||
|
"db_name": "PostgreSQL",
|
||||||
|
"query": "\n SELECT id, parent_id, role, content, created_at, tokens\n FROM chat.message\n WHERE conversation_id = $1\n AND ($2::timestamptz IS NULL OR created_at < $2)\n ORDER BY created_at ASC\n LIMIT $3\n ",
|
||||||
|
"describe": {
|
||||||
|
"columns": [
|
||||||
|
{
|
||||||
|
"ordinal": 0,
|
||||||
|
"name": "id",
|
||||||
|
"type_info": "Uuid"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ordinal": 1,
|
||||||
|
"name": "parent_id",
|
||||||
|
"type_info": "Uuid"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ordinal": 2,
|
||||||
|
"name": "role",
|
||||||
|
"type_info": "Text"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ordinal": 3,
|
||||||
|
"name": "content",
|
||||||
|
"type_info": "Text"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ordinal": 4,
|
||||||
|
"name": "created_at",
|
||||||
|
"type_info": "Timestamptz"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"ordinal": 5,
|
||||||
|
"name": "tokens",
|
||||||
|
"type_info": "Int4"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"parameters": {
|
||||||
|
"Left": [
|
||||||
|
"Uuid",
|
||||||
|
"Timestamptz",
|
||||||
|
"Int8"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"nullable": [
|
||||||
|
false,
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
false,
|
||||||
|
false,
|
||||||
|
true
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"hash": "aa2722be6d9aec100c73bf65ef14ce3e979afe85a6d9d2b6e0005b9848efbd77"
|
||||||
|
}
|
||||||
@@ -266,4 +266,3 @@ This project turns Ollama into:
|
|||||||
|
|
||||||
# TODO
|
# TODO
|
||||||
- open api doc for bearer token
|
- open api doc for bearer token
|
||||||
- db
|
|
||||||
|
|||||||
@@ -1,7 +1,10 @@
|
|||||||
use super::errors::DbError;
|
use super::errors::DbError;
|
||||||
|
use crate::dto::postgres;
|
||||||
use sqlx::{Acquire, PgPool};
|
use sqlx::{Acquire, PgPool};
|
||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
|
|
||||||
|
// ---- Creation ----
|
||||||
|
|
||||||
async fn validate_conversation<'e, E>(executor: E, conversation_id: Uuid) -> Result<(), DbError>
|
async fn validate_conversation<'e, E>(executor: E, conversation_id: Uuid) -> Result<(), DbError>
|
||||||
where
|
where
|
||||||
E: sqlx::Executor<'e, Database = sqlx::Postgres>,
|
E: sqlx::Executor<'e, Database = sqlx::Postgres>,
|
||||||
@@ -114,7 +117,7 @@ pub async fn insert_message(
|
|||||||
parent_id: Option<Uuid>,
|
parent_id: Option<Uuid>,
|
||||||
role: MessageRole,
|
role: MessageRole,
|
||||||
content: &str,
|
content: &str,
|
||||||
tokens: Option<i32>,
|
tokens: Option<u32>,
|
||||||
) -> Result<Uuid, DbError> {
|
) -> Result<Uuid, DbError> {
|
||||||
let mut tx = pool.begin().await?;
|
let mut tx = pool.begin().await?;
|
||||||
let conn = tx.acquire().await?;
|
let conn = tx.acquire().await?;
|
||||||
@@ -134,7 +137,7 @@ pub async fn insert_message(
|
|||||||
parent_id,
|
parent_id,
|
||||||
role as MessageRole,
|
role as MessageRole,
|
||||||
content,
|
content,
|
||||||
tokens
|
tokens.unwrap_or(0) as i32
|
||||||
)
|
)
|
||||||
.fetch_one(&mut *conn)
|
.fetch_one(&mut *conn)
|
||||||
.await?;
|
.await?;
|
||||||
@@ -147,7 +150,7 @@ pub async fn update_message_tokens(
|
|||||||
pool: &PgPool,
|
pool: &PgPool,
|
||||||
user_id: Uuid,
|
user_id: Uuid,
|
||||||
message_id: Uuid,
|
message_id: Uuid,
|
||||||
tokens: i32,
|
tokens: u32,
|
||||||
) -> Result<(), DbError> {
|
) -> Result<(), DbError> {
|
||||||
let mut tx = pool.begin().await?;
|
let mut tx = pool.begin().await?;
|
||||||
let conn = tx.acquire().await?;
|
let conn = tx.acquire().await?;
|
||||||
@@ -159,7 +162,7 @@ pub async fn update_message_tokens(
|
|||||||
|
|
||||||
sqlx::query!(
|
sqlx::query!(
|
||||||
r#"UPDATE chat.message SET tokens = $1 WHERE id = $2"#,
|
r#"UPDATE chat.message SET tokens = $1 WHERE id = $2"#,
|
||||||
tokens,
|
tokens as i32,
|
||||||
message_id
|
message_id
|
||||||
)
|
)
|
||||||
.execute(&mut *conn)
|
.execute(&mut *conn)
|
||||||
@@ -168,3 +171,73 @@ pub async fn update_message_tokens(
|
|||||||
tx.commit().await?;
|
tx.commit().await?;
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ---- Fetching ----
|
||||||
|
|
||||||
|
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> {
|
||||||
|
let mut tx = pool.begin().await?;
|
||||||
|
|
||||||
|
sqlx::query("SELECT set_config('app.current_user_id', $1::text, true)")
|
||||||
|
.bind(user_id.to_string())
|
||||||
|
.execute(&mut *tx)
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
let rows = sqlx::query_as!(
|
||||||
|
postgres::ConversationSummary,
|
||||||
|
r#"
|
||||||
|
SELECT id, title, created_at, updated_at
|
||||||
|
FROM chat.conversation
|
||||||
|
WHERE user_id = $1
|
||||||
|
AND ($2::timestamptz IS NULL OR updated_at < $2)
|
||||||
|
ORDER BY updated_at DESC
|
||||||
|
LIMIT $3
|
||||||
|
"#,
|
||||||
|
user_id,
|
||||||
|
before,
|
||||||
|
limit
|
||||||
|
)
|
||||||
|
.fetch_all(&mut *tx)
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
tx.commit().await?;
|
||||||
|
Ok(rows)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn get_conversation_messages(
|
||||||
|
pool: &PgPool,
|
||||||
|
user_id: Uuid,
|
||||||
|
conversation_id: Uuid,
|
||||||
|
limit: i64,
|
||||||
|
before: Option<chrono::DateTime<chrono::Utc>>,
|
||||||
|
) -> Result<Vec<postgres::MessageSummary>, DbError> {
|
||||||
|
let mut conn = pool.acquire().await?;
|
||||||
|
|
||||||
|
sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#)
|
||||||
|
.bind(user_id.to_string())
|
||||||
|
.execute(&mut *conn)
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
let rows = sqlx::query_as!(
|
||||||
|
postgres::MessageSummary,
|
||||||
|
r#"
|
||||||
|
SELECT id, parent_id, role, content, created_at, tokens
|
||||||
|
FROM chat.message
|
||||||
|
WHERE conversation_id = $1
|
||||||
|
AND ($2::timestamptz IS NULL OR created_at < $2)
|
||||||
|
ORDER BY created_at ASC
|
||||||
|
LIMIT $3
|
||||||
|
"#,
|
||||||
|
conversation_id,
|
||||||
|
before,
|
||||||
|
limit
|
||||||
|
)
|
||||||
|
.fetch_all(&mut *conn)
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
Ok(rows)
|
||||||
|
}
|
||||||
|
|||||||
+1
-1
@@ -44,7 +44,7 @@ use crate::routes;
|
|||||||
api::ChatChoice,
|
api::ChatChoice,
|
||||||
api::ChatCompletionChunk,
|
api::ChatCompletionChunk,
|
||||||
api::ChatChunkChoice,
|
api::ChatChunkChoice,
|
||||||
api::ChatDelta,
|
api::Delta,
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
tags(
|
tags(
|
||||||
|
|||||||
+57
-6
@@ -65,6 +65,8 @@ pub struct BaseLLMRequest {
|
|||||||
pub stop: Option<Vec<String>>,
|
pub stop: Option<Vec<String>>,
|
||||||
|
|
||||||
pub keep_alive: Option<String>,
|
pub keep_alive: Option<String>,
|
||||||
|
|
||||||
|
pub context_depth: Option<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
||||||
@@ -167,7 +169,7 @@ pub struct ChatChoice {
|
|||||||
pub finish_reason: FinishReason,
|
pub finish_reason: FinishReason,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Serialize, ToSchema)]
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
pub struct ChatCompletionChunk {
|
pub struct ChatCompletionChunk {
|
||||||
pub id: String,
|
pub id: String,
|
||||||
pub object: String,
|
pub object: String,
|
||||||
@@ -175,17 +177,42 @@ pub struct ChatCompletionChunk {
|
|||||||
pub usage: Option<Usage>,
|
pub usage: Option<Usage>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Serialize, ToSchema)]
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
pub struct ChatChunkChoice {
|
pub struct ChatChunkChoice {
|
||||||
pub index: u32,
|
pub index: u32,
|
||||||
pub delta: ChatDelta,
|
pub delta: Delta,
|
||||||
pub finish_reason: Option<FinishReason>,
|
pub finish_reason: Option<FinishReason>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Serialize, ToSchema)]
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
pub struct ChatDelta {
|
pub struct Delta {
|
||||||
pub role: Option<Role>,
|
|
||||||
pub content: Option<String>,
|
pub content: Option<String>,
|
||||||
|
pub role: Option<Role>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
|
pub struct StartEventData {
|
||||||
|
pub conversation_id: Uuid,
|
||||||
|
pub created: u64,
|
||||||
|
pub id: Uuid,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
|
pub struct EndEventData {
|
||||||
|
pub created: u64,
|
||||||
|
pub id: Uuid,
|
||||||
|
pub usage: Usage,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
|
#[serde(tag = "type", content = "data")]
|
||||||
|
pub enum StreamEvent {
|
||||||
|
#[serde(rename = "start")]
|
||||||
|
Start(StartEventData),
|
||||||
|
#[serde(rename = "end")]
|
||||||
|
End(EndEventData),
|
||||||
|
#[serde(rename = "delta")]
|
||||||
|
Delta(ChatCompletionChunk),
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(serde::Deserialize)]
|
#[derive(serde::Deserialize)]
|
||||||
@@ -198,3 +225,27 @@ pub struct CreateApiKeyRequest {
|
|||||||
pub struct CreateApiKeyResponse {
|
pub struct CreateApiKeyResponse {
|
||||||
pub api_key: String, // ONLY returned once
|
pub api_key: String, // ONLY returned once
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
pub struct ConversationQuery {
|
||||||
|
pub limit: Option<i64>,
|
||||||
|
pub before: Option<chrono::DateTime<chrono::Utc>>, // cursor
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize)]
|
||||||
|
pub struct ConversationListResponse {
|
||||||
|
pub conversations: Vec<super::postgres::ConversationSummary>,
|
||||||
|
pub has_more: bool, // client knows if there are more pages
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
pub struct MessageQuery {
|
||||||
|
pub limit: Option<i64>,
|
||||||
|
pub before: Option<chrono::DateTime<chrono::Utc>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize)]
|
||||||
|
pub struct MessageListResponse {
|
||||||
|
pub messages: Vec<super::postgres::MessageSummary>,
|
||||||
|
pub has_more: bool,
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,2 +1,3 @@
|
|||||||
pub mod api;
|
pub mod api;
|
||||||
pub mod ollama;
|
pub mod ollama;
|
||||||
|
pub mod postgres;
|
||||||
|
|||||||
@@ -0,0 +1,20 @@
|
|||||||
|
use serde::Serialize;
|
||||||
|
use uuid::Uuid;
|
||||||
|
|
||||||
|
#[derive(Debug, sqlx::FromRow, Serialize)]
|
||||||
|
pub struct ConversationSummary {
|
||||||
|
pub id: Uuid,
|
||||||
|
pub title: Option<String>,
|
||||||
|
pub created_at: chrono::DateTime<chrono::Utc>,
|
||||||
|
pub updated_at: chrono::DateTime<chrono::Utc>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, sqlx::FromRow, Serialize)]
|
||||||
|
pub struct MessageSummary {
|
||||||
|
pub id: Uuid,
|
||||||
|
pub parent_id: Option<Uuid>,
|
||||||
|
pub role: String,
|
||||||
|
pub content: String,
|
||||||
|
pub created_at: chrono::DateTime<chrono::Utc>,
|
||||||
|
pub tokens: Option<i32>,
|
||||||
|
}
|
||||||
@@ -416,7 +416,7 @@ impl OllamaProvider {
|
|||||||
object: "chat.completion.chunk".to_string(),
|
object: "chat.completion.chunk".to_string(),
|
||||||
choices: vec![api::ChatChunkChoice {
|
choices: vec![api::ChatChunkChoice {
|
||||||
index: 0,
|
index: 0,
|
||||||
delta: api::ChatDelta {
|
delta: api::Delta {
|
||||||
role: Some(parsed.message.role),
|
role: Some(parsed.message.role),
|
||||||
content: Some(parsed.message.content),
|
content: Some(parsed.message.content),
|
||||||
},
|
},
|
||||||
|
|||||||
+320
-191
@@ -4,7 +4,7 @@ use crate::{
|
|||||||
};
|
};
|
||||||
use axum::{
|
use axum::{
|
||||||
Json,
|
Json,
|
||||||
extract::{Extension, State},
|
extract::{Extension, Path, Query, State},
|
||||||
response::{
|
response::{
|
||||||
IntoResponse, Response,
|
IntoResponse, Response,
|
||||||
sse::{Event, KeepAlive, Sse},
|
sse::{Event, KeepAlive, Sse},
|
||||||
@@ -15,7 +15,8 @@ use tokio_stream::StreamExt;
|
|||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
|
|
||||||
use crate::databases::postgres::chat::{
|
use crate::databases::postgres::chat::{
|
||||||
get_or_create_conversation, set_conversation_title, update_message_tokens,
|
get_conversation_messages, get_conversations_entries, get_or_create_conversation,
|
||||||
|
set_conversation_title, update_message_tokens,
|
||||||
};
|
};
|
||||||
use crate::databases::postgres::{chat, errors};
|
use crate::databases::postgres::{chat, errors};
|
||||||
use crate::dto::api;
|
use crate::dto::api;
|
||||||
@@ -83,6 +84,236 @@ pub async fn completions(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn ensure_conversation(
|
||||||
|
state: &AppState,
|
||||||
|
auth: &Auth,
|
||||||
|
conversation_id: Option<Uuid>,
|
||||||
|
first_message: &str,
|
||||||
|
) -> Result<Uuid, (axum::http::StatusCode, String)> {
|
||||||
|
let conversation_state =
|
||||||
|
get_or_create_conversation(&state.postgres, conversation_id, auth.user_id())
|
||||||
|
.await
|
||||||
|
.map_err(crate::databases::postgres::errors::into_http_response)?;
|
||||||
|
|
||||||
|
let id = match conversation_state {
|
||||||
|
ConversationState::Existing(uuid) => uuid,
|
||||||
|
ConversationState::Created(uuid) => {
|
||||||
|
let pool = state.postgres.clone();
|
||||||
|
let ollama = state.ollama.clone();
|
||||||
|
let user_id = auth.user_id();
|
||||||
|
let first_message = first_message.to_string();
|
||||||
|
|
||||||
|
tokio::spawn(async move {
|
||||||
|
let title = generate_conversation_title(&ollama, &first_message).await;
|
||||||
|
if let Err(e) = set_conversation_title(&pool, uuid, user_id, &title).await {
|
||||||
|
tracing::warn!(
|
||||||
|
conversation_id = %uuid,
|
||||||
|
error = %e,
|
||||||
|
"Failed to set conversation title"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
tracing::debug!(conversation_id = %uuid, %title, "generated conversation title");
|
||||||
|
});
|
||||||
|
|
||||||
|
uuid
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
Ok(id)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn handle_stream(
|
||||||
|
state: AppState,
|
||||||
|
auth: Auth,
|
||||||
|
body: api::ChatRequest,
|
||||||
|
conversation_id: Option<Uuid>,
|
||||||
|
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||||
|
// 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 plain_stream = stream.map(
|
||||||
|
|item| -> Result<Event, crate::providers::ollama::errors::OllamaError> {
|
||||||
|
match item {
|
||||||
|
Ok(chunk) => Ok(
|
||||||
|
Event::default().data(serde_json::to_string(&chunk).unwrap_or_default())
|
||||||
|
),
|
||||||
|
Err(e) => Err(e),
|
||||||
|
}
|
||||||
|
},
|
||||||
|
);
|
||||||
|
return Ok(Sse::new(plain_stream)
|
||||||
|
.keep_alive(KeepAlive::default())
|
||||||
|
.into_response());
|
||||||
|
};
|
||||||
|
|
||||||
|
// From here conv_id is a plain Uuid — all variables stay in scope
|
||||||
|
let user_msg_id = log_user_message(
|
||||||
|
&state.postgres,
|
||||||
|
auth.user_id(),
|
||||||
|
conv_id,
|
||||||
|
body.parent_id,
|
||||||
|
body.messages
|
||||||
|
.last()
|
||||||
|
.map(|m| m.content.as_str())
|
||||||
|
.unwrap_or(""),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.map_err(errors::into_http_response)?;
|
||||||
|
|
||||||
|
let start_event = api::StreamEvent::Start(api::StartEventData {
|
||||||
|
conversation_id: conv_id,
|
||||||
|
created: chrono::Utc::now().timestamp() as u64,
|
||||||
|
id: user_msg_id,
|
||||||
|
});
|
||||||
|
|
||||||
|
let (tx, rx) = tokio::sync::mpsc::channel::<
|
||||||
|
Result<Event, crate::providers::ollama::errors::OllamaError>,
|
||||||
|
>(32);
|
||||||
|
|
||||||
|
// Send start event immediately, before Ollama is contacted
|
||||||
|
let _ = tx
|
||||||
|
.send(Ok(Event::default()
|
||||||
|
.event("metadata")
|
||||||
|
.data(serde_json::to_string(&start_event).unwrap())))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let pool = state.postgres.clone();
|
||||||
|
let user_id = auth.user_id();
|
||||||
|
|
||||||
|
tokio::spawn(async move {
|
||||||
|
// Ollama called inside spawn — start event already queued
|
||||||
|
let stream = match state.ollama.chat_completions_stream(&body).await {
|
||||||
|
Ok(s) => s,
|
||||||
|
Err(e) => {
|
||||||
|
let _ = tx.send(Err(e)).await;
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut stream = stream;
|
||||||
|
let mut accumulated = String::new();
|
||||||
|
|
||||||
|
while let Some(item) = futures::StreamExt::next(&mut stream).await {
|
||||||
|
match item {
|
||||||
|
Ok(chunk) => {
|
||||||
|
let is_done = chunk.choices[0].finish_reason == Some(api::FinishReason::Stop);
|
||||||
|
|
||||||
|
if let Some(content) = chunk.choices[0].delta.content.as_ref() {
|
||||||
|
accumulated.push_str(content);
|
||||||
|
}
|
||||||
|
|
||||||
|
if is_done {
|
||||||
|
let prompt_tokens = chunk.usage.as_ref().map(|u| u.prompt_tokens);
|
||||||
|
let completion_tokens = chunk.usage.as_ref().map(|u| u.completion_tokens);
|
||||||
|
|
||||||
|
if let Some(pt) = prompt_tokens {
|
||||||
|
let _ = update_message_tokens(&pool, user_id, user_msg_id, pt).await;
|
||||||
|
}
|
||||||
|
|
||||||
|
let assistant_msg_id = log_assistant_message(
|
||||||
|
&pool,
|
||||||
|
user_id,
|
||||||
|
conv_id,
|
||||||
|
user_msg_id,
|
||||||
|
&accumulated,
|
||||||
|
completion_tokens,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
if let Ok(msg_id) = assistant_msg_id {
|
||||||
|
let end_event = api::StreamEvent::End(api::EndEventData {
|
||||||
|
usage: api::Usage {
|
||||||
|
prompt_tokens: prompt_tokens.unwrap_or(0),
|
||||||
|
completion_tokens: completion_tokens.unwrap_or(0),
|
||||||
|
total_tokens: chunk
|
||||||
|
.usage
|
||||||
|
.as_ref()
|
||||||
|
.map(|u| u.total_tokens)
|
||||||
|
.unwrap_or(0),
|
||||||
|
},
|
||||||
|
id: msg_id,
|
||||||
|
created: chrono::Utc::now().timestamp() as u64,
|
||||||
|
});
|
||||||
|
|
||||||
|
let _ = tx
|
||||||
|
.send(Ok(Event::default()
|
||||||
|
.event("metadata")
|
||||||
|
.data(serde_json::to_string(&end_event).unwrap())))
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
let data = api::StreamEvent::Delta(chunk);
|
||||||
|
let json = serde_json::to_string(&data).unwrap();
|
||||||
|
if tx.send(Ok(Event::default().data(json))).await.is_err() {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
let _ = tx.send(Err(e)).await;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
Ok(Sse::new(tokio_stream::wrappers::ReceiverStream::new(rx))
|
||||||
|
.keep_alive(KeepAlive::default())
|
||||||
|
.into_response())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn handle_non_stream(
|
||||||
|
state: AppState,
|
||||||
|
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)?;
|
||||||
|
|
||||||
|
response.conversation_id = conversation_id;
|
||||||
|
|
||||||
|
if let Some(conversation_id) = conversation_id {
|
||||||
|
let user_msg_id = log_user_message(
|
||||||
|
&state.postgres,
|
||||||
|
auth.user_id(),
|
||||||
|
conversation_id,
|
||||||
|
body.parent_id,
|
||||||
|
body.messages
|
||||||
|
.last()
|
||||||
|
.map(|m| m.content.as_str())
|
||||||
|
.unwrap_or(""),
|
||||||
|
response.usage.map(|u| u.prompt_tokens),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.map_err(errors::into_http_response)?;
|
||||||
|
|
||||||
|
log_assistant_message(
|
||||||
|
&state.postgres,
|
||||||
|
auth.user_id(),
|
||||||
|
conversation_id,
|
||||||
|
user_msg_id,
|
||||||
|
&response.choices[0].message.content,
|
||||||
|
response.usage.map(|u| u.completion_tokens),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.map_err(errors::into_http_response)?;
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(Json(response).into_response())
|
||||||
|
}
|
||||||
|
|
||||||
#[utoipa::path(
|
#[utoipa::path(
|
||||||
post,
|
post,
|
||||||
path = "/chat/completions",
|
path = "/chat/completions",
|
||||||
@@ -128,210 +359,53 @@ 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>,
|
Extension(auth): Extension<Auth>,
|
||||||
Json(body): Json<api::ChatRequest>,
|
Json(mut body): Json<api::ChatRequest>,
|
||||||
) -> Result<Response, (axum::http::StatusCode, String)> {
|
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||||
tracing::debug!("Received /chat/completion with body {:?}", body);
|
tracing::debug!("Received /chat/completion with body {:?}", body);
|
||||||
|
|
||||||
// Get current conversation or create it
|
|
||||||
let conversation_id = if matches!(&auth, Auth::Jwt(_)) {
|
let conversation_id = if matches!(&auth, Auth::Jwt(_)) {
|
||||||
tracing::debug!(
|
let first_message = body.messages[0].content.clone();
|
||||||
"Is conversation_id existing: {:?}",
|
let conv_id =
|
||||||
body.conversation_id.is_some()
|
ensure_conversation(&state, &auth, body.conversation_id, &first_message).await?;
|
||||||
);
|
|
||||||
|
|
||||||
let conversation_state =
|
if let Some(depth) = body.base.context_depth {
|
||||||
get_or_create_conversation(&state.postgres, body.conversation_id, auth.user_id())
|
dbg!({ depth });
|
||||||
|
if depth > 0 {
|
||||||
|
let history = get_conversation_messages(
|
||||||
|
&state.postgres,
|
||||||
|
auth.user_id(),
|
||||||
|
conv_id,
|
||||||
|
depth as i64,
|
||||||
|
None, // no cursor — fetch the most recent N messages
|
||||||
|
)
|
||||||
.await
|
.await
|
||||||
.map_err(crate::databases::postgres::errors::into_http_response)?;
|
.map_err(errors::into_http_response)?;
|
||||||
|
|
||||||
let id = match &conversation_state {
|
// Map MessageSummary → api::Message and prepend to the outgoing request
|
||||||
ConversationState::Existing(uuid) => *uuid,
|
let history_messages: Vec<api::Message> = history
|
||||||
ConversationState::Created(uuid) => {
|
.into_iter()
|
||||||
let conversation_id = *uuid;
|
.map(|m| api::Message {
|
||||||
let title =
|
role: api::Role::Assistant, // TODO
|
||||||
generate_conversation_title(&state.ollama, body.messages[0].content.as_str())
|
content: m.content,
|
||||||
.await;
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
if let Err(e) =
|
// body.messages = [history_messages, body.messages].concat();
|
||||||
set_conversation_title(&state.postgres, conversation_id, auth.user_id(), &title)
|
body.messages.splice(0..0, history_messages);
|
||||||
.await
|
}
|
||||||
{
|
|
||||||
tracing::warn!(
|
|
||||||
conversation_id = %conversation_id,
|
|
||||||
error = %e,
|
|
||||||
"Failed to set conversation title, continuing with no title"
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
tracing::debug!(conversation_id = %conversation_id, %title, "generated conversation title");
|
dbg!("{:?}", &body.messages);
|
||||||
conversation_id
|
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
tracing::debug!("Using conversation_id: {:?}", id);
|
Some(conv_id)
|
||||||
Some(id)
|
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
|
|
||||||
if body.base.stream {
|
if body.base.stream {
|
||||||
let stream = state
|
handle_stream(state, auth, body, conversation_id).await
|
||||||
.ollama
|
|
||||||
.chat_completions_stream(&body)
|
|
||||||
.await
|
|
||||||
.map_err(into_http_response)?;
|
|
||||||
|
|
||||||
if let Some(conversation_id) = conversation_id {
|
|
||||||
let user_msg_id = log_user_message(
|
|
||||||
&state.postgres,
|
|
||||||
auth.user_id(),
|
|
||||||
conversation_id,
|
|
||||||
body.parent_id,
|
|
||||||
body.messages
|
|
||||||
.last()
|
|
||||||
.map(|m| m.content.as_str())
|
|
||||||
.unwrap_or(""),
|
|
||||||
None, // not known at this stage
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(errors::into_http_response)?;
|
|
||||||
|
|
||||||
let start_event = serde_json::json!({
|
|
||||||
"type": "start",
|
|
||||||
"model": body.base.model,
|
|
||||||
"conversation_id": conversation_id,
|
|
||||||
"created": chrono::Utc::now().timestamp() as u64,
|
|
||||||
});
|
|
||||||
|
|
||||||
let start_stream = futures::stream::once(async move {
|
|
||||||
Ok::<Event, crate::providers::ollama::errors::OllamaError>(
|
|
||||||
Event::default()
|
|
||||||
.event("metadata")
|
|
||||||
.data(start_event.to_string()),
|
|
||||||
)
|
|
||||||
});
|
|
||||||
|
|
||||||
let mut accumulated = String::new();
|
|
||||||
|
|
||||||
let wrapped_stream = stream.map(move |item| match item {
|
|
||||||
Ok(chunk) => {
|
|
||||||
let is_done = chunk.choices[0].finish_reason == Some(api::FinishReason::Stop);
|
|
||||||
|
|
||||||
if let Some(content) = chunk.choices[0].delta.content.as_ref() {
|
|
||||||
accumulated.push_str(content);
|
|
||||||
}
|
|
||||||
|
|
||||||
if is_done {
|
|
||||||
let content = accumulated.clone();
|
|
||||||
let pool = state.postgres.clone();
|
|
||||||
let user_id = auth.user_id();
|
|
||||||
let prompt_tokens = chunk.usage.as_ref().map(|u| u.prompt_tokens as i32);
|
|
||||||
let completion_tokens =
|
|
||||||
chunk.usage.as_ref().map(|u| u.completion_tokens as i32);
|
|
||||||
|
|
||||||
let end_event = serde_json::json!({
|
|
||||||
"type": "end",
|
|
||||||
"usage": {
|
|
||||||
"prompt_tokens": prompt_tokens,
|
|
||||||
"completion_tokens": completion_tokens,
|
|
||||||
"total_tokens": chunk.usage.as_ref().map(|u| u.total_tokens),
|
|
||||||
}
|
|
||||||
});
|
|
||||||
|
|
||||||
tokio::spawn(async move {
|
|
||||||
if let Some(prompt_tokens) = prompt_tokens {
|
|
||||||
let _ = update_message_tokens(
|
|
||||||
&pool,
|
|
||||||
user_id,
|
|
||||||
user_msg_id,
|
|
||||||
prompt_tokens,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
|
||||||
|
|
||||||
let _ = log_assistant_message(
|
|
||||||
&pool,
|
|
||||||
user_id,
|
|
||||||
conversation_id,
|
|
||||||
user_msg_id,
|
|
||||||
&content,
|
|
||||||
completion_tokens,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
});
|
|
||||||
|
|
||||||
return Ok(Event::default()
|
|
||||||
.event("metadata")
|
|
||||||
.data(end_event.to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
let data = serde_json::to_string(&chunk).unwrap_or_default();
|
|
||||||
Ok(Event::default().data(data))
|
|
||||||
}
|
|
||||||
Err(e) => Err(e),
|
|
||||||
});
|
|
||||||
|
|
||||||
let full_stream = start_stream.chain(wrapped_stream);
|
|
||||||
|
|
||||||
Ok(Sse::new(full_stream)
|
|
||||||
.keep_alive(KeepAlive::default())
|
|
||||||
.into_response())
|
|
||||||
} else {
|
} else {
|
||||||
let plain_stream = stream.map(
|
handle_non_stream(state, auth, body, conversation_id).await
|
||||||
|item| -> Result<Event, crate::providers::ollama::errors::OllamaError> {
|
|
||||||
match item {
|
|
||||||
Ok(chunk) => {
|
|
||||||
let data = serde_json::to_string(&chunk).unwrap_or_default();
|
|
||||||
Ok(Event::default().data(data))
|
|
||||||
}
|
|
||||||
Err(e) => Err(e),
|
|
||||||
}
|
|
||||||
},
|
|
||||||
);
|
|
||||||
|
|
||||||
Ok(Sse::new(plain_stream)
|
|
||||||
.keep_alive(KeepAlive::default())
|
|
||||||
.into_response())
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
let mut response = state
|
|
||||||
.ollama
|
|
||||||
.chat_completions(&body)
|
|
||||||
.await
|
|
||||||
.map_err(into_http_response)?;
|
|
||||||
|
|
||||||
response.conversation_id = conversation_id;
|
|
||||||
|
|
||||||
if let Some(conversation_id) = conversation_id {
|
|
||||||
// log user message, get its id as parent for assistant
|
|
||||||
let user_msg_id = log_user_message(
|
|
||||||
&state.postgres,
|
|
||||||
auth.user_id(),
|
|
||||||
conversation_id,
|
|
||||||
body.parent_id,
|
|
||||||
body.messages
|
|
||||||
.last()
|
|
||||||
.map(|m| m.content.as_str())
|
|
||||||
.unwrap_or(""),
|
|
||||||
response.usage.map(|u| u.prompt_tokens as i32),
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(errors::into_http_response)?;
|
|
||||||
|
|
||||||
// log assistant response
|
|
||||||
log_assistant_message(
|
|
||||||
&state.postgres,
|
|
||||||
auth.user_id(),
|
|
||||||
conversation_id,
|
|
||||||
user_msg_id,
|
|
||||||
&response.choices[0].message.content,
|
|
||||||
response.usage.map(|u| u.completion_tokens as i32),
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(errors::into_http_response)?;
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(Json(response).into_response())
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -362,7 +436,7 @@ async fn log_user_message(
|
|||||||
conversation_id: Uuid,
|
conversation_id: Uuid,
|
||||||
parent_id: Option<Uuid>,
|
parent_id: Option<Uuid>,
|
||||||
content: &str,
|
content: &str,
|
||||||
tokens: Option<i32>,
|
tokens: Option<u32>,
|
||||||
) -> Result<Uuid, errors::DbError> {
|
) -> Result<Uuid, errors::DbError> {
|
||||||
chat::insert_message(
|
chat::insert_message(
|
||||||
pool,
|
pool,
|
||||||
@@ -382,7 +456,7 @@ async fn log_assistant_message(
|
|||||||
conversation_id: Uuid,
|
conversation_id: Uuid,
|
||||||
parent_id: Uuid,
|
parent_id: Uuid,
|
||||||
content: &str,
|
content: &str,
|
||||||
tokens: Option<i32>,
|
tokens: Option<u32>,
|
||||||
) -> Result<Uuid, errors::DbError> {
|
) -> Result<Uuid, errors::DbError> {
|
||||||
chat::insert_message(
|
chat::insert_message(
|
||||||
pool,
|
pool,
|
||||||
@@ -395,3 +469,58 @@ async fn log_assistant_message(
|
|||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Conversation retrieveing
|
||||||
|
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>)>
|
||||||
|
{
|
||||||
|
tracing::debug!("Conversation hit: {:?}", auth);
|
||||||
|
|
||||||
|
let conversations = get_conversations_entries(
|
||||||
|
&state.postgres,
|
||||||
|
auth.user_id(),
|
||||||
|
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)))
|
||||||
|
})?;
|
||||||
|
|
||||||
|
let has_more = conversations.len() == params.limit.unwrap_or(20) as usize;
|
||||||
|
|
||||||
|
Ok(Json(api::ConversationListResponse {
|
||||||
|
conversations,
|
||||||
|
has_more,
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn get_messages(
|
||||||
|
State(state): State<AppState>,
|
||||||
|
Extension(auth): Extension<Auth>,
|
||||||
|
Path(conversation_id): Path<Uuid>,
|
||||||
|
Query(params): Query<api::MessageQuery>,
|
||||||
|
) -> Result<Json<api::MessageListResponse>, (axum::http::StatusCode, Json<api::ErrorResponse>)> {
|
||||||
|
tracing::debug!("Messages hit: {:?}", auth);
|
||||||
|
|
||||||
|
let messages = get_conversation_messages(
|
||||||
|
&state.postgres,
|
||||||
|
auth.user_id(),
|
||||||
|
conversation_id,
|
||||||
|
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)))
|
||||||
|
})?;
|
||||||
|
|
||||||
|
let has_more = messages.len() == params.limit.unwrap_or(50) as usize;
|
||||||
|
|
||||||
|
Ok(Json(api::MessageListResponse { messages, has_more }))
|
||||||
|
}
|
||||||
|
|||||||
@@ -30,6 +30,11 @@ pub fn protected_router() -> Router<AppState> {
|
|||||||
require_roles(req, next, None, Some("admin"))
|
require_roles(req, next, None, Some("admin"))
|
||||||
})),
|
})),
|
||||||
)
|
)
|
||||||
|
.route("/conversations", get(chat::get_conversations))
|
||||||
|
.route(
|
||||||
|
"/conversations/{conversation_id}/messages",
|
||||||
|
get(chat::get_messages),
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn router(state: AppState) -> Router<AppState> {
|
pub fn router(state: AppState) -> Router<AppState> {
|
||||||
|
|||||||
Reference in New Issue
Block a user