use crate::message::{Message, MessagePart};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::fmt;
pub mod handler;
pub mod heartbeat;
#[derive(Debug)]
#[non_exhaustive]
pub enum StreamError {
InvalidToolInputJson(serde_json::Error, String),
}
impl fmt::Display for StreamError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
StreamError::InvalidToolInputJson(err, raw) => {
write!(f, "invalid tool input JSON: {err} (raw_len={})", raw.len())
}
}
}
}
impl std::error::Error for StreamError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
StreamError::InvalidToolInputJson(err, _) => Some(err),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum StreamEvent {
MessageStart(MessageStart),
PartStart(PartStart),
IndexedDelta(IndexedDelta),
PartStop,
MessageDelta(MessageDelta),
MessageStop,
Ping,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageStart {
pub message: MessageMetadata,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageMetadata {
pub id: String,
pub role: String,
pub model: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PartStart {
pub index: usize,
pub part: Option<MessagePart>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct IndexedDelta {
pub index: usize,
pub delta: DeltaPart,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum DeltaPart {
#[serde(rename = "text_delta")]
Text {
text: String,
},
#[serde(rename = "tool_call_delta")]
ToolCall {
partial_json: Value,
},
#[serde(rename = "input_json_delta")]
InputJson {
partial_json: String,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum StreamStopReason {
ToolCall,
MaxTokens,
StopSequence,
EndTurn,
}
impl StreamStopReason {
#[must_use]
pub fn from_api_str(s: &str) -> Option<Self> {
match s {
"tool_call" => Some(Self::ToolCall),
"max_tokens" => Some(Self::MaxTokens),
"stop_sequence" => Some(Self::StopSequence),
"end_turn" => Some(Self::EndTurn),
_ => None,
}
}
#[must_use]
pub fn to_api_str(self) -> &'static str {
match self {
Self::ToolCall => "tool_call",
Self::MaxTokens => "max_tokens",
Self::StopSequence => "stop_sequence",
Self::EndTurn => "end_turn",
}
}
#[must_use]
pub fn should_continue_tool_loop(self) -> bool {
matches!(self, Self::ToolCall)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageDelta {
pub delta: MessageDeltaPayload,
pub usage: Option<Usage>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageDeltaPayload {
pub stop_reason: Option<String>,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default)]
pub struct Usage {
pub input_tokens: u32,
pub output_tokens: u32,
}
impl Usage {
#[must_use]
pub fn new(input_tokens: u32, output_tokens: u32) -> Self {
Self {
input_tokens,
output_tokens,
}
}
#[must_use]
pub fn total_tokens(self) -> u32 {
self.input_tokens.saturating_add(self.output_tokens)
}
}
#[derive(Debug, Default)]
pub struct StreamAccumulator {
parts: Vec<MessagePart>,
current_text: String,
current_tool_id: String,
current_tool_name: String,
current_tool_input: String,
current_index: Option<usize>,
model: Option<String>,
usage: Option<Usage>,
}
impl StreamAccumulator {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn process(&mut self, event: &StreamEvent) -> Result<(), StreamError> {
match event {
StreamEvent::MessageStart(msg_start) => {
self.model = Some(msg_start.message.model.clone());
Ok(())
}
StreamEvent::PartStart(part_start) => {
self.current_index = Some(part_start.index);
self.current_text.clear();
self.current_tool_id.clear();
self.current_tool_name.clear();
self.current_tool_input.clear();
if let Some(MessagePart::ToolCall { id, name, .. }) = &part_start.part {
self.current_tool_id.clone_from(id);
self.current_tool_name.clone_from(name);
}
Ok(())
}
StreamEvent::IndexedDelta(delta) => {
if self.current_index != Some(delta.index) {
return Ok(());
}
match &delta.delta {
DeltaPart::Text { text } => {
self.current_text.push_str(text);
}
DeltaPart::InputJson { partial_json } => {
self.current_tool_input.push_str(partial_json);
}
DeltaPart::ToolCall { partial_json } => {
if let Some(s) = partial_json.as_str() {
self.current_tool_input.push_str(s);
}
}
}
Ok(())
}
StreamEvent::PartStop => {
if !self.current_text.is_empty() {
self.parts.push(MessagePart::text(&self.current_text));
} else if !self.current_tool_name.is_empty() {
let input: Value = if self.current_tool_input.is_empty() {
Value::Object(serde_json::Map::new())
} else {
serde_json::from_str(&self.current_tool_input).map_err(|e| {
StreamError::InvalidToolInputJson(
e,
std::mem::take(&mut self.current_tool_input),
)
})?
};
self.parts.push(MessagePart::tool_call(
&self.current_tool_id,
&self.current_tool_name,
input,
));
}
self.current_text.clear();
Ok(())
}
StreamEvent::MessageDelta(delta) => {
self.usage = delta.usage;
Ok(())
}
StreamEvent::MessageStop | StreamEvent::Ping => Ok(()),
}
}
#[must_use]
pub fn peek_parts(&self) -> &[MessagePart] {
&self.parts
}
#[must_use]
pub fn build(self) -> Message {
Message {
role: crate::message::Role::Assistant,
parts: self.parts,
}
}
#[must_use]
pub fn usage(&self) -> Option<&Usage> {
self.usage.as_ref()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_stream_stop_reason_from_api_str() {
assert_eq!(
StreamStopReason::from_api_str("tool_call"),
Some(StreamStopReason::ToolCall)
);
assert_eq!(
StreamStopReason::from_api_str("max_tokens"),
Some(StreamStopReason::MaxTokens)
);
assert_eq!(
StreamStopReason::from_api_str("end_turn"),
Some(StreamStopReason::EndTurn)
);
assert_eq!(
StreamStopReason::from_api_str("stop_sequence"),
Some(StreamStopReason::StopSequence)
);
assert_eq!(StreamStopReason::from_api_str("unknown"), None);
}
#[test]
fn test_stream_stop_reason_to_api_str() {
assert_eq!(StreamStopReason::ToolCall.to_api_str(), "tool_call");
assert_eq!(StreamStopReason::MaxTokens.to_api_str(), "max_tokens");
assert_eq!(StreamStopReason::EndTurn.to_api_str(), "end_turn");
assert_eq!(StreamStopReason::StopSequence.to_api_str(), "stop_sequence");
}
#[test]
fn test_stream_stop_reason_should_continue() {
assert!(StreamStopReason::ToolCall.should_continue_tool_loop());
assert!(!StreamStopReason::EndTurn.should_continue_tool_loop());
assert!(!StreamStopReason::MaxTokens.should_continue_tool_loop());
}
#[test]
fn test_usage() {
let usage = Usage::new(100, 50);
assert_eq!(usage.input_tokens, 100);
assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.total_tokens(), 150);
}
#[test]
fn test_usage_default() {
let usage = Usage::default();
assert_eq!(usage.total_tokens(), 0);
}
#[test]
fn test_accumulator_text_message() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::MessageStart(MessageStart {
message: MessageMetadata {
id: "msg_1".to_string(),
role: "assistant".to_string(),
model: "test-model".to_string(),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: None,
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "Hello".to_string(),
},
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: " world".to_string(),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
acc.process(&StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("end_turn".to_string()),
},
usage: Some(Usage::new(10, 5)),
}))
.unwrap();
acc.process(&StreamEvent::MessageStop).unwrap();
let msg = acc.build();
assert_eq!(msg.role, crate::message::Role::Assistant);
assert_eq!(msg.parts.len(), 1);
assert_eq!(msg.parts[0].as_text(), Some("Hello world"));
}
#[test]
fn test_accumulator_tool_call() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call(
"tool_1",
"read_file",
Value::Object(serde_json::Map::new()),
)),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::InputJson {
partial_json: r#"{"path":"/tmp/test"}"#.to_string(),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
let msg = acc.build();
assert_eq!(msg.parts.len(), 1);
assert!(msg.parts[0].is_tool_call());
}
#[test]
fn test_accumulator_empty() {
let acc = StreamAccumulator::new();
let msg = acc.build();
assert_eq!(msg.parts.len(), 0);
}
#[test]
fn test_accumulator_usage() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload { stop_reason: None },
usage: Some(Usage::new(100, 50)),
}))
.unwrap();
assert_eq!(acc.usage().unwrap().total_tokens(), 150);
}
#[test]
fn test_stream_event_variants() {
let _ = StreamEvent::Ping;
let _ = StreamEvent::MessageStop;
let _ = StreamEvent::PartStop;
}
#[test]
fn test_accumulator_invalid_tool_json() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call(
"tool_1",
"bad_tool",
Value::Object(serde_json::Map::new()),
)),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::InputJson {
partial_json: "not valid json{".to_string(),
},
}))
.unwrap();
let result = acc.process(&StreamEvent::PartStop);
assert!(result.is_err());
let err = result.unwrap_err();
match &err {
StreamError::InvalidToolInputJson(_, raw) => {
assert_eq!(raw, "not valid json{");
}
}
}
#[test]
fn test_accumulator_tool_call_empty_input() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call(
"tool_1",
"no_args",
Value::Object(serde_json::Map::new()),
)),
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
let msg = acc.build();
assert_eq!(msg.parts.len(), 1);
assert!(msg.parts[0].is_tool_call());
}
#[test]
fn test_accumulator_ignores_delta_with_mismatched_index() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::text("")),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 1,
delta: DeltaPart::Text {
text: "ignored".into(),
},
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "hello".into(),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
let msg = acc.build();
assert_eq!(msg.parts.len(), 1);
assert_eq!(msg.parts[0].as_text(), Some("hello"));
}
#[test]
fn test_accumulator_ignores_input_json_with_mismatched_index() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::text("")),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 5,
delta: DeltaPart::InputJson {
partial_json: "{\"bad\":true}".into(),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
let msg = acc.build();
assert!(msg.parts.is_empty());
}
#[test]
fn test_accumulator_delta_tool_call_string_value() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call("id1", "search", Value::Null)),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::ToolCall {
partial_json: Value::String("{\"q\":\"rust\"}".into()),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
let msg = acc.build();
assert_eq!(msg.parts.len(), 1);
if let MessagePart::ToolCall { input, .. } = &msg.parts[0] {
assert_eq!(input["q"], "rust");
} else {
panic!("expected ToolCall");
}
}
#[test]
fn test_accumulator_delta_tool_call_non_string_ignored() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::tool_call("id1", "search", Value::Null)),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::ToolCall {
partial_json: Value::Number(42.into()),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
let msg = acc.build();
assert_eq!(msg.parts.len(), 1);
if let MessagePart::ToolCall { input, .. } = &msg.parts[0] {
assert!(input.is_object());
}
}
#[test]
fn test_accumulator_ping_no_op() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::Ping).unwrap();
assert!(acc.usage().is_none());
let msg = acc.build();
assert!(msg.parts.is_empty());
}
#[test]
fn test_accumulator_message_stop_no_op() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::text("")),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text { text: "hi".into() },
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
acc.process(&StreamEvent::MessageStop).unwrap();
let msg = acc.build();
assert_eq!(msg.parts.len(), 1);
assert_eq!(msg.parts[0].as_text(), Some("hi"));
}
#[test]
fn test_accumulator_multiple_text_parts() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::PartStart(PartStart {
index: 0,
part: Some(MessagePart::text("")),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 0,
delta: DeltaPart::Text {
text: "hello".into(),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
acc.process(&StreamEvent::PartStart(PartStart {
index: 1,
part: Some(MessagePart::text("")),
}))
.unwrap();
acc.process(&StreamEvent::IndexedDelta(IndexedDelta {
index: 1,
delta: DeltaPart::Text {
text: "world".into(),
},
}))
.unwrap();
acc.process(&StreamEvent::PartStop).unwrap();
let msg = acc.build();
assert_eq!(msg.parts.len(), 2);
assert_eq!(msg.parts[0].as_text(), Some("hello"));
assert_eq!(msg.parts[1].as_text(), Some("world"));
}
#[test]
fn test_accumulator_message_delta_overwrites_usage() {
let mut acc = StreamAccumulator::new();
acc.process(&StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("end_turn".into()),
},
usage: Some(Usage::new(100, 50)),
}))
.unwrap();
assert_eq!(acc.usage().unwrap().input_tokens, 100);
assert_eq!(acc.usage().unwrap().output_tokens, 50);
acc.process(&StreamEvent::MessageDelta(MessageDelta {
delta: MessageDeltaPayload {
stop_reason: Some("max_tokens".into()),
},
usage: Some(Usage::new(200, 75)),
}))
.unwrap();
assert_eq!(acc.usage().unwrap().input_tokens, 200);
assert_eq!(acc.usage().unwrap().output_tokens, 75);
}
}