#[cfg(test)]
use std::cell::RefCell;
use std::pin::Pin;
use std::task::{Context, Poll};
use bytes::Bytes;
#[cfg(test)]
use bytes::BytesMut;
use futures_core::Stream;
use memchr::memchr;
use pin_project_lite::pin_project;
use crate::error::{LiterLlmError, Result};
use crate::http::request::with_retry;
use crate::types::ChatCompletionChunk;
const MAX_BUFFER_BYTES: usize = 1024 * 1024;
#[cfg(test)]
const MAX_POOL_BUFFER_CAPACITY: usize = 64 * 1024;
#[cfg(test)]
thread_local! {
static BYTES_POOL: RefCell<Option<BytesMut>> = const { RefCell::new(None) };
}
#[cfg(test)]
fn pool_acquire() -> BytesMut {
BYTES_POOL.with(|cell| {
cell.borrow_mut()
.take()
.map(|mut buf| {
buf.clear();
buf
})
.unwrap_or_else(|| BytesMut::with_capacity(4096))
})
}
#[cfg(test)]
fn pool_release(buf: BytesMut) {
if buf.capacity() <= MAX_POOL_BUFFER_CAPACITY {
BYTES_POOL.with(|cell| {
*cell.borrow_mut() = Some(buf);
});
}
}
#[cfg(feature = "native-http")]
pub use tokio_util::sync::CancellationToken;
#[cfg_attr(
feature = "tracing",
tracing::instrument(
skip_all,
fields(
http.method = "POST",
http.url = %url,
http.status_code = tracing::field::Empty,
http.retry_count = tracing::field::Empty,
)
)
)]
pub async fn post_stream<P>(
client: &reqwest::Client,
url: &str,
auth_header: Option<(&str, &str)>,
extra_headers: &[(&str, &str)],
body: Bytes,
max_retries: u32,
parse_event: P,
) -> Result<crate::client::BoxStream<'static, Result<ChatCompletionChunk>>>
where
P: Fn(&str) -> Result<Option<ChatCompletionChunk>> + Send + 'static,
{
let mut retry_count = 0u32;
let resp = with_retry(max_retries, || {
let mut builder = client
.post(url)
.header(reqwest::header::CONTENT_TYPE, "application/json")
.body(body.clone());
if let Some((name, value)) = auth_header {
builder = builder.header(name, value);
}
for (name, value) in extra_headers {
builder = builder.header(*name, *value);
}
retry_count += 1;
builder.send()
})
.await?;
#[cfg(feature = "tracing")]
{
let span = tracing::Span::current();
span.record("http.status_code", resp.status().as_u16());
span.record("http.retry_count", retry_count.saturating_sub(1));
}
let byte_stream = resp.bytes_stream();
let stream = SseParser::new(byte_stream, parse_event, None);
Ok(Box::pin(stream))
}
#[cfg(feature = "native-http")]
#[allow(dead_code)] #[allow(clippy::too_many_arguments)] #[cfg_attr(
feature = "tracing",
tracing::instrument(
skip_all,
fields(
http.method = "POST",
http.url = %url,
http.status_code = tracing::field::Empty,
http.retry_count = tracing::field::Empty,
)
)
)]
pub async fn post_stream_with_cancel<P>(
client: &reqwest::Client,
url: &str,
auth_header: Option<(&str, &str)>,
extra_headers: &[(&str, &str)],
body: Bytes,
max_retries: u32,
parse_event: P,
cancel: CancellationToken,
) -> Result<crate::client::BoxStream<'static, Result<ChatCompletionChunk>>>
where
P: Fn(&str) -> Result<Option<ChatCompletionChunk>> + Send + 'static,
{
let mut retry_count = 0u32;
let resp = with_retry(max_retries, || {
let mut builder = client
.post(url)
.header(reqwest::header::CONTENT_TYPE, "application/json")
.body(body.clone());
if let Some((name, value)) = auth_header {
builder = builder.header(name, value);
}
for (name, value) in extra_headers {
builder = builder.header(*name, *value);
}
retry_count += 1;
builder.send()
})
.await?;
#[cfg(feature = "tracing")]
{
let span = tracing::Span::current();
span.record("http.status_code", resp.status().as_u16());
span.record("http.retry_count", retry_count.saturating_sub(1));
}
let byte_stream = resp.bytes_stream();
let stream = SseParser::new(byte_stream, parse_event, Some(cancel));
Ok(Box::pin(stream))
}
#[cfg(feature = "native-http")]
type CancelField = Option<CancellationToken>;
#[cfg(not(feature = "native-http"))]
type CancelField = Option<std::convert::Infallible>;
pin_project! {
struct SseParser<S, P> {
#[pin]
inner: S,
buffer: String,
cursor: usize,
done: bool,
parse_event: P,
cancel: CancelField,
}
impl<S, P> PinnedDrop for SseParser<S, P> {
fn drop(this: Pin<&mut Self>) {
let _ = this;
}
}
}
impl<S, P> SseParser<S, P>
where
P: Fn(&str) -> Result<Option<ChatCompletionChunk>>,
{
fn new(inner: S, parse_event: P, cancel: CancelField) -> Self {
Self {
inner,
buffer: String::with_capacity(4096),
cursor: 0,
done: false,
parse_event,
cancel,
}
}
}
impl<S, P> Stream for SseParser<S, P>
where
S: Stream<Item = std::result::Result<Bytes, reqwest::Error>>,
P: Fn(&str) -> Result<Option<ChatCompletionChunk>>,
{
type Item = Result<ChatCompletionChunk>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let mut this = self.project();
#[cfg(feature = "native-http")]
if this.cancel.as_ref().is_some_and(|t| t.is_cancelled()) {
#[cfg(feature = "tracing")]
tracing::debug!("SSE stream cancelled by downstream disconnect");
*this.done = true;
return Poll::Ready(None);
}
loop {
if let Some(offset) = memchr(b'\n', &this.buffer.as_bytes()[*this.cursor..]) {
let newline_pos = *this.cursor + offset;
let line = this.buffer[*this.cursor..newline_pos].trim_end_matches('\r').trim();
if line.is_empty() || line.starts_with(':') {
*this.cursor = newline_pos + 1;
compact_if_needed(this.buffer, this.cursor);
continue;
}
if let Some(raw) = line.strip_prefix("data:") {
let data = raw.strip_prefix(' ').unwrap_or(raw).trim();
if data == "[DONE]" {
*this.cursor = newline_pos + 1;
compact_if_needed(this.buffer, this.cursor);
return Poll::Ready(None);
}
let result = (this.parse_event)(data);
*this.cursor = newline_pos + 1;
compact_if_needed(this.buffer, this.cursor);
match result {
Ok(None) => continue,
Ok(Some(chunk)) => return Poll::Ready(Some(Ok(chunk))),
Err(e) => return Poll::Ready(Some(Err(e))),
}
}
*this.cursor = newline_pos + 1;
compact_if_needed(this.buffer, this.cursor);
continue;
}
if *this.done {
let remaining = this.buffer.len() - *this.cursor;
if remaining > 0 {
#[cfg(feature = "tracing")]
tracing::warn!(
leftover_bytes = remaining,
preview = &this.buffer[*this.cursor..(*this.cursor + remaining.min(64))],
"SSE stream ended with unterminated data in buffer; dropping partial line"
);
this.buffer.clear();
*this.cursor = 0;
}
return Poll::Ready(None);
}
#[cfg(feature = "native-http")]
if this.cancel.as_ref().is_some_and(|t| t.is_cancelled()) {
#[cfg(feature = "tracing")]
tracing::debug!("SSE stream cancelled while waiting for next chunk");
*this.done = true;
return Poll::Ready(None);
}
match this.inner.as_mut().poll_next(cx) {
Poll::Ready(Some(Ok(bytes))) => {
if this.buffer.len() + bytes.len() > MAX_BUFFER_BYTES {
*this.done = true;
return Poll::Ready(Some(Err(LiterLlmError::Streaming {
message: format!("SSE buffer exceeded {MAX_BUFFER_BYTES} bytes; stream aborted"),
})));
}
match std::str::from_utf8(&bytes) {
Ok(s) => this.buffer.push_str(s),
Err(e) => {
*this.done = true;
return Poll::Ready(Some(Err(LiterLlmError::Streaming {
message: format!("invalid UTF-8 in SSE stream: {e}"),
})));
}
}
}
Poll::Ready(Some(Err(e))) => {
return Poll::Ready(Some(Err(LiterLlmError::from(e))));
}
Poll::Ready(None) => {
*this.done = true;
continue;
}
Poll::Pending => {
return Poll::Pending;
}
}
}
}
}
fn compact_if_needed(buffer: &mut String, cursor: &mut usize) {
if *cursor > buffer.len() / 2 {
buffer.drain(..*cursor);
*cursor = 0;
}
}
#[cfg(test)]
pub(crate) fn parse_sse_line(line: &str) -> Option<Result<ChatCompletionChunk>> {
let raw = line.strip_prefix("data:")?;
let data = raw.strip_prefix(' ').unwrap_or(raw).trim();
if data == "[DONE]" {
return None;
}
Some(serde_json::from_str(data).map_err(|e| LiterLlmError::Streaming {
message: format!("failed to parse SSE data: {e}"),
}))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bytes_pool_reuses_buffer() {
let buf = pool_acquire();
let ptr_before = buf.as_ptr();
pool_release(buf);
let buf2 = pool_acquire();
let ptr_after = buf2.as_ptr();
assert_eq!(ptr_before, ptr_after, "pool should reuse the same BytesMut allocation");
pool_release(buf2);
}
#[test]
fn bytes_pool_discards_oversized_buffers() {
let mut big = BytesMut::with_capacity(MAX_POOL_BUFFER_CAPACITY + 1);
big.resize(MAX_POOL_BUFFER_CAPACITY + 1, 0u8);
pool_release(big);
let acquired = pool_acquire();
assert!(
acquired.capacity() <= 4096,
"oversized buffer should have been discarded; got capacity {}",
acquired.capacity()
);
pool_release(acquired);
}
#[test]
fn sse_parser_does_not_hold_idle_buffer_when_stream_idle() {
use std::pin::Pin;
use std::task::{Context, Poll};
use bytes::Bytes;
use futures_core::Stream;
struct NeverStream;
impl Stream for NeverStream {
type Item = std::result::Result<Bytes, reqwest::Error>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Pending
}
}
let sentinel = pool_acquire();
let sentinel_ptr = sentinel.as_ptr();
pool_release(sentinel);
let parser = SseParser::new(
NeverStream,
|_data: &str| -> Result<Option<crate::types::ChatCompletionChunk>> { Ok(None) },
None,
);
let reclaimed = pool_acquire();
assert_eq!(
reclaimed.as_ptr(),
sentinel_ptr,
"SseParser constructor must not acquire a BytesMut from the pool"
);
pool_release(reclaimed);
drop(parser);
let after_drop = pool_acquire();
assert_eq!(
after_drop.as_ptr(),
sentinel_ptr,
"SseParser Drop must not corrupt the pool slot"
);
pool_release(after_drop);
}
#[cfg(feature = "native-http")]
#[tokio::test]
async fn cancellation_aborts_upstream() {
use std::pin::Pin;
use std::task::{Context, Poll};
use bytes::Bytes;
use futures_core::Stream;
use crate::types::ChatCompletionChunk;
struct InfiniteStream;
impl Stream for InfiniteStream {
type Item = std::result::Result<Bytes, reqwest::Error>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Pending
}
}
let token = CancellationToken::new();
let token_clone = token.clone();
let mut parser: Pin<Box<SseParser<_, _>>> = Box::pin(SseParser::new(
InfiniteStream,
|_data: &str| -> Result<Option<ChatCompletionChunk>> { Ok(None) },
Some(token_clone),
));
token.cancel();
let result = futures_util::StreamExt::next(&mut parser).await;
assert!(result.is_none(), "cancelled stream should return None immediately");
}
}