use crate::chat::{ChatEvent, ToolCall};
use anyhow::{Context, Result, anyhow};
use tokio::sync::mpsc::Sender;
#[cfg(test)]
#[path = "sse_pump_tests.rs"]
mod tests;
#[derive(Debug, Default)]
pub(super) struct ToolCallAccumulator {
slots: Vec<Option<(String, String, String)>>,
}
impl ToolCallAccumulator {
pub(super) fn apply_delta(&mut self, tool_calls: &serde_json::Value) {
let Some(arr) = tool_calls.as_array() else {
return;
};
for tc in arr {
let idx = tc.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize;
while self.slots.len() <= idx {
self.slots.push(None);
}
let slot = self.slots[idx]
.get_or_insert_with(|| (String::new(), String::new(), String::new()));
if let Some(id) = tc
.get("id")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
{
slot.0 = id.to_string();
}
if let Some(func) = tc.get("function") {
if let Some(name) = func
.get("name")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
{
slot.1 = name.to_string();
}
if let Some(args) = func.get("arguments").and_then(|v| v.as_str()) {
slot.2.push_str(args);
}
}
}
}
pub(super) fn finalize(self) -> Vec<ToolCall> {
self.slots
.into_iter()
.filter_map(|opt| {
opt.and_then(|(id, name, arguments)| {
if name.is_empty() {
None
} else {
Some(ToolCall {
id,
name,
arguments,
})
}
})
})
.collect()
}
}
#[derive(Debug, PartialEq, Eq)]
enum Flow {
Continue,
Stop,
Failed(String),
}
fn error_message_from_frame(err: &serde_json::Value) -> String {
if let Some(s) = err.as_str().filter(|s| !s.is_empty()) {
return s.to_string();
}
let message = err
.get("message")
.and_then(|m| m.as_str())
.filter(|s| !s.is_empty())
.unwrap_or("provider streaming error");
let code = err.get("code");
let code_str = code
.and_then(|c| c.as_str())
.filter(|s| !s.is_empty())
.map(str::to_string)
.or_else(|| code.and_then(|c| c.as_u64()).map(|n| n.to_string()));
match code_str {
Some(c) => format!("{c}: {message}"),
None => message.to_string(),
}
}
fn is_event_stream(headers: &reqwest::header::HeaderMap) -> bool {
headers
.get(reqwest::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.map(|v| v.to_ascii_lowercase().contains("text/event-stream"))
.unwrap_or(false)
}
async fn flush_terminal(acc: ToolCallAccumulator, tx: &Sender<ChatEvent>) -> Flow {
for call in acc.finalize() {
if tx.send(ChatEvent::ToolCall(call)).await.is_err() {
return Flow::Stop;
}
}
let _ = tx.send(ChatEvent::Done).await;
Flow::Stop
}
async fn handle_line(line: &str, acc: &mut ToolCallAccumulator, tx: &Sender<ChatEvent>) -> Flow {
let line = line.trim();
let Some(payload) = line.strip_prefix("data:").map(str::trim) else {
return Flow::Continue;
};
if payload.is_empty() {
return Flow::Continue;
}
if payload == "[DONE]" {
return flush_terminal(std::mem::take(acc), tx).await;
}
let Ok(v) = serde_json::from_str::<serde_json::Value>(payload) else {
return Flow::Continue;
};
if let Some(err) = v.get("error").filter(|e| !e.is_null()) {
let message = error_message_from_frame(err);
let _ = tx.send(ChatEvent::Error(message.clone())).await;
return Flow::Failed(message);
}
let Some(delta) = v
.get("choices")
.and_then(|c| c.get(0))
.and_then(|c| c.get("delta"))
else {
return Flow::Continue;
};
let content = delta
.get("content")
.and_then(|c| c.as_str())
.filter(|s| !s.is_empty())
.map(|s| s.to_string());
if let Some(content) = content
&& tx.send(ChatEvent::Delta(content)).await.is_err()
{
return Flow::Stop;
}
if let Some(tc) = delta.get("tool_calls") {
acc.apply_delta(tc);
}
Flow::Continue
}
pub(super) async fn pump_openai_sse(resp: reqwest::Response, tx: Sender<ChatEvent>) -> Result<()> {
use futures_util::StreamExt;
if !is_event_stream(resp.headers()) {
return pump_non_sse_body(resp, tx).await;
}
let mut acc = ToolCallAccumulator::default();
let mut buf: Vec<u8> = Vec::new();
let mut stream = resp.bytes_stream();
while let Some(chunk) = stream.next().await {
let bytes = match chunk {
Ok(b) => b,
Err(e) => {
let _ = tx
.send(ChatEvent::Error(format!(
"chat stream transport error: {e}"
)))
.await;
return Err(e).context("read chat stream chunk");
}
};
buf.extend_from_slice(&bytes);
while let Some(idx) = buf.iter().position(|b| *b == b'\n') {
let line: Vec<u8> = buf.drain(..=idx).collect();
match handle_line(&String::from_utf8_lossy(&line), &mut acc, &tx).await {
Flow::Continue => {}
Flow::Stop => return Ok(()),
Flow::Failed(message) => return Err(anyhow!("{message}")),
}
}
}
let residual = String::from_utf8_lossy(&buf);
let residual = residual.trim();
if !residual.is_empty() {
if !is_complete_frame(residual) {
let message = "chat stream ended with an incomplete SSE frame (truncated response)";
let _ = tx.send(ChatEvent::Error(message.to_string())).await;
return Err(anyhow!("{message}"));
}
match handle_line(residual, &mut acc, &tx).await {
Flow::Continue => {}
Flow::Stop => return Ok(()),
Flow::Failed(message) => return Err(anyhow!("{message}")),
}
}
flush_terminal(acc, &tx).await;
Ok(())
}
async fn pump_non_sse_body(resp: reqwest::Response, tx: Sender<ChatEvent>) -> Result<()> {
let text = resp
.text()
.await
.context("read non-SSE chat completion body")?;
if text.lines().any(|l| l.trim_start().starts_with("data:")) {
let mut acc = ToolCallAccumulator::default();
for line in text.split_inclusive('\n') {
match handle_line(line, &mut acc, &tx).await {
Flow::Continue => {}
Flow::Stop => return Ok(()),
Flow::Failed(message) => return Err(anyhow!("{message}")),
}
}
flush_terminal(acc, &tx).await;
return Ok(());
}
let Ok(v) = serde_json::from_str::<serde_json::Value>(&text) else {
return fail(
&tx,
format!(
"chat stream: provider returned a non-SSE, non-JSON body ({} bytes): {}",
text.len(),
excerpt(&text)
),
)
.await;
};
if let Some(err) = v.get("error").filter(|e| !e.is_null()) {
return fail(&tx, error_message_from_frame(err)).await;
}
let message = v
.get("choices")
.and_then(|c| c.get(0))
.and_then(|c| c.get("message"));
let content = message
.and_then(|m| m.get("content"))
.and_then(|c| c.as_str())
.filter(|s| !s.is_empty());
let mut acc = ToolCallAccumulator::default();
if let Some(tc) = message.and_then(|m| m.get("tool_calls")) {
acc.apply_delta(tc);
}
let calls = acc.finalize();
if content.is_none() && calls.is_empty() {
return fail(
&tx,
format!(
"chat stream: provider returned a non-SSE body with no content and no tool \
calls ({} bytes): {}",
text.len(),
excerpt(&text)
),
)
.await;
}
if let Some(content) = content
&& tx
.send(ChatEvent::Delta(content.to_string()))
.await
.is_err()
{
return Ok(());
}
for call in calls {
if tx.send(ChatEvent::ToolCall(call)).await.is_err() {
return Ok(());
}
}
let _ = tx.send(ChatEvent::Done).await;
Ok(())
}
fn is_complete_frame(line: &str) -> bool {
let line = line.trim();
if line.is_empty() || line.starts_with(':') {
return true;
}
let Some(payload) = line.strip_prefix("data:").map(str::trim) else {
return false;
};
payload.is_empty()
|| payload == "[DONE]"
|| serde_json::from_str::<serde_json::Value>(payload).is_ok()
}
async fn fail(tx: &Sender<ChatEvent>, message: String) -> Result<()> {
let _ = tx.send(ChatEvent::Error(message.clone())).await;
Err(anyhow!("{message}"))
}
fn excerpt(text: &str) -> String {
const LIMIT: usize = 200;
if text.chars().count() <= LIMIT {
return text.to_string();
}
let head: String = text.chars().take(LIMIT).collect();
format!("{head}…")
}