Skip to main content

atman_runtime/
exec.rs

1use std::{collections::HashMap, path::PathBuf};
2
3use atman_dsl::ast::{CmpOp, Expr, FlowDecl, Node, Stmt, WatchAction, WatchDecl, WatchEvent};
4
5use crate::env::Env;
6use crate::error::RuntimeError;
7use crate::eval::{EvalCtx, eval_expr};
8use crate::event::NodeEvent;
9use crate::provider::LlmRequest;
10use crate::tool::{BoxFut, ToolCtx, ToolRegistry};
11use crate::value::Value;
12
13fn bind_pattern(
14    pattern: &atman_dsl::ast::Pattern,
15    value: Value,
16    env: &mut Env,
17) -> Result<(), RuntimeError> {
18    use atman_dsl::ast::{Pattern, PatternFieldBinding};
19    match pattern {
20        Pattern::Ident(id) => {
21            env.bind(id.name.clone(), value);
22            Ok(())
23        }
24        Pattern::Struct { fields } => {
25            let pairs = match value {
26                Value::Struct(pairs) => pairs,
27                other => {
28                    return Err(RuntimeError::TypeMismatch {
29                        expected: "struct for destructuring bind".into(),
30                        actual: other.kind_name().into(),
31                    });
32                }
33            };
34            for field in fields {
35                let Some((_, matched)) = pairs.iter().find(|(k, _)| k == &field.source.name) else {
36                    return Err(RuntimeError::MissingArg(format!(
37                        "destructure: struct has no field `{}`",
38                        field.source.name
39                    )));
40                };
41                match &field.binding {
42                    PatternFieldBinding::Same => {
43                        env.bind(field.source.name.clone(), matched.clone());
44                    }
45                    PatternFieldBinding::Rename(target) => {
46                        env.bind(target.name.clone(), matched.clone());
47                    }
48                    PatternFieldBinding::Nested(inner) => {
49                        bind_pattern(inner, matched.clone(), env)?;
50                    }
51                }
52            }
53            Ok(())
54        }
55    }
56}
57
58pub enum StmtOutcome {
59    Continue,
60    Return(Value),
61    Err(RuntimeError),
62}
63
64pub fn exec_stmts<'a>(
65    stmts: &'a [Stmt],
66    env: &'a mut Env,
67    ctx: &'a EvalCtx<'a>,
68) -> BoxFut<'a, StmtOutcome> {
69    exec_stmts_prefixed(stmts, env, ctx, String::new())
70}
71
72pub fn exec_stmts_prefixed<'a>(
73    stmts: &'a [Stmt],
74    env: &'a mut Env,
75    ctx: &'a EvalCtx<'a>,
76    prefix: String,
77) -> BoxFut<'a, StmtOutcome> {
78    Box::pin(async move {
79        let watches = collect_watches(stmts);
80        let parent_node_id = ctx.current_node_id.clone();
81        for (i, stmt) in stmts.iter().enumerate() {
82            let node_id = if prefix.is_empty() {
83                format!("{i}")
84            } else {
85                format!("{prefix}.{i}")
86            };
87            // Check for pending L4 stop / L3 redirect between statements.
88            if let Some(session) = ctx.session_runtime.as_ref()
89                && let Some(turn_id) = ctx.turn_id.as_ref()
90            {
91                if let Some(inj) = session.peek_pending_l2_or_higher(turn_id) {
92                    match inj.level {
93                        crate::injection::InjectionLevel::L4HardStop => {
94                            session.mark_injection_consumed(&inj.id);
95                            emit_flow_node_start(ctx, &node_id, stmt, parent_node_id.as_deref());
96                            emit_flow_node_end(
97                                ctx,
98                                &node_id,
99                                &StmtOutcome::Continue,
100                                parent_node_id.as_deref(),
101                                Some("cancelled: hard stop"),
102                            );
103                            return StmtOutcome::Err(RuntimeError::Cancelled(
104                                "hard stop from user".into(),
105                            ));
106                        }
107                        crate::injection::InjectionLevel::L3Redirect => {
108                            if let Some(target) = inj.redirect_target.clone() {
109                                session.mark_injection_consumed(&inj.id);
110                                return StmtOutcome::Err(RuntimeError::Redirect(target));
111                            }
112                        }
113                        _ => {}
114                    }
115                }
116            }
117            emit_flow_node_start(ctx, &node_id, stmt, parent_node_id.as_deref());
118            let stmt_ctx = ctx.with_node(&node_id);
119            let (outcome, preview) = exec_stmt(stmt, env, &stmt_ctx, &watches).await;
120            emit_flow_node_end(
121                ctx,
122                &node_id,
123                &outcome,
124                parent_node_id.as_deref(),
125                preview.as_deref(),
126            );
127            match outcome {
128                StmtOutcome::Continue => continue,
129                other => return other,
130            }
131        }
132        StmtOutcome::Continue
133    })
134}
135
136fn emit_flow_node_start(
137    ctx: &EvalCtx<'_>,
138    node_id: &str,
139    stmt: &Stmt,
140    parent_node_id: Option<&str>,
141) {
142    let Some(run_id) = ctx.flow_run_id.clone() else {
143        return;
144    };
145    let (kind, label) = stmt_to_node_kind_label(stmt);
146    if let Some(sink) = ctx.events
147        && ctx.session_runtime.is_some()
148    {
149        sink.emit(crate::event::Event::FlowNodeStart {
150            run_id: run_id.clone(),
151            node_id: node_id.to_string(),
152            kind: kind.clone(),
153            label: label.clone(),
154            parent_node_id: parent_node_id.map(String::from),
155        });
156    }
157    if let Some(session) = ctx.session_runtime.as_ref() {
158        let _ = session
159            .stream_tx()
160            .send(crate::stream::StreamFrame::FlowNodeStart {
161                run_id: run_id.0.to_string(),
162                node_id: node_id.to_string(),
163                kind,
164                label,
165                parent_node_id: parent_node_id.map(String::from),
166            });
167    }
168}
169
170fn value_preview(v: &Value) -> Option<String> {
171    let raw = match v {
172        Value::Str(s) => s.clone(),
173        Value::Message(m) => {
174            let text = m.text_concat();
175            let tool_uses: Vec<String> = m
176                .parts
177                .iter()
178                .filter_map(|p| match p {
179                    crate::message::MessagePart::ToolUse { name, .. } => Some(name.clone()),
180                    _ => None,
181                })
182                .collect();
183            match (text.trim().is_empty(), tool_uses.is_empty()) {
184                (false, true) => text,
185                (false, false) => format!("{}\n\n→ tool_uses: {}", text, tool_uses.join(", ")),
186                (true, false) => format!("→ tool_uses: {}", tool_uses.join(", ")),
187                (true, true) => return None,
188            }
189        }
190        Value::Path(p) => p.display().to_string(),
191        Value::Int(n) => n.to_string(),
192        Value::Float(n) => n.to_string(),
193        Value::Bool(b) => b.to_string(),
194        Value::Unit => return None,
195        Value::Err(e) => format!("err: {e}"),
196        Value::List(items) => format!("list[{}]", items.len()),
197        Value::Struct(fields) => format!(
198            "{{{}}}",
199            fields
200                .iter()
201                .map(|(k, _)| k.as_str())
202                .collect::<Vec<_>>()
203                .join(", ")
204        ),
205        Value::EditProposal(_) => "<edit proposal>".into(),
206    };
207    let trimmed = raw.trim();
208    if trimmed.is_empty() {
209        None
210    } else {
211        Some(trimmed.chars().take(4000).collect())
212    }
213}
214
215fn emit_flow_node_end(
216    ctx: &EvalCtx<'_>,
217    node_id: &str,
218    outcome: &StmtOutcome,
219    parent_node_id: Option<&str>,
220    output_preview: Option<&str>,
221) {
222    let Some(run_id) = ctx.flow_run_id.clone() else {
223        return;
224    };
225    let status = match outcome {
226        StmtOutcome::Err(_) => crate::event::FlowNodeStatus::Err,
227        _ => crate::event::FlowNodeStatus::Ok,
228    };
229    let preview_owned = output_preview.map(String::from);
230    if let Some(sink) = ctx.events
231        && ctx.session_runtime.is_some()
232    {
233        sink.emit(crate::event::Event::FlowNodeEnd {
234            run_id: run_id.clone(),
235            node_id: node_id.to_string(),
236            status: status.clone(),
237            output_preview: preview_owned.clone(),
238        });
239    }
240    if let Some(session) = ctx.session_runtime.as_ref() {
241        let _ = session
242            .stream_tx()
243            .send(crate::stream::StreamFrame::FlowNodeEnd {
244                run_id: run_id.0.to_string(),
245                node_id: node_id.to_string(),
246                status,
247                output_preview: preview_owned,
248                parent_node_id: parent_node_id.map(String::from),
249            });
250    }
251}
252
253fn stmt_to_node_kind_label(stmt: &Stmt) -> (crate::nodegraph::NodeKind, String) {
254    use crate::nodegraph::NodeKind;
255    match stmt {
256        Stmt::Bind { value, .. } | Stmt::Expr(value) => expr_to_node_kind_label(value),
257        Stmt::Return { .. } => (NodeKind::Return, "return".into()),
258        Stmt::When { .. } => (
259            NodeKind::When {
260                condition_preview: "when".into(),
261            },
262            "when …".into(),
263        ),
264        Stmt::Watch(_) => (NodeKind::Return, "watch".into()),
265    }
266}
267
268fn expr_to_node_kind_label(expr: &Expr) -> (crate::nodegraph::NodeKind, String) {
269    use crate::nodegraph::NodeKind;
270    match expr {
271        Expr::Node(Node::Llm { .. }) => (NodeKind::Llm { model: None }, "llm".into()),
272        Expr::Node(Node::ToolCall { path, .. }) => {
273            let p = path
274                .iter()
275                .map(|s| s.name.clone())
276                .collect::<Vec<_>>()
277                .join(".");
278            (NodeKind::ToolCall { path: p.clone() }, format!("⟶ {p}"))
279        }
280        Expr::Node(Node::Fanout { items, collect }) => (
281            NodeKind::Fanout {
282                collect: (*collect).into(),
283            },
284            format!("fanout ×{}", items.len()),
285        ),
286        Expr::Node(Node::Subflow { name, .. }) => (
287            NodeKind::Subflow {
288                name: name.name.clone(),
289            },
290            format!("subflow({})", name.name),
291        ),
292        _ => (NodeKind::Return, "expr".into()),
293    }
294}
295
296fn collect_watches(stmts: &[Stmt]) -> HashMap<String, Vec<&WatchDecl>> {
297    let mut out: HashMap<String, Vec<&WatchDecl>> = HashMap::new();
298    for stmt in stmts {
299        if let Stmt::Watch(w) = stmt {
300            out.entry(w.target.name.clone()).or_default().push(w);
301        }
302    }
303    out
304}
305
306fn exec_stmt<'a>(
307    stmt: &'a Stmt,
308    env: &'a mut Env,
309    ctx: &'a EvalCtx<'a>,
310    watches: &'a HashMap<String, Vec<&'a WatchDecl>>,
311) -> BoxFut<'a, (StmtOutcome, Option<String>)> {
312    Box::pin(async move {
313        match stmt {
314            Stmt::Bind { name, value } => {
315                let watch_target = name.as_single_ident().map(|id| id.name.clone());
316                let v = if let Some(target) = watch_target.as_ref()
317                    && let Some(ws) = watches.get(target)
318                {
319                    match eval_bind_with_watches(value, env, ctx, ws).await {
320                        Ok(v) => v,
321                        Err(e) => return (StmtOutcome::Err(e), None),
322                    }
323                } else {
324                    eval_expr(value, env, ctx).await
325                };
326                if let Value::Err(e) = v {
327                    return (StmtOutcome::Err(e), None);
328                }
329                let preview = value_preview(&v);
330                if let Err(e) = bind_pattern(name, v, env) {
331                    return (StmtOutcome::Err(e), None);
332                }
333                (StmtOutcome::Continue, preview)
334            }
335            Stmt::When { cond, body } => {
336                let c = eval_expr(cond, env, ctx).await;
337                match c {
338                    Value::Bool(true) => (exec_stmts(body, env, ctx).await, Some("true".into())),
339                    Value::Bool(false) => (StmtOutcome::Continue, Some("false".into())),
340                    Value::Err(e) => (StmtOutcome::Err(e), None),
341                    other => (
342                        StmtOutcome::Err(RuntimeError::TypeMismatch {
343                            expected: "bool".into(),
344                            actual: other.kind_name().into(),
345                        }),
346                        None,
347                    ),
348                }
349            }
350            Stmt::Return { value } => {
351                let v = eval_expr(value, env, ctx).await;
352                if let Value::Err(e) = v {
353                    return (StmtOutcome::Err(e), None);
354                }
355                let preview = value_preview(&v);
356                (StmtOutcome::Return(v), preview)
357            }
358            Stmt::Expr(e) => {
359                let v = eval_expr(e, env, ctx).await;
360                if let Value::Err(err) = v {
361                    return (StmtOutcome::Err(err), None);
362                }
363                let preview = value_preview(&v);
364                (StmtOutcome::Continue, preview)
365            }
366            Stmt::Watch(_) => (StmtOutcome::Continue, None),
367        }
368    })
369}
370
371async fn eval_bind_with_watches(
372    expr: &Expr,
373    env: &mut Env,
374    ctx: &EvalCtx<'_>,
375    watches: &[&WatchDecl],
376) -> Result<Value, RuntimeError> {
377    let Expr::Node(Node::Llm { kwargs }) = expr else {
378        return Ok(eval_expr(expr, env, ctx).await);
379    };
380
381    let mut model: Option<String> = None;
382    let mut prompt: Option<String> = None;
383    let mut input = Value::Unit;
384    let mut cache_prompt = false;
385    let mut context_budget: Option<u64> = None;
386    for (k, v) in kwargs {
387        if k.name == "schema" || k.name == "fallback" || k.name == "retry" {
388            continue;
389        }
390        let val = eval_expr(v, env, ctx).await;
391        if val.is_err() {
392            return Ok(val);
393        }
394        match k.name.as_str() {
395            "model" => match val {
396                Value::Str(s) => model = Some(s),
397                other => {
398                    return Ok(Value::Err(RuntimeError::TypeMismatch {
399                        expected: "string".into(),
400                        actual: other.kind_name().into(),
401                    }));
402                }
403            },
404            "prompt" => match val {
405                Value::Str(s) => prompt = Some(s),
406                other => {
407                    return Ok(Value::Err(RuntimeError::TypeMismatch {
408                        expected: "string".into(),
409                        actual: other.kind_name().into(),
410                    }));
411                }
412            },
413            "input" => input = val,
414            "cache" => match val {
415                Value::Bool(b) => cache_prompt = b,
416                other => {
417                    return Ok(Value::Err(RuntimeError::TypeMismatch {
418                        expected: "bool".into(),
419                        actual: other.kind_name().into(),
420                    }));
421                }
422            },
423            "context_budget" => match val {
424                Value::Int(n) if n > 0 => context_budget = Some(n as u64),
425                other => {
426                    return Ok(Value::Err(RuntimeError::TypeMismatch {
427                        expected: "positive int".into(),
428                        actual: other.kind_name().into(),
429                    }));
430                }
431            },
432            _ => {}
433        }
434    }
435    let Some(model) = model else {
436        return Ok(Value::Err(RuntimeError::MissingArg("llm.model".into())));
437    };
438    let Some(mut prompt) = prompt else {
439        return Ok(Value::Err(RuntimeError::MissingArg("llm.prompt".into())));
440    };
441    if let Some(budget) = context_budget {
442        let (truncated, stat) = crate::eval::truncate_prompt_to_budget_tracked(prompt, budget);
443        prompt = truncated;
444        if let (Some(sink), Some(stat)) = (ctx.events, stat) {
445            sink.emit(crate::event::Event::ContextTruncated {
446                turn_id: ctx.turn_id.clone(),
447                flow_run_id: ctx.flow_run_id.clone(),
448                original_chars: stat.original_chars as u64,
449                result_chars: stat.result_chars as u64,
450                dropped_chars: stat.dropped_chars as u64,
451                budget_tokens: stat.budget_tokens,
452            });
453        }
454    }
455    let Some(provider) = ctx.providers.resolve(&model) else {
456        return Ok(Value::Err(RuntimeError::ToolFailed(format!(
457            "no provider registered for model `{model}`"
458        ))));
459    };
460
461    let rules = collect_watch_rules(watches);
462    let mut restart_count = 0u32;
463    let mut correction: Option<String> = None;
464    let mut prior_partial: Option<String> = None;
465    loop {
466        let mut messages = Vec::with_capacity(3);
467        if let Some(partial) = &prior_partial {
468            messages.push(crate::message::Message::assistant_text(
469                ctx.turn_id
470                    .clone()
471                    .unwrap_or_else(crate::event::TurnId::now),
472                format!("[partial output before user correction]\n{partial}"),
473            ));
474        }
475        if let Some(corr) = &correction {
476            messages.push(crate::provider::user_text_message(format!(
477                "<user_correction>{corr}</user_correction>\n\n{prompt}"
478            )));
479        } else {
480            messages.push(crate::provider::user_text_message(prompt.clone()));
481        }
482        let req = LlmRequest {
483            model: model.clone(),
484            messages,
485            system: None,
486            input: input.clone(),
487            schema: None,
488            cache_prompt,
489            tools: Vec::new(),
490            thinking_enabled: false,
491            stall_timeout_secs: 120,
492        };
493        let outcome = run_streaming_once(provider.as_ref(), req, &rules, ctx).await;
494        match outcome {
495            StreamOutcome::Restart {
496                level: crate::injection::InjectionLevel::L3Redirect,
497                redirect_target: Some(target),
498                ..
499            } => {
500                return Ok(Value::Err(RuntimeError::Redirect(target)));
501            }
502            StreamOutcome::Restart {
503                text: correction_text,
504                partial_output,
505                partial_tokens,
506                ..
507            } if restart_count < 3 => {
508                if let Some(sink) = ctx.events {
509                    sink.emit(crate::event::Event::LlmPartialCall {
510                        turn_id: ctx.turn_id.clone(),
511                        flow_run_id: ctx.flow_run_id.clone(),
512                        model: model.clone(),
513                        provider: provider.name().to_string(),
514                        tokens_before_abort: partial_tokens,
515                        restart_reason: "l2_course_correct".to_string(),
516                    });
517                }
518                restart_count += 1;
519                correction = Some(correction_text);
520                prior_partial = Some(partial_output);
521                continue;
522            }
523            StreamOutcome::Restart { .. } => {
524                return Ok(Value::Err(RuntimeError::ToolFailed(
525                    "l2 restart exhausted 3x, giving up".into(),
526                )));
527            }
528            StreamOutcome::Done(value) => return Ok(value),
529        }
530    }
531}
532
533enum StreamOutcome {
534    Done(Value),
535    Restart {
536        level: crate::injection::InjectionLevel,
537        text: String,
538        redirect_target: Option<String>,
539        partial_output: String,
540        partial_tokens: u64,
541    },
542}
543
544async fn run_streaming_once<'a>(
545    provider: &dyn crate::provider::Provider,
546    req: LlmRequest,
547    rules: &WatchRules,
548    ctx: &EvalCtx<'a>,
549) -> StreamOutcome {
550    let mut inj_rx = ctx
551        .session_runtime
552        .as_ref()
553        .map(|s| s.subscribe_injections());
554    let stream_tx = ctx.session_runtime.as_ref().map(|s| s.stream_tx());
555    let model_name = req.model.clone();
556    let stall_secs = req.stall_timeout_secs;
557    let obs = provider.call_streaming(req);
558    let cancel = obs.cancel.clone();
559    let mut events = obs.events;
560    let output = obs.output;
561    tokio::pin!(output);
562
563    let stall_active = stall_secs > 0;
564    let stall_dur = std::time::Duration::from_secs(stall_secs);
565    let stall_sleep = tokio::time::sleep(stall_dur);
566    tokio::pin!(stall_sleep);
567
568    let mut state = StreamMonitor::new(rules, ctx);
569    let elapsed_active = rules.elapsed_ms_gt.is_some();
570    let elapsed_deadline_ms = rules.elapsed_ms_gt.unwrap_or(u64::MAX / 2);
571    let elapsed_sleep = tokio::time::sleep(tokio::time::Duration::from_millis(
572        elapsed_deadline_ms.saturating_add(1),
573    ));
574    tokio::pin!(elapsed_sleep);
575    let started = std::time::Instant::now();
576
577    let mut l1_nudge: Option<String> = None;
578    let mut l2_correction: Option<String> = None;
579    let mut l3_redirect: Option<String> = None;
580    let mut events_closed = false;
581    let final_result = loop {
582        tokio::select! {
583            ev = async {
584                if events_closed {
585                    std::future::pending().await
586                } else {
587                    events.recv().await
588                }
589            }, if !events_closed => {
590                match ev {
591                    Ok(NodeEvent::LlmChunk { text, cumulative_tokens }) => {
592                        if let Some(session) = ctx.session_runtime.as_ref() {
593                            session.mark_streamed();
594                        }
595                        if let Some(tx) = &stream_tx {
596                            let _ = tx.send(crate::stream::StreamFrame::LlmChunk {
597                                text: text.clone(),
598                                model: model_name.clone(),
599                            });
600                        }
601                        state.on_chunk(&text, cumulative_tokens, started, rules, &cancel);
602                        if stall_active {
603                            stall_sleep
604                                .as_mut()
605                                .reset(tokio::time::Instant::now() + stall_dur);
606                        }
607                    }
608                    Ok(NodeEvent::ThinkingChunk { text }) => {
609                        if let Some(tx) = &stream_tx {
610                            let _ = tx.send(crate::stream::StreamFrame::ThinkingChunk {
611                                text,
612                            });
613                        }
614                    }
615                    Ok(NodeEvent::LlmDone { total_tokens }) => {
616                        if let Some(tx) = &stream_tx {
617                            let _ = tx.send(crate::stream::StreamFrame::LlmDone { total_tokens });
618                        }
619                        state.on_done(total_tokens, started, rules, &cancel);
620                    }
621                    Ok(_) => {}
622                    Err(tokio::sync::broadcast::error::RecvError::Closed) => {
623                        events_closed = true;
624                    }
625                    Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {}
626                }
627            }
628            inj_msg = poll_injection(&mut inj_rx), if inj_rx.is_some() && l3_redirect.is_none() => {
629                if let Some(inj) = inj_msg
630                    && ctx.turn_id.as_ref().is_some_and(|t| &inj.turn_id == t)
631                {
632                    match inj.level {
633                        crate::injection::InjectionLevel::L1Nudge => {
634                            cancel.cancel();
635                            l1_nudge = Some(inj.text.clone());
636                        }
637                        crate::injection::InjectionLevel::L2CourseCorrect => {
638                            cancel.cancel();
639                            l2_correction = Some(inj.text.clone());
640                        }
641                        crate::injection::InjectionLevel::L3Redirect => {
642                            cancel.cancel();
643                            l3_redirect = inj.redirect_target.clone();
644                        }
645                        crate::injection::InjectionLevel::L4HardStop => {
646                            cancel.cancel();
647                            break Err(RuntimeError::Cancelled("hard stop from user".into()));
648                        }
649                    }
650                }
651            }
652            _ = &mut elapsed_sleep, if elapsed_active && state.abort_reason.is_none() => {
653                state.abort_reason = Some(format!("elapsed > {elapsed_deadline_ms}ms"));
654                cancel.cancel();
655                break Err(RuntimeError::Cancelled("elapsed".into()));
656            }
657            _ = &mut stall_sleep, if stall_active => {
658                cancel.cancel();
659                break Err(RuntimeError::ToolFailed(format!(
660                    "llm stall timeout after {}s",
661                    stall_secs
662                )));
663            }
664            result = &mut output => break result,
665        }
666    };
667
668    while let Ok(ev) = events.try_recv() {
669        match ev {
670            NodeEvent::LlmChunk {
671                text,
672                cumulative_tokens,
673            } => {
674                if let Some(session) = ctx.session_runtime.as_ref() {
675                    session.mark_streamed();
676                }
677                if let Some(tx) = &stream_tx {
678                    let _ = tx.send(crate::stream::StreamFrame::LlmChunk {
679                        text: text.clone(),
680                        model: model_name.clone(),
681                    });
682                }
683                state.on_chunk(&text, cumulative_tokens, started, rules, &cancel);
684            }
685            NodeEvent::ThinkingChunk { text } => {
686                if let Some(tx) = &stream_tx {
687                    let _ = tx.send(crate::stream::StreamFrame::ThinkingChunk { text });
688                }
689            }
690            NodeEvent::LlmDone { total_tokens } => {
691                if let Some(tx) = &stream_tx {
692                    let _ = tx.send(crate::stream::StreamFrame::LlmDone { total_tokens });
693                }
694                state.on_done(total_tokens, started, rules, &cancel);
695            }
696            _ => {}
697        }
698    }
699
700    if let Some(redirect) = l3_redirect {
701        return StreamOutcome::Restart {
702            level: crate::injection::InjectionLevel::L3Redirect,
703            text: String::new(),
704            redirect_target: Some(redirect),
705            partial_output: state.text_captured.clone(),
706            partial_tokens: state.tokens_seen,
707        };
708    }
709    if let Some(correction) = l2_correction {
710        return StreamOutcome::Restart {
711            level: crate::injection::InjectionLevel::L2CourseCorrect,
712            text: correction,
713            redirect_target: None,
714            partial_output: state.text_captured.clone(),
715            partial_tokens: state.tokens_seen,
716        };
717    }
718    if let Some(nudge) = l1_nudge {
719        return StreamOutcome::Restart {
720            level: crate::injection::InjectionLevel::L1Nudge,
721            text: nudge,
722            redirect_target: None,
723            partial_output: state.text_captured.clone(),
724            partial_tokens: state.tokens_seen,
725        };
726    }
727
728    let mut abort_reason = state.abort_reason;
729    match final_result {
730        _ if abort_reason.is_some() => StreamOutcome::Done(Value::Err(RuntimeError::Aborted(
731            abort_reason.take().unwrap_or_default(),
732        ))),
733        Ok(am) => StreamOutcome::Done(crate::provider::assistant_message_to_value(&am)),
734        Err(e) => StreamOutcome::Done(Value::Err(e)),
735    }
736}
737
738async fn poll_injection(
739    rx: &mut Option<tokio::sync::broadcast::Receiver<crate::injection::Injection>>,
740) -> Option<crate::injection::Injection> {
741    let rx = rx.as_mut()?;
742    loop {
743        match rx.recv().await {
744            Ok(inj) => return Some(inj),
745            Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue,
746            Err(_) => return None,
747        }
748    }
749}
750
751#[derive(Default)]
752struct WatchRules {
753    token_matches: Vec<(String, String)>,
754    tokens_gt: Option<u64>,
755    elapsed_ms_gt: Option<u64>,
756    warn_token: Vec<WarnRule>,
757    warn_tokens_gt: Vec<(u64, WarnRule)>,
758    warn_elapsed_ms_gt: Vec<(u64, WarnRule)>,
759}
760
761#[derive(Clone)]
762struct WarnRule {
763    target: String,
764    message: String,
765    pattern: String,
766}
767
768fn render_warn_msg(msg: &Option<Expr>, fallback: &str) -> String {
769    match msg {
770        Some(Expr::Literal(atman_dsl::ast::Literal::Str(s))) => s.clone(),
771        _ => fallback.to_string(),
772    }
773}
774
775struct StreamMonitor<'a> {
776    window: String,
777    text_captured: String,
778    tokens_seen: u64,
779    abort_reason: Option<String>,
780    fired_warn_token: std::collections::HashSet<String>,
781    fired_warn_tokens: std::collections::HashSet<u64>,
782    fired_warn_elapsed: std::collections::HashSet<u64>,
783    ctx: &'a EvalCtx<'a>,
784}
785
786impl<'a> StreamMonitor<'a> {
787    fn new(_rules: &WatchRules, ctx: &'a EvalCtx<'a>) -> Self {
788        Self {
789            window: String::new(),
790            text_captured: String::new(),
791            tokens_seen: 0,
792            abort_reason: None,
793            fired_warn_token: Default::default(),
794            fired_warn_tokens: Default::default(),
795            fired_warn_elapsed: Default::default(),
796            ctx,
797        }
798    }
799
800    fn push_window(&mut self, text: &str) {
801        self.window.push_str(text);
802        if self.window.len() > 512 {
803            let drop = self.window.len() - 512;
804            self.window.drain(..drop);
805        }
806        self.text_captured.push_str(text);
807    }
808
809    fn emit_warn(&self, rule: &WarnRule, trigger: &str) {
810        if let Some(sink) = self.ctx.events {
811            sink.emit(crate::event::Event::WatchWarn {
812                turn_id: self.ctx.turn_id.clone(),
813                flow_run_id: self.ctx.flow_run_id.clone(),
814                target: rule.target.clone(),
815                trigger: trigger.to_string(),
816                message: rule.message.clone(),
817            });
818        }
819    }
820
821    fn check_token_warns(&mut self, rules: &WatchRules) {
822        for rule in &rules.warn_token {
823            if !self.fired_warn_token.contains(&rule.pattern)
824                && self.window.contains(rule.pattern.as_str())
825            {
826                self.fired_warn_token.insert(rule.pattern.clone());
827                self.emit_warn(rule, &format!("token({})", rule.pattern));
828            }
829        }
830    }
831
832    fn check_tokens_consumed_warns(&mut self, rules: &WatchRules) {
833        for (threshold, rule) in &rules.warn_tokens_gt {
834            if !self.fired_warn_tokens.contains(threshold) && self.tokens_seen > *threshold {
835                self.fired_warn_tokens.insert(*threshold);
836                self.emit_warn(rule, &format!("tokens_consumed>{threshold}"));
837            }
838        }
839    }
840
841    fn check_elapsed_warns(&mut self, rules: &WatchRules, started: std::time::Instant) {
842        let elapsed = started.elapsed().as_millis() as u64;
843        for (threshold, rule) in &rules.warn_elapsed_ms_gt {
844            if !self.fired_warn_elapsed.contains(threshold) && elapsed > *threshold {
845                self.fired_warn_elapsed.insert(*threshold);
846                self.emit_warn(rule, &format!("elapsed>{threshold}ms"));
847            }
848        }
849    }
850
851    fn on_chunk(
852        &mut self,
853        text: &str,
854        cumulative_tokens: u64,
855        started: std::time::Instant,
856        rules: &WatchRules,
857        cancel: &tokio_util::sync::CancellationToken,
858    ) {
859        self.tokens_seen = cumulative_tokens.max(self.tokens_seen);
860        self.push_window(text);
861        if self.abort_reason.is_none() {
862            for (pat, reason) in &rules.token_matches {
863                if self.window.contains(pat.as_str()) {
864                    self.abort_reason = Some(reason.clone());
865                    cancel.cancel();
866                    break;
867                }
868            }
869        }
870        if self.abort_reason.is_none()
871            && let Some(limit) = rules.tokens_gt
872            && self.tokens_seen > limit
873        {
874            self.abort_reason = Some(format!("tokens_consumed > {limit}"));
875            cancel.cancel();
876        }
877        self.check_token_warns(rules);
878        self.check_tokens_consumed_warns(rules);
879        self.check_elapsed_warns(rules, started);
880    }
881
882    fn on_done(
883        &mut self,
884        total_tokens: u64,
885        started: std::time::Instant,
886        rules: &WatchRules,
887        _cancel: &tokio_util::sync::CancellationToken,
888    ) {
889        self.tokens_seen = total_tokens.max(self.tokens_seen);
890        if self.abort_reason.is_none()
891            && let Some(limit) = rules.tokens_gt
892            && self.tokens_seen > limit
893        {
894            self.abort_reason = Some(format!("tokens_consumed > {limit}"));
895        }
896        self.check_tokens_consumed_warns(rules);
897        self.check_elapsed_warns(rules, started);
898    }
899}
900
901fn collect_watch_rules(watches: &[&WatchDecl]) -> WatchRules {
902    let mut rules = WatchRules::default();
903    for w in watches {
904        for on in &w.on_blocks {
905            let has_abort = on
906                .actions
907                .iter()
908                .any(|a| matches!(a, WatchAction::Abort { .. }));
909            let warn_msg_expr = on.actions.iter().find_map(|a| match a {
910                WatchAction::Warn { msg } => Some(msg),
911                _ => None,
912            });
913            if !has_abort && warn_msg_expr.is_none() {
914                continue;
915            }
916            match &on.event {
917                WatchEvent::Token { patterns } => {
918                    for p in patterns {
919                        if has_abort {
920                            rules
921                                .token_matches
922                                .push((p.clone(), format!("token match: {p}")));
923                        }
924                        if let Some(msg_expr) = warn_msg_expr {
925                            rules.warn_token.push(WarnRule {
926                                target: w.target.name.clone(),
927                                message: render_warn_msg(
928                                    msg_expr,
929                                    &format!("watch warn: token `{p}`"),
930                                ),
931                                pattern: p.clone(),
932                            });
933                        }
934                    }
935                }
936                WatchEvent::TokensConsumed { cmp, value }
937                    if matches!(cmp, CmpOp::Gt | CmpOp::Ge) =>
938                {
939                    let threshold = if matches!(cmp, CmpOp::Ge) {
940                        value.saturating_sub(1)
941                    } else {
942                        *value
943                    };
944                    if has_abort {
945                        rules.tokens_gt = Some(match rules.tokens_gt {
946                            Some(existing) => existing.min(threshold),
947                            None => threshold,
948                        });
949                    }
950                    if let Some(msg_expr) = warn_msg_expr {
951                        rules.warn_tokens_gt.push((
952                            threshold,
953                            WarnRule {
954                                target: w.target.name.clone(),
955                                message: render_warn_msg(
956                                    msg_expr,
957                                    &format!("watch warn: tokens_consumed > {threshold}"),
958                                ),
959                                pattern: format!("tokens_consumed>{threshold}"),
960                            },
961                        ));
962                    }
963                }
964                WatchEvent::Elapsed { cmp, duration_ms }
965                    if matches!(cmp, CmpOp::Gt | CmpOp::Ge) =>
966                {
967                    let threshold = if matches!(cmp, CmpOp::Ge) {
968                        duration_ms.saturating_sub(1)
969                    } else {
970                        *duration_ms
971                    };
972                    if has_abort {
973                        rules.elapsed_ms_gt = Some(match rules.elapsed_ms_gt {
974                            Some(existing) => existing.min(threshold),
975                            None => threshold,
976                        });
977                    }
978                    if let Some(msg_expr) = warn_msg_expr {
979                        rules.warn_elapsed_ms_gt.push((
980                            threshold,
981                            WarnRule {
982                                target: w.target.name.clone(),
983                                message: render_warn_msg(
984                                    msg_expr,
985                                    &format!("watch warn: elapsed > {threshold}ms"),
986                                ),
987                                pattern: format!("elapsed>{threshold}ms"),
988                            },
989                        ));
990                    }
991                }
992                _ => {}
993            }
994        }
995    }
996    rules
997}
998
999pub async fn exec_flow(
1000    flow: &FlowDecl,
1001    args: Vec<(String, Value)>,
1002    tools: &ToolRegistry,
1003    tool_ctx: &ToolCtx,
1004    providers: &crate::provider::ProviderRegistry,
1005    source_dir: Option<PathBuf>,
1006) -> Result<Value, RuntimeError> {
1007    let flows = std::collections::HashMap::new();
1008    exec_flow_with_siblings(
1009        flow,
1010        args,
1011        tools,
1012        tool_ctx,
1013        providers,
1014        &flows,
1015        None,
1016        None,
1017        None,
1018        None,
1019        tokio_util::sync::CancellationToken::new(),
1020        None,
1021        source_dir,
1022    )
1023    .await
1024}
1025
1026#[allow(clippy::too_many_arguments)]
1027pub async fn exec_flow_with_siblings(
1028    flow: &FlowDecl,
1029    args: Vec<(String, Value)>,
1030    tools: &ToolRegistry,
1031    tool_ctx: &ToolCtx,
1032    providers: &crate::provider::ProviderRegistry,
1033    flows: &std::collections::HashMap<String, FlowDecl>,
1034    events: Option<&crate::event::EventSink>,
1035    turn_id: Option<crate::event::TurnId>,
1036    flow_run_id: Option<crate::event::FlowRunId>,
1037    session: Option<std::sync::Arc<crate::session::Session>>,
1038    flow_cancel: tokio_util::sync::CancellationToken,
1039    safety: Option<&crate::safety::SafetyConfig>,
1040    source_dir: Option<PathBuf>,
1041) -> Result<Value, RuntimeError> {
1042    let mut env = Env::new();
1043    for (name, value) in args {
1044        env.bind(name, value);
1045    }
1046    let ctx = EvalCtx {
1047        tools,
1048        tool_ctx,
1049        providers,
1050        flows,
1051        contract: flow.contract.as_ref(),
1052        events,
1053        turn_id,
1054        flow_run_id,
1055        session_runtime: session,
1056        flow_cancel,
1057        safety,
1058        current_node_id: None,
1059        source_dir,
1060    };
1061    match exec_stmts(&flow.body, &mut env, &ctx).await {
1062        StmtOutcome::Return(v) => Ok(v),
1063        StmtOutcome::Err(e) => Err(e),
1064        StmtOutcome::Continue => Ok(Value::Unit),
1065    }
1066}
1067
1068#[cfg(test)]
1069mod tests {
1070    use super::*;
1071    use atman_dsl::parse::parse_file;
1072
1073    async fn run(src: &str, args: Vec<(String, Value)>) -> Result<Value, RuntimeError> {
1074        let file = parse_file(src).expect("parse test src");
1075        let tools = ToolRegistry::new();
1076        let tool_ctx = ToolCtx::new();
1077        let providers = crate::provider::ProviderRegistry::new();
1078        exec_flow(&file.flows[0], args, &tools, &tool_ctx, &providers, None).await
1079    }
1080
1081    #[tokio::test]
1082    async fn bind_and_return() {
1083        let out = run(
1084            r#"flow t() -> Int {
1085    x = 1
1086    y = x + 2
1087    return y
1088}
1089"#,
1090            vec![],
1091        )
1092        .await
1093        .unwrap();
1094        assert!(matches!(out, Value::Int(3)));
1095    }
1096
1097    #[tokio::test]
1098    async fn when_true_executes_body() {
1099        let out = run(
1100            r#"flow t() -> Int {
1101    x = 5
1102    when x > 3 {
1103        return 42
1104    }
1105    return 0
1106}
1107"#,
1108            vec![],
1109        )
1110        .await
1111        .unwrap();
1112        assert!(matches!(out, Value::Int(42)));
1113    }
1114
1115    #[tokio::test]
1116    async fn when_false_skips_body() {
1117        let out = run(
1118            r#"flow t() -> Int {
1119    x = 1
1120    when x > 3 {
1121        return 42
1122    }
1123    return 0
1124}
1125"#,
1126            vec![],
1127        )
1128        .await
1129        .unwrap();
1130        assert!(matches!(out, Value::Int(0)));
1131    }
1132
1133    #[tokio::test]
1134    async fn err_in_bind_stops_flow() {
1135        let err = run(
1136            r#"flow t() -> Int {
1137    x = missing
1138    return 1
1139}
1140"#,
1141            vec![],
1142        )
1143        .await
1144        .unwrap_err();
1145        assert!(matches!(err, RuntimeError::UndefinedVar(n) if n == "missing"));
1146    }
1147
1148    #[tokio::test]
1149    async fn flow_args_bind_before_body() {
1150        let out = run(
1151            r#"flow t() -> Int {
1152    return n + 1
1153}
1154"#,
1155            vec![("n".into(), Value::Int(4))],
1156        )
1157        .await
1158        .unwrap();
1159        assert!(matches!(out, Value::Int(5)));
1160    }
1161
1162    #[tokio::test]
1163    async fn when_cond_non_bool_is_type_error() {
1164        let err = run(
1165            r#"flow t() -> Int {
1166    when 1 {
1167        return 1
1168    }
1169    return 0
1170}
1171"#,
1172            vec![],
1173        )
1174        .await
1175        .unwrap_err();
1176        assert!(matches!(err, RuntimeError::TypeMismatch { .. }));
1177    }
1178
1179    #[tokio::test]
1180    async fn flow_falls_through_to_unit_without_return() {
1181        let out = run(
1182            r#"flow t() {
1183    x = 1
1184}
1185"#,
1186            vec![],
1187        )
1188        .await
1189        .unwrap();
1190        assert!(matches!(out, Value::Unit));
1191    }
1192}