use std::num::NonZeroU32;
use serde::Deserialize;
use super::{CompletionError, ModelCatalog, ModelDescriptor, ModelId, ThinkingMode};
use crate::Error;
use crate::dialects::{ToolDialectId, ToolsMode};
#[derive(Debug, Deserialize)]
struct ModelsListEntry {
id: String,
description: String,
context: u32,
thinking: ThinkingMode,
#[serde(default = "default_tool_dialect")]
tool_dialect: ToolDialectId,
#[serde(default)]
tools_mode: Option<ToolsMode>,
}
fn default_tool_dialect() -> ToolDialectId {
ToolDialectId::OpenAi
}
#[derive(Debug, Deserialize)]
struct ModelsListResponse {
data: Vec<ModelsListEntry>,
}
pub(crate) const MAX_CATALOG_ERROR_BODY: usize = 2000;
pub(crate) const MAX_CATALOG_BODY: u64 = 16 * 1024 * 1024;
async fn read_catalog_body_capped(
mut response: reqwest::Response,
cap: u64,
) -> std::result::Result<Vec<u8>, CompletionError> {
if let Some(len) = response.content_length()
&& len > cap
{
return Err(CompletionError::from(Error::MalformedResponse(format!(
"model list body of {len} bytes exceeds the {cap}-byte limit"
))));
}
let mut body: Vec<u8> = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(Error::http)? {
if body.len() as u64 + chunk.len() as u64 > cap {
return Err(CompletionError::from(Error::MalformedResponse(format!(
"model list body exceeds the {cap}-byte limit"
))));
}
body.extend_from_slice(&chunk);
}
Ok(body)
}
async fn read_error_body_bounded(
mut response: reqwest::Response,
limit: usize,
) -> std::result::Result<String, reqwest::Error> {
let mut buffer: Vec<u8> = Vec::new();
while buffer.len() < limit {
match response.chunk().await? {
Some(chunk) => {
let take = (limit - buffer.len()).min(chunk.len());
buffer.extend_from_slice(&chunk[..take]);
if take < chunk.len() {
break;
}
}
None => break,
}
}
if buffer.is_empty() {
return Ok("(empty body)".to_owned());
}
let lossy = String::from_utf8_lossy(&buffer);
let mut escaped = String::with_capacity(lossy.len());
for ch in lossy.chars() {
if ch.is_control() {
escaped.extend(ch.escape_default());
} else {
escaped.push(ch);
}
}
Ok(escaped)
}
fn catalog_client() -> reqwest::Client {
static CATALOG_CLIENT: std::sync::OnceLock<reqwest::Client> = std::sync::OnceLock::new();
CATALOG_CLIENT.get_or_init(reqwest::Client::new).clone()
}
pub async fn fetch_model_catalog(
base_url: &str,
token: &str,
) -> std::result::Result<ModelCatalog, CompletionError> {
let base = base_url.trim_end_matches('/');
let http = catalog_client();
let response = http
.get(format!("{base}/models"))
.bearer_auth(token)
.send()
.await
.map_err(Error::http)?;
let status = response.status();
if !status.is_success() {
let body = match read_error_body_bounded(response, MAX_CATALOG_ERROR_BODY).await {
Ok(body) => body,
Err(source) => {
return Err(CompletionError::from(Error::BackendBodyRead {
status: status.as_u16(),
source: Box::new(source),
}));
}
};
return Err(CompletionError::from(Error::Backend {
status: status.as_u16(),
body,
}));
}
let body = read_catalog_body_capped(response, MAX_CATALOG_BODY).await?;
let list: ModelsListResponse = serde_json::from_slice(&body).map_err(|error| {
CompletionError::from(Error::MalformedResponseSource {
message: "model list response was not valid JSON".to_owned(),
source: Box::new(error),
})
})?;
let mut descriptors = Vec::with_capacity(list.data.len());
for entry in list.data {
let id = ModelId::gateway(entry.id).map_err(|error| {
CompletionError::from(Error::MalformedResponse(format!(
"model catalog entry has an invalid id: {error}"
)))
})?;
let context = NonZeroU32::new(entry.context).ok_or_else(|| {
CompletionError::from(Error::MalformedResponse(format!(
"model {} declares a zero-token context window",
id.name()
)))
})?;
if let Some(wire_mode) = entry.tools_mode {
let derived = entry.tool_dialect.tools_mode();
if wire_mode != derived {
return Err(CompletionError::from(Error::MalformedResponse(format!(
"model {} wire tools_mode {wire_mode} contradicts dialect-derived {derived}",
id.name()
))));
}
}
descriptors.push(
ModelDescriptor::new(id, entry.description, context, entry.thinking)
.with_dialect(entry.tool_dialect),
);
}
ModelCatalog::new(descriptors).map_err(|error| {
CompletionError::from(Error::MalformedResponse(format!(
"gateway returned an inconsistent model catalog: {error}"
)))
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::CompletionErrorKind;
#[test]
fn models_list_entry_parses_dialect_fields() {
let json = serde_json::json!({
"id": "gemma-local",
"description": "A Gemma model",
"context": 32768,
"thinking": "never",
"tool_dialect": "gemma3_tool_code",
"tools_mode": "emulated"
});
let entry: ModelsListEntry = serde_json::from_value(json).unwrap();
assert_eq!(entry.tool_dialect, ToolDialectId::Gemma3ToolCode);
assert_eq!(entry.tool_dialect.tools_mode(), ToolsMode::Emulated);
}
#[test]
fn models_list_entry_defaults_to_openai_native() {
let json = serde_json::json!({
"id": "remote",
"description": "A remote model",
"context": 8192,
"thinking": "never"
});
let entry: ModelsListEntry = serde_json::from_value(json).unwrap();
assert_eq!(entry.tool_dialect, ToolDialectId::OpenAi);
assert_eq!(entry.tool_dialect.tools_mode(), ToolsMode::Native);
}
#[tokio::test]
async fn fetch_model_catalog_rejects_a_wire_tools_mode_that_contradicts_the_dialect() {
use axum::Router;
use axum::routing::get;
async fn models() -> axum::Json<serde_json::Value> {
axum::Json(serde_json::json!({
"data": [{
"id": "remote",
"description": "a remote model",
"context": 8192,
"thinking": "never",
"tool_dialect": "openai",
"tools_mode": "emulated"
}]
}))
}
let app = Router::new().route("/models", get(models));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let err = fetch_model_catalog(&format!("http://{addr}"), "tok")
.await
.expect_err("a contradictory wire tools_mode must be rejected");
assert_eq!(err.kind(), CompletionErrorKind::MalformedResponse);
assert!(
err.to_string().contains("contradicts"),
"the rejection must name the contradiction, got {err}"
);
}
#[tokio::test]
async fn fetch_model_catalog_bounds_and_reports_non_success_body() {
use axum::Router;
use axum::routing::get;
async fn models() -> (axum::http::StatusCode, String) {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"e".repeat(MAX_CATALOG_ERROR_BODY * 4),
)
}
let app = Router::new().route("/models", get(models));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let err = fetch_model_catalog(&format!("http://{addr}"), "tok")
.await
.expect_err("a 500 response must surface as an error");
assert_eq!(err.kind(), CompletionErrorKind::Backend);
let msg = err.to_string();
assert!(
msg.len() < MAX_CATALOG_ERROR_BODY + 128,
"the error-path body must be bounded, got {} bytes",
msg.len()
);
}
#[tokio::test]
async fn fetch_model_catalog_bounds_an_oversized_success_body() {
use axum::Router;
use axum::routing::get;
async fn models() -> (axum::http::StatusCode, String) {
let oversized = usize::try_from(MAX_CATALOG_BODY).unwrap() + 1;
(axum::http::StatusCode::OK, "e".repeat(oversized))
}
let app = Router::new().route("/models", get(models));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let err = fetch_model_catalog(&format!("http://{addr}"), "tok")
.await
.expect_err("an oversized success body must be refused");
assert_eq!(err.kind(), CompletionErrorKind::MalformedResponse);
assert!(
err.to_string().contains("exceeds"),
"the bound must report the size limit, got {err}"
);
}
#[tokio::test]
async fn fetch_model_catalog_preserves_the_json_decode_source() {
use axum::Router;
use axum::routing::get;
async fn models() -> (axum::http::StatusCode, String) {
(axum::http::StatusCode::OK, "{ this is not json".to_owned())
}
let app = Router::new().route("/models", get(models));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let err = fetch_model_catalog(&format!("http://{addr}"), "tok")
.await
.expect_err("an undecodable body must surface as an error");
assert_eq!(err.kind(), CompletionErrorKind::MalformedResponse);
let source =
std::error::Error::source(&err).expect("the decode error must be a preserved source");
assert!(
source.downcast_ref::<serde_json::Error>().is_some(),
"the preserved source must be the JSON decode error, got {source}"
);
}
#[tokio::test]
async fn fetch_model_catalog_preserves_a_body_read_failure_source() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
if let Ok((mut sock, _)) = listener.accept().await {
let mut buf = [0u8; 1024];
let _ = sock.read(&mut buf).await;
let header = "HTTP/1.1 500 Internal Server Error\r\n\
Content-Length: 1000000\r\n\r\n";
let _ = sock.write_all(header.as_bytes()).await;
let _ = sock.write_all(b"abc").await;
}
});
let err = fetch_model_catalog(&format!("http://{addr}"), "tok")
.await
.expect_err("a truncated error body must surface as an error");
assert_eq!(err.kind(), CompletionErrorKind::Transport);
assert_eq!(err.status(), Some(500));
let source =
std::error::Error::source(&err).expect("the read failure must be a preserved source");
assert!(
source.downcast_ref::<reqwest::Error>().is_some(),
"the preserved source must be the reqwest read error, got {source}"
);
}
}