use crate::types::{RequestBuilder, CompletionResponse, MessageContent, ContentBlock};
use crate::error::Result;
use crate::client::LowLevelClient;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ChatterId {
User,
System,
Me,
Agent,
Custom(String),
Multiple(Vec<ChatterId>),
}
impl ChatterId {
pub fn to_role(&self) -> String {
match self {
ChatterId::User | ChatterId::Me => "user".to_string(),
ChatterId::System => "system".to_string(),
ChatterId::Agent => "assistant".to_string(),
ChatterId::Custom(name) => name.clone(),
ChatterId::Multiple(ids) => {
ids.first()
.map(|id| id.to_role())
.unwrap_or_else(|| "user".to_string())
}
}
}
pub fn custom(name: &str) -> Self {
ChatterId::Custom(name.to_string())
}
}
#[derive(Debug, Clone)]
pub struct ChatMessage {
pub chatter_id: ChatterId,
pub content: MessageContent,
}
pub struct ChatBuilder {
conversation: Vec<ChatMessage>,
request_builder: RequestBuilder,
current_chatter: ChatterId,
}
impl ChatBuilder {
pub(crate) fn new(request_builder: RequestBuilder, initial_chatter: ChatterId) -> Self {
Self {
conversation: Vec::new(),
request_builder,
current_chatter: initial_chatter,
}
}
pub fn add_message(mut self, chatter_id: ChatterId, content: &str) -> Self {
let message = ChatMessage {
chatter_id: chatter_id.clone(),
content: MessageContent::Text {
role: chatter_id.to_role(),
content: content.to_string(),
},
};
self.conversation.push(message);
self
}
pub fn add_multimodal_message(mut self, chatter_id: ChatterId, content_blocks: Vec<ContentBlock>) -> Self {
let message = ChatMessage {
chatter_id: chatter_id.clone(),
content: MessageContent::Multimodal {
role: chatter_id.to_role(),
content: content_blocks,
},
};
self.conversation.push(message);
self
}
pub fn system(mut self, prompt: &str) -> Self {
self.request_builder.request.system_prompt = Some(prompt.to_string());
self
}
pub async fn send(mut self, chatter_id: ChatterId, content: &str) -> Result<CompletionResponse> {
self = self.add_message(chatter_id, content);
let messages: Vec<MessageContent> = self.conversation.iter()
.map(|msg| msg.content.clone())
.collect();
self.request_builder.request.messages = messages;
let client = crate::client::HttpClient::from_model_id(
&*self.request_builder.config_provider,
&self.request_builder.model_id
)?;
let response = client.complete(&self.request_builder.request).await?;
Ok(response)
}
pub async fn send_multimodal(mut self, chatter_id: ChatterId, content_blocks: Vec<ContentBlock>) -> Result<CompletionResponse> {
self = self.add_multimodal_message(chatter_id, content_blocks);
let messages: Vec<MessageContent> = self.conversation.iter()
.map(|msg| msg.content.clone())
.collect();
self.request_builder.request.messages = messages;
let client = crate::client::HttpClient::from_model_id(
&*self.request_builder.config_provider,
&self.request_builder.model_id
)?;
client.complete(&self.request_builder.request).await
}
pub async fn continue_conversation(self, content: &str) -> Result<CompletionResponse> {
let current_chatter = self.current_chatter.clone();
self.send(current_chatter, content).await
}
pub fn conversation(&self) -> &[ChatMessage] {
&self.conversation
}
pub fn clear_history(mut self) -> Self {
self.conversation.clear();
self
}
pub fn message_count(&self) -> usize {
self.conversation.len()
}
pub fn temperature(mut self, temp: f64) -> Self {
self.request_builder = self.request_builder.temperature(temp);
self
}
pub fn max_tokens(mut self, tokens: u32) -> Self {
self.request_builder = self.request_builder.max_tokens(tokens);
self
}
pub fn top_p(mut self, top_p: f64) -> Self {
self.request_builder = self.request_builder.top_p(top_p);
self
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_chatter_id_to_role() {
assert_eq!(ChatterId::User.to_role(), "user");
assert_eq!(ChatterId::Me.to_role(), "user");
assert_eq!(ChatterId::System.to_role(), "system");
assert_eq!(ChatterId::Agent.to_role(), "assistant");
assert_eq!(ChatterId::Custom("narrator".to_string()).to_role(), "narrator");
let custom = ChatterId::custom("teacher");
assert_eq!(custom.to_role(), "teacher");
let multiple = ChatterId::Multiple(vec![ChatterId::User, ChatterId::Agent]);
assert_eq!(multiple.to_role(), "user");
}
#[test]
fn test_chatter_id_equality() {
assert_eq!(ChatterId::User, ChatterId::User);
assert_eq!(ChatterId::Custom("alice".to_string()), ChatterId::Custom("alice".to_string()));
assert_ne!(ChatterId::User, ChatterId::Agent);
assert_ne!(ChatterId::Custom("alice".to_string()), ChatterId::Custom("bob".to_string()));
}
}