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#[derive(Debug, Clone, PartialEq, Eq)]
12pub struct AccountTarget {
13 pub provider_id: String,
15 pub account_id: String,
17 pub label: Option<String>,
19 pub explicit: bool,
21}
22
23pub 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 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 pub fn register(&mut self, provider: Box<dyn UsageProvider>) {
71 let id = provider.metadata().id;
72 self.providers.insert(id, provider);
73 }
74
75 pub fn get(&self, id: &str) -> Option<&dyn UsageProvider> {
77 self.providers.get(id).map(|p| p.as_ref())
78 }
79
80 pub fn all(&self) -> Vec<&dyn UsageProvider> {
82 self.providers.values().map(|p| p.as_ref()).collect()
83 }
84
85 pub fn all_metadata(&self) -> Vec<&ProviderMetadata> {
87 self.providers.values().map(|p| p.metadata()).collect()
88 }
89
90 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 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 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 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 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 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 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 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#[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"))); reg.register(Box::new(MockProvider::new("auto-off"))); 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 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 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 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 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 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 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}