use std::collections::HashMap;
use agent_client_protocol as acp;
use super::ZephAcpAgentState;
pub(crate) struct ProviderSetOverride {
pub api_type: agent_client_protocol_schema::v1::LlmProtocol,
pub base_url: String,
#[allow(dead_code)]
pub headers: HashMap<String, String>,
}
impl ZephAcpAgentState {
#[cfg_attr(docsrs, doc(cfg(feature = "unstable-llm-providers")))]
#[must_use]
pub fn with_provider_names(
mut self,
names: Vec<(String, agent_client_protocol_schema::v1::LlmProtocol)>,
) -> Self {
self.provider_names = names;
self
}
pub(crate) fn ext_method_providers(
&self,
args: &acp::schema::v1::ExtRequest,
) -> acp::Result<Option<acp::schema::v1::ExtResponse>> {
use agent_client_protocol_schema::v1 as schema;
let method = args.method.as_ref();
match method {
"providers/list" => {
let req: schema::ListProvidersRequest = serde_json::from_str(args.params.get())
.map_err(|e| acp::Error::invalid_request().data(e.to_string()))?;
let resp = self.do_list_providers(req)?;
let json = serde_json::to_string(&resp)
.map_err(|e| acp::Error::internal_error().data(e.to_string()))?;
let raw = serde_json::value::RawValue::from_string(json)
.map_err(|e| acp::Error::internal_error().data(e.to_string()))?;
Ok(Some(acp::schema::v1::ExtResponse::new(raw.into())))
}
"providers/set" => {
let req: schema::SetProviderRequest = serde_json::from_str(args.params.get())
.map_err(|e| acp::Error::invalid_request().data(e.to_string()))?;
let resp = self.do_set_providers(req)?;
let json = serde_json::to_string(&resp)
.map_err(|e| acp::Error::internal_error().data(e.to_string()))?;
let raw = serde_json::value::RawValue::from_string(json)
.map_err(|e| acp::Error::internal_error().data(e.to_string()))?;
Ok(Some(acp::schema::v1::ExtResponse::new(raw.into())))
}
"providers/disable" => {
let req: schema::DisableProviderRequest =
serde_json::from_str(args.params.get())
.map_err(|e| acp::Error::invalid_request().data(e.to_string()))?;
let resp = self.do_disable_providers(&req)?;
let json = serde_json::to_string(&resp)
.map_err(|e| acp::Error::internal_error().data(e.to_string()))?;
let raw = serde_json::value::RawValue::from_string(json)
.map_err(|e| acp::Error::internal_error().data(e.to_string()))?;
Ok(Some(acp::schema::v1::ExtResponse::new(raw.into())))
}
_ => Ok(None),
}
}
#[tracing::instrument(skip_all, name = "acp.handler.list_providers")]
pub(crate) fn do_list_providers(
&self,
_req: agent_client_protocol_schema::v1::ListProvidersRequest,
) -> acp::Result<agent_client_protocol_schema::v1::ListProvidersResponse> {
let disabled = self.global_disabled_providers.lock();
let overrides = self.global_provider_overrides.lock();
let providers: Vec<agent_client_protocol_schema::v1::ProviderInfo> = self
.provider_names
.iter()
.map(|(name, protocol)| {
let is_disabled = disabled.contains(name.as_str());
let current = if is_disabled {
None
} else if let Some(ov) = overrides.get(name.as_str()) {
Some(
agent_client_protocol_schema::v1::ProviderCurrentConfig::new(
ov.api_type.clone(),
ov.base_url.clone(),
),
)
} else {
Some(
agent_client_protocol_schema::v1::ProviderCurrentConfig::new(
protocol.clone(),
String::new(),
),
)
};
agent_client_protocol_schema::v1::ProviderInfo::new(
name.clone(),
vec![protocol.clone()],
false,
current,
)
})
.collect();
Ok(agent_client_protocol_schema::v1::ListProvidersResponse::new(providers))
}
#[tracing::instrument(skip_all, name = "acp.handler.set_providers")]
pub(crate) fn do_set_providers(
&self,
req: agent_client_protocol_schema::v1::SetProviderRequest,
) -> acp::Result<agent_client_protocol_schema::v1::SetProviderResponse> {
if !self
.provider_names
.iter()
.any(|(name, _)| name.as_str() == req.provider_id.0.as_ref())
{
return Err(acp::Error::invalid_params()
.data(format!("unknown provider id: {}", req.provider_id)));
}
self.global_provider_overrides.lock().insert(
req.provider_id.to_string(),
ProviderSetOverride {
api_type: req.api_type,
base_url: req.base_url,
headers: req.headers,
},
);
tracing::debug!(provider_id = %req.provider_id, "provider override set");
Ok(agent_client_protocol_schema::v1::SetProviderResponse::new())
}
#[tracing::instrument(skip_all, name = "acp.handler.disable_providers")]
pub(crate) fn do_disable_providers(
&self,
req: &agent_client_protocol_schema::v1::DisableProviderRequest,
) -> acp::Result<agent_client_protocol_schema::v1::DisableProviderResponse> {
let id = req.provider_id.to_string();
tracing::debug!(provider_id = %id, "provider disabled");
self.global_disabled_providers.lock().insert(id);
Ok(agent_client_protocol_schema::v1::DisableProviderResponse::new())
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use agent_client_protocol_schema::v1 as schema;
use crate::agent::{AgentSpawner, ZephAcpAgentState};
fn make_state() -> ZephAcpAgentState {
let spawner: AgentSpawner = Arc::new(|_ch, _ctx, _sc| Box::pin(async {}));
ZephAcpAgentState::new(spawner, 4, 1800, None).with_provider_names(vec![
("openai".to_owned(), schema::LlmProtocol::OpenAi),
("claude".to_owned(), schema::LlmProtocol::Anthropic),
])
}
#[test]
fn list_providers_returns_all_registered() {
let state = make_state();
let resp = state
.do_list_providers(schema::ListProvidersRequest::new())
.unwrap();
assert_eq!(resp.providers.len(), 2);
let ids: Vec<&str> = resp
.providers
.iter()
.map(|p| p.provider_id.0.as_ref())
.collect();
assert!(ids.contains(&"openai"));
assert!(ids.contains(&"claude"));
}
#[test]
fn list_providers_empty_when_none_registered() {
let spawner: AgentSpawner = Arc::new(|_ch, _ctx, _sc| Box::pin(async {}));
let state = ZephAcpAgentState::new(spawner, 4, 1800, None).with_provider_names(vec![]);
let resp = state
.do_list_providers(schema::ListProvidersRequest::new())
.unwrap();
assert!(resp.providers.is_empty());
}
#[test]
fn protocol_type_reflected_in_default_current_config() {
let state = make_state();
let resp = state
.do_list_providers(schema::ListProvidersRequest::new())
.unwrap();
let openai = resp
.providers
.iter()
.find(|p| p.provider_id.0.as_ref() == "openai")
.unwrap();
let current = openai
.current
.as_ref()
.expect("openai must have current config");
assert_eq!(
current.api_type,
schema::LlmProtocol::OpenAi,
"openai provider must report OpenAi protocol"
);
let claude = resp
.providers
.iter()
.find(|p| p.provider_id.0.as_ref() == "claude")
.unwrap();
let current = claude
.current
.as_ref()
.expect("claude must have current config");
assert_eq!(
current.api_type,
schema::LlmProtocol::Anthropic,
"claude provider must report Anthropic protocol"
);
}
#[test]
fn disable_provider_hides_current_config_in_list() {
let state = make_state();
state
.do_disable_providers(&schema::DisableProviderRequest::new("openai"))
.unwrap();
let resp = state
.do_list_providers(schema::ListProvidersRequest::new())
.unwrap();
let openai = resp
.providers
.iter()
.find(|p| p.provider_id.0.as_ref() == "openai")
.unwrap();
assert!(
openai.current.is_none(),
"disabled provider must have no current config"
);
let claude = resp
.providers
.iter()
.find(|p| p.provider_id.0.as_ref() == "claude")
.unwrap();
assert!(
claude.current.is_some(),
"non-disabled provider must still have current config"
);
}
#[test]
fn disable_unknown_provider_succeeds() {
let state = make_state();
state
.do_disable_providers(&schema::DisableProviderRequest::new("nonexistent"))
.unwrap();
}
#[test]
fn set_provider_unknown_id_returns_error() {
let state = make_state();
let err = state
.do_set_providers(
schema::SetProviderRequest::new(
"unknown_provider",
schema::LlmProtocol::OpenAi,
"https://evil.example.com",
)
.headers(std::collections::HashMap::new()),
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("unknown provider id"),
"expected 'unknown provider id' in error, got: {msg}"
);
}
#[test]
fn set_provider_override_appears_in_list() {
let state = make_state();
state
.do_set_providers(
schema::SetProviderRequest::new(
"openai",
schema::LlmProtocol::OpenAi,
"https://custom.example.com",
)
.headers(std::collections::HashMap::new()),
)
.unwrap();
let resp = state
.do_list_providers(schema::ListProvidersRequest::new())
.unwrap();
let openai = resp
.providers
.iter()
.find(|p| p.provider_id.0.as_ref() == "openai")
.unwrap();
let current = openai.current.as_ref().expect("override must be present");
assert_eq!(current.base_url, "https://custom.example.com");
}
#[test]
fn disable_after_set_clears_current_config() {
let state = make_state();
state
.do_set_providers(
schema::SetProviderRequest::new(
"openai",
schema::LlmProtocol::OpenAi,
"https://custom.example.com",
)
.headers(std::collections::HashMap::new()),
)
.unwrap();
state
.do_disable_providers(&schema::DisableProviderRequest::new("openai"))
.unwrap();
let resp = state
.do_list_providers(schema::ListProvidersRequest::new())
.unwrap();
let openai = resp
.providers
.iter()
.find(|p| p.provider_id.0.as_ref() == "openai")
.unwrap();
assert!(
openai.current.is_none(),
"provider disabled after set must have no current config"
);
}
}