pub mod build_params;
pub mod cache_control;
pub mod cost;
pub mod json_parse;
pub mod mapper;
pub mod models;
pub mod retry;
pub mod sse;
use crate::error::AiError;
use crate::event_stream::{
create_assistant_message_event_stream, AssistantMessageEventStream,
AssistantMessageEventStreamProducer,
};
use crate::model::{Model, StreamingProtocolCompat};
use crate::provider::{CacheRetention, Provider, SimpleStreamOptions};
use crate::types::{AssistantMessageEvent, Content, Context, DoneReason, StopReason, Usage};
use async_trait::async_trait;
use std::sync::Arc;
use crate::providers::anthropic::build_params::build_params;
use crate::providers::anthropic::cost::calculate_cost;
use crate::providers::anthropic::mapper::{emit_terminal_error, run_mapper, MapperState};
use crate::providers::anthropic::models::anthropic_models;
use crate::providers::anthropic::retry::retry_provider_request;
use crate::providers::anthropic::sse::SseEventStream;
const ANTHROPIC_VERSION: &str = "2023-06-01";
pub struct AnthropicProvider {
api_key: Option<String>,
http: reqwest::Client,
models: Vec<Model>,
}
impl AnthropicProvider {
pub fn new(api_key: Option<String>, http: reqwest::Client) -> Self {
Self {
api_key,
http,
models: anthropic_models(),
}
}
pub fn with_models(api_key: Option<String>, http: reqwest::Client, models: Vec<Model>) -> Self {
Self {
api_key,
http,
models,
}
}
pub fn from_env() -> Self {
let http = reqwest::Client::new();
Self::new(None, http)
}
pub fn api_key(&self) -> Option<&str> {
self.api_key.as_deref()
}
}
#[async_trait]
impl Provider for AnthropicProvider {
fn id(&self) -> &str {
"anthropic"
}
fn models(&self) -> &[Model] {
&self.models
}
async fn stream_simple(
&self,
model: &Model,
ctx: &Context,
opts: &SimpleStreamOptions,
) -> AssistantMessageEventStream {
let (mut prod, stream) = create_assistant_message_event_stream();
let http = self.http.clone();
let provider_key = self.api_key.clone();
let model = model.clone();
let ctx = Arc::new(ctx.clone());
let opts = opts.clone();
tokio::spawn(async move {
run_anthropic_stream(&mut prod, http, provider_key, &model, &ctx, &opts).await;
});
stream
}
}
async fn run_anthropic_stream(
prod: &mut AssistantMessageEventStreamProducer,
http: reqwest::Client,
provider_key: Option<String>,
model: &Model,
ctx: &Context,
opts: &SimpleStreamOptions,
) {
let mut state = MapperState::new(
model.api.clone(),
model.provider.clone(),
model.id.clone(),
now_ms(),
);
let resolved_key = resolve_api_key(&provider_key, opts);
let header_owned_auth = has_header_auth(&opts.headers) || has_header_auth(&model.headers);
let api_key_for_header = match (resolved_key, header_owned_auth) {
(Some(k), _) => Some(k),
(None, true) => None, (None, false) => {
let msg = format!("No API key for provider: {}", model.provider);
emit_terminal_error(prod, &mut state, msg, false);
return;
}
};
let built = build_params(model, ctx, false, opts);
let mut body = match serde_json::to_value(&built.request) {
Ok(v) => v,
Err(e) => {
emit_terminal_error(
prod,
&mut state,
format!("failed to serialize request body: {e}"),
false,
);
return;
}
};
let non_stream = std::env::var("RPI_ANTHROPIC_NON_STREAM").ok().as_deref() == Some("1");
if non_stream {
body["stream"] = serde_json::Value::Bool(false);
strip_cache_control(&mut body);
simplify_non_stream_request(&mut body);
}
let headers = assemble_headers(
model,
opts,
built.beta_header.as_deref(),
api_key_for_header.as_deref(),
);
let url = format!("{}/v1/messages", model.base_url.trim_end_matches('/'));
let timeout = opts.timeout;
let signal = opts.signal.clone();
let response_result = retry_provider_request(
move || {
let http = http.clone();
let body = body.clone();
let url = url.clone();
let headers = headers.clone();
let signal = signal.clone();
async move {
let mut req = http.post(&url);
if let Some(t) = timeout {
req = req.timeout(t);
}
for (k, v) in &headers {
req = req.header(k.as_str(), v.as_str());
}
let send_fut = req.json(&body).send();
let resp = tokio::select! {
r = send_fut => r.map_err(|e| AiError::Http {
status: None,
message: format!("http transport error: {e}"),
})?,
_ = signal.cancelled() => return Err(AiError::Abort {
message: "Request aborted".to_string(),
}),
};
let status = resp.status();
if !status.is_success() {
let code = status.as_u16();
let text = resp.text().await.unwrap_or_default();
return Err(AiError::Http {
status: Some(code),
message: text,
});
}
Ok(resp)
}
},
opts.max_retries,
opts.max_retry_delay,
&opts.signal,
)
.await;
let response = match response_result {
Ok(r) => r,
Err(err) => {
let aborted = err.is_abort();
eprintln!("anthropic request failed: {err}");
emit_terminal_error(prod, &mut state, err.to_string(), aborted);
return;
}
};
if non_stream {
let response_body = match response.text().await {
Ok(body) => body,
Err(error) => {
emit_terminal_error(
prod,
&mut state,
format!("failed to read non-stream response: {error}"),
false,
);
return;
}
};
let value: serde_json::Value = match serde_json::from_str(&response_body) {
Ok(value) => value,
Err(error) => {
emit_terminal_error(
prod,
&mut state,
format!("failed to parse non-stream response: {error}"),
false,
);
return;
}
};
let text = value
.get("content")
.and_then(serde_json::Value::as_array)
.map(|blocks| {
blocks
.iter()
.filter(|block| {
block.get("type").and_then(serde_json::Value::as_str) == Some("text")
})
.filter_map(|block| block.get("text").and_then(serde_json::Value::as_str))
.collect::<String>()
})
.unwrap_or_default();
if text.trim().is_empty() {
emit_terminal_error(
prod,
&mut state,
"non-stream response contained no text content".into(),
false,
);
return;
}
state.output.content.push(Content::text(text.clone()));
state.output.response_id = value
.get("id")
.and_then(serde_json::Value::as_str)
.map(str::to_string);
state.output.response_model = value
.get("model")
.and_then(serde_json::Value::as_str)
.map(str::to_string);
if let Some(usage) = value.get("usage") {
state.output.usage.input = usage
.get("input_tokens")
.and_then(serde_json::Value::as_i64)
.unwrap_or(0);
state.output.usage.output = usage
.get("output_tokens")
.and_then(serde_json::Value::as_i64)
.unwrap_or(0);
state.output.usage.total_tokens = state.output.usage.input + state.output.usage.output;
}
state.output.stop_reason = StopReason::Stop;
let partial = std::sync::Arc::new(state.output.clone());
prod.push(AssistantMessageEvent::Start {
partial: partial.clone(),
});
prod.push(AssistantMessageEvent::TextStart {
content_index: 0,
partial: partial.clone(),
});
prod.push(AssistantMessageEvent::TextDelta {
content_index: 0,
delta: text.clone(),
partial: partial.clone(),
});
prod.push(AssistantMessageEvent::TextEnd {
content_index: 0,
content: text,
partial,
});
prod.push(AssistantMessageEvent::Done {
reason: DoneReason::Stop,
message: state.output,
});
return;
}
let mut sse = SseEventStream::new(response, opts.signal.clone());
let cost_model = model.cost.clone();
let cost_fn = move |usage: &Usage| calculate_cost(&cost_model, usage);
run_mapper(&mut sse, prod, &mut state, cost_fn).await;
}
fn strip_cache_control(value: &mut serde_json::Value) {
match value {
serde_json::Value::Object(map) => {
map.remove("cache_control");
for child in map.values_mut() {
strip_cache_control(child);
}
}
serde_json::Value::Array(items) => {
for child in items {
strip_cache_control(child);
}
}
_ => {}
}
}
fn simplify_non_stream_request(value: &mut serde_json::Value) {
let text_from_blocks = |blocks: &serde_json::Value| {
blocks
.as_array()
.map(|items| {
items
.iter()
.filter_map(|item| item.get("text").and_then(serde_json::Value::as_str))
.collect::<String>()
})
.unwrap_or_default()
};
if let Some(messages) = value
.get_mut("messages")
.and_then(serde_json::Value::as_array_mut)
{
for message in messages {
if let Some(content) = message.get_mut("content") {
if content.is_array() {
*content = serde_json::Value::String(text_from_blocks(content));
}
}
}
}
if let Some(system) = value.get_mut("system") {
if system.is_array() {
*system = serde_json::Value::String(text_from_blocks(system));
}
}
}
fn resolve_api_key(provider_key: &Option<String>, opts: &SimpleStreamOptions) -> Option<String> {
opts.api_key
.clone()
.or_else(|| provider_key.clone())
.or_else(|| {
std::env::var("ANTHROPIC_API_KEY")
.ok()
.filter(|s| !s.is_empty())
})
}
fn has_header_auth(headers: &Option<std::collections::BTreeMap<String, String>>) -> bool {
let Some(h) = headers else {
return false;
};
const NAMES: &[&str] = &["authorization", "x-api-key", "cf-aig-authorization"];
h.keys()
.any(|k| NAMES.contains(&k.to_ascii_lowercase().as_str()))
}
fn assemble_headers(
model: &Model,
opts: &SimpleStreamOptions,
beta_header: Option<&str>,
api_key: Option<&str>,
) -> Vec<(String, String)> {
let mut headers: Vec<(String, String)> = Vec::new();
headers.push(("accept".into(), "application/json".into()));
if std::env::var("RPI_ANTHROPIC_NON_STREAM").ok().as_deref() != Some("1") {
headers.push((
"anthropic-dangerous-direct-browser-access".into(),
"true".into(),
));
}
headers.push(("anthropic-version".into(), ANTHROPIC_VERSION.into()));
if let Some(beta) = beta_header {
headers.push(("anthropic-beta".into(), beta.into()));
}
if !matches!(opts.cache_retention, CacheRetention::None) {
if let Some(session_id) = &opts.session_id {
let send = model
.compat
.as_ref()
.and_then(|c| match c {
StreamingProtocolCompat::AnthropicMessages(a) => Some(a),
_ => None,
})
.map(|a| a.send_session_affinity_headers.unwrap_or(false))
.unwrap_or(false);
if send {
headers.push(("x-session-affinity".into(), session_id.clone()));
}
}
}
if let Some(model_headers) = &model.headers {
for (k, v) in model_headers {
merge_header(&mut headers, k.clone(), v.clone());
}
}
if let Some(opt_headers) = &opts.headers {
for (k, v) in opt_headers {
merge_header(&mut headers, k.clone(), v.clone());
}
}
if let Some(key) = api_key {
headers.retain(|(k, _)| k.to_ascii_lowercase() != "x-api-key");
headers.push(("x-api-key".into(), key.into()));
}
headers
}
fn merge_header(headers: &mut Vec<(String, String)>, name: String, value: String) {
let lname = name.to_ascii_lowercase();
if let Some(slot) = headers
.iter_mut()
.find(|(k, _)| k.to_ascii_lowercase() == lname)
{
slot.1 = value;
} else {
headers.push((name, value));
}
}
fn now_ms() -> i64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as i64)
.unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::providers::anthropic::models::claude_haiku_4_5;
use crate::types::{Api, InputModality, ModelCost};
fn test_model() -> Model {
let mut m = Model::new(
"claude-haiku-4-5",
"Claude Haiku 4.5",
Api::AnthropicMessages,
"anthropic",
"https://api.anthropic.com",
);
m.input = vec![InputModality::Text];
m.context_window = 200_000;
m.max_tokens = 8_192;
m.cost = ModelCost::default();
m.compat = Some(crate::model::StreamingProtocolCompat::AnthropicMessages(
crate::model::AnthropicMessagesCompat::default(),
));
m
}
#[test]
fn assemble_headers_includes_required_defaults() {
let model = test_model();
let opts = SimpleStreamOptions::default();
let headers = assemble_headers(&model, &opts, None, Some("sk-test"));
let names: Vec<&str> = headers.iter().map(|(k, _)| k.as_str()).collect();
assert!(names.contains(&"accept"));
assert!(names.contains(&"anthropic-version"));
assert!(names.contains(&"anthropic-dangerous-direct-browser-access"));
assert!(names.contains(&"x-api-key"));
assert!(!names.contains(&"anthropic-beta"));
let version = headers
.iter()
.find(|(k, _)| k == "anthropic-version")
.map(|(_, v)| v.as_str())
.unwrap();
assert_eq!(version, ANTHROPIC_VERSION);
}
#[test]
fn assemble_headers_applies_beta_when_present() {
let model = test_model();
let opts = SimpleStreamOptions::default();
let headers = assemble_headers(
&model,
&opts,
Some("fine-grained-tool-streaming-2025-05-14"),
Some("k"),
);
let beta = headers
.iter()
.find(|(k, _)| k == "anthropic-beta")
.map(|(_, v)| v.as_str())
.unwrap();
assert_eq!(beta, "fine-grained-tool-streaming-2025-05-14");
}
#[test]
fn assemble_headers_x_api_key_not_overridable_by_opts() {
let model = test_model();
let mut opts = SimpleStreamOptions::default();
let mut extra = std::collections::BTreeMap::new();
extra.insert("x-api-key".to_string(), "SK-ATTACK".to_string());
opts.headers = Some(extra);
let headers = assemble_headers(&model, &opts, None, Some("sk-resolved"));
let key = headers
.iter()
.find(|(k, _)| k == "x-api-key")
.map(|(_, v)| v.as_str())
.unwrap();
assert_eq!(key, "sk-resolved");
}
#[test]
fn assemble_headers_session_affinity_gated_on_compat_and_retention() {
let mut model = claude_haiku_4_5();
let _ = &mut model; let mut opts = SimpleStreamOptions::default();
opts.session_id = Some("sess-1".into());
opts.cache_retention = CacheRetention::Short;
let headers = assemble_headers(&model, &opts, None, Some("k"));
assert!(
!headers.iter().any(|(k, _)| k == "x-session-affinity"),
"default catalog should not send session affinity"
);
let mut model = test_model();
let mut compat = crate::model::AnthropicMessagesCompat::default();
compat.send_session_affinity_headers = Some(true);
model.compat = Some(StreamingProtocolCompat::AnthropicMessages(compat));
let headers = assemble_headers(&model, &opts, None, Some("k"));
assert_eq!(
headers
.iter()
.find(|(k, _)| k == "x-session-affinity")
.map(|(_, v)| v.as_str()),
Some("sess-1")
);
let mut opts2 = opts.clone();
opts2.cache_retention = CacheRetention::None;
let headers = assemble_headers(&model, &opts2, None, Some("k"));
assert!(
!headers.iter().any(|(k, _)| k == "x-session-affinity"),
"None retention suppresses session affinity"
);
}
#[test]
fn assemble_headers_model_and_opts_merge_last_wins() {
let mut model = test_model();
let mut mh = std::collections::BTreeMap::new();
mh.insert("x-custom".to_string(), "from-model".to_string());
model.headers = Some(mh);
let mut opts = SimpleStreamOptions::default();
let mut oh = std::collections::BTreeMap::new();
oh.insert("x-custom".to_string(), "from-opts".to_string());
oh.insert("x-extra".to_string(), "extra".to_string());
opts.headers = Some(oh);
let headers = assemble_headers(&model, &opts, None, None);
assert_eq!(
headers
.iter()
.find(|(k, _)| k == "x-custom")
.map(|(_, v)| v.as_str()),
Some("from-opts"),
"opts.headers wins over model.headers"
);
assert_eq!(
headers
.iter()
.find(|(k, _)| k == "x-extra")
.map(|(_, v)| v.as_str()),
Some("extra")
);
}
#[test]
fn has_header_auth_detects_owned_auth_headers() {
let mut h = std::collections::BTreeMap::new();
h.insert("X-API-Key".to_string(), "k".to_string());
assert!(has_header_auth(&Some(h.clone())));
let mut h2 = std::collections::BTreeMap::new();
h2.insert("authorization".to_string(), "Bearer t".to_string());
assert!(has_header_auth(&Some(h2)));
assert!(!has_header_auth(&None));
let mut h3 = std::collections::BTreeMap::new();
h3.insert("x-other".to_string(), "v".to_string());
assert!(!has_header_auth(&Some(h3)));
}
#[test]
fn resolve_api_key_prefers_opts_then_provider_then_env() {
let opts = SimpleStreamOptions::default().with_api_key("opts-key");
assert_eq!(
resolve_api_key(&Some("provider-key".into()), &opts),
Some("opts-key".into())
);
let opts = SimpleStreamOptions::default();
assert_eq!(
resolve_api_key(&Some("provider-key".into()), &opts),
Some("provider-key".into())
);
}
#[test]
fn provider_exposes_catalog_models() {
let p = AnthropicProvider::from_env();
assert_eq!(p.id(), "anthropic");
assert!(!p.models().is_empty());
assert!(p.models().iter().any(|m| m.id == "claude-haiku-4-5"));
assert!(p.api_key().is_none(), "from_env does not pre-read the key");
}
#[test]
fn header_auth_on_model_headers_counts_as_owned() {
let mut model = test_model();
let mut mh = std::collections::BTreeMap::new();
mh.insert("authorization".to_string(), "Bearer tok".to_string());
model.headers = Some(mh);
let owned = has_header_auth(&SimpleStreamOptions::default().headers)
|| has_header_auth(&model.headers);
assert!(
owned,
"a Bearer header on model.headers must satisfy header-owned auth"
);
let mut model2 = test_model();
let mut mh2 = std::collections::BTreeMap::new();
mh2.insert("x-custom".to_string(), "v".to_string());
model2.headers = Some(mh2);
let owned2 = has_header_auth(&SimpleStreamOptions::default().headers)
|| has_header_auth(&model2.headers);
assert!(
!owned2,
"non-auth headers on model.headers must not satisfy header-owned auth"
);
}
}