Skip to main content

machi_runtime/
workflow_host.rs

1//! Adapter: `machi-workflow` host channel → [`SessionHost`] nested runs.
2
3use 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
19/// Run a workflow script whose `agent` / `parallel` calls resolve through `host`.
20///
21/// Blocks the calling async task on a worker thread for the Rhai engine while
22/// servicing host requests on the current runtime.
23///
24/// # Errors
25///
26/// Propagates channel / join failures as [`HostError::Failed`]. Workflow
27/// terminal outcomes are returned as [`WorkflowOutcome`] (including
28/// `Failed` / `BudgetExceeded` variants) rather than `Err`.
29pub 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
37/// Like [`run_workflow_on_host`] with an explicit metrics sink.
38///
39/// # Errors
40///
41/// Same as [`run_workflow_on_host`].
42pub 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
58/// Full configuration: metrics + side-effect store (scratch / templates).
59///
60/// # Errors
61///
62/// Same as [`run_workflow_on_host`].
63pub 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            // SpawnAgent runs concurrent so parallel() fan-out is real concurrency.
90            // Other requests are cheap and handled inline.
91            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    // Dropping the sender (inside run_workflow when it finishes) ends the service loop.
117    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            // Concurrent path is handled in the service loop.
200            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
280/// Map workflow [`AgentOpts`] → host [`SpawnOpts`] without silent field drops.
281///
282/// `fork_context` requires host `parent_handle` or explicit `fork_messages`.
283/// `resume_from` requires host `run_store` and a completed [`WorkflowRunStore`] row.
284fn 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        // Budget 0 at adapter reserve layer
424        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}