Files
CC-Switch/src-tauri/src/services/usage_stats.rs
T
YoVinchen b1103c8a59 Feat/proxy server (#355)
* feat(proxy): implement local HTTP proxy server with multi-provider failover

Add a complete HTTP proxy server implementation built on Axum framework,
enabling local API request forwarding with automatic provider failover
and load balancing capabilities.

Backend Implementation (Rust):
- Add proxy server module with 7 core components:
  * server.rs: Axum HTTP server lifecycle management (start/stop/status)
  * router.rs: API routing configuration for Claude/OpenAI/Gemini endpoints
  * handlers.rs: Request/response handling and transformation
  * forwarder.rs: Upstream forwarding logic with retry mechanism (652 lines)
  * error.rs: Comprehensive error handling and HTTP status mapping
  * types.rs: Shared types (ProxyConfig, ProxyStatus, ProxyServerInfo)
  * health.rs: Provider health check infrastructure

Service Layer:
- Add ProxyService (services/proxy.rs, 157 lines):
  * Manage proxy server lifecycle
  * Handle configuration updates
  * Track runtime status and metrics

Database Layer:
- Add proxy configuration DAO (dao/proxy.rs, 242 lines):
  * Persist proxy settings (listen address, port, timeout)
  * Store provider priority and availability flags
- Update schema with proxy_config table (schema.rs):
  * Support runtime configuration persistence

Tauri Commands:
- Add 6 command endpoints (commands/proxy.rs):
  * start_proxy_server: Launch proxy server
  * stop_proxy_server: Gracefully shutdown server
  * get_proxy_status: Query runtime status
  * get_proxy_config: Retrieve current configuration
  * update_proxy_config: Modify settings without restart
  * is_proxy_running: Check server state

Frontend Implementation (React + TypeScript):
- Add ProxyPanel component (222 lines):
  * Real-time server status display
  * Start/stop controls
  * Provider availability monitoring
- Add ProxySettingsDialog component (420 lines):
  * Configuration editor (address, port, timeout)
  * Provider priority management
  * Settings validation
- Add React hooks:
  * useProxyConfig: Manage proxy configuration state
  * useProxyStatus: Poll and display server status
- Add TypeScript types (types/proxy.ts):
  * Define ProxyConfig, ProxyStatus interfaces

Provider Integration:
- Extend Provider model with availability field (providers.rs):
  * Track provider health for failover logic
- Update ProviderCard UI to display proxy status
- Integrate proxy controls in Settings page

Dependencies:
- Add Axum 0.7 (async web framework)
- Add Tower 0.4 (middleware and service abstractions)
- Add Tower-HTTP (CORS layer)
- Add Tokio sync primitives (oneshot, RwLock)

Technical Details:
- Graceful shutdown via oneshot channel
- Shared state with Arc<RwLock<T>> for thread-safe config updates
- CORS enabled for cross-origin frontend access
- Request/response streaming support
- Automatic retry with exponential backoff (forwarder)
- API key extraction from multiple config formats (Claude/Codex/Gemini)

File Statistics:
- 41 files changed
- 3491 insertions(+), 41 deletions(-)
- Core modules: 1393 lines (server + forwarder + handlers)
- Frontend UI: 642 lines (ProxyPanel + ProxySettingsDialog)
- Database/DAO: 326 lines

This implementation provides the foundation for advanced features like:
- Multi-provider load balancing
- Automatic failover on provider errors
- Request logging and analytics
- Usage tracking and cost monitoring

* fix(proxy): resolve UI/UX issues and database constraint error

Simplify proxy control interface and fix database persistence issues:

Backend Fixes:
- Fix NOT NULL constraint error in proxy_config.created_at field
  * Use COALESCE to preserve created_at on updates
  * Ensure proper INSERT OR REPLACE behavior
- Remove redundant enabled field validation on startup
  * Auto-enable when user clicks start button
  * Persist enabled state after successful start
- Preserve enabled state during config updates
  * Prevent accidental service shutdown on config save

Frontend Improvements:
- Remove duplicate proxy enable switch from settings dialog
  * Keep only runtime toggle in ProxyPanel
  * Simplify user experience with single control point
- Hide proxy target button when proxy service is stopped
  * Add isProxyRunning prop to ProviderCard
  * Conditionally render proxy controls based on service status
- Update form schema to omit enabled field
  * Managed automatically by backend

Files: 5 changed, 81 insertions(+), 94 deletions(-)

* fix(proxy): improve URL building and Gemini request handling

- Refactor URL construction with version path deduplication (/v1, /v1beta)
- Preserve query parameters for Gemini API requests
- Support GOOGLE_GEMINI_API_KEY field name (with fallback)
- Change default proxy port from 5000 to 15721
- Fix test: use Option type for is_proxy_target field

* refactor(proxy): remove unused request handlers and routes

- Remove unused GET/DELETE request forwarding methods
- Remove count_tokens, get/delete response handlers
- Simplify router by removing unused endpoints
- Keep only essential routes: /v1/messages, /v1/responses, /v1beta/*

* Merge branch 'main' into feat/proxy-server

* fix(proxy): resolve clippy warnings for dead code and uninlined format args

- Add #[allow(dead_code)] to unused ProviderUnhealthy variant
- Inline format string arguments in handlers.rs and codex.rs log macros
- Refactor error response handling to properly pass through upstream errors
- Add URL deduplication logic for /v1/v1 paths in CodexAdapter

* feat(proxy): implement provider adapter pattern with OpenRouter support

This major refactoring introduces a modular provider adapter architecture
to support format transformation between different AI API formats.

New features:
- Add ProviderAdapter trait for unified provider abstraction
- Implement Claude, Codex, and Gemini adapters with specific logic
- Add Anthropic ↔ OpenAI format transformation for OpenRouter compatibility
- Support model mapping from provider configuration (ANTHROPIC_MODEL, etc.)
- Add OpenRouter preset to Claude provider presets

Refactoring:
- Extract authentication logic into auth.rs with AuthInfo and AuthStrategy
- Move URL building and request transformation to individual adapters
- Simplify ProviderRouter to only use proxy target providers
- Refactor RequestForwarder to use adapter-based request/response handling
- Use whitelist mode for header forwarding (only pass necessary headers)

Architecture:
- providers/adapter.rs: ProviderAdapter trait definition
- providers/auth.rs: AuthInfo, AuthStrategy types
- providers/claude.rs: Claude adapter with OpenRouter detection
- providers/codex.rs: Codex (OpenAI) adapter
- providers/gemini.rs: Gemini (Google) adapter
- providers/models/: Anthropic and OpenAI API data models
- providers/transform.rs: Bidirectional format transformation

* feat(proxy): add streaming SSE transform and thinking parameter support

New features:
- Add OpenAI → Anthropic SSE streaming response transformation
- Support thinking parameter detection for reasoning model selection
- Add ANTHROPIC_REASONING_MODEL config option for extended thinking

Changes:
- streaming.rs: Implement SSE event parsing and Anthropic format conversion
- transform.rs: Add thinking detection logic and reasoning model mapping
- handlers.rs: Integrate streaming transform for OpenRouter compatibility
- Cargo.toml: Add async-stream and bytes dependencies

* feat(db): add usage tracking schema and types

Add database tables for proxy request logs and model pricing.
Extend Provider and error types to support usage statistics.

* 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.

* feat(proxy): integrate usage logging into request handlers

Add usage logging to forwarder and streaming handlers.
Track token usage and costs for each proxy request.

* feat(commands): add usage statistics Tauri commands

Register usage commands for summary, trends, logs, and pricing.
Expose usage stats service through Tauri command layer.

* feat(api): add frontend usage API and query hooks

Add TypeScript types for usage statistics.
Implement usage API with Tauri invoke calls.
Add TanStack Query hooks for usage data fetching.

* feat(ui): add usage dashboard components

Add UsageDashboard with summary cards, trend chart, and data tables.
Implement model pricing configuration panel.
Add request log viewer with filtering and detail panel.

* fix(ui): integrate usage dashboard and fix type errors

Add usage dashboard tab to settings page.
Fix UsageScriptModal TypeScript type annotations.

* deps: add recharts for charts and rust_decimal/uuid for usage tracking

- recharts: Chart visualization for usage trends
- rust_decimal: Precise cost calculations
- uuid: Request ID generation

* feat(proxy): add ProviderType enum for fine-grained provider detection

Introduce ProviderType enum to distinguish between different provider
implementations (Claude, ClaudeAuth, Codex, Gemini, GeminiCli, OpenRouter).
This enables proper authentication handling and request transformation
based on the actual provider type rather than just AppType.

- Add ProviderType enum with detection logic from config
- Enhance Claude adapter with OpenRouter detection
- Enhance Gemini adapter with CLI mode detection
- Add helper methods for provider type inference

* feat(database): extend schema with streaming and timing fields

Add new columns to proxy_request_logs table for enhanced usage tracking:
- first_token_ms and duration_ms for performance metrics
- provider_type and is_streaming for request classification
- cost_multiplier for flexible pricing

Update model pricing with accurate rates for Claude/GPT/Gemini models.
Add ensure_model_pricing_seeded() call on database initialization.
Add test for model pricing auto-seeding verification.

* feat(proxy/usage): enhance token parser and logger for multi-format support

Parser enhancements:
- Add OpenAI Chat Completions format parsing (prompt_tokens/completion_tokens)
- Add model field to TokenUsage for actual model name extraction
- Add from_codex_response_adjusted() for proper cache token handling
- Add debug logging for better stream event tracing

Logger enhancements:
- Add first_token_ms, provider_type, is_streaming, cost_multiplier fields
- Extend RequestLog struct with full metadata tracking
- Update log_with_calculation() signature for new fields

Calculator: Update tests with model field in TokenUsage.

* feat(proxy): enhance proxy server with session tracking and OpenAI route

Error handling:
- Add StreamIdleTimeout and AuthError variants for better error classification

Module exports:
- Export ResponseType, StreamHandler, NonStreamHandler from response_handler
- Export ProxySession, ClientFormat from session module

Server routing:
- Add /v1/chat/completions route for OpenAI Chat Completions API

Handlers:
- Add log_usage_with_session() for enhanced usage tracking with session context
- Add first_token_ms timing measurement for streaming responses
- Use SseUsageCollector with start_time for accurate latency calculation
- Track is_streaming flag in usage logs

* feat(services): add pagination and enhanced filtering for request logs

Usage stats service:
- Change get_request_logs() from limit/offset to page/page_size pagination
- Return PaginatedLogs with total count, page, and page_size
- Add appType and providerName filters with LIKE search
- Add is_streaming, first_token_ms, duration_ms to RequestLogDetail
- Join with providers table for provider name lookup

Commands:
- Update get_request_logs command signature for pagination params

Module exports:
- Export PaginatedLogs struct

* feat(frontend): update usage types and API for pagination support

Types (usage.ts):
- Add isStreaming, firstTokenMs, durationMs to RequestLog
- Add PaginatedLogs interface with data, total, page, pageSize
- Change LogFilters: providerId -> appType + providerName

API (usage.ts):
- Change getRequestLogs params from limit/offset to page/pageSize
- Return PaginatedLogs instead of RequestLog[]
- Pass filters object directly to backend

Query (usage.ts):
- Update usageKeys.logs key generation for pagination
- Update useRequestLogs hook signature

* refactor(ui): enhance RequestLogTable with filtering and pagination

UI improvements:
- Add filter bar with app type, provider name, model, status selectors
- Add date range picker (startDate/endDate)
- Add search/reset/refresh buttons

Pagination:
- Implement proper page-based pagination with page info display
- Show total count and current page range
- Add prev/next navigation buttons

Features:
- Default to last 24 hours filter
- Streamlined table columns layout
- Query invalidation on refresh

* style(config): format mcpPresets code style

Apply consistent formatting to createNpxCommand function and
sequential-thinking server configuration.

* fix(ui): update SettingsPage tab styles for improved appearance (#342)

* feat(model-test): add provider model availability testing

Implement standalone model testing feature to verify provider API connectivity:
- Add ModelTestService for Claude/Codex/Gemini endpoint testing
- Create model_test_logs table for test result persistence
- Add test button to ProviderCard with loading state
- Include ModelTestConfigPanel for customizing test parameters

* fix(proxy): resolve token parsing for OpenRouter streaming responses

Problem:
- OpenRouter and similar third-party services return streaming responses
  where input_tokens appear in message_delta instead of message_start
- The previous implementation only extracted input_tokens from message_start,
  causing input_tokens to be recorded as 0 for these providers

Changes:
- streaming.rs: Add prompt_tokens field to Usage struct and include
  input_tokens in the transformed message_delta event when converting
  OpenAI format to Anthropic format
- parser.rs: Update from_claude_stream_events() to handle input_tokens
  from both message_start (native Claude API) and message_delta (OpenRouter)
  - Use if-let pattern instead of direct unwrap for safer parsing
  - Only update input_tokens from message_delta if not already set
- logger.rs: Adjust test parameters to match updated function signature

Tests:
- Add test_openrouter_stream_parsing() for OpenRouter format validation
- Add test_native_claude_stream_parsing() for native Claude API validation

* fix(pricing): standardize model ID format for pricing lookup

Normalize model IDs by removing vendor prefixes and converting dots to hyphens to ensure consistent pricing lookups across different API response formats.

Changes:
- Update seed data to use hyphen format (e.g., gpt-5-1, gemini-2-5-pro)
- Add normalize_model_id() function to strip vendor prefixes (anthropic/, openai/)
- Convert dots to hyphens in model IDs (claude-haiku-4.5 → claude-haiku-4-5)
- Try both original and normalized IDs for exact matching
- Use normalized ID for suffix-based fallback matching
- Add comprehensive test cases for prefix and dot handling
- Add warning log when no pricing found

This ensures pricing lookups work correctly for:
- Models with vendor prefixes: anthropic/claude-haiku-4.5
- Models with dots in version: claude-sonnet-4.5
- Models with date suffixes: claude-haiku-4-5-20240229

* style(rust): apply clippy formatting suggestions

Apply automatic clippy fixes for uninlined_format_args warnings across Rust codebase. Replace format string placeholders with inline variable syntax for improved readability.

Changes:
- Convert format!("{}", var) to format!("{var}")
- Apply to model_test.rs, parser.rs, and usage_stats.rs
- Fix line length issues by breaking long function calls
- Improve code formatting consistency

All changes are automatic formatting with no functional impact.

* fix(ui): restore card borders in usage statistics panels

Restore proper card styling for ModelTestConfigPanel and PricingConfigPanel by adding back border and rounded-lg classes. The transparent background styling was causing visual inconsistency.

Changes:
- Replace border-none bg-transparent shadow-none with border rounded-lg
- Apply to both loading and error states for consistency
- Format TypeScript code for better readability
- Break long function signatures across multiple lines

This ensures the usage statistics panels have consistent visual appearance with proper borders and rounded corners.

* feat(pricing): add GPT-5 Codex model pricing presets

Add pricing configuration for GPT-5 Codex variants to support cost tracking for Codex-specific models.

Changes:
- Add gpt-5-codex model with standard GPT-5 pricing
- Add gpt-5-1-codex model with standard GPT-5.1 pricing
- Input: $1.25/M tokens, Output: $10/M tokens
- Cache read: $0.125/M tokens, Cache creation: $0

This ensures accurate cost calculation for Codex API requests using GPT-5 Codex models.
2025-12-05 11:26:41 +08:00

1085 lines
40 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 使用统计服务
//!
//! 提供使用量数据的聚合查询功能
use crate::database::{lock_conn, Database};
use crate::error::AppError;
use chrono::{Duration, Utc};
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,
pub request_count: u64,
pub total_tokens: u64,
pub total_cost: String,
pub success_rate: f32,
pub avg_latency_ms: u64,
}
/// 模型统计
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ModelStats {
pub model: String,
pub request_count: u64,
pub total_tokens: u64,
pub total_cost: String,
pub avg_cost_per_request: String,
}
/// 请求日志过滤器
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LogFilters {
pub app_type: Option<String>,
pub provider_name: Option<String>,
pub model: Option<String>,
pub status_code: Option<u16>,
pub start_date: Option<i64>,
pub end_date: Option<i64>,
}
/// 分页请求日志响应
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PaginatedLogs {
pub data: Vec<RequestLogDetail>,
pub total: u32,
pub page: u32,
pub page_size: u32,
}
/// 请求日志详情
#[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<String>,
pub app_type: String,
pub model: String,
pub input_tokens: u32,
pub output_tokens: u32,
pub cache_read_tokens: u32,
pub cache_creation_tokens: u32,
pub input_cost_usd: String,
pub output_cost_usd: String,
pub cache_read_cost_usd: String,
pub cache_creation_cost_usd: String,
pub total_cost_usd: String,
pub is_streaming: bool,
pub latency_ms: u64,
pub first_token_ms: Option<u64>,
pub duration_ms: Option<u64>,
pub status_code: u16,
pub error_message: Option<String>,
pub created_at: i64,
}
impl Database {
/// 获取使用量汇总
pub fn get_usage_summary(
&self,
start_date: Option<i64>,
end_date: Option<i64>,
) -> Result<UsageSummary, AppError> {
let conn = lock_conn!(self.conn);
let (where_clause, params_vec) = if start_date.is_some() || end_date.is_some() {
let mut conditions = Vec::new();
let mut params = Vec::new();
if let Some(start) = start_date {
conditions.push("created_at >= ?");
params.push(start);
}
if let Some(end) = end_date {
conditions.push("created_at <= ?");
params.push(end);
}
(format!("WHERE {}", conditions.join(" AND ")), params)
} else {
(String::new(), Vec::new())
};
let sql = format!(
"SELECT
COUNT(*) as total_requests,
COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost,
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}"
);
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_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
} else {
0.0
};
Ok(UsageSummary {
total_requests: total_requests as u64,
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,
})
})?;
Ok(result)
}
/// 获取每日趋势
pub fn get_daily_trends(&self, days: u32) -> Result<Vec<DailyStats>, AppError> {
let conn = lock_conn!(self.conn);
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([], |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 buckets: HashMap<String, DailyStats> = 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)
}
}
/// 获取 Provider 统计
pub fn get_provider_stats(&self) -> Result<Vec<ProviderStats>, AppError> {
let conn = lock_conn!(self.conn);
let sql = "SELECT
l.provider_id,
p.name as provider_name,
COUNT(*) as request_count,
COALESCE(SUM(l.input_tokens + l.output_tokens), 0) as total_tokens,
COALESCE(SUM(CAST(l.total_cost_usd AS REAL)), 0) as total_cost,
COALESCE(SUM(CASE WHEN l.status_code >= 200 AND l.status_code < 300 THEN 1 ELSE 0 END), 0) as success_count,
COALESCE(AVG(l.latency_ms), 0) as avg_latency
FROM proxy_request_logs l
LEFT JOIN providers p ON l.provider_id = p.id AND l.app_type = p.app_type
GROUP BY l.provider_id, l.app_type
ORDER BY total_cost DESC";
let mut stmt = conn.prepare(sql)?;
let rows = stmt.query_map([], |row| {
let request_count: i64 = row.get(2)?;
let success_count: i64 = row.get(5)?;
let success_rate = if request_count > 0 {
(success_count as f32 / request_count as f32) * 100.0
} else {
0.0
};
Ok(ProviderStats {
provider_id: row.get(0)?,
provider_name: row
.get::<_, Option<String>>(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)?),
success_rate,
avg_latency_ms: row.get::<_, f64>(6)? as u64,
})
})?;
let mut stats = Vec::new();
for row in rows {
stats.push(row?);
}
Ok(stats)
}
/// 获取模型统计
pub fn get_model_stats(&self) -> Result<Vec<ModelStats>, AppError> {
let conn = lock_conn!(self.conn);
let sql = "SELECT
model,
COUNT(*) as request_count,
COALESCE(SUM(input_tokens + output_tokens), 0) as total_tokens,
COALESCE(SUM(CAST(total_cost_usd AS REAL)), 0) as total_cost
FROM proxy_request_logs
GROUP BY model
ORDER BY total_cost DESC";
let mut stmt = conn.prepare(sql)?;
let rows = stmt.query_map([], |row| {
let request_count: i64 = row.get(1)?;
let total_cost: f64 = row.get(3)?;
let avg_cost = if request_count > 0 {
total_cost / request_count as f64
} else {
0.0
};
Ok(ModelStats {
model: row.get(0)?,
request_count: request_count as u64,
total_tokens: row.get::<_, i64>(2)? as u64,
total_cost: format!("{total_cost:.6}"),
avg_cost_per_request: format!("{avg_cost:.6}"),
})
})?;
let mut stats = Vec::new();
for row in rows {
stats.push(row?);
}
Ok(stats)
}
/// 获取请求日志列表(分页)
pub fn get_request_logs(
&self,
filters: &LogFilters,
page: u32,
page_size: u32,
) -> Result<PaginatedLogs, AppError> {
let conn = lock_conn!(self.conn);
let mut conditions = Vec::new();
let mut params: Vec<Box<dyn rusqlite::ToSql>> = Vec::new();
if let Some(ref app_type) = filters.app_type {
conditions.push("l.app_type = ?");
params.push(Box::new(app_type.clone()));
}
if let Some(ref provider_name) = filters.provider_name {
conditions.push("p.name LIKE ?");
params.push(Box::new(format!("%{provider_name}%")));
}
if let Some(ref model) = filters.model {
conditions.push("l.model LIKE ?");
params.push(Box::new(format!("%{model}%")));
}
if let Some(status) = filters.status_code {
conditions.push("l.status_code = ?");
params.push(Box::new(status as i64));
}
if let Some(start) = filters.start_date {
conditions.push("l.created_at >= ?");
params.push(Box::new(start));
}
if let Some(end) = filters.end_date {
conditions.push("l.created_at <= ?");
params.push(Box::new(end));
}
let where_clause = if conditions.is_empty() {
String::new()
} else {
format!("WHERE {}", conditions.join(" AND "))
};
// 获取总数
let count_sql = format!(
"SELECT COUNT(*) FROM proxy_request_logs l
LEFT JOIN providers p ON l.provider_id = p.id AND l.app_type = p.app_type
{where_clause}"
);
let count_params: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p.as_ref()).collect();
let total: u32 = conn.query_row(&count_sql, count_params.as_slice(), |row| {
row.get::<_, i64>(0).map(|v| v as u32)
})?;
// 获取数据
let offset = page * page_size;
params.push(Box::new(page_size as i64));
params.push(Box::new(offset as i64));
let sql = format!(
"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.is_streaming, l.latency_ms, l.first_token_ms, l.duration_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 l.created_at DESC
LIMIT ? OFFSET ?"
);
let mut stmt = conn.prepare(&sql)?;
let params_refs: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p.as_ref()).collect();
let rows = stmt.query_map(params_refs.as_slice(), |row| {
Ok(RequestLogDetail {
request_id: row.get(0)?,
provider_id: row.get(1)?,
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)?,
is_streaming: row.get::<_, i64>(14)? != 0,
latency_ms: row.get::<_, i64>(15)? as u64,
first_token_ms: row.get::<_, Option<i64>>(16)?.map(|v| v as u64),
duration_ms: row.get::<_, Option<i64>>(17)?.map(|v| v as u64),
status_code: row.get::<_, i64>(18)? as u16,
error_message: row.get(19)?,
created_at: row.get(20)?,
})
})?;
let mut logs = Vec::new();
let mut provider_cache = HashMap::new();
let mut pricing_cache = HashMap::new();
for row in rows {
let mut log = row?;
Self::maybe_backfill_log_costs(
&conn,
&mut log,
&mut provider_cache,
&mut pricing_cache,
)?;
logs.push(log);
}
Ok(PaginatedLogs {
data: logs,
total,
page,
page_size,
})
}
/// 获取单个请求详情
pub fn get_request_detail(
&self,
request_id: &str,
) -> Result<Option<RequestLogDetail>, AppError> {
let conn = lock_conn!(self.conn);
let result = conn.query_row(
"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,
is_streaming, latency_ms, first_token_ms, duration_ms,
status_code, error_message, 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 l.request_id = ?",
[request_id],
|row| {
Ok(RequestLogDetail {
request_id: row.get(0)?,
provider_id: row.get(1)?,
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)?,
is_streaming: row.get::<_, i64>(14)? != 0,
latency_ms: row.get::<_, i64>(15)? as u64,
first_token_ms: row.get::<_, Option<i64>>(16)?.map(|v| v as u64),
duration_ms: row.get::<_, Option<i64>>(17)?.map(|v| v as u64),
status_code: row.get::<_, i64>(18)? as u16,
error_message: row.get(19)?,
created_at: row.get(20)?,
})
},
);
match result {
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())),
}
}
/// 检查 Provider 使用限额
pub fn check_provider_limits(
&self,
provider_id: &str,
app_type: &str,
) -> Result<ProviderLimitStatus, AppError> {
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::<serde_json::Value>(&meta_str).ok())
.map(|meta| {
let daily = meta
.get("limitDailyUsd")
.and_then(|v| v.as_str())
.and_then(|s| s.parse::<f64>().ok());
let monthly = meta
.get("limitMonthlyUsd")
.and_then(|v| v.as_str())
.and_then(|s| s.parse::<f64>().ok());
(daily, monthly)
})
.unwrap_or((None, None));
// 计算今日使用量
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);
// 计算本月使用量
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);
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!("{daily_usage:.6}"),
daily_limit: limit_daily.map(|l| format!("{l:.2}")),
daily_exceeded,
monthly_usage: format!("{monthly_usage:.6}"),
monthly_limit: limit_monthly.map(|l| format!("{l:.2}")),
monthly_exceeded,
})
}
/// 更新每日统计聚合
///
/// 在请求完成后调用,更新 usage_daily_stats 表
#[allow(clippy::too_many_arguments)]
pub fn update_daily_stats(
&self,
provider_id: &str,
app_type: &str,
model: &str,
input_tokens: u32,
output_tokens: u32,
total_cost: &str,
is_success: bool,
) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
let date = Utc::now().format("%Y-%m-%d").to_string();
// 使用 UPSERT 更新或插入统计
conn.execute(
"INSERT INTO usage_daily_stats (
date, provider_id, app_type, model,
request_count, total_input_tokens, total_output_tokens,
total_cost_usd, success_count, error_count
) VALUES (?1, ?2, ?3, ?4, 1, ?5, ?6, ?7, ?8, ?9)
ON CONFLICT(date, provider_id, app_type, model) DO UPDATE SET
request_count = request_count + 1,
total_input_tokens = total_input_tokens + ?5,
total_output_tokens = total_output_tokens + ?6,
total_cost_usd = CAST(
CAST(total_cost_usd AS REAL) + CAST(?7 AS REAL) AS TEXT
),
success_count = success_count + ?8,
error_count = error_count + ?9",
params![
date,
provider_id,
app_type,
model,
input_tokens,
output_tokens,
total_cost,
if is_success { 1 } else { 0 },
if is_success { 0 } else { 1 },
],
)
.map_err(|e| AppError::Database(format!("更新每日统计失败: {e}")))?;
Ok(())
}
}
/// Provider 限额状态
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ProviderLimitStatus {
pub provider_id: String,
pub daily_usage: String,
pub daily_limit: Option<String>,
pub daily_exceeded: bool,
pub monthly_usage: String,
pub monthly_limit: Option<String>,
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<String, PricingInfo>,
) -> 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<rust_decimal::Decimal, AppError> {
let key = (provider_id.to_string(), app_type.to_string());
if let Some(multiplier) = cache.get(&key) {
return Ok(*multiplier);
}
let meta_json: Option<String> = 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::<Value>(&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<String, PricingInfo>,
model: &str,
) -> Result<Option<PricingInfo>, 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))
}
}
/// 标准化模型名称:去除供应商前缀并将点号替换为短横线
/// 例如:anthropic/claude-haiku-4.5 → claude-haiku-4-5
fn normalize_model_id(model_id: &str) -> String {
// 1. 去除供应商前缀(如 anthropic/、openai/
let stripped = if let Some(pos) = model_id.find('/') {
&model_id[pos + 1..]
} else {
model_id
};
// 2. 将点号替换为短横线(如 claude-haiku-4.5 → claude-haiku-4-5
stripped.replace('.', "-")
}
pub(crate) fn find_model_pricing_row(
conn: &Connection,
model_id: &str,
) -> Result<Option<(String, String, String, String)>, AppError> {
// 0. 标准化模型名称(去除前缀 + 点号转短横线)
// 例如:anthropic/claude-haiku-4.5 → claude-haiku-4-5
let normalized = normalize_model_id(model_id);
// 1. 精确匹配(先尝试原始名称,再尝试标准化后的名称)
for id in [model_id, normalized.as_str()] {
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",
[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() {
if id != model_id {
log::info!("模型 {model_id} 标准化后精确匹配到: {id}");
}
return Ok(exact);
}
}
// 2. 逐步删除后缀匹配(claude-haiku-4-5-20250929 → claude-haiku-4-5 → claude-haiku-4 → claude-haiku
// 使用标准化后的名称进行后缀匹配
let mut current = normalized;
while let Some(pos) = current.rfind('-') {
current = current[..pos].to_string();
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",
[&current],
|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 result.is_some() {
log::info!("模型 {model_id} 通过删除后缀匹配到: {current}");
return Ok(result);
}
}
log::warn!("模型 {model_id} 未找到定价信息,成本将记录为 0");
Ok(None)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_get_usage_summary() -> Result<(), AppError> {
let db = Database::memory()?;
// 插入测试数据
{
let conn = lock_conn!(db.conn);
conn.execute(
"INSERT INTO proxy_request_logs (
request_id, provider_id, app_type, model,
input_tokens, output_tokens, total_cost_usd,
latency_ms, status_code, created_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
params!["req1", "p1", "claude", "claude-3", 100, 50, "0.01", 100, 200, 1000],
)?;
conn.execute(
"INSERT INTO proxy_request_logs (
request_id, provider_id, app_type, model,
input_tokens, output_tokens, total_cost_usd,
latency_ms, status_code, created_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
params!["req2", "p1", "claude", "claude-3", 200, 100, "0.02", 150, 200, 2000],
)?;
}
let summary = db.get_usage_summary(None, None)?;
assert_eq!(summary.total_requests, 2);
assert_eq!(summary.success_rate, 100.0);
Ok(())
}
#[test]
fn test_get_model_stats() -> Result<(), AppError> {
let db = Database::memory()?;
// 插入测试数据
{
let conn = lock_conn!(db.conn);
conn.execute(
"INSERT INTO proxy_request_logs (
request_id, provider_id, app_type, model,
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
],
)?;
}
let stats = db.get_model_stats()?;
assert_eq!(stats.len(), 1);
assert_eq!(stats[0].model, "claude-3-sonnet");
assert_eq!(stats[0].request_count, 1);
Ok(())
}
#[test]
fn test_model_pricing_matching() -> Result<(), AppError> {
let db = Database::memory()?;
let conn = lock_conn!(db.conn);
// 测试精确匹配
let result = find_model_pricing_row(&conn, "claude-sonnet-4-5")?;
assert!(result.is_some(), "应该能精确匹配 claude-sonnet-4-5");
// 测试带供应商前缀的模型名称(anthropic/claude-haiku-4.5 → claude-haiku-4-5
let result = find_model_pricing_row(&conn, "anthropic/claude-haiku-4.5")?;
assert!(
result.is_some(),
"应该能匹配带前缀的模型 anthropic/claude-haiku-4.5"
);
// 测试带供应商前缀 + 点号的模型名称
let result = find_model_pricing_row(&conn, "anthropic/claude-sonnet-4.5")?;
assert!(
result.is_some(),
"应该能匹配带前缀的模型 anthropic/claude-sonnet-4.5"
);
// 测试逐步删除后缀匹配 - 日期后缀
let result = find_model_pricing_row(&conn, "claude-sonnet-4-5-20241022")?;
assert!(
result.is_some(),
"应该能通过删除后缀匹配 claude-sonnet-4-5-20241022"
);
// 测试逐步删除后缀匹配 - 多个后缀
let result = find_model_pricing_row(&conn, "claude-haiku-4-5-20240229-preview")?;
assert!(
result.is_some(),
"应该能通过删除后缀匹配 claude-haiku-4-5-20240229-preview"
);
// 测试 GPT 模型
let result = find_model_pricing_row(&conn, "gpt-5-2024-11-20")?;
assert!(result.is_some(), "应该能通过删除后缀匹配 gpt-5-2024-11-20");
// 测试 Gemini 模型
let result = find_model_pricing_row(&conn, "gemini-2.5-flash-exp")?;
assert!(
result.is_some(),
"应该能通过删除后缀匹配 gemini-2.5-flash-exp"
);
// 测试 claude-sonnet-4-5 命名格式
let result = find_model_pricing_row(&conn, "claude-sonnet-4-5-20250929")?;
assert!(
result.is_some(),
"应该能通过删除后缀匹配 claude-sonnet-4-5-20250929"
);
// 测试不存在的模型
let result = find_model_pricing_row(&conn, "unknown-model-123")?;
assert!(result.is_none(), "不应该匹配不存在的模型");
Ok(())
}
}