use std::collections::BTreeMap;
use color_eyre::{eyre::bail, Result};
use tokio::sync::mpsc;
use super::wire::{to_neutral_stop_reason, tool_block, StreamChunk};
use crate::agent::{ContentPart, StreamOutcome, TokenUsage};
use crate::sse::{EventSink, Flow};
pub struct StreamState<'a> {
update_tx: &'a mpsc::UnboundedSender<String>,
text: String,
calls: BTreeMap<usize, PartialCall>,
finish_reason: Option<String>,
usage: Option<TokenUsage>,
}
impl<'a> StreamState<'a> {
pub fn new(update_tx: &'a mpsc::UnboundedSender<String>) -> Self {
Self {
update_tx,
text: String::new(),
calls: BTreeMap::new(),
finish_reason: None,
usage: None,
}
}
pub fn into_outcome(self) -> StreamOutcome {
let mut blocks = Vec::new();
if !self.text.is_empty() {
blocks.push(ContentPart::Text { text: self.text });
}
for (_, call) in self.calls {
blocks.push(call.into_block());
}
StreamOutcome {
blocks,
stop_reason: self.finish_reason.as_deref().map(to_neutral_stop_reason),
usage: self.usage,
}
}
}
impl EventSink for StreamState<'_> {
fn absorb(&mut self, payload: &str) -> Result<Flow> {
if payload == "[DONE]" {
return Ok(Flow::Continue);
}
let Ok(event) = serde_json::from_str::<StreamChunk>(payload) else {
return Ok(Flow::Continue);
};
if let Some(error) = event.error {
bail!("Stream error: {}", error);
}
if let Some(wire_usage) = event.usage {
let reported: TokenUsage = wire_usage.into();
if reported.total() > 0 {
self.usage = Some(reported);
}
}
let Some(choice) = event.choices.into_iter().next() else {
return Ok(Flow::Continue);
};
if let Some(reason) = choice.finish_reason {
self.finish_reason = Some(reason);
}
if let Some(content) = choice.delta.content.filter(|c| !c.is_empty()) {
if self.update_tx.send(content.clone()).is_err() {
return Ok(Flow::Stop);
}
self.text.push_str(&content);
}
for fragment in choice.delta.tool_calls.into_iter().flatten() {
let call = self.calls.entry(fragment.index).or_default();
if let Some(id) = fragment.id {
call.id = id;
}
if let Some(function) = fragment.function {
if let Some(name) = function.name {
call.name = name;
}
if let Some(arguments) = function.arguments {
call.arguments.push_str(&arguments);
}
}
}
Ok(Flow::Continue)
}
}
#[derive(Debug, Default)]
struct PartialCall {
id: String,
name: String,
arguments: String,
}
impl PartialCall {
fn into_block(self) -> ContentPart {
tool_block(self.id, self.name, &self.arguments)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sse::EventReader;
fn run(chunks: &[&[u8]]) -> Result<(StreamOutcome, Vec<String>)> {
let (tx, mut rx) = mpsc::unbounded_channel();
let outcome = {
let mut reader = EventReader::new(StreamState::new(&tx));
for chunk in chunks {
if reader.feed(chunk)? == Flow::Stop {
break;
}
}
reader.finish()?;
reader.into_sink().into_outcome()
};
drop(tx);
let mut streamed = Vec::new();
while let Ok(chunk) = rx.try_recv() {
streamed.push(chunk);
}
Ok((outcome, streamed))
}
fn event(json: &str) -> Vec<u8> {
format!("data: {}\n\n", json).into_bytes()
}
fn run_events(events: &[&str]) -> (StreamOutcome, Vec<String>) {
let chunks: Vec<Vec<u8>> = events.iter().map(|e| event(e)).collect();
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
run(&refs).unwrap()
}
#[test]
fn content_deltas_assemble_and_stream_as_they_arrive() {
let (outcome, streamed) = run_events(&[
r#"{"choices":[{"delta":{"content":"he"}}]}"#,
r#"{"choices":[{"delta":{"content":"llo"}}]}"#,
r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#,
"[DONE]",
]);
assert_eq!(
outcome.blocks,
vec![ContentPart::Text {
text: "hello".to_string()
}]
);
assert_eq!(streamed, vec!["he", "llo"]);
assert_eq!(outcome.stop_reason.as_deref(), Some("end_turn"));
}
#[test]
fn tool_call_fragments_assemble_by_index() {
let (outcome, _) = run_events(&[
r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"c1","function":{"name":"read","arguments":"{\"path\":"}}]}}]}"#,
r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"a.rs\"}"}}]}}]}"#,
r#"{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}"#,
]);
assert_eq!(
outcome.blocks,
vec![ContentPart::ToolUse {
id: "c1".to_string(),
name: "read".to_string(),
input: serde_json::json!({"path": "a.rs"}),
}]
);
assert_eq!(outcome.stop_reason.as_deref(), Some("tool_use"));
}
#[test]
fn interleaved_parallel_calls_stay_separate() {
let (outcome, _) = run_events(&[
r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"c1","function":{"name":"read","arguments":"{}"}}]}}]}"#,
r#"{"choices":[{"delta":{"tool_calls":[{"index":1,"id":"c2","function":{"name":"list","arguments":"{"}}]}}]}"#,
r#"{"choices":[{"delta":{"tool_calls":[{"index":1,"function":{"arguments":"}"}}]}}]}"#,
]);
assert_eq!(
outcome.blocks,
vec![
ContentPart::ToolUse {
id: "c1".to_string(),
name: "read".to_string(),
input: serde_json::json!({}),
},
ContentPart::ToolUse {
id: "c2".to_string(),
name: "list".to_string(),
input: serde_json::json!({}),
},
]
);
}
#[test]
fn usage_is_read_from_the_final_chunk() {
let (outcome, _) = run_events(&[
r#"{"choices":[{"delta":{"content":"hi"}}]}"#,
r#"{"choices":[],"usage":{"prompt_tokens":1000,"completion_tokens":50,"prompt_cache_hit_tokens":800}}"#,
]);
let usage = outcome.usage.expect("usage reported");
assert_eq!(usage.input, 200);
assert_eq!(usage.cache_read, 800);
assert_eq!(usage.output, 50);
}
#[test]
fn a_stream_without_usage_reports_none() {
let (outcome, _) = run_events(&[r#"{"choices":[{"delta":{"content":"hi"}}]}"#]);
assert!(outcome.usage.is_none());
}
#[test]
fn an_empty_usage_chunk_does_not_erase_a_real_one() {
let (outcome, _) = run_events(&[
r#"{"choices":[],"usage":{"prompt_tokens":1000,"completion_tokens":50}}"#,
r#"{"choices":[],"usage":{"prompt_tokens":0,"completion_tokens":0}}"#,
]);
let usage = outcome.usage.expect("usage reported");
assert_eq!(usage.input, 1000);
assert_eq!(usage.output, 50);
}
#[test]
fn an_error_chunk_fails_the_request() {
let chunks = [event(
r#"{"error":{"message":"model not loaded","code":500}}"#,
)];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let err = run(&refs).unwrap_err().to_string();
assert!(err.contains("model not loaded"), "{}", err);
}
#[test]
fn a_bare_string_error_also_fails_the_request() {
let chunks = [event(r#"{"error":"context length exceeded"}"#)];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let err = run(&refs).unwrap_err().to_string();
assert!(err.contains("context length exceeded"), "{}", err);
}
#[test]
fn a_tool_call_without_an_index_is_still_assembled() {
let (outcome, _) = run_events(&[
r#"{"choices":[{"delta":{"tool_calls":[{"id":"c1","function":{"name":"read","arguments":"{\"path\":\"a.rs\"}"}}]}}]}"#,
]);
assert_eq!(
outcome.blocks,
vec![ContentPart::ToolUse {
id: "c1".to_string(),
name: "read".to_string(),
input: serde_json::json!({"path": "a.rs"}),
}]
);
}
#[test]
fn an_unparseable_payload_is_skipped() {
let (outcome, _) = run_events(&[
"not json at all",
r#"{"choices":[{"delta":{"content":"ok"}}]}"#,
]);
assert_eq!(
outcome.blocks,
vec![ContentPart::Text {
text: "ok".to_string()
}]
);
}
#[test]
fn a_closed_update_channel_stops_the_stream() {
let (tx, rx) = mpsc::unbounded_channel();
drop(rx);
let mut reader = EventReader::new(StreamState::new(&tx));
let flow = reader
.feed(&event(r#"{"choices":[{"delta":{"content":"hi"}}]}"#))
.unwrap();
assert_eq!(flow, Flow::Stop);
}
}