use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use serde::de::{self, Deserializer, MapAccess, SeqAccess, Visitor};
use serde::Deserialize;
use serde_json::{Map, Value};
use crate::auth::OAuthConfig;
use crate::config::partial::{PartialConfig, PartialProvider};
use crate::config::provider::{AuthId, HeaderSpec, ModelsOverride, ProtocolId, TransportSpec};
use crate::store::AmbientSpec;
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct ProviderRow {
name: String,
base_url: Option<String>,
exec: Option<String>,
transport: Option<TransportSpec>,
protocol: Option<ProtocolId>,
auth: Option<AuthId>,
beta_headers: Option<Vec<(String, String)>>,
generation_query: Option<Vec<(String, String)>>,
api_header: Option<HeaderSpec>,
model_aliases: Option<BTreeMap<String, String>>,
model_prefixes: Option<Vec<String>>,
#[serde(default)]
body_defaults: Map<String, Value>,
unsupported_body_keys: Option<Vec<String>>,
models: Option<ModelsOverride>,
oauth: Option<OAuthConfig>,
ambient: Option<AmbientSpec>,
}
impl ProviderRow {
fn into_pair(self) -> (String, PartialProvider) {
(
self.name,
PartialProvider {
base_url: self.base_url,
exec: self.exec,
transport: self.transport,
protocol: self.protocol,
auth: self.auth,
api_header: self.api_header,
beta_headers: self.beta_headers,
generation_query: self.generation_query,
model_aliases: self.model_aliases,
model_prefixes: self.model_prefixes,
body_defaults: self.body_defaults,
unsupported_body_keys: self.unsupported_body_keys,
models: self.models,
oauth: self.oauth,
ambient: self.ambient,
},
)
}
}
enum ProviderField {
Selector(String),
Rows(Vec<ProviderRow>),
}
impl<'de> Deserialize<'de> for ProviderField {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
d.deserialize_any(ProviderFieldVisitor)
}
}
struct ProviderFieldVisitor;
impl<'de> Visitor<'de> for ProviderFieldVisitor {
type Value = ProviderField;
fn expecting(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.write_str("a provider name or a list of [[provider]] tables")
}
fn visit_str<E: de::Error>(self, v: &str) -> Result<ProviderField, E> {
Ok(ProviderField::Selector(v.to_owned()))
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<ProviderField, A::Error> {
let mut rows = Vec::new();
while let Some(row) = seq.next_element()? {
rows.push(row);
}
Ok(ProviderField::Rows(rows))
}
}
impl<'de> Deserialize<'de> for PartialConfig {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
d.deserialize_map(PartialConfigVisitor)
}
}
struct PartialConfigVisitor;
impl<'de> Visitor<'de> for PartialConfigVisitor {
type Value = PartialConfig;
fn expecting(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.write_str("a brazen config table")
}
fn visit_map<M: MapAccess<'de>>(self, mut map: M) -> Result<PartialConfig, M::Error> {
let mut cfg = PartialConfig::default();
while let Some(key) = map.next_key::<String>()? {
match key.as_str() {
"provider" => match map.next_value::<ProviderField>()? {
ProviderField::Selector(name) => cfg.provider = Some(name),
ProviderField::Rows(rows) => {
let mut seen: BTreeSet<String> = BTreeSet::new();
for row in rows {
let (name, partial) = row.into_pair();
if !seen.insert(name.clone()) {
return Err(de::Error::custom(format!(
"duplicate provider name `{name}`"
)));
}
cfg.providers.push((name, partial));
}
}
},
"model" => cfg.model = Some(map.next_value()?),
"base_url" => cfg.base_url = Some(map.next_value()?),
"api_key" => cfg.api_key = Some(map.next_value()?),
"output" => cfg.output = Some(map.next_value()?),
"thinking" => cfg.thinking = Some(map.next_value()?),
"max_tokens" => cfg.max_tokens = Some(map.next_value()?),
"temperature" => cfg.temperature = Some(map.next_value()?),
"top_p" => cfg.top_p = Some(map.next_value()?),
"reasoning" => cfg.reasoning = Some(map.next_value()?),
"stream" => cfg.stream = Some(map.next_value()?),
"timeout" => cfg.timeout = Some(map.next_value()?),
"system" => cfg.system = Some(map.next_value()?),
"ingress" => cfg.ingress = Some(map.next_value()?),
_ => {
cfg.extra.insert(key, map.next_value()?);
}
}
}
Ok(cfg)
}
}