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 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}