Skip to main content

llm_agent_runtime/
dialogue.rs

1//! # Module: Dialogue
2//!
3//! ## Responsibility
4//! Provides multi-turn conversation memory for agents.  Each session tracks a
5//! chronological sequence of (role, content) turns and can be assembled into a
6//! prompt context that blends episodic memory with live dialogue history.
7//!
8//! ## Guarantees
9//! - Thread-safe: `DialogueStore` wraps its map in `Arc<Mutex<_>>`
10//! - Non-panicking: all operations return `Result`
11//! - Serialisable: all public types implement `serde::Serialize` /
12//!   `serde::Deserialize`
13//! - Lock-poisoning resilient: uses `crate::util::recover_lock`
14//!
15//! ## NOT Responsible For
16//! - Persistent storage of dialogue history (see `persistence` module)
17//! - Semantic search across past conversations (see `memory` module)
18
19use crate::error::AgentRuntimeError;
20use crate::util::recover_lock;
21use serde::{Deserialize, Serialize};
22use std::collections::HashMap;
23use std::sync::{Arc, Mutex};
24
25// ── Role ─────────────────────────────────────────────────────────────────────
26
27/// The speaker role for a single dialogue turn.
28#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
29#[serde(rename_all = "lowercase")]
30pub enum Role {
31    /// A message produced by a human user.
32    User,
33    /// A message produced by the AI assistant.
34    Assistant,
35    /// A system-level instruction injected at the start of a conversation.
36    System,
37}
38
39impl std::fmt::Display for Role {
40    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41        match self {
42            Role::User => f.write_str("user"),
43            Role::Assistant => f.write_str("assistant"),
44            Role::System => f.write_str("system"),
45        }
46    }
47}
48
49// ── DialogueTurn ─────────────────────────────────────────────────────────────
50
51/// A single turn in a multi-turn dialogue — a role/content pair.
52#[derive(Debug, Clone, Serialize, Deserialize)]
53pub struct DialogueTurn {
54    /// Who produced this message.
55    pub role: Role,
56    /// The text content of this message.
57    pub content: String,
58}
59
60impl DialogueTurn {
61    /// Construct a new turn with the given role and content.
62    pub fn new(role: Role, content: impl Into<String>) -> Self {
63        Self {
64            role,
65            content: content.into(),
66        }
67    }
68}
69
70// ── DialogueSession ───────────────────────────────────────────────────────────
71
72/// A single conversation session holding an ordered history of turns.
73///
74/// The session is keyed by a `session_id` string supplied by the caller (e.g.
75/// a UUID or a stable user-facing conversation ID).
76#[derive(Debug, Clone, Serialize, Deserialize)]
77pub struct DialogueSession {
78    /// Caller-assigned stable identifier for this session.
79    pub session_id: String,
80    /// Chronological list of dialogue turns.
81    turns: Vec<DialogueTurn>,
82    /// Optional single-sentence summary of turns that were truncated.
83    summary: Option<String>,
84}
85
86impl DialogueSession {
87    /// Create a new, empty `DialogueSession` with the given ID.
88    pub fn new(session_id: impl Into<String>) -> Self {
89        Self {
90            session_id: session_id.into(),
91            turns: Vec::new(),
92            summary: None,
93        }
94    }
95
96    /// Append a new turn to the session.
97    ///
98    /// Returns `Err` if either `role` or `content` would be empty after
99    /// trimming, to guard against accidental blank entries.
100    pub fn append_turn(
101        &mut self,
102        role: Role,
103        content: impl Into<String>,
104    ) -> Result<(), AgentRuntimeError> {
105        let content = content.into();
106        if content.trim().is_empty() {
107            return Err(AgentRuntimeError::Memory(
108                "dialogue turn content must not be blank".into(),
109            ));
110        }
111        self.turns.push(DialogueTurn::new(role, content));
112        Ok(())
113    }
114
115    /// Return a slice of all turns in chronological order.
116    pub fn get_history(&self) -> &[DialogueTurn] {
117        &self.turns
118    }
119
120    /// Return the number of turns currently stored.
121    pub fn turn_count(&self) -> usize {
122        self.turns.len()
123    }
124
125    /// Return the optional summary string, if one has been produced by
126    /// [`summarize_if_long`].
127    pub fn summary(&self) -> Option<&str> {
128        self.summary.as_deref()
129    }
130
131    /// Truncate the history to the most recent `keep_last` turns, setting
132    /// `summary` to a brief sentence describing how many turns were dropped.
133    ///
134    /// If the total number of turns is already ≤ `keep_last`, this is a no-op
135    /// and `Ok(false)` is returned.  When truncation occurs `Ok(true)` is
136    /// returned.
137    ///
138    /// # Errors
139    /// Returns `Err` if `keep_last` is zero, since a zero-length history
140    /// would discard all context.
141    pub fn summarize_if_long(&mut self, keep_last: usize) -> Result<bool, AgentRuntimeError> {
142        if keep_last == 0 {
143            return Err(AgentRuntimeError::Memory(
144                "keep_last must be greater than zero".into(),
145            ));
146        }
147        if self.turns.len() <= keep_last {
148            return Ok(false);
149        }
150        let dropped = self.turns.len() - keep_last;
151        self.summary = Some(format!(
152            "[{dropped} earlier turn(s) omitted for brevity]"
153        ));
154        self.turns = self.turns.split_off(dropped);
155        Ok(true)
156    }
157
158    /// Assemble the full session into a single prompt-ready context string.
159    ///
160    /// The optional `episodic_context` prefix (recalled memories) is prepended
161    /// before the dialogue history.  If a summary is present it is inserted
162    /// between the episodic context and the turn history.
163    pub fn build_context(&self, episodic_context: Option<&str>) -> String {
164        let mut parts: Vec<String> = Vec::new();
165
166        if let Some(ctx) = episodic_context {
167            if !ctx.trim().is_empty() {
168                parts.push(format!("[Memory context]\n{ctx}"));
169            }
170        }
171
172        if let Some(summary) = &self.summary {
173            parts.push(summary.clone());
174        }
175
176        for turn in &self.turns {
177            parts.push(format!("{}: {}", turn.role, turn.content));
178        }
179
180        parts.join("\n")
181    }
182}
183
184// ── DialogueStore ─────────────────────────────────────────────────────────────
185
186/// Concurrent map of session ID → [`DialogueSession`].
187///
188/// `DialogueStore` is cheap to clone — the inner map is wrapped in an
189/// `Arc<Mutex<_>>` so all clones share the same underlying storage.
190#[derive(Debug, Clone, Default)]
191pub struct DialogueStore {
192    sessions: Arc<Mutex<HashMap<String, DialogueSession>>>,
193}
194
195impl DialogueStore {
196    /// Create an empty `DialogueStore`.
197    pub fn new() -> Self {
198        Self::default()
199    }
200
201    /// Insert or completely replace the session for `session_id`.
202    pub fn upsert(&self, session: DialogueSession) -> Result<(), AgentRuntimeError> {
203        let mut map = recover_lock(self.sessions.lock(), "DialogueStore::upsert");
204        map.insert(session.session_id.clone(), session);
205        Ok(())
206    }
207
208    /// Retrieve a clone of the session for `session_id`, or `None` if absent.
209    pub fn get(&self, session_id: &str) -> Result<Option<DialogueSession>, AgentRuntimeError> {
210        let map = recover_lock(self.sessions.lock(), "DialogueStore::get");
211        Ok(map.get(session_id).cloned())
212    }
213
214    /// Append a turn to an existing session, creating the session if it does
215    /// not yet exist.
216    pub fn append_turn(
217        &self,
218        session_id: impl Into<String>,
219        role: Role,
220        content: impl Into<String>,
221    ) -> Result<(), AgentRuntimeError> {
222        let session_id = session_id.into();
223        let mut map = recover_lock(self.sessions.lock(), "DialogueStore::append_turn");
224        let session = map
225            .entry(session_id.clone())
226            .or_insert_with(|| DialogueSession::new(session_id));
227        session.append_turn(role, content)
228    }
229
230    /// Remove and return the session for `session_id`.
231    pub fn remove(&self, session_id: &str) -> Result<Option<DialogueSession>, AgentRuntimeError> {
232        let mut map = recover_lock(self.sessions.lock(), "DialogueStore::remove");
233        Ok(map.remove(session_id))
234    }
235
236    /// Return the number of sessions currently stored.
237    pub fn session_count(&self) -> usize {
238        let map = recover_lock(self.sessions.lock(), "DialogueStore::session_count");
239        map.len()
240    }
241
242    /// Return a sorted list of all session IDs.
243    pub fn session_ids(&self) -> Vec<String> {
244        let map = recover_lock(self.sessions.lock(), "DialogueStore::session_ids");
245        let mut ids: Vec<String> = map.keys().cloned().collect();
246        ids.sort();
247        ids
248    }
249}
250
251// ── DialogueContext ───────────────────────────────────────────────────────────
252
253/// A fully assembled prompt context combining episodic memory with the live
254/// dialogue history for a session.
255///
256/// Construct via [`DialogueContext::build`] and pass the resulting
257/// [`DialogueContext::prompt`] string to an inference closure or LLM provider.
258#[derive(Debug, Clone, Serialize, Deserialize)]
259pub struct DialogueContext {
260    /// The session ID this context was assembled for.
261    pub session_id: String,
262    /// Final assembled prompt string ready to be sent to an LLM.
263    pub prompt: String,
264    /// Number of dialogue turns included (after any summarization).
265    pub turn_count: usize,
266    /// Whether a summary of older turns is present in the prompt.
267    pub has_summary: bool,
268}
269
270impl DialogueContext {
271    /// Assemble a `DialogueContext` from a `DialogueSession` and an optional
272    /// episodic memory snippet.
273    ///
274    /// The `episodic_context` string is typically produced by recalling items
275    /// from an [`crate::memory::EpisodicStore`] and formatting them as text.
276    pub fn build(session: &DialogueSession, episodic_context: Option<&str>) -> Self {
277        Self {
278            session_id: session.session_id.clone(),
279            prompt: session.build_context(episodic_context),
280            turn_count: session.turn_count(),
281            has_summary: session.summary().is_some(),
282        }
283    }
284}
285
286// ── Tests ────────────────────────────────────────────────────────────────────
287
288#[cfg(test)]
289mod tests {
290    use super::*;
291
292    // ── Role ──────────────────────────────────────────────────────────────────
293
294    #[test]
295    fn test_role_display_user() {
296        assert_eq!(Role::User.to_string(), "user");
297    }
298
299    #[test]
300    fn test_role_display_assistant() {
301        assert_eq!(Role::Assistant.to_string(), "assistant");
302    }
303
304    #[test]
305    fn test_role_display_system() {
306        assert_eq!(Role::System.to_string(), "system");
307    }
308
309    #[test]
310    fn test_role_serialize_roundtrip() {
311        let json = serde_json::to_string(&Role::Assistant).unwrap();
312        let back: Role = serde_json::from_str(&json).unwrap();
313        assert_eq!(back, Role::Assistant);
314    }
315
316    // ── DialogueTurn ─────────────────────────────────────────────────────────
317
318    #[test]
319    fn test_dialogue_turn_new() {
320        let t = DialogueTurn::new(Role::User, "hello");
321        assert_eq!(t.role, Role::User);
322        assert_eq!(t.content, "hello");
323    }
324
325    #[test]
326    fn test_dialogue_turn_serialize_roundtrip() {
327        let t = DialogueTurn::new(Role::Assistant, "hi there");
328        let json = serde_json::to_string(&t).unwrap();
329        let back: DialogueTurn = serde_json::from_str(&json).unwrap();
330        assert_eq!(back.content, "hi there");
331    }
332
333    // ── DialogueSession ───────────────────────────────────────────────────────
334
335    #[test]
336    fn test_session_new_is_empty() {
337        let s = DialogueSession::new("sess-1");
338        assert_eq!(s.session_id, "sess-1");
339        assert_eq!(s.turn_count(), 0);
340        assert!(s.summary().is_none());
341    }
342
343    #[test]
344    fn test_append_turn_increments_count() {
345        let mut s = DialogueSession::new("s");
346        s.append_turn(Role::User, "hello").unwrap();
347        assert_eq!(s.turn_count(), 1);
348    }
349
350    #[test]
351    fn test_append_turn_rejects_blank_content() {
352        let mut s = DialogueSession::new("s");
353        let err = s.append_turn(Role::User, "   ").unwrap_err();
354        assert!(matches!(err, AgentRuntimeError::Memory(_)));
355    }
356
357    #[test]
358    fn test_get_history_returns_turns_in_order() {
359        let mut s = DialogueSession::new("s");
360        s.append_turn(Role::User, "first").unwrap();
361        s.append_turn(Role::Assistant, "second").unwrap();
362        let h = s.get_history();
363        assert_eq!(h[0].content, "first");
364        assert_eq!(h[1].content, "second");
365    }
366
367    #[test]
368    fn test_summarize_if_long_truncates_and_sets_summary() {
369        let mut s = DialogueSession::new("s");
370        for i in 0..5 {
371            s.append_turn(Role::User, format!("msg {i}")).unwrap();
372        }
373        let truncated = s.summarize_if_long(2).unwrap();
374        assert!(truncated);
375        assert_eq!(s.turn_count(), 2);
376        assert!(s.summary().is_some());
377        assert!(s.summary().unwrap().contains("3"));
378    }
379
380    #[test]
381    fn test_summarize_if_long_noop_when_under_limit() {
382        let mut s = DialogueSession::new("s");
383        s.append_turn(Role::User, "only turn").unwrap();
384        let truncated = s.summarize_if_long(10).unwrap();
385        assert!(!truncated);
386        assert!(s.summary().is_none());
387    }
388
389    #[test]
390    fn test_summarize_if_long_rejects_zero_keep() {
391        let mut s = DialogueSession::new("s");
392        s.append_turn(Role::User, "msg").unwrap();
393        assert!(s.summarize_if_long(0).is_err());
394    }
395
396    #[test]
397    fn test_summarize_retains_last_n_turns() {
398        let mut s = DialogueSession::new("s");
399        for i in 0..6 {
400            s.append_turn(Role::User, format!("turn {i}")).unwrap();
401        }
402        s.summarize_if_long(3).unwrap();
403        let h = s.get_history();
404        assert_eq!(h[0].content, "turn 3");
405        assert_eq!(h[2].content, "turn 5");
406    }
407
408    #[test]
409    fn test_build_context_includes_turns() {
410        let mut s = DialogueSession::new("s");
411        s.append_turn(Role::User, "hi").unwrap();
412        let ctx = s.build_context(None);
413        assert!(ctx.contains("user: hi"));
414    }
415
416    #[test]
417    fn test_build_context_includes_episodic_prefix() {
418        let mut s = DialogueSession::new("s");
419        s.append_turn(Role::User, "hi").unwrap();
420        let ctx = s.build_context(Some("Rust is fast."));
421        assert!(ctx.contains("Memory context"));
422        assert!(ctx.contains("Rust is fast."));
423    }
424
425    #[test]
426    fn test_build_context_includes_summary() {
427        let mut s = DialogueSession::new("s");
428        for i in 0..5 {
429            s.append_turn(Role::User, format!("m{i}")).unwrap();
430        }
431        s.summarize_if_long(2).unwrap();
432        let ctx = s.build_context(None);
433        assert!(ctx.contains("omitted"));
434    }
435
436    #[test]
437    fn test_session_serialize_roundtrip() {
438        let mut s = DialogueSession::new("rt");
439        s.append_turn(Role::User, "hello").unwrap();
440        let json = serde_json::to_string(&s).unwrap();
441        let back: DialogueSession = serde_json::from_str(&json).unwrap();
442        assert_eq!(back.session_id, "rt");
443        assert_eq!(back.turn_count(), 1);
444    }
445
446    // ── DialogueStore ─────────────────────────────────────────────────────────
447
448    #[test]
449    fn test_store_starts_empty() {
450        let store = DialogueStore::new();
451        assert_eq!(store.session_count(), 0);
452    }
453
454    #[test]
455    fn test_store_upsert_and_get() {
456        let store = DialogueStore::new();
457        let mut s = DialogueSession::new("abc");
458        s.append_turn(Role::User, "hello").unwrap();
459        store.upsert(s).unwrap();
460        let fetched = store.get("abc").unwrap().unwrap();
461        assert_eq!(fetched.turn_count(), 1);
462    }
463
464    #[test]
465    fn test_store_get_missing_returns_none() {
466        let store = DialogueStore::new();
467        assert!(store.get("missing").unwrap().is_none());
468    }
469
470    #[test]
471    fn test_store_append_turn_creates_session() {
472        let store = DialogueStore::new();
473        store
474            .append_turn("new-sess", Role::User, "first message")
475            .unwrap();
476        let s = store.get("new-sess").unwrap().unwrap();
477        assert_eq!(s.turn_count(), 1);
478    }
479
480    #[test]
481    fn test_store_append_turn_to_existing_session() {
482        let store = DialogueStore::new();
483        store.append_turn("s", Role::User, "a").unwrap();
484        store.append_turn("s", Role::Assistant, "b").unwrap();
485        let s = store.get("s").unwrap().unwrap();
486        assert_eq!(s.turn_count(), 2);
487    }
488
489    #[test]
490    fn test_store_remove_existing() {
491        let store = DialogueStore::new();
492        store.append_turn("x", Role::User, "hi").unwrap();
493        let removed = store.remove("x").unwrap();
494        assert!(removed.is_some());
495        assert_eq!(store.session_count(), 0);
496    }
497
498    #[test]
499    fn test_store_remove_missing_returns_none() {
500        let store = DialogueStore::new();
501        assert!(store.remove("ghost").unwrap().is_none());
502    }
503
504    #[test]
505    fn test_store_session_ids_sorted() {
506        let store = DialogueStore::new();
507        store.append_turn("z-sess", Role::User, "z").unwrap();
508        store.append_turn("a-sess", Role::User, "a").unwrap();
509        store.append_turn("m-sess", Role::User, "m").unwrap();
510        let ids = store.session_ids();
511        assert_eq!(ids, vec!["a-sess", "m-sess", "z-sess"]);
512    }
513
514    #[test]
515    fn test_store_clone_shares_state() {
516        let store = DialogueStore::new();
517        let clone = store.clone();
518        store.append_turn("shared", Role::User, "hi").unwrap();
519        assert_eq!(clone.session_count(), 1);
520    }
521
522    // ── DialogueContext ───────────────────────────────────────────────────────
523
524    #[test]
525    fn test_dialogue_context_build_no_episodic() {
526        let mut s = DialogueSession::new("ctx-test");
527        s.append_turn(Role::User, "what is 2+2?").unwrap();
528        let ctx = DialogueContext::build(&s, None);
529        assert_eq!(ctx.session_id, "ctx-test");
530        assert_eq!(ctx.turn_count, 1);
531        assert!(!ctx.has_summary);
532        assert!(ctx.prompt.contains("user: what is 2+2?"));
533    }
534
535    #[test]
536    fn test_dialogue_context_build_with_episodic() {
537        let mut s = DialogueSession::new("ctx-ep");
538        s.append_turn(Role::User, "hello").unwrap();
539        let ctx = DialogueContext::build(&s, Some("fact: Rust is fast"));
540        assert!(ctx.prompt.contains("Memory context"));
541        assert!(ctx.prompt.contains("Rust is fast"));
542    }
543
544    #[test]
545    fn test_dialogue_context_has_summary_flag() {
546        let mut s = DialogueSession::new("summ");
547        for i in 0..4 {
548            s.append_turn(Role::User, format!("m{i}")).unwrap();
549        }
550        s.summarize_if_long(2).unwrap();
551        let ctx = DialogueContext::build(&s, None);
552        assert!(ctx.has_summary);
553    }
554
555    #[test]
556    fn test_dialogue_context_serialize_roundtrip() {
557        let s = DialogueSession::new("ser");
558        let ctx = DialogueContext::build(&s, None);
559        let json = serde_json::to_string(&ctx).unwrap();
560        let back: DialogueContext = serde_json::from_str(&json).unwrap();
561        assert_eq!(back.session_id, "ser");
562    }
563}