1use std::sync::Arc;
4use std::sync::atomic::{AtomicU64, Ordering};
5
6use machi_obs::{NoopMetrics, SharedMetrics, record_workflow_agents, record_workflow_run};
7use machi_tools::registry::CapabilityMode;
8use machi_workflow::{
9 AgentOpts, AgentResult, BudgetState, HostError, WorkflowHostRequest, WorkflowOutcome,
10 WorkflowRunParams, run_workflow,
11};
12use tokio::sync::mpsc;
13use tokio_util::sync::CancellationToken;
14use tracing::{Instrument, info_span};
15
16use crate::host::{SessionHost, SpawnOpts};
17use crate::side_effects::WorkflowSideEffects;
18
19pub async fn run_workflow_on_host(
30 host: Arc<dyn SessionHost>,
31 params: WorkflowRunParams,
32 agent_budget: Option<u64>,
33) -> Result<WorkflowOutcome, HostError> {
34 run_workflow_on_host_with_metrics(host, params, agent_budget, Arc::new(NoopMetrics)).await
35}
36
37pub async fn run_workflow_on_host_with_metrics(
43 host: Arc<dyn SessionHost>,
44 params: WorkflowRunParams,
45 agent_budget: Option<u64>,
46 metrics: SharedMetrics,
47) -> Result<WorkflowOutcome, HostError> {
48 run_workflow_configured(
49 host,
50 params,
51 agent_budget,
52 metrics,
53 WorkflowSideEffects::shared(),
54 )
55 .await
56}
57
58pub async fn run_workflow_configured(
64 host: Arc<dyn SessionHost>,
65 mut params: WorkflowRunParams,
66 agent_budget: Option<u64>,
67 metrics: SharedMetrics,
68 effects: Arc<WorkflowSideEffects>,
69) -> Result<WorkflowOutcome, HostError> {
70 let (tx, mut rx) = mpsc::unbounded_channel::<WorkflowHostRequest>();
71 let spent = Arc::new(AtomicU64::new(0));
72 let reserved = Arc::new(AtomicU64::new(0));
73 let cancel = params.cancel.clone();
74
75 let budget = agent_budget;
76 let spent_h = Arc::clone(&spent);
77 let reserved_h = Arc::clone(&reserved);
78 let cancel_h = cancel.clone();
79 let host_svc = Arc::clone(&host);
80 let effects_svc = Arc::clone(&effects);
81
82 let service = tokio::spawn(async move {
83 let mut inflight = Vec::new();
84 while let Some(req) = rx.recv().await {
85 if cancel_h.is_cancelled() {
86 reply_cancelled(req);
87 continue;
88 }
89 match req {
92 WorkflowHostRequest::SpawnAgent { opts, reply } => {
93 let host = Arc::clone(&host_svc);
94 let spent = Arc::clone(&spent_h);
95 let reserved = Arc::clone(&reserved_h);
96 let cancel = cancel_h.clone();
97 inflight.push(tokio::spawn(async move {
98 handle_spawn(host.as_ref(), opts, reply, &spent, &reserved, &cancel).await;
99 }));
100 }
101 other => {
102 dispatch_inline(other, budget, &spent_h, &reserved_h, effects_svc.as_ref());
103 }
104 }
105 }
106 for t in inflight {
107 let _ = t.await;
108 }
109 });
110
111 params.host_tx = tx;
112 let outcome = tokio::task::spawn_blocking(move || run_workflow(params))
113 .await
114 .map_err(|e| HostError::Failed(format!("workflow join: {e}")))?;
115
116 let _ = service.await;
118
119 let spent_n = spent.load(Ordering::Relaxed);
120 record_workflow_agents(metrics.as_ref(), spent_n);
121 record_workflow_run(metrics.as_ref(), outcome_label(&outcome));
122 Ok(outcome)
123}
124
125fn outcome_label(outcome: &WorkflowOutcome) -> &'static str {
126 match outcome {
127 WorkflowOutcome::Completed { .. } => "completed",
128 WorkflowOutcome::Paused { .. } => "paused",
129 WorkflowOutcome::BudgetExceeded { .. } => "budget_exceeded",
130 WorkflowOutcome::Cancelled => "cancelled",
131 WorkflowOutcome::Failed { .. } => "failed",
132 _ => "other",
133 }
134}
135
136async fn handle_spawn(
137 host: &dyn SessionHost,
138 opts: AgentOpts,
139 reply: tokio::sync::oneshot::Sender<Result<AgentResult, HostError>>,
140 spent: &AtomicU64,
141 reserved: &AtomicU64,
142 cancel: &CancellationToken,
143) {
144 let span = info_span!(
145 "machi.workflow.host",
146 machi.workflow.kind = "spawn_agent",
147 machi.agent_label = opts.label.as_deref().unwrap_or(""),
148 );
149 let result = async {
150 if cancel.is_cancelled() {
151 return Err(HostError::Cancelled);
152 }
153 let spawn = to_spawn_opts(opts, cancel.child_token());
154 match host.spawn_agent(spawn).await {
155 Ok(run) => {
156 spent.fetch_add(1, Ordering::Relaxed);
157 let r = reserved.load(Ordering::Relaxed);
158 reserved.fetch_sub(r.min(1), Ordering::Relaxed);
159 let tokens = u64::from(run.usage.total_tokens);
160 Ok(AgentResult {
161 agent_id: run.agent_id.to_string(),
162 success: run.success && !run.cancelled,
163 output: run.output,
164 cancelled: run.cancelled,
165 tokens_used: tokens,
166 duration_ms: run.duration_ms,
167 })
168 }
169 Err(e) => {
170 let r = reserved.load(Ordering::Relaxed);
171 reserved.fetch_sub(r.min(1), Ordering::Relaxed);
172 Err(map_host_spawn_error(e))
173 }
174 }
175 }
176 .instrument(span)
177 .await;
178 let _ = reply.send(result);
179}
180
181fn dispatch_inline(
182 req: WorkflowHostRequest,
183 budget: Option<u64>,
184 spent: &AtomicU64,
185 reserved: &AtomicU64,
186 effects: &WorkflowSideEffects,
187) {
188 match req {
189 WorkflowHostRequest::ReserveAgentCalls { count, reply } => {
190 let result = reserve(budget, spent, reserved, count);
191 let _ = reply.send(result);
192 }
193 WorkflowHostRequest::ReleaseAgentCalls { count, reply } => {
194 let r = reserved.load(Ordering::Relaxed);
195 reserved.fetch_sub(count.min(r), Ordering::Relaxed);
196 let _ = reply.send(Ok(()));
197 }
198 WorkflowHostRequest::SpawnAgent { reply, .. } => {
199 let _ = reply.send(Err(HostError::Failed(
201 "internal: SpawnAgent must be handled concurrently".into(),
202 )));
203 }
204 WorkflowHostRequest::BudgetQuery { reply } => {
205 let s = spent.load(Ordering::Relaxed);
206 let r = reserved.load(Ordering::Relaxed);
207 let state = BudgetState {
208 total: budget,
209 spent: s,
210 reserved: r,
211 remaining: budget.map(|b| b.saturating_sub(s.saturating_add(r))),
212 };
213 let _ = reply.send(Ok(state));
214 }
215 WorkflowHostRequest::Phase { title, replayed } => {
216 tracing::info!(target: "machi.workflow", %title, replayed, "phase");
217 }
218 WorkflowHostRequest::Log { message, replayed } => {
219 tracing::info!(target: "machi.workflow", %message, replayed, "log");
220 }
221 WorkflowHostRequest::Telemetry {
222 name,
223 fields,
224 replayed,
225 } => {
226 tracing::info!(target: "machi.workflow", %name, %fields, replayed, "telemetry");
227 }
228 WorkflowHostRequest::RenderTemplate { reply, name, vars } => {
229 let _ = reply.send(effects.render_template(&name, &vars));
230 }
231 WorkflowHostRequest::WriteScratchFile {
232 reply,
233 name,
234 content,
235 } => {
236 let _ = reply.send(effects.write_scratch(&name, content));
237 }
238 WorkflowHostRequest::ReadScratchFile { reply, name } => {
239 let _ = reply.send(effects.read_scratch(&name));
240 }
241 WorkflowHostRequest::GitDiffSince { reply, commit } => {
242 let _ = reply.send(effects.git_diff_since(&commit));
243 }
244 }
245}
246
247fn reserve(
248 budget: Option<u64>,
249 spent: &AtomicU64,
250 reserved: &AtomicU64,
251 count: u64,
252) -> Result<(), HostError> {
253 if let Some(max) = budget {
254 loop {
255 let s = spent.load(Ordering::Acquire);
256 let r = reserved.load(Ordering::Acquire);
257 if s.saturating_add(r).saturating_add(count) > max {
258 return Err(HostError::AgentCallQuotaExceeded {
259 requested: s.saturating_add(r).saturating_add(count),
260 maximum: max,
261 });
262 }
263 if reserved
264 .compare_exchange(
265 r,
266 r.saturating_add(count),
267 Ordering::AcqRel,
268 Ordering::Acquire,
269 )
270 .is_ok()
271 {
272 return Ok(());
273 }
274 }
275 }
276 reserved.fetch_add(count, Ordering::Relaxed);
277 Ok(())
278}
279
280fn map_host_spawn_error(e: machi_types::MachiError) -> HostError {
285 use machi_types::ErrorCode;
286 match e.code() {
287 ErrorCode::HostBudget => HostError::BudgetExceeded,
288 ErrorCode::HostCancelled => HostError::Cancelled,
289 ErrorCode::HostUnsupported
290 | ErrorCode::HostDepth
291 | ErrorCode::HostConcurrency
292 | ErrorCode::AgentNotFound => HostError::Unsupported(e.message().to_owned()),
293 _ => HostError::Failed(e.to_string()),
294 }
295}
296
297fn to_spawn_opts(opts: AgentOpts, cancel: CancellationToken) -> SpawnOpts {
298 let mut spawn = SpawnOpts::new(opts.prompt).with_cancel(cancel);
299 if let Some(label) = opts.label {
300 spawn = spawn.with_label(label);
301 }
302 if let Some(model) = opts.model {
303 spawn.model = Some(model);
304 }
305 if let Some(mode) = opts.capability_mode.as_deref() {
306 spawn.capability_mode = parse_capability(mode);
307 }
308 if let Some(agent_type) = opts.agent_type {
309 spawn = spawn.with_agent_type(agent_type);
310 }
311 if let Some(schema) = opts.output_schema {
312 spawn = spawn.with_output_schema(schema);
313 }
314 if let Some(n) = opts.max_output_tokens {
315 spawn = spawn.with_max_output_tokens(n);
316 }
317 if opts.fork_context {
318 spawn = spawn.with_fork_context(true);
319 }
320 if let Some(id) = opts.resume_from {
321 spawn = spawn.with_resume_from(id);
322 }
323 spawn
324}
325
326fn parse_capability(mode: &str) -> CapabilityMode {
327 match mode {
328 "read_only" | "read-only" | "readonly" => CapabilityMode::ReadOnly,
329 "plan" => CapabilityMode::Plan,
330 _ => CapabilityMode::Full,
331 }
332}
333
334fn reply_cancelled(req: WorkflowHostRequest) {
335 match req {
336 WorkflowHostRequest::ReserveAgentCalls { reply, .. }
337 | WorkflowHostRequest::ReleaseAgentCalls { reply, .. } => {
338 let _ = reply.send(Err(HostError::Cancelled));
339 }
340 WorkflowHostRequest::SpawnAgent { reply, .. } => {
341 let _ = reply.send(Err(HostError::Cancelled));
342 }
343 WorkflowHostRequest::BudgetQuery { reply } => {
344 let _ = reply.send(Err(HostError::Cancelled));
345 }
346 WorkflowHostRequest::RenderTemplate { reply, .. }
347 | WorkflowHostRequest::WriteScratchFile { reply, .. }
348 | WorkflowHostRequest::ReadScratchFile { reply, .. }
349 | WorkflowHostRequest::GitDiffSince { reply, .. } => {
350 let _ = reply.send(Err(HostError::Cancelled));
351 }
352 WorkflowHostRequest::Phase { .. }
353 | WorkflowHostRequest::Log { .. }
354 | WorkflowHostRequest::Telemetry { .. } => {}
355 }
356}
357
358#[cfg(test)]
359mod tests {
360 use std::sync::Arc;
361
362 use machi_llm::MockSampler;
363 use machi_workflow::{Journal, WorkflowOutcome, WorkflowRunParams};
364 use tokio_util::sync::CancellationToken;
365
366 use super::*;
367 use crate::host::InProcessHost;
368
369 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
370 async fn workflow_parallel_on_session_host() {
371 let sampler = Arc::new(MockSampler::new());
372 sampler.map_user_text("a", "from-a");
373 sampler.map_user_text("b", "from-b");
374 let host: Arc<dyn SessionHost> = Arc::new(InProcessHost::new(sampler, vec![]));
375 let script = r#"
376 let meta = #{ name: "fanout", description: "test" };
377 phase("work");
378 let rs = parallel([
379 #{ prompt: "a", label: "wa" },
380 #{ prompt: "b", label: "wb" },
381 ]);
382 complete(#{ results: rs });
383 "#;
384 let (tx, _rx) = mpsc::unbounded_channel();
385 let outcome = run_workflow_on_host(
386 host,
387 WorkflowRunParams {
388 script: script.into(),
389 args: serde_json::json!({}),
390 journal: Journal::new(None),
391 host_tx: tx,
392 cancel: CancellationToken::new(),
393 max_ops: WorkflowRunParams::DEFAULT_MAX_OPS,
394 },
395 Some(16),
396 )
397 .await
398 .expect("run");
399 let WorkflowOutcome::Completed { result } = outcome else {
400 unreachable!("expected completed outcome");
401 };
402 let arr = result
403 .get("results")
404 .and_then(|v| v.as_array())
405 .expect("results array");
406 assert_eq!(arr.len(), 2);
407 assert_eq!(
408 arr.first().and_then(|v| v.get("output")),
409 Some(&serde_json::json!("from-a"))
410 );
411 assert_eq!(
412 arr.get(1).and_then(|v| v.get("output")),
413 Some(&serde_json::json!("from-b"))
414 );
415 }
416
417 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
418 async fn workflow_budget_on_host() {
419 let sampler = Arc::new(MockSampler::new());
420 sampler.push_text("x");
421 let host: Arc<dyn SessionHost> =
422 Arc::new(InProcessHost::new(sampler, vec![]).with_agent_budget(0));
423 let script = r#"
425 let meta = #{ name: "b", description: "b" };
426 agent("x");
427 complete(1);
428 "#;
429 let (tx, _rx) = mpsc::unbounded_channel();
430 let outcome = run_workflow_on_host(
431 host,
432 WorkflowRunParams {
433 script: script.into(),
434 args: serde_json::json!({}),
435 journal: Journal::new(None),
436 host_tx: tx,
437 cancel: CancellationToken::new(),
438 max_ops: WorkflowRunParams::DEFAULT_MAX_OPS,
439 },
440 Some(0),
441 )
442 .await
443 .expect("run");
444 assert!(
445 matches!(outcome, WorkflowOutcome::BudgetExceeded { .. }),
446 "{outcome:?}"
447 );
448 }
449}