use futures::{Stream, StreamExt};
use serde::{Deserialize, Serialize};
use crate::error::{Error, ProviderKind};
use crate::message::{ConversationTurn, Message, Response, StreamEvent, ToolCall, Usage};
pub const WIRE_PROTOCOL_VERSION: u32 = 1;
fn default_protocol_version() -> u32 {
WIRE_PROTOCOL_VERSION
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
#[non_exhaustive]
pub enum WireStreamEvent {
MessageStart {
#[serde(default = "default_protocol_version")]
protocol_version: u32,
model: String,
provider: ProviderKind,
},
TextDelta {
text: String,
},
ToolCallStart {
id: String,
name: String,
},
ToolCallDelta {
id: String,
arguments: String,
},
ToolCallEnd {
id: String,
name: String,
arguments: String,
},
ToolResult {
id: String,
result: String,
},
Usage {
usage: Usage,
},
MessageStop {
#[serde(default, skip_serializing_if = "Option::is_none")]
finish_reason: Option<String>,
},
TurnComplete {
turn: ConversationTurn,
},
Error {
error: WireError,
},
}
impl WireStreamEvent {
pub fn message_start(model: impl Into<String>, provider: ProviderKind) -> Self {
Self::MessageStart {
protocol_version: WIRE_PROTOCOL_VERSION,
model: model.into(),
provider,
}
}
pub fn error(error: &Error) -> Self {
Self::Error {
error: WireError::from(error),
}
}
pub fn tag(&self) -> &'static str {
match self {
Self::MessageStart { .. } => "message_start",
Self::TextDelta { .. } => "text_delta",
Self::ToolCallStart { .. } => "tool_call_start",
Self::ToolCallDelta { .. } => "tool_call_delta",
Self::ToolCallEnd { .. } => "tool_call_end",
Self::ToolResult { .. } => "tool_result",
Self::Usage { .. } => "usage",
Self::MessageStop { .. } => "message_stop",
Self::TurnComplete { .. } => "turn_complete",
Self::Error { .. } => "error",
}
}
pub fn is_terminal(&self) -> bool {
matches!(self, Self::MessageStop { .. } | Self::Error { .. })
}
}
impl From<StreamEvent> for WireStreamEvent {
fn from(event: StreamEvent) -> Self {
match event {
StreamEvent::TextDelta { text } => Self::TextDelta { text },
StreamEvent::ToolCall {
id,
name,
arguments,
} => Self::ToolCallEnd {
id,
name,
arguments,
},
StreamEvent::ToolResult { id, result } => Self::ToolResult { id, result },
StreamEvent::TurnComplete { turn } => Self::TurnComplete { turn },
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct UnrepresentableWireEvent {
pub tag: &'static str,
}
impl std::fmt::Display for UnrepresentableWireEvent {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "wire event '{}' has no StreamEvent equivalent", self.tag)
}
}
impl std::error::Error for UnrepresentableWireEvent {}
impl TryFrom<WireStreamEvent> for StreamEvent {
type Error = UnrepresentableWireEvent;
fn try_from(event: WireStreamEvent) -> std::result::Result<Self, UnrepresentableWireEvent> {
let tag = event.tag();
match event {
WireStreamEvent::TextDelta { text } => Ok(Self::TextDelta { text }),
WireStreamEvent::ToolCallEnd {
id,
name,
arguments,
} => Ok(Self::ToolCall {
id,
name,
arguments,
}),
WireStreamEvent::ToolResult { id, result } => Ok(Self::ToolResult { id, result }),
WireStreamEvent::TurnComplete { turn } => Ok(Self::TurnComplete { turn }),
_ => Err(UnrepresentableWireEvent { tag }),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct WireError {
pub kind: WireErrorKind,
pub message: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider: Option<ProviderKind>,
pub retryable: bool,
}
impl WireError {
pub fn new(kind: WireErrorKind, message: impl Into<String>) -> Self {
Self {
kind,
message: message.into(),
provider: None,
retryable: false,
}
}
}
impl From<&Error> for WireError {
fn from(error: &Error) -> Self {
Self {
kind: WireErrorKind::from(error.kind_str()),
message: error.to_string(),
provider: error.provider(),
retryable: error.is_retryable(),
}
}
}
impl std::fmt::Display for WireError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}: {}", self.kind.as_str(), self.message)
}
}
impl std::error::Error for WireError {}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum WireErrorKind {
Auth,
Request,
RateLimit,
InvalidRequest,
ModelNotAvailable,
ProviderNotConfigured,
ProviderNotEnabled,
CapabilityUnsupported,
ContentFiltered,
Config,
Serialization,
Http,
Stream,
Timeout,
ToolProviderUnsupported,
ToolArguments,
ToolNotFound,
ToolLoopLimitExceeded,
StructuredOutput,
#[serde(untagged)]
Other(String),
}
impl WireErrorKind {
pub fn as_str(&self) -> &str {
match self {
Self::Auth => "auth",
Self::Request => "request",
Self::RateLimit => "rate_limit",
Self::InvalidRequest => "invalid_request",
Self::ModelNotAvailable => "model_not_available",
Self::ProviderNotConfigured => "provider_not_configured",
Self::ProviderNotEnabled => "provider_not_enabled",
Self::CapabilityUnsupported => "capability_unsupported",
Self::ContentFiltered => "content_filtered",
Self::Config => "config",
Self::Serialization => "serialization",
Self::Http => "http",
Self::Stream => "stream",
Self::Timeout => "timeout",
Self::ToolProviderUnsupported => "tool_provider_unsupported",
Self::ToolArguments => "tool_arguments",
Self::ToolNotFound => "tool_not_found",
Self::ToolLoopLimitExceeded => "tool_loop_limit_exceeded",
Self::StructuredOutput => "structured_output",
Self::Other(kind) => kind,
}
}
}
impl From<&str> for WireErrorKind {
fn from(kind: &str) -> Self {
match kind {
"auth" => Self::Auth,
"request" => Self::Request,
"rate_limit" => Self::RateLimit,
"invalid_request" => Self::InvalidRequest,
"model_not_available" => Self::ModelNotAvailable,
"provider_not_configured" => Self::ProviderNotConfigured,
"provider_not_enabled" => Self::ProviderNotEnabled,
"capability_unsupported" => Self::CapabilityUnsupported,
"content_filtered" => Self::ContentFiltered,
"config" => Self::Config,
"serialization" => Self::Serialization,
"http" => Self::Http,
"stream" => Self::Stream,
"timeout" => Self::Timeout,
"tool_provider_unsupported" => Self::ToolProviderUnsupported,
"tool_arguments" => Self::ToolArguments,
"tool_not_found" => Self::ToolNotFound,
"tool_loop_limit_exceeded" => Self::ToolLoopLimitExceeded,
"structured_output" => Self::StructuredOutput,
other => Self::Other(other.to_string()),
}
}
}
impl std::fmt::Display for WireErrorKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone)]
struct PendingToolCall {
id: String,
name: String,
arguments: String,
}
impl PendingToolCall {
fn finish(self) -> ToolCall {
ToolCall {
id: self.id,
name: self.name,
arguments: serde_json::from_str(&self.arguments).unwrap_or(serde_json::Value::Null),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct StreamAccumulator {
protocol_version: Option<u32>,
model: Option<String>,
provider: Option<ProviderKind>,
text: String,
tool_calls: Vec<ToolCall>,
pending_tool_call: Option<PendingToolCall>,
tool_results: Vec<Message>,
usage: Option<Usage>,
finish_reason: Option<String>,
error: Option<WireError>,
stopped: bool,
}
impl StreamAccumulator {
pub fn new() -> Self {
Self::default()
}
pub fn protocol_version(&self) -> Option<u32> {
self.protocol_version
}
pub fn text(&self) -> &str {
&self.text
}
pub fn usage(&self) -> Option<&Usage> {
self.usage.as_ref()
}
pub fn is_complete(&self) -> bool {
self.stopped || self.error.is_some()
}
pub fn push(&mut self, event: WireStreamEvent) -> std::result::Result<(), WireError> {
match event {
WireStreamEvent::MessageStart {
protocol_version,
model,
provider,
} => {
self.protocol_version = Some(protocol_version);
self.model = Some(model);
self.provider = Some(provider);
}
WireStreamEvent::TextDelta { text } => self.text.push_str(&text),
WireStreamEvent::ToolCallStart { id, name } => {
self.flush_pending_tool_call();
self.pending_tool_call = Some(PendingToolCall {
id,
name,
arguments: String::new(),
});
}
WireStreamEvent::ToolCallDelta { id, arguments } => match &mut self.pending_tool_call {
Some(pending) if pending.id == id => pending.arguments.push_str(&arguments),
_ => {
self.flush_pending_tool_call();
self.pending_tool_call = Some(PendingToolCall {
id,
name: String::new(),
arguments,
});
}
},
WireStreamEvent::ToolCallEnd {
id,
name,
arguments,
} => {
if self
.pending_tool_call
.as_ref()
.is_some_and(|pending| pending.id == id)
{
self.pending_tool_call = None;
} else {
self.flush_pending_tool_call();
}
self.tool_calls.push(
PendingToolCall {
id,
name,
arguments,
}
.finish(),
);
}
WireStreamEvent::ToolResult { id, result } => {
self.tool_results.push(Message::tool(result, id));
}
WireStreamEvent::Usage { usage } => self.usage = Some(usage),
WireStreamEvent::MessageStop { finish_reason } => {
self.flush_pending_tool_call();
if finish_reason.is_some() {
self.finish_reason = finish_reason;
}
self.stopped = true;
}
WireStreamEvent::TurnComplete { turn } => {
self.flush_pending_tool_call();
if self.text.is_empty() {
self.text = turn.assistant_message.text_content();
}
if self.tool_calls.is_empty() {
self.tool_calls = turn.assistant_message.tool_calls.clone();
}
self.tool_results.extend(turn.tool_results);
}
WireStreamEvent::Error { error } => {
self.error = Some(error.clone());
return Err(error);
}
#[allow(unreachable_patterns)]
_ => {}
}
Ok(())
}
fn flush_pending_tool_call(&mut self) {
if let Some(pending) = self.pending_tool_call.take() {
self.tool_calls.push(pending.finish());
}
}
pub fn finish(mut self) -> std::result::Result<Response, WireError> {
if let Some(error) = self.error {
return Err(error);
}
self.flush_pending_tool_call();
let (Some(model), Some(provider)) = (self.model, self.provider) else {
return Err(WireError::new(
WireErrorKind::Stream,
"stream ended without a message_start event, so the model and provider are unknown",
));
};
if !self.stopped {
return Err(WireError::new(
WireErrorKind::Stream,
"stream ended before its terminal event (message_stop or error); \
the connection was most likely dropped mid-generation",
));
}
let mut assistant = Message::assistant(self.text);
assistant.tool_calls = self.tool_calls;
let mut messages = vec![assistant];
messages.extend(self.tool_results);
Ok(Response {
messages,
usage: self.usage,
model,
provider,
finish_reason: self.finish_reason,
})
}
pub async fn accumulate<S>(stream: S) -> std::result::Result<Response, WireError>
where
S: Stream<Item = WireStreamEvent>,
{
let mut stream = std::pin::pin!(stream);
let mut accumulator = Self::new();
while let Some(event) = stream.next().await {
accumulator.push(event)?;
}
accumulator.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::{ContentBlock, Message};
const FIXTURE_DIR: &str = concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/wire");
const UPDATE_ENV: &str = "RAI_SDK_UPDATE_WIRE_FIXTURES";
fn sample_usage() -> Usage {
Usage {
prompt_tokens: Some(1_024),
completion_tokens: Some(256),
total_tokens: Some(1_280),
}
}
fn sample_turn() -> ConversationTurn {
let mut assistant = Message::assistant("Sunny.");
assistant.tool_calls = vec![ToolCall {
id: "call_1".to_string(),
name: "get_weather".to_string(),
arguments: serde_json::json!({ "city": "Paris" }),
}];
ConversationTurn {
user_message: Message::user("Weather in Paris?"),
assistant_message: assistant,
tool_results: vec![Message::tool("{\"c\":21}", "call_1")],
}
}
fn every_variant() -> Vec<WireStreamEvent> {
vec![
WireStreamEvent::message_start("gpt-4o-mini", ProviderKind::OpenAI),
WireStreamEvent::TextDelta {
text: "Hello, world.".to_string(),
},
WireStreamEvent::ToolCallStart {
id: "call_1".to_string(),
name: "get_weather".to_string(),
},
WireStreamEvent::ToolCallDelta {
id: "call_1".to_string(),
arguments: "{\"city\":".to_string(),
},
WireStreamEvent::ToolCallEnd {
id: "call_1".to_string(),
name: "get_weather".to_string(),
arguments: "{\"city\":\"Paris\"}".to_string(),
},
WireStreamEvent::ToolResult {
id: "call_1".to_string(),
result: "{\"celsius\":21}".to_string(),
},
WireStreamEvent::Usage {
usage: sample_usage(),
},
WireStreamEvent::MessageStop {
finish_reason: Some("stop".to_string()),
},
WireStreamEvent::TurnComplete {
turn: sample_turn(),
},
WireStreamEvent::error(&Error::RateLimit {
provider: ProviderKind::Anthropic,
message: "too many requests".to_string(),
}),
]
}
fn fixture_name(event: &WireStreamEvent) -> &'static str {
match event {
WireStreamEvent::MessageStart { .. } => "message_start",
WireStreamEvent::TextDelta { .. } => "text_delta",
WireStreamEvent::ToolCallStart { .. } => "tool_call_start",
WireStreamEvent::ToolCallDelta { .. } => "tool_call_delta",
WireStreamEvent::ToolCallEnd { .. } => "tool_call_end",
WireStreamEvent::ToolResult { .. } => "tool_result",
WireStreamEvent::Usage { .. } => "usage",
WireStreamEvent::MessageStop { .. } => "message_stop",
WireStreamEvent::TurnComplete { .. } => "turn_complete",
WireStreamEvent::Error { .. } => "error",
}
}
#[test]
fn every_variant_is_covered_exactly_once() {
let mut names: Vec<&str> = every_variant().iter().map(fixture_name).collect();
let total = names.len();
names.sort_unstable();
names.dedup();
assert_eq!(
names.len(),
total,
"every_variant() must contain each variant exactly once"
);
}
#[test]
fn round_trip_preserves_every_variant() {
for event in every_variant() {
let json = serde_json::to_string(&event).expect("event should serialize");
let parsed: WireStreamEvent =
serde_json::from_str(&json).expect("event should deserialize");
assert_eq!(
parsed,
event,
"round trip changed the '{}' event: {json}",
fixture_name(&event)
);
}
}
#[test]
fn the_tag_matches_the_serialized_discriminant() {
for event in every_variant() {
let json = serde_json::to_value(&event).expect("event should serialize");
assert_eq!(
json.get("type").and_then(serde_json::Value::as_str),
Some(event.tag()),
"tag() disagrees with the serialized discriminant: {json}"
);
assert_eq!(event.tag(), fixture_name(&event));
}
}
#[test]
fn wire_format_matches_the_committed_fixtures() {
let updating = std::env::var_os(UPDATE_ENV).is_some();
for event in every_variant() {
let name = fixture_name(&event);
let path = std::path::Path::new(FIXTURE_DIR).join(format!("{name}.json"));
let actual = serde_json::to_value(&event).expect("event should serialize");
if updating {
std::fs::create_dir_all(FIXTURE_DIR).expect("create the fixture directory");
let pretty =
serde_json::to_string_pretty(&actual).expect("fixture should serialize");
std::fs::write(&path, format!("{pretty}\n")).expect("write the fixture");
continue;
}
let raw = std::fs::read_to_string(&path).unwrap_or_else(|error| {
panic!(
"missing wire-format fixture {}: {error}\n\
If this variant is new, regenerate the fixtures with:\n \
{UPDATE_ENV}=1 cargo test --all-features wire_format_matches_the_committed_fixtures",
path.display()
)
});
let expected: serde_json::Value =
serde_json::from_str(&raw).expect("fixture should be valid JSON");
assert_eq!(
actual,
expected,
"the '{name}' event no longer matches {}.\n\
Wire-format changes break clients built against an older rai-sdk. \
If this change is intentional, note it in CHANGELOG.md and regenerate with:\n \
{UPDATE_ENV}=1 cargo test --all-features wire_format_matches_the_committed_fixtures",
path.display()
);
let parsed: WireStreamEvent =
serde_json::from_value(expected).expect("fixture should deserialize");
assert_eq!(parsed, event);
}
}
#[test]
fn stream_event_conversion_round_trips() {
let events = vec![
StreamEvent::TextDelta {
text: "hi".to_string(),
},
StreamEvent::ToolCall {
id: "call_1".to_string(),
name: "get_weather".to_string(),
arguments: "{\"city\":\"Paris\"}".to_string(),
},
StreamEvent::ToolResult {
id: "call_1".to_string(),
result: "{\"celsius\":21}".to_string(),
},
StreamEvent::TurnComplete {
turn: sample_turn(),
},
];
for event in events {
let wire = WireStreamEvent::from(event.clone());
let back = StreamEvent::try_from(wire).expect("wire event should convert back");
assert_eq!(back, event);
}
}
#[test]
fn framing_events_have_no_stream_event_equivalent() {
let wire = WireStreamEvent::MessageStop {
finish_reason: None,
};
let error = StreamEvent::try_from(wire).expect_err("message_stop is wire-only");
assert_eq!(error.tag, "message_stop");
assert!(error.to_string().contains("message_stop"));
}
#[test]
fn message_start_defaults_the_protocol_version_when_absent() {
let parsed: WireStreamEvent = serde_json::from_str(
r#"{"type":"message_start","model":"gpt-4o-mini","provider":"openai"}"#,
)
.expect("legacy payload should still parse");
assert_eq!(
parsed,
WireStreamEvent::message_start("gpt-4o-mini", ProviderKind::OpenAI)
);
}
#[test]
fn is_terminal_marks_exactly_the_stream_enders() {
for event in every_variant() {
let expected = matches!(fixture_name(&event), "message_stop" | "error");
assert_eq!(event.is_terminal(), expected, "{}", event.tag());
}
}
#[test]
fn every_error_variant_maps_to_a_named_kind() {
let provider = ProviderKind::OpenAI;
let errors = vec![
Error::Auth {
provider,
message: "bad key".into(),
},
Error::Request {
provider,
message: "boom".into(),
},
Error::RateLimit {
provider,
message: "slow down".into(),
},
Error::InvalidRequest("nope".into()),
Error::ModelNotAvailable {
provider,
model: "gpt-9".into(),
},
Error::ProviderNotConfigured(provider),
Error::ProviderNotEnabled(provider),
Error::ContentFiltered {
provider,
reason: "policy".into(),
},
Error::Config("missing".into()),
Error::Serialization(
serde_json::from_str::<serde_json::Value>("{").expect_err("invalid JSON"),
),
Error::Stream("truncated".into()),
Error::Timeout { provider },
Error::ToolProviderUnsupported { provider },
Error::ToolArguments {
name: "echo".into(),
message: "bad".into(),
issues: Vec::new(),
},
Error::ToolNotFound {
name: "echo".into(),
},
Error::ToolLoopLimitExceeded { max_rounds: 8 },
Error::StructuredOutput {
provider,
model: "gpt-4o-mini".into(),
message: "invalid".into(),
},
];
for error in &errors {
let wire = WireError::from(error);
assert!(
!matches!(wire.kind, WireErrorKind::Other(_)),
"{} has no named WireErrorKind",
error.kind_str()
);
assert_eq!(wire.kind.as_str(), error.kind_str());
assert_eq!(wire.message, error.to_string());
assert_eq!(wire.provider, error.provider());
assert_eq!(wire.retryable, error.is_retryable());
}
assert_eq!(WireErrorKind::from("http"), WireErrorKind::Http);
}
#[test]
fn unknown_error_kinds_survive_a_round_trip() {
let json = r#"{"kind":"quantum_flux","message":"from the future","retryable":true}"#;
let parsed: WireError = serde_json::from_str(json).expect("unknown kind should parse");
assert_eq!(parsed.kind, WireErrorKind::Other("quantum_flux".into()));
assert_eq!(parsed.kind.as_str(), "quantum_flux");
assert_eq!(parsed.provider, None);
let reserialized = serde_json::to_value(&parsed).expect("should serialize");
assert_eq!(reserialized["kind"], "quantum_flux");
}
#[test]
fn accumulator_rebuilds_text_usage_and_finish_reason() {
let mut accumulator = StreamAccumulator::new();
for event in [
WireStreamEvent::message_start("gpt-4o-mini", ProviderKind::OpenAI),
WireStreamEvent::TextDelta {
text: "Hello, ".into(),
},
WireStreamEvent::TextDelta {
text: "world.".into(),
},
WireStreamEvent::Usage {
usage: sample_usage(),
},
WireStreamEvent::MessageStop {
finish_reason: Some("stop".into()),
},
] {
accumulator.push(event).expect("event should be absorbed");
}
assert_eq!(accumulator.protocol_version(), Some(WIRE_PROTOCOL_VERSION));
assert!(accumulator.is_complete());
let response = accumulator.finish().expect("stream was well formed");
assert_eq!(response.text(), "Hello, world.");
assert_eq!(response.model, "gpt-4o-mini");
assert_eq!(response.provider, ProviderKind::OpenAI);
assert_eq!(response.finish_reason.as_deref(), Some("stop"));
assert_eq!(response.usage, Some(sample_usage()));
}
#[test]
fn accumulator_assembles_tool_calls_from_fragments() {
let mut accumulator = StreamAccumulator::new();
for event in [
WireStreamEvent::message_start("gpt-4o-mini", ProviderKind::OpenAI),
WireStreamEvent::ToolCallStart {
id: "call_1".into(),
name: "get_weather".into(),
},
WireStreamEvent::ToolCallDelta {
id: "call_1".into(),
arguments: "{\"city\":".into(),
},
WireStreamEvent::ToolCallDelta {
id: "call_1".into(),
arguments: "\"Paris\"}".into(),
},
WireStreamEvent::MessageStop {
finish_reason: Some("tool_use".into()),
},
] {
accumulator.push(event).expect("event should be absorbed");
}
let response = accumulator.finish().expect("stream was well formed");
let tool_calls = &response.messages[0].tool_calls;
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].name, "get_weather");
assert_eq!(tool_calls[0].arguments["city"], "Paris");
}
#[test]
fn a_tool_call_end_supersedes_its_own_fragments() {
let mut accumulator = StreamAccumulator::new();
for event in [
WireStreamEvent::message_start("gpt-4o-mini", ProviderKind::OpenAI),
WireStreamEvent::ToolCallStart {
id: "call_1".into(),
name: "get_weather".into(),
},
WireStreamEvent::ToolCallDelta {
id: "call_1".into(),
arguments: "{\"city\":".into(),
},
WireStreamEvent::ToolCallEnd {
id: "call_1".into(),
name: "get_weather".into(),
arguments: "{\"city\":\"Paris\"}".into(),
},
WireStreamEvent::MessageStop {
finish_reason: Some("tool_use".into()),
},
] {
accumulator.push(event).expect("event should be absorbed");
}
let response = accumulator.finish().expect("stream was well formed");
assert_eq!(response.messages[0].tool_calls.len(), 1);
}
#[test]
fn accumulator_keeps_tool_results_as_messages() {
let mut accumulator = StreamAccumulator::new();
for event in [
WireStreamEvent::message_start("gpt-4o-mini", ProviderKind::OpenAI),
WireStreamEvent::ToolResult {
id: "call_1".into(),
result: "{\"celsius\":21}".into(),
},
WireStreamEvent::MessageStop {
finish_reason: Some("stop".into()),
},
] {
accumulator.push(event).expect("event should be absorbed");
}
let response = accumulator.finish().expect("stream was well formed");
assert_eq!(response.messages.len(), 2);
assert_eq!(response.messages[1].tool_call_id.as_deref(), Some("call_1"));
}
#[test]
fn an_error_event_stops_the_accumulator() {
let mut accumulator = StreamAccumulator::new();
accumulator
.push(WireStreamEvent::message_start(
"gpt-4o-mini",
ProviderKind::OpenAI,
))
.expect("event should be absorbed");
let error = accumulator
.push(WireStreamEvent::error(&Error::ContentFiltered {
provider: ProviderKind::OpenAI,
reason: "policy".into(),
}))
.expect_err("an error event should surface as an error");
assert_eq!(error.kind, WireErrorKind::ContentFiltered);
assert!(accumulator.is_complete());
assert_eq!(
accumulator
.finish()
.expect_err("finish reports it too")
.kind,
WireErrorKind::ContentFiltered
);
}
#[test]
fn a_truncated_stream_is_distinguishable_from_a_provider_error() {
let mut accumulator = StreamAccumulator::new();
for event in [
WireStreamEvent::message_start("gpt-4o-mini", ProviderKind::OpenAI),
WireStreamEvent::TextDelta {
text: "half a sen".into(),
},
] {
accumulator.push(event).expect("event should be absorbed");
}
assert!(!accumulator.is_complete());
let error = accumulator
.finish()
.expect_err("a truncated stream is an error");
assert_eq!(error.kind, WireErrorKind::Stream);
assert!(
error.message.contains("terminal event"),
"unhelpful truncation message: {}",
error.message
);
}
#[test]
fn a_stream_without_message_start_is_rejected() {
let mut accumulator = StreamAccumulator::new();
accumulator
.push(WireStreamEvent::MessageStop {
finish_reason: Some("stop".into()),
})
.expect("event should be absorbed");
let error = accumulator
.finish()
.expect_err("model and provider are unknown");
assert_eq!(error.kind, WireErrorKind::Stream);
assert!(error.message.contains("message_start"));
}
#[test]
fn turn_complete_backfills_an_otherwise_empty_response() {
let mut accumulator = StreamAccumulator::new();
for event in [
WireStreamEvent::message_start("claude-sonnet-4-6", ProviderKind::Anthropic),
WireStreamEvent::TurnComplete {
turn: sample_turn(),
},
WireStreamEvent::MessageStop {
finish_reason: Some("tool_use".into()),
},
] {
accumulator.push(event).expect("event should be absorbed");
}
let response = accumulator.finish().expect("stream was well formed");
assert_eq!(response.text(), "Sunny.");
assert_eq!(response.messages[0].tool_calls.len(), 1);
}
#[tokio::test]
async fn accumulate_drains_a_whole_stream() {
let response = StreamAccumulator::accumulate(futures::stream::iter(vec![
WireStreamEvent::message_start("gpt-4o-mini", ProviderKind::OpenAI),
WireStreamEvent::TextDelta { text: "ok".into() },
WireStreamEvent::MessageStop {
finish_reason: Some("stop".into()),
},
]))
.await
.expect("stream was well formed");
assert_eq!(response.text(), "ok");
}
#[test]
fn multimodal_turns_survive_the_wire() {
let turn = ConversationTurn {
user_message: Message::user_multimodal(vec![
ContentBlock::text("What is this?"),
ContentBlock::image_url("https://example.com/cat.png"),
]),
assistant_message: Message::assistant("A cat."),
tool_results: Vec::new(),
};
let event = WireStreamEvent::TurnComplete { turn };
let json = serde_json::to_string(&event).expect("event should serialize");
let parsed: WireStreamEvent = serde_json::from_str(&json).expect("should deserialize");
assert_eq!(parsed, event);
}
}