1use serde::{Deserialize, Serialize};
6use std::collections::HashMap;
7
8#[derive(Debug, Clone, Serialize, Deserialize)]
10pub struct AnthropicConfig {
11 pub api_key: String,
13 #[serde(default = "default_base_url")]
15 pub base_url: String,
16 #[serde(default = "default_timeout")]
18 pub timeout_secs: u64,
19 #[serde(default = "default_max_retries")]
21 pub max_retries: u32,
22 #[serde(default = "default_rate_limit")]
24 pub rate_limit_per_minute: u32,
25 #[serde(default = "default_api_version")]
27 pub api_version: String,
28}
29
30fn default_base_url() -> String {
31 "https://api.anthropic.com".to_string()
32}
33
34fn default_timeout() -> u64 {
35 60
36}
37
38fn default_max_retries() -> u32 {
39 3
40}
41
42fn default_rate_limit() -> u32 {
43 50
44}
45
46fn default_api_version() -> String {
47 "2023-06-01".to_string()
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
52pub enum ClaudeModel {
53 #[serde(rename = "claude-3-5-sonnet-20241022")]
55 Claude35Sonnet,
56 #[serde(rename = "claude-3-opus-20240229")]
58 Claude3Opus,
59 #[serde(rename = "claude-3-sonnet-20240229")]
61 Claude3Sonnet,
62 #[serde(rename = "claude-3-haiku-20240307")]
64 Claude3Haiku,
65}
66
67impl ClaudeModel {
68 pub fn as_str(&self) -> &'static str {
70 match self {
71 ClaudeModel::Claude35Sonnet => "claude-3-5-sonnet-20241022",
72 ClaudeModel::Claude3Opus => "claude-3-opus-20240229",
73 ClaudeModel::Claude3Sonnet => "claude-3-sonnet-20240229",
74 ClaudeModel::Claude3Haiku => "claude-3-haiku-20240307",
75 }
76 }
77
78 pub fn max_tokens(&self) -> u32 {
80 match self {
81 ClaudeModel::Claude35Sonnet => 200_000,
82 ClaudeModel::Claude3Opus => 200_000,
83 ClaudeModel::Claude3Sonnet => 200_000,
84 ClaudeModel::Claude3Haiku => 200_000,
85 }
86 }
87
88 pub fn input_cost_per_mtok(&self) -> f64 {
90 match self {
91 ClaudeModel::Claude35Sonnet => 3.0,
92 ClaudeModel::Claude3Opus => 15.0,
93 ClaudeModel::Claude3Sonnet => 3.0,
94 ClaudeModel::Claude3Haiku => 0.25,
95 }
96 }
97
98 pub fn output_cost_per_mtok(&self) -> f64 {
100 match self {
101 ClaudeModel::Claude35Sonnet => 15.0,
102 ClaudeModel::Claude3Opus => 75.0,
103 ClaudeModel::Claude3Sonnet => 15.0,
104 ClaudeModel::Claude3Haiku => 1.25,
105 }
106 }
107}
108
109#[derive(Debug, Clone, Serialize, Deserialize)]
111pub struct MessageRequest {
112 pub model: String,
114 pub messages: Vec<Message>,
116 pub max_tokens: u32,
118 #[serde(skip_serializing_if = "Option::is_none")]
120 pub system: Option<String>,
121 #[serde(skip_serializing_if = "Option::is_none")]
123 pub temperature: Option<f32>,
124 #[serde(skip_serializing_if = "Option::is_none")]
126 pub top_p: Option<f32>,
127 #[serde(skip_serializing_if = "Option::is_none")]
129 pub top_k: Option<u32>,
130 #[serde(skip_serializing_if = "Option::is_none")]
132 pub stop_sequences: Option<Vec<String>>,
133 #[serde(default)]
135 pub stream: bool,
136 #[serde(skip_serializing_if = "Option::is_none")]
138 pub metadata: Option<MessageMetadata>,
139}
140
141#[derive(Debug, Clone, Serialize, Deserialize)]
143pub struct Message {
144 pub role: Role,
146 pub content: MessageContent,
148}
149
150#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
152#[serde(rename_all = "lowercase")]
153pub enum Role {
154 User,
156 Assistant,
158}
159
160#[derive(Debug, Clone, Serialize, Deserialize)]
162#[serde(untagged)]
163pub enum MessageContent {
164 Text(String),
166 Parts(Vec<ContentBlock>),
168}
169
170#[derive(Debug, Clone, Serialize, Deserialize)]
172#[serde(tag = "type")]
173pub enum ContentBlock {
174 #[serde(rename = "text")]
176 Text { text: String },
177 #[serde(rename = "image")]
179 Image {
180 source: ImageSource,
181 },
182}
183
184#[derive(Debug, Clone, Serialize, Deserialize)]
186#[serde(tag = "type")]
187pub enum ImageSource {
188 #[serde(rename = "base64")]
190 Base64 {
191 media_type: String,
192 data: String,
193 },
194 #[serde(rename = "url")]
196 Url {
197 url: String,
198 },
199}
200
201#[derive(Debug, Clone, Serialize, Deserialize)]
203pub struct MessageMetadata {
204 #[serde(skip_serializing_if = "Option::is_none")]
206 pub user_id: Option<String>,
207 #[serde(flatten)]
209 pub custom: HashMap<String, serde_json::Value>,
210}
211
212#[derive(Debug, Clone, Serialize, Deserialize)]
214pub struct MessageResponse {
215 pub id: String,
217 #[serde(rename = "type")]
219 pub type_field: String,
220 pub role: Role,
222 pub content: Vec<ContentBlock>,
224 pub model: String,
226 pub stop_reason: Option<StopReason>,
228 pub stop_sequence: Option<String>,
230 pub usage: Usage,
232}
233
234#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
236#[serde(rename_all = "snake_case")]
237pub enum StopReason {
238 EndTurn,
240 MaxTokens,
242 StopSequence,
244}
245
246#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
248pub struct Usage {
249 pub input_tokens: u32,
251 pub output_tokens: u32,
253}
254
255impl Usage {
256 pub fn calculate_cost(&self, model: ClaudeModel) -> f64 {
258 let input_cost = (self.input_tokens as f64 / 1_000_000.0) * model.input_cost_per_mtok();
259 let output_cost = (self.output_tokens as f64 / 1_000_000.0) * model.output_cost_per_mtok();
260 input_cost + output_cost
261 }
262
263 pub fn total_tokens(&self) -> u32 {
265 self.input_tokens + self.output_tokens
266 }
267}
268
269#[derive(Debug, Clone, Serialize, Deserialize)]
271#[serde(tag = "type")]
272pub enum StreamEvent {
273 #[serde(rename = "message_start")]
275 MessageStart {
276 message: MessageStart,
277 },
278 #[serde(rename = "content_block_start")]
280 ContentBlockStart {
281 index: usize,
282 content_block: ContentBlockStart,
283 },
284 #[serde(rename = "ping")]
286 Ping,
287 #[serde(rename = "content_block_delta")]
289 ContentBlockDelta {
290 index: usize,
291 delta: Delta,
292 },
293 #[serde(rename = "content_block_stop")]
295 ContentBlockStop {
296 index: usize,
297 },
298 #[serde(rename = "message_delta")]
300 MessageDelta {
301 delta: MessageDeltaData,
302 usage: Usage,
303 },
304 #[serde(rename = "message_stop")]
306 MessageStop,
307 #[serde(rename = "error")]
309 Error {
310 error: ApiError,
311 },
312}
313
314#[derive(Debug, Clone, Serialize, Deserialize)]
316pub struct MessageStart {
317 pub id: String,
318 #[serde(rename = "type")]
319 pub type_field: String,
320 pub role: Role,
321 pub content: Vec<ContentBlock>,
322 pub model: String,
323 pub usage: Usage,
324}
325
326#[derive(Debug, Clone, Serialize, Deserialize)]
328#[serde(tag = "type")]
329pub enum ContentBlockStart {
330 #[serde(rename = "text")]
331 Text {
332 text: String,
333 },
334}
335
336#[derive(Debug, Clone, Serialize, Deserialize)]
338#[serde(tag = "type")]
339pub enum Delta {
340 #[serde(rename = "text_delta")]
341 TextDelta {
342 text: String,
343 },
344}
345
346#[derive(Debug, Clone, Serialize, Deserialize)]
348pub struct MessageDeltaData {
349 pub stop_reason: Option<StopReason>,
350 pub stop_sequence: Option<String>,
351}
352
353#[derive(Debug, Clone, Serialize, Deserialize)]
355pub struct ApiError {
356 #[serde(rename = "type")]
357 pub error_type: String,
358 pub message: String,
359}
360
361#[derive(Debug, Clone)]
363pub struct RateLimitInfo {
364 pub requests_remaining: Option<u32>,
366 pub requests_limit: Option<u32>,
368 pub tokens_remaining: Option<u32>,
370 pub tokens_limit: Option<u32>,
372 pub reset_at: Option<i64>,
374}
375
376#[derive(Debug, Clone, Default)]
378pub struct CostTracker {
379 pub total_input_tokens: u64,
381 pub total_output_tokens: u64,
383 pub total_cost: f64,
385 pub request_count: u64,
387}
388
389impl CostTracker {
390 pub fn new() -> Self {
392 Self::default()
393 }
394
395 pub fn record_usage(&mut self, usage: &Usage, model: ClaudeModel) {
397 self.total_input_tokens += usage.input_tokens as u64;
398 self.total_output_tokens += usage.output_tokens as u64;
399 self.total_cost += usage.calculate_cost(model);
400 self.request_count += 1;
401 }
402
403 pub fn avg_cost_per_request(&self) -> f64 {
405 if self.request_count == 0 {
406 0.0
407 } else {
408 self.total_cost / self.request_count as f64
409 }
410 }
411
412 pub fn avg_tokens_per_request(&self) -> f64 {
414 if self.request_count == 0 {
415 0.0
416 } else {
417 (self.total_input_tokens + self.total_output_tokens) as f64 / self.request_count as f64
418 }
419 }
420
421 pub fn reset(&mut self) {
423 self.total_input_tokens = 0;
424 self.total_output_tokens = 0;
425 self.total_cost = 0.0;
426 self.request_count = 0;
427 }
428}
429
430#[cfg(test)]
431mod tests {
432 use super::*;
433
434 #[test]
435 fn test_claude_model_costs() {
436 let usage = Usage {
437 input_tokens: 1000,
438 output_tokens: 500,
439 };
440
441 let cost_haiku = usage.calculate_cost(ClaudeModel::Claude3Haiku);
442 let cost_sonnet = usage.calculate_cost(ClaudeModel::Claude35Sonnet);
443 let cost_opus = usage.calculate_cost(ClaudeModel::Claude3Opus);
444
445 assert!(cost_haiku < cost_sonnet);
447 assert!(cost_haiku < cost_opus);
448
449 assert!(cost_opus > cost_sonnet);
451 }
452
453 #[test]
454 fn test_cost_tracker() {
455 let mut tracker = CostTracker::new();
456
457 let usage = Usage {
458 input_tokens: 1000,
459 output_tokens: 500,
460 };
461
462 tracker.record_usage(&usage, ClaudeModel::Claude3Haiku);
463 tracker.record_usage(&usage, ClaudeModel::Claude3Haiku);
464
465 assert_eq!(tracker.request_count, 2);
466 assert_eq!(tracker.total_input_tokens, 2000);
467 assert_eq!(tracker.total_output_tokens, 1000);
468 assert!(tracker.total_cost > 0.0);
469
470 let avg = tracker.avg_cost_per_request();
471 assert!(avg > 0.0);
472
473 tracker.reset();
474 assert_eq!(tracker.request_count, 0);
475 assert_eq!(tracker.total_cost, 0.0);
476 }
477}