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], ®istry).unwrap();
456 }
457}