use std::io::Write;
use crate::auth::encode_pairs;
use crate::canonical::{CanonicalError, ErrorKind, Model};
use crate::config::provider::ModelsOverride;
use crate::config::{
config_path, defaults, partial_from_env, read_config_file, OutMode, ResolvedConfig,
};
use crate::protocol::{decode_models, http_error, ModelsShape, WireRequest};
use crate::registry::Registry;
use crate::store::{Clock, CredStore, ModelCache};
use crate::transport::Transport;
use super::drain;
use super::events::is_2xx;
pub struct ListIo<'a> {
pub stdout: &'a mut dyn Write,
pub stderr: &'a mut dyn Write,
pub transport: &'a dyn Transport,
pub store: &'a dyn CredStore,
pub cache: &'a dyn ModelCache,
pub clock: &'a dyn Clock,
}
pub fn list_models(args: &crate::cli::Args, io: &mut ListIo) -> u8 {
match run_list(args, io) {
Ok(code) => code,
Err(e) => {
let _ = writeln!(io.stderr, "{}", e.message);
e.exit_code()
}
}
}
fn run_list(args: &crate::cli::Args, io: &mut ListIo) -> Result<u8, CanonicalError> {
let flags = crate::cli::parse_args(&args.argv)?;
if flags.help {
return Ok(super::emit(io.stdout, super::HELP));
}
if flags.version {
return Ok(super::emit(io.stdout, super::VERSION_LINE));
}
let file = read_config_file(&config_path(flags.config_path, &args.env))?;
let env = partial_from_env(&args.env).map_err(CanonicalError::from)?;
let merged = flags.config.or(env).or(file).or(defaults());
let cfg: ResolvedConfig = merged.into_resolved(None).map_err(CanonicalError::from)?;
let json = cfg.output == OutMode::Ndjson;
let models = fetch_models(&cfg, io.transport, io.store, io.clock)?;
io.cache.put(&cfg.provider.name, &models);
print_models(io.stdout, &models, json).map_err(write_failed)?;
if models.is_empty() {
let _ = writeln!(io.stderr, "no models returned for `{}`", cfg.provider.name);
}
Ok(0)
}
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 req = models_req(
proto.models_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);
}
wire.timeouts = cfg.timeouts();
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.array_key, req.id_key, req.strip)
}
pub(crate) fn models_req<'a>(
shape: ModelsShape,
over: Option<&'a ModelsOverride>,
base_url: &str,
) -> ModelsReq<'a> {
let path = over.and_then(|m| m.path.as_deref()).unwrap_or(shape.path);
let array_key = over
.and_then(|m| m.array_key.as_deref())
.unwrap_or(shape.array_key);
let id_key = over
.and_then(|m| m.id_key.as_deref())
.unwrap_or(shape.id_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,
array_key,
id_key,
strip: shape.strip,
}
}
pub(crate) struct ModelsReq<'a> {
pub(crate) url: String,
pub(crate) array_key: &'a str,
pub(crate) id_key: &'a str,
pub(crate) strip: &'a str,
}
fn print_models(out: &mut dyn Write, models: &[Model], json: bool) -> std::io::Result<()> {
if json {
let obj = serde_json::json!({ "models": models });
writeln!(out, "{obj}")
} else {
for m in models {
let suffix = if m.default { " (default)" } else { "" };
writeln!(out, "{}{suffix}", m.id)?;
}
Ok(())
}
}
fn read_failed(e: std::io::Error) -> CanonicalError {
CanonicalError {
kind: ErrorKind::Transport,
message: format!("failed to read models response body: {e}"),
provider_detail: None,
}
}
fn write_failed(e: std::io::Error) -> CanonicalError {
CanonicalError {
kind: ErrorKind::Transport,
message: format!("failed to write model list: {e}"),
provider_detail: None,
}
}
#[cfg(test)]
mod tests {
use super::print_models;
use crate::canonical::Model;
#[test]
fn text_suffixes_the_default_flagged_id() {
let models = [
Model {
id: "fast".into(),
default: false,
},
Model {
id: "smart".into(),
default: true,
},
];
let mut out = Vec::new();
print_models(&mut out, &models, false).unwrap();
assert_eq!(String::from_utf8(out).unwrap(), "fast\nsmart (default)\n");
}
}