Skip to main content

atman_runtime/tools/
mcp.rs

1use std::collections::BTreeSet;
2use std::time::Duration;
3
4use crate::error::RuntimeError;
5use crate::mcp::{McpServerState, McpServerStatus};
6use crate::tool::{ApprovalLevel, BoxFut, Tier, Tool, ToolArgs, ToolCtx, ToolResult};
7use crate::value::Value;
8
9pub const CONTROL_TOOL_NAMES: &[&str] = &["mcp.status", "mcp.tools", "mcp.await", "mcp.call"];
10
11pub struct McpStatus;
12pub struct McpTools;
13pub struct McpAwait;
14pub struct McpCall;
15
16pub async fn await_requested_tools(args: &ToolArgs, ctx: &ToolCtx) -> Result<(), RuntimeError> {
17    let Some(Value::List(items)) = args.named("tools") else {
18        return Ok(());
19    };
20    let selectors = items
21        .iter()
22        .filter_map(|item| match item {
23            Value::Str(selector) if selector.starts_with("mcp.") => Some(selector.clone()),
24            _ => None,
25        })
26        .collect::<Vec<_>>();
27    await_selectors(&selectors, ctx, Duration::from_secs(120)).await
28}
29
30async fn await_selectors(
31    selectors: &[String],
32    ctx: &ToolCtx,
33    timeout: Duration,
34) -> Result<(), RuntimeError> {
35    if selectors.is_empty() {
36        return Ok(());
37    }
38    let Some(session) = ctx.session_runtime.as_ref() else {
39        return Ok(());
40    };
41    let mut context = session.subscribe_context();
42    let deadline = tokio::time::Instant::now() + timeout;
43    loop {
44        let statuses = context.borrow().mcp_servers.clone();
45        let required = required_servers(selectors, &statuses);
46        if required.is_empty() {
47            return Ok(());
48        }
49        let mut pending = Vec::new();
50        for server in required {
51            let status = statuses.iter().find(|status| status.name == server);
52            match status.map(|status| &status.state) {
53                Some(McpServerState::Connected { .. }) => {}
54                Some(McpServerState::Pending | McpServerState::Connecting) => {
55                    pending.push(server);
56                }
57                Some(McpServerState::Disabled) => {
58                    return Err(RuntimeError::ToolFailed(format!(
59                        "MCP server `{server}` is disabled"
60                    )));
61                }
62                Some(McpServerState::Error { message })
63                | Some(McpServerState::Disconnected { message })
64                | Some(McpServerState::Timeout { message }) => {
65                    return Err(RuntimeError::ToolFailed(format!(
66                        "MCP server `{server}` is unavailable: {message}"
67                    )));
68                }
69                None => {}
70            }
71        }
72        if pending.is_empty() {
73            return Ok(());
74        }
75        tokio::select! {
76            _ = ctx.cancel.cancelled() => {
77                return Err(RuntimeError::Cancelled("MCP readiness wait cancelled".into()));
78            }
79            _ = tokio::time::sleep_until(deadline) => {
80                return Err(RuntimeError::ToolFailed(format!(
81                    "timed out waiting for MCP server(s): {}",
82                    pending.join(", ")
83                )));
84            }
85            changed = context.changed() => {
86                if changed.is_err() {
87                    return Err(RuntimeError::ToolFailed(
88                        "MCP readiness channel closed before connection completed".into(),
89                    ));
90                }
91            }
92        }
93    }
94}
95
96fn required_servers(selectors: &[String], statuses: &[McpServerStatus]) -> BTreeSet<String> {
97    let mut required = BTreeSet::new();
98    for selector in selectors {
99        if CONTROL_TOOL_NAMES.contains(&selector.as_str()) {
100            continue;
101        }
102        if selector == "mcp.*" {
103            required.extend(
104                statuses
105                    .iter()
106                    .filter(|status| !matches!(status.state, McpServerState::Disabled))
107                    .map(|status| status.name.clone()),
108            );
109            continue;
110        }
111        if let Some(status) = statuses
112            .iter()
113            .filter(|status| selector.starts_with(&format!("mcp.{}.", status.name)))
114            .max_by_key(|status| status.name.len())
115        {
116            required.insert(status.name.clone());
117        }
118    }
119    required
120}
121
122fn status_value(status: &McpServerStatus) -> Value {
123    let (state, tool_count, message) = match &status.state {
124        McpServerState::Disabled => ("disabled", 0, None),
125        McpServerState::Pending => ("pending", 0, None),
126        McpServerState::Connecting => ("connecting", 0, None),
127        McpServerState::Connected { tool_count, .. } => ("connected", *tool_count, None),
128        McpServerState::Error { message } => ("error", 0, Some(message.clone())),
129        McpServerState::Disconnected { message } => ("disconnected", 0, Some(message.clone())),
130        McpServerState::Timeout { message } => ("timeout", 0, Some(message.clone())),
131    };
132    Value::Struct(vec![
133        ("name".into(), Value::Str(status.name.clone())),
134        ("state".into(), Value::Str(state.into())),
135        ("tool_count".into(), Value::Int(tool_count as i64)),
136        (
137            "message".into(),
138            message.map(Value::Str).unwrap_or(Value::Unit),
139        ),
140    ])
141}
142
143fn string_arg(args: &ToolArgs, name: &str) -> Result<String, RuntimeError> {
144    match args.named(name) {
145        Some(Value::Str(value)) => Ok(value.clone()),
146        Some(value) => Err(RuntimeError::TypeMismatch {
147            expected: "string".into(),
148            actual: value.kind_name().into(),
149        }),
150        None => Err(RuntimeError::MissingArg(name.into())),
151    }
152}
153
154fn call_target(args: &ToolArgs) -> Result<(String, ToolArgs), RuntimeError> {
155    let server = string_arg(args, "server")?;
156    let tool = string_arg(args, "tool")?;
157    let target_args = match args.named("input") {
158        None | Some(Value::Unit) => ToolArgs::default(),
159        Some(Value::Struct(fields)) => ToolArgs {
160            positional: Vec::new(),
161            named: fields.clone(),
162        },
163        Some(value) => {
164            return Err(RuntimeError::TypeMismatch {
165                expected: "struct".into(),
166                actual: value.kind_name().into(),
167            });
168        }
169    };
170    Ok((format!("mcp.{server}.{tool}"), target_args))
171}
172
173impl Tool for McpStatus {
174    fn name(&self) -> &str {
175        "mcp.status"
176    }
177    fn tier(&self) -> Tier {
178        Tier::Zero
179    }
180    fn description(&self) -> Option<&str> {
181        Some("Inspect MCP connection readiness before selecting or calling a server tool.")
182    }
183    fn input_schema(&self) -> serde_json::Value {
184        serde_json::json!({"type":"object","properties":{"server":{"type":"string"}},"additionalProperties":false})
185    }
186    fn call<'a>(&'a self, args: ToolArgs, ctx: &'a ToolCtx) -> BoxFut<'a, ToolResult> {
187        Box::pin(async move {
188            let server = match args.named("server") {
189                Some(Value::Str(value)) => Some(value.as_str()),
190                Some(value) => {
191                    return Err(RuntimeError::TypeMismatch {
192                        expected: "string".into(),
193                        actual: value.kind_name().into(),
194                    });
195                }
196                None => None,
197            };
198            let session = ctx.session_runtime.as_ref().ok_or_else(|| {
199                RuntimeError::ToolFailed("mcp.status: no session available".into())
200            })?;
201            let snapshot = session.subscribe_context().borrow().clone();
202            let statuses = snapshot
203                .mcp_servers
204                .iter()
205                .filter(|status| server.is_none_or(|server| status.name == server))
206                .map(status_value)
207                .collect();
208            Ok(Value::List(statuses))
209        })
210    }
211}
212
213impl Tool for McpTools {
214    fn name(&self) -> &str {
215        "mcp.tools"
216    }
217    fn tier(&self) -> Tier {
218        Tier::Zero
219    }
220    fn description(&self) -> Option<&str> {
221        Some("List currently registered MCP tools, optionally restricted to one server.")
222    }
223    fn input_schema(&self) -> serde_json::Value {
224        serde_json::json!({"type":"object","properties":{"server":{"type":"string"}},"additionalProperties":false})
225    }
226    fn call<'a>(&'a self, args: ToolArgs, ctx: &'a ToolCtx) -> BoxFut<'a, ToolResult> {
227        Box::pin(async move {
228            let registry = ctx.registry.as_ref().ok_or_else(|| {
229                RuntimeError::ToolFailed("mcp.tools: no tool registry available".into())
230            })?;
231            let prefix = match args.named("server") {
232                Some(Value::Str(server)) => format!("mcp.{server}."),
233                Some(value) => {
234                    return Err(RuntimeError::TypeMismatch {
235                        expected: "string".into(),
236                        actual: value.kind_name().into(),
237                    });
238                }
239                None => "mcp.".into(),
240            };
241            let mut names = registry
242                .names()
243                .into_iter()
244                .filter(|name| {
245                    name.starts_with(&prefix) && !CONTROL_TOOL_NAMES.contains(&name.as_str())
246                })
247                .collect::<Vec<_>>();
248            names.sort();
249            Ok(Value::List(names.into_iter().map(Value::Str).collect()))
250        })
251    }
252}
253
254impl Tool for McpAwait {
255    fn name(&self) -> &str {
256        "mcp.await"
257    }
258    fn tier(&self) -> Tier {
259        Tier::Zero
260    }
261    fn description(&self) -> Option<&str> {
262        Some("Wait until one MCP server is connected or reports a terminal connection error.")
263    }
264    fn input_schema(&self) -> serde_json::Value {
265        serde_json::json!({"type":"object","properties":{"server":{"type":"string"},"timeout":{"type":"integer","minimum":1,"default":120}},"required":["server"],"additionalProperties":false})
266    }
267    fn call<'a>(&'a self, args: ToolArgs, ctx: &'a ToolCtx) -> BoxFut<'a, ToolResult> {
268        Box::pin(async move {
269            let server = string_arg(&args, "server")?;
270            let session = ctx.session_runtime.as_ref().ok_or_else(|| {
271                RuntimeError::ToolFailed("mcp.await: no session available".into())
272            })?;
273            if !session
274                .subscribe_context()
275                .borrow()
276                .mcp_servers
277                .iter()
278                .any(|status| status.name == server)
279            {
280                return Err(RuntimeError::ToolFailed(format!(
281                    "MCP server `{server}` is not configured"
282                )));
283            }
284            let timeout = match args.named("timeout") {
285                None => 120,
286                Some(Value::Int(value)) if *value > 0 => *value as u64,
287                Some(value) => {
288                    return Err(RuntimeError::TypeMismatch {
289                        expected: "positive int".into(),
290                        actual: value.kind_name().into(),
291                    });
292                }
293            };
294            await_selectors(
295                &[format!("mcp.{server}.*")],
296                ctx,
297                Duration::from_secs(timeout),
298            )
299            .await?;
300            Ok(Value::Bool(true))
301        })
302    }
303}
304
305impl Tool for McpCall {
306    fn name(&self) -> &str {
307        "mcp.call"
308    }
309    fn tier(&self) -> Tier {
310        Tier::Zero
311    }
312    fn approval_level(&self, args: &ToolArgs, ctx: &ToolCtx) -> ApprovalLevel {
313        let Ok((target, target_args)) = call_target(args) else {
314            return ApprovalLevel::Dangerous;
315        };
316        ctx.registry
317            .as_ref()
318            .and_then(|registry| registry.get(&target))
319            .map_or(ApprovalLevel::Dangerous, |tool| {
320                tool.approval_level(&target_args, ctx)
321            })
322    }
323    fn description(&self) -> Option<&str> {
324        Some(
325            "Call a connected MCP tool directly from At code without routing the operation through an LLM.",
326        )
327    }
328    fn input_schema(&self) -> serde_json::Value {
329        serde_json::json!({"type":"object","properties":{"server":{"type":"string"},"tool":{"type":"string"},"input":{"type":"object","default":{}}},"required":["server","tool"],"additionalProperties":false})
330    }
331    fn call<'a>(&'a self, args: ToolArgs, ctx: &'a ToolCtx) -> BoxFut<'a, ToolResult> {
332        Box::pin(async move {
333            let (target, target_args) = call_target(&args)?;
334            let server = string_arg(&args, "server")?;
335            await_selectors(&[format!("mcp.{server}.*")], ctx, Duration::from_secs(120)).await?;
336            let tool = ctx
337                .registry
338                .as_ref()
339                .and_then(|registry| registry.get(&target))
340                .ok_or_else(|| RuntimeError::UndefinedTool(target.clone()))?;
341            tool.call(target_args, ctx).await
342        })
343    }
344}
345
346#[cfg(test)]
347mod tests {
348    use super::*;
349
350    #[test]
351    fn wildcard_waits_for_all_enabled_servers() {
352        let statuses = vec![
353            McpServerStatus {
354                name: "alpha".into(),
355                transport: crate::mcp::TransportKind::Stdio,
356                state: McpServerState::Pending,
357            },
358            McpServerStatus {
359                name: "beta".into(),
360                transport: crate::mcp::TransportKind::Http,
361                state: McpServerState::Disabled,
362            },
363        ];
364        assert_eq!(
365            required_servers(&["mcp.*".into()], &statuses),
366            BTreeSet::from(["alpha".into()])
367        );
368    }
369
370    #[test]
371    fn exact_tool_waits_only_for_its_server() {
372        let statuses = vec![
373            McpServerStatus {
374                name: "mi-jira-phone".into(),
375                transport: crate::mcp::TransportKind::Stdio,
376                state: McpServerState::Pending,
377            },
378            McpServerStatus {
379                name: "other".into(),
380                transport: crate::mcp::TransportKind::Http,
381                state: McpServerState::Pending,
382            },
383        ];
384        assert_eq!(
385            required_servers(&["mcp.mi-jira-phone.jira_search".into()], &statuses),
386            BTreeSet::from(["mi-jira-phone".into()])
387        );
388    }
389
390    #[test]
391    fn exact_tool_uses_the_longest_matching_server_name() {
392        let statuses = vec![
393            McpServerStatus {
394                name: "jira".into(),
395                transport: crate::mcp::TransportKind::Stdio,
396                state: McpServerState::Pending,
397            },
398            McpServerStatus {
399                name: "jira.cloud".into(),
400                transport: crate::mcp::TransportKind::Http,
401                state: McpServerState::Pending,
402            },
403        ];
404        assert_eq!(
405            required_servers(&["mcp.jira.cloud.search".into()], &statuses),
406            BTreeSet::from(["jira.cloud".into()])
407        );
408    }
409
410    #[tokio::test]
411    async fn readiness_wait_unblocks_after_requested_server_connects() {
412        let session = std::sync::Arc::new(crate::session::Session::open_ephemeral());
413        session.update_mcp_server(McpServerStatus {
414            name: "alpha".into(),
415            transport: crate::mcp::TransportKind::Stdio,
416            state: McpServerState::Pending,
417        });
418        let ctx = ToolCtx::new().with_session_runtime(session.clone());
419        let update = session.clone();
420        let task = tokio::spawn(async move {
421            tokio::task::yield_now().await;
422            update.update_mcp_server(McpServerStatus {
423                name: "alpha".into(),
424                transport: crate::mcp::TransportKind::Stdio,
425                state: McpServerState::Connected {
426                    tool_count: 1,
427                    tools: Vec::new(),
428                },
429            });
430        });
431
432        await_selectors(&["mcp.alpha.search".into()], &ctx, Duration::from_secs(1))
433            .await
434            .unwrap();
435        task.await.unwrap();
436    }
437
438    #[test]
439    fn direct_mcp_calls_are_valid_at_nodes() {
440        let source = r#"
441flow invoke() {
442    ready = mcp.await(server: "jira")
443    tools = mcp.tools(server: "jira")
444    result = mcp.call(server: "jira", tool: "search", input: {query: "open"})
445    return {ready: ready, tools: tools, result: result}
446}
447"#;
448        let file = atman_dsl::parse::parse_file(source).unwrap();
449        let registry = crate::tool::ToolRegistry::new();
450        registry.register(std::sync::Arc::new(McpStatus));
451        registry.register(std::sync::Arc::new(McpTools));
452        registry.register(std::sync::Arc::new(McpAwait));
453        registry.register(std::sync::Arc::new(McpCall));
454
455        crate::validate::validate(&file.flows[0], &registry).unwrap();
456    }
457}