1use std::time::Duration;
7
8use clap::ValueEnum;
9
10use crate::usage::VendorSnapshot;
11use crate::widget::cli::Cli;
12
13pub const HTTP_CLIENT_TIMEOUT: Duration = Duration::from_secs(30);
16
17pub const MAX_BODY_BYTES: usize = 2 * 1024 * 1024;
22
23pub(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
50pub 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
74pub 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#[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 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 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#[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#[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 #[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 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 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}