use crate::api::ApiClient;
use crate::api::error::{ApiError, http_status_from_message, parse_retry_after};
use crate::cancel::CancelSignal;
use crate::message::Message;
use crate::stream::rate_limit;
use crate::stream::{StreamAccumulator, StreamEvent, StreamStopReason, Usage};
use futures::StreamExt;
use futures::stream::Stream;
use std::fmt;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct StreamTimeoutConfig {
pub initial_event_timeout: Duration,
pub per_event_timeout: Duration,
pub total_stream_timeout: Duration,
pub max_consecutive_timeouts: u32,
pub fallback_to_non_streaming: bool,
}
impl Default for StreamTimeoutConfig {
fn default() -> Self {
Self {
initial_event_timeout: Duration::from_mins(2),
per_event_timeout: Duration::from_mins(3),
total_stream_timeout: Duration::from_mins(5),
max_consecutive_timeouts: 10,
fallback_to_non_streaming: true,
}
}
}
impl StreamTimeoutConfig {
pub fn validate(&self) -> Result<(), String> {
if self.initial_event_timeout.is_zero() {
return Err("initial_event_timeout must be non-zero".to_string());
}
if self.per_event_timeout.is_zero() {
return Err("per_event_timeout must be non-zero".to_string());
}
if self.total_stream_timeout.is_zero() {
return Err("total_stream_timeout must be non-zero".to_string());
}
if self.total_stream_timeout == Duration::MAX {
return Err(
"total_stream_timeout must be finite — Duration::MAX silently disables the \
total deadline, both backoff clamps, and the non-streaming fallback deadline; \
construct the config directly (as passthrough does) to opt into that"
.to_string(),
);
}
if self.total_stream_timeout < self.initial_event_timeout {
return Err(format!(
"total_stream_timeout ({:?}) must be >= initial_event_timeout ({:?})",
self.total_stream_timeout, self.initial_event_timeout
));
}
if self.max_consecutive_timeouts == 0 {
return Err("max_consecutive_timeouts must be >= 1".to_string());
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct StreamRetryConfig {
pub max_retries: u32,
pub base_delay_ms: u64,
pub max_delay_ms: u64,
pub jitter_factor: f64,
}
impl Default for StreamRetryConfig {
fn default() -> Self {
Self {
max_retries: 3,
base_delay_ms: 100,
max_delay_ms: 10_000,
jitter_factor: 0.1,
}
}
}
impl StreamRetryConfig {
#[must_use]
pub fn base_delay(&self, attempt: u32) -> Duration {
let delay_ms = self
.base_delay_ms
.saturating_mul(1u64.checked_shl(attempt).unwrap_or(u64::MAX));
Duration::from_millis(delay_ms.min(self.max_delay_ms))
}
#[must_use]
pub fn jittered_base_delay(&self, attempt: u32) -> Duration {
let base = self.base_delay(attempt);
if self.jitter_factor == 0.0 {
return base;
}
let f = Self::random_signed_fraction() * self.jitter_factor;
base.mul_f64(1.0 + f)
}
#[must_use]
fn random_signed_fraction() -> f64 {
(fastrand::f64() - 0.5) * 2.0
}
pub fn validate(&self) -> Result<(), String> {
if self.base_delay_ms == 0 {
return Err("base_delay_ms must be non-zero".to_string());
}
if self.max_delay_ms == 0 {
return Err("max_delay_ms must be non-zero".to_string());
}
if self.max_delay_ms < self.base_delay_ms {
return Err(format!(
"max_delay_ms ({}) must be >= base_delay_ms ({})",
self.max_delay_ms, self.base_delay_ms
));
}
if !self.jitter_factor.is_finite() {
return Err(format!(
"jitter_factor must be finite, got {}",
self.jitter_factor
));
}
if !(0.0..=1.0).contains(&self.jitter_factor) {
return Err(format!(
"jitter_factor must be in 0.0..=1.0, got {}",
self.jitter_factor
));
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct RateLimitConfig {
pub respect_retry_after: bool,
pub default_delay: Duration,
pub max_delay: Duration,
pub requests_per_minute: u32,
pub fallback_after_retries: u32,
pub max_retries: u32,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
respect_retry_after: true,
default_delay: Duration::from_secs(5),
max_delay: Duration::from_mins(1),
requests_per_minute: 0,
fallback_after_retries: 3,
max_retries: 5,
}
}
}
impl RateLimitConfig {
pub fn validate(&self) -> Result<(), String> {
if self.default_delay == Duration::ZERO {
return Err("default_delay must be non-zero".into());
}
if self.max_delay < self.default_delay {
return Err("max_delay must be >= default_delay".into());
}
if self.max_retries == 0 {
return Err("max_retries must be >= 1".into());
}
if self.fallback_after_retries > self.max_retries {
return Err(format!(
"fallback_after_retries ({}) must be <= max_retries ({})",
self.fallback_after_retries, self.max_retries
));
}
Ok(())
}
#[must_use]
pub fn backoff(&self, server_hint: Option<Duration>) -> Duration {
match server_hint {
Some(d) if self.respect_retry_after => d.min(self.max_delay),
_ => self.default_delay.min(self.max_delay),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RateLimitKind {
RateLimited,
Overloaded,
}
#[derive(Debug, Clone)]
pub struct DetectedRateLimit {
pub kind: RateLimitKind,
pub retry_after: Option<Duration>,
pub message: String,
}
impl DetectedRateLimit {
#[must_use]
pub fn detect(err: &crate::api::error::ApiError) -> Option<Self> {
match err {
ApiError::RateLimit {
retry_after,
message,
} => Some(Self {
kind: match http_status_from_message(message) {
Some(503 | 529) => RateLimitKind::Overloaded,
_ => RateLimitKind::RateLimited,
},
retry_after: *retry_after,
message: message.clone(),
}),
ApiError::Api(msg) => {
let lower = msg.to_lowercase();
if lower.contains("rate limit") || lower.contains("429") {
Some(Self {
kind: RateLimitKind::RateLimited,
retry_after: parse_retry_after(msg),
message: msg.clone(),
})
} else {
None
}
}
ApiError::Http(msg) => {
let kind = match http_status_from_message(msg) {
Some(429) => RateLimitKind::RateLimited,
Some(503 | 529) => RateLimitKind::Overloaded,
_ => return None,
};
Some(Self {
kind,
retry_after: parse_retry_after(msg),
message: msg.clone(),
})
}
_ => None,
}
}
}
fn clamp_delay_to_deadline(delay: Duration, deadline: Option<Instant>) -> Duration {
let Some(deadline) = deadline else {
return delay;
};
let now = Instant::now();
let Some(remaining) = deadline.checked_duration_since(now) else {
return Duration::ZERO;
};
delay.min(remaining)
}
#[derive(Debug)]
enum RateLimitRetry {
Escalate {
attempts: u32,
retry_after: Option<Duration>,
},
HardStop,
Retry(Duration),
}
enum ErrorAction {
Fail(StreamHandlerError),
TryFallback(Option<StreamOutcome>),
Retry(Duration),
}
struct StreamFailure {
error: StreamHandlerError,
retryable: bool,
}
impl StreamFailure {
fn transient(error: StreamHandlerError) -> Self {
Self {
error,
retryable: true,
}
}
}
fn carried_outcome(error: &StreamHandlerError) -> Option<StreamOutcome> {
match error {
StreamHandlerError::InitFailed(o) | StreamHandlerError::StreamFailed(o) => {
Some(o.to_owned())
}
_ => None,
}
}
async fn sleep_cancellable(
delay: Duration,
cancel: &Arc<CancelSignal>,
) -> Result<(), StreamHandlerError> {
tokio::select! {
() = tokio::time::sleep(delay) => Ok(()),
() = cancel.notified() => Err(StreamHandlerError::Cancelled),
}
}
async fn deadline_future(deadline: Option<Instant>) {
match deadline {
Some(deadline) => tokio::time::sleep_until(deadline.into()).await,
None => std::future::pending::<()>().await,
}
}
enum EventPoll {
Next(Option<Result<crate::stream::StreamEvent, crate::api::error::ApiError>>),
TimedOut,
}
struct EventDiagnostics {
events_processed: u64,
stream_start: Instant,
has_partial_data: bool,
attempts_so_far: u32,
}
impl EventDiagnostics {
fn new(
events_processed: u64,
stream_start: Instant,
shadow: &StreamAccumulator,
attempts_so_far: u32,
) -> Self {
Self {
events_processed,
stream_start,
has_partial_data: !shadow.peek_parts().is_empty(),
attempts_so_far,
}
}
fn total_timeout(&self) -> StreamOutcome {
StreamOutcome::TotalTimeout {
has_partial_data: self.has_partial_data,
events_processed: self.events_processed,
duration: self.stream_start.elapsed(),
}
}
fn event_timeout(&self, consecutive_timeouts: u32) -> StreamOutcome {
StreamOutcome::EventTimeout {
has_partial_data: self.has_partial_data,
consecutive_timeouts,
}
}
fn api_error_failure(&self, error: &crate::api::error::ApiError) -> StreamFailure {
let retryable = error.is_retryable();
if let Some(detail) = DetectedRateLimit::detect(error) {
return StreamFailure {
error: StreamHandlerError::StreamFailed(StreamOutcome::RateLimited {
detail,
has_partial_data: self.has_partial_data,
events_processed: self.events_processed,
}),
retryable,
};
}
StreamFailure {
error: StreamHandlerError::StreamFailed(StreamOutcome::InitFailed {
attempts: self.attempts_so_far,
last_error: error.to_string(),
}),
retryable,
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum StreamOutcome {
Completed {
events_processed: u64,
duration: Duration,
},
TotalTimeout {
has_partial_data: bool,
events_processed: u64,
duration: Duration,
},
EventTimeout {
has_partial_data: bool,
consecutive_timeouts: u32,
},
RateLimited {
detail: DetectedRateLimit,
has_partial_data: bool,
events_processed: u64,
},
InitFailed {
last_error: String,
attempts: u32,
},
FallbackToNonStreaming,
Cancelled,
}
impl fmt::Display for StreamOutcome {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Completed {
events_processed,
duration,
} => {
write!(
f,
"stream completed ({events_processed} events in {:.1}s)",
duration.as_secs_f64()
)
}
Self::TotalTimeout {
has_partial_data,
events_processed,
duration,
} => {
let partial = if *has_partial_data {
" (partial data)"
} else {
""
};
write!(
f,
"total timeout after {:.1}s, {events_processed} events{partial}",
duration.as_secs_f64()
)
}
Self::EventTimeout {
has_partial_data,
consecutive_timeouts,
} => {
let partial = if *has_partial_data {
" (partial data)"
} else {
""
};
write!(
f,
"event timeout after {consecutive_timeouts} consecutive timeouts{partial}"
)
}
Self::RateLimited {
detail,
has_partial_data,
events_processed,
} => {
let kind = match detail.kind {
RateLimitKind::RateLimited => "rate limit",
RateLimitKind::Overloaded => "overloaded",
};
let retry = detail
.retry_after
.map(|d| format!(" (retry after {d:?})"))
.unwrap_or_default();
let partial = if *has_partial_data {
" (partial data)"
} else {
""
};
write!(
f,
"{kind}{retry}{partial}, {events_processed} events processed"
)
}
Self::InitFailed {
last_error,
attempts,
} => {
write!(
f,
"stream failed before completing after {attempts} attempts: {last_error}"
)
}
Self::FallbackToNonStreaming => {
write!(f, "fell back to non-streaming request")
}
Self::Cancelled => write!(f, "cancelled"),
}
}
}
#[derive(Debug)]
#[non_exhaustive]
pub enum StreamHandlerError {
InitFailed(StreamOutcome),
StreamFailed(StreamOutcome),
FallbackFailed {
stream_outcome: StreamOutcome,
fallback_error: String,
},
Cancelled,
Poisoned(&'static str),
RateLimitEscalation {
attempts: u32,
retry_after: Option<Duration>,
},
}
impl fmt::Display for StreamHandlerError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InitFailed(outcome) => write!(f, "stream failed before completing: {outcome}"),
Self::StreamFailed(outcome) => write!(f, "stream failed: {outcome}"),
Self::FallbackFailed {
stream_outcome,
fallback_error,
} => {
write!(
f,
"stream failed ({stream_outcome}) and fallback also failed: {fallback_error}"
)
}
Self::Cancelled => write!(f, "cancelled"),
Self::Poisoned(what) => write!(f, "lock poisoned: {what}"),
Self::RateLimitEscalation {
attempts,
retry_after,
} => write!(
f,
"rate-limit escalation after {attempts} retries (retry-after {retry_after:?})"
),
}
}
}
impl std::error::Error for StreamHandlerError {}
pub struct StreamHandler {
timeout_config: StreamTimeoutConfig,
retry_config: StreamRetryConfig,
rate_limit_config: RateLimitConfig,
rate_limiter: Option<Arc<crate::stream::rate_limit::RateLimiter>>,
rate_limit_max_wait: Duration,
}
impl fmt::Debug for StreamHandler {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StreamHandler")
.field("timeout_config", &self.timeout_config)
.field("retry_config", &self.retry_config)
.field("rate_limit_config", &self.rate_limit_config)
.field("rate_limiter", &self.rate_limiter)
.field("rate_limit_max_wait", &self.rate_limit_max_wait)
.finish()
}
}
impl Default for StreamHandler {
fn default() -> Self {
Self::new()
}
}
impl StreamHandler {
#[must_use]
pub fn passthrough() -> Self {
const NEVER_TIME_OUT: Duration = Duration::MAX;
Self {
timeout_config: StreamTimeoutConfig {
initial_event_timeout: NEVER_TIME_OUT,
per_event_timeout: NEVER_TIME_OUT,
total_stream_timeout: NEVER_TIME_OUT,
max_consecutive_timeouts: 1,
fallback_to_non_streaming: false,
},
retry_config: StreamRetryConfig {
max_retries: 0,
..Default::default()
},
rate_limit_config: RateLimitConfig {
max_retries: 0,
fallback_after_retries: 0,
..Default::default()
},
rate_limiter: None,
rate_limit_max_wait: Duration::from_secs(30),
}
}
#[must_use]
pub fn passthrough_default() -> &'static Self {
static PASSTHROUGH: std::sync::OnceLock<StreamHandler> = std::sync::OnceLock::new();
PASSTHROUGH.get_or_init(Self::passthrough)
}
#[must_use]
pub fn new() -> Self {
Self {
timeout_config: StreamTimeoutConfig::default(),
retry_config: StreamRetryConfig::default(),
rate_limit_config: RateLimitConfig::default(),
rate_limiter: None,
rate_limit_max_wait: Duration::from_secs(30),
}
}
#[must_use]
pub fn with_timeout_config(mut self, timeout: StreamTimeoutConfig) -> Self {
self.timeout_config = Self::sanitized_timeout_config(timeout);
self
}
fn sanitized_timeout_config(timeout: StreamTimeoutConfig) -> StreamTimeoutConfig {
let default = StreamTimeoutConfig::default();
let mut sanitized = timeout;
let mut repaired: Vec<&'static str> = Vec::new();
if sanitized.initial_event_timeout.is_zero()
|| sanitized.initial_event_timeout == Duration::MAX
{
sanitized.initial_event_timeout = default.initial_event_timeout;
repaired.push("initial_event_timeout");
}
if sanitized.per_event_timeout.is_zero() || sanitized.per_event_timeout == Duration::MAX {
sanitized.per_event_timeout = default.per_event_timeout;
repaired.push("per_event_timeout");
}
if sanitized.total_stream_timeout.is_zero()
|| sanitized.total_stream_timeout == Duration::MAX
{
sanitized.total_stream_timeout = default.total_stream_timeout;
repaired.push("total_stream_timeout");
}
if sanitized.total_stream_timeout < sanitized.initial_event_timeout {
sanitized.total_stream_timeout = sanitized.initial_event_timeout;
repaired.push("total_stream_timeout");
}
if sanitized.max_consecutive_timeouts == 0 {
sanitized.max_consecutive_timeouts = default.max_consecutive_timeouts;
repaired.push("max_consecutive_timeouts");
}
if !repaired.is_empty() {
tracing::warn!(
fields = repaired.join(","),
"invalid StreamTimeoutConfig fields substituted with defaults"
);
}
sanitized
}
#[must_use]
pub fn with_retry_config(mut self, retry: StreamRetryConfig) -> Self {
if let Err(e) = retry.validate() {
tracing::warn!(error = %e, "invalid StreamRetryConfig, falling back to default");
} else {
self.retry_config = retry;
}
self
}
#[must_use]
pub fn timeout_config(&self) -> &StreamTimeoutConfig {
&self.timeout_config
}
#[must_use]
pub fn retry_config(&self) -> &StreamRetryConfig {
&self.retry_config
}
#[must_use]
pub fn with_rate_limit_config(mut self, rl: RateLimitConfig) -> Self {
if let Err(e) = rl.validate() {
tracing::warn!(error = %e, "invalid RateLimitConfig, falling back to default");
return self;
}
self.rate_limit_config = rl;
self
}
#[must_use]
pub fn rate_limit_config(&self) -> &RateLimitConfig {
&self.rate_limit_config
}
#[must_use]
pub fn with_rate_limiter(
mut self,
limiter: Arc<crate::stream::rate_limit::RateLimiter>,
) -> Self {
self.rate_limiter = Some(limiter);
self
}
#[must_use]
pub fn with_rate_limit_max_wait(mut self, max_wait: Duration) -> Self {
self.rate_limit_max_wait = max_wait;
self
}
pub fn stream_turn<'a, C: ApiClient>(
&'a self,
client: &'a C,
request: &'a crate::api::StreamRequest,
options: crate::structured::RequestOptions,
cancel: &'a Arc<CancelSignal>,
) -> Pin<Box<dyn Stream<Item = Result<HandlerEvent, StreamHandlerError>> + Send + 'a>> {
let total_deadline = Instant::now().checked_add(self.timeout_config.total_stream_timeout);
let stream_start = Instant::now();
let max_attempts = self.retry_config.max_retries.saturating_add(1);
Box::pin(async_stream::try_stream! {
let mut rate_limit_retries: u32 = 0;
let mut transport_attempts: u32 = 0;
let mut first_attempt = true;
let mut shadow = StreamAccumulator::new();
loop {
if !first_attempt {
shadow = StreamAccumulator::new();
yield HandlerEvent::AttemptReset;
}
first_attempt = false;
self.gate_on_rate_limit(client, cancel, total_deadline).await?;
let mut stream =
client.stream_messages_with_options(request, options.clone());
let mut consecutive_timeouts: usize = 0;
let mut events_processed: u64 = 0;
let mut saw_terminal = false;
let action = loop {
let diagnostics = EventDiagnostics::new(
events_processed,
stream_start,
&shadow,
transport_attempts
.saturating_add(rate_limit_retries)
.saturating_add(1),
);
match self
.next_event(
&mut stream,
cancel,
&mut consecutive_timeouts,
total_deadline,
&diagnostics,
)
.await
{
Ok(Some(event)) => {
events_processed = events_processed.saturating_add(1);
consecutive_timeouts = 0;
if matches!(event, StreamEvent::MessageStop) {
saw_terminal = true;
}
if let Err(failure) =
Self::accumulate_event(&diagnostics, &mut shadow, &event)
{
break self.decide_failure_action(
failure,
&mut rate_limit_retries,
&mut transport_attempts,
max_attempts,
total_deadline,
);
}
yield HandlerEvent::Stream(event);
}
Ok(None) => {
if saw_terminal {
return;
}
let failure = StreamFailure::transient(
StreamHandlerError::StreamFailed(StreamOutcome::InitFailed {
attempts: diagnostics.attempts_so_far,
last_error: format!(
"stream ended without a terminal event after \
{events_processed} events (truncated?)"
),
}),
);
break self.decide_failure_action(
failure,
&mut rate_limit_retries,
&mut transport_attempts,
max_attempts,
total_deadline,
);
}
Err(failure) => break self.decide_failure_action(
failure,
&mut rate_limit_retries,
&mut transport_attempts,
max_attempts,
total_deadline,
),
}
};
match action {
ErrorAction::Fail(e) => {
Err(e)?;
return;
}
ErrorAction::TryFallback(outcome) => {
let (message, stop_reason, usage) = self
.fallback_non_streaming(
client,
request,
&options,
cancel,
total_deadline,
outcome,
)
.await?;
yield HandlerEvent::Fallback {
message,
stop_reason,
usage,
};
return;
}
ErrorAction::Retry(delay) => {
sleep_cancellable(delay, cancel).await?;
}
}
}
})
}
fn decide_failure_action(
&self,
failure: StreamFailure,
rate_limit_retries: &mut u32,
transport_attempts: &mut u32,
max_attempts: u32,
total_deadline: Option<Instant>,
) -> ErrorAction {
let last_stream_outcome = carried_outcome(&failure.error);
if let Some(StreamOutcome::RateLimited { detail, .. }) = &last_stream_outcome {
self.decide_rate_limit_error(
failure.error,
detail,
rate_limit_retries,
total_deadline,
last_stream_outcome.clone(),
)
} else {
self.decide_transport_error(
failure.error,
failure.retryable,
transport_attempts,
max_attempts,
last_stream_outcome.clone(),
total_deadline,
)
}
}
fn accumulate_event(
diagnostics: &EventDiagnostics,
shadow: &mut StreamAccumulator,
event: &StreamEvent,
) -> Result<(), StreamFailure> {
shadow.process(event).map_err(|e| {
tracing::warn!(
error = %e,
attempts = diagnostics.attempts_so_far,
events_processed = diagnostics.events_processed,
"malformed accumulator event rejected"
);
StreamFailure::transient(StreamHandlerError::StreamFailed(
StreamOutcome::InitFailed {
attempts: diagnostics.attempts_so_far,
last_error: e.to_string(),
},
))
})
}
fn decide_rate_limit_error(
&self,
err: StreamHandlerError,
detail: &DetectedRateLimit,
rate_limit_retries: &mut u32,
total_deadline: Option<Instant>,
last_outcome: Option<StreamOutcome>,
) -> ErrorAction {
match self.rate_limit_retry(detail, rate_limit_retries, total_deadline) {
RateLimitRetry::Escalate {
attempts,
retry_after,
} => ErrorAction::Fail(StreamHandlerError::RateLimitEscalation {
attempts,
retry_after,
}),
RateLimitRetry::HardStop => {
if self.timeout_config.fallback_to_non_streaming {
ErrorAction::TryFallback(last_outcome)
} else {
ErrorAction::Fail(err)
}
}
RateLimitRetry::Retry(delay) => ErrorAction::Retry(delay),
}
}
fn decide_transport_error(
&self,
err: StreamHandlerError,
retryable: bool,
transport_attempts: &mut u32,
max_attempts: u32,
last_outcome: Option<StreamOutcome>,
total_deadline: Option<Instant>,
) -> ErrorAction {
if !retryable {
return ErrorAction::Fail(err);
}
if matches!(last_outcome, Some(StreamOutcome::TotalTimeout { .. })) {
if self.timeout_config.fallback_to_non_streaming {
return ErrorAction::TryFallback(last_outcome);
}
return ErrorAction::Fail(err);
}
if *transport_attempts >= max_attempts.saturating_sub(1) {
if self.timeout_config.fallback_to_non_streaming {
return ErrorAction::TryFallback(last_outcome);
}
return ErrorAction::Fail(err);
}
let delay = self.retry_config.jittered_base_delay(*transport_attempts);
let delay = clamp_delay_to_deadline(delay, total_deadline);
*transport_attempts = transport_attempts.saturating_add(1);
ErrorAction::Retry(delay)
}
fn rate_limit_retry(
&self,
detail: &DetectedRateLimit,
count: &mut u32,
deadline: Option<Instant>,
) -> RateLimitRetry {
*count = count.saturating_add(1);
if *count > self.rate_limit_config.max_retries {
return RateLimitRetry::HardStop;
}
if *count > self.rate_limit_config.fallback_after_retries {
return RateLimitRetry::Escalate {
attempts: *count,
retry_after: detail.retry_after,
};
}
let delay =
clamp_delay_to_deadline(self.rate_limit_config.backoff(detail.retry_after), deadline);
RateLimitRetry::Retry(delay)
}
async fn gate_on_rate_limit<C: ApiClient>(
&self,
client: &C,
cancel: &Arc<CancelSignal>,
total_deadline: Option<Instant>,
) -> Result<(), StreamHandlerError> {
let Some(limiter) = &self.rate_limiter else {
return Ok(());
};
let key = client.base_url();
let max_wait = self.rate_limit_max_wait;
let mut waited = Duration::ZERO;
loop {
match limiter.acquire(&key) {
Ok(()) => return Ok(()),
Err(rate_limit::RateLimitError::Poisoned) => {
tracing::warn!("rate-limit bucket poisoned; pacing unavailable");
return Err(StreamHandlerError::Poisoned("rate_limit"));
}
Err(rate_limit::RateLimitError::Wait(wait)) => {
if waited >= max_wait {
return Ok(());
}
let max_wait_remaining = max_wait.checked_sub(waited).unwrap_or(Duration::ZERO);
let total_deadline_remaining = match total_deadline {
None => max_wait_remaining,
Some(deadline) => deadline
.checked_duration_since(Instant::now())
.unwrap_or(Duration::ZERO),
};
let capped = wait.min(max_wait_remaining).min(total_deadline_remaining);
if capped.is_zero() {
return Ok(());
}
tokio::select! {
() = tokio::time::sleep(capped) => {}
() = cancel.notified() => return Err(StreamHandlerError::Cancelled),
}
waited = waited.saturating_add(capped);
}
}
}
}
async fn next_event<S>(
&self,
stream: &mut S,
cancel: &Arc<CancelSignal>,
consecutive_timeouts: &mut usize,
total_deadline: Option<Instant>,
diagnostics: &EventDiagnostics,
) -> Result<Option<StreamEvent>, StreamFailure>
where
S: futures::Stream<Item = Result<crate::stream::StreamEvent, crate::api::error::ApiError>>
+ Unpin,
{
loop {
if Self::deadline_exceeded(total_deadline) {
return Err(StreamFailure::transient(StreamHandlerError::StreamFailed(
diagnostics.total_timeout(),
)));
}
if cancel.is_cancelled() {
return Err(StreamFailure {
error: StreamHandlerError::Cancelled,
retryable: false,
});
}
let event_deadline = self.event_deadline(diagnostics.events_processed);
let event_result = tokio::select! {
event = stream.next() => EventPoll::Next(event),
() = cancel.notified() => return Err(StreamFailure {
error: StreamHandlerError::Cancelled,
retryable: false,
}),
() = deadline_future(event_deadline) => EventPoll::TimedOut,
() = deadline_future(total_deadline) => {
return Err(StreamFailure::transient(
StreamHandlerError::StreamFailed(diagnostics.total_timeout()),
));
}
};
match event_result {
EventPoll::TimedOut => {
*consecutive_timeouts = consecutive_timeouts.saturating_add(1);
let max_consecutive = if diagnostics.events_processed == 0 {
self.timeout_config.max_consecutive_timeouts.min(2) as usize
} else {
self.timeout_config.max_consecutive_timeouts as usize
};
if *consecutive_timeouts >= max_consecutive {
return Err(StreamFailure::transient(StreamHandlerError::StreamFailed(
diagnostics.event_timeout(
u32::try_from(*consecutive_timeouts).unwrap_or(u32::MAX),
),
)));
}
}
EventPoll::Next(Some(Ok(event))) => return Ok(Some(event)),
EventPoll::Next(Some(Err(api_error))) => {
return Err(diagnostics.api_error_failure(&api_error));
}
EventPoll::Next(None) => return Ok(None),
}
}
}
fn event_deadline(&self, events_processed: u64) -> Option<Instant> {
let base_timeout = if events_processed == 0 {
self.timeout_config.initial_event_timeout
} else {
self.timeout_config.per_event_timeout
};
Instant::now().checked_add(base_timeout)
}
fn deadline_exceeded(total_deadline: Option<Instant>) -> bool {
match total_deadline {
Some(deadline) => Instant::now() >= deadline,
None => false,
}
}
async fn fallback_non_streaming<C: ApiClient>(
&self,
client: &C,
request: &crate::api::StreamRequest,
options: &crate::structured::RequestOptions,
cancel: &Arc<CancelSignal>,
total_deadline: Option<Instant>,
stream_outcome: Option<StreamOutcome>,
) -> Result<(Message, StreamStopReason, Option<Usage>), StreamHandlerError> {
if cancel.is_cancelled() {
return Err(StreamHandlerError::Cancelled);
}
let fallback_deadline = match total_deadline {
Some(deadline) if deadline > Instant::now() => total_deadline,
Some(_) => Instant::now().checked_add(self.timeout_config.initial_event_timeout),
None => None,
};
let result = tokio::select! {
biased;
() = cancel.notified() => {
return Err(StreamHandlerError::Cancelled);
}
res = client.create_message_with_options(request, options.clone()) => res,
() = deadline_future(fallback_deadline) => {
return Err(StreamHandlerError::FallbackFailed {
stream_outcome: stream_outcome.unwrap_or(StreamOutcome::InitFailed {
attempts: 0,
last_error: "unknown".to_string(),
}),
fallback_error: "fallback request exceeded its deadline".to_string(),
});
}
};
match result {
Ok(response) => Ok((response.message, response.stop_reason, response.usage)),
Err(e) => Err(StreamHandlerError::FallbackFailed {
stream_outcome: stream_outcome.unwrap_or(StreamOutcome::InitFailed {
attempts: 0,
last_error: "unknown".to_string(),
}),
fallback_error: e.to_string(),
}),
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum HandlerEvent {
Stream(StreamEvent),
AttemptReset,
Fallback {
message: Message,
stop_reason: StreamStopReason,
usage: Option<Usage>,
},
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug)]
#[allow(dead_code)]
struct DriveResult {
message: Message,
usage: Option<Usage>,
stop_reason: StreamStopReason,
from_fallback: bool,
}
impl StreamHandler {
async fn drive_turn<C: ApiClient>(
&self,
client: &C,
request: &crate::api::StreamRequest,
cancel: &Arc<CancelSignal>,
) -> Result<DriveResult, StreamHandlerError> {
let mut stream = self.stream_turn(
client,
request,
crate::structured::RequestOptions::default(),
cancel,
);
let mut accumulator = StreamAccumulator::new();
let mut stop_reason = StreamStopReason::EndTurn;
let mut from_fallback = false;
while let Some(item) = stream.next().await {
match item? {
HandlerEvent::Stream(ev) => {
if let StreamEvent::MessageDelta(delta) = &ev
&& let Some(ref reason_str) = delta.delta.stop_reason
{
stop_reason =
StreamStopReason::from_api_str(reason_str).unwrap_or(stop_reason);
}
accumulator.process(&ev).map_err(|e| {
StreamHandlerError::StreamFailed(StreamOutcome::InitFailed {
attempts: 1,
last_error: e.to_string(),
})
})?;
}
HandlerEvent::AttemptReset => {
accumulator = StreamAccumulator::new();
stop_reason = StreamStopReason::EndTurn;
}
HandlerEvent::Fallback {
message,
stop_reason: fallback_stop_reason,
usage: fallback_usage,
} => {
from_fallback = true;
return Ok(DriveResult {
message,
usage: fallback_usage,
stop_reason: fallback_stop_reason,
from_fallback,
});
}
}
}
let usage = accumulator.usage().copied();
Ok(DriveResult {
message: accumulator.build(),
usage,
stop_reason,
from_fallback,
})
}
}
#[test]
fn timeout_config_default_values() {
let config = StreamTimeoutConfig::default();
assert_eq!(config.initial_event_timeout, Duration::from_mins(2));
assert_eq!(config.per_event_timeout, Duration::from_mins(3));
assert_eq!(config.total_stream_timeout, Duration::from_mins(5));
assert_eq!(config.max_consecutive_timeouts, 10);
assert!(config.fallback_to_non_streaming);
}
#[test]
fn passthrough_sets_no_resilience_config() {
let h = StreamHandler::passthrough();
assert_eq!(h.timeout_config().initial_event_timeout, Duration::MAX);
assert_eq!(h.timeout_config().per_event_timeout, Duration::MAX);
assert_eq!(h.timeout_config().total_stream_timeout, Duration::MAX);
assert!(!h.timeout_config().fallback_to_non_streaming);
assert_eq!(h.retry_config().max_retries, 0);
assert_eq!(h.rate_limit_config().max_retries, 0);
assert_eq!(h.rate_limit_config().fallback_after_retries, 0);
}
#[test]
fn passthrough_default_returns_shared_static() {
let a = StreamHandler::passthrough_default();
let b = StreamHandler::passthrough_default();
assert!(
std::ptr::eq(a, b),
"passthrough_default must return the same static"
);
}
#[test]
fn timeout_config_custom_values() {
let config = StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(30),
per_event_timeout: Duration::from_mins(1),
total_stream_timeout: Duration::from_mins(5),
max_consecutive_timeouts: 5,
fallback_to_non_streaming: false,
};
assert_eq!(config.initial_event_timeout, Duration::from_secs(30));
assert!(!config.fallback_to_non_streaming);
}
#[test]
fn retry_config_default_values() {
let config = StreamRetryConfig::default();
assert_eq!(config.max_retries, 3);
assert_eq!(config.base_delay_ms, 100);
assert_eq!(config.max_delay_ms, 10_000);
assert!((config.jitter_factor - 0.1).abs() < f64::EPSILON);
}
#[test]
fn retry_config_base_delay_exponential() {
let config = StreamRetryConfig::default();
assert_eq!(config.base_delay(0), Duration::from_millis(100));
assert_eq!(config.base_delay(1), Duration::from_millis(200));
assert_eq!(config.base_delay(2), Duration::from_millis(400));
assert_eq!(config.base_delay(3), Duration::from_millis(800));
}
#[test]
fn retry_config_base_delay_capped_at_max() {
let config = StreamRetryConfig {
base_delay_ms: 1000,
max_delay_ms: 5000,
..Default::default()
};
assert_eq!(config.base_delay(3), Duration::from_secs(5));
}
#[test]
fn jittered_base_delay_zero_jitter_equals_raw() {
let config = StreamRetryConfig {
jitter_factor: 0.0,
..Default::default()
};
for attempt in 0..5 {
assert_eq!(
config.jittered_base_delay(attempt),
config.base_delay(attempt),
"zero jitter must reproduce the raw backoff exactly"
);
}
}
#[test]
fn jittered_base_delay_stays_within_jitter_band() {
let config = StreamRetryConfig {
base_delay_ms: 100,
max_delay_ms: 100_000,
jitter_factor: 0.2,
..Default::default()
};
for attempt in 0..64 {
let base = config.base_delay(attempt);
let delay = config.jittered_base_delay(attempt);
let lo = base.mul_f64(0.8);
let hi = base.mul_f64(1.2);
assert!(
delay >= lo && delay <= hi,
"attempt {attempt}: jittered delay {delay:?} outside [{lo:?}, {hi:?}]"
);
}
}
#[test]
fn jittered_base_delay_concurrent_calls_produce_different_delays() {
let config = StreamRetryConfig {
base_delay_ms: 100,
max_delay_ms: 100_000,
jitter_factor: 0.5,
..Default::default()
};
let attempt = 1;
let mut delays: Vec<_> = (0..10)
.map(|_| config.jittered_base_delay(attempt))
.collect();
delays.sort();
delays.dedup();
assert!(
delays.len() > 1,
"concurrent calls with the same attempt must produce varied delays"
);
}
#[test]
fn jittered_base_delay_max_jitter_stays_non_negative() {
let config = StreamRetryConfig {
base_delay_ms: 100,
max_delay_ms: 100_000,
jitter_factor: 1.0,
..Default::default()
};
for attempt in 0..256 {
let delay = config.jittered_base_delay(attempt);
let hi = config.base_delay(attempt).mul_f64(2.0);
assert!(
delay <= hi,
"attempt {attempt}: delay {delay:?} exceeds 2x base under max jitter"
);
}
}
#[test]
fn outcome_completed_display() {
let outcome = StreamOutcome::Completed {
events_processed: 42,
duration: Duration::from_secs(5),
};
let s = outcome.to_string();
assert!(s.contains("42 events"));
assert!(s.contains("5.0s"));
}
#[test]
fn outcome_total_timeout_display() {
let outcome = StreamOutcome::TotalTimeout {
has_partial_data: true,
events_processed: 10,
duration: Duration::from_mins(15),
};
let s = outcome.to_string();
assert!(s.contains("partial data"));
assert!(s.contains("900.0s"));
}
#[test]
fn outcome_event_timeout_display() {
let outcome = StreamOutcome::EventTimeout {
has_partial_data: false,
consecutive_timeouts: 10,
};
let s = outcome.to_string();
assert!(s.contains("10 consecutive"));
assert!(!s.contains("partial data"));
}
#[test]
fn outcome_init_failed_display() {
let outcome = StreamOutcome::InitFailed {
last_error: "connection refused".to_string(),
attempts: 3,
};
let s = outcome.to_string();
assert!(s.contains("3 attempts"));
assert!(s.contains("connection refused"));
assert!(
!s.contains("init failed"),
"the historical variant name must not leak into the rendered \
message — a mid-stream truncation is not an init failure: {s}"
);
}
#[test]
fn outcome_fallback_display() {
let outcome = StreamOutcome::FallbackToNonStreaming;
let s = outcome.to_string();
assert!(s.contains("non-streaming"));
}
#[test]
fn outcome_cancelled_display() {
let outcome = StreamOutcome::Cancelled;
assert_eq!(outcome.to_string(), "cancelled");
}
#[test]
fn error_init_failed_display() {
let outcome = StreamOutcome::InitFailed {
last_error: "timeout".to_string(),
attempts: 3,
};
let err = StreamHandlerError::InitFailed(outcome);
let s = err.to_string();
assert!(
s.contains("stream failed before completing"),
"the historical variant name must not leak into the message: {s}"
);
}
#[test]
fn error_stream_failed_display() {
let outcome = StreamOutcome::EventTimeout {
has_partial_data: true,
consecutive_timeouts: 5,
};
let err = StreamHandlerError::StreamFailed(outcome);
let s = err.to_string();
assert!(s.contains("stream failed"));
}
#[test]
fn error_fallback_failed_display() {
let stream_outcome = StreamOutcome::TotalTimeout {
has_partial_data: false,
events_processed: 0,
duration: Duration::from_mins(15),
};
let err = StreamHandlerError::FallbackFailed {
stream_outcome,
fallback_error: "api error 429".to_string(),
};
let s = err.to_string();
assert!(s.contains("fallback also failed"));
assert!(s.contains("429"));
}
#[test]
fn error_cancelled_display() {
let err = StreamHandlerError::Cancelled;
assert_eq!(err.to_string(), "cancelled");
}
#[test]
fn error_rate_limit_escalation_display() {
let err = StreamHandlerError::RateLimitEscalation {
attempts: 4,
retry_after: Some(Duration::from_secs(5)),
};
let s = err.to_string();
assert!(s.contains("rate-limit escalation"), "got: {s}");
assert!(s.contains("4 retries"), "got: {s}");
assert!(
s.contains("5s"),
"should render the retry-after duration, got: {s}"
);
}
#[test]
fn handler_new_defaults() {
let handler = StreamHandler::new();
assert_eq!(
handler.timeout_config().initial_event_timeout,
Duration::from_mins(2),
);
assert_eq!(handler.retry_config().max_retries, 3);
}
#[test]
fn handler_with_timeout_and_retry_config() {
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_mins(1),
..Default::default()
})
.with_retry_config(StreamRetryConfig {
max_retries: 5,
..Default::default()
});
assert_eq!(
handler.timeout_config().initial_event_timeout,
Duration::from_mins(1),
);
assert_eq!(handler.retry_config().max_retries, 5);
}
#[test]
fn handler_default_trait() {
let handler = StreamHandler::default();
assert_eq!(
handler.timeout_config().initial_event_timeout,
Duration::from_mins(2),
);
}
#[test]
fn handler_debug_format() {
let handler = StreamHandler::new();
let debug = format!("{handler:?}");
assert!(debug.contains("StreamHandler"));
assert!(debug.contains("timeout_config"));
}
#[test]
fn timeout_config_validate_rejects_infinite_total_timeout() {
let config = StreamTimeoutConfig {
total_stream_timeout: Duration::MAX,
..Default::default()
};
let err = config
.validate()
.expect_err("Duration::MAX must be rejected");
assert!(
err.contains("finite"),
"the error must name the silent-disable hazard: {err}"
);
}
#[test]
fn timeout_config_validate_default_ok() {
assert!(StreamTimeoutConfig::default().validate().is_ok());
}
#[test]
fn timeout_config_validate_zero_initial() {
let config = StreamTimeoutConfig {
initial_event_timeout: Duration::ZERO,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("initial_event_timeout"));
}
#[test]
fn timeout_config_validate_zero_per_event() {
let config = StreamTimeoutConfig {
per_event_timeout: Duration::ZERO,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("per_event_timeout"));
}
#[test]
fn timeout_config_validate_zero_total() {
let config = StreamTimeoutConfig {
total_stream_timeout: Duration::ZERO,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("total_stream_timeout"));
}
#[test]
fn timeout_config_validate_total_less_than_initial() {
let config = StreamTimeoutConfig {
initial_event_timeout: Duration::from_mins(2),
total_stream_timeout: Duration::from_mins(1),
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("total_stream_timeout"));
assert!(err.contains("initial_event_timeout"));
}
#[test]
fn retry_config_validate_default_ok() {
assert!(StreamRetryConfig::default().validate().is_ok());
}
#[test]
fn retry_config_validate_zero_base_delay() {
let config = StreamRetryConfig {
base_delay_ms: 0,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("base_delay_ms"));
}
#[test]
fn retry_config_validate_zero_max_delay() {
let config = StreamRetryConfig {
max_delay_ms: 0,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("max_delay_ms"));
}
#[test]
fn retry_config_validate_max_less_than_base() {
let config = StreamRetryConfig {
base_delay_ms: 1000,
max_delay_ms: 500,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("max_delay_ms"));
assert!(err.contains("base_delay_ms"));
}
#[test]
fn retry_config_validate_jitter_nan() {
let config = StreamRetryConfig {
jitter_factor: f64::NAN,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("finite"));
}
#[test]
fn retry_config_validate_jitter_infinity() {
let config = StreamRetryConfig {
jitter_factor: f64::INFINITY,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("finite"));
}
#[test]
fn retry_config_validate_jitter_above_one() {
let config = StreamRetryConfig {
jitter_factor: 1.5,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("0.0..=1.0"));
}
#[test]
fn retry_config_validate_jitter_negative() {
let config = StreamRetryConfig {
jitter_factor: -0.1,
..Default::default()
};
let err = config.validate().unwrap_err();
assert!(err.contains("0.0..=1.0"));
}
#[test]
fn retry_config_validate_jitter_boundaries() {
let config = StreamRetryConfig {
jitter_factor: 0.0,
..Default::default()
};
assert!(config.validate().is_ok());
let config = StreamRetryConfig {
jitter_factor: 1.0,
..Default::default()
};
assert!(config.validate().is_ok());
}
use crate::api::error::ApiError;
use crate::stream::{
DeltaPart, IndexedDelta, MessageDelta, MessageDeltaPayload, MessageMetadata, MessageStart,
PartStart, StreamEvent, Usage,
};
fn happy_stream_events() -> Vec<Result<StreamEvent, ApiError>> {
vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_test".to_string(),
role: "assistant".to_string(),
model: "test-model".to_string(),
},
})),
Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(crate::stream::MessagePart::text("")),
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "hi".to_string(),
},
})),
Ok(StreamEvent::PartStop { index: None }),
Ok(StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("end_turn".to_string()),
},
usage: None,
})),
Ok(StreamEvent::MessageStop),
]
}
struct HandlerMock {
create_error: Option<String>,
create_response: Option<Message>,
}
impl HandlerMock {
fn new() -> Self {
Self {
create_error: None,
create_response: None,
}
}
fn with_text_response(mut self, text: &str) -> Self {
self.create_response = Some(Message::assistant(text));
self
}
fn with_create_error(mut self, msg: &str) -> Self {
self.create_error = Some(msg.to_string());
self
}
}
impl ApiClient for HandlerMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::iter(happy_stream_events()))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
if let Some(ref err) = self.create_error {
let err = err.clone();
return Box::pin(async move { Err(ApiError::api(&err)) });
}
let message = self
.create_response
.clone()
.unwrap_or_else(|| Message::assistant("default"));
Box::pin(async move {
Ok(crate::api::NonStreamingResponse {
message,
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
#[tokio::test]
async fn fallback_non_streaming_success() {
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: true,
..Default::default()
});
let client = HandlerMock::new().with_text_response("fallback works");
let cancel = Arc::new(CancelSignal::new());
let (message, stop_reason, usage) = handler
.fallback_non_streaming(
&client,
&crate::api::StreamRequest::new(vec![]),
&crate::structured::RequestOptions::default(),
&cancel,
None,
Some(StreamOutcome::InitFailed {
last_error: "stream failed".to_string(),
attempts: 3,
}),
)
.await
.expect("fallback should succeed");
let text: String = message
.parts
.iter()
.filter_map(|p| match p {
crate::stream::MessagePart::Text { text } => Some(text.clone()),
_ => None,
})
.collect();
assert!(text.contains("fallback works"), "got: {text:?}");
assert_eq!(stop_reason, StreamStopReason::EndTurn);
assert_eq!(usage, Some(Usage::default()));
}
#[tokio::test]
async fn fallback_non_streaming_cancelled_before_start() {
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: true,
..Default::default()
});
let client = HandlerMock::new().with_text_response("fallback works");
let cancel = Arc::new(CancelSignal::new());
cancel.cancel();
let err = handler
.fallback_non_streaming(
&client,
&crate::api::StreamRequest::new(vec![]),
&crate::structured::RequestOptions::default(),
&cancel,
None,
None,
)
.await
.expect_err("should fail on cancellation");
assert!(
matches!(err, StreamHandlerError::Cancelled),
"expected Cancelled, got: {err}"
);
}
struct OptionsRecordingMock {
seen: std::sync::Mutex<Vec<crate::structured::RequestOptions>>,
}
impl ApiClient for OptionsRecordingMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::empty())
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: Message::assistant("unused"),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
fn create_message_with_options(
&self,
_request: &crate::api::StreamRequest,
options: crate::structured::RequestOptions,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
self.seen.lock().unwrap().push(options);
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: Message::assistant("fallback works"),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
#[tokio::test]
async fn fallback_non_streaming_forwards_request_options() {
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: true,
..Default::default()
});
let client = OptionsRecordingMock {
seen: std::sync::Mutex::new(Vec::new()),
};
let cancel = Arc::new(CancelSignal::new());
let mut options = crate::structured::RequestOptions::default();
options.response_format = Some(crate::structured::ResponseFormat::new(
"probe",
serde_json::json!({"type": "object"}),
));
let (message, _stop, _usage) = handler
.fallback_non_streaming(
&client,
&crate::api::StreamRequest::new(vec![]),
&options,
&cancel,
None,
Some(StreamOutcome::InitFailed {
last_error: "stream failed".to_string(),
attempts: 1,
}),
)
.await
.expect("fallback should succeed");
assert!(
message.text_content().contains("fallback works"),
"the options-aware response is the one used"
);
let seen = client.seen.lock().unwrap();
assert_eq!(seen.len(), 1, "exactly one options-aware call");
assert!(
seen[0]
.response_format
.as_ref()
.is_some_and(|format| format.name == "probe"),
"the fallback must receive the turn's RequestOptions verbatim"
);
}
struct HangingFallbackMock;
impl ApiClient for HangingFallbackMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::empty())
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(std::future::pending())
}
}
#[tokio::test]
async fn fallback_non_streaming_honors_the_total_deadline() {
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: true,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let deadline = Instant::now() + Duration::from_millis(10);
let err = handler
.fallback_non_streaming(
&HangingFallbackMock,
&crate::api::StreamRequest::new(vec![]),
&crate::structured::RequestOptions::default(),
&cancel,
Some(deadline),
None,
)
.await
.expect_err("a hanging fallback must be cut by the deadline");
match err {
StreamHandlerError::FallbackFailed { fallback_error, .. } => {
assert!(
fallback_error.contains("deadline"),
"the deadline arm must be the failure cause: {fallback_error}"
);
}
other => panic!("expected FallbackFailed, got: {other}"),
}
}
#[tokio::test]
async fn completed_fallback_response_racing_the_deadline_is_accepted() {
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: true,
..Default::default()
});
let client = HandlerMock::new().with_text_response("worth keeping");
let cancel = Arc::new(CancelSignal::new());
let deadline = Instant::now()
.checked_sub(Duration::from_millis(1))
.expect("a past instant");
let (message, _stop_reason, _usage) = handler
.fallback_non_streaming(
&client,
&crate::api::StreamRequest::new(vec![]),
&crate::structured::RequestOptions::default(),
&cancel,
Some(deadline),
None,
)
.await
.expect("a completed response outranks the expired deadline");
assert!(
message.text_content().contains("worth keeping"),
"the completed response is returned, not discarded"
);
}
#[tokio::test]
async fn fallback_non_streaming_error() {
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: true,
..Default::default()
});
let client = HandlerMock::new().with_create_error("service unavailable");
let cancel = Arc::new(CancelSignal::new());
let err = handler
.fallback_non_streaming(
&client,
&crate::api::StreamRequest::new(vec![]),
&crate::structured::RequestOptions::default(),
&cancel,
None,
Some(StreamOutcome::InitFailed {
last_error: "stream timeout".to_string(),
attempts: 2,
}),
)
.await
.expect_err("should fail when fallback also errors");
match err {
StreamHandlerError::FallbackFailed {
stream_outcome,
fallback_error,
} => {
let stream_s = stream_outcome.to_string();
assert!(
stream_s.contains("stream timeout"),
"unexpected: {stream_s}"
);
assert!(
fallback_error.contains("service unavailable"),
"unexpected: {fallback_error}"
);
}
other => panic!("expected FallbackFailed, got: {other}"),
}
}
#[tokio::test]
async fn stream_turn_yields_handler_event_stream_per_event() {
let handler = StreamHandler::new();
let client = HandlerMock::new().with_text_response("hello");
let cancel = Arc::new(CancelSignal::new());
let req = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&client,
&req,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut saw_stream_events = 0;
let mut saw_attempt_reset = false;
let mut saw_fallback = false;
while let Some(item) = stream.next().await {
match item.expect("stream item ok") {
HandlerEvent::Stream(_) => saw_stream_events += 1,
HandlerEvent::AttemptReset => saw_attempt_reset = true,
HandlerEvent::Fallback { .. } => saw_fallback = true,
}
}
assert!(saw_stream_events > 0, "should yield Stream events");
assert!(!saw_attempt_reset, "happy path must not emit AttemptReset");
assert!(!saw_fallback, "happy path must not emit Fallback");
}
#[tokio::test]
async fn empty_stream_fast_fails_after_lower_threshold() {
struct NeverYieldingMock;
impl ApiClient for NeverYieldingMock {
fn model(&self) -> String {
"stuck".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::pending())
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(10),
per_event_timeout: Duration::from_millis(10),
total_stream_timeout: Duration::from_secs(10),
max_consecutive_timeouts: 10,
fallback_to_non_streaming: false,
})
.with_retry_config(StreamRetryConfig {
max_retries: 0,
..Default::default()
});
let client = NeverYieldingMock;
let cancel = Arc::new(CancelSignal::new());
let req = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&client,
&req,
crate::structured::RequestOptions::default(),
&cancel,
);
let start = Instant::now();
let mut got = None;
while let Some(item) = stream.next().await {
if item.is_err() {
got = Some(item);
break;
}
}
let elapsed = start.elapsed();
match got.expect("stream must terminate with an error") {
Err(StreamHandlerError::StreamFailed(StreamOutcome::EventTimeout { .. })) => {}
other => panic!("expected EventTimeout on dead stream, got {other:?}"),
}
assert!(
elapsed < Duration::from_millis(60),
"empty-stream fast-fail (2×10ms) must beat the full threshold (10×10ms); \
elapsed {elapsed:?}",
);
}
#[tokio::test]
async fn fallback_preserves_tool_call_parts() {
struct ToolFallbackMock;
impl ApiClient for ToolFallbackMock {
fn model(&self) -> String {
"test".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::api("connection refused"))
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::new(
crate::message::Role::Assistant,
vec![
crate::message::MessagePart::text("Let me search"),
crate::message::MessagePart::tool_call(
"tc_1",
"search",
serde_json::json!({"q": "hello"}),
),
],
),
stop_reason: crate::stream::StreamStopReason::ToolCall,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: true,
..Default::default()
})
.with_retry_config(StreamRetryConfig {
max_retries: 0,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let req = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&ToolFallbackMock,
&req,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut got_fallback = false;
while let Some(item) = stream.next().await {
if let Ok(HandlerEvent::Fallback { message, .. }) = item {
got_fallback = true;
let has_tool = message
.parts
.iter()
.any(|p| matches!(p, crate::message::MessagePart::ToolCall { name, .. } if name == "search"));
assert!(
has_tool,
"fallback message must preserve the tool-call part, got: {:?}",
message.parts
);
let has_text = message
.parts
.iter()
.any(|p| matches!(p, crate::message::MessagePart::Text { text } if text == "Let me search"));
assert!(has_text, "fallback message must preserve the text part");
}
}
assert!(got_fallback, "must emit a Fallback event");
}
#[tokio::test]
async fn rate_limit_hard_stop_tries_fallback_when_enabled() {
struct RateLimitThenOkMock;
impl ApiClient for RateLimitThenOkMock {
fn model(&self) -> String {
"test".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::RateLimit {
retry_after: None,
message: "slow down".into(),
})
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant("fallback ok"),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new().with_rate_limit_config(RateLimitConfig {
fallback_after_retries: 2,
max_retries: 2,
default_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(1),
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let req = crate::api::StreamRequest::new(vec![]);
let result = handler
.drive_turn(&RateLimitThenOkMock, &req, &cancel)
.await;
assert!(
result.is_ok(),
"hard-stop must try fallback when enabled, got: {:?}",
result.err()
);
let drive = result.unwrap();
assert!(drive.from_fallback);
assert!(drive.message.text_content().contains("fallback ok"));
}
struct RetryingMock {
attempts: Arc<std::sync::atomic::AtomicUsize>,
}
impl ApiClient for RetryingMock {
fn model(&self) -> String {
"retry-test".to_string()
}
fn base_url(&self) -> String {
"retry-test".to_string()
}
fn set_model(&self, _: &str) -> bool {
false
}
fn stream_messages(
&self,
request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
self.stream_messages_with_options(request, crate::structured::RequestOptions::default())
}
fn create_message(
&self,
request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
self.create_message_with_options(request, crate::structured::RequestOptions::default())
}
fn stream_messages_with_options(
&self,
_request: &crate::api::StreamRequest,
_options: crate::structured::RequestOptions,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
use std::sync::atomic::Ordering;
let n = self.attempts.fetch_add(1, Ordering::SeqCst);
if n == 0 {
Box::pin(futures::stream::iter(vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: String::new(),
role: "assistant".into(),
model: String::new(),
},
})),
Err(ApiError::api("transient")),
]))
} else {
Box::pin(futures::stream::iter(happy_stream_events()))
}
}
fn create_message_with_options(
&self,
_request: &crate::api::StreamRequest,
_options: crate::structured::RequestOptions,
) -> Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
fn extract_structured(&self, _: &crate::message::Message) -> serde_json::Value {
serde_json::Value::Null
}
}
#[tokio::test]
async fn clean_first_attempt_emits_no_attempt_reset() {
struct OneShotMock;
impl ApiClient for OneShotMock {
fn model(&self) -> String {
"one-shot".to_string()
}
fn base_url(&self) -> String {
"one-shot".to_string()
}
fn set_model(&self, _: &str) -> bool {
false
}
fn stream_messages(
&self,
request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
self.stream_messages_with_options(
request,
crate::structured::RequestOptions::default(),
)
}
fn create_message(
&self,
request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
self.create_message_with_options(
request,
crate::structured::RequestOptions::default(),
)
}
fn stream_messages_with_options(
&self,
_request: &crate::api::StreamRequest,
_options: crate::structured::RequestOptions,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
Box::pin(futures::stream::iter(happy_stream_events()))
}
fn create_message_with_options(
&self,
_request: &crate::api::StreamRequest,
_options: crate::structured::RequestOptions,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
fn extract_structured(&self, _: &crate::message::Message) -> serde_json::Value {
serde_json::Value::Null
}
}
let handler = StreamHandler::new();
let cancel = Arc::new(CancelSignal::new());
let req = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&OneShotMock,
&req,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut events_seen = 0usize;
while let Some(item) = stream.next().await {
events_seen += 1;
assert!(
!matches!(item.expect("clean stream item"), HandlerEvent::AttemptReset),
"a clean first attempt never announces a reset"
);
}
assert!(
events_seen > 0,
"the silence assertion only counts on a stream that produced events"
);
}
#[tokio::test]
async fn stream_turn_yields_attempt_reset_on_retry() {
use std::sync::atomic::AtomicUsize;
let attempts = Arc::new(AtomicUsize::new(0));
let handler = StreamHandler::new();
let client = RetryingMock { attempts };
let cancel = Arc::new(CancelSignal::new());
let req = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&client,
&req,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut saw_attempt_reset = false;
while let Some(item) = stream.next().await {
if let HandlerEvent::AttemptReset = item.expect("stream item ok") {
saw_attempt_reset = true;
}
}
assert!(
saw_attempt_reset,
"second attempt must be preceded by AttemptReset"
);
}
#[tokio::test]
async fn stream_turn_happy_path() {
let handler = StreamHandler::new();
let client = HandlerMock::new().with_text_response("hello world");
let cancel = Arc::new(CancelSignal::new());
let result = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("stream_turn should succeed");
assert!(!result.from_fallback);
assert_eq!(result.stop_reason, StreamStopReason::EndTurn);
}
#[tokio::test]
async fn stream_turn_cancelled_at_start() {
let handler = StreamHandler::new();
let client = HandlerMock::new().with_text_response("hello");
let cancel = Arc::new(CancelSignal::new());
cancel.cancel();
let err = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect_err("should fail on cancellation");
assert!(
matches!(err, StreamHandlerError::Cancelled),
"expected Cancelled, got: {err}"
);
}
#[tokio::test]
async fn stream_turn_fallback_after_stream_error() {
struct ErrorMock;
impl ApiClient for ErrorMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::api("API down"))
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::api("unreachable")) })
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: false,
..Default::default()
})
.with_retry_config(StreamRetryConfig {
max_retries: 0,
..Default::default()
});
let client = ErrorMock;
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect_err("should fail when streaming errors and fallback is disabled");
match err {
StreamHandlerError::StreamFailed(outcome) => {
let s = outcome.to_string();
assert!(s.contains("API down"), "unexpected: {s}");
}
other => panic!("expected StreamFailed, got: {other}"),
}
}
struct StreamingFailingFallbackMock;
impl ApiClient for StreamingFailingFallbackMock {
fn model(&self) -> String {
"fallback-test".to_string()
}
fn base_url(&self) -> String {
"fallback-test".to_string()
}
fn set_model(&self, _: &str) -> bool {
false
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
Box::pin(futures::stream::once(async {
Err(ApiError::api("stream down"))
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant("fallback answer"),
stop_reason: crate::stream::StreamStopReason::MaxTokens,
usage: Some(crate::stream::Usage::new(42, 13)),
})
})
}
fn stream_messages_with_options(
&self,
_request: &crate::api::StreamRequest,
_options: crate::structured::RequestOptions,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
Box::pin(futures::stream::once(async {
Err(ApiError::api("stream down"))
}))
}
fn create_message_with_options(
&self,
_request: &crate::api::StreamRequest,
_options: crate::structured::RequestOptions,
) -> Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant("fallback answer"),
stop_reason: crate::stream::StreamStopReason::MaxTokens,
usage: Some(crate::stream::Usage::new(42, 13)),
})
})
}
fn extract_structured(&self, _: &crate::message::Message) -> serde_json::Value {
serde_json::Value::Null
}
}
#[tokio::test]
async fn drive_turn_returns_fallback_message_and_stop_reason() {
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: true,
..Default::default()
})
.with_retry_config(StreamRetryConfig {
max_retries: 0,
..Default::default()
});
let client = StreamingFailingFallbackMock;
let cancel = Arc::new(CancelSignal::new());
let result = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("fallback should succeed");
assert!(
result.from_fallback,
"result should be marked from_fallback"
);
assert_eq!(
result.stop_reason,
StreamStopReason::MaxTokens,
"fallback stop_reason must come from the JSON response"
);
let text: String = result
.message
.parts
.iter()
.filter_map(|p| match p {
crate::stream::MessagePart::Text { text } => Some(text.clone()),
_ => None,
})
.collect();
assert!(
text.contains("fallback answer"),
"fallback message text, got {text:?}"
);
assert_eq!(
result.usage,
Some(Usage::new(42, 13)),
"fallback path must propagate usage from the non-streaming response"
);
}
#[test]
fn rate_limit_config_default_values() {
let cfg = RateLimitConfig::default();
assert!(cfg.respect_retry_after);
assert_eq!(cfg.default_delay, Duration::from_secs(5));
assert_eq!(cfg.max_delay, Duration::from_mins(1));
assert_eq!(cfg.requests_per_minute, 0);
assert_eq!(cfg.fallback_after_retries, 3);
assert_eq!(cfg.max_retries, 5);
}
#[test]
fn rate_limit_config_validate_rejects_invalid() {
assert!(RateLimitConfig::default().validate().is_ok());
assert!(
RateLimitConfig {
default_delay: Duration::ZERO,
..Default::default()
}
.validate()
.is_err()
);
assert!(
RateLimitConfig {
max_delay: Duration::from_secs(1),
default_delay: Duration::from_secs(10),
..Default::default()
}
.validate()
.is_err()
);
assert!(
RateLimitConfig {
max_retries: 0,
..Default::default()
}
.validate()
.is_err()
);
}
#[test]
fn with_timeout_config_substitutes_only_invalid_fields() {
let bad_total = StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(45),
per_event_timeout: Duration::from_secs(45),
total_stream_timeout: Duration::MAX,
max_consecutive_timeouts: 7,
..Default::default()
};
let handler = StreamHandler::new().with_timeout_config(bad_total);
let config = handler.timeout_config();
assert_eq!(
config.initial_event_timeout,
Duration::from_secs(45),
"valid fields the caller supplied must survive an invalid sibling"
);
assert_eq!(config.per_event_timeout, Duration::from_secs(45));
assert_eq!(config.max_consecutive_timeouts, 7);
assert_eq!(
config.total_stream_timeout,
StreamTimeoutConfig::default().total_stream_timeout,
"an infinite total timeout is substituted with the default, not silently disabling every deadline"
);
let unordered = StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(400),
..Default::default()
};
let handler = StreamHandler::new().with_timeout_config(unordered);
assert_eq!(
handler.timeout_config().total_stream_timeout,
Duration::from_secs(400),
"a default total below a custom initial timeout is raised to it, \
keeping the caller's initial customization"
);
}
#[test]
fn sanitized_config_always_validates() {
let adversarial = [
StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(600),
total_stream_timeout: Duration::ZERO,
..Default::default()
},
StreamTimeoutConfig {
initial_event_timeout: Duration::MAX,
..Default::default()
},
StreamTimeoutConfig {
per_event_timeout: Duration::MAX,
..Default::default()
},
StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(45),
per_event_timeout: Duration::from_secs(45),
total_stream_timeout: Duration::MAX,
max_consecutive_timeouts: 7,
..Default::default()
},
];
for config in adversarial {
let handler = StreamHandler::new().with_timeout_config(config);
assert!(
handler.timeout_config().validate().is_ok(),
"the sanitized builder output must satisfy every validate rule: {:?}",
handler.timeout_config()
);
}
let zero_total = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(600),
total_stream_timeout: Duration::ZERO,
..Default::default()
});
assert_eq!(
zero_total.timeout_config().total_stream_timeout,
Duration::from_secs(600),
"a repaired total must still honor the ordering rule against a large initial"
);
}
#[test]
fn handler_error_display_never_says_init_failed() {
let outcome = StreamOutcome::InitFailed {
last_error: "stream ended without a terminal event".to_string(),
attempts: 3,
};
let rendered = StreamHandlerError::InitFailed(outcome).to_string();
assert!(
rendered.contains("without a terminal event"),
"the wrapper names the failure cause: {rendered}"
);
assert!(
!rendered.contains("init failed"),
"the historical variant name must not leak into the rendered message: {rendered}"
);
}
#[test]
fn with_timeout_config_keeps_valid() {
let good = StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(45),
..Default::default()
};
let handler = StreamHandler::new().with_timeout_config(good);
assert_eq!(
handler.timeout_config().initial_event_timeout,
Duration::from_secs(45)
);
}
#[test]
fn with_retry_config_rejects_invalid_falls_back_to_default() {
let bad = StreamRetryConfig {
base_delay_ms: 0,
..Default::default()
};
let handler = StreamHandler::new().with_retry_config(bad);
assert_eq!(
handler.retry_config().base_delay_ms,
StreamRetryConfig::default().base_delay_ms,
"invalid retry config must fall back to default"
);
}
#[test]
fn with_retry_config_keeps_valid() {
let good = StreamRetryConfig {
max_retries: 7,
..Default::default()
};
let handler = StreamHandler::new().with_retry_config(good);
assert_eq!(handler.retry_config().max_retries, 7);
}
#[test]
fn with_timeout_and_retry_config_are_independent() {
let good_timeout = StreamTimeoutConfig {
initial_event_timeout: Duration::from_mins(1),
..Default::default()
};
let bad_retry = StreamRetryConfig {
jitter_factor: 2.0,
..Default::default()
};
let handler = StreamHandler::new()
.with_timeout_config(good_timeout)
.with_retry_config(bad_retry);
assert_eq!(
handler.timeout_config().initial_event_timeout,
Duration::from_mins(1),
"valid timeout must be kept when retry config is invalid"
);
assert_eq!(
handler.retry_config().max_retries,
StreamRetryConfig::default().max_retries,
"invalid retry config must fall back to default"
);
}
#[test]
fn with_rate_limit_config_rejects_invalid_falls_back_to_default() {
let bad = RateLimitConfig {
max_retries: 0,
..Default::default()
};
let handler = StreamHandler::new().with_rate_limit_config(bad);
assert_eq!(
handler.rate_limit_config().max_retries,
RateLimitConfig::default().max_retries,
"invalid rate-limit config must fall back to default"
);
}
#[test]
fn rate_limit_config_backoff_honours_hint_and_caps() {
let cfg = RateLimitConfig::default();
assert_eq!(
cfg.backoff(Some(Duration::from_secs(12))),
Duration::from_secs(12)
);
assert_eq!(
cfg.backoff(Some(Duration::from_mins(2))),
cfg.max_delay,
"should cap at max_delay"
);
assert_eq!(cfg.backoff(None), cfg.default_delay);
let ignore = RateLimitConfig {
respect_retry_after: false,
..Default::default()
};
assert_eq!(
ignore.backoff(Some(Duration::from_secs(12))),
ignore.default_delay
);
}
#[test]
fn clamp_delay_to_deadline_none_deadline_returns_delay_unchanged() {
let delay = Duration::from_mins(10);
assert_eq!(clamp_delay_to_deadline(delay, None), delay);
}
#[test]
fn clamp_delay_to_deadline_future_deadline_fits() {
let delay = Duration::from_millis(10);
let deadline = Some(Instant::now() + Duration::from_mins(1));
assert_eq!(clamp_delay_to_deadline(delay, deadline), delay);
}
#[test]
fn clamp_delay_to_deadline_exceeds_remaining() {
let delay = Duration::from_mins(10);
let remaining = Duration::from_millis(50);
let deadline = Some(Instant::now() + remaining);
let clamped = clamp_delay_to_deadline(delay, deadline);
assert!(
clamped <= remaining,
"clamped {clamped:?} must not exceed remaining {remaining:?}"
);
assert!(
!clamped.is_zero(),
"deadline still in the future, so sleep should be positive"
);
}
#[test]
fn clamp_delay_to_deadline_past_deadline_zero() {
let delay = Duration::from_mins(10);
let deadline = Some(Instant::now().checked_sub(Duration::from_secs(1)).unwrap());
assert_eq!(clamp_delay_to_deadline(delay, deadline), Duration::ZERO);
}
#[test]
fn backoff_clamps_huge_hint_to_max_delay() {
let cfg = RateLimitConfig {
max_delay: Duration::from_mins(1),
..Default::default()
};
assert_eq!(
cfg.backoff(Some(Duration::from_secs(9_999_999))),
Duration::from_mins(1)
);
}
fn detected_limit(retry_after: Option<Duration>) -> DetectedRateLimit {
DetectedRateLimit {
kind: RateLimitKind::RateLimited,
retry_after,
message: "slow down".to_string(),
}
}
#[test]
fn rate_limit_retry_returns_clamped_delay_below_ceilings() {
let handler = StreamHandler::new().with_rate_limit_config(RateLimitConfig {
fallback_after_retries: 3,
max_retries: 5,
default_delay: Duration::from_millis(1),
max_delay: Duration::from_mins(1),
..Default::default()
});
let mut count = 0u32;
let detail = detected_limit(Some(Duration::from_mins(10)));
let decision = handler.rate_limit_retry(&detail, &mut count, None);
assert_eq!(count, 1);
match decision {
RateLimitRetry::Retry(delay) => assert_eq!(delay, Duration::from_mins(1)),
other => panic!("expected Retry, got {other:?}"),
}
}
#[test]
fn rate_limit_retry_escalates_after_fallback_ceiling() {
let handler = StreamHandler::new().with_rate_limit_config(RateLimitConfig {
fallback_after_retries: 2,
max_retries: 5,
..Default::default()
});
let mut count = 0u32;
let detail = detected_limit(Some(Duration::from_millis(5)));
let _ = handler.rate_limit_retry(&detail, &mut count, None);
let _ = handler.rate_limit_retry(&detail, &mut count, None);
assert_eq!(count, 2);
let decision = handler.rate_limit_retry(&detail, &mut count, None);
assert_eq!(count, 3);
match decision {
RateLimitRetry::Escalate {
attempts,
retry_after,
} => {
assert_eq!(attempts, 3);
assert_eq!(retry_after, Some(Duration::from_millis(5)));
}
other => panic!("expected Escalate, got {other:?}"),
}
}
#[test]
fn rate_limit_retry_hard_stops_after_max_retries() {
let handler = StreamHandler::new().with_rate_limit_config(RateLimitConfig {
fallback_after_retries: 1,
max_retries: 2,
..Default::default()
});
let mut count = 0u32;
let detail = detected_limit(None);
let _ = handler.rate_limit_retry(&detail, &mut count, None);
let _ = handler.rate_limit_retry(&detail, &mut count, None);
assert_eq!(count, 2);
assert!(matches!(
handler.rate_limit_retry(&detail, &mut count, None),
RateLimitRetry::HardStop
));
assert_eq!(count, 3);
}
#[test]
fn rate_limit_retry_max_retries_should_be_enforced_under_valid_config() {
let handler = StreamHandler::new();
let detail = detected_limit(None);
let mut count = 0u32;
for _ in 0..(handler.rate_limit_config().max_retries + 2) {
let _ = handler.rate_limit_retry(&detail, &mut count, None);
}
let max = handler.rate_limit_config().max_retries;
assert!(
count > max,
"count {count} must exceed max_retries {max} after enough calls"
);
let decision = handler.rate_limit_retry(&detail, &mut count, None);
assert!(
matches!(decision, RateLimitRetry::HardStop),
"max_retries={max} should be enforced as a hard ceiling, \
but Escalate shadows it — HardStop is dead code under valid config"
);
}
#[test]
fn with_rate_limit_config_should_reject_invalid() {
let invalid = RateLimitConfig {
fallback_after_retries: 10,
max_retries: 3,
..Default::default()
};
let result = StreamHandler::new().with_rate_limit_config(invalid);
let detail = detected_limit(None);
let mut count = 0u32;
for _ in 0..4 {
let _ = result.rate_limit_retry(&detail, &mut count, None);
}
let decision = result.rate_limit_retry(&detail, &mut count, None);
assert!(
!matches!(decision, RateLimitRetry::HardStop),
"invalid config (fallback_after=10 > max_retries=3) must not \
silently invert behavior — HardStop should never fire before Escalation"
);
}
#[test]
fn detected_rate_limit_from_structured_variant() {
let err = ApiError::RateLimit {
retry_after: Some(Duration::from_secs(7)),
message: "slow down".into(),
};
let detected = DetectedRateLimit::detect(&err).expect("RateLimit variant should detect");
assert_eq!(detected.kind, RateLimitKind::RateLimited);
assert_eq!(detected.retry_after, Some(Duration::from_secs(7)));
}
#[test]
fn detected_rate_limit_from_structured_variant_no_hint() {
let err = ApiError::RateLimit {
retry_after: None,
message: "slow down".into(),
};
let detected = DetectedRateLimit::detect(&err).expect("RateLimit variant should detect");
assert_eq!(detected.retry_after, None);
}
#[test]
fn detected_rate_limit_from_http_503() {
let err = ApiError::http_with_status(503, "overloaded");
let detected = DetectedRateLimit::detect(&err).expect("503 should detect as Overloaded");
assert_eq!(detected.kind, RateLimitKind::Overloaded);
}
#[test]
fn detected_rate_limit_http_500_is_not_overload() {
let err = ApiError::http_with_status(500, "boom");
assert!(DetectedRateLimit::detect(&err).is_none());
}
#[test]
fn detected_rate_limit_non_rate_errors_return_none() {
assert!(DetectedRateLimit::detect(&ApiError::api("connection reset")).is_none());
assert!(DetectedRateLimit::detect(&ApiError::auth("bad key")).is_none());
}
#[test]
fn is_rate_limited_matches_detect() {
let cases: &[ApiError] = &[
ApiError::RateLimit {
retry_after: None,
message: "x".into(),
},
ApiError::http_with_status(503, "overloaded"),
ApiError::http_with_status(500, "boom"),
ApiError::api("connection reset"),
ApiError::auth("bad key"),
];
for err in cases {
assert_eq!(
err.is_rate_limited(),
DetectedRateLimit::detect(err).is_some(),
"is_rate_limited disagree with detect on {err}",
);
}
}
#[test]
fn stream_outcome_rate_limited_display() {
let outcome = StreamOutcome::RateLimited {
detail: DetectedRateLimit {
kind: RateLimitKind::RateLimited,
retry_after: Some(Duration::from_secs(12)),
message: "slow down".into(),
},
has_partial_data: false,
events_processed: 5,
};
let s = outcome.to_string();
assert!(s.contains("rate limit"), "got: {s}");
assert!(s.contains("12"), "retry-after seconds missing: {s}");
}
#[test]
fn stream_handler_rate_limit_config_round_trip() {
let handler = StreamHandler::new();
assert_eq!(
handler.rate_limit_config().max_retries,
RateLimitConfig::default().max_retries
);
let custom = RateLimitConfig {
max_retries: 2,
fallback_after_retries: 1,
default_delay: Duration::from_secs(1),
..Default::default()
};
let handler = StreamHandler::new().with_rate_limit_config(custom);
assert_eq!(handler.rate_limit_config().max_retries, 2);
assert_eq!(
handler.rate_limit_config().default_delay,
Duration::from_secs(1)
);
}
struct GateMock {
url: &'static str,
}
impl ApiClient for GateMock {
fn model(&self) -> String {
"gate-model".to_string()
}
fn base_url(&self) -> String {
self.url.to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::iter(happy_stream_events()))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
#[tokio::test]
async fn gate_on_rate_limit_noop_without_limiter() {
let handler = StreamHandler::new();
let client = GateMock { url: "openai" };
let cancel = Arc::new(CancelSignal::new());
let result = handler.gate_on_rate_limit(&client, &cancel, None).await;
assert!(result.is_ok(), "no limiter => gate is a no-op");
}
#[tokio::test]
async fn gate_on_rate_limit_full_bucket_acquires_immediately() {
use crate::stream::rate_limit::RateLimiter;
let limiter = Arc::new(RateLimiter::new(60));
let handler = StreamHandler::new().with_rate_limiter(Arc::clone(&limiter));
let client = GateMock { url: "openai" };
let cancel = Arc::new(CancelSignal::new());
let start = Instant::now();
handler
.gate_on_rate_limit(&client, &cancel, None)
.await
.expect("full bucket should acquire");
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_millis(100),
"full bucket should not wait; elapsed {elapsed:?}"
);
}
#[tokio::test]
async fn gate_on_rate_limit_respects_total_deadline() {
use crate::stream::rate_limit::RateLimiter;
let limiter = Arc::new(RateLimiter::new(1));
let handler = StreamHandler::new()
.with_rate_limiter(Arc::clone(&limiter))
.with_rate_limit_max_wait(Duration::from_mins(2));
let client = GateMock { url: "openai" };
let cancel = Arc::new(CancelSignal::new());
handler
.gate_on_rate_limit(&client, &cancel, None)
.await
.expect("first acquire should succeed (full bucket)");
let expired = Some(
Instant::now()
.checked_sub(Duration::from_secs(1))
.unwrap_or(Instant::now()),
);
let start = Instant::now();
handler
.gate_on_rate_limit(&client, &cancel, expired)
.await
.expect("gate should proceed on an expired deadline, not hang or spin");
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_millis(500),
"gate should proceed immediately on an expired deadline; elapsed {elapsed:?}"
);
}
#[tokio::test]
async fn gate_on_rate_limit_clamps_sleep_to_remaining_deadline() {
use crate::stream::rate_limit::RateLimiter;
let limiter = Arc::new(RateLimiter::new(1));
let handler = StreamHandler::new()
.with_rate_limiter(Arc::clone(&limiter))
.with_rate_limit_max_wait(Duration::from_mins(2));
let client = GateMock { url: "openai" };
let cancel = Arc::new(CancelSignal::new());
handler
.gate_on_rate_limit(&client, &cancel, None)
.await
.expect("first acquire should succeed (full bucket)");
let near_deadline = Some(
Instant::now()
.checked_add(Duration::from_millis(80))
.unwrap_or(Instant::now()),
);
let start = Instant::now();
handler
.gate_on_rate_limit(&client, &cancel, near_deadline)
.await
.expect("gate should proceed after clamping to the deadline");
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(2),
"gate should proceed within the ~80ms deadline window, not wait 60s; elapsed {elapsed:?}"
);
}
#[tokio::test]
async fn proactive_throttle_slows_burst() {
use crate::stream::rate_limit::RateLimiter;
let limiter = Arc::new(RateLimiter::new(60));
let handler = StreamHandler::new()
.with_rate_limiter(Arc::clone(&limiter))
.with_rate_limit_max_wait(Duration::from_mins(2));
let client = GateMock { url: "openai" };
let cancel = Arc::new(CancelSignal::new());
let start = Instant::now();
for _ in 0..3 {
handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("turn should succeed");
}
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(5),
"three turns from a 60-burst should be fast; elapsed {elapsed:?}"
);
assert!(limiter.is_enabled());
}
#[tokio::test]
async fn proactive_throttle_cancel_interrupts_wait() {
use crate::stream::rate_limit::RateLimiter;
let limiter = Arc::new(RateLimiter::new(1));
let handler = StreamHandler::new()
.with_rate_limiter(limiter)
.with_rate_limit_max_wait(Duration::from_millis(50));
let client = GateMock { url: "openai" };
let cancel = Arc::new(CancelSignal::new());
handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("first turn should succeed");
let cancel2 = Arc::new(CancelSignal::new());
cancel2.cancel();
let start = Instant::now();
let err = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel2)
.await
.expect_err("should be cancelled, not hang for 60s");
let elapsed = start.elapsed();
assert!(
matches!(err, StreamHandlerError::Cancelled),
"expected Cancelled, got {err:?}"
);
assert!(
elapsed < Duration::from_secs(2),
"cancel should interrupt the wait promptly; elapsed {elapsed:?}"
);
}
#[tokio::test]
async fn proactive_throttle_max_wait_clamp_degrades_to_reactive() {
use crate::stream::rate_limit::RateLimiter;
let limiter = Arc::new(RateLimiter::new(1));
let handler = StreamHandler::new()
.with_rate_limiter(limiter)
.with_rate_limit_max_wait(Duration::from_millis(50));
let client = GateMock { url: "openai" };
let cancel = Arc::new(CancelSignal::new());
handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("first turn should succeed");
let start = Instant::now();
let result = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await;
let elapsed = start.elapsed();
assert!(
result.is_ok(),
"max_wait clamp should let the turn proceed, got {result:?}"
);
assert!(
elapsed < Duration::from_secs(2),
"should proceed after ~50ms, not wait 60s; elapsed {elapsed:?}"
);
}
#[tokio::test]
async fn stream_turn_uses_rate_limit_delay_on_rate_limited_outcome() {
use std::sync::atomic::{AtomicUsize, Ordering};
struct RateLimitOnceMock {
attempts: AtomicUsize,
}
impl ApiClient for RateLimitOnceMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let n = self.attempts.fetch_add(1, Ordering::SeqCst);
if n == 0 {
Box::pin(futures::stream::once(async {
Err(ApiError::RateLimit {
retry_after: None,
message: "slow down".into(),
})
}))
} else {
Box::pin(futures::stream::iter(happy_stream_events()))
}
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new().with_retry_config(StreamRetryConfig {
max_retries: 1,
base_delay_ms: 2_000,
..Default::default()
});
let handler = handler.with_rate_limit_config(RateLimitConfig {
default_delay: Duration::from_millis(1),
..Default::default()
});
let client = RateLimitOnceMock {
attempts: AtomicUsize::new(0),
};
let cancel = Arc::new(CancelSignal::new());
let start = Instant::now();
let result = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("second attempt should succeed");
let elapsed = start.elapsed();
assert!(!result.from_fallback);
assert!(
elapsed < Duration::from_secs(1),
"rate-limit retry should use RateLimitConfig delay, not the 2s transport delay; elapsed {elapsed:?}",
);
}
#[tokio::test]
async fn stream_turn_escalates_after_rate_limit_threshold() {
struct AlwaysRateLimitMock;
impl ApiClient for AlwaysRateLimitMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::RateLimit {
retry_after: Some(Duration::from_millis(1)),
message: "slow down".into(),
})
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new().with_rate_limit_config(RateLimitConfig {
fallback_after_retries: 2,
default_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(1),
..Default::default()
});
let client = AlwaysRateLimitMock;
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect_err("should escalate, not succeed");
match err {
StreamHandlerError::RateLimitEscalation {
attempts,
retry_after,
} => {
assert_eq!(attempts, 3);
assert_eq!(retry_after, Some(Duration::from_millis(1)));
}
other => panic!("expected RateLimitEscalation, got {other:?}"),
}
}
#[tokio::test]
async fn default_rate_limit_config_escalates_without_a_non_streaming_attempt() {
struct Counting429Mock {
non_streaming_calls: std::sync::atomic::AtomicUsize,
}
impl ApiClient for Counting429Mock {
fn model(&self) -> String {
"test".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::RateLimit {
retry_after: None,
message: "slow down".into(),
})
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
self.non_streaming_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant("fallback ok"),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let client = Counting429Mock {
non_streaming_calls: std::sync::atomic::AtomicUsize::new(0),
};
let handler = StreamHandler::new().with_rate_limit_config(RateLimitConfig {
default_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(1),
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect_err("the default ladder escalates rather than exhausting");
assert!(
matches!(err, StreamHandlerError::RateLimitEscalation { .. }),
"the default ladder (fallback_after=3 < max=5) escalates to the model \
breaker, got: {err:?}"
);
assert_eq!(
client
.non_streaming_calls
.load(std::sync::atomic::Ordering::SeqCst),
0,
"a rate limit is charged against the model's quota — a same-model \
non-streaming request is deliberately not attempted (the ceiling-equal \
ladder opts into it)"
);
}
#[tokio::test]
async fn rate_limit_after_partial_data_reports_has_partial_data() {
struct PartialThenRateLimitMock;
impl ApiClient for PartialThenRateLimitMock {
fn model(&self) -> String {
"partial-then-429".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "partial-then-429".to_string(),
},
})),
Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(crate::stream::MessagePart::text("")),
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "partial".to_string(),
},
})),
Ok(StreamEvent::PartStop { index: Some(0) }),
Err(ApiError::RateLimit {
retry_after: None,
message: "slow down".into(),
}),
];
Box::pin(futures::stream::iter(events))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::api("unused")) })
}
}
let handler = StreamHandler::new()
.with_rate_limit_config(RateLimitConfig {
fallback_after_retries: 1,
max_retries: 1,
default_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(1),
..Default::default()
})
.with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: false,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(
&PartialThenRateLimitMock,
&crate::api::StreamRequest::new(vec![]),
&cancel,
)
.await
.expect_err("the disabled fallback makes the hard stop terminal");
match err {
StreamHandlerError::StreamFailed(StreamOutcome::RateLimited {
has_partial_data,
events_processed,
..
}) => {
assert!(
has_partial_data,
"a 429 after accepted events must report salvageable partial data"
);
assert_eq!(
events_processed, 4,
"the outcome counts the events that got through before the 429"
);
}
other => panic!("expected a RateLimited terminal, got {other:?}"),
}
}
#[tokio::test]
async fn event_timeout_after_partial_data_reports_has_partial_data() {
struct PartialThenHangMock;
impl ApiClient for PartialThenHangMock {
fn model(&self) -> String {
"partial-then-hang".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "partial-then-hang".to_string(),
},
})),
Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(crate::stream::MessagePart::text("")),
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "partial".to_string(),
},
})),
Ok(StreamEvent::PartStop { index: Some(0) }),
];
let pending = futures::stream::pending();
Box::pin(futures::stream::iter(events).chain(pending))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::api("unused")) })
}
}
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(50),
per_event_timeout: Duration::from_millis(50),
max_consecutive_timeouts: 1,
fallback_to_non_streaming: false,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(
&PartialThenHangMock,
&crate::api::StreamRequest::new(vec![]),
&cancel,
)
.await
.expect_err("the hang must terminate via the event timeout");
match err {
StreamHandlerError::StreamFailed(StreamOutcome::EventTimeout {
has_partial_data,
consecutive_timeouts,
}) => {
assert!(
has_partial_data,
"a hang after accepted events must report salvageable partial data"
);
assert_eq!(consecutive_timeouts, 1);
}
other => panic!("expected an EventTimeout terminal, got {other:?}"),
}
}
#[tokio::test]
async fn retried_attempt_re_gates_on_the_rate_limiter() {
struct FailOnceThenAnswerMock {
calls: std::sync::atomic::AtomicUsize,
}
impl ApiClient for FailOnceThenAnswerMock {
fn model(&self) -> String {
"fail-once".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let calls = self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if calls == 0 {
return Box::pin(futures::stream::once(async {
Err(ApiError::http("transient transport failure"))
}));
}
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "fail-once".to_string(),
},
})),
Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(crate::stream::MessagePart::text("")),
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "recovered".to_string(),
},
})),
Ok(StreamEvent::PartStop { index: Some(0) }),
Ok(StreamEvent::MessageStop),
];
Box::pin(futures::stream::iter(events))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::api("unused")) })
}
}
let handler = StreamHandler::new()
.with_rate_limiter(Arc::new(crate::stream::rate_limit::RateLimiter::new(1)))
.with_rate_limit_max_wait(Duration::from_millis(1200))
.with_retry_config(crate::stream::handler::StreamRetryConfig {
base_delay_ms: 1,
max_delay_ms: 1,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let client = FailOnceThenAnswerMock {
calls: std::sync::atomic::AtomicUsize::new(0),
};
let started = std::time::Instant::now();
handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("the retried attempt must succeed");
let elapsed = started.elapsed();
assert!(
elapsed >= Duration::from_millis(1000),
"the retried attempt must re-gate on the limiter and wait out the max-wait \
ceiling (1 rpm = the first attempt drains the bucket); elapsed {elapsed:?}"
);
}
#[tokio::test]
async fn stream_turn_rate_limit_budget_independent_of_transport() {
use std::sync::atomic::{AtomicUsize, Ordering};
struct TransportThenRateLimitMock {
calls: AtomicUsize,
}
impl ApiClient for TransportThenRateLimitMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let n = self.calls.fetch_add(1, Ordering::SeqCst);
let result = if n == 0 {
Err(ApiError::api("connection refused"))
} else {
Err(ApiError::RateLimit {
retry_after: Some(Duration::from_millis(1)),
message: "slow down".into(),
})
};
Box::pin(futures::stream::once(async { result }))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: false,
..Default::default()
})
.with_retry_config(StreamRetryConfig {
max_retries: 1,
base_delay_ms: 1,
..Default::default()
})
.with_rate_limit_config(RateLimitConfig {
fallback_after_retries: 3,
default_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(1),
..Default::default()
});
let client = TransportThenRateLimitMock {
calls: AtomicUsize::new(0),
};
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect_err("should escalate after the rate-limit budget, not fall through");
match err {
StreamHandlerError::RateLimitEscalation { attempts, .. } => {
assert_eq!(attempts, 4);
}
other => panic!("expected RateLimitEscalation, got {other:?}"),
}
}
#[tokio::test]
async fn stream_turn_rate_limit_hard_stop_after_max_retries() {
struct AlwaysRateLimitMock;
impl ApiClient for AlwaysRateLimitMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::RateLimit {
retry_after: None,
message: "slow down".into(),
})
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: false,
..Default::default()
})
.with_rate_limit_config(RateLimitConfig {
fallback_after_retries: 2,
max_retries: 2,
default_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(1),
..Default::default()
});
let client = AlwaysRateLimitMock;
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect_err("hard-stop should fail the turn");
match err {
StreamHandlerError::StreamFailed(StreamOutcome::RateLimited { .. })
| StreamHandlerError::InitFailed(StreamOutcome::RateLimited { .. }) => {}
StreamHandlerError::RateLimitEscalation { .. } => {
panic!("escalation must not fire when max_retries == fallback_after_retries")
}
other => panic!("expected rate-limit outcome, got {other:?}"),
}
}
#[tokio::test]
async fn stream_turn_rate_limit_counter_does_not_leak_across_calls() {
use std::sync::atomic::{AtomicUsize, Ordering};
struct RateLimitOnceMock {
attempts: AtomicUsize,
}
impl ApiClient for RateLimitOnceMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let n = self.attempts.fetch_add(1, Ordering::SeqCst);
if n == 0 {
Box::pin(futures::stream::once(async {
Err(ApiError::RateLimit {
retry_after: None,
message: "slow down".into(),
})
}))
} else {
Box::pin(futures::stream::iter(happy_stream_events()))
}
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new().with_rate_limit_config(RateLimitConfig {
default_delay: Duration::from_millis(1),
max_delay: Duration::from_millis(1),
..Default::default()
});
let client = RateLimitOnceMock {
attempts: AtomicUsize::new(0),
};
let cancel = Arc::new(CancelSignal::new());
handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("first call should succeed after one rate-limit retry");
let client2 = RateLimitOnceMock {
attempts: AtomicUsize::new(0),
};
handler
.drive_turn(&client2, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect("second call should not see leaked rate-limit state");
}
#[tokio::test]
async fn stream_turn_non_rate_limit_error_path_unchanged() {
struct AlwaysFailingMock;
impl ApiClient for AlwaysFailingMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::api("connection refused"))
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: false,
..Default::default()
})
.with_retry_config(StreamRetryConfig {
max_retries: 1,
base_delay_ms: 1,
..Default::default()
});
let client = AlwaysFailingMock;
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await
.expect_err("transport errors should fail");
assert!(
!matches!(err, StreamHandlerError::RateLimitEscalation { .. }),
"non-rate-limit errors must not escalate"
);
}
#[tokio::test]
async fn stream_turn_rate_limit_delay_clamped_to_total_timeout() {
struct AlwaysRateLimitMock;
impl ApiClient for AlwaysRateLimitMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::RateLimit {
retry_after: Some(Duration::from_mins(10)),
message: "slow down".into(),
})
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(40),
per_event_timeout: Duration::from_millis(40),
total_stream_timeout: Duration::from_millis(80),
..Default::default()
})
.with_retry_config(StreamRetryConfig {
max_retries: 10,
base_delay_ms: 1,
..Default::default()
})
.with_rate_limit_config(RateLimitConfig {
max_delay: Duration::from_mins(10),
default_delay: Duration::from_millis(1),
fallback_after_retries: 100,
max_retries: 100,
..Default::default()
});
let client = AlwaysRateLimitMock;
let cancel = Arc::new(CancelSignal::new());
let start = Instant::now();
let result = handler
.drive_turn(&client, &crate::api::StreamRequest::new(vec![]), &cancel)
.await;
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(2),
"deadline clamp should prevent a 600s sleep; elapsed {elapsed:?}",
);
match result {
Ok(done) => assert!(
done.from_fallback,
"a prompt success here can only be the non-streaming fallback"
),
Err(err) => assert!(
!matches!(err, StreamHandlerError::RateLimitEscalation { .. }),
"timeout should fire before escalation"
),
}
}
#[tokio::test]
async fn stream_turn_cancel_during_backoff_returns_immediately() {
struct AlwaysFailingMock;
impl ApiClient for AlwaysFailingMock {
fn model(&self) -> String {
"test-model".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::once(async {
Err(ApiError::api("connection lost"))
}))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new().with_retry_config(StreamRetryConfig {
max_retries: 5,
base_delay_ms: 60_000,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let cancel_clone = Arc::clone(&cancel);
tokio::spawn(async move {
tokio::task::yield_now().await;
cancel_clone.cancel();
});
let start = Instant::now();
let err = handler
.drive_turn(
&AlwaysFailingMock,
&crate::api::StreamRequest::new(vec![]),
&cancel,
)
.await
.expect_err("should return Cancelled, not hang for 60s");
let elapsed = start.elapsed();
assert!(
matches!(err, StreamHandlerError::Cancelled),
"expected Cancelled, got {err:?}",
);
assert!(
elapsed < Duration::from_secs(5),
"cancellation during backoff should return immediately, not wait for the 60s sleep; elapsed {elapsed:?}",
);
}
#[tokio::test]
async fn malformed_event_fails_the_stream_with_attempt_context() {
struct GarbageToolInputMock;
impl ApiClient for GarbageToolInputMock {
fn model(&self) -> String {
"garbage-input".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "garbage-input".to_string(),
},
})),
Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(crate::stream::MessagePart::tool_call(
"t1",
"search",
serde_json::json!({}),
)),
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::InputJson {
partial_json: "not json".to_string(),
},
})),
Ok(StreamEvent::PartStop { index: None }),
];
Box::pin(futures::stream::iter(events))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::http("unused")) })
}
}
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(5),
per_event_timeout: Duration::from_secs(5),
total_stream_timeout: Duration::from_secs(60),
max_consecutive_timeouts: 3,
fallback_to_non_streaming: false,
});
let cancel = Arc::new(CancelSignal::new());
let request = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&GarbageToolInputMock,
&request,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut yielded = 0usize;
let mut terminal = None;
while let Some(item) = stream.next().await {
match item {
Ok(HandlerEvent::Stream(_)) => yielded += 1,
Ok(_) => {}
Err(e) => {
terminal = Some(e);
break;
}
}
}
assert_eq!(
yielded, 12,
"each ladder attempt replays the accepted events before the malformed one (4 attempts × 3)"
);
match terminal.expect("the malformed event must fail the stream") {
StreamHandlerError::StreamFailed(StreamOutcome::InitFailed {
attempts,
last_error,
}) => {
assert_eq!(
attempts, 4,
"the failure counts every attempt the ladder made"
);
assert!(
last_error.contains("invalid tool input JSON"),
"the accumulator's rejection surfaces verbatim, got: {last_error}"
);
}
other => panic!("expected a StreamFailed InitFailed terminal, got {other:?}"),
}
}
#[tokio::test]
async fn truncated_stream_is_not_a_completed_turn() {
struct CutStreamMock;
impl ApiClient for CutStreamMock {
fn model(&self) -> String {
"cut-stream".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "cut-stream".to_string(),
},
})),
Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(crate::stream::MessagePart::text("")),
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "partial".to_string(),
},
})),
];
Box::pin(futures::stream::iter(events))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::api("unused")) })
}
}
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
fallback_to_non_streaming: false,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let err = handler
.drive_turn(
&CutStreamMock,
&crate::api::StreamRequest::new(vec![]),
&cancel,
)
.await
.expect_err("a stream that ends without a terminal event is truncated");
let rendered = err.to_string();
assert!(
rendered.contains("without a terminal event") && !rendered.contains("init failed"),
"the engine-facing message must name the truncation, not the \
historical init framing: {rendered}"
);
}
#[tokio::test]
async fn truncated_stream_with_fallback_enabled_gets_the_ladder() {
struct CutStreamThenAnswerMock;
impl ApiClient for CutStreamThenAnswerMock {
fn model(&self) -> String {
"cut-then-answer".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "cut-then-answer".to_string(),
},
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "partial".to_string(),
},
})),
];
Box::pin(futures::stream::iter(events))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant("fallback ok"),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new();
let cancel = Arc::new(CancelSignal::new());
let request = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&CutStreamThenAnswerMock,
&request,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut fallback_message = None;
while let Some(item) = stream.next().await {
match item {
Ok(HandlerEvent::Fallback { message, .. }) => fallback_message = Some(message),
Err(e) => panic!(
"a truncated stream with the fallback enabled must not fail the turn: {e}"
),
_ => {}
}
}
assert_eq!(
fallback_message
.expect("the non-streaming fallback must serve the truncated turn")
.text_content(),
"fallback ok"
);
}
#[tokio::test]
async fn malformed_event_with_fallback_enabled_gets_the_ladder() {
struct GarbageThenAnswerMock;
impl ApiClient for GarbageThenAnswerMock {
fn model(&self) -> String {
"garbage-then-answer".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "garbage-then-answer".to_string(),
},
})),
Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(crate::stream::MessagePart::tool_call(
"t1",
"search",
serde_json::json!({}),
)),
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::InputJson {
partial_json: "not json".to_string(),
},
})),
Ok(StreamEvent::PartStop { index: Some(0) }),
];
Box::pin(futures::stream::iter(events))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant("fallback ok"),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
let handler = StreamHandler::new();
let cancel = Arc::new(CancelSignal::new());
let request = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&GarbageThenAnswerMock,
&request,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut fallback_message = None;
while let Some(item) = stream.next().await {
match item {
Ok(HandlerEvent::Fallback { message, .. }) => fallback_message = Some(message),
Err(e) => panic!(
"a malformed event with the fallback enabled must not fail the turn: {e}"
),
_ => {}
}
}
assert_eq!(
fallback_message
.expect("the non-streaming fallback must serve the turn")
.text_content(),
"fallback ok",
"exhausting the retry ladder on accumulation failures routes to the fallback"
);
}
#[test]
fn http_429_is_classified_as_rate_limited() {
let detected =
DetectedRateLimit::detect(&ApiError::http_with_status(429, "Too Many Requests"))
.expect("429 must be detected as a rate limit");
assert_eq!(
detected.kind,
RateLimitKind::RateLimited,
"doc: RateLimited is the HTTP 429 Too Many Requests kind"
);
}
#[test]
fn rate_limit_variant_kind_splits_by_message_status() {
let overload =
DetectedRateLimit::detect(&ApiError::rate_limited("HTTP 503: unavailable", None))
.expect("a 503-shaped RateLimit must be detected");
assert!(
matches!(overload.kind, RateLimitKind::Overloaded),
"503 is the Overloaded kind, got {:?}",
overload.kind
);
let overloaded_529 =
DetectedRateLimit::detect(&ApiError::rate_limited("HTTP 529: overloaded", None))
.expect("a 529-shaped RateLimit must be detected");
assert!(matches!(overloaded_529.kind, RateLimitKind::Overloaded));
let quota = DetectedRateLimit::detect(&ApiError::rate_limited("HTTP 429: slow down", None))
.expect("a 429-shaped RateLimit must be detected");
assert!(matches!(quota.kind, RateLimitKind::RateLimited));
let untyped =
DetectedRateLimit::detect(&ApiError::rate_limited("provider quota text", None))
.expect("a statusless RateLimit must be detected");
assert!(
matches!(untyped.kind, RateLimitKind::RateLimited),
"without an embedded status the default kind is RateLimited"
);
}
struct FailingStreamClient {
make_error: fn() -> ApiError,
stream_calls: std::sync::atomic::AtomicUsize,
non_streaming_calls: std::sync::atomic::AtomicUsize,
}
impl FailingStreamClient {
fn failing_with(make_error: fn() -> ApiError) -> Self {
Self {
make_error,
stream_calls: std::sync::atomic::AtomicUsize::new(0),
non_streaming_calls: std::sync::atomic::AtomicUsize::new(0),
}
}
fn stream_calls(&self) -> usize {
self.stream_calls.load(std::sync::atomic::Ordering::SeqCst)
}
fn non_streaming_calls(&self) -> usize {
self.non_streaming_calls
.load(std::sync::atomic::Ordering::SeqCst)
}
}
impl ApiClient for FailingStreamClient {
fn model(&self) -> String {
"failing".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
self.stream_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(futures::stream::iter(vec![Err((self.make_error)())]))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
self.non_streaming_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(async { Err((self.make_error)()) })
}
}
async fn terminal_error<C: ApiClient>(
handler: &StreamHandler,
client: &C,
cancel: &Arc<CancelSignal>,
) -> StreamHandlerError {
let request = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
client,
&request,
crate::structured::RequestOptions::default(),
cancel,
);
while let Some(item) = stream.next().await {
if let Err(e) = item {
return e;
}
}
panic!("the stream must terminate with an error");
}
#[tokio::test]
async fn unauthorized_stream_errors_are_not_retried() {
let client = FailingStreamClient::failing_with(|| {
ApiError::auth_invalid_key("HTTP 401: invalid api key")
});
let handler = StreamHandler::new()
.with_retry_config(StreamRetryConfig {
max_retries: 3,
base_delay_ms: 1,
max_delay_ms: 2,
..Default::default()
})
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(5),
per_event_timeout: Duration::from_secs(5),
total_stream_timeout: Duration::from_secs(60),
max_consecutive_timeouts: 3,
fallback_to_non_streaming: true,
});
let cancel = Arc::new(CancelSignal::new());
let err = terminal_error(&handler, &client, &cancel).await;
assert_eq!(
client.stream_calls(),
1,
"a permanent 401 must cost exactly one streaming attempt"
);
assert_eq!(
client.non_streaming_calls(),
0,
"a permanent 401 must not get a non-streaming fallback attempt"
);
match err {
StreamHandlerError::StreamFailed(StreamOutcome::InitFailed { last_error, .. }) => {
assert!(
last_error.contains("Invalid API key"),
"the auth failure must surface verbatim, got: {last_error}"
);
}
other => panic!("the 401 must fail the stream, got {other:?}"),
}
}
#[tokio::test]
async fn internal_server_error_is_still_retried() {
let client = FailingStreamClient::failing_with(|| ApiError::http_with_status(500, "boom"));
let handler = StreamHandler::new()
.with_retry_config(StreamRetryConfig {
max_retries: 3,
base_delay_ms: 1,
max_delay_ms: 2,
..Default::default()
})
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(5),
per_event_timeout: Duration::from_secs(5),
total_stream_timeout: Duration::from_secs(60),
max_consecutive_timeouts: 3,
fallback_to_non_streaming: false,
});
let cancel = Arc::new(CancelSignal::new());
let err = terminal_error(&handler, &client, &cancel).await;
assert_eq!(
client.stream_calls(),
4,
"a 500-class error keeps the full ladder: initial + max_retries retries"
);
assert!(
matches!(
err,
StreamHandlerError::StreamFailed(StreamOutcome::InitFailed { .. })
),
"with the fallback disabled the exhausted ladder fails the stream, got {err:?}"
);
}
#[tokio::test]
async fn not_found_stream_errors_are_not_retried() {
let client =
FailingStreamClient::failing_with(|| ApiError::http_with_status(404, "unknown model"));
let handler = StreamHandler::new()
.with_retry_config(StreamRetryConfig {
max_retries: 3,
base_delay_ms: 1,
max_delay_ms: 2,
..Default::default()
})
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(5),
per_event_timeout: Duration::from_secs(5),
total_stream_timeout: Duration::from_secs(60),
max_consecutive_timeouts: 3,
fallback_to_non_streaming: true,
});
let cancel = Arc::new(CancelSignal::new());
let err = terminal_error(&handler, &client, &cancel).await;
assert_eq!(
client.stream_calls(),
1,
"a permanent 404 must cost exactly one streaming attempt"
);
assert_eq!(
client.non_streaming_calls(),
0,
"a permanent 404 must not get a non-streaming fallback attempt"
);
match err {
StreamHandlerError::StreamFailed(StreamOutcome::InitFailed { last_error, .. }) => {
assert!(
last_error.contains("HTTP 404"),
"the permanent status must surface verbatim, got: {last_error}"
);
}
other => panic!("the 404 must fail the stream, got {other:?}"),
}
}
struct StalledStreamClient {
stream_calls: std::sync::atomic::AtomicUsize,
}
impl ApiClient for StalledStreamClient {
fn model(&self) -> String {
"stalled".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
self.stream_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(futures::stream::pending())
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::http("no non-streaming path")) })
}
}
#[tokio::test]
async fn mid_stream_total_timeout_takes_the_fallback_path() {
struct StallThenFallbackMock {
stream_calls: std::sync::atomic::AtomicUsize,
non_streaming_calls: std::sync::atomic::AtomicUsize,
}
impl StallThenFallbackMock {
fn counting() -> Self {
Self {
stream_calls: std::sync::atomic::AtomicUsize::new(0),
non_streaming_calls: std::sync::atomic::AtomicUsize::new(0),
}
}
}
impl ApiClient for StallThenFallbackMock {
fn model(&self) -> String {
"stall-fallback".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
self.stream_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "stall-fallback".to_string(),
},
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "partial".to_string(),
},
})),
];
Box::pin(futures::stream::iter(events).chain(futures::stream::pending()))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
self.non_streaming_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(async {
tokio::time::sleep(Duration::from_millis(50)).await;
Ok(crate::api::NonStreamingResponse {
message: Message::new(
crate::message::Role::Assistant,
vec![crate::stream::MessagePart::text("fallback answer")],
),
stop_reason: StreamStopReason::EndTurn,
usage: None,
})
})
}
}
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(200),
per_event_timeout: Duration::from_millis(200),
total_stream_timeout: Duration::from_millis(400),
max_consecutive_timeouts: 10,
fallback_to_non_streaming: true,
});
let cancel = Arc::new(CancelSignal::new());
let request = crate::api::StreamRequest::new(vec![]);
let client = StallThenFallbackMock::counting();
let started = Instant::now();
let mut stream = handler.stream_turn(
&client,
&request,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut fell_back = false;
while let Some(item) = stream.next().await {
match item.expect("an expired deadline with fallback configured must not error") {
HandlerEvent::Fallback { .. } => fell_back = true,
HandlerEvent::Stream(_) | HandlerEvent::AttemptReset => {}
}
}
assert!(
fell_back,
"a mid-stream total timeout must reach the non-streaming fallback, not a retry or a bare failure"
);
assert!(
started.elapsed() >= Duration::from_millis(400),
"the fallback must complete after the streaming deadline expired, at {started:?}+{elapsed:?}",
elapsed = started.elapsed()
);
assert_eq!(
client
.stream_calls
.load(std::sync::atomic::Ordering::SeqCst),
1,
"the expired deadline must cost exactly one streaming attempt"
);
assert_eq!(
client
.non_streaming_calls
.load(std::sync::atomic::Ordering::SeqCst),
1,
"the fallback must run exactly once"
);
}
#[tokio::test]
async fn mid_stream_total_timeout_is_not_retried() {
struct CountingStallMock {
stream_calls: std::sync::atomic::AtomicUsize,
}
impl ApiClient for CountingStallMock {
fn model(&self) -> String {
"counting-stall".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
self.stream_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "counting-stall".to_string(),
},
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "partial".to_string(),
},
})),
];
Box::pin(futures::stream::iter(events).chain(futures::stream::pending()))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::http("unused")) })
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(200),
per_event_timeout: Duration::from_millis(200),
total_stream_timeout: Duration::from_millis(400),
max_consecutive_timeouts: 10,
fallback_to_non_streaming: false,
})
.with_retry_config(StreamRetryConfig {
max_retries: 3,
base_delay_ms: 1,
max_delay_ms: 2,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let request = crate::api::StreamRequest::new(vec![]);
let client = CountingStallMock {
stream_calls: std::sync::atomic::AtomicUsize::new(0),
};
let mut stream = handler.stream_turn(
&client,
&request,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut resets = 0usize;
let mut terminal = None;
while let Some(item) = stream.next().await {
match item {
Ok(HandlerEvent::AttemptReset) => resets += 1,
Ok(_) => {}
Err(e) => {
terminal = Some(e);
break;
}
}
}
assert_eq!(
client
.stream_calls
.load(std::sync::atomic::Ordering::SeqCst),
1,
"an expired total deadline must never trigger a second streaming attempt"
);
assert_eq!(
resets, 0,
"no AttemptReset may be emitted when the timeout is terminal"
);
assert!(
matches!(
terminal,
Some(StreamHandlerError::StreamFailed(StreamOutcome::TotalTimeout {
events_processed,
..
})) if events_processed >= 2
),
"the terminal error must be the mid-stream TotalTimeout with real progress, got {terminal:?}"
);
}
#[tokio::test]
async fn hanging_fallback_is_cut_by_the_fresh_budget() {
struct StallWithHangingFallbackMock {
stream_calls: std::sync::atomic::AtomicUsize,
non_streaming_calls: std::sync::atomic::AtomicUsize,
}
impl ApiClient for StallWithHangingFallbackMock {
fn model(&self) -> String {
"stall-hanging-fallback".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
self.stream_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let events = vec![Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "stall-hanging-fallback".to_string(),
},
}))];
Box::pin(futures::stream::iter(events).chain(futures::stream::pending()))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
self.non_streaming_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(std::future::pending())
}
}
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(200),
per_event_timeout: Duration::from_millis(200),
total_stream_timeout: Duration::from_millis(400),
max_consecutive_timeouts: 10,
fallback_to_non_streaming: true,
});
let cancel = Arc::new(CancelSignal::new());
let client = StallWithHangingFallbackMock {
stream_calls: std::sync::atomic::AtomicUsize::new(0),
non_streaming_calls: std::sync::atomic::AtomicUsize::new(0),
};
let started = Instant::now();
let err = terminal_error(&handler, &client, &cancel).await;
let elapsed = started.elapsed();
assert_eq!(
client
.stream_calls
.load(std::sync::atomic::Ordering::SeqCst),
1,
"the stalled stream costs one attempt"
);
assert_eq!(
client
.non_streaming_calls
.load(std::sync::atomic::Ordering::SeqCst),
1,
"the fallback must actually start"
);
assert!(
elapsed >= Duration::from_millis(550),
"the fallback must run its fresh initial_event_timeout budget (200ms) after the \
streaming deadline (400ms), not be cut instantly by the expired deadline; elapsed {elapsed:?}"
);
assert!(
elapsed < Duration::from_secs(5),
"the fresh budget must still bound a hanging fallback; elapsed {elapsed:?}"
);
match err {
StreamHandlerError::FallbackFailed { fallback_error, .. } => assert!(
fallback_error.contains("deadline"),
"the fresh budget's expiry must be the failure cause: {fallback_error}"
),
other => panic!("a hanging fallback must fail as FallbackFailed, got {other:?}"),
}
}
#[tokio::test]
async fn expired_deadline_before_retry_takes_the_fallback() {
struct RetryErrorThenFallbackMock {
stream_calls: std::sync::atomic::AtomicUsize,
non_streaming_calls: std::sync::atomic::AtomicUsize,
}
impl ApiClient for RetryErrorThenFallbackMock {
fn model(&self) -> String {
"retry-then-fallback".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
self.stream_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(futures::stream::iter(vec![Err(ApiError::http(
"connection reset",
))]))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
self.non_streaming_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: Message::new(
crate::message::Role::Assistant,
vec![crate::stream::MessagePart::text("fallback answer")],
),
stop_reason: StreamStopReason::EndTurn,
usage: None,
})
})
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(5),
per_event_timeout: Duration::from_secs(5),
total_stream_timeout: Duration::from_millis(150),
max_consecutive_timeouts: 3,
fallback_to_non_streaming: true,
})
.with_retry_config(StreamRetryConfig {
max_retries: 1,
base_delay_ms: 400,
max_delay_ms: 400,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let client = RetryErrorThenFallbackMock {
stream_calls: std::sync::atomic::AtomicUsize::new(0),
non_streaming_calls: std::sync::atomic::AtomicUsize::new(0),
};
let request = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&client,
&request,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut fell_back = false;
while let Some(item) = stream.next().await {
match item {
Ok(HandlerEvent::Fallback { .. }) => fell_back = true,
Ok(_) => {}
Err(e) => panic!("the expiry must take the fallback, got {e:?}"),
}
}
assert!(
fell_back,
"a deadline expiring before the next retry must reach the non-streaming fallback"
);
let calls = client
.stream_calls
.load(std::sync::atomic::Ordering::SeqCst);
assert_eq!(
calls, 2,
"the retried attempt starts and is cut on its first poll — the expiry is terminal"
);
assert_eq!(
client
.non_streaming_calls
.load(std::sync::atomic::Ordering::SeqCst),
1,
"the fallback must run exactly once"
);
}
#[tokio::test]
async fn per_event_timeout_exhaustion_still_retries() {
struct AlwaysStalledMock {
stream_calls: std::sync::atomic::AtomicUsize,
}
impl ApiClient for AlwaysStalledMock {
fn model(&self) -> String {
"always-stalled".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
self.stream_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(futures::stream::pending())
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::http("no non-streaming path")) })
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(50),
per_event_timeout: Duration::from_millis(50),
total_stream_timeout: Duration::from_secs(60),
max_consecutive_timeouts: 2,
fallback_to_non_streaming: false,
})
.with_retry_config(StreamRetryConfig {
max_retries: 1,
base_delay_ms: 1,
max_delay_ms: 2,
..Default::default()
});
let cancel = Arc::new(CancelSignal::new());
let client = AlwaysStalledMock {
stream_calls: std::sync::atomic::AtomicUsize::new(0),
};
let err = terminal_error(&handler, &client, &cancel).await;
assert_eq!(
client
.stream_calls
.load(std::sync::atomic::Ordering::SeqCst),
2,
"per-event timeout exhaustion keeps the retry ladder: initial + one retry"
);
assert!(
matches!(
err,
StreamHandlerError::StreamFailed(StreamOutcome::EventTimeout { .. })
),
"the exhausted ladder terminates with the EventTimeout outcome, got {err:?}"
);
}
#[tokio::test]
async fn per_event_stall_still_uses_the_total_deadline() {
struct SlowButHealthyStream;
impl ApiClient for SlowButHealthyStream {
fn model(&self) -> String {
"slow-healthy".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
Box::pin(async_stream::stream! {
yield Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "slow-healthy".to_string(),
},
}));
for _ in 0..10 {
tokio::time::sleep(Duration::from_millis(30)).await;
yield Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text { text: "chunk".to_string() },
}));
}
yield Ok(StreamEvent::MessageStop);
})
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::http("unused")) })
}
}
struct FlakyThenOkStream {
calls: std::sync::atomic::AtomicUsize,
}
impl ApiClient for FlakyThenOkStream {
fn model(&self) -> String {
"flaky-then-ok".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>
{
let call = self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if call == 0 {
return Box::pin(futures::stream::iter(vec![Err(ApiError::http(
"connection reset",
))]));
}
Box::pin(futures::stream::iter(vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m2".to_string(),
role: "assistant".to_string(),
model: "flaky-then-ok".to_string(),
},
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "recovered".to_string(),
},
})),
Ok(StreamEvent::MessageStop),
]))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::http("unused")) })
}
}
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(1),
per_event_timeout: Duration::from_secs(1),
total_stream_timeout: Duration::from_secs(2),
max_consecutive_timeouts: 3,
fallback_to_non_streaming: false,
});
let cancel = Arc::new(CancelSignal::new());
let request = crate::api::StreamRequest::new(vec![]);
{
let mut stream = handler.stream_turn(
&SlowButHealthyStream,
&request,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut stopped = false;
while let Some(item) = stream.next().await {
if let HandlerEvent::Stream(StreamEvent::MessageStop) =
item.expect("a healthy stream within both budgets must not error")
{
stopped = true;
}
}
assert!(
stopped,
"a stream producing events under the total budget must complete, not be cut"
);
}
let client = FlakyThenOkStream {
calls: std::sync::atomic::AtomicUsize::new(0),
};
let handler = handler.with_retry_config(StreamRetryConfig {
max_retries: 1,
base_delay_ms: 1,
max_delay_ms: 2,
..Default::default()
});
let mut stream = handler.stream_turn(
&client,
&request,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut stopped = false;
while let Some(item) = stream.next().await {
if let HandlerEvent::Stream(StreamEvent::MessageStop) =
item.expect("a retried-then-successful stream must not error")
{
stopped = true;
}
}
assert!(stopped, "the recovered attempt must complete the turn");
assert_eq!(
client.calls.load(std::sync::atomic::Ordering::SeqCst),
2,
"exactly one retry, then success"
);
}
#[tokio::test]
async fn cancelled_stream_is_not_retried() {
let client = StalledStreamClient {
stream_calls: std::sync::atomic::AtomicUsize::new(0),
};
let handler = StreamHandler::new()
.with_retry_config(StreamRetryConfig {
max_retries: 3,
base_delay_ms: 1,
max_delay_ms: 2,
..Default::default()
})
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_secs(5),
per_event_timeout: Duration::from_secs(5),
total_stream_timeout: Duration::from_secs(60),
max_consecutive_timeouts: 3,
fallback_to_non_streaming: true,
});
let cancel = Arc::new(CancelSignal::new());
let cancel_for_task = Arc::clone(&cancel);
tokio::spawn(async move {
tokio::task::yield_now().await;
cancel_for_task.cancel();
});
let err = terminal_error(&handler, &client, &cancel).await;
assert_eq!(
client
.stream_calls
.load(std::sync::atomic::Ordering::SeqCst),
1,
"cancellation mid-stream must not re-enter the retry ladder"
);
assert!(
matches!(err, StreamHandlerError::Cancelled),
"the terminal error must be the cancellation, got {err:?}"
);
}
#[tokio::test]
async fn mid_stream_total_timeout_reports_real_progress() {
struct StallAfterEventsMock;
impl ApiClient for StallAfterEventsMock {
fn model(&self) -> String {
"stall".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let events = vec![
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "stall".to_string(),
},
})),
Ok(StreamEvent::PartStart(PartStart {
index: 0,
part: Some(crate::stream::MessagePart::text("")),
})),
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "hi".to_string(),
},
})),
];
Box::pin(futures::stream::iter(events).chain(futures::stream::pending()))
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::http_with_status(500, "no non-streaming")) })
}
}
let handler = StreamHandler::new().with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(100),
per_event_timeout: Duration::from_millis(100),
total_stream_timeout: Duration::from_millis(500),
max_consecutive_timeouts: 10,
fallback_to_non_streaming: false,
});
let cancel = Arc::new(CancelSignal::new());
let req = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&StallAfterEventsMock,
&req,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut streamed = 0usize;
let mut terminal = None;
while let Some(item) = stream.next().await {
match item {
Ok(HandlerEvent::Stream(_)) => streamed += 1,
Err(e) => {
terminal = Some(e);
break;
}
Ok(_) => {}
}
}
assert!(streamed >= 3, "the stream processed real events first");
match terminal.expect("stream must terminate with an error") {
StreamHandlerError::StreamFailed(StreamOutcome::TotalTimeout {
events_processed,
..
}) => assert!(
events_processed >= 3,
"doc: events_processed counts accepted events before the deadline — zero implies an immediate stall"
),
other => panic!(
"a mid-stream deadline is a StreamFailed TotalTimeout, got {other:?} after {streamed} events"
),
}
}
#[tokio::test]
async fn total_timeout_duration_covers_retried_attempts() {
use std::sync::atomic::{AtomicUsize, Ordering};
struct FailThenStallMock {
calls: AtomicUsize,
}
impl ApiClient for FailThenStallMock {
fn model(&self) -> String {
"flaky".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
if call == 0 {
let opening = futures::stream::once(async {
Ok(StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "m1".to_string(),
role: "assistant".to_string(),
model: "flaky".to_string(),
},
}))
});
let kept_alive = opening.chain(futures::stream::once(async {
tokio::time::sleep(Duration::from_millis(150)).await;
Ok(StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "chunk".to_string(),
},
}))
}));
Box::pin(kept_alive.chain(futures::stream::once(async {
tokio::time::sleep(Duration::from_millis(1200)).await;
Err(ApiError::http_with_status(500, "transient boom"))
})))
} else {
Box::pin(futures::stream::pending())
}
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<crate::api::NonStreamingResponse, ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async { Err(ApiError::http_with_status(500, "no non-streaming")) })
}
}
let handler = StreamHandler::new()
.with_timeout_config(StreamTimeoutConfig {
initial_event_timeout: Duration::from_millis(500),
per_event_timeout: Duration::from_millis(500),
total_stream_timeout: Duration::from_secs(2),
max_consecutive_timeouts: 10,
fallback_to_non_streaming: false,
})
.with_retry_config(StreamRetryConfig {
max_retries: 1,
..Default::default()
});
let client = FailThenStallMock {
calls: AtomicUsize::new(0),
};
let cancel = Arc::new(CancelSignal::new());
let req = crate::api::StreamRequest::new(vec![]);
let mut stream = handler.stream_turn(
&client,
&req,
crate::structured::RequestOptions::default(),
&cancel,
);
let mut terminal = None;
while let Some(item) = stream.next().await {
if let Err(e) = item {
terminal = Some(e);
break;
}
}
match terminal.expect("stream must terminate with an error") {
StreamHandlerError::StreamFailed(StreamOutcome::TotalTimeout { duration, .. }) => {
assert!(
duration >= Duration::from_millis(1500),
"doc: duration is the full stream lifetime, approximately the configured total (2s); got {duration:?}"
);
}
other => panic!("expected StreamFailed TotalTimeout, got {other:?}"),
}
}
}