use std::fmt;
use std::sync::Arc;
use locode_protocol::{ContentBlock, Usage};
use locode_provider::{
AnthropicProvider, Completion, MockProvider, OpenAiResponsesProvider, Provider, StopReason,
};
pub struct ProviderInit {
pub session_id: String,
}
pub struct BuiltProvider {
pub provider: Arc<dyn Provider>,
pub model: String,
}
impl fmt::Debug for BuiltProvider {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BuiltProvider")
.field("model", &self.model)
.finish_non_exhaustive() }
}
#[derive(Debug)]
pub struct ProviderBuildError(pub String);
impl fmt::Display for ProviderBuildError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl std::error::Error for ProviderBuildError {}
pub type ProviderFactory =
Box<dyn Fn(&ProviderInit) -> Result<BuiltProvider, ProviderBuildError> + Send + Sync>;
pub struct ProviderRegistry {
entries: Vec<(String, ProviderFactory)>,
}
impl ProviderRegistry {
#[must_use]
pub fn new() -> Self {
ProviderRegistry {
entries: Vec::new(),
}
}
#[must_use]
pub fn builtin() -> Self {
Self::new()
.register("anthropic", |_init| {
let provider = AnthropicProvider::from_env()
.map_err(|e| ProviderBuildError(format!("anthropic wire: {e}")))?;
let model = provider.config().model.clone();
Ok(BuiltProvider {
provider: Arc::new(provider),
model,
})
})
.register("openai-responses", |init| {
let mut provider = OpenAiResponsesProvider::from_env()
.map_err(|e| ProviderBuildError(format!("openai-responses wire: {e}")))?;
provider.config_mut().prompt_cache_key = Some(init.session_id.clone());
let model = provider.config().model.clone();
Ok(BuiltProvider {
provider: Arc::new(provider),
model,
})
})
.register("mock", |_init| {
let provider = MockProvider::new(vec![Completion {
content: vec![ContentBlock::Text {
text: "Mock run complete.".to_string(),
}],
usage: Usage::default(),
stop: StopReason::EndTurn,
}]);
Ok(BuiltProvider {
provider: Arc::new(provider),
model: "mock-1".to_string(),
})
})
}
#[must_use]
pub fn register<F>(mut self, name: impl Into<String>, factory: F) -> Self
where
F: Fn(&ProviderInit) -> Result<BuiltProvider, ProviderBuildError> + Send + Sync + 'static,
{
let name = name.into();
let factory: ProviderFactory = Box::new(factory);
match self.entries.iter_mut().find(|(n, _)| *n == name) {
Some(entry) => entry.1 = factory,
None => self.entries.push((name, factory)),
}
self
}
#[must_use]
pub fn names(&self) -> Vec<&str> {
self.entries.iter().map(|(n, _)| n.as_str()).collect()
}
pub fn build(
&self,
name: &str,
init: &ProviderInit,
) -> Result<BuiltProvider, ProviderBuildError> {
let factory = self
.entries
.iter()
.find(|(n, _)| n == name)
.map(|(_, f)| f)
.ok_or_else(|| {
ProviderBuildError(format!(
"unknown --api-schema `{name}`; available: {}",
self.names().join(", ")
))
})?;
factory(init)
}
}
impl Default for ProviderRegistry {
fn default() -> Self {
Self::builtin()
}
}