1use macp_core::session::Session;
2use std::collections::{BinaryHeap, HashMap};
3use std::fs;
4use std::path::{Path, PathBuf};
5use std::sync::Arc;
6use tokio::sync::RwLock;
7
8#[derive(serde::Serialize, serde::Deserialize)]
9pub struct PersistedRoot {
10 pub uri: String,
11 pub name: String,
12}
13
14#[derive(serde::Serialize, serde::Deserialize)]
15pub struct PersistedSession {
16 #[serde(default = "default_schema_version")]
17 pub schema_version: u32,
18 pub session_id: String,
19 pub state: macp_core::session::SessionState,
20 pub ttl_expiry: i64,
21 #[serde(default)]
22 pub ttl_ms: i64,
23 pub started_at_unix_ms: i64,
24 pub resolution: Option<Vec<u8>>,
25 pub mode: String,
26 pub mode_state: Vec<u8>,
27 pub participants: Vec<String>,
28 pub seen_message_ids: Vec<String>,
29 pub intent: String,
30 pub mode_version: String,
31 pub configuration_version: String,
32 pub policy_version: String,
33 #[serde(default)]
34 pub context_id: String,
35 #[serde(default)]
36 pub extensions: HashMap<String, Vec<u8>>,
37 pub roots: Vec<PersistedRoot>,
38 pub initiator_sender: String,
39 #[serde(default)]
40 pub policy_definition: Option<macp_core::policy::PolicyDefinition>,
41 #[serde(default)]
42 pub suspended_at_ms: Option<i64>,
43 #[serde(default)]
44 pub accumulated_suspended_ms: i64,
45 #[serde(default)]
48 pub semantics_rev: u32,
49 #[serde(default)]
52 pub max_suspend_ms: i64,
53}
54
55fn default_schema_version() -> u32 {
56 2
57}
58
59impl From<&Session> for PersistedSession {
60 fn from(session: &Session) -> Self {
61 Self {
62 schema_version: 2,
63 session_id: session.session_id.clone(),
64 state: session.state.clone(),
65 ttl_expiry: session.ttl_expiry,
66 ttl_ms: session.ttl_ms,
67 started_at_unix_ms: session.started_at_unix_ms,
68 resolution: session.resolution.clone(),
69 mode: session.mode.clone(),
70 mode_state: session.mode_state.clone(),
71 participants: session.participants.clone(),
72 seen_message_ids: session.seen_message_ids.iter().cloned().collect(),
73 intent: session.intent.clone(),
74 mode_version: session.mode_version.clone(),
75 configuration_version: session.configuration_version.clone(),
76 policy_version: session.policy_version.clone(),
77 context_id: session.context_id.clone(),
78 extensions: session.extensions.clone(),
79 roots: session
80 .roots
81 .iter()
82 .map(|root| PersistedRoot {
83 uri: root.uri.clone(),
84 name: root.name.clone(),
85 })
86 .collect(),
87 initiator_sender: session.initiator_sender.clone(),
88 policy_definition: session.policy_definition.clone(),
89 suspended_at_ms: session.suspended_at_ms,
90 accumulated_suspended_ms: session.accumulated_suspended_ms,
91 semantics_rev: session.semantics_rev,
92 max_suspend_ms: session.max_suspend_ms,
93 }
94 }
95}
96
97impl From<PersistedSession> for Session {
98 fn from(session: PersistedSession) -> Self {
99 let ttl_ms = if session.ttl_ms > 0 {
100 session.ttl_ms
101 } else {
102 session
104 .ttl_expiry
105 .saturating_sub(session.started_at_unix_ms)
106 };
107 Session::builder(session.session_id, session.mode, session.initiator_sender)
108 .state(session.state)
109 .ttl_expiry(session.ttl_expiry)
110 .ttl_ms(ttl_ms)
111 .started_at_unix_ms(session.started_at_unix_ms)
112 .resolution(session.resolution)
113 .mode_state(session.mode_state)
114 .participants(session.participants)
115 .seen_message_ids(session.seen_message_ids.into_iter().collect())
116 .intent(session.intent)
117 .mode_version(session.mode_version)
118 .configuration_version(session.configuration_version)
119 .policy_version(session.policy_version)
120 .context_id(session.context_id)
121 .extensions(session.extensions)
122 .roots(
123 session
124 .roots
125 .into_iter()
126 .map(|root| macp_pb::pb::Root {
127 uri: root.uri,
128 name: root.name,
129 })
130 .collect(),
131 )
132 .policy_definition(session.policy_definition)
133 .suspended_at_ms(session.suspended_at_ms)
134 .accumulated_suspended_ms(session.accumulated_suspended_ms)
135 .semantics_rev(session.semantics_rev)
136 .max_suspend_ms(session.max_suspend_ms)
137 .build()
138 }
139}
140
141pub type SharedSession = Arc<tokio::sync::Mutex<Session>>;
149
150pub struct SessionRegistry {
151 pub sessions: RwLock<HashMap<String, SharedSession>>,
152 persistence_path: Option<PathBuf>,
153}
154
155impl Default for SessionRegistry {
156 fn default() -> Self {
157 Self::new()
158 }
159}
160
161impl SessionRegistry {
162 pub fn new() -> Self {
163 Self {
164 sessions: RwLock::new(HashMap::new()),
165 persistence_path: None,
166 }
167 }
168
169 pub fn with_persistence<P: AsRef<Path>>(dir: P) -> std::io::Result<Self> {
170 let dir = dir.as_ref().to_path_buf();
171 fs::create_dir_all(&dir)?;
172 let path = dir.join("sessions.json");
173 let sessions = Self::load_sessions(&path)?;
174 Ok(Self {
175 sessions: RwLock::new(sessions),
176 persistence_path: Some(path),
177 })
178 }
179
180 fn load_sessions(path: &Path) -> std::io::Result<HashMap<String, SharedSession>> {
181 if !path.exists() {
182 return Ok(HashMap::new());
183 }
184 let bytes = fs::read(path)?;
185 let persisted: HashMap<String, PersistedSession> = match serde_json::from_slice(&bytes) {
186 Ok(v) => v,
187 Err(e) => {
188 eprintln!("warning: failed to deserialize sessions from {}: {e}; starting with empty state", path.display());
189 HashMap::new()
190 }
191 };
192 Ok(persisted
193 .into_iter()
194 .map(|(id, mut record)| {
195 if record.session_id != id {
205 tracing::warn!(
206 map_key = %id,
207 record_session_id = %record.session_id,
208 path = %path.display(),
209 "persisted session key disagrees with its session_id; \
210 repairing to the map key"
211 );
212 record.session_id.clone_from(&id);
213 }
214 let session: Session = record.into();
215 (id, Arc::new(tokio::sync::Mutex::new(session)))
216 })
217 .collect())
218 }
219
220 fn persist_map(
221 path: &Path,
222 sessions: &HashMap<String, PersistedSession>,
223 ) -> std::io::Result<()> {
224 let bytes = serde_json::to_vec_pretty(sessions)?;
225 let tmp_path = path.with_extension("json.tmp");
226 fs::write(&tmp_path, bytes)?;
227 fs::rename(&tmp_path, path)
228 }
229
230 pub async fn persist_snapshot(&self) -> std::io::Result<()> {
233 let Some(path) = self.persistence_path.clone() else {
234 return Ok(());
235 };
236 let arcs: Vec<(String, SharedSession)> = {
237 let guard = self.sessions.read().await;
238 guard
239 .iter()
240 .map(|(id, arc)| (id.clone(), Arc::clone(arc)))
241 .collect()
242 };
243 let mut persisted = HashMap::with_capacity(arcs.len());
244 for (id, arc) in arcs {
245 let session = arc.lock().await;
246 persisted.insert(id, PersistedSession::from(&*session));
247 }
248 Self::persist_map(&path, &persisted)
249 }
250
251 pub async fn get_shared(&self, session_id: &str) -> Option<SharedSession> {
253 let guard = self.sessions.read().await;
254 guard.get(session_id).cloned()
255 }
256
257 pub async fn get_session(&self, session_id: &str) -> Option<Session> {
258 let arc = self.get_shared(session_id).await?;
259 let session = arc.lock().await;
260 Some(session.clone())
261 }
262
263 pub async fn get_all_sessions(&self) -> Vec<Session> {
272 let arcs: Vec<SharedSession> = {
273 let guard = self.sessions.read().await;
274 guard.values().cloned().collect()
275 };
276 let mut out = Vec::with_capacity(arcs.len());
277 for arc in arcs {
278 out.push(arc.lock().await.clone());
279 }
280 out
281 }
282
283 pub async fn shared_sessions(&self) -> Vec<SharedSession> {
310 let guard = self.sessions.read().await;
311 guard.values().map(Arc::clone).collect()
312 }
313
314 pub async fn session_ids_after(&self, after: Option<&str>, limit: usize) -> Vec<String> {
334 if limit == 0 {
335 return Vec::new();
336 }
337 {
341 let guard = self.sessions.read().await;
342 let capacity = limit.saturating_add(1).min(guard.len().saturating_add(1));
349 let mut heap: BinaryHeap<&String> = BinaryHeap::with_capacity(capacity);
350 for key in guard.keys() {
351 if after.is_none_or(|a| key.as_str() > a) {
352 heap.push(key);
353 if heap.len() > limit {
354 heap.pop();
355 }
356 }
357 }
358 heap.into_sorted_vec().into_iter().cloned().collect()
359 }
360 }
361
362 pub async fn insert_recovered_session(&self, session_id: String, session: Session) {
363 debug_assert_eq!(
375 session.session_id, session_id,
376 "registry map key must equal Session::session_id — ListSessions paging \
377 orders by the key but emits the field (plan D1)"
378 );
379 {
380 let mut guard = self.sessions.write().await;
381 guard.insert(session_id, Arc::new(tokio::sync::Mutex::new(session)));
382 }
383 let _ = self.persist_snapshot().await;
384 }
385
386 pub async fn count_open_sessions_for_initiator(&self, sender: &str) -> usize {
387 let now = chrono::Utc::now().timestamp_millis();
388 let arcs: Vec<SharedSession> = {
389 let guard = self.sessions.read().await;
390 guard.values().cloned().collect()
391 };
392 let mut count = 0;
393 for arc in arcs {
394 let counts = match arc.try_lock() {
397 Ok(session) => {
398 session.initiator_sender == sender
399 && session.state == macp_core::session::SessionState::Open
400 && now <= session.ttl_expiry
401 }
402 Err(_) => true,
403 };
404 if counts {
405 count += 1;
406 }
407 }
408 count
409 }
410}
411
412#[cfg(test)]
413mod tests {
414 use super::*;
415 use macp_core::session::{Session, SessionState};
416 use std::collections::HashSet;
417 use std::time::{SystemTime, UNIX_EPOCH};
418
419 fn sample_session(id: &str) -> Session {
420 Session::builder(id, "macp.mode.decision.v1", "alice")
421 .ttl_expiry(10)
422 .ttl_ms(9)
423 .started_at_unix_ms(1)
424 .mode_state(vec![1, 2, 3])
425 .participants(vec!["alice".into()])
426 .seen_message_ids(HashSet::from(["m1".into()]))
427 .intent("intent")
428 .mode_version("1.0.0")
429 .configuration_version("cfg")
430 .policy_version("pol")
431 .context_id("test-ctx")
432 .roots(vec![macp_pb::pb::Root {
433 uri: "root://1".into(),
434 name: "r1".into(),
435 }])
436 .build()
437 }
438
439 async fn registry_with(ids: &[String]) -> SessionRegistry {
442 let registry = SessionRegistry::new();
443 for id in ids {
444 registry
445 .insert_recovered_session(id.clone(), sample_session(id))
446 .await;
447 }
448 registry
449 }
450
451 fn sort_then_truncate_reference(
454 ids: &[String],
455 after: Option<&str>,
456 limit: usize,
457 ) -> Vec<String> {
458 let mut sorted: Vec<String> = ids.to_vec();
459 sorted.sort();
460 sorted
461 .into_iter()
462 .filter(|id| after.is_none_or(|a| id.as_str() > a))
463 .take(limit)
464 .collect()
465 }
466
467 fn deterministic_ids(count: usize) -> Vec<String> {
471 let mut state: u64 = 0x2545_F491_4F6C_DD1D;
472 let mut ids = Vec::with_capacity(count);
473 for i in 0..count {
474 state = state
475 .wrapping_mul(6_364_136_223_846_793_005)
476 .wrapping_add(1_442_695_040_888_963_407);
477 ids.push(format!("sess-{:016x}-{i:04}", state >> 16));
479 }
480 ids
481 }
482
483 #[tokio::test]
484 async fn shared_sessions_snapshots_every_session_once() {
485 let ids: Vec<String> = ["delta", "alpha", "charlie", "bravo"]
486 .iter()
487 .map(|s| s.to_string())
488 .collect();
489 let registry = registry_with(&ids).await;
490
491 let handles = registry.shared_sessions().await;
492 assert_eq!(handles.len(), 4);
493 let mut seen = Vec::new();
494 for handle in &handles {
495 seen.push(handle.lock().await.session_id.clone());
496 }
497 seen.sort();
500 assert_eq!(seen, vec!["alpha", "bravo", "charlie", "delta"]);
501
502 {
506 let mut guard = registry.sessions.write().await;
507 guard.remove("alpha");
508 guard.remove("bravo");
509 }
510 assert!(registry.get_session("alpha").await.is_none());
511 let mut after = Vec::new();
512 for handle in &handles {
513 after.push(handle.lock().await.session_id.clone());
514 }
515 after.sort();
516 assert_eq!(after, vec!["alpha", "bravo", "charlie", "delta"]);
517
518 assert!(SessionRegistry::new().shared_sessions().await.is_empty());
519 }
520
521 #[tokio::test]
522 async fn session_ids_after_returns_ascending_ids() {
523 let ids: Vec<String> = ["delta", "alpha", "charlie", "bravo"]
524 .iter()
525 .map(|s| s.to_string())
526 .collect();
527 let registry = registry_with(&ids).await;
528
529 let page = registry.session_ids_after(None, 10).await;
530 assert_eq!(page, vec!["alpha", "bravo", "charlie", "delta"]);
531
532 let page = registry.session_ids_after(None, 2).await;
534 assert_eq!(page, vec!["alpha", "bravo"]);
535 }
536
537 #[tokio::test]
538 async fn session_ids_after_respects_limit() {
539 let ids: Vec<String> = (0..10).map(|i| format!("s{i:02}")).collect();
540 let registry = registry_with(&ids).await;
541
542 assert_eq!(registry.session_ids_after(None, 1).await, vec!["s00"]);
543 assert_eq!(
546 registry.session_ids_after(None, 3).await,
547 vec!["s00", "s01", "s02"]
548 );
549 assert_eq!(registry.session_ids_after(None, 100).await.len(), 10);
551 }
552
553 #[tokio::test]
554 async fn session_ids_after_is_exclusive_of_cursor() {
555 let ids: Vec<String> = ["a", "b", "c", "d"].iter().map(|s| s.to_string()).collect();
556 let registry = registry_with(&ids).await;
557
558 let page = registry.session_ids_after(Some("b"), 10).await;
559 assert_eq!(page, vec!["c", "d"]);
560 assert!(!page.contains(&"b".to_string()));
561 assert!(page.iter().all(|id| id.as_str() > "b"));
562
563 assert!(registry.session_ids_after(Some("d"), 10).await.is_empty());
565 assert!(registry.session_ids_after(Some("zzz"), 10).await.is_empty());
567 }
568
569 #[tokio::test]
570 async fn session_ids_after_tolerates_absent_cursor() {
571 let ids: Vec<String> = ["a", "c", "e"].iter().map(|s| s.to_string()).collect();
572 let registry = registry_with(&ids).await;
573
574 assert_eq!(
577 registry.session_ids_after(Some("b"), 10).await,
578 vec!["c", "e"]
579 );
580 assert_eq!(
582 registry.session_ids_after(Some("b"), 10).await,
583 registry.session_ids_after(Some("a"), 10).await
584 );
585 assert_eq!(
590 registry.session_ids_after(Some(""), 10).await,
591 vec!["a", "c", "e"]
592 );
593 }
594
595 #[tokio::test]
596 async fn session_ids_after_zero_limit_is_empty() {
597 let ids: Vec<String> = ["a", "b", "c"].iter().map(|s| s.to_string()).collect();
598 let registry = registry_with(&ids).await;
599
600 assert!(registry.session_ids_after(None, 0).await.is_empty());
601 assert!(registry.session_ids_after(Some("a"), 0).await.is_empty());
602
603 let empty = SessionRegistry::new();
605 assert!(empty.session_ids_after(None, 0).await.is_empty());
606 assert!(empty.session_ids_after(None, 10).await.is_empty());
607 assert!(empty.session_ids_after(Some("a"), 10).await.is_empty());
608 }
609
610 #[tokio::test]
615 async fn session_ids_after_handles_huge_limits() {
616 let ids: Vec<String> = ["a", "b", "c"].iter().map(|s| s.to_string()).collect();
617 let registry = registry_with(&ids).await;
618
619 for limit in [usize::MAX, usize::MAX - 1, 10_000_000, 1 << 40] {
620 assert_eq!(
621 registry.session_ids_after(None, limit).await,
622 vec!["a", "b", "c"],
623 "limit={limit}"
624 );
625 assert_eq!(
626 registry.session_ids_after(Some("a"), limit).await,
627 vec!["b", "c"],
628 "limit={limit}"
629 );
630 }
631
632 let empty = SessionRegistry::new();
634 assert!(empty.session_ids_after(None, usize::MAX).await.is_empty());
635 }
636
637 #[tokio::test]
638 async fn session_ids_after_matches_sort_then_truncate_reference() {
639 let ids = deterministic_ids(200);
640 let registry = registry_with(&ids).await;
641
642 let mut sorted = ids.clone();
643 sorted.sort();
644
645 let cursors: Vec<Option<String>> = std::iter::once(None)
646 .chain(std::iter::once(Some(String::new())))
647 .chain(std::iter::once(Some("sess-".to_string())))
648 .chain(std::iter::once(Some("zzzz".to_string())))
649 .chain(sorted.iter().step_by(17).cloned().map(Some))
651 .chain(std::iter::once(Some(sorted.last().unwrap().clone())))
652 .chain(sorted.iter().step_by(23).map(|k| Some(format!("{k}~"))))
654 .collect();
655
656 for cursor in &cursors {
657 for limit in [1usize, 2, 7, 50, 199, 200, 201, 1000] {
658 let got = registry.session_ids_after(cursor.as_deref(), limit).await;
659 let want = sort_then_truncate_reference(&ids, cursor.as_deref(), limit);
660 assert_eq!(got, want, "cursor={cursor:?} limit={limit}");
661 }
662 }
663 }
664
665 #[tokio::test]
666 async fn session_ids_after_full_traversal_covers_every_id_once() {
667 let ids = deterministic_ids(200);
668 let registry = registry_with(&ids).await;
669
670 for page_size in [1usize, 3, 7, 64, 199, 200, 500] {
671 let mut collected: Vec<String> = Vec::new();
672 let mut cursor: Option<String> = None;
673 loop {
674 let page = registry
675 .session_ids_after(cursor.as_deref(), page_size)
676 .await;
677 let short = page.len() < page_size;
678 assert!(
682 page.len() <= page_size,
683 "page_size={page_size}: page of {} exceeds the limit",
684 page.len()
685 );
686 if let (Some(last), Some(first)) = (collected.last(), page.first()) {
688 assert!(first > last, "page_size={page_size}: page did not advance");
689 }
690 collected.extend(page.iter().cloned());
691 cursor = page.last().cloned();
692 if short {
693 break;
694 }
695 }
696
697 let unique: HashSet<&String> = collected.iter().collect();
698 assert_eq!(
700 collected.len(),
701 unique.len(),
702 "page_size={page_size}: duplicate IDs across pages"
703 );
704 let expected: HashSet<&String> = ids.iter().collect();
705 assert_eq!(unique, expected, "page_size={page_size}: coverage mismatch");
706 assert_eq!(collected.len(), ids.len(), "page_size={page_size}");
707 }
708 }
709
710 #[tokio::test]
711 async fn expired_sessions_not_counted_against_limit() {
712 let registry = SessionRegistry::new();
713 let now = chrono::Utc::now().timestamp_millis();
714 let mut expired = sample_session("expired-s1");
716 expired.initiator_sender = "agent://alice".into();
717 expired.ttl_expiry = now - 1000; expired.state = SessionState::Open; registry
720 .insert_recovered_session("expired-s1".into(), expired)
721 .await;
722
723 let count = registry
725 .count_open_sessions_for_initiator("agent://alice")
726 .await;
727 assert_eq!(count, 0);
728
729 let mut active = sample_session("active-s1");
731 active.initiator_sender = "agent://alice".into();
732 active.ttl_expiry = now + 60_000; active.state = SessionState::Open;
734 registry
735 .insert_recovered_session("active-s1".into(), active)
736 .await;
737
738 let count = registry
739 .count_open_sessions_for_initiator("agent://alice")
740 .await;
741 assert_eq!(count, 1);
742 }
743
744 #[tokio::test]
750 async fn load_sessions_repairs_key_field_mismatch() {
751 let base = std::env::temp_dir().join(format!(
752 "macp-registry-mismatch-{}",
753 SystemTime::now()
754 .duration_since(UNIX_EPOCH)
755 .unwrap()
756 .as_nanos()
757 ));
758 fs::create_dir_all(&base).unwrap();
759
760 let mut persisted = HashMap::new();
761 persisted.insert(
762 "A".to_string(),
763 PersistedSession::from(&sample_session("B")),
764 );
765 SessionRegistry::persist_map(&base.join("sessions.json"), &persisted).unwrap();
766
767 let reopened = SessionRegistry::with_persistence(&base).unwrap();
768
769 let session = reopened.get_session("A").await.unwrap();
771 assert_eq!(session.session_id, "A");
772 assert!(reopened.get_session("B").await.is_none());
774 assert_eq!(reopened.session_ids_after(None, 10).await, vec!["A"]);
777 }
778
779 #[tokio::test]
780 async fn persistent_registry_round_trip() {
781 let base = std::env::temp_dir().join(format!(
782 "macp-registry-test-{}",
783 SystemTime::now()
784 .duration_since(UNIX_EPOCH)
785 .unwrap()
786 .as_nanos()
787 ));
788
789 let registry = SessionRegistry::with_persistence(&base).unwrap();
790 registry
791 .insert_recovered_session("s1".into(), sample_session("s1"))
792 .await;
793
794 let reopened = SessionRegistry::with_persistence(&base).unwrap();
795 let session = reopened.get_session("s1").await.unwrap();
796 assert_eq!(session.mode, "macp.mode.decision.v1");
797 assert_eq!(session.mode_version, "1.0.0");
798 assert!(session.seen_message_ids.contains("m1"));
799 }
800}