use crate::adapter::{AdapterDispatcher, AdapterKind};
use crate::chat::ChatOptions;
use crate::client::{ModelSpec, ServiceTarget};
use crate::embed::EmbedOptions;
use crate::resolver::{AuthData, AuthResolver, Endpoint, ModelMapper, ProviderConfig, ServiceTargetResolver};
use crate::{Error, ModelIden, Result, WebConfig};
use std::collections::HashMap;
#[derive(Debug, Default, Clone)]
pub struct ClientConfig {
pub(super) auth_resolver: Option<AuthResolver>,
pub(super) service_target_resolver: Option<ServiceTargetResolver>,
pub(super) model_mapper: Option<ModelMapper>,
pub(super) web_config: Option<WebConfig>,
pub(super) chat_options: Option<ChatOptions>,
pub(super) embed_options: Option<EmbedOptions>,
pub(super) adapter_kind: Option<AdapterKind>,
pub(super) provider_configs: HashMap<AdapterKind, ProviderConfig>,
}
impl ClientConfig {
pub fn with_auth_resolver(mut self, auth_resolver: AuthResolver) -> Self {
self.auth_resolver = Some(auth_resolver);
self
}
pub fn append_provider_config(
mut self,
adapter_kind: AdapterKind,
provider_config: impl Into<ProviderConfig>,
) -> Self {
let ProviderConfig { endpoint, auth } = provider_config.into();
let entry = self.provider_configs.entry(adapter_kind).or_default();
if endpoint.is_some() {
entry.endpoint = endpoint;
}
if auth.is_some() {
entry.auth = auth;
}
self
}
pub fn with_model_mapper(mut self, model_mapper: ModelMapper) -> Self {
self.model_mapper = Some(model_mapper);
self
}
pub fn with_service_target_resolver(mut self, service_target_resolver: ServiceTargetResolver) -> Self {
self.service_target_resolver = Some(service_target_resolver);
self
}
pub fn with_chat_options(mut self, options: ChatOptions) -> Self {
self.chat_options = Some(options);
self
}
pub fn with_embed_options(mut self, options: EmbedOptions) -> Self {
self.embed_options = Some(options);
self
}
pub fn with_web_config(mut self, web_config: WebConfig) -> Self {
self.web_config = Some(web_config);
self
}
pub fn with_adapter_kind(mut self, adapter_kind: AdapterKind) -> Self {
self.adapter_kind = Some(adapter_kind);
self
}
pub fn web_config(&self) -> Option<&WebConfig> {
self.web_config.as_ref()
}
}
impl ClientConfig {
pub fn auth_resolver(&self) -> Option<&AuthResolver> {
self.auth_resolver.as_ref()
}
pub fn provider_config(&self, adapter_kind: AdapterKind) -> Option<&ProviderConfig> {
self.provider_configs.get(&adapter_kind)
}
pub fn service_target_resolver(&self) -> Option<&ServiceTargetResolver> {
self.service_target_resolver.as_ref()
}
pub fn model_mapper(&self) -> Option<&ModelMapper> {
self.model_mapper.as_ref()
}
pub fn chat_options(&self) -> Option<&ChatOptions> {
self.chat_options.as_ref()
}
pub fn embed_options(&self) -> Option<&EmbedOptions> {
self.embed_options.as_ref()
}
pub fn adapter_kind(&self) -> Option<AdapterKind> {
self.adapter_kind
}
}
impl ClientConfig {
pub(crate) async fn resolve_adapter_config(&self, adapter_kind: AdapterKind) -> Result<(AuthData, Endpoint)> {
let model = ModelIden::new(adapter_kind, "");
let auth = self.run_auth_resolver(model.clone()).await?;
let endpoint = self.default_endpoint(adapter_kind);
let service_target = ServiceTarget { model, auth, endpoint };
let service_target = self.run_service_target_resolver(service_target).await?;
Ok((service_target.auth, service_target.endpoint))
}
pub async fn resolve_service_target(&self, model: ModelIden) -> Result<ServiceTarget> {
let model = self.run_model_mapper(model.clone())?;
let auth = self.run_auth_resolver(model.clone()).await?;
let endpoint = self.default_endpoint(model.adapter_kind);
let service_target = ServiceTarget {
model: model.clone(),
auth,
endpoint,
};
let service_target = self.run_service_target_resolver(service_target).await?;
Ok(service_target)
}
fn run_model_mapper(&self, model: ModelIden) -> Result<ModelIden> {
match self.model_mapper() {
Some(model_mapper) => model_mapper.map_model(model.clone()),
None => Ok(model.clone()),
}
.map_err(|resolver_error| Error::Resolver {
model_iden: model.clone(),
resolver_error,
})
}
async fn run_auth_resolver(&self, model: ModelIden) -> Result<AuthData> {
match self.auth_resolver() {
Some(auth_resolver) => {
let auth_data = auth_resolver
.resolve(model.clone())
.await
.map_err(|err| Error::Resolver {
model_iden: model.clone(),
resolver_error: err,
})?
.unwrap_or_else(|| self.default_auth(model.adapter_kind));
Ok(auth_data)
}
None => Ok(self.default_auth(model.adapter_kind)),
}
}
fn default_auth(&self, adapter_kind: AdapterKind) -> AuthData {
self.provider_configs
.get(&adapter_kind)
.and_then(|config| config.auth.clone())
.unwrap_or_else(|| AdapterDispatcher::default_auth(adapter_kind))
}
fn default_endpoint(&self, adapter_kind: AdapterKind) -> Endpoint {
self.provider_configs
.get(&adapter_kind)
.and_then(|config| config.endpoint.clone())
.unwrap_or_else(|| AdapterDispatcher::default_endpoint(adapter_kind))
}
async fn run_service_target_resolver(&self, service_target: ServiceTarget) -> Result<ServiceTarget> {
let model = service_target.model.clone();
match self.service_target_resolver() {
Some(service_target_resolver) => {
service_target_resolver
.resolve(service_target)
.await
.map_err(|resolver_error| Error::Resolver {
model_iden: model,
resolver_error,
})
}
None => Ok(service_target),
}
}
pub async fn resolve_model_spec(&self, spec: ModelSpec) -> Result<ServiceTarget> {
match spec {
ModelSpec::Name(name) => {
let resolved = AdapterKind::from_model(&name)?;
let adapter_kind = match self.adapter_kind {
Some(bound) => {
if resolved != bound && AdapterKind::from_model_namespace(&name).is_some() {
return Err(Error::AdapterKindMismatch {
bound,
requested: resolved,
model: name.to_string(),
});
}
bound
}
None => resolved,
};
let model = ModelIden::new(adapter_kind, name);
self.resolve_service_target(model).await
}
ModelSpec::Iden(model) => {
if let Some(bound) = self.adapter_kind
&& model.adapter_kind != bound
{
return Err(Error::AdapterKindMismatch {
bound,
requested: model.adapter_kind,
model: model.model_name.to_string(),
});
}
self.resolve_service_target(model).await
}
ModelSpec::Target(target) => self.run_service_target_resolver(target).await,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::resolver::{AuthData, AuthResolver, Endpoint, ServiceTargetResolver};
type TestResult = core::result::Result<(), Box<dyn std::error::Error>>;
fn bound_config(adapter_kind: AdapterKind, observed_endpoint: &'static str) -> ClientConfig {
ClientConfig::default()
.with_adapter_kind(adapter_kind)
.with_auth_resolver(AuthResolver::from_resolver_fn(
move |model_iden: ModelIden| -> std::result::Result<Option<AuthData>, crate::resolver::Error> {
if model_iden.adapter_kind == adapter_kind {
Ok(Some(AuthData::from_single("test-key")))
} else {
Ok(None)
}
},
))
.with_service_target_resolver(ServiceTargetResolver::from_resolver_fn(
move |mut service_target: crate::ServiceTarget| -> std::result::Result<crate::ServiceTarget, crate::resolver::Error> {
if service_target.model.adapter_kind == adapter_kind {
service_target.endpoint = Endpoint::from_static(observed_endpoint);
}
Ok(service_target)
},
))
}
#[tokio::test]
async fn adapter_config_applies_service_target_resolver() {
let config = bound_config(AdapterKind::OpenAI, "https://custom.example/v1");
let (auth, endpoint) = config
.resolve_adapter_config(AdapterKind::OpenAI)
.await
.expect("adapter config should resolve");
assert_eq!(endpoint.base_url(), "https://custom.example/v1");
assert!(matches!(auth, AuthData::Key(_)));
}
#[tokio::test]
async fn adapter_config_falls_back_to_the_default_endpoint() {
let config = ClientConfig::default();
let (_auth, endpoint) = config
.resolve_adapter_config(AdapterKind::OpenAI)
.await
.expect("adapter config should resolve");
assert_eq!(
endpoint.base_url(),
AdapterDispatcher::default_endpoint(AdapterKind::OpenAI).base_url()
);
}
#[tokio::test]
async fn provider_config_supplies_endpoint_and_auth() {
let config = ClientConfig::default().append_provider_config(
AdapterKind::OpenAI,
(
Endpoint::from_static("https://gateway.internal/v1/"),
AuthData::from_single("gateway-key"),
),
);
let target = config
.resolve_model_spec(ModelSpec::Iden(ModelIden::new(AdapterKind::OpenAI, "gpt-4o")))
.await
.expect("should resolve");
assert_eq!(target.endpoint.base_url(), "https://gateway.internal/v1/");
assert!(matches!(target.auth, AuthData::Key(ref key) if key == "gateway-key"));
}
#[tokio::test]
async fn provider_config_is_scoped_to_its_adapter() {
let config = ClientConfig::default().append_provider_config(
AdapterKind::OpenAI,
Endpoint::from_static("https://gateway.internal/v1/"),
);
let target = config
.resolve_model_spec(ModelSpec::Iden(ModelIden::new(
AdapterKind::Anthropic,
"claude-sonnet-4-6",
)))
.await
.expect("should resolve");
assert_eq!(
target.endpoint.base_url(),
AdapterDispatcher::default_endpoint(AdapterKind::Anthropic).base_url()
);
}
#[tokio::test]
async fn resolvers_still_win_over_provider_config() {
let config = bound_config(AdapterKind::OpenAI, "https://from-resolver/v1")
.append_provider_config(AdapterKind::OpenAI, Endpoint::from_static("https://from-config/v1/"));
let target = config
.resolve_model_spec(ModelSpec::Iden(ModelIden::new(AdapterKind::OpenAI, "gpt-4o")))
.await
.expect("should resolve");
assert_eq!(target.endpoint.base_url(), "https://from-resolver/v1");
}
#[tokio::test]
async fn provider_config_applies_to_adapter_config_resolution() {
let config = ClientConfig::default().append_provider_config(
AdapterKind::OpenAI,
Endpoint::from_static("https://gateway.internal/v1/"),
);
let (_auth, endpoint) = config
.resolve_adapter_config(AdapterKind::OpenAI)
.await
.expect("should resolve");
assert_eq!(endpoint.base_url(), "https://gateway.internal/v1/");
}
#[tokio::test]
async fn bound_client_routes_bare_name_through_bound_adapter() {
let config = bound_config(AdapterKind::OpenAI, "https://custom.example/v1");
let target = config
.resolve_model_spec(ModelSpec::Name("mini-max-m2.7".into()))
.await
.expect("bound name should resolve");
assert_eq!(target.model.adapter_kind, AdapterKind::OpenAI);
assert_eq!(target.endpoint.base_url(), "https://custom.example/v1");
}
#[tokio::test]
async fn bound_client_rejects_mismatched_namespace() {
let config = bound_config(AdapterKind::OpenAI, "https://custom.example/v1");
let err = config
.resolve_model_spec(ModelSpec::Name("anthropic::claude-3-5-sonnet".into()))
.await
.expect_err("mismatched namespace should error");
match err {
Error::AdapterKindMismatch { bound, requested, .. } => {
assert_eq!(bound, AdapterKind::OpenAI);
assert_eq!(requested, AdapterKind::Anthropic);
}
other => panic!("expected AdapterKindMismatch, got {other:?}"),
}
}
#[tokio::test]
async fn bound_client_accepts_matching_namespace() {
let config = bound_config(AdapterKind::OpenAI, "https://custom.example/v1");
let target = config
.resolve_model_spec(ModelSpec::Name("openai::gpt-4".into()))
.await
.expect("matching namespace should resolve");
assert_eq!(target.model.adapter_kind, AdapterKind::OpenAI);
}
#[tokio::test]
async fn bound_client_rejects_mismatched_iden() {
let config = bound_config(AdapterKind::OpenAI, "https://custom.example/v1");
let iden = ModelIden::new(AdapterKind::Gemini, "gemini-1.5-pro");
let err = config
.resolve_model_spec(ModelSpec::Iden(iden))
.await
.expect_err("mismatched iden should error");
assert!(matches!(
err,
Error::AdapterKindMismatch {
bound: AdapterKind::OpenAI,
requested: AdapterKind::Gemini,
..
}
));
}
#[tokio::test]
async fn unbound_client_preserves_inference() {
let config = bound_config(AdapterKind::OpenAI, "https://custom.example/v1");
let config = ClientConfig {
adapter_kind: None,
..config
};
let target = config
.resolve_model_spec(ModelSpec::Name("gpt-4".into()))
.await
.expect("unbound name should resolve via inference");
assert_eq!(target.model.adapter_kind, AdapterKind::OpenAI);
}
#[test]
fn bound_client_exposes_adapter_kind_via_getter() -> TestResult {
let bound = crate::Client::builder().with_adapter_kind(AdapterKind::OpenAI).build()?;
assert_eq!(bound.adapter_kind(), Some(AdapterKind::OpenAI));
let unbound = crate::Client::new()?;
assert_eq!(unbound.adapter_kind(), None);
Ok(())
}
}