Skip to main content

llm/providers/anthropic/
provider.rs

1use super::mappers::{map_messages, map_tools};
2use super::streaming::process_anthropic_stream;
3use super::types::{Request, Thinking};
4use crate::provider::{
5    LlmResponseStream, ProviderFactory, StreamingModelProvider, error_stream, get_context_window, validate_reasoning,
6};
7use crate::providers::http::{anthropic_code, rejected};
8use crate::{Context, LlmError, ProviderAuthMode, ProviderConnectionConfig, ProviderError, ReasoningEffort, Result};
9use async_stream;
10use eventsource_stream::Eventsource;
11use futures::StreamExt;
12use reqwest::header::{CONTENT_TYPE, HeaderMap, HeaderValue};
13use reqwest::{Client, header};
14use std::env;
15use std::future::ready;
16use std::time::Duration;
17use tracing::debug;
18
19/// Tokens requested when no `max_tokens` is set via `ModelSettings`. Anthropic requires
20/// `max_tokens` on every request, so a value is always sent.
21const DEFAULT_MAX_TOKENS: u32 = 16_384;
22
23#[derive(Clone)]
24pub struct AnthropicProvider {
25    client: Client,
26    model: String,
27    base_url: Option<String>,
28    auth_mode: ProviderAuthMode,
29    api_key: Option<String>,
30}
31
32impl AnthropicProvider {
33    pub fn new(api_key: Option<String>) -> Result<Self> {
34        let client = build_client()?;
35
36        Ok(Self {
37            client,
38            model: "claude-sonnet-4-5-20250929".to_string(),
39            base_url: Some("https://api.anthropic.com".to_string()),
40            auth_mode: ProviderAuthMode::Default,
41            api_key,
42        })
43    }
44
45    pub fn with_model(mut self, model: &str) -> Self {
46        self.model = model.to_string();
47        self
48    }
49
50    pub fn with_base_url(mut self, base_url: &str) -> Self {
51        self.base_url = Some(base_url.to_string());
52        self
53    }
54
55    pub fn with_connection(mut self, connection: ProviderConnectionConfig) -> Self {
56        if let Some(base_url) = connection.base_url {
57            self.base_url = Some(base_url);
58        }
59        self.auth_mode = connection.auth_mode;
60        self
61    }
62
63    pub(crate) fn build_request(&self, context: &Context) -> Result<Request> {
64        let (system_prompt, messages) = map_messages(context.messages())?;
65        let tools = if context.tools().is_empty() { None } else { Some(map_tools(context.tools())?) };
66
67        let settings = context.model_settings();
68
69        let mut request = Request::new(self.model.clone(), messages)
70            .with_max_tokens(settings.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS))
71            .with_stream(true)
72            .with_auto_caching();
73
74        if let Some(temp) = settings.temperature {
75            request = request.with_temperature(temp);
76        }
77
78        if let Some(top_p) = settings.top_p {
79            request = request.with_top_p(top_p);
80        }
81
82        if let Some(system) = system_prompt {
83            request = request.with_system_cached(system);
84        }
85
86        if let Some(tools) = tools {
87            request = request.with_tools(tools);
88        }
89
90        if context.reasoning_effort() == ReasoningEffort::Disabled {
91            request = request.with_thinking(Thinking::Disabled);
92        } else if let Some(budget_tokens) = effort_to_budget_tokens(context.reasoning_effort()) {
93            request = request.with_thinking(Thinking::new(budget_tokens));
94            // Anthropic requires temperature and top_p to be unset when thinking is enabled
95            request.temperature = None;
96            request.top_p = None;
97            // max_tokens must be > budget_tokens
98            if request.max_tokens <= budget_tokens {
99                request.max_tokens = budget_tokens + 1024;
100            }
101        }
102
103        debug!("Built Anthropic request for model: {}", request.model);
104        Ok(request)
105    }
106
107    fn get_api_key(&self) -> Result<String> {
108        if let Some(key) = &self.api_key {
109            return Ok(key.clone());
110        }
111
112        if let Ok(api_key) = env::var("ANTHROPIC_API_KEY") {
113            return Ok(api_key);
114        }
115
116        Err(LlmError::MissingApiKey(
117            "No Anthropic credentials found. Set ANTHROPIC_API_KEY environment variable.".to_string(),
118        ))
119    }
120
121    fn build_headers(&self) -> Result<HeaderMap> {
122        let mut headers = HeaderMap::new();
123        headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01"));
124        headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
125        if self.auth_mode != ProviderAuthMode::None {
126            let api_key = self.get_api_key()?;
127            headers.insert("x-api-key", HeaderValue::from_str(&api_key)?);
128        }
129        Ok(headers)
130    }
131
132    async fn send_request(
133        &self,
134        request: Request,
135        headers: header::HeaderMap,
136    ) -> Result<impl futures::Stream<Item = Result<String>>> {
137        let base_url = self.base_url.as_deref().unwrap_or("https://api.anthropic.com");
138        let url = format!("{base_url}/v1/messages");
139
140        debug!("Sending request to Anthropic API: {url}");
141        debug!(
142            "Anthropic request body: {}",
143            serde_json::to_string(&request).unwrap_or_else(|_| "<failed to serialize>".to_string())
144        );
145
146        debug!("Anthropic request headers: {}", format_headers(&headers));
147        let response = self.client.post(&url).headers(headers).json(&request).send().await?;
148
149        if !response.status().is_success() {
150            return Err(rejected(response, anthropic_code).await.into());
151        }
152
153        let event_stream = response.bytes_stream().eventsource();
154        let processed_stream = event_stream.filter_map(|result| {
155            std::future::ready(match result {
156                Ok(event) => {
157                    let data = event.data;
158                    if data == "[DONE]" { None } else { Some(Ok(data)) }
159                }
160                Err(e) => Some(Err(ProviderError::stream_interrupted(e.to_string()).into())),
161            })
162        });
163
164        Ok(processed_stream)
165    }
166}
167
168impl ProviderFactory for AnthropicProvider {
169    fn from_env() -> impl Future<Output = Result<Self>> + Send {
170        ready(Self::new(None))
171    }
172
173    fn from_env_with_connection(connection: ProviderConnectionConfig) -> impl Future<Output = Result<Self>> + Send {
174        ready(Self::new(None).map(|provider| provider.with_connection(connection)))
175    }
176
177    fn with_model(self, model: &str) -> Self {
178        self.with_model(model)
179    }
180}
181
182impl StreamingModelProvider for AnthropicProvider {
183    fn model(&self) -> Option<crate::LlmModel> {
184        format!("anthropic:{}", self.model).parse().ok()
185    }
186
187    fn context_window(&self) -> Option<u32> {
188        get_context_window("anthropic", &self.model)
189    }
190
191    fn stream_response<'a>(&self, context: &Context) -> LlmResponseStream {
192        if let Err(error) = validate_reasoning(context, self.model().as_ref()) {
193            return error_stream(error);
194        }
195        let provider = self.clone();
196        let context = context.clone();
197
198        Box::pin(async_stream::stream! {
199            let headers = match provider.build_headers() {
200                Ok(result) => result,
201                Err(e) => {
202                    yield Err(e);
203                    return;
204                }
205            };
206
207            let request = match provider.build_request(&context) {
208                Ok(req) => req,
209                Err(e) => {
210                    yield Err(e);
211                    return;
212                }
213            };
214
215            let stream = match provider.send_request(request, headers).await {
216                Ok(stream) => stream,
217                Err(e) => {
218                    yield Err(e);
219                    return;
220                }
221            };
222
223            let mut anthropic_stream = Box::pin(process_anthropic_stream(stream));
224            while let Some(result) = anthropic_stream.next().await {
225                yield result;
226            }
227        })
228    }
229
230    fn display_name(&self) -> String {
231        format!("Anthropic ({})", self.model)
232    }
233}
234
235fn build_client() -> Result<Client> {
236    Client::builder().timeout(Duration::from_mins(1)).build().map_err(|e| LlmError::HttpClientCreation(e.to_string()))
237}
238
239fn effort_to_budget_tokens(effort: ReasoningEffort) -> Option<u32> {
240    Some(match effort {
241        ReasoningEffort::Default | ReasoningEffort::Disabled => return None,
242        // 1024 is the Anthropic API's minimum thinking budget.
243        ReasoningEffort::Minimal | ReasoningEffort::Low => 1024,
244        ReasoningEffort::Medium => 4096,
245        ReasoningEffort::High | ReasoningEffort::Xhigh => 10240,
246        ReasoningEffort::Max => 32768,
247    })
248}
249
250fn should_redact_header(name: &str) -> bool {
251    let lower = name.to_ascii_lowercase();
252    lower == "authorization" || lower == "x-api-key" || lower.contains("secret") || lower.contains("token")
253}
254
255fn format_headers(headers: &header::HeaderMap) -> String {
256    let mut parts = Vec::new();
257    for (name, value) in headers {
258        let name_str = name.as_str();
259        let value_str = if should_redact_header(name_str) {
260            "<redacted>".to_string()
261        } else {
262            value.to_str().unwrap_or("<non-utf8>").to_string()
263        };
264        parts.push(format!("{name_str}={value_str}"));
265    }
266    parts.join(", ")
267}
268
269#[cfg(test)]
270mod tests {
271    use super::*;
272    use crate::ChatMessage;
273
274    use crate::ToolDefinition;
275    use crate::providers::anthropic::types::{SystemContent, SystemContentBlock};
276
277    use reqwest::header::AUTHORIZATION;
278
279    fn create_test_provider() -> AnthropicProvider {
280        AnthropicProvider::new(Some("test-api-key".to_string())).unwrap().with_model("claude-sonnet-4-5-20250929")
281    }
282
283    #[tokio::test]
284    async fn default_and_disabled_thinking_preserve_sampling_and_token_limit() {
285        use crate::providers::test_capture_server::CaptureServer;
286        let mut server =
287            CaptureServer::start_with_response(include_str!("../../../tests/fixtures/anthropic/01_minimal.sse")).await;
288        let model = crate::LlmModel::all()
289            .iter()
290            .find(|model| {
291                model.provider_enum() == crate::catalog::Provider::Anthropic && model.supports_reasoning_off()
292            })
293            .unwrap();
294        let provider = create_test_provider().with_model(&model.model_id()).with_base_url(&server.base_url);
295        for (effort, temperature) in [(ReasoningEffort::Default, 0.0), (ReasoningEffort::Disabled, 0.5)] {
296            let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
297            context.set_reasoning_effort(effort);
298            context.set_model_settings(crate::ModelSettings {
299                temperature: Some(temperature),
300                top_p: Some(0.5),
301                max_tokens: Some(128),
302            });
303            let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
304            assert!(responses.iter().all(Result::is_ok), "{responses:?}");
305            let body = server.captured().await.body;
306            assert_eq!(body["max_tokens"], 128);
307            assert_eq!(body["top_p"], 0.5);
308            assert_eq!(body["temperature"], serde_json::json!(temperature));
309            if effort == ReasoningEffort::Disabled {
310                assert_eq!(body["thinking"], serde_json::json!({"type": "disabled"}));
311            } else {
312                assert!(body.get("thinking").is_none());
313            }
314        }
315    }
316
317    #[test]
318    fn test_provider_creation() {
319        let provider = AnthropicProvider::new(Some("test-api-key".to_string()));
320        assert!(provider.is_ok());
321    }
322
323    #[test]
324    fn build_headers_uses_api_key() {
325        let provider = AnthropicProvider::new(Some("test-api-key".to_string())).unwrap();
326        let headers = provider.build_headers().expect("headers");
327        assert_eq!(headers.get("x-api-key").and_then(|value| value.to_str().ok()), Some("test-api-key"));
328        assert!(headers.get(AUTHORIZATION).is_none());
329        assert!(headers.get("anthropic-beta").is_none());
330    }
331
332    #[test]
333    fn build_headers_skips_api_key_when_auth_is_none() {
334        let provider = AnthropicProvider::new(None)
335            .unwrap()
336            .with_connection(ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() });
337        let headers = provider.build_headers().expect("headers");
338        assert!(headers.get("x-api-key").is_none());
339        assert_eq!(headers.get("anthropic-version").and_then(|value| value.to_str().ok()), Some("2023-06-01"));
340    }
341
342    #[test]
343    fn test_build_request_simple() {
344        let provider = create_test_provider();
345
346        let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
347
348        let request = provider.build_request(&context).unwrap();
349        assert_eq!(request.model, "claude-sonnet-4-5-20250929");
350        assert_eq!(request.max_tokens, DEFAULT_MAX_TOKENS);
351        assert_eq!(request.messages.len(), 1);
352        assert!(request.tools.is_none());
353        assert!(request.stream);
354    }
355
356    #[test]
357    fn test_build_request_with_system_and_tools() {
358        let provider = create_test_provider();
359
360        let context = Context::new(
361            vec![ChatMessage::system("You are helpful"), ChatMessage::user("Hello")],
362            vec![ToolDefinition::new(
363                "search",
364                "Search for information",
365                serde_json::from_str(r#"{"type": "object", "properties": {"query": {"type": "string"}}}"#).unwrap(),
366            )],
367        );
368
369        let request = provider.build_request(&context).unwrap();
370        if let Some(system) = &request.system {
371            match system {
372                SystemContent::Blocks(blocks) => {
373                    assert_eq!(blocks.len(), 1);
374                    let SystemContentBlock::Text { text, .. } = &blocks[0];
375                    assert_eq!(text, "You are helpful");
376                }
377                SystemContent::Text(_) => panic!("Expected blocks system content"),
378            }
379        } else {
380            panic!("Expected system prompt");
381        }
382        assert_eq!(request.messages.len(), 1);
383        assert!(request.tools.is_some());
384        assert_eq!(request.tools.unwrap().len(), 1);
385    }
386
387    #[test]
388    fn test_build_request_with_caching() {
389        let provider = AnthropicProvider::new(Some("test-api-key".to_string())).unwrap(); // Caching is enabled by default
390
391        let context = Context::new(
392            vec![ChatMessage::system("Hello"), ChatMessage::user("Hello")],
393            vec![ToolDefinition::new(
394                "search",
395                "Search for information",
396                serde_json::from_str(r#"{"type": "object", "properties": {"query": {"type": "string"}}}"#).unwrap(),
397            )],
398        );
399
400        let request = provider.build_request(&context).unwrap();
401
402        // With caching enabled, system prompt should be cached
403        if let Some(system) = &request.system {
404            match system {
405                SystemContent::Blocks(blocks) => {
406                    assert_eq!(blocks.len(), 1);
407                    let SystemContentBlock::Text { text, cache_control } = &blocks[0];
408                    assert_eq!(text, "Hello");
409                    assert!(cache_control.is_some());
410                }
411                SystemContent::Text(_) => panic!("Expected blocks system content for caching"),
412            }
413        } else {
414            panic!("Expected system prompt");
415        }
416
417        assert!(request.tools.is_some());
418
419        // Top-level cache_control enables automatic caching
420        assert!(request.cache_control.is_some());
421    }
422
423    #[test]
424    fn test_build_request_with_reasoning_effort() {
425        let provider = create_test_provider();
426
427        let mut context = Context::new(vec![ChatMessage::user("Think hard")], vec![]);
428        context.set_reasoning_effort(crate::ReasoningEffort::High);
429
430        let request = provider.build_request(&context).unwrap();
431        let Thinking::Enabled { budget_tokens } = request.thinking.unwrap() else {
432            panic!("expected enabled thinking")
433        };
434        assert_eq!(budget_tokens, 10240);
435        assert!(request.temperature.is_none());
436        assert!(request.max_tokens > budget_tokens);
437    }
438
439    #[test]
440    fn test_build_request_thinking_clears_sampling() {
441        let provider = create_test_provider();
442        let mut context = Context::new(vec![ChatMessage::user("Think")], vec![]);
443        context.set_model_settings(crate::ModelSettings { temperature: Some(0.2), top_p: Some(0.9), max_tokens: None });
444        context.set_reasoning_effort(crate::ReasoningEffort::High);
445
446        let request = provider.build_request(&context).unwrap();
447        assert!(request.temperature.is_none());
448        assert!(request.top_p.is_none());
449    }
450
451    #[test]
452    fn test_build_request_thinking_bumps_max_tokens_if_needed() {
453        let provider = AnthropicProvider::new(Some("test-api-key".to_string())).unwrap();
454
455        let mut context = Context::new(vec![ChatMessage::user("Hi")], vec![]);
456        context.set_model_settings(crate::ModelSettings { max_tokens: Some(500), ..Default::default() });
457        context.set_reasoning_effort(crate::ReasoningEffort::Low);
458
459        let request = provider.build_request(&context).unwrap();
460        let Thinking::Enabled { budget_tokens } = request.thinking.unwrap() else {
461            panic!("expected enabled thinking")
462        };
463        assert!(request.max_tokens > budget_tokens);
464    }
465
466    #[test]
467    fn test_anthropic_provider_display_name() {
468        let provider = create_test_provider();
469        assert_eq!(provider.display_name(), "Anthropic (claude-sonnet-4-5-20250929)");
470    }
471
472    #[test]
473    fn test_anthropic_provider_display_name_default() {
474        let provider = AnthropicProvider::new(Some("test-api-key".to_string())).unwrap();
475        assert_eq!(provider.display_name(), "Anthropic (claude-sonnet-4-5-20250929)");
476    }
477
478    #[test]
479    fn format_headers_redacts_x_api_key() {
480        let mut headers = HeaderMap::new();
481        headers.insert("x-api-key", HeaderValue::from_static("sk-secret-123"));
482        headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
483
484        let formatted = format_headers(&headers);
485        assert!(formatted.contains("x-api-key=<redacted>"));
486        assert!(formatted.contains("content-type=application/json"));
487        assert!(!formatted.contains("sk-secret-123"));
488    }
489
490    #[test]
491    fn format_headers_redacts_authorization() {
492        let mut headers = HeaderMap::new();
493        headers.insert(AUTHORIZATION, HeaderValue::from_static("Bearer token123"));
494
495        let formatted = format_headers(&headers);
496        assert!(formatted.contains("authorization=<redacted>"));
497        assert!(!formatted.contains("token123"));
498    }
499
500    #[test]
501    fn format_headers_redacts_secret_and_token_headers() {
502        let mut headers = HeaderMap::new();
503        headers.insert("x-client-secret", HeaderValue::from_static("mysecret"));
504        headers.insert("x-auth-token", HeaderValue::from_static("mytoken"));
505        headers.insert("accept", HeaderValue::from_static("text/plain"));
506
507        let formatted = format_headers(&headers);
508        assert!(formatted.contains("x-client-secret=<redacted>"));
509        assert!(formatted.contains("x-auth-token=<redacted>"));
510        assert!(formatted.contains("accept=text/plain"));
511        assert!(!formatted.contains("mysecret"));
512        assert!(!formatted.contains("mytoken"));
513    }
514}