use std::collections::HashMap;
use std::sync::LazyLock;
use anyhow::Context;
use astrid_capsule::manifest::OptionsFrom;
use regex::Regex;
const PROVIDER_BASE_URL_KEY: &str = "base_url";
const MAX_RESPONSE_BYTES: usize = 5 * 1024 * 1024;
static PLACEHOLDER: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"\{(\w+)\}").expect("static placeholder regex is valid"));
pub(crate) fn resolve_template(template: &str, values: &HashMap<String, String>) -> String {
PLACEHOLDER
.replace_all(template, |caps: ®ex::Captures<'_>| {
let key = &caps[1];
match values.get(key) {
Some(value) => value.clone(),
None => caps[0].to_string(),
}
})
.into_owned()
}
pub(crate) fn should_send_bearer(http_url: &str, base_url: &str) -> bool {
let Ok(http) = reqwest::Url::parse(http_url) else {
return false;
};
let Ok(base) = reqwest::Url::parse(base_url) else {
return false;
};
let (Some(http_host), Some(base_host)) = (http.host_str(), base.host_str()) else {
return false;
};
if http.scheme() != base.scheme() {
return false;
}
http_host.eq_ignore_ascii_case(base_host)
&& http.port_or_known_default() == base.port_or_known_default()
}
pub(crate) fn parse_options_response(body: &str, select_hint: &str) -> Vec<String> {
let Ok(json) = serde_json::from_str::<serde_json::Value>(body) else {
return Vec::new();
};
let _ = select_hint;
let Some(data) = json.get("data").and_then(serde_json::Value::as_array) else {
return Vec::new();
};
let mut seen = std::collections::HashSet::new();
let mut ids = Vec::new();
for entry in data {
let Some(id) = entry.get("id").and_then(serde_json::Value::as_str) else {
continue;
};
let id = id.trim();
if id.is_empty() {
continue;
}
if seen.insert(id.to_string()) {
ids.push(id.to_string());
}
}
ids
}
pub(crate) async fn fetch_options(
opts: &OptionsFrom,
values: &HashMap<String, String>,
) -> anyhow::Result<Vec<String>> {
let url = resolve_template(&opts.http, values);
anyhow::ensure!(
!url.contains('{'),
"endpoint still contains unresolved placeholders: {url}"
);
anyhow::ensure!(
url.starts_with("http://") || url.starts_with("https://"),
"endpoint is not an http(s) URL: {url}"
);
let bearer = opts
.bearer
.as_ref()
.map(|b| resolve_template(b, values))
.map(|b| b.trim().to_string())
.filter(|b| !b.is_empty());
let bearer = bearer.filter(|_| match values.get(PROVIDER_BASE_URL_KEY) {
Some(base_url) if should_send_bearer(&url, base_url) => true,
_ => {
tracing::warn!(
endpoint = %url,
"withholding options-from bearer: fetch host does not match the \
configured provider ({PROVIDER_BASE_URL_KEY}) host"
);
false
},
});
let client = reqwest::Client::builder()
.user_agent("astrid-cli")
.timeout(std::time::Duration::from_secs(15))
.build()?;
let mut request = client.get(&url);
if let Some(token) = bearer {
request = request.bearer_auth(token);
}
let response = request.send().await?;
anyhow::ensure!(
response.status().is_success(),
"models endpoint returned HTTP {}",
response.status()
);
anyhow::ensure!(
response
.content_length()
.is_none_or(|len| len <= MAX_RESPONSE_BYTES as u64),
"models response too large (advertised {} bytes; limit {MAX_RESPONSE_BYTES})",
response.content_length().unwrap_or_default()
);
let body = read_capped_body(response).await?;
let options = parse_options_response(&body, opts.select_or_default());
anyhow::ensure!(
!options.is_empty(),
"models endpoint returned no usable options"
);
Ok(options)
}
async fn read_capped_body(response: reqwest::Response) -> anyhow::Result<String> {
use futures::StreamExt;
let mut stream = response.bytes_stream();
let mut body: Vec<u8> = Vec::new();
while let Some(chunk) = stream.next().await {
let chunk = chunk.context("error reading models response body")?;
anyhow::ensure!(
body.len().saturating_add(chunk.len()) <= MAX_RESPONSE_BYTES,
"models response too large (exceeded {MAX_RESPONSE_BYTES} bytes; aborted mid-stream)"
);
body.extend_from_slice(&chunk);
}
String::from_utf8(body).context("models response was not valid UTF-8")
}
#[cfg(test)]
fn decode_capped_body(body: &[u8]) -> anyhow::Result<String> {
anyhow::ensure!(
body.len() <= MAX_RESPONSE_BYTES,
"models response too large ({} bytes; limit {MAX_RESPONSE_BYTES})",
body.len()
);
String::from_utf8(body.to_vec()).context("models response was not valid UTF-8")
}
pub(crate) fn fetch_options_blocking(
opts: &OptionsFrom,
values: &HashMap<String, String>,
) -> anyhow::Result<Vec<String>> {
std::thread::scope(|scope| {
scope
.spawn(|| {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.context("failed to build discovery runtime")?;
runtime.block_on(fetch_options(opts, values))
})
.join()
.map_err(|_| anyhow::anyhow!("model discovery thread panicked"))?
})
}
#[cfg(test)]
mod tests {
use super::*;
fn vals(pairs: &[(&str, &str)]) -> HashMap<String, String> {
pairs
.iter()
.map(|(k, v)| ((*k).to_string(), (*v).to_string()))
.collect()
}
#[test]
fn resolve_template_substitutes_known_keys() {
let v = vals(&[("base_url", "https://api.openai.com"), ("api_key", "sk-x")]);
assert_eq!(
resolve_template("{base_url}/v1/models", &v),
"https://api.openai.com/v1/models"
);
assert_eq!(resolve_template("{api_key}", &v), "sk-x");
}
#[test]
fn resolve_template_substring_keys_are_order_independent() {
let template = "{base_url}?a={api}&k={api_key}";
let expected = "https://h?a=AAA&k=KKK";
let mut v1 = HashMap::new();
v1.insert("api".to_string(), "AAA".to_string());
v1.insert("api_key".to_string(), "KKK".to_string());
v1.insert("base_url".to_string(), "https://h".to_string());
let mut v2 = HashMap::new();
v2.insert("api_key".to_string(), "KKK".to_string());
v2.insert("base_url".to_string(), "https://h".to_string());
v2.insert("api".to_string(), "AAA".to_string());
assert_eq!(resolve_template(template, &v1), expected);
assert_eq!(resolve_template(template, &v2), expected);
}
#[test]
fn resolve_template_does_not_rescan_substituted_value() {
let v = vals(&[("base_url", "https://h/{api_key}"), ("api_key", "secret")]);
assert_eq!(
resolve_template("{base_url}/v1", &v),
"https://h/{api_key}/v1"
);
}
#[test]
fn resolve_template_leaves_unknown_keys() {
let v = vals(&[("known", "x")]);
assert_eq!(
resolve_template("{known}/{unknown}", &v),
"x/{unknown}",
"unresolved placeholder must remain so the caller can detect the miss"
);
}
#[test]
fn parse_extracts_ids_in_server_order() {
let body = r#"{ "data": [ { "id": "gpt-4o" }, { "id": "gpt-4o-mini" }, { "id": "o1" } ] }"#;
assert_eq!(
parse_options_response(body, "data[].id"),
vec!["gpt-4o", "gpt-4o-mini", "o1"]
);
}
#[test]
fn parse_dedupes_preserving_first_occurrence() {
let body = r#"{ "data": [ { "id": "a" }, { "id": "b" }, { "id": "a" } ] }"#;
assert_eq!(parse_options_response(body, "data[].id"), vec!["a", "b"]);
}
#[test]
fn parse_drops_blank_and_missing_ids() {
let body = r#"{ "data": [ { "id": " " }, { "id": "real" }, { "name": "no-id" }, { "id": "" } ] }"#;
assert_eq!(parse_options_response(body, "data[].id"), vec!["real"]);
}
#[test]
fn parse_unknown_hint_falls_back_to_data_id() {
let body = r#"{ "data": [ { "id": "m1" } ] }"#;
assert_eq!(parse_options_response(body, "models[].name"), vec!["m1"]);
}
#[test]
fn parse_returns_empty_on_non_json() {
assert!(parse_options_response("<html>not json</html>", "data[].id").is_empty());
}
#[test]
fn parse_returns_empty_on_missing_data() {
assert!(parse_options_response(r#"{ "object": "list" }"#, "data[].id").is_empty());
}
#[test]
fn parse_returns_empty_when_data_not_array() {
assert!(parse_options_response(r#"{ "data": "oops" }"#, "data[].id").is_empty());
}
#[tokio::test]
async fn fetch_rejects_unresolved_template_before_network() {
let opts = OptionsFrom {
http: "{base_url}/v1/models".to_string(),
bearer: None,
select: None,
after: vec!["base_url".to_string()],
};
let err = fetch_options(&opts, &HashMap::new())
.await
.expect_err("unresolved placeholder must error");
assert!(
err.to_string().contains("unresolved placeholder"),
"got: {err}"
);
}
#[test]
fn bearer_sent_when_fetch_host_matches_configured_provider() {
assert!(should_send_bearer(
"https://api.openai.com/v1/models",
"https://api.openai.com"
));
assert!(should_send_bearer(
"https://API.OpenAI.com/v1/models",
"https://api.openai.com"
));
assert!(should_send_bearer(
"https://provider.example:8443/v1/models",
"https://provider.example:8443"
));
}
#[test]
fn bearer_withheld_when_fetch_host_differs_from_provider() {
assert!(!should_send_bearer(
"https://attacker.com/v1/models",
"https://api.openai.com"
));
assert!(!should_send_bearer(
"https://api.openai.com.attacker.com/v1/models",
"https://api.openai.com"
));
assert!(!should_send_bearer(
"https://api.openai.com:8443/v1/models",
"https://api.openai.com"
));
assert!(!should_send_bearer(
"http://api.openai.com/v1/models",
"https://api.openai.com"
));
assert!(!should_send_bearer(
"http://api.openai.com:443/v1/models",
"https://api.openai.com"
));
assert!(!should_send_bearer(
"https://api.openai.com:80/v1/models",
"http://api.openai.com"
));
}
#[test]
fn bearer_withheld_when_either_url_is_unparseable() {
assert!(!should_send_bearer(
"https://api.openai.com/v1/models",
"not a url"
));
assert!(!should_send_bearer("not a url", "https://api.openai.com"));
assert!(!should_send_bearer(
"https://api.openai.com/v1/models",
"mailto:ops@example.com"
));
}
#[test]
fn capped_body_decodes_within_limit() {
let body = br#"{ "data": [ { "id": "gpt-4o" } ] }"#;
let decoded = decode_capped_body(body).expect("within-limit body must decode");
assert!(decoded.contains("gpt-4o"));
}
#[test]
fn capped_body_rejects_over_limit() {
assert_eq!(MAX_RESPONSE_BYTES, 5 * 1024 * 1024);
let oversized = vec![b'x'; MAX_RESPONSE_BYTES + 1];
let err = decode_capped_body(&oversized).expect_err("over-cap body must error");
assert!(err.to_string().contains("too large"), "got: {err}");
let at_limit = vec![b'x'; MAX_RESPONSE_BYTES];
assert!(decode_capped_body(&at_limit).is_ok());
}
#[tokio::test]
async fn fetch_rejects_non_http_endpoint() {
let opts = OptionsFrom {
http: "file:///etc/passwd".to_string(),
bearer: None,
select: None,
after: vec![],
};
let err = fetch_options(&opts, &HashMap::new())
.await
.expect_err("non-http endpoint must error");
assert!(err.to_string().contains("not an http"), "got: {err}");
}
#[tokio::test]
async fn fetch_rejects_oversized_chunked_body_without_buffering_it_all() {
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let port = listener.local_addr().expect("local_addr").port();
let sent = Arc::new(AtomicUsize::new(0));
let sent_srv = Arc::clone(&sent);
let server = tokio::spawn(async move {
let (mut sock, _) = listener.accept().await.expect("accept");
let mut buf = [0u8; 1024];
let _ = sock.read(&mut buf).await;
let head = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nTransfer-Encoding: chunked\r\n\r\n";
if sock.write_all(head).await.is_err() {
return;
}
let chunk_payload = vec![b'x'; 64 * 1024];
let frame_header = format!("{:x}\r\n", chunk_payload.len());
let max_chunks = (MAX_RESPONSE_BYTES / chunk_payload.len()) + 16;
for _ in 0..max_chunks {
if sock.write_all(frame_header.as_bytes()).await.is_err() {
break;
}
if sock.write_all(&chunk_payload).await.is_err() {
break;
}
if sock.write_all(b"\r\n").await.is_err() {
break;
}
sent_srv.fetch_add(chunk_payload.len(), Ordering::SeqCst);
}
let _ = sock.write_all(b"0\r\n\r\n").await;
});
let opts = OptionsFrom {
http: format!("http://127.0.0.1:{port}/v1/models"),
bearer: None,
select: None,
after: vec![],
};
let err = fetch_options(&opts, &HashMap::new())
.await
.expect_err("oversized chunked body must error");
assert!(
err.to_string().contains("too large"),
"expected a size-cap error, got: {err}"
);
let _ = server.await;
let total_sent = sent.load(Ordering::SeqCst);
assert!(
total_sent < MAX_RESPONSE_BYTES + 4 * 1024 * 1024,
"server sent {total_sent} bytes; client did not abort the stream early enough"
);
}
}