use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AnthropicConfig {
pub api_key: String,
#[serde(default = "default_base_url")]
pub base_url: String,
#[serde(default = "default_timeout")]
pub timeout_secs: u64,
#[serde(default = "default_max_retries")]
pub max_retries: u32,
#[serde(default = "default_rate_limit")]
pub rate_limit_per_minute: u32,
#[serde(default = "default_api_version")]
pub api_version: String,
}
fn default_base_url() -> String {
"https://api.anthropic.com".to_string()
}
fn default_timeout() -> u64 {
60
}
fn default_max_retries() -> u32 {
3
}
fn default_rate_limit() -> u32 {
50
}
fn default_api_version() -> String {
"2023-06-01".to_string()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ClaudeModel {
#[serde(rename = "claude-3-5-sonnet-20241022")]
Claude35Sonnet,
#[serde(rename = "claude-3-opus-20240229")]
Claude3Opus,
#[serde(rename = "claude-3-sonnet-20240229")]
Claude3Sonnet,
#[serde(rename = "claude-3-haiku-20240307")]
Claude3Haiku,
}
impl ClaudeModel {
pub fn as_str(&self) -> &'static str {
match self {
ClaudeModel::Claude35Sonnet => "claude-3-5-sonnet-20241022",
ClaudeModel::Claude3Opus => "claude-3-opus-20240229",
ClaudeModel::Claude3Sonnet => "claude-3-sonnet-20240229",
ClaudeModel::Claude3Haiku => "claude-3-haiku-20240307",
}
}
pub fn max_tokens(&self) -> u32 {
match self {
ClaudeModel::Claude35Sonnet => 200_000,
ClaudeModel::Claude3Opus => 200_000,
ClaudeModel::Claude3Sonnet => 200_000,
ClaudeModel::Claude3Haiku => 200_000,
}
}
pub fn input_cost_per_mtok(&self) -> f64 {
match self {
ClaudeModel::Claude35Sonnet => 3.0,
ClaudeModel::Claude3Opus => 15.0,
ClaudeModel::Claude3Sonnet => 3.0,
ClaudeModel::Claude3Haiku => 0.25,
}
}
pub fn output_cost_per_mtok(&self) -> f64 {
match self {
ClaudeModel::Claude35Sonnet => 15.0,
ClaudeModel::Claude3Opus => 75.0,
ClaudeModel::Claude3Sonnet => 15.0,
ClaudeModel::Claude3Haiku => 1.25,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageRequest {
pub model: String,
pub messages: Vec<Message>,
pub max_tokens: u32,
#[serde(skip_serializing_if = "Option::is_none")]
pub system: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop_sequences: Option<Vec<String>>,
#[serde(default)]
pub stream: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<MessageMetadata>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Message {
pub role: Role,
pub content: MessageContent,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Role {
User,
Assistant,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum MessageContent {
Text(String),
Parts(Vec<ContentBlock>),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum ContentBlock {
#[serde(rename = "text")]
Text { text: String },
#[serde(rename = "image")]
Image {
source: ImageSource,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum ImageSource {
#[serde(rename = "base64")]
Base64 {
media_type: String,
data: String,
},
#[serde(rename = "url")]
Url {
url: String,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageMetadata {
#[serde(skip_serializing_if = "Option::is_none")]
pub user_id: Option<String>,
#[serde(flatten)]
pub custom: HashMap<String, serde_json::Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageResponse {
pub id: String,
#[serde(rename = "type")]
pub type_field: String,
pub role: Role,
pub content: Vec<ContentBlock>,
pub model: String,
pub stop_reason: Option<StopReason>,
pub stop_sequence: Option<String>,
pub usage: Usage,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum StopReason {
EndTurn,
MaxTokens,
StopSequence,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub struct Usage {
pub input_tokens: u32,
pub output_tokens: u32,
}
impl Usage {
pub fn calculate_cost(&self, model: ClaudeModel) -> f64 {
let input_cost = (self.input_tokens as f64 / 1_000_000.0) * model.input_cost_per_mtok();
let output_cost = (self.output_tokens as f64 / 1_000_000.0) * model.output_cost_per_mtok();
input_cost + output_cost
}
pub fn total_tokens(&self) -> u32 {
self.input_tokens + self.output_tokens
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum StreamEvent {
#[serde(rename = "message_start")]
MessageStart {
message: MessageStart,
},
#[serde(rename = "content_block_start")]
ContentBlockStart {
index: usize,
content_block: ContentBlockStart,
},
#[serde(rename = "ping")]
Ping,
#[serde(rename = "content_block_delta")]
ContentBlockDelta {
index: usize,
delta: Delta,
},
#[serde(rename = "content_block_stop")]
ContentBlockStop {
index: usize,
},
#[serde(rename = "message_delta")]
MessageDelta {
delta: MessageDeltaData,
usage: Usage,
},
#[serde(rename = "message_stop")]
MessageStop,
#[serde(rename = "error")]
Error {
error: ApiError,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageStart {
pub id: String,
#[serde(rename = "type")]
pub type_field: String,
pub role: Role,
pub content: Vec<ContentBlock>,
pub model: String,
pub usage: Usage,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum ContentBlockStart {
#[serde(rename = "text")]
Text {
text: String,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum Delta {
#[serde(rename = "text_delta")]
TextDelta {
text: String,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MessageDeltaData {
pub stop_reason: Option<StopReason>,
pub stop_sequence: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ApiError {
#[serde(rename = "type")]
pub error_type: String,
pub message: String,
}
#[derive(Debug, Clone)]
pub struct RateLimitInfo {
pub requests_remaining: Option<u32>,
pub requests_limit: Option<u32>,
pub tokens_remaining: Option<u32>,
pub tokens_limit: Option<u32>,
pub reset_at: Option<i64>,
}
#[derive(Debug, Clone, Default)]
pub struct CostTracker {
pub total_input_tokens: u64,
pub total_output_tokens: u64,
pub total_cost: f64,
pub request_count: u64,
}
impl CostTracker {
pub fn new() -> Self {
Self::default()
}
pub fn record_usage(&mut self, usage: &Usage, model: ClaudeModel) {
self.total_input_tokens += usage.input_tokens as u64;
self.total_output_tokens += usage.output_tokens as u64;
self.total_cost += usage.calculate_cost(model);
self.request_count += 1;
}
pub fn avg_cost_per_request(&self) -> f64 {
if self.request_count == 0 {
0.0
} else {
self.total_cost / self.request_count as f64
}
}
pub fn avg_tokens_per_request(&self) -> f64 {
if self.request_count == 0 {
0.0
} else {
(self.total_input_tokens + self.total_output_tokens) as f64 / self.request_count as f64
}
}
pub fn reset(&mut self) {
self.total_input_tokens = 0;
self.total_output_tokens = 0;
self.total_cost = 0.0;
self.request_count = 0;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_claude_model_costs() {
let usage = Usage {
input_tokens: 1000,
output_tokens: 500,
};
let cost_haiku = usage.calculate_cost(ClaudeModel::Claude3Haiku);
let cost_sonnet = usage.calculate_cost(ClaudeModel::Claude35Sonnet);
let cost_opus = usage.calculate_cost(ClaudeModel::Claude3Opus);
assert!(cost_haiku < cost_sonnet);
assert!(cost_haiku < cost_opus);
assert!(cost_opus > cost_sonnet);
}
#[test]
fn test_cost_tracker() {
let mut tracker = CostTracker::new();
let usage = Usage {
input_tokens: 1000,
output_tokens: 500,
};
tracker.record_usage(&usage, ClaudeModel::Claude3Haiku);
tracker.record_usage(&usage, ClaudeModel::Claude3Haiku);
assert_eq!(tracker.request_count, 2);
assert_eq!(tracker.total_input_tokens, 2000);
assert_eq!(tracker.total_output_tokens, 1000);
assert!(tracker.total_cost > 0.0);
let avg = tracker.avg_cost_per_request();
assert!(avg > 0.0);
tracker.reset();
assert_eq!(tracker.request_count, 0);
assert_eq!(tracker.total_cost, 0.0);
}
}