@@ -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
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
@@ -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())
|
|
||||||
// }
|
|
||||||
|
|||||||
@@ -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(
|
||||||
state.clone(),
|
protected_router().route_layer(middleware::from_fn_with_state(
|
||||||
auth_middleware,
|
state.clone(),
|
||||||
)))
|
auth_middleware,
|
||||||
|
)),
|
||||||
|
)
|
||||||
.with_state(state)
|
.with_state(state)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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),
|
||||||
|
|||||||
@@ -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,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -285,41 +285,62 @@ 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>> =
|
||||||
let chunk = match chunk_result {
|
Vec::new();
|
||||||
Ok(b) => b,
|
let chunk = match chunk_result {
|
||||||
Err(e) => {
|
Ok(b) => b,
|
||||||
out.push(Err(LlmError::Http(e)));
|
Err(e) => {
|
||||||
return futures::stream::iter(out);
|
tracing::debug!("Error: {:?}", 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,
|
|
||||||
};
|
};
|
||||||
|
|
||||||
if !parsed.message.content.is_empty() && !parsed.done {
|
tracing::debug!("Raw line: {:?}", std::str::from_utf8(&chunk));
|
||||||
out.push(Ok(super::types::OllamaChatStreamEvent::Token(
|
|
||||||
parsed.message.content.clone(),
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
|
|
||||||
if parsed.done {
|
for line in chunk.split(|&b| b == b'\n') {
|
||||||
out.push(Ok(super::types::OllamaChatStreamEvent::Final(parsed)));
|
if line.is_empty() {
|
||||||
return futures::stream::iter(out);
|
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))
|
Ok(Box::pin(stream))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)]
|
||||||
|
|||||||
@@ -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,72 +129,100 @@ 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 ollama_stream = self.ollama.chat_completions_stream(&request).await?;
|
let created_at = 3;
|
||||||
|
|
||||||
let mapped = ollama_stream.map(|item| {
|
let start_event = futures::stream::once(async move {
|
||||||
item.map(|event| match event {
|
Ok(core::llm::chat::ChatCompletionStreamEvent::Start {
|
||||||
OllamaChatStreamEvent::Token(tok) => {
|
conversation_id,
|
||||||
core::llm::chat::ChatCompletionStreamEvent::Token(tok)
|
message_id: user_msg_id,
|
||||||
}
|
created_at,
|
||||||
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 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(
|
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
|
||||||
auth.user_id(),
|
.conversation
|
||||||
conversation_id,
|
.log_assistant_message(
|
||||||
user_msg_id,
|
|
||||||
&response.message.content,
|
|
||||||
response.eval_count,
|
|
||||||
)
|
|
||||||
.await?;
|
|
||||||
self.conversation
|
|
||||||
.update_message_tokens(
|
|
||||||
auth.user_id(),
|
auth.user_id(),
|
||||||
|
conversation_id,
|
||||||
user_msg_id,
|
user_msg_id,
|
||||||
response.prompt_eval_count.unwrap_or_default(),
|
&response.message.content,
|
||||||
|
response.eval_count,
|
||||||
)
|
)
|
||||||
.await?;
|
.await?;
|
||||||
|
self.conversation
|
||||||
|
.update_message_tokens(auth.user_id(), user_msg_id, response.prompt_eval_count)
|
||||||
|
.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,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,
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
Reference in New Issue
Block a user