Skip to main content

anda_engine/
model.rs

1//! Model provider integration and label-based routing.
2//!
3//! This module adapts provider-specific completion APIs to the common
4//! [`CompletionRequest`] and [`AgentOutput`] contract used by Anda agents.
5//! Built-in providers currently include OpenAI-compatible APIs, Anthropic, and
6//! Google Gemini.
7//!
8//! The [`Models`] registry maps model labels such as `primary`, `pro`,
9//! `flash`, or `lite` to concrete [`Model`] instances. Labels let agents
10//! request capability tiers without hard-coding provider model names.
11//!
12//! Custom providers can implement [`CompletionFeaturesDyn`] and be wrapped with
13//! [`Model::with_completer`].
14
15use anda_core::{
16    AgentOutput, BoxError, BoxPinFut, CONTENT_TYPE_JSON, CompletionRequest, Json, ToolCall,
17};
18use arc_swap::ArcSwap;
19use futures_util::StreamExt;
20use serde::de::DeserializeOwned;
21use serde::{Deserialize, Serialize};
22use std::future::Future;
23use std::time::{Duration, Instant};
24use std::{
25    collections::{BTreeSet, HashMap, hash_map::Entry},
26    error::Error,
27    fmt,
28    sync::Arc,
29};
30
31pub mod anthropic;
32pub(crate) mod driver;
33pub mod gemini;
34pub mod openai;
35pub(crate) mod raw;
36#[cfg(test)]
37pub(crate) mod test_support;
38pub mod testing;
39
40/// Deserializes a JSON `null` as the type's default value, so providers that
41/// send explicit nulls for optional counters/objects do not fail typed parses.
42pub(crate) fn null_default<'de, D, T>(deserializer: D) -> Result<T, D::Error>
43where
44    D: serde::Deserializer<'de>,
45    T: Default + Deserialize<'de>,
46{
47    Ok(Option::<T>::deserialize(deserializer)?.unwrap_or_default())
48}
49
50/// Generates `as_str`, `Serialize`, and `Deserialize` for an open string enum:
51/// known wire values map to variants (extra input-only aliases may follow the
52/// canonical value with `|`), and unknown values are preserved in the given
53/// catch-all variant instead of failing the parse.
54macro_rules! string_enum_serde {
55    ($ty:ident, { $($wire:literal $(| $alias:literal)* => $variant:ident),+ $(,)? }, $unknown:ident) => {
56        impl $ty {
57            fn as_str(&self) -> &str {
58                match self {
59                    $(Self::$variant => $wire,)+
60                    Self::$unknown(value) => value.as_str(),
61                }
62            }
63        }
64
65        impl serde::Serialize for $ty {
66            fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
67            where
68                S: serde::Serializer,
69            {
70                serializer.serialize_str(self.as_str())
71            }
72        }
73
74        impl<'de> serde::Deserialize<'de> for $ty {
75            fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
76            where
77                D: serde::Deserializer<'de>,
78            {
79                let value = String::deserialize(deserializer)?;
80                Ok(match value.as_str() {
81                    $($wire $(| $alias)* => Self::$variant,)+
82                    _ => Self::$unknown(value),
83                })
84            }
85        }
86    };
87}
88pub(crate) use string_enum_serde;
89
90/// Resolves a configured endpoint, falling back to the provider default when
91/// the endpoint is absent or empty.
92pub(crate) fn resolve_endpoint(endpoint: Option<String>, default: &str) -> String {
93    match endpoint {
94        Some(endpoint) if !endpoint.is_empty() => endpoint,
95        _ => default.to_string(),
96    }
97}
98
99pub use reqwest;
100pub use reqwest::Proxy;
101
102use crate::APP_USER_AGENT;
103
104pub use anda_core::ModelEffort;
105
106const MODEL_REQUEST_MAX_RETRIES: usize = 3;
107const MODEL_RETRY_BACKOFF: Duration = Duration::from_secs(1);
108const MODEL_RETRY_MAX_BACKOFF: Duration = Duration::from_secs(300);
109const COMPLETION_HTTP2_KEEP_ALIVE_INTERVAL: Option<Duration> = None;
110const COMPLETION_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
111const COMPLETION_READ_TIMEOUT: Duration = Duration::from_secs(180);
112const COMPLETION_REQUEST_TIMEOUT: Duration = Duration::from_secs(600);
113
114/// Upper bound on the number of bytes buffered from a completion response.
115///
116/// Guards against a runaway or malicious provider streaming an unbounded body
117/// and exhausting memory. Text completion responses are far smaller than this
118/// in practice.
119const MAX_COMPLETION_RESPONSE_BYTES: usize = 64 * 1024 * 1024;
120
121/// Upper bound on the response-body excerpt embedded in a completion error.
122///
123/// Error messages end up in logs and in `failed_reason`, so a gateway's HTML
124/// error page or a malformed multi-megabyte body is cut to a diagnostic prefix.
125const MAX_ERROR_BODY_BYTES: usize = 8 * 1024;
126
127/// Returns a lossy UTF-8 excerpt of `data` bounded by [`MAX_ERROR_BODY_BYTES`],
128/// for embedding a provider response body in an error message.
129pub(crate) fn error_body_excerpt(data: &[u8]) -> String {
130    if data.len() <= MAX_ERROR_BODY_BYTES {
131        return String::from_utf8_lossy(data).into_owned();
132    }
133
134    let mut text = String::from_utf8_lossy(&data[..MAX_ERROR_BODY_BYTES]).into_owned();
135    text.push_str("… [truncated]");
136    text
137}
138
139/// Reads a non-success response body for diagnostics, stopping once
140/// [`MAX_ERROR_BODY_BYTES`] is exceeded instead of buffering the whole body.
141async fn read_error_body(response: reqwest::Response) -> Result<String, reqwest::Error> {
142    let mut stream = response.bytes_stream();
143    let mut body = Vec::new();
144    while let Some(chunk) = stream.next().await {
145        body.extend_from_slice(&chunk?);
146        if body.len() > MAX_ERROR_BODY_BYTES {
147            break;
148        }
149    }
150    Ok(error_body_excerpt(&body))
151}
152
153/// Serializable configuration for constructing a model adapter.
154///
155/// [`Debug`] is implemented manually so the `api_key` is redacted and never
156/// leaks through `{:?}` logging.
157#[derive(Default, Clone, Deserialize, Serialize)]
158pub struct ModelConfig {
159    /// Provider family, such as `gemini`, `anthropic`, `openai-response` or `openai`.
160    pub family: String,
161
162    /// Provider-specific model name.
163    pub model: String,
164
165    /// Base URL for the provider API.
166    pub api_base: String,
167
168    /// API key used by the provider adapter.
169    pub api_key: String,
170
171    /// Optional labels for selecting this model in the engine.
172    ///
173    /// If omitted, the provider model name is used as the only label. Common
174    /// labels include `primary`, `pro`, `flash`, `lite`, `audio`, `video`, `image`, `memory`.
175    #[serde(default)]
176    pub labels: Vec<String>,
177
178    #[serde(default)]
179    /// Provider context window in input tokens; `0` means unknown.
180    pub context_window: usize,
181
182    #[serde(default)]
183    /// Provider maximum output tokens; `0` means unknown.
184    ///
185    /// When known, it caps each request's explicit output budget (see
186    /// [`Model::completion`]) and replaces the default budget of adapters that
187    /// always send one (Anthropic `max_tokens`, Gemini `maxOutputTokens`).
188    pub max_output: usize,
189
190    /// Optional reasoning/thinking effort for providers and models that support it.
191    ///
192    /// Supported config values are `minimal`, `low`, `medium`, `high`, and `max`.
193    /// The effective set depends on the selected provider and model.
194    #[serde(default)]
195    pub effort: Option<ModelEffort>,
196
197    /// Skips this model when loading a list of configs.
198    #[serde(default)]
199    pub disabled: bool,
200
201    /// Sends Anthropic credentials with bearer authentication instead of the
202    /// provider-specific API-key header.
203    #[serde(default)]
204    pub bearer_auth: bool,
205
206    #[serde(default)]
207    /// Whether to request streaming completions from this model.
208    /// The OpenAI Responses adapter always streams and aggregates the result,
209    /// regardless of this setting.
210    pub stream: bool,
211}
212
213impl std::fmt::Debug for ModelConfig {
214    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
215        f.debug_struct("ModelConfig")
216            .field("family", &self.family)
217            .field("model", &self.model)
218            .field("api_base", &self.api_base)
219            .field(
220                "api_key",
221                &if self.api_key.is_empty() {
222                    ""
223                } else {
224                    "[REDACTED]"
225                },
226            )
227            .field("labels", &self.labels)
228            .field("context_window", &self.context_window)
229            .field("max_output", &self.max_output)
230            .field("effort", &self.effort)
231            .field("disabled", &self.disabled)
232            .field("bearer_auth", &self.bearer_auth)
233            .field("stream", &self.stream)
234            .finish()
235    }
236}
237
238impl ModelConfig {
239    /// Builds a [`Model`] from this configuration.
240    pub fn model(&self, http_client: reqwest::Client) -> Result<Model, BoxError> {
241        if self.disabled {
242            return Err("model is disabled".into());
243        }
244        if self.model.is_empty() {
245            return Err(format!("{}: model name is required", self.model).into());
246        }
247        if self.family.is_empty() {
248            return Err(format!("{}: model family is required", self.model).into());
249        }
250        if self.api_base.is_empty() {
251            return Err(format!("{}: api_base is required", self.model).into());
252        }
253        if self.api_key.is_empty() {
254            return Err(format!("{}: api_key is required", self.model).into());
255        }
256
257        let mut model = match self.family.as_str() {
258            "gemini" => Model::with_completer(Arc::new(
259                gemini::Client::new_with_client(
260                    &self.api_key,
261                    Some(self.api_base.clone()),
262                    http_client,
263                )
264                .completion_model(&self.model)
265                .with_stream(self.stream)
266                .with_effort(self.effort)
267                .with_max_output(self.max_output),
268            )),
269            "anthropic" => {
270                let mut cli = anthropic::Client::new_with_client(
271                    &self.api_key,
272                    Some(self.api_base.clone()),
273                    http_client,
274                );
275                if self.bearer_auth {
276                    cli = cli.with_bearer_auth(true);
277                }
278                Model::with_completer(Arc::new(
279                    cli.completion_model(&self.model)
280                        .with_stream(self.stream)
281                        .with_effort(self.effort)
282                        .with_max_output(self.max_output),
283                ))
284            }
285            "openai-response" => {
286                let cli = openai::Client::new_with_client(
287                    &self.api_key,
288                    Some(self.api_base.clone()),
289                    http_client,
290                );
291                Model::with_completer(Arc::new(
292                    cli.completion_model_v2(&self.model)
293                        .with_stream(self.stream)
294                        .with_effort(self.effort),
295                ))
296            }
297            "openai" => {
298                let cli = openai::Client::new_with_client(
299                    &self.api_key,
300                    Some(self.api_base.clone()),
301                    http_client,
302                );
303                if self.model.starts_with("gpt") {
304                    Model::with_completer(Arc::new(
305                        cli.completion_model_v2(&self.model)
306                            .with_stream(self.stream)
307                            .with_effort(self.effort),
308                    ))
309                } else {
310                    Model::with_completer(Arc::new(
311                        cli.completion_model(&self.model)
312                            .with_stream(self.stream)
313                            .with_effort(self.effort),
314                    ))
315                }
316            }
317            _ => return Err(format!("unsupported model family: {}", self.family).into()),
318        };
319
320        let labels = if self.labels.is_empty() {
321            vec![self.model.to_ascii_lowercase()]
322        } else {
323            self.labels.clone()
324        };
325        model.context_window = self.context_window;
326        model.max_output = self.max_output;
327        Ok(model.with_labels(labels))
328    }
329}
330
331/// Thread-safe model registry used by the engine.
332///
333/// It maintains two layers:
334/// - `model`: the primary default model for general requests
335/// - `models`: a label-based map for selecting specific models
336///
337/// The dedicated primary slot can be set explicitly via [`Models::set_model`]
338/// or derived from the special label `primary` in the label map. This keeps
339/// direct lookup (`get`) separate from default-routing (`get_model`).
340pub struct Models {
341    model: ArcSwap<Option<Model>>,
342    models: ArcSwap<HashMap<String, Vec<Model>>>,
343}
344
345impl Default for Models {
346    fn default() -> Self {
347        Self {
348            model: ArcSwap::new(Arc::new(None)),
349            models: ArcSwap::new(Arc::new(HashMap::new())),
350        }
351    }
352}
353
354impl Models {
355    /// Creates a new Models instance by cloning the internal state of another Models instance.
356    pub fn from_clone(other: &Models) -> Self {
357        Self {
358            model: ArcSwap::new(other.model.load_full()),
359            models: ArcSwap::new(Arc::new(other.models.load().as_ref().clone())),
360        }
361    }
362
363    /// Replaces this registry with a clone of another [`Models`] instance.
364    pub fn replace(&self, other: &Models) {
365        self.model.store(other.model.load_full());
366        self.models
367            .store(Arc::new(other.models.load().as_ref().clone()));
368    }
369
370    /// Builds a registry from model configs by registering every resolved label.
371    ///
372    /// Configs explicitly marked `disabled` are skipped quietly (info level);
373    /// configs that fail to build because of a misconfiguration are skipped with
374    /// a warning so the problem is not lost silently.
375    pub fn from_configs(configs: &[ModelConfig], http_client: reqwest::Client) -> Self {
376        let models = Self::default();
377        for config in configs {
378            if config.disabled {
379                log::info!(
380                    "skipping disabled model: family={}, model={}",
381                    config.family,
382                    config.model
383                );
384                continue;
385            }
386            match config.model(http_client.clone()) {
387                Ok(model) => models.inner_set(model.labels.clone(), model),
388                Err(err) => {
389                    log::warn!(
390                        "skipping misconfigured model: family={}, model={}, error={}",
391                        config.family,
392                        config.model,
393                        err
394                    );
395                }
396            }
397        }
398        models
399    }
400
401    /// Returns whether a label exists in the direct lookup table.
402    pub fn contains(&self, label: &str) -> bool {
403        self.models.load().contains_key(&label.to_ascii_lowercase())
404    }
405
406    /// Returns the set of all registered model names across all labels.
407    pub fn model_names(&self) -> BTreeSet<String> {
408        self.models
409            .load()
410            .values()
411            .flatten()
412            .map(|m| m.model_name())
413            .collect()
414    }
415
416    /// Sets the primary default model without mutating the label map.
417    pub fn set_model(&self, model: Model) {
418        self.inner_set(model.labels.clone(), model.clone());
419        self.model.store(Arc::new(Some(model)));
420    }
421
422    /// Inserts or updates a single labeled model.
423    ///
424    /// The special label `primary` also updates the dedicated routing slot.
425    /// If no primary exists yet, any inserted model is promoted
426    /// to become the primary default.
427    pub fn set(&self, label: String, model: Model) {
428        self.inner_set(vec![label], model);
429    }
430
431    fn inner_set(&self, mut labels: Vec<String>, model: Model) {
432        if self.model.load().is_none() {
433            self.model.store(Arc::new(Some(model.clone())));
434        }
435
436        let model_name = model.model_name();
437        labels.push(model_name.to_ascii_lowercase());
438        for label in labels.iter_mut() {
439            label.make_ascii_lowercase();
440            if label == "primary" {
441                self.model.store(Arc::new(Some(model.clone())));
442            }
443        }
444
445        // rcu keeps concurrent inserts from losing each other's labels.
446        self.models.rcu(|models| {
447            let mut models = models.as_ref().clone();
448            for label in &labels {
449                match models.entry(label.clone()) {
450                    Entry::Vacant(e) => {
451                        e.insert(vec![model.clone()]);
452                    }
453                    Entry::Occupied(mut e) => {
454                        e.get_mut().retain(|m| m.model_name() != model_name);
455                        e.get_mut().push(model.clone());
456                    }
457                }
458            }
459            models
460        });
461    }
462
463    /// Returns a model by lowercase label if it exists.
464    ///
465    /// This is a direct lookup only and never falls back to default routing.
466    pub fn get(&self, label: &str) -> Option<Model> {
467        self.models
468            .load()
469            .get(&label.to_ascii_lowercase())
470            .and_then(|v| v.last().cloned())
471    }
472
473    /// Returns the primary model if available; otherwise returns any remaining
474    /// labeled model.
475    pub fn get_model(&self) -> Option<Model> {
476        if let Some(m) = self.model.load().as_ref() {
477            return Some(m.clone());
478        }
479        self.models
480            .load()
481            .values()
482            .next()
483            .and_then(|v| v.last().cloned())
484    }
485
486    /// Resolves a model for lowercase-label-aware routing.
487    ///
488    /// Resolution order is:
489    /// - the exact label match when `label` is non-empty
490    /// - the default routing result from [`Models::get_model`]
491    pub fn resolve(&self, label: &str) -> Option<Model> {
492        if label.is_empty() {
493            return self.get_model();
494        }
495        self.get(label).or_else(|| self.get_model())
496    }
497}
498
499/// Object-safe completion provider interface.
500pub trait CompletionFeaturesDyn: Send + Sync + 'static {
501    /// Performs a completion request and returns the agent-facing output.
502    ///
503    /// Built-in adapters wrap exhausted transient provider failures in
504    /// [`ModelError`]. Use [`is_retryable_box_error`] to decide whether an upper
505    /// layer should schedule a delayed retry.
506    fn completion(&self, req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>>;
507
508    /// Returns the provider model name used for diagnostics and usage reports.
509    fn model_name(&self) -> String;
510
511    /// Removes unanswered tool-call requests from `raw_history[start..]`.
512    ///
513    /// `raw_history` holds this provider's own message JSON, so only the
514    /// provider can classify its items reliably. The runner calls this after an
515    /// interrupt (steering, discard, stop) so the next request does not inherit
516    /// a tool-call requirement the provider would reject as unanswered. Visible
517    /// text and reasoning must stay untouched.
518    ///
519    /// The default implementation is a conservative union over the built-in
520    /// provider wire shapes; override it when your provider's raw items are not
521    /// covered by it.
522    fn prune_unanswered_tool_calls(&self, raw_history: &mut Vec<Json>, start: usize) {
523        raw::prune_unanswered_tool_calls(raw_history, start);
524    }
525
526    /// Removes completed tool interactions (calls and results) from `raw_history`.
527    ///
528    /// Long-lived callers invoke this at an idle boundary to reclaim
529    /// context-window budget from tool payloads the model has already consumed.
530    /// Visible text and reasoning must stay untouched, and provider ordering
531    /// constraints (for example a reasoning item that must precede its sibling)
532    /// must be preserved.
533    ///
534    /// The default implementation is a conservative union over the built-in
535    /// provider wire shapes; override it when your provider's raw items are not
536    /// covered by it.
537    fn prune_tool_interactions(&self, raw_history: &mut Vec<Json>) {
538        raw::prune_tool_interactions(raw_history);
539    }
540}
541
542/// Placeholder implementation that returns errors for completion requests.
543#[derive(Clone, Debug)]
544pub struct NotImplemented;
545
546impl CompletionFeaturesDyn for NotImplemented {
547    fn model_name(&self) -> String {
548        "not_implemented".to_string()
549    }
550
551    fn completion(&self, _req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>> {
552        Box::pin(futures::future::ready(Err("not implemented".into())))
553    }
554}
555
556/// Mock implementation for tests and examples.
557#[derive(Clone, Debug)]
558pub struct MockImplemented;
559
560impl CompletionFeaturesDyn for MockImplemented {
561    fn model_name(&self) -> String {
562        "mock_implemented".to_string()
563    }
564
565    fn completion(&self, req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>> {
566        Box::pin(futures::future::ready(Ok(AgentOutput {
567            content: req.prompt.clone(),
568            tool_calls: req
569                .tools
570                .iter()
571                .filter_map(|tool| {
572                    if req.prompt.is_empty() {
573                        return None;
574                    }
575                    Some(ToolCall {
576                        name: tool.name.clone(),
577                        args: serde_json::from_str(&req.prompt).unwrap_or_default(),
578                        call_id: None,
579                        result: None,
580                        remote_id: None,
581                    })
582                })
583                .collect(),
584            ..Default::default()
585        })))
586    }
587}
588
589/// Concrete model entry registered with the engine.
590#[derive(Clone)]
591pub struct Model {
592    /// Completion provider implementation.
593    pub completer: Arc<dyn CompletionFeaturesDyn>,
594
595    /// Labels that can route requests to this model.
596    pub labels: Vec<String>,
597
598    /// Context window in input tokens; `0` means unknown.
599    pub context_window: usize,
600
601    /// Maximum output tokens; `0` means unknown. A known value caps each
602    /// request's explicit `max_output_tokens`.
603    pub max_output: usize,
604}
605
606impl Model {
607    /// Creates a model from a completion provider.
608    pub fn new(completer: Arc<dyn CompletionFeaturesDyn>) -> Self {
609        Self {
610            completer,
611            labels: Vec::new(),
612            context_window: 0,
613            max_output: 0,
614        }
615    }
616
617    /// Creates a model from a completion provider.
618    pub fn with_completer(completer: Arc<dyn CompletionFeaturesDyn>) -> Self {
619        Self::new(completer)
620    }
621
622    /// Assigns labels used by [`Models`] for routing.
623    pub fn with_labels(mut self, labels: Vec<String>) -> Self {
624        self.labels = labels;
625        self
626    }
627
628    /// Creates a model whose completion calls return `not implemented` errors.
629    pub fn not_implemented() -> Self {
630        Self::new(Arc::new(NotImplemented))
631    }
632
633    /// Creates a model with deterministic mock completion behavior for tests.
634    pub fn mock_implemented() -> Self {
635        Self::new(Arc::new(MockImplemented))
636    }
637
638    /// Returns the provider model name for this model.
639    pub fn model_name(&self) -> String {
640        self.completer.model_name()
641    }
642
643    /// Executes a completion request with the underlying provider.
644    ///
645    /// A known [`Model::max_output`] caps an explicit `max_output_tokens`, so a request never
646    /// asks the provider for more than the configured model accepts. An unset budget stays
647    /// unset: defaulting it is the provider adapter's decision (see
648    /// [`ModelConfig::model`]), since some APIs reject an output limit they do not need.
649    pub async fn completion(&self, mut req: CompletionRequest) -> Result<AgentOutput, BoxError> {
650        if self.max_output > 0
651            && let Some(tokens) = req.max_output_tokens.as_mut()
652        {
653            *tokens = (*tokens).min(self.max_output);
654        }
655        self.completer.completion(req).await
656    }
657
658    /// Removes unanswered tool-call requests from `raw_history[start..]` using
659    /// the provider's own wire-format knowledge.
660    ///
661    /// See [`CompletionFeaturesDyn::prune_unanswered_tool_calls`].
662    pub fn prune_unanswered_tool_calls(&self, raw_history: &mut Vec<Json>, start: usize) {
663        self.completer
664            .prune_unanswered_tool_calls(raw_history, start);
665    }
666
667    /// Removes completed tool interactions from `raw_history` using the
668    /// provider's own wire-format knowledge.
669    ///
670    /// See [`CompletionFeaturesDyn::prune_tool_interactions`].
671    pub fn prune_tool_interactions(&self, raw_history: &mut Vec<Json>) {
672        self.completer.prune_tool_interactions(raw_history);
673    }
674}
675
676/// Error returned by built-in model adapters when the caller can inspect retry
677/// semantics after the SDK-level retry has already been attempted.
678#[derive(Debug)]
679pub struct ModelError {
680    message: String,
681    retryable: bool,
682    status: Option<http::StatusCode>,
683    retry_after: Option<Duration>,
684    source: Option<BoxError>,
685}
686
687impl ModelError {
688    /// Creates a non-retryable model error.
689    pub fn new(message: impl Into<String>) -> Self {
690        Self {
691            message: message.into(),
692            retryable: false,
693            status: None,
694            retry_after: None,
695            source: None,
696        }
697    }
698
699    /// Marks whether this error should be considered retryable by the caller.
700    pub fn with_retryable(mut self, retryable: bool) -> Self {
701        self.retryable = retryable;
702        self
703    }
704
705    /// Attaches the upstream HTTP status, if the provider returned one.
706    pub fn with_status(mut self, status: http::StatusCode) -> Self {
707        self.status = Some(status);
708        self
709    }
710
711    /// Attaches a suggested retry delay from the upstream response.
712    pub fn with_retry_after(mut self, retry_after: Option<Duration>) -> Self {
713        self.retry_after = retry_after;
714        self
715    }
716
717    /// Attaches the lower-level transport/read error.
718    pub fn with_source(mut self, source: BoxError) -> Self {
719        self.source = Some(source);
720        self
721    }
722
723    /// Returns true when the upper layer may choose a delayed retry.
724    pub fn is_retryable(&self) -> bool {
725        self.retryable
726    }
727
728    /// Returns the upstream HTTP status, when present.
729    pub fn status(&self) -> Option<http::StatusCode> {
730        self.status
731    }
732
733    /// Returns the upstream retry delay, when present.
734    pub fn retry_after(&self) -> Option<Duration> {
735        self.retry_after
736    }
737}
738
739impl fmt::Display for ModelError {
740    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
741        f.write_str(&self.message)
742    }
743}
744
745impl Error for ModelError {
746    fn source(&self) -> Option<&(dyn Error + 'static)> {
747        self.source
748            .as_deref()
749            .map(|source| source as &(dyn Error + 'static))
750    }
751}
752
753/// Walks an error chain and returns the first value produced by `f`.
754fn find_in_error_chain<T>(
755    error: &(dyn Error + 'static),
756    f: impl Fn(&(dyn Error + 'static)) -> Option<T>,
757) -> Option<T> {
758    let mut current = Some(error);
759    while let Some(error) = current {
760        if let Some(found) = f(error) {
761            return Some(found);
762        }
763        current = error.source();
764    }
765    None
766}
767
768/// Returns true if the error chain carries a retryable model error signal.
769///
770/// This function and its siblings ([`is_retryable_box_error`],
771/// [`model_error_status`], [`model_error_retry_after`]) are the public
772/// inspection surface of [`ModelError`] for embedding applications: built-in
773/// adapters have already exhausted their own transport retries when they
774/// return one, and the engine deliberately does not retry again on its own —
775/// whether to schedule a delayed retry is an application-level policy.
776pub fn is_retryable_model_error(error: &(dyn Error + 'static)) -> bool {
777    find_in_error_chain(error, |error| {
778        let retryable = error
779            .downcast_ref::<ModelError>()
780            .is_some_and(ModelError::is_retryable)
781            || error
782                .downcast_ref::<reqwest::Error>()
783                .is_some_and(is_retryable_reqwest_error);
784        retryable.then_some(())
785    })
786    .is_some()
787}
788
789/// Convenience wrapper for callers that keep completion errors as [`BoxError`].
790pub fn is_retryable_box_error(error: &BoxError) -> bool {
791    is_retryable_model_error(error.as_ref() as &(dyn Error + 'static))
792}
793
794/// Returns the first [`ModelError`] HTTP status found in the error chain.
795pub fn model_error_status(error: &(dyn Error + 'static)) -> Option<http::StatusCode> {
796    find_in_error_chain(error, |error| error.downcast_ref::<ModelError>()?.status())
797}
798
799/// Returns the first upstream retry delay found in the error chain.
800pub fn model_error_retry_after(error: &(dyn Error + 'static)) -> Option<Duration> {
801    find_in_error_chain(error, |error| {
802        error.downcast_ref::<ModelError>()?.retry_after()
803    })
804}
805
806/// Statuses that are transient enough for one immediate SDK retry and for an
807/// upper-layer delayed retry after the SDK retry has been exhausted.
808pub fn is_retryable_status(status: http::StatusCode) -> bool {
809    matches!(
810        status,
811        http::StatusCode::REQUEST_TIMEOUT
812            | http::StatusCode::TOO_MANY_REQUESTS
813            | http::StatusCode::INTERNAL_SERVER_ERROR
814            | http::StatusCode::BAD_GATEWAY
815            | http::StatusCode::SERVICE_UNAVAILABLE
816            | http::StatusCode::GATEWAY_TIMEOUT
817    ) || status.as_u16() == 529
818}
819
820pub(crate) fn is_retryable_reqwest_error(err: &reqwest::Error) -> bool {
821    err.is_timeout()
822        || err.is_connect()
823        || err.is_body()
824        || err.is_decode()
825        || err.status().is_some_and(is_retryable_status)
826}
827
828/// Formats an error together with its source chain, e.g.
829/// "error decoding response body: request or response body error: operation
830/// timed out". reqwest 0.12.2+ stopped including sources in `Display`, so the
831/// top-level message alone hides the root cause of transport failures behind
832/// a generic phrase.
833pub(crate) fn format_error_chain(err: &(dyn Error + 'static)) -> String {
834    let mut message = err.to_string();
835    let mut source = err.source();
836    while let Some(err) = source {
837        let text = err.to_string();
838        // Some errors repeat their source in `Display`; skip duplicates.
839        if !message.contains(&text) {
840            message.push_str(": ");
841            message.push_str(&text);
842        }
843        source = err.source();
844    }
845    message
846}
847
848/// Extracts the upstream request id from response headers for error context.
849/// Covers the header names used by OpenAI-compatible APIs, Anthropic, and
850/// common gateway/CDN fronts.
851fn upstream_request_id(headers: &http::HeaderMap) -> Option<String> {
852    ["x-request-id", "request-id", "x-amzn-requestid", "cf-ray"]
853        .into_iter()
854        .find_map(|name| headers.get(name)?.to_str().ok())
855        .map(str::to_string)
856}
857
858pub(crate) fn completion_transport_error(
859    model: &str,
860    action: &str,
861    err: reqwest::Error,
862) -> BoxError {
863    let retryable = is_retryable_reqwest_error(&err);
864    let message = format!(
865        "{action}, model: {model}, error: {}",
866        format_error_chain(&err)
867    );
868    Box::new(
869        ModelError::new(message)
870            .with_retryable(retryable)
871            .with_source(Box::new(err)),
872    )
873}
874
875pub(crate) async fn read_completion_response_bytes(
876    response: reqwest::Response,
877    model: &str,
878) -> Result<bytes::Bytes, BoxError> {
879    let request_id = upstream_request_id(response.headers());
880    // Reject up front when the declared length already exceeds the cap.
881    if let Some(len) = response.content_length()
882        && len > MAX_COMPLETION_RESPONSE_BYTES as u64
883    {
884        return Err(completion_response_too_large(
885            model,
886            request_id.as_deref(),
887            len as usize,
888        ));
889    }
890
891    // Stream the body so the size cap is enforced even when no `Content-Length`
892    // is provided.
893    let mut stream = response.bytes_stream();
894    let mut body: Vec<u8> = Vec::new();
895    while let Some(chunk) = stream.next().await {
896        let chunk = chunk.map_err(|err| {
897            let action = format!(
898                "Failed to read completion response (request_id: {})",
899                request_id.as_deref().unwrap_or("-")
900            );
901            completion_transport_error(model, &action, err)
902        })?;
903        if body.len() + chunk.len() > MAX_COMPLETION_RESPONSE_BYTES {
904            return Err(completion_response_too_large(
905                model,
906                request_id.as_deref(),
907                body.len() + chunk.len(),
908            ));
909        }
910        body.extend_from_slice(&chunk);
911    }
912    Ok(body.into())
913}
914
915/// Builds the error returned when a completion response exceeds the size cap.
916fn completion_response_too_large(
917    model: &str,
918    request_id: Option<&str>,
919    received: usize,
920) -> BoxError {
921    Box::new(ModelError::new(format!(
922        "Completion response too large (model: {}, request_id: {}, received: {} bytes, limit: {} bytes)",
923        model,
924        request_id.unwrap_or("-"),
925        received,
926        MAX_COMPLETION_RESPONSE_BYTES,
927    )))
928}
929
930pub(crate) async fn execute_completion_request_with_retry<T, BuildRequest, HandleResponse, Fut>(
931    model: &str,
932    build_request: BuildRequest,
933    handle_response: HandleResponse,
934) -> Result<T, BoxError>
935where
936    BuildRequest: Fn() -> reqwest::RequestBuilder,
937    HandleResponse: Fn(reqwest::Response) -> Fut,
938    Fut: Future<Output = Result<T, BoxError>>,
939{
940    for attempt in 0..=MODEL_REQUEST_MAX_RETRIES {
941        let response = match build_request().send().await {
942            Ok(response) => response,
943            Err(err) => {
944                let retryable = is_retryable_reqwest_error(&err);
945                let message = format!(
946                    "Failed to send completion request, model: {}, error: {}",
947                    model,
948                    format_error_chain(&err)
949                );
950                if retryable && attempt < MODEL_REQUEST_MAX_RETRIES {
951                    log_completion_retry(model, attempt + 1);
952                    backoff_before_retry(None).await;
953                    continue;
954                }
955
956                return Err(Box::new(
957                    ModelError::new(message)
958                        .with_retryable(retryable)
959                        .with_source(Box::new(err)),
960                ));
961            }
962        };
963
964        let status = response.status();
965        if status.is_success() {
966            match handle_response(response).await {
967                Ok(output) => return Ok(output),
968                Err(err) if is_retryable_box_error(&err) && attempt < MODEL_REQUEST_MAX_RETRIES => {
969                    log_completion_retry(model, attempt + 1);
970                    backoff_before_retry(None).await;
971                    continue;
972                }
973                Err(err) => return Err(err),
974            }
975        }
976
977        let retryable = is_retryable_status(status);
978        let retry_after = retry_after_duration(response.headers());
979        let body = match read_error_body(response).await {
980            Ok(body) => body,
981            Err(err) => {
982                let retryable = retryable || is_retryable_reqwest_error(&err);
983                let message = format!(
984                    "Completion failed, model: {}, status: {}; failed to read error body: {}",
985                    model,
986                    status,
987                    format_error_chain(&err)
988                );
989                if retryable && attempt < MODEL_REQUEST_MAX_RETRIES {
990                    log_completion_retry(model, attempt + 1);
991                    backoff_before_retry(retry_after).await;
992                    continue;
993                }
994
995                return Err(Box::new(
996                    ModelError::new(message)
997                        .with_retryable(retryable)
998                        .with_status(status)
999                        .with_retry_after(retry_after)
1000                        .with_source(Box::new(err)),
1001                ));
1002            }
1003        };
1004        let message = format!(
1005            "Completion failed, model: {}, status: {}, body: {}",
1006            model, status, body
1007        );
1008
1009        if retryable && attempt < MODEL_REQUEST_MAX_RETRIES {
1010            log_completion_retry(model, attempt + 1);
1011            backoff_before_retry(retry_after).await;
1012            continue;
1013        }
1014
1015        return Err(Box::new(
1016            ModelError::new(message)
1017                .with_retryable(retryable)
1018                .with_status(status)
1019                .with_retry_after(retry_after),
1020        ));
1021    }
1022
1023    unreachable!("completion retry loop always returns before exhausting attempts")
1024}
1025
1026/// Sleeps before an in-SDK retry so transient overload (429/5xx/connection
1027/// flaps) has a chance to clear. The upstream `Retry-After` hint is honored up
1028/// to the configured retry cap; longer waits are the responsibility of upper
1029/// layers, which receive the hint via [`ModelError::retry_after`].
1030async fn backoff_before_retry(retry_after: Option<Duration>) {
1031    let delay = retry_after
1032        .unwrap_or(MODEL_RETRY_BACKOFF)
1033        .min(MODEL_RETRY_MAX_BACKOFF);
1034    tokio::time::sleep(delay).await;
1035}
1036
1037fn retry_after_duration(headers: &http::HeaderMap) -> Option<Duration> {
1038    let value = headers
1039        .get(http::header::RETRY_AFTER)?
1040        .to_str()
1041        .ok()?
1042        .trim();
1043    if let Ok(seconds) = value.parse::<u64>() {
1044        return Some(Duration::from_secs(seconds));
1045    }
1046
1047    // HTTP-date form, e.g. "Wed, 21 Oct 2026 07:28:00 GMT", common from
1048    // gateways and CDNs in front of model providers.
1049    let when = chrono::DateTime::parse_from_rfc2822(value).ok()?;
1050    (when.with_timezone(&chrono::Utc) - chrono::Utc::now())
1051        .to_std()
1052        .ok()
1053}
1054
1055fn log_completion_retry(model: &str, retry: usize) {
1056    // The returned error retains diagnostics; routine retry logs omit upstream bodies.
1057    log::warn!(
1058        "Retrying completion request, model: {}, retry: {}/{}",
1059        model,
1060        retry,
1061        MODEL_REQUEST_MAX_RETRIES
1062    );
1063}
1064
1065/// Host matcher that accepts every provider host.
1066#[derive(Clone, Copy, Debug)]
1067pub struct AnyHost;
1068
1069impl PartialEq<&str> for AnyHost {
1070    fn eq(&self, _other: &&str) -> bool {
1071        true
1072    }
1073}
1074
1075/// Creates a reqwest client builder with Anda Engine defaults.
1076pub fn request_client_builder() -> reqwest::ClientBuilder {
1077    reqwest::Client::builder()
1078        .use_rustls_tls()
1079        .https_only(true)
1080        .retry(
1081            reqwest::retry::for_host(AnyHost)
1082                .max_retries_per_request(1)
1083                .classify_fn(|req_rep| {
1084                    let is_idempotent = matches!(
1085                        req_rep.method(),
1086                        &http::Method::GET
1087                            | &http::Method::HEAD
1088                            | &http::Method::OPTIONS
1089                            | &http::Method::TRACE
1090                            | &http::Method::PUT
1091                            | &http::Method::DELETE
1092                    );
1093
1094                    if !is_idempotent {
1095                        return req_rep.success();
1096                    }
1097
1098                    if req_rep.error().is_some() {
1099                        return req_rep.retryable();
1100                    }
1101
1102                    match req_rep.status() {
1103                        Some(status) if is_retryable_status(status) => req_rep.retryable(),
1104                        _ => req_rep.success(),
1105                    }
1106                }),
1107        )
1108        // Do not use HTTP/2 PINGs as the liveness detector for model SSE
1109        // streams. Some provider/CDN edges can keep a long reasoning stream
1110        // alive while delaying PING ACKs; hyper then closes the connection and
1111        // reqwest reports "error decoding response body: ... operation timed
1112        // out" even though the body was still progressing. The per-read body
1113        // timeout below is the stall detector for completions.
1114        .http2_keep_alive_interval(COMPLETION_HTTP2_KEEP_ALIVE_INTERVAL)
1115        .connect_timeout(COMPLETION_CONNECT_TIMEOUT)
1116        // Read (idle) timeout is the authoritative stall detector for streamed
1117        // completions: it resets on every chunk, so a long-but-progressing
1118        // reasoning stream of unknown size is never killed, while a connection
1119        // that goes silent is failed promptly with clear attribution.
1120        .read_timeout(COMPLETION_READ_TIMEOUT)
1121        // Total request timeout, including the streamed body. Heavy reasoning
1122        // completions can run for many minutes; provider SDKs default to 10
1123        // minutes.
1124        .timeout(COMPLETION_REQUEST_TIMEOUT)
1125        .user_agent(APP_USER_AGENT)
1126        .default_headers({
1127            let mut headers = reqwest::header::HeaderMap::new();
1128            let ct: http::HeaderValue = http::HeaderValue::from_static(CONTENT_TYPE_JSON);
1129            headers.insert(http::header::CONTENT_TYPE, ct.clone());
1130            headers.insert(http::header::ACCEPT, ct);
1131            headers
1132        })
1133}
1134
1135const SSE_DONE_MARKER: &[u8] = b"data: [DONE]";
1136
1137pub(crate) async fn read_completion_stream<T>(
1138    response: reqwest::Response,
1139    model: &str,
1140) -> Result<(Vec<T>, bool), BoxError>
1141where
1142    T: DeserializeOwned,
1143{
1144    let request_id = upstream_request_id(response.headers());
1145    let started = Instant::now();
1146    let mut body = Vec::new();
1147    let mut scanned: usize = 0;
1148    let mut stream = response.bytes_stream();
1149
1150    while let Some(chunk) = stream.next().await {
1151        let chunk = chunk.map_err(|err| {
1152            // Mid-stream failures are reported with enough context to tell a
1153            // client-side timeout (elapsed near the request deadline) apart
1154            // from an upstream/gateway abort, and to follow up with the
1155            // provider via the request id.
1156            let action = format!(
1157                "Failed to read streaming completion response (request_id: {}, received: {} bytes, elapsed: {:.1?})",
1158                request_id.as_deref().unwrap_or("-"),
1159                body.len(),
1160                started.elapsed(),
1161            );
1162            completion_transport_error(model, &action, err)
1163        })?;
1164        if body.len() + chunk.len() > MAX_COMPLETION_RESPONSE_BYTES {
1165            return Err(completion_response_too_large(
1166                model,
1167                request_id.as_deref(),
1168                body.len() + chunk.len(),
1169            ));
1170        }
1171        body.extend_from_slice(&chunk);
1172        // Only scan the unscanned tail (with marker-sized overlap), so long
1173        // streams are not rescanned from the start on every chunk.
1174        let start = scanned.saturating_sub(SSE_DONE_MARKER.len());
1175        if body_contains_sse_done(&body, start) {
1176            return Ok((parse_streaming_json_events(&body, model)?, true));
1177        }
1178        scanned = body.len();
1179    }
1180
1181    Ok((parse_streaming_json_events(&body, model)?, false))
1182}
1183
1184#[cfg(test)]
1185async fn read_sse_json_events<T: DeserializeOwned>(
1186    response: reqwest::Response,
1187    model: &str,
1188) -> Result<Vec<T>, BoxError> {
1189    read_completion_stream(response, model)
1190        .await
1191        .map(|(events, _)| events)
1192}
1193
1194/// Returns true when a `data: [DONE]` line exists at or after `from`.
1195///
1196/// The marker must be anchored at the start of an SSE line (buffer start or a
1197/// preceding `\n`). Generated content inside a JSON string can legitimately
1198/// contain the marker text, and must not terminate the stream early.
1199fn body_contains_sse_done(body: &[u8], from: usize) -> bool {
1200    if from == 0 && body.starts_with(SSE_DONE_MARKER) {
1201        return true;
1202    }
1203    body[from..]
1204        .windows(SSE_DONE_MARKER.len() + 1)
1205        .any(|window| window[0] == b'\n' && &window[1..] == SSE_DONE_MARKER)
1206}
1207
1208fn parse_streaming_json_events<T>(body: &[u8], model: &str) -> Result<Vec<T>, BoxError>
1209where
1210    T: DeserializeOwned,
1211{
1212    let body = std::str::from_utf8(body).map_err(|err| {
1213        format!(
1214            "Invalid UTF-8 in streaming completion response, model: {}, error: {}",
1215            model, err
1216        )
1217    })?;
1218    let body = body.strip_prefix('\u{feff}').unwrap_or(body);
1219
1220    if !looks_like_sse(body) {
1221        return parse_json_event_payload(body, model);
1222    }
1223
1224    let mut data = String::new();
1225    let mut events = Vec::new();
1226
1227    for line in body.lines() {
1228        let line = line.strip_suffix('\r').unwrap_or(line);
1229        handle_sse_text_line(line, &mut data, &mut events, model)?;
1230    }
1231    flush_sse_data(&mut data, &mut events, model)?;
1232
1233    Ok(events)
1234}
1235
1236fn looks_like_sse(body: &str) -> bool {
1237    body.lines().any(|line| {
1238        let line = line.strip_prefix('\u{feff}').unwrap_or(line);
1239        line.starts_with("data:")
1240            || line.starts_with("event:")
1241            || line.starts_with("id:")
1242            || line.starts_with("retry:")
1243            || line.starts_with(':')
1244    })
1245}
1246
1247fn handle_sse_text_line<T>(
1248    line: &str,
1249    data: &mut String,
1250    events: &mut Vec<T>,
1251    model: &str,
1252) -> Result<(), BoxError>
1253where
1254    T: DeserializeOwned,
1255{
1256    if line.is_empty() {
1257        return flush_sse_data(data, events, model);
1258    }
1259    if line.starts_with(':') {
1260        return Ok(());
1261    }
1262
1263    let Some(value) = line.strip_prefix("data:") else {
1264        return Ok(());
1265    };
1266    let value = value.strip_prefix(' ').unwrap_or(value);
1267    if !data.is_empty() {
1268        data.push('\n');
1269    }
1270    data.push_str(value);
1271    Ok(())
1272}
1273
1274fn flush_sse_data<T>(data: &mut String, events: &mut Vec<T>, model: &str) -> Result<(), BoxError>
1275where
1276    T: DeserializeOwned,
1277{
1278    let value = data.trim_end();
1279    if value.is_empty() || value == "[DONE]" {
1280        data.clear();
1281        return Ok(());
1282    }
1283
1284    let event = serde_json::from_str::<T>(value).map_err(|err| {
1285        format!(
1286            "Invalid streaming completion event, model: {}, error: {}, body: {}",
1287            model, err, value
1288        )
1289    })?;
1290    events.push(event);
1291    data.clear();
1292    Ok(())
1293}
1294
1295fn parse_json_event_payload<T>(body: &str, model: &str) -> Result<Vec<T>, BoxError>
1296where
1297    T: DeserializeOwned,
1298{
1299    let value = body.trim().strip_prefix('\u{feff}').unwrap_or(body.trim());
1300    if value.is_empty() || value == "[DONE]" {
1301        return Ok(Vec::new());
1302    }
1303
1304    if value.starts_with('[')
1305        && let Ok(events) = serde_json::from_str::<Vec<T>>(value)
1306    {
1307        return Ok(events);
1308    }
1309
1310    match serde_json::from_str::<T>(value) {
1311        Ok(event) => Ok(vec![event]),
1312        Err(single_err) => match serde_json::from_str::<Vec<T>>(value) {
1313            Ok(events) => Ok(events),
1314            Err(array_err) => {
1315                let mut events = Vec::new();
1316                let mut saw_line = false;
1317                for line in value.lines() {
1318                    let line = line.trim();
1319                    if line.is_empty() || line == "[DONE]" {
1320                        continue;
1321                    }
1322                    saw_line = true;
1323                    let event = serde_json::from_str::<T>(line).map_err(|line_err| {
1324                        format!(
1325                            "Invalid streaming completion event, model: {}, error: {}, body: {}",
1326                            model, line_err, line
1327                        )
1328                    })?;
1329                    events.push(event);
1330                }
1331
1332                if saw_line {
1333                    return Ok(events);
1334                }
1335
1336                Err(format!(
1337                    "Invalid streaming completion event, model: {}, error: {}; array error: {}, body: {}",
1338                    model, single_err, array_err, value
1339                )
1340                .into())
1341            }
1342        },
1343    }
1344}
1345
1346pub(crate) fn streaming_completion_request(
1347    request: reqwest::RequestBuilder,
1348) -> reqwest::RequestBuilder {
1349    request
1350        .header(reqwest::header::ACCEPT, "text/event-stream")
1351        .header(reqwest::header::ACCEPT_ENCODING, "identity")
1352        // Streaming model completions can run much longer than generic HTTP
1353        // calls made with the same client. Set the request-level total timeout
1354        // explicitly so downstream clients with shorter shared timeouts do not
1355        // abort a progressing SSE body before the model timeout budget.
1356        .timeout(COMPLETION_REQUEST_TIMEOUT)
1357}
1358
1359#[cfg(test)]
1360mod tests {
1361    use super::*;
1362    use anda_core::FunctionDefinition;
1363    use http::{HeaderMap, HeaderValue, StatusCode};
1364    use tokio::io::{AsyncReadExt, AsyncWriteExt};
1365
1366    #[tokio::test]
1367    async fn prematurely_closed_provider_streams_are_retried() {
1368        use serde_json::json;
1369        let chat = json!({"id":"r","model":"test","choices":[{"index":0,"delta":{"role":"assistant","content":"ok"}}]});
1370        let gemini = json!({"candidates":[{"content":{"role":"model","parts":[{"text":"ok"}]}}]});
1371        let anthropic = json!({"type":"message_start","message":{"id":"r","type":"message","role":"assistant","model":"test","content":[{"type":"text","text":"ok"}],"usage":{}}});
1372        let cases = [
1373            (
1374                "openai",
1375                vec![chat],
1376                vec![json!({"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]})],
1377            ),
1378            (
1379                "gemini",
1380                vec![gemini],
1381                vec![json!({"candidates":[{"finishReason":"STOP"}]})],
1382            ),
1383            (
1384                "anthropic",
1385                vec![anthropic],
1386                vec![
1387                    json!({"type":"message_delta","delta":{"stop_reason":"end_turn"}}),
1388                    json!({"type":"message_stop"}),
1389                ],
1390            ),
1391        ];
1392        for (family, initial, terminal) in cases {
1393            let encode = |events: Vec<Json>| {
1394                events
1395                    .into_iter()
1396                    .map(|event| format!("data: {event}\n\n"))
1397                    .collect::<String>()
1398                    .into_bytes()
1399            };
1400            let incomplete = test_support::MockResponse {
1401                status: StatusCode::OK,
1402                headers: test_support::sse_headers(),
1403                body: encode(initial.clone()),
1404            };
1405            let complete = test_support::MockResponse {
1406                body: encode(initial.into_iter().chain(terminal).collect()),
1407                ..incomplete.clone()
1408            };
1409            let (endpoint, state) =
1410                test_support::spawn_retry_mock_server(vec![incomplete, complete]).await;
1411            let config = ModelConfig {
1412                family: family.into(),
1413                model: "test".into(),
1414                api_base: endpoint,
1415                api_key: "fake".into(),
1416                stream: true,
1417                ..Default::default()
1418            };
1419            let model = config.model(test_support::no_proxy_client()).unwrap();
1420            let output = model
1421                .completion(CompletionRequest {
1422                    prompt: "hello".into(),
1423                    ..Default::default()
1424                })
1425                .await
1426                .unwrap();
1427            assert_eq!(output.content, "ok", "{family}");
1428            assert!(output.failed_reason.is_none(), "{family}");
1429            assert_eq!(state.lock().unwrap().1, 2, "{family}");
1430        }
1431    }
1432
1433    #[derive(Clone)]
1434    struct TestCompleter {
1435        name: &'static str,
1436    }
1437
1438    impl CompletionFeaturesDyn for TestCompleter {
1439        fn completion(&self, _req: CompletionRequest) -> BoxPinFut<Result<AgentOutput, BoxError>> {
1440            Box::pin(futures::future::ready(Ok(AgentOutput::default())))
1441        }
1442
1443        fn model_name(&self) -> String {
1444            self.name.to_string()
1445        }
1446    }
1447
1448    fn test_model(name: &'static str) -> Model {
1449        Model::new(Arc::new(TestCompleter { name }))
1450    }
1451
1452    fn http_client() -> reqwest::Client {
1453        test_support::no_proxy_client()
1454    }
1455
1456    async fn spawn_truncated_sse_after_done_server() -> String {
1457        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1458        let addr = listener.local_addr().unwrap();
1459        tokio::spawn(async move {
1460            let (mut socket, _) = listener.accept().await.unwrap();
1461            let mut request = [0; 1024];
1462            let _ = socket.read(&mut request).await;
1463            socket
1464                .write_all(
1465                    b"HTTP/1.1 200 OK\r\n\
1466                      Content-Type: text/event-stream\r\n\
1467                      Content-Length: 4096\r\n\
1468                      Connection: close\r\n\
1469                      \r\n\
1470                      data: {\"a\":1}\n\n\
1471                      data: [DONE]\n\n",
1472                )
1473                .await
1474                .unwrap();
1475            let _ = socket.shutdown().await;
1476        });
1477        format!("http://{addr}")
1478    }
1479
1480    /// Sends an event whose JSON content embeds the literal `data: [DONE]`
1481    /// text in an early chunk, then a second event and the real terminator
1482    /// in a later chunk.
1483    async fn spawn_sse_with_done_marker_in_content_server() -> String {
1484        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1485        let addr = listener.local_addr().unwrap();
1486        tokio::spawn(async move {
1487            let (mut socket, _) = listener.accept().await.unwrap();
1488            let mut request = [0; 1024];
1489            let _ = socket.read(&mut request).await;
1490            socket
1491                .write_all(
1492                    b"HTTP/1.1 200 OK\r\n\
1493                      Content-Type: text/event-stream\r\n\
1494                      Connection: close\r\n\
1495                      \r\n\
1496                      data: {\"text\":\"sse ends with data: [DONE]\"}\n\n",
1497                )
1498                .await
1499                .unwrap();
1500            socket.flush().await.unwrap();
1501            tokio::time::sleep(Duration::from_millis(50)).await;
1502            socket
1503                .write_all(b"data: {\"b\":2}\n\ndata: [DONE]\n\n")
1504                .await
1505                .unwrap();
1506            let _ = socket.shutdown().await;
1507        });
1508        format!("http://{addr}")
1509    }
1510
1511    async fn spawn_stalling_sse_body_server() -> String {
1512        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1513        let addr = listener.local_addr().unwrap();
1514        tokio::spawn(async move {
1515            let (mut socket, _) = listener.accept().await.unwrap();
1516            let mut request = [0; 1024];
1517            let _ = socket.read(&mut request).await;
1518            socket
1519                .write_all(
1520                    b"HTTP/1.1 200 OK\r\n\
1521                      Content-Type: text/event-stream\r\n\
1522                      Connection: close\r\n\
1523                      \r\n\
1524                      data: {\"a\":1}\n\n",
1525                )
1526                .await
1527                .unwrap();
1528            socket.flush().await.unwrap();
1529            tokio::time::sleep(Duration::from_millis(500)).await;
1530            let _ = socket.shutdown().await;
1531        });
1532        format!("http://{addr}")
1533    }
1534
1535    fn retry_count(state: &test_support::RetryState) -> usize {
1536        state.lock().unwrap().1
1537    }
1538
1539    fn model_config(family: &str, model: &str) -> ModelConfig {
1540        ModelConfig {
1541            family: family.to_string(),
1542            model: model.to_string(),
1543            api_base: "https://example.com".to_string(),
1544            api_key: "test-key".to_string(),
1545            ..Default::default()
1546        }
1547    }
1548
1549    #[test]
1550    fn model_effort_serializes_config_values() {
1551        let config: ModelConfig = serde_json::from_value(serde_json::json!({
1552            "family": "openai",
1553            "model": "gpt-5",
1554            "api_base": "http://localhost",
1555            "api_key": "test-key",
1556            "effort": "max"
1557        }))
1558        .unwrap();
1559
1560        assert_eq!(config.effort, Some(ModelEffort::Max));
1561        assert_eq!(
1562            serde_json::to_value(ModelEffort::Minimal).unwrap(),
1563            "minimal"
1564        );
1565    }
1566
1567    #[test]
1568    fn models_default_is_empty() {
1569        let models = Models::default();
1570
1571        assert!(models.get_model().is_none());
1572        assert!(models.get("missing").is_none());
1573        assert!(models.resolve("missing").is_none());
1574    }
1575
1576    #[test]
1577    fn set_model_sets_primary_without_registering_a_label() {
1578        let models = Models::default();
1579        models.set_model(test_model("primary"));
1580
1581        assert_eq!(
1582            models
1583                .get_model()
1584                .expect("primary model should exist")
1585                .model_name(),
1586            "primary"
1587        );
1588        assert!(models.get("primary").is_some());
1589    }
1590
1591    #[test]
1592    fn set_promotes_first_inserted_model_to_primary() {
1593        let models = Models::default();
1594        models.set("x".to_string(), test_model("X"));
1595
1596        assert_eq!(
1597            models.get("x").expect("label x should exist").model_name(),
1598            "X"
1599        );
1600        assert_eq!(
1601            models
1602                .get_model()
1603                .expect("primary model should be initialized")
1604                .model_name(),
1605            "X"
1606        );
1607    }
1608
1609    #[test]
1610    fn fallback_label_has_no_default_routing_semantics() {
1611        let models = Models::default();
1612        models.set_model(test_model("primary"));
1613        models.set("fallback".to_string(), test_model("fallback"));
1614
1615        assert_eq!(
1616            models
1617                .get("fallback")
1618                .expect("fallback is still a normal label")
1619                .model_name(),
1620            "fallback"
1621        );
1622        assert_eq!(
1623            models
1624                .get_model()
1625                .expect("primary model should stay the default")
1626                .model_name(),
1627            "primary"
1628        );
1629        assert_eq!(
1630            models
1631                .resolve("unknown")
1632                .expect("missing label should use default routing")
1633                .model_name(),
1634            "primary"
1635        );
1636    }
1637
1638    #[test]
1639    fn resolve_prefers_exact_label_then_default() {
1640        let models = Models::default();
1641        models.set_model(test_model("primary"));
1642        models.set("flash".to_string(), test_model("flash"));
1643
1644        assert_eq!(
1645            models
1646                .resolve("flash")
1647                .expect("exact label should win")
1648                .model_name(),
1649            "flash"
1650        );
1651        assert_eq!(
1652            models
1653                .resolve("missing")
1654                .expect("missing label should use default routing")
1655                .model_name(),
1656            "primary"
1657        );
1658        assert_eq!(
1659            models
1660                .resolve("")
1661                .expect("empty label should use default routing")
1662                .model_name(),
1663            "primary"
1664        );
1665    }
1666
1667    #[test]
1668    fn model_config_validates_required_fields_and_builds_supported_families() {
1669        let client = http_client();
1670
1671        let mut config = model_config("openai", "gpt-5");
1672        config.disabled = true;
1673        let Err(err) = config.model(client.clone()) else {
1674            panic!("disabled model should fail");
1675        };
1676        assert!(err.to_string().contains("disabled"));
1677
1678        for (field, config) in [
1679            (
1680                "model name",
1681                ModelConfig {
1682                    model: String::new(),
1683                    ..model_config("openai", "gpt-5")
1684                },
1685            ),
1686            (
1687                "model family",
1688                ModelConfig {
1689                    family: String::new(),
1690                    ..model_config("openai", "gpt-5")
1691                },
1692            ),
1693            (
1694                "api_base",
1695                ModelConfig {
1696                    api_base: String::new(),
1697                    ..model_config("openai", "gpt-5")
1698                },
1699            ),
1700            (
1701                "api_key",
1702                ModelConfig {
1703                    api_key: String::new(),
1704                    ..model_config("openai", "gpt-5")
1705                },
1706            ),
1707        ] {
1708            let Err(err) = config.model(client.clone()) else {
1709                panic!("{field} should fail");
1710            };
1711            let err = err.to_string();
1712            assert!(err.contains(field), "{field}: {err}");
1713        }
1714
1715        let Err(err) = model_config("unknown", "m").model(client.clone()) else {
1716            panic!("unsupported family should fail");
1717        };
1718        assert!(err.to_string().contains("unsupported model family"));
1719
1720        let mut gemini = model_config("gemini", "gemini-2.5-pro");
1721        gemini.context_window = 123;
1722        gemini.max_output = 45;
1723        let model = gemini.model(client.clone()).unwrap();
1724        assert_eq!(model.model_name(), "gemini-2.5-pro");
1725        assert_eq!(model.labels, vec!["gemini-2.5-pro"]);
1726        assert_eq!(model.context_window, 123);
1727        assert_eq!(model.max_output, 45);
1728
1729        let mut anthropic = model_config("anthropic", "claude-sonnet-4-5");
1730        anthropic.labels = vec!["pro".to_string(), "primary".to_string()];
1731        anthropic.bearer_auth = true;
1732        anthropic.stream = true;
1733        anthropic.effort = Some(ModelEffort::High);
1734        let model = anthropic.model(client.clone()).unwrap();
1735        assert_eq!(model.model_name(), "claude-sonnet-4-5");
1736        assert_eq!(model.labels, vec!["pro", "primary"]);
1737
1738        let model = model_config("openai", "gpt-5")
1739            .model(client.clone())
1740            .unwrap();
1741        assert_eq!(model.model_name(), "gpt-5");
1742        let model = model_config("openai", "deepseek-chat")
1743            .model(client)
1744            .unwrap();
1745        assert_eq!(model.model_name(), "deepseek-chat");
1746    }
1747
1748    #[test]
1749    fn models_registry_clones_names_replaces_labels_and_loads_configs() {
1750        let models = Models::default();
1751        models.set_model(test_model("flash-v1").with_labels(vec!["FAST".into()]));
1752        assert!(models.contains("fast"));
1753        assert_eq!(
1754            models.model_names(),
1755            BTreeSet::from(["flash-v1".to_string()])
1756        );
1757
1758        models.set("flash".to_string(), test_model("flash-v2"));
1759        assert!(models.contains("flash"));
1760        assert_eq!(models.get("FLASH").unwrap().model_name(), "flash-v2");
1761        assert_eq!(
1762            models.model_names(),
1763            BTreeSet::from(["flash-v1".to_string(), "flash-v2".to_string()])
1764        );
1765
1766        models.set("primary".to_string(), test_model("primary-v2"));
1767        assert_eq!(models.get_model().unwrap().model_name(), "primary-v2");
1768
1769        let cloned = Models::from_clone(&models);
1770        assert_eq!(cloned.get("primary").unwrap().model_name(), "primary-v2");
1771        assert_eq!(
1772            cloned.resolve("missing").unwrap().model_name(),
1773            "primary-v2"
1774        );
1775
1776        let replacement = Models::default();
1777        replacement.set_model(test_model("replacement-primary").with_labels(vec!["next".into()]));
1778        let replaced = Models::default();
1779        replaced.set("old".to_string(), test_model("old"));
1780        replaced.replace(&replacement);
1781        assert!(!replaced.contains("old"));
1782        assert!(replaced.contains("next"));
1783        assert_eq!(
1784            replaced.get_model().unwrap().model_name(),
1785            "replacement-primary"
1786        );
1787
1788        replacement.set("later".to_string(), test_model("later"));
1789        assert!(!replaced.contains("later"));
1790
1791        let configs = vec![
1792            ModelConfig {
1793                labels: vec!["primary".to_string()],
1794                ..model_config("openai", "gpt-5")
1795            },
1796            ModelConfig {
1797                disabled: true,
1798                ..model_config("openai", "disabled")
1799            },
1800        ];
1801        let loaded = Models::from_configs(&configs, http_client());
1802        assert!(loaded.contains("primary"));
1803        assert!(!loaded.contains("disabled"));
1804        assert_eq!(loaded.get_model().unwrap().model_name(), "gpt-5");
1805    }
1806
1807    #[tokio::test]
1808    async fn model_max_output_caps_explicit_request_output_budget() {
1809        let completer = testing::ScriptedCompleter::new("capped").into_arc();
1810        let mut model = Model::with_completer(completer.clone());
1811
1812        // Unknown model limit: the request budget passes through untouched.
1813        model
1814            .completion(CompletionRequest::default())
1815            .await
1816            .unwrap();
1817
1818        model.max_output = 32_000;
1819        model
1820            .completion(CompletionRequest::default())
1821            .await
1822            .unwrap();
1823        for requested in [8_000, 64_000] {
1824            model
1825                .completion(CompletionRequest {
1826                    max_output_tokens: Some(requested),
1827                    ..Default::default()
1828                })
1829                .await
1830                .unwrap();
1831        }
1832
1833        let sent = completer
1834            .requests()
1835            .into_iter()
1836            .map(|req| req.max_output_tokens)
1837            .collect::<Vec<_>>();
1838        // An unset budget is left to the adapter; explicit budgets are capped.
1839        assert_eq!(sent, vec![None, None, Some(8_000), Some(32_000)]);
1840    }
1841
1842    #[tokio::test]
1843    async fn model_completion_placeholders_and_mock_tool_calls_are_stable() {
1844        let not_implemented = Model::not_implemented();
1845        assert_eq!(not_implemented.model_name(), "not_implemented");
1846        let err = not_implemented
1847            .completion(CompletionRequest::default())
1848            .await
1849            .unwrap_err();
1850        assert!(err.to_string().contains("not implemented"));
1851
1852        let mock = Model::mock_implemented().with_labels(vec!["mock".into()]);
1853        assert_eq!(mock.model_name(), "mock_implemented");
1854        let output = mock
1855            .completion(CompletionRequest {
1856                prompt: "{\"q\":\"anda\"}".to_string(),
1857                tools: vec![FunctionDefinition {
1858                    name: "lookup".to_string(),
1859                    ..Default::default()
1860                }],
1861                ..Default::default()
1862            })
1863            .await
1864            .unwrap();
1865        assert_eq!(output.content, "{\"q\":\"anda\"}");
1866        assert_eq!(output.tool_calls.len(), 1);
1867        assert_eq!(output.tool_calls[0].name, "lookup");
1868        assert_eq!(output.tool_calls[0].args["q"], "anda");
1869
1870        let output = mock
1871            .completion(CompletionRequest {
1872                prompt: String::new(),
1873                tools: vec![FunctionDefinition {
1874                    name: "lookup".to_string(),
1875                    ..Default::default()
1876                }],
1877                ..Default::default()
1878            })
1879            .await
1880            .unwrap();
1881        assert!(output.tool_calls.is_empty());
1882    }
1883
1884    #[test]
1885    fn streaming_json_event_parser_accepts_bom_sse_ndjson_and_arrays() {
1886        let events = parse_streaming_json_events::<serde_json::Value>(
1887            b"\xef\xbb\xbfdata: {\"a\":1}\n\ndata: [DONE]\n\n",
1888            "test-model",
1889        )
1890        .unwrap();
1891        assert_eq!(events, vec![serde_json::json!({"a": 1})]);
1892
1893        let events = parse_streaming_json_events::<serde_json::Value>(
1894            b"{\"a\":1}\n{\"b\":2}\n[DONE]\n",
1895            "test-model",
1896        )
1897        .unwrap();
1898        assert_eq!(
1899            events,
1900            vec![serde_json::json!({"a": 1}), serde_json::json!({"b": 2})]
1901        );
1902
1903        let events =
1904            parse_streaming_json_events::<serde_json::Value>(br#"[{"a":1},{"b":2}]"#, "test-model")
1905                .unwrap();
1906        assert_eq!(
1907            events,
1908            vec![serde_json::json!({"a": 1}), serde_json::json!({"b": 2})]
1909        );
1910    }
1911
1912    #[tokio::test]
1913    async fn streaming_reader_ignores_mislabelled_content_encoding() {
1914        let mut headers = HeaderMap::new();
1915        headers.insert(
1916            http::header::CONTENT_TYPE,
1917            HeaderValue::from_static("text/event-stream"),
1918        );
1919        let (endpoint, _) =
1920            test_support::spawn_retry_mock_server(vec![test_support::MockResponse {
1921                status: StatusCode::OK,
1922                headers,
1923                body: b"data: {\"a\":1}\n\ndata: [DONE]\n\n".to_vec(),
1924            }])
1925            .await;
1926        let client = request_client_builder()
1927            .https_only(false)
1928            .no_proxy()
1929            .build()
1930            .unwrap();
1931        let response = client.get(endpoint).send().await.unwrap();
1932
1933        let events = read_sse_json_events::<serde_json::Value>(response, "test-model")
1934            .await
1935            .unwrap();
1936
1937        assert_eq!(events, vec![serde_json::json!({"a": 1})]);
1938    }
1939
1940    #[tokio::test]
1941    async fn streaming_reader_returns_after_done_before_late_body_error() {
1942        let endpoint = spawn_truncated_sse_after_done_server().await;
1943        let client = request_client_builder()
1944            .https_only(false)
1945            .no_proxy()
1946            .build()
1947            .unwrap();
1948        let response = client.get(endpoint).send().await.unwrap();
1949
1950        let events = read_sse_json_events::<serde_json::Value>(response, "test-model")
1951            .await
1952            .unwrap();
1953
1954        assert_eq!(events, vec![serde_json::json!({"a": 1})]);
1955    }
1956
1957    #[test]
1958    fn sse_done_detection_is_line_anchored() {
1959        assert!(body_contains_sse_done(b"data: [DONE]\n\n", 0));
1960        assert!(body_contains_sse_done(
1961            b"data: {\"a\":1}\n\ndata: [DONE]\n\n",
1962            0
1963        ));
1964        // The marker text inside generated JSON content must not terminate
1965        // the stream.
1966        assert!(!body_contains_sse_done(
1967            b"data: {\"text\":\"sse ends with data: [DONE]\"}\n\n",
1968            0
1969        ));
1970    }
1971
1972    #[tokio::test]
1973    async fn streaming_reader_is_not_truncated_by_done_marker_in_content() {
1974        let endpoint = spawn_sse_with_done_marker_in_content_server().await;
1975        let client = request_client_builder()
1976            .https_only(false)
1977            .no_proxy()
1978            .build()
1979            .unwrap();
1980        let response = client.get(endpoint).send().await.unwrap();
1981
1982        let events = read_sse_json_events::<serde_json::Value>(response, "test-model")
1983            .await
1984            .unwrap();
1985
1986        assert_eq!(
1987            events,
1988            vec![
1989                serde_json::json!({"text": "sse ends with data: [DONE]"}),
1990                serde_json::json!({"b": 2})
1991            ]
1992        );
1993    }
1994
1995    #[test]
1996    fn completion_transport_timeouts_are_streaming_safe() {
1997        // The observed failure hit at ~118s with body bytes already received;
1998        // HTTP/2 PING ACK timeouts must not be able to abort such a stream
1999        // before the explicit idle body timeout can make that decision.
2000        assert_eq!(COMPLETION_HTTP2_KEEP_ALIVE_INTERVAL, None);
2001        assert!(COMPLETION_READ_TIMEOUT > Duration::from_secs(118));
2002        assert!(COMPLETION_READ_TIMEOUT < COMPLETION_REQUEST_TIMEOUT);
2003        assert_eq!(COMPLETION_REQUEST_TIMEOUT, Duration::from_secs(600));
2004    }
2005
2006    #[test]
2007    fn streaming_completion_request_overrides_short_client_total_timeout() {
2008        let client = reqwest::Client::builder()
2009            .no_proxy()
2010            .timeout(Duration::from_millis(100))
2011            .build()
2012            .unwrap();
2013        let request = streaming_completion_request(client.get("https://example.com"))
2014            .build()
2015            .unwrap();
2016
2017        assert_eq!(request.timeout(), Some(&COMPLETION_REQUEST_TIMEOUT));
2018    }
2019
2020    #[tokio::test]
2021    async fn streaming_reader_body_idle_timeout_is_retryable() {
2022        let endpoint = spawn_stalling_sse_body_server().await;
2023        let client = request_client_builder()
2024            .https_only(false)
2025            .no_proxy()
2026            .read_timeout(Duration::from_millis(100))
2027            .timeout(Duration::from_secs(5))
2028            .build()
2029            .unwrap();
2030        let response = client.get(endpoint).send().await.unwrap();
2031
2032        let err = tokio::time::timeout(
2033            Duration::from_secs(2),
2034            read_sse_json_events::<serde_json::Value>(response, "test-model"),
2035        )
2036        .await
2037        .expect("body read timeout should fire")
2038        .unwrap_err();
2039
2040        let message = err.to_string();
2041        assert!(
2042            message.contains("Failed to read streaming completion response"),
2043            "{message}"
2044        );
2045        assert!(message.contains("received:"), "{message}");
2046        assert!(message.contains("operation timed out"), "{message}");
2047        assert!(is_retryable_box_error(&err));
2048    }
2049
2050    #[test]
2051    fn retry_after_parses_seconds_and_http_date() {
2052        let mut headers = HeaderMap::new();
2053        headers.insert(http::header::RETRY_AFTER, HeaderValue::from_static("42"));
2054        assert_eq!(
2055            retry_after_duration(&headers),
2056            Some(Duration::from_secs(42))
2057        );
2058
2059        let when = chrono::Utc::now() + chrono::Duration::seconds(90);
2060        headers.insert(
2061            http::header::RETRY_AFTER,
2062            HeaderValue::from_str(&when.to_rfc2822()).unwrap(),
2063        );
2064        let parsed = retry_after_duration(&headers).expect("http-date should parse");
2065        assert!(parsed <= Duration::from_secs(90));
2066        assert!(parsed >= Duration::from_secs(80));
2067
2068        // A date in the past yields no delay hint.
2069        let when = chrono::Utc::now() - chrono::Duration::seconds(90);
2070        headers.insert(
2071            http::header::RETRY_AFTER,
2072            HeaderValue::from_str(&when.to_rfc2822()).unwrap(),
2073        );
2074        assert_eq!(retry_after_duration(&headers), None);
2075
2076        headers.insert(
2077            http::header::RETRY_AFTER,
2078            HeaderValue::from_static("not-a-date"),
2079        );
2080        assert_eq!(retry_after_duration(&headers), None);
2081    }
2082
2083    #[tokio::test]
2084    async fn custom_client_streaming_decode_errors_are_retryable() {
2085        let mut headers = HeaderMap::new();
2086        headers.insert(
2087            http::header::CONTENT_TYPE,
2088            HeaderValue::from_static("text/event-stream"),
2089        );
2090        headers.insert(
2091            http::header::CONTENT_ENCODING,
2092            HeaderValue::from_static("gzip"),
2093        );
2094        let (endpoint, _) =
2095            test_support::spawn_retry_mock_server(vec![test_support::MockResponse {
2096                status: StatusCode::OK,
2097                headers,
2098                body: b"data: {\"a\":1}\n\ndata: [DONE]\n\n".to_vec(),
2099            }])
2100            .await;
2101        let client = reqwest::Client::builder().no_proxy().build().unwrap();
2102        let response = streaming_completion_request(client.get(endpoint))
2103            .send()
2104            .await
2105            .unwrap();
2106
2107        let err = read_sse_json_events::<serde_json::Value>(response, "test-model")
2108            .await
2109            .unwrap_err();
2110
2111        let message = err.to_string();
2112        assert!(message.contains("error decoding response body"));
2113        // The reported error carries stream context and the decode source
2114        // chain, which reqwest's `Display` alone no longer exposes.
2115        assert!(message.contains("received: 0 bytes"), "{message}");
2116        assert!(message.contains("request_id: -"), "{message}");
2117        assert!(
2118            message.contains("error decoding response body: "),
2119            "{message}"
2120        );
2121        assert!(is_retryable_box_error(&err));
2122    }
2123
2124    #[test]
2125    fn error_chain_formatting_appends_unique_sources() {
2126        let root = std::io::Error::new(std::io::ErrorKind::TimedOut, "operation timed out");
2127        let outer =
2128            ModelError::new("error decoding response body".to_string()).with_source(Box::new(root));
2129        assert_eq!(
2130            format_error_chain(&outer),
2131            "error decoding response body: operation timed out"
2132        );
2133
2134        // A source already repeated in the message is not appended twice.
2135        let root = std::io::Error::new(std::io::ErrorKind::TimedOut, "operation timed out");
2136        let outer = ModelError::new("request failed: operation timed out".to_string())
2137            .with_source(Box::new(root));
2138        assert_eq!(
2139            format_error_chain(&outer),
2140            "request failed: operation timed out"
2141        );
2142    }
2143
2144    #[tokio::test]
2145    async fn completion_error_bodies_are_truncated_for_diagnostics() {
2146        assert_eq!(error_body_excerpt(b"short body"), "short body");
2147        let excerpt = error_body_excerpt(&vec![b'x'; MAX_ERROR_BODY_BYTES * 4]);
2148        assert!(excerpt.ends_with("… [truncated]"));
2149        assert!(excerpt.len() < MAX_ERROR_BODY_BYTES + 32);
2150
2151        let (endpoint, _) =
2152            test_support::spawn_retry_mock_server(vec![test_support::MockResponse {
2153                status: StatusCode::BAD_REQUEST,
2154                headers: HeaderMap::new(),
2155                body: vec![b'e'; 1024 * 1024],
2156            }])
2157            .await;
2158        let client = http_client();
2159        let err = execute_completion_request_with_retry(
2160            "error-body-test",
2161            || client.post(&endpoint),
2162            |response| async { read_completion_response_bytes(response, "error-body-test").await },
2163        )
2164        .await
2165        .unwrap_err();
2166        let message = err.to_string();
2167        assert!(message.contains("status: 400"));
2168        assert!(message.ends_with("… [truncated]"));
2169        assert!(message.len() < MAX_ERROR_BODY_BYTES + 256);
2170    }
2171
2172    #[test]
2173    fn upstream_request_id_checks_known_headers() {
2174        let mut headers = HeaderMap::new();
2175        assert_eq!(upstream_request_id(&headers), None);
2176
2177        headers.insert("cf-ray", HeaderValue::from_static("ray-123"));
2178        assert_eq!(upstream_request_id(&headers), Some("ray-123".to_string()));
2179
2180        headers.insert("x-request-id", HeaderValue::from_static("req-456"));
2181        assert_eq!(upstream_request_id(&headers), Some("req-456".to_string()));
2182    }
2183
2184    #[tokio::test]
2185    async fn completion_request_retries_transient_errors_and_exposes_retry_signal() {
2186        let mut headers = HeaderMap::new();
2187        headers.insert(http::header::RETRY_AFTER, HeaderValue::from_static("0"));
2188        let (endpoint, state) = test_support::spawn_retry_mock_server(vec![
2189            test_support::MockResponse {
2190                status: StatusCode::TOO_MANY_REQUESTS,
2191                headers,
2192                body: b"rate limited".to_vec(),
2193            },
2194            test_support::MockResponse {
2195                status: StatusCode::OK,
2196                headers: HeaderMap::new(),
2197                body: b"ok".to_vec(),
2198            },
2199        ])
2200        .await;
2201        let client = http_client();
2202
2203        let body = execute_completion_request_with_retry(
2204            "retry-test",
2205            || client.post(&endpoint),
2206            |response| async { read_completion_response_bytes(response, "retry-test").await },
2207        )
2208        .await
2209        .unwrap();
2210
2211        assert_eq!(&body[..], b"ok");
2212        assert_eq!(retry_count(&state), 2);
2213
2214        let mut retry_now = HeaderMap::new();
2215        retry_now.insert(http::header::RETRY_AFTER, HeaderValue::from_static("0"));
2216        let mut final_headers = HeaderMap::new();
2217        final_headers.insert(http::header::RETRY_AFTER, HeaderValue::from_static("45"));
2218        let (endpoint, state) = test_support::spawn_retry_mock_server(vec![
2219            test_support::MockResponse {
2220                status: StatusCode::TOO_MANY_REQUESTS,
2221                headers: retry_now.clone(),
2222                body: b"first limit".to_vec(),
2223            },
2224            test_support::MockResponse {
2225                status: StatusCode::TOO_MANY_REQUESTS,
2226                headers: retry_now.clone(),
2227                body: b"still limited".to_vec(),
2228            },
2229            test_support::MockResponse {
2230                status: StatusCode::TOO_MANY_REQUESTS,
2231                headers: retry_now,
2232                body: b"limited again".to_vec(),
2233            },
2234            test_support::MockResponse {
2235                status: StatusCode::TOO_MANY_REQUESTS,
2236                headers: final_headers,
2237                body: b"finally limited".to_vec(),
2238            },
2239        ])
2240        .await;
2241        let err = execute_completion_request_with_retry(
2242            "retry-test",
2243            || client.post(&endpoint),
2244            |response| async { read_completion_response_bytes(response, "retry-test").await },
2245        )
2246        .await
2247        .unwrap_err();
2248        let err_ref = err.as_ref() as &(dyn Error + 'static);
2249
2250        assert_eq!(retry_count(&state), MODEL_REQUEST_MAX_RETRIES + 1);
2251        assert!(is_retryable_box_error(&err));
2252        assert_eq!(
2253            model_error_status(err_ref),
2254            Some(StatusCode::TOO_MANY_REQUESTS)
2255        );
2256        assert_eq!(
2257            model_error_retry_after(err_ref),
2258            Some(Duration::from_secs(45))
2259        );
2260
2261        let (endpoint, state) =
2262            test_support::spawn_retry_mock_server(vec![test_support::MockResponse {
2263                status: StatusCode::BAD_REQUEST,
2264                headers: HeaderMap::new(),
2265                body: b"bad request".to_vec(),
2266            }])
2267            .await;
2268        let err = execute_completion_request_with_retry(
2269            "retry-test",
2270            || client.post(&endpoint),
2271            |response| async { read_completion_response_bytes(response, "retry-test").await },
2272        )
2273        .await
2274        .unwrap_err();
2275
2276        assert_eq!(retry_count(&state), 1);
2277        assert!(!is_retryable_box_error(&err));
2278    }
2279}