aster_forge_external_auth/
registry.rs1use 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
30pub struct ExternalAuthProviderRegistry {
38 drivers: HashMap<ExternalAuthProviderKind, Arc<dyn ExternalAuthProviderDriver>>,
39}
40
41impl ExternalAuthProviderRegistry {
42 pub fn empty() -> Self {
48 Self {
49 drivers: HashMap::new(),
50 }
51 }
52
53 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 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 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 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 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 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 pub fn supported_kinds(&self) -> impl Iterator<Item = ExternalAuthProviderKind> + '_ {
135 self.drivers.keys().copied()
136 }
137
138 pub fn contains(&self, kind: ExternalAuthProviderKind) -> bool {
140 self.drivers.contains_key(&kind)
141 }
142
143 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 pub fn descriptor_for(
156 &self,
157 kind: ExternalAuthProviderKind,
158 ) -> Result<ExternalAuthProviderDescriptor> {
159 Ok(self.get_driver(kind)?.descriptor())
160 }
161
162 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 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 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 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 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
266pub 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}