Luigit
repositories / smith

smith

There are many coding harnesses - but this one is fast

owned by admin

smith-ai/src/anthropic.rs

Raw
//! Anthropic Messages stream adapter (`SMH-SPEC-SPEC0001`, Providers).
//!
//! Translates typed server-sent events, block-indexed content, and partial
//! JSON tool arguments into the normalized vocabulary. Event types, block
//! bookkeeping, and the `message_stop` terminal stay inside this module.

use crate::conformance::StreamDecoder;
use smith::error::{ProviderFault, Result, SmithError};
use smith::id::{MessageId, ToolCallId};
use smith::stream::{StopReason, StreamEvent, Usage};
use std::collections::BTreeMap;

/// Incremental decoder for one Anthropic Messages stream.
#[derive(Debug)]
pub struct AnthropicDecoder {
    buffer: String,
    message_id: MessageId,
    tools: BTreeMap<u64, ToolAccumulator>,
    stop_reason: Option<StopReason>,
    usage: Usage,
    stopped: bool,
    terminal: bool,
}

#[derive(Debug)]
struct ToolAccumulator {
    call_id: ToolCallId,
    name: String,
    arguments: String,
}

impl AnthropicDecoder {
    /// A decoder for a fresh assistant message.
    #[must_use]
    pub fn new() -> Self {
        Self {
            buffer: String::new(),
            message_id: MessageId::new(),
            tools: BTreeMap::new(),
            stop_reason: None,
            usage: Usage::new(),
            stopped: false,
            terminal: false,
        }
    }

    fn process_event(&mut self, event: &str, events: &mut Vec<StreamEvent>) -> Result<()> {
        if self.terminal {
            return Err(protocol("event after message_stop"));
        }
        let value: serde_json::Value = serde_json::from_str(event.trim())
            .map_err(|e| protocol(format!("event payload is not JSON: {e}")))?;
        // The SSE event line and the payload type field agree in practice;
        // the payload is authoritative.
        match value.get("type").and_then(serde_json::Value::as_str) {
            Some("message_start") => {
                if let Some(input) = value
                    .pointer("/message/usage/input_tokens")
                    .and_then(serde_json::Value::as_u64)
                {
                    self.usage.input_tokens = input;
                }
                Ok(())
            }
            Some("content_block_start") => self.block_start(&value),
            Some("content_block_delta") => self.block_delta(&value, events),
            Some("content_block_stop") => self.block_stop(&value, events),
            Some("message_delta") => self.message_delta(&value),
            Some("message_stop") => {
                self.stop(events)?;
                Ok(())
            }
            Some("ping") => Ok(()),
            Some("error") => Err(map_error(&value)),
            other => Err(protocol(format!("unknown event type {other:?}"))),
        }
    }

    fn block_start(&mut self, value: &serde_json::Value) -> Result<()> {
        let index = block_index(value)?;
        let block = value
            .get("content_block")
            .ok_or_else(|| protocol("content_block_start without content_block"))?;
        if block.get("type").and_then(serde_json::Value::as_str) == Some("tool_use") {
            let name = block
                .get("name")
                .and_then(serde_json::Value::as_str)
                .ok_or_else(|| protocol("tool_use block without name"))?;
            self.tools.insert(
                index,
                ToolAccumulator {
                    call_id: ToolCallId::new(),
                    name: name.to_string(),
                    arguments: String::new(),
                },
            );
        }
        Ok(())
    }

    fn block_delta(
        &mut self,
        value: &serde_json::Value,
        events: &mut Vec<StreamEvent>,
    ) -> Result<()> {
        let index = block_index(value)?;
        let delta = value
            .get("delta")
            .ok_or_else(|| protocol("content_block_delta without delta"))?;
        match delta.get("type").and_then(serde_json::Value::as_str) {
            Some("text_delta") => {
                let text = delta
                    .get("text")
                    .and_then(serde_json::Value::as_str)
                    .unwrap_or_default();
                if !text.is_empty() {
                    events.push(StreamEvent::text_delta(self.message_id, text));
                }
            }
            Some("thinking_delta") => {
                let text = delta
                    .get("thinking")
                    .and_then(serde_json::Value::as_str)
                    .unwrap_or_default();
                if !text.is_empty() {
                    events.push(StreamEvent::thinking_delta(self.message_id, text));
                }
            }
            Some("input_json_delta") => {
                let fragment = delta
                    .get("partial_json")
                    .and_then(serde_json::Value::as_str)
                    .unwrap_or_default();
                if let Some(tool) = self.tools.get_mut(&index) {
                    tool.arguments.push_str(fragment);
                }
            }
            Some(other) => return Err(protocol(format!("unknown delta type {other}"))),
            None => return Err(protocol("content_block_delta without delta type")),
        }
        Ok(())
    }

    fn block_stop(
        &mut self,
        value: &serde_json::Value,
        events: &mut Vec<StreamEvent>,
    ) -> Result<()> {
        let index = block_index(value)?;
        if let Some(tool) = self.tools.remove(&index) {
            let input = if tool.arguments.trim().is_empty() {
                serde_json::json!({})
            } else {
                serde_json::from_str(&tool.arguments)
                    .map_err(|e| protocol(format!("partial_json arguments do not assemble: {e}")))?
            };
            events.push(StreamEvent::tool_use(
                self.message_id,
                tool.name,
                input,
                tool.call_id,
            ));
        }
        Ok(())
    }

    fn message_delta(&mut self, value: &serde_json::Value) -> Result<()> {
        if let Some(reason) = value
            .get("delta")
            .and_then(|delta| delta.get("stop_reason"))
            .and_then(serde_json::Value::as_str)
        {
            self.stop_reason = Some(map_stop_reason(reason)?);
        }
        if let Some(usage) = value.get("usage") {
            if let Some(output) = usage
                .get("output_tokens")
                .and_then(serde_json::Value::as_u64)
            {
                self.usage.output_tokens = output;
            }
            for (target, key) in [
                (
                    &mut self.usage.input_cache_hit_tokens,
                    "cache_read_input_tokens",
                ),
                (
                    &mut self.usage.input_cache_write_tokens,
                    "cache_creation_input_tokens",
                ),
            ] {
                if let Some(tokens) = usage.get(key).and_then(serde_json::Value::as_u64) {
                    *target = tokens;
                }
            }
            if let Some(input) = usage
                .get("input_tokens")
                .and_then(serde_json::Value::as_u64)
            {
                self.usage.input_tokens = input;
            }
        }
        Ok(())
    }

    fn stop(&mut self, events: &mut Vec<StreamEvent>) -> Result<()> {
        let Some(reason) = self.stop_reason else {
            return Err(protocol("message_stop without stop_reason"));
        };
        if !self.tools.is_empty() {
            return Err(protocol("message_stop with unfinished tool blocks"));
        }
        self.terminal = true;
        self.stopped = true;
        events.push(StreamEvent::stop(reason, Some(self.usage)));
        Ok(())
    }
}

impl Default for AnthropicDecoder {
    fn default() -> Self {
        Self::new()
    }
}

impl StreamDecoder for AnthropicDecoder {
    fn push(&mut self, bytes: &[u8]) -> Result<Vec<StreamEvent>> {
        self.buffer.push_str(&String::from_utf8_lossy(bytes));
        let mut events = Vec::new();
        while let Some((event, rest)) = take_sse_event(&self.buffer) {
            self.buffer = rest;
            self.process_event(&event, &mut events)?;
        }
        Ok(events)
    }

    fn finish(&mut self) -> Result<Vec<StreamEvent>> {
        let mut events = Vec::new();
        if !self.buffer.trim().is_empty() {
            let event = std::mem::take(&mut self.buffer);
            self.process_event(event.trim_end(), &mut events)?;
        }
        if !self.stopped {
            return Err(SmithError::Provider {
                fault: ProviderFault::Incomplete,
            });
        }
        Ok(events)
    }
}

/// Anthropic SSE carries `event:` and `data:` lines; the data payload rules.
fn take_sse_event(buffer: &str) -> Option<(String, String)> {
    let mut data = String::new();
    let mut consumed = 0;
    for line in buffer.split_inclusive('\n') {
        let trimmed = line.trim_end_matches(['\r', '\n']);
        if trimmed.is_empty() {
            if data.is_empty() {
                consumed += line.len();
                continue;
            }
            let rest = buffer[consumed + line.len()..].to_string();
            return Some((data, rest));
        }
        if let Some(payload) = trimmed.strip_prefix("data:") {
            data.push_str(payload.strip_prefix(' ').unwrap_or(payload));
        }
        consumed += line.len();
    }
    None
}

fn block_index(value: &serde_json::Value) -> Result<u64> {
    value
        .get("index")
        .and_then(serde_json::Value::as_u64)
        .ok_or_else(|| protocol("block event without index"))
}

fn map_stop_reason(reason: &str) -> Result<StopReason> {
    match reason {
        "end_turn" => Ok(StopReason::EndTurn),
        "max_tokens" => Ok(StopReason::Limit),
        "stop_sequence" => Ok(StopReason::StopSequence),
        "tool_use" => Ok(StopReason::ToolUse),
        "refusal" => Ok(StopReason::Refusal),
        // Server-side tool use is out of scope; pause is explicit, not silent.
        "pause_turn" => Err(protocol(
            "pause_turn requires server-side tools, which this build does not support",
        )),
        other => Err(protocol(format!("unmapped stop_reason {other}"))),
    }
}

fn map_error(value: &serde_json::Value) -> SmithError {
    let error = value.get("error").cloned().unwrap_or_else(|| value.clone());
    let kind = error.get("type").and_then(serde_json::Value::as_str);
    let message = error
        .get("message")
        .and_then(serde_json::Value::as_str)
        .unwrap_or("unknown provider error")
        .to_string();
    let fault = match kind {
        Some("rate_limit_error") => ProviderFault::RateLimit {
            retry_after_ms: None,
        },
        Some("api_error" | "timeout_error") => ProviderFault::Transient { message },
        Some("invalid_request_error") => ProviderFault::Invalid {
            field: None,
            message,
        },
        Some("authentication_error" | "permission_error") => {
            ProviderFault::Authentication { message }
        }
        Some("overloaded_error") => ProviderFault::Overloaded { message },
        _ => ProviderFault::Protocol { message },
    };
    SmithError::Provider { fault }
}

fn protocol(message: impl Into<String>) -> SmithError {
    SmithError::Provider {
        fault: ProviderFault::Protocol {
            message: message.into(),
        },
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::conformance::{Expected, Fixture, conformance_report, sse_boundary_fixtures};
    use smith::stream::StopReason;

    fn sse(lines: &[&str]) -> Vec<u8> {
        let mut out = String::new();
        for line in lines {
            out.push_str("data: ");
            out.push_str(line);
            out.push_str("\n\n");
        }
        out.into_bytes()
    }

    fn tool_fixture() -> Fixture {
        Fixture {
            name: "text, thinking, and assembled tool input",
            chunks: vec![sse(&[
                r#"{"type":"message_start","message":{"usage":{"input_tokens":12}}}"#,
                r#"{"type":"content_block_start","index":0,"content_block":{"type":"text"}}"#,
                r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Look"}}"#,
                r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"ing"}}"#,
                r#"{"type":"content_block_stop","index":0}"#,
                r#"{"type":"content_block_start","index":1,"content_block":{"type":"thinking"}}"#,
                r#"{"type":"content_block_delta","index":1,"delta":{"type":"thinking_delta","thinking":"hmm"}}"#,
                r#"{"type":"content_block_stop","index":1}"#,
                r#"{"type":"content_block_start","index":2,"content_block":{"type":"tool_use","id":"t","name":"write"}}"#,
                r#"{"type":"content_block_delta","index":2,"delta":{"type":"input_json_delta","partial_json":"{\"path\""}}"#,
                r#"{"type":"content_block_delta","index":2,"delta":{"type":"input_json_delta","partial_json":":[1,2]}"}}"#,
                r#"{"type":"content_block_stop","index":2}"#,
                r#"{"type":"message_delta","delta":{"stop_reason":"tool_use"},"usage":{"output_tokens":7,"cache_read_input_tokens":3}}"#,
                r#"{"type":"message_stop"}"#,
            ])],
            expected: vec![
                Expected::Text("Look".to_string()),
                Expected::Text("ing".to_string()),
                Expected::Thinking("hmm".to_string()),
                Expected::Tool {
                    name: "write".to_string(),
                    input: serde_json::json!({"path": [1, 2]}),
                },
                Expected::Stop {
                    reason: StopReason::ToolUse,
                    usage: Some(Usage {
                        input_tokens: 12,
                        output_tokens: 7,
                        input_cache_hit_tokens: 3,
                        input_cache_write_tokens: 0,
                    }),
                },
            ],
        }
    }

    fn fault_fixtures() -> Vec<Fixture> {
        vec![
            Fixture {
                name: "overloaded error",
                chunks: vec![sse(&[
                    r#"{"type":"error","error":{"type":"overloaded_error","message":"busy"}}"#,
                ])],
                expected: vec![Expected::Fault("PROVIDER_OVERLOADED".to_string())],
            },
            Fixture {
                name: "refusal stop",
                chunks: vec![sse(&[
                    r#"{"type":"message_delta","delta":{"stop_reason":"refusal"},"usage":{"output_tokens":1}}"#,
                    r#"{"type":"message_stop"}"#,
                ])],
                expected: vec![Expected::Stop {
                    reason: StopReason::Refusal,
                    usage: Some(Usage {
                        input_tokens: 0,
                        output_tokens: 1,
                        input_cache_hit_tokens: 0,
                        input_cache_write_tokens: 0,
                    }),
                }],
            },
            Fixture {
                name: "pause turn is explicit",
                chunks: vec![sse(&[
                    r#"{"type":"message_delta","delta":{"stop_reason":"pause_turn"}}"#,
                    r#"{"type":"message_stop"}"#,
                ])],
                expected: vec![Expected::Fault("PROVIDER_PROTOCOL".to_string())],
            },
            Fixture {
                name: "message_stop without stop reason",
                chunks: vec![sse(&[r#"{"type":"message_stop"}"#])],
                expected: vec![Expected::Fault("PROVIDER_PROTOCOL".to_string())],
            },
            Fixture {
                name: "truncated before message stop",
                chunks: vec![sse(&[
                    r#"{"type":"content_block_start","index":0,"content_block":{"type":"text"}}"#,
                ])],
                expected: vec![Expected::Fault("PROVIDER_INCOMPLETE".to_string())],
            },
            Fixture {
                name: "event after terminal",
                chunks: vec![
                    sse(&[
                        r#"{"type":"message_delta","delta":{"stop_reason":"end_turn"}}"#,
                        r#"{"type":"message_stop"}"#,
                    ]),
                    sse(&[
                        r#"{"type":"content_block_start","index":0,"content_block":{"type":"text"}}"#,
                    ]),
                ],
                expected: vec![
                    Expected::Stop {
                        reason: StopReason::EndTurn,
                        usage: Some(Usage::new()),
                    },
                    Expected::Fault("PROVIDER_PROTOCOL".to_string()),
                ],
            },
        ]
    }

    fn boundary_fixtures() -> Vec<Fixture> {
        sse_boundary_fixtures(
            r#"data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hi"}}"#,
            &sse(&[
                r#"{"type":"message_delta","delta":{"stop_reason":"end_turn"}}"#,
                r#"{"type":"message_stop"}"#,
            ]),
            &[
                Expected::Text("Hi".to_string()),
                Expected::Stop {
                    reason: StopReason::EndTurn,
                    usage: Some(Usage::new()),
                },
            ],
        )
    }

    fn interleaved_tools_fixture() -> Fixture {
        Fixture {
            name: "tool blocks stay keyed by their own index",
            chunks: vec![sse(&[
                r#"{"type":"content_block_start","index":2,"content_block":{"type":"tool_use","id":"a","name":"read"}}"#,
                r#"{"type":"content_block_start","index":3,"content_block":{"type":"tool_use","id":"b","name":"write"}}"#,
                r#"{"type":"content_block_delta","index":3,"delta":{"type":"input_json_delta","partial_json":"{\"n\":3}"}}"#,
                r#"{"type":"content_block_delta","index":2,"delta":{"type":"input_json_delta","partial_json":"{\"n\":2}"}}"#,
                r#"{"type":"content_block_stop","index":2}"#,
                r#"{"type":"content_block_stop","index":3}"#,
                r#"{"type":"message_delta","delta":{"stop_reason":"tool_use"}}"#,
                r#"{"type":"message_stop"}"#,
            ])],
            expected: vec![
                Expected::Tool {
                    name: "read".to_string(),
                    input: serde_json::json!({"n": 2}),
                },
                Expected::Tool {
                    name: "write".to_string(),
                    input: serde_json::json!({"n": 3}),
                },
                Expected::Stop {
                    reason: StopReason::ToolUse,
                    usage: Some(Usage::new()),
                },
            ],
        }
    }

    #[test]
    fn vendor_error_types_map_to_provider_faults() {
        use crate::conformance::faults::*;
        let cases = [
            ("rate_limit_error", rate_limit()),
            ("api_error", transient()),
            ("timeout_error", transient()),
            ("invalid_request_error", invalid()),
            ("authentication_error", authentication()),
            ("permission_error", authentication()),
            ("overloaded_error", overloaded()),
            ("not_a_vendor_type", protocol()),
        ];
        for (kind, fault) in cases {
            let value =
                serde_json::json!({"type": "error", "error": {"type": kind, "message": "m"}});
            assert_eq!(map_error(&value), SmithError::Provider { fault }, "{kind}");
        }
    }

    #[test]
    fn anthropic_adapter_satisfies_conformance() {
        let mut fixtures = vec![tool_fixture(), interleaved_tools_fixture()];
        fixtures.extend(boundary_fixtures());
        fixtures.extend(fault_fixtures());
        let failures = conformance_report(&|| Box::new(AnthropicDecoder::new()), &fixtures);
        assert!(failures.is_empty(), "{failures:?}");
    }
}