From bba652497926b5990541a613dab42c0f6f18829c Mon Sep 17 00:00:00 2001 From: YoVinchen Date: Mon, 26 Jan 2026 01:37:51 +0800 Subject: [PATCH] feat(db): add pricing config fields to proxy_config table - Add default_cost_multiplier field per app type - Add pricing_model_source field (request/response) - Add request_model field to proxy_request_logs table - Implement schema migration v5 --- src-tauri/src/database/dao/proxy.rs | 149 ++++++++++++++++++ src-tauri/src/database/mod.rs | 2 +- src-tauri/src/database/schema.rs | 35 ++++ src-tauri/src/database/tests.rs | 80 +++++++++- src-tauri/tests/proxy_commands.rs | 68 ++++++++ tests/components/GlobalProxySettings.test.tsx | 83 ++++++++++ tests/setupGlobals.ts | 26 +++ 7 files changed, 435 insertions(+), 8 deletions(-) create mode 100644 src-tauri/tests/proxy_commands.rs create mode 100644 tests/components/GlobalProxySettings.test.tsx create mode 100644 tests/setupGlobals.ts diff --git a/src-tauri/src/database/dao/proxy.rs b/src-tauri/src/database/dao/proxy.rs index 009f24af7..a985fcaa6 100644 --- a/src-tauri/src/database/dao/proxy.rs +++ b/src-tauri/src/database/dao/proxy.rs @@ -4,6 +4,7 @@ use crate::error::AppError; use crate::proxy::types::*; +use rust_decimal::Decimal; use super::super::{lock_conn, Database}; @@ -75,6 +76,101 @@ impl Database { Ok(()) } + /// 获取默认成本倍率 + pub async fn get_default_cost_multiplier(&self, app_type: &str) -> Result { + let result = { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT default_cost_multiplier FROM proxy_config WHERE app_type = ?1", + [app_type], + |row| row.get(0), + ) + }; + + match result { + Ok(value) => Ok(value), + Err(rusqlite::Error::QueryReturnedNoRows) => { + self.init_proxy_config_rows().await?; + Ok("1".to_string()) + } + Err(e) => Err(AppError::Database(e.to_string())), + } + } + + /// 设置默认成本倍率 + pub async fn set_default_cost_multiplier( + &self, + app_type: &str, + value: &str, + ) -> Result<(), AppError> { + let trimmed = value.trim(); + if trimmed.is_empty() { + return Err(AppError::InvalidInput("倍率不能为空".to_string())); + } + trimmed + .parse::() + .map_err(|e| AppError::InvalidInput(format!("无效倍率: {value} - {e}")))?; + + let conn = lock_conn!(self.conn); + conn.execute( + "UPDATE proxy_config SET + default_cost_multiplier = ?2, + updated_at = datetime('now') + WHERE app_type = ?1", + rusqlite::params![app_type, trimmed], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + Ok(()) + } + + /// 获取计费模式来源 + pub async fn get_pricing_model_source(&self, app_type: &str) -> Result { + let result = { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT pricing_model_source FROM proxy_config WHERE app_type = ?1", + [app_type], + |row| row.get(0), + ) + }; + + match result { + Ok(value) => Ok(value), + Err(rusqlite::Error::QueryReturnedNoRows) => { + self.init_proxy_config_rows().await?; + Ok("response".to_string()) + } + Err(e) => Err(AppError::Database(e.to_string())), + } + } + + /// 设置计费模式来源 + pub async fn set_pricing_model_source( + &self, + app_type: &str, + value: &str, + ) -> Result<(), AppError> { + let trimmed = value.trim(); + if !matches!(trimmed, "response" | "request") { + return Err(AppError::InvalidInput(format!( + "无效计费模式: {value}" + ))); + } + + let conn = lock_conn!(self.conn); + conn.execute( + "UPDATE proxy_config SET + pricing_model_source = ?2, + updated_at = datetime('now') + WHERE app_type = ?1", + rusqlite::params![app_type, trimmed], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + Ok(()) + } + /// 获取应用级代理配置 pub async fn get_proxy_config_for_app( &self, @@ -662,3 +758,56 @@ impl Database { Ok(()) } } + +#[cfg(test)] +mod tests { + use crate::database::Database; + use crate::error::AppError; + + #[tokio::test] + async fn test_default_cost_multiplier_round_trip() -> Result<(), AppError> { + let db = Database::memory()?; + + let default = db.get_default_cost_multiplier("claude").await?; + assert_eq!(default, "1"); + + db.set_default_cost_multiplier("claude", "1.5").await?; + let updated = db.get_default_cost_multiplier("claude").await?; + assert_eq!(updated, "1.5"); + + Ok(()) + } + + #[tokio::test] + async fn test_default_cost_multiplier_validation() -> Result<(), AppError> { + let db = Database::memory()?; + + let err = db + .set_default_cost_multiplier("claude", "not-a-number") + .await + .unwrap_err(); + assert!(matches!(err, AppError::InvalidInput(_))); + + Ok(()) + } + + #[tokio::test] + async fn test_pricing_model_source_round_trip_and_validation() -> Result<(), AppError> { + let db = Database::memory()?; + + let default = db.get_pricing_model_source("claude").await?; + assert_eq!(default, "response"); + + db.set_pricing_model_source("claude", "request").await?; + let updated = db.get_pricing_model_source("claude").await?; + assert_eq!(updated, "request"); + + let err = db + .set_pricing_model_source("claude", "invalid") + .await + .unwrap_err(); + assert!(matches!(err, AppError::InvalidInput(_))); + + Ok(()) + } +} diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs index 1c855a5c7..2ba3ed591 100644 --- a/src-tauri/src/database/mod.rs +++ b/src-tauri/src/database/mod.rs @@ -47,7 +47,7 @@ const DB_BACKUP_RETAIN: usize = 10; /// 当前 Schema 版本号 /// 每次修改表结构时递增,并在 schema.rs 中添加相应的迁移逻辑 -pub(crate) const SCHEMA_VERSION: i32 = 4; +pub(crate) const SCHEMA_VERSION: i32 = 5; /// 安全地序列化 JSON,避免 unwrap panic pub(crate) fn to_json_string(value: &T) -> Result { diff --git a/src-tauri/src/database/schema.rs b/src-tauri/src/database/schema.rs index 4bd334fd7..33a3f6c74 100644 --- a/src-tauri/src/database/schema.rs +++ b/src-tauri/src/database/schema.rs @@ -120,6 +120,8 @@ impl Database { circuit_failure_threshold INTEGER NOT NULL DEFAULT 4, circuit_success_threshold INTEGER NOT NULL DEFAULT 2, circuit_timeout_seconds INTEGER NOT NULL DEFAULT 60, circuit_error_rate_threshold REAL NOT NULL DEFAULT 0.6, circuit_min_requests INTEGER NOT NULL DEFAULT 10, + default_cost_multiplier TEXT NOT NULL DEFAULT '1', + pricing_model_source TEXT NOT NULL DEFAULT 'response', created_at TEXT NOT NULL DEFAULT (datetime('now')), updated_at TEXT NOT NULL DEFAULT (datetime('now')) )", []).map_err(|e| AppError::Database(e.to_string()))?; @@ -170,6 +172,7 @@ impl Database { // 10. Proxy Request Logs 表 conn.execute("CREATE TABLE IF NOT EXISTS proxy_request_logs ( request_id TEXT PRIMARY KEY, provider_id TEXT NOT NULL, app_type TEXT NOT NULL, model TEXT NOT NULL, + request_model TEXT, input_tokens INTEGER NOT NULL DEFAULT 0, output_tokens INTEGER NOT NULL DEFAULT 0, cache_read_tokens INTEGER NOT NULL DEFAULT 0, cache_creation_tokens INTEGER NOT NULL DEFAULT 0, input_cost_usd TEXT NOT NULL DEFAULT '0', output_cost_usd TEXT NOT NULL DEFAULT '0', @@ -352,6 +355,11 @@ impl Database { Self::migrate_v3_to_v4(conn)?; Self::set_user_version(conn, 4)?; } + 4 => { + log::info!("迁移数据库从 v4 到 v5(计费模式支持)"); + Self::migrate_v4_to_v5(conn)?; + Self::set_user_version(conn, 5)?; + } _ => { return Err(AppError::Database(format!( "未知的数据库版本 {version},无法迁移到 {SCHEMA_VERSION}" @@ -521,6 +529,7 @@ impl Database { // proxy_request_logs 表 conn.execute("CREATE TABLE IF NOT EXISTS proxy_request_logs ( request_id TEXT PRIMARY KEY, provider_id TEXT NOT NULL, app_type TEXT NOT NULL, model TEXT NOT NULL, + request_model TEXT, input_tokens INTEGER NOT NULL DEFAULT 0, output_tokens INTEGER NOT NULL DEFAULT 0, cache_read_tokens INTEGER NOT NULL DEFAULT 0, cache_creation_tokens INTEGER NOT NULL DEFAULT 0, input_cost_usd TEXT NOT NULL DEFAULT '0', output_cost_usd TEXT NOT NULL DEFAULT '0', @@ -677,6 +686,8 @@ impl Database { circuit_failure_threshold INTEGER NOT NULL DEFAULT 4, circuit_success_threshold INTEGER NOT NULL DEFAULT 2, circuit_timeout_seconds INTEGER NOT NULL DEFAULT 60, circuit_error_rate_threshold REAL NOT NULL DEFAULT 0.6, circuit_min_requests INTEGER NOT NULL DEFAULT 10, + default_cost_multiplier TEXT NOT NULL DEFAULT '1', + pricing_model_source TEXT NOT NULL DEFAULT 'response', created_at TEXT NOT NULL DEFAULT (datetime('now')), updated_at TEXT NOT NULL DEFAULT (datetime('now')) )", [])?; @@ -879,6 +890,30 @@ impl Database { Ok(()) } + /// v4 -> v5 迁移:新增计费模式配置与请求模型字段 + fn migrate_v4_to_v5(conn: &Connection) -> Result<(), AppError> { + if Self::table_exists(conn, "proxy_config")? { + Self::add_column_if_missing( + conn, + "proxy_config", + "default_cost_multiplier", + "TEXT NOT NULL DEFAULT '1'", + )?; + Self::add_column_if_missing( + conn, + "proxy_config", + "pricing_model_source", + "TEXT NOT NULL DEFAULT 'response'", + )?; + } + if Self::table_exists(conn, "proxy_request_logs")? { + Self::add_column_if_missing(conn, "proxy_request_logs", "request_model", "TEXT")?; + } + + log::info!("v4 -> v5 迁移完成:已添加计费模式与请求模型字段"); + Ok(()) + } + /// 插入默认模型定价数据 /// 格式: (model_id, display_name, input, output, cache_read, cache_creation) /// 注意: model_id 使用短横线格式(如 claude-haiku-4-5),与 API 返回的模型名称标准化后一致 diff --git a/src-tauri/src/database/tests.rs b/src-tauri/src/database/tests.rs index f4325470c..620bf25d8 100644 --- a/src-tauri/src/database/tests.rs +++ b/src-tauri/src/database/tests.rs @@ -151,7 +151,7 @@ fn normalize_default(default: &Option) -> Option { } #[test] -fn migration_sets_user_version_when_missing() { +fn schema_migration_sets_user_version_when_missing() { let conn = Connection::open_in_memory().expect("open memory db"); Database::create_tables_on_conn(&conn).expect("create tables"); @@ -169,7 +169,7 @@ fn migration_sets_user_version_when_missing() { } #[test] -fn migration_rejects_future_version() { +fn schema_migration_rejects_future_version() { let conn = Connection::open_in_memory().expect("open memory db"); Database::create_tables_on_conn(&conn).expect("create tables"); Database::set_user_version(&conn, SCHEMA_VERSION + 1).expect("set future version"); @@ -183,7 +183,7 @@ fn migration_rejects_future_version() { } #[test] -fn migration_adds_missing_columns_for_providers() { +fn schema_migration_adds_missing_columns_for_providers() { let conn = Connection::open_in_memory().expect("open memory db"); // 创建旧版 providers 表,缺少新增列 @@ -224,7 +224,7 @@ fn migration_adds_missing_columns_for_providers() { } #[test] -fn migration_aligns_column_defaults_and_types() { +fn schema_migration_aligns_column_defaults_and_types() { let conn = Connection::open_in_memory().expect("open memory db"); conn.execute_batch(LEGACY_SCHEMA_SQL) .expect("seed old schema"); @@ -268,7 +268,73 @@ fn migration_aligns_column_defaults_and_types() { } #[test] -fn create_tables_repairs_legacy_proxy_config_singleton_to_per_app() { +fn schema_create_tables_include_pricing_model_columns() { + let conn = Connection::open_in_memory().expect("open memory db"); + Database::create_tables_on_conn(&conn).expect("create tables"); + + let multiplier = get_column_info(&conn, "proxy_config", "default_cost_multiplier"); + assert_eq!(multiplier.r#type, "TEXT"); + assert_eq!(multiplier.notnull, 1); + assert_eq!( + normalize_default(&multiplier.default).as_deref(), + Some("1") + ); + + let pricing_source = get_column_info(&conn, "proxy_config", "pricing_model_source"); + assert_eq!(pricing_source.r#type, "TEXT"); + assert_eq!(pricing_source.notnull, 1); + assert_eq!( + normalize_default(&pricing_source.default).as_deref(), + Some("response") + ); + + let request_model = get_column_info(&conn, "proxy_request_logs", "request_model"); + assert_eq!(request_model.r#type, "TEXT"); + assert_eq!(request_model.notnull, 0); +} + +#[test] +fn schema_migration_v4_adds_pricing_model_columns() { + let conn = Connection::open_in_memory().expect("open memory db"); + conn.execute_batch( + r#" + CREATE TABLE proxy_config (app_type TEXT PRIMARY KEY); + CREATE TABLE proxy_request_logs (request_id TEXT PRIMARY KEY, model TEXT NOT NULL); + "#, + ) + .expect("seed v4 schema"); + + Database::set_user_version(&conn, 4).expect("set user_version=4"); + Database::apply_schema_migrations_on_conn(&conn).expect("apply migrations"); + + let multiplier = get_column_info(&conn, "proxy_config", "default_cost_multiplier"); + assert_eq!(multiplier.r#type, "TEXT"); + assert_eq!(multiplier.notnull, 1); + assert_eq!( + normalize_default(&multiplier.default).as_deref(), + Some("1") + ); + + let pricing_source = get_column_info(&conn, "proxy_config", "pricing_model_source"); + assert_eq!(pricing_source.r#type, "TEXT"); + assert_eq!(pricing_source.notnull, 1); + assert_eq!( + normalize_default(&pricing_source.default).as_deref(), + Some("response") + ); + + let request_model = get_column_info(&conn, "proxy_request_logs", "request_model"); + assert_eq!(request_model.r#type, "TEXT"); + assert_eq!(request_model.notnull, 0); + + assert_eq!( + Database::get_user_version(&conn).expect("version after migration"), + SCHEMA_VERSION + ); +} + +#[test] +fn schema_create_tables_repairs_legacy_proxy_config_singleton_to_per_app() { let conn = Connection::open_in_memory().expect("open memory db"); // 模拟测试版 v2:user_version=2,但 proxy_config 仍是单例结构(无 app_type) @@ -433,7 +499,7 @@ fn migration_from_v3_8_schema_v1_to_current_schema_v3() { } #[test] -fn dry_run_does_not_write_to_disk() { +fn schema_dry_run_does_not_write_to_disk() { // Create minimal valid config for migration let mut apps = HashMap::new(); apps.insert("claude".to_string(), ProviderManager::default()); @@ -507,7 +573,7 @@ fn dry_run_validates_schema_compatibility() { } #[test] -fn model_pricing_is_seeded_on_init() { +fn schema_model_pricing_is_seeded_on_init() { let db = Database::memory().expect("create memory db"); let conn = db.conn.lock().expect("lock conn"); diff --git a/src-tauri/tests/proxy_commands.rs b/src-tauri/tests/proxy_commands.rs new file mode 100644 index 000000000..d727e7ae7 --- /dev/null +++ b/src-tauri/tests/proxy_commands.rs @@ -0,0 +1,68 @@ +use cc_switch_lib::{ + get_default_cost_multiplier_test_hook, get_pricing_model_source_test_hook, + set_default_cost_multiplier_test_hook, set_pricing_model_source_test_hook, AppError, +}; + +#[path = "support.rs"] +mod support; +use support::{create_test_state, ensure_test_home, reset_test_fs, test_mutex}; + +#[tokio::test] +async fn default_cost_multiplier_commands_round_trip() { + let _guard = test_mutex().lock().expect("acquire test mutex"); + reset_test_fs(); + let _home = ensure_test_home(); + + let state = create_test_state().expect("create test state"); + + let default = get_default_cost_multiplier_test_hook(&state, "claude") + .await + .expect("read default multiplier"); + assert_eq!(default, "1"); + + set_default_cost_multiplier_test_hook(&state, "claude", "1.5") + .await + .expect("set multiplier"); + let updated = get_default_cost_multiplier_test_hook(&state, "claude") + .await + .expect("read updated multiplier"); + assert_eq!(updated, "1.5"); + + let err = set_default_cost_multiplier_test_hook(&state, "claude", "not-a-number") + .await + .expect_err("invalid multiplier should error"); + match err { + AppError::InvalidInput(_) => {} + other => panic!("expected invalid input error, got {other:?}"), + } +} + +#[tokio::test] +async fn pricing_model_source_commands_round_trip() { + let _guard = test_mutex().lock().expect("acquire test mutex"); + reset_test_fs(); + let _home = ensure_test_home(); + + let state = create_test_state().expect("create test state"); + + let default = get_pricing_model_source_test_hook(&state, "claude") + .await + .expect("read default pricing model source"); + assert_eq!(default, "response"); + + set_pricing_model_source_test_hook(&state, "claude", "request") + .await + .expect("set pricing model source"); + let updated = get_pricing_model_source_test_hook(&state, "claude") + .await + .expect("read updated pricing model source"); + assert_eq!(updated, "request"); + + let err = set_pricing_model_source_test_hook(&state, "claude", "invalid") + .await + .expect_err("invalid pricing model source should error"); + match err { + AppError::InvalidInput(_) => {} + other => panic!("expected invalid input error, got {other:?}"), + } +} diff --git a/tests/components/GlobalProxySettings.test.tsx b/tests/components/GlobalProxySettings.test.tsx new file mode 100644 index 000000000..f2147c32e --- /dev/null +++ b/tests/components/GlobalProxySettings.test.tsx @@ -0,0 +1,83 @@ +import { render, screen, fireEvent, waitFor } from "@testing-library/react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { GlobalProxySettings } from "@/components/settings/GlobalProxySettings"; + +vi.mock("react-i18next", () => ({ + useTranslation: () => ({ t: (key: string) => key }), +})); + +const mutateAsyncMock = vi.fn(); +const testMutateAsyncMock = vi.fn(); +const scanMutateAsyncMock = vi.fn(); + +vi.mock("@/hooks/useGlobalProxy", () => ({ + useGlobalProxyUrl: () => ({ data: "http://127.0.0.1:7890", isLoading: false }), + useSetGlobalProxyUrl: () => ({ + mutateAsync: mutateAsyncMock, + isPending: false, + }), + useTestProxy: () => ({ + mutateAsync: testMutateAsyncMock, + isPending: false, + }), + useScanProxies: () => ({ + mutateAsync: scanMutateAsyncMock, + isPending: false, + }), +})); + +describe("GlobalProxySettings", () => { + beforeEach(() => { + mutateAsyncMock.mockReset(); + testMutateAsyncMock.mockReset(); + scanMutateAsyncMock.mockReset(); + }); + + it("renders proxy URL input with saved value", async () => { + render(); + + const urlInput = screen.getByPlaceholderText( + "http://127.0.0.1:7890 / socks5://127.0.0.1:1080", + ); + // URL 对象会在末尾添加斜杠 + await waitFor(() => + expect(urlInput).toHaveValue("http://127.0.0.1:7890/"), + ); + }); + + it("saves proxy URL when save button is clicked", async () => { + render(); + + const urlInput = screen.getByPlaceholderText( + "http://127.0.0.1:7890 / socks5://127.0.0.1:1080", + ); + + fireEvent.change(urlInput, { target: { value: "http://localhost:8080" } }); + + const saveButton = screen.getByRole("button", { name: "common.save" }); + fireEvent.click(saveButton); + + await waitFor(() => expect(mutateAsyncMock).toHaveBeenCalled()); + // 没有用户名时,URL 不经过 URL 对象解析,所以没有尾部斜杠 + expect(mutateAsyncMock).toHaveBeenCalledWith("http://localhost:8080"); + }); + + it("clears proxy URL when clear button is clicked", async () => { + render(); + + const urlInput = screen.getByPlaceholderText( + "http://127.0.0.1:7890 / socks5://127.0.0.1:1080", + ); + + // Wait for initial value to load + await waitFor(() => + expect(urlInput).toHaveValue("http://127.0.0.1:7890/"), + ); + + // Click clear button + const clearButton = screen.getByTitle("settings.globalProxy.clear"); + fireEvent.click(clearButton); + + expect(urlInput).toHaveValue(""); + }); +}); diff --git a/tests/setupGlobals.ts b/tests/setupGlobals.ts new file mode 100644 index 000000000..01633a7f5 --- /dev/null +++ b/tests/setupGlobals.ts @@ -0,0 +1,26 @@ +const storage = new Map(); + +if ( + typeof globalThis.localStorage === "undefined" || + typeof globalThis.localStorage?.getItem !== "function" +) { + Object.defineProperty(globalThis, "localStorage", { + value: { + getItem: (key: string) => storage.get(key) ?? null, + setItem: (key: string, value: string) => { + storage.set(key, String(value)); + }, + removeItem: (key: string) => { + storage.delete(key); + }, + clear: () => { + storage.clear(); + }, + key: (index: number) => Array.from(storage.keys())[index] ?? null, + get length() { + return storage.size; + }, + }, + configurable: true, + }); +}