From 37e78166c13be430186af46c834080823c636d6d Mon Sep 17 00:00:00 2001 From: SaladDay Date: Sat, 1 Aug 2026 08:37:49 +0000 Subject: [PATCH] refactor(provider): complete prerequisite A ownership --- src-tauri/src/codex_history_migration.rs | 19 +- src-tauri/src/database/dao/provider_write.rs | 31 +- src-tauri/src/lib.rs | 8 +- src-tauri/src/proxy/switch_lock.rs | 23 +- src-tauri/src/services/provider/live.rs | 6 + src-tauri/src/services/provider/mod.rs | 291 ++++++++++-- src-tauri/src/services/proxy.rs | 463 +++++++++++-------- src/hooks/useProviderActions.ts | 6 +- src/lib/api/providers.test.ts | 40 ++ src/lib/api/providers.ts | 36 +- src/lib/query/mutations.ts | 8 +- 11 files changed, 667 insertions(+), 264 deletions(-) create mode 100644 src/lib/api/providers.test.ts diff --git a/src-tauri/src/codex_history_migration.rs b/src-tauri/src/codex_history_migration.rs index f4c986ba2..1547c83f9 100644 --- a/src-tauri/src/codex_history_migration.rs +++ b/src-tauri/src/codex_history_migration.rs @@ -8,9 +8,12 @@ use crate::codex_config::{ }; use crate::codex_state_db::codex_state_db_paths; use crate::config::{atomic_write, copy_file, get_app_config_dir}; -use crate::database::{is_official_seed_id, Database, ProviderKey, ProviderRowUpdate}; +use crate::database::{is_official_seed_id, Database}; use crate::error::AppError; -use crate::services::provider::provider_to_mutation_input; +use crate::services::provider::{ + provider_row_fingerprint, provider_to_mutation_input, + reconcile_provider_record_with_precondition, ReconcilePrecondition, +}; use crate::settings::{ CodexOfficialHistoryUnifyMigration, CodexProviderTemplateMigration, CodexThirdPartyHistoryProviderBucketMigration, @@ -665,6 +668,7 @@ fn migrate_codex_provider_templates_to_custom( let mut migrated_provider_ids = Vec::new(); for (_, mut provider) in providers { + let observed_fingerprint = provider_row_fingerprint(&provider); if provider.category.as_deref() == Some("official") || is_official_seed_id(&provider.id) || provider.is_codex_oauth() @@ -701,9 +705,14 @@ fn migrate_codex_provider_templates_to_custom( meta.custom_endpoints.clear(); } let input = provider_to_mutation_input(provider); - let key = ProviderKey::new("codex", &provider_id)?; - let row = ProviderRowUpdate::from_input(&input)?; - db.update_provider(&key, &row)?; + reconcile_provider_record_with_precondition( + db, + "codex", + input, + ReconcilePrecondition::ExpectPresent { + fingerprint: observed_fingerprint, + }, + )?; migrated_provider_ids.push(provider_id); } diff --git a/src-tauri/src/database/dao/provider_write.rs b/src-tauri/src/database/dao/provider_write.rs index 37b9b8264..554ecd605 100644 --- a/src-tauri/src/database/dao/provider_write.rs +++ b/src-tauri/src/database/dao/provider_write.rs @@ -190,10 +190,15 @@ impl RenameProvider { "provider target id must be non-empty".to_string(), )); } + let mut row = ProviderRowUpdate::from_input(input)?; + // A successful key change always remains DB-only. The service owns + // the corresponding live-file absence check, while the DAO persists + // the durable half of that invariant. + row.meta.live_config_managed = Some(false); Ok(Self { source, target_id: input.id.clone(), - row: ProviderRowUpdate::from_input(input)?, + row, }) } } @@ -456,14 +461,17 @@ impl Database { .map_err(|error| AppError::Database(error.to_string())) } - pub fn rename_db_only_additive_provider(&self, input: RenameProvider) -> Result<(), AppError> { + pub(crate) fn rename_db_only_additive_provider( + &self, + input: RenameProvider, + ) -> Result<(), AppError> { let mut conn = lock_conn!(self.conn); let tx = conn .transaction() .map_err(|error| AppError::Database(error.to_string()))?; let source_state = tx .query_row( - "SELECT sort_index, is_current, in_failover_queue, category, created_at + "SELECT sort_index, is_current, in_failover_queue, category, created_at, meta FROM providers WHERE id = ?1 AND app_type = ?2", params![input.source.id, input.source.app_type], @@ -474,6 +482,7 @@ impl Database { row.get::<_, bool>(2)?, row.get::<_, Option>(3)?, row.get::<_, Option>(4)?, + row.get::<_, String>(5)?, )) }, ) @@ -490,6 +499,22 @@ impl Database { "OMO/OMO Slim providers cannot be renamed".to_string(), )); } + let source_meta: ProviderMeta = if source_state.5.trim().is_empty() { + ProviderMeta::default() + } else { + serde_json::from_str(&source_state.5).map_err(|error| { + AppError::Database(format!( + "invalid meta for provider '{}/{}': {error}", + input.source.app_type, input.source.id + )) + })? + }; + if source_meta.live_config_managed == Some(true) { + return Err(AppError::Conflict(format!( + "provider '{}/{}' became live-managed before rename", + input.source.app_type, input.source.id + ))); + } let target = ProviderKey::new(&input.source.app_type, &input.target_id)?; insert_row( &tx, diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 1236b3bde..72fa77cf1 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -1993,6 +1993,7 @@ fn initialize_common_config_snippets(state: &store::AppState) { .unwrap_or(true); if should_run_legacy_migration { + let mut legacy_migration_succeeded = true; for app_type in [ crate::app_config::AppType::Claude, crate::app_config::AppType::Codex, @@ -2006,11 +2007,14 @@ fn initialize_common_config_snippets(state: &store::AppState) { "✗ Failed to migrate legacy common-config usage for {}: {e}", app_type.as_str() ); + legacy_migration_succeeded = false; } } - if let Err(e) = state.db.set_legacy_common_config_migrated(true) { - log::warn!("✗ Failed to persist legacy common-config migration flag: {e}"); + if legacy_migration_succeeded { + if let Err(e) = state.db.set_legacy_common_config_migrated(true) { + log::warn!("✗ Failed to persist legacy common-config migration flag: {e}"); + } } } } diff --git a/src-tauri/src/proxy/switch_lock.rs b/src-tauri/src/proxy/switch_lock.rs index 10a67ba65..95e9acf36 100644 --- a/src-tauri/src/proxy/switch_lock.rs +++ b/src-tauri/src/proxy/switch_lock.rs @@ -4,15 +4,32 @@ //! 防止并发切换导致 is_current 与 Live 备份不一致。 use std::collections::HashMap; -use std::sync::Arc; +use std::sync::{Arc, OnceLock}; use tokio::sync::{Mutex, OwnedMutexGuard, RwLock}; +type PerAppLocks = Arc>>>>; + /// 每个应用类型一把互斥锁,保证同一应用的切换操作串行执行。 /// /// 不同应用之间(如 Claude 和 Codex)可以并行切换。 -#[derive(Clone, Default)] +#[derive(Clone)] pub struct SwitchLockManager { - locks: Arc>>>>, + locks: PerAppLocks, +} + +impl Default for SwitchLockManager { + fn default() -> Self { + // Some commands construct a short-lived AppState around the shared + // database before running a blocking sync. A per-ProxyService map + // would give those paths a different lock and defeat serialization + // with provider rename/switch operations in the primary AppState. + static LOCKS: OnceLock = OnceLock::new(); + Self { + locks: LOCKS + .get_or_init(|| Arc::new(RwLock::new(HashMap::new()))) + .clone(), + } + } } impl SwitchLockManager { diff --git a/src-tauri/src/services/provider/live.rs b/src-tauri/src/services/provider/live.rs index 02379d057..555d48a21 100644 --- a/src-tauri/src/services/provider/live.rs +++ b/src-tauri/src/services/provider/live.rs @@ -1282,6 +1282,12 @@ pub fn sync_current_to_live(state: &AppState) -> Result<(), AppError> { // Sync providers based on mode for app_type in AppType::all() { if app_type.is_additive_mode() { + // Provider rename and every additive live mutation share this + // per-app lock. Acquire it before reading the catalog so a key + // cannot be renamed after this sync captured a stale provider map. + let _guard = futures::executor::block_on( + state.proxy_service.lock_switch_for_app(app_type.as_str()), + ); // Additive mode: sync ALL providers sync_all_providers_to_live(state, &app_type)?; } else { diff --git a/src-tauri/src/services/provider/mod.rs b/src-tauri/src/services/provider/mod.rs index 8d7d2a591..19afadd36 100644 --- a/src-tauri/src/services/provider/mod.rs +++ b/src-tauri/src/services/provider/mod.rs @@ -151,6 +151,37 @@ fn update_provider_record( state.db.update_provider(&key, &row) } +fn update_provider_record_if_unchanged( + state: &AppState, + app_type: &AppType, + observed_fingerprint: String, + input: ProviderMutationInput, +) -> Result<(), AppError> { + reconcile_provider_record_with_precondition( + state.db.as_ref(), + app_type.as_str(), + input, + ReconcilePrecondition::ExpectPresent { + fingerprint: observed_fingerprint, + }, + ) +} + +fn remove_hydrated_endpoints_from_row_update(provider: &mut Provider) { + if let Some(meta) = provider.meta.as_mut() { + meta.custom_endpoints.clear(); + } +} + +fn lock_additive_provider_mutation( + state: &AppState, + app_type: &AppType, +) -> Option> { + app_type.is_additive_mode().then(|| { + futures::executor::block_on(state.proxy_service.lock_switch_for_app(app_type.as_str())) + }) +} + /// Reconcile 的显式前置期望(前置工程 A 认证契约 T9)。 /// check-then-branch 的 TOCTOU 由调用方在观察时声明期望、由本层强制。 #[derive(Debug, Clone)] @@ -218,7 +249,9 @@ mod tests { use std::env; use std::fs; use std::path::{Path, PathBuf}; - use std::sync::{Arc, Mutex, OnceLock}; + use std::sync::{mpsc, Arc, Mutex, OnceLock}; + use std::thread; + use std::time::Duration; use tempfile::TempDir; struct TempHome { @@ -835,6 +868,145 @@ mod tests { }); } + #[test] + #[serial] + fn automated_row_transform_conflicts_instead_of_reverting_a_newer_edit() { + with_test_home(|state, _| { + let provider = opencode_provider("observed-row"); + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(provider), + false, + ) + .expect("create observed provider"); + + let observed = state + .db + .get_provider_by_id("observed-row", "opencode") + .expect("read observed provider") + .expect("observed provider exists"); + let fingerprint = provider_row_fingerprint(&observed); + let mut stale_transform = observed.clone(); + stale_transform.notes = Some("automatic transform".to_string()); + remove_hydrated_endpoints_from_row_update(&mut stale_transform); + + let mut user_edit = observed; + user_edit.name = "Concurrent user edit".to_string(); + user_edit + .meta + .get_or_insert_with(Default::default) + .custom_endpoints + .clear(); + ProviderService::update( + state, + AppType::OpenCode, + None, + provider_to_mutation_input(user_edit), + ) + .expect("persist concurrent user edit"); + + let error = update_provider_record_if_unchanged( + state, + &AppType::OpenCode, + fingerprint, + provider_to_mutation_input(stale_transform), + ) + .expect_err("stale automatic transform must conflict"); + assert!(matches!(error, AppError::Conflict(_))); + + let saved = state + .db + .get_provider_by_id("observed-row", "opencode") + .expect("read saved provider") + .expect("saved provider exists"); + assert_eq!(saved.name, "Concurrent user edit"); + assert_eq!(saved.notes, None); + }); + } + + #[test] + #[serial] + fn rename_waits_for_shared_live_lock_and_rechecks_source_ownership() { + let _test_guard = test_guard(); + let _home = TempHome::new(); + let db = Arc::new(Database::memory().expect("in-memory database")); + let lock_owner_state = AppState::new(db.clone()); + let rename_state = AppState::new(db.clone()); + + ProviderService::add( + &lock_owner_state, + AppType::OpenCode, + provider_to_mutation_input(opencode_provider("race-source")), + false, + ) + .expect("create DB-only source"); + + let live_guard = futures::executor::block_on( + lock_owner_state + .proxy_service + .lock_switch_for_app(AppType::OpenCode.as_str()), + ); + let (started_tx, started_rx) = mpsc::channel(); + let rename_thread = thread::spawn(move || { + started_tx.send(()).expect("signal rename start"); + ProviderService::update( + &rename_state, + AppType::OpenCode, + Some("race-source"), + provider_to_mutation_input(opencode_provider("race-target")), + ) + }); + started_rx + .recv_timeout(Duration::from_secs(1)) + .expect("rename thread started"); + thread::sleep(Duration::from_millis(30)); + assert!( + !rename_thread.is_finished(), + "a distinct AppState must share the same per-app live mutation lock" + ); + + let mut switched = db + .get_provider_by_id("race-source", "opencode") + .expect("read source during simulated switch") + .expect("source exists"); + ProviderService::set_provider_live_config_managed(&mut switched, true); + remove_hydrated_endpoints_from_row_update(&mut switched); + let key = ProviderKey::new("opencode", "race-source").expect("source key"); + let row = ProviderRowUpdate::from_input(&provider_to_mutation_input(switched)) + .expect("marker row"); + db.update_provider(&key, &row) + .expect("persist simulated switch marker"); + drop(live_guard); + + let error = rename_thread + .join() + .expect("join rename thread") + .expect_err("live-managed source must not be renamed"); + assert!( + matches!(&error, AppError::Conflict(_) | AppError::Message(_)), + "ownership recheck should reject with a structured conflict or service error: {error}" + ); + assert!(db + .get_provider_by_id("race-source", "opencode") + .expect("read source") + .is_some()); + assert!(db + .get_provider_by_id("race-target", "opencode") + .expect("read target") + .is_none()); + + let rename = RenameProvider::from_input( + ProviderKey::new("opencode", "race-source").expect("source key"), + &provider_to_mutation_input(opencode_provider("direct-target")), + ) + .expect("build direct rename"); + assert!(matches!( + db.rename_db_only_additive_provider(rename), + Err(AppError::Conflict(_)) + )); + } + #[test] #[serial] fn add_clears_usage_credentials_that_match_provider_config() { @@ -2961,6 +3133,38 @@ impl ProviderService { .live_config_managed = Some(managed); } + fn persist_live_config_managed( + state: &AppState, + app_type: &AppType, + provider_id: &str, + managed: bool, + ) -> Result<(), AppError> { + // This is a narrow metadata transformation, so a concurrent content + // edit can be preserved by rereading and reapplying it. Every attempt + // still uses the single-lock/single-transaction fingerprint primitive. + for attempt in 0..3 { + let mut provider = state + .db + .get_provider_aggregate(app_type.as_str(), provider_id)? + .ok_or_else(|| { + AppError::NotFound(format!("provider '{}/{}'", app_type.as_str(), provider_id)) + })? + .provider; + if Self::provider_live_config_managed(&provider) == Some(managed) { + return Ok(()); + } + let fingerprint = provider_row_fingerprint(&provider); + Self::set_provider_live_config_managed(&mut provider, managed); + remove_hydrated_endpoints_from_row_update(&mut provider); + let input = provider_to_mutation_input(provider); + match update_provider_record_if_unchanged(state, app_type, fingerprint, input) { + Err(AppError::Conflict(_)) if attempt < 2 => continue, + result => return result, + } + } + unreachable!("bounded live-config marker retry always returns") + } + fn normalize_usage_script_credential_overrides(app_type: &AppType, provider: &mut Provider) { let current_credentials = provider.resolve_usage_credentials(app_type); @@ -3058,6 +3262,7 @@ impl ProviderService { input: ProviderMutationInput, add_to_live: bool, ) -> Result { + let _provider_mutation_guard = lock_additive_provider_mutation(state, &app_type); let mut provider: Provider = input.into(); // Normalize Claude model keys Self::normalize_provider_if_claude(&app_type, &mut provider); @@ -3115,6 +3320,7 @@ impl ProviderService { // Reject endpoint-bearing edit payloads before any live or DB side // effect. Endpoints have their own typed mutation API. ProviderRowUpdate::from_input(&input)?; + let _provider_mutation_guard = lock_additive_provider_mutation(state, &app_type); let mut provider: Provider = input.into(); let original_id = original_id.unwrap_or(provider.id.as_str()).to_string(); let provider_id_changed = original_id != provider.id; @@ -3357,6 +3563,7 @@ impl ProviderService { /// 同时检查本地 settings 和数据库的当前供应商,防止删除任一端正在使用的供应商。 /// 对于累加模式应用(OpenCode, OpenClaw),可以随时删除任意供应商,同时从 live 配置中移除。 pub fn delete(state: &AppState, app_type: AppType, id: &str) -> Result<(), AppError> { + let _provider_mutation_guard = lock_additive_provider_mutation(state, &app_type); // Additive mode apps - no current provider concept if app_type.is_additive_mode() { // Single DB read shared across all additive-mode sub-paths below. @@ -3428,6 +3635,7 @@ impl ProviderService { app_type: AppType, id: &str, ) -> Result<(), AppError> { + let _provider_mutation_guard = lock_additive_provider_mutation(state, &app_type); match app_type { AppType::OpenCode => { let provider_category = state @@ -3471,13 +3679,12 @@ impl ProviderService { } } - if let Some(mut provider) = state + if state .db .get_provider_aggregate(app_type.as_str(), id)? - .map(|aggregate| aggregate.provider) + .is_some() { - Self::set_provider_live_config_managed(&mut provider, false); - update_provider_record(state, &app_type, &provider_to_mutation_input(provider))?; + Self::persist_live_config_managed(state, &app_type, id, false)?; } Ok(()) @@ -3496,6 +3703,21 @@ impl ProviderService { /// d. Write target provider config to live files /// e. Sync MCP configuration pub fn switch(state: &AppState, app_type: AppType, id: &str) -> Result { + // The same per-app lock also guards additive provider key changes and + // bulk live sync. Acquire it before observing the provider map so a + // queued rename cannot leave this switch holding a stale source key. + let _switch_guard = if app_type.is_additive_mode() + || matches!( + app_type, + AppType::Claude | AppType::Codex | AppType::Gemini | AppType::GrokBuild + ) { + Some(futures::executor::block_on( + state.proxy_service.lock_switch_for_app(app_type.as_str()), + )) + } else { + None + }; + // Check if provider exists let providers = state.db.get_all_providers(app_type.as_str())?; let _provider = providers @@ -3518,21 +3740,6 @@ impl ProviderService { return Self::switch_normal(state, app_type, id, &providers); } - // Provider switches and takeover toggles both mutate live config and the - // restore backup. Serialize them per app, then decide from the locked - // current state so a just-started takeover cannot be overwritten by a - // normal live write. - let _switch_guard = if matches!( - app_type, - AppType::Claude | AppType::Codex | AppType::Gemini | AppType::GrokBuild - ) { - Some(futures::executor::block_on( - state.proxy_service.lock_switch_for_app(app_type.as_str()), - )) - } else { - None - }; - // Backup or live placeholders mean the live file is owned by proxy // takeover, even if the proxy server is temporarily stopped or is in the // activation window before enabled=true is committed. @@ -3634,6 +3841,7 @@ impl ProviderService { .get_provider_aggregate(app_type.as_str(), ¤t_id)? .map(|aggregate| aggregate.provider) { + let fingerprint = provider_row_fingerprint(¤t_provider); // 切走前先把 live 里的可共享改动(含用户直接在应用内 // 装插件/加 hook/改偏好)同步进通用配置片段,再做剥离回填。 // 详见 sync_common_config_snippet_from_live 的文档。 @@ -3652,10 +3860,12 @@ impl ProviderService { ¤t_provider, live_config, ); - if let Err(e) = update_provider_record( + remove_hydrated_endpoints_from_row_update(&mut current_provider); + if let Err(e) = update_provider_record_if_unchanged( state, &app_type, - &provider_to_mutation_input(current_provider), + fingerprint, + provider_to_mutation_input(current_provider), ) { log::warn!("Backfill failed: {e}"); result @@ -3707,16 +3917,8 @@ impl ProviderService { // the provider in a silent inconsistent state (present in live, but still marked DB-only). if app_type.is_additive_mode() && Self::provider_live_config_managed(provider) != Some(true) { - let mut updated = state - .db - .get_provider_aggregate(app_type.as_str(), &provider.id)? - .ok_or_else(|| { - AppError::NotFound(format!("provider '{}/{}'", app_type.as_str(), provider.id)) - })? - .provider; - Self::set_provider_live_config_managed(&mut updated, true); - let update = provider_to_mutation_input(updated); - if let Err(e) = update_provider_record(state, &app_type, &update) { + if let Err(e) = Self::persist_live_config_managed(state, &app_type, &provider.id, true) + { let rollback_result = match app_type { AppType::OpenCode => remove_opencode_provider_from_live(&provider.id), AppType::OpenClaw => remove_openclaw_provider_from_live(&provider.id), @@ -3764,6 +3966,7 @@ impl ProviderService { state: &AppState, app_type: AppType, ) -> Result<(), AppError> { + let _provider_mutation_guard = lock_additive_provider_mutation(state, &app_type); if app_type.is_additive_mode() { return sync_current_provider_for_app_to_live(state, &app_type); } @@ -3835,6 +4038,7 @@ impl ProviderService { continue; } + let fingerprint = provider_row_fingerprint(provider); let mut updated_provider = provider.clone(); updated_provider .meta @@ -3856,10 +4060,12 @@ impl ProviderService { } } - update_provider_record( + remove_hydrated_endpoints_from_row_update(&mut updated_provider); + update_provider_record_if_unchanged( state, &app_type, - &provider_to_mutation_input(updated_provider), + fingerprint, + provider_to_mutation_input(updated_provider), )?; } @@ -4377,9 +4583,10 @@ impl ProviderService { // 1) 先算出各供应商清理后的配置,但**先不落库** let providers = state.db.get_all_provider_aggregates(app.as_str())?; - let mut pending: Vec<(String, Provider, Value)> = Vec::new(); + let mut pending: Vec<(String, Provider, Value, String)> = Vec::new(); for (id, aggregate) in providers { let provider = aggregate.provider; + let fingerprint = provider_row_fingerprint(&provider); let cleaned = match live::remove_common_config_from_settings( &app, &provider.settings_config, @@ -4392,7 +4599,7 @@ impl ProviderService { } }; if cleaned != provider.settings_config { - pending.push((id, provider, cleaned)); + pending.push((id, provider, cleaned, fingerprint)); } } @@ -4421,7 +4628,7 @@ impl ProviderService { "removedFromSnippet": poison_keys, "providers": pending .iter() - .map(|(id, provider, cleaned)| serde_json::json!({ + .map(|(id, provider, cleaned, _)| serde_json::json!({ "id": id, "removedKeys": removed_env_keys(&provider.settings_config, cleaned), })) @@ -4438,11 +4645,11 @@ impl ProviderService { } // 3) 各供应商 settings_config:按值相等定向删除扩散出去的副本 - for (id, provider, cleaned) in pending { - let mut updated = provider; - updated.settings_config = cleaned; - let update = provider_to_mutation_input(updated); - update_provider_record(state, &app, &update)?; + for (id, mut provider, cleaned, fingerprint) in pending { + provider.settings_config = cleaned; + remove_hydrated_endpoints_from_row_update(&mut provider); + let update = provider_to_mutation_input(provider); + update_provider_record_if_unchanged(state, &app, fingerprint, update)?; log::info!("已从 Gemini 供应商 '{id}' 中清除泄漏的共享凭据"); } diff --git a/src-tauri/src/services/proxy.rs b/src-tauri/src/services/proxy.rs index 7c651585f..a1c3e9ef5 100644 --- a/src-tauri/src/services/proxy.rs +++ b/src-tauri/src/services/proxy.rs @@ -4,14 +4,15 @@ use crate::app_config::AppType; use crate::config::{get_claude_settings_path, read_json_file, write_json_file}; -use crate::database::{Database, ProviderKey, ProviderRowUpdate}; +use crate::database::Database; use crate::provider::Provider; use crate::proxy::server::ProxyServer; use crate::proxy::switch_lock::SwitchLockManager; use crate::proxy::types::*; use crate::services::provider::{ - build_effective_settings_with_common_config, provider_to_mutation_input, - write_live_with_common_config, + build_effective_settings_with_common_config, provider_row_fingerprint, + provider_to_mutation_input, reconcile_provider_record_with_precondition, + write_live_with_common_config, ReconcilePrecondition, }; use serde_json::{json, Map, Value}; use std::str::FromStr; @@ -971,6 +972,27 @@ impl ProxyService { .await } + fn persist_synced_live_token( + &self, + app_type: &str, + provider_id: &str, + observed_fingerprint: String, + mut provider: Provider, + ) -> Result<(), String> { + if let Some(meta) = provider.meta.as_mut() { + meta.custom_endpoints.clear(); + } + reconcile_provider_record_with_precondition( + self.db.as_ref(), + app_type, + provider_to_mutation_input(provider), + ReconcilePrecondition::ExpectPresent { + fingerprint: observed_fingerprint, + }, + ) + .map_err(|error| format!("同步 {app_type}/{provider_id} Live Token 到数据库失败: {error}")) + } + async fn sync_live_config_to_provider( &self, app_type: &AppType, @@ -983,96 +1005,89 @@ impl ProxyService { .map_err(|e| format!("获取 Claude 当前供应商失败: {e}"))?; if let Some(provider_id) = provider_id { - if let Ok(Some(mut provider)) = - self.db.get_provider_by_id(&provider_id, "claude") - { - if let Some(env) = live_config.get("env").and_then(|v| v.as_object()) { - let token_pair = [ - "ANTHROPIC_AUTH_TOKEN", - "ANTHROPIC_API_KEY", - "OPENROUTER_API_KEY", - "OPENAI_API_KEY", - ] - .into_iter() - .find_map(|key| { - env.get(key) - .and_then(|v| v.as_str()) - .map(|s| (key, s.trim())) - }) - .filter(|(_, token)| { - !token.is_empty() && *token != PROXY_TOKEN_PLACEHOLDER - }); + let Some(mut provider) = self + .db + .get_provider_by_id(&provider_id, "claude") + .map_err(|error| { + format!("读取 Claude 供应商 '{provider_id}' 失败: {error}") + })? + else { + return Err(format!("Claude 当前供应商不存在: {provider_id}")); + }; + let observed_fingerprint = provider_row_fingerprint(&provider); + if let Some(env) = live_config.get("env").and_then(|v| v.as_object()) { + let token_pair = [ + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_KEY", + "OPENROUTER_API_KEY", + "OPENAI_API_KEY", + ] + .into_iter() + .find_map(|key| { + env.get(key) + .and_then(|v| v.as_str()) + .map(|s| (key, s.trim())) + }) + .filter(|(_, token)| { + !token.is_empty() && *token != PROXY_TOKEN_PLACEHOLDER + }); - if let Some((token_key, token)) = token_pair { - let env_obj = provider - .settings_config - .get_mut("env") - .and_then(|v| v.as_object_mut()); + if let Some((token_key, token)) = token_pair { + let env_obj = provider + .settings_config + .get_mut("env") + .and_then(|v| v.as_object_mut()); - match env_obj { - Some(obj) => { - if token_key == "ANTHROPIC_AUTH_TOKEN" - || token_key == "ANTHROPIC_API_KEY" - { - let mut updated = false; - if obj.contains_key("ANTHROPIC_AUTH_TOKEN") { - obj.insert( - "ANTHROPIC_AUTH_TOKEN".to_string(), - json!(token), - ); - updated = true; - } - if obj.contains_key("ANTHROPIC_API_KEY") { - obj.insert( - "ANTHROPIC_API_KEY".to_string(), - json!(token), - ); - updated = true; - } - if !updated { - obj.insert(token_key.to_string(), json!(token)); - } - } else { + match env_obj { + Some(obj) => { + if token_key == "ANTHROPIC_AUTH_TOKEN" + || token_key == "ANTHROPIC_API_KEY" + { + let mut updated = false; + if obj.contains_key("ANTHROPIC_AUTH_TOKEN") { + obj.insert( + "ANTHROPIC_AUTH_TOKEN".to_string(), + json!(token), + ); + updated = true; + } + if obj.contains_key("ANTHROPIC_API_KEY") { + obj.insert( + "ANTHROPIC_API_KEY".to_string(), + json!(token), + ); + updated = true; + } + if !updated { obj.insert(token_key.to_string(), json!(token)); } + } else { + obj.insert(token_key.to_string(), json!(token)); + } + } + None => { + // 至少写入一份可用的 Token + if provider.settings_config.is_null() { + provider.settings_config = json!({}); } - None => { - // 至少写入一份可用的 Token - if provider.settings_config.is_null() { - provider.settings_config = json!({}); - } - if let Some(root) = provider.settings_config.as_object_mut() - { - root.insert( - "env".to_string(), - json!({ token_key: token }), - ); - } else { - log::warn!( + if let Some(root) = provider.settings_config.as_object_mut() { + root.insert("env".to_string(), json!({ token_key: token })); + } else { + log::warn!( "Claude provider settings_config 格式异常(非对象),跳过写入 Token (provider: {provider_id})" ); - } } } - - if let Some(meta) = provider.meta.as_mut() { - meta.custom_endpoints.clear(); - } - let input = provider_to_mutation_input(provider); - let result = - ProviderKey::new("claude", &provider_id).and_then(|key| { - let row = ProviderRowUpdate::from_input(&input)?; - self.db.update_provider(&key, &row) - }); - if let Err(e) = result { - log::warn!("同步 Claude Token 到数据库失败: {e}"); - } else { - log::info!( - "已同步 Claude Token 到数据库 (provider: {provider_id})" - ); - } } + + self.persist_synced_live_token( + "claude", + &provider_id, + observed_fingerprint, + provider, + )?; + log::info!("已同步 Claude Token 到数据库 (provider: {provider_id})"); } } } @@ -1083,59 +1098,56 @@ impl ProxyService { .map_err(|e| format!("获取 Codex 当前供应商失败: {e}"))?; if let Some(provider_id) = provider_id { - if let Ok(Some(mut provider)) = - self.db.get_provider_by_id(&provider_id, "codex") + let Some(mut provider) = self + .db + .get_provider_by_id(&provider_id, "codex") + .map_err(|error| { + format!("读取 Codex 供应商 '{provider_id}' 失败: {error}") + })? + else { + return Err(format!("Codex 当前供应商不存在: {provider_id}")); + }; + let observed_fingerprint = provider_row_fingerprint(&provider); + // The built-in official row is a routing capability, not + // a credential store. Its auth must remain empty even + // when the live Codex login uses OPENAI_API_KEY mode. + if crate::proxy::providers::is_codex_official_provider(&provider) { + return Ok(()); + } + if let Some(token) = live_config + .get("auth") + .and_then(|v| v.get("OPENAI_API_KEY")) + .and_then(|v| v.as_str()) + .map(|s| s.trim()) + .filter(|s| !s.is_empty() && *s != PROXY_TOKEN_PLACEHOLDER) { - // The built-in official row is a routing capability, not - // a credential store. Its auth must remain empty even - // when the live Codex login uses OPENAI_API_KEY mode. - if crate::proxy::providers::is_codex_official_provider(&provider) { - return Ok(()); - } - if let Some(token) = live_config - .get("auth") - .and_then(|v| v.get("OPENAI_API_KEY")) - .and_then(|v| v.as_str()) - .map(|s| s.trim()) - .filter(|s| !s.is_empty() && *s != PROXY_TOKEN_PLACEHOLDER) + if let Some(auth_obj) = provider + .settings_config + .get_mut("auth") + .and_then(|v| v.as_object_mut()) { - if let Some(auth_obj) = provider - .settings_config - .get_mut("auth") - .and_then(|v| v.as_object_mut()) - { - auth_obj.insert("OPENAI_API_KEY".to_string(), json!(token)); - } else { - if provider.settings_config.is_null() { - provider.settings_config = json!({}); - } + auth_obj.insert("OPENAI_API_KEY".to_string(), json!(token)); + } else { + if provider.settings_config.is_null() { + provider.settings_config = json!({}); + } - if let Some(root) = provider.settings_config.as_object_mut() { - root.insert( - "auth".to_string(), - json!({ "OPENAI_API_KEY": token }), - ); - } else { - log::warn!( + if let Some(root) = provider.settings_config.as_object_mut() { + root.insert("auth".to_string(), json!({ "OPENAI_API_KEY": token })); + } else { + log::warn!( "Codex provider settings_config 格式异常(非对象),跳过写入 Token (provider: {provider_id})" ); - } - } - - if let Some(meta) = provider.meta.as_mut() { - meta.custom_endpoints.clear(); - } - let input = provider_to_mutation_input(provider); - let result = ProviderKey::new("codex", &provider_id).and_then(|key| { - let row = ProviderRowUpdate::from_input(&input)?; - self.db.update_provider(&key, &row) - }); - if let Err(e) = result { - log::warn!("同步 Codex Token 到数据库失败: {e}"); - } else { - log::info!("已同步 Codex Token 到数据库 (provider: {provider_id})"); } } + + self.persist_synced_live_token( + "codex", + &provider_id, + observed_fingerprint, + provider, + )?; + log::info!("已同步 Codex Token 到数据库 (provider: {provider_id})"); } } } @@ -1145,55 +1157,50 @@ impl ProxyService { .map_err(|e| format!("获取 Gemini 当前供应商失败: {e}"))?; if let Some(provider_id) = provider_id { - if let Ok(Some(mut provider)) = - self.db.get_provider_by_id(&provider_id, "gemini") + let Some(mut provider) = self + .db + .get_provider_by_id(&provider_id, "gemini") + .map_err(|error| { + format!("读取 Gemini 供应商 '{provider_id}' 失败: {error}") + })? + else { + return Err(format!("Gemini 当前供应商不存在: {provider_id}")); + }; + let observed_fingerprint = provider_row_fingerprint(&provider); + if let Some(token) = live_config + .get("env") + .and_then(|v| v.get("GEMINI_API_KEY")) + .and_then(|v| v.as_str()) + .map(|s| s.trim()) + .filter(|s| !s.is_empty() && *s != PROXY_TOKEN_PLACEHOLDER) { - if let Some(token) = live_config - .get("env") - .and_then(|v| v.get("GEMINI_API_KEY")) - .and_then(|v| v.as_str()) - .map(|s| s.trim()) - .filter(|s| !s.is_empty() && *s != PROXY_TOKEN_PLACEHOLDER) + if let Some(env_obj) = provider + .settings_config + .get_mut("env") + .and_then(|v| v.as_object_mut()) { - if let Some(env_obj) = provider - .settings_config - .get_mut("env") - .and_then(|v| v.as_object_mut()) - { - env_obj.insert("GEMINI_API_KEY".to_string(), json!(token)); - } else { - if provider.settings_config.is_null() { - provider.settings_config = json!({}); - } + env_obj.insert("GEMINI_API_KEY".to_string(), json!(token)); + } else { + if provider.settings_config.is_null() { + provider.settings_config = json!({}); + } - if let Some(root) = provider.settings_config.as_object_mut() { - root.insert( - "env".to_string(), - json!({ "GEMINI_API_KEY": token }), - ); - } else { - log::warn!( + if let Some(root) = provider.settings_config.as_object_mut() { + root.insert("env".to_string(), json!({ "GEMINI_API_KEY": token })); + } else { + log::warn!( "Gemini provider settings_config 格式异常(非对象),跳过写入 Token (provider: {provider_id})" ); - } - } - - if let Some(meta) = provider.meta.as_mut() { - meta.custom_endpoints.clear(); - } - let input = provider_to_mutation_input(provider); - let result = ProviderKey::new("gemini", &provider_id).and_then(|key| { - let row = ProviderRowUpdate::from_input(&input)?; - self.db.update_provider(&key, &row) - }); - if let Err(e) = result { - log::warn!("同步 Gemini Token 到数据库失败: {e}"); - } else { - log::info!( - "已同步 Gemini Token 到数据库 (provider: {provider_id})" - ); } } + + self.persist_synced_live_token( + "gemini", + &provider_id, + observed_fingerprint, + provider, + )?; + log::info!("已同步 Gemini Token 到数据库 (provider: {provider_id})"); } } } @@ -1203,43 +1210,41 @@ impl ProxyService { .map_err(|e| format!("获取 Grok Build 当前供应商失败: {e}"))?; if let Some(provider_id) = provider_id { - if let Ok(Some(mut provider)) = - self.db.get_provider_by_id(&provider_id, "grokbuild") + let Some(mut provider) = self + .db + .get_provider_by_id(&provider_id, "grokbuild") + .map_err(|error| { + format!("读取 Grok Build 供应商 '{provider_id}' 失败: {error}") + })? + else { + return Err(format!("Grok Build 当前供应商不存在: {provider_id}")); + }; + let observed_fingerprint = provider_row_fingerprint(&provider); + let live_config_toml = live_config + .get("config") + .and_then(Value::as_str) + .unwrap_or_default(); + if let Some(token) = + crate::grok_config::extract_inline_api_key(live_config_toml) { - let live_config_toml = live_config - .get("config") - .and_then(Value::as_str) - .unwrap_or_default(); - if let Some(token) = - crate::grok_config::extract_inline_api_key(live_config_toml) - { - if !token.is_empty() && token != PROXY_TOKEN_PLACEHOLDER { - if let Some(provider_config) = provider - .settings_config - .get("config") - .and_then(Value::as_str) - { - let updated = - crate::grok_config::update_api_key(provider_config, &token) - .map_err(|e| { - format!("更新 Grok Build API Key 失败: {e}") - })?; - provider.settings_config["config"] = json!(updated); - if let Some(meta) = provider.meta.as_mut() { - meta.custom_endpoints.clear(); - } - let input = provider_to_mutation_input(provider); - let key = ProviderKey::new("grokbuild", &provider_id).map_err( - |e| format!("同步 Grok Build Token 到数据库失败: {e}"), - )?; - let row = - ProviderRowUpdate::from_input(&input).map_err(|e| { - format!("同步 Grok Build Token 到数据库失败: {e}") + if !token.is_empty() && token != PROXY_TOKEN_PLACEHOLDER { + if let Some(provider_config) = provider + .settings_config + .get("config") + .and_then(Value::as_str) + { + let updated = + crate::grok_config::update_api_key(provider_config, &token) + .map_err(|e| { + format!("更新 Grok Build API Key 失败: {e}") })?; - self.db.update_provider(&key, &row).map_err(|e| { - format!("同步 Grok Build Token 到数据库失败: {e}") - })?; - } + provider.settings_config["config"] = json!(updated); + self.persist_synced_live_token( + "grokbuild", + &provider_id, + observed_fingerprint, + provider, + )?; } } } @@ -5373,6 +5378,54 @@ model = "gpt-5.1-codex" ); } + #[test] + fn synced_live_token_cannot_revert_a_concurrent_provider_edit() { + let db = Arc::new(Database::memory().expect("init db")); + let service = ProxyService::new(db.clone()); + let provider = Provider::with_id( + "p1".to_string(), + "Original".to_string(), + json!({ "env": { "ANTHROPIC_AUTH_TOKEN": "stale" } }), + None, + ); + db.reconcile_provider_fixture("claude", &provider) + .expect("save provider"); + + let observed = db + .get_provider_by_id("p1", "claude") + .expect("read observed provider") + .expect("provider exists"); + let fingerprint = provider_row_fingerprint(&observed); + let mut token_update = observed.clone(); + token_update.settings_config["env"]["ANTHROPIC_AUTH_TOKEN"] = json!("from-live"); + + let mut concurrent_edit = observed; + concurrent_edit.name = "Concurrent user edit".to_string(); + if let Some(meta) = concurrent_edit.meta.as_mut() { + meta.custom_endpoints.clear(); + } + let input = provider_to_mutation_input(concurrent_edit); + let key = crate::database::ProviderKey::new("claude", "p1").expect("provider key"); + let row = crate::database::ProviderRowUpdate::from_input(&input).expect("row update"); + db.update_provider(&key, &row) + .expect("persist concurrent user edit"); + + let error = service + .persist_synced_live_token("claude", "p1", fingerprint, token_update) + .expect_err("stale token sync must fail closed"); + assert!(error.contains("changed since it was read"), "{error}"); + + let saved = db + .get_provider_by_id("p1", "claude") + .expect("read saved provider") + .expect("saved provider exists"); + assert_eq!(saved.name, "Concurrent user edit"); + assert_eq!( + saved.settings_config["env"]["ANTHROPIC_AUTH_TOKEN"], + json!("stale") + ); + } + #[tokio::test] #[serial] async fn switch_proxy_target_updates_live_backup_when_taken_over() { diff --git a/src/hooks/useProviderActions.ts b/src/hooks/useProviderActions.ts index bd3fd74b9..3493acbfc 100644 --- a/src/hooks/useProviderActions.ts +++ b/src/hooks/useProviderActions.ts @@ -30,6 +30,7 @@ import { supportsOfficialProxyTakeover, } from "@/utils/providerCapabilities"; import { isOAuthProviderType } from "@/config/constants"; +import { toProviderUpdateInput } from "@/lib/api/providers"; /** * Hook for managing provider actions (add, update, delete, switch) @@ -362,7 +363,10 @@ export function useProviderActions( }, }; - await providersApi.update(updatedProvider, activeApp); + await providersApi.update( + toProviderUpdateInput(updatedProvider), + activeApp, + ); await queryClient.invalidateQueries({ queryKey: ["providers", activeApp], }); diff --git a/src/lib/api/providers.test.ts b/src/lib/api/providers.test.ts new file mode 100644 index 000000000..b3cb14e45 --- /dev/null +++ b/src/lib/api/providers.test.ts @@ -0,0 +1,40 @@ +import { describe, expect, it } from "vitest"; +import type { Provider } from "@/types"; +import { toProviderUpdateInput } from "./providers"; + +describe("toProviderUpdateInput", () => { + it("removes hydrated endpoints and row-state fields from update payloads", () => { + const hydrated: Provider = { + id: "endpoint-provider", + name: "Endpoint provider", + settingsConfig: { env: { API_KEY: "secret" } }, + createdAt: 1_700_000_000, + sortIndex: 7, + inFailoverQueue: true, + meta: { + custom_endpoints: { + "https://one.example": { + url: "https://one.example", + addedAt: null, + }, + }, + usage_script: { + enabled: true, + language: "javascript", + code: "{}", + }, + }, + }; + + const update = toProviderUpdateInput(hydrated); + + expect(update).not.toHaveProperty("createdAt"); + expect(update).not.toHaveProperty("sortIndex"); + expect(update).not.toHaveProperty("inFailoverQueue"); + expect(update.meta).not.toHaveProperty("custom_endpoints"); + expect(update.meta?.usage_script).toEqual(hydrated.meta?.usage_script); + expect( + hydrated.meta?.custom_endpoints?.["https://one.example"].addedAt, + ).toBeNull(); + }); +}); diff --git a/src/lib/api/providers.ts b/src/lib/api/providers.ts index 5f2e1828f..d334eb5a5 100644 --- a/src/lib/api/providers.ts +++ b/src/lib/api/providers.ts @@ -2,6 +2,7 @@ import { invoke } from "@tauri-apps/api/core"; import { listen, type UnlistenFn } from "@tauri-apps/api/event"; import type { Provider, + ProviderMeta, UniversalProvider, UniversalProvidersMap, } from "@/types"; @@ -12,6 +13,39 @@ export interface ProviderSortUpdate { sortIndex: number; } +export type ProviderUpdateMeta = Omit & { + custom_endpoints?: never; +}; + +export type ProviderUpdateInput = Omit< + Provider, + "createdAt" | "sortIndex" | "inFailoverQueue" | "meta" +> & { + meta?: ProviderUpdateMeta; +}; + +export function toProviderUpdateInput(provider: Provider): ProviderUpdateInput { + let meta: ProviderUpdateMeta | undefined; + if (provider.meta) { + const rowMeta = { ...provider.meta }; + delete rowMeta.custom_endpoints; + meta = rowMeta as ProviderUpdateMeta; + } + + return { + id: provider.id, + name: provider.name, + settingsConfig: provider.settingsConfig, + websiteUrl: provider.websiteUrl, + category: provider.category, + notes: provider.notes, + isPartner: provider.isPartner, + meta, + icon: provider.icon, + iconColor: provider.iconColor, + }; +} + export interface ProviderSwitchEvent { appType: AppId; providerId: string; @@ -64,7 +98,7 @@ export const providersApi = { }, async update( - provider: Provider, + provider: ProviderUpdateInput, appId: AppId, originalId?: string, ): Promise { diff --git a/src/lib/query/mutations.ts b/src/lib/query/mutations.ts index 8d6e8bd0f..753991e01 100644 --- a/src/lib/query/mutations.ts +++ b/src/lib/query/mutations.ts @@ -3,7 +3,7 @@ import { useTranslation } from "react-i18next"; import { toast } from "sonner"; import { providersApi, sessionsApi, settingsApi, type AppId } from "@/lib/api"; import type { DeleteSessionOptions } from "@/lib/api/sessions"; -import type { SwitchResult } from "@/lib/api/providers"; +import { toProviderUpdateInput, type SwitchResult } from "@/lib/api/providers"; import type { Provider, SessionMeta, Settings } from "@/types"; import { extractErrorMessage } from "@/utils/errorUtils"; import { generateUUID } from "@/utils/uuid"; @@ -168,7 +168,11 @@ export const useUpdateProviderMutation = (appId: AppId) => { provider: Provider; originalId?: string; }) => { - await providersApi.update(provider, appId, originalId); + await providersApi.update( + toProviderUpdateInput(provider), + appId, + originalId, + ); return provider; }, onSuccess: async (provider, variables) => {