use std::collections::HashMap;
use std::sync::Arc;
use futures::StreamExt;
use futures::stream::BoxStream;
use serde_json::Value;
use tokio::sync::mpsc::UnboundedSender;
use tokio_util::sync::CancellationToken;
use crate::core::goal::runtime::{GoalRuntime, StopReason};
use crate::core::types::{ChatRequest, ContentBlock, Message, Role, ToolOutput, Usage, WireEvent};
use crate::provider::Provider;
use crate::tools::{Registry, ToolCtx};
#[derive(Debug, Clone)]
pub enum AgentEvent {
TextDelta(String),
ThinkingDelta(String),
ToolStarted {
id: String,
name: String,
input: Value,
},
ToolFinished {
id: String,
name: String,
output: ToolOutput,
},
TurnComplete {
usage: Usage,
},
SubAgent {
agent: String,
note: String,
},
GoalContinued {
objective: String,
turn: u32,
},
Error(String),
}
pub struct Session {
pub system: String,
pub model: String,
pub messages: Vec<Message>,
pub total_usage: Usage,
}
impl Session {
pub fn new(system: impl Into<String>, model: impl Into<String>) -> Self {
Session {
system: system.into(),
model: model.into(),
messages: Vec::new(),
total_usage: Usage::default(),
}
}
pub fn push_user(&mut self, text: impl Into<String>) {
self.messages.push(Message::user(text));
}
pub fn push_user_content(&mut self, content: Vec<ContentBlock>) {
self.messages.push(Message {
role: Role::User,
content,
});
}
}
pub struct AgentLoop {
pub provider: Arc<dyn Provider>,
pub registry: Arc<Registry>,
pub ctx: ToolCtx,
pub max_tokens: Option<u32>,
pub goal: Option<Arc<GoalRuntime>>,
turn_seq: Arc<std::sync::atomic::AtomicU64>,
}
impl AgentLoop {
pub fn new(provider: Arc<dyn Provider>, registry: Arc<Registry>, ctx: ToolCtx) -> Self {
AgentLoop {
provider,
registry,
ctx,
max_tokens: None,
goal: None,
turn_seq: Arc::new(std::sync::atomic::AtomicU64::new(0)),
}
}
pub fn with_goal(mut self, goal: Arc<GoalRuntime>) -> Self {
self.goal = Some(goal);
self
}
pub fn inherit_from(mut self, other: &AgentLoop) -> Self {
self.goal = other.goal.clone();
self.turn_seq = other.turn_seq.clone();
self
}
fn next_turn_id(&self) -> String {
let n = self
.turn_seq
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
format!("turn-{n}")
}
pub async fn run_turn(
&self,
session: &mut Session,
events: &UnboundedSender<AgentEvent>,
cancel: &CancellationToken,
) -> anyhow::Result<()> {
let turn_id = self.next_turn_id();
if let Some(goal) = &self.goal {
goal.on_turn_start(&turn_id, session.total_usage);
}
let result = self.drive_turn(&turn_id, session, events, cancel).await;
if let Some(goal) = &self.goal {
goal.on_iteration(&turn_id);
match &result {
Err(e) => {
goal.on_turn_error(&turn_id, stop_reason_for(&e.to_string()))
.await
}
Ok(TurnEnd::Interrupted) => goal.on_turn_abort(&turn_id).await,
Ok(TurnEnd::Completed) => goal.on_turn_stop(&turn_id).await,
}
}
result.map(|_| ())
}
async fn drive_turn(
&self,
turn_id: &str,
session: &mut Session,
events: &UnboundedSender<AgentEvent>,
cancel: &CancellationToken,
) -> anyhow::Result<TurnEnd> {
loop {
if cancel.is_cancelled() {
let _ = events.send(AgentEvent::TurnComplete {
usage: Usage::default(),
});
return Ok(TurnEnd::Interrupted);
}
if let Some(goal) = &self.goal {
for text in goal.take_pending_steering() {
session.messages.push(Message::user(text));
}
}
let req = ChatRequest {
model: session.model.clone(),
system: session.system.clone(),
messages: session.messages.clone(),
tools: self.registry.specs(),
max_tokens: self.max_tokens,
temperature: None,
};
let stream = tokio::select! {
s = self.provider.stream(req) => s?,
_ = cancel.cancelled() => {
let _ = events.send(AgentEvent::TurnComplete { usage: Usage::default() });
return Ok(TurnEnd::Interrupted);
}
};
let tap = events.clone();
let (assistant, usage) = tokio::select! {
r = assemble_with(stream, move |ev| {
let mapped = match ev {
WireEvent::TextDelta(s) => Some(AgentEvent::TextDelta(s.clone())),
WireEvent::ThinkingDelta(s) => Some(AgentEvent::ThinkingDelta(s.clone())),
_ => None,
};
if let Some(m) = mapped {
let _ = tap.send(m);
}
}) => r?,
_ = cancel.cancelled() => {
let _ = events.send(AgentEvent::TurnComplete { usage: Usage::default() });
return Ok(TurnEnd::Interrupted);
}
};
session.total_usage = add_usage(session.total_usage, usage);
if let Some(goal) = &self.goal {
goal.on_token_usage(turn_id, session.total_usage);
}
let mut assistant = assistant;
let has_structured = assistant
.content
.iter()
.any(|b| matches!(b, ContentBlock::ToolUse { .. }));
if !has_structured {
crate::core::toolcall_text::recover_text_tool_calls(&mut assistant);
}
session.messages.push(assistant.clone());
let tool_uses: Vec<(String, String, Value)> = assistant
.content
.iter()
.filter_map(|b| match b {
ContentBlock::ToolUse { id, name, input } => {
Some((id.clone(), name.clone(), input.clone()))
}
_ => None,
})
.collect();
if tool_uses.is_empty() {
let has_text = assistant
.content
.iter()
.any(|b| matches!(b, ContentBlock::Text { text } if !text.is_empty()));
if !has_text {
let _ = events.send(AgentEvent::Error(
"empty response — check the model name and endpoint".into(),
));
}
let _ = events.send(AgentEvent::TurnComplete { usage });
return Ok(TurnEnd::Completed);
}
let mut results = Vec::with_capacity(tool_uses.len());
let mut interrupted = false;
for (id, name, input) in tool_uses {
if interrupted {
results.push(ContentBlock::ToolResult {
id,
out: ToolOutput::error("interrupted"),
});
continue;
}
let _ = events.send(AgentEvent::ToolStarted {
id: id.clone(),
name: name.clone(),
input: input.clone(),
});
let output = match self.registry.get(&name) {
Some(tool) => tokio::select! {
o = tool.run(input, &self.ctx) => o,
_ = cancel.cancelled() => {
interrupted = true;
ToolOutput::error("interrupted")
}
},
None => ToolOutput::error(format!("unknown tool: {name}")),
};
let _ = events.send(AgentEvent::ToolFinished {
id: id.clone(),
name: name.clone(),
output: output.clone(),
});
results.push(ContentBlock::ToolResult { id, out: output });
if let Some(goal) = &self.goal {
goal.on_tool_finish(turn_id, &name).await;
}
}
session.messages.push(Message {
role: Role::Tool,
content: results,
});
if interrupted {
let _ = events.send(AgentEvent::TurnComplete {
usage: Usage::default(),
});
return Ok(TurnEnd::Interrupted);
}
}
}
}
enum TurnEnd {
Completed,
Interrupted,
}
fn stop_reason_for(error: &str) -> StopReason {
let lowered = error.to_ascii_lowercase();
let usage_limit = ["usage limit", "quota", "insufficient_quota", "billing"]
.iter()
.any(|needle| lowered.contains(needle));
if usage_limit {
StopReason::UsageLimit
} else {
StopReason::TurnError
}
}
fn add_usage(a: Usage, b: Usage) -> Usage {
Usage {
input_tokens: a.input_tokens + b.input_tokens,
output_tokens: a.output_tokens + b.output_tokens,
cache_read: a.cache_read + b.cache_read,
cache_write: a.cache_write + b.cache_write,
}
}
pub async fn assemble(stream: BoxStream<'static, WireEvent>) -> anyhow::Result<(Message, Usage)> {
assemble_with(stream, |_| {}).await
}
pub async fn assemble_with<F: FnMut(&WireEvent)>(
mut stream: BoxStream<'static, WireEvent>,
mut tap: F,
) -> anyhow::Result<(Message, Usage)> {
let mut blocks: Vec<ContentBlock> = Vec::new();
let mut text = String::new();
let mut think = String::new();
let mut usage = Usage::default();
let mut tool_idx: HashMap<String, usize> = HashMap::new();
let mut tool_args: HashMap<String, String> = HashMap::new();
fn flush(text: &mut String, think: &mut String, blocks: &mut Vec<ContentBlock>) {
if !think.is_empty() {
blocks.push(ContentBlock::Thinking {
text: std::mem::take(think),
});
}
if !text.is_empty() {
blocks.push(ContentBlock::Text {
text: std::mem::take(text),
});
}
}
while let Some(ev) = stream.next().await {
tap(&ev);
match ev {
WireEvent::TextDelta(s) => text.push_str(&s),
WireEvent::ThinkingDelta(s) => think.push_str(&s),
WireEvent::ToolUseStart { id, name } => {
flush(&mut text, &mut think, &mut blocks);
let idx = blocks.len();
blocks.push(ContentBlock::ToolUse {
id: id.clone(),
name,
input: Value::Null,
});
tool_idx.insert(id.clone(), idx);
tool_args.insert(id, String::new());
}
WireEvent::ToolInputDelta { id, json } => {
if let Some(buf) = tool_args.get_mut(&id) {
buf.push_str(&json);
}
}
WireEvent::ToolUseEnd { id } => {
if let (Some(&idx), Some(args)) = (tool_idx.get(&id), tool_args.get(&id)) {
let input = if args.trim().is_empty() {
Value::Null
} else {
serde_json::from_str(args).unwrap_or(Value::Null)
};
if let ContentBlock::ToolUse { input: slot, .. } = &mut blocks[idx] {
*slot = input;
}
}
}
WireEvent::Usage(u) => usage = u,
WireEvent::Done => break,
WireEvent::Error(e) => return Err(anyhow::anyhow!(e)),
}
}
flush(&mut text, &mut think, &mut blocks);
Ok((Message::assistant(blocks), usage))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::capability::CapabilitySource;
use crate::core::types::Caps;
use crate::provider::Provider;
use crate::tools::builtins::BuiltinTools;
use crate::tools::optimize::Optimizer;
use async_trait::async_trait;
use std::collections::VecDeque;
use std::sync::Mutex;
fn stream_of(events: Vec<WireEvent>) -> BoxStream<'static, WireEvent> {
futures::stream::iter(events).boxed()
}
#[tokio::test]
async fn assembles_text_then_tool_call() {
let events = vec![
WireEvent::TextDelta("Editing ".into()),
WireEvent::TextDelta("the file.".into()),
WireEvent::ToolUseStart {
id: "c1".into(),
name: "edit".into(),
},
WireEvent::ToolInputDelta {
id: "c1".into(),
json: "{\"path\":".into(),
},
WireEvent::ToolInputDelta {
id: "c1".into(),
json: "\"a.rs\"}".into(),
},
WireEvent::ToolUseEnd { id: "c1".into() },
WireEvent::Done,
];
let (msg, _usage) = assemble(stream_of(events)).await.unwrap();
assert_eq!(msg.role, Role::Assistant);
assert_eq!(msg.content.len(), 2);
assert_eq!(msg.content[0], ContentBlock::text("Editing the file."));
}
#[tokio::test]
async fn error_event_propagates() {
let err = assemble(stream_of(vec![WireEvent::Error("boom".into())]))
.await
.unwrap_err();
assert_eq!(err.to_string(), "boom");
}
struct MockProvider {
scripts: Mutex<VecDeque<Vec<WireEvent>>>,
}
#[async_trait]
impl Provider for MockProvider {
async fn stream(&self, _req: ChatRequest) -> anyhow::Result<BoxStream<'static, WireEvent>> {
let script = self.scripts.lock().unwrap().pop_front().unwrap_or_default();
Ok(stream_of(script))
}
fn caps(&self) -> Caps {
Caps {
tools: true,
..Default::default()
}
}
}
#[tokio::test]
async fn full_turn_runs_tool_and_feeds_result_back() {
let dir = tempfile::tempdir().unwrap();
std::fs::write(dir.path().join("f.txt"), "hello world").unwrap();
let scripts = VecDeque::from(vec![
vec![
WireEvent::TextDelta("reading".into()),
WireEvent::ToolUseStart {
id: "t1".into(),
name: "read".into(),
},
WireEvent::ToolInputDelta {
id: "t1".into(),
json: "{\"path\":\"f.txt\"}".into(),
},
WireEvent::ToolUseEnd { id: "t1".into() },
WireEvent::Done,
],
vec![
WireEvent::TextDelta("the file says hello".into()),
WireEvent::Done,
],
]);
let provider = Arc::new(MockProvider {
scripts: Mutex::new(scripts),
});
let mut reg = Registry::new();
for t in BuiltinTools::new(Arc::new(Optimizer::new(true))).tools() {
reg.register(t);
}
let ctx = ToolCtx::new(dir.path());
let agent = AgentLoop::new(provider, Arc::new(reg), ctx);
let mut session = Session::new("sys", "mock");
session.push_user("what does f.txt say?");
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
agent
.run_turn(&mut session, &tx, &CancellationToken::new())
.await
.unwrap();
drop(tx);
let mut saw_tool_result_with_hello = false;
let mut completed = false;
while let Some(ev) = rx.recv().await {
match ev {
AgentEvent::ToolFinished { name, output, .. } => {
assert_eq!(name, "read");
if output.text.contains("hello world") {
saw_tool_result_with_hello = true;
}
}
AgentEvent::TurnComplete { .. } => completed = true,
_ => {}
}
}
assert!(
saw_tool_result_with_hello,
"read result should reach the UI"
);
assert!(completed, "turn should complete");
assert_eq!(session.messages.len(), 4);
assert_eq!(session.messages[3].role, Role::Assistant);
}
use crate::core::goal::runtime::GoalRuntime;
use crate::core::goal::{GoalLimits, GoalStatus, GoalStore};
fn read_then_reply(reply: &str) -> VecDeque<Vec<WireEvent>> {
VecDeque::from(vec![
vec![
WireEvent::ToolUseStart {
id: "t1".into(),
name: "read".into(),
},
WireEvent::ToolInputDelta {
id: "t1".into(),
json: "{\"path\":\"f.txt\"}".into(),
},
WireEvent::ToolUseEnd { id: "t1".into() },
WireEvent::Usage(Usage {
input_tokens: 400,
output_tokens: 100,
..Default::default()
}),
WireEvent::Done,
],
vec![WireEvent::TextDelta(reply.into()), WireEvent::Done],
])
}
fn goal_agent(
scripts: VecDeque<Vec<WireEvent>>,
goal: Arc<GoalRuntime>,
cwd: &std::path::Path,
) -> AgentLoop {
let provider = Arc::new(MockProvider {
scripts: Mutex::new(scripts),
});
let mut reg = Registry::new();
let tools = BuiltinTools::new(Arc::new(Optimizer::new(true))).with_goal(goal.clone());
for t in tools.tools() {
reg.register(t);
}
AgentLoop::new(provider, Arc::new(reg), ToolCtx::new(cwd)).with_goal(goal)
}
#[tokio::test]
async fn a_turn_charges_its_usage_to_the_active_goal() {
let dir = tempfile::tempdir().unwrap();
std::fs::write(dir.path().join("f.txt"), "hello world").unwrap();
let goal = Arc::new(GoalRuntime::new(Arc::new(GoalStore::ephemeral()), true));
goal.create_goal("ship the parser", GoalLimits::tokens(Some(10_000)))
.await
.unwrap();
let agent = goal_agent(read_then_reply("done for now"), goal.clone(), dir.path());
let mut session = Session::new("sys", "mock");
session.push_user("get going");
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
agent
.run_turn(&mut session, &tx, &CancellationToken::new())
.await
.unwrap();
let g = goal.goal().unwrap();
assert_eq!(g.tokens_used, 500, "fresh input + output are charged");
assert_eq!(g.iterations_used, 1);
assert_eq!(g.status, GoalStatus::Active);
assert!(goal.continue_if_idle().await.is_some());
}
#[tokio::test]
async fn crossing_the_budget_mid_turn_injects_the_wrap_up_prompt() {
let dir = tempfile::tempdir().unwrap();
std::fs::write(dir.path().join("f.txt"), "hello world").unwrap();
let goal = Arc::new(GoalRuntime::new(Arc::new(GoalStore::ephemeral()), true));
goal.create_goal("ship the parser", GoalLimits::tokens(Some(100)))
.await
.unwrap();
let agent = goal_agent(read_then_reply("wrapping up"), goal.clone(), dir.path());
let mut session = Session::new("sys", "mock");
session.push_user("get going");
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
agent
.run_turn(&mut session, &tx, &CancellationToken::new())
.await
.unwrap();
assert_eq!(goal.goal().unwrap().status, GoalStatus::BudgetLimited);
let injected = session.messages.iter().any(|m| {
m.role == Role::User
&& m.content.iter().any(|b| {
matches!(b, ContentBlock::Text { text }
if text.contains("do not start new substantive work"))
})
});
assert!(injected, "the model must be told to land the plane");
assert!(goal.continue_if_idle().await.is_none());
}
#[tokio::test]
async fn the_model_ends_the_loop_by_completing_the_goal() {
let dir = tempfile::tempdir().unwrap();
let goal = Arc::new(GoalRuntime::new(Arc::new(GoalStore::ephemeral()), true));
goal.create_goal("ship the parser", GoalLimits::default())
.await
.unwrap();
let scripts = VecDeque::from(vec![
vec![
WireEvent::TextDelta("made progress".into()),
WireEvent::Usage(Usage {
input_tokens: 100,
output_tokens: 20,
..Default::default()
}),
WireEvent::Done,
],
vec![
WireEvent::ToolUseStart {
id: "u1".into(),
name: "update_goal".into(),
},
WireEvent::ToolInputDelta {
id: "u1".into(),
json: "{\"status\":\"complete\"}".into(),
},
WireEvent::ToolUseEnd { id: "u1".into() },
WireEvent::Done,
],
vec![
WireEvent::TextDelta("goal achieved".into()),
WireEvent::Done,
],
]);
let agent = goal_agent(scripts, goal.clone(), dir.path());
let mut session = Session::new("sys", "mock");
session.push_user("get going");
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let mut turns = 0;
loop {
agent
.run_turn(&mut session, &tx, &CancellationToken::new())
.await
.unwrap();
turns += 1;
match goal.continue_if_idle().await {
Some(prompt) => session.push_user(prompt),
None => break,
}
assert!(turns < 5, "the loop must terminate");
}
assert_eq!(turns, 2, "one user turn plus one automatic continuation");
assert_eq!(goal.goal().unwrap().status, GoalStatus::Complete);
}
#[tokio::test]
async fn a_failed_turn_blocks_the_goal_instead_of_retrying_forever() {
let dir = tempfile::tempdir().unwrap();
let goal = Arc::new(GoalRuntime::new(Arc::new(GoalStore::ephemeral()), true));
goal.create_goal("ship the parser", GoalLimits::default())
.await
.unwrap();
let scripts = VecDeque::from(vec![vec![WireEvent::Error("stream exploded".into())]]);
let agent = goal_agent(scripts, goal.clone(), dir.path());
let mut session = Session::new("sys", "mock");
session.push_user("get going");
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
assert!(
agent
.run_turn(&mut session, &tx, &CancellationToken::new())
.await
.is_err()
);
assert_eq!(goal.goal().unwrap().status, GoalStatus::Blocked);
assert!(goal.continue_if_idle().await.is_none());
}
}