diff --git a/apps/desktop-tauri/src/components/providers/icons/ProviderIcon-deepinfra.svg b/apps/desktop-tauri/src/components/providers/icons/ProviderIcon-deepinfra.svg new file mode 100644 index 0000000000..b181a117d8 --- /dev/null +++ b/apps/desktop-tauri/src/components/providers/icons/ProviderIcon-deepinfra.svg @@ -0,0 +1,4 @@ + + DeepInfra + + diff --git a/apps/desktop-tauri/src/components/providers/providerIcons.ts b/apps/desktop-tauri/src/components/providers/providerIcons.ts index 6f6b31e684..4252a1d68e 100644 --- a/apps/desktop-tauri/src/components/providers/providerIcons.ts +++ b/apps/desktop-tauri/src/components/providers/providerIcons.ts @@ -17,6 +17,7 @@ import crof from "./icons/ProviderIcon-crof.svg?raw"; import crossmodel from "./icons/ProviderIcon-crossmodel.svg?raw"; import cursor from "./icons/ProviderIcon-cursor.svg?raw"; import deepgram from "./icons/ProviderIcon-deepgram.svg?raw"; +import deepinfra from "./icons/ProviderIcon-deepinfra.svg?raw"; import deepseek from "./icons/ProviderIcon-deepseek.svg?raw"; import doubao from "./icons/ProviderIcon-doubao.svg?raw"; import elevenlabs from "./icons/ProviderIcon-elevenlabs.svg?raw"; @@ -89,6 +90,7 @@ const RAW: Record = { crossmodel: tint(crossmodel), cursor: tint(cursor), deepgram: tint(deepgram), + deepinfra: tint(deepinfra), deepseek: tint(deepseek), doubao: tint(doubao), elevenlabs: tint(elevenlabs), @@ -139,6 +141,7 @@ export const PROVIDER_ICON_REGISTRY: Record = { copilot: { id: "copilot", brandColor: "#a855f7", fallbackLetter: "⬡", svgPath: RAW.copilot }, cursor: { id: "cursor", brandColor: "#00bfa5", fallbackLetter: "▸", svgPath: RAW.cursor }, deepgram: { id: "deepgram", brandColor: "#13ef93", fallbackLetter: "D", svgPath: RAW.deepgram }, + deepinfra: { id: "deepinfra", brandColor: "#2a3275", fallbackLetter: "D", svgPath: RAW.deepinfra }, deepseek: { id: "deepseek", brandColor: "#527df0", fallbackLetter: "D", svgPath: RAW.deepseek }, elevenlabs: { id: "elevenlabs", brandColor: "#111827", fallbackLetter: "E", svgPath: RAW.elevenlabs }, factory: { id: "factory", brandColor: "#ff6b35", fallbackLetter: "◎", svgPath: RAW.factory }, @@ -211,6 +214,9 @@ const ALIASES: Record = { manicode: "codebuff", "deep seek": "deepseek", "deep-seek": "deepseek", + "deep infra": "deepinfra", + "deep-infra": "deepinfra", + di: "deepinfra", codeium: "windsurf", "xiaomi mimo": "mimo", xiaomimimo: "mimo", diff --git a/apps/desktop-tauri/src/surfaces/TrayPanel.tsx b/apps/desktop-tauri/src/surfaces/TrayPanel.tsx index ea80c3b7ea..0a0e614524 100644 --- a/apps/desktop-tauri/src/surfaces/TrayPanel.tsx +++ b/apps/desktop-tauri/src/surfaces/TrayPanel.tsx @@ -38,7 +38,7 @@ import AgentSessions from "../components/AgentSessions"; const HAS_DASHBOARD = new Set([ "abacus", "alibaba", "alibabatokenplan", "amp", "augment", "azureopenai", "bedrock", "claude", "codex", "codebuff", - "commandcode", "copilot", "crof", "crossmodel", "cursor", "deepgram", "deepseek", + "commandcode", "copilot", "crof", "crossmodel", "cursor", "deepgram", "deepinfra", "deepseek", "doubao", "elevenlabs", "factory", "gemini", "grok", "groq", "infini", "jetbrains", "kilo", "kimi", "kimik2", "kiro", "manus", "mimo", "minimax", "mistral", "nanogpt", "ollama", "openaiapi", @@ -49,7 +49,7 @@ const HAS_DASHBOARD = new Set([ /** Provider IDs that have a status page URL in the backend */ const HAS_STATUS_PAGE = new Set([ "alibabatokenplan", "amp", "augment", "azureopenai", "bedrock", - "claude", "codex", "copilot", "deepgram", "deepseek", "elevenlabs", + "claude", "codex", "copilot", "deepgram", "deepinfra", "deepseek", "elevenlabs", "gemini", "grok", "groq", "kiro", "mistral", "openaiapi", "openrouter", "vertexai", "windsurf", ]); diff --git a/apps/desktop-tauri/src/surfaces/settings/tabs/ProvidersTab.tsx b/apps/desktop-tauri/src/surfaces/settings/tabs/ProvidersTab.tsx index 76d272fe90..cb4f05c312 100644 --- a/apps/desktop-tauri/src/surfaces/settings/tabs/ProvidersTab.tsx +++ b/apps/desktop-tauri/src/surfaces/settings/tabs/ProvidersTab.tsx @@ -229,6 +229,7 @@ function providerSourceHintShort( case "bedrock": case "nanogpt": case "warp": + case "deepinfra": case "doubao": case "crof": case "stepfun": diff --git a/apps/desktop-tauri/src/test/providerCatalog.ts b/apps/desktop-tauri/src/test/providerCatalog.ts index c20db8e493..972b0e4ba3 100644 --- a/apps/desktop-tauri/src/test/providerCatalog.ts +++ b/apps/desktop-tauri/src/test/providerCatalog.ts @@ -33,6 +33,7 @@ export const TEST_PROVIDER_CATALOG: Array<[string, string]> = [ ["bedrock", "AWS Bedrock"], ["codebuff", "Codebuff"], ["deepseek", "DeepSeek"], + ["deepinfra", "DeepInfra"], ["windsurf", "Windsurf"], ["manus", "Manus"], ["mimo", "Xiaomi MiMo"], diff --git a/apps/desktop-tauri/src/types/bridge.ts b/apps/desktop-tauri/src/types/bridge.ts index 6d5a0a01ad..13c3799f34 100644 --- a/apps/desktop-tauri/src/types/bridge.ts +++ b/apps/desktop-tauri/src/types/bridge.ts @@ -89,6 +89,7 @@ export type ProofProviderId = | "mistral" | "codebuff" | "deepseek" + | "deepinfra" | "windsurf" | "manus" | "mimo" diff --git a/rust/src/core/credential_migration.rs b/rust/src/core/credential_migration.rs index 73b2e65aba..32965bc0c5 100755 --- a/rust/src/core/credential_migration.rs +++ b/rust/src/core/credential_migration.rs @@ -340,6 +340,7 @@ const PROVIDER_ACCOUNT_NAMES: &[(ProviderId, &str)] = &[ (ProviderId::Bedrock, "bedrock-aws-credentials"), (ProviderId::Codebuff, "codebuff-api-token"), (ProviderId::DeepSeek, "deepseek-api-token"), + (ProviderId::DeepInfra, "deepinfra-api-token"), (ProviderId::Windsurf, "windsurf-local-cache"), (ProviderId::Manus, "manus-cookie"), (ProviderId::MiMo, "mimo-cookie"), diff --git a/rust/src/core/provider.rs b/rust/src/core/provider.rs index 63be3eeeda..18ad2cc9ad 100755 --- a/rust/src/core/provider.rs +++ b/rust/src/core/provider.rs @@ -45,6 +45,7 @@ pub enum ProviderId { Bedrock, Codebuff, DeepSeek, + DeepInfra, Windsurf, Manus, MiMo, @@ -109,6 +110,7 @@ impl ProviderId { ProviderId::Bedrock, ProviderId::Codebuff, ProviderId::DeepSeek, + ProviderId::DeepInfra, ProviderId::Windsurf, ProviderId::Manus, ProviderId::MiMo, @@ -173,6 +175,7 @@ impl ProviderId { ProviderId::Bedrock => "bedrock", ProviderId::Codebuff => "codebuff", ProviderId::DeepSeek => "deepseek", + ProviderId::DeepInfra => "deepinfra", ProviderId::Windsurf => "windsurf", ProviderId::Manus => "manus", ProviderId::MiMo => "mimo", @@ -237,6 +240,7 @@ impl ProviderId { ProviderId::Bedrock => "AWS Bedrock", ProviderId::Codebuff => "Codebuff", ProviderId::DeepSeek => "DeepSeek", + ProviderId::DeepInfra => "DeepInfra", ProviderId::Windsurf => "Windsurf", ProviderId::Manus => "Manus", ProviderId::MiMo => "Xiaomi MiMo", @@ -312,6 +316,7 @@ impl ProviderId { ProviderId::Bedrock => None, ProviderId::Codebuff => None, ProviderId::DeepSeek => None, + ProviderId::DeepInfra => None, ProviderId::Windsurf => None, ProviderId::Doubao => None, ProviderId::Crof => None, @@ -373,6 +378,7 @@ impl ProviderId { "bedrock" | "aws-bedrock" | "aws bedrock" => Some(ProviderId::Bedrock), "codebuff" | "manicode" => Some(ProviderId::Codebuff), "deepseek" | "deep-seek" | "ds" => Some(ProviderId::DeepSeek), + "deepinfra" | "deep-infra" | "di" => Some(ProviderId::DeepInfra), "windsurf" | "codeium" => Some(ProviderId::Windsurf), "manus" => Some(ProviderId::Manus), "mimo" | "xiaomi" | "xiaomimimo" | "xiaomi-mimo" | "xiaomi mimo" => { @@ -582,6 +588,8 @@ pub fn cli_name_map() -> HashMap<&'static str, ProviderId> { map.insert("manicode", ProviderId::Codebuff); map.insert("deep-seek", ProviderId::DeepSeek); map.insert("ds", ProviderId::DeepSeek); + map.insert("deep-infra", ProviderId::DeepInfra); + map.insert("di", ProviderId::DeepInfra); map.insert("codeium", ProviderId::Windsurf); map.insert("google", ProviderId::Gemini); map.insert("agy", ProviderId::Antigravity); @@ -634,7 +642,7 @@ mod tests { #[test] fn test_provider_id_all() { let all = ProviderId::all(); - assert_eq!(all.len(), 58); + assert_eq!(all.len(), 59); assert!(all.contains(&ProviderId::Claude)); assert!(all.contains(&ProviderId::Codex)); assert!(all.contains(&ProviderId::Kimi)); @@ -649,6 +657,7 @@ mod tests { assert!(all.contains(&ProviderId::Bedrock)); assert!(all.contains(&ProviderId::Codebuff)); assert!(all.contains(&ProviderId::DeepSeek)); + assert!(all.contains(&ProviderId::DeepInfra)); assert!(all.contains(&ProviderId::Windsurf)); assert!(all.contains(&ProviderId::Manus)); assert!(all.contains(&ProviderId::MiMo)); diff --git a/rust/src/core/provider_factory.rs b/rust/src/core/provider_factory.rs index e57669947e..84514ea40c 100644 --- a/rust/src/core/provider_factory.rs +++ b/rust/src/core/provider_factory.rs @@ -10,7 +10,7 @@ use crate::providers::{ AbacusProvider, AlibabaProvider, AlibabaTokenPlanProvider, AmpProvider, AntigravityProvider, AugmentProvider, AzureOpenAIProvider, BedrockProvider, ChutesProvider, ClaudeProvider, CodebuffProvider, CodexProvider, CommandCodeProvider, CopilotProvider, CrofProvider, - CrossModelProvider, CursorProvider, DeepSeekProvider, DeepgramProvider, DevinProvider, + CrossModelProvider, CursorProvider, DeepInfraProvider, DeepSeekProvider, DeepgramProvider, DevinProvider, DoubaoProvider, ElevenLabsProvider, FactoryProvider, GeminiProvider, GrokProvider, GroqProvider, InfiniProvider, JetBrainsProvider, KiloProvider, KimiK2Provider, KimiProvider, KiroProvider, LLMProxyProvider, LiteLLMProvider, ManusProvider, MiMoProvider, MiniMaxProvider, @@ -60,6 +60,7 @@ pub fn instantiate(id: ProviderId) -> Box { ProviderId::Bedrock => Box::new(BedrockProvider::new()), ProviderId::Codebuff => Box::new(CodebuffProvider::new()), ProviderId::DeepSeek => Box::new(DeepSeekProvider::new()), + ProviderId::DeepInfra => Box::new(DeepInfraProvider::new()), ProviderId::Windsurf => Box::new(WindsurfProvider::new()), ProviderId::Manus => Box::new(ManusProvider::new()), ProviderId::MiMo => Box::new(MiMoProvider::new()), diff --git a/rust/src/core/token_accounts.rs b/rust/src/core/token_accounts.rs index b29a4fd851..fa9a4ac74a 100755 --- a/rust/src/core/token_accounts.rs +++ b/rust/src/core/token_accounts.rs @@ -199,6 +199,16 @@ impl TokenAccountSupport { requires_manual_cookie_source: false, cookie_name: None, }), + ProviderId::DeepInfra => Some(TokenAccountSupport { + title: "API keys", + subtitle: "Store multiple DeepInfra API keys.", + placeholder: "API key from deepinfra.com/dash", + injection: TokenInjection::Environment { + key: "DEEPINFRA_API_KEY".to_string(), + }, + requires_manual_cookie_source: false, + cookie_name: None, + }), ProviderId::Copilot => Some(TokenAccountSupport { title: "GitHub accounts", subtitle: "Store GitHub OAuth tokens for Copilot plan usage.", diff --git a/rust/src/providers/deepinfra/mod.rs b/rust/src/providers/deepinfra/mod.rs new file mode 100644 index 0000000000..45a9633e5e --- /dev/null +++ b/rust/src/providers/deepinfra/mod.rs @@ -0,0 +1,426 @@ +//! DeepInfra provider implementation. +//! +//! Fetches prepaid balance and monthly spend from DeepInfra billing APIs: +//! - `GET https://api.deepinfra.com/payment/checklist?compute_owed=true` +//! - `GET https://api.deepinfra.com/payment/usage?from=current` +//! +//! Ported from steipete/CodexBar `DeepInfraUsageFetcher`. + +use async_trait::async_trait; +use reqwest::Client; +use serde::Deserialize; + +use crate::core::{ + CostSnapshot, FetchContext, Provider, ProviderError, ProviderFetchResult, ProviderId, + ProviderMetadata, RateWindow, SourceMode, UsageSnapshot, +}; + +const CHECKLIST_URL: &str = "https://api.deepinfra.com/payment/checklist?compute_owed=true"; +const USAGE_URL: &str = "https://api.deepinfra.com/payment/usage?from=current"; +const CREDENTIAL_TARGET: &str = "codexbar-deepinfra"; +const CENTS_PER_DOLLAR: f64 = 100.0; +const ENV_KEYS: &[&str] = &["DEEPINFRA_API_KEY", "DEEPINFRA_TOKEN"]; + +/// Checklist monetary fields are USD. Negative `stripe_balance` means prepaid funds. +#[derive(Debug, Deserialize, Clone)] +struct ChecklistResponse { + stripe_balance: f64, + recent: f64, + limit: Option, + #[serde(default)] + suspended: bool, + suspend_reason: Option, +} + +/// Usage endpoint reports `total_cost` in cents. +#[derive(Debug, Deserialize, Clone)] +struct UsageResponse { + months: Vec, + #[serde(default)] + initial_month: Option, +} + +#[derive(Debug, Deserialize, Clone)] +struct UsageMonth { + #[allow(dead_code)] + period: String, + /// Cost in cents (upstream field name is `total_cost`). + total_cost: f64, +} + +#[derive(Debug, Clone, PartialEq)] +struct DeepInfraSnapshot { + available_balance_usd: f64, + amount_owed_usd: f64, + current_month_cost_usd: f64, + recent_cost_usd: f64, + spending_limit_usd: Option, + suspended: bool, + suspend_reason: Option, +} + +impl DeepInfraSnapshot { + fn from_responses(checklist: &ChecklistResponse, usage: &UsageResponse) -> Self { + let recent_cost = checklist.recent.max(0.0); + let current_month_cost = usage + .months + .last() + .map(|m| (m.total_cost / CENTS_PER_DOLLAR).max(0.0)) + .unwrap_or(recent_cost); + let net_balance = checklist.stripe_balance + recent_cost; + let spending_limit = checklist + .limit + .and_then(|limit| (limit > 0.0).then_some(limit)); + + Self { + available_balance_usd: (-net_balance).max(0.0), + amount_owed_usd: net_balance.max(0.0), + current_month_cost_usd: current_month_cost, + recent_cost_usd: recent_cost, + spending_limit_usd: spending_limit, + suspended: checklist.suspended, + suspend_reason: checklist.suspend_reason.clone(), + } + } + + fn to_usage_snapshot(&self) -> UsageSnapshot { + let used_percent = + if self.suspended || self.amount_owed_usd > 0.0 || self.available_balance_usd <= 0.0 { + 100.0 + } else { + 0.0 + }; + + let balance_text = if self.amount_owed_usd > 0.0 { + format!("${:.2} owed", self.amount_owed_usd) + } else { + format!("${:.2} available", self.available_balance_usd) + }; + let spending_text = format!("${:.2} spent this month", self.current_month_cost_usd); + let detail = if self.suspended { + let reason = self + .suspend_reason + .as_deref() + .map(str::trim) + .filter(|s| !s.is_empty()); + match reason { + Some(reason) => format!("Suspended: {reason} · {balance_text} · {spending_text}"), + None => format!("Suspended · {balance_text} · {spending_text}"), + } + } else { + format!("{balance_text} · {spending_text}") + }; + + let mut primary = RateWindow::new(used_percent); + primary.reset_description = Some(detail); + + UsageSnapshot::new(primary).with_login_method(balance_text) + } + + fn to_cost_snapshot(&self) -> Option { + self.spending_limit_usd.map(|limit| { + CostSnapshot::new(self.recent_cost_usd, "USD", "Billing cycle").with_limit(limit) + }) + } +} + +pub struct DeepInfraProvider { + metadata: ProviderMetadata, + client: Client, +} + +impl DeepInfraProvider { + pub fn new() -> Self { + Self { + metadata: ProviderMetadata { + id: ProviderId::DeepInfra, + display_name: "DeepInfra", + session_label: "Balance", + weekly_label: "Balance", + supports_opus: false, + // Upstream marks supportsCredits=false; balance is shown via primary window text. + supports_credits: false, + default_enabled: false, + is_primary: false, + dashboard_url: Some("https://deepinfra.com/dash"), + status_page_url: Some("https://status.deepinfra.com"), + }, + client: crate::core::credentialed_http_client_builder() + .timeout(std::time::Duration::from_secs(30)) + .build() + .unwrap_or_else(|_| Client::new()), + } + } + + fn resolve_api_key(api_key: Option<&str>) -> Result { + let raw = crate::providers::resolve_api_key(api_key, CREDENTIAL_TARGET, ENV_KEYS)?; + clean_api_key(&raw).ok_or_else(|| { + ProviderError::NotInstalled( + "DeepInfra API key not found. Set DEEPINFRA_API_KEY / DEEPINFRA_TOKEN or Preferences → Providers." + .to_string(), + ) + }) + } + + async fn fetch_usage_api( + &self, + ctx: &FetchContext, + ) -> Result { + let api_key = Self::resolve_api_key(ctx.api_key.as_deref())?; + let checklist = self + .fetch_json::(CHECKLIST_URL, &api_key) + .await?; + let usage = self.fetch_json::(USAGE_URL, &api_key).await?; + let snapshot = DeepInfraSnapshot::from_responses(&checklist, &usage); + + let mut result = ProviderFetchResult::new(snapshot.to_usage_snapshot(), "api"); + if let Some(cost) = snapshot.to_cost_snapshot() { + result = result.with_cost(cost); + } + Ok(result) + } + + async fn fetch_json( + &self, + url: &str, + api_key: &str, + ) -> Result { + let resp = self + .client + .get(url) + .header("Authorization", format!("Bearer {api_key}")) + .header("Accept", "application/json") + .send() + .await?; + + let status = resp.status(); + if status == reqwest::StatusCode::UNAUTHORIZED { + return Err(ProviderError::Other( + "DeepInfra API key rejected (HTTP 401).".to_string(), + )); + } + if status == reqwest::StatusCode::FORBIDDEN { + return Err(ProviderError::Other( + "DeepInfra API key cannot access billing data (HTTP 403).".to_string(), + )); + } + if !status.is_success() { + return Err(ProviderError::Other(format!( + "DeepInfra API error: HTTP {status}" + ))); + } + + resp.json().await.map_err(|e| { + ProviderError::Parse(format!("Failed to parse DeepInfra response: {e}")) + }) + } +} + +impl Default for DeepInfraProvider { + fn default() -> Self { + Self::new() + } +} + +#[async_trait] +impl Provider for DeepInfraProvider { + fn id(&self) -> ProviderId { + ProviderId::DeepInfra + } + + fn metadata(&self) -> &ProviderMetadata { + &self.metadata + } + + async fn fetch_usage(&self, ctx: &FetchContext) -> Result { + match ctx.source_mode { + SourceMode::Auto | SourceMode::OAuth => self.fetch_usage_api(ctx).await, + SourceMode::Web | SourceMode::Cli => { + Err(ProviderError::UnsupportedSource(ctx.source_mode)) + } + } + } + + fn available_sources(&self) -> Vec { + vec![SourceMode::Auto, SourceMode::OAuth] + } +} + +fn clean_api_key(raw: &str) -> Option { + let mut value = raw.trim().to_string(); + if (value.starts_with('"') && value.ends_with('"')) + || (value.starts_with('\'') && value.ends_with('\'')) + { + value = value[1..value.len() - 1].trim().to_string(); + } + if let Some(stripped) = value + .strip_prefix("Bearer ") + .or_else(|| value.strip_prefix("bearer ")) + { + value = stripped.trim().to_string(); + } + (!value.is_empty()).then_some(value) +} + +/// Parse fixture JSON without network (used by unit tests). +fn parse_snapshot_for_testing( + checklist_json: &str, + usage_json: &str, +) -> Result { + let checklist: ChecklistResponse = serde_json::from_str(checklist_json) + .map_err(|e| ProviderError::Parse(format!("Failed to parse DeepInfra checklist: {e}")))?; + let usage: UsageResponse = serde_json::from_str(usage_json) + .map_err(|e| ProviderError::Parse(format!("Failed to parse DeepInfra usage: {e}")))?; + let _ = usage.initial_month.as_ref(); + Ok(DeepInfraSnapshot::from_responses(&checklist, &usage)) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn checklist_json( + stripe_balance: f64, + recent: f64, + limit: Option, + suspended: bool, + suspend_reason: Option<&str>, + ) -> String { + let limit_json = match limit { + Some(v) => v.to_string(), + None => "null".to_string(), + }; + let reason_json = match suspend_reason { + Some(r) => format!("\"{r}\""), + None => "null".to_string(), + }; + format!( + r#"{{ + "stripe_balance": {stripe_balance}, + "recent": {recent}, + "limit": {limit_json}, + "suspended": {suspended}, + "suspend_reason": {reason_json} + }}"# + ) + } + + fn usage_json(total_cost_cents: f64) -> String { + format!( + r#"{{ + "months": [ + {{ + "period": "2026.07", + "items": [], + "total_cost": {total_cost_cents} + }} + ], + "initial_month": "2026.07" + }}"# + ) + } + + #[test] + fn converts_monthly_cents_and_deducts_recent_usage_from_prepaid_balance() { + let snapshot = parse_snapshot_for_testing( + &checklist_json(-99.75, 3.94, Some(20.0), false, None), + &usage_json(394.0), + ) + .unwrap(); + + assert!((snapshot.available_balance_usd - 95.81).abs() < 1e-6); + assert_eq!(snapshot.amount_owed_usd, 0.0); + assert!((snapshot.current_month_cost_usd - 3.94).abs() < 1e-6); + assert_eq!(snapshot.recent_cost_usd, 3.94); + assert_eq!(snapshot.spending_limit_usd, Some(20.0)); + + let usage = snapshot.to_usage_snapshot(); + assert_eq!(usage.primary.used_percent, 0.0); + assert_eq!( + usage.primary.reset_description.as_deref(), + Some("$95.81 available · $3.94 spent this month") + ); + + let cost = snapshot.to_cost_snapshot().unwrap(); + assert_eq!(cost.used, 3.94); + assert_eq!(cost.limit, Some(20.0)); + assert_eq!(cost.period, "Billing cycle"); + } + + #[test] + fn positive_stripe_balance_is_reported_as_amount_owed() { + let snapshot = parse_snapshot_for_testing( + &checklist_json(2.75, 7.0, Some(-1.0), false, None), + &usage_json(650.0), + ) + .unwrap(); + + assert_eq!(snapshot.available_balance_usd, 0.0); + assert_eq!(snapshot.amount_owed_usd, 9.75); + assert_eq!(snapshot.spending_limit_usd, None); + + let usage = snapshot.to_usage_snapshot(); + assert_eq!(usage.primary.used_percent, 100.0); + assert_eq!( + usage.primary.reset_description.as_deref(), + Some("$9.75 owed · $6.50 spent this month") + ); + assert!(snapshot.to_cost_snapshot().is_none()); + } + + #[test] + fn suspended_account_is_marked_exhausted() { + let snapshot = parse_snapshot_for_testing( + &checklist_json(-5.0, 1.0, None, true, Some("Payment review")), + &usage_json(100.0), + ) + .unwrap() + .to_usage_snapshot(); + + assert_eq!(snapshot.primary.used_percent, 100.0); + assert!( + snapshot + .primary + .reset_description + .as_deref() + .unwrap_or("") + .starts_with("Suspended: Payment review") + ); + } + + #[test] + fn rejects_malformed_billing_response() { + let err = parse_snapshot_for_testing("{}", &usage_json(100.0)).unwrap_err(); + match err { + ProviderError::Parse(msg) => assert!(msg.contains("checklist")), + other => panic!("expected parse error, got {other:?}"), + } + } + + #[test] + fn cleans_quoted_and_bearer_prefixed_keys() { + assert_eq!( + clean_api_key(" \"Bearer sk-test\" ").as_deref(), + Some("sk-test") + ); + assert_eq!(clean_api_key("bearer sk-abc").as_deref(), Some("sk-abc")); + assert_eq!(clean_api_key(" ").as_deref(), None); + } + + #[test] + fn metadata_matches_upstream_descriptor() { + let provider = DeepInfraProvider::new(); + assert_eq!(provider.id(), ProviderId::DeepInfra); + assert_eq!(provider.metadata().display_name, "DeepInfra"); + assert_eq!( + provider.metadata().dashboard_url, + Some("https://deepinfra.com/dash") + ); + assert_eq!( + provider.metadata().status_page_url, + Some("https://status.deepinfra.com") + ); + assert!(!provider.metadata().supports_credits); + assert!(!provider.metadata().default_enabled); + } +} diff --git a/rust/src/providers/mod.rs b/rust/src/providers/mod.rs index ed7ffd7300..c626404be8 100755 --- a/rust/src/providers/mod.rs +++ b/rust/src/providers/mod.rs @@ -20,6 +20,7 @@ pub mod crof; pub mod crossmodel; pub mod cursor; pub mod deepgram; +pub mod deepinfra; pub mod deepseek; pub mod devin; pub mod doubao; @@ -81,6 +82,7 @@ pub use crof::CrofProvider; pub use crossmodel::CrossModelProvider; pub use cursor::CursorProvider; pub use deepgram::DeepgramProvider; +pub use deepinfra::DeepInfraProvider; pub use deepseek::DeepSeekProvider; pub use devin::DevinProvider; pub use doubao::DoubaoProvider; diff --git a/rust/src/settings/api_keys.rs b/rust/src/settings/api_keys.rs index b76172521d..61c2e42719 100644 --- a/rust/src/settings/api_keys.rs +++ b/rust/src/settings/api_keys.rs @@ -299,6 +299,17 @@ pub fn get_api_key_providers() -> Vec { config_file_path: None, dashboard_url: Some("https://platform.deepseek.com/usage"), }, + ProviderConfigInfo { + id: ProviderId::DeepInfra, + name: "DeepInfra", + requires_api_key: true, + api_key_env_var: Some("DEEPINFRA_API_KEY"), + api_key_help: Some( + "Get your API key from deepinfra.com/dash. Also accepts DEEPINFRA_TOKEN.", + ), + config_file_path: None, + dashboard_url: Some("https://deepinfra.com/dash"), + }, ProviderConfigInfo { id: ProviderId::Doubao, name: "Doubao / Volcengine Ark", diff --git a/rust/src/settings/tests.rs b/rust/src/settings/tests.rs index 18dcf9e3fc..e0e22061de 100644 --- a/rust/src/settings/tests.rs +++ b/rust/src/settings/tests.rs @@ -374,6 +374,7 @@ fn test_api_key_provider_catalog_includes_token_providers() { ProviderId::Bedrock, ProviderId::Codebuff, ProviderId::DeepSeek, + ProviderId::DeepInfra, ProviderId::ElevenLabs, ProviderId::Deepgram, ProviderId::Grok,