Skip to main content

macp_runtime/extensions/
registry.rs

1use super::provider::{SessionExtensionProvider, SessionOutcome};
2use std::collections::HashMap;
3
4pub struct ExtensionProviderRegistry {
5    providers: Vec<Box<dyn SessionExtensionProvider>>,
6}
7
8impl ExtensionProviderRegistry {
9    pub fn new() -> Self {
10        Self {
11            providers: Vec::new(),
12        }
13    }
14
15    pub fn register(&mut self, provider: Box<dyn SessionExtensionProvider>) {
16        tracing::info!(key = provider.key(), "registered extension provider");
17        self.providers.push(provider);
18    }
19
20    pub async fn on_session_start(&self, session_id: &str, extensions: &HashMap<String, Vec<u8>>) {
21        for provider in &self.providers {
22            if !extensions.contains_key(provider.key()) {
23                continue;
24            }
25            if let Err(e) = provider.on_session_start(session_id, extensions).await {
26                tracing::warn!(
27                    key = provider.key(),
28                    session_id,
29                    error = %e,
30                    "extension provider on_session_start failed (non-fatal)"
31                );
32            }
33        }
34    }
35
36    pub async fn on_session_terminal(&self, session_id: &str, outcome: SessionOutcome) {
37        for provider in &self.providers {
38            if let Err(e) = provider
39                .on_session_terminal(session_id, outcome_ref(&outcome))
40                .await
41            {
42                tracing::warn!(
43                    key = provider.key(),
44                    session_id,
45                    error = %e,
46                    "extension provider on_session_terminal failed (non-fatal)"
47                );
48            }
49        }
50    }
51}
52
53impl Default for ExtensionProviderRegistry {
54    fn default() -> Self {
55        Self::new()
56    }
57}
58
59fn outcome_ref(outcome: &SessionOutcome) -> SessionOutcome {
60    match outcome {
61        SessionOutcome::Resolved => SessionOutcome::Resolved,
62        SessionOutcome::Expired => SessionOutcome::Expired,
63    }
64}
65
66#[cfg(test)]
67mod tests {
68    use super::*;
69    use crate::extensions::provider::ExtensionError;
70    use std::sync::{Arc, Mutex};
71
72    /// Provider that records every lifecycle callback it receives.
73    struct RecordingProvider {
74        key: &'static str,
75        calls: Arc<Mutex<Vec<String>>>,
76        fail: bool,
77    }
78
79    #[async_trait::async_trait]
80    impl SessionExtensionProvider for RecordingProvider {
81        fn key(&self) -> &str {
82            self.key
83        }
84
85        async fn on_session_start(
86            &self,
87            session_id: &str,
88            _extensions: &HashMap<String, Vec<u8>>,
89        ) -> Result<(), ExtensionError> {
90            self.calls
91                .lock()
92                .unwrap()
93                .push(format!("{}:start:{session_id}", self.key));
94            if self.fail {
95                Err(ExtensionError::Internal("boom".into()))
96            } else {
97                Ok(())
98            }
99        }
100
101        async fn on_session_terminal(
102            &self,
103            session_id: &str,
104            outcome: SessionOutcome,
105        ) -> Result<(), ExtensionError> {
106            let outcome = match outcome {
107                SessionOutcome::Resolved => "resolved",
108                SessionOutcome::Expired => "expired",
109            };
110            self.calls
111                .lock()
112                .unwrap()
113                .push(format!("{}:terminal:{session_id}:{outcome}", self.key));
114            if self.fail {
115                Err(ExtensionError::Internal("boom".into()))
116            } else {
117                Ok(())
118            }
119        }
120    }
121
122    fn recording(
123        key: &'static str,
124        fail: bool,
125    ) -> (Box<RecordingProvider>, Arc<Mutex<Vec<String>>>) {
126        let calls = Arc::new(Mutex::new(Vec::new()));
127        (
128            Box::new(RecordingProvider {
129                key,
130                calls: Arc::clone(&calls),
131                fail,
132            }),
133            calls,
134        )
135    }
136
137    fn extensions_with_key(key: &str) -> HashMap<String, Vec<u8>> {
138        let mut map = HashMap::new();
139        map.insert(key.to_string(), b"cfg".to_vec());
140        map
141    }
142
143    #[tokio::test]
144    async fn on_session_start_dispatches_only_to_providers_with_matching_key() {
145        let mut registry = ExtensionProviderRegistry::new();
146        let (a, a_calls) = recording("ext.a", false);
147        let (b, b_calls) = recording("ext.b", false);
148        registry.register(a);
149        registry.register(b);
150
151        registry
152            .on_session_start("s1", &extensions_with_key("ext.a"))
153            .await;
154
155        assert_eq!(*a_calls.lock().unwrap(), vec!["ext.a:start:s1"]);
156        assert!(
157            b_calls.lock().unwrap().is_empty(),
158            "provider whose key is absent from the extensions map must not be invoked"
159        );
160    }
161
162    #[tokio::test]
163    async fn on_session_start_with_unknown_key_invokes_no_provider() {
164        let mut registry = ExtensionProviderRegistry::new();
165        let (a, a_calls) = recording("ext.a", false);
166        registry.register(a);
167
168        registry
169            .on_session_start("s1", &extensions_with_key("ext.unknown"))
170            .await;
171        registry.on_session_start("s2", &HashMap::new()).await;
172
173        assert!(a_calls.lock().unwrap().is_empty());
174    }
175
176    #[tokio::test]
177    async fn on_session_terminal_notifies_all_registered_providers() {
178        let mut registry = ExtensionProviderRegistry::new();
179        let (a, a_calls) = recording("ext.a", false);
180        let (b, b_calls) = recording("ext.b", false);
181        registry.register(a);
182        registry.register(b);
183
184        registry
185            .on_session_terminal("s1", SessionOutcome::Resolved)
186            .await;
187        registry
188            .on_session_terminal("s2", SessionOutcome::Expired)
189            .await;
190
191        assert_eq!(
192            *a_calls.lock().unwrap(),
193            vec!["ext.a:terminal:s1:resolved", "ext.a:terminal:s2:expired"]
194        );
195        assert_eq!(
196            *b_calls.lock().unwrap(),
197            vec!["ext.b:terminal:s1:resolved", "ext.b:terminal:s2:expired"]
198        );
199    }
200
201    #[tokio::test]
202    async fn provider_failure_is_non_fatal_and_later_providers_still_run() {
203        let mut registry = ExtensionProviderRegistry::new();
204        let (failing, failing_calls) = recording("ext.a", true);
205        let (ok, ok_calls) = recording("ext.b", false);
206        registry.register(failing);
207        registry.register(ok);
208
209        let mut extensions = extensions_with_key("ext.a");
210        extensions.insert("ext.b".to_string(), vec![]);
211        registry.on_session_start("s1", &extensions).await;
212        registry
213            .on_session_terminal("s1", SessionOutcome::Resolved)
214            .await;
215
216        // The failing provider was invoked and its error swallowed (E-1);
217        // the provider registered after it still ran for both callbacks.
218        assert_eq!(
219            *failing_calls.lock().unwrap(),
220            vec!["ext.a:start:s1", "ext.a:terminal:s1:resolved"]
221        );
222        assert_eq!(
223            *ok_calls.lock().unwrap(),
224            vec!["ext.b:start:s1", "ext.b:terminal:s1:resolved"]
225        );
226    }
227
228    #[tokio::test]
229    async fn duplicate_key_registration_keeps_both_providers() {
230        // The registry is an append-only dispatch list: registering a second
231        // provider under the same key does not replace the first — both are
232        // invoked, in registration order.
233        let mut registry = ExtensionProviderRegistry::new();
234        let (first, first_calls) = recording("ext.a", false);
235        let (second, second_calls) = recording("ext.a", false);
236        registry.register(first);
237        registry.register(second);
238
239        registry
240            .on_session_start("s1", &extensions_with_key("ext.a"))
241            .await;
242
243        assert_eq!(*first_calls.lock().unwrap(), vec!["ext.a:start:s1"]);
244        assert_eq!(*second_calls.lock().unwrap(), vec!["ext.a:start:s1"]);
245    }
246
247    #[tokio::test]
248    async fn empty_registry_dispatch_is_a_no_op() {
249        let registry = ExtensionProviderRegistry::default();
250        registry
251            .on_session_start("s1", &extensions_with_key("ext.a"))
252            .await;
253        registry
254            .on_session_terminal("s1", SessionOutcome::Expired)
255            .await;
256    }
257}