use std::collections::VecDeque;
use std::pin::Pin;
use futures_util::stream::{self, Stream, StreamExt};
use serde::Deserialize;
use crate::inference::error::InferenceError;
use crate::inference::types::{ChatResponse, StopReason, Usage, UsageBlock};
pub type ChatStream = Pin<Box<dyn Stream<Item = Result<ChatStreamEvent, InferenceError>> + Send>>;
#[derive(Debug, Clone, PartialEq)]
pub enum ChatStreamEvent {
Delta(String),
ToolCall(ToolCallDelta),
Done(StreamCompletion),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ToolCallDelta {
pub index: usize,
pub id: Option<String>,
pub name: Option<String>,
pub arguments: String,
}
#[derive(Debug, Clone, PartialEq)]
pub struct StreamCompletion {
pub finish_reason: Option<StopReason>,
pub usage: Usage,
}
#[derive(Debug, Default, Deserialize)]
struct StreamChunk {
#[serde(default)]
choices: Vec<StreamChoice>,
#[serde(default)]
usage: Option<UsageBlock>,
}
fn error_from_chunk(err: &serde_json::Value) -> InferenceError {
let message = err
.get("message")
.and_then(|m| m.as_str())
.filter(|s| !s.is_empty())
.unwrap_or("provider streaming error")
.to_string();
let code = err.get("code");
let status = code
.and_then(|c| c.as_u64())
.and_then(|n| u16::try_from(n).ok())
.unwrap_or(0);
let body = match code.and_then(|c| c.as_str()) {
Some(s) if status == 0 && !s.is_empty() => format!("{s}: {message}"),
_ => message,
};
InferenceError::Api { status, body }
}
#[derive(Debug, Default, Deserialize)]
struct StreamChoice {
#[serde(default)]
delta: WireDelta,
#[serde(default)]
finish_reason: Option<String>,
}
#[derive(Debug, Default, Deserialize)]
struct WireDelta {
#[serde(default)]
content: Option<String>,
#[serde(default)]
tool_calls: Vec<WireToolCall>,
}
#[derive(Debug, Default, Deserialize)]
struct WireToolCall {
#[serde(default)]
index: usize,
#[serde(default)]
id: Option<String>,
#[serde(default)]
function: Option<WireFunction>,
}
#[derive(Debug, Default, Deserialize)]
struct WireFunction {
#[serde(default)]
name: Option<String>,
#[serde(default)]
arguments: Option<String>,
}
#[derive(Debug, Default)]
pub struct SseDecoder {
buf: Vec<u8>,
data: String,
finish_reason: Option<StopReason>,
usage: Usage,
done: bool,
}
impl SseDecoder {
pub fn new() -> Self {
Self::default()
}
pub fn feed(&mut self, chunk: &[u8]) -> Vec<Result<ChatStreamEvent, InferenceError>> {
let mut out = Vec::new();
if self.done {
return out;
}
self.buf.extend_from_slice(chunk);
while let Some(nl) = self.buf.iter().position(|&b| b == b'\n') {
let line_bytes: Vec<u8> = self.buf.drain(..=nl).collect();
let line = String::from_utf8_lossy(&line_bytes);
let line = line.trim_end_matches(['\r', '\n']);
self.dispatch_line(line, &mut out);
if self.done {
break;
}
}
out
}
fn dispatch_line(
&mut self,
line: &str,
out: &mut Vec<Result<ChatStreamEvent, InferenceError>>,
) {
if line.is_empty() {
self.flush_event(out);
return;
}
if line.starts_with(':') {
return; }
let Some(rest) = line.strip_prefix("data:") else {
return; };
let value = rest.strip_prefix(' ').unwrap_or(rest);
if !self.data.is_empty() {
self.data.push('\n');
}
self.data.push_str(value);
self.flush_event(out);
}
fn flush_event(&mut self, out: &mut Vec<Result<ChatStreamEvent, InferenceError>>) {
if self.data.is_empty() {
return;
}
let payload = std::mem::take(&mut self.data);
let payload = payload.trim();
if payload.is_empty() {
return;
}
if payload == "[DONE]" {
out.push(Ok(self.terminal_event()));
self.done = true;
return;
}
let value: serde_json::Value = match serde_json::from_str(payload) {
Ok(v) => v,
Err(_) => return, };
if let Some(err) = value.get("error") {
out.push(Err(error_from_chunk(err)));
self.done = true;
return;
}
let chunk: StreamChunk = match serde_json::from_value(value) {
Ok(c) => c,
Err(_) => return, };
if let Some(block) = chunk.usage {
self.usage = block.into_usage();
}
for choice in chunk.choices {
if let Some(reason) = choice.finish_reason.as_deref() {
self.finish_reason = Some(StopReason::from_wire(reason));
}
if let Some(text) = choice.delta.content.filter(|s| !s.is_empty()) {
out.push(Ok(ChatStreamEvent::Delta(text)));
}
for tc in choice.delta.tool_calls {
let (name, arguments) = tc
.function
.map(|f| (f.name, f.arguments.unwrap_or_default()))
.unwrap_or((None, String::new()));
out.push(Ok(ChatStreamEvent::ToolCall(ToolCallDelta {
index: tc.index,
id: tc.id,
name,
arguments,
})));
}
}
}
fn terminal_event(&self) -> ChatStreamEvent {
ChatStreamEvent::Done(StreamCompletion {
finish_reason: self.finish_reason.clone(),
usage: self.usage,
})
}
pub fn finish(&mut self) -> Option<Result<ChatStreamEvent, InferenceError>> {
if self.done {
return None;
}
self.done = true;
let buf_incomplete = self.buf.iter().any(|b| !b.is_ascii_whitespace());
let data_incomplete = !self.data.trim().is_empty();
if buf_incomplete || data_incomplete {
return Some(Err(InferenceError::Transport(
"stream ended with an incomplete SSE frame".to_string(),
)));
}
Some(Ok(self.terminal_event()))
}
}
pub fn decode_event_stream<S, B, E>(byte_stream: S) -> ChatStream
where
S: Stream<Item = Result<B, E>> + Send + 'static,
B: AsRef<[u8]>,
E: std::fmt::Display,
{
struct State<S> {
inner: S,
decoder: SseDecoder,
queue: VecDeque<Result<ChatStreamEvent, InferenceError>>,
done: bool,
}
let init = State {
inner: Box::pin(byte_stream),
decoder: SseDecoder::new(),
queue: VecDeque::new(),
done: false,
};
let s = stream::unfold(init, |mut st| async move {
loop {
if let Some(ev) = st.queue.pop_front() {
return Some((ev, st));
}
if st.done {
return None;
}
match st.inner.next().await {
Some(Ok(bytes)) => {
st.queue.extend(st.decoder.feed(bytes.as_ref()));
}
Some(Err(e)) => {
st.done = true;
return Some((Err(InferenceError::Transport(e.to_string())), st));
}
None => {
st.done = true;
if let Some(res) = st.decoder.finish() {
return Some((res, st));
}
return None;
}
}
}
});
Box::pin(s)
}
pub fn buffered_stream(response: ChatResponse) -> ChatStream {
let mut events: Vec<Result<ChatStreamEvent, InferenceError>> = Vec::new();
if let Some(text) = response.first_text().filter(|s| !s.is_empty()) {
events.push(Ok(ChatStreamEvent::Delta(text)));
}
for (index, call) in response.first_tool_calls().iter().enumerate() {
events.push(Ok(ChatStreamEvent::ToolCall(ToolCallDelta {
index,
id: Some(call.id.clone()),
name: Some(call.function.name.clone()),
arguments: call.function.arguments.clone(),
})));
}
events.push(Ok(ChatStreamEvent::Done(StreamCompletion {
finish_reason: response.stop_reason(),
usage: response.usage(),
})));
Box::pin(stream::iter(events))
}
#[cfg(test)]
mod tests;