//! 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, /// Para los mensajes de rol `tool`: a qué llamada responden. pub tool_call_id: Option, } impl Message { pub fn system(content: impl Into) -> Self { Self::plain(Role::System, content) } pub fn user(content: impl Into) -> Self { Self::plain(Role::User, content) } pub fn assistant(content: impl Into) -> Self { Self::plain(Role::Assistant, content) } fn plain(role: Role, content: impl Into) -> 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) -> 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, max_turns: usize, } impl Conversation { pub fn new(system_prompt: impl Into, 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) { 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 = 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 = 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"); } }