Skip to main content

laminar_db/ai/backends/
anthropic.rs

1//! Anthropic Messages API provider. One bounded-concurrent call per input row.
2//! `ai_embed` is unsupported — use the OpenAI-compatible provider for embeddings.
3
4use std::time::Duration;
5
6use async_trait::async_trait;
7use futures::stream::{StreamExt, TryStreamExt};
8use serde::{Deserialize, Serialize};
9
10use crate::ai::backends::remote::{add_usage, chat_prompt, post_json};
11use crate::ai::provider::{
12    InferenceOutputs, InferenceProvider, InferenceRequest, InferenceResponse, ProviderError, Usage,
13};
14use crate::ai::registry::Task;
15
16const REQUEST_TIMEOUT_MS: u64 = 60_000;
17const MAX_RETRIES: u32 = 2;
18const ANTHROPIC_VERSION: &str = "2023-06-01";
19/// Messages API requires `max_tokens`.
20const DEFAULT_MAX_TOKENS: u32 = 1024;
21
22/// Anthropic Messages provider.
23pub struct AnthropicProvider {
24    client: reqwest::Client,
25    base_url: String,
26    api_key: String,
27    max_concurrency: usize,
28}
29
30impl AnthropicProvider {
31    /// Build an Anthropic provider.
32    ///
33    /// # Errors
34    ///
35    /// Returns [`ProviderError::Transport`] if the HTTP client cannot be built.
36    pub fn new(
37        base_url: impl Into<String>,
38        api_key: impl Into<String>,
39        max_concurrency: usize,
40    ) -> Result<Self, ProviderError> {
41        let client = reqwest::Client::builder()
42            .timeout(Duration::from_millis(REQUEST_TIMEOUT_MS))
43            .build()
44            .map_err(|e| ProviderError::Transport(e.to_string()))?;
45        Ok(Self {
46            client,
47            base_url: base_url.into().trim_end_matches('/').to_string(),
48            api_key: api_key.into(),
49            max_concurrency: max_concurrency.max(1),
50        })
51    }
52}
53
54#[async_trait]
55impl InferenceProvider for AnthropicProvider {
56    async fn infer_batch(
57        &self,
58        request: InferenceRequest,
59    ) -> Result<InferenceResponse, ProviderError> {
60        if request.task == Task::Embed {
61            return Err(ProviderError::UnsupportedTask(Task::Embed));
62        }
63
64        let url = format!("{}/v1/messages", self.base_url);
65        let bodies: Vec<MessageBody> = request
66            .inputs
67            .iter()
68            .map(|input| {
69                let (system, user) =
70                    chat_prompt(request.task, input, request.params.labels.as_deref());
71                MessageBody {
72                    model: request.model.clone(),
73                    max_tokens: DEFAULT_MAX_TOKENS,
74                    system,
75                    messages: vec![Message::user(user)],
76                }
77            })
78            .collect();
79
80        let url = &url;
81        let calls = bodies.into_iter().map(|body| async move {
82            let builder = self
83                .client
84                .post(url)
85                .header("x-api-key", &self.api_key)
86                .header("anthropic-version", ANTHROPIC_VERSION)
87                .json(&body);
88            let response: MessageResponse =
89                post_json(builder, MAX_RETRIES, REQUEST_TIMEOUT_MS).await?;
90            parse_message(response)
91        });
92
93        let results: Vec<(String, Usage)> = futures::stream::iter(calls)
94            .buffered(self.max_concurrency)
95            .try_collect()
96            .await?;
97
98        let mut texts = Vec::with_capacity(results.len());
99        let mut usage = Usage::ZERO;
100        for (text, call_usage) in results {
101            texts.push(text);
102            usage = add_usage(usage, call_usage);
103        }
104        Ok(InferenceResponse {
105            outputs: InferenceOutputs::Text(texts),
106            usage,
107        })
108    }
109
110    fn name(&self) -> &'static str {
111        "anthropic"
112    }
113}
114
115// --- wire shapes ---
116
117#[derive(Serialize)]
118struct MessageBody {
119    model: String,
120    max_tokens: u32,
121    system: String,
122    messages: Vec<Message>,
123}
124
125#[derive(Serialize)]
126struct Message {
127    role: &'static str,
128    content: String,
129}
130
131impl Message {
132    fn user(content: String) -> Self {
133        Self {
134            role: "user",
135            content,
136        }
137    }
138}
139
140#[derive(Deserialize)]
141struct MessageResponse {
142    content: Vec<ContentBlock>,
143    usage: Option<AnthropicUsage>,
144}
145
146#[derive(Deserialize)]
147struct ContentBlock {
148    #[serde(rename = "type")]
149    kind: String,
150    #[serde(default)]
151    text: String,
152}
153
154#[derive(Deserialize)]
155struct AnthropicUsage {
156    #[serde(default)]
157    input_tokens: u64,
158    #[serde(default)]
159    output_tokens: u64,
160}
161
162/// Extract the first text block and usage from a messages response.
163fn parse_message(response: MessageResponse) -> Result<(String, Usage), ProviderError> {
164    let text = response
165        .content
166        .into_iter()
167        .find(|b| b.kind == "text")
168        .map(|b| b.text)
169        .ok_or_else(|| {
170            ProviderError::BadResponse("messages response had no text block".to_string())
171        })?;
172    let usage = response.usage.map_or(Usage::ZERO, |u| Usage {
173        input_tokens: u.input_tokens,
174        output_tokens: u.output_tokens,
175        cost_micros: 0,
176    });
177    Ok((text, usage))
178}
179
180#[cfg(test)]
181mod tests {
182    use super::*;
183
184    #[test]
185    fn message_body_has_system_and_max_tokens() {
186        let body = MessageBody {
187            model: "claude-x".to_string(),
188            max_tokens: 256,
189            system: "be terse".to_string(),
190            messages: vec![Message::user("hello".to_string())],
191        };
192        let value = serde_json::to_value(&body).unwrap();
193        assert_eq!(value["model"], "claude-x");
194        assert_eq!(value["max_tokens"], 256);
195        assert_eq!(value["system"], "be terse");
196        assert_eq!(value["messages"][0]["role"], "user");
197    }
198
199    #[test]
200    fn parse_message_takes_first_text_block() {
201        let json = r#"{
202            "content": [{"type": "text", "text": "positive"}],
203            "usage": {"input_tokens": 9, "output_tokens": 1}
204        }"#;
205        let response: MessageResponse = serde_json::from_str(json).unwrap();
206        let (text, usage) = parse_message(response).unwrap();
207        assert_eq!(text, "positive");
208        assert_eq!(usage.input_tokens, 9);
209        assert_eq!(usage.output_tokens, 1);
210    }
211
212    #[test]
213    fn parse_message_errors_without_text() {
214        let response: MessageResponse = serde_json::from_str(r#"{"content": []}"#).unwrap();
215        assert!(parse_message(response).is_err());
216    }
217}