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