use std::{
collections::HashMap,
error::Error,
fmt,
sync::{Mutex, MutexGuard},
};
use subc_protocol::manifest::{CapabilityDeclarations, ModuleManifest, ProviderRole};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct ConnectionId(u64);
impl ConnectionId {
#[cfg(test)]
pub const LOCAL: Self = Self(0);
pub const fn new(raw: u64) -> Self {
Self(raw)
}
pub fn get(self) -> u64 {
self.0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ChannelState {
Active,
Closed,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ModuleRegistration {
pub manifest: ModuleManifest,
pub ready: bool,
pub negotiated_ver: u8,
pub state: ChannelState,
pub connection_id: ConnectionId,
pub control_ops: Vec<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RegistrationSlot<'a> {
Active(&'a str),
Candidate(&'a str),
Connection(ConnectionId),
}
#[derive(Debug, Clone, PartialEq)]
pub struct RegistryCutover {
pub promoted: ModuleRegistration,
pub superseded: Option<ModuleRegistration>,
}
#[derive(Debug, Default)]
pub struct Registry {
inner: Mutex<RegistryInner>,
}
#[derive(Debug, Default)]
struct RegistryInner {
modules: HashMap<String, ModuleRegistration>,
candidates: HashMap<String, ModuleRegistration>,
superseded: Vec<ModuleRegistration>,
generation: u64,
}
impl Registry {
pub fn register_with_control_ops(
&self,
manifest: ModuleManifest,
negotiated_ver: u8,
connection_id: ConnectionId,
control_ops: Vec<String>,
) -> Result<ModuleRegistration, RegistryError> {
let module_id = manifest.module_id.clone();
if let Err(reason) = module_id_path_hazard(&module_id) {
return Err(RegistryError::PathHazardModuleId { module_id, reason });
}
let mut inner = self.lock_inner()?;
if inner.modules.contains_key(&module_id) {
return Err(RegistryError::DuplicateModuleId { module_id });
}
let ready = manifest.ready.unwrap_or(true);
let registration = ModuleRegistration {
manifest,
ready,
negotiated_ver,
state: ChannelState::Active,
connection_id,
control_ops,
};
inner.modules.insert(module_id, registration.clone());
inner.bump_generation();
Ok(registration)
}
pub fn register_candidate_with_control_ops(
&self,
manifest: ModuleManifest,
negotiated_ver: u8,
connection_id: ConnectionId,
control_ops: Vec<String>,
) -> Result<ModuleRegistration, RegistryError> {
let module_id = manifest.module_id.clone();
if let Err(reason) = module_id_path_hazard(&module_id) {
return Err(RegistryError::PathHazardModuleId { module_id, reason });
}
let mut inner = self.lock_inner()?;
if inner.candidates.contains_key(&module_id) {
return Err(RegistryError::DuplicateModuleId { module_id });
}
let ready = manifest.ready.unwrap_or(true);
let registration = ModuleRegistration {
manifest,
ready,
negotiated_ver,
state: ChannelState::Active,
connection_id,
control_ops,
};
inner.candidates.insert(module_id, registration.clone());
Ok(registration)
}
pub fn promote_candidate(
&self,
module_id: &str,
) -> Result<Option<RegistryCutover>, RegistryError> {
let mut inner = self.lock_inner()?;
let Some(promoted) = inner.candidates.remove(module_id) else {
return Ok(None);
};
let superseded = inner
.modules
.insert(module_id.to_string(), promoted.clone());
if let Some(superseded) = superseded.clone() {
inner.superseded.push(superseded);
}
inner.bump_generation();
Ok(Some(RegistryCutover {
promoted,
superseded,
}))
}
pub fn get_module(&self, module_id: &str) -> Result<Option<ModuleRegistration>, RegistryError> {
Ok(self.lock_inner()?.modules.get(module_id).cloned())
}
pub fn get_candidate(
&self,
module_id: &str,
) -> Result<Option<ModuleRegistration>, RegistryError> {
Ok(self.lock_inner()?.candidates.get(module_id).cloned())
}
pub fn registration(
&self,
slot: RegistrationSlot<'_>,
) -> Result<Option<ModuleRegistration>, RegistryError> {
let inner = self.lock_inner()?;
Ok(match slot {
RegistrationSlot::Active(module_id) => inner.modules.get(module_id).cloned(),
RegistrationSlot::Candidate(module_id) => inner.candidates.get(module_id).cloned(),
RegistrationSlot::Connection(connection_id) => inner
.find_by_connection(connection_id)
.map(|(_, registration)| registration.clone()),
})
}
pub fn active_registration_count(&self) -> Result<usize, RegistryError> {
Ok(self.lock_inner()?.modules.len())
}
pub fn list_modules(&self) -> Result<(u64, Vec<ModuleRegistration>), RegistryError> {
let inner = self.lock_inner()?;
let mut modules = inner.modules.values().cloned().collect::<Vec<_>>();
modules.sort_by(|left, right| left.manifest.module_id.cmp(&right.manifest.module_id));
Ok((inner.generation, modules))
}
pub fn generation(&self) -> Result<u64, RegistryError> {
Ok(self.lock_inner()?.generation)
}
#[cfg(test)]
pub(crate) fn set_module_state_for_test(
&self,
module_id: &str,
state: ChannelState,
) -> Result<bool, RegistryError> {
let mut inner = self.lock_inner()?;
let Some(registration) = inner.modules.get_mut(module_id) else {
return Ok(false);
};
registration.state = state;
Ok(true)
}
pub fn get_module_by_connection(
&self,
connection_id: ConnectionId,
) -> Result<Option<ModuleRegistration>, RegistryError> {
Ok(self
.lock_inner()?
.find_by_connection(connection_id)
.map(|(_, registration)| registration.clone()))
}
pub fn replace_catalog_for_connection(
&self,
connection_id: ConnectionId,
provides: Vec<ProviderRole>,
capabilities: Option<CapabilityDeclarations>,
ready: Option<bool>,
) -> Result<Option<ModuleRegistration>, RegistryError> {
let mut inner = self.lock_inner()?;
let Some((slot, _)) = inner.find_by_connection(connection_id) else {
return Ok(None);
};
let registration = inner
.registration_mut(slot, connection_id)
.expect("registration discovered under the same registry lock must still exist");
registration.manifest.provides = provides;
if let Some(capabilities) = capabilities {
registration.manifest.capabilities = Some(capabilities);
}
if let Some(ready) = ready {
registration.ready = ready;
registration.manifest.ready = Some(ready);
}
let updated = registration.clone();
if matches!(slot, SlotKind::Active) {
inner.bump_generation();
}
Ok(Some(updated))
}
pub fn deregister_connection(
&self,
connection_id: ConnectionId,
) -> Result<Vec<ModuleRegistration>, RegistryError> {
let mut inner = self.lock_inner()?;
let module_ids: Vec<String> = inner
.modules
.iter()
.filter(|(_, registration)| registration.connection_id == connection_id)
.map(|(module_id, _)| module_id.clone())
.collect();
let mut closed: Vec<ModuleRegistration> = module_ids
.into_iter()
.filter_map(|module_id| inner.close_module(&module_id))
.collect();
let candidate_ids: Vec<String> = inner
.candidates
.iter()
.filter(|(_, registration)| registration.connection_id == connection_id)
.map(|(module_id, _)| module_id.clone())
.collect();
for module_id in candidate_ids {
if let Some(mut registration) = inner.candidates.remove(&module_id) {
registration.state = ChannelState::Closed;
closed.push(registration);
}
}
let (removed, kept): (Vec<_>, Vec<_>) = std::mem::take(&mut inner.superseded)
.into_iter()
.partition(|registration| registration.connection_id == connection_id);
inner.superseded = kept;
closed.extend(removed.into_iter().map(|mut registration| {
registration.state = ChannelState::Closed;
registration
}));
Ok(closed)
}
fn lock_inner(&self) -> Result<MutexGuard<'_, RegistryInner>, RegistryError> {
self.inner.lock().map_err(|_| RegistryError::Poisoned)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum SlotKind {
Active,
Candidate,
Superseded,
}
impl RegistryInner {
fn find_by_connection(
&self,
connection_id: ConnectionId,
) -> Option<(SlotKind, &ModuleRegistration)> {
let owned_by =
|registration: &&ModuleRegistration| registration.connection_id == connection_id;
self.modules
.values()
.find(owned_by)
.map(|registration| (SlotKind::Active, registration))
.or_else(|| {
self.candidates
.values()
.find(owned_by)
.map(|registration| (SlotKind::Candidate, registration))
})
.or_else(|| {
self.superseded
.iter()
.find(owned_by)
.map(|registration| (SlotKind::Superseded, registration))
})
}
fn registration_mut(
&mut self,
slot: SlotKind,
connection_id: ConnectionId,
) -> Option<&mut ModuleRegistration> {
let owned_by =
|registration: &&mut ModuleRegistration| registration.connection_id == connection_id;
match slot {
SlotKind::Active => self.modules.values_mut().find(owned_by),
SlotKind::Candidate => self.candidates.values_mut().find(owned_by),
SlotKind::Superseded => self.superseded.iter_mut().find(owned_by),
}
}
fn close_module(&mut self, module_id: &str) -> Option<ModuleRegistration> {
let mut registration = self.modules.remove(module_id)?;
registration.state = ChannelState::Closed;
self.bump_generation();
Some(registration)
}
fn bump_generation(&mut self) {
self.generation = self.generation.wrapping_add(1);
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RegistryError {
DuplicateModuleId {
module_id: String,
},
PathHazardModuleId {
module_id: String,
reason: String,
},
Poisoned,
}
impl fmt::Display for RegistryError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::DuplicateModuleId { module_id } => {
write!(f, "module_id '{module_id}' is already registered")
}
Self::PathHazardModuleId { module_id, reason } => {
write!(
f,
"module_id '{}' is not usable as a path component: {reason}",
module_id.escape_debug()
)
}
Self::Poisoned => write!(f, "registry lock was poisoned"),
}
}
}
impl Error for RegistryError {}
pub fn module_id_path_hazard(module_id: &str) -> Result<(), String> {
if module_id.is_empty() {
return Err("empty".to_string());
}
if module_id.contains('/') || module_id.contains('\\') {
return Err("contains a path separator".to_string());
}
if module_id == "." || module_id == ".." {
return Err("is a dot path component".to_string());
}
if module_id.chars().any(|c| c.is_control()) {
return Err("contains a control character".to_string());
}
if module_id.len() > 255 {
return Err("is longer than 255 bytes".to_string());
}
Ok(())
}
#[cfg(test)]
mod path_hazard_tests {
use super::*;
use crate::ConnectionId;
use subc_protocol::manifest::ModuleManifest;
fn manifest(module_id: &str) -> ModuleManifest {
ModuleManifest::builder(module_id, "0.1.0")
.protocol_ver(1)
.build()
}
#[test]
fn path_hazard_ids_are_refused_and_nothing_registers() {
let registry = Registry::default();
for (bad, reason_fragment) in [
("../escape", "path separator"),
("a/b", "path separator"),
("a\\b", "path separator"),
("..", "dot path component"),
(".", "dot path component"),
("", "empty"),
("evil\u{0}id", "control character"),
] {
let err = registry
.register_with_control_ops(manifest(bad), 1, ConnectionId::new(7), Vec::new())
.expect_err("path-hazard id must refuse");
assert!(
err.to_string().contains(reason_fragment),
"id {bad:?}: expected {reason_fragment:?} in {err}"
);
}
assert_eq!(registry.active_registration_count().unwrap(), 0);
assert_eq!(registry.generation().unwrap(), 0);
}
#[test]
fn module_id_path_component_length_matches_shared_refusal_vectors() {
let doc: serde_json::Value = serde_json::from_str(include_str!(
"../tests/golden/module_id_path_component_refusals.json"
))
.expect("refusal fixture parses");
for case in doc["vectors"].as_array().expect("vectors array") {
let name = case["name"].as_str().expect("name");
let module_id = case["module_id"]["unit"]
.as_str()
.expect("module_id unit")
.repeat(
case["module_id"]["repeat"]
.as_u64()
.expect("module_id repeat") as usize,
);
assert_eq!(
module_id.len(),
case["utf8_bytes"].as_u64().expect("utf8 bytes") as usize
);
let expected = case["expect_reason"].as_str().map(str::to_owned);
assert_eq!(
module_id_path_hazard(&module_id).err(),
expected,
"shared refusal vector {name:?} diverged"
);
}
}
#[test]
fn working_id_shapes_register_including_namespace_colons() {
let registry = Registry::default();
for (i, good) in ["magic-context", "mcp:everything", "v1.2-module"]
.iter()
.enumerate()
{
registry
.register_with_control_ops(
manifest(good),
1,
ConnectionId::new(10 + i as u64),
Vec::new(),
)
.unwrap_or_else(|err| panic!("id {good:?} must register: {err}"));
}
assert_eq!(registry.active_registration_count().unwrap(), 3);
}
}
#[cfg(test)]
mod swap_slot_tests {
use super::*;
fn manifest(module_id: &str, ready: Option<bool>) -> ModuleManifest {
let mut manifest = ModuleManifest::builder(module_id, "0.1.0").build();
manifest.ready = ready;
manifest
}
const INCUMBENT: ConnectionId = ConnectionId(1);
const CANDIDATE: ConnectionId = ConnectionId(2);
fn registry_with_candidate() -> Registry {
let registry = Registry::default();
registry
.register_with_control_ops(manifest("m", None), 1, INCUMBENT, Vec::new())
.unwrap();
registry
.register_candidate_with_control_ops(
manifest("m", Some(false)),
1,
CANDIDATE,
Vec::new(),
)
.unwrap();
registry
}
#[test]
fn candidate_is_invisible_to_by_id_lookups_and_listing() {
let registry = registry_with_candidate();
let generation = registry.generation().unwrap();
assert_eq!(
registry.get_module("m").unwrap().unwrap().connection_id,
INCUMBENT
);
let (listed_generation, listed) = registry.list_modules().unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].connection_id, INCUMBENT);
assert_eq!(listed_generation, generation);
assert_eq!(registry.active_registration_count().unwrap(), 1);
assert_eq!(
registry.get_candidate("m").unwrap().unwrap().connection_id,
CANDIDATE
);
assert_eq!(
registry
.register_candidate_with_control_ops(
manifest("m", None),
1,
ConnectionId(3),
Vec::new()
)
.unwrap_err(),
RegistryError::DuplicateModuleId {
module_id: "m".to_string()
}
);
}
#[test]
fn candidate_catalog_update_reaches_the_candidate_registration() {
let registry = registry_with_candidate();
assert!(!registry.get_candidate("m").unwrap().unwrap().ready);
let updated = registry
.replace_catalog_for_connection(CANDIDATE, Vec::new(), None, Some(true))
.unwrap()
.expect("the candidate's own connection finds its registration");
assert_eq!(updated.connection_id, CANDIDATE);
assert!(registry.get_candidate("m").unwrap().unwrap().ready);
assert_eq!(
registry.get_module_by_connection(CANDIDATE).unwrap(),
Some(updated)
);
assert_eq!(
registry.get_module("m").unwrap().unwrap().connection_id,
INCUMBENT,
"a candidate's update must not touch the active registration"
);
}
#[test]
fn promotion_swaps_slots_and_each_connection_still_deregisters_its_own() {
let registry = registry_with_candidate();
let before = registry.generation().unwrap();
let cutover = registry.promote_candidate("m").unwrap().unwrap();
assert_eq!(cutover.promoted.connection_id, CANDIDATE);
assert_eq!(cutover.superseded.unwrap().connection_id, INCUMBENT);
assert_ne!(registry.generation().unwrap(), before);
assert_eq!(registry.promote_candidate("m").unwrap(), None);
assert_eq!(
registry
.registration(RegistrationSlot::Active("m"))
.unwrap()
.unwrap()
.connection_id,
CANDIDATE
);
assert!(registry
.registration(RegistrationSlot::Candidate("m"))
.unwrap()
.is_none());
assert!(registry
.registration(RegistrationSlot::Connection(INCUMBENT))
.unwrap()
.is_some());
let closed = registry.deregister_connection(INCUMBENT).unwrap();
assert_eq!(closed.len(), 1);
assert_eq!(closed[0].connection_id, INCUMBENT);
assert_eq!(closed[0].state, ChannelState::Closed);
assert!(registry
.registration(RegistrationSlot::Connection(INCUMBENT))
.unwrap()
.is_none());
assert_eq!(
registry.get_module("m").unwrap().unwrap().connection_id,
CANDIDATE
);
}
#[test]
fn a_dropped_candidate_deregisters_from_the_candidate_slot_only() {
let registry = registry_with_candidate();
let closed = registry.deregister_connection(CANDIDATE).unwrap();
assert_eq!(closed.len(), 1);
assert_eq!(closed[0].connection_id, CANDIDATE);
assert!(registry.get_candidate("m").unwrap().is_none());
assert_eq!(
registry.get_module("m").unwrap().unwrap().connection_id,
INCUMBENT
);
}
}