1use std::sync::Arc;
4use std::time::{Duration, Instant};
5
6use futures::future::join_all;
7use machi_obs::{NoopMetrics, SharedMetrics, record_tool_call};
8use machi_types::{ToolCall, ToolCallId};
9use tokio::time::timeout;
10use tracing::{Instrument, info_span};
11
12use crate::approval::{ApprovalDecision, ApprovalGate, AutoApprove};
13use crate::context::ToolCallContext;
14use crate::error::{ToolError, codes};
15use crate::metadata::{ConcurrencyMode, Destructiveness, ToolMetadata};
16use crate::registry::{CapabilityMode, ToolRegistry};
17use crate::stream::drain_terminal;
18use crate::tool::{DynTool, SharedTool, ToolResult};
19
20#[derive(Debug, Clone)]
22pub struct DispatchRequest {
23 pub call: ToolCall,
25}
26
27#[derive(Debug, Clone)]
29pub struct DispatchOutcome {
30 pub id: ToolCallId,
32 pub name: String,
34 pub result: Result<ToolResult, ToolError>,
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
40#[non_exhaustive]
41pub enum ApprovalPolicy {
42 Never,
44 #[default]
46 Destructive,
47 Always,
49}
50
51#[derive(Clone)]
53pub struct ToolDispatch {
54 pub max_concurrency: usize,
56 pub capability_mode: CapabilityMode,
58 pub approval: Arc<dyn ApprovalGate>,
60 pub approval_policy: ApprovalPolicy,
62 pub metrics: SharedMetrics,
64}
65
66impl std::fmt::Debug for ToolDispatch {
67 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
68 f.debug_struct("ToolDispatch")
69 .field("max_concurrency", &self.max_concurrency)
70 .field("capability_mode", &self.capability_mode)
71 .field("approval_policy", &self.approval_policy)
72 .finish_non_exhaustive()
73 }
74}
75
76impl Default for ToolDispatch {
77 fn default() -> Self {
78 Self {
79 max_concurrency: 32,
80 capability_mode: CapabilityMode::Full,
81 approval: Arc::new(AutoApprove),
82 approval_policy: ApprovalPolicy::Destructive,
83 metrics: Arc::new(NoopMetrics),
84 }
85 }
86}
87
88impl ToolDispatch {
89 #[must_use]
91 pub fn with_capability(mut self, mode: CapabilityMode) -> Self {
92 self.capability_mode = mode;
93 self
94 }
95
96 #[must_use]
98 pub const fn with_max_concurrency(mut self, n: usize) -> Self {
99 self.max_concurrency = n;
100 self
101 }
102
103 #[must_use]
105 pub fn with_approval(mut self, gate: Arc<dyn ApprovalGate>) -> Self {
106 self.approval = gate;
107 self
108 }
109
110 #[must_use]
112 pub const fn with_approval_policy(mut self, policy: ApprovalPolicy) -> Self {
113 self.approval_policy = policy;
114 self
115 }
116
117 #[must_use]
119 pub fn with_metrics(mut self, metrics: SharedMetrics) -> Self {
120 self.metrics = metrics;
121 self
122 }
123
124 pub async fn execute_batch(
126 &self,
127 registry: &ToolRegistry,
128 ctx: ToolCallContext,
129 requests: Vec<DispatchRequest>,
130 ) -> Vec<DispatchOutcome> {
131 if requests.is_empty() {
132 return Vec::new();
133 }
134
135 let mut outcomes: Vec<Option<DispatchOutcome>> =
136 (0..requests.len()).map(|_| None).collect();
137 let mut index = 0usize;
138
139 while index < requests.len() {
140 if ctx.is_cancelled() {
141 fill_cancelled(&requests, &mut outcomes, index);
142 break;
143 }
144
145 let Some(req) = requests.get(index) else {
146 break;
147 };
148
149 match prepare_call(registry, self.capability_mode, req) {
150 Prepare::Deny(out) | Prepare::Missing(out) => {
151 set_outcome(&mut outcomes, index, out);
152 index = index.saturating_add(1);
153 }
154 Prepare::Ready(tool)
155 if tool.metadata().concurrency == ConcurrencyMode::Exclusive =>
156 {
157 let out = self.run_one(tool.as_ref(), ctx.clone(), req).await;
158 set_outcome(&mut outcomes, index, out);
159 index = index.saturating_add(1);
160 }
161 Prepare::Ready(_) => {
162 index = self
163 .run_concurrent_window(registry, &ctx, &requests, &mut outcomes, index)
164 .await;
165 }
166 }
167 }
168
169 finalize_outcomes(&requests, outcomes)
170 }
171
172 async fn run_concurrent_window(
173 &self,
174 registry: &ToolRegistry,
175 ctx: &ToolCallContext,
176 requests: &[DispatchRequest],
177 outcomes: &mut [Option<DispatchOutcome>],
178 index: usize,
179 ) -> usize {
180 let window = collect_concurrent_window(
181 registry,
182 self.capability_mode,
183 requests,
184 index,
185 self.max_concurrency.max(1),
186 );
187 let next = window.last().map_or(index + 1, |i| i.saturating_add(1));
188 let futs = window.into_iter().filter_map(|win_i| {
189 let win_req = requests.get(win_i)?.clone();
190 let win_tool = registry.require(&win_req.call.name).ok()?;
191 let win_ctx = ctx.clone();
192 Some(async move {
193 (
194 win_i,
195 self.run_one(win_tool.as_ref(), win_ctx, &win_req).await,
196 )
197 })
198 });
199 for (i, out) in join_all(futs).await {
200 set_outcome(outcomes, i, out);
201 }
202 next
203 }
204
205 async fn run_one(
206 &self,
207 tool: &dyn DynTool,
208 ctx: ToolCallContext,
209 req: &DispatchRequest,
210 ) -> DispatchOutcome {
211 let span = info_span!(
212 "machi.tool",
213 machi.tool_name = tool.name(),
214 machi.tool_call_id = %req.call.id,
215 );
216 let meta = tool.metadata();
217 let started = Instant::now();
218 let result = async { self.execute_tool(tool, &meta, ctx, req).await }
219 .instrument(span)
220 .await;
221 let ms = started.elapsed().as_secs_f64() * 1000.0;
222 let status = match &result {
223 Ok(r) if r.is_error => "tool_error",
224 Ok(_) => "ok",
225 Err(e) if e.code() == machi_types::ErrorCode::ToolCancelled => "cancelled",
226 Err(e) if e.code() == machi_types::ErrorCode::ToolApprovalDenied => "denied",
227 Err(_) => "error",
228 };
229 record_tool_call(self.metrics.as_ref(), tool.name(), status, ms);
230
231 DispatchOutcome {
232 id: req.call.id.clone(),
233 name: req.call.name.clone(),
234 result,
235 }
236 }
237
238 async fn execute_tool(
239 &self,
240 tool: &dyn DynTool,
241 meta: &ToolMetadata,
242 ctx: ToolCallContext,
243 req: &DispatchRequest,
244 ) -> Result<ToolResult, ToolError> {
245 if ctx.is_cancelled() {
246 return Err(codes::cancelled());
247 }
248 self.check_approval(tool, meta, &req.call.arguments).await?;
249 let fut = async {
250 let stream = tool.execute(ctx.clone(), req.call.arguments.clone()).await;
251 drain_terminal(stream).await
252 };
253 let limit = meta
254 .timeout
255 .or_else(|| ctx.deadline.map(|d| d.remaining()).filter(|d| !d.is_zero()));
256 match limit {
257 Some(limit) => match timeout(limit.max(Duration::from_millis(1)), fut).await {
258 Ok(r) => r,
259 Err(_) => Err(codes::timeout(format!("tool '{}' timed out", tool.name()))),
260 },
261 None => fut.await,
262 }
263 }
264
265 async fn check_approval(
266 &self,
267 tool: &dyn DynTool,
268 meta: &ToolMetadata,
269 arguments: &serde_json::Value,
270 ) -> Result<(), ToolError> {
271 if !needs_approval(self.approval_policy, meta) {
272 return Ok(());
273 }
274 match self.approval.approve(tool, meta, arguments).await? {
275 ApprovalDecision::Allow => Ok(()),
276 ApprovalDecision::Deny => Err(codes::approval_denied(format!(
277 "approval denied for tool {}",
278 tool.name()
279 ))),
280 }
281 }
282}
283
284fn set_outcome(outcomes: &mut [Option<DispatchOutcome>], index: usize, out: DispatchOutcome) {
285 if let Some(slot) = outcomes.get_mut(index) {
286 *slot = Some(out);
287 }
288}
289
290fn needs_approval(policy: ApprovalPolicy, meta: &ToolMetadata) -> bool {
291 match policy {
292 ApprovalPolicy::Never => false,
293 ApprovalPolicy::Always => true,
294 ApprovalPolicy::Destructive => {
295 meta.destructiveness != Destructiveness::None
296 || meta.capabilities.iter().any(|c| {
297 matches!(
298 c,
299 crate::metadata::CapabilityFlag::Write
300 | crate::metadata::CapabilityFlag::Execute
301 )
302 })
303 }
304 }
305}
306
307enum Prepare {
308 Ready(SharedTool),
309 Missing(DispatchOutcome),
310 Deny(DispatchOutcome),
311}
312
313fn prepare_call(registry: &ToolRegistry, mode: CapabilityMode, req: &DispatchRequest) -> Prepare {
314 match registry.require(&req.call.name) {
315 Err(err) => Prepare::Missing(DispatchOutcome {
316 id: req.call.id.clone(),
317 name: req.call.name.clone(),
318 result: Err(err),
319 }),
320 Ok(tool) if !registry.allows(tool.as_ref(), mode) => Prepare::Deny(DispatchOutcome {
321 id: req.call.id.clone(),
322 name: req.call.name.clone(),
323 result: Err(codes::denied(format!(
324 "tool '{}' denied by capability mode {mode:?}",
325 req.call.name
326 ))),
327 }),
328 Ok(tool) => Prepare::Ready(tool),
329 }
330}
331
332fn collect_concurrent_window(
333 registry: &ToolRegistry,
334 mode: CapabilityMode,
335 requests: &[DispatchRequest],
336 start: usize,
337 max: usize,
338) -> Vec<usize> {
339 let mut window = Vec::new();
340 let mut per_tool: std::collections::HashMap<String, usize> = std::collections::HashMap::new();
341 let mut j = start;
342 while j < requests.len() && window.len() < max {
343 let Some(req) = requests.get(j) else {
344 break;
345 };
346 let Ok(tool) = registry.require(&req.call.name) else {
347 break;
348 };
349 if !registry.allows(tool.as_ref(), mode) {
350 break;
351 }
352 let meta = tool.metadata();
353 if meta.concurrency == ConcurrencyMode::Exclusive {
354 if window.is_empty() {
356 window.push(j);
357 }
358 break;
359 }
360 if let Some(cap) = meta.max_concurrency {
362 let count = per_tool.entry(req.call.name.clone()).or_insert(0);
363 if *count >= cap.max(1) {
364 if window.is_empty() {
366 window.push(j);
368 }
369 break;
370 }
371 *count = count.saturating_add(1);
372 }
373 window.push(j);
374 j = j.saturating_add(1);
375 }
376 if window.is_empty() {
377 window.push(start);
379 }
380 window
381}
382
383fn fill_cancelled(
384 requests: &[DispatchRequest],
385 outcomes: &mut [Option<DispatchOutcome>],
386 from: usize,
387) {
388 for (i, req) in requests.iter().enumerate().skip(from) {
389 if let Some(slot) = outcomes.get_mut(i)
390 && slot.is_none()
391 {
392 *slot = Some(DispatchOutcome {
393 id: req.call.id.clone(),
394 name: req.call.name.clone(),
395 result: Err(codes::cancelled()),
396 });
397 }
398 }
399}
400
401fn finalize_outcomes(
402 requests: &[DispatchRequest],
403 outcomes: Vec<Option<DispatchOutcome>>,
404) -> Vec<DispatchOutcome> {
405 outcomes
406 .into_iter()
407 .enumerate()
408 .map(|(i, o)| {
409 o.unwrap_or_else(|| {
410 let req = requests.get(i);
411 DispatchOutcome {
412 id: req.map_or_else(ToolCallId::generate, |r| r.call.id.clone()),
413 name: req.map_or_else(|| "unknown".into(), |r| r.call.name.clone()),
414 result: Err(codes::execution("dispatch internal gap")),
415 }
416 })
417 })
418 .collect()
419}
420
421#[cfg(test)]
422#[allow(clippy::expect_used, clippy::unwrap_used, reason = "unit tests")]
423mod tests {
424 use super::*;
425 use crate::tool::{DynTool, ToolResult};
426 use async_trait::async_trait;
427 use machi_types::{ToolCall, ToolCallId};
428 use serde_json::json;
429
430 struct CapTool {
431 name: String,
432 cap: usize,
433 }
434
435 #[async_trait]
436 impl DynTool for CapTool {
437 fn name(&self) -> &str {
438 &self.name
439 }
440 fn description(&self) -> &str {
441 "cap"
442 }
443 fn parameters(&self) -> serde_json::Value {
444 json!({})
445 }
446 fn metadata(&self) -> ToolMetadata {
447 ToolMetadata {
448 concurrency: ConcurrencyMode::Concurrent,
449 max_concurrency: Some(self.cap),
450 ..Default::default()
451 }
452 }
453 async fn call(
454 &self,
455 _ctx: ToolCallContext,
456 _args: serde_json::Value,
457 ) -> Result<ToolResult, ToolError> {
458 Ok(ToolResult::text("ok"))
459 }
460 }
461
462 #[test]
463 fn per_tool_max_concurrency_limits_window() {
464 let reg = ToolRegistry::from_tools(vec![Arc::new(CapTool {
465 name: "a".into(),
466 cap: 1,
467 })]);
468 let reqs: Vec<DispatchRequest> = (0..3)
469 .map(|i| DispatchRequest {
470 call: ToolCall {
471 id: ToolCallId::new(format!("c{i}")).expect("id"),
472 name: "a".into(),
473 arguments: json!({}),
474 },
475 })
476 .collect();
477 let window = collect_concurrent_window(®, CapabilityMode::Full, &reqs, 0, 32);
478 assert_eq!(
479 window.len(),
480 1,
481 "cap=1 must not fan out three concurrent a()"
482 );
483 }
484
485 use std::sync::Arc;
486 use std::sync::atomic::{AtomicUsize, Ordering};
487
488 use machi_types::ErrorCode;
489 use tokio::sync::Barrier;
490
491 use crate::approval::AlwaysDeny;
492 use crate::metadata::ToolMetadata;
493
494 struct CountingTool {
495 name: String,
496 meta: ToolMetadata,
497 active: Arc<AtomicUsize>,
498 max_active: Arc<AtomicUsize>,
499 barrier: Option<Arc<Barrier>>,
500 }
501
502 #[async_trait]
503 impl DynTool for CountingTool {
504 fn name(&self) -> &str {
505 &self.name
506 }
507 fn description(&self) -> &str {
508 "test"
509 }
510 fn parameters(&self) -> serde_json::Value {
511 json!({"type":"object","properties":{}})
512 }
513 fn metadata(&self) -> ToolMetadata {
514 self.meta.clone()
515 }
516 async fn call(
517 &self,
518 _ctx: ToolCallContext,
519 _arguments: serde_json::Value,
520 ) -> Result<ToolResult, ToolError> {
521 let n = self.active.fetch_add(1, Ordering::SeqCst) + 1;
522 self.max_active.fetch_max(n, Ordering::SeqCst);
523 if let Some(b) = &self.barrier {
524 b.wait().await;
525 }
526 self.active.fetch_sub(1, Ordering::SeqCst);
527 Ok(ToolResult::text("ok"))
528 }
529 }
530
531 fn call(name: &str, id: &str) -> DispatchRequest {
532 DispatchRequest {
533 call: ToolCall {
534 id: ToolCallId::new(id).expect("id"),
535 name: name.into(),
536 arguments: json!({}),
537 },
538 }
539 }
540
541 #[tokio::test]
542 async fn concurrent_readonly_overlap() {
543 let active = Arc::new(AtomicUsize::new(0));
544 let max_active = Arc::new(AtomicUsize::new(0));
545 let barrier = Arc::new(Barrier::new(2));
546 let t1 = Arc::new(CountingTool {
547 name: "r1".into(),
548 meta: ToolMetadata {
549 concurrency: ConcurrencyMode::ReadOnly,
550 ..ToolMetadata::read_only()
551 },
552 active: Arc::clone(&active),
553 max_active: Arc::clone(&max_active),
554 barrier: Some(Arc::clone(&barrier)),
555 });
556 let t2 = Arc::new(CountingTool {
557 name: "r2".into(),
558 meta: ToolMetadata {
559 concurrency: ConcurrencyMode::ReadOnly,
560 ..ToolMetadata::read_only()
561 },
562 active: Arc::clone(&active),
563 max_active: Arc::clone(&max_active),
564 barrier: Some(barrier),
565 });
566 let reg = ToolRegistry::from_tools(vec![t1, t2]);
567 let outs = ToolDispatch::default()
568 .execute_batch(
569 ®,
570 ToolCallContext::default(),
571 vec![call("r1", "c1"), call("r2", "c2")],
572 )
573 .await;
574 assert_eq!(outs.len(), 2);
575 assert!(outs.iter().all(|o| o.result.is_ok()));
576 assert!(
577 max_active.load(Ordering::SeqCst) >= 2,
578 "expected overlap, max={}",
579 max_active.load(Ordering::SeqCst)
580 );
581 }
582
583 #[tokio::test]
584 async fn exclusive_serial() {
585 let active = Arc::new(AtomicUsize::new(0));
586 let max_active = Arc::new(AtomicUsize::new(0));
587 let t1 = Arc::new(CountingTool {
588 name: "e1".into(),
589 meta: ToolMetadata::exclusive_write(),
590 active: Arc::clone(&active),
591 max_active: Arc::clone(&max_active),
592 barrier: None,
593 });
594 let t2 = Arc::new(CountingTool {
595 name: "e2".into(),
596 meta: ToolMetadata::exclusive_write(),
597 active,
598 max_active: Arc::clone(&max_active),
599 barrier: None,
600 });
601 let reg = ToolRegistry::from_tools(vec![t1, t2]);
602 let outs = ToolDispatch::default()
603 .execute_batch(
604 ®,
605 ToolCallContext::default(),
606 vec![call("e1", "c1"), call("e2", "c2")],
607 )
608 .await;
609 assert!(outs.iter().all(|o| o.result.is_ok()));
610 assert_eq!(max_active.load(Ordering::SeqCst), 1);
611 }
612
613 #[tokio::test]
614 async fn readonly_mode_denies_write() {
615 let tool = Arc::new(CountingTool {
616 name: "w".into(),
617 meta: ToolMetadata::exclusive_write(),
618 active: Arc::new(AtomicUsize::new(0)),
619 max_active: Arc::new(AtomicUsize::new(0)),
620 barrier: None,
621 });
622 let reg = ToolRegistry::from_tools(vec![tool]);
623 let dispatch = ToolDispatch::default().with_capability(CapabilityMode::ReadOnly);
624 let outs = dispatch
625 .execute_batch(®, ToolCallContext::default(), vec![call("w", "c1")])
626 .await;
627 let err = outs
628 .first()
629 .expect("one outcome")
630 .result
631 .as_ref()
632 .expect_err("denied");
633 assert_eq!(err.code(), ErrorCode::ToolDenied);
634 }
635
636 #[tokio::test]
637 async fn approval_blocks_destructive() {
638 let tool = Arc::new(CountingTool {
639 name: "w".into(),
640 meta: ToolMetadata::exclusive_write(),
641 active: Arc::new(AtomicUsize::new(0)),
642 max_active: Arc::new(AtomicUsize::new(0)),
643 barrier: None,
644 });
645 let reg = ToolRegistry::from_tools(vec![tool]);
646 let dispatch = ToolDispatch::default().with_approval(Arc::new(AlwaysDeny));
647 let outs = dispatch
648 .execute_batch(®, ToolCallContext::default(), vec![call("w", "c1")])
649 .await;
650 let err = outs
651 .first()
652 .expect("one")
653 .result
654 .as_ref()
655 .expect_err("approval");
656 assert_eq!(err.code(), ErrorCode::ToolApprovalDenied);
657 }
658
659 struct SlowTool;
660
661 #[async_trait]
662 impl DynTool for SlowTool {
663 fn name(&self) -> &str {
664 "slow"
665 }
666 fn description(&self) -> &str {
667 "sleeps"
668 }
669 fn parameters(&self) -> serde_json::Value {
670 json!({"type":"object","properties":{}})
671 }
672 fn metadata(&self) -> ToolMetadata {
673 ToolMetadata {
674 timeout: Some(Duration::from_millis(20)),
675 ..ToolMetadata::read_only()
676 }
677 }
678 async fn call(
679 &self,
680 _ctx: ToolCallContext,
681 _arguments: serde_json::Value,
682 ) -> Result<ToolResult, ToolError> {
683 tokio::time::sleep(Duration::from_secs(5)).await;
684 Ok(ToolResult::text("late"))
685 }
686 }
687
688 #[tokio::test]
689 async fn tool_timeout_matrix() {
690 let reg = ToolRegistry::from_tools(vec![Arc::new(SlowTool)]);
691 let outs = ToolDispatch::default()
692 .execute_batch(®, ToolCallContext::default(), vec![call("slow", "c1")])
693 .await;
694 let err = outs
695 .first()
696 .expect("one")
697 .result
698 .as_ref()
699 .expect_err("timeout");
700 assert_eq!(err.code(), ErrorCode::ToolTimeout);
701 }
702
703 struct CancelAwareTool;
704
705 #[async_trait]
706 impl DynTool for CancelAwareTool {
707 fn name(&self) -> &str {
708 "cancel_me"
709 }
710 fn description(&self) -> &str {
711 "waits for cancel"
712 }
713 fn parameters(&self) -> serde_json::Value {
714 json!({"type":"object","properties":{}})
715 }
716 fn metadata(&self) -> ToolMetadata {
717 ToolMetadata::read_only()
718 }
719 async fn call(
720 &self,
721 ctx: ToolCallContext,
722 _arguments: serde_json::Value,
723 ) -> Result<ToolResult, ToolError> {
724 ctx.cancel.cancelled().await;
725 Err(codes::cancelled())
726 }
727 }
728
729 #[tokio::test]
730 async fn tool_cancel_matrix() {
731 use tokio_util::sync::CancellationToken;
732
733 let reg = ToolRegistry::from_tools(vec![Arc::new(CancelAwareTool)]);
734 let cancel = CancellationToken::new();
735 let ctx = ToolCallContext::default().with_cancel(cancel.clone());
736 let dispatch = ToolDispatch::default();
737 let handle = tokio::spawn(async move {
738 dispatch
739 .execute_batch(®, ctx, vec![call("cancel_me", "c1")])
740 .await
741 });
742 tokio::time::sleep(Duration::from_millis(10)).await;
744 cancel.cancel();
745 let outs = handle.await.expect("join");
746 let err = outs
747 .first()
748 .expect("one")
749 .result
750 .as_ref()
751 .expect_err("cancelled");
752 assert_eq!(err.code(), ErrorCode::ToolCancelled);
753 }
754
755 #[tokio::test]
756 async fn batch_cancel_fills_remaining() {
757 use tokio_util::sync::CancellationToken;
758
759 let reg = ToolRegistry::from_tools(vec![
760 Arc::new(CountingTool {
761 name: "r1".into(),
762 meta: ToolMetadata::read_only(),
763 active: Arc::new(AtomicUsize::new(0)),
764 max_active: Arc::new(AtomicUsize::new(0)),
765 barrier: None,
766 }),
767 Arc::new(CountingTool {
768 name: "r2".into(),
769 meta: ToolMetadata::read_only(),
770 active: Arc::new(AtomicUsize::new(0)),
771 max_active: Arc::new(AtomicUsize::new(0)),
772 barrier: None,
773 }),
774 ]);
775 let cancel = CancellationToken::new();
776 cancel.cancel();
777 let outs = ToolDispatch::default()
778 .execute_batch(
779 ®,
780 ToolCallContext::default().with_cancel(cancel),
781 vec![call("r1", "c1"), call("r2", "c2")],
782 )
783 .await;
784 assert_eq!(outs.len(), 2);
785 for o in &outs {
786 let err = o.result.as_ref().expect_err("cancelled");
787 assert_eq!(err.code(), ErrorCode::ToolCancelled);
788 }
789 }
790}