From 05973bf9591da273c4c00c4f1453b80fd5827d3c Mon Sep 17 00:00:00 2001 From: YoVinchen Date: Wed, 3 Dec 2025 00:17:34 +0800 Subject: [PATCH] feat(proxy): implement usage tracking subsystem Add request logger with automatic cost calculation. Implement token parser for Claude/OpenAI/Gemini responses. Add cost calculator based on model pricing configuration. --- src-tauri/src/proxy/usage/calculator.rs | 7 +- src-tauri/src/proxy/usage/logger.rs | 76 ++- src-tauri/src/proxy/usage/mod.rs | 4 + src-tauri/src/proxy/usage/parser.rs | 17 +- src-tauri/src/services/usage_stats.rs | 627 +++++++++++++++++++----- 5 files changed, 562 insertions(+), 169 deletions(-) diff --git a/src-tauri/src/proxy/usage/calculator.rs b/src-tauri/src/proxy/usage/calculator.rs index 4de0bfe71..e7cd335fa 100644 --- a/src-tauri/src/proxy/usage/calculator.rs +++ b/src-tauri/src/proxy/usage/calculator.rs @@ -48,10 +48,9 @@ impl CostCalculator { let output_cost = Decimal::from(usage.output_tokens) * pricing.output_cost_per_million / million * cost_multiplier; - let cache_read_cost = Decimal::from(usage.cache_read_tokens) - * pricing.cache_read_cost_per_million - / million - * cost_multiplier; + let cache_read_cost = + Decimal::from(usage.cache_read_tokens) * pricing.cache_read_cost_per_million / million + * cost_multiplier; let cache_creation_cost = Decimal::from(usage.cache_creation_tokens) * pricing.cache_creation_cost_per_million / million diff --git a/src-tauri/src/proxy/usage/logger.rs b/src-tauri/src/proxy/usage/logger.rs index 0a77fb7ce..8fb995e3f 100644 --- a/src-tauri/src/proxy/usage/logger.rs +++ b/src-tauri/src/proxy/usage/logger.rs @@ -4,6 +4,7 @@ use super::calculator::{CostBreakdown, CostCalculator, ModelPricing}; use super::parser::TokenUsage; use crate::database::Database; use crate::error::AppError; +use crate::services::usage_stats::find_model_pricing_row; use rust_decimal::Decimal; use std::time::SystemTime; @@ -94,6 +95,7 @@ impl<'a> UsageLogger<'a> { } /// 记录失败的请求 + #[allow(dead_code, clippy::too_many_arguments)] pub fn log_error( &self, request_id: String, @@ -123,31 +125,19 @@ impl<'a> UsageLogger<'a> { /// 获取模型定价 pub fn get_model_pricing(&self, model_id: &str) -> Result, AppError> { let conn = crate::database::lock_conn!(self.db.conn); - - let result = conn.query_row( - "SELECT input_cost_per_million, output_cost_per_million, - cache_read_cost_per_million, cache_creation_cost_per_million - FROM model_pricing WHERE model_id = ?1", - [model_id], - |row| { - Ok(ModelPricing::from_strings( - &row.get::<_, String>(0)?, - &row.get::<_, String>(1)?, - &row.get::<_, String>(2)?, - &row.get::<_, String>(3)?, - )) - }, - ); - - match result { - Ok(Ok(pricing)) => Ok(Some(pricing)), - Ok(Err(e)) => Err(AppError::Database(format!("解析定价数据失败: {e}"))), - Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(AppError::Database(format!("查询定价失败: {e}"))), + let row = find_model_pricing_row(&conn, model_id)?; + match row { + Some((input, output, cache_read, cache_creation)) => { + ModelPricing::from_strings(&input, &output, &cache_read, &cache_creation) + .map(Some) + .map_err(|e| AppError::Database(format!("解析定价数据失败: {e}"))) + } + None => Ok(None), } } /// 计算并记录请求 + #[allow(clippy::too_many_arguments)] pub fn log_with_calculation( &self, request_id: String, @@ -163,7 +153,7 @@ impl<'a> UsageLogger<'a> { let pricing = self.get_model_pricing(&model)?; if pricing.is_none() { - log::warn!("模型 {} 的定价信息未找到,成本将记录为 0", model); + log::warn!("模型 {model} 的定价信息未找到,成本将记录为 0"); } let cost = CostCalculator::try_calculate(&usage, pricing.as_ref(), cost_multiplier); @@ -213,18 +203,17 @@ mod tests { cache_creation_tokens: 0, }; - logger - .log_with_calculation( - "req-123".to_string(), - "provider-1".to_string(), - "claude".to_string(), - "test-model".to_string(), - usage, - Decimal::from(1), - 100, - 200, - None, - )?; + logger.log_with_calculation( + "req-123".to_string(), + "provider-1".to_string(), + "claude".to_string(), + "test-model".to_string(), + usage, + Decimal::from(1), + 100, + 200, + None, + )?; // 验证记录已插入 let conn = crate::database::lock_conn!(db.conn); @@ -244,16 +233,15 @@ mod tests { let db = Database::memory()?; let logger = UsageLogger::new(&db); - logger - .log_error( - "req-error".to_string(), - "provider-1".to_string(), - "claude".to_string(), - "unknown-model".to_string(), - 500, - "Internal Server Error".to_string(), - 50, - )?; + logger.log_error( + "req-error".to_string(), + "provider-1".to_string(), + "claude".to_string(), + "unknown-model".to_string(), + 500, + "Internal Server Error".to_string(), + 50, + )?; // 验证错误记录已插入 let conn = crate::database::lock_conn!(db.conn); diff --git a/src-tauri/src/proxy/usage/mod.rs b/src-tauri/src/proxy/usage/mod.rs index 38a356fbc..052ec8e08 100644 --- a/src-tauri/src/proxy/usage/mod.rs +++ b/src-tauri/src/proxy/usage/mod.rs @@ -6,6 +6,10 @@ pub mod calculator; pub mod logger; pub mod parser; +// 仅导出内部使用的类型,避免未使用警告 +#[allow(unused_imports)] pub use calculator::{CostBreakdown, CostCalculator, ModelPricing}; +#[allow(unused_imports)] pub use logger::{RequestLog, UsageLogger}; +#[allow(unused_imports)] pub use parser::{ApiType, TokenUsage}; diff --git a/src-tauri/src/proxy/usage/parser.rs b/src-tauri/src/proxy/usage/parser.rs index ef9c1f551..7945ff549 100644 --- a/src-tauri/src/proxy/usage/parser.rs +++ b/src-tauri/src/proxy/usage/parser.rs @@ -20,6 +20,7 @@ pub struct TokenUsage { /// API 类型 #[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[allow(dead_code)] pub enum ApiType { Claude, OpenRouter, @@ -46,6 +47,7 @@ impl TokenUsage { } /// 从 Claude API 流式响应解析 + #[allow(dead_code)] pub fn from_claude_stream_events(events: &[Value]) -> Option { let mut usage = Self::default(); @@ -53,18 +55,18 @@ impl TokenUsage { if let Some(event_type) = event.get("type").and_then(|v| v.as_str()) { match event_type { "message_start" => { - if let Some(msg_usage) = event.get("message").and_then(|m| m.get("usage")) - { - usage.input_tokens = - msg_usage.get("input_tokens")?.as_u64()? as u32; + if let Some(msg_usage) = event.get("message").and_then(|m| m.get("usage")) { + usage.input_tokens = msg_usage.get("input_tokens")?.as_u64()? as u32; usage.cache_read_tokens = msg_usage .get("cache_read_input_tokens") .and_then(|v| v.as_u64()) - .unwrap_or(0) as u32; + .unwrap_or(0) + as u32; usage.cache_creation_tokens = msg_usage .get("cache_creation_input_tokens") .and_then(|v| v.as_u64()) - .unwrap_or(0) as u32; + .unwrap_or(0) + as u32; } } "message_delta" => { @@ -86,6 +88,7 @@ impl TokenUsage { } /// 从 OpenRouter 响应解析 (OpenAI 格式) + #[allow(dead_code)] pub fn from_openrouter_response(body: &Value) -> Option { let usage = body.get("usage")?; Some(Self { @@ -114,6 +117,7 @@ impl TokenUsage { } /// 从 Codex API 流式响应解析 + #[allow(dead_code)] pub fn from_codex_stream_events(events: &[Value]) -> Option { for event in events { if let Some(event_type) = event.get("type").and_then(|v| v.as_str()) { @@ -142,6 +146,7 @@ impl TokenUsage { } /// 从 Gemini API 流式响应解析 + #[allow(dead_code)] pub fn from_gemini_stream_chunks(chunks: &[Value]) -> Option { let mut total_input = 0u32; let mut total_output = 0u32; diff --git a/src-tauri/src/services/usage_stats.rs b/src-tauri/src/services/usage_stats.rs index f74e80b8f..e55e6eb79 100644 --- a/src-tauri/src/services/usage_stats.rs +++ b/src-tauri/src/services/usage_stats.rs @@ -4,30 +4,44 @@ use crate::database::{lock_conn, Database}; use crate::error::AppError; -use rusqlite::params; +use chrono::{Duration, Utc}; +use regex::Regex; +use rusqlite::{params, Connection, OptionalExtension}; use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::collections::HashMap; +use std::str::FromStr; /// 使用量汇总 #[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] pub struct UsageSummary { pub total_requests: u64, pub total_cost: String, pub total_input_tokens: u64, pub total_output_tokens: u64, + pub total_cache_creation_tokens: u64, + pub total_cache_read_tokens: u64, pub success_rate: f32, } /// 每日统计 #[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] pub struct DailyStats { pub date: String, pub request_count: u64, pub total_cost: String, pub total_tokens: u64, + pub total_input_tokens: u64, + pub total_output_tokens: u64, + pub total_cache_creation_tokens: u64, + pub total_cache_read_tokens: u64, } /// Provider 统计 #[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] pub struct ProviderStats { pub provider_id: String, pub provider_name: String, @@ -40,6 +54,7 @@ pub struct ProviderStats { /// 模型统计 #[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] pub struct ModelStats { pub model: String, pub request_count: u64, @@ -49,7 +64,8 @@ pub struct ModelStats { } /// 请求日志过滤器 -#[derive(Debug, Clone, Default)] +#[derive(Debug, Clone, Default, Deserialize)] +#[serde(rename_all = "camelCase")] pub struct LogFilters { pub provider_id: Option, pub model: Option, @@ -60,9 +76,12 @@ pub struct LogFilters { /// 请求日志详情 #[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] pub struct RequestLogDetail { pub request_id: String, pub provider_id: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub provider_name: Option, pub app_type: String, pub model: String, pub input_tokens: u32, @@ -111,7 +130,10 @@ impl Database { "SELECT COUNT(*) as total_requests, COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost, - COALESCE(SUM(input_tokens + output_tokens), 0) as total_tokens, + COALESCE(SUM(input_tokens), 0) as total_input_tokens, + COALESCE(SUM(output_tokens), 0) as total_output_tokens, + COALESCE(SUM(cache_creation_tokens), 0) as total_cache_creation_tokens, + COALESCE(SUM(cache_read_tokens), 0) as total_cache_read_tokens, COALESCE(SUM(CASE WHEN status_code >= 200 AND status_code < 300 THEN 1 ELSE 0 END), 0) as success_count FROM proxy_request_logs {where_clause}" @@ -120,8 +142,11 @@ impl Database { let result = conn.query_row(&sql, rusqlite::params_from_iter(params_vec), |row| { let total_requests: i64 = row.get(0)?; let total_cost: f64 = row.get(1)?; - let total_tokens: i64 = row.get(2)?; - let success_count: i64 = row.get(3)?; + let total_input_tokens: i64 = row.get(2)?; + let total_output_tokens: i64 = row.get(3)?; + let total_cache_creation_tokens: i64 = row.get(4)?; + let total_cache_read_tokens: i64 = row.get(5)?; + let success_count: i64 = row.get(6)?; let success_rate = if total_requests > 0 { (success_count as f32 / total_requests as f32) * 100.0 @@ -131,9 +156,11 @@ impl Database { Ok(UsageSummary { total_requests: total_requests as u64, - total_cost: format!("{:.6}", total_cost), - total_input_tokens: 0, - total_output_tokens: 0, + total_cost: format!("{total_cost:.6}"), + total_input_tokens: total_input_tokens as u64, + total_output_tokens: total_output_tokens as u64, + total_cache_creation_tokens: total_cache_creation_tokens as u64, + total_cache_read_tokens: total_cache_read_tokens as u64, success_rate, }) })?; @@ -142,38 +169,128 @@ impl Database { } /// 获取每日趋势 - pub fn get_daily_trends( - &self, - days: u32, - ) -> Result, AppError> { + pub fn get_daily_trends(&self, days: u32) -> Result, AppError> { let conn = lock_conn!(self.conn); - let sql = "SELECT - date(created_at, 'unixepoch') as date, - COUNT(*) as request_count, - COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost, - COALESCE(SUM(input_tokens + output_tokens), 0) as total_tokens - FROM proxy_request_logs - WHERE created_at >= strftime('%s', 'now', ?) - GROUP BY date - ORDER BY date DESC"; + if days <= 1 { + let sql = "SELECT + strftime('%Y-%m-%dT%H:00:00Z', datetime(created_at, 'unixepoch')) as bucket, + COUNT(*) as request_count, + COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost, + COALESCE(SUM(input_tokens + output_tokens), 0) as total_tokens, + COALESCE(SUM(input_tokens), 0) as total_input_tokens, + COALESCE(SUM(output_tokens), 0) as total_output_tokens, + COALESCE(SUM(cache_creation_tokens), 0) as total_cache_creation_tokens, + COALESCE(SUM(cache_read_tokens), 0) as total_cache_read_tokens + FROM proxy_request_logs + WHERE created_at >= strftime('%s', 'now', '-1 day') + GROUP BY bucket + ORDER BY bucket ASC"; - let mut stmt = conn.prepare(sql)?; - let rows = stmt.query_map([format!("-{} days", days)], |row| { - Ok(DailyStats { - date: row.get(0)?, - request_count: row.get::<_, i64>(1)? as u64, - total_cost: format!("{:.6}", row.get::<_, f64>(2)?), - total_tokens: row.get::<_, i64>(3)? as u64, - }) - })?; + let mut stmt = conn.prepare(sql)?; + let rows = stmt.query_map([], |row| { + Ok(DailyStats { + date: row.get(0)?, + request_count: row.get::<_, i64>(1)? as u64, + total_cost: format!("{:.6}", row.get::<_, f64>(2)?), + total_tokens: row.get::<_, i64>(3)? as u64, + total_input_tokens: row.get::<_, i64>(4)? as u64, + total_output_tokens: row.get::<_, i64>(5)? as u64, + total_cache_creation_tokens: row.get::<_, i64>(6)? as u64, + total_cache_read_tokens: row.get::<_, i64>(7)? as u64, + }) + })?; - let mut stats = Vec::new(); - for row in rows { - stats.push(row?); + let mut buckets: HashMap = HashMap::new(); + for row in rows { + let stat = row?; + buckets.insert(stat.date.clone(), stat); + } + + let mut stats = Vec::new(); + let today = Utc::now().date_naive(); + for hour in 0..24 { + let bucket = today + .and_hms_opt(hour, 0, 0) + .unwrap() + .format("%Y-%m-%dT%H:00:00Z") + .to_string(); + + if let Some(stat) = buckets.remove(&bucket) { + stats.push(stat); + } else { + stats.push(DailyStats { + date: bucket, + request_count: 0, + total_cost: "0.000000".to_string(), + total_tokens: 0, + total_input_tokens: 0, + total_output_tokens: 0, + total_cache_creation_tokens: 0, + total_cache_read_tokens: 0, + }); + } + } + Ok(stats) + } else { + let sql = "SELECT + date(created_at, 'unixepoch') as bucket, + COUNT(*) as request_count, + COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost, + COALESCE(SUM(input_tokens + output_tokens), 0) as total_tokens, + COALESCE(SUM(input_tokens), 0) as total_input_tokens, + COALESCE(SUM(output_tokens), 0) as total_output_tokens, + COALESCE(SUM(cache_creation_tokens), 0) as total_cache_creation_tokens, + COALESCE(SUM(cache_read_tokens), 0) as total_cache_read_tokens + FROM proxy_request_logs + WHERE created_at >= strftime('%s', 'now', ?) + GROUP BY bucket + ORDER BY bucket ASC"; + + let mut stmt = conn.prepare(sql)?; + let rows = stmt.query_map([format!("-{days} days")], |row| { + Ok(DailyStats { + date: row.get(0)?, + request_count: row.get::<_, i64>(1)? as u64, + total_cost: format!("{:.6}", row.get::<_, f64>(2)?), + total_tokens: row.get::<_, i64>(3)? as u64, + total_input_tokens: row.get::<_, i64>(4)? as u64, + total_output_tokens: row.get::<_, i64>(5)? as u64, + total_cache_creation_tokens: row.get::<_, i64>(6)? as u64, + total_cache_read_tokens: row.get::<_, i64>(7)? as u64, + }) + })?; + + let mut map = HashMap::new(); + for row in rows { + let stat = row?; + map.insert(stat.date.clone(), stat); + } + + let mut stats = Vec::new(); + let start_day = + Utc::now().date_naive() - Duration::days((days.saturating_sub(1)) as i64); + + for i in 0..days { + let day = start_day + Duration::days(i as i64); + let key = day.format("%Y-%m-%d").to_string(); + if let Some(stat) = map.remove(&key) { + stats.push(stat); + } else { + stats.push(DailyStats { + date: key, + request_count: 0, + total_cost: "0.000000".to_string(), + total_tokens: 0, + total_input_tokens: 0, + total_output_tokens: 0, + total_cache_creation_tokens: 0, + total_cache_read_tokens: 0, + }); + } + } + Ok(stats) } - - Ok(stats) } /// 获取 Provider 统计 @@ -205,7 +322,9 @@ impl Database { Ok(ProviderStats { provider_id: row.get(0)?, - provider_name: row.get::<_, Option>(1)?.unwrap_or_else(|| "Unknown".to_string()), + provider_name: row + .get::<_, Option>(1)? + .unwrap_or_else(|| "Unknown".to_string()), request_count: request_count as u64, total_tokens: row.get::<_, i64>(3)? as u64, total_cost: format!("{:.6}", row.get::<_, f64>(4)?), @@ -249,8 +368,8 @@ impl Database { model: row.get(0)?, request_count: request_count as u64, total_tokens: row.get::<_, i64>(2)? as u64, - total_cost: format!("{:.6}", total_cost), - avg_cost_per_request: format!("{:.6}", avg_cost), + total_cost: format!("{total_cost:.6}"), + avg_cost_per_request: format!("{avg_cost:.6}"), }) })?; @@ -305,13 +424,14 @@ impl Database { params.push(Box::new(offset as i64)); let sql = format!( - "SELECT request_id, provider_id, app_type, model, - input_tokens, output_tokens, cache_read_tokens, cache_creation_tokens, - input_cost_usd, output_cost_usd, cache_read_cost_usd, cache_creation_cost_usd, total_cost_usd, - latency_ms, status_code, error_message, created_at - FROM proxy_request_logs + "SELECT l.request_id, l.provider_id, p.name as provider_name, l.app_type, l.model, + l.input_tokens, l.output_tokens, l.cache_read_tokens, l.cache_creation_tokens, + l.input_cost_usd, l.output_cost_usd, l.cache_read_cost_usd, l.cache_creation_cost_usd, l.total_cost_usd, + l.latency_ms, l.status_code, l.error_message, l.created_at + FROM proxy_request_logs l + LEFT JOIN providers p ON l.provider_id = p.id AND l.app_type = p.app_type {where_clause} - ORDER BY created_at DESC + ORDER BY l.created_at DESC LIMIT ? OFFSET ?" ); @@ -321,69 +441,95 @@ impl Database { Ok(RequestLogDetail { request_id: row.get(0)?, provider_id: row.get(1)?, - app_type: row.get(2)?, - model: row.get(3)?, - input_tokens: row.get::<_, i64>(4)? as u32, - output_tokens: row.get::<_, i64>(5)? as u32, - cache_read_tokens: row.get::<_, i64>(6)? as u32, - cache_creation_tokens: row.get::<_, i64>(7)? as u32, - input_cost_usd: row.get(8)?, - output_cost_usd: row.get(9)?, - cache_read_cost_usd: row.get(10)?, - cache_creation_cost_usd: row.get(11)?, - total_cost_usd: row.get(12)?, - latency_ms: row.get::<_, i64>(13)? as u64, - status_code: row.get::<_, i64>(14)? as u16, - error_message: row.get(15)?, - created_at: row.get(16)?, + provider_name: row.get(2)?, + app_type: row.get(3)?, + model: row.get(4)?, + input_tokens: row.get::<_, i64>(5)? as u32, + output_tokens: row.get::<_, i64>(6)? as u32, + cache_read_tokens: row.get::<_, i64>(7)? as u32, + cache_creation_tokens: row.get::<_, i64>(8)? as u32, + input_cost_usd: row.get(9)?, + output_cost_usd: row.get(10)?, + cache_read_cost_usd: row.get(11)?, + cache_creation_cost_usd: row.get(12)?, + total_cost_usd: row.get(13)?, + latency_ms: row.get::<_, i64>(14)? as u64, + status_code: row.get::<_, i64>(15)? as u16, + error_message: row.get(16)?, + created_at: row.get(17)?, }) })?; let mut logs = Vec::new(); + let mut provider_cache = HashMap::new(); + let mut pricing_cache = HashMap::new(); + for row in rows { - logs.push(row?); + let mut log = row?; + Self::maybe_backfill_log_costs( + &conn, + &mut log, + &mut provider_cache, + &mut pricing_cache, + )?; + logs.push(log); } Ok(logs) } /// 获取单个请求详情 - pub fn get_request_detail(&self, request_id: &str) -> Result, AppError> { + pub fn get_request_detail( + &self, + request_id: &str, + ) -> Result, AppError> { let conn = lock_conn!(self.conn); let result = conn.query_row( - "SELECT request_id, provider_id, app_type, model, + "SELECT l.request_id, l.provider_id, p.name as provider_name, l.app_type, l.model, input_tokens, output_tokens, cache_read_tokens, cache_creation_tokens, input_cost_usd, output_cost_usd, cache_read_cost_usd, cache_creation_cost_usd, total_cost_usd, latency_ms, status_code, error_message, created_at - FROM proxy_request_logs - WHERE request_id = ?", + FROM proxy_request_logs l + LEFT JOIN providers p ON l.provider_id = p.id AND l.app_type = p.app_type + WHERE l.request_id = ?", [request_id], |row| { Ok(RequestLogDetail { request_id: row.get(0)?, provider_id: row.get(1)?, - app_type: row.get(2)?, - model: row.get(3)?, - input_tokens: row.get::<_, i64>(4)? as u32, - output_tokens: row.get::<_, i64>(5)? as u32, - cache_read_tokens: row.get::<_, i64>(6)? as u32, - cache_creation_tokens: row.get::<_, i64>(7)? as u32, - input_cost_usd: row.get(8)?, - output_cost_usd: row.get(9)?, - cache_read_cost_usd: row.get(10)?, - cache_creation_cost_usd: row.get(11)?, - total_cost_usd: row.get(12)?, - latency_ms: row.get::<_, i64>(13)? as u64, - status_code: row.get::<_, i64>(14)? as u16, - error_message: row.get(15)?, - created_at: row.get(16)?, + provider_name: row.get(2)?, + app_type: row.get(3)?, + model: row.get(4)?, + input_tokens: row.get::<_, i64>(5)? as u32, + output_tokens: row.get::<_, i64>(6)? as u32, + cache_read_tokens: row.get::<_, i64>(7)? as u32, + cache_creation_tokens: row.get::<_, i64>(8)? as u32, + input_cost_usd: row.get(9)?, + output_cost_usd: row.get(10)?, + cache_read_cost_usd: row.get(11)?, + cache_creation_cost_usd: row.get(12)?, + total_cost_usd: row.get(13)?, + latency_ms: row.get::<_, i64>(14)? as u64, + status_code: row.get::<_, i64>(15)? as u16, + error_message: row.get(16)?, + created_at: row.get(17)?, }) }, ); match result { - Ok(detail) => Ok(Some(detail)), + Ok(mut detail) => { + let mut provider_cache = HashMap::new(); + let mut pricing_cache = HashMap::new(); + Self::maybe_backfill_log_costs( + &conn, + &mut detail, + &mut provider_cache, + &mut pricing_cache, + )?; + Ok(Some(detail)) + } Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), Err(e) => Err(AppError::Database(e.to_string())), } @@ -398,55 +544,68 @@ impl Database { let conn = lock_conn!(self.conn); // 获取 provider 的限额设置 - let (limit_daily, limit_monthly) = conn.query_row( - "SELECT meta FROM providers WHERE id = ? AND app_type = ?", - params![provider_id, app_type], - |row| { - let meta_str: String = row.get(0)?; - Ok(meta_str) - }, - ).ok().and_then(|meta_str| { - serde_json::from_str::(&meta_str).ok() - }).and_then(|meta| { - let daily = meta.get("limitDailyUsd") - .and_then(|v| v.as_str()) - .and_then(|s| s.parse::().ok()); - let monthly = meta.get("limitMonthlyUsd") - .and_then(|v| v.as_str()) - .and_then(|s| s.parse::().ok()); - Some((daily, monthly)) - }).unwrap_or((None, None)); + let (limit_daily, limit_monthly) = conn + .query_row( + "SELECT meta FROM providers WHERE id = ? AND app_type = ?", + params![provider_id, app_type], + |row| { + let meta_str: String = row.get(0)?; + Ok(meta_str) + }, + ) + .ok() + .and_then(|meta_str| serde_json::from_str::(&meta_str).ok()) + .map(|meta| { + let daily = meta + .get("limitDailyUsd") + .and_then(|v| v.as_str()) + .and_then(|s| s.parse::().ok()); + let monthly = meta + .get("limitMonthlyUsd") + .and_then(|v| v.as_str()) + .and_then(|s| s.parse::().ok()); + (daily, monthly) + }) + .unwrap_or((None, None)); // 计算今日使用量 - let daily_usage: f64 = conn.query_row( - "SELECT COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) + let daily_usage: f64 = conn + .query_row( + "SELECT COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) FROM proxy_request_logs WHERE provider_id = ? AND app_type = ? AND date(created_at, 'unixepoch') = date('now')", - params![provider_id, app_type], - |row| row.get(0), - ).unwrap_or(0.0); + params![provider_id, app_type], + |row| row.get(0), + ) + .unwrap_or(0.0); // 计算本月使用量 - let monthly_usage: f64 = conn.query_row( - "SELECT COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) + let monthly_usage: f64 = conn + .query_row( + "SELECT COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) FROM proxy_request_logs WHERE provider_id = ? AND app_type = ? AND strftime('%Y-%m', created_at, 'unixepoch') = strftime('%Y-%m', 'now')", - params![provider_id, app_type], - |row| row.get(0), - ).unwrap_or(0.0); + params![provider_id, app_type], + |row| row.get(0), + ) + .unwrap_or(0.0); - let daily_exceeded = limit_daily.map(|limit| daily_usage >= limit).unwrap_or(false); - let monthly_exceeded = limit_monthly.map(|limit| monthly_usage >= limit).unwrap_or(false); + let daily_exceeded = limit_daily + .map(|limit| daily_usage >= limit) + .unwrap_or(false); + let monthly_exceeded = limit_monthly + .map(|limit| monthly_usage >= limit) + .unwrap_or(false); Ok(ProviderLimitStatus { provider_id: provider_id.to_string(), - daily_usage: format!("{:.6}", daily_usage), - daily_limit: limit_daily.map(|l| format!("{:.2}", l)), + daily_usage: format!("{daily_usage:.6}"), + daily_limit: limit_daily.map(|l| format!("{l:.2}")), daily_exceeded, - monthly_usage: format!("{:.6}", monthly_usage), - monthly_limit: limit_monthly.map(|l| format!("{:.2}", l)), + monthly_usage: format!("{monthly_usage:.6}"), + monthly_limit: limit_monthly.map(|l| format!("{l:.2}")), monthly_exceeded, }) } @@ -454,6 +613,7 @@ impl Database { /// Provider 限额状态 #[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] pub struct ProviderLimitStatus { pub provider_id: String, pub daily_usage: String, @@ -464,6 +624,232 @@ pub struct ProviderLimitStatus { pub monthly_exceeded: bool, } +#[derive(Clone)] +struct PricingInfo { + input: rust_decimal::Decimal, + output: rust_decimal::Decimal, + cache_read: rust_decimal::Decimal, + cache_creation: rust_decimal::Decimal, +} + +impl Database { + fn maybe_backfill_log_costs( + conn: &Connection, + log: &mut RequestLogDetail, + provider_cache: &mut HashMap<(String, String), rust_decimal::Decimal>, + pricing_cache: &mut HashMap, + ) -> Result<(), AppError> { + let total_cost = rust_decimal::Decimal::from_str(&log.total_cost_usd) + .unwrap_or(rust_decimal::Decimal::ZERO); + let has_cost = total_cost > rust_decimal::Decimal::ZERO; + let has_usage = log.input_tokens > 0 + || log.output_tokens > 0 + || log.cache_read_tokens > 0 + || log.cache_creation_tokens > 0; + + if has_cost || !has_usage { + return Ok(()); + } + + let pricing = match Self::get_model_pricing_cached(conn, pricing_cache, &log.model)? { + Some(info) => info, + None => return Ok(()), + }; + let multiplier = Self::get_cost_multiplier_cached( + conn, + provider_cache, + &log.provider_id, + &log.app_type, + )?; + + let million = rust_decimal::Decimal::from(1_000_000u64); + let input_cost = rust_decimal::Decimal::from(log.input_tokens as u64) * pricing.input + / million + * multiplier; + let output_cost = rust_decimal::Decimal::from(log.output_tokens as u64) * pricing.output + / million + * multiplier; + let cache_read_cost = rust_decimal::Decimal::from(log.cache_read_tokens as u64) + * pricing.cache_read + / million + * multiplier; + let cache_creation_cost = rust_decimal::Decimal::from(log.cache_creation_tokens as u64) + * pricing.cache_creation + / million + * multiplier; + let total_cost = input_cost + output_cost + cache_read_cost + cache_creation_cost; + + log.input_cost_usd = format!("{input_cost:.6}"); + log.output_cost_usd = format!("{output_cost:.6}"); + log.cache_read_cost_usd = format!("{cache_read_cost:.6}"); + log.cache_creation_cost_usd = format!("{cache_creation_cost:.6}"); + log.total_cost_usd = format!("{total_cost:.6}"); + + conn.execute( + "UPDATE proxy_request_logs + SET input_cost_usd = ?1, + output_cost_usd = ?2, + cache_read_cost_usd = ?3, + cache_creation_cost_usd = ?4, + total_cost_usd = ?5 + WHERE request_id = ?6", + params![ + log.input_cost_usd, + log.output_cost_usd, + log.cache_read_cost_usd, + log.cache_creation_cost_usd, + log.total_cost_usd, + log.request_id + ], + ) + .map_err(|e| AppError::Database(format!("更新请求成本失败: {e}")))?; + + Ok(()) + } + + fn get_cost_multiplier_cached( + conn: &Connection, + cache: &mut HashMap<(String, String), rust_decimal::Decimal>, + provider_id: &str, + app_type: &str, + ) -> Result { + let key = (provider_id.to_string(), app_type.to_string()); + if let Some(multiplier) = cache.get(&key) { + return Ok(*multiplier); + } + + let meta_json: Option = conn + .query_row( + "SELECT meta FROM providers WHERE id = ? AND app_type = ?", + params![provider_id, app_type], + |row| row.get(0), + ) + .optional() + .map_err(|e| AppError::Database(format!("查询 provider meta 失败: {e}")))?; + + let multiplier = meta_json + .and_then(|meta| serde_json::from_str::(&meta).ok()) + .and_then(|value| value.get("costMultiplier").cloned()) + .and_then(|val| { + val.as_str() + .and_then(|s| rust_decimal::Decimal::from_str(s).ok()) + }) + .unwrap_or(rust_decimal::Decimal::ONE); + + cache.insert(key, multiplier); + Ok(multiplier) + } + + fn get_model_pricing_cached( + conn: &Connection, + cache: &mut HashMap, + model: &str, + ) -> Result, AppError> { + if let Some(info) = cache.get(model) { + return Ok(Some(info.clone())); + } + + let row = find_model_pricing_row(conn, model)?; + let Some((input, output, cache_read, cache_creation)) = row else { + return Ok(None); + }; + + let pricing = PricingInfo { + input: rust_decimal::Decimal::from_str(&input) + .map_err(|e| AppError::Database(format!("解析输入价格失败: {e}")))?, + output: rust_decimal::Decimal::from_str(&output) + .map_err(|e| AppError::Database(format!("解析输出价格失败: {e}")))?, + cache_read: rust_decimal::Decimal::from_str(&cache_read) + .map_err(|e| AppError::Database(format!("解析缓存读取价格失败: {e}")))?, + cache_creation: rust_decimal::Decimal::from_str(&cache_creation) + .map_err(|e| AppError::Database(format!("解析缓存写入价格失败: {e}")))?, + }; + + cache.insert(model.to_string(), pricing.clone()); + Ok(Some(pricing)) + } +} + +pub(crate) fn find_model_pricing_row( + conn: &Connection, + model_id: &str, +) -> Result, AppError> { + let exact = conn + .query_row( + "SELECT input_cost_per_million, output_cost_per_million, + cache_read_cost_per_million, cache_creation_cost_per_million + FROM model_pricing + WHERE model_id = ?1", + [model_id], + |row| { + Ok(( + row.get::<_, String>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + row.get::<_, String>(3)?, + )) + }, + ) + .optional() + .map_err(|e| AppError::Database(format!("查询模型定价失败: {e}")))?; + + if exact.is_some() { + return Ok(exact); + } + + let mut stmt = conn.prepare( + "SELECT model_id, input_cost_per_million, output_cost_per_million, + cache_read_cost_per_million, cache_creation_cost_per_million + FROM model_pricing + WHERE model_id LIKE '%*%' OR model_id LIKE '%?%' + ORDER BY LENGTH(model_id) DESC", + )?; + + let rows = stmt + .query_map([], |row| { + Ok(( + row.get::<_, String>(0)?, + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + row.get::<_, String>(3)?, + row.get::<_, String>(4)?, + )) + })? + .filter_map(|res| res.ok()); + + for (pattern, input, output, cache_read, cache_creation) in rows { + if matches_model_pattern(&pattern, model_id) { + return Ok(Some((input, output, cache_read, cache_creation))); + } + } + + Ok(None) +} + +fn matches_model_pattern(pattern: &str, model_id: &str) -> bool { + if pattern == model_id { + return true; + } + + if !pattern.contains('*') && !pattern.contains('?') { + return false; + } + + let mut regex_pattern = String::from("^"); + for ch in pattern.chars() { + match ch { + '*' => regex_pattern.push_str(".*"), + '?' => regex_pattern.push('.'), + _ => regex_pattern.push_str(®ex::escape(&ch.to_string())), + } + } + regex_pattern.push('$'); + + Regex::new(®ex_pattern) + .map(|re| re.is_match(model_id)) + .unwrap_or(false) +} + #[cfg(test)] mod tests { use super::*; @@ -513,7 +899,18 @@ mod tests { input_tokens, output_tokens, total_cost_usd, latency_ms, status_code, created_at ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", - params!["req1", "p1", "claude", "claude-3-sonnet", 100, 50, "0.01", 100, 200, 1000], + params![ + "req1", + "p1", + "claude", + "claude-3-sonnet", + 100, + 50, + "0.01", + 100, + 200, + 1000 + ], )?; }