mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
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:
parent
88fd15315c
commit
3c93ea1697
13 changed files with 657 additions and 275 deletions
16
litellm-rust/Cargo.lock
generated
16
litellm-rust/Cargo.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)?)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
|
|
|||
21
litellm-rust/crates/framer/src/framed.rs
Normal file
21
litellm-rust/crates/framer/src/framed.rs
Normal 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()
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue