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 model: Option<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 mut provider = AnthropicProvider::from_env()
.map_err(|e| ProviderBuildError(format!("anthropic wire: {e}")))?;
if let Some(model) = &init.model {
provider.config_mut().model.clone_from(model);
}
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());
if let Some(model) = &init.model {
provider.config_mut().model.clone_from(model);
}
let model = provider.config().model.clone();
Ok(BuiltProvider {
provider: Arc::new(provider),
model,
})
})
.register("mock", |init| {
let script = match std::env::var("LOCODE_MOCK_SCRIPT") {
Ok(json) => mock_script(&json)
.map_err(|e| ProviderBuildError(format!("LOCODE_MOCK_SCRIPT: {e}")))?,
Err(_) => vec![Completion {
content: vec![ContentBlock::Text {
text: "Mock run complete.".to_string(),
}],
usage: Usage::default(),
stop: StopReason::EndTurn,
}],
};
Ok(BuiltProvider {
provider: Arc::new(MockProvider::new(script)),
model: init.model.clone().unwrap_or_else(|| "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()
}
}
fn mock_script(json: &str) -> Result<Vec<Completion>, String> {
let value: serde_json::Value =
serde_json::from_str(json).map_err(|e| format!("invalid JSON: {e}"))?;
let turns = value
.as_array()
.ok_or_else(|| "expected a JSON array of turns".to_string())?;
if turns.is_empty() {
return Err("expected at least one turn".to_string());
}
turns
.iter()
.enumerate()
.map(|(i, turn)| {
if let Some(text) = turn.get("text").and_then(serde_json::Value::as_str) {
Ok(Completion {
content: vec![ContentBlock::Text {
text: text.to_string(),
}],
usage: Usage::default(),
stop: StopReason::EndTurn,
})
} else if let Some(tool) = turn.get("tool").and_then(serde_json::Value::as_str) {
Ok(Completion {
content: vec![ContentBlock::ToolUse {
id: format!("call_{i}"),
name: tool.to_string(),
input: turn
.get("input")
.cloned()
.unwrap_or_else(|| serde_json::Value::Object(serde_json::Map::new())),
}],
usage: Usage::default(),
stop: StopReason::ToolUse,
})
} else {
Err(format!(
"turn {i}: expected {{\"text\": …}} or {{\"tool\": …, \"input\": …}}"
))
}
})
.collect()
}