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