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 <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-25 12:18:32 -07:00 • committed by GitHub
parent 88fd15315c
commit 3c93ea1697
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 657 additions and 275 deletions

View file

@ -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"

View file

@ -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()
},

View file

@ -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

View file

@ -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<Header>,
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<S, B, E>(self, input: S) -> impl Stream<Item = Result<Self::Frame, Error>> + Send
where
S: Stream<Item = Result<B, E>> + 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<Option<Message>, 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<Option<Message>, EventStreamError> {
match self.decode(src)? {
Some(message) => Ok(Some(message)),
None if src.is_empty() => Ok(None),
None => Err(EventStreamError::Truncated),
}
}
}
impl Encoder<Message> for AwsEventStreamCodec {
type Error = EventStreamError;
fn encode(&mut self, message: Message, dst: &mut BytesMut) -> Result<(), EventStreamError> {
Ok(write_message_to(&message, dst)?)
}
}

View file

@ -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<dyn std::error::Error + Send + Sync>),
#[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),
}

View file

@ -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<S, B, E, D>(
input: S,
codec: D,
) -> impl Stream<Item = Result<D::Item, D::Error>> + Send
where
S: Stream<Item = Result<B, E>> + 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()
}

View file

@ -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;

View file

@ -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<String>,
pub data: Option<String>,
pub data: String,
pub id: Option<String>,
pub retry: Option<u64>,
}
#[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<S, B, E>(self, input: S) -> impl Stream<Item = Result<SseFrame, Error>> + Send
where
S: Stream<Item = Result<B, E>> + 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<Option<SseEvent>, 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<Option<SseEvent>, 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<usize> {
lines(bytes)
.find(|(line, _)| line.is_empty())
.map(|(_, end)| end)
}
fn lines(bytes: &[u8]) -> impl Iterator<Item = (&[u8], usize)> {
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<String>,
data: Option<String>,
id: Option<String>,
retry: Option<u64>,
}
impl Pending {
fn apply(self, line: &[u8]) -> Result<Self, SseError> {
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<SseEvent> {
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<String>, line: &str) -> String {
match buffer {
Some(existing) => format!("{existing}\n{line}"),
None => line.to_owned(),
}
}
fn parse_retry(value: &[u8]) -> Option<u64> {
if !value.iter().all(u8::is_ascii_digit) {
return None;
}
str::from_utf8(value).ok()?.parse().ok()
}
impl Encoder<SseEvent> 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(())
}
}

View file

@ -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<Vec<AwsEventStreamFrame>, Error> {
AwsEventStreamFramer
.frame(futures_util::stream::iter(
bytes.chunks(chunk_size).map(Ok::<_, io::Error>),
))
async fn collect(pieces: Vec<Bytes>) -> Result<Vec<Message>, EventStreamError> {
frames(input(pieces), AwsEventStreamCodec)
.try_collect()
.await
}
#[fixture]
fn two_frames() -> Vec<u8> {
[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<u8> {
encode(b"payload")
encode_all(AwsEventStreamCodec, [message(b"payload")])
}
fn header_value() -> impl Strategy<Value = HeaderValue> {
prop_oneof![
"[a-z]{0,8}".prop_map(|text| HeaderValue::String(text.into())),
any::<i32>().prop_map(HeaderValue::Int32),
any::<bool>().prop_map(HeaderValue::Bool),
proptest::collection::vec(any::<u8>(), 0..8)
.prop_map(|bytes| HeaderValue::ByteArray(bytes.into())),
]
}
fn arbitrary_message() -> impl Strategy<Value = Message> {
(
proptest::collection::vec(("[a-z:-]{1,12}", header_value()), 0..3),
proptest::collection::vec(any::<u8>(), 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<u8>,
#[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<u8>, #[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<u8>, #[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<u8>, #[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<u8>, #[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::<io::Error>(&body).unwrap().kind(),
io::ErrorKind::ConnectionReset
);
assert!(messages.next().await.is_none());
}

View file

@ -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<Bytes>) -> Vec<u8> {
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::<Vec<_>>().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::<EventStreamError>(&body),
Some(EventStreamError::Truncated)
));
let frames = SseFramer
.frame(envelopes.map_ok(|frame: AwsEventStreamFrame| frame.payload))
.try_collect::<Vec<_>>()
.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());
}

View file

@ -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<Vec<SseFrame>, Error> {
SseFramer
.frame(futures_util::stream::iter(
chunks.iter().copied().map(Ok::<_, io::Error>),
))
async fn collect(pieces: Vec<Bytes>) -> Result<Vec<SseEvent>, 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<Value = SseEvent> {
(
proptest::option::of("[^\r\n\0]{0,8}"),
"[^\r\0]{0,16}",
proptest::option::of("[^\r\n\0]{0,8}"),
proptest::option::of(any::<u64>()),
)
.prop_map(|(event, data, id, retry)| SseEvent {
event,
data,
id,
retry,
})
}
fn terminators() -> impl Strategy<Value = &'static [u8]> {
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::<bool>(),
) {
let lf_wire = encode_all(SseCodec::default(), events.clone());
let body: Vec<u8> = 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<SseFrame>,
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<SseEvent>) {
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<SseEvent>,
) {
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::<io::Error>().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::<io::Error>(&body).unwrap().kind(), kind);
assert!(events.next().await.is_none());
assert!(events.next().await.is_none());
}

View file

@ -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<u8> {
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<C, I>(mut codec: C, items: impl IntoIterator<Item = I>) -> Vec<u8>
where
C: Encoder<I>,
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<Item = usize>) -> Vec<Bytes> {
let mut sorted: Vec<usize> = 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::<Vec<_>>();
bounds
.windows(2)
.map(|pair| Bytes::copy_from_slice(&bytes[pair[0]..pair[1]]))
.collect()
}
pub fn every(bytes: &[u8], size: usize) -> Vec<Bytes> {
bytes
.chunks(size.max(1))
.map(Bytes::copy_from_slice)
.collect()
}
pub fn input(pieces: Vec<Bytes>) -> impl Stream<Item = Result<Bytes, io::Error>> + Send {
stream::iter(pieces.into_iter().map(Ok))
}
pub fn body_cause<T: Error + 'static>(body: &io::Error) -> Option<&T> {
body.get_ref()?.downcast_ref::<T>()
}
pub fn runtime() -> tokio::runtime::Runtime {
tokio::runtime::Builder::new_current_thread()
.build()
.unwrap()
}

View file

@ -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<AnthropicMessagesStreamEvent, Error> {
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<AnthropicMessagesStreamEvent, Error> {
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<AnthropicMessagesStreamEvent, Error> {
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,
})