use std::time::{Duration, Instant};
use futures::StreamExt;
use tokio::time::Instant as TokioInstant;
use crate::driver_registry::{
LlmCompletionMetadata, LlmResponse, LlmResponseStream, LlmStreamEvent,
};
use crate::error::{AgentLoopError, LlmErrorKind, Result};
use crate::execution_phase::ExecutionPhase;
use crate::reasoning::ReasoningContentPart;
use crate::tool_types::ToolCall;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct TurnLimits {
pub total: Option<Duration>,
pub first_event: Option<Duration>,
pub max_response_bytes: Option<u64>,
pub require_terminal_event: bool,
}
impl TurnLimits {
#[must_use]
pub fn with_total(mut self, total: Duration) -> Self {
self.total = Some(total);
self
}
#[must_use]
pub fn with_first_event(mut self, first_event: Duration) -> Self {
self.first_event = Some(first_event);
self
}
#[must_use]
pub fn with_max_response_bytes(mut self, bytes: u64) -> Self {
self.max_response_bytes = Some(bytes);
self
}
#[must_use]
pub fn requiring_terminal_event(mut self) -> Self {
self.require_terminal_event = true;
self
}
pub fn is_unbounded(&self) -> bool {
*self == TurnLimits::default()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub struct TurnTiming {
pub time_to_first_event: Option<Duration>,
pub mean_inter_event: Option<Duration>,
pub total: Duration,
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct CollectedTurn {
pub text: String,
pub reasoning: Vec<ReasoningContentPart>,
pub tool_calls: Vec<ToolCall>,
pub phase: Option<ExecutionPhase>,
pub metadata: LlmCompletionMetadata,
pub complete: bool,
pub timing: TurnTiming,
}
impl CollectedTurn {
pub fn reasoning_text(&self) -> String {
self.reasoning
.iter()
.filter_map(|part| part.display_text())
.collect::<Vec<_>>()
.join("")
}
pub fn into_response(self) -> LlmResponse {
LlmResponse {
text: self.text,
reasoning: self.reasoning,
tool_calls: (!self.tool_calls.is_empty()).then_some(self.tool_calls),
metadata: self.metadata,
}
}
}
pub async fn collect_turn(
mut stream: LlmResponseStream,
limits: &TurnLimits,
mut observe: impl FnMut(&LlmStreamEvent),
) -> Result<CollectedTurn> {
let started = Instant::now();
let total_deadline = limits.total.map(|total| TokioInstant::now() + total);
let first_event_deadline = limits
.first_event
.map(|first_event| TokioInstant::now() + first_event);
let mut turn = CollectedTurn::default();
let mut accumulated: u64 = 0;
let mut first_event_at: Option<Instant> = None;
let mut last_event_at: Option<Instant> = None;
let mut content_events: u32 = 0;
loop {
let deadline = match (total_deadline, first_event_deadline) {
(Some(total), Some(first)) if first_event_at.is_none() => Some(total.min(first)),
(total, first) => {
if first_event_at.is_none() {
total.or(first)
} else {
total
}
}
};
let next = match deadline {
Some(deadline) => match tokio::time::timeout_at(deadline, stream.next()).await {
Ok(next) => next,
Err(_) => return Err(timed_out(limits, first_event_at.is_none(), started)),
},
None => stream.next().await,
};
let Some(event) = next else { break };
let event = event?;
let now = Instant::now();
if first_event_at.is_none() {
first_event_at = Some(now);
}
observe(&event);
match event {
LlmStreamEvent::TextDelta(delta) => {
if delta.is_empty() {
continue;
}
accumulated += delta.len() as u64;
check_cap(accumulated, limits)?;
turn.text.push_str(&delta);
content_events += 1;
last_event_at = Some(now);
}
LlmStreamEvent::ReasoningDelta { delta, .. } => {
if delta.is_empty() {
continue;
}
accumulated += delta.len() as u64;
check_cap(accumulated, limits)?;
content_events += 1;
last_event_at = Some(now);
}
LlmStreamEvent::ReasoningItem(item) => {
if let Some(text) = item.display_text() {
accumulated += text.len() as u64;
check_cap(accumulated, limits)?;
}
turn.reasoning.push(item);
}
LlmStreamEvent::ToolCalls(calls) => turn.tool_calls = calls,
LlmStreamEvent::NativeToolCall(_) => {
return Err(AgentLoopError::config(
"native async/custom calls require a streaming coordinator",
));
}
LlmStreamEvent::MessagePhase(phase) => turn.phase = Some(phase),
LlmStreamEvent::ProviderCompactionStarted | LlmStreamEvent::HostedToolCall(_) => {}
LlmStreamEvent::Done(metadata) => {
turn.metadata = *metadata;
turn.complete = true;
}
LlmStreamEvent::Error(error) => return Err(error.into_agent_error()),
}
}
if !turn.complete && limits.require_terminal_event {
return Err(AgentLoopError::llm_kind(
LlmErrorKind::MalformedResponse,
"provider stream ended before its terminal event: the turn is incomplete",
));
}
turn.timing = TurnTiming {
time_to_first_event: first_event_at.map(|at| at.duration_since(started)),
mean_inter_event: match (first_event_at, last_event_at) {
(Some(first), Some(last)) if content_events >= 2 => {
Some(last.duration_since(first) / (content_events - 1))
}
_ => None,
},
total: started.elapsed(),
};
Ok(turn)
}
pub fn limit_stream(stream: LlmResponseStream, limits: TurnLimits) -> LlmResponseStream {
if limits.is_unbounded() {
return stream;
}
struct State {
stream: LlmResponseStream,
limits: TurnLimits,
started: Instant,
total_deadline: Option<TokioInstant>,
first_event_deadline: Option<TokioInstant>,
seen_event: bool,
accumulated: u64,
done: bool,
}
let state = State {
stream,
limits,
started: Instant::now(),
total_deadline: limits.total.map(|total| TokioInstant::now() + total),
first_event_deadline: limits.first_event.map(|first| TokioInstant::now() + first),
seen_event: false,
accumulated: 0,
done: false,
};
Box::pin(futures::stream::unfold(state, |mut state| async move {
if state.done {
return None;
}
let deadline = match (state.total_deadline, state.first_event_deadline) {
(Some(total), Some(first)) if !state.seen_event => Some(total.min(first)),
(total, first) if !state.seen_event => total.or(first),
(total, _) => total,
};
let next = match deadline {
Some(deadline) => match tokio::time::timeout_at(deadline, state.stream.next()).await {
Ok(next) => next,
Err(_) => {
let error = timed_out(&state.limits, !state.seen_event, state.started);
state.done = true;
return Some((Err(error), state));
}
},
None => state.stream.next().await,
};
let item = next?;
state.seen_event = true;
if let Ok(event) = &item {
let produced = match event {
LlmStreamEvent::TextDelta(delta) => delta.len() as u64,
LlmStreamEvent::ReasoningDelta { delta, .. } => delta.len() as u64,
LlmStreamEvent::ReasoningItem(item) => {
item.display_text().map_or(0, |text| text.len() as u64)
}
_ => 0,
};
state.accumulated += produced;
if let Err(error) = check_cap(state.accumulated, &state.limits) {
state.done = true;
return Some((Err(error), state));
}
}
Some((item, state))
}))
}
fn check_cap(accumulated: u64, limits: &TurnLimits) -> Result<()> {
match limits.max_response_bytes {
Some(cap) if accumulated > cap => Err(AgentLoopError::llm_kind(
LlmErrorKind::MalformedResponse,
format!("provider response exceeded the {cap}-byte limit for this call"),
)),
_ => Ok(()),
}
}
fn timed_out(limits: &TurnLimits, before_first_event: bool, started: Instant) -> AgentLoopError {
let elapsed = started.elapsed();
if before_first_event && limits.first_event.is_some_and(|first| elapsed >= first) {
first_event_timeout(limits.first_event.unwrap_or(elapsed))
} else {
turn_timeout(limits.total.unwrap_or(elapsed))
}
}
pub fn turn_timeout(after: Duration) -> AgentLoopError {
AgentLoopError::llm_kind(
LlmErrorKind::Unavailable,
format!("provider turn did not finish within {after:?}"),
)
}
pub fn first_event_timeout(after: Duration) -> AgentLoopError {
AgentLoopError::llm_kind(
LlmErrorKind::Unavailable,
format!("provider sent no response within {after:?}"),
)
}
pub async fn connect_within<T, F>(
limits: &TurnLimits,
future: F,
) -> Result<(T, std::time::Duration)>
where
F: std::future::Future<Output = Result<T>>,
{
let budget = match (limits.total, limits.first_event) {
(Some(total), Some(first)) => Some(total.min(first)),
(total, first) => total.or(first),
};
let started = Instant::now();
let value = match budget {
Some(budget) => match tokio::time::timeout(budget, future).await {
Ok(value) => value?,
Err(_) => {
return Err(if limits.first_event == Some(budget) {
first_event_timeout(budget)
} else {
turn_timeout(budget)
});
}
},
None => future.await?,
};
Ok((value, started.elapsed()))
}
impl TurnLimits {
#[must_use]
pub fn after(mut self, spent: std::time::Duration) -> Self {
self.total = self.total.map(|total| total.saturating_sub(spent));
self.first_event = self.first_event.map(|first| first.saturating_sub(spent));
self
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::driver_registry::LlmStreamError;
use crate::reasoning::{ReasoningContentPart, ReasoningText};
use futures::stream;
use serde_json::json;
fn reasoning(text: &str) -> ReasoningContentPart {
ReasoningContentPart::opaque("test").with_text(ReasoningText::Plain { text: text.into() })
}
fn done(finish: &str) -> LlmStreamEvent {
LlmStreamEvent::Done(Box::new(LlmCompletionMetadata {
finish_reason: Some(finish.to_owned()),
total_tokens: Some(7),
..Default::default()
}))
}
fn streamed(events: Vec<LlmStreamEvent>) -> LlmResponseStream {
Box::pin(stream::iter(events.into_iter().map(Ok)))
}
#[tokio::test]
async fn folds_text_reasoning_and_tool_calls_into_one_turn() {
let events = vec![
LlmStreamEvent::TextDelta("Hel".into()),
LlmStreamEvent::TextDelta(String::new()),
LlmStreamEvent::TextDelta("lo".into()),
LlmStreamEvent::ReasoningDelta {
delta: "thin".into(),
summary: false,
},
LlmStreamEvent::ReasoningItem(reasoning("thinking")),
LlmStreamEvent::ToolCalls(vec![ToolCall {
id: "call_1".into(),
name: "search".into(),
arguments: json!({"q": "everruns"}),
}]),
done("tool_calls"),
];
let mut seen = 0;
let turn = collect_turn(streamed(events), &TurnLimits::default(), |_| seen += 1)
.await
.unwrap();
assert_eq!(turn.text, "Hello");
assert_eq!(turn.reasoning_text(), "thinking");
assert_eq!(turn.tool_calls.len(), 1);
assert_eq!(turn.metadata.finish_reason.as_deref(), Some("tool_calls"));
assert!(turn.complete);
assert_eq!(
seen, 7,
"every event reaches the observer, empties included"
);
}
#[tokio::test]
async fn a_cut_short_stream_is_reported_and_only_refused_on_request() {
let events = || vec![LlmStreamEvent::TextDelta("partial".into())];
let lenient = collect_turn(streamed(events()), &TurnLimits::default(), |_| {})
.await
.unwrap();
assert_eq!(lenient.text, "partial");
assert!(!lenient.complete, "the missing terminal event is visible");
let strict = collect_turn(
streamed(events()),
&TurnLimits::default().requiring_terminal_event(),
|_| {},
)
.await
.expect_err("a required terminal event that never arrived is an error");
assert_eq!(
strict.llm_error_kind(),
Some(LlmErrorKind::MalformedResponse)
);
}
#[tokio::test]
async fn the_byte_cap_trips_while_the_answer_is_still_arriving() {
let events = vec![
LlmStreamEvent::TextDelta("a".repeat(8)),
LlmStreamEvent::TextDelta("b".repeat(8)),
done("stop"),
];
let error = collect_turn(
streamed(events),
&TurnLimits::default().with_max_response_bytes(10),
|_| {},
)
.await
.expect_err("an over-cap answer must not pass for a turn");
assert_eq!(
error.llm_error_kind(),
Some(LlmErrorKind::MalformedResponse)
);
assert!(error.to_string().contains("10-byte limit"), "{error}");
}
#[tokio::test]
async fn reasoning_text_counts_against_the_same_cap_as_answer_text() {
let events = vec![
LlmStreamEvent::TextDelta("hi".into()),
LlmStreamEvent::ReasoningItem(reasoning(&"r".repeat(20))),
done("stop"),
];
let error = collect_turn(
streamed(events),
&TurnLimits::default().with_max_response_bytes(10),
|_| {},
)
.await
.expect_err("reasoning is answer text too");
assert_eq!(
error.llm_error_kind(),
Some(LlmErrorKind::MalformedResponse)
);
}
#[tokio::test]
async fn reasoning_deltas_count_against_the_response_byte_cap() {
let events = vec![
LlmStreamEvent::ReasoningDelta {
delta: "r".repeat(11),
summary: false,
},
done("stop"),
];
let error = collect_turn(
streamed(events),
&TurnLimits::default().with_max_response_bytes(10),
|_| {},
)
.await
.expect_err("reasoning deltas must not bypass the response byte cap");
assert_eq!(
error.llm_error_kind(),
Some(LlmErrorKind::MalformedResponse)
);
}
#[tokio::test]
async fn a_stream_error_keeps_its_kind_and_status() {
let events = vec![Err(AgentLoopError::llm("upstream reset"))];
let stream: LlmResponseStream = Box::pin(stream::iter(events));
let error = collect_turn(stream, &TurnLimits::default(), |_| {})
.await
.expect_err("a stream failure propagates");
assert_eq!(error.to_string(), "LLM error: upstream reset");
let mut inline = LlmStreamError::new("rate limited");
inline.status = Some(429);
let stream = streamed(vec![LlmStreamEvent::Error(inline)]);
let error = collect_turn(stream, &TurnLimits::default(), |_| {})
.await
.expect_err("an inline error event fails the turn");
assert_eq!(error.http_status(), Some(429));
assert!(error.is_rate_limited());
}
#[tokio::test]
async fn a_silent_provider_trips_the_first_event_limit() {
let stream: LlmResponseStream = Box::pin(stream::once(async {
tokio::time::sleep(Duration::from_secs(30)).await;
Ok(LlmStreamEvent::TextDelta("late".into()))
}));
let limits = TurnLimits::default().with_first_event(Duration::from_millis(20));
let error = collect_turn(stream, &limits, |_| {})
.await
.expect_err("a provider that never answers must not hang the caller");
assert_eq!(error.llm_error_kind(), Some(LlmErrorKind::Unavailable));
assert!(error.to_string().contains("no response within"), "{error}");
assert!(error.is_transient_llm_error());
}
#[tokio::test]
async fn a_slow_turn_trips_the_total_limit_after_it_has_started() {
let stream: LlmResponseStream = Box::pin(
stream::once(async { Ok(LlmStreamEvent::TextDelta("start".into())) }).chain(
stream::once(async {
tokio::time::sleep(Duration::from_secs(30)).await;
Ok(done("stop"))
}),
),
);
let limits = TurnLimits::default()
.with_first_event(Duration::from_secs(30))
.with_total(Duration::from_millis(20));
let error = collect_turn(stream, &limits, |_| {})
.await
.expect_err("a turn that never finishes must not hang the caller");
assert!(
error.to_string().contains("did not finish within"),
"{error}"
);
}
#[tokio::test]
async fn timing_reports_first_event_and_the_mean_gap_between_content_events() {
let stream: LlmResponseStream = Box::pin(
stream::once(async {
tokio::time::sleep(Duration::from_millis(20)).await;
Ok(LlmStreamEvent::TextDelta("a".into()))
})
.chain(stream::once(async {
tokio::time::sleep(Duration::from_millis(20)).await;
Ok(LlmStreamEvent::TextDelta("b".into()))
}))
.chain(stream::once(async { Ok(done("stop")) })),
);
let turn = collect_turn(stream, &TurnLimits::default(), |_| {})
.await
.unwrap();
let ttfe = turn.timing.time_to_first_event.expect("an event arrived");
assert!(ttfe >= Duration::from_millis(20), "{ttfe:?}");
let mean = turn.timing.mean_inter_event.expect("two content events");
assert!(mean >= Duration::from_millis(20), "{mean:?}");
assert!(turn.timing.total >= ttfe);
}
#[tokio::test]
async fn a_single_content_event_has_no_gap_to_average() {
let turn = collect_turn(
streamed(vec![LlmStreamEvent::TextDelta("one".into()), done("stop")]),
&TurnLimits::default(),
|_| {},
)
.await
.unwrap();
assert_eq!(turn.timing.mean_inter_event, None);
}
#[tokio::test]
async fn limit_stream_ends_the_stream_at_the_first_breach() {
let stream = streamed(vec![
LlmStreamEvent::TextDelta("a".repeat(8)),
LlmStreamEvent::TextDelta("b".repeat(8)),
done("stop"),
]);
let limited = limit_stream(stream, TurnLimits::default().with_max_response_bytes(10));
let items: Vec<_> = limited.collect().await;
assert_eq!(items.len(), 2, "one good event, then the breach");
assert!(items[0].is_ok());
let error = items[1].as_ref().expect_err("the cap trips");
assert_eq!(
error.llm_error_kind(),
Some(LlmErrorKind::MalformedResponse)
);
}
#[tokio::test]
async fn limit_stream_counts_reasoning_deltas() {
let stream = streamed(vec![
LlmStreamEvent::ReasoningDelta {
delta: "r".repeat(11),
summary: false,
},
done("stop"),
]);
let items: Vec<_> = limit_stream(stream, TurnLimits::default().with_max_response_bytes(10))
.collect()
.await;
assert_eq!(items.len(), 1);
assert!(items[0].is_err());
}
#[tokio::test]
async fn limit_stream_bounds_a_provider_that_stops_sending() {
let stream: LlmResponseStream = Box::pin(
stream::once(async { Ok(LlmStreamEvent::TextDelta("start".into())) }).chain(
stream::once(async {
tokio::time::sleep(Duration::from_secs(30)).await;
Ok(done("stop"))
}),
),
);
let limited = limit_stream(
stream,
TurnLimits::default().with_total(Duration::from_millis(20)),
);
let items: Vec<_> = limited.collect().await;
assert_eq!(items.len(), 2);
let error = items[1].as_ref().expect_err("the turn ran out of time");
assert_eq!(error.llm_error_kind(), Some(LlmErrorKind::Unavailable));
}
#[tokio::test]
async fn limit_stream_passes_an_unbounded_stream_straight_through() {
let stream = streamed(vec![LlmStreamEvent::TextDelta("hi".into()), done("stop")]);
let items: Vec<_> = limit_stream(stream, TurnLimits::default()).collect().await;
assert_eq!(items.len(), 2);
assert!(items.iter().all(Result::is_ok));
}
#[test]
fn default_limits_are_unbounded() {
assert!(TurnLimits::default().is_unbounded());
assert!(
!TurnLimits::default()
.requiring_terminal_event()
.is_unbounded()
);
}
}