use std::collections::BTreeMap;
use std::time::{Duration, Instant};
use serde_json::json;
use super::{read_sse, CompletionRequest, CompletionResponse, Provider, ToolCall, Usage};
use crate::error::{Error, Result};
pub use crate::net::REQUEST_TIMEOUT;
const ENDPOINT: &str = "https://api.anthropic.com/v1/messages";
const API_VERSION: &str = "2023-06-01";
const MAX_TOKENS: u64 = 8192;
pub struct Anthropic {
client: reqwest::Client,
api_key: String,
model: String,
endpoint: String,
}
impl Anthropic {
pub fn new(api_key: impl Into<String>, model: impl Into<String>) -> Self {
Self {
client: crate::net::http_client(),
api_key: api_key.into(),
model: model.into(),
endpoint: ENDPOINT.to_string(),
}
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.client = crate::net::http_client_with_timeout(timeout);
self
}
#[cfg(test)]
pub(crate) fn at(endpoint: impl Into<String>, timeout: std::time::Duration) -> Self {
Self {
client: crate::net::http_client_with_timeout(timeout),
api_key: "test-key".into(),
model: "test-model".into(),
endpoint: endpoint.into(),
}
}
pub fn from_env() -> Result<Self> {
let api_key = std::env::var("ANTHROPIC_API_KEY")
.map_err(|_| Error::Config("ANTHROPIC_API_KEY is not set".into()))?;
let model = std::env::var("ANTHROPIC_MODEL")
.map_err(|_| Error::Config("ANTHROPIC_MODEL is not set".into()))?;
Ok(Self::new(api_key, model))
}
fn body(&self, request: &CompletionRequest) -> serde_json::Value {
let tools: Vec<serde_json::Value> = request
.tools
.iter()
.map(|t| {
json!({
"name": t.name,
"description": t.description,
"input_schema": t.parameters,
})
})
.collect();
json!({
"model": self.model,
"max_tokens": MAX_TOKENS,
"stream": true,
"system": request.system,
"messages": [
{ "role": "user", "content": Self::user_content(request) },
],
"tools": tools,
})
}
#[cfg(feature = "media")]
fn user_content(request: &CompletionRequest) -> serde_json::Value {
if request.media.is_empty() {
return json!(request.user);
}
let mut parts: Vec<serde_json::Value> = request
.media
.iter()
.map(|m| {
json!({
"type": "image",
"source": {
"type": "base64",
"media_type": m.media_type,
"data": m.base64,
},
})
})
.collect();
parts.push(json!({ "type": "text", "text": request.user }));
json!(parts)
}
#[cfg(not(feature = "media"))]
fn user_content(request: &CompletionRequest) -> serde_json::Value {
json!(request.user)
}
}
impl Provider for Anthropic {
fn name(&self) -> &str {
"anthropic"
}
fn endpoint(&self) -> Option<&str> {
Some(&self.endpoint)
}
#[cfg(feature = "media")]
fn accepts_images(&self) -> bool {
true
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse> {
#[cfg(feature = "media")]
super::ensure_media_accepted(self.name(), self.accepts_images(), &request)?;
let sent = Instant::now();
let resp = self
.client
.post(&self.endpoint)
.header("x-api-key", &self.api_key)
.header("anthropic-version", API_VERSION)
.json(&self.body(&request))
.send()
.await?;
let resp = super::ensure_success(resp).await?;
let mut acc = Accumulator::since(sent);
read_sse(resp, |data| {
if let Ok(value) = serde_json::from_str::<serde_json::Value>(data) {
if value.get("type").and_then(|t| t.as_str()) == Some("message_stop") {
return true;
}
acc.ingest(&value);
}
false
})
.await?;
super::ensure_parsed(acc.finish())
}
}
#[derive(Default)]
struct Accumulator {
text: String,
tool_calls: BTreeMap<u64, (String, String)>,
input_tokens: u64,
output_tokens: u64,
cache_write_tokens: u64,
cache_read_tokens: u64,
server_tool_requests: u64,
model: Option<String>,
finish_reason: Option<String>,
sent: Option<Instant>,
ttft_ms: Option<u64>,
}
impl Accumulator {
fn since(sent: Instant) -> Self {
Self {
sent: Some(sent),
..Default::default()
}
}
fn mark_first_token(&mut self) {
if let Some(sent) = self.sent {
self.ttft_ms
.get_or_insert(sent.elapsed().as_millis() as u64);
}
}
fn ingest(&mut self, value: &serde_json::Value) {
let index = || value.get("index").and_then(|i| i.as_u64()).unwrap_or(0);
match value.get("type").and_then(|t| t.as_str()) {
Some("message_start") => {
if let Some(n) = value
.pointer("/message/usage/input_tokens")
.and_then(|v| v.as_u64())
{
self.input_tokens = n;
}
if let Some(m) = value
.pointer("/message/model")
.and_then(|v| v.as_str())
.filter(|m| !m.is_empty())
{
self.model = Some(m.to_string());
}
self.ingest_usage(value.pointer("/message/usage"));
}
Some("content_block_start") => {
self.mark_first_token();
if let Some(cb) = value.get("content_block") {
if cb.get("type").and_then(|t| t.as_str()) == Some("tool_use") {
let name = cb
.get("name")
.and_then(|n| n.as_str())
.unwrap_or_default()
.to_string();
self.tool_calls.entry(index()).or_default().0 = name;
}
}
}
Some("content_block_delta") => {
self.mark_first_token();
let delta = value.get("delta");
match delta.and_then(|d| d.get("type")).and_then(|t| t.as_str()) {
Some("text_delta") => {
if let Some(t) = delta.and_then(|d| d.get("text")).and_then(|t| t.as_str())
{
self.text.push_str(t);
}
}
Some("input_json_delta") => {
if let Some(p) = delta
.and_then(|d| d.get("partial_json"))
.and_then(|p| p.as_str())
{
self.tool_calls.entry(index()).or_default().1.push_str(p);
}
}
_ => {}
}
}
Some("message_delta") => {
if let Some(n) = value
.pointer("/usage/output_tokens")
.and_then(|v| v.as_u64())
{
self.output_tokens = n;
}
if let Some(r) = value
.pointer("/delta/stop_reason")
.and_then(|v| v.as_str())
.filter(|r| !r.is_empty())
{
self.finish_reason = Some(r.to_string());
}
self.ingest_usage(value.pointer("/usage"));
}
_ => {}
}
}
fn ingest_usage(&mut self, usage: Option<&serde_json::Value>) {
let Some(usage) = usage else { return };
let get = |k: &str| usage.get(k).and_then(|v| v.as_u64());
if let Some(n) = get("cache_creation_input_tokens") {
self.cache_write_tokens = n;
}
if let Some(n) = get("cache_read_input_tokens") {
self.cache_read_tokens = n;
}
if let Some(counts) = usage.get("server_tool_use").and_then(|v| v.as_object()) {
let sum = counts.values().filter_map(|v| v.as_u64()).sum();
if sum > 0 {
self.server_tool_requests = sum;
}
}
}
fn finish(self) -> CompletionResponse {
let tool_calls = self
.tool_calls
.into_values()
.filter(|(name, _)| !name.is_empty())
.map(|(name, args)| ToolCall {
name,
arguments: serde_json::from_str(if args.is_empty() { "{}" } else { &args })
.unwrap_or(serde_json::Value::Null),
})
.collect();
let prompt = self.input_tokens + self.cache_read_tokens + self.cache_write_tokens;
let total = prompt + self.output_tokens;
CompletionResponse {
text: if self.text.is_empty() {
None
} else {
Some(self.text)
},
tool_calls,
usage: (total > 0).then_some(Usage {
prompt_tokens: prompt,
completion_tokens: self.output_tokens,
total_tokens: total,
cache_read_tokens: self.cache_read_tokens,
cache_write_tokens: self.cache_write_tokens,
reasoning_tokens: 0,
server_tool_requests: self.server_tool_requests,
}),
model: self.model,
finish_reason: self.finish_reason,
ttft_ms: self.ttft_ms,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::provider::ToolSpec;
#[test]
fn body_maps_tools_to_input_schema_and_system_top_level() {
let a = Anthropic::new("k", "claude-x");
#[allow(clippy::needless_update)] let req = CompletionRequest {
system: "sys".into(),
user: "hi".into(),
tools: vec![ToolSpec {
name: "write_file".into(),
description: "w".into(),
parameters: json!({"type":"object"}),
}],
..Default::default()
};
let b = a.body(&req);
assert_eq!(b["system"], "sys");
assert_eq!(b["messages"][0]["content"], "hi");
assert_eq!(b["tools"][0]["name"], "write_file");
assert_eq!(b["tools"][0]["input_schema"], json!({"type":"object"}));
assert!(b["max_tokens"].is_u64());
}
#[test]
fn accumulates_tool_use_from_input_json_deltas_and_usage() {
let mut acc = Accumulator::default();
acc.ingest(&json!({"type":"message_start","message":{"usage":{"input_tokens":11}}}));
acc.ingest(&json!({"type":"content_block_start","index":0,
"content_block":{"type":"tool_use","id":"t1","name":"write_file","input":{}}}));
acc.ingest(&json!({"type":"content_block_delta","index":0,
"delta":{"type":"input_json_delta","partial_json":"{\"path\":\"a"}}));
acc.ingest(&json!({"type":"content_block_delta","index":0,
"delta":{"type":"input_json_delta","partial_json":".rs\",\"content\":\"x\"}"}}));
acc.ingest(&json!({"type":"message_delta","usage":{"output_tokens":7}}));
let out = acc.finish();
assert_eq!(out.tool_calls.len(), 1);
assert_eq!(out.tool_calls[0].name, "write_file");
assert_eq!(out.tool_calls[0].arguments["path"], "a.rs");
assert_eq!(out.tool_calls[0].arguments["content"], "x");
let u = out.usage.unwrap();
assert_eq!(u.prompt_tokens, 11);
assert_eq!(u.completion_tokens, 7);
assert_eq!(u.total_tokens, 18);
}
#[test]
fn cache_tokens_the_model_reports_reach_usage_with_the_model_and_stop_reason() {
let mut acc = Accumulator::default();
acc.ingest(&json!({"type":"message_start","message":{
"model":"claude-sonnet-4-5",
"usage":{"input_tokens":11,"cache_creation_input_tokens":300,
"cache_read_input_tokens":1_200}}}));
acc.ingest(
&json!({"type":"message_delta","delta":{"stop_reason":"max_tokens"},
"usage":{"output_tokens":7,"server_tool_use":{"web_search_requests":2}}}),
);
let out = acc.finish();
let u = out.usage.unwrap();
assert_eq!(u.cache_write_tokens, 300);
assert_eq!(u.cache_read_tokens, 1_200);
assert_eq!(u.server_tool_requests, 2);
assert_eq!(u.prompt_tokens, 11 + 300 + 1_200);
assert_eq!(u.total_tokens, 11 + 300 + 1_200 + 7);
assert_eq!(u.reasoning_tokens, 0);
assert_eq!(out.model.as_deref(), Some("claude-sonnet-4-5"));
assert_eq!(out.finish_reason.as_deref(), Some("max_tokens"));
}
#[test]
fn a_stream_that_reports_no_cache_or_stop_reason_yields_zeros_and_nones() {
let mut acc = Accumulator::default();
acc.ingest(&json!({"type":"message_start","message":{"usage":{"input_tokens":11}}}));
acc.ingest(&json!({"type":"message_delta","usage":{"output_tokens":7}}));
let out = acc.finish();
let u = out.usage.unwrap();
assert_eq!((u.cache_read_tokens, u.cache_write_tokens), (0, 0));
assert_eq!(u.server_tool_requests, 0);
assert_eq!(u.prompt_tokens, 11);
assert_eq!(u.total_tokens, 18);
assert_eq!(out.model, None);
assert_eq!(out.finish_reason, None);
assert_eq!(out.ttft_ms, None);
}
#[test]
fn the_ttft_clock_stops_at_the_first_content_event() {
let mut acc = Accumulator::since(Instant::now() - Duration::from_millis(40));
acc.ingest(&json!({"type":"content_block_delta","index":0,
"delta":{"type":"text_delta","text":"a"}}));
let first = acc.ttft_ms.expect("measured");
std::thread::sleep(Duration::from_millis(15));
acc.ingest(&json!({"type":"content_block_delta","index":0,
"delta":{"type":"text_delta","text":"b"}}));
assert_eq!(acc.ttft_ms, Some(first), "a later chunk moved the clock");
assert!(
first >= 40,
"the clock runs from the request, got {first}ms"
);
}
#[test]
fn accumulates_plain_text() {
let mut acc = Accumulator::default();
acc.ingest(&json!({"type":"content_block_delta","index":0,
"delta":{"type":"text_delta","text":"hello "}}));
acc.ingest(&json!({"type":"content_block_delta","index":0,
"delta":{"type":"text_delta","text":"world"}}));
let out = acc.finish();
assert_eq!(out.text.as_deref(), Some("hello world"));
assert!(out.tool_calls.is_empty());
}
}
#[cfg(all(test, feature = "media"))]
mod media_wire {
use super::*;
use crate::provider::Media;
#[allow(clippy::needless_update)] fn req_with_image() -> CompletionRequest {
CompletionRequest {
system: "sys".into(),
user: "what is this".into(),
media: vec![Media::image("image/png", &[1, 2, 3]).unwrap()],
..Default::default()
}
}
#[test]
fn an_image_becomes_a_base64_source_block_before_the_text() {
let b = Anthropic::new("k", "claude-x").body(&req_with_image());
let content = &b["messages"][0]["content"];
assert!(content.is_array(), "content must be blocks, got {content}");
assert_eq!(content[0]["type"], "image");
assert_eq!(content[0]["source"]["type"], "base64");
assert_eq!(content[0]["source"]["media_type"], "image/png");
assert_eq!(content[0]["source"]["data"], "AQID");
assert_eq!(content[1]["type"], "text");
assert_eq!(content[1]["text"], "what is this");
}
#[test]
fn a_request_without_an_image_still_sends_a_bare_string() {
#[allow(clippy::needless_update)] let b = Anthropic::new("k", "claude-x").body(&CompletionRequest {
system: "sys".into(),
user: "no picture".into(),
..Default::default()
});
assert_eq!(b["messages"][0]["content"], "no picture");
}
}