diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 3ef9f1f70..c484796f2 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -799,6 +799,7 @@ dependencies = [ "serde_yaml", "serial_test", "sha2", + "syn 2.0.117", "sys-locale", "tauri", "tauri-build", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index ff69824cf..56e223c91 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -116,3 +116,4 @@ strip = "symbols" [dev-dependencies] serial_test = "3" tempfile = "3" +syn = { version = "2", features = ["full", "visit"] } diff --git a/src-tauri/src/codex_history_migration.rs b/src-tauri/src/codex_history_migration.rs index 9e4f0e023..1547c83f9 100644 --- a/src-tauri/src/codex_history_migration.rs +++ b/src-tauri/src/codex_history_migration.rs @@ -10,6 +10,10 @@ 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}; use crate::error::AppError; +use crate::services::provider::{ + provider_row_fingerprint, provider_to_mutation_input, + reconcile_provider_record_with_precondition, ReconcilePrecondition, +}; use crate::settings::{ CodexOfficialHistoryUnifyMigration, CodexProviderTemplateMigration, CodexThirdPartyHistoryProviderBucketMigration, @@ -663,7 +667,8 @@ fn migrate_codex_provider_templates_to_custom( let providers = db.get_all_providers("codex")?; let mut migrated_provider_ids = Vec::new(); - for (_, provider) in providers { + 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() @@ -694,8 +699,21 @@ fn migrate_codex_provider_templates_to_custom( }; backup_provider_settings_config(&provider.id, &provider.settings_config, backup_root)?; obj.insert("config".to_string(), Value::String(migrated_config_text)); - db.update_provider_settings_config("codex", &provider.id, &settings)?; - migrated_provider_ids.push(provider.id); + let provider_id = provider.id.clone(); + provider.settings_config = settings; + if let Some(meta) = provider.meta.as_mut() { + meta.custom_endpoints.clear(); + } + let input = provider_to_mutation_input(provider); + reconcile_provider_record_with_precondition( + db, + "codex", + input, + ReconcilePrecondition::ExpectPresent { + fingerprint: observed_fingerprint, + }, + )?; + migrated_provider_ids.push(provider_id); } Ok(CodexProviderTemplateBucketMigrationOutcome { @@ -1439,7 +1457,8 @@ base_url = "https://proxy.example/v1" ), ]; for provider in providers { - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); } let mut official = Provider::with_id( @@ -1449,7 +1468,8 @@ base_url = "https://proxy.example/v1" None, ); official.category = Some("official".to_string()); - db.save_provider("codex", &official).expect("save official"); + db.reconcile_provider_fixture("codex", &official) + .expect("save official"); let source_provider_ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert_eq!( @@ -2171,9 +2191,10 @@ base_url = "https://proxy.example/v1" ); official.category = Some("official".to_string()); - db.save_provider("codex", &third_party) + db.reconcile_provider_fixture("codex", &third_party) .expect("save third-party"); - db.save_provider("codex", &official).expect("save official"); + db.reconcile_provider_fixture("codex", &official) + .expect("save official"); let ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert!(ids.contains("rightcode")); @@ -2196,7 +2217,8 @@ base_url = "https://proxy.example/v1" ); provider.category = Some("aggregator".to_string()); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert!(!ids.contains("my-private-relay")); @@ -2216,7 +2238,8 @@ base_url = "https://proxy.example/v1" ); provider.category = Some("aggregator".to_string()); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert!(!ids.contains("my-private-relay")); @@ -2244,7 +2267,8 @@ model_provider = "my-private-relay" ); provider.category = Some("aggregator".to_string()); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert!(!ids.contains("my-private-relay")); @@ -2264,7 +2288,8 @@ model_provider = "my-private-relay" ); provider.category = Some("aggregator".to_string()); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert!(ids.contains("aihubmix")); @@ -2285,7 +2310,8 @@ model_provider = "my-private-relay" ); provider.category = Some("aggregator".to_string()); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert!(ids.contains("ccswitch")); @@ -2317,7 +2343,8 @@ model = "gpt-5.4" }), None, ); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let (outcome, backup_dir) = migrate_provider_templates_for_test(&db); assert_eq!(outcome.migrated_provider_ids, vec!["legacy".to_string()]); @@ -2390,7 +2417,8 @@ base_url = "https://aihubmix.example/v1" }), None, ); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db); assert_eq!( @@ -2446,7 +2474,8 @@ base_url = "http://localhost:8080/v1" }), None, ); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db); assert!(outcome.migrated_provider_ids.is_empty()); @@ -2495,7 +2524,8 @@ base_url = "https://proxy.example/v1" }), None, ); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db); assert!(outcome.migrated_provider_ids.is_empty()); @@ -2552,7 +2582,8 @@ model_provider = "aihubmix" }), None, ); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let (outcome, _backup_dir) = migrate_provider_templates_for_test(&db); assert_eq!(outcome.migrated_provider_ids, vec!["profiled".to_string()]); @@ -2601,7 +2632,8 @@ model_provider = "aihubmix" provider.category = Some("custom".to_string()); provider.created_at = Some(1); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert!(!ids.contains("my-private-relay")); @@ -2622,7 +2654,8 @@ model_provider = "aihubmix" ); provider.category = Some("custom".to_string()); - db.save_provider("codex", &provider).expect("save provider"); + db.reconcile_provider_fixture("codex", &provider) + .expect("save provider"); let ids = collect_source_model_provider_ids(&db).expect("collect ids"); assert!(!ids.contains("my-local-relay")); diff --git a/src-tauri/src/commands/provider.rs b/src-tauri/src/commands/provider.rs index 54243dfaa..376cf0f73 100644 --- a/src-tauri/src/commands/provider.rs +++ b/src-tauri/src/commands/provider.rs @@ -4,8 +4,9 @@ use tauri::{Emitter, Manager, State}; use crate::app_config::AppType; use crate::commands::copilot::CopilotAuthState; use crate::commands::xai_oauth::XaiOAuthState; +use crate::database::NewProviderAggregate; use crate::error::AppError; -use crate::provider::{ClaudeDesktopMode, Provider}; +use crate::provider::{ClaudeDesktopMode, Provider, ProviderMutationInput}; use crate::services::{ EndpointLatency, ProviderService, ProviderSortUpdate, SpeedtestService, SwitchResult, }; @@ -39,7 +40,7 @@ pub fn get_current_provider(state: State<'_, AppState>, app: String) -> Result, app: String, - provider: Provider, + provider: ProviderMutationInput, #[allow(non_snake_case)] addToLive: Option, ) -> Result { let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?; @@ -51,7 +52,7 @@ pub fn add_provider( pub fn update_provider( state: State<'_, AppState>, app: String, - provider: Provider, + provider: ProviderMutationInput, #[allow(non_snake_case)] originalId: Option, ) -> Result { let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?; @@ -250,7 +251,13 @@ pub fn import_claude_desktop_providers_from_claude( state .db - .save_provider(AppType::ClaudeDesktop.as_str(), &desktop_provider) + .create_provider( + NewProviderAggregate::from_input( + AppType::ClaudeDesktop.as_str(), + crate::services::provider::provider_to_mutation_input(desktop_provider), + ) + .map_err(|e| e.to_string())?, + ) .map_err(|e| e.to_string())?; imported += 1; } diff --git a/src-tauri/src/database/dao/mod.rs b/src-tauri/src/database/dao/mod.rs index bb78f3075..f2b7ed1b6 100644 --- a/src-tauri/src/database/dao/mod.rs +++ b/src-tauri/src/database/dao/mod.rs @@ -6,6 +6,9 @@ pub mod failover; pub mod mcp; pub mod profiles; pub mod prompts; +pub mod provider_write; +#[cfg(test)] +mod provider_write_certification; pub mod providers; pub mod providers_seed; pub mod proxy; diff --git a/src-tauri/src/database/dao/provider_write.rs b/src-tauri/src/database/dao/provider_write.rs new file mode 100644 index 000000000..3ddab2ac2 --- /dev/null +++ b/src-tauri/src/database/dao/provider_write.rs @@ -0,0 +1,640 @@ +use crate::database::{lock_conn, Database}; +use crate::error::AppError; +use crate::provider::{ProviderMeta, ProviderMutationInput}; +use crate::settings::CustomEndpoint; +use rusqlite::{params, OptionalExtension, Transaction}; +use serde_json::Value; +use std::collections::HashSet; + +use super::providers::{StoredProviderRow, PROVIDER_SELECT}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ProviderKey { + app_type: String, + id: String, +} + +impl ProviderKey { + pub fn new(app_type: impl Into, id: impl Into) -> Result { + let app_type = app_type.into(); + let id = id.into(); + if app_type.trim().is_empty() || id.trim().is_empty() { + return Err(AppError::InvalidInput( + "provider app type and id must be non-empty".to_string(), + )); + } + Ok(Self { app_type, id }) + } + + pub fn app_type(&self) -> &str { + &self.app_type + } + + pub fn id(&self) -> &str { + &self.id + } +} + +#[derive(Debug, Clone)] +pub struct ProviderRowUpdate { + name: String, + settings_config: Value, + website_url: Option, + category: Option, + notes: Option, + meta: ProviderMeta, + icon: Option, + icon_color: Option, +} + +impl ProviderRowUpdate { + pub fn from_input(input: &ProviderMutationInput) -> Result { + let meta = input.meta.clone().unwrap_or_default(); + if !meta.custom_endpoints.is_empty() { + return Err(AppError::InvalidInput( + "provider update must not contain customEndpoints; use endpoint operations" + .to_string(), + )); + } + Ok(Self { + name: input.name.clone(), + settings_config: input.settings_config.clone(), + website_url: input.website_url.clone(), + category: input.category.clone(), + notes: input.notes.clone(), + meta, + icon: input.icon.clone(), + icon_color: input.icon_color.clone(), + }) + } +} + +#[derive(Debug, Clone)] +pub struct ProviderRowCreate { + content: ProviderRowUpdate, + created_at: Option, +} + +#[derive(Debug, Clone)] +pub struct NewEndpoint { + url: String, + added_at: Option, + last_used: Option, +} + +impl NewEndpoint { + pub fn new( + url: impl Into, + added_at: Option, + last_used: Option, + ) -> Result { + let url = url.into(); + if url.trim().is_empty() { + return Err(AppError::InvalidInput( + "provider endpoint URL cannot be empty".to_string(), + )); + } + Ok(Self { + url, + added_at, + last_used, + }) + } + + pub fn now(url: impl Into) -> Result { + Self::new(url, Some(chrono::Utc::now().timestamp_millis()), None) + } +} + +impl TryFrom for NewEndpoint { + type Error = AppError; + + fn try_from(endpoint: CustomEndpoint) -> Result { + Self::new(endpoint.url, endpoint.added_at, endpoint.last_used) + } +} + +#[derive(Debug, Clone)] +pub struct NewProviderAggregate { + key: ProviderKey, + row: ProviderRowCreate, + sort_index: Option, + in_failover_queue: bool, + initial_endpoints: Vec, +} + +impl NewProviderAggregate { + pub fn from_input(app_type: &str, mut input: ProviderMutationInput) -> Result { + let endpoints = input + .meta + .as_mut() + .map(|meta| std::mem::take(&mut meta.custom_endpoints)) + .unwrap_or_default(); + let mut seen = HashSet::with_capacity(endpoints.len()); + let mut initial_endpoints = Vec::with_capacity(endpoints.len()); + for (key, endpoint) in endpoints { + let normalized_key = key.trim().trim_end_matches('/').to_string(); + let normalized_url = endpoint.url.trim().trim_end_matches('/').to_string(); + if normalized_key != normalized_url { + return Err(AppError::InvalidInput(format!( + "provider endpoint key '{key}' must match endpoint URL '{}'", + endpoint.url + ))); + } + if !seen.insert(normalized_url.clone()) { + return Err(AppError::InvalidInput(format!( + "duplicate initial provider endpoint '{}'", + endpoint.url + ))); + } + initial_endpoints.push(NewEndpoint::new( + normalized_url, + endpoint.added_at, + endpoint.last_used, + )?); + } + let key = ProviderKey::new(app_type, input.id.clone())?; + let row = ProviderRowCreate { + content: ProviderRowUpdate::from_input(&input)?, + created_at: input.created_at, + }; + Ok(Self { + key, + row, + sort_index: input.sort_index, + in_failover_queue: input.in_failover_queue, + initial_endpoints, + }) + } +} + +#[derive(Debug, Clone)] +pub struct RenameProvider { + source: ProviderKey, + target_id: String, + row: ProviderRowUpdate, +} + +impl RenameProvider { + pub fn from_input( + source: ProviderKey, + input: &ProviderMutationInput, + ) -> Result { + if !matches!(source.app_type(), "opencode" | "openclaw") { + return Err(AppError::InvalidInput( + "provider key changes are restricted to additive OpenCode/OpenClaw providers" + .to_string(), + )); + } + if source.id() == input.id { + return Err(AppError::InvalidInput( + "provider rename requires a different target id".to_string(), + )); + } + if input.id.trim().is_empty() { + return Err(AppError::InvalidInput( + "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, + }) + } +} + +fn encode_row(row: &ProviderRowUpdate) -> Result<(String, String), AppError> { + let settings_config = serde_json::to_string(&row.settings_config).map_err(|error| { + AppError::Database(format!("failed to serialize settings_config: {error}")) + })?; + let meta = serde_json::to_string(&row.meta).map_err(|error| { + AppError::Database(format!("failed to serialize provider meta: {error}")) + })?; + Ok((settings_config, meta)) +} + +fn insert_row( + tx: &Transaction<'_>, + key: &ProviderKey, + row: &ProviderRowUpdate, + created_at: Option, + sort_index: Option, + is_current: bool, + in_failover_queue: bool, +) -> Result<(), AppError> { + let (settings_config, meta) = encode_row(row)?; + tx.execute( + "INSERT INTO providers ( + id, app_type, name, settings_config, website_url, category, + created_at, sort_index, notes, icon, icon_color, meta, + is_current, in_failover_queue + ) VALUES ( + ?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14 + )", + params![ + key.id, + key.app_type, + row.name, + settings_config, + row.website_url, + row.category, + created_at, + sort_index, + row.notes, + row.icon, + row.icon_color, + meta, + is_current, + in_failover_queue, + ], + ) + .map_err(|error| match &error { + rusqlite::Error::SqliteFailure(code, _) + if matches!( + code.extended_code, + rusqlite::ffi::SQLITE_CONSTRAINT_PRIMARYKEY + | rusqlite::ffi::SQLITE_CONSTRAINT_UNIQUE + ) => + { + AppError::Conflict(format!( + "provider '{}/{}' already exists", + key.app_type, key.id + )) + } + _ => AppError::Database(error.to_string()), + })?; + Ok(()) +} + +fn insert_endpoint( + tx: &Transaction<'_>, + key: &ProviderKey, + endpoint: &NewEndpoint, +) -> Result<(), AppError> { + tx.execute( + "INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at, last_used) + VALUES (?1, ?2, ?3, ?4, ?5)", + params![ + key.id, + key.app_type, + endpoint.url, + endpoint.added_at, + endpoint.last_used + ], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(()) +} + +/// Exact aggregate replacement is sealed inside the DAO parent module. The +/// catalog compensation coordinator introduced with the ordered mutation +/// pipeline is the only intended caller. +#[allow(dead_code)] +// The certification contract keeps immutable creation time separate from the +// mutable row DTO and calls this sealed helper directly with the full snapshot. +#[allow(clippy::too_many_arguments)] +pub(super) fn restore_provider_aggregate_on_tx( + tx: &Transaction<'_>, + key: &ProviderKey, + row: &ProviderRowUpdate, + created_at: Option, + sort_index: Option, + is_current: bool, + in_failover_queue: bool, + endpoints: &[NewEndpoint], +) -> Result<(), AppError> { + let updated = update_row(tx, key, row)?; + if updated == 0 { + insert_row( + tx, + key, + row, + created_at, + sort_index, + is_current, + in_failover_queue, + )?; + } else { + // Exact compensation is the only path allowed to restore immutable + // creation time after a prior aggregate mutation. + tx.execute( + "UPDATE providers SET created_at = ?1 WHERE id = ?2 AND app_type = ?3", + params![created_at, key.id, key.app_type], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + } + tx.execute( + "DELETE FROM provider_endpoints WHERE provider_id = ?1 AND app_type = ?2", + params![key.id, key.app_type], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + for endpoint in endpoints { + insert_endpoint(tx, key, endpoint)?; + } + // State and order are maintained by their dedicated authorities. Exact + // compensation may restore their captured values without exposing them in + // ProviderRowUpdate. + tx.execute( + "UPDATE providers + SET sort_index = ?1, is_current = ?2, in_failover_queue = ?3 + WHERE id = ?4 AND app_type = ?5", + params![ + sort_index, + is_current, + in_failover_queue, + key.id, + key.app_type + ], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(()) +} + +fn update_row( + tx: &Transaction<'_>, + key: &ProviderKey, + row: &ProviderRowUpdate, +) -> Result { + let (settings_config, meta) = encode_row(row)?; + tx.execute( + "UPDATE providers SET + name = ?1, + settings_config = ?2, + website_url = ?3, + category = ?4, + notes = ?5, + icon = ?6, + icon_color = ?7, + meta = ?8 + WHERE id = ?9 AND app_type = ?10", + params![ + row.name, + settings_config, + row.website_url, + row.category, + row.notes, + row.icon, + row.icon_color, + meta, + key.id, + key.app_type, + ], + ) + .map_err(|error| AppError::Database(error.to_string())) +} + +impl Database { + pub fn create_provider(&self, input: NewProviderAggregate) -> Result<(), AppError> { + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + insert_row( + &tx, + &input.key, + &input.row.content, + input.row.created_at, + input.sort_index, + false, + input.in_failover_queue, + )?; + for endpoint in &input.initial_endpoints { + insert_endpoint(&tx, &input.key, endpoint)?; + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub fn update_provider( + &self, + key: &ProviderKey, + row: &ProviderRowUpdate, + ) -> Result<(), AppError> { + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + if update_row(&tx, key, row)? != 1 { + return Err(AppError::NotFound(format!( + "provider '{}/{}'", + key.app_type, key.id + ))); + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn update_provider_if_content_fingerprint( + &self, + key: &ProviderKey, + expected_fingerprint: &str, + row: &ProviderRowUpdate, + ) -> Result<(), AppError> { + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + let current = tx + .query_row( + &format!("{PROVIDER_SELECT} WHERE id = ?1 AND app_type = ?2"), + params![key.id, key.app_type], + StoredProviderRow::from_row, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))? + .ok_or_else(|| AppError::NotFound(format!("provider '{}/{}'", key.app_type, key.id)))? + .decode(key.app_type())?; + if current.row_content_fingerprint() != expected_fingerprint { + return Err(AppError::Conflict(format!( + "provider '{}/{}' changed since it was read", + key.app_type, key.id + ))); + } + if update_row(&tx, key, row)? != 1 { + return Err(AppError::NotFound(format!( + "provider '{}/{}'", + key.app_type, key.id + ))); + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string())) + } + + 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, meta + FROM providers + WHERE id = ?1 AND app_type = ?2", + params![input.source.id, input.source.app_type], + |row| { + Ok(( + row.get::<_, Option>(0)?, + row.get::<_, bool>(1)?, + row.get::<_, bool>(2)?, + row.get::<_, Option>(3)?, + row.get::<_, Option>(4)?, + row.get::<_, String>(5)?, + )) + }, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))? + .ok_or_else(|| { + AppError::NotFound(format!( + "provider '{}/{}'", + input.source.app_type, input.source.id + )) + })?; + if matches!(source_state.3.as_deref(), Some("omo" | "omo-slim")) { + return Err(AppError::InvalidInput( + "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, + &target, + &input.row, + source_state.4, + source_state.0, + source_state.1, + source_state.2, + )?; + tx.execute( + "INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at, last_used) + SELECT ?1, app_type, url, added_at, last_used + FROM provider_endpoints + WHERE provider_id = ?2 AND app_type = ?3 + ORDER BY id", + params![target.id, input.source.id, input.source.app_type], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if tx + .execute( + "DELETE FROM providers WHERE id = ?1 AND app_type = ?2", + params![input.source.id, input.source.app_type], + ) + .map_err(|error| AppError::Database(error.to_string()))? + != 1 + { + return Err(AppError::NotFound(format!( + "provider '{}/{}'", + input.source.app_type, input.source.id + ))); + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub fn add_provider_endpoint( + &self, + key: &ProviderKey, + endpoint: NewEndpoint, + ) -> Result<(), AppError> { + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + insert_endpoint(&tx, key, &endpoint)?; + tx.commit() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub fn remove_provider_endpoint(&self, key: &ProviderKey, url: &str) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + if conn + .execute( + "DELETE FROM provider_endpoints + WHERE provider_id = ?1 AND app_type = ?2 AND url = ?3", + params![key.id, key.app_type, url], + ) + .map_err(|error| AppError::Database(error.to_string()))? + != 1 + { + return Err(AppError::NotFound(format!( + "provider endpoint '{}/{}/{}'", + key.app_type, key.id, url + ))); + } + Ok(()) + } + + pub fn touch_provider_endpoint( + &self, + key: &ProviderKey, + url: &str, + at: i64, + ) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + if conn + .execute( + "UPDATE provider_endpoints + SET last_used = ?1 + WHERE provider_id = ?2 AND app_type = ?3 AND url = ?4", + params![at, key.id, key.app_type, url], + ) + .map_err(|error| AppError::Database(error.to_string()))? + != 1 + { + return Err(AppError::NotFound(format!( + "provider endpoint '{}/{}/{}'", + key.app_type, key.id, url + ))); + } + Ok(()) + } + + pub(crate) fn update_provider_sort_index( + &self, + key: &ProviderKey, + sort_index: usize, + ) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + if conn + .execute( + "UPDATE providers SET sort_index = ?1 WHERE id = ?2 AND app_type = ?3", + params![sort_index, key.id, key.app_type], + ) + .map_err(|error| AppError::Database(error.to_string()))? + != 1 + { + return Err(AppError::NotFound(format!( + "provider '{}/{}'", + key.app_type, key.id + ))); + } + Ok(()) + } +} diff --git a/src-tauri/src/database/dao/provider_write_certification.rs b/src-tauri/src/database/dao/provider_write_certification.rs new file mode 100644 index 000000000..693be90ae --- /dev/null +++ b/src-tauri/src/database/dao/provider_write_certification.rs @@ -0,0 +1,2204 @@ +#![cfg(test)] +//! 前置工程 A:Provider 写面认证测试套件 v5(测试先行) +//! +//! 本文件是认证契约的可执行字面,固化 R1–R4 盲审与三轮对抗审查揭示的全部 +//! 写面故障场景。规则: +//! - 全绿是进入前置 A 盲审的前置条件,但不是充分条件; +//! - 对写面(provider_write.rs)新增任何函数、对 infra 三文件的任何改动、 +//! 对本文件清单的任何调整,均须先经裁决。 +//! +//! ## 已裁决的语义决定 +//! 1. `created_at` 不可变:update 不得改写创建时间;`ProviderRowUpdate` +//! 必须不含 `created_at` 字段。create/restore 所需创建时间由各自入参单独 +//! 携带(restore 已由裁决方加 `created_at` 参数;create 拆分是实现方职责)。 +//! `tests/fixtures/pi/provider-write-api-v1.json` 与 +//! `architecture_tests.rs` 的 snapshot 生成器随 DTO 拆分同步更新(生成器 +//! 必须收录新建的 create/restore 专属类型)——这是实现方义务;拆分落地前 +//! 旧 fixture 保持一致属预期。 +//! 2. 结构化冲突:重复 create(含并发输家)必须返回 `AppError::Conflict`。 +//! 注:`AppError` 的 `Serialize` 目前把错误序列化为字符串,IPC 层的结构化 +//! discriminant 是后续裁决项(P1),不在本 PR 强制。 +//! 3. reconcile 显式前置期望(T9):脚手架已落地(`ReconcilePrecondition`、 +//! `provider_row_fingerprint`(规范化排序哈希,不含 endpoint)、 +//! `reconcile_provider_record_with_precondition`,故意保留旧语义使 T9 红)。 +//! 实现方必须以**单事务原语**实现:ExpectAbsent → `create_provider` +//! (冲突 → Conflict);ExpectPresent → 新 DAO 原语 +//! `update_provider_if_content_fingerprint`(单事务内读-比-写,过期 → +//! Conflict)。reconcile 函数体内禁止内联 aggregate 读取后再分支 +//! (`certify_reconcile_uses_single_transaction_primitives` 机械强制)。 +//! **盲审重点核查项**:`update_provider_if_content_fingerprint` 内部必须 +//! 在单次连接锁/单事务内完成读-比-写(本仓库为单连接 Mutex,持锁即全局 +//! 串行);静态测试只能约束委托关系,原语内部"读后释放锁再写"的变体由 +//! 组件盲审逐行核查——reviewer 材料必须包含本条。 +//! 完成后迁移全部调用方并删除旧 `reconcile_provider_record`,由裁决方将 +//! 旧符号加入禁止清单。 +//! +//! ## 扫描器 authority 表(精确相对路径 × DML 种类 × 列集合) +//! - `database/dao/provider_write.rs`:全部 provider DML 允许(写面本体); +//! - `database/dao/providers.rs`:仅 `UPDATE providers`(列 ⊆ {is_current}) +//! 与 `DELETE FROM providers`; +//! - `database/dao/failover.rs`:仅 `UPDATE providers`(列 ⊆ {in_failover_queue}); +//! - infra 三文件:Deferred to 前置工程 B,由 SHA-256 基线冻结兜底; +//! - 测试专属文件必须自带文件级 `#![cfg(test)]`(注册元测试机械强制;借用 +//! 他处注册的伪测试名生产文件在此失败),扫描器凭该属性天然跳过其内容; +//! - cfg 判定按布尔语义:仅当谓词蕴含 test 才跳过;`not(test)`、 +//! `any(test, unix)` 一律扫描; +//! - 宏 token 纳入扫描:宏内字符串经 `syn::LitStr::value()` 解码(覆盖 +//! `\xNN`/`\u{}`);纯字符串宏(`concat!`)拼接整体参与分类;含非字面量 +//! token 且出现 provider DML 锚点的宏(`format!`、`stringify!` 构造) +//! 一律 fail-closed 记为违规;`include!` 全生产源禁止,`include_str!` +//! token 含 `.sql` 时禁止; +//! - 解析 fail-closed:SET 子句引号/括号不闭合或列集无法确定时产出 +//! `!unparseable` 哨兵列,任何 authority 不放行; +//! - 已知残余风险(接受,由盲审与前置 B 兜底,须向 reviewer 声明): +//! 完全运行时构造、无任何可识别字面锚点的动态 SQL;trigger/view 间接写; +//! SQLite Backup API 整库复制;`r#"..."#` 多井号原始字符串宏字面量; +//! `#[path]`/非 `.rs` 重定向包含;非 `.sql` 扩展名文件装载 SQL 文本; +//! 宏展开生成的 impl(inventory 已禁 item 级宏与 out-of-line 子模块, +//! 属性宏路径由盲审兜底);T10 的 identifier 探测只证明"标识符存在", +//! 真实调用行为由盲审核对。 +//! 本清单为对抗加固的**收口边界**:静态扫描是护栏,组件盲审才是认证; +//! 清单外的新绕过按盲审 finding 处理,不再无限扩充扫描器。 +//! + +use crate::database::dao::provider_write::{ + self, NewEndpoint, NewProviderAggregate, ProviderKey, ProviderRowUpdate, RenameProvider, +}; +use crate::database::Database; +use crate::error::AppError; +use crate::provider::{ProviderMeta, ProviderMutationInput}; +use crate::services::provider::{ + provider_row_fingerprint, reconcile_provider_record_with_precondition, ReconcilePrecondition, +}; +use crate::settings::CustomEndpoint; +use regex::Regex; +use serde_json::json; +use std::collections::{BTreeSet, HashMap}; +use std::fs; +use std::path::{Path, PathBuf}; +use std::sync::LazyLock; +use syn::visit::{self, Visit}; +use syn::{Attribute, ExprLit, ImplItem, Item, Lit, Meta}; + +// --------------------------------------------------------------------------- +// 测试基建 +// --------------------------------------------------------------------------- + +fn db() -> Database { + Database::memory().expect("memory db") +} + +fn base_input(id: &str, name: &str) -> ProviderMutationInput { + ProviderMutationInput { + id: id.to_string(), + name: name.to_string(), + settings_config: json!({"env": {"KEY": "v"}}), + website_url: None, + category: None, + created_at: Some(1_700_000_000), + sort_index: None, + notes: None, + meta: None, + icon: None, + icon_color: None, + in_failover_queue: false, + } +} + +fn with_endpoints( + mut input: ProviderMutationInput, + endpoints: &[(&str, Option, Option)], +) -> ProviderMutationInput { + let mut map = HashMap::new(); + for (url, added_at, last_used) in endpoints { + map.insert( + url.to_string(), + CustomEndpoint { + url: url.to_string(), + added_at: *added_at, + last_used: *last_used, + }, + ); + } + let mut meta = input.meta.take().unwrap_or_default(); + meta.custom_endpoints = map; + input.meta = Some(meta); + input +} + +type RowSnapshot = ( + String, // name + String, // settings_config + Option, // website_url + Option, // category + Option, // created_at + Option, // sort_index + Option, // notes + Option, // icon + Option, // icon_color + String, // meta + i64, // is_current + i64, // in_failover_queue +); + +type EndpointRows = Vec<(String, Option, Option)>; + +/// 逐列快照,用于"零副作用"断言。绕过 hydration 直接读库,以免 hydration +/// 自身的有损转换掩盖破坏;查询错误必须炸出来,不得伪装成"不存在"。 +fn snapshot(database: &Database, app_type: &str, id: &str) -> (Option, EndpointRows) { + use rusqlite::OptionalExtension; + let conn = database.conn.lock().expect("lock certification database"); + let row = conn + .query_row( + "SELECT name, settings_config, website_url, category, created_at, + sort_index, notes, icon, icon_color, meta, is_current, in_failover_queue + FROM providers WHERE id = ?1 AND app_type = ?2", + rusqlite::params![id, app_type], + |r| { + Ok(( + r.get(0)?, + r.get(1)?, + r.get(2)?, + r.get(3)?, + r.get(4)?, + r.get(5)?, + r.get(6)?, + r.get(7)?, + r.get(8)?, + r.get(9)?, + r.get(10)?, + r.get(11)?, + )) + }, + ) + .optional() + .expect("snapshot row query must not error"); + let mut stmt = conn + .prepare( + "SELECT url, added_at, last_used FROM provider_endpoints + WHERE provider_id = ?1 AND app_type = ?2 ORDER BY url", + ) + .expect("prepare endpoint snapshot"); + let endpoints = stmt + .query_map(rusqlite::params![id, app_type], |r| { + Ok((r.get(0)?, r.get(1)?, r.get(2)?)) + }) + .expect("query endpoints") + .collect::, _>>() + .expect("collect endpoints"); + (row, endpoints) +} + +const ENDPOINT_REJECT_MESSAGE: &str = "certification injected endpoint failure"; + +fn install_endpoint_reject_trigger(database: &Database) { + let conn = database.conn.lock().expect("lock certification database"); + conn.execute_batch( + "CREATE TRIGGER certification_reject_endpoint_insert + BEFORE INSERT ON provider_endpoints + BEGIN SELECT RAISE(ABORT, 'certification injected endpoint failure'); END;", + ) + .expect("install endpoint reject trigger"); +} + +fn source_root() -> PathBuf { + Path::new(env!("CARGO_MANIFEST_DIR")).join("src") +} + +fn relative_source_path(root: &Path, file: &Path) -> String { + file.strip_prefix(root) + .expect("source file under root") + .to_string_lossy() + .replace('\\', "/") +} + +fn collect_rs_files(dir: &Path, out: &mut Vec) { + let Ok(entries) = fs::read_dir(dir) else { + return; + }; + for entry in entries.flatten() { + let path = entry.path(); + if path.is_dir() { + collect_rs_files(&path, out); + } else if path.extension().and_then(|ext| ext.to_str()) == Some("rs") { + out.push(path); + } + } +} + +// --------------------------------------------------------------------------- +// cfg 布尔语义:仅当谓词蕴含 test 才视为 test-only +// --------------------------------------------------------------------------- + +fn split_top_level(args: &str) -> Vec { + let mut parts = Vec::new(); + let mut depth: i32 = 0; + let mut current = String::new(); + for c in args.chars() { + match c { + '(' => { + depth += 1; + current.push(c); + } + ')' => { + depth -= 1; + current.push(c); + } + ',' if depth == 0 => { + parts.push(current.trim().to_string()); + current = String::new(); + } + _ => current.push(c), + } + } + if !current.trim().is_empty() { + parts.push(current.trim().to_string()); + } + parts +} + +fn strip_call<'a>(expr: &'a str, name: &str) -> Option<&'a str> { + let rest = expr.strip_prefix(name)?.trim_start(); + let rest = rest.strip_prefix('(')?; + rest.strip_suffix(')') +} + +/// `test` → true;`all(..)` 任一分支蕴含 test → true;`any(..)` 需全部分支 +/// 蕴含 test;`not(..)` 与其他谓词一律 false(保守:继续扫描)。 +fn cfg_expr_requires_test(expr: &str) -> bool { + let expr = expr.trim(); + if expr == "test" { + return true; + } + if let Some(args) = strip_call(expr, "all") { + return split_top_level(args) + .iter() + .any(|part| cfg_expr_requires_test(part)); + } + if let Some(args) = strip_call(expr, "any") { + let parts = split_top_level(args); + return !parts.is_empty() && parts.iter().all(|part| cfg_expr_requires_test(part)); + } + false +} + +fn attrs_mark_test_only(attrs: &[Attribute]) -> bool { + attrs.iter().any(|attribute| { + attribute.path().is_ident("cfg") + && matches!( + &attribute.meta, + Meta::List(list) if cfg_expr_requires_test(&list.tokens.to_string()) + ) + }) +} + +// --------------------------------------------------------------------------- +// 扫描器 v4:syn AST(含宏 token)+ 列敏感 DML 分类 +// --------------------------------------------------------------------------- + +const STATE_COLUMNS_PROVIDERS_RS: [&str; 1] = ["is_current"]; +const STATE_COLUMNS_FAILOVER_RS: [&str; 1] = ["in_failover_queue"]; +/// restore 面基础设施文件:DML 列权限扫描对它们另有归属规则。 +const INFRA_FILES: [&str; 3] = [ + "database/schema.rs", + "database/migration.rs", + "database/backup.rs", +]; + +#[derive(Debug, Clone, PartialEq, Eq)] +enum Dml { + Insert { table: String }, + Delete { table: String }, + Update { table: String, columns: Vec }, +} + +/// 表 token:引号成对匹配的交替(裸形式带 `\b`)。R8 终审:引号表名在 +/// "可选闭引号 + \b" 的写法上必然失配,必须成对交替。 +const TABLE_TOKEN: &str = r#"("provider_endpoints"|"providers"|'provider_endpoints'|'providers'|`provider_endpoints`|`providers`|\[provider_endpoints\]|\[providers\]|provider_endpoints\b|providers\b)"#; +const NAME_PREFIX: &str = + r#"(?:(?:"(?:[^"]|"")*"|'(?:[^']|'')*'|`[^`]*`|\[[^\]]*\]|\w+)\s*\.\s*)?"#; +const NAME_TOKEN: &str = r#"(?:"(?:[^"]|"")*"|'(?:[^']|'')*'|`[^`]*`|\[[^\]]*\]|\w+)"#; + +fn table_from_capture(raw: &str) -> String { + raw.trim_matches(|c: char| !c.is_ascii_alphanumeric() && c != '_') + .to_lowercase() +} + +static INSERT_HEAD: LazyLock = LazyLock::new(|| { + Regex::new(&format!( + r#"(?is)\b(?:REPLACE|INSERT(?:\s+OR\s+(?:ABORT|FAIL|IGNORE|REPLACE|ROLLBACK))?)\s+INTO\s+{NAME_PREFIX}{TABLE_TOKEN}"# + )) + .expect("compile insert head") +}); +static DELETE_HEAD: LazyLock = LazyLock::new(|| { + Regex::new(&format!( + r#"(?is)\bDELETE\s+FROM\s+{NAME_PREFIX}{TABLE_TOKEN}"# + )) + .expect("compile delete head") +}); +static UPDATE_HEAD: LazyLock = LazyLock::new(|| { + Regex::new(&format!( + r#"(?is)\bUPDATE(?:\s+OR\s+(?:ABORT|FAIL|IGNORE|REPLACE|ROLLBACK))?\s+{NAME_PREFIX}{TABLE_TOKEN}(?:\s+(?:AS\s+{NAME_TOKEN}|NOT\s+INDEXED|INDEXED\s+BY\s+{NAME_TOKEN}|{NAME_TOKEN}))*?\s+SET\b"# + )) + .expect("compile update head") +}); +static MACRO_STRING: LazyLock = LazyLock::new(|| { + Regex::new(r#""(?:[^"\\]|\\.)*"|r"[^"]*""#).expect("compile macro string extractor") +}); +static FORBIDDEN_SYMBOL: LazyLock = LazyLock::new(|| { + Regex::new(r"\bupdate_provider_settings_config\b").expect("compile forbidden symbol") +}); + +fn contains_provider_dml_anchor(text: &str) -> bool { + INSERT_HEAD.is_match(text) || DELETE_HEAD.is_match(text) || UPDATE_HEAD.is_match(text) +} + +/// 去掉 SQL 注释;引号感知(单引号/双引号/反引号/方括号内的 `--`、`/*` +/// 不是注释)。 +fn strip_sql_comments(sql: &str) -> String { + let mut out = String::with_capacity(sql.len()); + let bytes = sql.as_bytes(); + let mut i = 0; + let mut quote: Option = None; + while i < bytes.len() { + let b = bytes[i]; + if let Some(q) = quote { + out.push(b as char); + let closing = match q { + b'[' => b']', + other => other, + }; + if b == closing { + quote = None; + } + i += 1; + continue; + } + match b { + b'\'' | b'"' | b'`' | b'[' => { + quote = Some(b); + out.push(b as char); + i += 1; + } + b'/' if i + 1 < bytes.len() && bytes[i + 1] == b'*' => { + i += 2; + while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') { + i += 1; + } + i = (i + 2).min(bytes.len()); + out.push(' '); + } + b'-' if i + 1 < bytes.len() && bytes[i + 1] == b'-' => { + while i < bytes.len() && bytes[i] != b'\n' { + i += 1; + } + } + _ => { + out.push(b as char); + i += 1; + } + } + } + out +} + +const UNPARSEABLE: &str = "!unparseable"; + +/// SET 列解析:括号深度 + 四类引号感知,顶层 `WHERE`(前一字符不得是 +/// `:@$?` 参数记号)或 `;` 终止;列名剥离 alias 前缀与引号;引号/括号 +/// 不闭合或列集为空时 fail-closed 产出哨兵列。 +fn parse_set_columns(tail: &str) -> Vec { + let upper = tail.to_ascii_uppercase(); + let bytes = upper.as_bytes(); + let raw = tail.as_bytes(); + let mut depth: i32 = 0; + let mut quote: Option = None; + let mut end = tail.len(); + let mut i = 0; + while i < bytes.len() { + let b = raw[i]; + if let Some(q) = quote { + let closing = match q { + b'[' => b']', + other => other, + }; + if b == closing { + quote = None; + } + i += 1; + continue; + } + match b { + b'\'' | b'"' | b'`' | b'[' => quote = Some(b), + b'(' => depth += 1, + b')' => { + depth -= 1; + if depth < 0 { + return vec![UNPARSEABLE.to_string()]; + } + } + b'W' | b'w' if depth == 0 => { + let prev = if i == 0 { b' ' } else { raw[i - 1] }; + let boundary_before = !(prev.is_ascii_alphanumeric() + || prev == b'_' + || matches!(prev, b':' | b'@' | b'$' | b'?')); + if boundary_before && upper[i..].starts_with("WHERE") { + let after = i + 5; + let boundary_after = after >= bytes.len() + || !(bytes[after].is_ascii_alphanumeric() || bytes[after] == b'_'); + if boundary_after { + end = i; + break; + } + } + } + b';' if depth == 0 => { + end = i; + break; + } + _ => {} + } + i += 1; + } + if quote.is_some() || depth != 0 { + return vec![UNPARSEABLE.to_string()]; + } + let clause = &tail[..end]; + let mut columns = Vec::new(); + let mut segment_start = 0; + let mut depth: i32 = 0; + let mut quote: Option = None; + let clause_bytes = clause.as_bytes(); + let push_segment = |segment: &str, columns: &mut Vec| { + if let Some(identifier) = segment.split('=').next() { + let identifier = identifier + .trim() + .rsplit('.') + .next() + .unwrap_or("") + .trim_matches(|c: char| !c.is_ascii_alphanumeric() && c != '_') + .to_lowercase(); + if !identifier.is_empty() { + columns.push(identifier); + } + } + }; + for (i, &b) in clause_bytes.iter().enumerate() { + if let Some(q) = quote { + let closing = match q { + b'[' => b']', + other => other, + }; + if b == closing { + quote = None; + } + continue; + } + match b { + b'\'' | b'"' | b'`' | b'[' => quote = Some(b), + b'(' => depth += 1, + b')' => depth -= 1, + b',' if depth == 0 => { + push_segment(&clause[segment_start..i], &mut columns); + segment_start = i + 1; + } + _ => {} + } + } + push_segment(&clause[segment_start..], &mut columns); + if columns.is_empty() { + return vec![UNPARSEABLE.to_string()]; + } + columns +} + +fn classify_sql(literal: &str) -> Vec { + let sql = strip_sql_comments(literal); + let mut found = Vec::new(); + for capture in INSERT_HEAD.captures_iter(&sql) { + found.push(Dml::Insert { + table: table_from_capture(&capture[1]), + }); + } + for capture in DELETE_HEAD.captures_iter(&sql) { + found.push(Dml::Delete { + table: table_from_capture(&capture[1]), + }); + } + for capture in UPDATE_HEAD.captures_iter(&sql) { + let whole = capture.get(0).expect("capture 0"); + found.push(Dml::Update { + table: table_from_capture(&capture[1]), + columns: parse_set_columns(&sql[whole.end()..]), + }); + } + found +} + +#[derive(Default)] +struct ProductionCollector { + literals: Vec, + ident_text: String, + macro_violations: Vec, +} + +impl ProductionCollector { + fn record_macro_tokens(&mut self, macro_name: &str, tokens: &str) { + self.ident_text.push_str(tokens); + self.ident_text.push(' '); + + if macro_name == "include" { + self.macro_violations + .push("include! smuggles unscanned production code".to_string()); + } + + let mut pieces = Vec::new(); + let mut stripped = String::with_capacity(tokens.len()); + let mut cursor = 0; + for matched in MACRO_STRING.find_iter(tokens) { + stripped.push_str(&tokens[cursor..matched.start()]); + cursor = matched.end(); + let raw = matched.as_str(); + // 经 syn 解码转义(覆盖 \xNN、\u{});r"..." 直接取内容。 + let value = syn::parse_str::(raw) + .map(|lit| lit.value()) + .unwrap_or_else(|_| { + raw.trim_start_matches("r\"") + .trim_start_matches('"') + .trim_end_matches('"') + .to_string() + }); + self.literals.push(value.clone()); + pieces.push(value); + } + stripped.push_str(&tokens[cursor..]); + // include_str!/include_bytes! 的 .sql 判定作用于原 token 与拼接体 + // (覆盖 concat!("query.", "sql") 拆分),大小写不敏感。 + let joined_pieces = pieces.join(""); + if matches!(macro_name, "include_str" | "include_bytes") + && (tokens.to_ascii_lowercase().contains(".sql") + || joined_pieces.to_ascii_lowercase().contains(".sql")) + { + self.macro_violations + .push(format!("{macro_name}! loads external SQL")); + } + // 纯字面量:剥离字符串后仅剩标点,且字面量内无 format 插值花括号 + // (`format!("... SET {col} = ...")` 的隐式捕获只有一个字符串 token, + // 必须按非纯字面量处理——R9 终审绕过)。 + let pure_literal = stripped + .chars() + .all(|c| c.is_whitespace() || c == ',' || c == '(' || c == ')') + && !pieces.iter().any(|piece| piece.contains('{')); + + if pure_literal { + // concat! 相邻拼接:拼接体整体参与常规分类。 + if pieces.len() > 1 { + self.literals.push(pieces.join("")); + } + } else { + // 含非字面量 token 的宏(format!/stringify! 构造):一旦出现 + // provider DML 锚点即 fail-closed,不猜插值后的语义。 + let joined = pieces.join(""); + if contains_provider_dml_anchor(tokens) + || contains_provider_dml_anchor(&joined) + || pieces.iter().any(|p| contains_provider_dml_anchor(p)) + { + self.macro_violations.push(format!( + "{macro_name}! builds provider DML from non-literal tokens" + )); + } + } + } +} + +impl<'ast> Visit<'ast> for ProductionCollector { + fn visit_item(&mut self, item: &'ast Item) { + let attrs = match item { + Item::Const(item) => Some(&item.attrs), + Item::Enum(item) => Some(&item.attrs), + Item::Fn(item) => Some(&item.attrs), + Item::Impl(item) => Some(&item.attrs), + Item::Macro(item) => Some(&item.attrs), + Item::Mod(item) => Some(&item.attrs), + Item::Static(item) => Some(&item.attrs), + Item::Struct(item) => Some(&item.attrs), + Item::Trait(item) => Some(&item.attrs), + Item::Type(item) => Some(&item.attrs), + Item::Union(item) => Some(&item.attrs), + Item::Use(item) => Some(&item.attrs), + _ => None, + }; + if attrs.is_some_and(|attrs| attrs_mark_test_only(attrs)) { + return; + } + visit::visit_item(self, item); + } + + fn visit_impl_item(&mut self, item: &'ast ImplItem) { + let attrs = match item { + ImplItem::Const(item) => Some(&item.attrs), + ImplItem::Fn(item) => Some(&item.attrs), + ImplItem::Type(item) => Some(&item.attrs), + ImplItem::Macro(item) => Some(&item.attrs), + _ => None, + }; + if attrs.is_some_and(|attrs| attrs_mark_test_only(attrs)) { + return; + } + visit::visit_impl_item(self, item); + } + + fn visit_expr_lit(&mut self, expression: &'ast ExprLit) { + if let Lit::Str(literal) = &expression.lit { + self.literals.push(literal.value()); + } + visit::visit_expr_lit(self, expression); + } + + fn visit_ident(&mut self, identifier: &'ast syn::Ident) { + self.ident_text.push_str(&identifier.to_string()); + self.ident_text.push(' '); + } + + // syn 默认不遍历宏 token:自行提取字符串与符号。 + fn visit_macro(&mut self, mac: &'ast syn::Macro) { + let name = mac + .path + .segments + .last() + .map(|segment| segment.ident.to_string()) + .unwrap_or_default(); + self.record_macro_tokens(&name, &mac.tokens.to_string()); + } +} + +fn collect_production(source: &str) -> Result { + let syntax = syn::parse_file(source).map_err(|error| error.to_string())?; + if attrs_mark_test_only(&syntax.attrs) { + return Ok(ProductionCollector::default()); + } + let mut collector = ProductionCollector::default(); + collector.visit_file(&syntax); + Ok(collector) +} + +fn is_test_convention_file(relative: &str) -> bool { + let stem = Path::new(relative) + .file_stem() + .and_then(|s| s.to_str()) + .unwrap_or(""); + stem == "tests" || stem.ends_with("_tests") || stem.ends_with("_certification") +} + +fn is_infra_file(relative: &str) -> bool { + INFRA_FILES.contains(&relative) +} + +/// authority 判定使用精确相对路径,杜绝 `ends_with` 伪路径冒充。 +fn dml_allowed(relative: &str, dml: &Dml) -> bool { + match relative { + "database/dao/provider_write.rs" => true, + "database/dao/providers.rs" => match dml { + Dml::Delete { table } => table == "providers", + Dml::Update { table, columns } => { + table == "providers" + && !columns.is_empty() + && columns + .iter() + .all(|c| STATE_COLUMNS_PROVIDERS_RS.contains(&c.as_str())) + } + Dml::Insert { .. } => false, + }, + "database/dao/failover.rs" => match dml { + Dml::Update { table, columns } => { + table == "providers" + && !columns.is_empty() + && columns + .iter() + .all(|c| STATE_COLUMNS_FAILOVER_RS.contains(&c.as_str())) + } + _ => false, + }, + _ => false, + } +} + +#[test] +fn certify_provider_dml_column_authority() { + let root = source_root(); + let mut files = Vec::new(); + collect_rs_files(&root, &mut files); + assert!( + files.len() > 100, + "scanner must see the full source tree, found only {} files", + files.len() + ); + let mut violations = Vec::new(); + for file in &files { + let relative = relative_source_path(&root, file); + if is_infra_file(&relative) { + continue; + } + let source = fs::read_to_string(file).expect("read source file"); + let collector = match collect_production(&source) { + Ok(collector) => collector, + Err(error) => { + violations.push(format!("{relative}: syn parse error: {error}")); + continue; + } + }; + for literal in &collector.literals { + for dml in classify_sql(literal) { + if !dml_allowed(&relative, &dml) { + violations.push(format!("{relative}: {dml:?}")); + } + } + } + for violation in &collector.macro_violations { + violations.push(format!("{relative}: {violation}")); + } + } + assert!( + violations.is_empty(), + "provider DML outside the column-granular authority table:\n{}", + violations.join("\n") + ); +} + +#[test] +fn certify_forbidden_symbols_are_zero_treewide() { + // update_provider_settings_config 是 R4 认定的绕面 mutator:目标是符号 + // 全树归零(定义与调用点一并消失),不是把它搬进写面文件让扫描器沉默。 + // 宏 token 中的出现同样命中(macro_rules 隐藏)。 + let root = source_root(); + let mut files = Vec::new(); + collect_rs_files(&root, &mut files); + let mut hits = Vec::new(); + for file in &files { + let relative = relative_source_path(&root, file); + let source = fs::read_to_string(file).expect("read source file"); + if collect_production(&source) + .is_ok_and(|collector| FORBIDDEN_SYMBOL.is_match(&collector.ident_text)) + { + hits.push(relative.clone()); + } + } + assert!( + hits.is_empty(), + "forbidden mutator symbol still present in: {hits:?}" + ); +} + +#[test] +fn certify_update_dto_has_no_created_at() { + // 裁决 1:created_at 不可变,update DTO 不得携带该字段。 + let root = source_root(); + let source = fs::read_to_string(root.join("database/dao/provider_write.rs")) + .expect("read provider_write.rs"); + let syntax = syn::parse_file(&source).expect("parse provider_write.rs"); + for item in &syntax.items { + let Item::Struct(item_struct) = item else { + continue; + }; + if item_struct.ident != "ProviderRowUpdate" { + continue; + } + let has_created_at = item_struct.fields.iter().any(|field| { + field + .ident + .as_ref() + .is_some_and(|ident| ident == "created_at") + }); + assert!( + !has_created_at, + "ProviderRowUpdate must not carry created_at (immutability ruling)" + ); + return; + } + panic!("ProviderRowUpdate struct not found in provider_write.rs"); +} + +fn find_fn<'a>(items: &'a [Item], name: &str) -> Option<&'a syn::ItemFn> { + for item in items { + match item { + Item::Fn(function) if function.sig.ident == name => return Some(function), + Item::Mod(item_mod) => { + let nested = item_mod.content.as_ref().map(|(_, items)| items.as_slice()); + if let Some(found) = nested.and_then(|items| find_fn(items, name)) { + return Some(found); + } + } + _ => {} + } + } + None +} + +#[derive(Default)] +struct IdentProbe { + found: BTreeSet, +} + +impl<'ast> Visit<'ast> for IdentProbe { + fn visit_ident(&mut self, identifier: &'ast syn::Ident) { + self.found.insert(identifier.to_string()); + } +} + +#[test] +fn certify_reconcile_precondition_enum_shape() { + fn find_enum<'a>(items: &'a [Item], name: &str) -> Option<&'a syn::ItemEnum> { + for item in items { + match item { + Item::Enum(item_enum) if item_enum.ident == name => return Some(item_enum), + Item::Mod(item_mod) => { + let nested = item_mod.content.as_ref().map(|(_, items)| items.as_slice()); + if let Some(found) = nested.and_then(|items| find_enum(items, name)) { + return Some(found); + } + } + _ => {} + } + } + None + } + let root = source_root(); + let source = fs::read_to_string(root.join("services/provider/mod.rs")) + .expect("read services/provider/mod.rs"); + let syntax = syn::parse_file(&source).expect("parse services/provider/mod.rs"); + let precondition = + find_enum(&syntax.items, "ReconcilePrecondition").expect("ReconcilePrecondition exists"); + let variants: Vec = precondition + .variants + .iter() + .map(|variant| variant.ident.to_string()) + .collect(); + assert_eq!( + variants, + vec!["ExpectAbsent".to_string(), "ExpectPresent".to_string()], + "ReconcilePrecondition variants drifted from the adjudicated contract" + ); + let expect_present = precondition + .variants + .iter() + .find(|variant| variant.ident == "ExpectPresent") + .expect("ExpectPresent variant"); + let field_names: Vec = expect_present + .fields + .iter() + .filter_map(|field| field.ident.as_ref().map(|ident| ident.to_string())) + .collect(); + assert_eq!( + field_names, + vec!["fingerprint".to_string()], + "ExpectPresent must carry exactly a fingerprint field" + ); +} + +#[test] +fn certify_reconcile_uses_single_transaction_primitives() { + // 裁决 3:reconcile 不得在函数体内内联读取 aggregate 再分支(那是 + // check-then-act 的新形态);必须委托单事务 DAO 原语。脚手架当前委托旧 + // 函数 → 本测试红,实现方按裁决实现后转绿。 + let root = source_root(); + let source = fs::read_to_string(root.join("services/provider/mod.rs")) + .expect("read services/provider/mod.rs"); + let syntax = syn::parse_file(&source).expect("parse services/provider/mod.rs"); + let function = find_fn(&syntax.items, "reconcile_provider_record_with_precondition") + .expect("reconcile_provider_record_with_precondition exists"); + let mut probe = IdentProbe::default(); + probe.visit_block(&function.block); + for banned in [ + "get_provider_aggregate", + "get_provider_by_id", + "get_all_providers", + "get_all_provider_aggregates", + "reconcile_provider_record", + ] { + assert!( + !probe.found.contains(banned), + "reconcile body must not use '{banned}'; delegate to single-transaction DAO primitives" + ); + } + for required in ["create_provider", "update_provider_if_content_fingerprint"] { + assert!( + probe.found.contains(required), + "reconcile body must delegate to '{required}'" + ); + } +} + +#[test] +fn certify_test_convention_files_are_cfg_test_gated() { + // 命名约定的测试文件必须自带文件级 #![cfg(test)]:借用他处注册的伪测试名 + // 生产文件在此失败;带该属性的文件在任何构建里都不进入生产目标。 + let root = source_root(); + let mut files = Vec::new(); + collect_rs_files(&root, &mut files); + for file in &files { + let relative = relative_source_path(&root, file); + if !is_test_convention_file(&relative) { + continue; + } + let source = fs::read_to_string(file).expect("read source file"); + let syntax = syn::parse_file(&source).expect("parse convention file"); + assert!( + attrs_mark_test_only(&syntax.attrs), + "{relative} uses a test naming convention but lacks a file-level #![cfg(test)]" + ); + } +} + +#[test] +fn certify_scanner_negative_matrix() { + let content = |sql: &str| -> bool { + classify_sql(sql).iter().any(|dml| match dml { + Dml::Insert { table } | Dml::Delete { table } => table == "providers", + Dml::Update { table, columns } => { + table == "providers" + && columns.iter().any(|c| { + c == "settings_config" || c == "name" || c == "meta" || c == UNPARSEABLE + }) + } + }) + }; + // 大小写 + assert!(content( + "update providers set settings_config = ?1 where id = ?2" + )); + // 引号/反引号/方括号表名(R8:引号表名曾在 \b 上失配) + assert!(content(r#"UPDATE "providers" SET name = ?1 WHERE id = ?2"#)); + assert!(content("UPDATE `providers` SET name = ?1 WHERE id = ?2")); + assert!(content("UPDATE [providers] SET name = ?1 WHERE id = ?2")); + assert!(content(r#"DELETE FROM "providers" WHERE id = ?1"#)); + assert!(content(r#"INSERT INTO "providers" (id) VALUES (?1)"#)); + // schema 前缀(含引号 schema) + assert!(content("UPDATE main.providers SET meta = ?1 WHERE id = ?2")); + assert!(content( + r#"UPDATE "main".providers SET meta = ?1 WHERE id = ?2"# + )); + // 别名:AS、AS 带引号(含非 \w 字符)、裸别名、INDEXED BY / NOT INDEXED + assert!(content( + "UPDATE providers AS p SET settings_config = ?1 WHERE p.id = ?2" + )); + assert!(content( + r#"UPDATE providers AS "p-x" SET name = ?1 WHERE id = ?2"# + )); + assert!(content( + r#"UPDATE providers AS "p""x" SET name = ?1 WHERE id = ?2"# + )); + assert!(content( + r#"UPDATE providers "bare-alias" SET name = ?1 WHERE id = ?2"# + )); + assert!(content( + "UPDATE providers p SET p.settings_config = ?1 WHERE p.id = ?2" + )); + assert!(content( + "UPDATE providers INDEXED BY idx SET name = ?1 WHERE id = ?2" + )); + assert!(content( + r#"UPDATE providers INDEXED BY "i-1" SET name = ?1 WHERE id = ?2"# + )); + assert!(content( + "UPDATE providers NOT INDEXED SET name = ?1 WHERE id = ?2" + )); + // providers_seed 等相邻表名不得误报 + assert!(!content("INSERT INTO providers_seed (id) VALUES (?1)")); + assert!(!content( + "UPDATE universal_providers SET name = ?1 WHERE id = ?2" + )); + // OR 冲突子句与 REPLACE INTO + assert!(content( + "UPDATE OR REPLACE providers SET name = ?1 WHERE id = ?2" + )); + assert!(content( + "REPLACE INTO providers (id, app_type) VALUES (?1, ?2)" + )); + assert!(content("INSERT OR REPLACE INTO providers (id) VALUES (?1)")); + // 注释拆词 + assert!(content( + "UPDATE /* sneak */ providers SET settings_config = ?1 WHERE id = ?2" + )); + assert!(content( + "UPDATE providers -- x\n SET name = ?1 WHERE id = ?2" + )); + // 引号内的注释记号不是注释('--' 字符串吞列绕过) + let quoted_comment = "UPDATE providers SET is_current = '--', name = ?1 WHERE id = ?2"; + let classified = classify_sql(quoted_comment); + assert!( + classified.iter().any(|dml| matches!( + dml, + Dml::Update { columns, .. } if columns.contains(&"name".to_string()) + )), + "comment markers inside SQL strings must not swallow columns: {classified:?}" + ); + // 命名参数 :where 不得截断列解析 + let named_param = "UPDATE providers SET is_current = :where, name = ?1 WHERE id = ?2"; + let classified = classify_sql(named_param); + assert!( + classified.iter().any(|dml| matches!( + dml, + Dml::Update { columns, .. } if columns.contains(&"name".to_string()) + )), + "named parameter :where must not terminate column parsing: {classified:?}" + ); + // 双引号内容破坏深度 → fail-closed 哨兵 + let sabotage = r#"UPDATE providers SET is_current = ")", name = ?1 WHERE id = ?2"#; + assert!(content(sabotage), "quoted parens must not hide columns"); + // 子查询误导:内层 WHERE 不得截断列解析 + let subquery = "UPDATE providers SET is_current = (SELECT max(id) FROM t WHERE y = 1), settings_config = ?1 WHERE id = ?2"; + let classified = classify_sql(subquery); + assert!( + classified.iter().any(|dml| matches!( + dml, + Dml::Update { table, columns } + if table == "providers" + && columns.contains(&"is_current".to_string()) + && columns.contains(&"settings_config".to_string()) + )), + "subquery WHERE must not truncate column parsing: {classified:?}" + ); + // 多语句 + let batch = "UPDATE providers SET is_current = 1 WHERE id = 1; UPDATE providers SET settings_config = 'x' WHERE id = 2"; + assert_eq!(classify_sql(batch).len(), 2); + assert!(content(batch)); + // 不闭合引号 → fail-closed + assert!(content( + "UPDATE providers SET is_current = 'unterminated WHERE id = 1" + )); + // 状态列合法写法必须放行(防过杀) + let state_only = classify_sql("UPDATE providers SET is_current = 0 WHERE app_type = ?1"); + assert!(state_only + .iter() + .all(|dml| dml_allowed("database/dao/providers.rs", dml))); + // 内容列即使在状态 authority 内也必须拦下(R4 逃逸场景) + let escaped = classify_sql("UPDATE providers SET settings_config = ?1 WHERE id = ?2"); + assert!(escaped + .iter() + .any(|dml| !dml_allowed("database/dao/providers.rs", dml))); + // endpoints:touch-only 放行于写面,内容列到处拦 + let touch = classify_sql("UPDATE provider_endpoints SET last_used = ?1 WHERE provider_id = ?2"); + assert!(touch + .iter() + .all(|dml| dml_allowed("database/dao/provider_write.rs", dml))); + let ep_content = classify_sql("UPDATE provider_endpoints SET url = ?1 WHERE provider_id = ?2"); + assert!(ep_content + .iter() + .any(|dml| !dml_allowed("database/dao/providers.rs", dml))); + // 精确路径:伪路径不得冒充 authority + let state = Dml::Update { + table: "providers".to_string(), + columns: vec!["is_current".to_string()], + }; + assert!(dml_allowed("database/dao/providers.rs", &state)); + assert!(!dml_allowed("services/database/dao/providers.rs", &state)); + assert!(!dml_allowed( + "evil/database/dao/provider_write.rs", + &Dml::Insert { + table: "providers".to_string() + } + )); + // cfg 语义:not(test) 与 any(test, unix) 是生产代码 + assert!(!cfg_expr_requires_test("not (test)")); + assert!(!cfg_expr_requires_test("any (test , unix)")); + assert!(cfg_expr_requires_test("test")); + assert!(cfg_expr_requires_test("all (test , unix)")); + assert!(cfg_expr_requires_test("any (test , all (test , unix))")); + // 宏:concat! 纯字面量拼接 → 正常分类,不误报违规 + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens( + "concat", + r#""UPDATE providers " , "SET name = ?1 WHERE id = ?2""#, + ); + assert!( + collector.literals.iter().any(|lit| content(lit)), + "concat!-joined SQL must be classified" + ); + assert!(collector.macro_violations.is_empty()); + // 宏:format! 含非字面量 + DML 锚点 → 无条件违规(即使可见列全是状态列) + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens( + "format", + r#""UPDATE providers SET is_current = 0 , {} = ?1 WHERE id = ?2" , column"#, + ); + assert!( + !collector.macro_violations.is_empty(), + "format!-built provider DML must fail closed even when visible columns look like state" + ); + // 宏:stringify! 式 ident 构造 SQL → token 锚点命中 + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens("stringify", "UPDATE providers SET name = x WHERE id = y"); + assert!( + !collector.macro_violations.is_empty(), + "ident-built provider DML must fail closed" + ); + // 宏:无 DML 锚点的普通 format! 不误报 + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens("format", r#""hello {}" , name"#); + assert!(collector.macro_violations.is_empty()); + // 宏:include! 全禁,include_str!(.sql) 禁,普通资源不误报 + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens("include", r#""../generated.rs""#); + assert!(!collector.macro_violations.is_empty()); + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens("include_str", r#""queries/update.sql""#); + assert!(!collector.macro_violations.is_empty()); + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens("include_bytes", r#""queries/update.SQL""#); + assert!( + !collector.macro_violations.is_empty(), + "include_bytes and case variants must be banned for SQL" + ); + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens("include_str", r#"concat ! ("query." , "sql")"#); + assert!( + !collector.macro_violations.is_empty(), + "extension split via concat! must still be detected" + ); + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens("include_str", r#""resources/template.json""#); + assert!(collector.macro_violations.is_empty()); + // 隐式 format 捕获:单字符串 + 花括号插值不得被当纯字面量放行(R9) + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens( + "format", + r#""UPDATE providers SET {is_current} = ?1 WHERE id = ?2""#, + ); + assert!( + !collector.macro_violations.is_empty(), + "implicit format captures must fail closed" + ); + // 宏字符串转义:\x55(U)解码后仍识别 + let mut collector = ProductionCollector::default(); + collector.record_macro_tokens( + "concat", + r#""\x55PDATE providers " , "SET name = ?1 WHERE id = ?2""#, + ); + assert!( + collector.literals.iter().any(|lit| content(lit)), + "escaped SQL must be decoded via LitStr::value before classification" + ); +} + +// --------------------------------------------------------------------------- +// T3:create 冲突原子性与结构化 Conflict(裁决 2) +// --------------------------------------------------------------------------- + +#[test] +fn certify_duplicate_create_returns_structured_conflict() { + let database = db(); + let first = with_endpoints( + { + let mut input = base_input("dup", "第一次创建"); + input.sort_index = Some(5); + input.in_failover_queue = true; + input + }, + &[("https://a.example", Some(11), None)], + ); + database + .create_provider(NewProviderAggregate::from_input("claude", first).unwrap()) + .expect("first create"); + let before = snapshot(&database, "claude", "dup"); + + let second = with_endpoints( + base_input("dup", "冒名顶替"), + &[("https://b.example", Some(22), None)], + ); + let err = database + .create_provider(NewProviderAggregate::from_input("claude", second).unwrap()) + .expect_err("duplicate create must fail, not upsert"); + assert!( + matches!(err, AppError::Conflict(_)), + "duplicate create must surface a structured Conflict, got: {err:?}" + ); + assert_eq!( + snapshot(&database, "claude", "dup"), + before, + "duplicate create must leave row, endpoints and state untouched" + ); +} + +#[test] +fn certify_create_does_not_touch_current_state() { + let database = db(); + database + .create_provider( + NewProviderAggregate::from_input("claude", base_input("first", "既有")).unwrap(), + ) + .expect("create first"); + database + .set_current_provider("claude", "first") + .expect("set current"); + database + .create_provider( + NewProviderAggregate::from_input("claude", base_input("second", "新建")).unwrap(), + ) + .expect("create second"); + let (first_row, _) = snapshot(&database, "claude", "first"); + let (second_row, _) = snapshot(&database, "claude", "second"); + assert_eq!( + first_row.expect("first row").10, + 1, + "create must not clear another provider's is_current" + ); + assert_eq!( + second_row.expect("second row").10, + 0, + "create must never set is_current on the new row" + ); +} + +#[test] +fn certify_create_endpoint_failure_rolls_back_row() { + let database = db(); + install_endpoint_reject_trigger(&database); + let input = with_endpoints( + base_input("halfway", "半途失败"), + &[("https://blocked.example", Some(1), None)], + ); + database + .create_provider(NewProviderAggregate::from_input("claude", input).unwrap()) + .expect_err("endpoint insert failure must fail the create"); + let (row, endpoints) = snapshot(&database, "claude", "halfway"); + assert!( + row.is_none() && endpoints.is_empty(), + "failed create must leave no partial row" + ); +} + +// --------------------------------------------------------------------------- +// T4:update 严格单行、created_at 不可变、状态列保全、全内容往返 +// --------------------------------------------------------------------------- + +#[test] +fn certify_update_missing_provider_is_notfound_and_creates_nothing() { + let database = db(); + let key = ProviderKey::new("claude", "ghost").unwrap(); + let row = ProviderRowUpdate::from_input(&base_input("ghost", "幽灵")).unwrap(); + let err = database.update_provider(&key, &row).expect_err("must fail"); + assert!(matches!(err, AppError::NotFound(_)), "got: {err:?}"); + let (row_after, endpoints_after) = snapshot(&database, "claude", "ghost"); + assert!(row_after.is_none() && endpoints_after.is_empty()); +} + +#[test] +fn certify_update_cannot_change_created_at() { + let database = db(); + let mut created = base_input("epoch", "创建时间"); + created.created_at = Some(111); + database + .create_provider(NewProviderAggregate::from_input("claude", created).unwrap()) + .expect("create"); + let key = ProviderKey::new("claude", "epoch").unwrap(); + let mut edited = base_input("epoch", "被编辑"); + edited.created_at = Some(222); + database + .update_provider(&key, &ProviderRowUpdate::from_input(&edited).unwrap()) + .expect("update"); + let (row, _) = snapshot(&database, "claude", "epoch"); + assert_eq!( + row.expect("row").4, + Some(111), + "update must never rewrite created_at" + ); +} + +#[test] +fn certify_update_preserves_all_state_columns() { + let database = db(); + let created = { + let mut input = base_input("stately", "状态在身"); + input.sort_index = Some(9); + input.in_failover_queue = true; + input + }; + database + .create_provider(NewProviderAggregate::from_input("claude", created).unwrap()) + .expect("create"); + database + .set_current_provider("claude", "stately") + .expect("set current"); + let key = ProviderKey::new("claude", "stately").unwrap(); + database + .update_provider( + &key, + &ProviderRowUpdate::from_input(&base_input("stately", "改名")).unwrap(), + ) + .expect("update"); + let (row, _) = snapshot(&database, "claude", "stately"); + let row = row.expect("row"); + assert_eq!(row.5, Some(9), "sort_index must survive row update"); + assert_eq!(row.10, 1, "is_current must survive row update"); + assert_eq!(row.11, 1, "in_failover_queue must survive row update"); +} + +#[test] +fn certify_full_content_roundtrip_via_create_and_update() { + // 全列往返认证(含 meta):忽略任一内容列的实现都不得变绿。 + let database = db(); + let mut created = base_input("full", "全字段"); + created.settings_config = json!({"base_url": "https://one.example", "model": "m1"}); + created.website_url = Some("https://site.example".to_string()); + created.category = Some("cat-a".to_string()); + created.notes = Some("初始备注".to_string()); + created.icon = Some("icon-a".to_string()); + created.icon_color = Some("#111111".to_string()); + created.meta = Some(ProviderMeta { + common_config_enabled: Some(true), + ..Default::default() + }); + database + .create_provider(NewProviderAggregate::from_input("claude", created).unwrap()) + .expect("create"); + let (row, _) = snapshot(&database, "claude", "full"); + let row = row.expect("row"); + assert_eq!(row.0, "全字段"); + assert!(row.1.contains("https://one.example")); + assert_eq!(row.2.as_deref(), Some("https://site.example")); + assert_eq!(row.3.as_deref(), Some("cat-a")); + assert_eq!(row.6.as_deref(), Some("初始备注")); + assert_eq!(row.7.as_deref(), Some("icon-a")); + assert_eq!(row.8.as_deref(), Some("#111111")); + let meta_json: serde_json::Value = + serde_json::from_str(&row.9).expect("stored meta must be valid JSON"); + assert_eq!( + meta_json["commonConfigEnabled"], + json!(true), + "meta content must round-trip through create, got: {}", + row.9 + ); + + let key = ProviderKey::new("claude", "full").unwrap(); + let mut edited = base_input("full", "全字段二版"); + edited.settings_config = json!({"base_url": "https://two.example", "model": "m2"}); + edited.website_url = Some("https://site2.example".to_string()); + edited.category = Some("cat-b".to_string()); + edited.notes = Some("二版备注".to_string()); + edited.icon = Some("icon-b".to_string()); + edited.icon_color = Some("#222222".to_string()); + edited.meta = Some(ProviderMeta { + common_config_enabled: Some(false), + ..Default::default() + }); + database + .update_provider(&key, &ProviderRowUpdate::from_input(&edited).unwrap()) + .expect("update"); + let (row, _) = snapshot(&database, "claude", "full"); + let row = row.expect("row"); + assert_eq!(row.0, "全字段二版"); + assert!(row.1.contains("https://two.example")); + assert_eq!(row.2.as_deref(), Some("https://site2.example")); + assert_eq!(row.3.as_deref(), Some("cat-b")); + assert_eq!(row.6.as_deref(), Some("二版备注")); + assert_eq!(row.7.as_deref(), Some("icon-b")); + assert_eq!(row.8.as_deref(), Some("#222222")); + let meta_json: serde_json::Value = + serde_json::from_str(&row.9).expect("stored meta must be valid JSON"); + assert_eq!( + meta_json["commonConfigEnabled"], + json!(false), + "meta content must round-trip through update, got: {}", + row.9 + ); +} + +// --------------------------------------------------------------------------- +// T5:陈旧快照下并发 endpoint 变更存活(R3 核心场景)+ endpoint 严格性 +// --------------------------------------------------------------------------- + +#[test] +fn certify_concurrent_endpoint_changes_survive_row_update() { + let database = db(); + let created = with_endpoints( + base_input("surv", "并发存活"), + &[("https://old.example", Some(1), None)], + ); + database + .create_provider(NewProviderAggregate::from_input("claude", created).unwrap()) + .expect("create"); + let key = ProviderKey::new("claude", "surv").unwrap(); + + database + .add_provider_endpoint( + &key, + NewEndpoint::new("https://new.example", Some(2), None).unwrap(), + ) + .expect("add"); + database + .remove_provider_endpoint(&key, "https://old.example") + .expect("remove"); + database + .touch_provider_endpoint(&key, "https://new.example", 99) + .expect("touch"); + + let row = ProviderRowUpdate::from_input(&base_input("surv", "改名")).unwrap(); + database.update_provider(&key, &row).expect("row update"); + + let (_, endpoints) = snapshot(&database, "claude", "surv"); + assert_eq!( + endpoints, + vec![("https://new.example".to_string(), Some(2), Some(99))], + "all concurrent endpoint mutations must survive a row update" + ); +} + +#[test] +fn certify_update_payload_with_endpoints_is_rejected_explicitly() { + let stale = with_endpoints( + base_input("surv", "夹带"), + &[("https://smuggle.example", Some(3), None)], + ); + let err = ProviderRowUpdate::from_input(&stale).expect_err("must reject"); + assert!(matches!(err, AppError::InvalidInput(_)), "got: {err:?}"); +} + +#[test] +fn certify_endpoint_mutations_are_strict() { + let database = db(); + database + .create_provider( + NewProviderAggregate::from_input( + "claude", + with_endpoints( + base_input("strict", "严格"), + &[("https://one.example", Some(7), None)], + ), + ) + .unwrap(), + ) + .expect("create"); + let key = ProviderKey::new("claude", "strict").unwrap(); + let before = snapshot(&database, "claude", "strict"); + + database + .add_provider_endpoint( + &key, + NewEndpoint::new("https://one.example", Some(8), None).unwrap(), + ) + .expect_err("duplicate endpoint add must fail"); + assert_eq!(snapshot(&database, "claude", "strict"), before); + + assert!(matches!( + database.remove_provider_endpoint(&key, "https://none.example"), + Err(AppError::NotFound(_)) + )); + assert!(matches!( + database.touch_provider_endpoint(&key, "https://none.example", 1), + Err(AppError::NotFound(_)) + )); + + database + .touch_provider_endpoint(&key, "https://one.example", 55) + .expect("touch"); + let (_, endpoints) = snapshot(&database, "claude", "strict"); + assert_eq!( + endpoints, + vec![("https://one.example".to_string(), Some(7), Some(55))], + "touch must change last_used only" + ); +} + +// --------------------------------------------------------------------------- +// T6:added_at NULL 全链路无损 +// --------------------------------------------------------------------------- + +#[test] +fn certify_null_added_at_roundtrips_losslessly() { + let database = db(); + let created = with_endpoints( + base_input("nulls", "空值"), + &[("https://n.example", None, None)], + ); + database + .create_provider(NewProviderAggregate::from_input("claude", created).unwrap()) + .expect("create"); + + let (_, raw) = snapshot(&database, "claude", "nulls"); + assert_eq!( + raw, + vec![("https://n.example".to_string(), None, None)], + "storage must keep NULL, not 0" + ); + + let aggregate = database + .get_provider_aggregate("claude", "nulls") + .expect("hydrate") + .expect("exists"); + let endpoint = aggregate + .endpoints + .get("https://n.example") + .expect("endpoint present in hydration"); + assert_eq!( + endpoint.added_at, None, + "hydration must not coerce NULL added_at to 0" + ); + assert_eq!(endpoint.last_used, None); +} + +// --------------------------------------------------------------------------- +// T7:rename 认证矩阵 +// --------------------------------------------------------------------------- + +fn create_opencode_provider(database: &Database, id: &str) { + let created = with_endpoints( + { + let mut input = base_input(id, "opencode 源"); + input.sort_index = Some(3); + input.in_failover_queue = true; + input + }, + &[("https://keep.example", None, Some(42))], + ); + database + .create_provider(NewProviderAggregate::from_input("opencode", created).unwrap()) + .expect("create opencode provider"); +} + +#[test] +fn certify_rename_preserves_endpoints_nulls_state_and_current() { + let database = db(); + create_opencode_provider(&database, "old-key"); + database + .set_current_provider("opencode", "old-key") + .expect("set current"); + let source = ProviderKey::new("opencode", "old-key").unwrap(); + let rename = + RenameProvider::from_input(source, &base_input("new-key", "改键")).expect("build rename"); + database + .rename_db_only_additive_provider(rename) + .expect("rename"); + + let (old_row, old_eps) = snapshot(&database, "opencode", "old-key"); + assert!( + old_row.is_none() && old_eps.is_empty(), + "source must be gone" + ); + + let (new_row, new_eps) = snapshot(&database, "opencode", "new-key"); + let new_row = new_row.expect("target row"); + assert_eq!(new_row.0, "改键", "row content must come from rename input"); + assert_eq!(new_row.5, Some(3), "sort_index must carry over"); + assert_eq!(new_row.10, 1, "is_current must carry over"); + assert_eq!(new_row.11, 1, "in_failover_queue must carry over"); + assert_eq!( + new_eps, + vec![("https://keep.example".to_string(), None, Some(42))], + "endpoints must carry over with NULL timestamps intact" + ); +} + +#[test] +fn certify_rename_target_conflict_has_zero_side_effects() { + let database = db(); + create_opencode_provider(&database, "src"); + create_opencode_provider(&database, "dst"); + let before_src = snapshot(&database, "opencode", "src"); + let before_dst = snapshot(&database, "opencode", "dst"); + + let source = ProviderKey::new("opencode", "src").unwrap(); + let rename = RenameProvider::from_input(source, &base_input("dst", "撞车")).expect("build"); + database + .rename_db_only_additive_provider(rename) + .expect_err("rename onto an existing key must fail"); + + assert_eq!(snapshot(&database, "opencode", "src"), before_src); + assert_eq!(snapshot(&database, "opencode", "dst"), before_dst); +} + +#[test] +fn certify_rename_endpoint_copy_failure_is_atomic() { + let database = db(); + create_opencode_provider(&database, "guarded"); + let before = snapshot(&database, "opencode", "guarded"); + install_endpoint_reject_trigger(&database); + + let source = ProviderKey::new("opencode", "guarded").unwrap(); + let rename = RenameProvider::from_input(source, &base_input("moved", "搬家")).expect("build"); + database + .rename_db_only_additive_provider(rename) + .expect_err("endpoint copy failure must fail the rename"); + + assert_eq!( + snapshot(&database, "opencode", "guarded"), + before, + "failed rename must leave the source fully intact" + ); + let (moved_row, moved_eps) = snapshot(&database, "opencode", "moved"); + assert!( + moved_row.is_none() && moved_eps.is_empty(), + "failed rename must leave no partial target" + ); +} + +#[test] +fn certify_rename_scope_restrictions() { + let claude_source = ProviderKey::new("claude", "any").unwrap(); + assert!(matches!( + RenameProvider::from_input(claude_source, &base_input("other", "x")), + Err(AppError::InvalidInput(_)) + )); + + let database = db(); + for category in ["omo", "omo-slim"] { + let id = format!("omo-{category}"); + let mut omo = base_input(&id, "omo"); + omo.category = Some(category.to_string()); + database + .create_provider(NewProviderAggregate::from_input("opencode", omo).unwrap()) + .expect("create omo provider"); + let source = ProviderKey::new("opencode", &id).unwrap(); + let rename = + RenameProvider::from_input(source, &base_input("omo-target", "y")).expect("build"); + assert!( + matches!( + database.rename_db_only_additive_provider(rename), + Err(AppError::InvalidInput(_)) + ), + "{category} providers must not be renamable" + ); + } + + let ghost = ProviderKey::new("opencode", "ghost").unwrap(); + let rename = RenameProvider::from_input(ghost, &base_input("anywhere", "z")).expect("build"); + assert!(matches!( + database.rename_db_only_additive_provider(rename), + Err(AppError::NotFound(_)) + )); +} + +// --------------------------------------------------------------------------- +// D1:delete 补偿原语必须能重建完整 aggregate +// --------------------------------------------------------------------------- + +#[test] +fn certify_delete_compensation_recreates_exact_aggregate() { + let database = db(); + let created = with_endpoints( + { + let mut input = base_input("comp", "补偿对象"); + input.sort_index = Some(5); + input.in_failover_queue = true; + input.notes = Some("完整字段".to_string()); + input + }, + &[ + ("https://a.example", Some(11), Some(20)), + ("https://b.example", None, None), + ], + ); + database + .create_provider(NewProviderAggregate::from_input("claude", created).unwrap()) + .expect("create"); + database + .set_current_provider("claude", "comp") + .expect("set current"); + let before = snapshot(&database, "claude", "comp"); + + database + .delete_provider("claude", "comp") + .expect("delete provider"); + let (gone_row, gone_eps) = snapshot(&database, "claude", "comp"); + assert!( + gone_row.is_none() && gone_eps.is_empty(), + "delete must cascade endpoints" + ); + + // 补偿:必须能从快照原样重建已删除的 aggregate(update-first 语义在此 + // 必败——补偿原语要求 insert-or-restore 语义)。created_at 由专属参数 + // 携带(裁决 1)。 + let key = ProviderKey::new("claude", "comp").unwrap(); + let row = ProviderRowUpdate::from_input(&{ + let mut input = base_input("comp", "补偿对象"); + input.notes = Some("完整字段".to_string()); + input + }) + .unwrap(); + let endpoints = [ + NewEndpoint::new("https://a.example", Some(11), Some(20)).unwrap(), + NewEndpoint::new("https://b.example", None, None).unwrap(), + ]; + { + let mut conn = database.conn.lock().expect("lock certification database"); + let tx = conn.transaction().expect("open compensation transaction"); + provider_write::restore_provider_aggregate_on_tx( + &tx, + &key, + &row, + Some(1_700_000_000), + Some(5), + true, + true, + &endpoints, + ) + .expect("compensation must recreate a deleted aggregate"); + tx.commit().expect("commit compensation"); + } + assert_eq!( + snapshot(&database, "claude", "comp"), + before, + "restored aggregate must be byte-identical to the pre-delete snapshot" + ); +} + +#[test] +fn certify_delete_compensation_failure_leaves_no_partial_state() { + let database = db(); + let created = with_endpoints( + base_input("comp2", "补偿失败"), + &[("https://c.example", Some(1), None)], + ); + database + .create_provider(NewProviderAggregate::from_input("claude", created).unwrap()) + .expect("create"); + database + .delete_provider("claude", "comp2") + .expect("delete provider"); + install_endpoint_reject_trigger(&database); + + let key = ProviderKey::new("claude", "comp2").unwrap(); + let row = ProviderRowUpdate::from_input(&base_input("comp2", "补偿失败")).unwrap(); + let endpoints = [NewEndpoint::new("https://c.example", Some(1), None).unwrap()]; + let err = { + let mut conn = database.conn.lock().expect("lock certification database"); + let tx = conn.transaction().expect("open compensation transaction"); + provider_write::restore_provider_aggregate_on_tx( + &tx, + &key, + &row, + Some(1_700_000_000), + None, + false, + false, + &endpoints, + ) + .expect_err("endpoint restore failure must fail the compensation") + // 事务随 drop 回滚 + }; + assert!( + err.to_string().contains(ENDPOINT_REJECT_MESSAGE), + "compensation must fail at the injected endpoint restore, not before it; got: {err}" + ); + let (row_after, eps_after) = snapshot(&database, "claude", "comp2"); + assert!( + row_after.is_none() && eps_after.is_empty(), + "failed compensation must not leave a row without its endpoints" + ); +} + +// --------------------------------------------------------------------------- +// T8/T9:reconcile 显式前置期望 +// --------------------------------------------------------------------------- + +#[test] +fn certify_reconcile_expect_present_preserves_endpoints_and_state() { + let database = db(); + let created = with_endpoints( + { + let mut input = base_input("recon", "用户创建"); + input.sort_index = Some(7); + input.in_failover_queue = true; + input + }, + &[("https://user.example", Some(5), None)], + ); + database + .create_provider(NewProviderAggregate::from_input("claude", created).unwrap()) + .expect("create"); + database + .set_current_provider("claude", "recon") + .expect("set current"); + let aggregate = database + .get_provider_aggregate("claude", "recon") + .expect("hydrate") + .expect("exists"); + let fingerprint = provider_row_fingerprint(&aggregate.provider); + + reconcile_provider_record_with_precondition( + &database, + "claude", + base_input("recon", "同步覆盖"), + ReconcilePrecondition::ExpectPresent { fingerprint }, + ) + .expect("reconcile existing with fresh fingerprint"); + + let (row, endpoints) = snapshot(&database, "claude", "recon"); + let row = row.expect("row"); + assert_eq!(row.0, "同步覆盖", "row content may be reconciled"); + assert_eq!( + row.5, + Some(7), + "sort_index is state, reconcile must not clear it" + ); + assert_eq!( + row.10, 1, + "is_current is state, reconcile must not clear it" + ); + assert_eq!( + row.11, 1, + "failover membership is state, reconcile must not clear it" + ); + assert_eq!( + endpoints, + vec![("https://user.example".to_string(), Some(5), None)], + "reconcile of an existing provider must never touch endpoints" + ); +} + +#[test] +fn certify_reconcile_expect_absent_creates_with_initial_endpoints() { + let database = db(); + reconcile_provider_record_with_precondition( + &database, + "claude", + with_endpoints( + base_input("fresh", "同步新建"), + &[("https://seed.example", Some(9), None)], + ), + ReconcilePrecondition::ExpectAbsent, + ) + .expect("reconcile missing"); + let (row, endpoints) = snapshot(&database, "claude", "fresh"); + assert!(row.is_some()); + assert_eq!( + endpoints, + vec![("https://seed.example".to_string(), Some(9), None)] + ); +} + +#[test] +fn certify_reconcile_expect_absent_loser_cannot_overwrite_winner() { + // T9(TOCTOU 本体):观察为 Absent 后输掉竞争,必须结构化 Conflict, + // 绝不退化为覆盖更新。 + let database = db(); + let winner = with_endpoints( + base_input("race-slot", "竞争赢家"), + &[("https://winner.example", Some(1), None)], + ); + database + .create_provider(NewProviderAggregate::from_input("claude", winner).unwrap()) + .expect("winner create"); + let before = snapshot(&database, "claude", "race-slot"); + + let err = reconcile_provider_record_with_precondition( + &database, + "claude", + base_input("race-slot", "迟到输家"), + ReconcilePrecondition::ExpectAbsent, + ) + .expect_err("losing an ExpectAbsent race must surface an error"); + assert!( + matches!(err, AppError::Conflict(_)), + "race loser must get a structured Conflict, got: {err:?}" + ); + assert_eq!( + snapshot(&database, "claude", "race-slot"), + before, + "the winner's row must remain byte-identical" + ); +} + +#[test] +fn certify_reconcile_expect_present_stale_fingerprint_conflicts() { + let database = db(); + database + .create_provider( + NewProviderAggregate::from_input("claude", base_input("staleful", "第一版")).unwrap(), + ) + .expect("create"); + let aggregate = database + .get_provider_aggregate("claude", "staleful") + .expect("hydrate") + .expect("exists"); + let stale_fingerprint = provider_row_fingerprint(&aggregate.provider); + + // 其他写者更新了行内容,持旧指纹的 reconcile 必须 Conflict。 + let key = ProviderKey::new("claude", "staleful").unwrap(); + database + .update_provider( + &key, + &ProviderRowUpdate::from_input(&base_input("staleful", "第二版")).unwrap(), + ) + .expect("interleaved update"); + + let err = reconcile_provider_record_with_precondition( + &database, + "claude", + base_input("staleful", "第三版"), + ReconcilePrecondition::ExpectPresent { + fingerprint: stale_fingerprint, + }, + ) + .expect_err("stale fingerprint must surface an error"); + assert!( + matches!(err, AppError::Conflict(_)), + "stale fingerprint must get a structured Conflict, got: {err:?}" + ); + let (row, _) = snapshot(&database, "claude", "staleful"); + assert_eq!( + row.expect("row").0, + "第二版", + "stale reconcile must not overwrite the interleaved writer" + ); +} + +#[test] +fn certify_fingerprint_is_deterministic_and_endpoint_blind() { + // preserve_order + HashMap 意味着朴素序列化指纹不稳定(伪 Conflict); + // 指纹必须走规范化排序哈希,且不受 endpoint 填充差异影响。 + let database = db(); + database + .create_provider( + NewProviderAggregate::from_input( + "claude", + with_endpoints( + { + let mut input = base_input("fp", "指纹"); + input.meta = Some(ProviderMeta { + common_config_enabled: Some(true), + ..Default::default() + }); + input + }, + &[("https://e.example", Some(1), None)], + ), + ) + .unwrap(), + ) + .expect("create"); + let via_aggregate = database + .get_provider_aggregate("claude", "fp") + .expect("hydrate") + .expect("exists"); + let fp1 = provider_row_fingerprint(&via_aggregate.provider); + let fp2 = provider_row_fingerprint(&via_aggregate.provider); + assert_eq!(fp1, fp2, "fingerprint must be deterministic"); + + // endpoint 填充差异(get_provider_by_id 会把 endpoints 合回 meta)不得 + // 改变指纹。 + let mut with_endpoints_in_meta = via_aggregate.provider.clone(); + let mut meta = with_endpoints_in_meta.meta.take().unwrap_or_default(); + meta.custom_endpoints.insert( + "https://e.example".to_string(), + CustomEndpoint { + url: "https://e.example".to_string(), + added_at: Some(1), + last_used: None, + }, + ); + with_endpoints_in_meta.meta = Some(meta); + assert_eq!( + fp1, + provider_row_fingerprint(&with_endpoints_in_meta), + "endpoint hydration differences must not change the content fingerprint" + ); + + // preserve_order 下键插入顺序不同但逻辑相等的对象必须同指纹 + // (旧的朴素序列化对同一实例稳定,骗得过"哈希两次"断言,骗不过这个)。 + let mut ordered_a = via_aggregate.provider.clone(); + ordered_a.settings_config = + serde_json::from_str(r#"{"alpha": 1, "zeta": {"x": 1, "y": 2}}"#).unwrap(); + let mut ordered_b = via_aggregate.provider.clone(); + ordered_b.settings_config = + serde_json::from_str(r#"{"zeta": {"y": 2, "x": 1}, "alpha": 1}"#).unwrap(); + assert_eq!( + provider_row_fingerprint(&ordered_a), + provider_row_fingerprint(&ordered_b), + "logically equal objects with different key insertion order must share a fingerprint" + ); + + // 长度前缀:边界粘连的不同内容必须得到不同指纹(碰撞对)。 + let mut collide_a = via_aggregate.provider.clone(); + collide_a.settings_config = json!(["a", "b"]); + let mut collide_b = via_aggregate.provider.clone(); + collide_b.settings_config = json!(["a\u{0}sb"]); + assert_ne!( + provider_row_fingerprint(&collide_a), + provider_row_fingerprint(&collide_b), + "canonical encoding must be collision-free across value boundaries" + ); +} + +// --------------------------------------------------------------------------- +// T11:并发线性化 +// --------------------------------------------------------------------------- + +#[test] +fn certify_concurrent_create_single_winner() { + let database = db(); + let barrier = std::sync::Barrier::new(2); + let contenders = [ + ("赢家甲", "https://alpha.example"), + ("赢家乙", "https://beta.example"), + ]; + let results: Vec> = std::thread::scope(|scope| { + contenders + .iter() + .map(|(name, url)| { + let database = &database; + let barrier = &barrier; + scope.spawn(move || { + let input = with_endpoints(base_input("race", name), &[(url, Some(1), None)]); + let aggregate = NewProviderAggregate::from_input("claude", input).unwrap(); + barrier.wait(); + database.create_provider(aggregate).map(|_| *url) + }) + }) + .collect::>() + .into_iter() + .map(|handle| handle.join().expect("thread join")) + .collect() + }); + let winners: Vec<&str> = results + .iter() + .filter_map(|r| r.as_ref().ok().copied()) + .collect(); + assert_eq!(winners.len(), 1, "exactly one concurrent create must win"); + let (row, endpoints) = snapshot(&database, "claude", "race"); + let row = row.expect("winner row"); + let winner_url = winners[0]; + let winner_name = contenders + .iter() + .find(|(_, url)| *url == winner_url) + .map(|(name, _)| *name) + .expect("winner name"); + assert_eq!(row.0, winner_name, "row must belong entirely to the winner"); + assert_eq!( + endpoints, + vec![(winner_url.to_string(), Some(1), None)], + "endpoints must belong entirely to the same winner" + ); +} + +#[test] +fn certify_concurrent_full_updates_do_not_tear() { + let database = db(); + database + .create_provider( + NewProviderAggregate::from_input("claude", base_input("tear", "初始")).unwrap(), + ) + .expect("create"); + let barrier = std::sync::Barrier::new(2); + std::thread::scope(|scope| { + for suffix in ["一号", "二号"] { + let database = &database; + let barrier = &barrier; + scope.spawn(move || { + let mut input = base_input("tear", &format!("名-{suffix}")); + input.website_url = Some(format!("https://site-{suffix}.example")); + input.notes = Some(format!("注-{suffix}")); + input.icon = Some(format!("icon-{suffix}")); + let key = ProviderKey::new("claude", "tear").unwrap(); + let row = ProviderRowUpdate::from_input(&input).unwrap(); + barrier.wait(); + database.update_provider(&key, &row).expect("update"); + }); + } + }); + let (row, _) = snapshot(&database, "claude", "tear"); + let row = row.expect("row"); + let suffix = row.0.strip_prefix("名-").expect("name written by a writer"); + assert_eq!( + row.2.as_deref(), + Some(format!("https://site-{suffix}.example").as_str()), + "row content must come from a single writer, not interleaved" + ); + assert_eq!(row.6.as_deref(), Some(format!("注-{suffix}").as_str())); + assert_eq!(row.7.as_deref(), Some(format!("icon-{suffix}").as_str())); +} + +#[test] +fn certify_concurrent_endpoint_interleaving_is_consistent() { + let database = db(); + database + .create_provider( + NewProviderAggregate::from_input( + "claude", + with_endpoints( + base_input("weave", "交错"), + &[("https://c.example", Some(1), None)], + ), + ) + .unwrap(), + ) + .expect("create"); + let key = ProviderKey::new("claude", "weave").unwrap(); + let barrier = std::sync::Barrier::new(2); + std::thread::scope(|scope| { + { + let database = &database; + let key = &key; + let barrier = &barrier; + scope.spawn(move || { + barrier.wait(); + database + .add_provider_endpoint( + key, + NewEndpoint::new("https://a.example", Some(2), None).unwrap(), + ) + .expect("add a"); + database + .touch_provider_endpoint(key, "https://a.example", 7) + .expect("touch a"); + }); + } + { + let database = &database; + let key = &key; + let barrier = &barrier; + scope.spawn(move || { + barrier.wait(); + database + .add_provider_endpoint( + key, + NewEndpoint::new("https://b.example", Some(3), None).unwrap(), + ) + .expect("add b"); + database + .remove_provider_endpoint(key, "https://c.example") + .expect("remove c"); + }); + } + }); + let (_, endpoints) = snapshot(&database, "claude", "weave"); + assert_eq!( + endpoints, + vec![ + ("https://a.example".to_string(), Some(2), Some(7)), + ("https://b.example".to_string(), Some(3), None), + ], + "interleaved endpoint operations must all land exactly once" + ); +} + +// --------------------------------------------------------------------------- +// T10:服务入口认证绑定(syn 级:真实 #[test] 函数且真的触达 ProviderService) +// --------------------------------------------------------------------------- + +#[test] +fn certify_service_entry_tests_present() { + let source = fs::read_to_string(source_root().join("services/provider/mod.rs")) + .expect("read services/provider/mod.rs"); + let syntax = syn::parse_file(&source).expect("parse services/provider/mod.rs"); + for required in [ + "provider_service_create_owns_initial_endpoints_and_duplicate_is_atomic", + "provider_service_stale_edit_payload_cannot_overwrite_endpoint_operations", + "provider_service_db_only_rename_matrix_is_atomic_and_lossless", + ] { + let function = find_fn(&syntax.items, required) + .unwrap_or_else(|| panic!("bound service-entry test '{required}' is missing")); + assert!( + function + .attrs + .iter() + .any(|attribute| attribute.path().is_ident("test")), + "'{required}' must be a #[test] function" + ); + let mut probe = IdentProbe::default(); + probe.visit_block(&function.block); + assert!( + probe.found.contains("ProviderService"), + "'{required}' must exercise ProviderService (empty stubs cannot pass)" + ); + } +} diff --git a/src-tauri/src/database/dao/providers.rs b/src-tauri/src/database/dao/providers.rs index f64154b06..1d094348b 100644 --- a/src-tauri/src/database/dao/providers.rs +++ b/src-tauri/src/database/dao/providers.rs @@ -1,111 +1,212 @@ -use crate::database::{lock_conn, Database}; +use crate::database::{lock_conn, Database, NewProviderAggregate}; use crate::error::AppError; -use crate::provider::{Provider, ProviderMeta}; +use crate::provider::{Provider, ProviderAggregate, ProviderMeta, ProviderMutationInput}; +use crate::settings::CustomEndpoint; use indexmap::IndexMap; -use rusqlite::params; +use rusqlite::{params, OptionalExtension, Row}; use std::collections::{HashMap, HashSet}; -type OmoProviderRow = ( - String, - String, - String, - Option, - Option, - Option, - Option, - String, -); +pub(super) struct StoredProviderRow { + id: String, + name: String, + settings_config: String, + website_url: Option, + category: Option, + created_at: Option, + sort_index: Option, + notes: Option, + icon: Option, + icon_color: Option, + meta: String, + in_failover_queue: bool, +} + +impl StoredProviderRow { + pub(super) fn from_row(row: &Row<'_>) -> rusqlite::Result { + Ok(Self { + id: row.get(0)?, + name: row.get(1)?, + settings_config: row.get(2)?, + website_url: row.get(3)?, + category: row.get(4)?, + created_at: row.get(5)?, + sort_index: row.get(6)?, + notes: row.get(7)?, + icon: row.get(8)?, + icon_color: row.get(9)?, + meta: row.get(10)?, + in_failover_queue: row.get(11)?, + }) + } + + pub(super) fn decode(self, app_type: &str) -> Result { + let (settings_config, mut meta) = + decode_provider_json(app_type, &self.id, &self.settings_config, &self.meta)?; + // Child rows are the sole endpoint authority. Do not expose a stale + // legacy copy that happens to remain embedded in provider metadata. + meta.custom_endpoints.clear(); + Ok(Provider { + id: self.id, + name: self.name, + settings_config, + website_url: self.website_url, + category: self.category, + created_at: self.created_at, + sort_index: self.sort_index, + notes: self.notes, + meta: Some(meta), + icon: self.icon, + icon_color: self.icon_color, + in_failover_queue: self.in_failover_queue, + }) + } +} + +fn decode_provider_json( + app_type: &str, + provider_id: &str, + settings_config: &str, + meta: &str, +) -> Result<(serde_json::Value, ProviderMeta), AppError> { + let settings_config = serde_json::from_str(settings_config).map_err(|error| { + AppError::Database(format!( + "invalid settings_config for provider '{app_type}/{provider_id}': {error}" + )) + })?; + let meta = if meta.trim().is_empty() { + ProviderMeta::default() + } else { + serde_json::from_str(meta).map_err(|error| { + AppError::Database(format!( + "invalid meta for provider '{app_type}/{provider_id}': {error}" + )) + })? + }; + Ok((settings_config, meta)) +} + +pub(super) const PROVIDER_SELECT: &str = + "SELECT id, name, settings_config, website_url, category, created_at, sort_index, + notes, icon, icon_color, meta, in_failover_queue + FROM providers"; + +fn load_endpoints( + conn: &rusqlite::Connection, + app_type: &str, + provider_id: Option<&str>, +) -> Result>, AppError> { + let mut grouped: HashMap> = HashMap::new(); + if let Some(provider_id) = provider_id { + let mut stmt = conn + .prepare( + "SELECT provider_id, url, added_at, last_used + FROM provider_endpoints + WHERE app_type = ?1 AND provider_id = ?2 + ORDER BY added_at, url, id", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let rows = stmt + .query_map(params![app_type, provider_id], decode_endpoint_row) + .map_err(|error| AppError::Database(error.to_string()))?; + collect_endpoints(rows, app_type, &mut grouped)?; + } else { + let mut stmt = conn + .prepare( + "SELECT provider_id, url, added_at, last_used + FROM provider_endpoints + WHERE app_type = ?1 + ORDER BY provider_id, added_at, url, id", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let rows = stmt + .query_map([app_type], decode_endpoint_row) + .map_err(|error| AppError::Database(error.to_string()))?; + collect_endpoints(rows, app_type, &mut grouped)?; + } + Ok(grouped) +} + +type StoredEndpoint = (String, String, CustomEndpoint); + +fn decode_endpoint_row(row: &Row<'_>) -> rusqlite::Result { + let provider_id: String = row.get(0)?; + let url: String = row.get(1)?; + Ok(( + provider_id, + url.clone(), + CustomEndpoint { + url, + added_at: row.get(2)?, + last_used: row.get(3)?, + }, + )) +} + +fn collect_endpoints( + rows: rusqlite::MappedRows<'_, impl FnMut(&Row<'_>) -> rusqlite::Result>, + app_type: &str, + grouped: &mut HashMap>, +) -> Result<(), AppError> { + for row in rows { + let (provider_id, url, endpoint) = + row.map_err(|error| AppError::Database(error.to_string()))?; + if grouped + .entry(provider_id.clone()) + .or_default() + .insert(url.clone(), endpoint) + .is_some() + { + return Err(AppError::Database(format!( + "duplicate endpoint '{url}' for provider '{app_type}/{provider_id}'" + ))); + } + } + Ok(()) +} impl Database { + pub fn get_all_provider_aggregates( + &self, + app_type: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + let mut stmt = conn + .prepare(&format!( + "{PROVIDER_SELECT} + WHERE app_type = ?1 + ORDER BY COALESCE(sort_index, 999999), created_at, id" + )) + .map_err(|error| AppError::Database(error.to_string()))?; + let rows = stmt + .query_map([app_type], StoredProviderRow::from_row) + .map_err(|error| AppError::Database(error.to_string()))?; + let mut endpoints = load_endpoints(&conn, app_type, None)?; + let mut aggregates = IndexMap::new(); + for row in rows { + let provider = row + .map_err(|error| AppError::Database(error.to_string()))? + .decode(app_type)?; + let provider_id = provider.id.clone(); + aggregates.insert( + provider_id.clone(), + ProviderAggregate { + provider, + endpoints: endpoints.remove(&provider_id).unwrap_or_default(), + }, + ); + } + Ok(aggregates) + } + pub fn get_all_providers( &self, app_type: &str, ) -> Result, AppError> { - let conn = lock_conn!(self.conn); - let mut stmt = conn.prepare( - "SELECT id, name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta, in_failover_queue - FROM providers WHERE app_type = ?1 - ORDER BY COALESCE(sort_index, 999999), created_at ASC, id ASC" - ).map_err(|e| AppError::Database(e.to_string()))?; - - let provider_iter = stmt - .query_map(params![app_type], |row| { - let id: String = row.get(0)?; - let name: String = row.get(1)?; - let settings_config_str: String = row.get(2)?; - let website_url: Option = row.get(3)?; - let category: Option = row.get(4)?; - let created_at: Option = row.get(5)?; - let sort_index: Option = row.get(6)?; - let notes: Option = row.get(7)?; - let icon: Option = row.get(8)?; - let icon_color: Option = row.get(9)?; - let meta_str: String = row.get(10)?; - let in_failover_queue: bool = row.get(11)?; - - let settings_config = - serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null); - let meta: ProviderMeta = serde_json::from_str(&meta_str).unwrap_or_default(); - - Ok(( - id, - Provider { - id: "".to_string(), // Placeholder, set below - name, - settings_config, - website_url, - category, - created_at, - sort_index, - notes, - meta: Some(meta), - icon, - icon_color, - in_failover_queue, - }, - )) - }) - .map_err(|e| AppError::Database(e.to_string()))?; - - let mut providers = IndexMap::new(); - for provider_res in provider_iter { - let (id, mut provider) = provider_res.map_err(|e| AppError::Database(e.to_string()))?; - provider.id = id.clone(); - - let mut stmt_endpoints = conn.prepare( - "SELECT url, added_at FROM provider_endpoints WHERE provider_id = ?1 AND app_type = ?2 ORDER BY added_at ASC, url ASC" - ).map_err(|e| AppError::Database(e.to_string()))?; - - let endpoints_iter = stmt_endpoints - .query_map(params![id, app_type], |row| { - let url: String = row.get(0)?; - let added_at: Option = row.get(1)?; - Ok(( - url, - crate::settings::CustomEndpoint { - url: "".to_string(), - added_at: added_at.unwrap_or(0), - last_used: None, - }, - )) - }) - .map_err(|e| AppError::Database(e.to_string()))?; - - let mut custom_endpoints = HashMap::new(); - for ep_res in endpoints_iter { - let (url, mut ep) = ep_res.map_err(|e| AppError::Database(e.to_string()))?; - ep.url = url.clone(); - custom_endpoints.insert(url, ep); - } - - if let Some(meta) = &mut provider.meta { - meta.custom_endpoints = custom_endpoints; - } - - providers.insert(id, provider); - } - - Ok(providers) + Ok(self + .get_all_provider_aggregates(app_type)? + .into_iter() + .map(|(id, aggregate)| (id, aggregate.into_provider())) + .collect()) } pub fn get_current_provider(&self, app_type: &str) -> Result, AppError> { @@ -132,149 +233,34 @@ impl Database { id: &str, app_type: &str, ) -> Result, AppError> { - let conn = lock_conn!(self.conn); - let result = conn.query_row( - "SELECT name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta, in_failover_queue - FROM providers WHERE id = ?1 AND app_type = ?2", - params![id, app_type], - |row| { - let name: String = row.get(0)?; - let settings_config_str: String = row.get(1)?; - let website_url: Option = row.get(2)?; - let category: Option = row.get(3)?; - let created_at: Option = row.get(4)?; - let sort_index: Option = row.get(5)?; - let notes: Option = row.get(6)?; - let icon: Option = row.get(7)?; - let icon_color: Option = row.get(8)?; - let meta_str: String = row.get(9)?; - let in_failover_queue: bool = row.get(10)?; - - let settings_config = serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null); - let meta: ProviderMeta = serde_json::from_str(&meta_str).unwrap_or_default(); - - Ok(Provider { - id: id.to_string(), - name, - settings_config, - website_url, - category, - created_at, - sort_index, - notes, - meta: Some(meta), - icon, - icon_color, - in_failover_queue, - }) - }, - ); - - match result { - Ok(provider) => Ok(Some(provider)), - Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(AppError::Database(e.to_string())), - } + Ok(self + .get_provider_aggregate(app_type, id)? + .map(ProviderAggregate::into_provider)) } - pub fn save_provider(&self, app_type: &str, provider: &Provider) -> Result<(), AppError> { - let mut conn = lock_conn!(self.conn); - let tx = conn - .transaction() - .map_err(|e| AppError::Database(e.to_string()))?; - - let mut meta_clone = provider.meta.clone().unwrap_or_default(); - let endpoints = std::mem::take(&mut meta_clone.custom_endpoints); - - let existing: Option<(bool, bool)> = tx + pub fn get_provider_aggregate( + &self, + app_type: &str, + id: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + let row = conn .query_row( - "SELECT is_current, in_failover_queue FROM providers WHERE id = ?1 AND app_type = ?2", - params![provider.id, app_type], - |row| Ok((row.get(0)?, row.get(1)?)), + &format!("{PROVIDER_SELECT} WHERE id = ?1 AND app_type = ?2"), + params![id, app_type], + StoredProviderRow::from_row, ) - .ok(); - - let is_update = existing.is_some(); - let (is_current, in_failover_queue) = - existing.unwrap_or((false, provider.in_failover_queue)); - - if is_update { - tx.execute( - "UPDATE providers SET - name = ?1, - settings_config = ?2, - website_url = ?3, - category = ?4, - created_at = ?5, - sort_index = ?6, - notes = ?7, - icon = ?8, - icon_color = ?9, - meta = ?10, - is_current = ?11, - in_failover_queue = ?12 - WHERE id = ?13 AND app_type = ?14", - params![ - provider.name, - serde_json::to_string(&provider.settings_config).map_err(|e| { - AppError::Database(format!("Failed to serialize settings_config: {e}")) - })?, - provider.website_url, - provider.category, - provider.created_at, - provider.sort_index, - provider.notes, - provider.icon, - provider.icon_color, - serde_json::to_string(&meta_clone).map_err(|e| AppError::Database(format!( - "Failed to serialize meta: {e}" - )))?, - is_current, - in_failover_queue, - provider.id, - app_type, - ], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - } else { - tx.execute( - "INSERT INTO providers ( - id, app_type, name, settings_config, website_url, category, - created_at, sort_index, notes, icon, icon_color, meta, is_current, in_failover_queue - ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)", - params![ - provider.id, - app_type, - provider.name, - serde_json::to_string(&provider.settings_config) - .map_err(|e| AppError::Database(format!("Failed to serialize settings_config: {e}")))?, - provider.website_url, - provider.category, - provider.created_at, - provider.sort_index, - provider.notes, - provider.icon, - provider.icon_color, - serde_json::to_string(&meta_clone) - .map_err(|e| AppError::Database(format!("Failed to serialize meta: {e}")))?, - is_current, - in_failover_queue, - ], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - - for (url, endpoint) in endpoints { - tx.execute( - "INSERT INTO provider_endpoints (provider_id, app_type, url, added_at) - VALUES (?1, ?2, ?3, ?4)", - params![provider.id, app_type, url, endpoint.added_at], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - } - } - - tx.commit().map_err(|e| AppError::Database(e.to_string()))?; - Ok(()) + .optional() + .map_err(|error| AppError::Database(error.to_string()))?; + let Some(row) = row else { + return Ok(None); + }; + let provider = row.decode(app_type)?; + let mut endpoints = load_endpoints(&conn, app_type, Some(id))?; + Ok(Some(ProviderAggregate { + provider, + endpoints: endpoints.remove(id).unwrap_or_default(), + })) } pub fn delete_provider(&self, app_type: &str, id: &str) -> Result<(), AppError> { @@ -309,57 +295,6 @@ impl Database { Ok(()) } - pub fn update_provider_settings_config( - &self, - app_type: &str, - provider_id: &str, - settings_config: &serde_json::Value, - ) -> Result<(), AppError> { - let conn = lock_conn!(self.conn); - conn.execute( - "UPDATE providers SET settings_config = ?1 WHERE id = ?2 AND app_type = ?3", - params![ - serde_json::to_string(settings_config).map_err(|e| AppError::Database(format!( - "Failed to serialize settings_config: {e}" - )))?, - provider_id, - app_type - ], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - Ok(()) - } - - pub fn add_custom_endpoint( - &self, - app_type: &str, - provider_id: &str, - url: &str, - ) -> Result<(), AppError> { - let conn = lock_conn!(self.conn); - let added_at = chrono::Utc::now().timestamp_millis(); - conn.execute( - "INSERT INTO provider_endpoints (provider_id, app_type, url, added_at) VALUES (?1, ?2, ?3, ?4)", - params![provider_id, app_type, url, added_at], - ).map_err(|e| AppError::Database(e.to_string()))?; - Ok(()) - } - - pub fn remove_custom_endpoint( - &self, - app_type: &str, - provider_id: &str, - url: &str, - ) -> Result<(), AppError> { - let conn = lock_conn!(self.conn); - conn.execute( - "DELETE FROM provider_endpoints WHERE provider_id = ?1 AND app_type = ?2 AND url = ?3", - params![provider_id, app_type, url], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - Ok(()) - } - pub fn set_omo_provider_current( &self, app_type: &str, @@ -443,63 +378,22 @@ impl Database { app_type: &str, category: &str, ) -> Result, AppError> { - let conn = lock_conn!(self.conn); - let row_data: Result = conn.query_row( - "SELECT id, name, settings_config, category, created_at, sort_index, notes, meta - FROM providers - WHERE app_type = ?1 AND category = ?2 AND is_current = 1 - LIMIT 1", - params![app_type, category], - |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get(4)?, - row.get(5)?, - row.get(6)?, - row.get(7)?, - )) - }, - ); - - let (id, name, settings_config_str, _row_category, created_at, sort_index, notes, meta_str) = - match row_data { - Ok(v) => v, - Err(rusqlite::Error::QueryReturnedNoRows) => return Ok(None), - Err(e) => return Err(AppError::Database(e.to_string())), - }; - - let settings_config = serde_json::from_str(&settings_config_str).map_err(|e| { - AppError::Database(format!( - "Failed to parse {category} provider settings_config (provider_id={id}): {e}" - )) - })?; - let meta: crate::provider::ProviderMeta = if meta_str.trim().is_empty() { - crate::provider::ProviderMeta::default() - } else { - serde_json::from_str(&meta_str).map_err(|e| { - AppError::Database(format!( - "Failed to parse {category} provider meta (provider_id={id}): {e}" - )) - })? + let provider_id = { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT id FROM providers + WHERE app_type = ?1 AND category = ?2 AND is_current = 1 + LIMIT 1", + params![app_type, category], + |row| row.get::<_, String>(0), + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))? }; - - Ok(Some(Provider { - id, - name, - settings_config, - website_url: None, - category: Some(category.to_string()), - created_at, - sort_index, - notes, - meta: Some(meta), - icon: None, - icon_color: None, - in_failover_queue: false, - })) + provider_id + .map(|provider_id| self.get_provider_by_id(&provider_id, app_type)) + .transpose() + .map(Option::flatten) } /// 判断 providers 表是否为空(全 app_type 一起算)。 @@ -593,8 +487,8 @@ impl Database { /// - 老用户升级:同样会触发一次(flag 不存在),追加到末尾,不影响已有排序 /// - 用户删除 seed 后:不再重建(flag 已为 true),尊重用户意图 /// - /// 与 `Database::save_provider` 的 UPSERT 语义配合,即使被意外重复调用 - /// 也不会覆盖用户当前激活的供应商(is_current 字段会被保留)。 + /// 每条 seed 都先读存在性,再走严格 create;并发冲突向上传播,不会覆盖 + /// 用户已有的同名供应商或当前状态。 pub fn init_default_official_providers(&self) -> Result { use crate::database::dao::providers_seed::OFFICIAL_SEEDS; @@ -623,19 +517,23 @@ impl Database { AppError::Database(format!("Seed JSON parse failed for {}: {e}", seed.id)) })?; - let mut provider = Provider::with_id( - seed.id.to_string(), - seed.name.to_string(), - settings_config, - Some(seed.website_url.to_string()), - ); - provider.category = Some("official".to_string()); - provider.icon = Some(seed.icon.to_string()); - provider.icon_color = Some(seed.icon_color.to_string()); - provider.sort_index = Some(next_sort_index); - provider.created_at = Some(now_ms); - - self.save_provider(app_type_str, &provider)?; + self.create_provider(NewProviderAggregate::from_input( + app_type_str, + ProviderMutationInput { + id: seed.id.to_string(), + name: seed.name.to_string(), + settings_config, + website_url: Some(seed.website_url.to_string()), + category: Some("official".to_string()), + created_at: Some(now_ms), + sort_index: Some(next_sort_index), + notes: None, + meta: None, + icon: Some(seed.icon.to_string()), + icon_color: Some(seed.icon_color.to_string()), + in_failover_queue: false, + }, + )?)?; inserted += 1; log::info!( "✓ Seeded official provider: {} ({})", @@ -689,19 +587,23 @@ impl Database { let next_sort_index = self.next_sort_index_for_app(app_type_str)?; let now_ms = chrono::Utc::now().timestamp_millis(); - let mut provider = Provider::with_id( - seed.id.to_string(), - seed.name.to_string(), - settings_config, - Some(seed.website_url.to_string()), - ); - provider.category = Some("official".to_string()); - provider.icon = Some(seed.icon.to_string()); - provider.icon_color = Some(seed.icon_color.to_string()); - provider.sort_index = Some(next_sort_index); - provider.created_at = Some(now_ms); - - self.save_provider(app_type_str, &provider)?; + self.create_provider(NewProviderAggregate::from_input( + app_type_str, + ProviderMutationInput { + id: seed.id.to_string(), + name: seed.name.to_string(), + settings_config, + website_url: Some(seed.website_url.to_string()), + category: Some("official".to_string()), + created_at: Some(now_ms), + sort_index: Some(next_sort_index), + notes: None, + meta: None, + icon: Some(seed.icon.to_string()), + icon_color: Some(seed.icon_color.to_string()), + in_failover_queue: false, + }, + )?)?; Ok(true) } @@ -751,7 +653,7 @@ mod ensure_official_seed_tests { .expect("query ok") .expect("seed present"); renamed.name = "My Custom Backup".to_string(); - db.save_provider(AppType::ClaudeDesktop.as_str(), &renamed) + db.reconcile_provider_fixture(AppType::ClaudeDesktop.as_str(), &renamed) .expect("save customization"); let inserted = db @@ -826,3 +728,344 @@ mod ensure_official_seed_tests { assert!(result.is_err(), "(id, app_type) mismatch should be Err"); } } + +#[cfg(test)] +mod aggregate_tests { + use crate::database::dao::provider_write; + use crate::database::{ + Database, NewEndpoint, NewProviderAggregate, ProviderKey, ProviderRowUpdate, + }; + use crate::error::AppError; + use crate::provider::{Provider, ProviderAggregate, ProviderMeta, ProviderMutationInput}; + use crate::settings::CustomEndpoint; + use indexmap::IndexMap; + use serde_json::json; + + fn aggregate() -> ProviderAggregate { + ProviderAggregate { + provider: Provider::with_id( + "pi-provider".into(), + "Pi Provider".into(), + json!({"models": [{"id": "m"}]}), + Some("https://example.test".into()), + ), + endpoints: IndexMap::from([ + ( + "https://one.test".into(), + CustomEndpoint { + url: "https://one.test".into(), + added_at: Some(10), + last_used: Some(11), + }, + ), + ( + "https://two.test".into(), + CustomEndpoint { + url: "https://two.test".into(), + added_at: Some(20), + last_used: Some(21), + }, + ), + ]), + } + } + + fn mutation_input(provider: Provider) -> ProviderMutationInput { + ProviderMutationInput { + id: provider.id, + name: provider.name, + settings_config: provider.settings_config, + website_url: provider.website_url, + category: provider.category, + created_at: provider.created_at, + sort_index: provider.sort_index, + notes: provider.notes, + meta: provider.meta, + icon: provider.icon, + icon_color: provider.icon_color, + in_failover_queue: provider.in_failover_queue, + } + } + + fn create_aggregate(db: &Database, app_type: &str) -> Result<(), AppError> { + db.create_provider(NewProviderAggregate::from_input( + app_type, + mutation_input(aggregate().into_provider()), + )?) + } + + #[test] + fn aggregate_single_and_all_hydration_match() -> Result<(), AppError> { + let db = Database::memory()?; + create_aggregate(&db, "pi")?; + + let single = db + .get_provider_aggregate("pi", "pi-provider")? + .expect("single aggregate"); + let all = db.get_all_provider_aggregates("pi")?; + assert_eq!( + serde_json::to_value(&single).expect("serialize single"), + serde_json::to_value(&all["pi-provider"]).expect("serialize all") + ); + assert_eq!(single.endpoints["https://one.test"].last_used, Some(11)); + let legacy = db + .get_provider_by_id("pi-provider", "pi")? + .expect("legacy projection"); + assert_eq!( + legacy + .meta + .expect("meta") + .custom_endpoints + .get("https://two.test") + .and_then(|endpoint| endpoint.last_used), + Some(21) + ); + Ok(()) + } + + #[test] + fn strict_create_rolls_back_and_never_upserts() -> Result<(), AppError> { + let db = Database::memory()?; + let mut malformed = aggregate(); + malformed.provider.id = "duplicate-payload".into(); + malformed.endpoints.insert( + "wrong-map-key".into(), + CustomEndpoint { + url: "https://duplicate.test".into(), + added_at: Some(1), + last_used: None, + }, + ); + assert!( + NewProviderAggregate::from_input("pi", mutation_input(malformed.into_provider())) + .and_then(|input| db.create_provider(input)) + .is_err() + ); + assert!(db + .get_provider_aggregate("pi", "duplicate-payload")? + .is_none()); + + { + let conn = crate::database::lock_conn!(db.conn); + conn.execute_batch( + "CREATE TRIGGER reject_bad_endpoint + BEFORE INSERT ON provider_endpoints + WHEN NEW.url = 'https://reject.test' + BEGIN SELECT RAISE(ABORT, 'injected endpoint failure'); END;", + )?; + } + let mut rejected = aggregate(); + rejected.provider.id = "rejected".into(); + rejected.endpoints = IndexMap::from([( + "https://reject.test".into(), + CustomEndpoint { + url: "https://reject.test".into(), + added_at: Some(99), + last_used: None, + }, + )]); + assert!( + NewProviderAggregate::from_input("pi", mutation_input(rejected.into_provider())) + .and_then(|input| db.create_provider(input)) + .is_err() + ); + assert!(db.get_provider_aggregate("pi", "rejected")?.is_none()); + + { + let conn = crate::database::lock_conn!(db.conn); + conn.execute_batch("DROP TRIGGER reject_bad_endpoint;")?; + } + create_aggregate(&db, "pi")?; + let original = db + .get_provider_aggregate("pi", "pi-provider")? + .expect("baseline aggregate"); + let mut conflicting = original.clone(); + conflicting.provider.meta = Some(ProviderMeta { + custom_endpoints: conflicting.endpoints.clone().into_iter().collect(), + ..conflicting.provider.meta.clone().unwrap_or_default() + }); + let mut conflicting = conflicting.into_provider(); + conflicting.name = "Must not upsert".into(); + assert!( + NewProviderAggregate::from_input("pi", mutation_input(conflicting)) + .and_then(|input| db.create_provider(input)) + .is_err() + ); + let after = db + .get_provider_aggregate("pi", "pi-provider")? + .expect("unchanged aggregate"); + assert_eq!( + serde_json::to_value(after).expect("serialize after"), + serde_json::to_value(original).expect("serialize original") + ); + Ok(()) + } + + #[test] + fn stale_row_update_cannot_overwrite_endpoint_mutations() -> Result<(), AppError> { + let db = Database::memory()?; + create_aggregate(&db, "pi")?; + let mut stale = db + .get_provider_aggregate("pi", "pi-provider")? + .expect("stale aggregate"); + + let key = ProviderKey::new("pi", "pi-provider")?; + db.add_provider_endpoint(&key, NewEndpoint::now("https://three.test")?)?; + db.remove_provider_endpoint(&key, "https://one.test")?; + db.touch_provider_endpoint(&key, "https://two.test", 222)?; + + stale.provider.name = "Row-only edit".into(); + stale + .provider + .meta + .get_or_insert_with(ProviderMeta::default) + .custom_endpoints = stale.endpoints.clone().into_iter().collect(); + stale + .provider + .meta + .as_mut() + .expect("meta") + .custom_endpoints + .clear(); + db.update_provider( + &key, + &ProviderRowUpdate::from_input(&mutation_input(stale.provider))?, + )?; + + let after = db + .get_provider_aggregate("pi", "pi-provider")? + .expect("provider after row-only update"); + assert_eq!(after.provider.name, "Row-only edit"); + assert_eq!( + after.endpoints.keys().cloned().collect::>(), + vec!["https://two.test", "https://three.test"] + ); + assert_eq!(after.endpoints["https://two.test"].last_used, Some(222)); + assert!(matches!( + db.touch_provider_endpoint(&key, "https://missing.test", 1), + Err(AppError::NotFound(_)) + )); + let stored_meta: String = { + let conn = crate::database::lock_conn!(db.conn); + conn.query_row( + "SELECT meta FROM providers WHERE id = 'pi-provider' AND app_type = 'pi'", + [], + |row| row.get(0), + )? + }; + let stored_meta: ProviderMeta = serde_json::from_str(&stored_meta) + .map_err(|error| AppError::Database(error.to_string()))?; + assert!(stored_meta.custom_endpoints.is_empty()); + Ok(()) + } + + #[test] + fn projection_failure_compensation_restores_exact_aggregate() -> Result<(), AppError> { + let db = Database::memory()?; + create_aggregate(&db, "pi")?; + let snapshot = db + .get_provider_aggregate("pi", "pi-provider")? + .expect("rollback snapshot"); + + { + let mut conn = crate::database::lock_conn!(db.conn); + conn.execute_batch( + "CREATE TRIGGER reject_projection + BEFORE INSERT ON pi_provider_projections + BEGIN SELECT RAISE(ABORT, 'injected projection failure'); END;", + )?; + let tx = conn.transaction()?; + tx.execute( + "DELETE FROM provider_endpoints + WHERE provider_id = 'pi-provider' + AND app_type = 'pi' + AND url = 'https://one.test'", + [], + )?; + assert!(tx + .execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES ('pi-provider', 'native-key', 1, 1)", + [], + ) + .is_err()); + let key = ProviderKey::new("pi", "pi-provider")?; + let row = ProviderRowUpdate::from_input(&mutation_input(snapshot.provider.clone()))?; + let endpoints = snapshot + .endpoints + .values() + .cloned() + .map(NewEndpoint::try_from) + .collect::, _>>()?; + provider_write::restore_provider_aggregate_on_tx( + &tx, + &key, + &row, + snapshot.provider.created_at, + snapshot.provider.sort_index, + false, + snapshot.provider.in_failover_queue, + &endpoints, + )?; + tx.commit()?; + } + + let single = db + .get_provider_aggregate("pi", "pi-provider")? + .expect("restored single aggregate"); + let all = db.get_all_provider_aggregates("pi")?; + assert_eq!( + serde_json::to_value(&single).expect("serialize single"), + serde_json::to_value(&snapshot).expect("serialize snapshot") + ); + assert_eq!( + serde_json::to_value(&all["pi-provider"]).expect("serialize all"), + serde_json::to_value(&snapshot).expect("serialize snapshot") + ); + Ok(()) + } + + #[test] + fn custom_endpoint_add_is_strict_for_one_logical_url() -> Result<(), AppError> { + let db = Database::memory()?; + create_aggregate(&db, "pi")?; + let key = ProviderKey::new("pi", "pi-provider")?; + db.add_provider_endpoint(&key, NewEndpoint::now("https://repeat.test")?)?; + assert!(db + .add_provider_endpoint(&key, NewEndpoint::now("https://repeat.test")?) + .is_err()); + + let saved = db + .get_provider_aggregate("pi", "pi-provider")? + .expect("aggregate"); + assert_eq!( + saved + .endpoints + .keys() + .filter(|url| url.as_str() == "https://repeat.test") + .count(), + 1 + ); + Ok(()) + } + + #[test] + fn aggregate_read_rejects_corrupt_json_consistently() -> Result<(), AppError> { + let db = Database::memory()?; + { + let conn = crate::database::lock_conn!(db.conn); + conn.execute( + "INSERT INTO providers + (id, app_type, name, settings_config, meta) + VALUES ('corrupt', 'pi', 'Corrupt', '{', '{}')", + [], + )?; + } + assert!(db.get_provider_aggregate("pi", "corrupt").is_err()); + assert!(db.get_all_provider_aggregates("pi").is_err()); + assert!(db.get_provider_by_id("corrupt", "pi").is_err()); + assert!(db.get_all_providers("pi").is_err()); + Ok(()) + } +} diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs index 25cd162e8..23a05284a 100644 --- a/src-tauri/src/database/mod.rs +++ b/src-tauri/src/database/mod.rs @@ -32,6 +32,9 @@ mod schema; mod tests; // DAO 类型导出供外部使用 +pub use dao::provider_write::{ + NewEndpoint, NewProviderAggregate, ProviderKey, ProviderRowUpdate, RenameProvider, +}; pub(crate) use dao::providers_seed::{ is_official_seed_id, CLAUDE_DESKTOP_OFFICIAL_PROVIDER_ID, CODEX_OFFICIAL_PROVIDER_ID, GROKBUILD_OFFICIAL_PROVIDER_ID, @@ -53,7 +56,7 @@ use std::sync::Mutex; /// 当前 Schema 版本号 /// 每次修改表结构时递增,并在 schema.rs 中添加相应的迁移逻辑 -pub(crate) const SCHEMA_VERSION: i32 = 16; +pub(crate) const SCHEMA_VERSION: i32 = 17; /// 安全地序列化 JSON,避免 unwrap panic pub(crate) fn to_json_string(value: &T) -> Result { @@ -197,6 +200,11 @@ impl Database { conn: Mutex::new(conn), }; db.create_tables()?; + // Keep the test database structurally identical to a fresh production + // database. Marking the base DDL as current without running the + // migration chain creates a false-current schema and makes restore + // tests certify columns that do not actually exist. + db.apply_schema_migrations()?; db.ensure_model_pricing_seeded()?; Ok(db) @@ -293,3 +301,39 @@ impl Database { Ok(count == 0) } } + +#[cfg(test)] +impl Database { + /// Test-fixture reconciliation helper. Production code cannot call this: + /// provider writes there must choose a typed create or update operation. + pub(crate) fn reconcile_provider_fixture( + &self, + app_type: &str, + provider: &crate::provider::Provider, + ) -> Result<(), AppError> { + let mut input = crate::provider::ProviderMutationInput { + id: provider.id.clone(), + name: provider.name.clone(), + settings_config: provider.settings_config.clone(), + website_url: provider.website_url.clone(), + category: provider.category.clone(), + created_at: provider.created_at, + sort_index: provider.sort_index, + notes: provider.notes.clone(), + meta: provider.meta.clone(), + icon: provider.icon.clone(), + icon_color: provider.icon_color.clone(), + in_failover_queue: provider.in_failover_queue, + }; + if self.get_provider_aggregate(app_type, &input.id)?.is_some() { + if let Some(meta) = input.meta.as_mut() { + meta.custom_endpoints.clear(); + } + let key = ProviderKey::new(app_type, input.id.clone())?; + let row = ProviderRowUpdate::from_input(&input)?; + self.update_provider(&key, &row) + } else { + self.create_provider(NewProviderAggregate::from_input(app_type, input)?) + } + } +} diff --git a/src-tauri/src/database/schema.rs b/src-tauri/src/database/schema.rs index ff408a945..78a2bc646 100644 --- a/src-tauri/src/database/schema.rs +++ b/src-tauri/src/database/schema.rs @@ -53,7 +53,10 @@ impl Database { app_type TEXT NOT NULL, url TEXT NOT NULL, added_at INTEGER, - FOREIGN KEY (provider_id, app_type) REFERENCES providers(id, app_type) ON DELETE CASCADE + last_used INTEGER, + FOREIGN KEY (provider_id, app_type) + REFERENCES providers(id, app_type) ON DELETE CASCADE, + UNIQUE (provider_id, app_type, url) )", [], ) @@ -97,6 +100,7 @@ impl Database { enabled_grokbuild BOOLEAN NOT NULL DEFAULT 0, enabled_opencode BOOLEAN NOT NULL DEFAULT 0, enabled_hermes BOOLEAN NOT NULL DEFAULT 0, + enabled_pi BOOLEAN NOT NULL DEFAULT 0, installed_at INTEGER NOT NULL DEFAULT 0, content_hash TEXT, updated_at INTEGER NOT NULL DEFAULT 0 @@ -105,6 +109,36 @@ impl Database { ) .map_err(|e| AppError::Database(e.to_string()))?; + // Reserve the v17 device-local ledgers here so later stacked features + // never mutate the semantics of an already-published migration. + conn.execute( + "CREATE TABLE IF NOT EXISTS pi_provider_projections ( + provider_id TEXT PRIMARY KEY, + provider_key TEXT NOT NULL UNIQUE, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + )", + [], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + conn.execute( + "CREATE TABLE IF NOT EXISTS skill_deployments ( + app_type TEXT NOT NULL CHECK (app_type = 'pi'), + skill_id TEXT NOT NULL, + destination TEXT NOT NULL, + destination_key TEXT NOT NULL, + method TEXT NOT NULL CHECK (method IN ('symlink', 'copy')), + source_identity TEXT NOT NULL, + deployed_digest TEXT, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (app_type, skill_id, destination_key), + UNIQUE (app_type, destination_key) + )", + [], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + // 6. Skill Repos 表 conn.execute( "CREATE TABLE IF NOT EXISTS skill_repos ( @@ -511,6 +545,13 @@ impl Database { Self::migrate_v15_to_v16(conn)?; Self::set_user_version(conn, 16)?; } + 16 => { + log::info!( + "迁移数据库从 v16 到 v17(规范化 provider endpoint 并预留设备本地 ledger)" + ); + Self::migrate_v16_to_v17(conn)?; + Self::set_user_version(conn, 17)?; + } _ => { return Err(AppError::Database(format!( "未知的数据库版本 {version},无法迁移到 {SCHEMA_VERSION}" @@ -1523,6 +1564,112 @@ impl Database { crate::services::session_usage_codex::reset_codex_usage_on_conn(conn, &codex_dir) } + /// v16 -> v17: make endpoint rows a lossless, uniquely owned child + /// collection. The device-local ledger DDL is reserved in the same + /// migration because later stacked PRs must not rewrite a released + /// user_version step. + fn migrate_v16_to_v17(conn: &Connection) -> Result<(), AppError> { + if Self::table_exists(conn, "provider_endpoints")? { + Self::add_column_if_missing(conn, "provider_endpoints", "last_used", "INTEGER")?; + conn.execute_batch( + "UPDATE provider_endpoints AS kept + SET added_at = ( + SELECT MIN(other.added_at) + FROM provider_endpoints AS other + WHERE other.provider_id = kept.provider_id + AND other.app_type = kept.app_type + AND other.url = kept.url + ), + last_used = ( + SELECT MAX(other.last_used) + FROM provider_endpoints AS other + WHERE other.provider_id = kept.provider_id + AND other.app_type = kept.app_type + AND other.url = kept.url + ) + WHERE kept.id = ( + SELECT MIN(other.id) + FROM provider_endpoints AS other + WHERE other.provider_id = kept.provider_id + AND other.app_type = kept.app_type + AND other.url = kept.url + ); + DELETE FROM provider_endpoints + WHERE id NOT IN ( + SELECT MIN(id) + FROM provider_endpoints + GROUP BY provider_id, app_type, url + ); + DROP TABLE IF EXISTS provider_endpoints_v17_canonical; + CREATE TABLE provider_endpoints_v17_canonical ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + provider_id TEXT NOT NULL, + app_type TEXT NOT NULL, + url TEXT NOT NULL, + added_at INTEGER, + last_used INTEGER, + FOREIGN KEY (provider_id, app_type) + REFERENCES providers(id, app_type) ON DELETE CASCADE, + UNIQUE (provider_id, app_type, url) + ); + INSERT INTO provider_endpoints_v17_canonical + (id, provider_id, app_type, url, added_at, last_used) + SELECT id, provider_id, app_type, url, added_at, last_used + FROM provider_endpoints; + DROP TABLE provider_endpoints; + ALTER TABLE provider_endpoints_v17_canonical + RENAME TO provider_endpoints;", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + } else { + conn.execute( + "CREATE TABLE provider_endpoints ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + provider_id TEXT NOT NULL, + app_type TEXT NOT NULL, + url TEXT NOT NULL, + added_at INTEGER, + last_used INTEGER, + FOREIGN KEY (provider_id, app_type) + REFERENCES providers(id, app_type) ON DELETE CASCADE, + UNIQUE (provider_id, app_type, url) + )", + [], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + } + if Self::table_exists(conn, "skills")? { + Self::add_column_if_missing( + conn, + "skills", + "enabled_pi", + "BOOLEAN NOT NULL DEFAULT 0", + )?; + } + conn.execute_batch( + "CREATE TABLE IF NOT EXISTS pi_provider_projections ( + provider_id TEXT PRIMARY KEY, + provider_key TEXT NOT NULL UNIQUE, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS skill_deployments ( + app_type TEXT NOT NULL CHECK (app_type = 'pi'), + skill_id TEXT NOT NULL, + destination TEXT NOT NULL, + destination_key TEXT NOT NULL, + method TEXT NOT NULL CHECK (method IN ('symlink', 'copy')), + source_identity TEXT NOT NULL, + deployed_digest TEXT, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (app_type, skill_id, destination_key), + UNIQUE (app_type, destination_key) + );", + ) + .map_err(|error| AppError::Database(error.to_string())) + } + /// 插入默认模型定价数据 /// 格式: (model_id, display_name, input, output, cache_read, cache_creation) /// 注意: model_id 使用短横线格式(如 claude-haiku-4-5),与 API 返回的模型名称标准化后一致 @@ -3222,7 +3369,7 @@ mod tests { Database::apply_schema_migrations_on_conn(&conn)?; - assert_eq!(Database::get_user_version(&conn)?, 16); + assert_eq!(Database::get_user_version(&conn)?, SCHEMA_VERSION); let counts: (i64, i64, i64, i64) = conn.query_row( "SELECT (SELECT COUNT(*) FROM proxy_request_logs WHERE data_source = 'codex_session'), @@ -3235,4 +3382,67 @@ mod tests { assert_eq!(counts, (0, 1, 0, 1)); Ok(()) } + + #[test] + fn migrate_v16_to_v17_preserves_endpoint_metadata_and_starts_ledgers_empty( + ) -> Result<(), AppError> { + let conn = Connection::open_in_memory()?; + conn.execute_batch( + "CREATE TABLE providers ( + id TEXT NOT NULL, + app_type TEXT NOT NULL, + PRIMARY KEY (id, app_type) + ); + CREATE TABLE provider_endpoints ( + id INTEGER PRIMARY KEY, + provider_id TEXT NOT NULL, + app_type TEXT NOT NULL, + url TEXT NOT NULL, + added_at INTEGER + ); + CREATE TABLE skills ( + id TEXT PRIMARY KEY, + enabled_codex BOOLEAN NOT NULL DEFAULT 0 + ); + INSERT INTO providers (id, app_type) VALUES ('provider', 'pi'); + INSERT INTO provider_endpoints + (id, provider_id, app_type, url, added_at) + VALUES + (1, 'provider', 'pi', 'https://duplicate.test', 20), + (2, 'provider', 'pi', 'https://duplicate.test', 10); + INSERT INTO skills (id, enabled_codex) VALUES ('existing', 1);", + )?; + Database::set_user_version(&conn, 16)?; + + Database::apply_schema_migrations_on_conn(&conn)?; + + assert_eq!(Database::get_user_version(&conn)?, SCHEMA_VERSION); + assert!(Database::has_column( + &conn, + "provider_endpoints", + "last_used" + )?); + assert!(Database::has_column(&conn, "skills", "enabled_pi")?); + assert!(Database::table_exists(&conn, "pi_provider_projections")?); + assert!(Database::table_exists(&conn, "skill_deployments")?); + let endpoint: (i64, Option) = conn.query_row( + "SELECT COUNT(*), MIN(added_at) + FROM provider_endpoints + WHERE provider_id = 'provider' + AND app_type = 'pi' + AND url = 'https://duplicate.test'", + [], + |row| Ok((row.get(0)?, row.get(1)?)), + )?; + assert_eq!(endpoint, (1, Some(10))); + let ledgers: (i64, i64) = conn.query_row( + "SELECT + (SELECT COUNT(*) FROM pi_provider_projections), + (SELECT COUNT(*) FROM skill_deployments)", + [], + |row| Ok((row.get(0)?, row.get(1)?)), + )?; + assert_eq!(ledgers, (0, 0)); + Ok(()) + } } diff --git a/src-tauri/src/database/tests.rs b/src-tauri/src/database/tests.rs index e6c39dc9b..2c895dfae 100644 --- a/src-tauri/src/database/tests.rs +++ b/src-tauri/src/database/tests.rs @@ -1,3 +1,5 @@ +#![cfg(test)] + //! 数据库模块测试 //! //! 包含 Schema 迁移和基本功能的测试。 diff --git a/src-tauri/src/deeplink/provider.rs b/src-tauri/src/deeplink/provider.rs index 7adaf346b..4ab1fd13e 100644 --- a/src-tauri/src/deeplink/provider.rs +++ b/src-tauri/src/deeplink/provider.rs @@ -109,27 +109,35 @@ pub fn import_provider_from_deeplink( let provider_id = provider.id.clone(); - // Use ProviderService to add the provider - ProviderService::add(state, app_type.clone(), provider, true)?; - - // Add extra endpoints as custom endpoints (skip first one as it's the primary) - for ep in all_endpoints.iter().skip(1) { - let normalized = ep.trim().trim_end_matches('/').to_string(); + // All endpoints supplied by one import request belong to the same create + // intent. Put the non-primary endpoints into the initial aggregate so the + // provider row and its complete endpoint set commit atomically. + let initial_endpoints = &mut provider + .meta + .get_or_insert_with(ProviderMeta::default) + .custom_endpoints; + for endpoint in all_endpoints.iter().skip(1) { + let normalized = endpoint.trim().trim_end_matches('/').to_string(); if !normalized.is_empty() { - if let Err(e) = ProviderService::add_custom_endpoint( - state, - app_type.clone(), - &provider_id, + initial_endpoints.insert( normalized.clone(), - ) { - log::warn!( - "Failed to add custom endpoint '{}': {e}", - crate::url_for_log(&normalized) - ); - } + crate::settings::CustomEndpoint { + url: normalized, + added_at: Some(timestamp), + last_used: None, + }, + ); } } + // ProviderService owns the strict aggregate create. + ProviderService::add( + state, + app_type.clone(), + crate::services::provider::provider_to_mutation_input(provider), + true, + )?; + // If enabled=true, set as current provider if merged_request.enabled.unwrap_or(false) { ProviderService::switch(state, app_type.clone(), &provider_id)?; diff --git a/src-tauri/src/deeplink/tests.rs b/src-tauri/src/deeplink/tests.rs index 332086084..49fcea1d1 100644 --- a/src-tauri/src/deeplink/tests.rs +++ b/src-tauri/src/deeplink/tests.rs @@ -1,9 +1,11 @@ +#![cfg(test)] + //! Deep link module tests use super::mcp::parse_mcp_apps; use super::parser::parse_deeplink_url; use super::prompt::import_prompt_from_deeplink; -use super::provider::parse_and_merge_config; +use super::provider::{import_provider_from_deeplink, parse_and_merge_config}; use super::utils::{infer_homepage_from_endpoint, validate_url}; use super::DeepLinkImportRequest; use crate::AppType; @@ -952,6 +954,39 @@ fn test_parse_multiple_endpoints_comma_separated() { assert!(endpoint.contains("https://api3.example.com")); } +#[test] +#[serial_test::serial] +fn provider_deeplink_creates_all_initial_endpoints_in_one_aggregate() { + let _test_home = TestHomeGuard::new(); + let request = parse_deeplink_url( + "ccswitch://v1/import?resource=provider&app=claude&name=Endpoint%20Aggregate&endpoint=https%3A%2F%2Fprimary.example.com,https%3A%2F%2Fsecond.example.com%2F,https%3A%2F%2Fthird.example.com&apiKey=sk-test", + ) + .expect("parse provider deeplink"); + let state = AppState::new(Arc::new(Database::memory().expect("create memory db"))); + + let provider_id = + import_provider_from_deeplink(&state, request).expect("import provider aggregate"); + let aggregate = state + .db + .get_provider_aggregate(AppType::Claude.as_str(), &provider_id) + .expect("read provider aggregate") + .expect("provider exists"); + + assert_eq!(aggregate.endpoints.len(), 2); + assert_eq!( + aggregate.endpoints["https://second.example.com"].url, + "https://second.example.com" + ); + assert_eq!( + aggregate.endpoints["https://third.example.com"].url, + "https://third.example.com" + ); + assert!(aggregate + .endpoints + .values() + .all(|endpoint| endpoint.added_at.is_some() && endpoint.last_used.is_none())); +} + #[test] fn test_parse_single_endpoint_backward_compatible() { // Old format with single endpoint should still work diff --git a/src-tauri/src/error.rs b/src-tauri/src/error.rs index 04509626a..00bcc6aff 100644 --- a/src-tauri/src/error.rs +++ b/src-tauri/src/error.rs @@ -9,6 +9,13 @@ pub enum AppError { Config(String), #[error("无效输入: {0}")] InvalidInput(String), + #[error("未找到: {0}")] + NotFound(String), + /// 结构化冲突:并发前置期望失败(如 reconcile 的 ExpectAbsent 撞上竞争 + /// 创建、ExpectPresent 的指纹过期)。调用方据此重读重试或上浮,不得解析 + /// Database(String) 文本。由前置工程 A 认证契约引入(T9)。 + #[error("并发冲突: {0}")] + Conflict(String), #[error("IO 错误: {path}: {source}")] Io { path: String, diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 9e8d0c99e..2e0844105 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -45,7 +45,10 @@ pub use codex_config::{ pub use commands::open_provider_terminal; pub use commands::*; pub use config::{get_claude_mcp_path, get_claude_settings_path, read_json_file}; -pub use database::{Database, Profile}; +pub use database::{ + Database, NewEndpoint, NewProviderAggregate, Profile, ProviderKey, ProviderRowUpdate, + RenameProvider, +}; pub use deeplink::{import_provider_from_deeplink, parse_deeplink_url, DeepLinkImportRequest}; pub use error::AppError; pub use grok_config::get_grok_config_path; @@ -57,7 +60,7 @@ pub use mcp::{ sync_single_server_to_gemini, sync_single_server_to_grokbuild, }; pub use prompt::Prompt; -pub use provider::{Provider, ProviderMeta}; +pub use provider::{Provider, ProviderAggregate, ProviderMeta, ProviderMutationInput}; pub use services::{ profile::{ProfilePayload, ProfileScope, ProfileService}, provider::reapply_current_codex_official_live, @@ -1988,6 +1991,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, @@ -2001,11 +2005,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/provider.rs b/src-tauri/src/provider.rs index 1f62dbc34..5c9dd1bc6 100644 --- a/src-tauri/src/provider.rs +++ b/src-tauri/src/provider.rs @@ -43,6 +43,84 @@ pub struct Provider { pub in_failover_queue: bool, } +/// IPC/service input for creating or editing a provider. +/// +/// This deliberately is not the hydrated [`Provider`] read projection. In +/// particular, callers cannot pass a DAO aggregate back into the provider-row +/// writer without first crossing the service boundary, where endpoint +/// ownership is checked. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderMutationInput { + pub id: String, + pub name: String, + #[serde(rename = "settingsConfig")] + pub settings_config: Value, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(rename = "websiteUrl")] + pub website_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub category: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(rename = "createdAt")] + pub created_at: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(rename = "sortIndex")] + pub sort_index: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub notes: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub meta: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub icon: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(rename = "iconColor")] + pub icon_color: Option, + #[serde(default)] + #[serde(rename = "inFailoverQueue")] + pub in_failover_queue: bool, +} + +impl From for Provider { + fn from(input: ProviderMutationInput) -> Self { + Self { + id: input.id, + name: input.name, + settings_config: input.settings_config, + website_url: input.website_url, + category: input.category, + created_at: input.created_at, + sort_index: input.sort_index, + notes: input.notes, + meta: input.meta, + icon: input.icon, + icon_color: input.icon_color, + in_failover_queue: input.in_failover_queue, + } + } +} + +/// A provider row and every endpoint owned by that row. +/// +/// SQLite stores endpoints separately from provider metadata. This aggregate +/// is the only lossless DAO boundary; legacy `Provider` reads are projections +/// of it for API compatibility. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderAggregate { + pub provider: Provider, + #[serde(default)] + pub endpoints: IndexMap, +} + +impl ProviderAggregate { + pub(crate) fn into_provider(mut self) -> Provider { + self.provider + .meta + .get_or_insert_with(ProviderMeta::default) + .custom_endpoints = self.endpoints.into_iter().collect(); + self.provider + } +} + impl Provider { /// 从现有ID创建供应商 pub fn with_id( @@ -67,6 +145,72 @@ impl Provider { } } + pub(crate) fn row_content_fingerprint(&self) -> String { + use sha2::{Digest, Sha256}; + + fn hash_canonical(value: &serde_json::Value, hasher: &mut Sha256) { + match value { + serde_json::Value::Null => hasher.update(b"n"), + serde_json::Value::Bool(value) => { + hasher.update(b"b"); + hasher.update([*value as u8]); + } + serde_json::Value::Number(value) => { + let text = value.to_string(); + hasher.update(b"#"); + hasher.update((text.len() as u64).to_le_bytes()); + hasher.update(text.as_bytes()); + } + serde_json::Value::String(value) => { + hasher.update(b"s"); + hasher.update((value.len() as u64).to_le_bytes()); + hasher.update(value.as_bytes()); + } + serde_json::Value::Array(items) => { + hasher.update(b"["); + hasher.update((items.len() as u64).to_le_bytes()); + for item in items { + hash_canonical(item, hasher); + } + hasher.update(b"]"); + } + serde_json::Value::Object(map) => { + hasher.update(b"{"); + hasher.update((map.len() as u64).to_le_bytes()); + let mut keys: Vec<&String> = map.keys().collect(); + keys.sort(); + for key in keys { + hasher.update((key.len() as u64).to_le_bytes()); + hasher.update(key.as_bytes()); + hash_canonical(&map[key.as_str()], hasher); + } + hasher.update(b"}"); + } + } + } + + let mut meta = serde_json::to_value(&self.meta).unwrap_or(serde_json::Value::Null); + if let serde_json::Value::Object(map) = &mut meta { + map.remove("custom_endpoints"); + map.remove("customEndpoints"); + } + let mut hasher = Sha256::new(); + for part in [ + serde_json::Value::String(self.name.clone()), + self.settings_config.clone(), + serde_json::to_value(&self.website_url).unwrap_or(serde_json::Value::Null), + serde_json::to_value(&self.category).unwrap_or(serde_json::Value::Null), + serde_json::to_value(&self.notes).unwrap_or(serde_json::Value::Null), + serde_json::to_value(&self.icon).unwrap_or(serde_json::Value::Null), + serde_json::to_value(&self.icon_color).unwrap_or(serde_json::Value::Null), + meta, + ] { + hash_canonical(&part, &mut hasher); + hasher.update([0u8]); + } + format!("{:x}", hasher.finalize()) + } + pub fn is_codex_oauth(&self) -> bool { self.provider_type() == Some("codex_oauth") } diff --git a/src-tauri/src/proxy/provider_router.rs b/src-tauri/src/proxy/provider_router.rs index 28d2b8a2d..2baa11fa6 100644 --- a/src-tauri/src/proxy/provider_router.rs +++ b/src-tauri/src/proxy/provider_router.rs @@ -348,8 +348,10 @@ mod tests { let provider_b = Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); - db.save_provider("claude", &provider_a).unwrap(); - db.save_provider("claude", &provider_b).unwrap(); + db.reconcile_provider_fixture("claude", &provider_a) + .unwrap(); + db.reconcile_provider_fixture("claude", &provider_b) + .unwrap(); db.set_current_provider("claude", "a").unwrap(); db.add_to_failover_queue("claude", "b").unwrap(); @@ -374,8 +376,10 @@ mod tests { Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); provider_b.sort_index = Some(1); - db.save_provider("claude", &provider_a).unwrap(); - db.save_provider("claude", &provider_b).unwrap(); + db.reconcile_provider_fixture("claude", &provider_a) + .unwrap(); + db.reconcile_provider_fixture("claude", &provider_b) + .unwrap(); db.set_current_provider("claude", "a").unwrap(); db.add_to_failover_queue("claude", "b").unwrap(); @@ -407,8 +411,10 @@ mod tests { Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); provider_b.sort_index = Some(1); - db.save_provider("claude", &provider_a).unwrap(); - db.save_provider("claude", &provider_b).unwrap(); + db.reconcile_provider_fixture("claude", &provider_a) + .unwrap(); + db.reconcile_provider_fixture("claude", &provider_b) + .unwrap(); db.set_current_provider("claude", "a").unwrap(); // 只把 b 加入故障转移队列(模拟“当前供应商不在队列里”的常见配置) @@ -444,8 +450,10 @@ mod tests { let provider_b = Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); - db.save_provider("claude", &provider_a).unwrap(); - db.save_provider("claude", &provider_b).unwrap(); + db.reconcile_provider_fixture("claude", &provider_a) + .unwrap(); + db.reconcile_provider_fixture("claude", &provider_b) + .unwrap(); db.add_to_failover_queue("claude", "a").unwrap(); db.add_to_failover_queue("claude", "b").unwrap(); @@ -485,7 +493,8 @@ mod tests { let provider_a = Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None); - db.save_provider("claude", &provider_a).unwrap(); + db.reconcile_provider_fixture("claude", &provider_a) + .unwrap(); db.add_to_failover_queue("claude", "a").unwrap(); // 启用自动故障转移 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/omo.rs b/src-tauri/src/services/omo.rs index a91a35a5a..26743dd80 100644 --- a/src-tauri/src/services/omo.rs +++ b/src-tauri/src/services/omo.rs @@ -1,4 +1,5 @@ use crate::config::{atomic_write, write_json_file}; +use crate::database::NewProviderAggregate; use crate::error::AppError; use crate::opencode_config::get_opencode_dir; use crate::provider::Provider; @@ -288,7 +289,10 @@ impl OmoService { in_failover_queue: false, }; - state.db.save_provider("opencode", &provider)?; + state.db.create_provider(NewProviderAggregate::from_input( + "opencode", + crate::services::provider::provider_to_mutation_input(provider.clone()), + )?)?; state .db .set_omo_provider_current("opencode", &provider.id, v.category)?; diff --git a/src-tauri/src/services/provider/endpoints.rs b/src-tauri/src/services/provider/endpoints.rs index 4a7894aa0..d57ee75c6 100644 --- a/src-tauri/src/services/provider/endpoints.rs +++ b/src-tauri/src/services/provider/endpoints.rs @@ -5,6 +5,7 @@ use std::time::{SystemTime, UNIX_EPOCH}; use crate::app_config::AppType; +use crate::database::{NewEndpoint, ProviderKey}; use crate::error::AppError; use crate::settings::CustomEndpoint; use crate::store::AppState; @@ -47,9 +48,10 @@ pub fn add_custom_endpoint( )); } + let key = ProviderKey::new(app_type.as_str(), provider_id)?; state .db - .add_custom_endpoint(app_type.as_str(), provider_id, &normalized)?; + .add_provider_endpoint(&key, NewEndpoint::now(normalized)?)?; Ok(()) } @@ -61,9 +63,8 @@ pub fn remove_custom_endpoint( url: String, ) -> Result<(), AppError> { let normalized = url.trim().trim_end_matches('/').to_string(); - state - .db - .remove_custom_endpoint(app_type.as_str(), provider_id, &normalized)?; + let key = ProviderKey::new(app_type.as_str(), provider_id)?; + state.db.remove_provider_endpoint(&key, &normalized)?; Ok(()) } @@ -76,17 +77,10 @@ pub fn update_endpoint_last_used( ) -> Result<(), AppError> { let normalized = url.trim().trim_end_matches('/').to_string(); - // Get provider, update last_used, save back - let mut providers = state.db.get_all_providers(app_type.as_str())?; - if let Some(provider) = providers.get_mut(provider_id) { - if let Some(meta) = provider.meta.as_mut() { - if let Some(endpoint) = meta.custom_endpoints.get_mut(&normalized) { - endpoint.last_used = Some(now_millis()); - state.db.save_provider(app_type.as_str(), provider)?; - } - } - } - Ok(()) + let key = ProviderKey::new(app_type.as_str(), provider_id)?; + state + .db + .touch_provider_endpoint(&key, &normalized, now_millis()) } /// Get current timestamp in milliseconds diff --git a/src-tauri/src/services/provider/live.rs b/src-tauri/src/services/provider/live.rs index 4d26e5caf..555d48a21 100644 --- a/src-tauri/src/services/provider/live.rs +++ b/src-tauri/src/services/provider/live.rs @@ -19,7 +19,10 @@ use crate::store::AppState; use super::gemini_auth::{ detect_gemini_auth_type, ensure_google_oauth_security_flag, GeminiAuthType, }; -use super::normalize_claude_models_in_value; +use super::{ + normalize_claude_models_in_value, provider_row_fingerprint, provider_to_mutation_input, + reconcile_provider_record_with_precondition, ReconcilePrecondition, +}; /// ChatGPT Codex catalogs gpt-5.6 at a 372K context window with a ~353K /// effective budget (openai/codex#31860), far below the 1.05M API spec. @@ -1279,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 { @@ -1564,7 +1573,12 @@ pub fn import_default_config(state: &AppState, app_type: AppType) -> Result Result { + let existing = existing.provider; let display_name = config.name.clone().unwrap_or_else(|| existing.name.clone()); if existing.settings_config != settings_config || existing.name != display_name { + let fingerprint = provider_row_fingerprint(&existing); let mut provider = existing; provider.name = display_name; provider.settings_config = settings_config; - if let Err(e) = state.db.save_provider("opencode", &provider) { + if let Some(meta) = provider.meta.as_mut() { + meta.custom_endpoints.clear(); + } + if let Err(e) = reconcile_provider_record_with_precondition( + &state.db, + "opencode", + provider_to_mutation_input(provider), + ReconcilePrecondition::ExpectPresent { fingerprint }, + ) { log::warn!( "Failed to update OpenCode provider '{id}' from live config: {e}" ); @@ -1767,7 +1791,12 @@ pub fn import_opencode_providers_from_live(state: &AppState) -> Result Result { + let existing = existing.provider; if existing.settings_config != settings_config { + let fingerprint = provider_row_fingerprint(&existing); let mut provider = existing; provider.settings_config = settings_config; - if let Err(e) = state.db.save_provider("openclaw", &provider) { + if let Some(meta) = provider.meta.as_mut() { + meta.custom_endpoints.clear(); + } + if let Err(e) = reconcile_provider_record_with_precondition( + &state.db, + "openclaw", + provider_to_mutation_input(provider), + ReconcilePrecondition::ExpectPresent { fingerprint }, + ) { log::warn!( "Failed to update OpenClaw provider '{id}' from live config: {e}" ); @@ -1855,7 +1894,12 @@ pub fn import_openclaw_providers_from_live(state: &AppState) -> Result Result { + let existing = existing.provider; if existing.settings_config != config { + let fingerprint = provider_row_fingerprint(&existing); let mut provider = existing; provider.settings_config = config; - if let Err(e) = state.db.save_provider("hermes", &provider) { + if let Some(meta) = provider.meta.as_mut() { + meta.custom_endpoints.clear(); + } + if let Err(e) = reconcile_provider_record_with_precondition( + &state.db, + "hermes", + provider_to_mutation_input(provider), + ReconcilePrecondition::ExpectPresent { fingerprint }, + ) { log::warn!( "Failed to update Hermes provider '{name}' from live config: {e}" ); @@ -1923,7 +1977,12 @@ pub fn import_hermes_providers_from_live(state: &AppState) -> Result Result` implementation: a +/// hydrated read projection cannot silently become a write DTO via `.into()`. +pub(crate) fn provider_to_mutation_input(provider: Provider) -> ProviderMutationInput { + ProviderMutationInput { + id: provider.id, + name: provider.name, + settings_config: provider.settings_config, + website_url: provider.website_url, + category: provider.category, + created_at: provider.created_at, + sort_index: provider.sort_index, + notes: provider.notes, + meta: provider.meta, + icon: provider.icon, + icon_color: provider.icon_color, + in_failover_queue: provider.in_failover_queue, + } +} + +fn create_provider_record( + state: &AppState, + app_type: &AppType, + input: ProviderMutationInput, +) -> Result<(), AppError> { + state + .db + .create_provider(NewProviderAggregate::from_input(app_type.as_str(), input)?) +} + +fn update_provider_record( + state: &AppState, + app_type: &AppType, + input: &ProviderMutationInput, +) -> Result<(), AppError> { + let key = ProviderKey::new(app_type.as_str(), input.id.clone())?; + let row = ProviderRowUpdate::from_input(input)?; + 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)] +pub(crate) enum ReconcilePrecondition { + /// 调用方观察到目标不存在;若已被竞争者创建,必须返回 + /// [`AppError::Conflict`],绝不退化为覆盖更新。 + ExpectAbsent, + /// 调用方观察到目标存在且内容指纹为 `fingerprint`;指纹过期必须返回 + /// [`AppError::Conflict`],由调用方重读重试。 + ExpectPresent { fingerprint: String }, +} + +/// 行内容指纹:并发前置期望的版本标记(纯函数,不含状态列与 endpoint)。 +/// +/// 决定性要求:仓库启用了 serde_json `preserve_order`,且 `ProviderMeta` +/// 内含 HashMap——直接序列化的键序随机,会产生伪 Conflict。因此必须走 +/// 递归排序的规范化哈希;`meta.custom_endpoints` 属 endpoint authority, +/// 不参与内容指纹(不同读 API 对其填充不一致)。 +pub(crate) fn provider_row_fingerprint(provider: &crate::provider::Provider) -> String { + provider.row_content_fingerprint() +} + +/// Reconcile paths must carry the caller's observed state into the write. +/// Creation is strict, while updates compare the observed row fingerprint and +/// write under one database lock and transaction. +pub(crate) fn reconcile_provider_record_with_precondition( + db: &crate::database::Database, + app_type: &str, + input: ProviderMutationInput, + precondition: ReconcilePrecondition, +) -> Result<(), AppError> { + match precondition { + ReconcilePrecondition::ExpectAbsent => { + db.create_provider(NewProviderAggregate::from_input(app_type, input)?) + } + ReconcilePrecondition::ExpectPresent { fingerprint } => { + let key = ProviderKey::new(app_type, input.id.clone())?; + let row = ProviderRowUpdate::from_input(&input)?; + db.update_provider_if_content_fingerprint(&key, &fingerprint, &row) + } + } +} + /// Result of a provider switch operation, including any non-fatal warnings #[derive(Debug, serde::Serialize, Default)] #[serde(rename_all = "camelCase")] @@ -132,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 { @@ -437,6 +556,610 @@ mod tests { }) } + fn endpoint(url: &str, added_at: Option, last_used: Option) -> CustomEndpoint { + CustomEndpoint { + url: url.to_string(), + added_at, + last_used, + } + } + + fn provider_snapshot(state: &AppState, app_type: &str, id: &str) -> (Value, i64, i64) { + let aggregate = state + .db + .get_provider_aggregate(app_type, id) + .expect("read aggregate") + .map(|aggregate| serde_json::to_value(aggregate).expect("serialize aggregate")) + .unwrap_or(Value::Null); + let conn = state.db.conn.lock().expect("lock test database"); + let state_bits = conn + .query_row( + "SELECT is_current, in_failover_queue + FROM providers + WHERE app_type = ?1 AND id = ?2", + rusqlite::params![app_type, id], + |row| Ok((row.get(0)?, row.get(1)?)), + ) + .unwrap_or((0, 0)); + (aggregate, state_bits.0, state_bits.1) + } + + #[test] + #[serial] + fn provider_service_create_owns_initial_endpoints_and_duplicate_is_atomic() { + with_test_home(|state, _| { + let mut provider = opencode_provider("typed-create"); + provider.in_failover_queue = true; + let expected_endpoints = HashMap::from([ + ( + "https://one.example".to_string(), + endpoint("https://one.example", None, Some(11)), + ), + ( + "https://two.example".to_string(), + endpoint("https://two.example", Some(20), None), + ), + ]); + provider.meta = Some(ProviderMeta { + custom_endpoints: expected_endpoints.clone(), + ..Default::default() + }); + let input = provider_to_mutation_input(provider); + + ProviderService::add(state, AppType::OpenCode, input.clone(), false) + .expect("strict service create"); + let aggregate = state + .db + .get_provider_aggregate("opencode", "typed-create") + .expect("read") + .expect("aggregate"); + let hydrated_endpoints = aggregate.endpoints.into_iter().collect::>(); + assert_eq!( + hydrated_endpoints, expected_endpoints, + "the public create entry must hydrate the complete initial endpoint set losslessly" + ); + + let before = provider_snapshot(state, "opencode", "typed-create"); + assert!( + ProviderService::add(state, AppType::OpenCode, input, false).is_err(), + "duplicate create must not reconcile as update" + ); + assert_eq!( + provider_snapshot(state, "opencode", "typed-create"), + before, + "row, endpoints, current and failover state remain byte-logically unchanged" + ); + }); + } + + #[test] + #[serial] + fn provider_service_create_canonicalizes_initial_endpoint_identity() { + with_test_home(|state, _| { + let raw_url = " https://canonical.example/// "; + let mut provider = opencode_provider("canonical-endpoint"); + provider.meta = Some(ProviderMeta { + custom_endpoints: HashMap::from([( + raw_url.to_string(), + endpoint(raw_url, None, Some(11)), + )]), + ..Default::default() + }); + + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(provider), + false, + ) + .expect("create with a non-canonical initial endpoint"); + let aggregate = state + .db + .get_provider_aggregate("opencode", "canonical-endpoint") + .expect("read canonical aggregate") + .expect("canonical aggregate"); + assert_eq!(aggregate.endpoints.len(), 1); + assert!(aggregate + .endpoints + .contains_key("https://canonical.example")); + + ProviderService::update_endpoint_last_used( + state, + AppType::OpenCode, + "canonical-endpoint", + " https://canonical.example/ ".to_string(), + ) + .expect("touch must resolve the same canonical endpoint"); + ProviderService::remove_custom_endpoint( + state, + AppType::OpenCode, + "canonical-endpoint", + "https://canonical.example///".to_string(), + ) + .expect("remove must resolve the same canonical endpoint"); + assert!(state + .db + .get_provider_aggregate("opencode", "canonical-endpoint") + .expect("read after remove") + .expect("provider after remove") + .endpoints + .is_empty()); + + let mut duplicate = opencode_provider("duplicate-canonical-endpoint"); + duplicate.meta = Some(ProviderMeta { + custom_endpoints: HashMap::from([ + ( + "https://duplicate.example".to_string(), + endpoint("https://duplicate.example", None, None), + ), + ( + " https://duplicate.example/ ".to_string(), + endpoint(" https://duplicate.example/ ", Some(1), None), + ), + ]), + ..Default::default() + }); + assert!(matches!( + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(duplicate), + false, + ), + Err(AppError::InvalidInput(_)) + )); + assert!(state + .db + .get_provider_aggregate("opencode", "duplicate-canonical-endpoint") + .expect("read duplicate candidate") + .is_none()); + }); + } + + #[test] + #[serial] + fn provider_service_stale_edit_payload_cannot_overwrite_endpoint_operations() { + with_test_home(|state, _| { + let mut provider = opencode_provider("stale-edit"); + provider.meta = Some(ProviderMeta { + custom_endpoints: HashMap::from([ + ( + "https://remove.example".to_string(), + endpoint("https://remove.example", Some(1), None), + ), + ( + "https://touch.example".to_string(), + endpoint("https://touch.example", None, None), + ), + ]), + ..Default::default() + }); + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(provider), + false, + ) + .expect("create"); + + // This is the existing-provider form snapshot: endpoints are + // intentionally absent from the update IPC. + let mut edit = state + .db + .get_provider_by_id("stale-edit", "opencode") + .expect("read") + .expect("provider"); + edit.name = "Edited row".to_string(); + edit.meta + .get_or_insert_with(Default::default) + .custom_endpoints + .clear(); + let stale_row_payload = provider_to_mutation_input(edit); + + ProviderService::add_custom_endpoint( + state, + AppType::OpenCode, + "stale-edit", + "https://added.example".to_string(), + ) + .expect("concurrent add"); + ProviderService::remove_custom_endpoint( + state, + AppType::OpenCode, + "stale-edit", + "https://remove.example".to_string(), + ) + .expect("concurrent remove"); + ProviderService::update_endpoint_last_used( + state, + AppType::OpenCode, + "stale-edit", + "https://touch.example".to_string(), + ) + .expect("concurrent touch"); + + ProviderService::update(state, AppType::OpenCode, None, stale_row_payload) + .expect("row-only service update"); + let aggregate = state + .db + .get_provider_aggregate("opencode", "stale-edit") + .expect("read") + .expect("aggregate"); + assert_eq!(aggregate.provider.name, "Edited row"); + assert!(!aggregate.endpoints.contains_key("https://remove.example")); + assert!(aggregate.endpoints.contains_key("https://added.example")); + assert!(aggregate.endpoints["https://touch.example"] + .last_used + .is_some()); + + let mut forbidden = provider_to_mutation_input(aggregate.into_provider()); + forbidden + .meta + .get_or_insert_with(Default::default) + .custom_endpoints + .insert( + "https://forbidden.example".to_string(), + endpoint("https://forbidden.example", None, None), + ); + let before = provider_snapshot(state, "opencode", "stale-edit"); + assert!( + ProviderService::update(state, AppType::OpenCode, None, forbidden).is_err(), + "endpoint-bearing update IPC is rejected" + ); + assert_eq!(provider_snapshot(state, "opencode", "stale-edit"), before); + }); + } + + #[test] + #[serial] + fn provider_service_db_only_rename_matrix_is_atomic_and_lossless() { + with_test_home(|state, _| { + let mut source = opencode_provider("rename-source"); + let expected_endpoints = HashMap::from([ + ( + "https://nullable.example".to_string(), + endpoint("https://nullable.example", None, None), + ), + ( + "https://timed.example".to_string(), + endpoint("https://timed.example", Some(10), Some(11)), + ), + ]); + source.meta = Some(ProviderMeta { + custom_endpoints: expected_endpoints.clone(), + ..Default::default() + }); + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(source), + false, + ) + .expect("DB-only source"); + let mut renamed = opencode_provider("rename-target"); + renamed.name = "Renamed".to_string(); + ProviderService::update( + state, + AppType::OpenCode, + Some("rename-source"), + provider_to_mutation_input(renamed), + ) + .expect("DB-only additive rename"); + assert!(state + .db + .get_provider_aggregate("opencode", "rename-source") + .expect("old read") + .is_none()); + let renamed = state + .db + .get_provider_aggregate("opencode", "rename-target") + .expect("new read") + .expect("renamed"); + let renamed_endpoints = renamed.endpoints.into_iter().collect::>(); + assert_eq!( + renamed_endpoints, expected_endpoints, + "rename preserves every endpoint field, including NULL timestamps" + ); + + for id in ["conflict-source", "conflict-target"] { + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(opencode_provider(id)), + false, + ) + .expect("conflict fixture"); + } + let source_before = provider_snapshot(state, "opencode", "conflict-source"); + let target_before = provider_snapshot(state, "opencode", "conflict-target"); + assert!(ProviderService::update( + state, + AppType::OpenCode, + Some("conflict-source"), + provider_to_mutation_input(opencode_provider("conflict-target")), + ) + .is_err()); + assert_eq!( + provider_snapshot(state, "opencode", "conflict-source"), + source_before + ); + assert_eq!( + provider_snapshot(state, "opencode", "conflict-target"), + target_before + ); + + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(opencode_provider("live-source")), + true, + ) + .expect("live source"); + let live_before = provider_snapshot(state, "opencode", "live-source"); + assert!(ProviderService::update( + state, + AppType::OpenCode, + Some("live-source"), + provider_to_mutation_input(opencode_provider("live-target")), + ) + .is_err()); + assert_eq!( + provider_snapshot(state, "opencode", "live-source"), + live_before + ); + + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(opencode_omo_provider("omo-source", "omo")), + false, + ) + .expect("OMO source"); + let omo_before = provider_snapshot(state, "opencode", "omo-source"); + let mut omo_target = opencode_omo_provider("omo-target", "omo"); + omo_target.name = "Forbidden OMO rename".to_string(); + assert!(ProviderService::update( + state, + AppType::OpenCode, + Some("omo-source"), + provider_to_mutation_input(omo_target), + ) + .is_err()); + assert_eq!( + provider_snapshot(state, "opencode", "omo-source"), + omo_before + ); + + ProviderService::add( + state, + AppType::Hermes, + provider_to_mutation_input(hermes_provider("hermes-source")), + false, + ) + .expect("Hermes source"); + let hermes_before = provider_snapshot(state, "hermes", "hermes-source"); + assert!(ProviderService::update( + state, + AppType::Hermes, + Some("hermes-source"), + provider_to_mutation_input(hermes_provider("hermes-target")), + ) + .is_err()); + assert_eq!( + provider_snapshot(state, "hermes", "hermes-source"), + hermes_before + ); + }); + } + + #[test] + #[serial] + fn provider_service_rename_fails_closed_for_malformed_additive_live_config() { + with_test_home(|state, home| { + let cases = [ + ( + AppType::OpenCode, + opencode_provider("malformed-source"), + opencode_provider("malformed-target"), + home.join(".config").join("opencode").join("opencode.json"), + r#"{"provider":{"malformed-source":{"npm":"@ai-sdk/openai-compatible"}"#, + ), + ( + AppType::OpenClaw, + openclaw_provider("corrupt-source"), + openclaw_provider("corrupt-target"), + home.join(".openclaw").join("openclaw.json"), + r#"{"models":{"providers":{"corrupt-target":{"baseUrl":"https://example.test"}"#, + ), + ]; + + for (app_type, source, target, live_path, malformed_live) in cases { + let app_name = app_type.as_str().to_string(); + let source_id = source.id.clone(); + let target_id = target.id.clone(); + ProviderService::add( + state, + app_type.clone(), + provider_to_mutation_input(source), + false, + ) + .expect("create DB-only rename source"); + + fs::create_dir_all(live_path.parent().expect("live config parent")) + .expect("create live config directory"); + fs::write(&live_path, malformed_live).expect("write malformed live config"); + let source_before = provider_snapshot(state, &app_name, &source_id); + let target_before = provider_snapshot(state, &app_name, &target_id); + + let error = ProviderService::update( + state, + app_type, + Some(&source_id), + provider_to_mutation_input(target), + ) + .expect_err("rename must fail closed when live identity cannot be inspected"); + assert!( + matches!(error, AppError::Config(_)), + "rename should surface the live parse error, got {error:?}" + ); + assert_eq!( + provider_snapshot(state, &app_name, &source_id), + source_before, + "source aggregate must remain unchanged" + ); + assert_eq!( + provider_snapshot(state, &app_name, &target_id), + target_before, + "target aggregate must remain unchanged" + ); + assert_eq!( + fs::read_to_string(&live_path).expect("reread malformed live config"), + malformed_live, + "failed rename must not rewrite the live config" + ); + } + }); + } + + #[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() { @@ -450,7 +1173,13 @@ mod tests { None, ); - ProviderService::add(state, AppType::Codex, provider, false).expect("add provider"); + ProviderService::add( + state, + AppType::Codex, + provider_to_mutation_input(provider), + false, + ) + .expect("add provider"); let saved = state .db @@ -482,14 +1211,19 @@ mod tests { ); state .db - .save_provider(AppType::Codex.as_str(), &provider) + .reconcile_provider_fixture(AppType::Codex.as_str(), &provider) .expect("seed provider with explicit usage credentials"); let mut updated = provider.clone(); updated.settings_config = codex_settings("https://api.b.example/v1/", "sk-b"); - ProviderService::update(state, AppType::Codex, None, updated) - .expect("update provider main credentials"); + ProviderService::update( + state, + AppType::Codex, + None, + provider_to_mutation_input(updated), + ) + .expect("update provider main credentials"); let saved = state .db @@ -527,8 +1261,13 @@ mod tests { None, ); - ProviderService::add(state, AppType::Codex, copied_provider, false) - .expect("add copied provider"); + ProviderService::add( + state, + AppType::Codex, + provider_to_mutation_input(copied_provider), + false, + ) + .expect("add copied provider"); let saved_after_add = state .db @@ -546,8 +1285,13 @@ mod tests { let mut edited_provider = saved_after_add.clone(); edited_provider.settings_config = codex_settings("https://api.b.example/v1/", "sk-b"); - ProviderService::update(state, AppType::Codex, None, edited_provider) - .expect("edit copied provider credentials"); + ProviderService::update( + state, + AppType::Codex, + None, + provider_to_mutation_input(edited_provider), + ) + .expect("edit copied provider credentials"); let saved_after_update = state .db @@ -583,7 +1327,7 @@ mod tests { ); state .db - .save_provider(AppType::Codex.as_str(), &provider) + .reconcile_provider_fixture(AppType::Codex.as_str(), &provider) .expect("seed provider with distinct usage credentials"); let mut updated = provider.clone(); @@ -597,8 +1341,13 @@ mod tests { ..Default::default() }); - ProviderService::update(state, AppType::Codex, None, updated) - .expect("update provider with redundant usage credentials"); + ProviderService::update( + state, + AppType::Codex, + None, + provider_to_mutation_input(updated), + ) + .expect("update provider with redundant usage credentials"); let saved = state .db @@ -629,7 +1378,13 @@ mod tests { None, ); - ProviderService::add(state, AppType::Codex, provider, false).expect("add provider"); + ProviderService::add( + state, + AppType::Codex, + provider_to_mutation_input(provider), + false, + ) + .expect("add provider"); let saved = state .db @@ -663,7 +1418,13 @@ mod tests { Some("token_plan"), ); - ProviderService::add(state, AppType::Codex, provider, false).expect("add provider"); + ProviderService::add( + state, + AppType::Codex, + provider_to_mutation_input(provider), + false, + ) + .expect("add provider"); let saved = state .db @@ -810,7 +1571,8 @@ mod tests { }}), None, ); - db.save_provider("gemini", &victim).expect("save victim"); + db.reconcile_provider_fixture("gemini", &victim) + .expect("save victim"); // 供应商 C:自己写了同名键但值不同,不能被误删 let unrelated = Provider::with_id( @@ -822,7 +1584,8 @@ mod tests { }}), None, ); - db.save_provider("gemini", &unrelated).expect("save c"); + db.reconcile_provider_fixture("gemini", &unrelated) + .expect("save c"); } #[tokio::test] @@ -1469,7 +2232,7 @@ command = "legacy-cmd" }), None, ); - db.save_provider("claude", &original) + db.reconcile_provider_fixture("claude", &original) .expect("save provider"); db.set_current_provider("claude", "p1") .expect("set current provider"); @@ -1527,8 +2290,13 @@ command = "legacy-cmd" None, ); - ProviderService::update(&state, AppType::Claude, None, updated.clone()) - .expect("update current provider"); + ProviderService::update( + &state, + AppType::Claude, + None, + provider_to_mutation_input(updated.clone()), + ) + .expect("update current provider"); let backup = db .get_live_backup("claude") @@ -1604,7 +2372,8 @@ requires_openai_auth = true api_format: Some("openai_responses".into()), ..Default::default() }); - db.save_provider("codex", &original).expect("save provider"); + db.reconcile_provider_fixture("codex", &original) + .expect("save provider"); db.set_current_provider("codex", "p1") .expect("set current provider"); crate::settings::set_current_provider(&AppType::Codex, Some("p1")) @@ -1667,8 +2436,13 @@ requires_openai_auth = true "models": [{ "model": "gpt-5.4", "displayName": "GPT 5.4" }] }); - ProviderService::update(&state, AppType::Codex, None, updated.clone()) - .expect("update current Codex provider mapping"); + ProviderService::update( + &state, + AppType::Codex, + None, + provider_to_mutation_input(updated.clone()), + ) + .expect("update current Codex provider mapping"); let catalog_path = crate::codex_config::get_codex_model_catalog_path(); let catalog: Value = read_json_file(&catalog_path).expect("read generated catalog"); @@ -1683,8 +2457,13 @@ requires_openai_auth = true assert!(live_config.contains("model_catalog_json")); updated.settings_config["modelCatalog"] = json!({ "models": [] }); - ProviderService::update(&state, AppType::Codex, None, updated) - .expect("remove current Codex provider mapping"); + ProviderService::update( + &state, + AppType::Codex, + None, + provider_to_mutation_input(updated), + ) + .expect("remove current Codex provider mapping"); let live_config = fs::read_to_string(crate::codex_config::get_codex_config_path()) .expect("read Codex config.toml after mapping removal"); @@ -1734,7 +2513,7 @@ requires_openai_auth = true )]), ..Default::default() }); - db.save_provider("claude-desktop", &original) + db.reconcile_provider_fixture("claude-desktop", &original) .expect("save provider"); db.set_current_provider("claude-desktop", "p1") .expect("set current provider"); @@ -1821,8 +2600,13 @@ requires_openai_auth = true fn rename_rejects_missing_original_provider() { with_test_home(|state, _| { let original = openclaw_provider("deepseek"); - ProviderService::add(state, AppType::OpenClaw, original.clone(), false) - .expect("seed db-only provider"); + ProviderService::add( + state, + AppType::OpenClaw, + provider_to_mutation_input(original.clone()), + false, + ) + .expect("seed db-only provider"); let mut renamed = original.clone(); renamed.id = "deepseek-copy".to_string(); @@ -1831,7 +2615,7 @@ requires_openai_auth = true state, AppType::OpenClaw, Some("missing-provider"), - renamed, + provider_to_mutation_input(renamed), ) .expect_err("stale originalId should be rejected"); @@ -1855,8 +2639,13 @@ requires_openai_auth = true fn db_only_additive_update_survives_live_config_parse_errors() { with_test_home(|state, home| { let provider = openclaw_provider("deepseek"); - ProviderService::add(state, AppType::OpenClaw, provider.clone(), false) - .expect("seed db-only provider"); + ProviderService::add( + state, + AppType::OpenClaw, + provider_to_mutation_input(provider.clone()), + false, + ) + .expect("seed db-only provider"); let stored = state .db @@ -1881,8 +2670,13 @@ requires_openai_auth = true updated.name = "DeepSeek Edited".to_string(); updated.meta.get_or_insert_with(ProviderMeta::default); - ProviderService::update(state, AppType::OpenClaw, None, updated) - .expect("db-only update should ignore live parse errors"); + ProviderService::update( + state, + AppType::OpenClaw, + None, + provider_to_mutation_input(updated), + ) + .expect("db-only update should ignore live parse errors"); let saved = state .db @@ -1898,8 +2692,13 @@ requires_openai_auth = true fn sync_current_provider_for_app_skips_db_only_opencode_provider() { with_test_home(|state, _| { let provider = opencode_provider("db-only-opencode"); - ProviderService::add(state, AppType::OpenCode, provider.clone(), false) - .expect("seed db-only opencode provider"); + ProviderService::add( + state, + AppType::OpenCode, + provider_to_mutation_input(provider.clone()), + false, + ) + .expect("seed db-only opencode provider"); ProviderService::sync_current_provider_for_app(state, AppType::OpenCode) .expect("sync additive opencode providers"); @@ -1918,8 +2717,13 @@ requires_openai_auth = true fn sync_current_provider_for_app_skips_db_only_openclaw_provider() { with_test_home(|state, _| { let provider = openclaw_provider("db-only-openclaw"); - ProviderService::add(state, AppType::OpenClaw, provider.clone(), false) - .expect("seed db-only openclaw provider"); + ProviderService::add( + state, + AppType::OpenClaw, + provider_to_mutation_input(provider.clone()), + false, + ) + .expect("seed db-only openclaw provider"); ProviderService::sync_current_provider_for_app(state, AppType::OpenClaw) .expect("sync additive openclaw providers"); @@ -1942,14 +2746,14 @@ requires_openai_auth = true .expect("seed opencode live provider"); state .db - .save_provider(AppType::OpenCode.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenCode.as_str(), &provider) .expect("seed legacy opencode provider in db"); let mut updated = provider.clone(); updated.settings_config["options"]["apiKey"] = Value::String("updated-key".to_string()); state .db - .save_provider(AppType::OpenCode.as_str(), &updated) + .reconcile_provider_fixture(AppType::OpenCode.as_str(), &updated) .expect("update legacy opencode provider in db"); ProviderService::sync_current_provider_for_app(state, AppType::OpenCode) @@ -1975,7 +2779,7 @@ requires_openai_auth = true let provider = opencode_provider("legacy-opencode-reset"); state .db - .save_provider(AppType::OpenCode.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenCode.as_str(), &provider) .expect("seed legacy opencode provider in db"); ProviderService::sync_current_provider_for_app(state, AppType::OpenCode) @@ -2003,7 +2807,7 @@ requires_openai_auth = true ]); state .db - .save_provider(AppType::OpenClaw.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenClaw.as_str(), &provider) .expect("seed legacy openclaw provider in db"); ProviderService::sync_current_provider_for_app(state, AppType::OpenClaw) @@ -2053,7 +2857,7 @@ requires_openai_auth = true let provider = opencode_provider("existing-opencode"); state .db - .save_provider(AppType::OpenCode.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenCode.as_str(), &provider) .expect("seed existing opencode provider"); let mut live_settings = provider.settings_config.clone(); @@ -2127,7 +2931,7 @@ requires_openai_auth = true ]); state .db - .save_provider(AppType::OpenClaw.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenClaw.as_str(), &provider) .expect("seed existing openclaw provider"); let mut live_settings = provider.settings_config.clone(); @@ -2164,7 +2968,7 @@ requires_openai_auth = true let provider = hermes_provider("existing-hermes"); state .db - .save_provider(AppType::Hermes.as_str(), &provider) + .reconcile_provider_fixture(AppType::Hermes.as_str(), &provider) .expect("seed existing hermes provider"); let mut live_settings = provider.settings_config.clone(); @@ -2204,7 +3008,7 @@ requires_openai_auth = true let provider = openclaw_provider("legacy-provider"); state .db - .save_provider(AppType::OpenClaw.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenClaw.as_str(), &provider) .expect("seed legacy provider without live_config_managed marker"); let openclaw_dir = home.join(".openclaw"); @@ -2215,8 +3019,13 @@ requires_openai_auth = true let mut updated = provider.clone(); updated.name = "Legacy Edited".to_string(); - let err = ProviderService::update(state, AppType::OpenClaw, None, updated) - .expect_err("legacy providers should still surface live parse errors"); + let err = ProviderService::update( + state, + AppType::OpenClaw, + None, + provider_to_mutation_input(updated), + ) + .expect_err("legacy providers should still surface live parse errors"); assert!( err.to_string().contains("Failed to parse OpenClaw config"), "expected parse error, got {err:?}" @@ -2232,7 +3041,7 @@ requires_openai_auth = true let provider = opencode_omo_provider(&format!("{category}-provider"), category); state .db - .save_provider(AppType::OpenCode.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenCode.as_str(), &provider) .unwrap_or_else(|err| panic!("seed {category} provider: {err}")); let mut updated = provider.clone(); @@ -2240,8 +3049,13 @@ requires_openai_auth = true updated.settings_config["agents"]["writer"]["model"] = Value::String(format!("{category}-next-model")); - ProviderService::update(state, AppType::OpenCode, None, updated) - .unwrap_or_else(|err| panic!("update {category} provider: {err}")); + ProviderService::update( + state, + AppType::OpenCode, + None, + provider_to_mutation_input(updated), + ) + .unwrap_or_else(|err| panic!("update {category} provider: {err}")); let saved = state .db @@ -2267,7 +3081,7 @@ requires_openai_auth = true let provider = opencode_omo_provider(&format!("{category}-current"), category); state .db - .save_provider(AppType::OpenCode.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenCode.as_str(), &provider) .unwrap_or_else(|err| panic!("seed current {category} provider: {err}")); state .db @@ -2281,8 +3095,13 @@ requires_openai_auth = true updated.settings_config["otherFields"]["theme"] = Value::String(format!("{category}-light")); - ProviderService::update(state, AppType::OpenCode, None, updated) - .unwrap_or_else(|err| panic!("update current {category} provider: {err}")); + ProviderService::update( + state, + AppType::OpenCode, + None, + provider_to_mutation_input(updated), + ) + .unwrap_or_else(|err| panic!("update current {category} provider: {err}")); let saved = state .db @@ -2317,7 +3136,7 @@ requires_openai_auth = true let provider = opencode_omo_provider("omo-current", "omo"); state .db - .save_provider(AppType::OpenCode.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenCode.as_str(), &provider) .unwrap_or_else(|err| panic!("seed current omo provider: {err}")); state .db @@ -2334,8 +3153,13 @@ requires_openai_auth = true updated.settings_config["agents"]["writer"]["model"] = Value::String("omo-saved-model".to_string()); - ProviderService::update(state, AppType::OpenCode, None, updated) - .expect_err("update should fail when current omo file write fails"); + ProviderService::update( + state, + AppType::OpenCode, + None, + provider_to_mutation_input(updated), + ) + .expect_err("update should fail when current omo file write fails"); let saved = state .db @@ -2359,7 +3183,7 @@ requires_openai_auth = true let provider = opencode_omo_provider("omo-current", "omo"); state .db - .save_provider(AppType::OpenCode.as_str(), &provider) + .reconcile_provider_fixture(AppType::OpenCode.as_str(), &provider) .unwrap_or_else(|err| panic!("seed current omo provider: {err}")); state .db @@ -2393,8 +3217,13 @@ requires_openai_auth = true updated.settings_config["otherFields"]["theme"] = Value::String("omo-light".to_string()); - ProviderService::update(state, AppType::OpenCode, None, updated) - .expect_err("update should fail when plugin sync fails"); + ProviderService::update( + state, + AppType::OpenCode, + None, + provider_to_mutation_input(updated), + ) + .expect_err("update should fail when plugin sync fails"); let saved = state .db @@ -2457,6 +3286,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); @@ -2551,10 +3412,11 @@ impl ProviderService { pub fn add( state: &AppState, app_type: AppType, - provider: Provider, + input: ProviderMutationInput, add_to_live: bool, ) -> Result { - let mut provider = provider; + 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); Self::validate_provider_settings(&app_type, &provider)?; @@ -2564,8 +3426,12 @@ impl ProviderService { Self::set_provider_live_config_managed(&mut provider, add_to_live); } - // Save to database - state.db.save_provider(app_type.as_str(), &provider)?; + // Strict create owns both the provider row and initial endpoints. + create_provider_record( + state, + &app_type, + provider_to_mutation_input(provider.clone()), + )?; // Additive mode apps (OpenCode, OpenClaw): optionally write to live config. if app_type.is_additive_mode() { @@ -2602,9 +3468,13 @@ impl ProviderService { state: &AppState, app_type: AppType, original_id: Option<&str>, - provider: Provider, + input: ProviderMutationInput, ) -> Result { - let mut provider = provider; + // 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; let existing_provider = state @@ -2646,11 +3516,12 @@ impl ProviderService { )); } - let original_in_live = Self::check_live_config_exists( - &app_type, - &original_id, - Self::provider_live_config_managed(&existing_provider), - )?; + // A rename changes durable identity, so "cannot inspect live + // config" must never be treated as "not live". DB-only + // same-ID edits deliberately retain their tolerant path below, + // but rename proves both identities absent from a readable live + // config before committing the SQLite transaction. + let original_in_live = provider_exists_in_live_config(&app_type, &original_id)?; if original_in_live { return Err(AppError::Message( "Provider key cannot be changed after the provider has been added to the app config" @@ -2658,11 +3529,7 @@ impl ProviderService { )); } - let next_id_in_live = Self::check_live_config_exists( - &app_type, - &provider.id, - Self::provider_live_config_managed(&existing_provider), - )?; + let next_id_in_live = provider_exists_in_live_config(&app_type, &provider.id)?; if state .db .get_provider_by_id(&provider.id, app_type.as_str())? @@ -2677,8 +3544,13 @@ impl ProviderService { } Self::set_provider_live_config_managed(&mut provider, false); - state.db.save_provider(app_type.as_str(), &provider)?; - state.db.delete_provider(app_type.as_str(), &original_id)?; + let source = ProviderKey::new(app_type.as_str(), original_id.clone())?; + state + .db + .rename_db_only_additive_provider(RenameProvider::from_input( + source, + &provider_to_mutation_input(provider.clone()), + )?)?; if crate::settings::get_current_provider(&app_type).as_deref() == Some(&original_id) { crate::settings::set_current_provider(&app_type, Some(provider.id.as_str()))?; @@ -2708,7 +3580,11 @@ impl ProviderService { if is_current { crate::services::OmoService::write_provider_config_to_file(&provider, variant)?; } - if let Err(err) = state.db.save_provider(app_type.as_str(), &provider) { + if let Err(err) = update_provider_record( + state, + &app_type, + &provider_to_mutation_input(provider.clone()), + ) { if is_current { if let Err(rollback_err) = crate::services::OmoService::write_config_to_file(state, variant) @@ -2737,7 +3613,11 @@ impl ProviderService { // Save to database after live-config presence is resolved so parse errors // do not report failure after already mutating DB state. - state.db.save_provider(app_type.as_str(), &provider)?; + update_provider_record( + state, + &app_type, + &provider_to_mutation_input(provider.clone()), + )?; if !live_config_managed { return Ok(true); @@ -2747,7 +3627,11 @@ impl ProviderService { } // Save to database - state.db.save_provider(app_type.as_str(), &provider)?; + update_provider_record( + state, + &app_type, + &provider_to_mutation_input(provider.clone()), + )?; // For other apps: Check if this is current provider (use effective current, not just DB) let effective_current = @@ -2829,6 +3713,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. @@ -2900,6 +3785,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 @@ -2943,9 +3829,12 @@ impl ProviderService { } } - if let Some(mut provider) = state.db.get_provider_by_id(id, app_type.as_str())? { - Self::set_provider_live_config_managed(&mut provider, false); - state.db.save_provider(app_type.as_str(), &provider)?; + if state + .db + .get_provider_aggregate(app_type.as_str(), id)? + .is_some() + { + Self::persist_live_config_managed(state, &app_type, id, false)?; } Ok(()) @@ -2964,6 +3853,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 @@ -2986,21 +3890,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. @@ -3098,7 +3987,12 @@ impl ProviderService { if !app_type.is_additive_mode() { // Only backfill when switching to a different provider if let Ok(live_config) = read_live_settings(app_type.clone()) { - if let Some(mut current_provider) = providers.get(¤t_id).cloned() { + if let Some(mut current_provider) = state + .db + .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 的文档。 @@ -3117,9 +4011,13 @@ impl ProviderService { ¤t_provider, live_config, ); - if let Err(e) = - state.db.save_provider(app_type.as_str(), ¤t_provider) - { + remove_hydrated_endpoints_from_row_update(&mut current_provider); + if let Err(e) = update_provider_record_if_unchanged( + state, + &app_type, + fingerprint, + provider_to_mutation_input(current_provider), + ) { log::warn!("Backfill failed: {e}"); result .warnings @@ -3196,9 +4094,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 = provider.clone(); - Self::set_provider_live_config_managed(&mut updated, true); - if let Err(e) = state.db.save_provider(app_type.as_str(), &updated) { + 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), @@ -3246,6 +4143,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); } @@ -3300,9 +4198,10 @@ impl ProviderService { return Ok(()); } - let providers = state.db.get_all_providers(app_type.as_str())?; + let providers = state.db.get_all_provider_aggregates(app_type.as_str())?; - for provider in providers.values() { + for aggregate in providers.values() { + let provider = &aggregate.provider; if provider .meta .as_ref() @@ -3316,6 +4215,7 @@ impl ProviderService { continue; } + let fingerprint = provider_row_fingerprint(provider); let mut updated_provider = provider.clone(); updated_provider .meta @@ -3337,9 +4237,13 @@ impl ProviderService { } } - state - .db - .save_provider(app_type.as_str(), &updated_provider)?; + remove_hydrated_endpoints_from_row_update(&mut updated_provider); + update_provider_record_if_unchanged( + state, + &app_type, + fingerprint, + provider_to_mutation_input(updated_provider), + )?; } Ok(()) @@ -3855,9 +4759,11 @@ impl ProviderService { .map_err(|e| AppError::Message(format!("Serialization failed: {e}")))?; // 1) 先算出各供应商清理后的配置,但**先不落库** - let providers = state.db.get_all_providers(app.as_str())?; - let mut pending: Vec<(String, Provider, Value)> = Vec::new(); - for (id, provider) in providers { + let providers = state.db.get_all_provider_aggregates(app.as_str())?; + 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, @@ -3870,7 +4776,7 @@ impl ProviderService { } }; if cleaned != provider.settings_config { - pending.push((id, provider, cleaned)); + pending.push((id, provider, cleaned, fingerprint)); } } @@ -3899,7 +4805,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), })) @@ -3907,7 +4813,7 @@ impl ProviderService { }); let audit_text = serde_json::to_string(&audit) .map_err(|e| AppError::Message(format!("Serialization failed: {e}")))?; - // 只在没有记录时写。provider 的写入不是一个事务(每次 save_provider 各自 + // 只在没有记录时写。provider 的写入不是一个事务(每次类型化行更新各自 // 提交),上一轮可能改到一半就中止;此时完成标记没置位,下次启动会重跑, // 而重跑看到的"原始状态"已经残缺。无条件 INSERT OR REPLACE 会拿这份残缺 // 记录盖掉第一轮那份完整的。 @@ -3916,10 +4822,11 @@ impl ProviderService { } // 3) 各供应商 settings_config:按值相等定向删除扩散出去的副本 - for (id, provider, cleaned) in pending { - let mut updated = provider; - updated.settings_config = cleaned; - state.db.save_provider(app.as_str(), &updated)?; + 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}' 中清除泄漏的共享凭据"); } @@ -4101,13 +5008,11 @@ impl ProviderService { app_type: AppType, updates: Vec, ) -> Result { - let mut providers = state.db.get_all_providers(app_type.as_str())?; - for update in updates { - if let Some(provider) = providers.get_mut(&update.id) { - provider.sort_index = Some(update.sort_index); - state.db.save_provider(app_type.as_str(), provider)?; - } + let key = ProviderKey::new(app_type.as_str(), update.id)?; + state + .db + .update_provider_sort_index(&key, update.sort_index)?; } Ok(true) @@ -4655,12 +5560,23 @@ impl ProviderService { // 同步到 Claude if let Some(mut claude_provider) = provider.to_claude_provider() { // 合并已有配置 - if let Some(existing) = state.db.get_provider_by_id(&claude_provider.id, "claude")? { + let precondition = if let Some(existing) = + state.db.get_provider_by_id(&claude_provider.id, "claude")? + { + let fingerprint = provider_row_fingerprint(&existing); let mut merged = existing.settings_config.clone(); Self::merge_json(&mut merged, &claude_provider.settings_config); claude_provider.settings_config = merged; - } - state.db.save_provider("claude", &claude_provider)?; + ReconcilePrecondition::ExpectPresent { fingerprint } + } else { + ReconcilePrecondition::ExpectAbsent + }; + reconcile_provider_record_with_precondition( + &state.db, + "claude", + provider_to_mutation_input(claude_provider), + precondition, + )?; } else { // 如果禁用了 Claude,删除对应的子供应商 let claude_id = format!("universal-claude-{id}"); @@ -4670,12 +5586,22 @@ impl ProviderService { // 同步到 Codex if let Some(mut codex_provider) = provider.to_codex_provider() { // 合并已有配置 - if let Some(existing) = state.db.get_provider_by_id(&codex_provider.id, "codex")? { - let mut merged = existing.settings_config.clone(); - Self::merge_json(&mut merged, &codex_provider.settings_config); - codex_provider.settings_config = merged; - } - state.db.save_provider("codex", &codex_provider)?; + let precondition = + if let Some(existing) = state.db.get_provider_by_id(&codex_provider.id, "codex")? { + let fingerprint = provider_row_fingerprint(&existing); + let mut merged = existing.settings_config.clone(); + Self::merge_json(&mut merged, &codex_provider.settings_config); + codex_provider.settings_config = merged; + ReconcilePrecondition::ExpectPresent { fingerprint } + } else { + ReconcilePrecondition::ExpectAbsent + }; + reconcile_provider_record_with_precondition( + &state.db, + "codex", + provider_to_mutation_input(codex_provider), + precondition, + )?; } else { let codex_id = format!("universal-codex-{id}"); let _ = state.db.delete_provider("codex", &codex_id); @@ -4684,12 +5610,23 @@ impl ProviderService { // 同步到 Gemini if let Some(mut gemini_provider) = provider.to_gemini_provider() { // 合并已有配置 - if let Some(existing) = state.db.get_provider_by_id(&gemini_provider.id, "gemini")? { + let precondition = if let Some(existing) = + state.db.get_provider_by_id(&gemini_provider.id, "gemini")? + { + let fingerprint = provider_row_fingerprint(&existing); let mut merged = existing.settings_config.clone(); Self::merge_json(&mut merged, &gemini_provider.settings_config); gemini_provider.settings_config = merged; - } - state.db.save_provider("gemini", &gemini_provider)?; + ReconcilePrecondition::ExpectPresent { fingerprint } + } else { + ReconcilePrecondition::ExpectAbsent + }; + reconcile_provider_record_with_precondition( + &state.db, + "gemini", + provider_to_mutation_input(gemini_provider), + precondition, + )?; } else { let gemini_id = format!("universal-gemini-{id}"); let _ = state.db.delete_provider("gemini", &gemini_id); diff --git a/src-tauri/src/services/proxy.rs b/src-tauri/src/services/proxy.rs index e1c870c11..9dcaaa392 100644 --- a/src-tauri/src/services/proxy.rs +++ b/src-tauri/src/services/proxy.rs @@ -10,7 +10,9 @@ 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, 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; @@ -980,6 +982,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, @@ -992,91 +1015,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 Err(e) = self.db.update_provider_settings_config( - "claude", - &provider_id, - &provider.settings_config, - ) { - 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})"); } } } @@ -1087,55 +1108,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 Err(e) = self.db.update_provider_settings_config( - "codex", - &provider_id, - &provider.settings_config, - ) { - 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,51 +1167,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 Err(e) = self.db.update_provider_settings_config( - "gemini", - &provider_id, - &provider.settings_config, - ) { - 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})"); } } } @@ -1199,38 +1220,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); - self.db - .update_provider_settings_config( - "grokbuild", - &provider_id, - &provider.settings_config, - ) + 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 Token 到数据库失败: {e}") + format!("更新 Grok Build API Key 失败: {e}") })?; - } + provider.settings_config["config"] = json!(updated); + self.persist_synced_live_token( + "grokbuild", + &provider_id, + observed_fingerprint, + provider, + )?; } } } @@ -3813,7 +3837,7 @@ mod tests { }), None, ); - db.save_provider("claude", &provider) + db.reconcile_provider_fixture("claude", &provider) .expect("save provider"); db.set_current_provider("claude", "p1") .expect("set db current provider"); @@ -3999,7 +4023,7 @@ wire_api = "responses" None, ); provider.category = Some("cn_official".to_string()); - db.save_provider("codex", &provider) + db.reconcile_provider_fixture("codex", &provider) .expect("save DeepSeek provider"); db.set_current_provider("codex", "deepseek") .expect("set current provider"); @@ -4085,7 +4109,7 @@ wire_api = "responses" None, ); provider.category = Some("official".to_string()); - db.save_provider("codex", &provider) + db.reconcile_provider_fixture("codex", &provider) .expect("save misclassified DeepSeek provider"); db.set_current_provider("codex", "deepseek") .expect("set current provider"); @@ -4146,7 +4170,7 @@ wire_api = "responses" None, ); official.category = Some("official".to_string()); - db.save_provider("codex", &official) + db.reconcile_provider_fixture("codex", &official) .expect("save official provider"); let mut third_party = Provider::with_id( @@ -4165,7 +4189,7 @@ wire_api = "responses" None, ); third_party.category = Some("custom".to_string()); - db.save_provider("codex", &third_party) + db.reconcile_provider_fixture("codex", &third_party) .expect("save third-party provider"); db.set_current_provider("codex", "codex-official") .expect("set current provider"); @@ -4314,7 +4338,8 @@ wire_api = "responses" None, ); official.category = Some("official".to_string()); - db.save_provider("codex", &official).expect("save official"); + db.reconcile_provider_fixture("codex", &official) + .expect("save official"); db.set_current_provider("codex", crate::database::CODEX_OFFICIAL_PROVIDER_ID) .expect("set current"); crate::settings::set_current_provider( @@ -4392,7 +4417,7 @@ wire_api = "responses" None, ); provider.category = Some("official".to_string()); - db.save_provider("codex", &provider) + db.reconcile_provider_fixture("codex", &provider) .expect("save misclassified DeepSeek provider"); db.set_current_provider("codex", "deepseek") .expect("set current provider"); @@ -4472,7 +4497,7 @@ wire_api = "responses" None, ); provider.category = Some("official".to_string()); - db.save_provider("codex", &provider) + db.reconcile_provider_fixture("codex", &provider) .expect("save misclassified DeepSeek provider"); db.set_current_provider("codex", "deepseek") .expect("set current provider"); @@ -4584,7 +4609,7 @@ wire_api = "responses" None, ); provider.category = Some("official".to_string()); - db.save_provider("codex", &provider) + db.reconcile_provider_fixture("codex", &provider) .expect("save misclassified DeepSeek provider"); db.set_current_provider("codex", "deepseek") .expect("set current provider"); @@ -4702,7 +4727,7 @@ wire_api = "responses" None, ); provider.category = Some("official".to_string()); - db.save_provider("codex", &provider) + db.reconcile_provider_fixture("codex", &provider) .expect("save misclassified DeepSeek provider"); db.set_current_provider("codex", "deepseek") .expect("set current provider"); @@ -4838,7 +4863,7 @@ wire_api = "responses" None, ); provider.category = Some("cn_official".to_string()); - db.save_provider("codex", &provider) + db.reconcile_provider_fixture("codex", &provider) .expect("save DeepSeek provider"); db.set_current_provider("codex", "deepseek") .expect("set current provider"); @@ -5315,7 +5340,7 @@ model = "gpt-5.1-codex" }), None, ); - db.save_provider("claude", &provider) + db.reconcile_provider_fixture("claude", &provider) .expect("save provider"); db.set_current_provider("claude", "p1") .expect("set current provider"); @@ -5371,7 +5396,7 @@ model = "gpt-5.1-codex" }), None, ); - db.save_provider("claude", &provider) + db.reconcile_provider_fixture("claude", &provider) .expect("save provider"); db.set_current_provider("claude", "p1") .expect("set current provider"); @@ -5407,6 +5432,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() { @@ -5436,9 +5509,9 @@ model = "gpt-5.1-codex" }), None, ); - db.save_provider("claude", &provider_a) + db.reconcile_provider_fixture("claude", &provider_a) .expect("save provider a"); - db.save_provider("claude", &provider_b) + db.reconcile_provider_fixture("claude", &provider_b) .expect("save provider b"); db.set_current_provider("claude", "a") .expect("set current provider"); @@ -5511,9 +5584,9 @@ model = "gpt-5.1-codex" None, ); - db.save_provider("claude", &provider_a) + db.reconcile_provider_fixture("claude", &provider_a) .expect("save provider a"); - db.save_provider("claude", &provider_b) + db.reconcile_provider_fixture("claude", &provider_b) .expect("save provider b"); db.set_current_provider("claude", "a") .expect("set current provider"); @@ -5662,11 +5735,11 @@ model = "gpt-5.1-codex" None, ); - db.save_provider("claude", &provider_a) + db.reconcile_provider_fixture("claude", &provider_a) .expect("save provider a"); - db.save_provider("claude", &provider_b) + db.reconcile_provider_fixture("claude", &provider_b) .expect("save provider b"); - db.save_provider("claude", &provider_c) + db.reconcile_provider_fixture("claude", &provider_c) .expect("save provider c"); db.set_current_provider("claude", "a") .expect("set current provider"); @@ -5749,9 +5822,9 @@ model = "gpt-5.1-codex" None, ); - db.save_provider("claude", &provider_a) + db.reconcile_provider_fixture("claude", &provider_a) .expect("save provider a"); - db.save_provider("claude", &provider_b) + db.reconcile_provider_fixture("claude", &provider_b) .expect("save provider b"); db.set_current_provider("claude", "a") .expect("set current provider"); @@ -6049,9 +6122,9 @@ requires_openai_auth = true None, ); - db.save_provider("codex", &provider_a) + db.reconcile_provider_fixture("codex", &provider_a) .expect("save provider a"); - db.save_provider("codex", &provider_b) + db.reconcile_provider_fixture("codex", &provider_b) .expect("save provider b"); db.set_current_provider("codex", "a") .expect("set current provider"); @@ -6224,9 +6297,9 @@ requires_openai_auth = true ..Default::default() }); - db.save_provider("codex", &provider_a) + db.reconcile_provider_fixture("codex", &provider_a) .expect("save provider a"); - db.save_provider("codex", &provider_b) + db.reconcile_provider_fixture("codex", &provider_b) .expect("save provider b"); db.set_current_provider("codex", "a") .expect("set current provider"); @@ -6468,9 +6541,9 @@ requires_openai_auth = true ..Default::default() }); - db.save_provider("codex", &provider_a) + db.reconcile_provider_fixture("codex", &provider_a) .expect("save provider a"); - db.save_provider("codex", &provider_b) + db.reconcile_provider_fixture("codex", &provider_b) .expect("save provider b"); db.set_current_provider("codex", "a") .expect("set current provider a"); @@ -6604,9 +6677,9 @@ requires_openai_auth = true None, ); - db.save_provider("codex", &provider_a) + db.reconcile_provider_fixture("codex", &provider_a) .expect("save provider a"); - db.save_provider("codex", &provider_b) + db.reconcile_provider_fixture("codex", &provider_b) .expect("save provider b"); db.set_current_provider("codex", "a") .expect("set current provider a"); @@ -6686,9 +6759,9 @@ requires_openai_auth = true }), None, ); - db.save_provider("codex", &provider_a) + db.reconcile_provider_fixture("codex", &provider_a) .expect("save provider a"); - db.save_provider("codex", &provider_b) + db.reconcile_provider_fixture("codex", &provider_b) .expect("save provider b"); db.set_current_provider("codex", "a") .expect("set current provider a"); @@ -6970,7 +7043,7 @@ requires_openai_auth = true }), None, ); - db.save_provider("claude", &provider) + db.reconcile_provider_fixture("claude", &provider) .expect("save provider"); db.set_current_provider("claude", "p1") .expect("set current provider"); @@ -7224,9 +7297,9 @@ experimental_bearer_token = "PROXY_MANAGED" grok_provider_config("https://b.example.com/v1", "b-key"), None, ); - db.save_provider("grokbuild", &provider_a) + db.reconcile_provider_fixture("grokbuild", &provider_a) .expect("save provider a"); - db.save_provider("grokbuild", &provider_b) + db.reconcile_provider_fixture("grokbuild", &provider_b) .expect("save provider b"); db.set_current_provider("grokbuild", "grok-a") .expect("set db current"); @@ -7291,9 +7364,9 @@ experimental_bearer_token = "PROXY_MANAGED" json!({ "config": "not valid toml = [" }), None, ); - db.save_provider("grokbuild", &provider_a) + db.reconcile_provider_fixture("grokbuild", &provider_a) .expect("save provider a"); - db.save_provider("grokbuild", &provider_b) + db.reconcile_provider_fixture("grokbuild", &provider_b) .expect("save provider b"); db.set_current_provider("grokbuild", "grok-a") .expect("set db current"); diff --git a/src-tauri/src/settings.rs b/src-tauri/src/settings.rs index 98ac3d7ea..30010ff85 100644 --- a/src-tauri/src/settings.rs +++ b/src-tauri/src/settings.rs @@ -8,11 +8,11 @@ use crate::error::AppError; use crate::services::skill::{SkillStorageLocation, SyncMethod}; /// 自定义端点配置(历史兼容,实际存储在 provider.meta.custom_endpoints) -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct CustomEndpoint { pub url: String, - pub added_at: i64, + pub added_at: Option, #[serde(skip_serializing_if = "Option::is_none")] pub last_used: Option, } diff --git a/src-tauri/tests/profile_roundtrip.rs b/src-tauri/tests/profile_roundtrip.rs index 1f85f4987..80d443313 100644 --- a/src-tauri/tests/profile_roundtrip.rs +++ b/src-tauri/tests/profile_roundtrip.rs @@ -7,13 +7,14 @@ use std::fs; use serde_json::json; use cc_switch_lib::{ - AppType, InstalledSkill, McpServer, McpService, ProfilePayload, ProfileScope, ProfileService, - Prompt, PromptService, Provider, ProviderService, SkillApps, SkillService, + AppType, InstalledSkill, McpServer, McpService, NewProviderAggregate, ProfilePayload, + ProfileScope, ProfileService, Prompt, PromptService, Provider, ProviderService, SkillApps, + SkillService, }; #[path = "support.rs"] mod support; -use support::{create_test_state, ensure_test_home, reset_test_fs, test_mutex}; +use support::{create_test_state, ensure_test_home, new_provider_input, reset_test_fs, test_mutex}; fn claude_provider(id: &str, token: &str) -> Provider { Provider::with_id( @@ -107,34 +108,41 @@ fn profile_snapshot_apply_roundtrip_restores_configuration() { let state = create_test_state().expect("create test state"); // ---- 种子数据:2 个 Claude 供应商(p1 为当前)+ 2 个 MCP + 1 个 Skill + 2 个 Prompt ---- - state - .db - .save_provider(AppType::Claude.as_str(), &claude_provider("p1", "key-1")) - .expect("save provider p1"); - state - .db - .save_provider(AppType::Claude.as_str(), &claude_provider("p2", "key-2")) - .expect("save provider p2"); + ProviderService::add( + &state, + AppType::Claude, + new_provider_input(claude_provider("p1", "key-1")), + false, + ) + .expect("create provider p1"); + ProviderService::add( + &state, + AppType::Claude, + new_provider_input(claude_provider("p2", "key-2")), + false, + ) + .expect("create provider p2"); state .db .set_current_provider(AppType::Claude.as_str(), "p1") .expect("set current provider p1"); // Claude Desktop 只有供应商一个活跃维度(MCP/Skills/Prompt 对它不适用) - state - .db - .save_provider( - AppType::ClaudeDesktop.as_str(), - &desktop_provider("d1", "dk-1"), - ) - .expect("save desktop provider d1"); - state - .db - .save_provider( - AppType::ClaudeDesktop.as_str(), - &desktop_provider("d2", "dk-2"), - ) - .expect("save desktop provider d2"); + for provider in [ + desktop_provider("d1", "dk-1"), + desktop_provider("d2", "dk-2"), + ] { + state + .db + .create_provider( + NewProviderAggregate::from_input( + AppType::ClaudeDesktop.as_str(), + new_provider_input(provider), + ) + .expect("build typed desktop create"), + ) + .expect("create desktop provider"); + } state .db .set_current_provider(AppType::ClaudeDesktop.as_str(), "d1") @@ -287,10 +295,13 @@ fn shared_profile_sides_are_isolated_and_mergeable() { let state = create_test_state().expect("create test state"); // 种子:Claude 侧有当前供应商 + 启用的 MCP - state - .db - .save_provider(AppType::Claude.as_str(), &claude_provider("p1", "key-1")) - .expect("save provider p1"); + ProviderService::add( + &state, + AppType::Claude, + new_provider_input(claude_provider("p1", "key-1")), + false, + ) + .expect("create provider p1"); state .db .set_current_provider(AppType::Claude.as_str(), "p1") @@ -496,14 +507,20 @@ fn switching_profile_autosaves_previous_profile_state() { let state = create_test_state().expect("create test state"); // ---- 种子:Claude 侧两套供应商 / MCP / Prompt ---- - state - .db - .save_provider(AppType::Claude.as_str(), &claude_provider("p1", "key-1")) - .expect("save provider p1"); - state - .db - .save_provider(AppType::Claude.as_str(), &claude_provider("p2", "key-2")) - .expect("save provider p2"); + ProviderService::add( + &state, + AppType::Claude, + new_provider_input(claude_provider("p1", "key-1")), + false, + ) + .expect("create provider p1"); + ProviderService::add( + &state, + AppType::Claude, + new_provider_input(claude_provider("p2", "key-2")), + false, + ) + .expect("create provider p2"); state .db .set_current_provider(AppType::Claude.as_str(), "p1") @@ -665,17 +682,13 @@ fn profile_switch_auto_disables_takeover_before_apply() { // ---- 两个 Claude 供应商:custom1 与 custom2 ---- let mut custom1 = claude_provider("custom1", "custom-key-1"); custom1.category = Some("custom".to_string()); - state - .db - .save_provider(AppType::Claude.as_str(), &custom1) - .expect("save custom1 provider"); + ProviderService::add(&state, AppType::Claude, new_provider_input(custom1), false) + .expect("create custom1 provider"); let mut custom2 = claude_provider("custom2", "custom-key-2"); custom2.category = Some("custom".to_string()); - state - .db - .save_provider(AppType::Claude.as_str(), &custom2) - .expect("save custom2 provider"); + ProviderService::add(&state, AppType::Claude, new_provider_input(custom2), false) + .expect("create custom2 provider"); // 初始状态:custom1 + 代理接管 ProviderService::switch(&state, AppType::Claude, "custom1").expect("switch to custom1"); @@ -757,20 +770,20 @@ fn claude_desktop_profile_scope_is_independent() { let state = create_test_state().expect("create test state"); - state - .db - .save_provider( - AppType::ClaudeDesktop.as_str(), - &desktop_provider("d1", "dk-1"), - ) - .expect("save desktop provider d1"); - state - .db - .save_provider( - AppType::ClaudeDesktop.as_str(), - &desktop_provider("d2", "dk-2"), - ) - .expect("save desktop provider d2"); + ProviderService::add( + &state, + AppType::ClaudeDesktop, + new_provider_input(desktop_provider("d1", "dk-1")), + false, + ) + .expect("create desktop provider d1"); + ProviderService::add( + &state, + AppType::ClaudeDesktop, + new_provider_input(desktop_provider("d2", "dk-2")), + false, + ) + .expect("create desktop provider d2"); state .db .set_current_provider(AppType::ClaudeDesktop.as_str(), "d1") diff --git a/src-tauri/tests/provider_commands.rs b/src-tauri/tests/provider_commands.rs index baf532419..f8a1345c7 100644 --- a/src-tauri/tests/provider_commands.rs +++ b/src-tauri/tests/provider_commands.rs @@ -12,7 +12,7 @@ mod support; use std::collections::HashMap; use support::{ create_test_state, create_test_state_with_config, enable_codex_official_auth_preservation, - ensure_test_home, reset_test_fs, test_mutex, + ensure_test_home, new_provider_input, reset_test_fs, test_mutex, }; fn settings_path(home: &Path) -> PathBuf { @@ -64,18 +64,18 @@ fn grokbuild_import_and_switch_write_live_config() { ); let next_config = grokbuild_config("Relay", "https://new.example/v1", "new-key"); - state - .db - .save_provider( - AppType::GrokBuild.as_str(), - &Provider::with_id( - "relay".to_string(), - "Relay".to_string(), - json!({ "config": next_config }), - None, - ), - ) - .expect("save second Grok Build provider"); + ProviderService::add( + &state, + AppType::GrokBuild, + new_provider_input(Provider::with_id( + "relay".to_string(), + "Relay".to_string(), + json!({ "config": next_config }), + None, + )), + false, + ) + .expect("create second Grok Build provider"); switch_provider_test_hook(&state, AppType::GrokBuild, "relay") .expect("switch Grok Build provider"); diff --git a/src-tauri/tests/provider_service.rs b/src-tauri/tests/provider_service.rs index 5cb09e0f5..ff00e7dd9 100644 --- a/src-tauri/tests/provider_service.rs +++ b/src-tauri/tests/provider_service.rs @@ -9,7 +9,7 @@ use cc_switch_lib::{ mod support; use support::{ create_test_state, create_test_state_with_config, enable_codex_official_auth_preservation, - ensure_test_home, reset_test_fs, test_mutex, + ensure_test_home, new_provider_input, reset_test_fs, test_mutex, }; fn sanitize_provider_name(name: &str) -> String { @@ -3084,10 +3084,8 @@ fn recover_from_crash_without_backup_cleans_placeholder_instead_of_writing_it_ba taken_over_live.clone(), None, ); - state - .db - .save_provider(AppType::Claude.as_str(), &provider) - .expect("save placeholder provider"); + ProviderService::add(&state, AppType::Claude, new_provider_input(provider), false) + .expect("create placeholder provider"); state .db .set_current_provider(AppType::Claude.as_str(), "default") diff --git a/src-tauri/tests/support.rs b/src-tauri/tests/support.rs index d9b2f0259..457a8e0af 100644 --- a/src-tauri/tests/support.rs +++ b/src-tauri/tests/support.rs @@ -1,7 +1,31 @@ use std::path::{Path, PathBuf}; use std::sync::{Arc, Mutex, OnceLock}; -use cc_switch_lib::{update_settings, AppSettings, AppState, Database, MultiAppConfig}; +use cc_switch_lib::{ + update_settings, AppSettings, AppState, Database, MultiAppConfig, Provider, + ProviderMutationInput, +}; + +/// Build the public write DTO explicitly for integration tests. Keeping this +/// conversion test-only avoids reintroducing a production `From` +/// path from hydrated read projections to provider mutations. +#[allow(dead_code)] +pub fn new_provider_input(provider: Provider) -> ProviderMutationInput { + ProviderMutationInput { + id: provider.id, + name: provider.name, + settings_config: provider.settings_config, + website_url: provider.website_url, + category: provider.category, + created_at: provider.created_at, + sort_index: provider.sort_index, + notes: provider.notes, + meta: provider.meta, + icon: provider.icon, + icon_color: provider.icon_color, + in_failover_queue: provider.in_failover_queue, + } +} /// 为测试设置隔离的 HOME 目录,避免污染真实用户数据。 pub fn ensure_test_home() -> &'static Path { diff --git a/src/components/providers/forms/ProviderForm.tsx b/src/components/providers/forms/ProviderForm.tsx index 323ba1a46..ffc0e05e2 100644 --- a/src/components/providers/forms/ProviderForm.tsx +++ b/src/components/providers/forms/ProviderForm.tsx @@ -1537,8 +1537,16 @@ function ProviderFormFull({ } } - const baseMeta: ProviderMeta | undefined = - payload.meta ?? (initialData?.meta ? { ...initialData.meta } : undefined); + const metaSource = payload.meta ?? initialData?.meta; + const baseMeta: ProviderMeta | undefined = metaSource + ? { ...metaSource } + : undefined; + // Existing-provider edits never own endpoint membership. The backend + // rejects endpoint-bearing update payloads; add/remove/touch use their + // dedicated commands and remain safe from stale form snapshots. + if (isEditMode && baseMeta) { + delete baseMeta.custom_endpoints; + } // 确定 providerType(新建时从预设获取,编辑时从现有数据获取) const providerType = presetProviderType || initialData?.meta?.providerType; 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) => { diff --git a/src/types.ts b/src/types.ts index 585e95429..e316961c5 100644 --- a/src/types.ts +++ b/src/types.ts @@ -38,7 +38,7 @@ export interface AppConfig { // 自定义端点配置 export interface CustomEndpoint { url: string; - addedAt: number; + addedAt: number | null; lastUsed?: number; } diff --git a/tests/fixtures/pi/provider-write-api-v1.json b/tests/fixtures/pi/provider-write-api-v1.json new file mode 100644 index 000000000..c8776f3bd --- /dev/null +++ b/tests/fixtures/pi/provider-write-api-v1.json @@ -0,0 +1,35 @@ +{ + "manifestVersion": 1, + "codeAuthority": "src-tauri/src/database/dao/provider_write.rs", + "types": { + "ProviderKey": ["app_type", "id"], + "ProviderRowCreate": ["content", "created_at"], + "ProviderRowUpdate": [ + "name", + "settings_config", + "website_url", + "category", + "notes", + "meta", + "icon", + "icon_color" + ], + "NewEndpoint": ["url", "added_at", "last_used"], + "NewProviderAggregate": [ + "key", + "row", + "sort_index", + "in_failover_queue", + "initial_endpoints" + ], + "RenameProvider": ["source", "target_id", "row"] + }, + "databaseMethods": { + "create_provider": ["NewProviderAggregate"], + "update_provider": ["ProviderKey", "ProviderRowUpdate"], + "rename_db_only_additive_provider": ["RenameProvider"], + "add_provider_endpoint": ["ProviderKey", "NewEndpoint"], + "remove_provider_endpoint": ["ProviderKey", "str"], + "touch_provider_endpoint": ["ProviderKey", "str", "i64"] + } +}