diff options
Diffstat (limited to 'crates/asist-llm/src/chat.rs')
| -rw-r--r-- | crates/asist-llm/src/chat.rs | 275 |
1 files changed, 275 insertions, 0 deletions
diff --git a/crates/asist-llm/src/chat.rs b/crates/asist-llm/src/chat.rs new file mode 100644 index 0000000..01a9a74 --- /dev/null +++ b/crates/asist-llm/src/chat.rs @@ -0,0 +1,275 @@ +//! Historial de conversación en el formato que espera la API de chat. + +use serde::{Deserialize, Serialize}; +use serde_json::{json, Value}; + +use asist_core::tools::{ToolCall, ToolOutcome}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum Role { + System, + User, + Assistant, + Tool, +} + +impl Role { + fn as_str(self) -> &'static str { + match self { + Role::System => "system", + Role::User => "user", + Role::Assistant => "assistant", + Role::Tool => "tool", + } + } +} + +#[derive(Debug, Clone)] +pub struct Message { + pub role: Role, + pub content: String, + /// Herramientas que el asistente pidió en este turno. + pub tool_calls: Vec<ToolCall>, + /// Para los mensajes de rol `tool`: a qué llamada responden. + pub tool_call_id: Option<String>, +} + +impl Message { + pub fn system(content: impl Into<String>) -> Self { + Self::plain(Role::System, content) + } + + pub fn user(content: impl Into<String>) -> Self { + Self::plain(Role::User, content) + } + + pub fn assistant(content: impl Into<String>) -> Self { + Self::plain(Role::Assistant, content) + } + + fn plain(role: Role, content: impl Into<String>) -> Self { + Self { + role, + content: content.into(), + tool_calls: Vec::new(), + tool_call_id: None, + } + } + + /// Turno del asistente que en vez de hablar pidió herramientas. + pub fn tool_request(content: String, tool_calls: Vec<ToolCall>) -> Self { + Self { + role: Role::Assistant, + content, + tool_calls, + tool_call_id: None, + } + } + + /// Resultado devuelto al modelo. + pub fn tool_result(outcome: &ToolOutcome) -> Self { + Self { + role: Role::Tool, + content: outcome.output.clone(), + tool_calls: Vec::new(), + tool_call_id: Some(outcome.id.clone()), + } + } + + pub fn to_json(&self) -> Value { + let mut object = json!({ "role": self.role.as_str(), "content": self.content }); + if !self.tool_calls.is_empty() { + object["tool_calls"] = Value::Array( + self.tool_calls + .iter() + .map(|call| { + json!({ + "id": call.id, + "type": "function", + "function": { "name": call.name, "arguments": call.arguments } + }) + }) + .collect(), + ); + } + if let Some(id) = &self.tool_call_id { + object["tool_call_id"] = json!(id); + } + object + } +} + +/// Historial con la instrucción de sistema fija y una ventana deslizante de +/// turnos, para que la conversación no crezca sin fin. +#[derive(Debug, Clone)] +pub struct Conversation { + system: Message, + turns: Vec<Message>, + max_turns: usize, +} + +impl Conversation { + pub fn new(system_prompt: impl Into<String>, max_turns: usize) -> Self { + Self { + system: Message::system(system_prompt), + turns: Vec::new(), + max_turns, + } + } + + /// Cambia la instrucción de sistema sin tocar el historial. + /// + /// El turno alterna entre la guía de herramientas y la de estilo, y el + /// historial tiene que sobrevivir al cambio: si se reiniciara, el modelo + /// perdería los resultados que acaba de pedir. + pub fn set_system(&mut self, prompt: &str) { + self.system = Message::system(prompt); + } + + pub fn push(&mut self, message: Message) { + self.turns.push(message); + self.trim(); + } + + pub fn extend(&mut self, messages: impl IntoIterator<Item = Message>) { + self.turns.extend(messages); + self.trim(); + } + + /// Recorta a `max_turns` intervenciones de usuario, sin dejar nunca un + /// mensaje de rol `tool` huérfano al principio: la API lo rechaza si no + /// va precedido de la llamada que lo originó. + fn trim(&mut self) { + if self.max_turns == 0 { + self.turns.clear(); + return; + } + let user_positions: Vec<usize> = self + .turns + .iter() + .enumerate() + .filter(|(_, m)| m.role == Role::User) + .map(|(i, _)| i) + .collect(); + if user_positions.len() <= self.max_turns { + return; + } + let cut = user_positions[user_positions.len() - self.max_turns]; + self.turns.drain(..cut); + } + + pub fn messages(&self) -> Vec<&Message> { + std::iter::once(&self.system) + .chain(self.turns.iter()) + .collect() + } + + pub fn to_json(&self) -> Value { + Value::Array(self.messages().into_iter().map(Message::to_json).collect()) + } + + pub fn len(&self) -> usize { + self.turns.len() + } + + pub fn is_empty(&self) -> bool { + self.turns.is_empty() + } + + pub fn clear(&mut self) { + self.turns.clear(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn la_instruccion_de_sistema_va_siempre_delante() { + let mut chat = Conversation::new("sé breve", 4); + chat.push(Message::user("hola")); + let json = chat.to_json(); + assert_eq!(json[0]["role"], "system"); + assert_eq!(json[0]["content"], "sé breve"); + assert_eq!(json[1]["role"], "user"); + } + + #[test] + fn cambiar_la_instruccion_de_sistema_conserva_el_historial() { + let mut chat = Conversation::new("primera", 4); + chat.push(Message::user("hola")); + chat.set_system("segunda"); + let json = chat.to_json(); + assert_eq!(json[0]["content"], "segunda"); + assert_eq!( + json[1]["content"], "hola", + "el historial no puede perderse al cambiar" + ); + } + + #[test] + fn el_historial_se_recorta_por_turnos_de_usuario() { + let mut chat = Conversation::new("s", 2); + for i in 0..5 { + chat.push(Message::user(format!("p{i}"))); + chat.push(Message::assistant(format!("r{i}"))); + } + let messages = chat.messages(); + let users: Vec<&str> = messages + .iter() + .filter(|m| m.role == Role::User) + .map(|m| m.content.as_str()) + .collect(); + assert_eq!(users, vec!["p3", "p4"]); + } + + #[test] + fn el_recorte_no_deja_un_resultado_de_herramienta_huerfano() { + // La API rechaza un mensaje `tool` que no venga detrás de la llamada + // que lo pidió, así que el corte tiene que caer en un turno de usuario. + let mut chat = Conversation::new("s", 1); + chat.push(Message::user("p0")); + chat.push(Message::tool_request( + String::new(), + vec![ToolCall { + id: "a".into(), + name: "t".into(), + arguments: "{}".into(), + }], + )); + chat.push(Message { + role: Role::Tool, + content: "ok".into(), + tool_calls: vec![], + tool_call_id: Some("a".into()), + }); + chat.push(Message::user("p1")); + + let roles: Vec<Role> = chat.messages().iter().map(|m| m.role).collect(); + assert_eq!(roles, vec![Role::System, Role::User]); + } + + #[test] + fn con_cero_turnos_no_hay_memoria() { + let mut chat = Conversation::new("s", 0); + chat.push(Message::user("hola")); + assert_eq!(chat.messages().len(), 1); + } + + #[test] + fn una_llamada_a_herramienta_se_serializa_como_espera_la_api() { + let message = Message::tool_request( + String::new(), + vec![ToolCall { + id: "call_1".into(), + name: "hora_actual".into(), + arguments: "{}".into(), + }], + ); + let json = message.to_json(); + assert_eq!(json["tool_calls"][0]["type"], "function"); + assert_eq!(json["tool_calls"][0]["function"]["name"], "hora_actual"); + } +} |