use super::*; use super::access_token::{distrust_access_token, get_access_token}; use crate::config::{RenderMode, paths}; use crate::{ config::{AppConfig, Input, RequestContext}, function::{FunctionDeclaration, ToolCall, ToolResult, eval_tool_calls}, render::render_stream, utils::*, }; use crate::vault::Vault; use anyhow::{Context, Result, bail}; use fancy_regex::Regex; use indexmap::IndexMap; use inquire::{ MultiSelect, Select, Text, list_option::ListOption, required, validator::Validation, }; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{Value, json}; use std::sync::LazyLock; use std::time::Duration; use tokio::sync::mpsc::unbounded_channel; pub const MODELS_YAML: &str = include_str!("../../models.yaml"); pub static ALL_PROVIDER_MODELS: LazyLock> = LazyLock::new(|| { paths::local_models_override() .ok() .unwrap_or_else(|| serde_yaml::from_str(MODELS_YAML).unwrap()) }); static EMBEDDING_MODEL_RE: LazyLock = LazyLock::new(|| { Regex::new(r"((^|/)(bge-|e5-|uae-|gte-|text-)|embed|multilingual|minilm)").unwrap() }); static ESCAPE_SLASH_RE: LazyLock = LazyLock::new(|| Regex::new(r"(? &AppConfig; fn extra_config(&self) -> Option<&ExtraConfig>; fn patch_config(&self) -> Option<&RequestPatch>; fn name(&self) -> &str; fn model(&self) -> &Model; fn supports_oauth(&self) -> bool { false } fn build_client(&self) -> Result { let mut builder = ReqwestClient::builder(); let extra = self.extra_config(); let timeout = extra.and_then(|v| v.connect_timeout).unwrap_or(10); let read_timeout = extra.and_then(|v| v.read_timeout).unwrap_or(300); if let Some(proxy) = extra.and_then(|v| v.proxy.as_deref()) { builder = set_proxy(builder, proxy)?; } if let Some(user_agent) = self.app_config().user_agent.as_ref() { builder = builder.user_agent(user_agent); } if read_timeout > 0 { builder = builder.read_timeout(Duration::from_secs(read_timeout)); } let client = builder .connect_timeout(Duration::from_secs(timeout)) .build() .with_context(|| "Failed to build client")?; Ok(client) } /// On a 401 the cached access token is distrusted and the call retried /// exactly once; the retry re-runs the per-client prepare step, which /// sees the rejection marker, force-refreshes the token, and rebuilds /// the whole request. A second 401 propagates the original error; any /// other retry failure propagates as-is. async fn chat_completions(&self, input: Input) -> Result { if self.app_config().dry_run { let content = input.echo_messages(); return Ok(ChatCompletionsOutput::new(&content)); } let client = self.build_client()?; let data = input.prepare_completion_data(self.model(), false)?; let err = match self.chat_completions_inner(&client, data).await { Ok(output) => return Ok(output), Err(err) => err, }; let ret = if should_retry_auth(&err, self.name()) { debug!( "provider '{}' rejected access token (401); refreshing and retrying once", self.name() ); let data = input.prepare_completion_data(self.model(), false)?; match self.chat_completions_inner(&client, data).await { Err(retry_err) if is_auth_error(&retry_err) => Err(err), ret => ret, } } else { Err(err) }; ret.with_context(|| "Failed to call chat-completions api") } /// Same retry-once-on-401 semantics as [`Self::chat_completions`], but /// only while the handler has received nothing yet: retrying after /// partial output has streamed would render it to the user twice. The /// retry lives inside the same `select!` arm so abort stays responsive. async fn chat_completions_streaming( &self, input: &Input, handler: &mut SseHandler, ) -> Result<()> { let abort_signal = handler.abort(); let input = input.clone(); tokio::select! { ret = async { if self.app_config().dry_run { let content = input.echo_messages(); handler.text(&content)?; return Ok(()); } let client = self.build_client()?; let data = input.prepare_completion_data(self.model(), true)?; let err = match self.chat_completions_streaming_inner(&client, handler, data).await { Ok(()) => return Ok(()), Err(err) => err, }; if handler.has_received_content() || !should_retry_auth(&err, self.name()) { return Err(err); } debug!( "provider '{}' rejected access token (401); refreshing and retrying once", self.name() ); let data = input.prepare_completion_data(self.model(), true)?; match self.chat_completions_streaming_inner(&client, handler, data).await { Err(retry_err) if is_auth_error(&retry_err) => Err(err), ret => ret, } } => { handler.done(); ret.with_context(|| "Failed to call chat-completions api") } _ = wait_abort_signal(&abort_signal) => { handler.done(); Ok(()) }, } } /// Same retry-once-on-401 semantics as [`Self::chat_completions`] /// (gemini OAuth embeddings route here). async fn embeddings(&self, data: &EmbeddingsData) -> Result>> { let client = self.build_client()?; let err = match self.embeddings_inner(&client, data).await { Ok(output) => return Ok(output), Err(err) => err, }; let ret = if should_retry_auth(&err, self.name()) { debug!( "provider '{}' rejected access token (401); refreshing and retrying once", self.name() ); match self.embeddings_inner(&client, data).await { Err(retry_err) if is_auth_error(&retry_err) => Err(err), ret => ret, } } else { Err(err) }; ret.context("Failed to call embeddings api") } async fn rerank(&self, data: &RerankData) -> Result { let client = self.build_client()?; self.rerank_inner(&client, data) .await .context("Failed to call rerank api") } async fn chat_completions_inner( &self, client: &ReqwestClient, data: ChatCompletionsData, ) -> Result; async fn chat_completions_streaming_inner( &self, client: &ReqwestClient, handler: &mut SseHandler, data: ChatCompletionsData, ) -> Result<()>; async fn embeddings_inner( &self, _client: &ReqwestClient, _data: &EmbeddingsData, ) -> Result { bail!("The client doesn't support embeddings api") } async fn rerank_inner( &self, _client: &ReqwestClient, _data: &RerankData, ) -> Result { bail!("The client doesn't support rerank api") } fn request_builder( &self, client: &reqwest::Client, mut request_data: RequestData, ) -> RequestBuilder { self.patch_request_data(&mut request_data); request_data.into_builder(client) } fn patch_request_data(&self, request_data: &mut RequestData) { let model_type = self.model().model_type(); if let Some(patch) = self.model().patch() { request_data.apply_patch(patch.clone()); } let patch_map = std::env::var(get_env_name(&format!( "patch_{}_{}", self.model().client_name(), model_type.api_name(), ))) .ok() .and_then(|v| serde_json::from_str(&v).ok()) .or_else(|| { self.patch_config() .and_then(|v| model_type.extract_patch(v)) .cloned() }); let patch_map = match patch_map { Some(v) => v, _ => return, }; for (key, patch) in patch_map { let key = ESCAPE_SLASH_RE.replace_all(&key, r"\/"); if let Ok(regex) = Regex::new(&format!("^({key})$")) && let Ok(true) = regex.is_match(self.model().name()) { request_data.apply_patch(patch); return; } } } } impl Default for ClientConfig { fn default() -> Self { Self::OpenAIConfig(OpenAIConfig::default()) } } #[derive(Debug, Clone, Deserialize, Default)] pub struct ExtraConfig { pub proxy: Option, pub connect_timeout: Option, pub read_timeout: Option, } #[derive(Debug, Clone, Deserialize, Default)] pub struct RequestPatch { pub chat_completions: Option, pub embeddings: Option, pub rerank: Option, } pub type ApiPatch = IndexMap; pub struct RequestData { pub url: String, pub headers: IndexMap, pub body: Value, } impl RequestData { pub fn new(url: T, body: Value) -> Self where T: std::fmt::Display, { Self { url: url.to_string(), headers: Default::default(), body, } } pub fn bearer_auth(&mut self, auth: T) where T: std::fmt::Display, { self.headers .insert("authorization".into(), format!("Bearer {auth}")); } pub fn header(&mut self, key: K, value: V) where K: std::fmt::Display, V: std::fmt::Display, { self.headers.insert(key.to_string(), value.to_string()); } pub fn into_builder(self, client: &ReqwestClient) -> RequestBuilder { let RequestData { url, headers, body } = self; debug!("Request {url} {body}"); let mut builder = client.post(url); for (key, value) in headers { builder = builder.header(key, value); } builder = builder.json(&body); builder } pub fn apply_patch(&mut self, patch: Value) { if let Some(patch_url) = patch["url"].as_str() { self.url = patch_url.into(); } if let Some(patch_body) = patch.get("body") { json_patch::merge(&mut self.body, patch_body) } if let Some(patch_headers) = patch["headers"].as_object() { for (key, value) in patch_headers { if let Some(value) = value.as_str() { self.header(key, value) } else if value.is_null() { self.headers.swap_remove(key); } } } } } #[derive(Debug)] pub struct ChatCompletionsData { pub messages: Vec, pub temperature: Option, pub top_p: Option, pub reasoning_effort: Option, pub functions: Option>, pub stream: bool, } #[derive(Debug, Clone, Default)] pub struct ChatCompletionsOutput { pub text: String, pub tool_calls: Vec, pub thinking: Vec, } impl ChatCompletionsOutput { pub fn new(text: &str) -> Self { Self { text: text.to_string(), ..Default::default() } } } #[derive(Debug)] pub struct EmbeddingsData { pub texts: Vec, pub query: bool, } impl EmbeddingsData { pub fn new(texts: Vec, query: bool) -> Self { Self { texts, query } } } pub type EmbeddingsOutput = Vec>; #[derive(Debug)] pub struct RerankData { pub query: String, pub documents: Vec, pub top_n: usize, } impl RerankData { pub fn new(query: String, documents: Vec, top_n: usize) -> Self { Self { query, documents, top_n, } } } pub type RerankOutput = Vec; #[derive(Debug, Deserialize)] pub struct RerankResult { pub index: usize, } pub type PromptAction<'a> = (&'a str, &'a str, Option<&'a str>, bool); pub async fn create_config( prompts: &[PromptAction<'static>], client: &str, vault: &Vault, ) -> Result<(String, Value)> { let mut config = json!({ "type": client, }); for (key, desc, help_message, is_secret) in prompts { let env_name = format!("{client}-{key}") .to_ascii_uppercase() .replace("_", "-"); let required = std::env::var(&env_name).is_err(); let value = if !is_secret { prompt_input_string(desc, required, *help_message)? } else { vault.add_secret(&env_name)?; format!("{{{{{}}}}}", env_name) }; if !value.is_empty() { config[key] = value.into(); } } let model = set_client_models_config(&mut config, client).await?; let clients = json!(vec![config]); Ok((model, clients)) } pub async fn create_openai_compatible_client_config( client: &str, ) -> Result> { let api_base = OPENAI_COMPATIBLE_PROVIDERS .into_iter() .find(|(name, _)| client == *name) .map(|(_, api_base)| api_base) .unwrap_or("http(s)://{API_ADDR}/v1"); let name = if client == OpenAICompatibleClient::NAME { let value = prompt_input_string("Provider Name", true, None)?; value.replace(' ', "-") } else { client.to_string() }; let mut config = json!({ "type": OpenAICompatibleClient::NAME, "name": &name, }); let api_base = if api_base.contains('{') { prompt_input_string("API Base", true, Some(&format!("e.g. {api_base}")))? } else { api_base.to_string() }; config["api_base"] = api_base.into(); let has_bundled_oauth = ALL_PROVIDER_MODELS .iter() .any(|p| p.provider == client && p.oauth.is_some()); let use_oauth = if has_bundled_oauth { let choice = Select::new("Authentication method:", vec!["API Key", "OAuth"]).prompt()?; choice == "OAuth" } else { false }; if use_oauth { config["auth"] = "oauth".into(); } else { let api_key = prompt_input_string("API Key", false, None)?; if !api_key.is_empty() { config["api_key"] = api_key.into(); } } let model = set_client_models_config(&mut config, &name).await?; let clients = json!(vec![config]); Ok(Some((model, clients))) } pub async fn call_chat_completions( input: &Input, print: bool, extract_code: bool, client: &dyn Client, ctx: &mut RequestContext, abort_signal: AbortSignal, ) -> Result<(String, Vec)> { let is_child_agent = ctx.current_depth > 0; let suppress_spinner = is_child_agent || ctx.render_mode == RenderMode::Silent; let spinner_message = if suppress_spinner { "" } else { "Generating" }; let ret = abortable_run_with_spinner( client.chat_completions(input.clone()), spinner_message, abort_signal, ) .await; match ret { Ok(ret) => { let ChatCompletionsOutput { mut text, tool_calls, thinking, .. } = ret; if !text.is_empty() { if extract_code { text = extract_code_block(&strip_think_tag(&text)).to_string(); } if print { ctx.app.config.print_markdown(&text)?; } } let mut tool_results = eval_tool_calls(ctx, tool_calls).await?; if let Some(first) = tool_results.first_mut() { first.thinking = thinking; } tool_results .iter() .for_each(|res| ctx.tool_scope.tool_tracker.record_call(res.call.clone())); Ok((text, tool_results)) } Err(err) => Err(err), } } pub async fn call_chat_completions_streaming( input: &Input, client: &dyn Client, ctx: &mut RequestContext, abort_signal: AbortSignal, ) -> Result<(String, Vec)> { let (tx, rx) = unbounded_channel(); let mut handler = SseHandler::new(tx, abort_signal.clone()); let silent = ctx.render_mode == RenderMode::Silent; if silent { handler.set_silent(true); } let (send_ret, render_ret) = tokio::join!( client.chat_completions_streaming(input, &mut handler), render_stream(rx, client.app_config(), abort_signal.clone(), silent), ); let aborted_ctrlc = handler.abort().aborted_ctrlc(); let aborted_ctrld = handler.abort().aborted_ctrld(); if aborted_ctrld { bail!("Aborted."); } render_ret?; let (text, tool_calls, thinking) = handler.take(); if aborted_ctrlc { if !ctx.working_mode.is_repl() || ctx.session.is_none() { bail!("Aborted."); } if text.is_empty() { if !silent && *IS_STDOUT_TERMINAL { println!(); eprintln!("{}", error_text("Response interrupted")); } return Ok(("".to_string(), vec![])); } if !silent && *IS_STDOUT_TERMINAL { println!(); eprintln!("{}", error_text("Response interrupted")); } return Ok((text, vec![])); } match send_ret { Ok(_) => { if !silent && !text.is_empty() && !text.ends_with('\n') { println!(); } let mut tool_results = eval_tool_calls(ctx, tool_calls).await?; if let Some(first) = tool_results.first_mut() { first.thinking = thinking; } tool_results .iter() .for_each(|res| ctx.tool_scope.tool_tracker.record_call(res.call.clone())); Ok((text, tool_results)) } Err(err) => { if !silent && !text.is_empty() { println!(); } Err(err) } } } pub fn noop_prepare_rerank(_client: &T, _data: &RerankData) -> Result { bail!("The client doesn't support rerank api") } pub async fn noop_rerank(_builder: RequestBuilder, _model: &Model) -> Result { bail!("The client doesn't support rerank api") } #[derive(Debug)] pub struct ApiStatusError { pub status: u16, pub message: String, } impl std::fmt::Display for ApiStatusError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "{}", self.message) } } impl std::error::Error for ApiStatusError {} /// True when the error chain bottoms out in an [`ApiStatusError`] with /// status 401 EXACTLY. 403 (entitlement) and 429 (rate limit) are never /// auth failures, and message text is never inspected. fn is_auth_error(err: &anyhow::Error) -> bool { err.downcast_ref::() .is_some_and(|api_err| api_err.status == 401) } /// Decides whether a 401 from `client_name` warrants a single retry after a /// forced token refresh: the error must be a 401 [`ApiStatusError`], and the /// client must have a cached access token to distrust (API-key clients have /// none and never retry). Distrusting marks the exact rejected token so the /// retry's prepare step force-refreshes it. There is deliberately no backoff: /// the blast radius is bounded at one extra request per user-visible call. /// /// Note: vertexai shares the ACCESS_TOKENS cache, so a 401 there also /// triggers distrust+retry — deliberate. fn should_retry_auth(err: &anyhow::Error, client_name: &str) -> bool { if !is_auth_error(err) { return false; } let Ok(token) = get_access_token(client_name) else { return false; }; distrust_access_token(client_name, &token) } pub fn catch_error(data: &Value, status: u16) -> Result<()> { if (200..300).contains(&status) { return Ok(()); } debug!("Invalid response, status: {status}, data: {data}"); let api_error = |message: String| anyhow::Error::new(ApiStatusError { status, message }); if let Some(error) = data["error"].as_object() { if let (Some(typ), Some(message)) = ( json_str_from_map(error, "type"), json_str_from_map(error, "message"), ) { return Err(api_error(format!("{message} (type: {typ})"))); } else if let (Some(typ), Some(message)) = ( json_str_from_map(error, "code"), json_str_from_map(error, "message"), ) { return Err(api_error(format!("{message} (code: {typ})"))); } } else if let Some(error) = data["errors"][0].as_object() { if let (Some(code), Some(message)) = ( error.get("code").and_then(|v| v.as_u64()), json_str_from_map(error, "message"), ) { return Err(api_error(format!("{message} (status: {code})"))); } } else if let Some(error) = data[0]["error"].as_object() { if let (Some(status), Some(message)) = ( json_str_from_map(error, "status"), json_str_from_map(error, "message"), ) { return Err(api_error(format!("{message} (status: {status})"))); } } else if let (Some(detail), Some(status)) = (data["detail"].as_str(), data["status"].as_i64()) { return Err(api_error(format!("{detail} (status: {status})"))); } else if let Some(error) = data["error"].as_str() { return Err(api_error(error.to_string())); } else if let Some(message) = data["message"].as_str() { return Err(api_error(message.to_string())); } Err(api_error(format!( "Invalid response data: {data} (status: {status})" ))) } pub fn json_str_from_map<'a>( map: &'a serde_json::Map, field_name: &str, ) -> Option<&'a str> { map.get(field_name).and_then(|v| v.as_str()) } pub async fn set_client_models_config(client_config: &mut Value, client: &str) -> Result { if let Some(provider) = ALL_PROVIDER_MODELS.iter().find(|v| v.provider == client) { let models: Vec = provider .models .iter() .filter(|v| v.model_type == "chat") .map(|v| v.name.clone()) .collect(); let model_name = select_model(models)?; return Ok(format!("{client}:{model_name}")); } let mut model_names = vec![]; if let (Some(true), Some(api_base), api_key) = ( client_config["type"] .as_str() .map(|v| v == OpenAICompatibleClient::NAME), client_config["api_base"].as_str(), client_config["api_key"] .as_str() .map(|v| v.to_string()) .or_else(|| { let env_name = format!("{client}_api_key").to_ascii_uppercase(); std::env::var(&env_name).ok() }), ) { match abortable_run_with_spinner( fetch_models(api_base, api_key.as_deref()), "Fetching models", create_abort_signal(), ) .await { Ok(fetched_models) => { model_names = MultiSelect::new("LLMs to include (required):", fetched_models) .with_validator(|list: &[ListOption<&String>]| { if list.is_empty() { Ok(Validation::Invalid( "At least one item must be selected".into(), )) } else { Ok(Validation::Valid) } }) .prompt()?; } Err(err) => { eprintln!("✗ Fetch models failed: {err}"); } } } if model_names.is_empty() { model_names = prompt_input_string( "LLMs to add", true, Some("Separated by commas, e.g. llama3.3,qwen2.5"), )? .split(',') .filter_map(|v| { let v = v.trim(); if v.is_empty() { None } else { Some(v.to_string()) } }) .collect::>(); } if model_names.is_empty() { bail!("No models"); } let models: Vec = model_names .iter() .map(|v| { let l = v.to_lowercase(); if l.contains("rank") { json!({ "name": v, "type": "reranker", }) } else if let Ok(true) = EMBEDDING_MODEL_RE.is_match(&l) { json!({ "name": v, "type": "embedding", "default_chunk_size": 1000, "max_batch_size": 100 }) } else if v.contains("vision") { json!({ "name": v, "supports_vision": true }) } else { json!({ "name": v, }) } }) .collect(); client_config["models"] = models.into(); let model_name = select_model(model_names)?; Ok(format!("{client}:{model_name}")) } fn select_model(model_names: Vec) -> Result { if model_names.is_empty() { bail!("No models"); } let model = if model_names.len() == 1 { model_names[0].clone() } else { Select::new("Default Model (required):", model_names).prompt()? }; Ok(model) } fn prompt_input_string(desc: &str, required: bool, help_message: Option<&str>) -> Result { let desc = if required { format!("{desc} (required):") } else { format!("{desc} (optional):") }; let mut text = Text::new(&desc); if required { text = text.with_validator(required!("This field is required")) } if let Some(help_message) = help_message { text = text.with_help_message(help_message); } let text = text.prompt()?; Ok(text) } #[cfg(test)] mod tests { use super::*; use super::super::access_token::{is_rejected, set_access_token}; fn catch_error_message(data: &Value, status: u16) -> String { catch_error(data, status).unwrap_err().to_string() } #[test] fn test_catch_error_display_json_with_type() { let data = json!({"error": {"type": "invalid_request_error", "message": "Bad request"}}); assert_eq!( catch_error_message(&data, 400), "Bad request (type: invalid_request_error)" ); } #[test] fn test_catch_error_display_json_with_code() { let data = json!({"error": {"code": "rate_limited", "message": "Too many requests"}}); assert_eq!( catch_error_message(&data, 429), "Too many requests (code: rate_limited)" ); } #[test] fn test_catch_error_display_errors_array() { let data = json!({"errors": [{"code": 7000, "message": "No route"}]}); assert_eq!(catch_error_message(&data, 404), "No route (status: 7000)"); } #[test] fn test_catch_error_display_array_error_status() { let data = json!([{"error": {"status": "PERMISSION_DENIED", "message": "Denied"}}]); assert_eq!( catch_error_message(&data, 403), "Denied (status: PERMISSION_DENIED)" ); } #[test] fn test_catch_error_display_detail_status() { let data = json!({"detail": "Not found", "status": 404}); assert_eq!(catch_error_message(&data, 404), "Not found (status: 404)"); } #[test] fn test_catch_error_display_error_string() { let data = json!({"error": "Something went wrong"}); assert_eq!(catch_error_message(&data, 500), "Something went wrong"); } #[test] fn test_catch_error_display_message_string() { let data = json!({"message": "Unauthorized"}); assert_eq!(catch_error_message(&data, 401), "Unauthorized"); } #[test] fn test_catch_error_display_fallback() { let data = json!({"unexpected": true}); assert_eq!( catch_error_message(&data, 500), format!("Invalid response data: {data} (status: 500)") ); } #[test] fn test_catch_error_ok_on_success_status() { let data = json!({"error": {"type": "x", "message": "y"}}); assert!(catch_error(&data, 200).is_ok()); assert!(catch_error(&data, 299).is_ok()); } #[test] fn test_catch_error_downcast_through_context_chain() { let data = json!({"error": {"type": "authentication_error", "message": "Invalid key"}}); let err = catch_error(&data, 401) .context("Failed to call chat-completions api") .unwrap_err(); let api_err = err .downcast_ref::() .expect("should downcast through context chain"); assert_eq!(api_err.status, 401); assert_eq!(api_err.message, "Invalid key (type: authentication_error)"); } #[test] fn test_catch_error_preserves_status() { let data = json!({"message": "Unauthorized"}); let err = catch_error(&data, 401).unwrap_err(); assert_eq!(err.downcast_ref::().unwrap().status, 401); let data = json!({"detail": "Rate limited", "status": 429}); let err = catch_error(&data, 429).unwrap_err(); assert_eq!(err.downcast_ref::().unwrap().status, 429); // The struct carries the outer HTTP status even when the body embeds another code let data = json!({"errors": [{"code": 7000, "message": "No route"}]}); let err = catch_error(&data, 429).unwrap_err(); assert_eq!(err.downcast_ref::().unwrap().status, 429); } /// Wrapped in `.context(...)` so every test below proves the downcast /// works through an anyhow context chain, as in the trait methods. fn api_status_error(status: u16) -> anyhow::Error { anyhow::Error::new(ApiStatusError { status, message: format!("error (status: {status})"), }) .context("Failed to call chat-completions api") } fn cache_token(client: &str, token: &str) { set_access_token( client, token.into(), chrono::Utc::now().timestamp() + 3600, None, ); } #[test] fn test_should_retry_auth_401_with_cached_token() { let client = "should-retry-auth-401"; cache_token(client, "at-1"); assert!(should_retry_auth(&api_status_error(401), client)); assert!(is_rejected(client, "at-1"), "rejected marker not set"); } #[test] fn test_should_retry_auth_non_401_statuses() { let client = "should-retry-auth-non-401"; cache_token(client, "at-1"); for status in [403, 429, 500] { assert!( !should_retry_auth(&api_status_error(status), client), "retried on {status}" ); } assert_eq!(get_access_token(client).unwrap(), "at-1"); assert!(!is_rejected(client, "at-1"), "marker set without a 401"); } #[test] fn test_should_retry_auth_non_api_status_error() { let client = "should-retry-auth-non-api"; cache_token(client, "at-1"); let err = anyhow::anyhow!("connection reset").context("Failed to call embeddings api"); assert!(!should_retry_auth(&err, client)); assert_eq!(get_access_token(client).unwrap(), "at-1"); assert!(!is_rejected(client, "at-1")); } #[test] fn test_should_retry_auth_401_without_cached_token() { let client = "should-retry-auth-no-token"; assert!(!should_retry_auth(&api_status_error(401), client)); assert!(!is_rejected(client, "at-1")); } }