use crate::capabilities::requirement::{
ModelRequirement, ProviderRequirement, ShellCommandRequirement,
};
use crate::providers::ApiType;
use crate::registry::{ConfigConstructable, Secret};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet};
pub use crate::launchers::{EnvBinding, LaunchContext};
macro_rules! define_bindings {
($(
$variant:ident {
request: $request_ty:ty,
result: $result_ty:ty,
display: $display:literal,
}
),+ $(,)?) => {
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum BindingType {
$($variant),+
}
impl std::fmt::Display for BindingType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
$(BindingType::$variant => write!(f, $display),)+
}
}
}
#[derive(Debug, Clone)]
pub enum BindingRequest {
$($variant($request_ty)),+
}
impl BindingRequest {
pub fn binding_type(&self) -> BindingType {
match self {
$(BindingRequest::$variant(_) => BindingType::$variant,)+
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum Binding {
$($variant($result_ty)),+
}
impl Binding {
pub fn binding_type(&self) -> BindingType {
match self {
$(Binding::$variant(_) => BindingType::$variant,)+
}
}
}
};
}
#[derive(Debug, Clone)]
pub struct AgentModelBindingRequest {
pub api_type: ApiType,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AgentModelBinding {
pub api_type: ApiType,
pub provider_name: String,
pub base_url: String,
pub model_name: String,
pub endpoint_path: String,
pub api_key: Option<Secret>,
pub verify_ssl: bool,
pub context_length: Option<u64>,
pub custom_headers: Option<HashMap<String, Secret>>,
}
#[derive(Debug, Clone)]
pub struct SubAgentBindingRequest {
pub api_type: ApiType,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)]
pub enum ToolName {
FileRead,
FileWrite,
FileEdit,
Search,
FileSearch,
Shell,
WebFetch,
WebSearch,
Mcp {
server: String,
tool: Option<String>,
},
Other(String),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)]
pub enum KnownSubAgent {
Explore,
Plan,
Code,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SubAgentBinding {
pub description: String,
pub prompt: String,
pub tools: Vec<ToolName>,
pub model: AgentModelBinding,
pub known_type: Option<KnownSubAgent>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum McpTransportKind {
Stdio,
Http,
Sse,
}
#[derive(Debug, Clone)]
pub struct McpBindingRequest {
pub supported_transports: HashSet<McpTransportKind>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum McpBinding {
Stdio {
command: String,
args: Vec<String>,
env: HashMap<String, String>,
timeout: Option<u64>,
},
Http {
url: String,
headers: HashMap<String, String>,
timeout: Option<u64>,
},
Sse {
url: String,
headers: HashMap<String, String>,
timeout: Option<u64>,
},
}
define_bindings! {
AgentModel {
request: AgentModelBindingRequest,
result: AgentModelBinding,
display: "Agent Model",
},
Mcp {
request: McpBindingRequest,
result: McpBinding,
display: "MCP Server",
},
SubAgent {
request: SubAgentBindingRequest,
result: SubAgentBinding,
display: "Sub-Agent",
},
}
impl McpBinding {
pub fn to_canonical_json(&self) -> serde_json::Value {
let mut map = serde_json::Map::new();
match self {
McpBinding::Stdio {
command,
args,
env,
timeout,
} => {
map.insert("type".into(), serde_json::json!("stdio"));
map.insert("command".into(), serde_json::json!(command));
map.insert("args".into(), serde_json::json!(args));
map.insert("env".into(), serde_json::json!(env));
if let Some(t) = timeout {
map.insert("timeout".into(), serde_json::json!(t));
}
}
McpBinding::Http {
url,
headers,
timeout,
} => {
map.insert("type".into(), serde_json::json!("http"));
map.insert("url".into(), serde_json::json!(url));
map.insert("headers".into(), serde_json::json!(headers));
if let Some(t) = timeout {
map.insert("timeout".into(), serde_json::json!(t));
}
}
McpBinding::Sse {
url,
headers,
timeout,
} => {
map.insert("type".into(), serde_json::json!("sse"));
map.insert("url".into(), serde_json::json!(url));
map.insert("headers".into(), serde_json::json!(headers));
if let Some(t) = timeout {
map.insert("timeout".into(), serde_json::json!(t));
}
}
}
serde_json::Value::Object(map)
}
}
#[async_trait]
pub trait Capability: crate::registry::Named + Send + Sync {
fn name(&self) -> &str;
fn description(&self) -> &str;
fn binding_types(&self) -> HashSet<BindingType>;
async fn bind(&self, request: BindingRequest) -> anyhow::Result<Binding>;
async fn on_setup(&self) -> anyhow::Result<()> {
Ok(())
}
async fn on_pre_launch(&self, _context: &LaunchContext) -> anyhow::Result<()> {
Ok(())
}
async fn on_post_launch(&self, _context: &LaunchContext) -> anyhow::Result<()> {
Ok(())
}
async fn on_shutdown(&self, _context: &LaunchContext) -> anyhow::Result<()> {
Ok(())
}
fn runtime_bindings(&self) -> Vec<EnvBinding> {
vec![]
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CapabilityMetadata {
pub name: String,
pub description: String,
pub dependencies: Vec<Dependency>,
pub tags: Vec<String>,
pub supported_binding_types: HashSet<BindingType>,
}
impl std::fmt::Display for CapabilityMetadata {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.description)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum Dependency {
Model {
config_key: String,
requirement: ModelRequirement,
resolved_id: Option<String>,
required: bool,
},
Provider {
config_key: String,
requirement: ProviderRequirement,
resolved_id: Option<String>,
required: bool,
},
ExternalTool {
requirement: ShellCommandRequirement,
required: bool,
},
}
impl std::fmt::Display for Dependency {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Dependency::Model {
resolved_id,
required,
..
} => {
write!(
f,
"Model: {}{}",
resolved_id.as_deref().unwrap_or("<unresolved>"),
if *required { " (required)" } else { "" }
)
}
Dependency::Provider {
resolved_id,
required,
..
} => {
write!(
f,
"Provider: {}{}",
resolved_id.as_deref().unwrap_or("<unresolved>"),
if *required { " (required)" } else { "" }
)
}
Dependency::ExternalTool {
requirement,
required,
} => {
write!(
f,
"ExternalTool: {}{}",
requirement.command,
if *required { " (required)" } else { "" }
)
}
}
}
}
use crate::define_factory;
define_factory!(Capability, CapabilityMetadata, CapabilityFactory);
#[cfg(test)]
mod mcp_binding_tests {
use super::*;
fn stdio_binding() -> McpBinding {
McpBinding::Stdio {
command: "/usr/local/bin/granite-cli".to_string(),
args: vec!["__mcp-serve".to_string(), "vision".to_string()],
env: HashMap::from([("FOO".to_string(), "bar".to_string())]),
timeout: None,
}
}
fn http_binding() -> McpBinding {
McpBinding::Http {
url: "http://127.0.0.1:54321/mcp".to_string(),
headers: HashMap::from([("X-Test".to_string(), "1".to_string())]),
timeout: None,
}
}
#[test]
fn canonical_json_stdio_matches_mcp_add_json_shape() {
let json = stdio_binding().to_canonical_json();
assert_eq!(json["type"], "stdio");
assert_eq!(json["command"], "/usr/local/bin/granite-cli");
assert_eq!(json["args"][0], "__mcp-serve");
assert_eq!(json["args"][1], "vision");
assert_eq!(json["env"]["FOO"], "bar");
}
#[test]
fn canonical_json_http_matches_mcp_add_json_shape() {
let json = http_binding().to_canonical_json();
assert_eq!(json["type"], "http");
assert_eq!(json["url"], "http://127.0.0.1:54321/mcp");
assert_eq!(json["headers"]["X-Test"], "1");
}
#[test]
fn canonical_json_sse_uses_sse_type() {
let json = McpBinding::Sse {
url: "http://127.0.0.1:1/sse".to_string(),
headers: HashMap::new(),
timeout: None,
}
.to_canonical_json();
assert_eq!(json["type"], "sse");
}
#[test]
fn canonical_json_includes_timeout_when_set() {
let json = McpBinding::Http {
url: "http://127.0.0.1:1/mcp".to_string(),
headers: HashMap::new(),
timeout: Some(300_000),
}
.to_canonical_json();
assert_eq!(json["timeout"], 300_000u64);
}
#[test]
fn canonical_json_omits_timeout_when_none() {
let json = McpBinding::Stdio {
command: "my-command".to_string(),
args: vec![],
env: HashMap::new(),
timeout: None,
}
.to_canonical_json();
assert!(!json.as_object().unwrap().contains_key("timeout"));
}
}