Files
chat-api/src/services/chat_service.rs
T

318 lines
10 KiB
Rust

use crate::core;
use crate::core::llm;
use crate::core::llm::completions::{
CompletionResult, CompletionResultNoStream, CompletionStreamEvent,
};
use crate::providers::ollama::types::OllamaChatStreamEvent;
use crate::providers::{ollama::client::OllamaProvider, ollama::types::OllamaGenerateStreamEvent};
use crate::services::errors::ServiceError;
use super::ConversationService;
use futures::StreamExt;
use std::boxed::Box;
use uuid::Uuid;
#[derive(Clone)]
pub struct ChatService {
ollama: OllamaProvider,
conversation: ConversationService,
}
impl ChatService {
pub fn new(ollama: OllamaProvider, conversation: ConversationService) -> Self {
Self {
ollama,
conversation,
}
}
pub async fn list_models(&self) -> Result<llm::models::Models, ServiceError> {
let models = self.ollama.list_models().await?;
Ok(models.into())
}
pub async fn load_model(
&self,
body: crate::core::llm::models::LoadModelRequest,
) -> Result<crate::core::llm::models::LoadModelResponse, ServiceError> {
let b = crate::providers::ollama::types::OllamaGenerateRequest {
model: body.model.clone(),
prompt: "load".to_string(),
stream: false,
keep_alive: body.keep_alive,
options: None,
};
self.ollama.completions(&b).await?;
Ok(crate::core::llm::models::LoadModelResponse { model: body.model })
}
pub async fn complete(
&self,
body: core::llm::completions::CompletionRequest,
) -> Result<core::llm::completions::CompletionResult, ServiceError> {
let stream = body.options.stream;
let request: crate::providers::ollama::types::OllamaGenerateRequest = body.into();
if stream {
let ollama_stream = self.ollama.completions_stream(&request).await?;
let mapped = ollama_stream.map(|item| {
item.map(|event| match event {
OllamaGenerateStreamEvent::Token(tok) => CompletionStreamEvent::Token(tok),
OllamaGenerateStreamEvent::Final(resp) => {
CompletionStreamEvent::Final(CompletionResultNoStream {
id: Uuid::new_v4(),
created_at: resp.created_at,
model: resp.model,
text: resp.response,
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,
})
}
})
});
Ok(CompletionResult::Stream(Box::pin(mapped)))
} else {
let response = self.ollama.completions(&request).await?;
let enriched = core::llm::completions::CompletionResultNoStream {
id: Uuid::new_v4(),
created_at: response.created_at,
model: response.model,
text: response.response,
prompt_tokens: response.prompt_eval_count.unwrap_or(0),
completion_tokens: response.eval_count.unwrap_or(0),
done_reason: response.done_reason,
total_duration: response.total_duration,
load_duration: response.load_duration,
};
Ok(core::llm::completions::CompletionResult::NoStream(enriched))
}
}
pub async fn chat_complete(
&self,
body: core::llm::chat::ChatCompletionRequest,
auth: &core::auth::Auth,
) -> Result<core::llm::chat::ChatCompletionResult, ServiceError> {
let conversation_id = self
.resolve_conversation_with_title(auth.user_id(), &body)
.await?;
let user_msg_id = self
.log_user_message(
auth.user_id(),
conversation_id,
body.parent_id,
&body.message.content,
None,
)
.await?;
let stream = body.options.stream;
let history = self
.build_chat_history(
auth.user_id(),
conversation_id,
body.message.clone(),
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 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,
},
)
}
})
});
Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin(
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(
auth.user_id(),
user_msg_id,
response.prompt_eval_count.unwrap_or_default(),
)
.await?;
let enriched = core::llm::chat::ChatCompletionResultNoStream {
id: Uuid::new_v4(),
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),
done_reason: response.done_reason,
total_duration: response.total_duration,
load_duration: response.load_duration,
};
Ok(core::llm::chat::ChatCompletionResult::NoStream(enriched))
}
}
// ------ Helpers ------
async fn resolve_conversation_with_title(
&self,
user_id: Uuid,
body: &core::llm::chat::ChatCompletionRequest,
) -> Result<Uuid, ServiceError> {
let conversation_id = match self
.conversation
.get_or_create_conversation(user_id, body.conversation_id)
.await?
{
core::databases::conversations::ConversationResult::Existing(id) => id,
core::databases::conversations::ConversationResult::Created(id) => {
let last_message = body.message.content.as_str();
let prompt = format!(
"Generate a title using next message in maximum 6 words: {}",
last_message
);
let title_result = self
.complete(core::llm::completions::CompletionRequest {
model: body.model.clone(),
prompt,
options: core::llm::completions::CompletionOptions {
stream: false,
..Default::default()
},
})
.await?;
if let CompletionResult::NoStream(t) = title_result {
self.conversation
.set_conversation_title(user_id, id, &t.text)
.await?;
}
id
}
};
Ok(conversation_id)
}
async fn build_chat_history(
&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
.conversation
.get_messages_entries(
auth_user_id,
conversation_id,
crate::core::databases::conversations::CursorPage {
limit: context_depth,
before: None,
},
)
.await?;
let mut history: Vec<_> = messages
.messages
.into_iter()
.map(crate::core::llm::chat::Message::from_summary)
.collect();
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
}
}