use super::types::{ContentBlockDeltaData, ContentBlockStartData, StreamEvent};
use crate::provider_connection::DEFAULT_STREAM_IDLE_TIMEOUT;
use crate::providers::response_stream::{OpenedStream, StreamAssembler, response_stream};
use crate::{LlmError, LlmResponse, LlmResponseStream, ProviderError, Result, StopReason};
use futures::Stream;
use std::future::ready;
use tracing::debug;
pub fn process_anthropic_stream(lines: impl Stream<Item = Result<String>> + Send + 'static) -> LlmResponseStream {
response_stream(
ready(Ok(OpenedStream::new(lines))),
|line, turn| decode_line(&line, turn),
DEFAULT_STREAM_IDLE_TIMEOUT,
)
}
pub(super) fn decode_line(line: &str, turn: &mut StreamAssembler<u32>) -> Result<Vec<LlmResponse>> {
match serde_json::from_str(line) {
Ok(event) => decode_event(event, turn),
Err(e) => {
debug!("Failed to parse SSE line: {line} - Error: {e}");
Ok(vec![])
}
}
}
fn decode_event(event: StreamEvent, turn: &mut StreamAssembler<u32>) -> Result<Vec<LlmResponse>> {
let response = match event {
StreamEvent::ContentBlockStart { data } => match data.content_block {
ContentBlockStartData::ToolUse { id, name } => Some(turn.start_tool(data.index, id, name)),
ContentBlockStartData::Text { .. } | ContentBlockStartData::Thinking { .. } => None,
},
StreamEvent::ContentBlockDelta { data } => match data.delta {
ContentBlockDeltaData::TextDelta { text } => {
(!text.is_empty()).then_some(LlmResponse::Text { chunk: text })
}
ContentBlockDeltaData::ThinkingDelta { thinking } => {
(!thinking.is_empty()).then_some(LlmResponse::Reasoning { chunk: thinking })
}
ContentBlockDeltaData::InputJsonDelta { partial_json } => turn.append_tool_args(&data.index, partial_json),
},
StreamEvent::ContentBlockStop { data } => turn.complete_tool(&data.index),
StreamEvent::MessageDelta { data } => {
if let Some(stop_reason) = data.delta.stop_reason.as_deref() {
turn.stop(map_anthropic_stop_reason(stop_reason));
}
data.usage.as_ref().map(|usage| LlmResponse::Usage { tokens: usage.into() })
}
StreamEvent::MessageStop { .. } => {
turn.allow_eof();
None
}
StreamEvent::Error { data } => {
return Err(map_anthropic_stream_error(&data.error.error_type, &data.error.message));
}
StreamEvent::MessageStart { .. } | StreamEvent::Ping => None,
};
Ok(response.into_iter().collect())
}
fn map_anthropic_stream_error(error_type: &str, message: &str) -> LlmError {
let kind = match error_type {
"rate_limit_error" => crate::ProviderErrorKind::RateLimit,
"overloaded_error" | "internal_server_error" | "api_error" => crate::ProviderErrorKind::Server,
_ => crate::ProviderErrorKind::Api,
};
ProviderError::new(kind, format!("Anthropic API error: {error_type} - {message}"))
.with_code(Some(error_type.to_string()))
.into()
}
fn map_anthropic_stop_reason(reason: &str) -> StopReason {
match reason {
"end_turn" | "stop_sequence" => StopReason::EndTurn,
"tool_use" => StopReason::ToolCalls,
"max_tokens" => StopReason::Length,
_ => StopReason::Unknown(reason.to_string()),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::testing::llm_response;
use crate::{ProviderErrorKind, TokenUsage};
use futures::{StreamExt, stream};
use serde_json::{Value, json};
#[tokio::test]
async fn test_process_text_stream() {
let responses = collect_responses(
anthropic_stream()
.text(0, &["Hello", " world"])
.message_delta("end_turn", &usage(10, 25))
.message_stop()
.build(),
)
.await;
assert_eq!(
responses,
llm_response().text(&["Hello", " world"]).usage(10, 25).build_with_stop_reason(StopReason::EndTurn)
);
}
#[tokio::test]
async fn test_process_tool_use_stream() {
let deltas = [r#"{"query":"#, r#""test"}"#];
let responses = collect_responses(
anthropic_stream()
.tool_call(0, "tool_123", "search", &deltas)
.message_delta("tool_use", &usage(10, 15))
.message_stop()
.build(),
)
.await;
assert_eq!(
responses,
llm_response()
.tool_call("tool_123", "search", &deltas)
.usage(10, 15)
.build_with_stop_reason(StopReason::ToolCalls)
);
}
#[tokio::test]
async fn unknown_tool_call_arguments_and_stops_are_ignored() {
let responses =
collect_responses(anthropic_stream().tool_delta(7, "{}").block_stop(7).message_stop().build()).await;
assert_eq!(responses, llm_response().build());
}
#[tokio::test]
async fn test_sequential_tool_calls_reusing_an_index_complete_separately() {
let responses = collect_responses(
anthropic_stream()
.tool_call(0, "tool_123", "search", &[r#"{"query":"test1"}"#])
.tool_call(0, "tool_456", "calculate", &[r#"{"expression":"2+2"}"#])
.message_stop()
.build(),
)
.await;
assert_eq!(
responses,
llm_response()
.tool_call("tool_123", "search", &[r#"{"query":"test1"}"#])
.tool_call("tool_456", "calculate", &[r#"{"expression":"2+2"}"#])
.build()
);
}
#[tokio::test]
async fn stream_closed_mid_tool_call_does_not_complete_it() {
let responses = process_lines(
anthropic_stream().tool_start(0, "tool_a", "bash").tool_delta(0, r#"{"command":"ls"#).build(),
)
.await;
assert!(
matches!(
responses.as_slice(),
[
Ok(LlmResponse::Start),
Ok(LlmResponse::ToolRequestStart { .. }),
Ok(LlmResponse::ToolRequestArg { .. }),
Err(error)
] if error.provider().map(|provider| provider.kind) == Some(ProviderErrorKind::StreamInterrupted)
),
"{responses:?}"
);
}
#[tokio::test]
async fn test_process_thinking_stream() {
let responses = collect_responses(
anthropic_stream()
.thinking(0, &["Let me think", " about this"])
.text(1, &["Here is my answer"])
.message_delta("end_turn", &usage(10, 50))
.message_stop()
.build(),
)
.await;
assert_eq!(
responses,
llm_response()
.reasoning(&["Let me think", " about this"])
.text(&["Here is my answer"])
.usage(10, 50)
.build_with_stop_reason(StopReason::EndTurn)
);
}
#[tokio::test]
async fn test_message_delta_forwards_both_cache_read_and_creation() {
let responses = collect_responses(
anthropic_stream()
.text(0, &["ok"])
.message_delta(
"end_turn",
&json!({
"input_tokens": 100,
"output_tokens": 25,
"cache_creation_input_tokens": 40,
"cache_read_input_tokens": 60
}),
)
.message_stop()
.build(),
)
.await;
let usage = responses.iter().find_map(|r| match r {
LlmResponse::Usage { tokens } => Some(*tokens),
_ => None,
});
assert_eq!(
usage,
Some(TokenUsage {
input_tokens: 200.into(),
output_tokens: 25.into(),
cache_read_tokens: Some(60.into()),
cache_creation_tokens: Some(40.into()),
..TokenUsage::default()
}),
"cached tokens count toward the prompt"
);
}
#[tokio::test]
async fn error_event_ends_the_stream_with_its_kind() {
for (error_type, kind) in [
("rate_limit_error", ProviderErrorKind::RateLimit),
("overloaded_error", ProviderErrorKind::Server),
("invalid_request_error", ProviderErrorKind::Api),
] {
let responses = process_lines(anthropic_stream().error(error_type, "boom").message_stop().build()).await;
assert!(
matches!(
responses.as_slice(),
[Ok(LlmResponse::Start), Err(error)] if error.provider().map(|provider| provider.kind) == Some(kind)
),
"{error_type}: {responses:?}"
);
}
}
#[tokio::test]
async fn test_anthropic_stream_event_enum_deserialization() {
use super::super::types::StreamEvent;
let message_start_json = r#"{"type": "message_start", "message": {"id": "msg_123", "type": "message", "role": "assistant", "content": [], "model": "claude-3", "stop_reason": null, "stop_sequence": null, "usage": {"input_tokens": 10, "output_tokens": 0}}}"#;
let event: StreamEvent = serde_json::from_str(message_start_json).unwrap();
assert!(matches!(event, StreamEvent::MessageStart { .. }));
let content_block_start_json =
r#"{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}"#;
let event: StreamEvent = serde_json::from_str(content_block_start_json).unwrap();
assert!(matches!(event, StreamEvent::ContentBlockStart { .. }));
let content_block_delta_json =
r#"{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello"}}"#;
let event: StreamEvent = serde_json::from_str(content_block_delta_json).unwrap();
assert!(matches!(event, StreamEvent::ContentBlockDelta { .. }));
let ping_json = r#"{"type": "ping"}"#;
let event: StreamEvent = serde_json::from_str(ping_json).unwrap();
assert!(matches!(event, StreamEvent::Ping));
let error_json =
r#"{"type": "error", "error": {"type": "rate_limit_error", "message": "Rate limit exceeded"}}"#;
let event: StreamEvent = serde_json::from_str(error_json).unwrap();
assert!(matches!(event, StreamEvent::Error { .. }));
}
async fn collect_responses(lines: Vec<String>) -> Vec<LlmResponse> {
process_lines(lines).await.into_iter().map(Result::unwrap).collect()
}
async fn process_lines(lines: Vec<String>) -> Vec<Result<LlmResponse>> {
process_anthropic_stream(stream::iter(lines.into_iter().map(Ok))).collect().await
}
fn anthropic_stream() -> AnthropicStreamBuilder {
AnthropicStreamBuilder::default().push(json!({
"type": "message_start",
"message": {
"id": "msg_123",
"type": "message",
"role": "assistant",
"content": [],
"model": "claude-3",
"stop_reason": null,
"stop_sequence": null,
"usage": usage(10, 0)
}
}))
}
#[derive(Default)]
struct AnthropicStreamBuilder {
events: Vec<Value>,
}
impl AnthropicStreamBuilder {
fn text(self, index: u32, chunks: &[&str]) -> Self {
chunks
.iter()
.fold(self.block_start(index, &json!({ "type": "text", "text": "" })), |builder, text| {
builder.block_delta(index, &json!({ "type": "text_delta", "text": text }))
})
.block_stop(index)
}
fn thinking(self, index: u32, chunks: &[&str]) -> Self {
chunks
.iter()
.fold(self.block_start(index, &json!({ "type": "thinking", "thinking": "" })), |builder, thinking| {
builder.block_delta(index, &json!({ "type": "thinking_delta", "thinking": thinking }))
})
.block_stop(index)
}
fn tool_call(self, index: u32, id: &str, name: &str, argument_deltas: &[&str]) -> Self {
argument_deltas
.iter()
.fold(self.tool_start(index, id, name), |builder, delta| builder.tool_delta(index, delta))
.block_stop(index)
}
fn tool_start(self, index: u32, id: &str, name: &str) -> Self {
self.block_start(index, &json!({ "type": "tool_use", "id": id, "name": name }))
}
fn tool_delta(self, index: u32, partial_json: &str) -> Self {
self.block_delta(index, &json!({ "type": "input_json_delta", "partial_json": partial_json }))
}
fn message_delta(self, stop_reason: &str, usage: &Value) -> Self {
self.push(json!({
"type": "message_delta",
"delta": { "stop_reason": stop_reason, "stop_sequence": null },
"usage": usage
}))
}
fn message_stop(self) -> Self {
self.push(json!({ "type": "message_stop" }))
}
fn error(self, error_type: &str, message: &str) -> Self {
self.push(json!({ "type": "error", "error": { "type": error_type, "message": message } }))
}
fn block_start(self, index: u32, content_block: &Value) -> Self {
self.push(json!({ "type": "content_block_start", "index": index, "content_block": content_block }))
}
fn block_delta(self, index: u32, delta: &Value) -> Self {
self.push(json!({ "type": "content_block_delta", "index": index, "delta": delta }))
}
fn block_stop(self, index: u32) -> Self {
self.push(json!({ "type": "content_block_stop", "index": index }))
}
fn push(mut self, event: Value) -> Self {
self.events.push(event);
self
}
fn build(self) -> Vec<String> {
self.events.iter().map(Value::to_string).collect()
}
}
fn usage(input_tokens: u32, output_tokens: u32) -> Value {
json!({ "input_tokens": input_tokens, "output_tokens": output_tokens })
}
}