use serde_json::{Map, Value};
use crate::config::errors::ConfigError;
use crate::config::partial::{or_map, OutMode, PartialConfig, PartialProvider};
use crate::config::provider::{AuthId, Provider};
use crate::config::resolved::ResolvedConfig;
impl PartialConfig {
pub fn into_resolved(self, req_model: Option<&str>) -> Result<ResolvedConfig, ConfigError> {
self.check_scalars()?;
let routing_model = req_model.or(self.model.as_deref());
let (name, mut partial) = self.route(routing_model)?;
let mut bd = std::mem::take(&mut partial.body_defaults);
let max_tokens = self.max_tokens.or(take_u32(&mut bd, "max_tokens")?);
let temperature = self.temperature.or(take_f32(&mut bd, "temperature")?);
let top_p = self.top_p.or(take_f32(&mut bd, "top_p")?);
let stream = self.stream.or(take_bool(&mut bd, "stream")?);
let extra = or_map(bd, self.extra);
let provider = complete(name, partial)?;
let model = match routing_model {
Some(m) => provider
.model_aliases
.get(m)
.cloned()
.unwrap_or_else(|| m.to_owned()),
None => String::new(),
};
Ok(ResolvedConfig {
provider,
model,
model_from_cache: false,
output: self.output.unwrap_or(OutMode::Text),
thinking: self.thinking.unwrap_or(false),
inline_key: self.api_key,
max_tokens,
temperature,
top_p,
reasoning: self.reasoning,
stream,
timeout_connect: self.timeout_connect,
timeout_response: self.timeout_response,
timeout_idle: self.timeout_idle,
system: self.system,
extra,
})
}
fn check_scalars(&self) -> Result<(), ConfigError> {
if self.max_tokens == Some(0) {
return Err(bad("max_tokens", "must be greater than zero"));
}
if self.temperature.is_some_and(f32::is_nan) {
return Err(bad("temperature", "must be a number"));
}
if self.top_p.is_some_and(f32::is_nan) {
return Err(bad("top_p", "must be a number"));
}
Ok(())
}
fn route(&self, routing_model: Option<&str>) -> Result<(String, PartialProvider), ConfigError> {
if let Some(name) = &self.provider {
let row = self
.providers
.get(name)
.ok_or_else(|| ConfigError::UnknownProvider { name: name.clone() })?;
return Ok((name.clone(), row.clone()));
}
let Some(model) = routing_model else {
return self
.default_provider
.as_deref()
.and_then(|name| self.providers.get_key_value(name))
.map(|(name, row)| (name.clone(), row.clone()))
.ok_or(ConfigError::NoProvider);
};
let mut matches: Vec<(String, PartialProvider)> = self
.providers
.iter()
.filter(|(_, row)| row_owns(row, model))
.map(|(name, row)| (name.clone(), row.clone()))
.collect();
if matches.is_empty() {
return Err(ConfigError::NoProvider);
}
if matches.len() > 1 {
return Err(ConfigError::AmbiguousModel {
model: model.to_owned(),
providers: matches.into_iter().map(|(name, _)| name).collect(),
});
}
Ok(matches.swap_remove(0))
}
}
fn row_owns(row: &PartialProvider, model: &str) -> bool {
let aliased = row
.model_aliases
.as_ref()
.is_some_and(|a| a.contains_key(model));
let prefixed = row
.model_prefixes
.as_ref()
.is_some_and(|ps| ps.iter().any(|p| model.starts_with(p.as_str())));
aliased || prefixed
}
fn bad(key: &str, detail: &str) -> ConfigError {
ConfigError::BadValue {
key: key.to_owned(),
detail: detail.to_owned(),
}
}
fn complete(name: String, row: PartialProvider) -> Result<Provider, ConfigError> {
let need = |field| ConfigError::IncompleteProvider {
name: name.clone(),
field,
};
let base_url = row.base_url.ok_or_else(|| need("base_url"))?;
let protocol = row.protocol.ok_or_else(|| need("protocol"))?;
let auth = row.auth.ok_or_else(|| need("auth"))?;
let api_header = row.api_header;
if auth != AuthId::None && api_header.is_none() {
return Err(need("api_header"));
}
if auth == AuthId::OAuth2 && row.oauth.is_none() {
return Err(need("oauth"));
}
Ok(Provider {
base_url,
protocol,
auth,
api_header,
beta_headers: row.beta_headers.unwrap_or_default(),
model_aliases: row.model_aliases.unwrap_or_default(),
unsupported_body_keys: row.unsupported_body_keys.unwrap_or_default(),
models: row.models,
oauth: row.oauth,
ambient: row.ambient,
name,
})
}
fn take_u32(bd: &mut Map<String, Value>, key: &str) -> Result<Option<u32>, ConfigError> {
match bd.remove(key) {
None => Ok(None),
Some(v) => v
.as_u64()
.filter(|n| *n > 0 && *n <= u64::from(u32::MAX))
.map(|n| Some(n as u32))
.ok_or_else(|| bad(key, "must be a positive integer")),
}
}
fn take_f32(bd: &mut Map<String, Value>, key: &str) -> Result<Option<f32>, ConfigError> {
match bd.remove(key) {
None => Ok(None),
Some(v) => v
.as_f64()
.map(|f| Some(f as f32))
.ok_or_else(|| bad(key, "must be a number")),
}
}
fn take_bool(bd: &mut Map<String, Value>, key: &str) -> Result<Option<bool>, ConfigError> {
match bd.remove(key) {
None => Ok(None),
Some(v) => v
.as_bool()
.map(Some)
.ok_or_else(|| bad(key, "must be a boolean")),
}
}