1use std::{
2 collections::HashMap,
3 error::Error,
4 fmt,
5 sync::{Mutex, MutexGuard},
6};
7
8use subc_protocol::manifest::{CapabilityDeclarations, ModuleManifest, ProviderRole};
9use tokio::sync::watch;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13pub struct ConnectionId(u64);
14
15impl ConnectionId {
16 #[cfg(test)]
19 pub const LOCAL: Self = Self(0);
20
21 pub const fn new(raw: u64) -> Self {
22 Self(raw)
23 }
24
25 pub fn get(self) -> u64 {
26 self.0
27 }
28}
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
32pub enum ChannelState {
33 Active,
34 Closed,
35}
36
37#[derive(Debug, Clone, PartialEq)]
39pub struct ModuleRegistration {
40 pub manifest: ModuleManifest,
41 pub ready: bool,
42 pub negotiated_ver: u8,
43 pub state: ChannelState,
44 pub connection_id: ConnectionId,
45 pub control_ops: Vec<String>,
46}
47
48#[derive(Debug, Clone, Copy, PartialEq, Eq)]
57pub enum RegistrationSlot<'a> {
58 Active(&'a str),
61 Candidate(&'a str),
63 Connection(ConnectionId),
65}
66
67#[derive(Debug, Clone, PartialEq)]
69pub struct RegistryCutover {
70 pub promoted: ModuleRegistration,
72 pub superseded: Option<ModuleRegistration>,
75}
76
77#[derive(Debug, Default)]
92pub struct Registry {
93 inner: Mutex<RegistryInner>,
94}
95
96#[derive(Debug, Default)]
97struct RegistryInner {
98 module_changes: HashMap<String, watch::Sender<()>>,
101 modules: HashMap<String, ModuleRegistration>,
102 candidates: HashMap<String, ModuleRegistration>,
104 superseded: Vec<ModuleRegistration>,
107 generation: u64,
108}
109
110impl Registry {
111 pub(crate) fn subscribe_module_changes(
114 &self,
115 module_id: &str,
116 ) -> Result<watch::Receiver<()>, RegistryError> {
117 let mut inner = self.lock_inner()?;
118 inner
120 .module_changes
121 .retain(|_, sender| sender.receiver_count() > 0);
122 Ok(inner
123 .module_changes
124 .entry(module_id.to_string())
125 .or_insert_with(|| watch::channel(()).0)
126 .subscribe())
127 }
128
129 pub fn register_with_control_ops(
131 &self,
132 manifest: ModuleManifest,
133 negotiated_ver: u8,
134 connection_id: ConnectionId,
135 control_ops: Vec<String>,
136 ) -> Result<ModuleRegistration, RegistryError> {
137 let module_id = manifest.module_id.clone();
138 if let Err(reason) = module_id_path_hazard(&module_id) {
139 return Err(RegistryError::PathHazardModuleId { module_id, reason });
140 }
141 let mut inner = self.lock_inner()?;
142 if inner.modules.contains_key(&module_id) {
143 return Err(RegistryError::DuplicateModuleId { module_id });
144 }
145
146 let ready = manifest.ready.unwrap_or(true);
147 let registration = ModuleRegistration {
148 manifest,
149 ready,
150 negotiated_ver,
151 state: ChannelState::Active,
152 connection_id,
153 control_ops,
154 };
155
156 inner.modules.insert(module_id, registration.clone());
157 inner.bump_generation();
158 inner.notify_module_changed(®istration.manifest.module_id);
159 Ok(registration)
160 }
161
162 pub fn register_candidate_with_control_ops(
171 &self,
172 manifest: ModuleManifest,
173 negotiated_ver: u8,
174 connection_id: ConnectionId,
175 control_ops: Vec<String>,
176 ) -> Result<ModuleRegistration, RegistryError> {
177 let module_id = manifest.module_id.clone();
178 if let Err(reason) = module_id_path_hazard(&module_id) {
179 return Err(RegistryError::PathHazardModuleId { module_id, reason });
180 }
181 let mut inner = self.lock_inner()?;
182 if inner.candidates.contains_key(&module_id) {
183 return Err(RegistryError::DuplicateModuleId { module_id });
184 }
185 let ready = manifest.ready.unwrap_or(true);
186 let registration = ModuleRegistration {
187 manifest,
188 ready,
189 negotiated_ver,
190 state: ChannelState::Active,
191 connection_id,
192 control_ops,
193 };
194 inner.candidates.insert(module_id, registration.clone());
195 inner.notify_module_changed(®istration.manifest.module_id);
196 Ok(registration)
197 }
198
199 pub fn promote_candidate(
205 &self,
206 module_id: &str,
207 ) -> Result<Option<RegistryCutover>, RegistryError> {
208 let mut inner = self.lock_inner()?;
209 let Some(promoted) = inner.candidates.remove(module_id) else {
210 return Ok(None);
211 };
212 let superseded = inner
213 .modules
214 .insert(module_id.to_string(), promoted.clone());
215 if let Some(superseded) = superseded.clone() {
216 inner.superseded.push(superseded);
217 }
218 inner.bump_generation();
219 inner.notify_module_changed(module_id);
220 Ok(Some(RegistryCutover {
221 promoted,
222 superseded,
223 }))
224 }
225
226 pub fn get_module(&self, module_id: &str) -> Result<Option<ModuleRegistration>, RegistryError> {
229 Ok(self.lock_inner()?.modules.get(module_id).cloned())
230 }
231
232 pub fn get_candidate(
234 &self,
235 module_id: &str,
236 ) -> Result<Option<ModuleRegistration>, RegistryError> {
237 Ok(self.lock_inner()?.candidates.get(module_id).cloned())
238 }
239
240 pub fn registration(
242 &self,
243 slot: RegistrationSlot<'_>,
244 ) -> Result<Option<ModuleRegistration>, RegistryError> {
245 let inner = self.lock_inner()?;
246 Ok(match slot {
247 RegistrationSlot::Active(module_id) => inner.modules.get(module_id).cloned(),
248 RegistrationSlot::Candidate(module_id) => inner.candidates.get(module_id).cloned(),
249 RegistrationSlot::Connection(connection_id) => inner
250 .find_by_connection(connection_id)
251 .map(|(_, registration)| registration.clone()),
252 })
253 }
254
255 pub fn active_registration_count(&self) -> Result<usize, RegistryError> {
256 Ok(self.lock_inner()?.modules.len())
257 }
258
259 pub fn list_modules(&self) -> Result<(u64, Vec<ModuleRegistration>), RegistryError> {
260 let inner = self.lock_inner()?;
261 let mut modules = inner.modules.values().cloned().collect::<Vec<_>>();
262 modules.sort_by(|left, right| left.manifest.module_id.cmp(&right.manifest.module_id));
263 Ok((inner.generation, modules))
264 }
265
266 pub fn generation(&self) -> Result<u64, RegistryError> {
267 Ok(self.lock_inner()?.generation)
268 }
269
270 #[cfg(test)]
271 pub(crate) fn set_module_state_for_test(
272 &self,
273 module_id: &str,
274 state: ChannelState,
275 ) -> Result<bool, RegistryError> {
276 let mut inner = self.lock_inner()?;
277 let Some(registration) = inner.modules.get_mut(module_id) else {
278 return Ok(false);
279 };
280 registration.state = state;
281 inner.notify_module_changed(module_id);
282 Ok(true)
283 }
284
285 pub fn get_module_by_connection(
288 &self,
289 connection_id: ConnectionId,
290 ) -> Result<Option<ModuleRegistration>, RegistryError> {
291 Ok(self
292 .lock_inner()?
293 .find_by_connection(connection_id)
294 .map(|(_, registration)| registration.clone()))
295 }
296
297 pub fn replace_catalog_for_connection(
306 &self,
307 connection_id: ConnectionId,
308 provides: Vec<ProviderRole>,
309 capabilities: Option<CapabilityDeclarations>,
310 ready: Option<bool>,
311 ) -> Result<Option<ModuleRegistration>, RegistryError> {
312 let mut inner = self.lock_inner()?;
313 let Some((slot, _)) = inner.find_by_connection(connection_id) else {
314 return Ok(None);
315 };
316 let registration = inner
317 .registration_mut(slot, connection_id)
318 .expect("registration discovered under the same registry lock must still exist");
319 registration.manifest.provides = provides;
320 if let Some(capabilities) = capabilities {
321 registration.manifest.capabilities = Some(capabilities);
322 }
323 if let Some(ready) = ready {
324 registration.ready = ready;
325 registration.manifest.ready = Some(ready);
326 }
327 let updated = registration.clone();
328 if matches!(slot, SlotKind::Active) {
329 inner.bump_generation();
330 }
331 inner.notify_module_changed(&updated.manifest.module_id);
332 Ok(Some(updated))
333 }
334
335 pub fn deregister_connection(
337 &self,
338 connection_id: ConnectionId,
339 ) -> Result<Vec<ModuleRegistration>, RegistryError> {
340 let mut inner = self.lock_inner()?;
341 let module_ids: Vec<String> = inner
342 .modules
343 .iter()
344 .filter(|(_, registration)| registration.connection_id == connection_id)
345 .map(|(module_id, _)| module_id.clone())
346 .collect();
347
348 let mut closed: Vec<ModuleRegistration> = module_ids
349 .into_iter()
350 .filter_map(|module_id| inner.close_module(&module_id))
351 .collect();
352
353 let candidate_ids: Vec<String> = inner
354 .candidates
355 .iter()
356 .filter(|(_, registration)| registration.connection_id == connection_id)
357 .map(|(module_id, _)| module_id.clone())
358 .collect();
359 for module_id in candidate_ids {
360 if let Some(mut registration) = inner.candidates.remove(&module_id) {
361 registration.state = ChannelState::Closed;
362 closed.push(registration);
363 }
364 }
365
366 let (removed, kept): (Vec<_>, Vec<_>) = std::mem::take(&mut inner.superseded)
367 .into_iter()
368 .partition(|registration| registration.connection_id == connection_id);
369 inner.superseded = kept;
370 closed.extend(removed.into_iter().map(|mut registration| {
371 registration.state = ChannelState::Closed;
372 registration
373 }));
374 for registration in &closed {
375 inner.notify_module_changed(®istration.manifest.module_id);
376 }
377 Ok(closed)
378 }
379
380 fn lock_inner(&self) -> Result<MutexGuard<'_, RegistryInner>, RegistryError> {
381 self.inner.lock().map_err(|_| RegistryError::Poisoned)
382 }
383}
384
385#[derive(Debug, Clone, Copy, PartialEq, Eq)]
387enum SlotKind {
388 Active,
389 Candidate,
390 Superseded,
391}
392
393impl RegistryInner {
394 fn notify_module_changed(&self, module_id: &str) {
395 if let Some(sender) = self.module_changes.get(module_id) {
396 sender.send_replace(());
397 }
398 }
399
400 fn find_by_connection(
401 &self,
402 connection_id: ConnectionId,
403 ) -> Option<(SlotKind, &ModuleRegistration)> {
404 let owned_by =
405 |registration: &&ModuleRegistration| registration.connection_id == connection_id;
406 self.modules
407 .values()
408 .find(owned_by)
409 .map(|registration| (SlotKind::Active, registration))
410 .or_else(|| {
411 self.candidates
412 .values()
413 .find(owned_by)
414 .map(|registration| (SlotKind::Candidate, registration))
415 })
416 .or_else(|| {
417 self.superseded
418 .iter()
419 .find(owned_by)
420 .map(|registration| (SlotKind::Superseded, registration))
421 })
422 }
423
424 fn registration_mut(
425 &mut self,
426 slot: SlotKind,
427 connection_id: ConnectionId,
428 ) -> Option<&mut ModuleRegistration> {
429 let owned_by =
430 |registration: &&mut ModuleRegistration| registration.connection_id == connection_id;
431 match slot {
432 SlotKind::Active => self.modules.values_mut().find(owned_by),
433 SlotKind::Candidate => self.candidates.values_mut().find(owned_by),
434 SlotKind::Superseded => self.superseded.iter_mut().find(owned_by),
435 }
436 }
437
438 fn close_module(&mut self, module_id: &str) -> Option<ModuleRegistration> {
439 let mut registration = self.modules.remove(module_id)?;
440 registration.state = ChannelState::Closed;
441 self.bump_generation();
442 Some(registration)
443 }
444
445 fn bump_generation(&mut self) {
446 self.generation = self.generation.wrapping_add(1);
447 }
448}
449
450#[derive(Debug, Clone, PartialEq, Eq)]
451pub enum RegistryError {
452 DuplicateModuleId {
453 module_id: String,
454 },
455 PathHazardModuleId {
464 module_id: String,
465 reason: String,
466 },
467 Poisoned,
468}
469
470impl fmt::Display for RegistryError {
471 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
472 match self {
473 Self::DuplicateModuleId { module_id } => {
474 write!(f, "module_id '{module_id}' is already registered")
475 }
476 Self::PathHazardModuleId { module_id, reason } => {
477 write!(
478 f,
479 "module_id '{}' is not usable as a path component: {reason}",
480 module_id.escape_debug()
481 )
482 }
483 Self::Poisoned => write!(f, "registry lock was poisoned"),
484 }
485 }
486}
487
488impl Error for RegistryError {}
489
490pub fn module_id_path_hazard(module_id: &str) -> Result<(), String> {
500 if module_id.is_empty() {
501 return Err("empty".to_string());
502 }
503 if module_id.contains('/') || module_id.contains('\\') {
504 return Err("contains a path separator".to_string());
505 }
506 if module_id == "." || module_id == ".." {
507 return Err("is a dot path component".to_string());
508 }
509 if module_id.chars().any(|c| c.is_control()) {
510 return Err("contains a control character".to_string());
511 }
512 if module_id.len() > 255 {
516 return Err("is longer than 255 bytes".to_string());
517 }
518 Ok(())
519}
520
521#[cfg(test)]
522mod change_notification_tests {
523 use super::*;
524
525 fn manifest() -> ModuleManifest {
526 ModuleManifest::builder("watched", "1.0.0")
527 .protocol_ver(1)
528 .build()
529 }
530
531 fn changed(receiver: &mut watch::Receiver<()>) {
532 assert!(
533 receiver.has_changed().unwrap(),
534 "registry write must notify"
535 );
536 receiver.borrow_and_update();
537 }
538
539 #[test]
540 fn every_registration_write_notifies_only_its_module() {
541 let registry = Registry::default();
542 let mut events = registry.subscribe_module_changes("watched").unwrap();
543 let other = registry.subscribe_module_changes("other").unwrap();
544 let active = ConnectionId::new(1);
545 let candidate = ConnectionId::new(2);
546 registry
547 .register_with_control_ops(manifest(), 1, active, vec![])
548 .unwrap();
549 changed(&mut events);
550 registry
551 .set_module_state_for_test("watched", ChannelState::Active)
552 .unwrap();
553 changed(&mut events);
554 registry
555 .replace_catalog_for_connection(active, vec![], None, Some(true))
556 .unwrap();
557 changed(&mut events);
558 registry
559 .register_candidate_with_control_ops(manifest(), 1, candidate, vec![])
560 .unwrap();
561 changed(&mut events);
562 registry
563 .replace_catalog_for_connection(candidate, vec![], None, Some(true))
564 .unwrap();
565 changed(&mut events);
566 registry.promote_candidate("watched").unwrap().unwrap();
567 changed(&mut events);
568 registry
570 .replace_catalog_for_connection(active, vec![], None, Some(false))
571 .unwrap();
572 changed(&mut events);
573 registry.deregister_connection(active).unwrap();
574 changed(&mut events);
575 registry.deregister_connection(candidate).unwrap();
576 changed(&mut events);
577 registry
578 .register_candidate_with_control_ops(manifest(), 1, candidate, vec![])
579 .unwrap();
580 changed(&mut events);
581 registry.deregister_connection(candidate).unwrap();
582 changed(&mut events);
583 assert!(
584 !other.has_changed().unwrap(),
585 "other module must stay parked"
586 );
587 }
588
589 #[test]
590 fn subscription_before_read_keeps_a_change_until_waited_on() {
591 let registry = Registry::default();
592 let mut events = registry.subscribe_module_changes("watched").unwrap();
593 assert!(registry.get_module("watched").unwrap().is_none());
594 registry
595 .register_with_control_ops(manifest(), 1, ConnectionId::new(1), vec![])
596 .unwrap();
597 registry.get_module("watched").unwrap().unwrap();
599 registry
600 .replace_catalog_for_connection(ConnectionId::new(1), vec![], None, Some(true))
601 .unwrap();
602 changed(&mut events);
603 assert!(!events.has_changed().unwrap());
604 }
605}
606
607#[cfg(test)]
608mod path_hazard_tests {
609 use super::*;
610 use crate::ConnectionId;
611 use subc_protocol::manifest::ModuleManifest;
612
613 fn manifest(module_id: &str) -> ModuleManifest {
614 ModuleManifest::builder(module_id, "0.1.0")
615 .protocol_ver(1)
616 .build()
617 }
618
619 #[test]
620 fn path_hazard_ids_are_refused_and_nothing_registers() {
621 let registry = Registry::default();
622 for (bad, reason_fragment) in [
623 ("../escape", "path separator"),
624 ("a/b", "path separator"),
625 ("a\\b", "path separator"),
626 ("..", "dot path component"),
627 (".", "dot path component"),
628 ("", "empty"),
629 ("evil\u{0}id", "control character"),
630 ] {
631 let err = registry
632 .register_with_control_ops(manifest(bad), 1, ConnectionId::new(7), Vec::new())
633 .expect_err("path-hazard id must refuse");
634 assert!(
636 err.to_string().contains(reason_fragment),
637 "id {bad:?}: expected {reason_fragment:?} in {err}"
638 );
639 }
640 assert_eq!(registry.active_registration_count().unwrap(), 0);
643 assert_eq!(registry.generation().unwrap(), 0);
644 }
645
646 #[test]
647 fn module_id_path_component_length_matches_shared_refusal_vectors() {
648 let doc: serde_json::Value = serde_json::from_str(include_str!(
649 "../tests/golden/module_id_path_component_refusals.json"
650 ))
651 .expect("refusal fixture parses");
652
653 for case in doc["vectors"].as_array().expect("vectors array") {
654 let name = case["name"].as_str().expect("name");
655 let module_id = case["module_id"]["unit"]
656 .as_str()
657 .expect("module_id unit")
658 .repeat(
659 case["module_id"]["repeat"]
660 .as_u64()
661 .expect("module_id repeat") as usize,
662 );
663 assert_eq!(
664 module_id.len(),
665 case["utf8_bytes"].as_u64().expect("utf8 bytes") as usize
666 );
667
668 let expected = case["expect_reason"].as_str().map(str::to_owned);
669 assert_eq!(
670 module_id_path_hazard(&module_id).err(),
671 expected,
672 "shared refusal vector {name:?} diverged"
673 );
674 }
675 }
676
677 #[test]
678 fn working_id_shapes_register_including_namespace_colons() {
679 let registry = Registry::default();
680 for (i, good) in ["magic-context", "mcp:everything", "v1.2-module"]
681 .iter()
682 .enumerate()
683 {
684 registry
685 .register_with_control_ops(
686 manifest(good),
687 1,
688 ConnectionId::new(10 + i as u64),
689 Vec::new(),
690 )
691 .unwrap_or_else(|err| panic!("id {good:?} must register: {err}"));
692 }
693 assert_eq!(registry.active_registration_count().unwrap(), 3);
694 }
695}
696
697#[cfg(test)]
698mod swap_slot_tests {
699 use super::*;
700
701 fn manifest(module_id: &str, ready: Option<bool>) -> ModuleManifest {
702 let mut manifest = ModuleManifest::builder(module_id, "0.1.0").build();
703 manifest.ready = ready;
704 manifest
705 }
706
707 const INCUMBENT: ConnectionId = ConnectionId(1);
708 const CANDIDATE: ConnectionId = ConnectionId(2);
709
710 fn registry_with_candidate() -> Registry {
711 let registry = Registry::default();
712 registry
713 .register_with_control_ops(manifest("m", None), 1, INCUMBENT, Vec::new())
714 .unwrap();
715 registry
716 .register_candidate_with_control_ops(
717 manifest("m", Some(false)),
718 1,
719 CANDIDATE,
720 Vec::new(),
721 )
722 .unwrap();
723 registry
724 }
725
726 #[test]
727 fn candidate_is_invisible_to_by_id_lookups_and_listing() {
728 let registry = registry_with_candidate();
729 let generation = registry.generation().unwrap();
730 assert_eq!(
731 registry.get_module("m").unwrap().unwrap().connection_id,
732 INCUMBENT
733 );
734 let (listed_generation, listed) = registry.list_modules().unwrap();
735 assert_eq!(listed.len(), 1);
736 assert_eq!(listed[0].connection_id, INCUMBENT);
737 assert_eq!(listed_generation, generation);
738 assert_eq!(registry.active_registration_count().unwrap(), 1);
739 assert_eq!(
740 registry.get_candidate("m").unwrap().unwrap().connection_id,
741 CANDIDATE
742 );
743 assert_eq!(
744 registry
745 .register_candidate_with_control_ops(
746 manifest("m", None),
747 1,
748 ConnectionId(3),
749 Vec::new()
750 )
751 .unwrap_err(),
752 RegistryError::DuplicateModuleId {
753 module_id: "m".to_string()
754 }
755 );
756 }
757
758 #[test]
761 fn candidate_catalog_update_reaches_the_candidate_registration() {
762 let registry = registry_with_candidate();
763 assert!(!registry.get_candidate("m").unwrap().unwrap().ready);
764
765 let updated = registry
766 .replace_catalog_for_connection(CANDIDATE, Vec::new(), None, Some(true))
767 .unwrap()
768 .expect("the candidate's own connection finds its registration");
769
770 assert_eq!(updated.connection_id, CANDIDATE);
771 assert!(registry.get_candidate("m").unwrap().unwrap().ready);
772 assert_eq!(
773 registry.get_module_by_connection(CANDIDATE).unwrap(),
774 Some(updated)
775 );
776 assert_eq!(
777 registry.get_module("m").unwrap().unwrap().connection_id,
778 INCUMBENT,
779 "a candidate's update must not touch the active registration"
780 );
781 }
782
783 #[test]
784 fn promotion_swaps_slots_and_each_connection_still_deregisters_its_own() {
785 let registry = registry_with_candidate();
786 let before = registry.generation().unwrap();
787 let cutover = registry.promote_candidate("m").unwrap().unwrap();
788 assert_eq!(cutover.promoted.connection_id, CANDIDATE);
789 assert_eq!(cutover.superseded.unwrap().connection_id, INCUMBENT);
790 assert_ne!(registry.generation().unwrap(), before);
791 assert_eq!(registry.promote_candidate("m").unwrap(), None);
792
793 assert_eq!(
794 registry
795 .registration(RegistrationSlot::Active("m"))
796 .unwrap()
797 .unwrap()
798 .connection_id,
799 CANDIDATE
800 );
801 assert!(registry
802 .registration(RegistrationSlot::Candidate("m"))
803 .unwrap()
804 .is_none());
805 assert!(registry
806 .registration(RegistrationSlot::Connection(INCUMBENT))
807 .unwrap()
808 .is_some());
809
810 let closed = registry.deregister_connection(INCUMBENT).unwrap();
811 assert_eq!(closed.len(), 1);
812 assert_eq!(closed[0].connection_id, INCUMBENT);
813 assert_eq!(closed[0].state, ChannelState::Closed);
814 assert!(registry
815 .registration(RegistrationSlot::Connection(INCUMBENT))
816 .unwrap()
817 .is_none());
818 assert_eq!(
819 registry.get_module("m").unwrap().unwrap().connection_id,
820 CANDIDATE
821 );
822 }
823
824 #[test]
825 fn a_dropped_candidate_deregisters_from_the_candidate_slot_only() {
826 let registry = registry_with_candidate();
827 let closed = registry.deregister_connection(CANDIDATE).unwrap();
828 assert_eq!(closed.len(), 1);
829 assert_eq!(closed[0].connection_id, CANDIDATE);
830 assert!(registry.get_candidate("m").unwrap().is_none());
831 assert_eq!(
832 registry.get_module("m").unwrap().unwrap().connection_id,
833 INCUMBENT
834 );
835 }
836}