mod resolver;
pub use resolver::{ResolvedProvider, provider_for};
use std::collections::HashMap;
use crate::inference::adapter::InferenceAdapter;
use crate::inference::credentials::KeyStore;
use crate::inference::error::InferenceError;
use crate::inference::registry::ProviderId;
pub trait AdapterFactory: Send + Sync {
fn build(
&self,
resolved: &ResolvedProvider,
) -> Result<Box<dyn InferenceAdapter>, InferenceError>;
}
impl<F> AdapterFactory for F
where
F: Fn(&ResolvedProvider) -> Result<Box<dyn InferenceAdapter>, InferenceError> + Send + Sync,
{
fn build(
&self,
resolved: &ResolvedProvider,
) -> Result<Box<dyn InferenceAdapter>, InferenceError> {
self(resolved)
}
}
#[derive(Default)]
pub struct Configurator {
factories: HashMap<ProviderId, Box<dyn AdapterFactory>>,
}
impl Configurator {
pub fn new() -> Self {
Self::default()
}
pub fn register(&mut self, id: ProviderId, factory: Box<dyn AdapterFactory>) {
self.factories.insert(id, factory);
}
pub fn build(
&self,
slug: &str,
store: &dyn KeyStore,
) -> Result<Box<dyn InferenceAdapter>, InferenceError> {
let resolved = provider_for(slug, store)?;
let factory = self.factories.get(&resolved.provider()).ok_or(
InferenceError::NoAdapterRegistered {
provider: resolved.provider(),
},
)?;
factory.build(&resolved)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::credentials::MemoryKeyStore;
use crate::inference::registry::capabilities;
use crate::inference::test_support::ScriptedAdapter;
use serial_test::serial;
fn clear_env() {
for var in ["OPENROUTER_API_KEY", "ANTHROPIC_API_KEY"] {
unsafe { std::env::remove_var(var) };
}
}
#[test]
#[serial(dotenv_credential_env)]
fn build_uses_registered_factory() {
clear_env();
let store = MemoryKeyStore::new();
store.set("openrouter", "sk-or-abc").unwrap();
let mut cfg = Configurator::new();
cfg.register(
ProviderId::OpenRouter,
Box::new(|resolved: &ResolvedProvider| {
let caps = capabilities(resolved.provider());
Ok(Box::new(ScriptedAdapter::echo("openrouter", caps))
as Box<dyn InferenceAdapter>)
}),
);
let adapter = cfg.build("some/model", &store).expect("built");
assert_eq!(adapter.name(), "openrouter");
}
#[test]
#[serial(dotenv_credential_env)]
fn build_unregistered_provider_errors() {
clear_env();
let store = MemoryKeyStore::new();
store.set("openrouter", "sk-or-abc").unwrap();
let cfg = Configurator::new(); let Err(err) = cfg.build("some/model", &store) else {
panic!("expected NoAdapterRegistered");
};
assert!(err.is_alarm());
assert!(matches!(
err,
InferenceError::NoAdapterRegistered {
provider: ProviderId::OpenRouter
}
));
}
}