use crate::auth::encode_pairs;
use crate::canonical::{CanonicalError, ErrorKind, Model};
use crate::config::provider::ModelsOverride;
use crate::config::ResolvedConfig;
use crate::protocol::{decode_models, http_error, ModelKeys, ModelsShape, WireRequest};
use crate::registry::Registry;
use crate::run::{drain, events::is_2xx};
use crate::store::{Clock, CredStore};
use crate::transport::Transport;
pub(super) fn fetch_models(
cfg: &ResolvedConfig,
transport: &dyn Transport,
store: &dyn CredStore,
clock: &dyn Clock,
) -> Result<Vec<Model>, CanonicalError> {
let registry = Registry::builtin();
let proto = registry.protocol(cfg.provider.protocol);
let auth = registry.auth(cfg.provider.auth);
let beta: Vec<(&str, &str)> = cfg
.provider
.beta_headers
.iter()
.map(|(k, v)| (k.as_str(), v.as_str()))
.collect();
let ctx = cfg.provider_ctx(&beta);
let authc = cfg.auth_ctx();
let Some(shape) = proto.models_shape() else {
return Err(no_listing(&cfg.provider.name));
};
let req = models_req(shape, cfg.provider.models.as_ref(), ctx.base_url);
let mut wire = WireRequest::get(req.url);
for (k, v) in &beta {
wire.set_header(k, v);
}
cfg.stamp_transport(&mut wire);
auth.apply(&mut wire, &ctx, &authc, store, clock, transport)?;
let resp = transport.send(wire)?;
let status = resp.status;
if !is_2xx(status) {
let body = drain(resp.body).unwrap_or_default();
return Err(http_error(&body, status));
}
let body = drain(resp.body).map_err(read_failed)?;
decode_models(&body, &req.keys)
}
pub(crate) fn models_req<'a>(
shape: ModelsShape,
over: Option<&'a ModelsOverride>,
base_url: &str,
) -> ModelsReq<'a> {
let d = shape.keys;
let pick = |o: Option<&'a String>, def: &'a str| o.map(String::as_str).unwrap_or(def);
let path = over.and_then(|m| m.path.as_deref()).unwrap_or(shape.path);
let keys = ModelKeys {
array_key: pick(over.and_then(|m| m.array_key.as_ref()), d.array_key),
id_key: pick(over.and_then(|m| m.id_key.as_ref()), d.id_key),
strip: d.strip, context_key: pick(over.and_then(|m| m.context_key.as_ref()), d.context_key),
max_output_key: pick(
over.and_then(|m| m.max_output_key.as_ref()),
d.max_output_key,
),
display_name_key: pick(
over.and_then(|m| m.display_name_key.as_ref()),
d.display_name_key,
),
};
let query = over.map(|m| m.query.as_slice()).unwrap_or(&[]);
let url = if query.is_empty() {
format!("{base_url}{path}")
} else {
let pairs: Vec<(&str, &str)> = query
.iter()
.map(|(k, v)| (k.as_str(), v.as_str()))
.collect();
format!("{base_url}{path}?{}", encode_pairs(&pairs))
};
ModelsReq { url, keys }
}
pub(crate) struct ModelsReq<'a> {
pub(crate) url: String,
pub(crate) keys: ModelKeys<'a>,
}
fn no_listing(provider: &str) -> CanonicalError {
CanonicalError {
kind: ErrorKind::Config,
message: format!(
"provider `{provider}` has no models listing; pass --model verbatim — \
a model that succeeds is learned into the cache"
),
provider_detail: None,
retry_after_seconds: None,
}
}
fn read_failed(e: std::io::Error) -> CanonicalError {
CanonicalError {
kind: ErrorKind::Transport,
message: format!("failed to read models response body: {e}"),
provider_detail: None,
retry_after_seconds: None,
}
}