Skip to main content

aster_forge_external_auth/
registry.rs

1//! Runtime registry for feature-enabled external authentication provider drivers.
2//!
3//! The default registry registers only drivers compiled into the crate through Cargo features.
4//! Applications can create their own registry and call [`ExternalAuthProviderRegistry::add`] to
5//! append product or plugin-provided drivers without replacing built-ins. Tests and advanced
6//! application code can still call [`ExternalAuthProviderRegistry::register`] when intentional
7//! replacement is required.
8
9use std::collections::HashMap;
10use std::sync::{Arc, OnceLock};
11
12use super::driver::{
13    ExternalAuthProviderConfig, ExternalAuthProviderDescriptor, ExternalAuthProviderDriver,
14};
15#[cfg(feature = "github")]
16use super::providers::github::GitHubProviderDriver;
17#[cfg(feature = "google")]
18use super::providers::google::GoogleProviderDriver;
19#[cfg(feature = "microsoft")]
20use super::providers::microsoft::MicrosoftProviderDriver;
21#[cfg(feature = "oauth2")]
22use super::providers::oauth2::OAuth2ProviderDriver;
23#[cfg(feature = "oidc")]
24use super::providers::oidc::OidcProviderDriver;
25#[cfg(feature = "qq")]
26use super::providers::qq::QqProviderDriver;
27use crate::types::ExternalAuthProviderKind;
28use crate::{ExternalAuthError, Result};
29
30/// Registry of external authentication provider drivers keyed by provider kind.
31///
32/// The registry is the shared capability boundary for application services. Callers can use it to
33/// list descriptors for admin surfaces, gate stored provider configs before a login flow starts,
34/// and retrieve the runtime driver for a validated provider. The registry deliberately does not
35/// know how providers are stored in a product database; applications adapt their own rows into
36/// [`ExternalAuthProviderConfig`] immediately before using this type.
37pub struct ExternalAuthProviderRegistry {
38    drivers: HashMap<ExternalAuthProviderKind, Arc<dyn ExternalAuthProviderDriver>>,
39}
40
41impl ExternalAuthProviderRegistry {
42    /// Creates an empty registry with no built-in provider drivers.
43    ///
44    /// This is useful for applications that want a fully explicit provider list, tests that need
45    /// deterministic registration behavior across Cargo feature sets, or plugin hosts that build a
46    /// registry from externally supplied drivers.
47    pub fn empty() -> Self {
48        Self {
49            drivers: HashMap::new(),
50        }
51    }
52
53    /// Creates a registry populated with all feature-enabled built-in drivers.
54    pub fn new() -> Self {
55        #[allow(unused_mut)]
56        let mut registry = Self::empty();
57        #[cfg(feature = "oidc")]
58        registry.register_builtin(OidcProviderDriver::new());
59        #[cfg(feature = "oauth2")]
60        registry.register_builtin(OAuth2ProviderDriver::new());
61        #[cfg(feature = "github")]
62        registry.register_builtin(GitHubProviderDriver::new());
63        #[cfg(feature = "google")]
64        registry.register_builtin(GoogleProviderDriver::new());
65        #[cfg(feature = "microsoft")]
66        registry.register_builtin(MicrosoftProviderDriver::new());
67        #[cfg(feature = "qq")]
68        registry.register_builtin(QqProviderDriver::new());
69        registry
70    }
71
72    /// Creates a built-in registry and lets an external system append registrations.
73    ///
74    /// This is the intended integration point for application-level extension systems: the caller
75    /// receives a mutable registry, calls [`ExternalAuthProviderRegistry::add`] for each external
76    /// driver it wants to expose, and returns any setup error. Built-in drivers are registered
77    /// before the callback runs, so external systems cannot accidentally replace them through the
78    /// non-replacing add API.
79    pub fn with_external_registrations<F>(configure: F) -> Result<Self>
80    where
81        F: FnOnce(&mut Self) -> Result<()>,
82    {
83        let mut registry = Self::new();
84        configure(&mut registry)?;
85        Ok(registry)
86    }
87
88    /// Adds a driver if its provider kind is not already registered.
89    ///
90    /// Use this for product-specific or plugin-provided drivers. Duplicate provider kinds return a
91    /// configuration error instead of replacing the existing driver, which keeps built-in behavior
92    /// stable when external systems are enabled.
93    pub fn add<D>(&mut self, driver: D) -> Result<()>
94    where
95        D: ExternalAuthProviderDriver + 'static,
96    {
97        self.add_arc(Arc::new(driver))
98    }
99
100    /// Adds an already shared driver if its provider kind is not already registered.
101    pub fn add_arc(&mut self, driver: Arc<dyn ExternalAuthProviderDriver>) -> Result<()> {
102        let kind = driver.kind();
103        Self::validate_driver_descriptor(kind, driver.descriptor())?;
104        if self.drivers.contains_key(&kind) {
105            return Err(ExternalAuthError::config_error(format!(
106                "external auth provider driver '{}' is already registered",
107                kind.as_str()
108            )));
109        }
110        self.drivers.insert(kind, driver);
111        Ok(())
112    }
113
114    /// Registers or replaces a driver for its provider kind.
115    ///
116    /// Use this only when replacement is intentional, such as tests or product-level overrides.
117    /// The driver must still report a descriptor for the same provider kind that it registers.
118    pub fn register<D>(&mut self, driver: D) -> Result<()>
119    where
120        D: ExternalAuthProviderDriver + 'static,
121    {
122        self.register_arc(Arc::new(driver))
123    }
124
125    /// Registers or replaces an already shared driver for its provider kind.
126    pub fn register_arc(&mut self, driver: Arc<dyn ExternalAuthProviderDriver>) -> Result<()> {
127        let kind = driver.kind();
128        Self::validate_driver_descriptor(kind, driver.descriptor())?;
129        self.drivers.insert(kind, driver);
130        Ok(())
131    }
132
133    /// Iterates over registered provider kinds.
134    pub fn supported_kinds(&self) -> impl Iterator<Item = ExternalAuthProviderKind> + '_ {
135        self.drivers.keys().copied()
136    }
137
138    /// Returns whether a driver for `kind` is registered.
139    pub fn contains(&self, kind: ExternalAuthProviderKind) -> bool {
140        self.drivers.contains_key(&kind)
141    }
142
143    /// Returns registered provider descriptors sorted by provider kind.
144    pub fn descriptors(&self) -> Vec<ExternalAuthProviderDescriptor> {
145        let mut descriptors = self
146            .drivers
147            .values()
148            .map(|driver| driver.descriptor())
149            .collect::<Vec<_>>();
150        descriptors.sort_by_key(|descriptor| descriptor.kind.as_str());
151        descriptors
152    }
153
154    /// Returns the descriptor for a registered provider kind.
155    pub fn descriptor_for(
156        &self,
157        kind: ExternalAuthProviderKind,
158    ) -> Result<ExternalAuthProviderDescriptor> {
159        Ok(self.get_driver(kind)?.descriptor())
160    }
161
162    /// Ensures that a provider kind is enabled in this registry.
163    ///
164    /// This is useful for service-layer guards that need to reject disabled provider kinds without
165    /// constructing a login flow or exposing the underlying driver.
166    pub fn ensure_provider_supported(&self, kind: ExternalAuthProviderKind) -> Result<()> {
167        if self.contains(kind) {
168            return Ok(());
169        }
170        Err(ExternalAuthError::config_error(format!(
171            "external auth provider driver '{}' is not registered",
172            kind.as_str()
173        )))
174    }
175
176    /// Validates that a product-owned provider config matches a registered driver descriptor.
177    ///
178    /// This catches configuration drift before provider-specific network calls run. The method
179    /// checks that the provider kind is enabled and that the stored protocol matches the driver's
180    /// declared protocol. It returns the descriptor so callers can keep using the capability data
181    /// without another registry lookup.
182    pub fn validate_provider_config(
183        &self,
184        provider: &ExternalAuthProviderConfig,
185    ) -> Result<ExternalAuthProviderDescriptor> {
186        let descriptor = self.descriptor_for(provider.provider_kind)?;
187        if provider.protocol != descriptor.protocol {
188            return Err(ExternalAuthError::validation_error(format!(
189                "external auth provider '{}' is configured with protocol '{}' but driver expects '{}'",
190                provider.provider_kind.as_str(),
191                provider.protocol.as_str(),
192                descriptor.protocol.as_str()
193            )));
194        }
195        Ok(descriptor)
196    }
197
198    /// Returns the registered driver for a provider config after validating the config boundary.
199    ///
200    /// Product services should prefer this method when starting authorization, exchanging a
201    /// callback, or testing a provider because it applies the same registry-level gates for every
202    /// flow.
203    pub fn driver_for_provider(
204        &self,
205        provider: &ExternalAuthProviderConfig,
206    ) -> Result<Arc<dyn ExternalAuthProviderDriver>> {
207        self.validate_provider_config(provider)?;
208        self.get_driver(provider.provider_kind)
209    }
210
211    /// Returns a registered driver by provider kind.
212    pub fn get_driver(
213        &self,
214        kind: ExternalAuthProviderKind,
215    ) -> Result<Arc<dyn ExternalAuthProviderDriver>> {
216        self.drivers.get(&kind).cloned().ok_or_else(|| {
217            ExternalAuthError::config_error(format!(
218                "external auth provider driver '{}' is not registered",
219                kind.as_str()
220            ))
221        })
222    }
223
224    /// Returns the OIDC driver from this registry.
225    pub fn oidc(&self) -> Result<Arc<dyn ExternalAuthProviderDriver>> {
226        self.get_driver(ExternalAuthProviderKind::Oidc)
227    }
228
229    fn validate_driver_descriptor(
230        kind: ExternalAuthProviderKind,
231        descriptor: ExternalAuthProviderDescriptor,
232    ) -> Result<()> {
233        if descriptor.kind == kind {
234            return Ok(());
235        }
236        Err(ExternalAuthError::config_error(format!(
237            "external auth provider driver '{}' returned descriptor for '{}'",
238            kind.as_str(),
239            descriptor.kind.as_str()
240        )))
241    }
242
243    #[cfg(any(
244        feature = "github",
245        feature = "google",
246        feature = "microsoft",
247        feature = "oauth2",
248        feature = "oidc",
249        feature = "qq"
250    ))]
251    fn register_builtin<D>(&mut self, driver: D)
252    where
253        D: ExternalAuthProviderDriver + 'static,
254    {
255        let kind = driver.kind();
256        self.drivers.insert(kind, Arc::new(driver));
257    }
258}
259
260impl Default for ExternalAuthProviderRegistry {
261    fn default() -> Self {
262        Self::new()
263    }
264}
265
266/// Returns a process-wide default registry populated with feature-enabled drivers.
267pub fn default_registry() -> &'static ExternalAuthProviderRegistry {
268    static REGISTRY: OnceLock<ExternalAuthProviderRegistry> = OnceLock::new();
269    REGISTRY.get_or_init(ExternalAuthProviderRegistry::new)
270}
271
272#[cfg(test)]
273mod tests {
274    use super::*;
275    use crate::{
276        ExternalAuthAuthorizationStart, ExternalAuthCallback, ExternalAuthProfile,
277        ExternalAuthProviderConfig, ExternalAuthProviderTestResult,
278    };
279    use async_trait::async_trait;
280
281    #[derive(Default)]
282    struct TestOidcDriver;
283
284    #[async_trait]
285    impl ExternalAuthProviderDriver for TestOidcDriver {
286        fn kind(&self) -> ExternalAuthProviderKind {
287            ExternalAuthProviderKind::Oidc
288        }
289
290        fn descriptor(&self) -> ExternalAuthProviderDescriptor {
291            ExternalAuthProviderDescriptor {
292                kind: ExternalAuthProviderKind::Oidc,
293                protocol: crate::types::ExternalAuthProtocol::Oidc,
294                display_name: "Test OIDC",
295                description: "Test OIDC driver",
296                default_scopes: "openid email profile",
297                issuer_url_required: true,
298                manual_endpoint_configuration_supported: false,
299                authorization_url_required: false,
300                token_url_required: false,
301                userinfo_url_required: false,
302                supports_discovery: true,
303                supports_pkce: true,
304                supports_email_verified_claim: true,
305            }
306        }
307
308        async fn start_authorization(
309            &self,
310            _provider: &ExternalAuthProviderConfig,
311            _redirect_uri: &str,
312        ) -> Result<ExternalAuthAuthorizationStart> {
313            unreachable!("registry tests only inspect driver registration")
314        }
315
316        async fn exchange_callback(
317            &self,
318            _provider: &ExternalAuthProviderConfig,
319            _callback: ExternalAuthCallback,
320        ) -> Result<ExternalAuthProfile> {
321            unreachable!("registry tests only inspect driver registration")
322        }
323
324        async fn test_provider(
325            &self,
326            _provider: &ExternalAuthProviderConfig,
327        ) -> Result<ExternalAuthProviderTestResult> {
328            unreachable!("registry tests only inspect driver registration")
329        }
330    }
331
332    #[derive(Default)]
333    struct MismatchedDescriptorDriver;
334
335    #[async_trait]
336    impl ExternalAuthProviderDriver for MismatchedDescriptorDriver {
337        fn kind(&self) -> ExternalAuthProviderKind {
338            ExternalAuthProviderKind::Oidc
339        }
340
341        fn descriptor(&self) -> ExternalAuthProviderDescriptor {
342            ExternalAuthProviderDescriptor {
343                kind: ExternalAuthProviderKind::GenericOAuth2,
344                protocol: crate::types::ExternalAuthProtocol::OAuth2,
345                display_name: "Mismatched driver",
346                description: "Driver with inconsistent registry metadata",
347                default_scopes: "email",
348                issuer_url_required: false,
349                manual_endpoint_configuration_supported: true,
350                authorization_url_required: true,
351                token_url_required: true,
352                userinfo_url_required: true,
353                supports_discovery: false,
354                supports_pkce: true,
355                supports_email_verified_claim: false,
356            }
357        }
358
359        async fn start_authorization(
360            &self,
361            _provider: &ExternalAuthProviderConfig,
362            _redirect_uri: &str,
363        ) -> Result<ExternalAuthAuthorizationStart> {
364            unreachable!("registry tests only inspect driver registration")
365        }
366
367        async fn exchange_callback(
368            &self,
369            _provider: &ExternalAuthProviderConfig,
370            _callback: ExternalAuthCallback,
371        ) -> Result<ExternalAuthProfile> {
372            unreachable!("registry tests only inspect driver registration")
373        }
374
375        async fn test_provider(
376            &self,
377            _provider: &ExternalAuthProviderConfig,
378        ) -> Result<ExternalAuthProviderTestResult> {
379            unreachable!("registry tests only inspect driver registration")
380        }
381    }
382
383    fn oidc_provider_config() -> ExternalAuthProviderConfig {
384        ExternalAuthProviderConfig {
385            id: 1,
386            key: "test-oidc".to_string(),
387            provider_kind: ExternalAuthProviderKind::Oidc,
388            protocol: crate::types::ExternalAuthProtocol::Oidc,
389            options: crate::types::ExternalAuthProviderOptions::default(),
390            issuer_url: Some("https://issuer.example.com".to_string()),
391            authorization_url: None,
392            token_url: None,
393            userinfo_url: None,
394            client_id: "client-id".to_string(),
395            client_secret: Some("client-secret".to_string()),
396            scopes: "openid email profile".to_string(),
397            subject_claim: None,
398            username_claim: None,
399            display_name_claim: None,
400            email_claim: None,
401            email_verified_claim: None,
402            groups_claim: None,
403            avatar_url_claim: None,
404            outbound_http_user_agent: None,
405        }
406    }
407
408    #[cfg(feature = "oidc")]
409    #[test]
410    fn registry_returns_oidc_driver_by_kind() {
411        let registry = ExternalAuthProviderRegistry::new();
412        let driver = registry
413            .get_driver(ExternalAuthProviderKind::Oidc)
414            .expect("OIDC driver should be registered");
415
416        assert_eq!(driver.kind(), ExternalAuthProviderKind::Oidc);
417    }
418
419    #[test]
420    fn registry_allows_driver_replacement_by_kind() {
421        let mut registry = ExternalAuthProviderRegistry::new();
422        registry
423            .register(TestOidcDriver)
424            .expect("replacement driver should register");
425
426        assert!(registry.contains(ExternalAuthProviderKind::Oidc));
427        #[cfg(feature = "oauth2")]
428        assert!(registry.contains(ExternalAuthProviderKind::GenericOAuth2));
429        #[cfg(feature = "github")]
430        assert!(registry.contains(ExternalAuthProviderKind::GitHub));
431        #[cfg(feature = "google")]
432        assert!(registry.contains(ExternalAuthProviderKind::Google));
433        #[cfg(feature = "microsoft")]
434        assert!(registry.contains(ExternalAuthProviderKind::Microsoft));
435        #[cfg(feature = "qq")]
436        assert!(registry.contains(ExternalAuthProviderKind::Qq));
437    }
438
439    #[test]
440    fn registry_register_rejects_driver_descriptor_kind_mismatch() {
441        let mut registry = ExternalAuthProviderRegistry {
442            drivers: HashMap::new(),
443        };
444
445        let error = registry
446            .register(MismatchedDescriptorDriver)
447            .expect_err("mismatched descriptor should fail");
448
449        assert!(error.to_string().contains(
450            "external auth provider driver 'oidc' returned descriptor for 'generic_oauth2'"
451        ));
452        assert!(!registry.contains(ExternalAuthProviderKind::Oidc));
453    }
454
455    #[test]
456    fn registry_add_rejects_duplicate_kind_without_replacing_existing_driver() {
457        let mut registry = ExternalAuthProviderRegistry {
458            drivers: HashMap::new(),
459        };
460        registry
461            .add(TestOidcDriver)
462            .expect("initial add should work");
463
464        let error = registry
465            .add(TestOidcDriver)
466            .expect_err("duplicate add should fail");
467
468        assert!(
469            error
470                .to_string()
471                .contains("external auth provider driver 'oidc' is already registered")
472        );
473    }
474
475    #[test]
476    fn registry_add_rejects_driver_descriptor_kind_mismatch() {
477        let mut registry = ExternalAuthProviderRegistry {
478            drivers: HashMap::new(),
479        };
480
481        let error = registry
482            .add(MismatchedDescriptorDriver)
483            .expect_err("mismatched descriptor should fail");
484
485        assert!(error.to_string().contains(
486            "external auth provider driver 'oidc' returned descriptor for 'generic_oauth2'"
487        ));
488        assert!(!registry.contains(ExternalAuthProviderKind::Oidc));
489    }
490
491    #[test]
492    fn registry_descriptor_for_returns_registered_descriptor() {
493        let mut registry = ExternalAuthProviderRegistry {
494            drivers: HashMap::new(),
495        };
496        registry
497            .add(TestOidcDriver)
498            .expect("test driver should register");
499
500        let descriptor = registry
501            .descriptor_for(ExternalAuthProviderKind::Oidc)
502            .expect("descriptor should exist");
503
504        assert_eq!(descriptor.kind, ExternalAuthProviderKind::Oidc);
505        assert_eq!(
506            descriptor.protocol,
507            crate::types::ExternalAuthProtocol::Oidc
508        );
509    }
510
511    #[test]
512    fn registry_validate_provider_config_rejects_protocol_mismatch() {
513        let mut registry = ExternalAuthProviderRegistry {
514            drivers: HashMap::new(),
515        };
516        registry
517            .add(TestOidcDriver)
518            .expect("test driver should register");
519        let mut provider = oidc_provider_config();
520        provider.protocol = crate::types::ExternalAuthProtocol::OAuth2;
521
522        let error = registry
523            .validate_provider_config(&provider)
524            .expect_err("protocol mismatch should fail");
525
526        assert!(error.to_string().contains(
527            "external auth provider 'oidc' is configured with protocol 'oauth2' but driver expects 'oidc'"
528        ));
529    }
530
531    #[test]
532    fn registry_driver_for_provider_returns_validated_driver() {
533        let mut registry = ExternalAuthProviderRegistry {
534            drivers: HashMap::new(),
535        };
536        registry
537            .add(TestOidcDriver)
538            .expect("test driver should register");
539        let provider = oidc_provider_config();
540
541        let driver = registry
542            .driver_for_provider(&provider)
543            .expect("valid provider should resolve driver");
544
545        assert_eq!(driver.kind(), ExternalAuthProviderKind::Oidc);
546    }
547
548    #[test]
549    fn registry_with_external_registrations_exposes_configure_hook() {
550        let registry = ExternalAuthProviderRegistry::with_external_registrations(|registry| {
551            if registry.contains(ExternalAuthProviderKind::Oidc) {
552                Ok(())
553            } else {
554                registry.add(TestOidcDriver)
555            }
556        })
557        .expect("external registration hook should add driver");
558
559        assert!(registry.contains(ExternalAuthProviderKind::Oidc));
560    }
561
562    #[test]
563    fn default_registry_is_singleton() {
564        let first = default_registry() as *const ExternalAuthProviderRegistry;
565        let second = default_registry() as *const ExternalAuthProviderRegistry;
566
567        assert_eq!(first, second);
568    }
569}