Skip to main content

aether_cli/
mcp_command.rs

1use clap::{ArgAction, Args};
2use mcp_utils::ServiceExt;
3use mcp_utils::tool_gateway::{AETHER_MCP_IPC_SOCKET, LIST_SERVERS_TOOL, UnixSocketPath, connect};
4use rmcp::model::{CallToolRequestParams, CallToolResponse, CallToolResult, Tool};
5use serde::Deserialize;
6use serde_json::{Map, Value};
7use std::env::var_os;
8use std::fmt::Display;
9use std::future::Future;
10use std::io::{self, IsTerminal, Read};
11use std::path::PathBuf;
12use std::time::Duration;
13
14const DISCOVERY_TIMEOUT: Duration = Duration::from_secs(5);
15const DEFAULT_CALL_TIMEOUT_SECONDS: u64 = 600;
16
17#[derive(Debug, Args)]
18#[command(disable_help_flag = true)]
19pub struct McpArgs {
20    #[arg(value_name = "SERVER")]
21    server: Option<String>,
22
23    #[arg(value_name = "TOOL", requires = "server")]
24    tool: Option<String>,
25
26    #[arg(long, action = ArgAction::SetTrue, conflicts_with_all = ["json", "timeout_seconds"])]
27    help: bool,
28
29    #[arg(long, value_name = "OBJECT", requires = "tool")]
30    json: Option<String>,
31
32    #[arg(
33        long = "timeout",
34        value_name = "SECONDS",
35        default_value_t = DEFAULT_CALL_TIMEOUT_SECONDS,
36        value_parser = clap::value_parser!(u64).range(1..),
37        requires = "tool"
38    )]
39    timeout_seconds: u64,
40}
41
42#[derive(Debug, thiserror::Error)]
43pub enum McpCommandError {
44    #[error("{0}")]
45    Usage(String),
46    #[error("`aether mcp` requires the inherited {AETHER_MCP_IPC_SOCKET} from an active Aether session")]
47    SessionUnavailable,
48    #[error("invalid {AETHER_MCP_IPC_SOCKET}: {0}")]
49    InvalidSocket(String),
50    #[error("failed to connect to the active Aether session: {0}")]
51    Connect(String),
52    #[error("MCP request timed out after {0} seconds")]
53    Timeout(u64),
54    #[error("MCP request failed: {0}")]
55    Request(String),
56    #[error("deferred tool returned an error: {0}")]
57    Tool(String),
58    #[error("failed to read JSON from stdin: {0}")]
59    Stdin(#[source] io::Error),
60    #[error("failed to print JSON result: {0}")]
61    Output(#[source] serde_json::Error),
62}
63
64#[derive(Debug)]
65enum Request {
66    Help(HelpLevel),
67    Call { server: String, tool: String, json: Map<String, Value>, timeout_seconds: u64 },
68}
69
70#[derive(Debug)]
71enum HelpLevel {
72    Servers,
73    Server(String),
74    Tool { server: String, tool: String },
75}
76
77#[derive(Deserialize)]
78struct ServerSummary {
79    name: String,
80    description: String,
81}
82
83impl McpCommandError {
84    pub fn exit_code(&self) -> u8 {
85        match self {
86            Self::Usage(_) | Self::Stdin(_) => 2,
87            Self::SessionUnavailable
88            | Self::InvalidSocket(_)
89            | Self::Connect(_)
90            | Self::Timeout(_)
91            | Self::Request(_)
92            | Self::Tool(_)
93            | Self::Output(_) => 1,
94        }
95    }
96}
97
98pub async fn run(args: McpArgs) -> Result<(), McpCommandError> {
99    let request = Request::try_from(args)?;
100    let socket = inherited_socket()?;
101    let transport = connect(socket.path()).await.map_err(|error| McpCommandError::Connect(error.to_string()))?;
102    let client = ().serve(transport).await.map_err(|error| McpCommandError::Connect(error.to_string()))?;
103    execute_request(&client, request).await
104}
105
106impl TryFrom<McpArgs> for Request {
107    type Error = McpCommandError;
108
109    fn try_from(args: McpArgs) -> Result<Self, Self::Error> {
110        match (args.server, args.tool, args.help) {
111            (None, None, _) => Ok(Request::Help(HelpLevel::Servers)),
112            (Some(server), None, true) => Ok(Request::Help(HelpLevel::Server(server))),
113            (Some(server), Some(tool), true) => Ok(Request::Help(HelpLevel::Tool { server, tool })),
114            (Some(server), Some(tool), false) => {
115                let json = if let Some(input) = args.json.as_deref() {
116                    parse_json_object(Some(input))?
117                } else {
118                    let stdin = read_stdin()?;
119                    parse_json_object(stdin.as_deref())?
120                };
121                Ok(Request::Call { server, tool, json, timeout_seconds: args.timeout_seconds })
122            }
123            (Some(_), None, false) | (None, Some(_), _) => Err(McpCommandError::Usage(
124                "usage: aether mcp <server> <tool> [--json <object>] [--timeout <seconds>]".into(),
125            )),
126        }
127    }
128}
129
130fn inherited_socket() -> Result<UnixSocketPath, McpCommandError> {
131    let path = var_os(AETHER_MCP_IPC_SOCKET).ok_or(McpCommandError::SessionUnavailable)?;
132    UnixSocketPath::from_path(PathBuf::from(path)).map_err(|error| McpCommandError::InvalidSocket(error.to_string()))
133}
134
135fn read_stdin() -> Result<Option<String>, McpCommandError> {
136    let mut stdin = io::stdin();
137    if stdin.is_terminal() {
138        return Ok(None);
139    }
140    let mut input = String::new();
141    stdin.read_to_string(&mut input).map_err(McpCommandError::Stdin)?;
142    Ok((!input.trim().is_empty()).then_some(input))
143}
144
145fn parse_json_object(input: Option<&str>) -> Result<Map<String, Value>, McpCommandError> {
146    let Some(input) = input else { return Ok(Map::new()) };
147    serde_json::from_str(input).map_err(|error| match error.classify() {
148        serde_json::error::Category::Data => McpCommandError::Usage("tool input must be a JSON object".into()),
149        _ => McpCommandError::Usage(format!("invalid JSON input: {error}")),
150    })
151}
152
153async fn execute_request(
154    client: &rmcp::service::RunningService<rmcp::RoleClient, ()>,
155    request: Request,
156) -> Result<(), McpCommandError> {
157    match request {
158        Request::Help(HelpLevel::Servers) => show_servers_help(client).await,
159        Request::Help(HelpLevel::Server(server)) => show_server_help(client, &server).await,
160        Request::Help(HelpLevel::Tool { server, tool }) => show_tool_help(client, &server, &tool).await,
161        Request::Call { server, tool, json, timeout_seconds } => {
162            call_tool(client, &server, &tool, json, timeout_seconds).await
163        }
164    }
165}
166
167async fn show_servers_help(
168    client: &rmcp::service::RunningService<rmcp::RoleClient, ()>,
169) -> Result<(), McpCommandError> {
170    let result = timed(DISCOVERY_TIMEOUT, client.call_tool_once(CallToolRequestParams::new(LIST_SERVERS_TOOL))).await?;
171    let result = complete(result)?;
172    let servers: Vec<ServerSummary> = serde_json::from_value(
173        result
174            .structured_content
175            .ok_or_else(|| McpCommandError::Request("server discovery returned no JSON".into()))?,
176    )
177    .map_err(|error| McpCommandError::Request(format!("invalid server discovery response: {error}")))?;
178    print_servers_help(&servers);
179    Ok(())
180}
181
182async fn show_server_help(
183    client: &rmcp::service::RunningService<rmcp::RoleClient, ()>,
184    server: &str,
185) -> Result<(), McpCommandError> {
186    let tools = list_tools(client).await?;
187    let tools = tools_for_server(&tools, server);
188    if tools.is_empty() {
189        return Err(McpCommandError::Usage(format!("unknown deferred server `{server}`")));
190    }
191    print_server_help(server, tools);
192    Ok(())
193}
194
195async fn show_tool_help(
196    client: &rmcp::service::RunningService<rmcp::RoleClient, ()>,
197    server: &str,
198    tool: &str,
199) -> Result<(), McpCommandError> {
200    let tools = list_tools(client).await?;
201    let namespaced = format!("{server}__{tool}");
202    let definition = tools
203        .into_iter()
204        .find(|definition| definition.name == namespaced)
205        .ok_or_else(|| McpCommandError::Usage(format!("unknown deferred tool `{server} {tool}`")))?;
206    print_tool_help(server, tool, &definition)
207}
208
209async fn call_tool(
210    client: &rmcp::service::RunningService<rmcp::RoleClient, ()>,
211    server: &str,
212    tool: &str,
213    json: Map<String, Value>,
214    timeout_seconds: u64,
215) -> Result<(), McpCommandError> {
216    let params = CallToolRequestParams::new(format!("{server}__{tool}")).with_arguments(json);
217    let result = timed(Duration::from_secs(timeout_seconds), client.call_tool_once(params)).await?;
218    let result = complete(result)?;
219    if result.is_error.unwrap_or(false) {
220        return Err(McpCommandError::Tool(result_json(&result).to_string()));
221    }
222    println!("{}", serde_json::to_string(&result_json(&result)).map_err(McpCommandError::Output)?);
223    Ok(())
224}
225
226async fn list_tools(
227    client: &rmcp::service::RunningService<rmcp::RoleClient, ()>,
228) -> Result<Vec<Tool>, McpCommandError> {
229    timed(DISCOVERY_TIMEOUT, client.list_all_tools()).await
230}
231
232async fn timed<T, U: Display>(
233    duration: Duration,
234    future: impl Future<Output = Result<T, U>>,
235) -> Result<T, McpCommandError> {
236    tokio::time::timeout(duration, future)
237        .await
238        .map_err(|_| McpCommandError::Timeout(duration.as_secs()))?
239        .map_err(|error| McpCommandError::Request(error.to_string()))
240}
241
242fn complete(response: CallToolResponse) -> Result<CallToolResult, McpCommandError> {
243    match response {
244        CallToolResponse::Complete(result) => Ok(result),
245        other => Err(McpCommandError::Request(format!("gateway returned an incomplete response: {other:?}"))),
246    }
247}
248
249fn tools_for_server<'a>(tools: &'a [Tool], server: &str) -> Vec<&'a Tool> {
250    let prefix = format!("{server}__");
251    tools.iter().filter(|tool| tool.name.starts_with(&prefix)).collect()
252}
253
254fn print_servers_help(servers: &[ServerSummary]) {
255    println!("Discover and call deferred MCP tools through the active Aether session.\n");
256    println!("Usage: aether mcp <server> --help\n");
257    println!("Deferred MCP servers:");
258    for server in servers {
259        println!("  {:<20} {}", server.name, server.description);
260    }
261}
262
263fn print_server_help(server: &str, tools: Vec<&Tool>) {
264    println!("Deferred tools from `{server}`.\n");
265    println!("Usage: aether mcp {server} <tool> --help\n");
266    println!("Tools:");
267    for tool in tools {
268        let local_name = tool.name.strip_prefix(&format!("{server}__")).unwrap_or(tool.name.as_ref());
269        println!("  {:<20} {}", local_name, tool.description.as_deref().unwrap_or_default());
270    }
271}
272
273fn print_tool_help(server: &str, tool: &str, definition: &Tool) -> Result<(), McpCommandError> {
274    println!("{}\n", definition.description.as_deref().unwrap_or_default());
275    println!("Usage:");
276    println!("  aether mcp {server} {tool} --json '{{...}}'");
277    println!("  printf '%s' '{{...}}' | aether mcp {server} {tool}\n");
278    println!("Input schema:");
279    println!("{}", serde_json::to_string_pretty(definition.input_schema.as_ref()).map_err(McpCommandError::Output)?);
280    Ok(())
281}
282
283fn result_json(result: &CallToolResult) -> Value {
284    result.structured_content.clone().unwrap_or_else(|| serde_json::to_value(&result.content).unwrap_or(Value::Null))
285}