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
|
||||
- open api doc for bearer token
|
||||
- db
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
use super::errors::DbError;
|
||||
use crate::dto::postgres;
|
||||
use sqlx::{Acquire, PgPool};
|
||||
use uuid::Uuid;
|
||||
|
||||
// ---- Creation ----
|
||||
|
||||
async fn validate_conversation<'e, E>(executor: E, conversation_id: Uuid) -> Result<(), DbError>
|
||||
where
|
||||
E: sqlx::Executor<'e, Database = sqlx::Postgres>,
|
||||
@@ -114,7 +117,7 @@ pub async fn insert_message(
|
||||
parent_id: Option<Uuid>,
|
||||
role: MessageRole,
|
||||
content: &str,
|
||||
tokens: Option<i32>,
|
||||
tokens: Option<u32>,
|
||||
) -> Result<Uuid, DbError> {
|
||||
let mut tx = pool.begin().await?;
|
||||
let conn = tx.acquire().await?;
|
||||
@@ -134,7 +137,7 @@ pub async fn insert_message(
|
||||
parent_id,
|
||||
role as MessageRole,
|
||||
content,
|
||||
tokens
|
||||
tokens.unwrap_or(0) as i32
|
||||
)
|
||||
.fetch_one(&mut *conn)
|
||||
.await?;
|
||||
@@ -147,7 +150,7 @@ pub async fn update_message_tokens(
|
||||
pool: &PgPool,
|
||||
user_id: Uuid,
|
||||
message_id: Uuid,
|
||||
tokens: i32,
|
||||
tokens: u32,
|
||||
) -> Result<(), DbError> {
|
||||
let mut tx = pool.begin().await?;
|
||||
let conn = tx.acquire().await?;
|
||||
@@ -159,7 +162,7 @@ pub async fn update_message_tokens(
|
||||
|
||||
sqlx::query!(
|
||||
r#"UPDATE chat.message SET tokens = $1 WHERE id = $2"#,
|
||||
tokens,
|
||||
tokens as i32,
|
||||
message_id
|
||||
)
|
||||
.execute(&mut *conn)
|
||||
@@ -168,3 +171,73 @@ pub async fn update_message_tokens(
|
||||
tx.commit().await?;
|
||||
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::ChatCompletionChunk,
|
||||
api::ChatChunkChoice,
|
||||
api::ChatDelta,
|
||||
api::Delta,
|
||||
)
|
||||
),
|
||||
tags(
|
||||
|
||||
+57
-6
@@ -65,6 +65,8 @@ pub struct BaseLLMRequest {
|
||||
pub stop: Option<Vec<String>>,
|
||||
|
||||
pub keep_alive: Option<String>,
|
||||
|
||||
pub context_depth: Option<u32>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
||||
@@ -167,7 +169,7 @@ pub struct ChatChoice {
|
||||
pub finish_reason: FinishReason,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
pub struct ChatCompletionChunk {
|
||||
pub id: String,
|
||||
pub object: String,
|
||||
@@ -175,17 +177,42 @@ pub struct ChatCompletionChunk {
|
||||
pub usage: Option<Usage>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
pub struct ChatChunkChoice {
|
||||
pub index: u32,
|
||||
pub delta: ChatDelta,
|
||||
pub delta: Delta,
|
||||
pub finish_reason: Option<FinishReason>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ChatDelta {
|
||||
pub role: Option<Role>,
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
pub struct Delta {
|
||||
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)]
|
||||
@@ -198,3 +225,27 @@ pub struct CreateApiKeyRequest {
|
||||
pub struct CreateApiKeyResponse {
|
||||
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 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(),
|
||||
choices: vec![api::ChatChunkChoice {
|
||||
index: 0,
|
||||
delta: api::ChatDelta {
|
||||
delta: api::Delta {
|
||||
role: Some(parsed.message.role),
|
||||
content: Some(parsed.message.content),
|
||||
},
|
||||
|
||||
+320
-191
@@ -4,7 +4,7 @@ use crate::{
|
||||
};
|
||||
use axum::{
|
||||
Json,
|
||||
extract::{Extension, State},
|
||||
extract::{Extension, Path, Query, State},
|
||||
response::{
|
||||
IntoResponse, Response,
|
||||
sse::{Event, KeepAlive, Sse},
|
||||
@@ -15,7 +15,8 @@ use tokio_stream::StreamExt;
|
||||
use uuid::Uuid;
|
||||
|
||||
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::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(
|
||||
post,
|
||||
path = "/chat/completions",
|
||||
@@ -128,210 +359,53 @@ pub async fn completions(
|
||||
pub async fn chat_completions(
|
||||
State(state): State<AppState>,
|
||||
Extension(auth): Extension<Auth>,
|
||||
Json(body): Json<api::ChatRequest>,
|
||||
Json(mut body): Json<api::ChatRequest>,
|
||||
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||
tracing::debug!("Received /chat/completion with body {:?}", body);
|
||||
|
||||
// Get current conversation or create it
|
||||
let conversation_id = if matches!(&auth, Auth::Jwt(_)) {
|
||||
tracing::debug!(
|
||||
"Is conversation_id existing: {:?}",
|
||||
body.conversation_id.is_some()
|
||||
);
|
||||
let first_message = body.messages[0].content.clone();
|
||||
let conv_id =
|
||||
ensure_conversation(&state, &auth, body.conversation_id, &first_message).await?;
|
||||
|
||||
let conversation_state =
|
||||
get_or_create_conversation(&state.postgres, body.conversation_id, auth.user_id())
|
||||
if let Some(depth) = body.base.context_depth {
|
||||
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
|
||||
.map_err(crate::databases::postgres::errors::into_http_response)?;
|
||||
.map_err(errors::into_http_response)?;
|
||||
|
||||
let id = match &conversation_state {
|
||||
ConversationState::Existing(uuid) => *uuid,
|
||||
ConversationState::Created(uuid) => {
|
||||
let conversation_id = *uuid;
|
||||
let title =
|
||||
generate_conversation_title(&state.ollama, body.messages[0].content.as_str())
|
||||
.await;
|
||||
// Map MessageSummary → api::Message and prepend to the outgoing request
|
||||
let history_messages: Vec<api::Message> = history
|
||||
.into_iter()
|
||||
.map(|m| api::Message {
|
||||
role: api::Role::Assistant, // TODO
|
||||
content: m.content,
|
||||
})
|
||||
.collect();
|
||||
|
||||
if let Err(e) =
|
||||
set_conversation_title(&state.postgres, conversation_id, auth.user_id(), &title)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(
|
||||
conversation_id = %conversation_id,
|
||||
error = %e,
|
||||
"Failed to set conversation title, continuing with no title"
|
||||
);
|
||||
// body.messages = [history_messages, body.messages].concat();
|
||||
body.messages.splice(0..0, history_messages);
|
||||
}
|
||||
}
|
||||
|
||||
tracing::debug!(conversation_id = %conversation_id, %title, "generated conversation title");
|
||||
conversation_id
|
||||
}
|
||||
};
|
||||
dbg!("{:?}", &body.messages);
|
||||
|
||||
tracing::debug!("Using conversation_id: {:?}", id);
|
||||
Some(id)
|
||||
Some(conv_id)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
if body.base.stream {
|
||||
let stream = state
|
||||
.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())
|
||||
handle_stream(state, auth, body, conversation_id).await
|
||||
} else {
|
||||
let plain_stream = stream.map(
|
||||
|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())
|
||||
handle_non_stream(state, auth, body, conversation_id).await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -362,7 +436,7 @@ async fn log_user_message(
|
||||
conversation_id: Uuid,
|
||||
parent_id: Option<Uuid>,
|
||||
content: &str,
|
||||
tokens: Option<i32>,
|
||||
tokens: Option<u32>,
|
||||
) -> Result<Uuid, errors::DbError> {
|
||||
chat::insert_message(
|
||||
pool,
|
||||
@@ -382,7 +456,7 @@ async fn log_assistant_message(
|
||||
conversation_id: Uuid,
|
||||
parent_id: Uuid,
|
||||
content: &str,
|
||||
tokens: Option<i32>,
|
||||
tokens: Option<u32>,
|
||||
) -> Result<Uuid, errors::DbError> {
|
||||
chat::insert_message(
|
||||
pool,
|
||||
@@ -395,3 +469,58 @@ async fn log_assistant_message(
|
||||
)
|
||||
.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"))
|
||||
})),
|
||||
)
|
||||
.route("/conversations", get(chat::get_conversations))
|
||||
.route(
|
||||
"/conversations/{conversation_id}/messages",
|
||||
get(chat::get_messages),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn router(state: AppState) -> Router<AppState> {
|
||||
|
||||
Reference in New Issue
Block a user