use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use crate::bridge::ids::ToolPolicy;
use systemprompt_identifiers::{McpServerId, McpToolName, ValidatedUrl};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(from = "ManagedMcpServerWire")]
pub struct ManagedMcpServer {
pub id: McpServerId,
pub name: McpServerId,
pub url: ValidatedUrl,
#[serde(skip_serializing_if = "Option::is_none")]
pub transport: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub headers: Option<BTreeMap<String, String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub oauth: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_policy: Option<BTreeMap<McpToolName, ToolPolicy>>,
}
impl ManagedMcpServer {
pub const TOOL_POLICY_WILDCARD: &'static str = "*";
#[must_use]
pub fn policy_for_tool(&self, tool: &str) -> Option<ToolPolicy> {
let map = self.tool_policy.as_ref()?;
map.iter()
.find(|(name, _)| name.as_str() == tool)
.or_else(|| {
map.iter()
.find(|(name, _)| name.as_str() == Self::TOOL_POLICY_WILDCARD)
})
.map(|(_, policy)| *policy)
}
#[must_use]
pub fn default_tool_policy(&self) -> Option<ToolPolicy> {
self.policy_for_tool(Self::TOOL_POLICY_WILDCARD)
}
}
#[derive(Deserialize)]
struct ManagedMcpServerWire {
#[serde(default)]
id: Option<McpServerId>,
name: McpServerId,
url: ValidatedUrl,
#[serde(default)]
transport: Option<String>,
#[serde(default)]
headers: Option<BTreeMap<String, String>>,
#[serde(default)]
oauth: Option<bool>,
#[serde(default)]
tool_policy: Option<BTreeMap<McpToolName, ToolPolicy>>,
}
impl From<ManagedMcpServerWire> for ManagedMcpServer {
fn from(wire: ManagedMcpServerWire) -> Self {
Self {
id: wire.id.unwrap_or_else(|| wire.name.clone()),
name: wire.name,
url: wire.url,
transport: wire.transport,
headers: wire.headers,
oauth: wire.oauth,
tool_policy: wire.tool_policy,
}
}
}