Skip to main content

llm_optimizer_integrations/anthropic/
types.rs

1//! Anthropic Claude API type definitions
2//!
3//! This module provides comprehensive type definitions for the Claude API.
4
5use serde::{Deserialize, Serialize};
6use std::collections::HashMap;
7
8/// Anthropic API configuration
9#[derive(Debug, Clone, Serialize, Deserialize)]
10pub struct AnthropicConfig {
11    /// API key for authentication
12    pub api_key: String,
13    /// API base URL
14    #[serde(default = "default_base_url")]
15    pub base_url: String,
16    /// Request timeout in seconds
17    #[serde(default = "default_timeout")]
18    pub timeout_secs: u64,
19    /// Maximum retry attempts
20    #[serde(default = "default_max_retries")]
21    pub max_retries: u32,
22    /// Rate limit: requests per minute (tier-specific)
23    #[serde(default = "default_rate_limit")]
24    pub rate_limit_per_minute: u32,
25    /// API version
26    #[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/// Claude model identifiers
51#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
52pub enum ClaudeModel {
53    /// Claude 3.5 Sonnet (latest)
54    #[serde(rename = "claude-3-5-sonnet-20241022")]
55    Claude35Sonnet,
56    /// Claude 3 Opus
57    #[serde(rename = "claude-3-opus-20240229")]
58    Claude3Opus,
59    /// Claude 3 Sonnet
60    #[serde(rename = "claude-3-sonnet-20240229")]
61    Claude3Sonnet,
62    /// Claude 3 Haiku
63    #[serde(rename = "claude-3-haiku-20240307")]
64    Claude3Haiku,
65}
66
67impl ClaudeModel {
68    /// Get the model identifier string
69    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    /// Get maximum tokens for the model
79    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    /// Get input token cost per million tokens (in USD)
89    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    /// Get output token cost per million tokens (in USD)
99    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/// Message request to Claude API
110#[derive(Debug, Clone, Serialize, Deserialize)]
111pub struct MessageRequest {
112    /// Model to use
113    pub model: String,
114    /// List of messages in the conversation
115    pub messages: Vec<Message>,
116    /// Maximum tokens to generate
117    pub max_tokens: u32,
118    /// Optional system prompt
119    #[serde(skip_serializing_if = "Option::is_none")]
120    pub system: Option<String>,
121    /// Optional temperature (0.0 - 1.0)
122    #[serde(skip_serializing_if = "Option::is_none")]
123    pub temperature: Option<f32>,
124    /// Optional top-p sampling
125    #[serde(skip_serializing_if = "Option::is_none")]
126    pub top_p: Option<f32>,
127    /// Optional top-k sampling
128    #[serde(skip_serializing_if = "Option::is_none")]
129    pub top_k: Option<u32>,
130    /// Optional stop sequences
131    #[serde(skip_serializing_if = "Option::is_none")]
132    pub stop_sequences: Option<Vec<String>>,
133    /// Enable streaming
134    #[serde(default)]
135    pub stream: bool,
136    /// Optional metadata
137    #[serde(skip_serializing_if = "Option::is_none")]
138    pub metadata: Option<MessageMetadata>,
139}
140
141/// Message in a conversation
142#[derive(Debug, Clone, Serialize, Deserialize)]
143pub struct Message {
144    /// Role of the message sender
145    pub role: Role,
146    /// Content of the message
147    pub content: MessageContent,
148}
149
150/// Message role
151#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
152#[serde(rename_all = "lowercase")]
153pub enum Role {
154    /// User message
155    User,
156    /// Assistant message
157    Assistant,
158}
159
160/// Message content (can be text or multi-modal)
161#[derive(Debug, Clone, Serialize, Deserialize)]
162#[serde(untagged)]
163pub enum MessageContent {
164    /// Simple text content
165    Text(String),
166    /// Multi-part content (text, images, etc.)
167    Parts(Vec<ContentBlock>),
168}
169
170/// Content block (text or image)
171#[derive(Debug, Clone, Serialize, Deserialize)]
172#[serde(tag = "type")]
173pub enum ContentBlock {
174    /// Text content
175    #[serde(rename = "text")]
176    Text { text: String },
177    /// Image content
178    #[serde(rename = "image")]
179    Image {
180        source: ImageSource,
181    },
182}
183
184/// Image source
185#[derive(Debug, Clone, Serialize, Deserialize)]
186#[serde(tag = "type")]
187pub enum ImageSource {
188    /// Base64-encoded image
189    #[serde(rename = "base64")]
190    Base64 {
191        media_type: String,
192        data: String,
193    },
194    /// Image URL (not supported in all contexts)
195    #[serde(rename = "url")]
196    Url {
197        url: String,
198    },
199}
200
201/// Message metadata
202#[derive(Debug, Clone, Serialize, Deserialize)]
203pub struct MessageMetadata {
204    /// User ID for tracking
205    #[serde(skip_serializing_if = "Option::is_none")]
206    pub user_id: Option<String>,
207    /// Custom metadata fields
208    #[serde(flatten)]
209    pub custom: HashMap<String, serde_json::Value>,
210}
211
212/// Response from Claude API
213#[derive(Debug, Clone, Serialize, Deserialize)]
214pub struct MessageResponse {
215    /// Unique identifier for the response
216    pub id: String,
217    /// Object type (always "message")
218    #[serde(rename = "type")]
219    pub type_field: String,
220    /// Role (always "assistant")
221    pub role: Role,
222    /// Response content
223    pub content: Vec<ContentBlock>,
224    /// Model used
225    pub model: String,
226    /// Stop reason
227    pub stop_reason: Option<StopReason>,
228    /// Stop sequence that was matched
229    pub stop_sequence: Option<String>,
230    /// Token usage statistics
231    pub usage: Usage,
232}
233
234/// Reason why generation stopped
235#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
236#[serde(rename_all = "snake_case")]
237pub enum StopReason {
238    /// Reached natural end of message
239    EndTurn,
240    /// Hit max_tokens limit
241    MaxTokens,
242    /// Matched a stop sequence
243    StopSequence,
244}
245
246/// Token usage statistics
247#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
248pub struct Usage {
249    /// Number of input tokens
250    pub input_tokens: u32,
251    /// Number of output tokens
252    pub output_tokens: u32,
253}
254
255impl Usage {
256    /// Calculate total cost in USD
257    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    /// Total tokens used
264    pub fn total_tokens(&self) -> u32 {
265        self.input_tokens + self.output_tokens
266    }
267}
268
269/// Streaming event from Claude API
270#[derive(Debug, Clone, Serialize, Deserialize)]
271#[serde(tag = "type")]
272pub enum StreamEvent {
273    /// Message start event
274    #[serde(rename = "message_start")]
275    MessageStart {
276        message: MessageStart,
277    },
278    /// Content block start
279    #[serde(rename = "content_block_start")]
280    ContentBlockStart {
281        index: usize,
282        content_block: ContentBlockStart,
283    },
284    /// Ping event (keep-alive)
285    #[serde(rename = "ping")]
286    Ping,
287    /// Content block delta (incremental content)
288    #[serde(rename = "content_block_delta")]
289    ContentBlockDelta {
290        index: usize,
291        delta: Delta,
292    },
293    /// Content block stop
294    #[serde(rename = "content_block_stop")]
295    ContentBlockStop {
296        index: usize,
297    },
298    /// Message delta (final statistics)
299    #[serde(rename = "message_delta")]
300    MessageDelta {
301        delta: MessageDeltaData,
302        usage: Usage,
303    },
304    /// Message stop
305    #[serde(rename = "message_stop")]
306    MessageStop,
307    /// Error event
308    #[serde(rename = "error")]
309    Error {
310        error: ApiError,
311    },
312}
313
314/// Message start data
315#[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/// Content block start
327#[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/// Delta (incremental content)
337#[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/// Message delta data
347#[derive(Debug, Clone, Serialize, Deserialize)]
348pub struct MessageDeltaData {
349    pub stop_reason: Option<StopReason>,
350    pub stop_sequence: Option<String>,
351}
352
353/// API error response
354#[derive(Debug, Clone, Serialize, Deserialize)]
355pub struct ApiError {
356    #[serde(rename = "type")]
357    pub error_type: String,
358    pub message: String,
359}
360
361/// Rate limit information
362#[derive(Debug, Clone)]
363pub struct RateLimitInfo {
364    /// Requests remaining in current window
365    pub requests_remaining: Option<u32>,
366    /// Request limit per window
367    pub requests_limit: Option<u32>,
368    /// Tokens remaining in current window
369    pub tokens_remaining: Option<u32>,
370    /// Token limit per window
371    pub tokens_limit: Option<u32>,
372    /// Time when rate limit resets (Unix timestamp)
373    pub reset_at: Option<i64>,
374}
375
376/// Cost tracking information
377#[derive(Debug, Clone, Default)]
378pub struct CostTracker {
379    /// Total input tokens used
380    pub total_input_tokens: u64,
381    /// Total output tokens used
382    pub total_output_tokens: u64,
383    /// Total cost in USD
384    pub total_cost: f64,
385    /// Number of requests made
386    pub request_count: u64,
387}
388
389impl CostTracker {
390    /// Create a new cost tracker
391    pub fn new() -> Self {
392        Self::default()
393    }
394
395    /// Record usage from a response
396    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    /// Get average cost per request
404    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    /// Get average tokens per request
413    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    /// Reset all counters
422    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        // Haiku should be cheapest
446        assert!(cost_haiku < cost_sonnet);
447        assert!(cost_haiku < cost_opus);
448
449        // Opus should be most expensive
450        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}