Skip to main content

usage_monitor_cli/provider/
registry.rs

1use std::collections::HashMap;
2
3use futures_util::future::join_all;
4
5use crate::config::{AppConfig, DEFAULT_ACCOUNT, ProviderState};
6use crate::error::SpendPanelError;
7
8use super::{ProviderContext, ProviderMetadata, UsageProvider};
9
10/// One unit of fetch work: a single account of a provider.
11#[derive(Debug, Clone, PartialEq, Eq)]
12pub struct AccountTarget {
13    /// Provider this account belongs to.
14    pub provider_id: String,
15    /// Account name (`"default"` for the implicit single account).
16    pub account_id: String,
17    /// Human-friendly label, when configured.
18    pub label: Option<String>,
19    /// Whether the account is explicitly configured (vs. the implicit default).
20    pub explicit: bool,
21}
22
23/// Registry of all available providers.
24pub struct ProviderRegistry {
25    providers: HashMap<&'static str, Box<dyn UsageProvider>>,
26}
27
28impl ProviderRegistry {
29    pub fn new() -> Self {
30        Self {
31            providers: HashMap::new(),
32        }
33    }
34
35    /// Registry with all built-in providers registered.
36    pub fn with_defaults() -> Self {
37        let mut reg = Self::new();
38        reg.register(Box::new(super::abacus::AbacusProvider::new()));
39        reg.register(Box::new(super::anthropic::AnthropicProvider::new()));
40        reg.register(Box::new(super::antigravity::AntigravityProvider::new()));
41        reg.register(Box::new(super::claude::ClaudeProvider::new()));
42        reg.register(Box::new(super::codex::CodexProvider::new()));
43        reg.register(Box::new(super::copilot::CopilotProvider::new()));
44        reg.register(Box::new(super::cursor::CursorProvider::new()));
45        reg.register(Box::new(super::deepseek::DeepSeekProvider::new()));
46        reg.register(Box::new(super::deepgram::DeepgramProvider::new()));
47        reg.register(Box::new(super::devin::DevinProvider::new()));
48        reg.register(Box::new(super::elevenlabs::ElevenLabsProvider::new()));
49        reg.register(Box::new(super::gemini::GeminiProvider::new()));
50        reg.register(Box::new(super::grok::GrokProvider::new()));
51        reg.register(Box::new(super::groq::GroqProvider::new()));
52        reg.register(Box::new(super::kimi::KimiProvider::new()));
53        reg.register(Box::new(super::kimik2::KimiK2Provider::new()));
54        reg.register(Box::new(super::llmproxy::LlmProxyProvider::new()));
55        reg.register(Box::new(super::minimax::MiniMaxProvider::new()));
56        reg.register(Box::new(super::mistral::MistralProvider::new()));
57        reg.register(Box::new(super::moonshot::MoonshotProvider::new()));
58        reg.register(Box::new(super::ollama::OllamaProvider::new()));
59        reg.register(Box::new(super::opencode_go::OpenCodeGoProvider::new()));
60        reg.register(Box::new(super::openai::OpenAIProvider::new()));
61        reg.register(Box::new(super::openrouter::OpenRouterProvider::new()));
62        reg.register(Box::new(super::perplexity::PerplexityProvider::new()));
63        reg.register(Box::new(super::venice::VeniceProvider::new()));
64        reg.register(Box::new(super::windsurf::WindsurfProvider::new()));
65        reg.register(Box::new(super::zai::ZaiProvider::new()));
66        reg
67    }
68
69    /// Registers a provider.
70    pub fn register(&mut self, provider: Box<dyn UsageProvider>) {
71        let id = provider.metadata().id;
72        self.providers.insert(id, provider);
73    }
74
75    /// Returns a provider by ID.
76    pub fn get(&self, id: &str) -> Option<&dyn UsageProvider> {
77        self.providers.get(id).map(|p| p.as_ref())
78    }
79
80    /// Lists all registered providers.
81    pub fn all(&self) -> Vec<&dyn UsageProvider> {
82        self.providers.values().map(|p| p.as_ref()).collect()
83    }
84
85    /// Returns metadata for all providers.
86    pub fn all_metadata(&self) -> Vec<&ProviderMetadata> {
87        self.providers.values().map(|p| p.metadata()).collect()
88    }
89
90    /// Fetches usage from a specific provider.
91    pub async fn fetch(
92        &self,
93        id: &str,
94        ctx: &ProviderContext,
95    ) -> Result<crate::model::UsageSnapshot, SpendPanelError> {
96        match self.get(id) {
97            Some(provider) => provider.fetch_usage(ctx).await,
98            None => Err(SpendPanelError::ProviderNotFound(id.to_string())),
99        }
100    }
101
102    /// Resolves the enablement state of a provider: explicit config toggle
103    /// wins, otherwise credential detection decides.
104    pub fn provider_state(&self, id: &str, config: &AppConfig) -> Option<ProviderState> {
105        let provider = self.get(id)?;
106        Some(config.resolve_state(id, provider.detect_credentials()))
107    }
108
109    /// IDs of all enabled providers (explicitly or by credential detection).
110    pub fn enabled_ids(&self, config: &AppConfig) -> Vec<String> {
111        let mut ids: Vec<String> = self
112            .all()
113            .iter()
114            .filter(|p| {
115                config
116                    .resolve_state(p.metadata().id, p.detect_credentials())
117                    .is_enabled()
118            })
119            .map(|p| p.metadata().id.to_string())
120            .collect();
121        ids.sort();
122        ids
123    }
124
125    /// Account targets for a single provider.
126    ///
127    /// Each configured (enabled) account becomes a target. In addition, the
128    /// implicit auto-detected `default` account is included — so a named
129    /// account lives *alongside* the auto-detected login rather than replacing
130    /// it — unless an explicit `default` account is configured. The implicit
131    /// default is added when the provider detects credentials, or when no other
132    /// account is configured (so a bare `fetch <provider>` still tries it).
133    ///
134    /// To drop the auto-detected default while keeping named accounts, disable
135    /// it: `<provider> account disable default`.
136    pub fn provider_targets(&self, id: &str, config: &AppConfig) -> Vec<AccountTarget> {
137        let mut targets: Vec<AccountTarget> = config
138            .account_ids(id)
139            .into_iter()
140            .filter(|acct| config.account_is_enabled(id, acct))
141            .map(|acct| AccountTarget {
142                label: config.account_label(id, &acct).map(str::to_string),
143                provider_id: id.to_string(),
144                account_id: acct,
145                explicit: true,
146            })
147            .collect();
148
149        // An explicit `default` account (even if disabled) takes over the
150        // default slot; otherwise add the implicit auto-detected one.
151        if config.account(id, DEFAULT_ACCOUNT).is_none() {
152            let detected = self.get(id).is_some_and(|p| p.detect_credentials());
153            if detected || targets.is_empty() {
154                targets.insert(
155                    0,
156                    AccountTarget {
157                        provider_id: id.to_string(),
158                        account_id: DEFAULT_ACCOUNT.to_string(),
159                        label: None,
160                        explicit: false,
161                    },
162                );
163            }
164        }
165        targets
166    }
167
168    /// All account targets across enabled providers.
169    pub fn enabled_targets(&self, config: &AppConfig) -> Vec<AccountTarget> {
170        self.enabled_ids(config)
171            .into_iter()
172            .flat_map(|id| self.provider_targets(&id, config))
173            .collect()
174    }
175
176    /// Fetches a list of account targets concurrently. `ctx_for` builds the
177    /// fetch context for each target. Successful snapshots are stamped with the
178    /// target's account id/label.
179    pub async fn fetch_targets<F>(
180        &self,
181        targets: Vec<AccountTarget>,
182        ctx_for: F,
183    ) -> Vec<(
184        AccountTarget,
185        Result<crate::model::UsageSnapshot, SpendPanelError>,
186    )>
187    where
188        F: Fn(&AccountTarget) -> ProviderContext,
189    {
190        let fetches = targets.into_iter().map(|target| {
191            let ctx = ctx_for(&target);
192            async move {
193                let mut result = self.fetch(&target.provider_id, &ctx).await;
194                if let Ok(snapshot) = &mut result {
195                    if target.explicit {
196                        snapshot.account_id = Some(target.account_id.clone());
197                    }
198                    snapshot.account_label = target.label.clone();
199                }
200                (target, result)
201            }
202        });
203        join_all(fetches).await
204    }
205
206    /// Fetches usage from all registered providers concurrently.
207    pub async fn fetch_all(
208        &self,
209        ctx_overrides: Option<&HashMap<String, ProviderContext>>,
210    ) -> Vec<(String, Result<crate::model::UsageSnapshot, SpendPanelError>)> {
211        let fetches = self.all().into_iter().map(|provider| {
212            let id = provider.metadata().id.to_string();
213            let ctx = ctx_overrides
214                .and_then(|o| o.get(id.as_str()))
215                .cloned()
216                .unwrap_or_default();
217            async move {
218                let result = provider.fetch_usage(&ctx).await;
219                (id, result)
220            }
221        });
222        join_all(fetches).await
223    }
224}
225
226impl Default for ProviderRegistry {
227    fn default() -> Self {
228        Self::new()
229    }
230}
231
232// Default context
233#[cfg(test)]
234mod tests {
235    use super::*;
236    use crate::model::UsageSnapshot;
237    use async_trait::async_trait;
238
239    struct MockProvider {
240        meta: ProviderMetadata,
241        should_fail: bool,
242    }
243
244    impl MockProvider {
245        fn new(id: &'static str) -> Self {
246            Self {
247                meta: ProviderMetadata {
248                    id,
249                    name: id,
250                    description: "mock",
251                    auth_methods: &["mock"],
252                    website: None,
253                },
254                should_fail: false,
255            }
256        }
257
258        fn failing(id: &'static str) -> Self {
259            Self {
260                meta: ProviderMetadata {
261                    id,
262                    name: id,
263                    description: "mock",
264                    auth_methods: &["mock"],
265                    website: None,
266                },
267                should_fail: true,
268            }
269        }
270    }
271
272    #[async_trait]
273    impl UsageProvider for MockProvider {
274        fn metadata(&self) -> &ProviderMetadata {
275            &self.meta
276        }
277
278        async fn fetch_usage(
279            &self,
280            _ctx: &ProviderContext,
281        ) -> Result<UsageSnapshot, SpendPanelError> {
282            if self.should_fail {
283                Err(SpendPanelError::ProviderError(
284                    self.id().into(),
285                    "mock fail".into(),
286                ))
287            } else {
288                Ok(UsageSnapshot::new(self.id()))
289            }
290        }
291    }
292
293    impl MockProvider {
294        fn id(&self) -> &'static str {
295            self.meta.id
296        }
297    }
298
299    #[test]
300    fn test_registry_new() {
301        let reg = ProviderRegistry::new();
302        assert!(reg.all().is_empty());
303    }
304
305    #[test]
306    fn test_registry_register_and_get() {
307        let mut reg = ProviderRegistry::new();
308        reg.register(Box::new(MockProvider::new("mock-provider")));
309
310        assert!(reg.get("mock-provider").is_some());
311        assert!(reg.get("nonexistent").is_none());
312    }
313
314    #[test]
315    fn test_registry_all_metadata() {
316        let mut reg = ProviderRegistry::new();
317        reg.register(Box::new(MockProvider::new("p1")));
318        reg.register(Box::new(MockProvider::new("p2")));
319
320        let meta = reg.all_metadata();
321        assert_eq!(meta.len(), 2);
322        let ids: Vec<&str> = meta.iter().map(|m| m.id).collect();
323        assert!(ids.contains(&"p1"));
324        assert!(ids.contains(&"p2"));
325    }
326
327    #[tokio::test]
328    async fn test_fetch_success() {
329        let mut reg = ProviderRegistry::new();
330        reg.register(Box::new(MockProvider::new("ok")));
331
332        let result = reg.fetch("ok", &ProviderContext::new()).await;
333        assert!(result.is_ok());
334        assert_eq!(result.unwrap().provider_id, "ok");
335    }
336
337    #[tokio::test]
338    async fn test_fetch_not_found() {
339        let reg = ProviderRegistry::new();
340        let result = reg.fetch("ghost", &ProviderContext::new()).await;
341        assert!(matches!(result, Err(SpendPanelError::ProviderNotFound(_))));
342    }
343
344    #[tokio::test]
345    async fn test_fetch_failure() {
346        let mut reg = ProviderRegistry::new();
347        reg.register(Box::new(MockProvider::failing("bad")));
348
349        let result = reg.fetch("bad", &ProviderContext::new()).await;
350        assert!(result.is_err());
351    }
352
353    struct DetectableProvider {
354        meta: ProviderMetadata,
355    }
356
357    struct DelayedProvider {
358        meta: ProviderMetadata,
359        delay: std::time::Duration,
360    }
361
362    #[async_trait]
363    impl UsageProvider for DetectableProvider {
364        fn metadata(&self) -> &ProviderMetadata {
365            &self.meta
366        }
367
368        fn detect_credentials(&self) -> bool {
369            true
370        }
371
372        async fn fetch_usage(
373            &self,
374            _ctx: &ProviderContext,
375        ) -> Result<UsageSnapshot, SpendPanelError> {
376            Ok(UsageSnapshot::new(self.meta.id))
377        }
378    }
379
380    fn detectable(id: &'static str) -> DetectableProvider {
381        DetectableProvider {
382            meta: ProviderMetadata {
383                id,
384                name: id,
385                description: "mock",
386                auth_methods: &["mock"],
387                website: None,
388            },
389        }
390    }
391
392    fn delayed(id: &'static str, delay: std::time::Duration) -> DelayedProvider {
393        DelayedProvider {
394            meta: ProviderMetadata {
395                id,
396                name: id,
397                description: "delayed",
398                auth_methods: &["mock"],
399                website: None,
400            },
401            delay,
402        }
403    }
404
405    #[async_trait]
406    impl UsageProvider for DelayedProvider {
407        fn metadata(&self) -> &ProviderMetadata {
408            &self.meta
409        }
410
411        fn detect_credentials(&self) -> bool {
412            true
413        }
414
415        async fn fetch_usage(
416            &self,
417            _ctx: &ProviderContext,
418        ) -> Result<UsageSnapshot, SpendPanelError> {
419            tokio::time::sleep(self.delay).await;
420            Ok(UsageSnapshot::new(self.meta.id))
421        }
422    }
423
424    #[test]
425    fn test_provider_state_and_enabled_ids() {
426        use crate::config::{AppConfig, ProviderState};
427
428        let mut reg = ProviderRegistry::new();
429        reg.register(Box::new(detectable("auto-on"))); // detect = true
430        reg.register(Box::new(MockProvider::new("auto-off"))); // detect = false
431        reg.register(Box::new(MockProvider::new("forced-on")));
432        reg.register(Box::new(detectable("forced-off")));
433
434        let mut cfg = AppConfig::default();
435        cfg.set_provider_enabled("forced-on", true);
436        cfg.set_provider_enabled("forced-off", false);
437
438        assert_eq!(
439            reg.provider_state("auto-on", &cfg),
440            Some(ProviderState::AutoEnabled)
441        );
442        assert_eq!(
443            reg.provider_state("auto-off", &cfg),
444            Some(ProviderState::AutoDisabled)
445        );
446        assert_eq!(
447            reg.provider_state("forced-on", &cfg),
448            Some(ProviderState::Enabled)
449        );
450        assert_eq!(
451            reg.provider_state("forced-off", &cfg),
452            Some(ProviderState::Disabled)
453        );
454        assert_eq!(reg.provider_state("ghost", &cfg), None);
455
456        assert_eq!(reg.enabled_ids(&cfg), vec!["auto-on", "forced-on"]);
457    }
458
459    #[tokio::test]
460    async fn test_enabled_targets_skips_disabled() {
461        use crate::config::AppConfig;
462
463        let mut reg = ProviderRegistry::new();
464        reg.register(Box::new(detectable("on")));
465        reg.register(Box::new(detectable("off")));
466
467        let mut cfg = AppConfig::default();
468        cfg.set_provider_enabled("off", false);
469
470        let targets = reg.enabled_targets(&cfg);
471        assert_eq!(targets.len(), 1);
472        assert_eq!(targets[0].provider_id, "on");
473        assert_eq!(targets[0].account_id, "default");
474        assert!(!targets[0].explicit);
475
476        let results = reg.fetch_targets(targets, |_| ProviderContext::new()).await;
477        assert_eq!(results.len(), 1);
478        assert!(results[0].1.is_ok());
479    }
480
481    #[tokio::test]
482    async fn test_provider_targets_expand_accounts() {
483        use crate::config::AppConfig;
484
485        // Non-detecting provider: no implicit default is added.
486        let mut reg = ProviderRegistry::new();
487        reg.register(Box::new(MockProvider::new("p")));
488
489        let mut cfg = AppConfig::default();
490        cfg.set_account_label("p", "work", "Work");
491        cfg.set_account_config("p", "home", "api_key", "x");
492        cfg.set_account_enabled("p", "home", false);
493
494        let targets = reg.provider_targets("p", &cfg);
495        // Only the enabled "work" account remains (no creds → no auto default).
496        assert_eq!(targets.len(), 1);
497        assert_eq!(targets[0].account_id, "work");
498        assert_eq!(targets[0].label.as_deref(), Some("Work"));
499        assert!(targets[0].explicit);
500
501        let results = reg.fetch_targets(targets, |_| ProviderContext::new()).await;
502        let snap = results[0].1.as_ref().unwrap();
503        assert_eq!(snap.account_id.as_deref(), Some("work"));
504        assert_eq!(snap.account_label.as_deref(), Some("Work"));
505    }
506
507    #[test]
508    fn test_auto_default_coexists_with_named_accounts() {
509        use crate::config::AppConfig;
510
511        // Detecting provider: implicit default is added alongside named accounts.
512        let mut reg = ProviderRegistry::new();
513        reg.register(Box::new(detectable("p")));
514
515        let mut cfg = AppConfig::default();
516        cfg.set_account_config("p", "work", "credentials_path", "/tmp/w.json");
517
518        let targets = reg.provider_targets("p", &cfg);
519        assert_eq!(targets.len(), 2);
520        // Auto default comes first, then the named account.
521        assert_eq!(targets[0].account_id, "default");
522        assert!(!targets[0].explicit);
523        assert_eq!(targets[1].account_id, "work");
524        assert!(targets[1].explicit);
525    }
526
527    #[test]
528    fn test_explicit_default_account_replaces_auto() {
529        use crate::config::AppConfig;
530
531        let mut reg = ProviderRegistry::new();
532        reg.register(Box::new(detectable("p")));
533
534        // An explicit `default` account takes over the default slot — no
535        // duplicate implicit target.
536        let mut cfg = AppConfig::default();
537        cfg.set_account_config("p", "default", "credentials_path", "/tmp/d.json");
538        cfg.set_account_config("p", "work", "credentials_path", "/tmp/w.json");
539
540        let targets = reg.provider_targets("p", &cfg);
541        assert_eq!(targets.len(), 2, "no duplicate default");
542        assert!(targets.iter().all(|t| t.explicit));
543
544        // Disabling the explicit default drops it without re-adding the auto one.
545        cfg.set_account_enabled("p", "default", false);
546        let targets = reg.provider_targets("p", &cfg);
547        assert_eq!(targets.len(), 1);
548        assert_eq!(targets[0].account_id, "work");
549    }
550
551    #[tokio::test]
552    async fn test_fetch_targets_runs_concurrently() {
553        let mut reg = ProviderRegistry::new();
554        let delay = std::time::Duration::from_millis(250);
555        reg.register(Box::new(delayed("slow-a", delay)));
556        reg.register(Box::new(delayed("slow-b", delay)));
557
558        let start = std::time::Instant::now();
559        let targets = reg.enabled_targets(&AppConfig::default());
560        let results = reg.fetch_targets(targets, |_| ProviderContext::new()).await;
561        let elapsed = start.elapsed();
562
563        assert_eq!(results.len(), 2);
564        assert!(results.iter().all(|(_, result)| result.is_ok()));
565        assert!(
566            elapsed < std::time::Duration::from_millis(450),
567            "fetch_targets should run concurrently; took {:?}",
568            elapsed
569        );
570    }
571
572    #[tokio::test]
573    async fn test_fetch_all() {
574        let mut reg = ProviderRegistry::new();
575        reg.register(Box::new(MockProvider::new("ok")));
576        reg.register(Box::new(MockProvider::failing("bad")));
577
578        let results = reg.fetch_all(None).await;
579        assert_eq!(results.len(), 2);
580
581        let ok_result = results.iter().find(|(id, _)| id == "ok").unwrap();
582        assert!(ok_result.1.is_ok());
583
584        let bad_result = results.iter().find(|(id, _)| id == "bad").unwrap();
585        assert!(bad_result.1.is_err());
586    }
587}