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}