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
This commit is contained in:
Thefool
2026-07-28 23:37:55 +08:00
committed by GitHub
parent ff3bc242cc
commit 12b972a66e
20 changed files with 2672 additions and 195 deletions
+47 -86
View File
@@ -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<Vec<ModelPricingInfo>, 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<ModelPricingInfo>,
) -> Result<usize, AppError> {
crate::services::model_pricing::update_model_pricing_batch(&state.db, entries)
}
#[tauri::command]
pub fn get_models_dev_sync_config(
state: State<'_, AppState>,
) -> Result<ModelsDevSyncState, AppError> {
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<i64>,
error: Option<String>,
) -> 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::*;
+3
View File
@@ -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) {
+4
View File
@@ -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,
+1
View File
@@ -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;
+761
View File
@@ -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<Mutex<()>> = 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<String>,
#[serde(default)]
pub excluded_common_model_keys: Vec<String>,
#[serde(default)]
pub last_sync_at: Option<i64>,
#[serde(default)]
pub last_sync_error: Option<String>,
}
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<ModelPricingInfo>,
#[serde(default)]
deleted_model_ids: Vec<String>,
}
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<String, AppError> {
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<ModelPricingInfo, AppError> {
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<String>) -> Vec<String> {
values
.into_iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<BTreeSet<_>>()
.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<ModelPricingFile, AppError> {
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::<BTreeSet<_>>();
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<Option<ModelPricingFile>, 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<ModelPricingFile, AppError> {
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<usize, AppError> {
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<usize, AppError> {
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<ModelsDevSyncState, AppError> {
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<i64>,
error: Option<String>,
) -> 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<ModelPricingInfo>,
backfill_all: bool,
) -> Result<usize, AppError> {
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::<Vec<_>>();
let model_ids = entries
.iter()
.map(|entry| entry.model_id.clone())
.collect::<Vec<_>>();
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::<BTreeMap<_, _>>();
let updated_ids = entries
.iter()
.map(|entry| entry.model_id.clone())
.collect::<BTreeSet<_>>();
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<usize, AppError> {
update_model_pricing_batch_inner(db, vec![entry], false)
}
pub fn update_model_pricing_batch(
db: &Database,
entries: Vec<ModelPricingInfo>,
) -> Result<usize, AppError> {
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"));
});
}
}