diff --git a/rust/src/locale/en-US.ftl b/rust/src/locale/en-US.ftl index fb96ecd4b8..ec1baafc9b 100644 --- a/rust/src/locale/en-US.ftl +++ b/rust/src/locale/en-US.ftl @@ -784,7 +784,7 @@ OpenAiProjectIdHelp = Leave blank for organization-wide usage. Set a project ID LiteLlmApiTitle = LiteLLM API LiteLlmBaseUrlLabel = Base URL LiteLlmBaseUrlPlaceholder = https://litellm.example.com -LiteLlmBaseUrlHelp = Used with the saved API key for LiteLLM /key/info. +LiteLlmBaseUrlHelp = Used with the saved API key for LiteLLM key, user, and team info. Use HTTPS, or HTTP on a loopback or private-network address. DevinApiTitle = Devin API DevinOrganizationLabel = Organization DevinOrganizationPlaceholder = org/acme diff --git a/rust/src/providers/litellm/endpoint.rs b/rust/src/providers/litellm/endpoint.rs new file mode 100644 index 0000000000..b175ed137e --- /dev/null +++ b/rust/src/providers/litellm/endpoint.rs @@ -0,0 +1,172 @@ +//! LiteLLM base-URL policy and management-route URLs. +//! +//! Upstream `litellm.ts` declares the `LITELLM_BASE_URL` endpoint with the +//! `https-or-private-network-http` policy: HTTPS anywhere, plain HTTP only for +//! loopback, RFC 1918, link-local, IPv6 unique-local, and `.local` hosts. + +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; + +use reqwest::Url; + +use crate::core::ProviderError; + +const INVALID_BASE: &str = "LiteLLM base URL must use HTTPS, or HTTP on a loopback or private-network address, without embedded credentials."; + +/// Validate a LiteLLM base URL. A scheme-less value is treated as HTTPS. +pub(crate) fn validated_base_url(raw: &str) -> Result { + let trimmed = raw.trim(); + if trimmed.is_empty() { + return Err(ProviderError::Other("LiteLLM base URL is empty".into())); + } + let lower = trimmed.to_ascii_lowercase(); + if ["%2f", "%5c", "%3f", "%23", "%40", "%3a"] + .iter() + .any(|encoded| lower.contains(encoded)) + { + return Err(ProviderError::Other( + "LiteLLM base URL must not contain encoded host delimiters".into(), + )); + } + let candidate = if trimmed.contains("://") { + trimmed.to_string() + } else { + format!("https://{trimmed}") + }; + let url = Url::parse(&candidate) + .map_err(|e| ProviderError::Other(format!("Invalid LiteLLM base URL: {e}")))?; + let host = url + .host_str() + .ok_or_else(|| ProviderError::Other("LiteLLM base URL must include a host".into()))?; + let scheme_ok = match url.scheme() { + "https" => true, + "http" => is_private_network_host(host), + _ => false, + }; + if !scheme_ok + || !url.username().is_empty() + || url.password().is_some() + || host.contains('%') + || host.chars().any(|c| c.is_control() || c.is_whitespace()) + { + return Err(ProviderError::Other(INVALID_BASE.into())); + } + Ok(url) +} + +/// Build `{base}/{path}` for a management route. A trailing `/v1` on the base +/// is dropped, and the base path and query are otherwise preserved. `query` +/// replaces the base query when given. +pub(super) fn management_url( + base: &str, + path: &str, + query: Option<(&str, &str)>, +) -> Result { + let mut url = validated_base_url(base)?; + let trimmed = url.path().trim_end_matches('/'); + let root = trimmed.strip_suffix("/v1").unwrap_or(trimmed).to_string(); + url.set_path(&format!("{root}/{path}")); + url.set_fragment(None); + if let Some((key, value)) = query { + url.query_pairs_mut().clear().append_pair(key, value); + } + Ok(url) +} + +fn is_private_network_host(host: &str) -> bool { + let normalized = host.trim_end_matches('.').to_ascii_lowercase(); + if normalized == "localhost" + || normalized.ends_with(".localhost") + || normalized.ends_with(".local") + { + return true; + } + let ip_candidate = normalized + .strip_prefix('[') + .and_then(|value| value.strip_suffix(']')) + .unwrap_or(&normalized); + match ip_candidate.parse::() { + Ok(IpAddr::V4(ip)) => is_private_ipv4(ip), + Ok(IpAddr::V6(ip)) => is_private_ipv6(ip), + Err(_) => false, + } +} + +fn is_private_ipv4(ip: Ipv4Addr) -> bool { + ip.is_loopback() || ip.is_private() || ip.is_link_local() +} + +fn is_private_ipv6(ip: Ipv6Addr) -> bool { + if let Some(mapped) = ip.to_ipv4_mapped() { + return is_private_ipv4(mapped); + } + ip.is_loopback() + || (ip.segments()[0] & 0xfe00) == 0xfc00 + || (ip.segments()[0] & 0xffc0) == 0xfe80 +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn allows_https_anywhere_and_private_network_http() { + for value in [ + "https://litellm.example.com", + "litellm.example.com", + "http://localhost:4000", + "http://127.0.0.1:4000", + "http://[::1]:4000", + "http://10.1.2.3", + "http://172.16.0.9", + "http://192.168.1.20:4000", + "http://169.254.10.10", + "http://[fd12:3456::1]", + "http://[fe80::1]", + "http://proxy.local:4000", + ] { + assert!(validated_base_url(value).is_ok(), "rejected {value}"); + } + } + + #[test] + fn rejects_public_http_credentials_and_encoded_delimiters() { + for value in [ + "", + "http://litellm.example.com", + "http://8.8.8.8", + "http://172.32.0.1", + "http://[2001:db8::1]", + "http://example.com.evil.test", + "ftp://10.0.0.1", + "https://user:pass@litellm.example.com", + "http://user@10.0.0.1", + "https://example.com%2f.evil.test", + ] { + assert!(validated_base_url(value).is_err(), "accepted {value}"); + } + } + + #[test] + fn management_url_strips_v1_and_keeps_subpath() { + let url = management_url("https://h.example.com/litellm/v1/", "key/info", None).unwrap(); + assert_eq!(url.as_str(), "https://h.example.com/litellm/key/info"); + let url = management_url("http://10.0.0.2:4000/v1", "key/info", None).unwrap(); + assert_eq!(url.as_str(), "http://10.0.0.2:4000/key/info"); + } + + #[test] + fn management_url_encodes_query_and_replaces_base_query() { + let url = management_url( + "https://h.example.com?token=abc", + "user/info", + Some(("user_id", "a b&c")), + ) + .unwrap(); + assert_eq!( + url.as_str(), + "https://h.example.com/user/info?user_id=a+b%26c" + ); + let kept = management_url("https://h.example.com?token=abc", "key/info", None).unwrap(); + assert_eq!(kept.as_str(), "https://h.example.com/key/info?token=abc"); + } +} diff --git a/rust/src/providers/litellm/info.rs b/rust/src/providers/litellm/info.rs new file mode 100644 index 0000000000..c308c41ede --- /dev/null +++ b/rust/src/providers/litellm/info.rs @@ -0,0 +1,271 @@ +//! LiteLLM management-route payloads and their projection into a usage result. +//! +//! Wire shapes follow upstream `litellm.ts`: `/key/info` names the key's +//! `user_id` / `team_id`, then `/user/info` or `/team/info` supplies the +//! budgets. Returned IDs must match the key's IDs before anything is shown. + +use chrono::{DateTime, NaiveDateTime, Utc}; +use serde::Deserialize; +use serde_json::Value; + +use crate::core::{ + CostSnapshot, ProviderError, ProviderFetchResult, RateWindow, SubscriptionMetadata, + UsageSnapshot, +}; + +#[derive(Deserialize)] +pub(super) struct KeyInfoResponse { + info: KeyInfo, +} + +#[derive(Deserialize)] +struct KeyInfo { + user_id: Option, + team_id: Option, + expires: Option, +} + +#[derive(Deserialize)] +pub(super) struct UserInfoResponse { + user_id: Option, + user_info: UserInfo, + teams: Option>, +} + +#[derive(Deserialize)] +struct UserInfo { + user_id: Option, + user_email: Option, + user_alias: Option, + spend: Option, + max_budget: Option, + budget_reset_at: Option, + metadata: Option, +} + +#[derive(Deserialize)] +struct UserMetadata { + preferred_username: Option, +} + +#[derive(Deserialize)] +pub(super) struct TeamInfoResponse { + team_id: Option, + team_info: Budget, +} + +#[derive(Deserialize)] +struct Budget { + team_id: Option, + team_alias: Option, + spend: Option, + max_budget: Option, + budget_reset_at: Option, +} + +/// Identity and routing data read from `/key/info`. +pub(super) struct KeyBinding { + pub user_id: Option, + pub team_id: Option, + expires: Option>, +} + +/// A spend/budget pair with its optional reset instant. +struct Spend { + spend: f64, + limit: Option, + reset: Option>, +} + +impl Spend { + fn budget(&self) -> Option { + self.limit.filter(|limit| *limit > 0.0) + } + + fn window(&self, label: Option<&str>) -> Option { + let limit = self.budget()?; + let mut window = RateWindow::new(self.spend / limit * 100.0); + window.resets_at = self.reset; + let detail = format!("${:.2} / ${limit:.2}", self.spend); + window.reset_description = Some(match label { + Some(label) => format!("{label}: {detail}"), + None => detail, + }); + Some(window) + } +} + +struct TeamBudget { + alias: Option, + spend: Spend, +} + +impl TeamBudget { + fn from_wire(budget: &Budget) -> Self { + Self { + alias: budget.team_alias.clone(), + spend: Spend { + spend: budget.spend.unwrap_or(0.0), + limit: budget.max_budget, + reset: parse_date(budget.budget_reset_at.as_deref()), + }, + } + } + + fn window(&self) -> Option { + let label = match &self.alias { + Some(alias) => format!("Team {alias}"), + None => "Team".to_string(), + }; + self.spend.window(Some(&label)) + } +} + +pub(super) fn parse_error(message: impl std::fmt::Display) -> ProviderError { + ProviderError::Parse(format!("LiteLLM parse error: {message}")) +} + +pub(super) fn bind_key(response: KeyInfoResponse) -> Result { + let info = response.info; + let user_id = nonempty(info.user_id); + let team_id = nonempty(info.team_id); + if user_id.is_none() && team_id.is_none() { + return Err(parse_error( + "LiteLLM key info did not include a user_id or team_id.", + )); + } + Ok(KeyBinding { + user_id, + team_id, + expires: parse_date(info.expires.as_deref()), + }) +} + +/// Project a user-bound key: personal budget plus the key's matching team. +pub(super) fn result_from_user( + key: &KeyBinding, + user_id: &str, + response: UserInfoResponse, +) -> Result { + let user = response.user_info; + let response_id = user.user_id.as_deref().or(response.user_id.as_deref()); + if response_id.is_some_and(|id| id != user_id) { + return Err(parse_error("user_id did not match /key/info")); + } + let preferred = user + .metadata + .and_then(|metadata| metadata.preferred_username) + .and_then(|value| value.as_str().map(str::to_owned)); + let email = nonempty(user.user_email) + .or_else(|| nonempty(user.user_alias)) + .or_else(|| nonempty(preferred)); + let mut team = None; + for wire in response.teams.unwrap_or_default() { + let id = wire + .team_id + .as_deref() + .ok_or_else(|| parse_error("missing team_id"))?; + if team.is_none() && key.team_id.as_deref() == Some(id) { + team = Some(TeamBudget::from_wire(&wire)); + } + } + let personal = Spend { + spend: user.spend.unwrap_or(0.0), + limit: user.max_budget, + reset: parse_date(user.budget_reset_at.as_deref()), + }; + let primary = personal + .window(None) + .unwrap_or_else(|| RateWindow::new(0.0)); + let mut snapshot = UsageSnapshot::new(primary); + if let Some(email) = email { + snapshot = snapshot.with_email(email); + } + if let Some(team) = &team { + if let Some(alias) = &team.alias { + snapshot = snapshot.with_organization(alias); + } + if let Some(window) = team.window() { + snapshot = snapshot.with_extra_rate_window("team", "Team budget", window); + } + } + Ok(finish(snapshot, key, &personal, "Personal")) +} + +/// Project a team-only key: the team budget is the sole usage window. +pub(super) fn result_from_team( + key: &KeyBinding, + team_id: &str, + response: TeamInfoResponse, +) -> Result { + let response_id = response + .team_info + .team_id + .as_deref() + .map(str::trim) + .filter(|id| !id.is_empty()) + .or(response + .team_id + .as_deref() + .filter(|id| !id.trim().is_empty())); + if response_id.is_some_and(|id| id != team_id) { + return Err(parse_error("team_id did not match /key/info")); + } + let team = TeamBudget::from_wire(&response.team_info); + let window = team.window().unwrap_or_else(|| RateWindow::new(0.0)); + let mut snapshot = UsageSnapshot::new(window).with_primary_label("Team budget"); + if let Some(alias) = &team.alias { + snapshot = snapshot.with_organization(alias); + } + Ok(finish(snapshot, key, &team.spend, "Team")) +} + +fn finish( + snapshot: UsageSnapshot, + key: &KeyBinding, + spend: &Spend, + scope: &str, +) -> ProviderFetchResult { + let mut snapshot = snapshot.with_login_method("api"); + if key.expires.is_some() { + snapshot = + snapshot.with_subscription(Some(SubscriptionMetadata::new(None, key.expires, None))); + } + let mut result = ProviderFetchResult::new(snapshot, "api"); + let limit = spend.budget(); + if spend.spend > 0.0 || limit.is_some() { + let kind = if limit.is_some() { "budget" } else { "spend" }; + let mut cost = CostSnapshot::new(spend.spend, "USD", format!("{scope} {kind}")); + if let Some(limit) = limit { + cost = cost.with_limit(limit); + } + if let Some(reset) = spend.reset { + cost = cost.with_resets_at(reset); + } + result = result.with_cost(cost); + } + result +} + +fn nonempty(value: Option) -> Option { + value + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) +} + +/// Parse an ISO-8601 instant; naive timestamps are read as UTC and anything +/// unparseable is dropped, matching upstream's tolerant `date()` helper. +fn parse_date(value: Option<&str>) -> Option> { + let raw = value?.trim(); + if raw.is_empty() { + return None; + } + DateTime::parse_from_rfc3339(raw) + .map(|date| date.with_timezone(&Utc)) + .ok() + .or_else(|| { + NaiveDateTime::parse_from_str(raw, "%Y-%m-%dT%H:%M:%S%.f") + .ok() + .map(|date| date.and_utc()) + }) +} diff --git a/rust/src/providers/litellm/mod.rs b/rust/src/providers/litellm/mod.rs index d27a573668..c4e479a67a 100644 --- a/rust/src/providers/litellm/mod.rs +++ b/rust/src/providers/litellm/mod.rs @@ -1,13 +1,27 @@ use async_trait::async_trait; -use reqwest::{Client, Url}; -use serde_json::Value; +use reqwest::{Client, StatusCode, Url}; +use serde::de::DeserializeOwned; use crate::core::{ - CostSnapshot, FetchContext, Provider, ProviderError, ProviderFetchResult, ProviderId, - ProviderMetadata, RateWindow, SourceMode, UsageSnapshot, + FetchContext, Provider, ProviderError, ProviderFetchResult, ProviderId, ProviderMetadata, + SourceMode, +}; +use crate::providers::{BoundedBodyError, read_bounded_response}; + +mod endpoint; +mod info; +#[cfg(test)] +mod tests; + +use endpoint::management_url; +pub(crate) use endpoint::validated_base_url; +use info::{ + KeyInfoResponse, TeamInfoResponse, UserInfoResponse, bind_key, parse_error, result_from_team, + result_from_user, }; const CREDENTIAL_TARGET: &str = "codexbar-litellm"; +const MAX_RESPONSE_BYTES: usize = 1024 * 1024; pub struct LiteLLMProvider { metadata: ProviderMetadata, @@ -38,6 +52,40 @@ impl LiteLLMProvider { } } +impl LiteLLMProvider { + async fn get_json(&self, url: Url, key: &str) -> Result { + let route = url.path().to_string(); + let response = self + .client + .get(url) + .bearer_auth(key) + .header("Accept", "application/json") + .send() + .await?; + let status = response.status(); + if status == StatusCode::UNAUTHORIZED || status == StatusCode::FORBIDDEN { + return Err(ProviderError::AuthRequired); + } + if status == StatusCode::TOO_MANY_REQUESTS { + return Err(ProviderError::Other( + "LiteLLM rate limited the request (HTTP 429).".into(), + )); + } + if !status.is_success() { + return Err(ProviderError::Other(format!( + "LiteLLM {route} returned status {status}" + ))); + } + let body = read_bounded_response(response, MAX_RESPONSE_BYTES) + .await + .map_err(|error| match error { + BoundedBodyError::TooLarge => parse_error("response too large"), + BoundedBodyError::Read(error) => ProviderError::Network(error), + })?; + serde_json::from_slice(&body).map_err(|e| parse_error(format!("{route}: {e}"))) + } +} + impl Default for LiteLLMProvider { fn default() -> Self { Self::new() @@ -58,28 +106,23 @@ impl Provider for LiteLLMProvider { match ctx.source_mode { SourceMode::Auto | SourceMode::OAuth => { let (base, key) = resolve_base_and_key(ctx)?; - let response = self - .client - .get(management_url(&base, "key/info")?) - .bearer_auth(key) - .header("Accept", "application/json") - .send() + let key_info: KeyInfoResponse = self + .get_json(management_url(&base, "key/info", None)?, &key) .await?; - if response.status() == reqwest::StatusCode::UNAUTHORIZED - || response.status() == reqwest::StatusCode::FORBIDDEN - { - return Err(ProviderError::AuthRequired); + let binding = bind_key(key_info)?; + if let Some(user_id) = binding.user_id.as_deref() { + let url = management_url(&base, "user/info", Some(("user_id", user_id)))?; + let response: UserInfoResponse = self.get_json(url, &key).await?; + result_from_user(&binding, user_id, response) + } else if let Some(team_id) = binding.team_id.as_deref() { + let url = management_url(&base, "team/info", Some(("team_id", team_id)))?; + let response: TeamInfoResponse = self.get_json(url, &key).await?; + result_from_team(&binding, team_id, response) + } else { + Err(parse_error( + "LiteLLM key info did not include a user_id or team_id.", + )) } - if !response.status().is_success() { - return Err(ProviderError::Other(format!( - "LiteLLM key/info returned status {}", - response.status() - ))); - } - let value: Value = response.json().await.map_err(|e| { - ProviderError::Parse(format!("Failed to parse LiteLLM key/info: {e}")) - })?; - Ok(result_from_key_info(&value)) } SourceMode::Web | SourceMode::Cli => { Err(ProviderError::UnsupportedSource(ctx.source_mode)) @@ -117,131 +160,3 @@ fn resolve_base_and_key(ctx: &FetchContext) -> Result<(String, String), Provider })?; Ok((base, key)) } - -fn management_url(base: &str, path: &str) -> Result { - let mut url = crate::providers::validated_https_url(base, "LiteLLM base")?; - if url.path().trim_end_matches('/').ends_with("/v1") { - let stripped = url - .path() - .trim_end_matches('/') - .trim_end_matches("/v1") - .to_string(); - url.set_path(&stripped); - } - url.join(path) - .map_err(|e| ProviderError::Other(format!("Invalid LiteLLM URL: {e}"))) -} - -fn result_from_key_info(value: &Value) -> ProviderFetchResult { - let root = value - .get("info") - .or_else(|| value.get("key")) - .unwrap_or(value); - let spend = number(root, &["spend", "spend_usd", "spendUSD"]).unwrap_or(0.0); - let limit = number(root, &["max_budget", "maxBudget", "budget", "limit"]); - let percent = limit - .filter(|v| *v > 0.0) - .map_or(0.0, |limit| spend / limit * 100.0); - let mut primary = RateWindow::new(percent); - if let Some(limit) = limit.filter(|value| *value > 0.0) { - primary.reset_description = Some(budget_detail(spend, limit)); - } - let mut snapshot = UsageSnapshot::new(primary).with_login_method(format!("Spend ${spend:.2}")); - if let Some(team) = root.get("team_info").or_else(|| root.get("teamInfo")) - && let Some(team_spend) = number(team, &["spend", "team_spend", "teamSpend"]) - { - let team_limit = number(team, &["max_budget", "budget", "limit"]); - let team_percent = team_limit - .filter(|v| *v > 0.0) - .map_or(0.0, |limit| team_spend / limit * 100.0); - let mut team_window = RateWindow::new(team_percent); - if let Some(team_limit) = team_limit.filter(|value| *value > 0.0) { - let alias = string(team, &["team_alias", "teamAlias", "alias"]) - .map(|value| format!("Team {value}: ")) - .unwrap_or_default(); - team_window.reset_description = - Some(format!("{alias}{}", budget_detail(team_spend, team_limit))); - } - snapshot = snapshot.with_extra_rate_window("team", "Team budget", team_window); - } - let mut result = ProviderFetchResult::new(snapshot, "api"); - if spend > 0.0 { - let mut cost = CostSnapshot::new(spend, "USD", "Spend"); - if let Some(limit) = limit { - cost = cost.with_limit(limit); - } - result = result.with_cost(cost); - } - result -} - -fn number(value: &Value, keys: &[&str]) -> Option { - keys.iter() - .find_map(|key| value.get(*key).and_then(Value::as_f64)) -} - -fn string(value: &Value, keys: &[&str]) -> Option { - keys.iter() - .find_map(|key| value.get(*key).and_then(Value::as_str)) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) -} - -fn budget_detail(spend: f64, budget: f64) -> String { - format!("${spend:.2} / ${budget:.2}") -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parses_spend_budget() { - let result = - result_from_key_info(&serde_json::json!({"info":{"spend":25.0,"max_budget":100.0}})); - assert_eq!(result.usage.primary.used_percent, 25.0); - assert_eq!( - result.usage.primary.reset_description.as_deref(), - Some("$25.00 / $100.00") - ); - } - - #[test] - fn preserves_team_budget_detail_with_alias() { - let result = result_from_key_info(&serde_json::json!({ - "info": { - "team_info": { - "team_alias": "Platform", - "spend": 70.0, - "max_budget": 1000.0 - } - } - })); - assert_eq!(result.usage.extra_rate_windows.len(), 1); - assert_eq!( - result.usage.extra_rate_windows[0] - .window - .reset_description - .as_deref(), - Some("Team Platform: $70.00 / $1000.00") - ); - } - - #[test] - fn saved_base_url_uses_only_app_saved_key() { - let mut ctx = FetchContext { - workspace_id: Some("https://litellm.example.com".to_string()), - ..Default::default() - }; - assert!(matches!( - resolve_base_and_key(&ctx), - Err(ProviderError::AuthRequired) - )); - - ctx.api_key = Some("sk-app".to_string()); - let (base, key) = resolve_base_and_key(&ctx).unwrap(); - assert_eq!(base, "https://litellm.example.com"); - assert_eq!(key, "sk-app"); - } -} diff --git a/rust/src/providers/litellm/tests.rs b/rust/src/providers/litellm/tests.rs new file mode 100644 index 0000000000..3cd74e01ae --- /dev/null +++ b/rust/src/providers/litellm/tests.rs @@ -0,0 +1,278 @@ +use serde_json::{Value, json}; + +use super::info::{ + KeyBinding, KeyInfoResponse, TeamInfoResponse, UserInfoResponse, bind_key, result_from_team, + result_from_user, +}; +use super::*; + +fn binding(info: Value) -> KeyBinding { + bind_key(serde_json::from_value::(json!({ "info": info })).unwrap()).unwrap() +} + +fn user_result(key: &KeyBinding, body: Value) -> Result { + let user_id = key.user_id.clone().unwrap(); + result_from_user( + key, + &user_id, + serde_json::from_value::(body).unwrap(), + ) +} + +fn team_result(key: &KeyBinding, body: Value) -> Result { + let team_id = key.team_id.clone().unwrap(); + result_from_team( + key, + &team_id, + serde_json::from_value::(body).unwrap(), + ) +} + +fn assert_parse_error(result: Result, expected: &str) { + match result { + Err(ProviderError::Parse(message)) => { + assert!(message.contains(expected), "unexpected message: {message}") + } + Err(other) => panic!("expected parse error, got {other}"), + Ok(_) => panic!("expected parse error"), + } +} + +#[test] +fn key_info_without_user_or_team_id_fails() { + let response: KeyInfoResponse = + serde_json::from_value(json!({"info": {"user_id": " ", "spend": 1.0}})).unwrap(); + match bind_key(response) { + Err(ProviderError::Parse(message)) => assert!( + message.contains("LiteLLM key info did not include a user_id or team_id."), + "unexpected message: {message}" + ), + _ => panic!("expected parse error"), + } +} + +#[test] +fn key_info_requires_info_object() { + assert!(serde_json::from_value::(json!({"user_id": "u"})).is_err()); +} + +#[test] +fn personal_budget_is_primary_with_identity() { + let key = binding(json!({"user_id": "user-1", "expires": "2026-12-31T00:00:00Z"})); + let result = user_result( + &key, + json!({ + "user_id": "user-1", + "user_info": { + "user_id": "user-1", + "user_email": "dev@example.com", + "spend": 25.0, + "max_budget": 100.0, + "budget_reset_at": "2026-10-01T00:00:00Z" + }, + "teams": [] + }), + ) + .unwrap(); + assert_eq!(result.usage.primary.used_percent, 25.0); + assert_eq!( + result.usage.primary.reset_description.as_deref(), + Some("$25.00 / $100.00") + ); + assert!(result.usage.primary.resets_at.is_some()); + assert_eq!( + result.usage.account_email.as_deref(), + Some("dev@example.com") + ); + assert_eq!(result.usage.login_method.as_deref(), Some("api")); + assert!(result.usage.extra_rate_windows.is_empty()); + assert!( + result + .usage + .subscription + .as_ref() + .is_some_and(|sub| sub.expires_at.is_some()) + ); + let cost = result.cost.expect("personal cost"); + assert_eq!(cost.used, 25.0); + assert_eq!(cost.limit, Some(100.0)); + assert_eq!(cost.period, "Personal budget"); +} + +#[test] +fn identity_falls_back_to_alias_then_preferred_username() { + let key = binding(json!({"user_id": "user-1"})); + let alias = user_result( + &key, + json!({"user_info": {"user_alias": "alias", "metadata": {"preferred_username": "pref"}}}), + ) + .unwrap(); + assert_eq!(alias.usage.account_email.as_deref(), Some("alias")); + let pref = user_result( + &key, + json!({"user_info": {"user_email": " ", "metadata": {"preferred_username": "pref"}}}), + ) + .unwrap(); + assert_eq!(pref.usage.account_email.as_deref(), Some("pref")); +} + +#[test] +fn matching_team_budget_is_a_separate_row() { + let key = binding(json!({"user_id": "user-1", "team_id": "team-b"})); + let result = user_result( + &key, + json!({ + "user_info": {"user_id": "user-1", "spend": 3.0}, + "teams": [ + {"team_id": "team-a", "team_alias": "Other", "spend": 1.0, "max_budget": 10.0}, + {"team_id": "team-b", "team_alias": "Platform", "spend": 70.0, "max_budget": 1000.0} + ] + }), + ) + .unwrap(); + assert_eq!(result.usage.extra_rate_windows.len(), 1); + let team = &result.usage.extra_rate_windows[0].window; + assert!((team.used_percent - 7.0).abs() < 1e-9); + assert_eq!( + team.reset_description.as_deref(), + Some("Team Platform: $70.00 / $1000.00") + ); + assert_eq!( + result.usage.account_organization.as_deref(), + Some("Platform") + ); + let cost = result.cost.expect("spend-only cost"); + assert_eq!(cost.period, "Personal spend"); + assert_eq!(cost.limit, None); +} + +#[test] +fn team_without_a_matching_entry_is_omitted() { + let key = binding(json!({"user_id": "user-1", "team_id": "team-x"})); + let result = user_result( + &key, + json!({ + "user_info": {"spend": 1.0}, + "teams": [{"team_id": "team-a", "spend": 1.0, "max_budget": 10.0}] + }), + ) + .unwrap(); + assert!(result.usage.extra_rate_windows.is_empty()); + assert_eq!(result.usage.account_organization, None); +} + +#[test] +fn mismatched_user_id_is_rejected() { + let key = binding(json!({"user_id": "user-1"})); + assert_parse_error( + user_result(&key, json!({"user_info": {"user_id": "user-2"}})), + "user_id did not match /key/info", + ); + assert_parse_error( + user_result(&key, json!({"user_id": "user-2", "user_info": {}})), + "user_id did not match /key/info", + ); +} + +#[test] +fn team_entries_without_team_id_are_rejected() { + let key = binding(json!({"user_id": "user-1", "team_id": "team-a"})); + assert_parse_error( + user_result(&key, json!({"user_info": {}, "teams": [{"spend": 1.0}]})), + "missing team_id", + ); +} + +#[test] +fn wrongly_typed_fields_fail_to_parse() { + assert!( + serde_json::from_value::(json!({"user_info": {"spend": "12"}})).is_err() + ); + assert!(serde_json::from_value::(json!({"teams": []})).is_err()); +} + +#[test] +fn team_only_key_shows_team_budget_as_sole_window() { + let key = binding(json!({"team_id": "team-a"})); + let result = team_result( + &key, + json!({ + "team_id": "team-a", + "team_info": { + "team_id": "team-a", + "team_alias": "Platform", + "spend": 70.0, + "max_budget": 1000.0, + "budget_reset_at": "2026-10-01T00:00:00" + } + }), + ) + .unwrap(); + assert!((result.usage.primary.used_percent - 7.0).abs() < 1e-9); + assert_eq!( + result.usage.primary.reset_description.as_deref(), + Some("Team Platform: $70.00 / $1000.00") + ); + assert!(result.usage.primary.resets_at.is_some()); + assert_eq!(result.usage.primary_label.as_deref(), Some("Team budget")); + assert!(result.usage.extra_rate_windows.is_empty()); + assert_eq!( + result.usage.account_organization.as_deref(), + Some("Platform") + ); + assert_eq!(result.usage.account_email, None); + assert_eq!(result.cost.expect("team cost").period, "Team budget"); +} + +#[test] +fn mismatched_team_id_is_rejected() { + let key = binding(json!({"team_id": "team-a"})); + assert_parse_error( + team_result( + &key, + json!({"team_id": "team-b", "team_info": {"spend": 1.0}}), + ), + "team_id did not match /key/info", + ); + assert_parse_error( + team_result(&key, json!({"team_info": {"team_id": "team-b"}})), + "team_id did not match /key/info", + ); +} + +#[test] +fn spend_above_budget_clamps_percent_and_zero_budget_is_spend_only() { + let key = binding(json!({"user_id": "user-1"})); + let over = user_result( + &key, + json!({"user_info": {"spend": 150.0, "max_budget": 100.0}}), + ) + .unwrap(); + assert_eq!(over.usage.primary.used_percent, 100.0); + let unbudgeted = user_result( + &key, + json!({"user_info": {"spend": 4.0, "max_budget": 0.0}}), + ) + .unwrap(); + assert_eq!(unbudgeted.usage.primary.reset_description, None); + assert_eq!(unbudgeted.cost.expect("cost").period, "Personal spend"); + let empty = user_result(&key, json!({"user_info": {}})).unwrap(); + assert!(empty.cost.is_none()); +} + +#[test] +fn saved_base_url_uses_only_app_saved_key() { + let mut ctx = FetchContext { + workspace_id: Some("https://litellm.example.com".to_string()), + ..Default::default() + }; + assert!(matches!( + resolve_base_and_key(&ctx), + Err(ProviderError::AuthRequired) + )); + + ctx.api_key = Some("sk-app".to_string()); + let (base, key) = resolve_base_and_key(&ctx).unwrap(); + assert_eq!(base, "https://litellm.example.com"); + assert_eq!(key, "sk-app"); +} diff --git a/rust/src/settings/provider_workspace.rs b/rust/src/settings/provider_workspace.rs index 338747c8ef..e0e48fb4be 100644 --- a/rust/src/settings/provider_workspace.rs +++ b/rust/src/settings/provider_workspace.rs @@ -55,12 +55,19 @@ pub fn validate_provider_workspace_value( Err("Helmcode tenant must be 'helmcode' or 'nanBuilders'".to_string()) } } - ProviderId::LiteLLM => validate_token_endpoint(trimmed, "LiteLLM base URL", |_| true), + ProviderId::LiteLLM => validate_litellm_base_url(trimmed), ProviderId::Sub2Api => validate_sub2api_base_url(trimmed), _ => Ok(trimmed.to_string()), } } +fn validate_litellm_base_url(raw: &str) -> Result { + match crate::providers::litellm::validated_base_url(raw) { + Ok(url) => Ok(url.to_string().trim_end_matches('/').to_string()), + Err(err) => Err(err.to_string()), + } +} + fn validate_sub2api_base_url(raw: &str) -> Result { match crate::providers::sub2api::validated_sub2api_base_url(raw) { Ok(url) => Ok(url.to_string().trim_end_matches('/').to_string()), @@ -195,7 +202,7 @@ mod tests { } #[test] - fn validates_token_endpoint_hosts() { + fn validates_litellm_base_url_policy() { assert_eq!( validate_provider_workspace_value( ProviderId::LiteLLM, @@ -204,12 +211,23 @@ mod tests { .unwrap(), "https://litellm.example.com/v1" ); + for value in [ + "http://127.0.0.1:4000", + "http://10.0.0.5:4000", + "http://192.168.1.4", + "http://[::1]:4000", + "http://proxy.local:4000", + "https://10.0.0.5", + ] { + assert!( + validate_provider_workspace_value(ProviderId::LiteLLM, value).is_ok(), + "rejected {value}" + ); + } for value in [ "http://litellm.example.com", + "http://8.8.8.8", "https://user@litellm.example.com", - "https://127.0.0.1", - "https://10.0.0.5", - "https://[::1]", "https://example.com%2f.evil.test", ] { assert!(