aboutsummaryrefslogtreecommitdiffstats
path: root/crates/asist-llm/src/chat.rs
diff options
context:
space:
mode:
Diffstat (limited to 'crates/asist-llm/src/chat.rs')
-rw-r--r--crates/asist-llm/src/chat.rs275
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");
+ }
+}