From 12b972a66e423897e92cc915ec680bfcd9156a8b Mon Sep 17 00:00:00 2001 From: Thefool <114680394+YUZHEthefool@users.noreply.github.com> Date: Tue, 28 Jul 2026 23:37:55 +0800 Subject: [PATCH] feat(usage): add automatic models.dev pricing sync (#5734) * feat(usage): persist model pricing in local config * feat(usage): sync selected models.dev pricing on startup * fix(usage): address models.dev sync review feedback * fix(usage): harden local pricing synchronization --- src-tauri/src/commands/usage.rs | 133 ++- src-tauri/src/database/mod.rs | 3 + src-tauri/src/lib.rs | 4 + src-tauri/src/services/mod.rs | 1 + src-tauri/src/services/model_pricing.rs | 761 ++++++++++++++++++ .../usage/ModelsDevAutoSyncPanel.tsx | 668 +++++++++++++++ .../usage/ModelsDevPickerDialog.tsx | 122 +-- src/components/usage/PricingConfigPanel.tsx | 3 + src/i18n/locales/en.json | 34 + src/i18n/locales/ja.json | 34 + src/i18n/locales/zh-TW.json | 34 + src/i18n/locales/zh.json | 34 + src/lib/api/usage.ts | 23 + src/lib/modelsDevAutoSync.ts | 93 +++ src/lib/modelsDevPricing.ts | 277 +++++++ src/main.tsx | 23 + src/types/usage.ts | 14 + .../ModelsDevAutoSyncPanel.test.tsx | 226 ++++++ .../components/ModelsDevPickerDialog.test.ts | 184 +++++ tests/lib/modelsDevAutoSync.test.ts | 196 +++++ 20 files changed, 2672 insertions(+), 195 deletions(-) create mode 100644 src-tauri/src/services/model_pricing.rs create mode 100644 src/components/usage/ModelsDevAutoSyncPanel.tsx create mode 100644 src/lib/modelsDevAutoSync.ts create mode 100644 src/lib/modelsDevPricing.ts create mode 100644 tests/components/ModelsDevAutoSyncPanel.test.tsx create mode 100644 tests/lib/modelsDevAutoSync.test.ts diff --git a/src-tauri/src/commands/usage.rs b/src-tauri/src/commands/usage.rs index fe577f9f4..0649c70aa 100644 --- a/src-tauri/src/commands/usage.rs +++ b/src-tauri/src/commands/usage.rs @@ -1,10 +1,9 @@ //! 使用统计相关命令 use crate::error::AppError; +use crate::services::model_pricing::{ModelPricingInfo, ModelsDevSyncConfig, ModelsDevSyncState}; use crate::services::usage_stats::*; use crate::store::AppState; -use rust_decimal::Decimal; -use std::str::FromStr; use tauri::State; /// 获取使用量汇总 @@ -125,6 +124,7 @@ pub fn get_request_detail( pub fn get_model_pricing(state: State<'_, AppState>) -> Result, AppError> { log::info!("获取模型定价列表"); state.db.ensure_model_pricing_seeded()?; + crate::services::model_pricing::sync_local_model_pricing(&state.db)?; let db = state.db.clone(); let conn = crate::database::lock_conn!(db.conn); @@ -181,72 +181,53 @@ pub fn update_model_pricing( cache_read_cost: String, cache_creation_cost: String, ) -> Result<(), AppError> { - let db = state.db.clone(); - let model_id = model_id.trim().to_string(); - let display_name = display_name.trim().to_string(); - if model_id.is_empty() { - return Err(AppError::localized( - "usage.modelIdRequired", - "模型 ID 不能为空", - "Model ID is required", - )); - } - if display_name.is_empty() { - return Err(AppError::localized( - "usage.displayNameRequired", - "显示名称不能为空", - "Display name is required", - )); - } - - for (label, value) in [ - ("input_cost", &input_cost), - ("output_cost", &output_cost), - ("cache_read_cost", &cache_read_cost), - ("cache_creation_cost", &cache_creation_cost), - ] { - let parsed = Decimal::from_str(value.trim()).map_err(|e| { - AppError::localized( - "usage.invalidPrice", - format!("{label} 价格无效: {value} - {e}"), - format!("{label} price is invalid: {value} - {e}"), - ) - })?; - if parsed < Decimal::ZERO { - return Err(AppError::localized( - "usage.invalidPrice", - format!("{label} 价格必须为非负数: {value}"), - format!("{label} price must be non-negative: {value}"), - )); - } - } - - { - let conn = crate::database::lock_conn!(db.conn); - conn.execute( - "INSERT OR REPLACE INTO model_pricing ( - model_id, display_name, input_cost_per_million, output_cost_per_million, - cache_read_cost_per_million, cache_creation_cost_per_million - ) VALUES (?1, ?2, ?3, ?4, ?5, ?6)", - rusqlite::params![ - model_id, - display_name, - input_cost.trim(), - output_cost.trim(), - cache_read_cost.trim(), - cache_creation_cost.trim() - ], - ) - .map_err(|e| AppError::Database(format!("更新模型定价失败: {e}")))?; - } - - if let Err(e) = db.backfill_missing_usage_costs_for_model(&model_id) { - log::warn!("模型定价更新后回填历史用量成本失败 (model_id={model_id}): {e}"); - } - + crate::services::model_pricing::update_model_pricing( + &state.db, + ModelPricingInfo { + model_id, + display_name, + input_cost_per_million: input_cost, + output_cost_per_million: output_cost, + cache_read_cost_per_million: cache_read_cost, + cache_creation_cost_per_million: cache_creation_cost, + }, + )?; Ok(()) } +/// 批量更新模型定价(models.dev 自动同步仅触发一次历史成本回填) +#[tauri::command] +pub fn update_model_pricing_batch( + state: State<'_, AppState>, + entries: Vec, +) -> Result { + crate::services::model_pricing::update_model_pricing_batch(&state.db, entries) +} + +#[tauri::command] +pub fn get_models_dev_sync_config( + state: State<'_, AppState>, +) -> Result { + crate::services::model_pricing::get_models_dev_sync_state(&state.db) +} + +#[tauri::command] +pub fn save_models_dev_sync_config( + state: State<'_, AppState>, + config: ModelsDevSyncConfig, +) -> Result<(), AppError> { + crate::services::model_pricing::save_models_dev_sync_config(&state.db, config) +} + +#[tauri::command] +pub fn record_models_dev_sync_result( + state: State<'_, AppState>, + synced_at: Option, + error: Option, +) -> Result<(), AppError> { + crate::services::model_pricing::record_models_dev_sync_result(&state.db, synced_at, error) +} + /// 检查 Provider 使用限额 #[tauri::command] pub fn check_provider_limits( @@ -260,15 +241,7 @@ pub fn check_provider_limits( /// 删除模型定价 #[tauri::command] pub fn delete_model_pricing(state: State<'_, AppState>, model_id: String) -> Result<(), AppError> { - let db = state.db.clone(); - let conn = crate::database::lock_conn!(db.conn); - - conn.execute( - "DELETE FROM model_pricing WHERE model_id = ?1", - rusqlite::params![model_id], - ) - .map_err(|e| AppError::Database(format!("删除模型定价失败: {e}")))?; - + crate::services::model_pricing::delete_model_pricing(&state.db, &model_id)?; log::info!("已删除模型定价: {model_id}"); Ok(()) } @@ -326,18 +299,6 @@ pub fn get_usage_data_sources( crate::services::session_usage::get_data_source_breakdown(&state.db) } -/// 模型定价信息 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct ModelPricingInfo { - pub model_id: String, - pub display_name: String, - pub input_cost_per_million: String, - pub output_cost_per_million: String, - pub cache_read_cost_per_million: String, - pub cache_creation_cost_per_million: String, -} - #[cfg(test)] mod tests { use super::*; diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs index ca946dd31..25cd162e8 100644 --- a/src-tauri/src/database/mod.rs +++ b/src-tauri/src/database/mod.rs @@ -144,6 +144,9 @@ impl Database { log::warn!("Failed to ensure incremental auto-vacuum: {e}"); } db.ensure_model_pricing_seeded()?; + if let Err(e) = crate::services::model_pricing::sync_local_model_pricing(&db) { + log::warn!("Failed to sync local model pricing file: {e}"); + } // Startup cleanup: prune old logs and reclaim space if let Err(e) = db.cleanup_old_stream_check_logs(7) { diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 13b3b7ba8..1e6cd62fa 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -1532,7 +1532,11 @@ pub fn run() { commands::get_request_detail, commands::get_model_pricing, commands::update_model_pricing, + commands::update_model_pricing_batch, commands::delete_model_pricing, + commands::get_models_dev_sync_config, + commands::save_models_dev_sync_config, + commands::record_models_dev_sync_result, commands::check_provider_limits, // Session usage sync commands::sync_session_usage, diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index 1d824cec0..2fce3bec7 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -6,6 +6,7 @@ pub mod env_checker; pub mod env_manager; pub mod mcp; pub mod model_fetch; +pub mod model_pricing; pub mod omo; pub mod profile; pub mod prompt; diff --git a/src-tauri/src/services/model_pricing.rs b/src-tauri/src/services/model_pricing.rs new file mode 100644 index 000000000..377838462 --- /dev/null +++ b/src-tauri/src/services/model_pricing.rs @@ -0,0 +1,761 @@ +use crate::config::{atomic_write, get_app_config_dir}; +use crate::database::{lock_conn, Database}; +use crate::error::AppError; +use rusqlite::{params, Transaction}; +use rust_decimal::Decimal; +use serde::{Deserialize, Serialize}; +use std::collections::{BTreeMap, BTreeSet}; +use std::fs; +use std::path::PathBuf; +use std::str::FromStr; +use std::sync::{Mutex, OnceLock}; + +const MODEL_PRICING_FILE_NAME: &str = "model-pricing.json"; +const MODEL_PRICING_FILE_VERSION: u32 = 1; + +static MODEL_PRICING_FILE_LOCK: OnceLock> = OnceLock::new(); + +fn file_lock() -> &'static Mutex<()> { + MODEL_PRICING_FILE_LOCK.get_or_init(|| Mutex::new(())) +} + +fn default_true() -> bool { + true +} + +fn default_file_version() -> u32 { + MODEL_PRICING_FILE_VERSION +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ModelPricingInfo { + pub model_id: String, + pub display_name: String, + pub input_cost_per_million: String, + pub output_cost_per_million: String, + pub cache_read_cost_per_million: String, + pub cache_creation_cost_per_million: String, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ModelsDevSyncConfig { + #[serde(default)] + pub auto_sync_enabled: bool, + #[serde(default = "default_true")] + pub include_common_models: bool, + #[serde(default)] + pub selected_model_keys: Vec, + #[serde(default)] + pub excluded_common_model_keys: Vec, + #[serde(default)] + pub last_sync_at: Option, + #[serde(default)] + pub last_sync_error: Option, +} + +impl Default for ModelsDevSyncConfig { + fn default() -> Self { + Self { + auto_sync_enabled: false, + include_common_models: true, + selected_model_keys: Vec::new(), + excluded_common_model_keys: Vec::new(), + last_sync_at: None, + last_sync_error: None, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +struct ModelPricingFile { + #[serde(default = "default_file_version")] + version: u32, + #[serde(default)] + models_dev_sync: ModelsDevSyncConfig, + #[serde(default)] + models: Vec, + #[serde(default)] + deleted_model_ids: Vec, +} + +impl Default for ModelPricingFile { + fn default() -> Self { + Self { + version: MODEL_PRICING_FILE_VERSION, + models_dev_sync: ModelsDevSyncConfig::default(), + models: Vec::new(), + deleted_model_ids: Vec::new(), + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct ModelsDevSyncState { + pub config: ModelsDevSyncConfig, + pub config_path: String, +} + +pub fn model_pricing_file_path() -> PathBuf { + get_app_config_dir().join(MODEL_PRICING_FILE_NAME) +} + +fn normalize_decimal(label: &str, value: &str) -> Result { + let value = value.trim(); + let parsed = Decimal::from_str(value).map_err(|error| { + AppError::localized( + "usage.invalidPrice", + format!("{label} 价格无效: {value} - {error}"), + format!("{label} price is invalid: {value} - {error}"), + ) + })?; + if parsed < Decimal::ZERO { + return Err(AppError::localized( + "usage.invalidPrice", + format!("{label} 价格必须为非负数: {value}"), + format!("{label} price must be non-negative: {value}"), + )); + } + Ok(value.to_string()) +} + +fn normalize_pricing(entry: ModelPricingInfo) -> Result { + let model_id = entry.model_id.trim().to_string(); + let display_name = entry.display_name.trim().to_string(); + if model_id.is_empty() { + return Err(AppError::localized( + "usage.modelIdRequired", + "模型 ID 不能为空", + "Model ID is required", + )); + } + if display_name.is_empty() { + return Err(AppError::localized( + "usage.displayNameRequired", + "显示名称不能为空", + "Display name is required", + )); + } + + Ok(ModelPricingInfo { + model_id, + display_name, + input_cost_per_million: normalize_decimal("input_cost", &entry.input_cost_per_million)?, + output_cost_per_million: normalize_decimal("output_cost", &entry.output_cost_per_million)?, + cache_read_cost_per_million: normalize_decimal( + "cache_read_cost", + &entry.cache_read_cost_per_million, + )?, + cache_creation_cost_per_million: normalize_decimal( + "cache_creation_cost", + &entry.cache_creation_cost_per_million, + )?, + }) +} + +fn normalize_key_list(values: Vec) -> Vec { + values + .into_iter() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .collect::>() + .into_iter() + .collect() +} + +fn normalize_sync_config(mut config: ModelsDevSyncConfig) -> ModelsDevSyncConfig { + config.selected_model_keys = normalize_key_list(config.selected_model_keys); + config.excluded_common_model_keys = normalize_key_list(config.excluded_common_model_keys); + config.last_sync_error = config.last_sync_error.and_then(|error| { + let trimmed = error.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed.chars().take(1000).collect()) + } + }); + config +} + +fn normalize_file(mut file: ModelPricingFile) -> Result { + if file.version > MODEL_PRICING_FILE_VERSION { + return Err(AppError::Config(format!( + "model-pricing.json version {} is newer than supported version {}", + file.version, MODEL_PRICING_FILE_VERSION + ))); + } + + let deleted = normalize_key_list(file.deleted_model_ids) + .into_iter() + .collect::>(); + let mut models = BTreeMap::new(); + for entry in file.models { + let entry = normalize_pricing(entry)?; + if !deleted.contains(&entry.model_id) { + models.insert(entry.model_id.clone(), entry); + } + } + + file.version = MODEL_PRICING_FILE_VERSION; + file.models_dev_sync = normalize_sync_config(file.models_dev_sync); + file.models = models.into_values().collect(); + file.deleted_model_ids = deleted.into_iter().collect(); + Ok(file) +} + +fn read_file_unlocked() -> Result, AppError> { + let path = model_pricing_file_path(); + if !path.exists() { + return Ok(None); + } + let content = fs::read_to_string(&path).map_err(|error| AppError::io(&path, error))?; + let file = serde_json::from_str(&content).map_err(|error| AppError::json(&path, error))?; + normalize_file(file).map(Some) +} + +fn write_file_unlocked(file: &ModelPricingFile) -> Result<(), AppError> { + let path = model_pricing_file_path(); + let mut data = serde_json::to_vec_pretty(file) + .map_err(|error| AppError::Config(format!("序列化模型定价配置失败: {error}")))?; + data.push(b'\n'); + atomic_write(&path, &data) +} + +fn load_or_create_file_unlocked() -> Result { + if let Some(file) = read_file_unlocked()? { + return Ok(file); + } + + // The local file stores user/models.dev overrides only. Exporting the + // complete seeded table here would turn built-in prices into overrides and + // roll back future repair_current_model_pricing corrections on startup. + let file = ModelPricingFile::default(); + write_file_unlocked(&file)?; + Ok(file) +} + +fn upsert_pricing( + transaction: &Transaction<'_>, + entry: &ModelPricingInfo, +) -> Result { + transaction + .execute( + "INSERT INTO model_pricing ( + model_id, display_name, input_cost_per_million, output_cost_per_million, + cache_read_cost_per_million, cache_creation_cost_per_million + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6) + ON CONFLICT(model_id) DO UPDATE SET + display_name = excluded.display_name, + input_cost_per_million = excluded.input_cost_per_million, + output_cost_per_million = excluded.output_cost_per_million, + cache_read_cost_per_million = excluded.cache_read_cost_per_million, + cache_creation_cost_per_million = excluded.cache_creation_cost_per_million + WHERE display_name <> excluded.display_name + OR input_cost_per_million <> excluded.input_cost_per_million + OR output_cost_per_million <> excluded.output_cost_per_million + OR cache_read_cost_per_million <> excluded.cache_read_cost_per_million + OR cache_creation_cost_per_million <> excluded.cache_creation_cost_per_million", + params![ + entry.model_id, + entry.display_name, + entry.input_cost_per_million, + entry.output_cost_per_million, + entry.cache_read_cost_per_million, + entry.cache_creation_cost_per_million + ], + ) + .map_err(|error| AppError::Database(format!("更新模型定价失败: {error}"))) +} + +fn apply_file_to_database( + db: &Database, + file: &ModelPricingFile, +) -> Result<(usize, usize), AppError> { + let mut conn = lock_conn!(db.conn); + let transaction = conn.transaction()?; + let mut upserted = 0; + for entry in &file.models { + upserted += upsert_pricing(&transaction, entry)?; + } + let mut deleted = 0; + for model_id in &file.deleted_model_ids { + deleted += transaction.execute( + "DELETE FROM model_pricing WHERE model_id = ?1", + params![model_id], + )?; + } + transaction.commit()?; + Ok((upserted, deleted)) +} + +/// Load user-maintained overrides from `~/.cc-switch/model-pricing.json`. +/// Built-in rows remain database-owned so application updates can repair them; +/// the file contains only explicit overrides and deletion tombstones. +pub fn sync_local_model_pricing(db: &Database) -> Result { + let (upserted, deleted) = { + let _file_guard = file_lock() + .lock() + .map_err(|error| AppError::Config(format!("模型定价文件锁失败: {error}")))?; + let file = load_or_create_file_unlocked()?; + apply_file_to_database(db, &file)? + }; + + // Deleting pricing cannot make a zero-cost usage row calculable. In + // particular, seeded rows covered by tombstones may be reinserted and + // deleted on every startup; they must not trigger a full-table backfill. + if upserted > 0 { + if let Err(error) = db.backfill_missing_usage_costs() { + log::warn!("本地模型定价同步后回填历史用量成本失败: {error}"); + } + } + Ok(upserted + deleted) +} + +pub fn get_models_dev_sync_state(db: &Database) -> Result { + sync_local_model_pricing(db)?; + let _file_guard = file_lock() + .lock() + .map_err(|error| AppError::Config(format!("模型定价文件锁失败: {error}")))?; + let file = load_or_create_file_unlocked()?; + Ok(ModelsDevSyncState { + config: file.models_dev_sync, + config_path: model_pricing_file_path().display().to_string(), + }) +} + +pub fn save_models_dev_sync_config( + db: &Database, + config: ModelsDevSyncConfig, +) -> Result<(), AppError> { + sync_local_model_pricing(db)?; + let _file_guard = file_lock() + .lock() + .map_err(|error| AppError::Config(format!("模型定价文件锁失败: {error}")))?; + let mut file = load_or_create_file_unlocked()?; + file.models_dev_sync = normalize_sync_config(config); + write_file_unlocked(&file) +} + +/// Persist only the outcome of a models.dev sync. Keeping this separate from +/// `save_models_dev_sync_config` prevents a slow startup fetch from restoring +/// stale switches or model selections that the user changed in the meantime. +pub fn record_models_dev_sync_result( + db: &Database, + synced_at: Option, + error: Option, +) -> Result<(), AppError> { + sync_local_model_pricing(db)?; + let _file_guard = file_lock() + .lock() + .map_err(|lock_error| AppError::Config(format!("模型定价文件锁失败: {lock_error}")))?; + let mut file = load_or_create_file_unlocked()?; + if let Some(synced_at) = synced_at { + file.models_dev_sync.last_sync_at = Some(synced_at); + } + file.models_dev_sync.last_sync_error = error; + file.models_dev_sync = normalize_sync_config(file.models_dev_sync); + write_file_unlocked(&file) +} + +fn update_model_pricing_batch_inner( + db: &Database, + entries: Vec, + backfill_all: bool, +) -> Result { + if entries.is_empty() { + return Ok(0); + } + let mut normalized = BTreeMap::new(); + for entry in entries { + let entry = normalize_pricing(entry)?; + normalized.insert(entry.model_id.clone(), entry); + } + let entries = normalized.into_values().collect::>(); + let model_ids = entries + .iter() + .map(|entry| entry.model_id.clone()) + .collect::>(); + + sync_local_model_pricing(db)?; + let changed = { + let _file_guard = file_lock() + .lock() + .map_err(|error| AppError::Config(format!("模型定价文件锁失败: {error}")))?; + let mut file = load_or_create_file_unlocked()?; + let mut file_models = file + .models + .into_iter() + .map(|entry| (entry.model_id.clone(), entry)) + .collect::>(); + let updated_ids = entries + .iter() + .map(|entry| entry.model_id.clone()) + .collect::>(); + for entry in &entries { + file_models.insert(entry.model_id.clone(), entry.clone()); + } + file.models = file_models.into_values().collect(); + file.deleted_model_ids + .retain(|model_id| !updated_ids.contains(model_id)); + + let mut conn = lock_conn!(db.conn); + let transaction = conn.transaction()?; + let mut changed = 0; + for entry in &entries { + changed += upsert_pricing(&transaction, entry)?; + } + write_file_unlocked(&file)?; + transaction.commit()?; + changed + }; + + if changed > 0 { + if backfill_all { + if let Err(error) = db.backfill_missing_usage_costs() { + log::warn!("批量更新模型定价后回填历史用量成本失败: {error}"); + } + } else { + for model_id in model_ids { + if let Err(error) = db.backfill_missing_usage_costs_for_model(&model_id) { + log::warn!("模型定价更新后回填历史用量成本失败 (model_id={model_id}): {error}"); + } + } + } + } + Ok(changed) +} + +pub fn update_model_pricing(db: &Database, entry: ModelPricingInfo) -> Result { + update_model_pricing_batch_inner(db, vec![entry], false) +} + +pub fn update_model_pricing_batch( + db: &Database, + entries: Vec, +) -> Result { + update_model_pricing_batch_inner(db, entries, true) +} + +pub fn delete_model_pricing(db: &Database, model_id: &str) -> Result<(), AppError> { + let model_id = model_id.trim(); + if model_id.is_empty() { + return Err(AppError::localized( + "usage.modelIdRequired", + "模型 ID 不能为空", + "Model ID is required", + )); + } + + sync_local_model_pricing(db)?; + let _file_guard = file_lock() + .lock() + .map_err(|error| AppError::Config(format!("模型定价文件锁失败: {error}")))?; + let mut file = load_or_create_file_unlocked()?; + file.models.retain(|entry| entry.model_id != model_id); + if !file.deleted_model_ids.iter().any(|entry| entry == model_id) { + file.deleted_model_ids.push(model_id.to_string()); + file.deleted_model_ids.sort(); + } + + let mut conn = lock_conn!(db.conn); + let transaction = conn.transaction()?; + transaction.execute( + "DELETE FROM model_pricing WHERE model_id = ?1", + params![model_id], + )?; + write_file_unlocked(&file)?; + transaction.commit()?; + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use serial_test::serial; + + fn with_test_home(test: impl FnOnce(&Database, &PathBuf)) { + let temp = tempfile::tempdir().expect("tempdir"); + let previous = std::env::var_os("CC_SWITCH_TEST_HOME"); + std::env::set_var("CC_SWITCH_TEST_HOME", temp.path()); + + let db = Database::memory().expect("memory database"); + let path = model_pricing_file_path(); + test(&db, &path); + + match previous { + Some(value) => std::env::set_var("CC_SWITCH_TEST_HOME", value), + None => std::env::remove_var("CC_SWITCH_TEST_HOME"), + } + } + + fn sample_pricing() -> ModelPricingInfo { + ModelPricingInfo { + model_id: "custom-model".to_string(), + display_name: "Custom Model".to_string(), + input_cost_per_million: "1.25".to_string(), + output_cost_per_million: "5".to_string(), + cache_read_cost_per_million: "0.1".to_string(), + cache_creation_cost_per_million: "1.5".to_string(), + } + } + + #[test] + #[serial] + fn creates_local_file_with_auto_sync_disabled_by_default() { + with_test_home(|db, path| { + let state = get_models_dev_sync_state(db).expect("sync state"); + assert!(path.exists()); + assert!(!state.config.auto_sync_enabled); + assert!(state.config.include_common_models); + assert_eq!(state.config_path, path.display().to_string()); + + let content = fs::read_to_string(path).expect("read pricing file"); + let file: ModelPricingFile = serde_json::from_str(&content).expect("parse file"); + assert!(file.models.is_empty()); + }); + } + + #[test] + #[serial] + fn empty_override_file_does_not_roll_back_builtin_pricing_repairs() { + with_test_home(|db, path| { + get_models_dev_sync_state(db).expect("create override file"); + { + let conn = db.conn.lock().expect("lock test database"); + assert_eq!( + conn.execute( + "UPDATE model_pricing + SET input_cost_per_million = '99' + WHERE model_id = 'claude-sonnet-5'", + [], + ) + .expect("simulate built-in pricing repair"), + 1 + ); + } + + assert_eq!(sync_local_model_pricing(db).expect("reload overrides"), 0); + + let conn = db.conn.lock().expect("lock test database"); + let input: String = conn + .query_row( + "SELECT input_cost_per_million + FROM model_pricing WHERE model_id = 'claude-sonnet-5'", + [], + |row| row.get(0), + ) + .expect("query repaired pricing"); + drop(conn); + assert_eq!(input, "99"); + + let content = fs::read_to_string(path).expect("read override file"); + let file: ModelPricingFile = serde_json::from_str(&content).expect("parse file"); + assert!(file.models.is_empty()); + }); + } + + #[test] + #[serial] + fn models_dev_batch_sync_overwrites_existing_manual_pricing() { + with_test_home(|db, path| { + let mut manual = sample_pricing(); + manual.input_cost_per_million = "9".to_string(); + manual.output_cost_per_million = "18".to_string(); + update_model_pricing(db, manual).expect("save manual pricing"); + + let synced = sample_pricing(); + update_model_pricing_batch(db, vec![synced.clone()]).expect("sync models.dev pricing"); + + let conn = db.conn.lock().expect("lock test database"); + let input: String = conn + .query_row( + "SELECT input_cost_per_million FROM model_pricing WHERE model_id = ?1", + params!["custom-model"], + |row| row.get(0), + ) + .expect("query synced pricing"); + drop(conn); + assert_eq!(input, synced.input_cost_per_million); + + let content = fs::read_to_string(path).expect("read pricing file"); + let file: ModelPricingFile = serde_json::from_str(&content).expect("parse file"); + let saved = file + .models + .iter() + .find(|entry| entry.model_id == "custom-model") + .expect("saved synced pricing"); + assert_eq!(saved, &synced); + }); + } + + #[test] + #[serial] + fn batch_update_and_delete_are_persisted_to_local_file() { + with_test_home(|db, path| { + assert_eq!( + update_model_pricing_batch(db, vec![sample_pricing()]).expect("batch update"), + 1 + ); + let content = fs::read_to_string(path).expect("read pricing file"); + let file: ModelPricingFile = serde_json::from_str(&content).expect("parse file"); + assert!(file + .models + .iter() + .any(|entry| entry.model_id == "custom-model")); + + delete_model_pricing(db, "custom-model").expect("delete pricing"); + let content = fs::read_to_string(path).expect("read updated file"); + let file: ModelPricingFile = + serde_json::from_str(&content).expect("parse updated file"); + assert!(!file + .models + .iter() + .any(|entry| entry.model_id == "custom-model")); + assert!(file + .deleted_model_ids + .iter() + .any(|entry| entry == "custom-model")); + }); + } + + #[test] + #[serial] + fn reloads_manual_file_edits_and_deletion_tombstones() { + with_test_home(|db, path| { + get_models_dev_sync_state(db).expect("create pricing file"); + let content = fs::read_to_string(path).expect("read pricing file"); + let mut file: ModelPricingFile = + serde_json::from_str(&content).expect("parse pricing file"); + file.models.push(sample_pricing()); + fs::write( + path, + serde_json::to_vec_pretty(&file).expect("serialize file"), + ) + .expect("write manual edit"); + + assert_eq!(sync_local_model_pricing(db).expect("reload file"), 1); + { + let conn = db.conn.lock().expect("lock test database"); + let input: String = conn + .query_row( + "SELECT input_cost_per_million FROM model_pricing WHERE model_id = ?1", + params!["custom-model"], + |row| row.get(0), + ) + .expect("query manually added pricing"); + assert_eq!(input, "1.25"); + } + + let content = fs::read_to_string(path).expect("read updated pricing file"); + let mut file: ModelPricingFile = + serde_json::from_str(&content).expect("parse updated pricing file"); + file.deleted_model_ids.push("custom-model".to_string()); + fs::write( + path, + serde_json::to_vec_pretty(&file).expect("serialize tombstone"), + ) + .expect("write tombstone"); + + assert_eq!(sync_local_model_pricing(db).expect("apply tombstone"), 1); + let conn = db.conn.lock().expect("lock test database"); + let count: i64 = conn + .query_row( + "SELECT COUNT(*) FROM model_pricing WHERE model_id = ?1", + params!["custom-model"], + |row| row.get(0), + ) + .expect("query deleted pricing"); + assert_eq!(count, 0); + }); + } + + #[test] + #[serial] + fn repeated_seeded_tombstone_deletion_does_not_backfill_unrelated_usage() { + with_test_home(|db, _path| { + get_models_dev_sync_state(db).expect("create override file"); + { + let conn = db.conn.lock().expect("lock test database"); + conn.execute( + "INSERT INTO proxy_request_logs ( + request_id, provider_id, app_type, model, request_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, created_at, data_source + ) VALUES ( + 'pending-cost', 'test-provider', 'codex', 'gpt-5', 'gpt-5', + 1000000, 0, 0, 0, '0', '0', '0', '0', '0', 100, 200, 1, 'proxy' + )", + [], + ) + .expect("insert zero-cost usage"); + } + + delete_model_pricing(db, "claude-sonnet-5").expect("create tombstone"); + db.ensure_model_pricing_seeded() + .expect("reseed built-in pricing"); + assert_eq!(sync_local_model_pricing(db).expect("apply tombstone"), 1); + + let conn = db.conn.lock().expect("lock test database"); + let deleted_count: i64 = conn + .query_row( + "SELECT COUNT(*) FROM model_pricing + WHERE model_id = 'claude-sonnet-5'", + [], + |row| row.get(0), + ) + .expect("query tombstoned pricing"); + let total_cost: f64 = conn + .query_row( + "SELECT CAST(total_cost_usd AS REAL) + FROM proxy_request_logs WHERE request_id = 'pending-cost'", + [], + |row| row.get(0), + ) + .expect("query pending usage cost"); + assert_eq!(deleted_count, 0); + assert_eq!(total_cost, 0.0); + }); + } + + #[test] + #[serial] + fn recording_sync_result_preserves_user_selection_and_switches() { + with_test_home(|db, _path| { + let config = ModelsDevSyncConfig { + auto_sync_enabled: false, + include_common_models: false, + selected_model_keys: vec!["relay/custom-model".to_string()], + excluded_common_model_keys: vec!["openai/gpt-5".to_string()], + last_sync_at: Some(123), + last_sync_error: Some("old error".to_string()), + }; + save_models_dev_sync_config(db, config.clone()).expect("save sync config"); + + record_models_dev_sync_result(db, Some(456), None).expect("record success"); + let state = get_models_dev_sync_state(db).expect("read sync state"); + assert_eq!(state.config.auto_sync_enabled, config.auto_sync_enabled); + assert_eq!( + state.config.include_common_models, + config.include_common_models + ); + assert_eq!(state.config.selected_model_keys, config.selected_model_keys); + assert_eq!( + state.config.excluded_common_model_keys, + config.excluded_common_model_keys + ); + assert_eq!(state.config.last_sync_at, Some(456)); + assert_eq!(state.config.last_sync_error, None); + + record_models_dev_sync_result(db, None, Some("offline".to_string())) + .expect("record failure"); + let state = get_models_dev_sync_state(db).expect("read failure state"); + assert_eq!(state.config.last_sync_at, Some(456)); + assert_eq!(state.config.last_sync_error.as_deref(), Some("offline")); + }); + } +} diff --git a/src/components/usage/ModelsDevAutoSyncPanel.tsx b/src/components/usage/ModelsDevAutoSyncPanel.tsx new file mode 100644 index 000000000..89a5d2182 --- /dev/null +++ b/src/components/usage/ModelsDevAutoSyncPanel.tsx @@ -0,0 +1,668 @@ +import { useMemo, useState } from "react"; +import { useTranslation } from "react-i18next"; +import { useQuery, useQueryClient } from "@tanstack/react-query"; +import { toast } from "sonner"; +import { + Check, + FolderOpen, + Loader2, + RefreshCw, + Search, + Settings2, +} from "lucide-react"; +import { Alert, AlertDescription } from "@/components/ui/alert"; +import { Button } from "@/components/ui/button"; +import { ConfirmDialog } from "@/components/ConfirmDialog"; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from "@/components/ui/dialog"; +import { Input } from "@/components/ui/input"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/components/ui/select"; +import { Switch } from "@/components/ui/switch"; +import { settingsApi } from "@/lib/api/settings"; +import { usageApi } from "@/lib/api/usage"; +import { + MODELS_DEV_SYNC_CONFIG_QUERY_KEY, + syncModelsDevPricing, +} from "@/lib/modelsDevAutoSync"; +import { + fetchModelsDevPricing, + flattenModels, + formatPrice, + getCommonModelKeys, + type ModelsDevEntry, +} from "@/lib/modelsDevPricing"; +import { usageKeys } from "@/lib/query/usage"; +import type { ModelsDevSyncConfig, ModelsDevSyncState } from "@/types/usage"; +import { isTextEditableTarget } from "@/utils/domUtils"; + +const MODELS_DEV_QUERY_KEY = ["models-dev-pricing"] as const; +const DEFAULT_VISIBLE_ROWS = 80; +const MAX_VISIBLE_ROWS = 300; + +interface AutoSyncDialogProps { + state: ModelsDevSyncState; + onClose: () => void; + onSaved: (state: ModelsDevSyncState) => void; +} + +function AutoSyncDialog({ state, onClose, onSaved }: AutoSyncDialogProps) { + const { t } = useTranslation(); + const [search, setSearch] = useState(""); + const [providerFilter, setProviderFilter] = useState("all"); + const [includeCommonModels, setIncludeCommonModels] = useState( + state.config.includeCommonModels, + ); + const [selectedModelKeys, setSelectedModelKeys] = useState( + () => new Set(state.config.selectedModelKeys), + ); + const [excludedCommonModelKeys, setExcludedCommonModelKeys] = useState( + () => new Set(state.config.excludedCommonModelKeys), + ); + const [isSaving, setIsSaving] = useState(false); + + const { data, isLoading, error, refetch } = useQuery({ + queryKey: MODELS_DEV_QUERY_KEY, + queryFn: fetchModelsDevPricing, + staleTime: 60 * 60 * 1000, + retry: 1, + }); + const entries = useMemo(() => (data ? flattenModels(data) : []), [data]); + const commonModelKeys = useMemo(() => getCommonModelKeys(entries), [entries]); + + const effectiveSelectedKeys = useMemo(() => { + const selected = new Set(selectedModelKeys); + if (includeCommonModels) { + for (const key of commonModelKeys) { + if (!excludedCommonModelKeys.has(key)) selected.add(key); + } + } + return selected; + }, [ + commonModelKeys, + excludedCommonModelKeys, + includeCommonModels, + selectedModelKeys, + ]); + + const providers = useMemo(() => { + const map = new Map(); + for (const entry of entries) { + if (!map.has(entry.providerId)) { + map.set(entry.providerId, entry.providerName); + } + } + return Array.from(map, ([id, name]) => ({ id, name })).sort((a, b) => + a.name.localeCompare(b.name), + ); + }, [entries]); + + const isFiltering = search.trim() !== "" || providerFilter !== "all"; + const filtered = useMemo(() => { + const query = search.trim().toLowerCase(); + return entries.filter( + (entry) => + (providerFilter === "all" || entry.providerId === providerFilter) && + (!query || + entry.modelId.toLowerCase().includes(query) || + entry.normalizedId.includes(query) || + entry.modelName.toLowerCase().includes(query) || + entry.providerName.toLowerCase().includes(query)), + ); + }, [entries, providerFilter, search]); + const visible = useMemo( + () => + filtered.slice(0, isFiltering ? MAX_VISIBLE_ROWS : DEFAULT_VISIBLE_ROWS), + [filtered, isFiltering], + ); + + const toggleEntry = (entry: ModelsDevEntry) => { + const isSelected = effectiveSelectedKeys.has(entry.key); + setSelectedModelKeys((previous) => { + const next = new Set(previous); + if (isSelected) next.delete(entry.key); + else next.add(entry.key); + return next; + }); + setExcludedCommonModelKeys((previous) => { + const next = new Set(previous); + if (isSelected && includeCommonModels && commonModelKeys.has(entry.key)) { + next.add(entry.key); + } else { + next.delete(entry.key); + } + return next; + }); + }; + + const selectFiltered = () => { + setSelectedModelKeys((previous) => { + const next = new Set(previous); + for (const entry of filtered) next.add(entry.key); + return next; + }); + setExcludedCommonModelKeys((previous) => { + const next = new Set(previous); + for (const entry of filtered) next.delete(entry.key); + return next; + }); + }; + + const clearFiltered = () => { + setSelectedModelKeys((previous) => { + const next = new Set(previous); + for (const entry of filtered) next.delete(entry.key); + return next; + }); + if (includeCommonModels) { + setExcludedCommonModelKeys((previous) => { + const next = new Set(previous); + for (const entry of filtered) { + if (commonModelKeys.has(entry.key)) next.add(entry.key); + } + return next; + }); + } + }; + + const save = async () => { + setIsSaving(true); + try { + const config: ModelsDevSyncConfig = { + ...state.config, + includeCommonModels, + selectedModelKeys: Array.from(selectedModelKeys).sort(), + excludedCommonModelKeys: Array.from(excludedCommonModelKeys).sort(), + }; + await usageApi.saveModelsDevSyncConfig(config); + onSaved({ ...state, config }); + toast.success(t("usage.modelsDevAutoSync.selectionSaved")); + onClose(); + } catch (saveError) { + toast.error(String(saveError)); + } finally { + setIsSaving(false); + } + }; + + const priceColumns = (entry: ModelsDevEntry) => + [ + { label: t("usage.inputCost"), value: entry.input }, + { label: t("usage.outputCost"), value: entry.output }, + { label: t("usage.cacheReadCost"), value: entry.cacheRead }, + { label: t("usage.cacheWriteCost"), value: entry.cacheWrite }, + ] as const; + + return ( + !open && !isSaving && onClose()}> + { + if (isTextEditableTarget(event.target)) event.preventDefault(); + }} + > + + + {t("usage.modelsDevAutoSync.configureTitle")} + + + {t("usage.modelsDevAutoSync.configureDescription")} + + + +
+
+
+
+ {t("usage.modelsDevAutoSync.commonModels")} +
+
+ {t("usage.modelsDevAutoSync.commonModelsDescription", { + count: commonModelKeys.size, + })} +
+
+ +
+ + {isLoading ? ( +
+ +
+ ) : error ? ( + + + + {t("usage.modelsDevLoadError")}: {String(error)} + + + + + ) : ( + <> +
+ +
+ + setSearch(event.target.value)} + placeholder={t("usage.modelsDevSearchPlaceholder")} + className="pl-8" + /> +
+ + +
+ +
+ + {t("usage.modelsDevAutoSync.selectedCount", { + count: effectiveSelectedKeys.size, + })} + + {t("usage.modelsDevAutoSync.selectionHint")} +
+ +
+ {filtered.length === 0 ? ( +
+ {t("usage.modelsDevNoResults")} +
+ ) : ( +
+ {visible.map((entry) => { + const selected = effectiveSelectedKeys.has(entry.key); + const common = commonModelKeys.has(entry.key); + return ( + + ); + })} + {filtered.length > visible.length && ( +
+ {isFiltering + ? t("usage.modelsDevTruncated", { + shown: visible.length, + total: filtered.length, + }) + : t("usage.modelsDevDefaultHint", { + shown: visible.length, + total: filtered.length, + })} +
+ )} +
+ )} +
+ + )} +
+ + + + + +
+
+ ); +} + +export function ModelsDevAutoSyncPanel() { + const { t, i18n } = useTranslation(); + const queryClient = useQueryClient(); + const [isDialogOpen, setIsDialogOpen] = useState(false); + const [isSaving, setIsSaving] = useState(false); + const [isSyncing, setIsSyncing] = useState(false); + const [isReloading, setIsReloading] = useState(false); + const [showEnableConfirm, setShowEnableConfirm] = useState(false); + + const { data, isLoading, error, refetch } = useQuery({ + queryKey: MODELS_DEV_SYNC_CONFIG_QUERY_KEY, + queryFn: usageApi.getModelsDevSyncConfig, + staleTime: Number.POSITIVE_INFINITY, + }); + + const updateCachedState = (state: ModelsDevSyncState) => { + queryClient.setQueryData(MODELS_DEV_SYNC_CONFIG_QUERY_KEY, state); + }; + + const saveConfig = async (config: ModelsDevSyncConfig) => { + if (!data) return; + setIsSaving(true); + try { + await usageApi.saveModelsDevSyncConfig(config); + updateCachedState({ ...data, config }); + } catch (saveError) { + toast.error(String(saveError)); + } finally { + setIsSaving(false); + } + }; + + const syncNow = async () => { + if (!data) return; + setIsSyncing(true); + try { + const result = await syncModelsDevPricing(data, true); + await Promise.all([ + refetch(), + queryClient.invalidateQueries({ queryKey: usageKeys.all }), + ]); + toast.success( + t("usage.modelsDevAutoSync.syncSuccess", { + imported: result.imported, + changed: result.changed, + }), + ); + } catch (syncError) { + await refetch(); + toast.error( + t("usage.modelsDevAutoSync.syncFailed", { error: String(syncError) }), + ); + } finally { + setIsSyncing(false); + } + }; + + const reloadLocalFile = async () => { + setIsReloading(true); + try { + await usageApi.getModelPricing(); + await Promise.all([ + refetch(), + queryClient.invalidateQueries({ queryKey: usageKeys.all }), + ]); + toast.success(t("usage.modelsDevAutoSync.localFileReloaded")); + } catch (reloadError) { + toast.error( + t("usage.modelsDevAutoSync.localFileReloadFailed", { + error: String(reloadError), + }), + ); + } finally { + setIsReloading(false); + } + }; + + const openLocalFileFolder = async () => { + try { + await settingsApi.openAppConfigFolder(); + } catch (openError) { + toast.error( + t("usage.modelsDevAutoSync.openFolderFailed", { + error: String(openError), + }), + ); + } + }; + + if (isLoading) { + return ( +
+ +
+ ); + } + + if (error || !data) { + return ( + + + + {t("usage.modelsDevAutoSync.configLoadFailed", { + error: String(error), + })} + + + + + ); + } + + const lastSync = data.config.lastSyncAt + ? new Date(data.config.lastSyncAt).toLocaleString(i18n.resolvedLanguage) + : t("usage.modelsDevAutoSync.neverSynced"); + + const handleAutoSyncChange = (autoSyncEnabled: boolean) => { + if (autoSyncEnabled) { + setShowEnableConfirm(true); + return; + } + void saveConfig({ ...data.config, autoSyncEnabled: false }); + }; + + return ( + <> +
+
+
+
+ {t("usage.modelsDevAutoSync.title")} +
+

+ {t("usage.modelsDevAutoSync.description")} +

+
+
+ + {data.config.autoSyncEnabled + ? t("usage.modelsDevAutoSync.enabled") + : t("usage.modelsDevAutoSync.disabled")} + + +
+
+ +
+
+ {t("usage.modelsDevAutoSync.lastSync")}: {lastSync} +
+
+ {t("usage.modelsDevAutoSync.commonStatus")}:{" "} + {data.config.includeCommonModels + ? t("usage.modelsDevAutoSync.enabled") + : t("usage.modelsDevAutoSync.disabled")} +
+
+ + {data.config.lastSyncError && ( + + + {t("usage.modelsDevAutoSync.lastError", { + error: data.config.lastSyncError, + })} + + + )} + +
+
+ {t("usage.modelsDevAutoSync.localFile")} +
+
+ {data.configPath} +
+
+ +
+ + + + +
+
+ + {isDialogOpen && ( + setIsDialogOpen(false)} + onSaved={updateCachedState} + /> + )} + { + setShowEnableConfirm(false); + void saveConfig({ ...data.config, autoSyncEnabled: true }); + }} + onCancel={() => setShowEnableConfirm(false)} + /> + + ); +} diff --git a/src/components/usage/ModelsDevPickerDialog.tsx b/src/components/usage/ModelsDevPickerDialog.tsx index 03025e160..17927eef2 100644 --- a/src/components/usage/ModelsDevPickerDialog.tsx +++ b/src/components/usage/ModelsDevPickerDialog.tsx @@ -22,114 +22,24 @@ import { SelectValue, } from "@/components/ui/select"; import { useUpdateModelPricing } from "@/lib/query/usage"; +import { + fetchModelsDevPricing, + flattenModels, + formatPrice, + type ModelsDevEntry, +} from "@/lib/modelsDevPricing"; import { isTextEditableTarget } from "@/utils/domUtils"; -const MODELS_DEV_API_URL = "https://models.dev/api.json"; +export { + flattenModels, + formatPrice, + normalizeModelIdForPricing, +} from "@/lib/modelsDevPricing"; + // 全量约 5000 条:默认只展示最新发布的一批,搜索时才做全量匹配 const DEFAULT_VISIBLE_ROWS = 50; const MAX_VISIBLE_ROWS = 200; -interface ModelsDevCost { - input?: number; - output?: number; - cache_read?: number; - cache_write?: number; -} - -interface ModelsDevModel { - id?: string; - name?: string; - release_date?: string; - cost?: ModelsDevCost; -} - -interface ModelsDevProvider { - id?: string; - name?: string; - models?: Record; -} - -type ModelsDevResponse = Record; - -interface ModelsDevEntry { - /** providerId/modelId,同一模型可能出现在多个供应商下 */ - key: string; - providerId: string; - providerName: string; - modelId: string; - /** 实际入库的 ID,与后端 clean_model_id_for_pricing 的归一化规则一致 */ - normalizedId: string; - modelName: string; - /** YYYY-MM-DD 或 YYYY-MM,缺失时为空串 */ - releaseDate: string; - input: number; - output: number; - cacheRead: number; - cacheWrite: number; -} - -/** - * 与后端 clean_model_id_for_pricing(usage_stats.rs)保持一致: - * 取最后一个 '/' 之后的段、去掉 ':' 后缀、'@' 换成 '-'、转小写、去掉 [1m] 标记。 - * 成本归因查询用的就是这种归一化形式,原样入库的 ID 永远匹配不上。 - */ -export function normalizeModelIdForPricing(modelId: string): string { - const afterSlash = modelId.slice(modelId.lastIndexOf("/") + 1); - const beforeColon = afterSlash.split(":")[0] ?? ""; - let normalized = beforeColon.trim().replace(/@/g, "-").toLowerCase(); - if (normalized.endsWith("[1m]")) { - normalized = normalized.slice(0, -"[1m]".length).trim(); - } - return normalized; -} - -/** 转成后端可解析的非负十进制字符串(不能用 String(),小数可能变成科学计数法) */ -export function formatPrice(value: number): string { - if (!Number.isFinite(value) || value <= 0) return "0"; - // toFixed 对 >=1e21 会退化成科学计数法;这种量级的"价格"只可能是脏数据,按 0 处理 - if (value >= 1e12) return "0"; - const trimmed = value.toFixed(6).replace(/0+$/, "").replace(/\.$/, ""); - return trimmed || "0"; -} - -export function flattenModels(data: ModelsDevResponse): ModelsDevEntry[] { - const entries: ModelsDevEntry[] = []; - for (const [providerId, provider] of Object.entries(data)) { - if (!provider || typeof provider !== "object") continue; - const providerName = provider.name || providerId; - for (const [modelId, model] of Object.entries(provider.models ?? {})) { - const cost = model?.cost; - const input = typeof cost?.input === "number" ? cost.input : null; - const output = typeof cost?.output === "number" ? cost.output : null; - if (input === null && output === null) continue; - const normalizedId = normalizeModelIdForPricing(modelId); - if (!normalizedId) continue; - entries.push({ - key: `${providerId}/${modelId}`, - providerId, - providerName, - modelId, - normalizedId, - modelName: model?.name || modelId, - releaseDate: - typeof model?.release_date === "string" ? model.release_date : "", - input: input ?? 0, - output: output ?? 0, - cacheRead: typeof cost?.cache_read === "number" ? cost.cache_read : 0, - cacheWrite: - typeof cost?.cache_write === "number" ? cost.cache_write : 0, - }); - } - } - // 最新发布的排在前面 - entries.sort( - (a, b) => - b.releaseDate.localeCompare(a.releaseDate) || - a.modelName.localeCompare(b.modelName), - ); - return entries; -} - interface ModelsDevPickerDialogProps { open: boolean; onClose: () => void; @@ -160,13 +70,7 @@ export function ModelsDevPickerDialog({ const { data, isLoading, error, refetch } = useQuery({ queryKey: ["models-dev-pricing"], - queryFn: async (): Promise => { - const res = await fetch(MODELS_DEV_API_URL); - if (!res.ok) { - throw new Error(`HTTP ${res.status}`); - } - return res.json(); - }, + queryFn: fetchModelsDevPricing, enabled: open, staleTime: 60 * 60 * 1000, retry: 1, diff --git a/src/components/usage/PricingConfigPanel.tsx b/src/components/usage/PricingConfigPanel.tsx index 1d3a939a2..239652ce0 100644 --- a/src/components/usage/PricingConfigPanel.tsx +++ b/src/components/usage/PricingConfigPanel.tsx @@ -32,6 +32,7 @@ import { isNonNegativeDecimalString, type ModelPricing } from "@/types/usage"; import { Plus, Pencil, Trash2, Loader2 } from "lucide-react"; import { toast } from "sonner"; import { proxyApi } from "@/lib/api/proxy"; +import { ModelsDevAutoSyncPanel } from "./ModelsDevAutoSyncPanel"; const PRICING_APPS = ["claude", "codex", "gemini", "grokbuild"] as const; type PricingApp = (typeof PRICING_APPS)[number]; @@ -341,6 +342,8 @@ export function PricingConfigPanel() { {/* 模型定价配置 */}
+ +

{t("usage.modelPricingDesc")} {t("usage.perMillion")} diff --git a/src/i18n/locales/en.json b/src/i18n/locales/en.json index 6e682ab8a..71a02d342 100644 --- a/src/i18n/locales/en.json +++ b/src/i18n/locales/en.json @@ -1633,6 +1633,40 @@ "modelsDevNoResults": "No matching models", "modelsDevTruncated": "Showing first {{shown}} of {{total}} results — refine your search", "modelsDevDefaultHint": "Showing the {{shown}} most recently released models (of {{total}}) — type to search all", + "modelsDevAutoSync": { + "title": "Automatic models.dev pricing sync", + "description": "When enabled, periodically fetch selected prices when CC Switch launches and keep them in your local pricing file.", + "enabled": "Enabled", + "disabled": "Disabled", + "enableConfirmTitle": "Enable automatic pricing sync?", + "enableConfirmMessage": "After enabling, CC Switch will periodically update prices for the selected models from models.dev when it launches (at most once every 6 hours). Built-in and manually configured prices with the same model ID will be overwritten.", + "enableConfirmAction": "Enable automatic sync", + "configure": "Choose models", + "configureTitle": "Choose models for automatic pricing sync", + "configureDescription": "Search and filter models to update automatically. Your choices are stored locally.", + "commonModels": "Automatically include common models", + "commonModelsDescription": "Selects {{count}} recent Claude, GPT, Gemini, Grok, DeepSeek, Qwen, MiMo, LongCat, Kimi, MiniMax, and GLM models.", + "selectFiltered": "Select filtered ({{count}})", + "clearFiltered": "Clear filtered", + "selectedCount": "{{count}} models selected", + "selectionHint": "Selections with the same normalized model ID are imported once.", + "commonBadge": "Common", + "selectionSaved": "Automatic pricing selection saved", + "syncSuccess": "Pricing sync complete: {{imported}} models checked, {{changed}} updated", + "syncFailed": "Pricing sync failed: {{error}}", + "configLoadFailed": "Failed to load local pricing configuration: {{error}}", + "neverSynced": "Never", + "lastSync": "Last sync", + "commonStatus": "Common models", + "lastError": "Last automatic sync failed: {{error}}", + "localFile": "Local pricing file", + "openFolder": "Open folder", + "openFolderFailed": "Failed to open the local pricing folder: {{error}}", + "reloadLocalFile": "Reload file", + "localFileReloaded": "Local pricing file reloaded", + "localFileReloadFailed": "Failed to reload local pricing file: {{error}}", + "syncNow": "Sync now" + }, "cacheReadCostPerMillion": "Cache Read Cost (per million tokens, USD)", "cacheCreationCostPerMillion": "Cache Write Cost (per million tokens, USD)" }, diff --git a/src/i18n/locales/ja.json b/src/i18n/locales/ja.json index af05d7f15..db7f37b25 100644 --- a/src/i18n/locales/ja.json +++ b/src/i18n/locales/ja.json @@ -1633,6 +1633,40 @@ "modelsDevNoResults": "一致するモデルがありません", "modelsDevTruncated": "{{total}} 件中、先頭 {{shown}} 件のみ表示しています。検索条件を絞り込んでください", "modelsDevDefaultHint": "最新リリース順に {{shown}} 件を表示しています(全 {{total}} 件)。キーワード入力で全件検索できます", + "modelsDevAutoSync": { + "title": "models.dev 料金の自動同期", + "description": "有効にすると、CC Switch の起動時に選択したモデルの最新料金を定期的に取得し、ローカル料金ファイルに保存します。", + "enabled": "有効", + "disabled": "無効", + "enableConfirmTitle": "料金の自動同期を有効にしますか?", + "enableConfirmMessage": "有効にすると、CC Switch の起動時に models.dev から選択したモデルの料金を定期的に更新します(最大 6 時間に 1 回)。同じモデル ID の内蔵料金と手動設定した料金は上書きされます。", + "enableConfirmAction": "自動同期を有効にする", + "configure": "モデルを選択", + "configureTitle": "料金を自動同期するモデルを選択", + "configureDescription": "検索とプロバイダーの絞り込みで自動更新するモデルを選択します。選択内容はローカルに保存されます。", + "commonModels": "よく使うモデルを自動的に含める", + "commonModelsDescription": "最近の Claude、GPT、Gemini、Grok、DeepSeek、Qwen、MiMo、LongCat、Kimi、MiniMax、GLM から {{count}} モデルを選択します。", + "selectFiltered": "絞り込み結果を選択({{count}})", + "clearFiltered": "絞り込み結果を解除", + "selectedCount": "{{count}} モデルを選択中", + "selectionHint": "正規化後のモデル ID が同じ項目は一度だけインポートされます。", + "commonBadge": "よく使う", + "selectionSaved": "自動料金同期の選択を保存しました", + "syncSuccess": "料金同期が完了しました:{{imported}} モデルを確認、{{changed}} モデルを更新", + "syncFailed": "料金同期に失敗しました:{{error}}", + "configLoadFailed": "ローカル料金設定の読み込みに失敗しました:{{error}}", + "neverSynced": "未同期", + "lastSync": "最終同期", + "commonStatus": "よく使うモデル", + "lastError": "前回の自動同期に失敗しました:{{error}}", + "localFile": "ローカル料金ファイル", + "openFolder": "フォルダーを開く", + "openFolderFailed": "ローカル料金フォルダーを開けませんでした: {{error}}", + "reloadLocalFile": "ファイルを再読み込み", + "localFileReloaded": "ローカル料金ファイルを再読み込みしました", + "localFileReloadFailed": "ローカル料金ファイルの再読み込みに失敗しました:{{error}}", + "syncNow": "今すぐ同期" + }, "cacheReadCostPerMillion": "キャッシュ読み取りコスト(100万トークンあたり、USD)", "cacheCreationCostPerMillion": "キャッシュ書き込みコスト(100万トークンあたり、USD)" }, diff --git a/src/i18n/locales/zh-TW.json b/src/i18n/locales/zh-TW.json index c9673be55..f6bc8c30d 100644 --- a/src/i18n/locales/zh-TW.json +++ b/src/i18n/locales/zh-TW.json @@ -1604,6 +1604,40 @@ "modelsDevNoResults": "沒有符合的模型", "modelsDevTruncated": "僅顯示前 {{shown}} 條,共 {{total}} 條結果,請縮小搜尋範圍", "modelsDevDefaultHint": "預設展示最新發布的 {{shown}} 個模型(共 {{total}} 個),輸入關鍵字可全量搜尋", + "modelsDevAutoSync": { + "title": "自動同步 models.dev 定價", + "description": "開啟後,CC Switch 會在啟動時定期擷取所選模型的最新價格,並儲存到本機定價檔案。", + "enabled": "已開啟", + "disabled": "已關閉", + "enableConfirmTitle": "開啟自動同步定價?", + "enableConfirmMessage": "開啟後,CC Switch 會在啟動時定期從 models.dev 更新所選模型的價格(最多每 6 小時一次)。相同模型 ID 的軟體內建價格和手動設定價格都會被覆寫。", + "enableConfirmAction": "開啟自動同步", + "configure": "選擇模型", + "configureTitle": "選擇自動同步定價的模型", + "configureDescription": "透過搜尋和供應商篩選選擇需要自動更新的模型,選擇結果儲存在本機。", + "commonModels": "自動包含常用模型", + "commonModelsDescription": "自動選擇 {{count}} 個近期 Claude、GPT、Gemini、Grok、DeepSeek、Qwen、MiMo、LongCat、Kimi、MiniMax 和 GLM 模型。", + "selectFiltered": "全選篩選結果({{count}})", + "clearFiltered": "清除篩選結果", + "selectedCount": "已選擇 {{count}} 個模型", + "selectionHint": "標準化模型 ID 相同的選項只會匯入一次。", + "commonBadge": "常用", + "selectionSaved": "自動定價同步選擇已儲存", + "syncSuccess": "定價同步完成:檢查 {{imported}} 個模型,更新 {{changed}} 個", + "syncFailed": "定價同步失敗:{{error}}", + "configLoadFailed": "載入本機定價設定失敗:{{error}}", + "neverSynced": "從未同步", + "lastSync": "上次同步", + "commonStatus": "常用模型", + "lastError": "上次自動同步失敗:{{error}}", + "localFile": "本機定價檔案", + "openFolder": "開啟資料夾", + "openFolderFailed": "開啟本機定價檔案資料夾失敗:{{error}}", + "reloadLocalFile": "重新載入檔案", + "localFileReloaded": "本機定價檔案已重新載入", + "localFileReloadFailed": "重新載入本機定價檔案失敗:{{error}}", + "syncNow": "立即同步" + }, "cacheReadCostPerMillion": "快取讀取成本 (每百萬 tokens, USD)", "cacheCreationCostPerMillion": "快取寫入成本 (每百萬 tokens, USD)" }, diff --git a/src/i18n/locales/zh.json b/src/i18n/locales/zh.json index 550f4a882..271db76bc 100644 --- a/src/i18n/locales/zh.json +++ b/src/i18n/locales/zh.json @@ -1633,6 +1633,40 @@ "modelsDevNoResults": "没有匹配的模型", "modelsDevTruncated": "仅显示前 {{shown}} 条,共 {{total}} 条结果,请缩小搜索范围", "modelsDevDefaultHint": "默认展示最新发布的 {{shown}} 个模型(共 {{total}} 个),输入关键字可全量搜索", + "modelsDevAutoSync": { + "title": "自动同步 models.dev 定价", + "description": "开启后,CC Switch 会在启动时定期拉取所选模型的最新价格,并保存到本地定价文件。", + "enabled": "已开启", + "disabled": "已关闭", + "enableConfirmTitle": "开启自动同步定价?", + "enableConfirmMessage": "开启后,CC Switch 会在启动时定期从 models.dev 更新所选模型的价格(最多每 6 小时一次)。同一模型 ID 的软件内置价格和手动设置价格都会被覆盖。", + "enableConfirmAction": "开启自动同步", + "configure": "选择模型", + "configureTitle": "选择自动同步定价的模型", + "configureDescription": "通过搜索和供应商筛选选择需要自动更新的模型,选择结果保存在本地。", + "commonModels": "自动包含常用模型", + "commonModelsDescription": "自动选择 {{count}} 个近期 Claude、GPT、Gemini、Grok、DeepSeek、Qwen、MiMo、LongCat、Kimi、MiniMax 和 GLM 模型。", + "selectFiltered": "全选筛选结果({{count}})", + "clearFiltered": "清空筛选结果", + "selectedCount": "已选择 {{count}} 个模型", + "selectionHint": "归一化模型 ID 相同的选项只会导入一次。", + "commonBadge": "常用", + "selectionSaved": "自动定价同步选择已保存", + "syncSuccess": "定价同步完成:检查 {{imported}} 个模型,更新 {{changed}} 个", + "syncFailed": "定价同步失败:{{error}}", + "configLoadFailed": "加载本地定价配置失败:{{error}}", + "neverSynced": "从未同步", + "lastSync": "上次同步", + "commonStatus": "常用模型", + "lastError": "上次自动同步失败:{{error}}", + "localFile": "本地定价文件", + "openFolder": "打开目录", + "openFolderFailed": "打开本地定价文件目录失败:{{error}}", + "reloadLocalFile": "重新加载文件", + "localFileReloaded": "本地定价文件已重新加载", + "localFileReloadFailed": "重新加载本地定价文件失败:{{error}}", + "syncNow": "立即同步" + }, "cacheReadCostPerMillion": "缓存读取成本 (每百万 tokens, USD)", "cacheCreationCostPerMillion": "缓存写入成本 (每百万 tokens, USD)" }, diff --git a/src/lib/api/usage.ts b/src/lib/api/usage.ts index 85c09acb8..fe59c3798 100644 --- a/src/lib/api/usage.ts +++ b/src/lib/api/usage.ts @@ -8,6 +8,8 @@ import type { RequestLog, LogFilters, ModelPricing, + ModelsDevSyncConfig, + ModelsDevSyncState, ProviderLimitStatus, PaginatedLogs, SessionSyncResult, @@ -164,6 +166,27 @@ export const usageApi = { }); }, + updateModelPricingBatch: async (entries: ModelPricing[]): Promise => { + return invoke("update_model_pricing_batch", { entries }); + }, + + getModelsDevSyncConfig: async (): Promise => { + return invoke("get_models_dev_sync_config"); + }, + + saveModelsDevSyncConfig: async ( + config: ModelsDevSyncConfig, + ): Promise => { + return invoke("save_models_dev_sync_config", { config }); + }, + + recordModelsDevSyncResult: async ( + syncedAt: number | null, + error: string | null, + ): Promise => { + return invoke("record_models_dev_sync_result", { syncedAt, error }); + }, + deleteModelPricing: async (modelId: string): Promise => { return invoke("delete_model_pricing", { modelId }); }, diff --git a/src/lib/modelsDevAutoSync.ts b/src/lib/modelsDevAutoSync.ts new file mode 100644 index 000000000..db81df5ba --- /dev/null +++ b/src/lib/modelsDevAutoSync.ts @@ -0,0 +1,93 @@ +import { usageApi } from "@/lib/api/usage"; +import { + fetchModelsDevPricing, + flattenModels, + resolveModelsDevSelection, + toModelPricing, +} from "@/lib/modelsDevPricing"; +import type { ModelsDevSyncState } from "@/types/usage"; + +export interface ModelsDevSyncResult { + skipped: boolean; + selected: number; + imported: number; + changed: number; + syncedAt: number | null; +} + +export const MODELS_DEV_SYNC_CONFIG_QUERY_KEY = [ + "models-dev-sync-config", +] as const; +export const MODELS_DEV_STARTUP_SYNC_INTERVAL_MS = 6 * 60 * 60 * 1000; + +const errorMessage = (error: unknown) => + error instanceof Error ? error.message : String(error); + +export async function syncModelsDevPricing( + state?: ModelsDevSyncState, + force = false, +): Promise { + const initialState = state ?? (await usageApi.getModelsDevSyncConfig()); + const recentlySynced = + initialState.config.lastSyncAt !== null && + Date.now() - initialState.config.lastSyncAt < + MODELS_DEV_STARTUP_SYNC_INTERVAL_MS; + if (!force && (!initialState.config.autoSyncEnabled || recentlySynced)) { + return { + skipped: true, + selected: 0, + imported: 0, + changed: 0, + syncedAt: initialState.config.lastSyncAt, + }; + } + + try { + const data = await fetchModelsDevPricing(); + const latestState = await usageApi.getModelsDevSyncConfig(); + if (!force && !latestState.config.autoSyncEnabled) { + return { + skipped: true, + selected: 0, + imported: 0, + changed: 0, + syncedAt: latestState.config.lastSyncAt, + }; + } + const selectedEntries = resolveModelsDevSelection( + flattenModels(data), + latestState.config, + ); + const pricing = toModelPricing(selectedEntries); + const changed = pricing.length + ? await usageApi.updateModelPricingBatch(pricing) + : 0; + const syncedAt = Date.now(); + await usageApi.recordModelsDevSyncResult(syncedAt, null); + return { + skipped: false, + selected: selectedEntries.length, + imported: pricing.length, + changed, + syncedAt, + }; + } catch (error) { + try { + await usageApi.recordModelsDevSyncResult(null, errorMessage(error)); + } catch (saveError) { + console.warn( + "[models.dev] Failed to persist automatic sync error", + saveError, + ); + } + throw error; + } +} + +let startupSync: Promise | null = null; + +/** Run once per renderer and at most once per interval across WebView rebuilds. */ +export function syncModelsDevPricingOnStartup(): Promise { + startupSync ??= syncModelsDevPricing(); + return startupSync; +} diff --git a/src/lib/modelsDevPricing.ts b/src/lib/modelsDevPricing.ts new file mode 100644 index 000000000..6a572649a --- /dev/null +++ b/src/lib/modelsDevPricing.ts @@ -0,0 +1,277 @@ +import type { ModelPricing, ModelsDevSyncConfig } from "@/types/usage"; + +export const MODELS_DEV_API_URL = "https://models.dev/api.json"; +const MODELS_DEV_FETCH_TIMEOUT_MS = 15_000; + +export interface ModelsDevCost { + input?: number; + output?: number; + cache_read?: number; + cache_write?: number; +} + +export interface ModelsDevModalities { + input?: string[]; + output?: string[]; +} + +export interface ModelsDevModel { + id?: string; + name?: string; + release_date?: string; + cost?: ModelsDevCost; + modalities?: ModelsDevModalities; + status?: string; +} + +export interface ModelsDevProvider { + id?: string; + name?: string; + models?: Record; +} + +export type ModelsDevResponse = Record; + +export interface ModelsDevEntry { + key: string; + providerId: string; + providerName: string; + modelId: string; + normalizedId: string; + modelName: string; + releaseDate: string; + input: number; + output: number; + cacheRead: number; + cacheWrite: number; +} + +const NON_TEXT_MODEL_MARKERS = [ + "audio", + "deprecated", + "embedding", + "image", + "moderation", + "realtime", + "transcribe", + "tts", + "video", +]; +const NON_TEXT_OUTPUT_MODALITIES = new Set(["audio", "image", "video"]); + +const isTextPricingModel = (modelId: string, model?: ModelsDevModel) => { + if (model?.status?.toLowerCase() === "deprecated") return false; + + const outputModalities = model?.modalities?.output + ?.filter((modality): modality is string => typeof modality === "string") + .map((modality) => modality.toLowerCase()); + if ( + outputModalities?.length && + (!outputModalities.includes("text") || + outputModalities.some((modality) => + NON_TEXT_OUTPUT_MODALITIES.has(modality), + )) + ) { + return false; + } + + const searchableName = `${modelId} ${model?.name ?? ""}`.toLowerCase(); + return !NON_TEXT_MODEL_MARKERS.some((marker) => + searchableName.includes(marker), + ); +}; + +export function normalizeModelIdForPricing(modelId: string): string { + const afterSlash = modelId.slice(modelId.lastIndexOf("/") + 1); + const beforeColon = afterSlash.split(":")[0] ?? ""; + let normalized = beforeColon.trim().replace(/@/g, "-").toLowerCase(); + if (normalized.endsWith("[1m]")) { + normalized = normalized.slice(0, -"[1m]".length).trim(); + } + return normalized; +} + +export function formatPrice(value: number): string { + if (!Number.isFinite(value) || value <= 0) return "0"; + if (value >= 1e12) return "0"; + const trimmed = value.toFixed(6).replace(/0+$/, "").replace(/\.$/, ""); + return trimmed || "0"; +} + +export function flattenModels(data: ModelsDevResponse): ModelsDevEntry[] { + const entries: ModelsDevEntry[] = []; + for (const [providerId, provider] of Object.entries(data)) { + if (!provider || typeof provider !== "object") continue; + const providerName = provider.name || providerId; + for (const [modelId, model] of Object.entries(provider.models ?? {})) { + if (!isTextPricingModel(modelId, model)) continue; + const cost = model?.cost; + const input = typeof cost?.input === "number" ? cost.input : null; + const output = typeof cost?.output === "number" ? cost.output : null; + if (input === null && output === null) continue; + const normalizedId = normalizeModelIdForPricing(modelId); + if (!normalizedId) continue; + entries.push({ + key: `${providerId}/${modelId}`, + providerId, + providerName, + modelId, + normalizedId, + modelName: model?.name || modelId, + releaseDate: + typeof model?.release_date === "string" ? model.release_date : "", + input: input ?? 0, + output: output ?? 0, + cacheRead: typeof cost?.cache_read === "number" ? cost.cache_read : 0, + cacheWrite: + typeof cost?.cache_write === "number" ? cost.cache_write : 0, + }); + } + } + entries.sort( + (a, b) => + b.releaseDate.localeCompare(a.releaseDate) || + a.modelName.localeCompare(b.modelName), + ); + return entries; +} + +export async function fetchModelsDevPricing(): Promise { + const controller = new AbortController(); + const timeout = window.setTimeout( + () => controller.abort(), + MODELS_DEV_FETCH_TIMEOUT_MS, + ); + try { + const response = await fetch(MODELS_DEV_API_URL, { + signal: controller.signal, + }); + if (!response.ok) { + throw new Error(`HTTP ${response.status}`); + } + return (await response.json()) as ModelsDevResponse; + } finally { + window.clearTimeout(timeout); + } +} + +const COMMON_MODEL_LIMIT_PER_FAMILY = 6; + +interface CommonFamilyRule { + id: string; + providers: ReadonlySet; + matches: (modelId: string) => boolean; +} + +const COMMON_FAMILY_RULES: CommonFamilyRule[] = [ + { + id: "claude", + providers: new Set(["anthropic"]), + matches: (modelId) => modelId.startsWith("claude-"), + }, + { + id: "gpt", + providers: new Set(["openai"]), + matches: (modelId) => + modelId.startsWith("gpt-") || + modelId.startsWith("o1-") || + modelId.startsWith("o3-") || + modelId.startsWith("o4-"), + }, + { + id: "gemini", + providers: new Set(["google"]), + matches: (modelId) => modelId.startsWith("gemini-"), + }, + { + id: "grok", + providers: new Set(["xai"]), + matches: (modelId) => modelId.startsWith("grok-"), + }, + { + id: "deepseek", + providers: new Set(["deepseek"]), + matches: (modelId) => modelId.startsWith("deepseek-"), + }, + { + id: "qwen", + providers: new Set(["alibaba"]), + matches: (modelId) => modelId.startsWith("qwen"), + }, + { + id: "mimo", + providers: new Set(["xiaomi"]), + matches: (modelId) => modelId.startsWith("mimo-"), + }, + { + id: "longcat", + providers: new Set(["longcat"]), + matches: (modelId) => modelId.startsWith("longcat-"), + }, + { + id: "kimi", + providers: new Set(["moonshotai"]), + matches: (modelId) => modelId.startsWith("kimi-"), + }, + { + id: "minimax", + providers: new Set(["minimax-cn"]), + matches: (modelId) => modelId.startsWith("minimax-m"), + }, + { + id: "glm", + providers: new Set(["zai"]), + matches: (modelId) => modelId.startsWith("glm-"), + }, +]; + +/** Pick a bounded, canonical set of recent chat/coding models per family. */ +export function getCommonModelKeys(entries: ModelsDevEntry[]): Set { + const keys = new Set(); + for (const rule of COMMON_FAMILY_RULES) { + let count = 0; + for (const entry of entries) { + if ( + rule.providers.has(entry.providerId) && + rule.matches(entry.modelId.toLowerCase()) + ) { + keys.add(entry.key); + count += 1; + if (count >= COMMON_MODEL_LIMIT_PER_FAMILY) break; + } + } + } + return keys; +} + +export function resolveModelsDevSelection( + entries: ModelsDevEntry[], + config: ModelsDevSyncConfig, +): ModelsDevEntry[] { + const explicit = new Set(config.selectedModelKeys); + const excluded = new Set(config.excludedCommonModelKeys); + const common = config.includeCommonModels + ? getCommonModelKeys(entries) + : new Set(); + return entries.filter( + (entry) => + explicit.has(entry.key) || + (common.has(entry.key) && !excluded.has(entry.key)), + ); +} + +export function toModelPricing(entries: ModelsDevEntry[]): ModelPricing[] { + const byModelId = new Map(); + for (const entry of entries) { + if (byModelId.has(entry.normalizedId)) continue; + byModelId.set(entry.normalizedId, { + modelId: entry.normalizedId, + displayName: entry.modelName, + inputCostPerMillion: formatPrice(entry.input), + outputCostPerMillion: formatPrice(entry.output), + cacheReadCostPerMillion: formatPrice(entry.cacheRead), + cacheCreationCostPerMillion: formatPrice(entry.cacheWrite), + }); + } + return Array.from(byModelId.values()); +} diff --git a/src/main.tsx b/src/main.tsx index fe4af84be..bc15c43ee 100644 --- a/src/main.tsx +++ b/src/main.tsx @@ -19,6 +19,10 @@ import { installGlobalErrorHandlers, reportFrontendError, } from "./lib/frontendLogger"; +import { + MODELS_DEV_SYNC_CONFIG_QUERY_KEY, + syncModelsDevPricingOnStartup, +} from "./lib/modelsDevAutoSync"; installGlobalErrorHandlers(); @@ -124,6 +128,25 @@ async function bootstrap() { , ); + + void syncModelsDevPricingOnStartup() + .then((result) => { + if (!result.skipped) { + return Promise.all([ + queryClient.invalidateQueries({ queryKey: ["usage"] }), + queryClient.invalidateQueries({ + queryKey: MODELS_DEV_SYNC_CONFIG_QUERY_KEY, + }), + ]); + } + }) + .catch((error) => { + // 离线或 models.dev 暂时不可用不应阻塞应用启动。 + reportFrontendError("models_dev_startup_sync", error); + void queryClient.invalidateQueries({ + queryKey: MODELS_DEV_SYNC_CONFIG_QUERY_KEY, + }); + }); } void bootstrap(); diff --git a/src/types/usage.ts b/src/types/usage.ts index 840fdf89f..08f73c262 100644 --- a/src/types/usage.ts +++ b/src/types/usage.ts @@ -67,6 +67,20 @@ export interface ModelPricing { cacheCreationCostPerMillion: string; } +export interface ModelsDevSyncConfig { + autoSyncEnabled: boolean; + includeCommonModels: boolean; + selectedModelKeys: string[]; + excludedCommonModelKeys: string[]; + lastSyncAt: number | null; + lastSyncError: string | null; +} + +export interface ModelsDevSyncState { + config: ModelsDevSyncConfig; + configPath: string; +} + export interface UsageSummary { totalRequests: number; totalCost: string; diff --git a/tests/components/ModelsDevAutoSyncPanel.test.tsx b/tests/components/ModelsDevAutoSyncPanel.test.tsx new file mode 100644 index 000000000..5dd7e00bb --- /dev/null +++ b/tests/components/ModelsDevAutoSyncPanel.test.tsx @@ -0,0 +1,226 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const { + getModelsDevSyncConfig, + saveModelsDevSyncConfig, + getModelPricing, + openAppConfigFolder, + syncModelsDevPricing, +} = vi.hoisted(() => ({ + getModelsDevSyncConfig: vi.fn(), + saveModelsDevSyncConfig: vi.fn(), + getModelPricing: vi.fn(), + openAppConfigFolder: vi.fn(), + syncModelsDevPricing: vi.fn(), +})); + +vi.mock("react-i18next", () => ({ + useTranslation: () => ({ + t: (key: string, options?: { count?: number }) => + options?.count == null ? key : `${key}:${options.count}`, + i18n: { resolvedLanguage: "en" }, + }), +})); + +vi.mock("sonner", () => ({ + toast: { success: vi.fn(), error: vi.fn() }, +})); + +vi.mock("@/lib/api/usage", () => ({ + usageApi: { + getModelsDevSyncConfig, + saveModelsDevSyncConfig, + getModelPricing, + }, +})); + +vi.mock("@/lib/api/settings", () => ({ + settingsApi: { openAppConfigFolder }, +})); + +vi.mock("@/lib/modelsDevAutoSync", () => ({ + MODELS_DEV_SYNC_CONFIG_QUERY_KEY: ["models-dev-sync-config"], + syncModelsDevPricing, +})); + +import { ModelsDevAutoSyncPanel } from "@/components/usage/ModelsDevAutoSyncPanel"; + +const state = { + configPath: "C:/Users/test/.cc-switch/model-pricing.json", + config: { + autoSyncEnabled: false, + includeCommonModels: true, + selectedModelKeys: [], + excludedCommonModelKeys: [], + lastSyncAt: null, + lastSyncError: null, + }, +}; + +function renderPanel() { + const client = new QueryClient({ + defaultOptions: { queries: { retry: false } }, + }); + return render( + + + , + ); +} + +describe("ModelsDevAutoSyncPanel", () => { + beforeEach(() => { + vi.clearAllMocks(); + getModelsDevSyncConfig.mockResolvedValue(state); + saveModelsDevSyncConfig.mockResolvedValue(undefined); + getModelPricing.mockResolvedValue([]); + openAppConfigFolder.mockResolvedValue(undefined); + syncModelsDevPricing.mockResolvedValue({ + skipped: false, + selected: 2, + imported: 2, + changed: 1, + syncedAt: Date.now(), + }); + vi.stubGlobal( + "fetch", + vi.fn().mockResolvedValue({ + ok: true, + json: async () => ({ + openai: { + name: "OpenAI", + models: { + "gpt-5": { + name: "GPT-5", + release_date: "2025-08-01", + cost: { input: 1, output: 2 }, + }, + }, + }, + deepseek: { + name: "DeepSeek", + models: { + "deepseek-chat": { + name: "DeepSeek Chat", + release_date: "2025-12-01", + cost: { input: 0.3, output: 1.2 }, + }, + }, + }, + }), + }), + ); + }); + + it("loads automatic sync as disabled by default", async () => { + renderPanel(); + + expect( + await screen.findByText("usage.modelsDevAutoSync.title"), + ).toBeInTheDocument(); + expect(screen.getByText(state.configPath)).toBeInTheDocument(); + expect(screen.getByRole("switch")).not.toBeChecked(); + expect(saveModelsDevSyncConfig).not.toHaveBeenCalled(); + }); + + it("persists disabling without showing the overwrite warning", async () => { + const enabledState = { + ...state, + config: { ...state.config, autoSyncEnabled: true }, + }; + getModelsDevSyncConfig.mockResolvedValue(enabledState); + renderPanel(); + + fireEvent.click(await screen.findByRole("switch")); + await waitFor(() => + expect(saveModelsDevSyncConfig).toHaveBeenCalledWith({ + ...enabledState.config, + autoSyncEnabled: false, + }), + ); + expect( + screen.queryByText("usage.modelsDevAutoSync.enableConfirmTitle"), + ).not.toBeInTheDocument(); + }); + + it("warns about price overwrites before enabling automatic sync", async () => { + renderPanel(); + + fireEvent.click(await screen.findByRole("switch")); + + expect(saveModelsDevSyncConfig).not.toHaveBeenCalled(); + expect( + await screen.findByText("usage.modelsDevAutoSync.enableConfirmTitle"), + ).toBeInTheDocument(); + fireEvent.click( + screen.getByRole("button", { + name: "usage.modelsDevAutoSync.enableConfirmAction", + }), + ); + + await waitFor(() => + expect(saveModelsDevSyncConfig).toHaveBeenCalledWith({ + ...state.config, + autoSyncEnabled: true, + }), + ); + }); + + it("keeps automatic sync disabled when the overwrite warning is cancelled", async () => { + renderPanel(); + + fireEvent.click(await screen.findByRole("switch")); + fireEvent.click( + await screen.findByRole("button", { name: "common.cancel" }), + ); + + expect(saveModelsDevSyncConfig).not.toHaveBeenCalled(); + expect(screen.getByRole("switch")).not.toBeChecked(); + }); + + it("reloads the automatic sync config after reading the local pricing file", async () => { + const initialState = { + ...state, + config: { ...state.config, autoSyncEnabled: true }, + }; + getModelsDevSyncConfig + .mockResolvedValueOnce(initialState) + .mockResolvedValue(state); + renderPanel(); + + fireEvent.click( + await screen.findByRole("button", { + name: "usage.modelsDevAutoSync.reloadLocalFile", + }), + ); + + await waitFor(() => + expect(getModelsDevSyncConfig).toHaveBeenCalledTimes(2), + ); + expect(getModelPricing).toHaveBeenCalledTimes(1); + await waitFor(() => expect(screen.getByRole("switch")).not.toBeChecked()); + }); + + it("opens the searchable multi-select dialog with common models selected", async () => { + renderPanel(); + fireEvent.click( + await screen.findByRole("button", { + name: "usage.modelsDevAutoSync.configure", + }), + ); + + expect( + await screen.findByText("usage.modelsDevAutoSync.configureTitle"), + ).toBeInTheDocument(); + expect(await screen.findByText("GPT-5")).toBeInTheDocument(); + expect(screen.getByText("DeepSeek Chat")).toBeInTheDocument(); + expect( + screen.getByText("usage.modelsDevAutoSync.selectedCount:2"), + ).toBeInTheDocument(); + expect( + screen.getAllByText("usage.modelsDevAutoSync.commonBadge"), + ).toHaveLength(2); + }); +}); diff --git a/tests/components/ModelsDevPickerDialog.test.ts b/tests/components/ModelsDevPickerDialog.test.ts index a162759d8..831c03f59 100644 --- a/tests/components/ModelsDevPickerDialog.test.ts +++ b/tests/components/ModelsDevPickerDialog.test.ts @@ -5,6 +5,11 @@ import { formatPrice, normalizeModelIdForPricing, } from "@/components/usage/ModelsDevPickerDialog"; +import { + getCommonModelKeys, + resolveModelsDevSelection, + toModelPricing, +} from "@/lib/modelsDevPricing"; describe("normalizeModelIdForPricing", () => { it("keeps already-normalized ids unchanged", () => { @@ -141,4 +146,183 @@ describe("flattenModels", () => { // 完全没有定价的模型被过滤 expect(entries.some((e) => e.modelId === "free-model")).toBe(false); }); + + it("filters deprecated and non-text output models while keeping multimodal input models", () => { + const entries = flattenModels({ + acme: { + models: { + "multimodal-chat": { + name: "Multimodal Chat", + modalities: { + input: ["text", "image", "audio", "video"], + output: ["text"], + }, + cost: { input: 1, output: 2 }, + }, + "legacy-chat": { + status: "deprecated", + modalities: { output: ["text"] }, + cost: { input: 1, output: 2 }, + }, + "speech-model": { + modalities: { output: ["audio"] }, + cost: { input: 1, output: 2 }, + }, + "mixed-output-model": { + modalities: { output: ["text", "audio"] }, + cost: { input: 1, output: 2 }, + }, + "movie-generator": { + modalities: { output: ["video"] }, + cost: { input: 1, output: 2 }, + }, + "fallback-video-model": { + cost: { input: 1, output: 2 }, + }, + }, + }, + }); + + expect(entries.map((entry) => entry.modelId)).toEqual(["multimodal-chat"]); + }); + + it("selects a bounded canonical set of common model families", () => { + const openAiModels = Object.fromEntries( + Array.from({ length: 7 }, (_, index) => { + const version = index + 1; + return [ + `gpt-${version}`, + { + name: `GPT ${version}`, + release_date: `2025-0${version}-01`, + cost: { input: version, output: version * 2 }, + }, + ]; + }), + ); + const entries = flattenModels({ + openai: { + name: "OpenAI", + models: { + ...openAiModels, + "gpt-image-1": { + release_date: "2026-01-01", + cost: { input: 1, output: 2 }, + }, + }, + }, + aggregator: { + name: "Aggregator", + models: { + "gpt-7": { + release_date: "2026-02-01", + cost: { input: 9, output: 18 }, + }, + }, + }, + anthropic: { + name: "Anthropic", + models: { + "claude-sonnet-5": { + release_date: "2026-06-01", + cost: { input: 3, output: 15 }, + }, + }, + }, + deepseek: { + name: "DeepSeek", + models: { + "deepseek-chat": { + release_date: "2025-12-01", + cost: { input: 0.3, output: 1.2 }, + }, + }, + }, + xiaomi: { + name: "Xiaomi", + models: { + "mimo-v2.5": { + release_date: "2026-04-01", + cost: { input: 0.2, output: 1 }, + }, + "mimo-v2.5-tts": { + release_date: "2026-05-01", + cost: { input: 0.1, output: 0.5 }, + }, + }, + }, + longcat: { + name: "LongCat", + models: { + "LongCat-2.0": { + release_date: "2026-03-01", + cost: { input: 0.4, output: 1.6 }, + }, + }, + }, + }); + + const common = getCommonModelKeys(entries); + expect(common.has("openai/gpt-image-1")).toBe(false); + expect(common.has("aggregator/gpt-7")).toBe(false); + expect(common.has("openai/gpt-7")).toBe(true); + expect(common.has("openai/gpt-1")).toBe(false); + expect(common.has("anthropic/claude-sonnet-5")).toBe(true); + expect(common.has("deepseek/deepseek-chat")).toBe(true); + expect(common.has("xiaomi/mimo-v2.5")).toBe(true); + expect(common.has("xiaomi/mimo-v2.5-tts")).toBe(false); + expect(common.has("longcat/LongCat-2.0")).toBe(true); + }); + + it("combines common and explicit selections and deduplicates normalized ids", () => { + const entries = flattenModels({ + openai: { + models: { + "gpt-5": { + name: "GPT-5 Official", + release_date: "2025-08-01", + cost: { input: 1, output: 2 }, + }, + }, + }, + relay: { + models: { + "vendor/GPT-5": { + name: "GPT-5 Relay", + release_date: "2025-07-01", + cost: { input: 9, output: 18 }, + }, + "custom-model": { + name: "Custom", + release_date: "2025-06-01", + cost: { input: 0.5, output: 1 }, + }, + }, + }, + }); + const selected = resolveModelsDevSelection(entries, { + autoSyncEnabled: true, + includeCommonModels: true, + selectedModelKeys: ["relay/vendor/GPT-5", "relay/custom-model"], + excludedCommonModelKeys: ["openai/gpt-5"], + lastSyncAt: null, + lastSyncError: null, + }); + + expect(selected.map((entry) => entry.key)).toEqual([ + "relay/vendor/GPT-5", + "relay/custom-model", + ]); + + const pricing = toModelPricing([ + entries.find((entry) => entry.key === "openai/gpt-5")!, + entries.find((entry) => entry.key === "relay/vendor/GPT-5")!, + ]); + expect(pricing).toHaveLength(1); + expect(pricing[0]).toMatchObject({ + modelId: "gpt-5", + displayName: "GPT-5 Official", + inputCostPerMillion: "1", + }); + }); }); diff --git a/tests/lib/modelsDevAutoSync.test.ts b/tests/lib/modelsDevAutoSync.test.ts new file mode 100644 index 000000000..a7256a02c --- /dev/null +++ b/tests/lib/modelsDevAutoSync.test.ts @@ -0,0 +1,196 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; + +const { + getModelsDevSyncConfig, + updateModelPricingBatch, + recordModelsDevSyncResult, +} = vi.hoisted(() => ({ + getModelsDevSyncConfig: vi.fn(), + updateModelPricingBatch: vi.fn(), + recordModelsDevSyncResult: vi.fn(), +})); + +vi.mock("@/lib/api/usage", () => ({ + usageApi: { + getModelsDevSyncConfig, + updateModelPricingBatch, + recordModelsDevSyncResult, + }, +})); + +import { + MODELS_DEV_STARTUP_SYNC_INTERVAL_MS, + syncModelsDevPricing, +} from "@/lib/modelsDevAutoSync"; + +const state = { + configPath: "C:/Users/test/.cc-switch/model-pricing.json", + config: { + autoSyncEnabled: true, + includeCommonModels: true, + selectedModelKeys: ["relay/custom-model"], + excludedCommonModelKeys: [], + lastSyncAt: null, + lastSyncError: null, + }, +}; + +describe("syncModelsDevPricing", () => { + beforeEach(() => { + vi.clearAllMocks(); + updateModelPricingBatch.mockResolvedValue(2); + recordModelsDevSyncResult.mockResolvedValue(undefined); + vi.stubGlobal( + "fetch", + vi.fn().mockResolvedValue({ + ok: true, + json: async () => ({ + openai: { + models: { + "gpt-5": { + name: "GPT-5", + release_date: "2025-08-01", + cost: { input: 1, output: 2 }, + }, + }, + }, + relay: { + models: { + "custom-model": { + name: "Custom Model", + release_date: "2025-07-01", + cost: { input: 0.5, output: 1 }, + }, + }, + }, + }), + }), + ); + }); + + it("skips network access when automatic sync is disabled", async () => { + getModelsDevSyncConfig.mockResolvedValue({ + ...state, + config: { ...state.config, autoSyncEnabled: false }, + }); + + const result = await syncModelsDevPricing(); + + expect(result.skipped).toBe(true); + expect(fetch).not.toHaveBeenCalled(); + expect(updateModelPricingBatch).not.toHaveBeenCalled(); + }); + + it("skips startup network access when pricing synced within the interval", async () => { + const lastSyncAt = Date.now() - MODELS_DEV_STARTUP_SYNC_INTERVAL_MS + 1; + getModelsDevSyncConfig.mockResolvedValue({ + ...state, + config: { ...state.config, lastSyncAt }, + }); + + const result = await syncModelsDevPricing(); + + expect(result).toEqual({ + skipped: true, + selected: 0, + imported: 0, + changed: 0, + syncedAt: lastSyncAt, + }); + expect(fetch).not.toHaveBeenCalled(); + expect(updateModelPricingBatch).not.toHaveBeenCalled(); + }); + + it("imports common and explicitly selected models in one batch", async () => { + getModelsDevSyncConfig.mockResolvedValue(state); + + const result = await syncModelsDevPricing(); + + expect(updateModelPricingBatch).toHaveBeenCalledTimes(1); + expect(updateModelPricingBatch).toHaveBeenCalledWith([ + expect.objectContaining({ modelId: "gpt-5", inputCostPerMillion: "1" }), + expect.objectContaining({ + modelId: "custom-model", + inputCostPerMillion: "0.5", + }), + ]); + expect(result).toMatchObject({ + skipped: false, + selected: 2, + imported: 2, + changed: 2, + }); + expect(recordModelsDevSyncResult).toHaveBeenCalledWith( + expect.any(Number), + null, + ); + const fetchOptions = vi.mocked(fetch).mock.calls[0]?.[1]; + expect(fetchOptions).toEqual({ signal: expect.any(AbortSignal) }); + expect(fetchOptions).not.toHaveProperty("cache"); + }); + + it("stops a startup sync when automatic sync is disabled during download", async () => { + getModelsDevSyncConfig.mockResolvedValueOnce(state).mockResolvedValueOnce({ + ...state, + config: { ...state.config, autoSyncEnabled: false, lastSyncAt: 123 }, + }); + + const result = await syncModelsDevPricing(); + + expect(result).toEqual({ + skipped: true, + selected: 0, + imported: 0, + changed: 0, + syncedAt: 123, + }); + expect(updateModelPricingBatch).not.toHaveBeenCalled(); + expect(recordModelsDevSyncResult).not.toHaveBeenCalled(); + }); + + it("uses the latest model selection after the download completes", async () => { + getModelsDevSyncConfig.mockResolvedValueOnce(state).mockResolvedValueOnce({ + ...state, + config: { + ...state.config, + includeCommonModels: false, + selectedModelKeys: ["relay/custom-model"], + }, + }); + + await syncModelsDevPricing(); + + expect(updateModelPricingBatch).toHaveBeenCalledWith([ + expect.objectContaining({ modelId: "custom-model" }), + ]); + }); + + it("uses the latest selection for a forced sync even when automatic sync is disabled", async () => { + getModelsDevSyncConfig.mockResolvedValueOnce({ + ...state, + config: { + ...state.config, + autoSyncEnabled: false, + includeCommonModels: false, + selectedModelKeys: ["relay/custom-model"], + lastSyncAt: Date.now(), + }, + }); + + const result = await syncModelsDevPricing(state, true); + + expect(result.skipped).toBe(false); + expect(updateModelPricingBatch).toHaveBeenCalledWith([ + expect.objectContaining({ modelId: "custom-model" }), + ]); + }); + + it("persists the last error without replacing the previous success time", async () => { + const previous = { ...state, config: { ...state.config, lastSyncAt: 123 } }; + getModelsDevSyncConfig.mockResolvedValue(previous); + vi.mocked(fetch).mockRejectedValueOnce(new Error("offline")); + + await expect(syncModelsDevPricing()).rejects.toThrow("offline"); + expect(recordModelsDevSyncResult).toHaveBeenCalledWith(null, "offline"); + }); +});