1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
14pub struct Session {
15 pub id: Uuid,
17 pub title: Option<String>,
19 pub created_at: DateTime<Utc>,
21 pub messages: Vec<SessionMessage>,
23 #[serde(default, skip_serializing_if = "Option::is_none")]
25 pub user_id: Option<String>,
26 #[serde(default, skip_serializing_if = "Option::is_none")]
28 pub tenant_id: Option<String>,
29}
30
31#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct SessionMessage {
34 pub role: SessionRole,
36 pub content: String,
38 pub timestamp: DateTime<Utc>,
40}
41
42#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
44#[serde(rename_all = "snake_case")]
45pub enum SessionRole {
46 User,
48 Assistant,
50}
51
52const MAX_HISTORY_MESSAGES: usize = 100;
57
58pub 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
85pub trait SessionStore: Send + Sync {
87 fn create(&self, title: Option<String>) -> Result<Session, Error>;
89 fn get(&self, id: Uuid) -> Result<Option<Session>, Error>;
91 fn list(&self) -> Result<Vec<Session>, Error>;
93 fn delete(&self, id: Uuid) -> Result<bool, Error>;
95 fn add_message(&self, id: Uuid, message: SessionMessage) -> Result<(), Error>;
97
98 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 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
123pub struct InMemorySessionStore {
128 sessions: RwLock<HashMap<Uuid, Session>>,
129}
130
131impl InMemorySessionStore {
132 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 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 {
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 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 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 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 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 #[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 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 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 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 #[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(); 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 let all = store.list().unwrap();
560 assert_eq!(all.len(), 4);
561 }
562
563 #[test]
564 fn session_serde_backward_compat() {
565 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}