Skip to main content

atman_runtime/
exec.rs

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