From 3c93ea1697a41aa432a7b69aba5d24767391c67a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 12:18:32 -0700 Subject: [PATCH] refactor(framer): replace Framer trait with tokio-util codecs (#43193) * refactor(framer): replace Framer trait with tokio-util codecs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(framer): port SSE and AWS event stream framing to codecs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): drop clone on Copy capabilities in messages request test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): use field init shorthand in messages request test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/Cargo.lock | 16 +- .../crates/core/tests/messages/request.rs | 2 +- litellm-rust/crates/framer/Cargo.toml | 5 +- .../crates/framer/src/aws_event_stream.rs | 95 ++++---- litellm-rust/crates/framer/src/error.rs | 24 +- litellm-rust/crates/framer/src/framed.rs | 21 ++ litellm-rust/crates/framer/src/lib.rs | 4 +- litellm-rust/crates/framer/src/sse.rs | 189 +++++++++++++--- .../crates/framer/tests/aws_event_stream.rs | 209 ++++++++++++------ litellm-rust/crates/framer/tests/chaining.rs | 74 +++++-- litellm-rust/crates/framer/tests/sse.rs | 188 ++++++++++++---- .../crates/framer/tests/support/mod.rs | 68 ++++-- .../messages/streaming_iterator.rs | 37 ++-- 13 files changed, 657 insertions(+), 275 deletions(-) create mode 100644 litellm-rust/crates/framer/src/framed.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index e7d911f5fd9..3677d1d654f 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3166,10 +3166,11 @@ dependencies = [ "aws-smithy-types", "bytes", "futures-util", + "proptest", "rstest", - "sse-stream", "thiserror 2.0.19", "tokio", + "tokio-util", ] [[package]] @@ -5468,19 +5469,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "sse-stream" -version = "0.2.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c25ac7aff0abd1dbc474536e40416e1102c7dd9bfba0b9861c6d357f835dcfb4" -dependencies = [ - "bytes", - "futures-util", - "http-body 1.1.0", - "http-body-util", - "pin-project-lite", -] - [[package]] name = "stable_deref_trait" version = "1.2.1" diff --git a/litellm-rust/crates/core/tests/messages/request.rs b/litellm-rust/crates/core/tests/messages/request.rs index 2927356b773..d37910d4ac4 100644 --- a/litellm-rust/crates/core/tests/messages/request.rs +++ b/litellm-rust/crates/core/tests/messages/request.rs @@ -398,7 +398,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i api_key: Some("sk".into()), api_base: Some(upstream.uri()), shaping: MessagesShaping { - capabilities: capabilities.clone(), + capabilities, drop_params, ..MessagesShaping::default() }, diff --git a/litellm-rust/crates/framer/Cargo.toml b/litellm-rust/crates/framer/Cargo.toml index 62bfcc7da3d..e11f2c02a97 100644 --- a/litellm-rust/crates/framer/Cargo.toml +++ b/litellm-rust/crates/framer/Cargo.toml @@ -8,16 +8,17 @@ repository.workspace = true [features] default = ["aws", "sse"] aws = ["dep:aws-smithy-eventstream", "dep:aws-smithy-types"] -sse = ["dep:sse-stream"] +sse = [] [dependencies] aws-smithy-eventstream = { version = "=0.61.4", optional = true } aws-smithy-types = { version = "1.6.1", optional = true } bytes = "1" futures-util.workspace = true -sse-stream = { version = "=0.2.6", optional = true } thiserror.workspace = true +tokio-util = { version = "0.7", features = ["codec", "io"] } [dev-dependencies] +proptest.workspace = true rstest.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/framer/src/aws_event_stream.rs b/litellm-rust/crates/framer/src/aws_event_stream.rs index efd7adeb64b..405ec2d5ad2 100644 --- a/litellm-rust/crates/framer/src/aws_event_stream.rs +++ b/litellm-rust/crates/framer/src/aws_event_stream.rs @@ -1,66 +1,47 @@ -use bytes::{Buf, Bytes, BytesMut}; -use futures_util::{Stream, StreamExt}; +use aws_smithy_eventstream::frame::{read_message_from, write_message_to}; +pub use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; +use bytes::BytesMut; +use tokio_util::codec::{Decoder, Encoder}; -use aws_smithy_eventstream::frame::read_message_from; -use aws_smithy_types::event_stream::Header; - -use crate::{Error, Framer}; +use crate::EventStreamError; +const MIN_FRAME_BYTES: usize = 16; const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024; -#[derive(Clone, Debug, PartialEq)] -pub struct AwsEventStreamFrame { - pub headers: Vec
, - pub payload: Bytes, -} - #[derive(Clone, Copy, Debug, Default)] -pub struct AwsEventStreamFramer; +pub struct AwsEventStreamCodec; -impl Framer for AwsEventStreamFramer { - type Frame = AwsEventStreamFrame; +impl Decoder for AwsEventStreamCodec { + type Item = Message; + type Error = EventStreamError; - fn frame(self, input: S) -> impl Stream> + Send - where - S: Stream> + Send, - B: Buf + Send, - E: std::error::Error + Send + Sync + 'static, - { - futures_util::stream::try_unfold( - (Box::pin(input), BytesMut::new()), - |(mut input, mut buffer)| async move { - loop { - if buffer.len() >= 4 { - let length = (&buffer[..4]).get_u32() as usize; - if !(16..=MAX_FRAME_BYTES).contains(&length) { - return Err(Error::InvalidLength(length)); - } - if buffer.len() >= length { - let raw = buffer.split_to(length).freeze(); - let message = read_message_from(raw)?; - let frame = AwsEventStreamFrame { - headers: message.headers().to_vec(), - payload: message.payload().clone(), - }; - return Ok(Some((frame, (input, buffer)))); - } - } - match input.next().await { - Some(Ok(mut chunk)) => { - while chunk.has_remaining() { - let bytes = chunk.chunk(); - buffer.extend_from_slice(bytes); - let length = bytes.len(); - chunk.advance(length); - } - } - Some(Err(error)) => return Err(Error::Body(Box::new(error))), - None if buffer.is_empty() => return Ok(None), - None => return Err(Error::Truncated), - } - } - }, - ) - .fuse() + fn decode(&mut self, src: &mut BytesMut) -> Result, EventStreamError> { + let Some(prefix) = src.first_chunk::<4>() else { + return Ok(None); + }; + let length = u32::from_be_bytes(*prefix) as usize; + if !(MIN_FRAME_BYTES..=MAX_FRAME_BYTES).contains(&length) { + return Err(EventStreamError::InvalidLength(length)); + } + if src.len() < length { + return Ok(None); + } + Ok(Some(read_message_from(src.split_to(length).freeze())?)) + } + + fn decode_eof(&mut self, src: &mut BytesMut) -> Result, EventStreamError> { + match self.decode(src)? { + Some(message) => Ok(Some(message)), + None if src.is_empty() => Ok(None), + None => Err(EventStreamError::Truncated), + } + } +} + +impl Encoder for AwsEventStreamCodec { + type Error = EventStreamError; + + fn encode(&mut self, message: Message, dst: &mut BytesMut) -> Result<(), EventStreamError> { + Ok(write_message_to(&message, dst)?) } } diff --git a/litellm-rust/crates/framer/src/error.rs b/litellm-rust/crates/framer/src/error.rs index b1f7ed96c5a..879d7557671 100644 --- a/litellm-rust/crates/framer/src/error.rs +++ b/litellm-rust/crates/framer/src/error.rs @@ -1,17 +1,21 @@ +#[cfg(feature = "sse")] #[derive(Debug, thiserror::Error)] -pub enum Error { - #[cfg(feature = "sse")] - #[error("SSE framing failed: {0}")] - Sse(#[from] sse_stream::Error), - #[cfg(feature = "aws")] - #[error("AWS EventStream framing failed: {0}")] - Aws(#[from] aws_smithy_eventstream::error::Error), +pub enum SseError { #[error("body stream failed: {0}")] - Body(#[source] Box), - #[cfg(feature = "aws")] + Body(#[from] std::io::Error), + #[error("SSE field is not UTF-8: {0}")] + InvalidUtf8(#[from] std::str::Utf8Error), +} + +#[cfg(feature = "aws")] +#[derive(Debug, thiserror::Error)] +pub enum EventStreamError { + #[error("body stream failed: {0}")] + Body(#[from] std::io::Error), #[error("invalid AWS EventStream frame length: {0}")] InvalidLength(usize), - #[cfg(feature = "aws")] #[error("truncated AWS EventStream frame")] Truncated, + #[error("malformed AWS EventStream frame: {0}")] + Malformed(#[from] aws_smithy_eventstream::error::Error), } diff --git a/litellm-rust/crates/framer/src/framed.rs b/litellm-rust/crates/framer/src/framed.rs new file mode 100644 index 00000000000..7a19dd40e13 --- /dev/null +++ b/litellm-rust/crates/framer/src/framed.rs @@ -0,0 +1,21 @@ +use std::io; + +use bytes::Buf; +use futures_util::{Stream, StreamExt, TryStreamExt}; +use tokio_util::{ + codec::{Decoder, FramedRead}, + io::StreamReader, +}; + +pub fn frames( + input: S, + codec: D, +) -> impl Stream> + Send +where + S: Stream> + Send, + B: Buf + Send, + E: std::error::Error + Send + Sync + 'static, + D: Decoder + Send, +{ + FramedRead::new(StreamReader::new(input.map_err(io::Error::other)), codec).fuse() +} diff --git a/litellm-rust/crates/framer/src/lib.rs b/litellm-rust/crates/framer/src/lib.rs index 552de419984..223f2f64120 100644 --- a/litellm-rust/crates/framer/src/lib.rs +++ b/litellm-rust/crates/framer/src/lib.rs @@ -1,8 +1,8 @@ mod error; -mod framer; +mod framed; pub use error::*; -pub use framer::*; +pub use framed::frames; #[cfg(feature = "aws")] pub mod aws_event_stream; diff --git a/litellm-rust/crates/framer/src/sse.rs b/litellm-rust/crates/framer/src/sse.rs index 79659f6ce13..6fee1cfab7f 100644 --- a/litellm-rust/crates/framer/src/sse.rs +++ b/litellm-rust/crates/framer/src/sse.rs @@ -1,43 +1,170 @@ -use futures_util::{Stream, StreamExt}; +use std::str; -use crate::{Error, Framer}; +use bytes::{Buf, BufMut, BytesMut}; +use tokio_util::codec::{Decoder, Encoder}; -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct SseFrame { +use crate::SseError; + +const BOM: &[u8] = b"\xEF\xBB\xBF"; + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct SseEvent { pub event: Option, - pub data: Option, + pub data: String, pub id: Option, pub retry: Option, } #[derive(Clone, Copy, Debug, Default)] -pub struct SseFramer; +pub struct SseCodec { + past_bom: bool, +} -impl Framer for SseFramer { - type Frame = SseFrame; +impl Decoder for SseCodec { + type Item = SseEvent; + type Error = SseError; - fn frame(self, input: S) -> impl Stream> + Send - where - S: Stream> + Send, - B: bytes::Buf + Send, - E: std::error::Error + Send + Sync + 'static, - { - let frames = Box::pin(sse_stream::SseStream::from_bytes_stream(input)); - futures_util::stream::try_unfold(frames, |mut frames| async move { - let Some(frame) = frames.next().await else { - return Ok(None); - }; - let frame = frame?; - Ok(Some(( - SseFrame { - event: frame.event, - data: frame.data, - id: frame.id, - retry: frame.retry, - }, - frames, - ))) - }) - .fuse() + fn decode(&mut self, src: &mut BytesMut) -> Result, SseError> { + if !self.skip_bom(src) { + return Ok(None); + } + while let Some(end) = block_end(src) { + let block = src.split_to(end); + let pending = lines(&block) + .map(|(line, _)| line) + .take_while(|line| !line.is_empty()) + .try_fold(Pending::default(), Pending::apply)?; + if let Some(event) = pending.dispatch() { + return Ok(Some(event)); + } + } + Ok(None) + } + + fn decode_eof(&mut self, _pending: &mut BytesMut) -> Result, SseError> { + Ok(None) + } +} + +impl SseCodec { + fn skip_bom(&mut self, src: &mut BytesMut) -> bool { + if self.past_bom { + return true; + } + if src.starts_with(BOM) { + src.advance(BOM.len()); + } else if BOM.starts_with(src) { + return false; + } + self.past_bom = true; + true + } +} + +fn block_end(bytes: &[u8]) -> Option { + lines(bytes) + .find(|(line, _)| line.is_empty()) + .map(|(_, end)| end) +} + +fn lines(bytes: &[u8]) -> impl Iterator { + let mut cursor: usize = 0; + std::iter::from_fn(move || { + let rest = &bytes[cursor..]; + let end = rest.iter().position(|byte| matches!(byte, b'\n' | b'\r'))?; + cursor += end + terminator_len(&rest[end..]); + Some((&rest[..end], cursor)) + }) +} + +fn terminator_len(terminated: &[u8]) -> usize { + match terminated { + [b'\r', b'\n', ..] => 2, + _ => 1, + } +} + +#[derive(Default)] +struct Pending { + event: Option, + data: Option, + id: Option, + retry: Option, +} + +impl Pending { + fn apply(self, line: &[u8]) -> Result { + let (name, value) = split_field(line); + Ok(match name { + b"event" => Self { + event: Some(str::from_utf8(value)?.to_owned()), + ..self + }, + b"data" => Self { + data: Some(append_data(self.data, str::from_utf8(value)?)), + ..self + }, + b"id" if !value.contains(&0) => Self { + id: Some(str::from_utf8(value)?.to_owned()), + ..self + }, + b"retry" => Self { + retry: parse_retry(value).or(self.retry), + ..self + }, + _ => self, + }) + } + + fn dispatch(self) -> Option { + Some(SseEvent { + event: self.event, + data: self.data?, + id: self.id, + retry: self.retry, + }) + } +} + +fn split_field(line: &[u8]) -> (&[u8], &[u8]) { + let Some(colon) = line.iter().position(|byte| *byte == b':') else { + return (line, &[]); + }; + let value = &line[colon + 1..]; + (&line[..colon], value.strip_prefix(b" ").unwrap_or(value)) +} + +fn append_data(buffer: Option, line: &str) -> String { + match buffer { + Some(existing) => format!("{existing}\n{line}"), + None => line.to_owned(), + } +} + +fn parse_retry(value: &[u8]) -> Option { + if !value.iter().all(u8::is_ascii_digit) { + return None; + } + str::from_utf8(value).ok()?.parse().ok() +} + +impl Encoder for SseCodec { + type Error = SseError; + + fn encode(&mut self, event: SseEvent, dst: &mut BytesMut) -> Result<(), SseError> { + if let Some(name) = event.event { + dst.put_slice(format!("event: {name}\n").as_bytes()); + } + for line in event.data.split('\n') { + dst.put_slice(format!("data: {line}\n").as_bytes()); + } + if let Some(id) = event.id { + dst.put_slice(format!("id: {id}\n").as_bytes()); + } + if let Some(retry) = event.retry { + dst.put_slice(format!("retry: {retry}\n").as_bytes()); + } + dst.put_u8(b'\n'); + Ok(()) } } diff --git a/litellm-rust/crates/framer/tests/aws_event_stream.rs b/litellm-rust/crates/framer/tests/aws_event_stream.rs index c90a15a2b0e..d16caa39948 100644 --- a/litellm-rust/crates/framer/tests/aws_event_stream.rs +++ b/litellm-rust/crates/framer/tests/aws_event_stream.rs @@ -4,89 +4,174 @@ mod support; use std::io; -use futures_util::TryStreamExt; -use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer}; -use litellm_framing::{Error, Framer}; +use bytes::Bytes; +use futures_util::{StreamExt, TryStreamExt, stream}; +use litellm_framing::{ + EventStreamError, + aws_event_stream::{AwsEventStreamCodec, Header, HeaderValue, Message}, + frames, +}; +use proptest::prelude::*; use rstest::{fixture, rstest}; +use support::{body_cause, cut_at, encode_all, every, input, runtime}; -use support::encode; - -async fn collect_aws(bytes: &[u8], chunk_size: usize) -> Result, Error> { - AwsEventStreamFramer - .frame(futures_util::stream::iter( - bytes.chunks(chunk_size).map(Ok::<_, io::Error>), - )) +async fn collect(pieces: Vec) -> Result, EventStreamError> { + frames(input(pieces), AwsEventStreamCodec) .try_collect() .await } -#[fixture] -fn two_frames() -> Vec { - [encode(b"\xff\x00"), encode(b"second")].concat() +fn message(payload: &[u8]) -> Message { + Message::new(Bytes::copy_from_slice(payload)) + .add_header(Header::new( + ":event-type", + HeaderValue::String("payload".into()), + )) + .add_header(Header::new("sequence", HeaderValue::Int32(7))) } #[fixture] fn payload_frame() -> Vec { - encode(b"payload") + encode_all(AwsEventStreamCodec, [message(b"payload")]) +} + +fn header_value() -> impl Strategy { + prop_oneof![ + "[a-z]{0,8}".prop_map(|text| HeaderValue::String(text.into())), + any::().prop_map(HeaderValue::Int32), + any::().prop_map(HeaderValue::Bool), + proptest::collection::vec(any::(), 0..8) + .prop_map(|bytes| HeaderValue::ByteArray(bytes.into())), + ] +} + +fn arbitrary_message() -> impl Strategy { + ( + proptest::collection::vec(("[a-z:-]{1,12}", header_value()), 0..3), + proptest::collection::vec(any::(), 0..32), + ) + .prop_map(|(headers, payload)| { + headers.into_iter().fold( + Message::new(Bytes::from(payload)), + |message, (name, value)| message.add_header(Header::new(name, value)), + ) + }) +} + +proptest! { + #[test] + fn any_messages_survive_a_round_trip_through_any_cuts( + messages in proptest::collection::vec(arbitrary_message(), 1..4), + cuts in proptest::collection::vec(0_usize..512, 0..4), + ) { + let wire = encode_all(AwsEventStreamCodec, messages.clone()); + let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap(); + prop_assert_eq!(decoded, messages); + } } #[rstest] -#[case(1)] -#[case(3)] -#[case(12)] -#[case(usize::MAX)] +#[case::prelude_crc(8)] +#[case::message_crc(usize::MAX)] #[tokio::test] -async fn fragmented_and_coalesced_frames_preserve_typed_headers_and_binary_payloads( - two_frames: Vec, - #[case] chunk_size: usize, -) { - let chunk_size = chunk_size.min(two_frames.len()); - let frames = collect_aws(&two_frames, chunk_size).await.unwrap(); - assert_eq!(frames.len(), 2); - assert_eq!(frames[0].payload, &b"\xff\x00"[..]); - assert_eq!(frames[1].payload, "second"); - assert_eq!( - frames[0].headers[0].value().as_string().unwrap().as_str(), - "payload" - ); - assert_eq!(frames[0].headers[1].value().as_int32(), Ok(7)); -} - -#[rstest] -#[case(8)] -#[case(0)] -#[tokio::test] -async fn rejects_corrupt_crcs(payload_frame: Vec, #[case] index: usize) { - let corrupt_index = if index == 0 { - payload_frame.len() - 1 - } else { - index - }; +async fn a_corrupt_crc_is_malformed(payload_frame: Vec, #[case] index: usize) { let mut corrupt = payload_frame; - corrupt[corrupt_index] ^= 1; - assert!(matches!(collect_aws(&corrupt, 3).await, Err(Error::Aws(_)))); -} - -#[rstest] -#[case(0_u32)] -#[case(15)] -#[case(u32::MAX)] -#[tokio::test] -async fn rejects_invalid_lengths(#[case] length: u32) { + let flipped = index.min(corrupt.len() - 1); + corrupt[flipped] ^= 1; assert!(matches!( - collect_aws(&length.to_be_bytes(), 1).await, - Err(Error::InvalidLength(_)) + collect(every(&corrupt, 3)).await, + Err(EventStreamError::Malformed(_)) )); } #[rstest] -#[case(1)] -#[case(3)] -#[case(5)] +#[case::zero(0)] +#[case::below_minimum(15)] +#[case::above_maximum(16 * 1024 * 1024 + 1)] +#[case::u32_max(u32::MAX)] #[tokio::test] -async fn rejects_truncation(payload_frame: Vec, #[case] end: usize) { +async fn a_length_outside_the_frame_bounds_fails_before_buffering(#[case] length: u32) { assert!(matches!( - collect_aws(&payload_frame[..end], 1).await, - Err(Error::Truncated) + collect(every(&length.to_be_bytes(), 1)).await, + Err(EventStreamError::InvalidLength(seen)) if seen == length as usize )); } + +#[rstest] +#[case::before_the_length(1)] +#[case::inside_the_prelude(5)] +#[case::one_byte_short(usize::MAX)] +#[tokio::test] +async fn eof_inside_a_frame_is_truncation(payload_frame: Vec, #[case] end: usize) { + let end = end.min(payload_frame.len() - 1); + assert!(matches!( + collect(every(&payload_frame[..end], 1)).await, + Err(EventStreamError::Truncated) + )); +} + +const FRAME_OVERHEAD_BYTES: usize = 16; +const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024; + +#[tokio::test] +async fn a_frame_at_exactly_the_maximum_length_decodes() { + let largest = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES]); + let wire = encode_all(AwsEventStreamCodec, [largest.clone()]); + assert_eq!(wire.len(), MAX_FRAME_BYTES); + assert_eq!(collect(every(&wire, 1 << 20)).await.unwrap(), vec![largest]); +} + +#[tokio::test] +async fn a_frame_one_byte_over_the_maximum_length_is_rejected_by_its_prelude() { + let oversized = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES + 1]); + let wire = encode_all(AwsEventStreamCodec, [oversized]); + assert!(matches!( + collect(every(&wire[..4], 1)).await, + Err(EventStreamError::InvalidLength(length)) if length == MAX_FRAME_BYTES + 1 + )); +} + +#[tokio::test] +async fn an_empty_body_yields_nothing() { + assert_eq!(collect(vec![]).await.unwrap(), vec![]); +} + +#[tokio::test] +async fn a_complete_frame_precedes_a_truncated_following_frame() { + let wire = encode_all(AwsEventStreamCodec, [message(b"first"), message(b"second")]); + let mut messages = Box::pin(frames( + input(every(&wire[..wire.len() - 1], 3)), + AwsEventStreamCodec, + )); + + assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first")); + assert!(matches!( + messages.next().await, + Some(Err(EventStreamError::Truncated)) + )); + assert!(messages.next().await.is_none()); +} + +#[tokio::test] +async fn a_body_error_after_a_complete_frame_preserves_its_cause() { + let first = encode_all(AwsEventStreamCodec, [message(b"first")]); + let mut messages = Box::pin(frames( + stream::iter([ + Ok(cut_at(&first, [5])[0].clone()), + Ok(cut_at(&first, [5])[1].clone()), + Ok(Bytes::from_static(b"\0\0\0")), + Err(io::Error::new(io::ErrorKind::ConnectionReset, "reset")), + ]), + AwsEventStreamCodec, + )); + + assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first")); + let Some(Err(EventStreamError::Body(body))) = messages.next().await else { + panic!("the body error surfaces"); + }; + assert_eq!( + body_cause::(&body).unwrap().kind(), + io::ErrorKind::ConnectionReset + ); + assert!(messages.next().await.is_none()); +} diff --git a/litellm-rust/crates/framer/tests/chaining.rs b/litellm-rust/crates/framer/tests/chaining.rs index afd24a90704..81884d58ba1 100644 --- a/litellm-rust/crates/framer/tests/chaining.rs +++ b/litellm-rust/crates/framer/tests/chaining.rs @@ -2,28 +2,64 @@ mod support; -use std::io; +use bytes::Bytes; +use futures_util::{StreamExt, TryStreamExt}; +use litellm_framing::{ + EventStreamError, SseError, + aws_event_stream::{AwsEventStreamCodec, Message}, + frames, + sse::{SseCodec, SseEvent}, +}; +use proptest::prelude::*; +use support::{body_cause, cut_at, encode_all, every, input, runtime}; -use futures_util::TryStreamExt; -use litellm_framing::Framer; -use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer}; -use litellm_framing::sse::SseFramer; +fn delta(data: &str) -> SseEvent { + SseEvent { + event: Some("delta".into()), + data: data.into(), + id: Some("7".into()), + retry: None, + } +} -use support::encode; +fn envelopes(payloads: Vec) -> Vec { + encode_all(AwsEventStreamCodec, payloads.into_iter().map(Message::new)) +} + +proptest! { + #[test] + fn an_sse_event_cut_anywhere_across_envelopes_is_reassembled(cut in 0_usize..64, chunk in 1_usize..8) { + let sse = encode_all(SseCodec::default(), [delta("hello")]); + let wire = envelopes(cut_at(&sse, [cut.min(sse.len())])); + let events = runtime().block_on(async { + let payloads = frames(input(every(&wire, chunk)), AwsEventStreamCodec) + .map_ok(|message| message.payload().clone()); + frames(payloads, SseCodec::default()).try_collect::>().await + }) + .unwrap(); + prop_assert_eq!(events, vec![delta("hello")]); + } +} #[tokio::test] -async fn hosting_payloads_feed_the_same_sse_framer_across_envelope_boundaries() { - let bytes = [encode(b"event: delta\ndata: hel"), encode(b"lo\nid: 7\n\n")].concat(); - let envelopes = AwsEventStreamFramer.frame(futures_util::stream::iter( - bytes.chunks(3).map(Ok::<_, io::Error>), +async fn a_truncated_envelope_after_an_sse_event_keeps_the_event_and_its_cause() { + let complete = encode_all(SseCodec::default(), [delta("complete")]); + let incomplete = encode_all(SseCodec::default(), [delta("incomplete")]); + let wire = envelopes(vec![complete.into(), incomplete.into()]); + let payloads = frames( + input(every(&wire[..wire.len() - 1], 3)), + AwsEventStreamCodec, + ) + .map_ok(|message| message.payload().clone()); + let mut events = Box::pin(frames(payloads, SseCodec::default())); + + assert_eq!(events.next().await.unwrap().unwrap(), delta("complete")); + let Some(Err(SseError::Body(body))) = events.next().await else { + panic!("the envelope error surfaces through the SSE layer"); + }; + assert!(matches!( + body_cause::(&body), + Some(EventStreamError::Truncated) )); - let frames = SseFramer - .frame(envelopes.map_ok(|frame: AwsEventStreamFrame| frame.payload)) - .try_collect::>() - .await - .unwrap(); - assert_eq!(frames.len(), 1); - assert_eq!(frames[0].event.as_deref(), Some("delta")); - assert_eq!(frames[0].data.as_deref(), Some("hello")); - assert_eq!(frames[0].id.as_deref(), Some("7")); + assert!(events.next().await.is_none()); } diff --git a/litellm-rust/crates/framer/tests/sse.rs b/litellm-rust/crates/framer/tests/sse.rs index 66339dfbfd2..2fa064653a6 100644 --- a/litellm-rust/crates/framer/tests/sse.rs +++ b/litellm-rust/crates/framer/tests/sse.rs @@ -1,67 +1,169 @@ #![cfg(feature = "sse")] +mod support; + use std::io; -use futures_util::{StreamExt, TryStreamExt}; -use litellm_framing::sse::{SseFrame, SseFramer}; -use litellm_framing::{Error, Framer}; +use bytes::Bytes; +use futures_util::{StreamExt, TryStreamExt, stream}; +use litellm_framing::{ + SseError, frames, + sse::{SseCodec, SseEvent}, +}; +use proptest::prelude::*; use rstest::rstest; +use support::{body_cause, cut_at, encode_all, every, input, runtime}; -async fn collect_sse(chunks: &[&[u8]]) -> Result, Error> { - SseFramer - .frame(futures_util::stream::iter( - chunks.iter().copied().map(Ok::<_, io::Error>), - )) +async fn collect(pieces: Vec) -> Result, SseError> { + frames(input(pieces), SseCodec::default()) .try_collect() .await } +fn event(name: Option<&str>, data: &str) -> SseEvent { + SseEvent { + event: name.map(str::to_owned), + data: data.to_owned(), + id: None, + retry: None, + } +} + +fn sse_event() -> impl Strategy { + ( + proptest::option::of("[^\r\n\0]{0,8}"), + "[^\r\0]{0,16}", + proptest::option::of("[^\r\n\0]{0,8}"), + proptest::option::of(any::()), + ) + .prop_map(|(event, data, id, retry)| SseEvent { + event, + data, + id, + retry, + }) +} + +fn terminators() -> impl Strategy { + prop_oneof![Just(&b"\n"[..]), Just(&b"\r\n"[..]), Just(&b"\r"[..])] +} + +proptest! { + #[test] + fn any_events_survive_a_round_trip_through_any_terminator_and_any_cuts( + events in proptest::collection::vec(sse_event(), 1..4), + terminator in terminators(), + cuts in proptest::collection::vec(0_usize..256, 0..4), + bom in any::(), + ) { + let lf_wire = encode_all(SseCodec::default(), events.clone()); + let body: Vec = lf_wire + .iter() + .flat_map(|byte| if *byte == b'\n' { terminator.to_vec() } else { vec![*byte] }) + .collect(); + let wire = if bom { [&b"\xEF\xBB\xBF"[..], &body].concat() } else { body }; + let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap(); + prop_assert_eq!(decoded, events); + } +} + #[rstest] -#[case( - &[&b":ping\r\nevent: delta\r\nid: 7\r\nretry: 10\r\ndata: \xe2"[..], &b"\x82"[..], &b"\xac\r"[..], &b"\ndata: next\r\n\r"[..], &b"\ndata: [DONE]\n\n"[..]], - vec![ - SseFrame { - event: Some("delta".into()), - data: Some("€\nnext".into()), - id: Some("7".into()), - retry: Some(10), - }, - SseFrame { - event: None, - data: Some("[DONE]".into()), - id: None, - retry: None, - }, - ] -)] +#[case::comment(b":ping\ndata: x\n\n")] +#[case::unknown_field(b"vendor: 1\ndata: x\n\n")] +#[case::field_without_colon(b"garbage\ndata: x\n\n")] +#[case::retry_with_non_digits(b"retry: soon\ndata: x\n\n")] +#[case::retry_with_a_sign(b"retry: +5\ndata: x\n\n")] +#[case::retry_without_a_value(b"retry:\ndata: x\n\n")] +#[case::id_with_nul(b"id: a\0b\ndata: x\n\n")] #[tokio::test] -async fn fragmented_utf8_crlf_and_multiline_data_retain_metadata_and_sentinel( - #[case] chunks: &[&[u8]], - #[case] expected: Vec, +async fn lines_the_spec_ignores_do_not_change_the_event(#[case] wire: &[u8]) { + assert_eq!( + collect(every(wire, 1)).await.unwrap(), + vec![event(None, "x")] + ); +} + +#[rstest] +#[case::no_data_at_all(b"event: ping\nid: 1\n\ndata: x\n\n", vec![event(None, "x")])] +#[case::empty_data_field(b"data:\n\n", vec![event(None, "")])] +#[case::one_leading_space_stripped(b"data: x\n\n", vec![event(None, " x")])] +#[case::multiline_data(b"data: a\ndata: b\ndata:\n\n", vec![event(None, "a\nb\n")])] +#[case::last_event_name_wins(b"event: a\nevent: b\ndata: x\n\n", vec![event(Some("b"), "x")])] +#[case::last_retry_wins(b"retry: 1\nretry: 2\ndata: x\n\n", vec![SseEvent { retry: Some(2), ..event(None, "x") }])] +#[case::split_utf8_across_lines_is_not_joined(b"data: \xe2\x82\xac\ndata: \xe2\x82\xac\n\n", vec![event(None, "€\n€")])] +#[tokio::test] +async fn dispatch_follows_the_data_buffer(#[case] wire: &[u8], #[case] expected: Vec) { + assert_eq!(collect(every(wire, 1)).await.unwrap(), expected); +} + +#[rstest] +#[case::unterminated_single(b"data: partial\n", vec![])] +#[case::unterminated_tail_after_complete(b"data: complete\n\ndata: unfinished\n", vec![event(None, "complete")])] +#[case::lone_cr_terminates_at_eof(b"data: x\r\r", vec![event(None, "x")])] +#[case::lone_cr_line_then_eof(b"data: x\r", vec![])] +#[tokio::test] +async fn eof_dispatches_only_terminated_events( + #[case] wire: &[u8], + #[case] expected: Vec, ) { - assert_eq!(collect_sse(chunks).await.unwrap(), expected); + assert_eq!( + collect(vec![Bytes::copy_from_slice(wire)]).await.unwrap(), + expected + ); +} + +#[rstest] +#[case::inside_the_first_line(vec![&b"data: a\r"[..], &b"\ndata: b\r\n\r\n"[..]])] +#[case::inside_the_blank_line(vec![&b"data: a\r\ndata: b\r\n\r"[..], &b"\n"[..]])] +#[tokio::test] +async fn a_crlf_split_across_chunks_is_one_terminator(#[case] pieces: Vec<&[u8]>) { + let pieces = pieces.into_iter().map(Bytes::copy_from_slice).collect(); + assert_eq!(collect(pieces).await.unwrap(), vec![event(None, "a\nb")]); } #[tokio::test] -async fn eof_does_not_dispatch_an_unterminated_frame() { - assert!(collect_sse(&[b"data: partial\n"]).await.unwrap().is_empty()); +async fn a_bom_is_stripped_only_at_the_start_of_the_stream() { + let wire = b"\xEF\xBB\xBFdata: a\n\n\xEF\xBB\xBFdata: b\ndata: c\n\n"; + let decoded = collect(every(wire, 2)).await.unwrap(); + assert_eq!(decoded, vec![event(None, "a"), event(None, "c")]); +} + +#[tokio::test] +async fn invalid_utf8_in_a_field_fails_after_earlier_events_and_terminates() { + let mut events = Box::pin(frames( + input(every(b"data: ok\n\ndata: \xff\n\n", 3)), + SseCodec::default(), + )); + + assert_eq!(events.next().await.unwrap().unwrap(), event(None, "ok")); + assert!(matches!( + events.next().await, + Some(Err(SseError::InvalidUtf8(_))) + )); + assert!(events.next().await.is_none()); } #[rstest] #[case(io::ErrorKind::ConnectionReset)] #[case(io::ErrorKind::UnexpectedEof)] #[tokio::test] -async fn framing_errors_terminate_and_preserve_input_error_causes(#[case] kind: io::ErrorKind) { - let mut frames = Box::pin(SseFramer.frame(futures_util::stream::iter([ - Err(io::Error::new(kind, "reset")), - Ok(&b"data: later\n\n"[..]), - ]))); - let error = frames.next().await.unwrap().unwrap_err(); - assert!(matches!( - error, - Error::Sse(sse_stream::Error::Body(ref cause)) - if cause.downcast_ref::().unwrap().kind() == kind +async fn a_body_error_keeps_earlier_events_and_its_cause_then_terminates( + #[case] kind: io::ErrorKind, +) { + let mut events = Box::pin(frames( + stream::iter([ + Ok(&b"data: first\n\ndata: partial"[..]), + Err(io::Error::new(kind, "reset")), + Ok(&b"\n\n"[..]), + ]), + SseCodec::default(), )); - assert!(frames.next().await.is_none()); - assert!(frames.next().await.is_none()); + + assert_eq!(events.next().await.unwrap().unwrap(), event(None, "first")); + let Some(Err(SseError::Body(body))) = events.next().await else { + panic!("the body error surfaces"); + }; + assert_eq!(body_cause::(&body).unwrap().kind(), kind); + assert!(events.next().await.is_none()); + assert!(events.next().await.is_none()); } diff --git a/litellm-rust/crates/framer/tests/support/mod.rs b/litellm-rust/crates/framer/tests/support/mod.rs index 9db305af073..9ff67aef149 100644 --- a/litellm-rust/crates/framer/tests/support/mod.rs +++ b/litellm-rust/crates/framer/tests/support/mod.rs @@ -1,15 +1,57 @@ -use aws_smithy_eventstream::frame::write_message_to; -use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; -use bytes::Bytes; +#![allow(dead_code)] -pub fn encode(payload: &'static [u8]) -> Vec { - let message = Message::new(Bytes::from_static(payload)) - .add_header(Header::new( - ":event-type", - HeaderValue::String("payload".into()), - )) - .add_header(Header::new("sequence", HeaderValue::Int32(7))); - let mut bytes = Vec::new(); - write_message_to(&message, &mut bytes).unwrap(); - bytes +use std::{error::Error, io}; + +use bytes::{Bytes, BytesMut}; +use futures_util::{Stream, stream}; +use tokio_util::codec::Encoder; + +pub fn encode_all(mut codec: C, items: impl IntoIterator) -> Vec +where + C: Encoder, + C::Error: std::fmt::Debug, +{ + let mut wire = BytesMut::new(); + for item in items { + codec.encode(item, &mut wire).unwrap(); + } + wire.to_vec() +} + +pub fn cut_at(bytes: &[u8], offsets: impl IntoIterator) -> Vec { + let mut sorted: Vec = offsets + .into_iter() + .filter(|offset| *offset <= bytes.len()) + .collect(); + sorted.sort_unstable(); + sorted.dedup(); + let bounds = std::iter::once(0) + .chain(sorted) + .chain(std::iter::once(bytes.len())) + .collect::>(); + bounds + .windows(2) + .map(|pair| Bytes::copy_from_slice(&bytes[pair[0]..pair[1]])) + .collect() +} + +pub fn every(bytes: &[u8], size: usize) -> Vec { + bytes + .chunks(size.max(1)) + .map(Bytes::copy_from_slice) + .collect() +} + +pub fn input(pieces: Vec) -> impl Stream> + Send { + stream::iter(pieces.into_iter().map(Ok)) +} + +pub fn body_cause(body: &io::Error) -> Option<&T> { + body.get_ref()?.downcast_ref::() +} + +pub fn runtime() -> tokio::runtime::Runtime { + tokio::runtime::Builder::new_current_thread() + .build() + .unwrap() } diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs index 35e7d5820b0..3f1b7ed9bcc 100644 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs +++ b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs @@ -2,9 +2,9 @@ use base64::Engine; use bytes::Buf; use futures_util::{Stream, StreamExt}; use litellm_framing::{ - Framer, - aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer}, - sse::{SseFrame, SseFramer}, + aws_event_stream::{AwsEventStreamCodec, Message}, + frames, + sse::{SseCodec, SseEvent}, }; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; @@ -13,8 +13,6 @@ use serde_json::{Map, Value}; pub enum Error { #[error("stream framing failed: {0}")] StreamFraming(String), - #[error("Anthropic SSE frame has no data")] - MissingStreamData, #[error("Anthropic stream event is invalid: {0}")] InvalidStreamEvent(String), #[error("Bedrock event payload is invalid: {0}")] @@ -165,15 +163,14 @@ struct BedrockChunkPayload { bytes: String, } -pub fn decode_anthropic_sse_frame(frame: SseFrame) -> Result { - let data = frame.data.ok_or(Error::MissingStreamData)?; - serde_json::from_str(&data).map_err(|error| Error::InvalidStreamEvent(error.to_string())) +pub fn decode_anthropic_sse_frame(event: SseEvent) -> Result { + serde_json::from_str(&event.data).map_err(|error| Error::InvalidStreamEvent(error.to_string())) } pub fn decode_bedrock_anthropic_frame( - frame: AwsEventStreamFrame, + message: Message, ) -> Result { - let payload: BedrockChunkPayload = serde_json::from_slice(&frame.payload) + let payload: BedrockChunkPayload = serde_json::from_slice(message.payload()) .map_err(|error| Error::InvalidBedrockPayload(error.to_string()))?; let event = base64::engine::general_purpose::STANDARD .decode(payload.bytes) @@ -189,9 +186,8 @@ where B: Buf + Send, E: std::error::Error + Send + Sync + 'static, { - SseFramer.frame(input).map(|frame| { - let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?; - decode_anthropic_sse_frame(frame) + frames(input, SseCodec::default()).map(|event| { + decode_anthropic_sse_frame(event.map_err(|error| Error::StreamFraming(error.to_string()))?) }) } @@ -203,9 +199,10 @@ where B: Buf + Send, E: std::error::Error + Send + Sync + 'static, { - AwsEventStreamFramer.frame(input).map(|frame| { - let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?; - decode_bedrock_anthropic_frame(frame) + frames(input, AwsEventStreamCodec).map(|message| { + decode_bedrock_anthropic_frame( + message.map_err(|error| Error::StreamFraming(error.to_string()))?, + ) }) } @@ -247,12 +244,10 @@ mod tests { #[test] fn decodes_citations_delta_events() { - let event = decode_anthropic_sse_frame(SseFrame { + let event = decode_anthropic_sse_frame(SseEvent { event: Some("content_block_delta".into()), - data: Some( - r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"# - .into(), - ), + data: r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"# + .into(), id: None, retry: None, })