1use std::{
31 collections::HashMap,
32 sync::Arc,
33 time::{Duration, SystemTime},
34};
35use tokio::sync::RwLock;
36
37use crate::SessionId;
38
39#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
45#[serde(rename_all = "lowercase")]
46pub enum Role {
47 System,
48 User,
49 Assistant,
50 Tool,
52}
53
54impl std::fmt::Display for Role {
55 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
56 let s = match self {
57 Role::System => "system",
58 Role::User => "user",
59 Role::Assistant => "assistant",
60 Role::Tool => "tool",
61 };
62 f.write_str(s)
63 }
64}
65
66#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
68pub struct Turn {
69 pub id: String,
70 pub role: Role,
71 pub content: String,
72 pub timestamp: SystemTime,
73 pub token_estimate: usize,
75 pub tags: Vec<String>,
77 pub is_summary: bool,
79}
80
81impl Turn {
82 fn new(role: Role, content: impl Into<String>) -> Self {
83 let content = content.into();
84 let token_estimate = estimate_tokens(&content);
85 Turn {
86 id: new_id(),
87 role,
88 content,
89 timestamp: SystemTime::now(),
90 token_estimate,
91 tags: Vec::new(),
92 is_summary: false,
93 }
94 }
95}
96
97#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
99pub struct Conversation {
100 pub id: String,
101 pub session_id: String,
102 pub turns: Vec<Turn>,
103 pub total_tokens: usize,
104 pub created_at: SystemTime,
105 pub last_active: SystemTime,
106 pub compressions: u32,
108 pub meta: HashMap<String, String>,
110}
111
112impl Conversation {
113 fn new(session_id: &str) -> Self {
114 let now = SystemTime::now();
115 Conversation {
116 id: new_id(),
117 session_id: session_id.to_owned(),
118 turns: Vec::new(),
119 total_tokens: 0,
120 created_at: now,
121 last_active: now,
122 compressions: 0,
123 meta: HashMap::new(),
124 }
125 }
126
127 fn push(&mut self, turn: Turn) {
128 self.total_tokens += turn.token_estimate;
129 self.last_active = SystemTime::now();
130 self.turns.push(turn);
131 }
132}
133
134#[derive(Debug, Clone)]
140pub struct ConversationConfig {
141 pub max_tokens: usize,
144
145 pub recency_keep: usize,
148
149 pub system_prompt: Option<String>,
152
153 pub ttl: Duration,
156
157 pub format: PromptFormat,
160}
161
162impl Default for ConversationConfig {
163 fn default() -> Self {
164 ConversationConfig {
165 max_tokens: 6_000,
166 recency_keep: 4,
167 system_prompt: None,
168 ttl: Duration::from_secs(2 * 3600),
169 format: PromptFormat::ChatMl,
170 }
171 }
172}
173
174#[derive(Debug, Clone, PartialEq, Eq)]
176pub enum PromptFormat {
177 ChatMl,
182
183 Markdown,
188
189 Inline,
191}
192
193#[derive(Clone)]
203pub struct ConversationManager {
204 store: Arc<RwLock<HashMap<String, Conversation>>>,
205 config: Arc<ConversationConfig>,
206}
207
208impl ConversationManager {
209 pub fn new(config: ConversationConfig) -> Self {
211 ConversationManager {
212 store: Arc::new(RwLock::new(HashMap::new())),
213 config: Arc::new(config),
214 }
215 }
216
217 pub async fn push_user(&self, session: &SessionId, content: impl Into<String>) {
223 self.push(session, Role::User, content).await;
224 }
225
226 pub async fn push_assistant(&self, session: &SessionId, content: impl Into<String>) {
228 self.push(session, Role::Assistant, content).await;
229 }
230
231 pub async fn push_tool(&self, session: &SessionId, content: impl Into<String>) {
233 self.push(session, Role::Tool, content).await;
234 }
235
236 pub async fn push_system(&self, session: &SessionId, content: impl Into<String>) {
238 self.push(session, Role::System, content).await;
239 }
240
241 async fn push(&self, session: &SessionId, role: Role, content: impl Into<String>) {
242 let key = session.as_str().to_owned();
243 let turn = Turn::new(role, content);
244 let max_tokens = self.config.max_tokens;
245 let recency_keep = self.config.recency_keep;
246
247 let mut guard = self.store.write().await;
248 let conv = guard.entry(key).or_insert_with(|| Conversation::new(session.as_str()));
249 conv.push(turn);
250
251 if conv.total_tokens > max_tokens {
252 compress_conversation(conv, recency_keep);
253 }
254 }
255
256 pub async fn build_prompt(&self, session: &SessionId, new_input: &str) -> String {
265 let key = session.as_str();
266 let guard = self.store.read().await;
267
268 let history: Vec<(Role, String)> = if let Some(conv) = guard.get(key) {
269 conv.turns
270 .iter()
271 .map(|t| (t.role.clone(), t.content.clone()))
272 .collect()
273 } else {
274 Vec::new()
275 };
276
277 drop(guard);
278 self.format_prompt(&history, new_input)
279 }
280
281 fn format_prompt(&self, history: &[(Role, String)], new_input: &str) -> String {
282 match self.config.format {
283 PromptFormat::ChatMl => {
284 let mut out = String::with_capacity(1024);
285 if let Some(sys) = &self.config.system_prompt {
286 out.push_str("<|system|>\n");
287 out.push_str(sys);
288 out.push_str("\n<|end|>\n");
289 }
290 for (role, content) in history {
291 match role {
292 Role::System => {
293 out.push_str("<|system|>\n");
294 out.push_str(content);
295 out.push_str("\n<|end|>\n");
296 }
297 Role::User => {
298 out.push_str("<|user|>\n");
299 out.push_str(content);
300 out.push_str("\n<|end|>\n");
301 }
302 Role::Assistant => {
303 out.push_str("<|assistant|>\n");
304 out.push_str(content);
305 out.push_str("\n<|end|>\n");
306 }
307 Role::Tool => {
308 out.push_str("<|tool|>\n");
309 out.push_str(content);
310 out.push_str("\n<|end|>\n");
311 }
312 }
313 }
314 out.push_str("<|user|>\n");
315 out.push_str(new_input);
316 out.push_str("\n<|end|>\n<|assistant|>\n");
317 out
318 }
319
320 PromptFormat::Markdown => {
321 let mut out = String::with_capacity(1024);
322 if let Some(sys) = &self.config.system_prompt {
323 out.push_str("## System\n\n");
324 out.push_str(sys);
325 out.push_str("\n\n---\n\n");
326 }
327 for (role, content) in history {
328 let header = match role {
329 Role::System => "## System",
330 Role::User => "## User",
331 Role::Assistant => "## Assistant",
332 Role::Tool => "## Tool",
333 };
334 out.push_str(header);
335 out.push_str("\n\n");
336 out.push_str(content);
337 out.push_str("\n\n");
338 }
339 out.push_str("## User\n\n");
340 out.push_str(new_input);
341 out.push_str("\n\n## Assistant\n\n");
342 out
343 }
344
345 PromptFormat::Inline => {
346 let mut out = String::with_capacity(1024);
347 if let Some(sys) = &self.config.system_prompt {
348 out.push_str("system: ");
349 out.push_str(sys);
350 out.push('\n');
351 }
352 for (role, content) in history {
353 out.push_str(&role.to_string());
354 out.push_str(": ");
355 out.push_str(content);
356 out.push('\n');
357 }
358 out.push_str("user: ");
359 out.push_str(new_input);
360 out.push('\n');
361 out
362 }
363 }
364 }
365
366 pub async fn get(&self, session: &SessionId) -> Option<Conversation> {
373 let guard = self.store.read().await;
374 guard.get(session.as_str()).cloned()
375 }
376
377 pub async fn turn_count(&self, session: &SessionId) -> usize {
379 let guard = self.store.read().await;
380 guard
381 .get(session.as_str())
382 .map(|c| c.turns.len())
383 .unwrap_or(0)
384 }
385
386 pub async fn token_count(&self, session: &SessionId) -> usize {
388 let guard = self.store.read().await;
389 guard
390 .get(session.as_str())
391 .map(|c| c.total_tokens)
392 .unwrap_or(0)
393 }
394
395 pub async fn set_meta(&self, session: &SessionId, key: impl Into<String>, value: impl Into<String>) {
397 let mut guard = self.store.write().await;
398 let conv = guard
399 .entry(session.as_str().to_owned())
400 .or_insert_with(|| Conversation::new(session.as_str()));
401 conv.meta.insert(key.into(), value.into());
402 }
403
404 pub async fn clear(&self, session: &SessionId) {
406 let mut guard = self.store.write().await;
407 guard.remove(session.as_str());
408 }
409
410 pub async fn evict_stale(&self) {
414 let ttl = self.config.ttl;
415 let mut guard = self.store.write().await;
416 guard.retain(|_, conv| {
417 conv.last_active
418 .elapsed()
419 .map(|d| d < ttl)
420 .unwrap_or(true)
421 });
422 }
423
424 pub async fn export_json(&self, session: &SessionId) -> Option<String> {
426 let guard = self.store.read().await;
427 guard
428 .get(session.as_str())
429 .and_then(|c| serde_json::to_string_pretty(c).ok())
430 }
431
432 pub async fn active_sessions(&self) -> Vec<String> {
434 let guard = self.store.read().await;
435 guard.keys().cloned().collect()
436 }
437
438 pub async fn len(&self) -> usize {
440 self.store.read().await.len()
441 }
442
443 pub async fn is_empty(&self) -> bool {
445 self.store.read().await.is_empty()
446 }
447}
448
449fn compress_conversation(conv: &mut Conversation, recency_keep: usize) {
456 let len = conv.turns.len();
457 if len <= recency_keep {
458 return;
459 }
460
461 let split = len.saturating_sub(recency_keep);
462 let old_turns: Vec<Turn> = conv.turns.drain(..split).collect();
463
464 let mut summary_parts = Vec::new();
466 let mut user_count = 0usize;
467 let mut assistant_count = 0usize;
468
469 for t in &old_turns {
470 match t.role {
471 Role::User => {
472 user_count += 1;
473 if t.token_estimate > 20 && summary_parts.len() < 5 {
475 let excerpt = truncate_str(&t.content, 120);
476 summary_parts.push(format!("User said: {excerpt}"));
477 }
478 }
479 Role::Assistant => {
480 assistant_count += 1;
481 if t.token_estimate > 30 && summary_parts.len() < 8 {
482 let excerpt = truncate_str(&t.content, 150);
483 summary_parts.push(format!("Assistant replied: {excerpt}"));
484 }
485 }
486 Role::Tool => {
487 if summary_parts.len() < 8 {
488 let excerpt = truncate_str(&t.content, 80);
489 summary_parts.push(format!("Tool returned: {excerpt}"));
490 }
491 }
492 Role::System => {}
493 }
494 }
495
496 let header = format!(
497 "[Conversation summary — {user_count} user turns, {assistant_count} assistant turns compressed]"
498 );
499
500 let summary_text = if summary_parts.is_empty() {
501 header
502 } else {
503 format!("{header}\n{}", summary_parts.join("\n"))
504 };
505
506 let summary_turn = Turn {
507 id: new_id(),
508 role: Role::System,
509 content: summary_text,
510 timestamp: SystemTime::now(),
511 token_estimate: estimate_tokens(&{
512
513 "[Conversation summary]".to_owned()
514 }),
515 tags: vec!["summary".to_owned()],
516 is_summary: true,
517 };
518
519 conv.total_tokens = summary_turn.token_estimate
521 + conv.turns.iter().map(|t| t.token_estimate).sum::<usize>();
522
523 conv.turns.insert(0, summary_turn);
524 conv.compressions += 1;
525}
526
527pub fn estimate_tokens(s: &str) -> usize {
533 (s.len() / 4).max(1)
534}
535
536fn truncate_str(s: &str, max_chars: usize) -> String {
537 if s.len() <= max_chars {
538 s.to_owned()
539 } else {
540 let mut end = max_chars;
541 while !s.is_char_boundary(end) {
542 end -= 1;
543 }
544 format!("{}…", &s[..end])
545 }
546}
547
548fn new_id() -> String {
549 use std::time::{SystemTime, UNIX_EPOCH};
550 let t = SystemTime::now()
551 .duration_since(UNIX_EPOCH)
552 .unwrap_or_default()
553 .subsec_nanos();
554 format!("{:08x}", t ^ (t.wrapping_mul(0x9e37_79b9)))
555}
556
557#[cfg(test)]
562mod tests {
563 use super::*;
564
565 fn sid(s: &str) -> SessionId {
566 SessionId::new(s)
567 }
568
569 #[tokio::test]
570 async fn test_push_and_count() {
571 let mgr = ConversationManager::new(ConversationConfig::default());
572 let s = sid("s1");
573 mgr.push_user(&s, "Hello").await;
574 mgr.push_assistant(&s, "Hi there").await;
575 assert_eq!(mgr.turn_count(&s).await, 2);
576 }
577
578 #[tokio::test]
579 async fn test_build_prompt_chatml() {
580 let mgr = ConversationManager::new(ConversationConfig {
581 format: PromptFormat::ChatMl,
582 system_prompt: Some("You are helpful.".into()),
583 ..Default::default()
584 });
585 let s = sid("s2");
586 mgr.push_user(&s, "Q1").await;
587 mgr.push_assistant(&s, "A1").await;
588 let prompt = mgr.build_prompt(&s, "Q2").await;
589 assert!(prompt.contains("<|system|>"));
590 assert!(prompt.contains("Q1"));
591 assert!(prompt.contains("A1"));
592 assert!(prompt.contains("Q2"));
593 assert!(prompt.ends_with("<|assistant|>\n"));
594 }
595
596 #[tokio::test]
597 async fn test_build_prompt_markdown() {
598 let mgr = ConversationManager::new(ConversationConfig {
599 format: PromptFormat::Markdown,
600 ..Default::default()
601 });
602 let s = sid("s3");
603 mgr.push_user(&s, "What?").await;
604 let prompt = mgr.build_prompt(&s, "Why?").await;
605 assert!(prompt.contains("## User"));
606 assert!(prompt.contains("What?"));
607 assert!(prompt.contains("## Assistant"));
608 }
609
610 #[tokio::test]
611 async fn test_compression_triggered() {
612 let mgr = ConversationManager::new(ConversationConfig {
613 max_tokens: 50,
614 recency_keep: 2,
615 ..Default::default()
616 });
617 let s = sid("s4");
618 for i in 0..10 {
620 mgr.push_user(&s, format!("This is user message number {i} which has some content in it."))
621 .await;
622 mgr.push_assistant(&s, format!("This is the assistant reply to message {i}."))
623 .await;
624 }
625 let conv = mgr.get(&s).await.unwrap();
626 assert!(conv.compressions > 0, "should have triggered compression");
627 }
628
629 #[tokio::test]
630 async fn test_clear() {
631 let mgr = ConversationManager::new(ConversationConfig::default());
632 let s = sid("s5");
633 mgr.push_user(&s, "hello").await;
634 mgr.clear(&s).await;
635 assert_eq!(mgr.turn_count(&s).await, 0);
636 }
637
638 #[tokio::test]
639 async fn test_export_json() {
640 let mgr = ConversationManager::new(ConversationConfig::default());
641 let s = sid("s6");
642 mgr.push_user(&s, "test").await;
643 let json = mgr.export_json(&s).await.unwrap();
644 assert!(json.contains("\"content\""));
645 assert!(json.contains("test"));
646 }
647
648 #[tokio::test]
649 async fn test_evict_stale_leaves_active() {
650 let mgr = ConversationManager::new(ConversationConfig {
651 ttl: Duration::from_secs(3600),
652 ..Default::default()
653 });
654 let s = sid("s7");
655 mgr.push_user(&s, "hello").await;
656 mgr.evict_stale().await;
657 assert_eq!(mgr.len().await, 1);
658 }
659
660 #[test]
661 fn test_estimate_tokens() {
662 assert_eq!(estimate_tokens("hello"), 1);
663 assert_eq!(estimate_tokens("hello world"), 2);
664 assert!(estimate_tokens("a".repeat(100).as_str()) == 25);
665 }
666}