use std::collections::BTreeMap;
use serde_json::json;
use super::{read_sse, CompletionRequest, CompletionResponse, Provider, ToolCall, Usage};
use crate::error::{Error, Result};
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,
}
impl Anthropic {
pub fn new(api_key: impl Into<String>, model: impl Into<String>) -> Self {
Self {
client: reqwest::Client::new(),
api_key: api_key.into(),
model: model.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": request.user },
],
"tools": tools,
})
}
}
impl Provider for Anthropic {
fn name(&self) -> &str {
"anthropic"
}
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse> {
let resp = self
.client
.post(ENDPOINT)
.header("x-api-key", &self.api_key)
.header("anthropic-version", API_VERSION)
.json(&self.body(&request))
.send()
.await
.map_err(|e| Error::Provider(e.to_string()))?;
if !resp.status().is_success() {
let status = resp.status();
let detail = resp.text().await.unwrap_or_default();
return Err(Error::Provider(format!("HTTP {status}: {detail}")));
}
let mut acc = Accumulator::default();
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?;
Ok(acc.finish())
}
}
#[derive(Default)]
struct Accumulator {
text: String,
tool_calls: BTreeMap<u64, (String, String)>,
input_tokens: u64,
output_tokens: u64,
}
impl Accumulator {
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;
}
}
Some("content_block_start") => {
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") => {
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;
}
}
_ => {}
}
}
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 total = self.input_tokens + 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: self.input_tokens,
completion_tokens: self.output_tokens,
total_tokens: total,
}),
}
}
}
#[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");
let req = CompletionRequest {
system: "sys".into(),
user: "hi".into(),
tools: vec![ToolSpec {
name: "write_file".into(),
description: "w".into(),
parameters: json!({"type":"object"}),
}],
};
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 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());
}
}