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 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 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 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}