Skip to main content

heartbit_core/channel/
session.rs

1//! Session management for WebSocket-connected agent interactions.
2
3use parking_lot::RwLock;
4use std::collections::HashMap;
5
6use chrono::{DateTime, Utc};
7use serde::{Deserialize, Serialize};
8use uuid::Uuid;
9
10use crate::error::Error;
11
12/// A conversation session containing message history.
13#[derive(Debug, Clone, Serialize, Deserialize)]
14pub struct Session {
15    /// Session identifier.
16    pub id: Uuid,
17    /// Optional human-readable title.
18    pub title: Option<String>,
19    /// Creation timestamp (UTC).
20    pub created_at: DateTime<Utc>,
21    /// Ordered message history.
22    pub messages: Vec<SessionMessage>,
23    /// User who owns this session (multi-tenant isolation).
24    #[serde(default, skip_serializing_if = "Option::is_none")]
25    pub user_id: Option<String>,
26    /// Tenant that owns this session (multi-tenant isolation).
27    #[serde(default, skip_serializing_if = "Option::is_none")]
28    pub tenant_id: Option<String>,
29}
30
31/// A single message within a session.
32#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct SessionMessage {
34    /// Who authored the message.
35    pub role: SessionRole,
36    /// Message text.
37    pub content: String,
38    /// Authoring timestamp (UTC).
39    pub timestamp: DateTime<Utc>,
40}
41
42/// Role of a session message participant.
43#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
44#[serde(rename_all = "snake_case")]
45pub enum SessionRole {
46    /// User-authored message.
47    User,
48    /// Assistant-authored (agent) message.
49    Assistant,
50}
51
52/// Cap on how many prior messages are replayed into the prompt. R5: a long-lived
53/// daemon session would otherwise concatenate its ENTIRE history into every
54/// request, inflating context size/cost without bound and eventually risking
55/// context overflow. Only the most recent messages are included.
56const MAX_HISTORY_MESSAGES: usize = 100;
57
58/// Format session history as context to prepend to a new message.
59///
60/// When there is prior conversation history, returns the new message prefixed
61/// with a formatted history section (windowed to the last
62/// [`MAX_HISTORY_MESSAGES`]). When history is empty, returns the message
63/// unchanged.
64pub fn format_session_context(history: &[SessionMessage], message: &str) -> String {
65    if history.is_empty() {
66        return message.to_string();
67    }
68
69    let start = history.len().saturating_sub(MAX_HISTORY_MESSAGES);
70    let mut ctx = String::from("## Conversation history\n");
71    if start > 0 {
72        ctx.push_str(&format!("[... {start} earlier message(s) omitted ...]\n"));
73    }
74    for msg in &history[start..] {
75        let role = match msg.role {
76            SessionRole::User => "User",
77            SessionRole::Assistant => "Assistant",
78        };
79        ctx.push_str(&format!("{role}: {}\n", msg.content));
80    }
81    ctx.push_str(&format!("\n## Current message\n{message}"));
82    ctx
83}
84
85/// Trait for session persistence.
86pub trait SessionStore: Send + Sync {
87    /// Create a new session with an optional title.
88    fn create(&self, title: Option<String>) -> Result<Session, Error>;
89    /// Get a session by ID. Returns `None` if not found.
90    fn get(&self, id: Uuid) -> Result<Option<Session>, Error>;
91    /// List all sessions (most recent first).
92    fn list(&self) -> Result<Vec<Session>, Error>;
93    /// Delete a session. Returns true if found and deleted.
94    fn delete(&self, id: Uuid) -> Result<bool, Error>;
95    /// Append a message to an existing session.
96    fn add_message(&self, id: Uuid, message: SessionMessage) -> Result<(), Error>;
97
98    /// Create a session with user/tenant context for multi-tenant isolation.
99    /// Default: delegates to `create()` and patches user/tenant fields.
100    fn create_with_user(
101        &self,
102        title: Option<String>,
103        user_id: &str,
104        tenant_id: &str,
105    ) -> Result<Session, Error> {
106        let mut session = self.create(title)?;
107        session.user_id = Some(user_id.to_string());
108        session.tenant_id = Some(tenant_id.to_string());
109        Ok(session)
110    }
111
112    /// List sessions scoped to a tenant (most recent first).
113    /// Default: calls `list()` and filters in-memory.
114    fn list_for_tenant(&self, tenant_id: &str) -> Result<Vec<Session>, Error> {
115        let all = self.list()?;
116        Ok(all
117            .into_iter()
118            .filter(|s| s.tenant_id.as_deref() == Some(tenant_id))
119            .collect())
120    }
121}
122
123/// In-memory session store using `parking_lot::RwLock` (not tokio — matches
124/// codebase pattern for locks never held across `.await`; `parking_lot` is
125/// adopted on the channel hot path for ~2× faster uncontended reads, see T2
126/// in `tasks/performance-audit-heartbit-core-2026-05-06.md`).
127pub struct InMemorySessionStore {
128    sessions: RwLock<HashMap<Uuid, Session>>,
129}
130
131impl InMemorySessionStore {
132    /// Create an empty session store.
133    pub fn new() -> Self {
134        Self {
135            sessions: RwLock::new(HashMap::new()),
136        }
137    }
138}
139
140impl Default for InMemorySessionStore {
141    fn default() -> Self {
142        Self::new()
143    }
144}
145
146impl SessionStore for InMemorySessionStore {
147    fn create(&self, title: Option<String>) -> Result<Session, Error> {
148        let session = Session {
149            id: Uuid::new_v4(),
150            title,
151            created_at: Utc::now(),
152            messages: Vec::new(),
153            user_id: None,
154            tenant_id: None,
155        };
156        self.sessions.write().insert(session.id, session.clone());
157        Ok(session)
158    }
159
160    fn create_with_user(
161        &self,
162        title: Option<String>,
163        user_id: &str,
164        tenant_id: &str,
165    ) -> Result<Session, Error> {
166        let session = Session {
167            id: Uuid::new_v4(),
168            title,
169            created_at: Utc::now(),
170            messages: Vec::new(),
171            user_id: Some(user_id.to_string()),
172            tenant_id: Some(tenant_id.to_string()),
173        };
174        self.sessions.write().insert(session.id, session.clone());
175        Ok(session)
176    }
177
178    fn get(&self, id: Uuid) -> Result<Option<Session>, Error> {
179        Ok(self.sessions.read().get(&id).cloned())
180    }
181
182    fn list(&self) -> Result<Vec<Session>, Error> {
183        let mut list: Vec<Session> = self.sessions.read().values().cloned().collect();
184        // Most recent first
185        list.sort_by_key(|s| std::cmp::Reverse(s.created_at));
186        Ok(list)
187    }
188
189    fn delete(&self, id: Uuid) -> Result<bool, Error> {
190        Ok(self.sessions.write().remove(&id).is_some())
191    }
192
193    fn add_message(&self, id: Uuid, message: SessionMessage) -> Result<(), Error> {
194        match self.sessions.write().get_mut(&id) {
195            Some(session) => {
196                session.messages.push(message);
197                Ok(())
198            }
199            None => Err(Error::Channel(format!("session {id} not found"))),
200        }
201    }
202}
203
204#[cfg(test)]
205mod tests {
206    use super::*;
207
208    fn make_message(role: SessionRole, content: &str) -> SessionMessage {
209        SessionMessage {
210            role,
211            content: content.to_string(),
212            timestamp: Utc::now(),
213        }
214    }
215
216    #[test]
217    fn create_session() {
218        let store = InMemorySessionStore::new();
219        let session = store.create(None).unwrap();
220        assert!(session.title.is_none());
221        assert!(session.messages.is_empty());
222        assert!(session.created_at <= Utc::now());
223    }
224
225    #[test]
226    fn create_session_with_title() {
227        let store = InMemorySessionStore::new();
228        let session = store.create(Some("My Chat".to_string())).unwrap();
229        assert_eq!(session.title.as_deref(), Some("My Chat"));
230        assert!(session.messages.is_empty());
231    }
232
233    #[test]
234    fn get_existing_session() {
235        let store = InMemorySessionStore::new();
236        let created = store.create(Some("Test".to_string())).unwrap();
237        let fetched = store
238            .get(created.id)
239            .unwrap()
240            .expect("session should exist");
241        assert_eq!(fetched.id, created.id);
242        assert_eq!(fetched.title, created.title);
243        assert_eq!(fetched.messages.len(), created.messages.len());
244    }
245
246    #[test]
247    fn get_missing_session() {
248        let store = InMemorySessionStore::new();
249        let result = store.get(Uuid::new_v4()).unwrap();
250        assert!(result.is_none());
251    }
252
253    #[test]
254    fn list_empty() {
255        let store = InMemorySessionStore::new();
256        let list = store.list().unwrap();
257        assert!(list.is_empty());
258    }
259
260    #[test]
261    fn list_multiple() {
262        let store = InMemorySessionStore::new();
263        store.create(None).unwrap();
264        store.create(None).unwrap();
265        store.create(None).unwrap();
266        let list = store.list().unwrap();
267        assert_eq!(list.len(), 3);
268    }
269
270    #[test]
271    fn list_ordered_by_created_at() {
272        let store = InMemorySessionStore::new();
273        // Create sessions — they get Utc::now() timestamps so ordering depends on
274        // insertion order. To test sorting, manually insert with controlled timestamps.
275        {
276            let mut sessions = store.sessions.write();
277
278            let old = Session {
279                id: Uuid::new_v4(),
280                title: Some("old".to_string()),
281                created_at: Utc::now() - chrono::Duration::hours(2),
282                messages: Vec::new(),
283                user_id: None,
284                tenant_id: None,
285            };
286            let mid = Session {
287                id: Uuid::new_v4(),
288                title: Some("mid".to_string()),
289                created_at: Utc::now() - chrono::Duration::hours(1),
290                messages: Vec::new(),
291                user_id: None,
292                tenant_id: None,
293            };
294            let new = Session {
295                id: Uuid::new_v4(),
296                title: Some("new".to_string()),
297                created_at: Utc::now(),
298                messages: Vec::new(),
299                user_id: None,
300                tenant_id: None,
301            };
302
303            // Insert in non-sorted order
304            sessions.insert(mid.id, mid);
305            sessions.insert(old.id, old);
306            sessions.insert(new.id, new);
307        }
308
309        let list = store.list().unwrap();
310        assert_eq!(list.len(), 3);
311        assert_eq!(list[0].title.as_deref(), Some("new"));
312        assert_eq!(list[1].title.as_deref(), Some("mid"));
313        assert_eq!(list[2].title.as_deref(), Some("old"));
314    }
315
316    #[test]
317    fn delete_existing() {
318        let store = InMemorySessionStore::new();
319        let session = store.create(None).unwrap();
320        assert!(store.delete(session.id).unwrap());
321        assert!(store.get(session.id).unwrap().is_none());
322    }
323
324    #[test]
325    fn delete_missing() {
326        let store = InMemorySessionStore::new();
327        assert!(!store.delete(Uuid::new_v4()).unwrap());
328    }
329
330    #[test]
331    fn add_message_to_existing() {
332        let store = InMemorySessionStore::new();
333        let session = store.create(None).unwrap();
334        let msg = make_message(SessionRole::User, "hello");
335        store.add_message(session.id, msg).unwrap();
336
337        let fetched = store.get(session.id).unwrap().unwrap();
338        assert_eq!(fetched.messages.len(), 1);
339        assert_eq!(fetched.messages[0].content, "hello");
340        assert_eq!(fetched.messages[0].role, SessionRole::User);
341    }
342
343    #[test]
344    fn add_message_to_missing() {
345        let store = InMemorySessionStore::new();
346        let msg = make_message(SessionRole::User, "hello");
347        let err = store.add_message(Uuid::new_v4(), msg).unwrap_err();
348        assert!(err.to_string().contains("not found"));
349    }
350
351    #[test]
352    fn add_multiple_messages() {
353        let store = InMemorySessionStore::new();
354        let session = store.create(None).unwrap();
355
356        store
357            .add_message(session.id, make_message(SessionRole::User, "first"))
358            .unwrap();
359        store
360            .add_message(session.id, make_message(SessionRole::Assistant, "second"))
361            .unwrap();
362        store
363            .add_message(session.id, make_message(SessionRole::User, "third"))
364            .unwrap();
365
366        let fetched = store.get(session.id).unwrap().unwrap();
367        assert_eq!(fetched.messages.len(), 3);
368        assert_eq!(fetched.messages[0].content, "first");
369        assert_eq!(fetched.messages[1].content, "second");
370        assert_eq!(fetched.messages[2].content, "third");
371        assert_eq!(fetched.messages[0].role, SessionRole::User);
372        assert_eq!(fetched.messages[1].role, SessionRole::Assistant);
373        assert_eq!(fetched.messages[2].role, SessionRole::User);
374    }
375
376    #[test]
377    fn session_role_serde() {
378        let user_json = serde_json::to_string(&SessionRole::User).unwrap();
379        assert_eq!(user_json, "\"user\"");
380
381        let assistant_json = serde_json::to_string(&SessionRole::Assistant).unwrap();
382        assert_eq!(assistant_json, "\"assistant\"");
383
384        let user: SessionRole = serde_json::from_str("\"user\"").unwrap();
385        assert_eq!(user, SessionRole::User);
386
387        let assistant: SessionRole = serde_json::from_str("\"assistant\"").unwrap();
388        assert_eq!(assistant, SessionRole::Assistant);
389    }
390
391    #[test]
392    fn session_message_roundtrip() {
393        let msg = SessionMessage {
394            role: SessionRole::Assistant,
395            content: "Hello, world!".to_string(),
396            timestamp: Utc::now(),
397        };
398        let json = serde_json::to_string(&msg).unwrap();
399        let deserialized: SessionMessage = serde_json::from_str(&json).unwrap();
400        assert_eq!(deserialized.role, msg.role);
401        assert_eq!(deserialized.content, msg.content);
402        assert_eq!(deserialized.timestamp, msg.timestamp);
403    }
404
405    #[test]
406    fn concurrent_access() {
407        use std::sync::Arc;
408        use std::thread;
409
410        let store = Arc::new(InMemorySessionStore::new());
411        let mut handles = Vec::new();
412
413        // Spawn threads that create sessions
414        for i in 0..10 {
415            let store = Arc::clone(&store);
416            handles.push(thread::spawn(move || {
417                let session = store
418                    .create(Some(format!("thread-{i}")))
419                    .expect("create should succeed");
420                // Add a message to the session we just created
421                let msg = SessionMessage {
422                    role: SessionRole::User,
423                    content: format!("msg from thread {i}"),
424                    timestamp: Utc::now(),
425                };
426                store
427                    .add_message(session.id, msg)
428                    .expect("add_message should succeed");
429                session.id
430            }));
431        }
432
433        let ids: Vec<Uuid> = handles.into_iter().map(|h| h.join().unwrap()).collect();
434
435        // All sessions should exist with one message each
436        for id in &ids {
437            let session = store.get(*id).unwrap().expect("session should exist");
438            assert_eq!(session.messages.len(), 1);
439        }
440
441        let list = store.list().unwrap();
442        assert_eq!(list.len(), 10);
443    }
444
445    // --- format_session_context tests ---
446
447    #[test]
448    fn format_context_no_history() {
449        let result = format_session_context(&[], "Hello");
450        assert_eq!(result, "Hello");
451    }
452
453    #[test]
454    fn format_context_with_history() {
455        let history = vec![
456            make_message(SessionRole::User, "What is Rust?"),
457            make_message(SessionRole::Assistant, "A systems programming language."),
458        ];
459        let result = format_session_context(&history, "Tell me more");
460        assert!(result.contains("## Conversation history"));
461        assert!(result.contains("User: What is Rust?"));
462        assert!(result.contains("Assistant: A systems programming language."));
463        assert!(result.contains("## Current message"));
464        assert!(result.contains("Tell me more"));
465    }
466
467    #[test]
468    fn format_context_preserves_message_order() {
469        let history = vec![
470            make_message(SessionRole::User, "First"),
471            make_message(SessionRole::Assistant, "Second"),
472            make_message(SessionRole::User, "Third"),
473            make_message(SessionRole::Assistant, "Fourth"),
474        ];
475        let result = format_session_context(&history, "Fifth");
476        let first_pos = result.find("First").unwrap();
477        let second_pos = result.find("Second").unwrap();
478        let third_pos = result.find("Third").unwrap();
479        let fourth_pos = result.find("Fourth").unwrap();
480        let fifth_pos = result.find("Fifth").unwrap();
481        assert!(first_pos < second_pos);
482        assert!(second_pos < third_pos);
483        assert!(third_pos < fourth_pos);
484        assert!(fourth_pos < fifth_pos);
485    }
486
487    #[test]
488    fn format_context_windows_long_history() {
489        // R5: a very long history must not be replayed in full.
490        let history: Vec<SessionMessage> = (0..MAX_HISTORY_MESSAGES + 20)
491            .map(|i| make_message(SessionRole::User, &format!("msg-{i}")))
492            .collect();
493        let result = format_session_context(&history, "now");
494        // Oldest messages are omitted; the most recent are kept.
495        assert!(result.contains("earlier message(s) omitted"));
496        assert!(
497            !result.contains("msg-0:"),
498            "oldest message should be windowed out"
499        );
500        assert!(!result.contains("User: msg-0\n"));
501        assert!(result.contains(&format!("msg-{}", MAX_HISTORY_MESSAGES + 19)));
502        // Exactly MAX_HISTORY_MESSAGES history lines are emitted.
503        let history_lines = result.matches("User: msg-").count();
504        assert_eq!(history_lines, MAX_HISTORY_MESSAGES);
505    }
506
507    #[test]
508    fn format_context_single_message_history() {
509        let history = vec![make_message(SessionRole::User, "Prior question")];
510        let result = format_session_context(&history, "Follow-up");
511        assert!(result.contains("User: Prior question"));
512        assert!(result.contains("Follow-up"));
513    }
514
515    // --- Multi-tenant session tests ---
516
517    #[test]
518    fn create_with_user_sets_fields() {
519        let store = InMemorySessionStore::new();
520        let session = store
521            .create_with_user(Some("Test".into()), "alice", "acme")
522            .unwrap();
523        assert_eq!(session.user_id.as_deref(), Some("alice"));
524        assert_eq!(session.tenant_id.as_deref(), Some("acme"));
525        assert_eq!(session.title.as_deref(), Some("Test"));
526    }
527
528    #[test]
529    fn create_without_user_has_none_fields() {
530        let store = InMemorySessionStore::new();
531        let session = store.create(None).unwrap();
532        assert!(session.user_id.is_none());
533        assert!(session.tenant_id.is_none());
534    }
535
536    #[test]
537    fn list_for_tenant_filters_by_tenant() {
538        let store = InMemorySessionStore::new();
539        store
540            .create_with_user(Some("acme-1".into()), "alice", "acme")
541            .unwrap();
542        store
543            .create_with_user(Some("acme-2".into()), "bob", "acme")
544            .unwrap();
545        store
546            .create_with_user(Some("globex-1".into()), "charlie", "globex")
547            .unwrap();
548        store.create(Some("legacy".into())).unwrap(); // no tenant
549
550        let acme = store.list_for_tenant("acme").unwrap();
551        assert_eq!(acme.len(), 2);
552        assert!(acme.iter().all(|s| s.tenant_id.as_deref() == Some("acme")));
553
554        let globex = store.list_for_tenant("globex").unwrap();
555        assert_eq!(globex.len(), 1);
556        assert_eq!(globex[0].tenant_id.as_deref(), Some("globex"));
557
558        // Legacy sessions (no tenant) are not returned by list_for_tenant
559        let all = store.list().unwrap();
560        assert_eq!(all.len(), 4);
561    }
562
563    #[test]
564    fn session_serde_backward_compat() {
565        // Old JSON without user_id/tenant_id should deserialize with None
566        let json = r#"{"id":"00000000-0000-0000-0000-000000000000","title":"old","created_at":"2026-01-01T00:00:00Z","messages":[]}"#;
567        let session: Session = serde_json::from_str(json).unwrap();
568        assert!(session.user_id.is_none());
569        assert!(session.tenant_id.is_none());
570        assert_eq!(session.title.as_deref(), Some("old"));
571    }
572
573    #[test]
574    fn session_serde_with_tenant() {
575        let session = Session {
576            id: Uuid::nil(),
577            title: None,
578            created_at: Utc::now(),
579            messages: Vec::new(),
580            user_id: Some("alice".into()),
581            tenant_id: Some("acme".into()),
582        };
583        let json = serde_json::to_string(&session).unwrap();
584        assert!(json.contains(r#""user_id":"alice""#));
585        assert!(json.contains(r#""tenant_id":"acme""#));
586
587        let deserialized: Session = serde_json::from_str(&json).unwrap();
588        assert_eq!(deserialized.user_id.as_deref(), Some("alice"));
589        assert_eq!(deserialized.tenant_id.as_deref(), Some("acme"));
590    }
591}