laminar_db/ai/backends/
anthropic.rs1use 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";
19const DEFAULT_MAX_TOKENS: u32 = 1024;
21
22pub struct AnthropicProvider {
24 client: reqwest::Client,
25 base_url: String,
26 api_key: String,
27 max_concurrency: usize,
28}
29
30impl AnthropicProvider {
31 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#[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
162fn 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}