Skip to main content

rig_core/
provider_response.rs

1//! Preserved provider error responses and their machine codes.
2//!
3//! ```
4//! use rig_core::ProviderResponseError;
5//!
6//! let reply = ProviderResponseError::new(http::StatusCode::TOO_MANY_REQUESTS, "slow down");
7//! assert!(reply.is_retryable());
8//! ```
9use http::StatusCode;
10
11/// A raw provider error body with captured transport metadata.
12/// Callers must supply the provider's actual payload, not a generated diagnostic.
13/// Serialization omits headers and route; deserialization restores them as `None`.
14#[derive(Debug, Clone, PartialEq, Eq)]
15pub struct ProviderResponseError {
16    /// HTTP status of the provider response, when it was captured alongside the body.
17    pub status: Option<StatusCode>,
18    /// Raw response body as returned by the provider.
19    pub body: String,
20    /// Transport request ID from response headers or SDK metadata, when captured.
21    pub provider_request_id: Option<String>,
22    /// Captured response headers, including rate-limit metadata. `None` means
23    /// not captured, rather than an empty header set. Omitted during serialization.
24    pub headers: Option<http::HeaderMap>,
25    /// The provider's own machine-readable code for the failure, when the
26    /// transport reported one apart from the body: a gRPC status code name
27    /// (`UNAVAILABLE`), an AWS exception type (`ThrottlingException`).
28    /// `None` when the reply carried only a status and a body.
29    pub code: Option<String>,
30    /// Transport retry verdict used when status is absent or successful.
31    /// A non-success HTTP status takes precedence; refusals are never retryable.
32    pub transient: Option<bool>,
33    /// Whether this error represents a content refusal. Refusals are never
34    /// retryable; refusal content in a successful model answer is separate.
35    pub refusal: bool,
36    /// The provider and request path the reply answered, when the operation
37    /// records them for diagnostics (model listing does). Omitted during
38    /// serialization.
39    pub route: Option<String>,
40}
41
42impl ProviderResponseError {
43    /// Preserve a provider error response captured with its HTTP status.
44    pub fn new(status: StatusCode, body: impl Into<String>) -> Self {
45        Self {
46            status: Some(status),
47            body: body.into(),
48            provider_request_id: None,
49            headers: None,
50            code: None,
51            transient: None,
52            refusal: false,
53            route: None,
54        }
55    }
56
57    /// Preserve a provider error body that has no HTTP status (gRPC / SDK
58    /// transports).
59    pub fn without_status(body: impl Into<String>) -> Self {
60        Self {
61            status: None,
62            body: body.into(),
63            provider_request_id: None,
64            headers: None,
65            code: None,
66            transient: None,
67            refusal: false,
68            route: None,
69        }
70    }
71
72    /// Mark the reply as the provider's verdict on the content: a refusal,
73    /// final, never retried.
74    pub fn with_refusal(mut self, refusal: bool) -> Self {
75        self.refusal = refusal;
76        self
77    }
78
79    /// Attach the HTTP status a transport reported beside a reply that was
80    /// first preserved without one (an SDK that hands back the raw HTTP
81    /// response next to its typed exception). A status already set is kept.
82    pub fn with_status(mut self, status: Option<StatusCode>) -> Self {
83        if self.status.is_none() {
84            self.status = status;
85        }
86        self
87    }
88
89    /// Attach the provider's own machine-readable code for the failure.
90    pub fn with_code(mut self, code: Option<String>) -> Self {
91        self.code = code.filter(|code| !code.is_empty());
92        self
93    }
94
95    /// Replaces the transport retry verdict. Used for absent or successful
96    /// HTTP statuses unless the response is a refusal.
97    pub fn with_transient(mut self, transient: Option<bool>) -> Self {
98        self.transient = transient;
99        self
100    }
101
102    /// Returns false for refusals. Otherwise classifies non-success HTTP statuses
103    /// through [`crate::error::retryable_status`], or uses `transient` for absent
104    /// or successful statuses. Missing verdicts default to false.
105    pub fn is_retryable(&self) -> bool {
106        if self.refusal {
107            return false;
108        }
109        match self.status {
110            Some(status) if !status.is_success() => {
111                crate::error::retryable_status(Some(status.as_u16()))
112            }
113            _ => self.transient.unwrap_or(false),
114        }
115    }
116
117    /// Returns the explicit code, falling back to nonempty string fields
118    /// `error.code`, `error.status`, then `error.type` in the JSON body.
119    pub fn machine_code(&self) -> Option<String> {
120        self.code.clone().or_else(|| body_code(&self.body))
121    }
122
123    /// Attach the transport request id the failed response reported.
124    pub fn with_provider_request_id(mut self, request_id: Option<String>) -> Self {
125        self.provider_request_id = request_id.filter(|id| !id.is_empty());
126        self
127    }
128
129    /// Replaces captured response headers, including rate-limit metadata.
130    pub fn with_headers(mut self, headers: Option<http::HeaderMap>) -> Self {
131        self.headers = headers;
132        self
133    }
134}
135
136impl std::fmt::Display for ProviderResponseError {
137    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
138        match self.status {
139            Some(status) => write!(f, "status {status}: {}", self.body)?,
140            None => write!(f, "{}", self.body)?,
141        }
142        // The id support asks for belongs in the message a caller logs.
143        if let Some(request_id) = &self.provider_request_id {
144            write!(f, " (request id: {request_id})")?;
145        }
146        if let Some(route) = &self.route {
147            write!(f, " [{route}]")?;
148        }
149        Ok(())
150    }
151}
152
153impl std::error::Error for ProviderResponseError {}
154
155/// Serialized error metadata with numeric HTTP status and no headers.
156/// Volatile transport headers are excluded to keep replay records stable.
157#[derive(serde::Serialize, serde::Deserialize)]
158#[serde(deny_unknown_fields)]
159struct ProviderResponseErrorWire {
160    status: Option<u16>,
161    body: String,
162    provider_request_id: Option<String>,
163    #[serde(default, skip_serializing_if = "Option::is_none")]
164    code: Option<String>,
165    #[serde(default, skip_serializing_if = "Option::is_none")]
166    transient: Option<bool>,
167    #[serde(default, skip_serializing_if = "std::ops::Not::not")]
168    refusal: bool,
169}
170
171impl serde::Serialize for ProviderResponseError {
172    fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
173        ProviderResponseErrorWire {
174            status: self.status.map(|status| status.as_u16()),
175            body: self.body.clone(),
176            provider_request_id: self.provider_request_id.clone(),
177            code: self.code.clone(),
178            transient: self.transient,
179            refusal: self.refusal,
180        }
181        .serialize(serializer)
182    }
183}
184
185impl<'de> serde::Deserialize<'de> for ProviderResponseError {
186    fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
187        use serde::de::Error as _;
188        let wire = ProviderResponseErrorWire::deserialize(deserializer)?;
189        let status = wire
190            .status
191            .map(StatusCode::from_u16)
192            .transpose()
193            .map_err(D::Error::custom)?;
194        Ok(Self {
195            status,
196            body: wire.body,
197            provider_request_id: wire.provider_request_id,
198            headers: None,
199            code: wire.code,
200            transient: wire.transient,
201            refusal: wire.refusal,
202            route: None,
203        })
204    }
205}
206
207/// Returns the first nonempty string at `error.code`, `error.status`, or
208/// `error.type`, in that order. Invalid JSON, non-string fields, and missing
209/// envelopes yield no code.
210pub fn body_code(body: &str) -> Option<String> {
211    let value: serde_json::Value = serde_json::from_str(body).ok()?;
212    let error = value.get("error")?;
213    ["code", "status", "type"].iter().find_map(|field| {
214        error
215            .get(field)
216            .and_then(serde_json::Value::as_str)
217            .filter(|code| !code.is_empty())
218            .map(str::to_owned)
219    })
220}
221
222/// An id or model name as a response carries it: an empty one is none.
223pub(crate) fn reported(id: Option<String>) -> Option<String> {
224    id.filter(|id| !id.is_empty())
225}
226
227/// Parses an optional response body as JSON.
228///
229/// Returns:
230/// - `Ok(Some(value))` when a body is present and valid JSON.
231/// - `Ok(None)` when the body is absent or empty.
232/// - `Err(error)` when a body is present but isn't valid JSON.
233pub(crate) fn json(body: Option<&str>) -> Result<Option<serde_json::Value>, serde_json::Error> {
234    body.filter(|body| !body.is_empty())
235        .map(serde_json::from_str)
236        .transpose()
237}
238
239#[cfg(test)]
240mod tests;