use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use futures_core::Stream;
use serde_json::{json, Value};
use tracing::{debug, warn};
use uuid::Uuid;
use crate::backends::dispatch::{dispatch_post_turn, dispatch_tool_call, gate_pre_turn};
use crate::builtins::FINISH_TOOL_NAME;
use crate::backends::loop_util::extract_canonical_path;
use crate::backends::state::LoopState;
use crate::backends::stream_timeout::{idle_timeout_ms, next_with_idle_timeout, NextChunk};
use crate::content::Content;
use crate::error::{Error, Result};
use crate::hooks::{HookRunner, SessionContext};
use crate::tools::ToolRunner;
use crate::types::{
Step, StepStatus, StreamChunk, ToolCall as NeutralToolCall, ToolResult, UsageMetadata,
};
pub(crate) const MAX_TOOL_ROUNDS: u32 = 16;
pub(crate) struct ResolvedCall {
pub id: Option<String>,
pub name: String,
pub args: Value,
pub parse_error: Option<String>,
}
pub(crate) struct DispatchedResult {
pub call: ResolvedCall,
pub value: Value,
#[allow(dead_code)]
pub is_error: bool,
}
#[allow(dead_code)]
pub(crate) enum StreamEnd {
Proceed,
Resume,
ProceedAndEndTurn,
}
pub(crate) struct EmitCtx<'a, M> {
state: &'a LoopState<M>,
trajectory_id: &'a str,
step_index: u32,
text: String,
}
impl<M> EmitCtx<'_, M> {
pub fn push_text(&mut self, t: &str) {
if !t.is_empty() {
self.text.push_str(t);
self.state
.emit(Step::text_delta(self.trajectory_id, self.step_index, t));
}
}
pub fn push_thought(&mut self, t: &str) {
if !t.is_empty() {
self.state
.emit(Step::thought_delta(self.trajectory_id, self.step_index, t));
}
}
}
pub(crate) trait TurnProvider {
type Message: Clone;
type Config;
type Request: Clone;
type Event;
type Accum: Default;
fn build_request(config: &Self::Config, history: &[Self::Message]) -> Self::Request;
fn compaction_threshold(config: &Self::Config) -> Option<u32>;
fn fold_event(
acc: &mut Self::Accum,
ctx: &mut EmitCtx<'_, Self::Message>,
event: Self::Event,
) -> Result<()>;
fn resolve_pending_calls(acc: &mut Self::Accum) -> Vec<ResolvedCall>;
fn round_usage(acc: &Self::Accum) -> UsageMetadata;
fn map_finish_reason(acc: &Self::Accum) -> (StepStatus, &'static str);
fn assemble_assistant_message(
acc: Self::Accum,
text: &str,
calls: &[ResolvedCall],
) -> Option<Self::Message>;
fn tool_result_messages(results: Vec<DispatchedResult>) -> Vec<Self::Message>;
fn on_stream_end(_acc: &mut Self::Accum, _pause_resumes: u32) -> StreamEnd {
StreamEnd::Proceed
}
fn on_cancel_with_pending_calls(_calls: &[ResolvedCall]) -> Vec<Self::Message> {
Vec::new()
}
}
pub(crate) struct EngineDeps<P: TurnProvider> {
pub config: P::Config,
pub state: Arc<LoopState<P::Message>>,
pub tool_runner: Option<Arc<ToolRunner>>,
pub hook_runner: Option<Arc<HookRunner>>,
pub session_ctx: Option<SessionContext>,
}
fn turn_fail<M>(state: &LoopState<M>, e: Error) -> Error {
state.emit_error(e.to_string());
state.idle.store(true, Ordering::Release);
state.idle_notify.notify_waiters();
e
}
pub(crate) async fn run_turn<P, St, Open, OFut, Compact, CFut>(
deps: EngineDeps<P>,
user: P::Message,
prompt: Content,
open: Open,
compact: Compact,
) -> Result<()>
where
P: TurnProvider,
St: Stream<Item = Result<P::Event>> + Unpin,
Open: Fn(P::Request) -> OFut,
OFut: Future<Output = Result<St>>,
Compact: FnOnce() -> CFut,
CFut: Future<Output = ()>,
{
deps.state.idle.store(false, Ordering::Release);
deps.state.cancel.store(false, Ordering::Release);
let turn_ctx = deps
.session_ctx
.as_ref()
.map(|s| s.child())
.unwrap_or_default();
if let Some(denied) = gate_pre_turn(deps.hook_runner.as_ref(), &turn_ctx, &prompt).await {
return Err(turn_fail(&deps.state, Error::other(denied)));
}
deps.state.history.lock().push(user);
*deps.state.last_turn_usage.lock() = Some(UsageMetadata::default());
*deps.state.last_structured_output.lock() = None;
let mut rounds = 0u32;
let mut last_text = String::new();
let mut last_status: (StepStatus, &'static str) = (StepStatus::Done, "");
let mut finished_turn = false;
let mut finish_summary: Option<String> = None;
let trajectory_id = Uuid::new_v4().to_string();
loop {
rounds += 1;
if rounds > MAX_TOOL_ROUNDS {
warn!(rounds, "exceeded MAX_TOOL_ROUNDS; forcing turn end");
break;
}
if deps.state.cancel.load(Ordering::Acquire) {
debug!("turn cancelled before model call");
break;
}
let step_index = deps.state.alloc_step_index();
let mut acc = P::Accum::default();
let mut ctx = EmitCtx {
state: &deps.state,
trajectory_id: &trajectory_id,
step_index,
text: String::new(),
};
let mut pause_resumes = 0u32;
let end_after_persist = 'request: loop {
let request = P::build_request(&deps.config, &deps.state.history.lock());
let mut stream = match crate::backends::retry::open_stream_with_retry(|| {
open(request.clone())
})
.await
{
Ok(s) => s,
Err(e) => return Err(turn_fail(&deps.state, e)),
};
let idle_ms = idle_timeout_ms();
loop {
let ev_res = match next_with_idle_timeout(&mut stream, idle_ms).await {
NextChunk::Item(item) => item,
NextChunk::End => break,
NextChunk::IdleTimeout => {
let e = Error::other(format!(
"model stream stalled — no data for {}s",
idle_ms / 1000
));
return Err(turn_fail(&deps.state, e));
}
};
if deps.state.cancel.load(Ordering::Acquire) {
break;
}
let ev = match ev_res {
Ok(c) => c,
Err(e) => return Err(turn_fail(&deps.state, e)),
};
if let Err(e) = P::fold_event(&mut acc, &mut ctx, ev) {
return Err(turn_fail(&deps.state, e));
}
}
let cancelled = deps.state.cancel.load(Ordering::Acquire);
match P::on_stream_end(&mut acc, pause_resumes) {
StreamEnd::Resume if !cancelled => {
pause_resumes += 1;
debug!(pause_resumes, "provider resumed the stream");
continue 'request;
}
StreamEnd::Resume | StreamEnd::ProceedAndEndTurn => break 'request true,
StreamEnd::Proceed => break 'request false,
}
};
let pending_calls = P::resolve_pending_calls(&mut acc);
last_status = P::map_finish_reason(&acc);
let usage = P::round_usage(&acc);
if let Some(msg) = P::assemble_assistant_message(acc, &ctx.text, &pending_calls) {
deps.state.history.lock().push(msg);
}
if usage != UsageMetadata::default() {
let mut slot = deps.state.last_turn_usage.lock();
match slot.as_mut() {
Some(a) => a.merge_round(&usage),
None => *slot = Some(usage),
}
}
last_text = ctx.text;
if pending_calls.is_empty() || end_after_persist {
break;
}
if deps.state.cancel.load(Ordering::Acquire) {
debug!("turn cancelled before tool dispatch");
let balance = P::on_cancel_with_pending_calls(&pending_calls);
deps.state.history.lock().extend(balance);
break;
}
let mut results: Vec<DispatchedResult> = Vec::with_capacity(pending_calls.len());
let mut saw_finish = false;
for call in pending_calls {
if let Some(msg) = call.parse_error.clone() {
let post_result = ToolResult {
name: call.name.clone(),
id: call.id.clone(),
result: Some(json!({ "error": msg.clone() })),
error: Some(msg.clone()),
};
deps.state
.emit_chunk_step(StreamChunk::ToolResult(post_result));
results.push(DispatchedResult {
call,
value: json!({ "error": msg }),
is_error: true,
});
continue;
}
if call.name == FINISH_TOOL_NAME {
if let Some(out) = call.args.get("output").cloned() {
*deps.state.last_structured_output.lock() = Some(out);
}
if let Some(sm) = call.args.get("summary").and_then(|v| v.as_str()) {
if !sm.is_empty() {
finish_summary = Some(sm.to_string());
}
}
saw_finish = true;
results.push(DispatchedResult {
call,
value: json!({ "ok": true }),
is_error: false,
});
continue;
}
let tool_call = NeutralToolCall {
name: call.name.clone(),
args: call.args.clone(),
id: call.id.clone(),
canonical_path: extract_canonical_path(&call.args),
};
deps.state
.emit_chunk_step(StreamChunk::ToolCall(tool_call.clone()));
let post_result = dispatch_tool_call(
deps.tool_runner.as_ref(),
deps.hook_runner.as_ref(),
&turn_ctx,
&tool_call,
)
.await;
let value = post_result.result.clone().unwrap_or(Value::Null);
let is_error = post_result.error.is_some();
deps.state
.emit_chunk_step(StreamChunk::ToolResult(post_result));
results.push(DispatchedResult {
call,
value,
is_error,
});
}
deps.state
.history
.lock()
.extend(P::tool_result_messages(results));
if saw_finish {
finished_turn = true;
break;
}
}
let usage = deps.state.last_turn_usage.lock().clone().unwrap_or_default();
let usage_opt = if usage == UsageMetadata::default() {
None
} else {
Some(usage.clone())
};
let (status, error_msg) = last_status;
let structured = deps.state.last_structured_output.lock().clone();
let terminal = Step::turn_complete(
trajectory_id,
deps.state.alloc_step_index(),
status,
last_text.as_str(),
error_msg,
finished_turn,
structured,
usage_opt,
)
.with_finish_summary(finish_summary);
deps.state.emit(terminal);
dispatch_post_turn(deps.hook_runner.as_ref(), &turn_ctx, &last_text).await;
let used = usage.prompt_token_count;
if crate::backends::compaction::should_compact(
used,
P::compaction_threshold(&deps.config),
) {
debug!(used, "compaction triggered");
compact().await;
}
deps.state.idle.store(true, Ordering::Release);
deps.state.idle_notify.notify_waiters();
debug!(rounds, "turn complete");
Ok(())
}
#[cfg(test)]
pub(crate) fn test_fold_events<P: TurnProvider>(
state: &LoopState<P::Message>,
acc: &mut P::Accum,
events: Vec<P::Event>,
) {
let mut ctx = EmitCtx {
state,
trajectory_id: "test",
step_index: 0,
text: String::new(),
};
for ev in events {
P::fold_event(acc, &mut ctx, ev).expect("fold_event ok");
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hooks::TurnContext;
use crate::types::{HookResult, StepSource, StepType};
use parking_lot::Mutex;
use std::collections::VecDeque;
use tokio::sync::broadcast;
#[derive(Clone)]
enum Ev {
Text(&'static str),
Call {
id: &'static str,
name: &'static str,
args: &'static str,
},
Resume,
Cancel,
}
#[derive(Default)]
struct Accum {
calls: Vec<(String, String, String)>,
resume: bool,
}
struct MockProvider;
impl TurnProvider for MockProvider {
type Message = String;
type Config = ();
type Request = usize;
type Event = Ev;
type Accum = Accum;
fn build_request(_c: &(), history: &[String]) -> usize {
history.len()
}
fn compaction_threshold(_c: &()) -> Option<u32> {
None
}
fn fold_event(acc: &mut Accum, ctx: &mut EmitCtx<'_, String>, ev: Ev) -> Result<()> {
match ev {
Ev::Text(t) => ctx.push_text(t),
Ev::Call { id, name, args } => {
acc.calls.push((id.into(), name.into(), args.into()))
}
Ev::Resume => acc.resume = true,
Ev::Cancel => ctx.state.cancel.store(true, Ordering::Release),
}
Ok(())
}
fn resolve_pending_calls(acc: &mut Accum) -> Vec<ResolvedCall> {
std::mem::take(&mut acc.calls)
.into_iter()
.map(|(id, name, args)| {
let (args, parse_error) =
crate::backends::loop_util::resolve_tool_args(&name, &args);
ResolvedCall {
id: Some(id),
name,
args,
parse_error,
}
})
.collect()
}
fn round_usage(_acc: &Accum) -> UsageMetadata {
UsageMetadata::default()
}
fn map_finish_reason(_acc: &Accum) -> (StepStatus, &'static str) {
(StepStatus::Done, "")
}
fn assemble_assistant_message(
_acc: Accum,
text: &str,
calls: &[ResolvedCall],
) -> Option<String> {
(!text.is_empty() || !calls.is_empty())
.then(|| format!("assistant:{text}:{}", calls.len()))
}
fn tool_result_messages(results: Vec<DispatchedResult>) -> Vec<String> {
results
.into_iter()
.map(|r| format!("tool:{}:{}", r.call.id.unwrap_or_default(), r.value))
.collect()
}
fn on_stream_end(acc: &mut Accum, pause_resumes: u32) -> StreamEnd {
if std::mem::take(&mut acc.resume) {
if pause_resumes < 2 {
StreamEnd::Resume
} else {
StreamEnd::ProceedAndEndTurn
}
} else {
StreamEnd::Proceed
}
}
fn on_cancel_with_pending_calls(calls: &[ResolvedCall]) -> Vec<String> {
calls
.iter()
.map(|c| format!("cancelled:{}", c.id.as_deref().unwrap_or_default()))
.collect()
}
}
type Steps = broadcast::Receiver<Step>;
async fn run(
streams: Vec<Vec<Ev>>,
hook_runner: Option<Arc<HookRunner>>,
) -> (Arc<LoopState<String>>, Steps, u32) {
let (tx, rx) = broadcast::channel::<Step>(64);
let state = Arc::new(LoopState::new(tx));
let deps = EngineDeps::<MockProvider> {
config: (),
state: state.clone(),
tool_runner: None,
hook_runner,
session_ctx: None,
};
let script = Mutex::new(streams.into_iter().collect::<VecDeque<_>>());
let opens = std::sync::atomic::AtomicU32::new(0);
let prompt = Content::text("hi");
let res = run_turn::<MockProvider, _, _, _, _, _>(
deps,
"user:hi".to_string(),
prompt,
|_req| {
opens.fetch_add(1, Ordering::SeqCst);
let evs = script.lock().pop_front().unwrap_or_default();
async move {
Ok(futures_util::stream::iter(
evs.into_iter().map(Ok::<_, Error>),
))
}
},
|| async {},
)
.await;
let _ = res;
(state, rx, opens.load(Ordering::SeqCst))
}
struct DenyAllTurns;
#[async_trait::async_trait]
impl crate::hooks::PreTurnHook for DenyAllTurns {
fn name(&self) -> &str {
"test::deny_all_turns"
}
async fn run(&self, _ctx: &TurnContext, _prompt: &Content) -> Result<HookResult> {
Ok(HookResult::deny("nope"))
}
}
#[tokio::test]
async fn pre_turn_deny_keeps_prompt_out_of_history() {
let hooks = Arc::new(HookRunner::new());
hooks.register_pre_turn(Arc::new(DenyAllTurns));
let (state, mut rx, opens) = run(vec![vec![Ev::Text("never")]], Some(hooks)).await;
assert!(state.history.lock().is_empty(), "denied prompt must not enter history");
assert_eq!(opens, 0, "the model must never be called on deny");
assert!(state.idle.load(Ordering::Acquire), "idle guard must release");
let step = rx.recv().await.expect("a step was broadcast");
assert_eq!(step.source, StepSource::System);
assert_eq!(step.status, StepStatus::Error);
assert!(step.error.contains("turn denied by hook: nope"));
}
#[tokio::test]
async fn finish_tool_ends_turn_with_summary_and_structured_output() {
let (state, mut rx, opens) = run(
vec![vec![
Ev::Text("working"),
Ev::Call {
id: "c1",
name: FINISH_TOOL_NAME,
args: r#"{"summary":"all done","output":{"x":1}}"#,
},
]],
None,
)
.await;
assert_eq!(opens, 1);
let hist = state.history.lock().clone();
assert_eq!(
hist,
vec![
"user:hi".to_string(),
"assistant:working:1".to_string(),
"tool:c1:{\"ok\":true}".to_string(),
],
"user + assistant + finish result persisted in order"
);
let mut terminal = None;
while let Ok(s) = rx.try_recv() {
if s.is_complete_response == Some(true) {
terminal = Some(s);
}
}
let t = terminal.expect("terminal step emitted");
assert_eq!(t.kind, StepType::Finish);
assert_eq!(t.finish_summary.as_deref(), Some("all done"));
assert_eq!(t.structured_output, Some(json!({"x": 1})));
assert_eq!(t.content, "working");
assert!(state.idle.load(Ordering::Acquire));
}
#[tokio::test]
async fn stream_end_resume_reopens_then_end_turn_skips_dispatch() {
let (state, _rx, opens) = run(
vec![
vec![Ev::Text("a"), Ev::Resume],
vec![Ev::Text("b"), Ev::Resume],
vec![
Ev::Call { id: "c9", name: "view_file", args: "{}" },
Ev::Resume, ],
],
None,
)
.await;
assert_eq!(opens, 3, "one open + two resumes");
let hist = state.history.lock().clone();
assert_eq!(
hist,
vec!["user:hi".to_string(), "assistant:ab:1".to_string()],
"text accumulates across resumes; the pending call persists but is NOT dispatched"
);
}
#[tokio::test]
async fn cancel_with_pending_calls_appends_the_providers_balance() {
let (state, _rx, _opens) = run(
vec![vec![
Ev::Call { id: "c1", name: "view_file", args: r#"{"path":"a.rs"}"# },
Ev::Cancel,
]],
None,
)
.await;
let hist = state.history.lock().clone();
assert_eq!(
hist,
vec![
"user:hi".to_string(),
"assistant::1".to_string(),
"cancelled:c1".to_string(),
],
"the pending call is balanced, never dispatched"
);
assert!(state.idle.load(Ordering::Acquire));
}
}