1use crate::error::AgentRuntimeError;
20use crate::util::recover_lock;
21use serde::{Deserialize, Serialize};
22use std::collections::HashMap;
23use std::sync::{Arc, Mutex};
24
25#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
29#[serde(rename_all = "lowercase")]
30pub enum Role {
31 User,
33 Assistant,
35 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#[derive(Debug, Clone, Serialize, Deserialize)]
53pub struct DialogueTurn {
54 pub role: Role,
56 pub content: String,
58}
59
60impl DialogueTurn {
61 pub fn new(role: Role, content: impl Into<String>) -> Self {
63 Self {
64 role,
65 content: content.into(),
66 }
67 }
68}
69
70#[derive(Debug, Clone, Serialize, Deserialize)]
77pub struct DialogueSession {
78 pub session_id: String,
80 turns: Vec<DialogueTurn>,
82 summary: Option<String>,
84}
85
86impl DialogueSession {
87 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 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 pub fn get_history(&self) -> &[DialogueTurn] {
117 &self.turns
118 }
119
120 pub fn turn_count(&self) -> usize {
122 self.turns.len()
123 }
124
125 pub fn summary(&self) -> Option<&str> {
128 self.summary.as_deref()
129 }
130
131 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 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#[derive(Debug, Clone, Default)]
191pub struct DialogueStore {
192 sessions: Arc<Mutex<HashMap<String, DialogueSession>>>,
193}
194
195impl DialogueStore {
196 pub fn new() -> Self {
198 Self::default()
199 }
200
201 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 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 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 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 pub fn session_count(&self) -> usize {
238 let map = recover_lock(self.sessions.lock(), "DialogueStore::session_count");
239 map.len()
240 }
241
242 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#[derive(Debug, Clone, Serialize, Deserialize)]
259pub struct DialogueContext {
260 pub session_id: String,
262 pub prompt: String,
264 pub turn_count: usize,
266 pub has_summary: bool,
268}
269
270impl DialogueContext {
271 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#[cfg(test)]
289mod tests {
290 use super::*;
291
292 #[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 #[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 #[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 #[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 #[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}