use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use futures_util::stream::BoxStream;
use hotl_types::{Item, StopReason, TokenUsage};
use serde::{Deserialize, Serialize};
use serde_json::Value;
pub mod timeouts {
use std::time::Duration;
pub const CONNECT: Duration = Duration::from_secs(10);
pub const HEADERS: Duration = Duration::from_secs(120);
pub const STREAM_IDLE: Duration = Duration::from_secs(300);
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolDef {
pub name: String,
pub description: String,
pub input_schema: Value,
}
#[derive(Debug, Clone)]
pub struct SamplingRequest {
pub model: String,
pub max_tokens: u32,
pub system: Arc<str>,
pub items: Arc<Vec<Item>>,
pub ephemeral_tail: Arc<Vec<Item>>,
pub tools: Arc<[ToolDef]>,
pub thinking: bool,
pub cache: CachePolicy,
pub turn_context: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CacheTtl {
FiveMinutes,
OneHour,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CachePolicy {
Off,
Static { prefix_ttl: CacheTtl },
}
impl CachePolicy {
pub fn marks_breakpoints(self) -> bool {
matches!(self, Self::Static { .. })
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "event", rename_all = "snake_case")]
pub enum StreamEvent {
Started,
BlockStart {
index: usize,
kind: String,
},
TextDelta {
index: usize,
text: String,
},
ThinkingDelta {
index: usize,
text: String,
},
ToolInputDelta {
index: usize,
json: String,
},
BlockEnd {
index: usize,
},
Retrying {
attempt: u32,
reason: String,
},
Completed {
stop: StopReason,
usage: TokenUsage,
blocks: Vec<Value>,
},
}
#[derive(Debug, Clone, thiserror::Error)]
pub enum ProviderError {
#[error("authentication failed: {0}")]
Auth(String),
#[error("HTTP {status}: {message}")]
Http {
status: u16,
message: String,
retry_after: Option<u64>,
},
#[error("transport error: {0}")]
Transport(String),
#[error("stream parse error: {0}")]
Parse(String),
}
pub trait Provider: Send + Sync {
fn stream(
&self,
req: SamplingRequest,
) -> BoxStream<'static, Result<StreamEvent, ProviderError>>;
fn arm(&self) -> ArmGuard {
ArmGuard::noop()
}
}
pub trait Warmable {
fn arm(&self) -> ArmGuard;
}
#[must_use]
pub struct ArmGuard {
cancel: Option<Box<dyn FnOnce() + Send>>,
}
impl ArmGuard {
pub fn noop() -> Self {
Self { cancel: None }
}
pub fn new(cancel: impl FnOnce() + Send + 'static) -> Self {
Self {
cancel: Some(Box::new(cancel)),
}
}
pub fn detach(mut self) {
self.cancel = None;
}
}
impl Drop for ArmGuard {
fn drop(&mut self) {
if let Some(cancel) = self.cancel.take() {
cancel();
}
}
}
#[cfg(test)]
mod arm_guard_tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
#[test]
fn drop_invokes_the_cancel_callback() {
let called = Arc::new(AtomicBool::new(false));
let flag = called.clone();
let guard = ArmGuard::new(move || flag.store(true, Ordering::SeqCst));
assert!(!called.load(Ordering::SeqCst));
drop(guard);
assert!(called.load(Ordering::SeqCst));
}
#[test]
fn noop_cancels_nothing() {
drop(ArmGuard::noop());
}
#[test]
fn detach_suppresses_the_cancel_callback() {
let called = Arc::new(AtomicBool::new(false));
let flag = called.clone();
let guard = ArmGuard::new(move || flag.store(true, Ordering::SeqCst));
guard.detach();
assert!(!called.load(Ordering::SeqCst));
}
#[test]
fn provider_default_arm_is_a_noop() {
let provider = ScriptedProvider::new(vec![]);
drop(provider.arm());
}
}
pub struct ScriptedProvider {
scripts: ScriptQueue,
requests: Mutex<Vec<SamplingRequest>>,
}
type Script = Vec<Result<StreamEvent, ProviderError>>;
type ScriptQueue = Arc<Mutex<VecDeque<Script>>>;
struct ScriptedStream {
script: Script,
pos: usize,
exhausted: bool,
scripts: ScriptQueue,
}
impl futures_util::Stream for ScriptedStream {
type Item = Result<StreamEvent, ProviderError>;
fn poll_next(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
let this = self.get_mut();
match this.script.get(this.pos) {
Some(event) => {
this.pos += 1;
std::task::Poll::Ready(Some(event.clone()))
}
None => {
this.exhausted = true;
std::task::Poll::Ready(None)
}
}
}
}
impl Drop for ScriptedStream {
fn drop(&mut self) {
if self.exhausted {
return;
}
self.scripts
.lock()
.expect("scripted provider mutex")
.push_front(std::mem::take(&mut self.script));
}
}
impl ScriptedProvider {
pub fn new(scripts: Vec<Vec<Result<StreamEvent, ProviderError>>>) -> Self {
Self {
scripts: Arc::new(Mutex::new(scripts.into())),
requests: Mutex::new(Vec::new()),
}
}
pub fn requests(&self) -> Vec<SamplingRequest> {
self.requests.lock().expect("requests mutex").clone()
}
pub fn last_request(&self) -> Option<SamplingRequest> {
self.requests
.lock()
.expect("requests mutex")
.last()
.cloned()
}
pub fn request_count(&self) -> usize {
self.requests.lock().expect("requests mutex").len()
}
pub fn push_script(&self, script: Vec<Result<StreamEvent, ProviderError>>) {
self.scripts
.lock()
.expect("scripted provider mutex")
.push_back(script);
}
pub fn text_reply(text: &str) -> Vec<Result<StreamEvent, ProviderError>> {
vec![
Ok(StreamEvent::Started),
Ok(StreamEvent::BlockStart {
index: 0,
kind: "text".into(),
}),
Ok(StreamEvent::TextDelta {
index: 0,
text: text.into(),
}),
Ok(StreamEvent::BlockEnd { index: 0 }),
Ok(StreamEvent::Completed {
stop: StopReason::EndTurn,
usage: TokenUsage {
input_tokens: 10,
output_tokens: 5,
..Default::default()
},
blocks: vec![serde_json::json!({"type": "text", "text": text})],
}),
]
}
pub fn tool_call(
id: &str,
name: &str,
input: Value,
) -> Vec<Result<StreamEvent, ProviderError>> {
let block = serde_json::json!({"type": "tool_use", "id": id, "name": name, "input": input});
vec![
Ok(StreamEvent::Started),
Ok(StreamEvent::BlockStart {
index: 0,
kind: "tool_use".into(),
}),
Ok(StreamEvent::BlockEnd { index: 0 }),
Ok(StreamEvent::Completed {
stop: StopReason::ToolUse,
usage: TokenUsage {
input_tokens: 10,
output_tokens: 8,
..Default::default()
},
blocks: vec![block],
}),
]
}
}
impl Provider for ScriptedProvider {
fn stream(
&self,
req: SamplingRequest,
) -> BoxStream<'static, Result<StreamEvent, ProviderError>> {
self.requests.lock().expect("requests mutex").push(req);
let script = self
.scripts
.lock()
.expect("scripted provider mutex")
.pop_front()
.unwrap_or_else(|| {
vec![Err(ProviderError::Transport(
"scripted provider exhausted".into(),
))]
});
Box::pin(ScriptedStream {
script,
pos: 0,
exhausted: false,
scripts: Arc::clone(&self.scripts),
})
}
}
pub fn v1_base(base: &str) -> String {
let base = base.trim_end_matches('/');
if base.ends_with("/v1") {
base.to_string()
} else {
format!("{base}/v1")
}
}
#[cfg(test)]
mod base_url_tests {
use super::v1_base;
#[test]
fn both_spellings_and_trailing_slashes_resolve_alike() {
for input in [
"http://127.0.0.1:3456",
"http://127.0.0.1:3456/",
"http://127.0.0.1:3456/v1",
"http://127.0.0.1:3456/v1/",
] {
assert_eq!(v1_base(input), "http://127.0.0.1:3456/v1", "input: {input}");
}
}
}
#[derive(Default)]
pub struct SseParser {
buf: Vec<u8>,
data: Vec<String>,
data_len: usize,
}
pub const SSE_MAX_BUFFER: usize = 1024 * 1024;
pub const SSE_MAX_EVENT: usize = 4 * 1024 * 1024;
impl SseParser {
pub fn feed(&mut self, chunk: &[u8]) -> Result<Vec<String>, ProviderError> {
self.buf.extend_from_slice(chunk);
let mut out = Vec::new();
let mut start = 0;
while let Some(pos) = self.buf[start..].iter().position(|&b| b == b'\n') {
let line = String::from_utf8_lossy(&self.buf[start..start + pos])
.trim_end_matches('\r')
.to_string();
start += pos + 1;
self.line(&line, &mut out)?;
}
self.buf.drain(..start);
if self.buf.len() > SSE_MAX_BUFFER {
return Err(ProviderError::Parse(format!(
"SSE line exceeded {SSE_MAX_BUFFER} bytes without a newline"
)));
}
Ok(out)
}
pub fn finish(&mut self) -> Result<Vec<String>, ProviderError> {
let mut out = Vec::new();
if !self.buf.is_empty() {
let tail = std::mem::take(&mut self.buf);
let line = String::from_utf8_lossy(&tail)
.trim_end_matches('\r')
.to_string();
self.line(&line, &mut out)?;
}
self.dispatch(&mut out);
Ok(out)
}
fn line(&mut self, line: &str, out: &mut Vec<String>) -> Result<(), ProviderError> {
if line.is_empty() {
self.dispatch(out);
return Ok(());
}
if line.starts_with(':') {
return Ok(()); }
let Some(value) = line.strip_prefix("data:") else {
return Ok(());
};
let value = value.strip_prefix(' ').unwrap_or(value);
self.data_len += value.len() + 1;
if self.data_len > SSE_MAX_EVENT {
return Err(ProviderError::Parse(format!(
"SSE event exceeded {SSE_MAX_EVENT} bytes without a blank line"
)));
}
self.data.push(value.to_string());
Ok(())
}
fn dispatch(&mut self, out: &mut Vec<String>) {
self.data_len = 0;
if self.data.is_empty() {
return;
}
let payload = std::mem::take(&mut self.data).join("\n");
if !payload.is_empty() && payload != "[DONE]" {
out.push(payload);
}
}
}
#[cfg(test)]
mod sse_parser_tests {
use super::*;
#[test]
fn multi_line_data_fields_are_joined() {
let mut p = SseParser::default();
let out = p
.feed(b"event: x\ndata: {\"a\":\ndata: 1}\n\ndata: {\"b\":2}\n\n")
.unwrap();
assert_eq!(
out,
vec!["{\"a\":\n1}".to_string(), "{\"b\":2}".to_string()]
);
for payload in &out {
serde_json::from_str::<serde_json::Value>(payload).expect(payload);
}
}
#[test]
fn the_unterminated_final_line_is_flushed() {
let mut p = SseParser::default();
assert!(p
.feed(b"data: {\"type\":\"message_stop\"}")
.unwrap()
.is_empty());
assert_eq!(p.finish().unwrap(), vec!["{\"type\":\"message_stop\"}"]);
}
#[test]
fn a_trailing_event_without_a_blank_line_is_flushed() {
let mut p = SseParser::default();
assert!(p.feed(b"data: {\"n\":1}\n").unwrap().is_empty());
assert_eq!(p.finish().unwrap(), vec!["{\"n\":1}"]);
}
#[test]
fn comments_other_fields_and_done_are_filtered() {
let mut p = SseParser::default();
let out = p
.feed(b": keepalive\nevent: ping\nid: 7\nretry: 100\ndata: [DONE]\n\ndata: {}\n\n")
.unwrap();
assert_eq!(out, vec!["{}".to_string()]);
}
#[test]
fn over_cap_input_is_a_parse_error_not_an_oom() {
let mut p = SseParser::default();
let big = vec![b'x'; SSE_MAX_BUFFER + 1];
assert!(matches!(p.feed(&big), Err(ProviderError::Parse(_))));
let mut q = SseParser::default();
let line = format!("data: {}\n", "y".repeat(64 * 1024));
let err = loop {
if let Err(e) = q.feed(line.as_bytes()) {
break e;
}
};
assert!(matches!(err, ProviderError::Parse(_)), "{err:?}");
}
#[test]
fn chunk_boundaries_and_utf8_splits_survive() {
let wire = "data: {\"t\":\"héllo → wörld\"}\n\n".as_bytes();
let mut p = SseParser::default();
let mut out = Vec::new();
for c in wire.chunks(3) {
out.extend(p.feed(c).unwrap());
}
out.extend(p.finish().unwrap());
assert_eq!(out.len(), 1);
let v: serde_json::Value = serde_json::from_str(&out[0]).unwrap();
assert_eq!(v["t"], "héllo → wörld");
}
}
pub mod retry {
use super::ProviderError;
use std::time::Duration;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Decision {
Retry { delay: Duration },
Fatal,
}
pub const MAX_ATTEMPTS: u32 = 5;
pub const RETRY_AFTER_CAP: Duration = Duration::from_secs(60);
pub fn classify(err: &ProviderError, attempt: u32) -> Decision {
if attempt >= MAX_ATTEMPTS {
return Decision::Fatal;
}
let backoff = Duration::from_secs(1u64 << (attempt - 1));
match err {
ProviderError::Http {
status,
retry_after,
..
} if *status == 429 || *status >= 500 => Decision::Retry {
delay: retry_after
.map(Duration::from_secs)
.unwrap_or(backoff)
.min(RETRY_AFTER_CAP),
},
ProviderError::Transport(_) => Decision::Retry {
delay: backoff.min(RETRY_AFTER_CAP),
},
_ => Decision::Fatal,
}
}
pub fn with_jitter(base: Duration) -> Duration {
use std::hash::{BuildHasher, Hasher};
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.subsec_nanos())
.unwrap_or(0);
let mut h = std::collections::hash_map::RandomState::new().build_hasher();
h.write_u32(nanos);
h.write_u32(std::process::id());
let half = base / 2;
let span = base.saturating_sub(half).as_nanos() as u64;
if span == 0 {
return base;
}
half + Duration::from_nanos(h.finish() % (span + 1))
}
pub fn now_unix() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
pub fn parse_retry_after(value: &str, now_unix: u64) -> Option<u64> {
let v = value.trim();
if let Ok(secs) = v.parse::<u64>() {
return Some(secs);
}
let rest = v.split_once(", ")?.1;
let mut parts = rest.split(' ');
let day: u32 = parts.next()?.parse().ok()?;
let month = match parts.next()? {
"Jan" => 1,
"Feb" => 2,
"Mar" => 3,
"Apr" => 4,
"May" => 5,
"Jun" => 6,
"Jul" => 7,
"Aug" => 8,
"Sep" => 9,
"Oct" => 10,
"Nov" => 11,
"Dec" => 12,
_ => return None,
};
let year: i64 = parts.next()?.parse().ok()?;
let mut hms = parts.next()?.split(':');
let h: u64 = hms.next()?.parse().ok()?;
let m: u64 = hms.next()?.parse().ok()?;
let s: u64 = hms.next()?.parse().ok()?;
if h > 23 || m > 59 || s > 60 || !(1..=31).contains(&day) {
return None;
}
let secs =
days_from_civil(year, month, day).checked_mul(86_400)? as u64 + h * 3600 + m * 60 + s;
Some(secs.saturating_sub(now_unix))
}
fn days_from_civil(y: i64, m: u32, d: u32) -> i64 {
let y = if m <= 2 { y - 1 } else { y };
let era = if y >= 0 { y } else { y - 399 } / 400;
let yoe = y - era * 400; let mp = if m > 2 { m - 3 } else { m + 9 } as i64; let doy = (153 * mp + 2) / 5 + d as i64 - 1;
let doe = yoe * 365 + yoe / 4 - yoe / 100 + doy;
era * 146_097 + doe - 719_468
}
pub fn is_availability(err: &ProviderError) -> bool {
matches!(
err,
ProviderError::Http { status, .. } if *status == 429 || *status >= 500
) || matches!(err, ProviderError::Transport(_))
}
pub fn is_context_overflow(err: &ProviderError) -> bool {
let ProviderError::Http {
status: 400,
message,
..
} = err
else {
return false;
};
let m = message.to_lowercase();
[
"too long",
"context length",
"context window",
"tokens exceed",
]
.iter()
.any(|needle| m.contains(needle))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn overflow_detection() {
let overflow = ProviderError::Http {
status: 400,
message:
r#"{"error":{"message":"prompt is too long: 210000 tokens > 200000 maximum"}}"#
.into(),
retry_after: None,
};
assert!(is_context_overflow(&overflow));
let oai = ProviderError::Http {
status: 400,
message: "This model's maximum context length is 128000 tokens".into(),
retry_after: None,
};
assert!(is_context_overflow(&oai));
let plain_400 = ProviderError::Http {
status: 400,
message: "bad schema".into(),
retry_after: None,
};
assert!(!is_context_overflow(&plain_400));
}
#[test]
fn retry_after_is_capped() {
let hostile = ProviderError::Http {
status: 429,
message: String::new(),
retry_after: Some(3600),
};
assert_eq!(
classify(&hostile, 1),
Decision::Retry {
delay: RETRY_AFTER_CAP
}
);
}
#[test]
fn budget_is_deep_enough_for_an_overload() {
let overload = ProviderError::Http {
status: 529,
message: String::new(),
retry_after: None,
};
assert_eq!(MAX_ATTEMPTS, 5);
for (attempt, secs) in [(1u32, 1u64), (2, 2), (3, 4), (4, 8)] {
assert_eq!(
classify(&overload, attempt),
Decision::Retry {
delay: std::time::Duration::from_secs(secs)
},
"attempt {attempt}"
);
}
assert_eq!(classify(&overload, MAX_ATTEMPTS), Decision::Fatal);
}
#[test]
fn jitter_stays_in_range_and_actually_varies() {
let base = std::time::Duration::from_secs(8);
let mut seen = std::collections::HashSet::new();
for _ in 0..500 {
let d = with_jitter(base);
assert!(d >= base / 2 && d <= base, "{d:?} outside [base/2, base]");
seen.insert(d.as_millis());
}
assert!(
seen.len() > 10,
"jitter is not varying: {} values",
seen.len()
);
}
#[test]
fn retry_after_parses_seconds_and_http_date() {
const WHEN: u64 = 1_445_412_480;
assert_eq!(parse_retry_after("120", WHEN), Some(120));
assert_eq!(parse_retry_after(" 120 ", WHEN), Some(120));
assert_eq!(
parse_retry_after("Wed, 21 Oct 2015 07:30:00 GMT", WHEN),
Some(120)
);
assert_eq!(
parse_retry_after("Wed, 21 Oct 2015 07:00:00 GMT", WHEN),
Some(0)
);
assert_eq!(parse_retry_after("not a date", WHEN), None);
assert_eq!(
parse_retry_after("Mon, 29 Feb 2016 00:00:01 GMT", 1_456_704_000),
Some(1)
);
}
#[test]
fn classify_rules() {
let overload = ProviderError::Http {
status: 529,
message: String::new(),
retry_after: Some(7),
};
assert_eq!(
classify(&overload, 1),
Decision::Retry {
delay: std::time::Duration::from_secs(7)
}
);
assert_eq!(classify(&overload, MAX_ATTEMPTS), Decision::Fatal);
let auth = ProviderError::Auth("bad".into());
assert_eq!(classify(&auth, 1), Decision::Fatal);
assert!(!is_availability(&auth));
let transport = ProviderError::Transport("reset".into());
assert_eq!(
classify(&transport, 2),
Decision::Retry {
delay: std::time::Duration::from_secs(2)
}
);
assert!(is_availability(&transport));
let bad_req = ProviderError::Http {
status: 400,
message: String::new(),
retry_after: None,
};
assert_eq!(classify(&bad_req, 1), Decision::Fatal);
}
}
}
pub mod transform {
use serde_json::Value;
pub fn strip_foreign_reasoning(blocks: &[Value]) -> Vec<Value> {
blocks
.iter()
.filter(|b| {
!matches!(
b.get("type").and_then(Value::as_str),
Some("thinking") | Some("redacted_thinking")
)
})
.cloned()
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn strips_thinking_keeps_rest() {
let blocks = vec![
json!({"type":"thinking","thinking":"x","signature":"s"}),
json!({"type":"redacted_thinking","data":"d"}),
json!({"type":"text","text":"hi"}),
json!({"type":"tool_use","id":"1","name":"read","input":{}}),
];
let out = strip_foreign_reasoning(&blocks);
assert_eq!(out.len(), 2);
assert_eq!(out[0]["type"], "text");
assert_eq!(out[1]["type"], "tool_use");
}
}
}
pub mod repair {
use serde_json::Value;
pub fn parse_or_repair(raw: &str) -> Option<Value> {
if let Ok(v) = serde_json::from_str(raw) {
return Some(v);
}
let without_commas = strip_trailing_commas(raw);
if let Ok(v) = serde_json::from_str(&without_commas) {
return Some(v);
}
serde_json::from_str(&close_truncation(&without_commas)).ok()
}
fn strip_trailing_commas(s: &str) -> String {
let mut out = String::with_capacity(s.len());
let mut in_string = false;
let mut escaped = false;
for c in s.chars() {
if in_string {
out.push(c);
if escaped {
escaped = false;
} else if c == '\\' {
escaped = true;
} else if c == '"' {
in_string = false;
}
continue;
}
match c {
'"' => {
in_string = true;
out.push(c);
}
'}' | ']' => {
while out.ends_with(char::is_whitespace) || out.ends_with(',') {
if out.ends_with(',') {
out.pop();
break;
}
out.pop();
}
out.push(c);
}
_ => out.push(c),
}
}
out
}
fn close_truncation(s: &str) -> String {
let mut stack = Vec::new();
let mut in_string = false;
let mut escaped = false;
for c in s.chars() {
if in_string {
if escaped {
escaped = false;
} else if c == '\\' {
escaped = true;
} else if c == '"' {
in_string = false;
}
continue;
}
match c {
'"' => in_string = true,
'{' => stack.push('}'),
'[' => stack.push(']'),
'}' | ']' => {
stack.pop();
}
_ => {}
}
}
let mut out = s.to_string();
if in_string {
out.push('"');
}
while let Some(closer) = stack.pop() {
out.push(closer);
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn repairs_common_damage_and_rejects_garbage() {
assert_eq!(
parse_or_repair(r#"{"path": "a.rs"}"#).unwrap()["path"],
"a.rs"
);
assert_eq!(
parse_or_repair(r#"{"path": "a.rs",}"#).unwrap()["path"],
"a.rs"
);
assert_eq!(
parse_or_repair(r#"{"items": [1, 2,]}"#).unwrap()["items"][1],
2
);
let v = parse_or_repair(r#"{"command": "cargo tes"#).unwrap();
assert_eq!(v["command"], "cargo tes");
assert_eq!(parse_or_repair(r#"{"t": "a,}"}"#).unwrap()["t"], "a,}");
let v = parse_or_repair(r#"{"t": "say \"hi\"",}"#).unwrap();
assert_eq!(v["t"], "say \"hi\"");
assert!(parse_or_repair("not json at all").is_none());
}
}
}
pub mod api_error;
pub mod catalog;
pub mod key;
pub const MAX_BLOCK_INDEX: usize = 512;
pub trait SseAssembler {
fn handle(&mut self, data: &str) -> Result<Vec<StreamEvent>, ProviderError>;
fn finish(self) -> Result<StreamEvent, ProviderError>;
}
pub fn drive_sse<B, E, A>(
bytes: B,
mut assembler: A,
idle: std::time::Duration,
) -> impl futures_util::Stream<Item = Result<StreamEvent, ProviderError>>
where
B: futures_util::Stream<Item = Result<bytes::Bytes, E>>,
E: std::fmt::Display,
A: SseAssembler,
{
async_stream::stream! {
let mut parser = SseParser::default();
futures_util::pin_mut!(bytes);
use futures_util::StreamExt;
loop {
let next = match tokio::time::timeout(idle, bytes.next()).await {
Ok(next) => next,
Err(_) => {
yield Err(ProviderError::Transport(format!(
"stream stalled: no data for {}s. The connection is likely dead \
(a proxy dropped it without closing); retry the request.",
idle.as_secs()
)));
return;
}
};
let Some(chunk) = next else { break };
let chunk = match chunk {
Ok(c) => c,
Err(e) => {
yield Err(ProviderError::Transport(format!("stream interrupted: {e}")));
return;
}
};
let payloads = match parser.feed(&chunk) {
Ok(payloads) => payloads,
Err(e) => { yield Err(e); return; }
};
for data in payloads {
match assembler.handle(&data) {
Ok(events) => for ev in events { yield Ok(ev); },
Err(e) => { yield Err(e); return; }
}
}
}
match parser.finish() {
Ok(payloads) => {
for data in payloads {
match assembler.handle(&data) {
Ok(events) => for ev in events { yield Ok(ev); },
Err(e) => { yield Err(e); return; }
}
}
}
Err(e) => { yield Err(e); return; }
}
yield assembler.finish();
}
}
#[cfg(test)]
mod drive_sse_tests {
use super::*;
use futures_util::StreamExt;
use std::time::Duration;
struct NeverEnds;
impl SseAssembler for NeverEnds {
fn handle(&mut self, _: &str) -> Result<Vec<StreamEvent>, ProviderError> {
Ok(vec![])
}
fn finish(self) -> Result<StreamEvent, ProviderError> {
Err(ProviderError::Parse("unreachable".into()))
}
}
#[tokio::test(start_paused = true)]
async fn a_silent_stream_times_out_instead_of_hanging() {
let stalled = futures_util::stream::pending::<Result<bytes::Bytes, std::io::Error>>();
let s = drive_sse(stalled, NeverEnds, Duration::from_secs(5));
futures_util::pin_mut!(s);
let first = s.next().await.expect("an event, not a hang");
match first {
Err(ProviderError::Transport(m)) => {
assert!(m.contains("stalled"), "{m}");
assert!(m.contains("retry"), "errors are prompts: {m}");
}
other => panic!("expected a transport timeout, got {other:?}"),
}
assert!(
s.next().await.is_none(),
"the stream must end after the timeout"
);
}
#[tokio::test(start_paused = true)]
async fn a_slow_but_live_stream_is_not_cut() {
let chunks = futures_util::stream::unfold(0u32, |n| async move {
if n == 6 {
return None;
}
tokio::time::sleep(Duration::from_secs(4)).await;
Some((
Ok::<_, std::io::Error>(bytes::Bytes::from_static(b":keepalive\n\n")),
n + 1,
))
});
let s = drive_sse(chunks, NeverEnds, Duration::from_secs(5));
futures_util::pin_mut!(s);
let evs: Vec<_> = s.collect().await;
assert_eq!(evs.len(), 1);
assert!(matches!(evs[0], Err(ProviderError::Parse(_))), "{evs:?}");
}
}