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