Skip to main content

harness_gateway_client/
catalog.rs

1//! Catalog transport: fetching and decoding gateway `GET /v1/models`.
2
3use std::num::NonZeroU32;
4
5use promptforge::model::{CompletionError, ModelCatalog, ModelDescriptor, ModelId, ThinkingMode};
6use serde::Deserialize;
7
8use crate::failure::{malformed, transport_failure};
9use crate::wire::classify::classify_http_failure;
10use crate::wire::stream::escape_controls;
11
12/// Wire shape of one entry from gateway `GET /v1/models`.
13///
14/// The list mixes inference models with the gateway's speech-to-text models,
15/// which have only `id`, `object`, and `kind` because they answer no
16/// completion request. The inference fields are therefore optional at the
17/// wire, and an entry without a context window is skipped rather than
18/// failing the whole catalog.
19#[derive(Debug, Deserialize)]
20struct ModelsListEntry {
21    id: String,
22    #[serde(default)]
23    description: String,
24    #[serde(default)]
25    context: Option<u32>,
26    #[serde(default)]
27    thinking: Option<ThinkingMode>,
28}
29
30/// Wire shape of gateway `GET /v1/models`.
31#[derive(Debug, Deserialize)]
32struct ModelsListResponse {
33    data: Vec<ModelsListEntry>,
34}
35
36/// The largest gateway error body kept for a catalog-fetch diagnostic, in bytes.
37const MAX_CATALOG_ERROR_BODY: usize = 2000;
38
39/// The largest success-path model-catalog body accepted before decoding, in
40/// bytes. A gateway that returns more than this is refused rather than buffered
41/// unbounded, mirroring the bound the error path already applies. Sized well
42/// above any realistic model list (16 MiB) so legitimate catalogs are unaffected.
43const MAX_CATALOG_BODY: u64 = 16 * 1024 * 1024;
44
45/// Reads a success-path response body, refusing it once it would exceed `cap`
46/// bytes so a decode cannot buffer an unbounded body first.
47///
48/// The advertised `Content-Length` short-circuits an oversize body, and the
49/// streamed chunks are bounded so a gateway that omits or lies about the length
50/// still cannot force an unbounded allocation.
51async fn read_catalog_body_capped(
52    mut response: reqwest::Response,
53    cap: u64,
54) -> std::result::Result<Vec<u8>, CompletionError> {
55    if let Some(len) = response.content_length()
56        && len > cap
57    {
58        return Err(malformed(format!(
59            "model list body of {len} bytes exceeds the {cap}-byte limit"
60        )));
61    }
62    let mut body: Vec<u8> = Vec::new();
63    while let Some(chunk) = response.chunk().await.map_err(transport_failure)? {
64        if body.len() as u64 + chunk.len() as u64 > cap {
65            return Err(malformed(format!(
66                "model list body exceeds the {cap}-byte limit"
67            )));
68        }
69        body.extend_from_slice(&chunk);
70    }
71    Ok(body)
72}
73
74/// Reads at most `limit` bytes of a non-success response body, stopping early so
75/// an oversized error body cannot exhaust memory.
76///
77/// A read failure is returned as the concrete [`reqwest::Error`] (MODEL-010) so
78/// the caller can retain it as an error-chain `#[source]`, rather than being
79/// flattened into display text that severs the cause.
80async fn read_error_body_bounded(
81    mut response: reqwest::Response,
82    limit: usize,
83) -> std::result::Result<String, reqwest::Error> {
84    let mut buffer: Vec<u8> = Vec::new();
85    while buffer.len() < limit {
86        match response.chunk().await? {
87            Some(chunk) => {
88                let take = (limit - buffer.len()).min(chunk.len());
89                buffer.extend_from_slice(&chunk[..take]);
90                if take < chunk.len() {
91                    break;
92                }
93            }
94            None => break,
95        }
96    }
97    if buffer.is_empty() {
98        return Ok("(empty body)".to_owned());
99    }
100    // F5: escape control characters so a hostile catalog error body cannot forge
101    // log lines or smuggle terminal control sequences into a diagnostic.
102    let lossy = String::from_utf8_lossy(&buffer);
103    let mut escaped = String::with_capacity(lossy.len());
104    for ch in lossy.chars() {
105        if ch.is_control() {
106            escaped.extend(ch.escape_default());
107        } else {
108            escaped.push(ch);
109        }
110    }
111    Ok(escaped)
112}
113
114/// Returns the process-wide catalog HTTP client, building it once on first use.
115///
116/// A single reusable client (MODEL-018) lets catalog fetches share one
117/// connection pool and transport configuration rather than each constructing a
118/// throwaway client with its own pool. The returned handle is a cheap clone of
119/// the shared client (its state is reference-counted internally).
120fn catalog_client() -> reqwest::Client {
121    static CATALOG_CLIENT: std::sync::OnceLock<reqwest::Client> = std::sync::OnceLock::new();
122    CATALOG_CLIENT.get_or_init(reqwest::Client::new).clone()
123}
124
125/// Sends a bearer-authed GET through the shared client (MODEL-018) and
126/// returns the success response, classifying every failure the same way for
127/// each gateway endpoint: `Transport` (or `Timeout`) when the send fails, the
128/// classified kind with a bounded, control-escaped body as its detail on a
129/// non-success status (MODEL-010: no unbounded buffering), and `Transport`
130/// (or `Timeout`) when that error body cannot be read, keeping the
131/// [`reqwest::Error`] as a typed source under the same timeout marking as a
132/// send failure.
133async fn get_authed(
134    url: String,
135    token: &str,
136) -> std::result::Result<reqwest::Response, CompletionError> {
137    let response = catalog_client()
138        .get(url)
139        .bearer_auth(token)
140        .send()
141        .await
142        .map_err(transport_failure)?;
143    let status = response.status();
144    if status.is_success() {
145        return Ok(response);
146    }
147    let body = match read_error_body_bounded(response, MAX_CATALOG_ERROR_BODY).await {
148        Ok(body) => body,
149        Err(source) => return Err(transport_failure(source)),
150    };
151    Err(classify_http_failure(status.as_u16(), &body))
152}
153
154/// Fetches a [`ModelCatalog`] from the Gateway's `/models` endpoint.
155///
156/// `base_url` is the root of the Gateway's OpenAI-compatible API, for
157/// example `http://127.0.0.1:8081/v1`. The request sends `token` as a
158/// bearer token.
159///
160/// # Errors
161/// Returns a [`CompletionError`]. Its [`kind`](CompletionError::kind) is:
162///
163/// - `Transport` or `Timeout` when the HTTP request or a response read fails;
164/// - the kind [`classify_http_failure`] picks for the response when the
165///   Gateway returns a status outside the 2xx range;
166/// - `MalformedResponse` when the body is not a valid model list.
167pub async fn fetch_model_catalog(
168    base_url: &str,
169    token: &str,
170) -> std::result::Result<ModelCatalog, CompletionError> {
171    let base = base_url.trim_end_matches('/');
172    let response = get_authed(format!("{base}/models"), token).await?;
173    // Bound the success body BEFORE decoding so an oversized (or unbounded)
174    // model list cannot exhaust memory, matching the bound the error path applies.
175    let body = read_catalog_body_capped(response, MAX_CATALOG_BODY).await?;
176    // A body that does not decode as a model list is a malformed response, not a
177    // transport failure - matching this function's documented error contract.
178    let list: ModelsListResponse = serde_json::from_slice(&body).map_err(|error| {
179        // MODEL-009: keep the decode error as a private `#[source]` cause instead
180        // of flattening it into the message, while the classification stays
181        // `MalformedResponse`.
182        malformed("model list response was not valid JSON").with_source(error)
183    })?;
184    let mut descriptors = Vec::with_capacity(list.data.len());
185    for entry in list.data {
186        // An entry with no context window is not an inference model (the
187        // gateway lists its transcription models here too); it is not a
188        // descriptor and must not fail the catalog.
189        let Some(context) = entry.context else {
190            continue;
191        };
192        let id = ModelId::gateway(entry.id).map_err(|error| {
193            malformed(format!("model catalog entry has an invalid id: {error}"))
194        })?;
195        let context = NonZeroU32::new(context).ok_or_else(|| {
196            malformed("a model declares a zero-token context window")
197                .with_detail(escape_controls(id.name(), MAX_CATALOG_ERROR_BODY))
198        })?;
199        let thinking = entry.thinking.ok_or_else(|| {
200            malformed("a model declares a context window but no thinking mode")
201                .with_detail(escape_controls(id.name(), MAX_CATALOG_ERROR_BODY))
202        })?;
203        descriptors.push(ModelDescriptor::new(
204            id,
205            entry.description,
206            context,
207            thinking,
208        ));
209    }
210    ModelCatalog::new(descriptors).map_err(|error| {
211        malformed("gateway returned an inconsistent model catalog")
212            .with_detail(escape_controls(&error.to_string(), MAX_CATALOG_ERROR_BODY))
213    })
214}
215
216#[cfg(test)]
217#[path = "catalog-tests.rs"]
218mod tests;