fix: streaming
CI / Rust CI (push) Successful in 5m19s

This commit is contained in:
2026-06-05 18:16:48 +02:00
parent 28352d8bdb
commit 77bd729d95
16 changed files with 408 additions and 385 deletions
+4
View File
@@ -275,6 +275,10 @@ This project turns Ollama into:
👉 A controllable model runtime 👉 A controllable model runtime
👉 A foundation for a full LLM gateway 👉 A foundation for a full LLM gateway
# Bugs
## Ollama
- When model answer onoly with 1 tokens, answer is empty
# TODO # TODO
- open api doc for bearer token - open api doc for bearer token
- Unify check before sending to ollama payload - Unify check before sending to ollama payload
+15 -15
View File
@@ -24,24 +24,24 @@ use crate::api::routes;
components( components(
schemas( schemas(
api::errors::ErrorResponse, api::errors::ErrorResponse,
api::types::ModelsResponse, api::types::ApiModelsResponse,
api::types::ModelInfo, api::types::ApiModelInfo,
api::types::LoadModelResponse, api::types::ApiLoadModelResponse,
api::types::LoadModelRequest, api::types::ApiLoadModelRequest,
api::types::UnloadModelResponse, api::types::ApiUnloadModelResponse,
api::types::LLMOptions, api::types::ApiLlmOptions,
api::types::CompletionRequest, api::types::ApiCompletionRequest,
api::types::CompletionObject, api::types::ApiCompletionObject,
api::types::FinishReason, api::types::ApiFinishReason,
api::types::CompletionResponse, api::types::ApiCompletionResponse,
api::types::Choice, api::types::Choice,
api::types::Usage, api::types::Usage,
api::types::CompletionChunk, api::types::CompletionChunk,
api::types::ChatRequest, api::types::ApiChatRequest,
api::types::Message, api::types::ApiMessage,
api::types::Role, api::types::ApiRole,
api::types::ChatResponse, api::types::ApiChatResponse,
api::types::ChatChoice, api::types::ApiChatChoice,
api::types::ChatCompletionChunk, api::types::ChatCompletionChunk,
api::types::ChatChunkChoice, api::types::ChatChunkChoice,
api::types::Delta, api::types::Delta,
+2
View File
@@ -17,6 +17,8 @@ pub struct ErrorResponse {
pub code: String, pub code: String,
} }
#[derive(Debug, thiserror::Error)]
#[error("{message}")]
pub struct ApiError { pub struct ApiError {
pub status: StatusCode, pub status: StatusCode,
pub code: &'static str, pub code: &'static str,
+8 -2
View File
@@ -58,8 +58,14 @@ pub async fn auth_middleware(
}; };
match handle_auth(&state, req, next, auth).await { match handle_auth(&state, req, next, auth).await {
Ok(response) => response, Ok(response) => {
Err(err) => err.into_response(), tracing::debug!("User authentified");
response
}
Err(err) => {
tracing::debug!("Error during authentification {:?}", err);
err.into_response()
}
} }
} }
+68 -159
View File
@@ -20,7 +20,7 @@ use uuid::Uuid;
path = "/completions", path = "/completions",
tag = "chat", tag = "chat",
request_body( request_body(
content = api::types::CompletionRequest, content = api::types::ApiCompletionRequest,
description = "Text completion request", description = "Text completion request",
content_type = "application/json" content_type = "application/json"
), ),
@@ -28,7 +28,7 @@ use uuid::Uuid;
( (
status = 200, status = 200,
description = "Text completion response. If stream=true, response is SSE stream of chunks ending in [DONE].", description = "Text completion response. If stream=true, response is SSE stream of chunks ending in [DONE].",
body = api::types::CompletionResponse, body = api::types::ApiCompletionResponse,
content_type = "application/json" content_type = "application/json"
), ),
( (
@@ -53,17 +53,16 @@ use uuid::Uuid;
)] )]
pub async fn completions( pub async fn completions(
State(state): State<SharedState>, State(state): State<SharedState>,
Json(body): Json<api::types::CompletionRequest>, Json(body): Json<api::types::ApiCompletionRequest>,
) -> Result<impl IntoResponse, ApiError> { ) -> Result<impl IntoResponse, ApiError> {
tracing::debug!("Received /completion with body {:?}", body); tracing::debug!("Received /completion with body {:?}", body);
let response = state.chat_service.complete(body.into()).await?; let response = state.chat_service.complete(body.into()).await?;
match response { match response {
CompletionResult::NoStream(res) => Ok(Json::< CompletionResult::NoStream(res) => {
crate::core::llm::completions::CompletionResultNoStream, Ok(Json::<crate::api::types::ApiCompletionResponse>(res.into()).into_response())
>(res) }
.into_response()),
CompletionResult::Stream(stream) => { CompletionResult::Stream(stream) => {
let sse_stream = stream.map(|item| match item { let sse_stream = stream.map(|item| match item {
Ok(event) => { Ok(event) => {
@@ -85,7 +84,7 @@ pub async fn completions(
path = "/chat/completions", path = "/chat/completions",
tag = "chat", tag = "chat",
request_body( request_body(
content = api::types::ChatRequest, content = api::types::ApiChatRequest,
description = "Chat completion request with message history", description = "Chat completion request with message history",
content_type = "application/json" content_type = "application/json"
), ),
@@ -93,7 +92,7 @@ pub async fn completions(
( (
status = 200, status = 200,
description = "Chat completion response. If stream=false returns JSON. If stream=true returns SSE stream of chunks ending with [DONE].", description = "Chat completion response. If stream=false returns JSON. If stream=true returns SSE stream of chunks ending with [DONE].",
body = api::types::ChatResponse, body = api::types::ApiChatResponse,
content_type = "application/json" content_type = "application/json"
), ),
( (
@@ -125,7 +124,7 @@ pub async fn completions(
pub async fn chat_completions( pub async fn chat_completions(
State(state): State<SharedState>, State(state): State<SharedState>,
Extension(auth): Extension<Auth>, Extension(auth): Extension<Auth>,
Json(body): Json<api::types::ChatRequest>, Json(body): Json<api::types::ApiChatRequest>,
) -> Result<impl IntoResponse, ApiError> { ) -> Result<impl IntoResponse, ApiError> {
tracing::debug!("Received /chat/completion with body {:?}", body); tracing::debug!("Received /chat/completion with body {:?}", body);
@@ -133,15 +132,68 @@ pub async fn chat_completions(
match response { match response {
crate::core::llm::chat::ChatCompletionResult::NoStream(res) => { crate::core::llm::chat::ChatCompletionResult::NoStream(res) => {
Ok(Json::<crate::core::llm::chat::ChatCompletionResultNoStream>(res).into_response()) Ok(Json::<crate::api::types::ApiChatResponse>(res.into()).into_response())
} }
crate::core::llm::chat::ChatCompletionResult::Stream(stream) => { crate::core::llm::chat::ChatCompletionResult::Stream(stream) => {
let sse_stream = stream.map(|item| match item { let sse_stream = stream.map(|item| -> Result<Event, ApiError> {
Ok(event) => { tracing::debug!("{:?}", item);
let data = serde_json::to_string(&event).unwrap_or_default(); match item {
Ok(Event::default().data(data)) Err(e) => {
tracing::debug!("Error in chat_completions: {:?}", e);
Err(ApiError::from(e))
}
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Start {
conversation_id,
message_id,
created_at,
}) => {
let payload = serde_json::to_string(&api::types::StreamEvent::Start(
api::types::StartEventData {
conversation_id,
created: created_at,
id: message_id,
},
))
.unwrap_or_default();
Ok(Event::default().data(payload))
}
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Token(tok)) => {
let chunk = api::types::ChatCompletionChunk {
id: String::new(),
object: "chat.completion.chunk".to_string(),
choices: vec![api::types::ChatChunkChoice {
index: 0,
delta: api::types::Delta {
content: Some(tok),
role: None,
},
finish_reason: None,
}],
usage: None,
};
let payload = serde_json::to_string(&api::types::StreamEvent::Delta(chunk))
.unwrap_or_default();
Ok(Event::default().data(payload))
}
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Final(res)) => {
let payload = serde_json::to_string(&api::types::StreamEvent::End(
api::types::EndEventData {
created: res.created_at.parse().unwrap_or(0),
id: res.id,
usage: api::types::Usage {
prompt_tokens: res.prompt_tokens,
completion_tokens: res.completion_tokens,
total_tokens: res.prompt_tokens + res.completion_tokens,
},
},
))
.unwrap_or_default();
Ok(Event::default().data(payload))
}
} }
Err(e) => Err(e),
}); });
Ok(Sse::new(sse_stream) Ok(Sse::new(sse_stream)
@@ -192,146 +244,3 @@ pub async fn get_messages(
Ok(Json(messages.into())) Ok(Json(messages.into()))
} }
// async fn handle_stream(
// state: AppState,
// auth: Auth,
// body: api::ChatRequest,
// conversation_id: Option<Uuid>,
// ) -> Result<Response, ApiError> {
// // 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?;
// 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?;
// 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())
// }
+4 -2
View File
@@ -43,9 +43,11 @@ pub fn protected_router() -> Router<SharedState> {
pub fn router(state: SharedState) -> Router<SharedState> { pub fn router(state: SharedState) -> Router<SharedState> {
Router::new() Router::new()
.merge(public_router()) .merge(public_router())
.merge(protected_router().layer(middleware::from_fn_with_state( .merge(
protected_router().route_layer(middleware::from_fn_with_state(
state.clone(), state.clone(),
auth_middleware, auth_middleware,
))) )),
)
.with_state(state) .with_state(state)
} }
+7 -7
View File
@@ -13,7 +13,7 @@ use axum::{
( (
status = 200, status = 200,
description = "List of locally available Ollama models", description = "List of locally available Ollama models",
body = api::types::ModelsResponse, body = api::types::ApiModelsResponse,
content_type = "application/json", content_type = "application/json",
), ),
( (
@@ -27,7 +27,7 @@ use axum::{
// #[axum::debug_handler] // #[axum::debug_handler]
pub async fn list_models( pub async fn list_models(
State(state): State<SharedState>, State(state): State<SharedState>,
) -> Result<Json<api::types::ModelsResponse>, api::errors::ApiError> { ) -> Result<Json<api::types::ApiModelsResponse>, api::errors::ApiError> {
let models = state.chat_service.list_models().await?; let models = state.chat_service.list_models().await?;
Ok(Json(models.into())) Ok(Json(models.into()))
@@ -41,7 +41,7 @@ pub async fn list_models(
("model" = String, Path, description = "Name of the model to load into memory (e.g. 'llama3')") ("model" = String, Path, description = "Name of the model to load into memory (e.g. 'llama3')")
), ),
request_body( request_body(
content = api::types::LoadModelRequest, content = api::types::ApiLoadModelRequest,
description = "Load model request", description = "Load model request",
content_type = "application/json", content_type = "application/json",
example = json!({ "keep_alive": "10m" }) example = json!({ "keep_alive": "10m" })
@@ -50,7 +50,7 @@ pub async fn list_models(
( (
status = 200, status = 200,
description = "Model successfully loaded into memory", description = "Model successfully loaded into memory",
body = api::types::LoadModelResponse, body = api::types::ApiLoadModelResponse,
content_type = "application/json", content_type = "application/json",
), ),
( (
@@ -79,8 +79,8 @@ pub async fn list_models(
pub async fn load_model( pub async fn load_model(
State(state): State<SharedState>, State(state): State<SharedState>,
Path(model): Path<String>, Path(model): Path<String>,
Json(body): Json<api::types::LoadModelRequest>, Json(body): Json<api::types::ApiLoadModelRequest>,
) -> Result<Json<api::types::LoadModelResponse>, api::errors::ApiError> { ) -> Result<Json<api::types::ApiLoadModelResponse>, api::errors::ApiError> {
let response = state let response = state
.chat_service .chat_service
.load_model(crate::core::llm::models::LoadModelRequest { .load_model(crate::core::llm::models::LoadModelRequest {
@@ -89,7 +89,7 @@ pub async fn load_model(
}) })
.await?; .await?;
Ok(Json(api::types::LoadModelResponse { Ok(Json(api::types::ApiLoadModelResponse {
model: response.model, model: response.model,
keep_alive: body.keep_alive, keep_alive: body.keep_alive,
status: "loaded".to_string(), status: "loaded".to_string(),
+54 -54
View File
@@ -6,43 +6,43 @@ use uuid::Uuid;
// ------ Models ------ // ------ Models ------
#[derive(Debug, Serialize, ToSchema)] #[derive(Debug, Serialize, ToSchema)]
pub struct ModelInfo { pub struct ApiModelInfo {
pub name: String, pub name: String,
pub family: Option<String>, pub family: Option<String>,
pub parameter_size: Option<String>, pub parameter_size: Option<String>,
pub metadata: ModelMetadata, pub metadata: ApiModelMetadata,
} }
#[derive(Debug, Clone, Default, Serialize, ToSchema)] #[derive(Debug, Clone, Default, Serialize, ToSchema)]
pub struct ModelMetadata { pub struct ApiModelMetadata {
pub extra: HashMap<String, String>, pub extra: HashMap<String, String>,
} }
// ------ Endpoint: /models ------ // ------ Endpoint: /models ------
#[derive(Debug, Serialize, ToSchema)] #[derive(Debug, Serialize, ToSchema)]
pub struct ModelsResponse { pub struct ApiModelsResponse {
pub models: Vec<ModelInfo>, pub models: Vec<ApiModelInfo>,
} }
// ------ Endpoint: /models/{model}/load ------ // ------ Endpoint: /models/{model}/load ------
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct LoadModelResponse { pub struct ApiLoadModelResponse {
pub model: String, pub model: String,
pub status: String, pub status: String,
pub keep_alive: String, pub keep_alive: String,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct LoadModelRequest { pub struct ApiLoadModelRequest {
pub keep_alive: String, pub keep_alive: String,
} }
// ------ Endpoint: /models/{model}/unload ------ // ------ Endpoint: /models/{model}/unload ------
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct UnloadModelResponse { pub struct ApiUnloadModelResponse {
pub model: String, pub model: String,
pub status: String, pub status: String,
} }
@@ -50,7 +50,7 @@ pub struct UnloadModelResponse {
// ------ Completions ------ // ------ Completions ------
#[derive(Debug, Deserialize, Serialize, ToSchema, Default)] #[derive(Debug, Deserialize, Serialize, ToSchema, Default)]
pub struct LLMOptions { pub struct ApiLlmOptions {
#[serde(default)] #[serde(default)]
pub stream: bool, pub stream: bool,
@@ -74,18 +74,18 @@ pub struct LLMOptions {
// ------ Endpoint: /completions ------ // ------ Endpoint: /completions ------
#[derive(Debug, Deserialize, Serialize, ToSchema)] #[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct CompletionRequest { pub struct ApiCompletionRequest {
#[serde(flatten)] #[serde(flatten)]
pub options: LLMOptions, pub options: ApiLlmOptions,
pub model: String, pub model: String,
pub prompt: String, pub prompt: String,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct CompletionResponse { pub struct ApiCompletionResponse {
pub id: Uuid, pub id: Uuid,
pub object: CompletionObject, pub object: ApiCompletionObject,
pub created: String, pub created: String,
pub model: String, pub model: String,
pub choices: Vec<Choice>, pub choices: Vec<Choice>,
@@ -94,7 +94,7 @@ pub struct CompletionResponse {
#[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)] #[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
pub enum CompletionObject { pub enum ApiCompletionObject {
TextCompletion, TextCompletion,
} }
@@ -102,7 +102,7 @@ pub enum CompletionObject {
#[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)] #[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
pub enum FinishReason { pub enum ApiFinishReason {
Stop, Stop,
Length, Length,
ContentFilter, ContentFilter,
@@ -114,7 +114,7 @@ pub enum FinishReason {
pub struct Choice { pub struct Choice {
pub text: String, pub text: String,
pub index: u32, pub index: u32,
pub finish_reason: FinishReason, pub finish_reason: ApiFinishReason,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema, Clone, Copy)] #[derive(Debug, Serialize, Deserialize, ToSchema, Clone, Copy)]
@@ -132,12 +132,12 @@ pub struct CompletionChunk {
} }
#[derive(Debug, Deserialize, Serialize, ToSchema)] #[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct ChatRequest { pub struct ApiChatRequest {
#[serde(flatten)] #[serde(flatten)]
pub base: LLMOptions, pub base: ApiLlmOptions,
pub model: String, pub model: String,
pub message: Message, pub message: ApiMessage,
// Non standard Open AI // Non standard Open AI
pub conversation_id: Option<Uuid>, pub conversation_id: Option<Uuid>,
@@ -145,36 +145,36 @@ pub struct ChatRequest {
} }
#[derive(Debug, Deserialize, Serialize, ToSchema)] #[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct Message { pub struct ApiMessage {
pub role: Role, pub role: ApiRole,
pub content: String, pub content: String,
} }
#[derive(Debug, Deserialize, Serialize, ToSchema, PartialEq, Eq)] #[derive(Debug, Deserialize, Serialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "lowercase")] #[serde(rename_all = "lowercase")]
pub enum Role { pub enum ApiRole {
System, System,
User, User,
Assistant, Assistant,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ChatResponse { pub struct ApiChatResponse {
pub id: String, pub id: Uuid,
pub object: String, pub object: ApiCompletionObject,
pub created: u64, pub created: String,
pub model: String, pub model: String,
pub choices: Vec<ChatChoice>, pub choices: Vec<ApiChatChoice>,
pub usage: Usage, pub usage: Usage,
pub conversation_id: Uuid, pub conversation_id: Uuid,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ChatChoice { pub struct ApiChatChoice {
pub index: u32, pub index: u32,
pub message: Message, pub message: ApiMessage,
pub finish_reason: FinishReason, pub finish_reason: ApiFinishReason,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
@@ -189,39 +189,39 @@ pub struct ChatCompletionChunk {
pub struct ChatChunkChoice { pub struct ChatChunkChoice {
pub index: u32, pub index: u32,
pub delta: Delta, pub delta: Delta,
pub finish_reason: Option<FinishReason>, pub finish_reason: Option<ApiFinishReason>,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct Delta { pub struct Delta {
pub content: Option<String>, pub content: Option<String>,
pub role: Option<Role>, pub role: Option<ApiRole>,
} }
// #[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
// pub struct StartEventData { pub struct StartEventData {
// pub conversation_id: Uuid, pub conversation_id: Uuid,
// pub created: u64, pub created: u64,
// pub id: Uuid, pub id: Uuid,
// } }
// #[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
// pub struct EndEventData { pub struct EndEventData {
// pub created: u64, pub created: u64,
// pub id: Uuid, pub id: Uuid,
// pub usage: Usage, pub usage: Usage,
// } }
// #[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
// #[serde(tag = "type", content = "data")] #[serde(tag = "type", content = "data")]
// pub enum StreamEvent { pub enum StreamEvent {
// #[serde(rename = "start")] #[serde(rename = "start")]
// Start(StartEventData), Start(StartEventData),
// #[serde(rename = "end")] #[serde(rename = "end")]
// End(EndEventData), End(EndEventData),
// #[serde(rename = "delta")] #[serde(rename = "delta")]
// Delta(ChatCompletionChunk), Delta(ChatCompletionChunk),
// } }
// ------ Api Key ------ // ------ Api Key ------
+13 -6
View File
@@ -1,10 +1,11 @@
use super::ChatRole; use super::ChatRole;
use crate::providers::ollama::errors::LlmError; use crate::services::errors::ServiceError;
use futures::Stream; use futures::Stream;
use serde::Serialize; use serde::Serialize;
use std::pin::Pin; use std::pin::Pin;
use uuid::Uuid;
#[derive(Debug, Clone, Default)] #[derive(Debug, Clone, Default)]
pub struct ChatCompletionOptions { pub struct ChatCompletionOptions {
@@ -45,24 +46,30 @@ pub struct ChatCompletionResultNoStream {
pub created_at: String, pub created_at: String,
pub id: uuid::Uuid, pub id: uuid::Uuid,
pub conversation_id: Uuid,
pub prompt_tokens: u32, pub prompt_tokens: u32,
pub completion_tokens: u32, pub completion_tokens: u32,
pub done_reason: Option<String>, pub done_reason: String,
pub total_duration: Option<u64>, pub total_duration: u64,
pub load_duration: Option<u64>, pub load_duration: u64,
} }
#[derive(Serialize)] #[derive(Debug, Serialize)]
pub enum ChatCompletionStreamEvent { pub enum ChatCompletionStreamEvent {
Start {
conversation_id: Uuid,
message_id: Uuid,
created_at: u64,
},
Token(String), Token(String),
Final(ChatCompletionResultNoStream), Final(ChatCompletionResultNoStream),
} }
pub type ChatCompletionStream = pub type ChatCompletionStream =
Pin<Box<dyn Stream<Item = Result<ChatCompletionStreamEvent, LlmError>> + Send>>; Pin<Box<dyn Stream<Item = Result<ChatCompletionStreamEvent, ServiceError>> + Send>>;
pub enum ChatCompletionResult { pub enum ChatCompletionResult {
Stream(ChatCompletionStream), Stream(ChatCompletionStream),
+9 -9
View File
@@ -1,7 +1,7 @@
use crate::{api, core}; use crate::{api, core};
impl From<api::types::CompletionRequest> for core::llm::completions::CompletionRequest { impl From<api::types::ApiCompletionRequest> for core::llm::completions::CompletionRequest {
fn from(m: api::types::CompletionRequest) -> Self { fn from(m: api::types::ApiCompletionRequest) -> Self {
Self { Self {
model: m.model, model: m.model,
prompt: m.prompt, prompt: m.prompt,
@@ -20,8 +20,8 @@ impl From<api::types::CompletionRequest> for core::llm::completions::CompletionR
} }
} }
impl From<api::types::ChatRequest> for core::llm::chat::ChatCompletionRequest { impl From<api::types::ApiChatRequest> for core::llm::chat::ChatCompletionRequest {
fn from(m: api::types::ChatRequest) -> Self { fn from(m: api::types::ApiChatRequest) -> Self {
Self { Self {
model: m.model, model: m.model,
message: m.message.into(), message: m.message.into(),
@@ -43,13 +43,13 @@ impl From<api::types::ChatRequest> for core::llm::chat::ChatCompletionRequest {
} }
} }
impl From<api::types::Message> for core::llm::chat::Message { impl From<api::types::ApiMessage> for core::llm::chat::Message {
fn from(m: api::types::Message) -> Self { fn from(m: api::types::ApiMessage) -> Self {
Self { Self {
role: match m.role { role: match m.role {
api::types::Role::System => core::llm::ChatRole::System, api::types::ApiRole::System => core::llm::ChatRole::System,
api::types::Role::User => core::llm::ChatRole::User, api::types::ApiRole::User => core::llm::ChatRole::User,
api::types::Role::Assistant => core::llm::ChatRole::Assistant, api::types::ApiRole::Assistant => core::llm::ChatRole::Assistant,
}, },
content: m.content, content: m.content,
} }
+41 -6
View File
@@ -1,6 +1,6 @@
use crate::{api, core}; use crate::{api, core};
impl From<core::llm::models::Models> for api::types::ModelsResponse { impl From<core::llm::models::Models> for api::types::ApiModelsResponse {
fn from(m: core::llm::models::Models) -> Self { fn from(m: core::llm::models::Models) -> Self {
Self { Self {
models: m.models.into_iter().map(Into::into).collect(), models: m.models.into_iter().map(Into::into).collect(),
@@ -8,7 +8,7 @@ impl From<core::llm::models::Models> for api::types::ModelsResponse {
} }
} }
impl From<core::llm::models::ModelMetadata> for api::types::ModelMetadata { impl From<core::llm::models::ModelMetadata> for api::types::ApiModelMetadata {
fn from(m: core::llm::models::ModelMetadata) -> Self { fn from(m: core::llm::models::ModelMetadata) -> Self {
Self { Self {
extra: m extra: m
@@ -20,7 +20,7 @@ impl From<core::llm::models::ModelMetadata> for api::types::ModelMetadata {
} }
} }
impl From<core::llm::models::Model> for api::types::ModelInfo { impl From<core::llm::models::Model> for api::types::ApiModelInfo {
fn from(m: core::llm::models::Model) -> Self { fn from(m: core::llm::models::Model) -> Self {
Self { Self {
name: m.name, name: m.name,
@@ -79,15 +79,28 @@ impl From<core::databases::conversations::MessageList> for api::types::MessageLi
} }
} }
impl From<core::llm::completions::CompletionResultNoStream> for api::types::CompletionResponse { impl From<core::llm::chat::Message> for api::types::ApiMessage {
fn from(c: core::llm::chat::Message) -> Self {
Self {
role: match c.role {
core::llm::ChatRole::User => api::types::ApiRole::User,
core::llm::ChatRole::Assistant => api::types::ApiRole::Assistant,
core::llm::ChatRole::System => api::types::ApiRole::System,
},
content: c.content,
}
}
}
impl From<core::llm::completions::CompletionResultNoStream> for api::types::ApiCompletionResponse {
fn from(m: core::llm::completions::CompletionResultNoStream) -> Self { fn from(m: core::llm::completions::CompletionResultNoStream) -> Self {
Self { Self {
id: m.id, id: m.id,
object: api::types::CompletionObject::TextCompletion, object: api::types::ApiCompletionObject::TextCompletion,
created: m.created_at, created: m.created_at,
model: m.model, model: m.model,
choices: vec![api::types::Choice { choices: vec![api::types::Choice {
finish_reason: api::types::FinishReason::Stop, finish_reason: api::types::ApiFinishReason::Stop,
text: m.text, text: m.text,
index: 0, index: 0,
}], }],
@@ -99,3 +112,25 @@ impl From<core::llm::completions::CompletionResultNoStream> for api::types::Comp
} }
} }
} }
impl From<core::llm::chat::ChatCompletionResultNoStream> for api::types::ApiChatResponse {
fn from(m: core::llm::chat::ChatCompletionResultNoStream) -> Self {
Self {
id: m.id,
conversation_id: m.conversation_id,
object: api::types::ApiCompletionObject::TextCompletion,
created: m.created_at,
model: m.model,
choices: vec![api::types::ApiChatChoice {
finish_reason: api::types::ApiFinishReason::Stop,
message: m.message.into(),
index: 0,
}],
usage: api::types::Usage {
completion_tokens: m.completion_tokens,
prompt_tokens: m.prompt_tokens,
total_tokens: m.completion_tokens + m.prompt_tokens,
},
}
}
}
+27 -6
View File
@@ -285,28 +285,35 @@ impl OllamaProvider {
.await? .await?
.bytes_stream(); .bytes_stream();
let stream = byte_stream.flat_map(|chunk_result| { let stream = byte_stream
let mut out: Vec<Result<super::types::OllamaChatStreamEvent, LlmError>> = Vec::new(); .flat_map(|chunk_result| {
let mut out: Vec<Result<super::types::OllamaChatStreamEvent, LlmError>> =
Vec::new();
let chunk = match chunk_result { let chunk = match chunk_result {
Ok(b) => b, Ok(b) => b,
Err(e) => { Err(e) => {
tracing::debug!("Error: {:?}", e);
out.push(Err(LlmError::Http(e))); out.push(Err(LlmError::Http(e)));
return futures::stream::iter(out); return futures::stream::iter(out);
} }
}; };
tracing::debug!("Raw line: {:?}", std::str::from_utf8(&chunk));
for line in chunk.split(|&b| b == b'\n') { for line in chunk.split(|&b| b == b'\n') {
if line.is_empty() { if line.is_empty() {
continue; continue;
} }
let parsed: ollama::types::OllamaChatResponse = match serde_json::from_slice(line) { let parsed: ollama::types::OllamaChatResponse =
match serde_json::from_slice(line) {
Ok(v) => v, Ok(v) => v,
Err(_) => continue, Err(_) => continue,
}; };
if !parsed.message.content.is_empty() && !parsed.done { tracing::debug!("Parsed: {:?}", parsed);
if !parsed.message.content.is_empty() {
out.push(Ok(super::types::OllamaChatStreamEvent::Token( out.push(Ok(super::types::OllamaChatStreamEvent::Token(
parsed.message.content.clone(), parsed.message.content.clone(),
))); )));
@@ -317,8 +324,22 @@ impl OllamaProvider {
return futures::stream::iter(out); return futures::stream::iter(out);
} }
} }
futures::stream::iter(out) futures::stream::iter(out)
})
.scan(String::new(), |acc, event| {
let result = match event {
Ok(ollama::types::OllamaChatStreamEvent::Token(ref tok)) => {
acc.push_str(tok);
Some(event)
}
Ok(ollama::types::OllamaChatStreamEvent::Final(mut resp)) => {
tracing::debug!("Acc: {:?}", acc);
resp.message.content = std::mem::take(acc);
Some(Ok(ollama::types::OllamaChatStreamEvent::Final(resp)))
}
Err(_) => Some(event),
};
futures::future::ready(result)
}); });
Ok(Box::pin(stream)) Ok(Box::pin(stream))
+5 -5
View File
@@ -117,13 +117,13 @@ pub struct OllamaChatResponse {
pub message: OllamaMessage, pub message: OllamaMessage,
pub done: bool, pub done: bool,
pub done_reason: Option<String>, pub done_reason: String,
pub total_duration: Option<u64>, pub total_duration: u64,
pub load_duration: Option<u64>, pub load_duration: u64,
pub prompt_eval_count: Option<u32>, pub prompt_eval_count: u32,
pub eval_count: Option<u32>, pub eval_count: u32,
} }
#[derive(Debug)] #[derive(Debug)]
+62 -70
View File
@@ -109,7 +109,14 @@ impl ChatService {
.resolve_conversation_with_title(auth.user_id(), &body) .resolve_conversation_with_title(auth.user_id(), &body)
.await?; .await?;
tracing::debug!(
"Received conversation_id={:?}, parent_id={:?}",
conversation_id,
body.parent_id
);
let user_msg_id = self let user_msg_id = self
.conversation
.log_user_message( .log_user_message(
auth.user_id(), auth.user_id(),
conversation_id, conversation_id,
@@ -122,50 +129,81 @@ impl ChatService {
let stream = body.options.stream; let stream = body.options.stream;
let history = self let history = self
.build_chat_history( .build_chat_history(auth.user_id(), conversation_id, body.options.context_depth)
auth.user_id(),
conversation_id,
body.message.clone(),
body.options.context_depth,
)
.await?; .await?;
let mut request: crate::providers::ollama::types::OllamaChatRequest = body.into(); let mut request: crate::providers::ollama::types::OllamaChatRequest = body.into();
request.messages = history.into_iter().map(Into::into).collect(); request.messages = history.into_iter().map(Into::into).collect();
if stream { if stream {
let created_at = 3;
let start_event = futures::stream::once(async move {
Ok(core::llm::chat::ChatCompletionStreamEvent::Start {
conversation_id,
message_id: user_msg_id,
created_at,
})
});
let ollama_stream = self.ollama.chat_completions_stream(&request).await?; let ollama_stream = self.ollama.chat_completions_stream(&request).await?;
let mapped = ollama_stream.map(|item| { let conversation_svc = self.conversation.clone();
item.map(|event| match event { let user_id = auth.user_id();
OllamaChatStreamEvent::Token(tok) => {
core::llm::chat::ChatCompletionStreamEvent::Token(tok) let mapped = ollama_stream.then(move |item| {
let conversation_svc = conversation_svc.clone();
async move {
match item {
Err(e) => Err(e.into()),
Ok(OllamaChatStreamEvent::Token(tok)) => {
Ok(core::llm::chat::ChatCompletionStreamEvent::Token(tok))
} }
OllamaChatStreamEvent::Final(resp) => { Ok(OllamaChatStreamEvent::Final(resp)) => {
core::llm::chat::ChatCompletionStreamEvent::Final( tracing::debug!("Inserting {:?}", resp.message.content);
let assistant_message_id = conversation_svc
.log_assistant_message(
user_id,
conversation_id,
user_msg_id,
&resp.message.content,
resp.eval_count,
)
.await?;
conversation_svc
.update_message_tokens(user_id, user_msg_id, resp.prompt_eval_count)
.await?;
Ok(core::llm::chat::ChatCompletionStreamEvent::Final(
core::llm::chat::ChatCompletionResultNoStream { core::llm::chat::ChatCompletionResultNoStream {
id: Uuid::new_v4(), id: assistant_message_id,
conversation_id,
created_at: resp.created_at, created_at: resp.created_at,
model: resp.model, model: resp.model,
message: resp.message.into(), message: resp.message.into(),
prompt_tokens: resp.prompt_eval_count.unwrap_or(0), prompt_tokens: resp.prompt_eval_count,
completion_tokens: resp.eval_count.unwrap_or(0), completion_tokens: resp.eval_count,
done_reason: resp.done_reason, done_reason: resp.done_reason,
total_duration: resp.total_duration, total_duration: resp.total_duration,
load_duration: resp.load_duration, load_duration: resp.load_duration,
}, },
) ))
}
}
} }
})
}); });
Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin( Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin(
mapped, start_event.chain(mapped),
))) )))
} else { } else {
let response = self.ollama.chat_completions(&request).await?; let response = self.ollama.chat_completions(&request).await?;
self.log_assistant_message( let assistant_message_id = self
.conversation
.log_assistant_message(
auth.user_id(), auth.user_id(),
conversation_id, conversation_id,
user_msg_id, user_msg_id,
@@ -174,20 +212,17 @@ impl ChatService {
) )
.await?; .await?;
self.conversation self.conversation
.update_message_tokens( .update_message_tokens(auth.user_id(), user_msg_id, response.prompt_eval_count)
auth.user_id(),
user_msg_id,
response.prompt_eval_count.unwrap_or_default(),
)
.await?; .await?;
let enriched = core::llm::chat::ChatCompletionResultNoStream { let enriched = core::llm::chat::ChatCompletionResultNoStream {
id: Uuid::new_v4(), id: assistant_message_id,
conversation_id,
created_at: response.created_at, created_at: response.created_at,
model: response.model, model: response.model,
message: response.message.into(), message: response.message.into(),
prompt_tokens: response.prompt_eval_count.unwrap_or(0), prompt_tokens: response.prompt_eval_count,
completion_tokens: response.eval_count.unwrap_or(0), completion_tokens: response.eval_count,
done_reason: response.done_reason, done_reason: response.done_reason,
total_duration: response.total_duration, total_duration: response.total_duration,
load_duration: response.load_duration, load_duration: response.load_duration,
@@ -247,7 +282,6 @@ impl ChatService {
&self, &self,
auth_user_id: Uuid, auth_user_id: Uuid,
conversation_id: Uuid, conversation_id: Uuid,
body_messages: crate::core::llm::chat::Message,
context_depth: u32, context_depth: u32,
) -> Result<Vec<crate::core::llm::chat::Message>, ServiceError> { ) -> Result<Vec<crate::core::llm::chat::Message>, ServiceError> {
let messages = self let messages = self
@@ -270,48 +304,6 @@ impl ChatService {
history.reverse(); history.reverse();
history.push(body_messages);
Ok(history) Ok(history)
} }
async fn log_user_message(
&self,
user_id: Uuid,
conversation_id: Uuid,
parent_id: Option<Uuid>,
content: &str,
tokens: Option<u32>,
) -> Result<Uuid, ServiceError> {
self.conversation
.log_message(
user_id,
conversation_id,
parent_id,
llm::ChatRole::User,
content,
tokens,
)
.await
}
async fn log_assistant_message(
&self,
user_id: Uuid,
conversation_id: Uuid,
parent_id: Uuid,
content: &str,
tokens: Option<u32>,
) -> Result<Uuid, ServiceError> {
self.conversation
.log_message(
user_id,
conversation_id,
Some(parent_id),
llm::ChatRole::Assistant,
content,
tokens,
)
.await
}
} }
+46 -2
View File
@@ -46,17 +46,23 @@ impl ConversationService {
) -> Result<core::databases::conversations::MessageList, ServiceError> { ) -> Result<core::databases::conversations::MessageList, ServiceError> {
let limit = pointer.limit; let limit = pointer.limit;
let messages = postgres::chat::queries::get_conversation_messages( let mut messages = postgres::chat::queries::get_conversation_messages(
&self.postgres, &self.postgres,
user_id, user_id,
conversation_id, conversation_id,
limit, limit + 1,
pointer.before, pointer.before,
) )
.await?; .await?;
let has_more = messages.len() == limit as usize; let has_more = messages.len() == limit as usize;
if has_more {
messages.pop();
}
messages.reverse();
Ok(crate::core::databases::conversations::MessageList { Ok(crate::core::databases::conversations::MessageList {
messages: messages.into_iter().map(Into::into).collect(), messages: messages.into_iter().map(Into::into).collect(),
has_more, has_more,
@@ -118,6 +124,44 @@ impl ConversationService {
Ok(id) Ok(id)
} }
pub async fn log_user_message(
&self,
user_id: Uuid,
conversation_id: Uuid,
parent_id: Option<Uuid>,
content: &str,
tokens: Option<u32>,
) -> Result<Uuid, ServiceError> {
self.log_message(
user_id,
conversation_id,
parent_id,
crate::core::llm::ChatRole::User,
content,
tokens,
)
.await
}
pub async fn log_assistant_message(
&self,
user_id: Uuid,
conversation_id: Uuid,
parent_id: Uuid,
content: &str,
tokens: u32,
) -> Result<Uuid, ServiceError> {
self.log_message(
user_id,
conversation_id,
Some(parent_id),
crate::core::llm::ChatRole::Assistant,
content,
Some(tokens),
)
.await
}
pub async fn update_message_tokens( pub async fn update_message_tokens(
&self, &self,
user_id: Uuid, user_id: Uuid,
+1
View File
@@ -2,6 +2,7 @@ use crate::databases::errors::DbError;
use crate::providers::keycloak::errors::AuthError; use crate::providers::keycloak::errors::AuthError;
use crate::providers::ollama::errors::LlmError; use crate::providers::ollama::errors::LlmError;
#[derive(Debug)]
pub enum ServiceError { pub enum ServiceError {
Db(DbError), Db(DbError),
Llm(LlmError), Llm(LlmError),