use std::collections::BTreeMap;
use color_eyre::{eyre::bail, Result};
use tokio::sync::mpsc;
use super::wire::{Delta, StreamEvent};
use crate::agent::{ContentPart, StreamOutcome, TokenUsage};
use crate::sse::{EventSink, Flow};
enum Partial {
Text(String),
Tool {
id: String,
name: String,
json: String,
},
}
impl Partial {
fn into_block(self) -> Option<ContentPart> {
match self {
Partial::Text(text) if text.is_empty() => None,
Partial::Text { 0: text } => Some(ContentPart::Text { text }),
Partial::Tool { name, .. } if name.is_empty() => None,
Partial::Tool { id, name, json } => {
let input = crate::agent::tool_input(&name, &json);
Some(ContentPart::ToolUse { id, name, input })
}
}
}
}
pub struct StreamState<'a> {
update_tx: &'a mpsc::UnboundedSender<String>,
blocks: BTreeMap<usize, Partial>,
stop_reason: Option<String>,
usage: TokenUsage,
saw_usage: bool,
}
impl<'a> StreamState<'a> {
pub fn new(update_tx: &'a mpsc::UnboundedSender<String>) -> Self {
Self {
update_tx,
blocks: BTreeMap::new(),
stop_reason: None,
usage: TokenUsage::default(),
saw_usage: false,
}
}
pub fn into_outcome(self) -> StreamOutcome {
StreamOutcome {
blocks: self
.blocks
.into_values()
.filter_map(Partial::into_block)
.collect(),
stop_reason: self.stop_reason,
usage: self.saw_usage.then_some(self.usage),
}
}
}
impl EventSink for StreamState<'_> {
fn absorb(&mut self, payload: &str) -> Result<Flow> {
let Ok(event) = serde_json::from_str::<StreamEvent>(payload) else {
return Ok(Flow::Continue);
};
match event {
StreamEvent::ContentBlockStart {
index,
content_block,
} => {
let partial = match content_block {
ContentPart::ToolUse { id, name, .. } => Partial::Tool {
id,
name,
json: String::new(),
},
ContentPart::Text { text } => Partial::Text(text),
ContentPart::ToolResult { .. } => return Ok(Flow::Continue),
};
self.blocks.insert(index, partial);
}
StreamEvent::ContentBlockDelta { index, delta } => match delta {
Delta::TextDelta { text } => {
match self
.blocks
.entry(index)
.or_insert_with(|| Partial::Text(String::new()))
{
Partial::Text(buffer) => buffer.push_str(&text),
Partial::Tool { name, .. } => crate::diag::warn(format!(
"text delta on tool block {} ({}), discarded",
index, name
)),
}
if self.update_tx.send(text).is_err() {
return Ok(Flow::Stop);
}
}
Delta::InputJsonDelta { partial_json } => match self.blocks.get_mut(&index) {
Some(Partial::Tool { json, .. }) => json.push_str(&partial_json),
_ => crate::diag::warn(format!(
"argument fragment for unopened tool block {}, discarded",
index
)),
},
},
StreamEvent::ContentBlockStop { .. } => {}
StreamEvent::MessageStart { message } => {
if let Some(reported) = message.usage {
self.usage.input = reported.input_tokens.unwrap_or(0);
self.usage.cache_read = reported.cache_read_input_tokens.unwrap_or(0);
self.usage.cache_write = reported.cache_creation_input_tokens.unwrap_or(0);
self.usage.output = reported.output_tokens.unwrap_or(0);
self.saw_usage = true;
}
}
StreamEvent::MessageDelta {
delta,
usage: reported,
} => {
self.stop_reason = delta.stop_reason;
if let Some(reported) = reported {
if let Some(output) = reported.output_tokens {
self.usage.output = output;
self.saw_usage = true;
}
}
}
StreamEvent::Error { error } => bail!("Stream error: {}", error.message),
_ => {}
}
Ok(Flow::Continue)
}
}
#[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()
}
#[test]
fn text_deltas_assemble_into_one_block_and_stream_as_they_arrive() {
let chunks = [
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"he"}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"llo"}}"#,
),
event(r#"{"type":"content_block_stop","index":0}"#),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, streamed) = run(&refs).unwrap();
assert_eq!(
outcome.blocks,
vec![ContentPart::Text {
text: "hello".to_string()
}]
);
assert_eq!(streamed, vec!["he", "llo"], "deltas must reach the UI live");
}
#[test]
fn a_tool_call_assembles_from_its_argument_fragments() {
let chunks = [
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"t1","name":"read","input":{}}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{\"path\":"}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"\"a.rs\"}"}}"#,
),
event(r#"{"type":"content_block_stop","index":0}"#),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
assert_eq!(
outcome.blocks,
vec![ContentPart::ToolUse {
id: "t1".to_string(),
name: "read".to_string(),
input: serde_json::json!({"path": "a.rs"}),
}]
);
}
#[test]
fn a_tool_call_with_no_arguments_gets_an_empty_object() {
let chunks = [
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"t1","name":"list","input":{}}}"#,
),
event(r#"{"type":"content_block_stop","index":0}"#),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
assert_eq!(
outcome.blocks,
vec![ContentPart::ToolUse {
id: "t1".to_string(),
name: "list".to_string(),
input: serde_json::json!({}),
}]
);
}
#[test]
fn unparseable_arguments_still_yield_a_tool_use() {
let _guard = crate::diag::test_lock();
crate::diag::drain();
let chunks = [
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"t1","name":"read","input":{}}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{invalid"}}"#,
),
event(r#"{"type":"content_block_stop","index":0}"#),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
assert_eq!(
outcome.blocks,
vec![ContentPart::ToolUse {
id: "t1".to_string(),
name: "read".to_string(),
input: serde_json::json!({}),
}]
);
let warnings = crate::diag::drain();
assert!(
warnings.iter().any(|w| w.contains("did not parse")),
"expected a recorded warning, got {:?}",
warnings
);
}
#[test]
fn usage_is_read_from_message_start_and_updated_by_message_delta() {
let chunks = [
event(
r#"{"type":"message_start","message":{"id":"m","usage":{"input_tokens":1200,"cache_read_input_tokens":400,"cache_creation_input_tokens":30,"output_tokens":1}}}"#,
),
event(
r#"{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":915}}"#,
),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
let usage = outcome.usage.expect("usage reported");
assert_eq!(usage.input, 1200);
assert_eq!(usage.cache_read, 400);
assert_eq!(usage.cache_write, 30);
assert_eq!(
usage.output, 915,
"message_delta must win over message_start"
);
assert_eq!(outcome.stop_reason.as_deref(), Some("end_turn"));
}
#[test]
fn a_stream_without_usage_reports_none() {
let chunks = [event(r#"{"type":"message_stop"}"#)];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
assert!(run(&refs).unwrap().0.usage.is_none());
}
#[test]
fn an_error_event_fails_the_request() {
let chunks = [event(
r#"{"type":"error","error":{"message":"overloaded"}}"#,
)];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let err = run(&refs).unwrap_err().to_string();
assert!(err.contains("overloaded"), "{}", err);
}
#[test]
fn a_payload_this_adapter_does_not_model_is_skipped() {
let chunks = [
event(r#"{"type":"ping"}"#),
event(r#"{"type":"something_new","weird":true}"#),
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"ok"}}"#,
),
event(r#"{"type":"content_block_stop","index":0}"#),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
assert_eq!(
outcome.blocks,
vec![ContentPart::Text {
text: "ok".to_string()
}]
);
}
#[test]
fn interleaved_tool_blocks_keep_their_own_arguments() {
let chunks = [
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"t0","name":"read","input":{}}}"#,
),
event(
r#"{"type":"content_block_start","index":1,"content_block":{"type":"tool_use","id":"t1","name":"grep","input":{}}}"#,
),
event(
r#"{"type":"content_block_delta","index":1,"delta":{"type":"input_json_delta","partial_json":"{\"pattern\":"}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{\"path\":"}}"#,
),
event(
r#"{"type":"content_block_delta","index":1,"delta":{"type":"input_json_delta","partial_json":"\"fn main\"}"}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"\"a.rs\"}"}}"#,
),
event(r#"{"type":"content_block_stop","index":0}"#),
event(r#"{"type":"content_block_stop","index":1}"#),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
assert_eq!(
outcome.blocks,
vec![
ContentPart::ToolUse {
id: "t0".to_string(),
name: "read".to_string(),
input: serde_json::json!({"path": "a.rs"}),
},
ContentPart::ToolUse {
id: "t1".to_string(),
name: "grep".to_string(),
input: serde_json::json!({"pattern": "fn main"}),
},
],
"each block must keep the fragments addressed to its own index"
);
}
#[test]
fn blocks_are_ordered_by_index_not_by_arrival() {
let chunks = [
event(
r#"{"type":"content_block_start","index":1,"content_block":{"type":"tool_use","id":"t1","name":"grep","input":{}}}"#,
),
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"text","text":"thinking"}}"#,
),
event(r#"{"type":"content_block_stop","index":1}"#),
event(r#"{"type":"content_block_stop","index":0}"#),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
assert!(
matches!(outcome.blocks.first(), Some(ContentPart::Text { .. })),
"index 0 must come first, got {:?}",
outcome.blocks
);
assert_eq!(outcome.blocks.len(), 2);
}
#[test]
fn a_stream_cut_before_the_stop_event_keeps_what_it_had() {
let chunks = [
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"half a sen"}}"#,
),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
assert_eq!(
outcome.blocks,
vec![ContentPart::Text {
text: "half a sen".to_string()
}],
"a truncated stream must not report an empty turn"
);
}
#[test]
fn a_tool_call_cut_before_the_stop_event_still_reaches_the_loop() {
let chunks = [
event(
r#"{"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"t1","name":"read","input":{}}}"#,
),
event(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{\"path\":\"a.rs\"}"}}"#,
),
];
let refs: Vec<&[u8]> = chunks.iter().map(|c| c.as_slice()).collect();
let (outcome, _) = run(&refs).unwrap();
assert_eq!(
outcome.blocks,
vec![ContentPart::ToolUse {
id: "t1".to_string(),
name: "read".to_string(),
input: serde_json::json!({"path": "a.rs"}),
}]
);
}
#[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#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}"#,
))
.unwrap();
assert_eq!(flow, Flow::Stop);
}
}