use async_trait::async_trait;
use serde::Deserialize;
use serde_json::{Value, json};
use crate::core::Secret;
use super::wire::{RESPOND_TOOL, classify_status, classify_transport, structured};
use super::{
Completion, ModelError, ModelId, ModelProvider, Request, SchemaMode, Usage, anthropic_stream,
sse,
};
pub struct Anthropic {
http: reqwest::Client,
key: Secret,
base: String,
version: String,
max_tokens: u32,
default_schema_mode: SchemaMode,
schema_modes: std::collections::BTreeMap<String, SchemaMode>,
stream: bool,
egress: Option<crate::core::Egress>,
}
impl std::fmt::Debug for Anthropic {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Anthropic")
.field("base", &self.base)
.field("version", &self.version)
.field("max_tokens", &self.max_tokens)
.field("key", &"<redacted>")
.finish_non_exhaustive()
}
}
impl Anthropic {
pub const VERSION: &'static str = "2023-06-01";
pub const DEFAULT_MAX_TOKENS: u32 = 4096;
pub fn new(key: impl Into<String>) -> Result<Self, ModelError> {
let http = reqwest::Client::builder()
.build()
.map_err(|e| ModelError::Unreachable {
model: ModelId::new("anthropic", "*"),
detail: format!("could not build an HTTP client: {e}"),
})?;
Ok(Self {
http,
key: Secret::new(key),
base: "https://api.anthropic.com".to_owned(),
version: Self::VERSION.to_owned(),
max_tokens: Self::DEFAULT_MAX_TOKENS,
default_schema_mode: SchemaMode::Native,
schema_modes: std::collections::BTreeMap::new(),
stream: true,
egress: None,
})
}
#[must_use]
pub fn base(mut self, base: impl Into<String>) -> Self {
self.base = base.into();
self
}
#[must_use]
pub const fn max_tokens(mut self, n: u32) -> Self {
self.max_tokens = n;
self
}
#[must_use]
pub fn structured_via(mut self, mode: SchemaMode) -> Self {
self.default_schema_mode = mode;
self
}
#[must_use]
pub fn structured_via_for(mut self, model: impl Into<String>, mode: SchemaMode) -> Self {
self.schema_modes.insert(model.into(), mode);
self
}
#[must_use]
pub fn egress(mut self, egress: crate::core::Egress) -> Self {
self.egress = Some(egress);
self
}
#[must_use]
pub const fn buffered(mut self) -> Self {
self.stream = false;
self
}
fn check_egress(&self, model: &ModelId) -> Result<(), ModelError> {
let Some(egress) = &self.egress else {
return Ok(());
};
let host = reqwest::Url::parse(&self.base)
.ok()
.and_then(|u| u.host_str().map(ToOwned::to_owned));
egress
.permits(host.as_deref())
.map_err(|e| ModelError::Refused {
model: model.clone(),
detail: e.to_string(),
})
}
fn mode_for(&self, model: &ModelId) -> SchemaMode {
self.schema_modes
.get(&model.model)
.copied()
.unwrap_or(self.default_schema_mode)
}
}
#[allow(clippy::struct_field_names)]
#[derive(Debug, Deserialize)]
struct ApiUsage {
#[serde(default)]
input_tokens: u64,
#[serde(default)]
output_tokens: u64,
#[serde(default)]
cache_creation_input_tokens: u64,
#[serde(default)]
cache_read_input_tokens: u64,
}
#[derive(Debug, Deserialize)]
struct ApiBlock {
#[serde(rename = "type")]
kind: String,
#[serde(default)]
text: String,
#[serde(default)]
input: Option<Value>,
}
#[derive(Debug, Deserialize)]
struct ApiResponse {
#[serde(default)]
content: Vec<ApiBlock>,
#[serde(default)]
usage: Option<ApiUsage>,
#[serde(default)]
stop_reason: Option<String>,
}
impl ApiResponse {
fn usage(&self) -> Usage {
let u = self.usage.as_ref();
let write = u.map_or(0, |u| u.cache_creation_input_tokens);
let read = u.map_or(0, |u| u.cache_read_input_tokens);
Usage {
input_tokens: u.map_or(0, |u| u.input_tokens) + write + read,
output_tokens: u.map_or(0, |u| u.output_tokens),
cache_write_tokens: write,
cache_read_tokens: read,
minor_units: 0,
}
}
fn text(&self) -> String {
self.content
.iter()
.filter(|b| b.kind == "text")
.map(|b| b.text.as_str())
.collect::<Vec<_>>()
.join("")
}
fn forced_tool_input(&self) -> Option<&Value> {
self.content
.iter()
.find(|b| b.kind == "tool_use")
.and_then(|b| b.input.as_ref())
}
}
fn messages(prompt: &Value) -> Value {
match prompt {
Value::String(s) => json!([{ "role": "user", "content": s }]),
Value::Array(_) => prompt.clone(),
other => other.get("messages").cloned().unwrap_or_else(|| {
let mut rest = other.clone();
if let Some(map) = rest.as_object_mut() {
map.remove("system");
}
json!([{ "role": "user", "content": rest.to_string() }])
}),
}
}
fn system(prompt: &Value) -> Option<Value> {
prompt.get("system").cloned().filter(|s| !s.is_null())
}
fn interpret(
model: &ModelId,
schema: Option<&Value>,
emulating: bool,
text: String,
forced: Option<Value>,
usage: Usage,
stop_reason: Option<String>,
) -> Result<Completion, ModelError> {
if stop_reason.as_deref() == Some("refusal") {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: "the model declined to answer".to_owned(),
});
}
let (text, structured_value) = if emulating {
let Some(value) = forced else {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: "a tool call was forced and no usable arguments came back — \
the model did not honour `tool_choice`, or its streamed \
fragments did not reassemble into JSON"
.to_owned(),
});
};
(value.to_string(), Some(value))
} else {
if text.is_empty() {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: "the answer carried no text content".to_owned(),
});
}
let parsed_schema = structured(schema, &text, model, usage)?;
(text, parsed_schema)
};
Ok(Completion {
structured: structured_value,
text,
usage,
truncated: stop_reason.as_deref() == Some("max_tokens"),
stop_reason,
})
}
impl Anthropic {
fn body(&self, model: &ModelId, prompt: &Value, schema: Option<&Value>) -> Value {
let mut body = json!({
"model": model.model,
"max_tokens": self.max_tokens,
"messages": messages(prompt),
});
if let Some(system) = system(prompt) {
body["system"] = system;
}
if let Some(schema) = schema {
match self.mode_for(model) {
SchemaMode::Native => {
body["output_config"] = json!({
"format": { "type": "json_schema", "schema": schema }
});
}
SchemaMode::ForcedTool => {
body["tools"] = json!([{
"name": RESPOND_TOOL,
"description": "Return the answer in the required shape.",
"input_schema": schema,
}]);
body["tool_choice"] = json!({ "type": "tool", "name": RESPOND_TOOL });
}
}
}
if self.stream {
body["stream"] = json!(true);
}
body
}
async fn read_buffered(
&self,
response: reqwest::Response,
model: &ModelId,
schema: Option<&Value>,
) -> Result<Completion, ModelError> {
let parsed: ApiResponse = response.json().await.map_err(|e| ModelError::Unusable {
model: model.clone(),
usage: Usage::default(),
detail: format!("the response body did not parse: {e}"),
})?;
let emulating = schema.is_some() && self.mode_for(model) == SchemaMode::ForcedTool;
interpret(
model,
schema,
emulating,
parsed.text(),
parsed.forced_tool_input().cloned(),
parsed.usage(),
parsed.stop_reason.clone(),
)
}
async fn read_streamed(
&self,
response: reqwest::Response,
model: &ModelId,
schema: Option<&Value>,
) -> Result<Completion, ModelError> {
use futures_util::StreamExt;
let mut decoder = sse::Decoder::new();
let mut acc = anthropic_stream::Accumulator::new();
let mut body = response.bytes_stream();
while let Some(chunk) = body.next().await {
let chunk = match chunk {
Ok(chunk) => chunk,
Err(e) => return Err(severed(model, &acc, &e.to_string())),
};
for event in decoder.push(&chunk) {
acc.event(&event.name, &event.data);
}
if let Some(err) = acc.error() {
return Err(stream_error(model, &acc, err));
}
}
if !acc.complete() {
return Err(severed(
model,
&acc,
"the stream ended before `message_stop`",
));
}
let emulating = schema.is_some() && self.mode_for(model) == SchemaMode::ForcedTool;
interpret(
model,
schema,
emulating,
acc.text().to_owned(),
acc.forced_tool_input(),
acc.billed(),
acc.stop_reason().map(ToOwned::to_owned),
)
}
}
fn severed(model: &ModelId, acc: &anthropic_stream::Accumulator, detail: &str) -> ModelError {
if acc.started() {
return ModelError::Interrupted {
model: model.clone(),
usage: acc.billed(),
detail: detail.to_owned(),
};
}
ModelError::Unavailable {
model: model.clone(),
detail: format!("the stream ended before it began: {detail}"),
}
}
fn stream_error(
model: &ModelId,
acc: &anthropic_stream::Accumulator,
err: &anthropic_stream::StreamError,
) -> ModelError {
let detail = format!("{}: {}", err.kind, err.message);
if acc.started() {
return ModelError::Interrupted {
model: model.clone(),
usage: acc.billed(),
detail,
};
}
match err.kind.as_str() {
"overloaded_error" | "rate_limit_error" => ModelError::RateLimited {
model: model.clone(),
detail,
},
"invalid_request_error"
| "authentication_error"
| "permission_error"
| "not_found_error" => ModelError::Refused {
model: model.clone(),
detail,
},
_ => ModelError::Unavailable {
model: model.clone(),
detail,
},
}
}
#[async_trait]
impl ModelProvider for Anthropic {
async fn complete(&self, request: Request<'_>) -> Result<Completion, ModelError> {
let Request {
model,
prompt,
schema,
} = request;
self.check_egress(model)?;
let response = self
.http
.post(format!("{}/v1/messages", self.base))
.header("x-api-key", self.key.expose())
.header("anthropic-version", &self.version)
.json(&self.body(model, prompt, schema))
.send()
.await
.map_err(|e| classify_transport(model, &e))?;
let status = response.status();
if !status.is_success() {
let detail = response.text().await.unwrap_or_default();
return Err(classify_status(model, status.as_u16(), &detail));
}
if self.stream {
self.read_streamed(response, model, schema).await
} else {
self.read_buffered(response, model, schema).await
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn driver() -> Anthropic {
Anthropic::new("test-key").expect("build the driver")
}
#[test]
fn a_system_instruction_rides_beside_the_messages() {
let body = driver().body(
&ModelId::new("anthropic", "claude-x"),
&json!({ "system": "answer only in French", "messages": [{"role": "user", "content": "hi"}] }),
None,
);
assert_eq!(
body["system"], "answer only in French",
"the system instruction must be a top-level parameter: {body}"
);
assert_eq!(
body["messages"],
json!([{ "role": "user", "content": "hi" }]),
"it must not also be pushed into the conversation: {body}"
);
}
#[test]
fn a_system_instruction_is_not_shown_as_the_question() {
let body = driver().body(
&ModelId::new("anthropic", "claude-x"),
&json!({ "system": "be terse", "ticket": "printer on fire" }),
None,
);
let asked = body["messages"][0]["content"].as_str().unwrap_or_default();
assert!(
!asked.contains("be terse"),
"the instruction leaked into the question: {asked}"
);
assert!(
asked.contains("printer on fire"),
"the actual content went missing: {asked}"
);
}
#[test]
fn a_multimodal_message_is_passed_through_verbatim() {
let parts = json!([{
"role": "user",
"content": [
{ "type": "text", "text": "what is in this image?" },
{ "type": "image", "source": {
"type": "base64", "media_type": "image/png", "data": "iVBORw0KGgo=" } }
]
}]);
let body = driver().body(
&ModelId::new("anthropic", "claude-x"),
&json!({ "messages": parts }),
None,
);
assert_eq!(
body["messages"], parts,
"content blocks must survive untouched: {body}"
);
}
#[test]
fn a_prompt_without_a_system_sends_no_system() {
let body = driver().body(&ModelId::new("anthropic", "claude-x"), &json!("hi"), None);
assert!(
body.get("system").is_none(),
"an unset instruction must not become an empty one: {body}"
);
}
}