Skip to main content

ai_usagebar/
vendor.rs

1//! Shared vendor IDs and renderer/fetcher structs used by the widget and TUI.
2//!
3//! Snapshots remain a discriminated `VendorSnapshot` enum because the vendors
4//! have genuinely different shapes — see `usage.rs`.
5
6use std::time::Duration;
7
8use clap::ValueEnum;
9
10use crate::usage::VendorSnapshot;
11use crate::widget::cli::Cli;
12
13/// Outer reqwest client timeout shared by widget and TUI entry points.
14/// Vendor fetchers still apply their own tighter per-request timeouts.
15pub const HTTP_CLIENT_TIMEOUT: Duration = Duration::from_secs(30);
16
17/// Upper bound on a vendor response body. Every one of these endpoints returns
18/// a small JSON document — the largest observed is a few kilobytes — so this is
19/// generous by three orders of magnitude while still bounding the damage from a
20/// misbehaving proxy or a hijacked endpoint.
21pub const MAX_BODY_BYTES: usize = 2 * 1024 * 1024;
22
23/// Credential-bearing environment variables owned by ai-usagebar vendors.
24/// Subprocesses receive only the entries that belong to their own provider.
25pub(crate) const VENDOR_SECRET_ENV_VARS: &[&str] = &[
26    "ZAI_API_KEY",
27    "OPENROUTER_API_KEY",
28    "DEEPSEEK_API_KEY",
29    "KIMI_API_KEY",
30    "KILO_API_KEY",
31    "NOVITA_API_KEY",
32    "MINIMAX_API_KEY",
33    "MOONSHOT_API_KEY",
34    "XAI_MANAGEMENT_KEY",
35    "ANTHROPIC_ADMIN_KEY",
36    "XAI_API_KEY",
37    "GROK_API_KEY",
38    "OPENCODE_GO_API_KEY",
39    "COMMANDCODE_API_KEY",
40];
41
42pub(crate) fn vendor_secret_env_vars_to_remove(keep: &[&str]) -> Vec<&'static str> {
43    VENDOR_SECRET_ENV_VARS
44        .iter()
45        .copied()
46        .filter(|var| !keep.contains(var))
47        .collect()
48}
49
50/// Follow ordinary vendor redirects without forwarding non-standard API-key
51/// headers to a different origin. Reqwest strips `Authorization` on sensitive
52/// redirects, but vendors also use headers such as `x-api-key`, which are not
53/// covered by that built-in list.
54pub fn same_origin_redirect_policy() -> reqwest::redirect::Policy {
55    reqwest::redirect::Policy::custom(|attempt| {
56        if attempt.previous().len() >= 10 {
57            return attempt.error("too many redirects");
58        }
59        let Some(origin) = attempt.previous().first() else {
60            return attempt.stop();
61        };
62        let target = attempt.url();
63        if target.scheme() == origin.scheme()
64            && target.host_str() == origin.host_str()
65            && target.port_or_known_default() == origin.port_or_known_default()
66        {
67            attempt.follow()
68        } else {
69            attempt.stop()
70        }
71    })
72}
73
74/// Read a response body with an upper bound.
75///
76/// Every vendor buffered the whole body with `resp.bytes()` *before* anything
77/// validated it. The widget is re-executed by Waybar every 60s, so an endpoint
78/// answering with an unbounded stream had a free hand at the machine's memory.
79/// `Content-Length` is checked first when present, then the body is read in
80/// chunks so a lying or absent length cannot get past the cap either.
81pub async fn read_body_capped(
82    mut resp: reqwest::Response,
83    max: usize,
84) -> crate::error::Result<Vec<u8>> {
85    let too_big = |n: u64| {
86        crate::error::AppError::Schema(format!(
87            "response body exceeds the {max}-byte limit ({n} bytes); refusing to buffer it"
88        ))
89    };
90    if let Some(len) = resp.content_length()
91        && len > max as u64
92    {
93        return Err(too_big(len));
94    }
95    let mut buf: Vec<u8> = Vec::new();
96    while let Some(chunk) = resp.chunk().await? {
97        if chunk.len() > max.saturating_sub(buf.len()) {
98            return Err(too_big(buf.len().saturating_add(chunk.len()) as u64));
99        }
100        buf.extend_from_slice(&chunk);
101    }
102    Ok(buf)
103}
104
105/// Stable enum used by `--vendor` and in config files.
106#[derive(
107    Debug, Clone, Copy, ValueEnum, PartialEq, Eq, Hash, serde::Deserialize, serde::Serialize,
108)]
109#[serde(rename_all = "lowercase")]
110pub enum VendorId {
111    Anthropic,
112    #[serde(rename = "anthropic_api")]
113    AnthropicApi,
114    Openai,
115    Zai,
116    Openrouter,
117    Deepseek,
118    Kimi,
119    Kilo,
120    Novita,
121    Moonshot,
122    Grok,
123    Supergrok,
124    Antigravity,
125    Cursor,
126    Minimax,
127    Kiro,
128    #[serde(rename = "nous")]
129    NousResearch,
130    #[serde(rename = "opencode-go")]
131    OpenCodeGo,
132    #[serde(rename = "commandcode")]
133    CommandCode,
134}
135
136impl VendorId {
137    pub fn slug(self) -> &'static str {
138        match self {
139            VendorId::Anthropic => "anthropic",
140            VendorId::AnthropicApi => "anthropic_api",
141            VendorId::Openai => "openai",
142            VendorId::Zai => "zai",
143            VendorId::Openrouter => "openrouter",
144            VendorId::Deepseek => "deepseek",
145            VendorId::Kimi => "kimi",
146            VendorId::Kilo => "kilo",
147            VendorId::Novita => "novita",
148            VendorId::Moonshot => "moonshot",
149            VendorId::Grok => "grok",
150            VendorId::Supergrok => "supergrok",
151            VendorId::Antigravity => "antigravity",
152            VendorId::Cursor => "cursor",
153            VendorId::Minimax => "minimax",
154            VendorId::Kiro => "kiro",
155            VendorId::NousResearch => "nous",
156            VendorId::OpenCodeGo => "opencode-go",
157            VendorId::CommandCode => "commandcode",
158        }
159    }
160
161    /// Canonical human-readable name for shared reports and compact UI labels.
162    /// Platform frontends may add context (for example, "GLM (Z.AI)" in a
163    /// wide TUI tab), but should not carry their own full vendor-name table.
164    pub fn display_name(self) -> &'static str {
165        match self {
166            VendorId::Anthropic => "Claude",
167            VendorId::AnthropicApi => "Anthropic API",
168            VendorId::Openai => "Codex",
169            VendorId::Zai => "Z.AI",
170            VendorId::Openrouter => "OpenRouter",
171            VendorId::Deepseek => "DeepSeek",
172            VendorId::Kimi => "Kimi",
173            VendorId::Kilo => "Kilo",
174            VendorId::Novita => "Novita",
175            VendorId::Moonshot => "Moonshot",
176            VendorId::Grok => "Grok",
177            VendorId::Supergrok => "SuperGrok",
178            VendorId::Antigravity => "Antigravity",
179            VendorId::Cursor => "Cursor",
180            VendorId::Minimax => "MiniMax",
181            VendorId::Kiro => "Kiro",
182            VendorId::NousResearch => "Nous Research",
183            VendorId::OpenCodeGo => "OpenCode Go",
184            VendorId::CommandCode => "Command Code",
185        }
186    }
187
188    /// Compact three-letter code for the bar. This is the single source for
189    /// `{vendor_short}` in every renderer, the `usage --json` `short_name`
190    /// field, and any frontend that wants a Waybar-style provider tag; a
191    /// second copy in a placeholder map or a QML file is how the table forks.
192    pub const fn short_name(self) -> &'static str {
193        match self {
194            VendorId::Anthropic => "cld",
195            VendorId::AnthropicApi => "aac",
196            VendorId::Openai => "gpt",
197            VendorId::Zai => "zai",
198            VendorId::Openrouter => "opr",
199            VendorId::Deepseek => "dsk",
200            VendorId::Kimi => "kmi",
201            VendorId::Kilo => "klo",
202            VendorId::Novita => "nvt",
203            VendorId::Moonshot => "msh",
204            VendorId::Grok => "grk",
205            VendorId::Supergrok => "sgk",
206            VendorId::Antigravity => "agy",
207            VendorId::Cursor => "cur",
208            VendorId::Minimax => "mmx",
209            VendorId::Kiro => "kir",
210            VendorId::NousResearch => "nrs",
211            VendorId::OpenCodeGo => "ocg",
212            VendorId::CommandCode => "cmc",
213        }
214    }
215
216    pub fn all() -> &'static [VendorId] {
217        &[
218            VendorId::Anthropic,
219            VendorId::AnthropicApi,
220            VendorId::Openai,
221            VendorId::Zai,
222            VendorId::Openrouter,
223            VendorId::Deepseek,
224            VendorId::Kimi,
225            VendorId::Kilo,
226            VendorId::Novita,
227            VendorId::Moonshot,
228            VendorId::Grok,
229            VendorId::Supergrok,
230            VendorId::Antigravity,
231            VendorId::Cursor,
232            VendorId::Minimax,
233            VendorId::Kiro,
234            VendorId::NousResearch,
235            VendorId::OpenCodeGo,
236            VendorId::CommandCode,
237        ]
238    }
239}
240
241/// What a vendor returns from a successful fetch — snapshot + meta. Mirrors
242/// `anthropic::fetch::FetchOutcome` but vendor-agnostic.
243#[derive(Debug, Clone)]
244pub struct VendorOutcome {
245    pub snapshot: VendorSnapshot,
246    pub stale: bool,
247    pub last_error: Option<(u16, String)>,
248    pub cache_age: Option<std::time::Duration>,
249}
250
251/// Options forwarded to renderers from the CLI.
252#[derive(Debug, Clone)]
253pub struct RenderOpts {
254    pub format: Option<String>,
255    pub tooltip_format: Option<String>,
256    pub icon: Option<String>,
257    pub pace_tolerance: u32,
258    pub format_pace_color: bool,
259    pub tooltip_pace_pts: bool,
260}
261
262impl RenderOpts {
263    pub fn from_cli(cli: &Cli) -> Self {
264        Self {
265            format: cli.format.clone(),
266            tooltip_format: cli.tooltip_format.clone(),
267            icon: cli.icon.clone(),
268            pace_tolerance: cli.pace_tolerance,
269            format_pace_color: cli.format_pace_color,
270            tooltip_pace_pts: cli.tooltip_pace_pts,
271        }
272    }
273}
274
275#[cfg(test)]
276mod tests {
277    use super::*;
278
279    #[test]
280    fn every_vendor_has_stable_machine_and_display_names() {
281        for vendor in VendorId::all() {
282            assert!(!vendor.slug().is_empty());
283            assert!(!vendor.display_name().is_empty());
284        }
285        assert_eq!(VendorId::Anthropic.slug(), "anthropic");
286        assert_eq!(VendorId::Anthropic.display_name(), "Claude");
287        assert_eq!(VendorId::Openai.display_name(), "Codex");
288        assert_eq!(VendorId::Zai.display_name(), "Z.AI");
289    }
290
291    /// `{vendor_short}` is a documented format placeholder and now also rides
292    /// the `usage --json` report, so a duplicate or a re-typed code would make
293    /// two providers indistinguishable in a bar that shows nothing else.
294    #[test]
295    fn every_vendor_short_name_is_a_unique_three_letter_code() {
296        let mut seen = std::collections::BTreeSet::new();
297        for vendor in VendorId::all() {
298            let short = vendor.short_name();
299            assert_eq!(short.len(), 3, "{} is not three letters", vendor.slug());
300            assert!(
301                short.chars().all(|c| c.is_ascii_lowercase()),
302                "{} is not lowercase ascii",
303                vendor.slug()
304            );
305            assert!(seen.insert(short), "{short} is used by two vendors");
306        }
307        assert_eq!(VendorId::Anthropic.short_name(), "cld");
308        assert_eq!(VendorId::Openai.short_name(), "gpt");
309        assert_eq!(VendorId::Zai.short_name(), "zai");
310        assert_eq!(VendorId::Antigravity.short_name(), "agy");
311    }
312
313    #[test]
314    fn new_vendor_contracts_keep_public_names_and_slugs() {
315        assert_eq!(VendorId::NousResearch.slug(), "nous");
316        assert_eq!(VendorId::NousResearch.display_name(), "Nous Research");
317        assert_eq!(VendorId::OpenCodeGo.slug(), "opencode-go");
318        assert_eq!(VendorId::OpenCodeGo.display_name(), "OpenCode Go");
319        assert_eq!(
320            serde_json::to_value(VendorId::OpenCodeGo).unwrap(),
321            serde_json::json!("opencode-go")
322        );
323    }
324
325    #[test]
326    fn vendor_secret_env_vars_cover_config_defaults() {
327        let configured_defaults = [
328            "ZAI_API_KEY",
329            "OPENROUTER_API_KEY",
330            "DEEPSEEK_API_KEY",
331            "KIMI_API_KEY",
332            "KILO_API_KEY",
333            "NOVITA_API_KEY",
334            "MINIMAX_API_KEY",
335            "MOONSHOT_API_KEY",
336            "XAI_MANAGEMENT_KEY",
337            "ANTHROPIC_ADMIN_KEY",
338        ];
339        for name in configured_defaults {
340            assert!(VENDOR_SECRET_ENV_VARS.contains(&name), "missing {name}");
341        }
342    }
343
344    #[test]
345    fn vars_to_remove_preserves_only_requested_grok_credentials() {
346        let removed = vendor_secret_env_vars_to_remove(&["XAI_API_KEY", "GROK_API_KEY"]);
347        assert!(!removed.contains(&"XAI_API_KEY"));
348        assert!(!removed.contains(&"GROK_API_KEY"));
349        assert!(removed.contains(&"ANTHROPIC_ADMIN_KEY"));
350        assert!(removed.contains(&"OPENROUTER_API_KEY"));
351        assert_eq!(removed.len(), VENDOR_SECRET_ENV_VARS.len() - 2);
352    }
353
354    #[tokio::test]
355    async fn body_over_the_cap_is_refused_and_under_it_round_trips() {
356        let mut server = mockito::Server::new_async().await;
357        server
358            .mock("GET", "/big")
359            .with_status(200)
360            .with_body("x".repeat(4096))
361            .create_async()
362            .await;
363        server
364            .mock("GET", "/small")
365            .with_status(200)
366            .with_body("hello")
367            .create_async()
368            .await;
369
370        let client = reqwest::Client::new();
371
372        // Over the cap: refused rather than buffered.
373        let resp = client
374            .get(format!("{}/big", server.url()))
375            .send()
376            .await
377            .unwrap();
378        let err = read_body_capped(resp, 1024).await.unwrap_err();
379        assert!(
380            err.to_string().contains("exceeds"),
381            "unexpected error: {err}"
382        );
383
384        // Under the cap: identical to the previous `resp.bytes()` behaviour.
385        let resp = client
386            .get(format!("{}/small", server.url()))
387            .send()
388            .await
389            .unwrap();
390        assert_eq!(read_body_capped(resp, 1024).await.unwrap(), b"hello");
391    }
392
393    #[tokio::test]
394    async fn chunked_body_without_content_length_still_hits_the_cap() {
395        let mut server = mockito::Server::new_async().await;
396        server
397            .mock("GET", "/chunked")
398            .with_status(200)
399            .with_chunked_body(|writer| writer.write_all(&[b'x'; 4096]))
400            .create_async()
401            .await;
402
403        let response = reqwest::Client::new()
404            .get(format!("{}/chunked", server.url()))
405            .send()
406            .await
407            .unwrap();
408        assert!(response.content_length().is_none());
409        let error = read_body_capped(response, 1024).await.unwrap_err();
410        assert!(error.to_string().contains("exceeds"), "{error}");
411    }
412
413    #[tokio::test]
414    async fn same_origin_redirects_still_work_with_vendor_headers() {
415        let mut server = mockito::Server::new_async().await;
416        let redirect = server
417            .mock("GET", "/start")
418            .match_header("x-api-key", "secret")
419            .with_status(302)
420            .with_header("location", "/finish")
421            .create_async()
422            .await;
423        let finish = server
424            .mock("GET", "/finish")
425            .match_header("x-api-key", "secret")
426            .with_status(200)
427            .create_async()
428            .await;
429        let client = reqwest::Client::builder()
430            .redirect(same_origin_redirect_policy())
431            .build()
432            .unwrap();
433
434        let response = client
435            .get(format!("{}/start", server.url()))
436            .header("x-api-key", "secret")
437            .send()
438            .await
439            .unwrap();
440
441        assert_eq!(response.status(), reqwest::StatusCode::OK);
442        redirect.assert_async().await;
443        finish.assert_async().await;
444    }
445
446    #[tokio::test]
447    async fn cross_origin_redirects_are_not_followed_with_vendor_headers() {
448        let mut origin = mockito::Server::new_async().await;
449        let mut target = mockito::Server::new_async().await;
450        let target_url = format!("{}/capture", target.url());
451        let redirect = origin
452            .mock("GET", "/start")
453            .match_header("x-api-key", "secret")
454            .with_status(302)
455            .with_header("location", &target_url)
456            .create_async()
457            .await;
458        let capture = target
459            .mock("GET", "/capture")
460            .expect(0)
461            .create_async()
462            .await;
463        let client = reqwest::Client::builder()
464            .redirect(same_origin_redirect_policy())
465            .build()
466            .unwrap();
467
468        let response = client
469            .get(format!("{}/start", origin.url()))
470            .header("x-api-key", "secret")
471            .send()
472            .await
473            .unwrap();
474
475        assert_eq!(response.status(), reqwest::StatusCode::FOUND);
476        redirect.assert_async().await;
477        capture.assert_async().await;
478    }
479
480    #[test]
481    fn vendor_id_slug_round_trip() {
482        for id in VendorId::all() {
483            assert_eq!(
484                id.slug(),
485                serde_json::to_value(id).unwrap().as_str().unwrap()
486            );
487        }
488    }
489}