318 lines
10 KiB
Rust
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
|
|
}
|
|
}
|