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)]
35#[non_exhaustive]
36pub struct PersistedSession {
37 #[serde(default = "default_schema_version")]
38 pub schema_version: u32,
39 pub session_id: String,
40 pub state: macp_core::session::SessionState,
41 pub ttl_expiry: i64,
42 #[serde(default)]
43 pub ttl_ms: i64,
44 pub started_at_unix_ms: i64,
45 pub resolution: Option<Vec<u8>>,
46 pub mode: String,
47 pub mode_state: Vec<u8>,
48 pub participants: Vec<String>,
49 pub seen_message_ids: Vec<String>,
50 pub intent: String,
51 pub mode_version: String,
52 pub configuration_version: String,
53 pub policy_version: String,
54 #[serde(default)]
55 pub context_id: String,
56 #[serde(default)]
57 pub extensions: HashMap<String, Vec<u8>>,
58 pub roots: Vec<PersistedRoot>,
59 pub initiator_sender: String,
60 #[serde(default)]
61 pub policy_definition: Option<macp_core::policy::PolicyDefinition>,
62 #[serde(default)]
63 pub suspended_at_ms: Option<i64>,
64 #[serde(default)]
65 pub accumulated_suspended_ms: i64,
66 #[serde(default)]
71 pub suspension_intervals: Vec<(i64, i64)>,
72 #[serde(default)]
75 pub semantics_rev: u32,
76 #[serde(default)]
79 pub max_suspend_ms: i64,
80}
81
82fn default_schema_version() -> u32 {
83 2
84}
85
86impl From<&Session> for PersistedSession {
87 fn from(session: &Session) -> Self {
88 Self {
89 schema_version: 2,
90 session_id: session.session_id.clone(),
91 state: session.state.clone(),
92 ttl_expiry: session.ttl_expiry,
93 ttl_ms: session.ttl_ms,
94 started_at_unix_ms: session.started_at_unix_ms,
95 resolution: session.resolution.clone(),
96 mode: session.mode.clone(),
97 mode_state: session.mode_state.clone(),
98 participants: session.participants.clone(),
99 seen_message_ids: session.seen_message_ids.iter().cloned().collect(),
100 intent: session.intent.clone(),
101 mode_version: session.mode_version.clone(),
102 configuration_version: session.configuration_version.clone(),
103 policy_version: session.policy_version.clone(),
104 context_id: session.context_id.clone(),
105 extensions: session.extensions.clone(),
106 roots: session
107 .roots
108 .iter()
109 .map(|root| PersistedRoot {
110 uri: root.uri.clone(),
111 name: root.name.clone(),
112 })
113 .collect(),
114 initiator_sender: session.initiator_sender.clone(),
115 policy_definition: session.policy_definition.clone(),
116 suspended_at_ms: session.suspended_at_ms,
117 accumulated_suspended_ms: session.accumulated_suspended_ms,
118 suspension_intervals: session.suspension_intervals.clone(),
119 semantics_rev: session.semantics_rev,
120 max_suspend_ms: session.max_suspend_ms,
121 }
122 }
123}
124
125impl From<PersistedSession> for Session {
126 fn from(session: PersistedSession) -> Self {
127 let ttl_ms = if session.ttl_ms > 0 {
128 session.ttl_ms
129 } else {
130 session
132 .ttl_expiry
133 .saturating_sub(session.started_at_unix_ms)
134 };
135 Session::builder(session.session_id, session.mode, session.initiator_sender)
136 .state(session.state)
137 .ttl_expiry(session.ttl_expiry)
138 .ttl_ms(ttl_ms)
139 .started_at_unix_ms(session.started_at_unix_ms)
140 .resolution(session.resolution)
141 .mode_state(session.mode_state)
142 .participants(session.participants)
143 .seen_message_ids(session.seen_message_ids.into_iter().collect())
144 .intent(session.intent)
145 .mode_version(session.mode_version)
146 .configuration_version(session.configuration_version)
147 .policy_version(session.policy_version)
148 .context_id(session.context_id)
149 .extensions(session.extensions)
150 .roots(
151 session
152 .roots
153 .into_iter()
154 .map(|root| macp_pb::pb::Root {
155 uri: root.uri,
156 name: root.name,
157 })
158 .collect(),
159 )
160 .policy_definition(session.policy_definition)
161 .suspended_at_ms(session.suspended_at_ms)
162 .accumulated_suspended_ms(session.accumulated_suspended_ms)
163 .suspension_intervals(session.suspension_intervals)
164 .semantics_rev(session.semantics_rev)
165 .max_suspend_ms(session.max_suspend_ms)
166 .build()
167 }
168}
169
170pub type SharedSession = Arc<tokio::sync::Mutex<Session>>;
178
179pub struct SessionRegistry {
180 pub sessions: RwLock<HashMap<String, SharedSession>>,
181 persistence_path: Option<PathBuf>,
182}
183
184impl Default for SessionRegistry {
185 fn default() -> Self {
186 Self::new()
187 }
188}
189
190impl SessionRegistry {
191 pub fn new() -> Self {
192 Self {
193 sessions: RwLock::new(HashMap::new()),
194 persistence_path: None,
195 }
196 }
197
198 pub fn with_persistence<P: AsRef<Path>>(dir: P) -> std::io::Result<Self> {
199 let dir = dir.as_ref().to_path_buf();
200 fs::create_dir_all(&dir)?;
201 let path = dir.join("sessions.json");
202 let sessions = Self::load_sessions(&path)?;
203 Ok(Self {
204 sessions: RwLock::new(sessions),
205 persistence_path: Some(path),
206 })
207 }
208
209 fn load_sessions(path: &Path) -> std::io::Result<HashMap<String, SharedSession>> {
210 if !path.exists() {
211 return Ok(HashMap::new());
212 }
213 let bytes = fs::read(path)?;
214 let persisted: HashMap<String, PersistedSession> = match serde_json::from_slice(&bytes) {
215 Ok(v) => v,
216 Err(e) => {
217 eprintln!("warning: failed to deserialize sessions from {}: {e}; starting with empty state", path.display());
218 HashMap::new()
219 }
220 };
221 Ok(persisted
222 .into_iter()
223 .map(|(id, mut record)| {
224 if record.session_id != id {
234 tracing::warn!(
235 map_key = %id,
236 record_session_id = %record.session_id,
237 path = %path.display(),
238 "persisted session key disagrees with its session_id; \
239 repairing to the map key"
240 );
241 record.session_id.clone_from(&id);
242 }
243 let session: Session = record.into();
244 (id, Arc::new(tokio::sync::Mutex::new(session)))
245 })
246 .collect())
247 }
248
249 fn persist_map(
250 path: &Path,
251 sessions: &HashMap<String, PersistedSession>,
252 ) -> std::io::Result<()> {
253 let bytes = serde_json::to_vec_pretty(sessions)?;
254 let tmp_path = path.with_extension("json.tmp");
255 fs::write(&tmp_path, bytes)?;
256 fs::rename(&tmp_path, path)
257 }
258
259 pub async fn persist_snapshot(&self) -> std::io::Result<()> {
262 let Some(path) = self.persistence_path.clone() else {
263 return Ok(());
264 };
265 let arcs: Vec<(String, SharedSession)> = {
266 let guard = self.sessions.read().await;
267 guard
268 .iter()
269 .map(|(id, arc)| (id.clone(), Arc::clone(arc)))
270 .collect()
271 };
272 let mut persisted = HashMap::with_capacity(arcs.len());
273 for (id, arc) in arcs {
274 let session = arc.lock().await;
275 persisted.insert(id, PersistedSession::from(&*session));
276 }
277 Self::persist_map(&path, &persisted)
278 }
279
280 pub async fn get_shared(&self, session_id: &str) -> Option<SharedSession> {
282 let guard = self.sessions.read().await;
283 guard.get(session_id).cloned()
284 }
285
286 pub async fn get_session(&self, session_id: &str) -> Option<Session> {
287 let arc = self.get_shared(session_id).await?;
288 let session = arc.lock().await;
289 Some(session.clone())
290 }
291
292 pub async fn get_all_sessions(&self) -> Vec<Session> {
301 let arcs: Vec<SharedSession> = {
302 let guard = self.sessions.read().await;
303 guard.values().cloned().collect()
304 };
305 let mut out = Vec::with_capacity(arcs.len());
306 for arc in arcs {
307 out.push(arc.lock().await.clone());
308 }
309 out
310 }
311
312 pub async fn shared_sessions(&self) -> Vec<SharedSession> {
339 let guard = self.sessions.read().await;
340 guard.values().map(Arc::clone).collect()
341 }
342
343 pub async fn session_ids_after(&self, after: Option<&str>, limit: usize) -> Vec<String> {
363 if limit == 0 {
364 return Vec::new();
365 }
366 {
370 let guard = self.sessions.read().await;
371 let capacity = limit.saturating_add(1).min(guard.len().saturating_add(1));
378 let mut heap: BinaryHeap<&String> = BinaryHeap::with_capacity(capacity);
379 for key in guard.keys() {
380 if after.is_none_or(|a| key.as_str() > a) {
381 heap.push(key);
382 if heap.len() > limit {
383 heap.pop();
384 }
385 }
386 }
387 heap.into_sorted_vec().into_iter().cloned().collect()
388 }
389 }
390
391 pub async fn insert_recovered_session(&self, session_id: String, session: Session) {
392 debug_assert_eq!(
404 session.session_id, session_id,
405 "registry map key must equal Session::session_id — ListSessions paging \
406 orders by the key but emits the field (plan D1)"
407 );
408 {
409 let mut guard = self.sessions.write().await;
410 guard.insert(session_id, Arc::new(tokio::sync::Mutex::new(session)));
411 }
412 let _ = self.persist_snapshot().await;
413 }
414
415 pub async fn count_open_sessions_for_initiator(&self, sender: &str) -> usize {
416 let now = chrono::Utc::now().timestamp_millis();
417 let arcs: Vec<SharedSession> = {
418 let guard = self.sessions.read().await;
419 guard.values().cloned().collect()
420 };
421 let mut count = 0;
422 for arc in arcs {
423 let counts = match arc.try_lock() {
426 Ok(session) => {
427 session.initiator_sender == sender
428 && session.state == macp_core::session::SessionState::Open
429 && now <= session.ttl_expiry
430 }
431 Err(_) => true,
432 };
433 if counts {
434 count += 1;
435 }
436 }
437 count
438 }
439}
440
441#[cfg(test)]
442mod tests {
443 use super::*;
444 use macp_core::session::{Session, SessionState};
445 use std::collections::HashSet;
446 use std::time::{SystemTime, UNIX_EPOCH};
447
448 fn sample_session(id: &str) -> Session {
449 Session::builder(id, "macp.mode.decision.v1", "alice")
450 .ttl_expiry(10)
451 .ttl_ms(9)
452 .started_at_unix_ms(1)
453 .mode_state(vec![1, 2, 3])
454 .participants(vec!["alice".into()])
455 .seen_message_ids(HashSet::from(["m1".into()]))
456 .intent("intent")
457 .mode_version("1.0.0")
458 .configuration_version("cfg")
459 .policy_version("pol")
460 .context_id("test-ctx")
461 .roots(vec![macp_pb::pb::Root {
462 uri: "root://1".into(),
463 name: "r1".into(),
464 }])
465 .build()
466 }
467
468 async fn registry_with(ids: &[String]) -> SessionRegistry {
471 let registry = SessionRegistry::new();
472 for id in ids {
473 registry
474 .insert_recovered_session(id.clone(), sample_session(id))
475 .await;
476 }
477 registry
478 }
479
480 fn sort_then_truncate_reference(
483 ids: &[String],
484 after: Option<&str>,
485 limit: usize,
486 ) -> Vec<String> {
487 let mut sorted: Vec<String> = ids.to_vec();
488 sorted.sort();
489 sorted
490 .into_iter()
491 .filter(|id| after.is_none_or(|a| id.as_str() > a))
492 .take(limit)
493 .collect()
494 }
495
496 fn deterministic_ids(count: usize) -> Vec<String> {
500 let mut state: u64 = 0x2545_F491_4F6C_DD1D;
501 let mut ids = Vec::with_capacity(count);
502 for i in 0..count {
503 state = state
504 .wrapping_mul(6_364_136_223_846_793_005)
505 .wrapping_add(1_442_695_040_888_963_407);
506 ids.push(format!("sess-{:016x}-{i:04}", state >> 16));
508 }
509 ids
510 }
511
512 #[tokio::test]
513 async fn shared_sessions_snapshots_every_session_once() {
514 let ids: Vec<String> = ["delta", "alpha", "charlie", "bravo"]
515 .iter()
516 .map(|s| s.to_string())
517 .collect();
518 let registry = registry_with(&ids).await;
519
520 let handles = registry.shared_sessions().await;
521 assert_eq!(handles.len(), 4);
522 let mut seen = Vec::new();
523 for handle in &handles {
524 seen.push(handle.lock().await.session_id.clone());
525 }
526 seen.sort();
529 assert_eq!(seen, vec!["alpha", "bravo", "charlie", "delta"]);
530
531 {
535 let mut guard = registry.sessions.write().await;
536 guard.remove("alpha");
537 guard.remove("bravo");
538 }
539 assert!(registry.get_session("alpha").await.is_none());
540 let mut after = Vec::new();
541 for handle in &handles {
542 after.push(handle.lock().await.session_id.clone());
543 }
544 after.sort();
545 assert_eq!(after, vec!["alpha", "bravo", "charlie", "delta"]);
546
547 assert!(SessionRegistry::new().shared_sessions().await.is_empty());
548 }
549
550 #[tokio::test]
551 async fn session_ids_after_returns_ascending_ids() {
552 let ids: Vec<String> = ["delta", "alpha", "charlie", "bravo"]
553 .iter()
554 .map(|s| s.to_string())
555 .collect();
556 let registry = registry_with(&ids).await;
557
558 let page = registry.session_ids_after(None, 10).await;
559 assert_eq!(page, vec!["alpha", "bravo", "charlie", "delta"]);
560
561 let page = registry.session_ids_after(None, 2).await;
563 assert_eq!(page, vec!["alpha", "bravo"]);
564 }
565
566 #[tokio::test]
567 async fn session_ids_after_respects_limit() {
568 let ids: Vec<String> = (0..10).map(|i| format!("s{i:02}")).collect();
569 let registry = registry_with(&ids).await;
570
571 assert_eq!(registry.session_ids_after(None, 1).await, vec!["s00"]);
572 assert_eq!(
575 registry.session_ids_after(None, 3).await,
576 vec!["s00", "s01", "s02"]
577 );
578 assert_eq!(registry.session_ids_after(None, 100).await.len(), 10);
580 }
581
582 #[tokio::test]
583 async fn session_ids_after_is_exclusive_of_cursor() {
584 let ids: Vec<String> = ["a", "b", "c", "d"].iter().map(|s| s.to_string()).collect();
585 let registry = registry_with(&ids).await;
586
587 let page = registry.session_ids_after(Some("b"), 10).await;
588 assert_eq!(page, vec!["c", "d"]);
589 assert!(!page.contains(&"b".to_string()));
590 assert!(page.iter().all(|id| id.as_str() > "b"));
591
592 assert!(registry.session_ids_after(Some("d"), 10).await.is_empty());
594 assert!(registry.session_ids_after(Some("zzz"), 10).await.is_empty());
596 }
597
598 #[tokio::test]
599 async fn session_ids_after_tolerates_absent_cursor() {
600 let ids: Vec<String> = ["a", "c", "e"].iter().map(|s| s.to_string()).collect();
601 let registry = registry_with(&ids).await;
602
603 assert_eq!(
606 registry.session_ids_after(Some("b"), 10).await,
607 vec!["c", "e"]
608 );
609 assert_eq!(
611 registry.session_ids_after(Some("b"), 10).await,
612 registry.session_ids_after(Some("a"), 10).await
613 );
614 assert_eq!(
619 registry.session_ids_after(Some(""), 10).await,
620 vec!["a", "c", "e"]
621 );
622 }
623
624 #[tokio::test]
625 async fn session_ids_after_zero_limit_is_empty() {
626 let ids: Vec<String> = ["a", "b", "c"].iter().map(|s| s.to_string()).collect();
627 let registry = registry_with(&ids).await;
628
629 assert!(registry.session_ids_after(None, 0).await.is_empty());
630 assert!(registry.session_ids_after(Some("a"), 0).await.is_empty());
631
632 let empty = SessionRegistry::new();
634 assert!(empty.session_ids_after(None, 0).await.is_empty());
635 assert!(empty.session_ids_after(None, 10).await.is_empty());
636 assert!(empty.session_ids_after(Some("a"), 10).await.is_empty());
637 }
638
639 #[tokio::test]
644 async fn session_ids_after_handles_huge_limits() {
645 let ids: Vec<String> = ["a", "b", "c"].iter().map(|s| s.to_string()).collect();
646 let registry = registry_with(&ids).await;
647
648 for limit in [usize::MAX, usize::MAX - 1, 10_000_000, 1 << 40] {
649 assert_eq!(
650 registry.session_ids_after(None, limit).await,
651 vec!["a", "b", "c"],
652 "limit={limit}"
653 );
654 assert_eq!(
655 registry.session_ids_after(Some("a"), limit).await,
656 vec!["b", "c"],
657 "limit={limit}"
658 );
659 }
660
661 let empty = SessionRegistry::new();
663 assert!(empty.session_ids_after(None, usize::MAX).await.is_empty());
664 }
665
666 #[tokio::test]
667 async fn session_ids_after_matches_sort_then_truncate_reference() {
668 let ids = deterministic_ids(200);
669 let registry = registry_with(&ids).await;
670
671 let mut sorted = ids.clone();
672 sorted.sort();
673
674 let cursors: Vec<Option<String>> = std::iter::once(None)
675 .chain(std::iter::once(Some(String::new())))
676 .chain(std::iter::once(Some("sess-".to_string())))
677 .chain(std::iter::once(Some("zzzz".to_string())))
678 .chain(sorted.iter().step_by(17).cloned().map(Some))
680 .chain(std::iter::once(Some(sorted.last().unwrap().clone())))
681 .chain(sorted.iter().step_by(23).map(|k| Some(format!("{k}~"))))
683 .collect();
684
685 for cursor in &cursors {
686 for limit in [1usize, 2, 7, 50, 199, 200, 201, 1000] {
687 let got = registry.session_ids_after(cursor.as_deref(), limit).await;
688 let want = sort_then_truncate_reference(&ids, cursor.as_deref(), limit);
689 assert_eq!(got, want, "cursor={cursor:?} limit={limit}");
690 }
691 }
692 }
693
694 #[tokio::test]
695 async fn session_ids_after_full_traversal_covers_every_id_once() {
696 let ids = deterministic_ids(200);
697 let registry = registry_with(&ids).await;
698
699 for page_size in [1usize, 3, 7, 64, 199, 200, 500] {
700 let mut collected: Vec<String> = Vec::new();
701 let mut cursor: Option<String> = None;
702 loop {
703 let page = registry
704 .session_ids_after(cursor.as_deref(), page_size)
705 .await;
706 let short = page.len() < page_size;
707 assert!(
711 page.len() <= page_size,
712 "page_size={page_size}: page of {} exceeds the limit",
713 page.len()
714 );
715 if let (Some(last), Some(first)) = (collected.last(), page.first()) {
717 assert!(first > last, "page_size={page_size}: page did not advance");
718 }
719 collected.extend(page.iter().cloned());
720 cursor = page.last().cloned();
721 if short {
722 break;
723 }
724 }
725
726 let unique: HashSet<&String> = collected.iter().collect();
727 assert_eq!(
729 collected.len(),
730 unique.len(),
731 "page_size={page_size}: duplicate IDs across pages"
732 );
733 let expected: HashSet<&String> = ids.iter().collect();
734 assert_eq!(unique, expected, "page_size={page_size}: coverage mismatch");
735 assert_eq!(collected.len(), ids.len(), "page_size={page_size}");
736 }
737 }
738
739 #[tokio::test]
740 async fn expired_sessions_not_counted_against_limit() {
741 let registry = SessionRegistry::new();
742 let now = chrono::Utc::now().timestamp_millis();
743 let mut expired = sample_session("expired-s1");
745 expired.initiator_sender = "agent://alice".into();
746 expired.ttl_expiry = now - 1000; expired.state = SessionState::Open; registry
749 .insert_recovered_session("expired-s1".into(), expired)
750 .await;
751
752 let count = registry
754 .count_open_sessions_for_initiator("agent://alice")
755 .await;
756 assert_eq!(count, 0);
757
758 let mut active = sample_session("active-s1");
760 active.initiator_sender = "agent://alice".into();
761 active.ttl_expiry = now + 60_000; active.state = SessionState::Open;
763 registry
764 .insert_recovered_session("active-s1".into(), active)
765 .await;
766
767 let count = registry
768 .count_open_sessions_for_initiator("agent://alice")
769 .await;
770 assert_eq!(count, 1);
771 }
772
773 #[tokio::test]
779 async fn load_sessions_repairs_key_field_mismatch() {
780 let base = std::env::temp_dir().join(format!(
781 "macp-registry-mismatch-{}",
782 SystemTime::now()
783 .duration_since(UNIX_EPOCH)
784 .unwrap()
785 .as_nanos()
786 ));
787 fs::create_dir_all(&base).unwrap();
788
789 let mut persisted = HashMap::new();
790 persisted.insert(
791 "A".to_string(),
792 PersistedSession::from(&sample_session("B")),
793 );
794 SessionRegistry::persist_map(&base.join("sessions.json"), &persisted).unwrap();
795
796 let reopened = SessionRegistry::with_persistence(&base).unwrap();
797
798 let session = reopened.get_session("A").await.unwrap();
800 assert_eq!(session.session_id, "A");
801 assert!(reopened.get_session("B").await.is_none());
803 assert_eq!(reopened.session_ids_after(None, 10).await, vec!["A"]);
806 }
807
808 #[tokio::test]
809 async fn persistent_registry_round_trip() {
810 let base = std::env::temp_dir().join(format!(
811 "macp-registry-test-{}",
812 SystemTime::now()
813 .duration_since(UNIX_EPOCH)
814 .unwrap()
815 .as_nanos()
816 ));
817
818 let registry = SessionRegistry::with_persistence(&base).unwrap();
819 registry
820 .insert_recovered_session("s1".into(), sample_session("s1"))
821 .await;
822
823 let reopened = SessionRegistry::with_persistence(&base).unwrap();
824 let session = reopened.get_session("s1").await.unwrap();
825 assert_eq!(session.mode, "macp.mode.decision.v1");
826 assert_eq!(session.mode_version, "1.0.0");
827 assert!(session.seen_message_ids.contains("m1"));
828 }
829}