From 4f78451405575158ff6562c7021c7f31f2860780 Mon Sep 17 00:00:00 2001 From: SaladDay Date: Fri, 31 Jul 2026 17:58:07 +0000 Subject: [PATCH] storage(pi): isolate typed provider writes and canonical restore Replace the generic Provider save/upsert surface with strict typed create/update/rename/endpoint operations. Keep aggregate hydration read-only, preserve nullable endpoint timestamps end to end, and make service create own the initial endpoint set atomically. Restore SQL and binary backups only through UntrustedScratch, migrate and copy fixed data columns into a fresh CanonicalStage, validate the canonical result, then publish through the Backup API. The imported schema is never eligible to become the live schema. Old save_provider callsite classification ========================================== Inventory authority: abandoned 5a385fc8 tree. The old definition at src-tauri/src/database/dao/providers.rs:180 is deleted and is not a callsite. Production callsites: - src-tauri/src/commands/provider.rs:253 [create] Claude Desktop import creates one absent aggregate; it now strict-inserts the row and initial endpoints in one transaction. - src-tauri/src/database/dao/providers.rs:638 [create] official seed first proves absence, then strict-creates; a racing insert is a conflict. - src-tauri/src/database/dao/providers.rs:704 [create] on-demand seed first proves absence, then strict-creates; it cannot overwrite an existing row. - src-tauri/src/services/omo.rs:291 [create] OMO import constructs a new aggregate and strict-creates it; OMO is not eligible for rename. - src-tauri/src/services/provider/endpoints.rs:85 [update] endpoint last-used is not a Provider-row save; it now calls the exact touch endpoint operation. - src-tauri/src/services/provider/live.rs:1567 [create/update] default live import is reconciliation: read first, then strict create or strict update. - src-tauri/src/services/provider/live.rs:1743 [update] an existing OpenCode live provider follows the strict row-update branch. - src-tauri/src/services/provider/live.rs:1770 [create] a new OpenCode live provider follows the strict aggregate-create branch. - src-tauri/src/services/provider/live.rs:1825 [update] an existing OpenClaw live provider follows the strict row-update branch. - src-tauri/src/services/provider/live.rs:1858 [create] a new OpenClaw live provider follows the strict aggregate-create branch. - src-tauri/src/services/provider/live.rs:1900 [update] an existing Hermes live provider follows the strict row-update branch. - src-tauri/src/services/provider/live.rs:1926 [create] a new Hermes live provider follows the strict aggregate-create branch. - src-tauri/src/services/provider/mod.rs:2568 [create] ProviderService::add owns strict aggregate creation and all initial endpoints. - src-tauri/src/services/provider/mod.rs:2680 [rename] an additive DB-only key change now uses the dedicated transactional rename after eligibility checks. - src-tauri/src/services/provider/mod.rs:2711 [update] OMO edit updates exactly the existing main row after its live-file coordination. - src-tauri/src/services/provider/mod.rs:2740 [update] additive-provider edit updates exactly the existing main row after resolving live ownership. - src-tauri/src/services/provider/mod.rs:2750 [update] switch-mode edit updates exactly the existing main row and never inserts. - src-tauri/src/services/provider/mod.rs:2948 [update] remove-from-live changes only the existing provider's live-managed marker. - src-tauri/src/services/provider/mod.rs:3120 [update] switch backfill updates only the existing current provider row. - src-tauri/src/services/provider/mod.rs:3174 [update] successful additive switch changes only the existing live-managed marker. - src-tauri/src/services/provider/mod.rs:3315 [update] common-config migration updates only each already-read existing row. - src-tauri/src/services/provider/mod.rs:3895 [update] Gemini credential scrub updates only each already-read existing row. - src-tauri/src/services/provider/mod.rs:4082 [update] sort ordering is routed to the dedicated sort-index state operation, not row replacement. - src-tauri/src/services/provider/mod.rs:4636 [create/update] universal-to- Claude reconciliation reads the target and selects strict create or update. - src-tauri/src/services/provider/mod.rs:4651 [create/update] universal-to- Codex reconciliation reads the target and selects strict create or update. - src-tauri/src/services/provider/mod.rs:4665 [create/update] universal-to- Gemini reconciliation reads the target and selects strict create or update. Required indirect ownership paths: - src-tauri/src/deeplink/provider.rs [create] the old indirect flow called ProviderService::add and then appended endpoints one by one. It now supplies every non-primary endpoint to one strict aggregate create, so hydration is complete atomically and a duplicate is zero-side-effect. - [restore] no old generic-save callsite is reclassified as restore. Exact aggregate replacement exists only as the sealed restore_provider_aggregate_on_tx compensation primitive. Test-only callsites: Every item below is classified [test]. Each is fixture setup, not a production write authority, and is migrated to a real ProviderService entry where the behavior is under test or to the cfg(test)-only typed fixture reconciler where the test merely needs pre-existing rows. - src-tauri/src/codex_history_migration.rs:1442 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:1452 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2174 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2176 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2199 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2219 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2247 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2267 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2288 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2320 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2393 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2449 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2498 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2555 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2604 [test] migration fixture setup. - src-tauri/src/codex_history_migration.rs:2625 [test] migration fixture setup. - src-tauri/src/database/dao/providers.rs:754 [test] DAO fixture setup. - src-tauri/src/proxy/provider_router.rs:351 [test] router fixture setup. - src-tauri/src/proxy/provider_router.rs:352 [test] router fixture setup. - src-tauri/src/proxy/provider_router.rs:377 [test] router fixture setup. - src-tauri/src/proxy/provider_router.rs:378 [test] router fixture setup. - src-tauri/src/proxy/provider_router.rs:410 [test] router fixture setup. - src-tauri/src/proxy/provider_router.rs:411 [test] router fixture setup. - src-tauri/src/proxy/provider_router.rs:447 [test] router fixture setup. - src-tauri/src/proxy/provider_router.rs:448 [test] router fixture setup. - src-tauri/src/proxy/provider_router.rs:488 [test] router fixture setup. - src-tauri/src/services/provider/mod.rs:485 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:586 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:813 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:825 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:1472 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:1607 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:1737 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:1945 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:1952 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:1978 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2006 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2056 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2130 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2167 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2207 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2235 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2270 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2320 [test] service fixture setup. - src-tauri/src/services/provider/mod.rs:2362 [test] service fixture setup. - src-tauri/src/services/proxy.rs:3762 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:3948 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4034 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4095 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4114 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4263 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4341 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4421 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4533 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4651 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:4787 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5264 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5320 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5385 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5387 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5460 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5462 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5611 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5613 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5615 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5698 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5700 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:5998 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6000 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6173 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6175 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6417 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6419 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6553 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6555 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6635 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6637 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:6919 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:7173 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:7175 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:7240 [test] proxy fixture setup. - src-tauri/src/services/proxy.rs:7242 [test] proxy fixture setup. - src-tauri/tests/profile_roundtrip.rs:112 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:116 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:126 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:133 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:292 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:501 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:505 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:670 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:677 [test] profile fixture create. - src-tauri/tests/profile_roundtrip.rs:762 [test] Linux Desktop fixture create. - src-tauri/tests/profile_roundtrip.rs:769 [test] Linux Desktop fixture create. - src-tauri/tests/provider_commands.rs:69 [test] command fixture create. - src-tauri/tests/provider_service.rs:2927 [test] service fixture create. --- src-tauri/Cargo.lock | 1 + src-tauri/Cargo.toml | 3 +- src-tauri/src/codex_history_migration.rs | 47 +- src-tauri/src/commands/provider.rs | 15 +- src-tauri/src/database/backup.rs | 2642 ++++++++++++++++- src-tauri/src/database/dao/mod.rs | 3 + src-tauri/src/database/dao/pi_projections.rs | 204 ++ src-tauri/src/database/dao/provider_write.rs | 531 ++++ src-tauri/src/database/dao/providers.rs | 978 +++--- .../src/database/dao/skill_deployments.rs | 208 ++ src-tauri/src/database/dao/skills.rs | 70 +- src-tauri/src/database/mod.rs | 46 +- src-tauri/src/database/schema.rs | 843 +++++- src-tauri/src/deeplink/provider.rs | 40 +- src-tauri/src/deeplink/tests.rs | 35 +- src-tauri/src/error.rs | 2 + src-tauri/src/lib.rs | 7 +- src-tauri/src/provider.rs | 78 + src-tauri/src/proxy/provider_router.rs | 27 +- src-tauri/src/services/omo.rs | 6 +- src-tauri/src/services/provider/endpoints.rs | 24 +- src-tauri/src/services/provider/live.rs | 49 +- src-tauri/src/services/provider/mod.rs | 725 ++++- src-tauri/src/services/proxy.rs | 75 +- src-tauri/src/settings.rs | 4 +- src-tauri/tests/profile_roundtrip.rs | 131 +- src-tauri/tests/provider_commands.rs | 26 +- src-tauri/tests/provider_service.rs | 8 +- src-tauri/tests/support.rs | 26 +- .../providers/forms/ProviderForm.tsx | 12 +- src/types.ts | 2 +- .../pi/canonical-schema-manifest-v1.json | 107 + tests/fixtures/pi/provider-write-api-v1.json | 35 + tests/fixtures/pi/restore-policy-v1.json | 31 + 34 files changed, 6210 insertions(+), 831 deletions(-) create mode 100644 src-tauri/src/database/dao/pi_projections.rs create mode 100644 src-tauri/src/database/dao/provider_write.rs create mode 100644 src-tauri/src/database/dao/skill_deployments.rs create mode 100644 tests/fixtures/pi/canonical-schema-manifest-v1.json create mode 100644 tests/fixtures/pi/provider-write-api-v1.json create mode 100644 tests/fixtures/pi/restore-policy-v1.json diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 295039b04..7c9bd484b 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -783,6 +783,7 @@ dependencies = [ "indexmap 2.13.0", "json-five", "json5", + "libc", "log", "objc2 0.5.2", "objc2-app-kit 0.2.2", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index b9cfc800e..22e3c442b 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -73,11 +73,12 @@ url = "2.5" auto-launch = "0.5" once_cell = "1.21.3" base64 = "0.22" -rusqlite = { version = "0.31", features = ["bundled", "backup", "hooks"] } +rusqlite = { version = "0.31", features = ["bundled", "backup", "hooks", "limits"] } indexmap = { version = "2", features = ["serde"] } rust_decimal = "1.33" uuid = { version = "1.11", features = ["v4"] } sha2 = "0.10" +libc = "0.2" hmac = "0.12" json5 = "0.4" json-five = "0.3.1" diff --git a/src-tauri/src/codex_history_migration.rs b/src-tauri/src/codex_history_migration.rs index 9e4f0e023..b7a10683e 100644 --- a/src-tauri/src/codex_history_migration.rs +++ b/src-tauri/src/codex_history_migration.rs @@ -1439,7 +1439,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 +1450,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 +2173,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 +2199,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 +2220,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 +2249,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 +2270,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 +2292,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 +2325,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 +2399,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 +2456,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 +2506,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 +2564,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 +2614,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 +2636,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/backup.rs b/src-tauri/src/database/backup.rs index 4808adf46..9abd30347 100644 --- a/src-tauri/src/database/backup.rs +++ b/src-tauri/src/database/backup.rs @@ -2,19 +2,57 @@ //! //! 提供 SQL 导出/导入和二进制快照备份功能。 -use super::{lock_conn, Database}; +use super::schema::CanonicalStage; +use super::{lock_conn, Database, SCHEMA_VERSION}; use crate::config::get_app_config_dir; use crate::error::AppError; use chrono::{Local, Utc}; use rusqlite::backup::Backup; -use rusqlite::types::ValueRef; -use rusqlite::Connection; -use std::fs; +use rusqlite::config::DbConfig; +use rusqlite::limits::Limit; +use rusqlite::types::{Value, ValueRef}; +use rusqlite::{Connection, OpenFlags}; +use std::fs::{self, File, Metadata, OpenOptions}; +use std::io::{Read, Take}; use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::Arc; use tempfile::NamedTempFile; const CC_SWITCH_SQL_EXPORT_HEADER: &str = "-- CC Switch SQLite 导出"; +pub(crate) const MAX_SQL_IMPORT_BYTES: u64 = 256 * 1024 * 1024; +pub(crate) const MAX_BINARY_RESTORE_BYTES: u64 = 2 * 1024 * 1024 * 1024; +pub(crate) const MAX_SCRATCH_BYTES: u64 = 2 * 1024 * 1024 * 1024; +const MAX_SQL_VALUE_BYTES: i32 = 64 * 1024 * 1024; +const MAX_VM_STEPS: u64 = 50_000_000; +const PROGRESS_GRANULARITY: u64 = 1_000; +const MAX_PAGE_COUNT: u64 = 524_288; + +#[cfg(test)] +thread_local! { + static TEST_MAX_VM_STEPS: std::cell::Cell> = + const { std::cell::Cell::new(None) }; + static TEST_MAX_PAGE_COUNT: std::cell::Cell> = + const { std::cell::Cell::new(None) }; +} + +fn max_vm_steps() -> u64 { + #[cfg(test)] + if let Some(limit) = TEST_MAX_VM_STEPS.with(std::cell::Cell::get) { + return limit; + } + MAX_VM_STEPS +} + +fn max_page_count() -> u64 { + #[cfg(test)] + if let Some(limit) = TEST_MAX_PAGE_COUNT.with(std::cell::Cell::get) { + return limit; + } + MAX_PAGE_COUNT +} + /// `dump_sql` 会写出的 PRAGMA。其余 PRAGMA 一律拒绝——`temp_store_directory` /// 能把临时文件重定向到任意目录,`writable_schema` 能绕过 schema 完整性检查。 const IMPORT_ALLOWED_PRAGMAS: &[&str] = &["foreign_keys", "user_version"]; @@ -24,7 +62,7 @@ const IMPORT_ALLOWED_PRAGMAS: &[&str] = &["foreign_keys", "user_version"]; /// 头部校验(`validate_cc_switch_sql_export`)只比较一个注释前缀,任何人都能在 /// 合法前缀后面接着写别的语句。`ATTACH DATABASE '/path/x.db'` 的副作用发生在 /// `validate_basic_state` 之前,导入即使最终失败,文件也已经被创建;而 `settings` -/// 表不在 `SYNC_SKIP_TABLES` / `SYNC_PRESERVE_TABLES` 之列,WebDAV/S3 同步会走 +/// 表不在同步 skip/commit-boundary overlay 之列,WebDAV/S3 同步会走 /// 同一条 `import_sql_string_inner`,所以这条路径的输入不可信。 /// /// 为什么是 authorizer 而不是「扫描 ATTACH 关键字」:字符串扫描会被 `/*x*/ATTACH`、 @@ -49,6 +87,14 @@ fn import_authorizer(context: rusqlite::hooks::AuthContext<'_>) -> rusqlite::hoo let escapes_temp_db = match context.action { AuthAction::Attach { .. } | AuthAction::Detach { .. } => true, AuthAction::CreateVtable { .. } | AuthAction::DropVtable { .. } => true, + // Genuine exports can contain expression indexes (for example, + // COALESCE in the request-log dedupe index), so ordinary built-ins + // must remain usable while the untrusted schema is assembled. + // No application functions are registered on this connection and + // extension loading is never enabled; deny the SQL entry point too. + AuthAction::Function { function_name } => { + function_name.eq_ignore_ascii_case("load_extension") + } AuthAction::Unknown { .. } => true, AuthAction::Pragma { pragma_name, .. } => !IMPORT_ALLOWED_PRAGMAS .iter() @@ -72,17 +118,1168 @@ const SYNC_SKIP_TABLES: &[&str] = &[ "provider_health", "proxy_live_backup", "usage_daily_rollups", + "pi_provider_projections", + "skill_deployments", ]; -/// Tables whose local data is preserved (restored from local snapshot) during WebDAV import. -/// Excludes ephemeral tables like provider_health that can safely rebuild at runtime. -const SYNC_PRESERVE_TABLES: &[&str] = &[ +/// Exact file/deployment ownership belongs to the current device. Portable SQL +/// exports carry table shape but never rows, and imports restore live rows. +const DEVICE_LOCAL_TABLES: &[&str] = &["pi_provider_projections", "skill_deployments"]; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum RestorePolicy { + PortableIncoming, + PreserveLive, + RebuildRuntime, + SeedCanonical, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum StorageKind { + Text, + Integer, + Real, +} + +#[derive(Debug, Clone, Copy)] +struct RestoreColumnSpec { + name: &'static str, + storage: StorageKind, + nullable: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum RestoreRowValidator { + /// Storage/nullability is the complete portable contract for this row; + /// no JSON or decimal domain decoding is intentionally required. + OpaqueStorage, + Provider, + Mcp, + Profile, + DecimalColumns(&'static [usize]), +} + +#[derive(Debug, Clone, Copy)] +struct RestoreTableSpec { + name: &'static str, + policy: RestorePolicy, + columns: &'static [RestoreColumnSpec], + validator: RestoreRowValidator, + parents: &'static [&'static str], +} + +macro_rules! text_col { + ($name:literal) => { + RestoreColumnSpec { + name: $name, + storage: StorageKind::Text, + nullable: false, + } + }; +} + +macro_rules! nullable_text_col { + ($name:literal) => { + RestoreColumnSpec { + name: $name, + storage: StorageKind::Text, + nullable: true, + } + }; +} + +macro_rules! integer_col { + ($name:literal) => { + RestoreColumnSpec { + name: $name, + storage: StorageKind::Integer, + nullable: false, + } + }; +} + +macro_rules! nullable_integer_col { + ($name:literal) => { + RestoreColumnSpec { + name: $name, + storage: StorageKind::Integer, + nullable: true, + } + }; +} + +macro_rules! real_col { + ($name:literal) => { + RestoreColumnSpec { + name: $name, + storage: StorageKind::Real, + nullable: false, + } + }; +} + +const PROVIDERS_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("id"), + text_col!("app_type"), + text_col!("name"), + text_col!("settings_config"), + nullable_text_col!("website_url"), + nullable_text_col!("category"), + nullable_integer_col!("created_at"), + nullable_integer_col!("sort_index"), + nullable_text_col!("notes"), + nullable_text_col!("icon"), + nullable_text_col!("icon_color"), + text_col!("meta"), + integer_col!("is_current"), + integer_col!("in_failover_queue"), + text_col!("cost_multiplier"), + nullable_text_col!("limit_daily_usd"), + nullable_text_col!("limit_monthly_usd"), + nullable_text_col!("provider_type"), +]; + +const PROVIDER_ENDPOINTS_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + integer_col!("id"), + text_col!("provider_id"), + text_col!("app_type"), + text_col!("url"), + nullable_integer_col!("added_at"), + nullable_integer_col!("last_used"), +]; + +const MCP_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("id"), + text_col!("name"), + text_col!("server_config"), + nullable_text_col!("description"), + nullable_text_col!("homepage"), + nullable_text_col!("docs"), + text_col!("tags"), + integer_col!("enabled_claude"), + integer_col!("enabled_codex"), + integer_col!("enabled_gemini"), + integer_col!("enabled_grokbuild"), + integer_col!("enabled_opencode"), + integer_col!("enabled_hermes"), +]; + +const PROMPTS_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("id"), + text_col!("app_type"), + text_col!("name"), + text_col!("content"), + nullable_text_col!("description"), + integer_col!("enabled"), + nullable_integer_col!("created_at"), + nullable_integer_col!("updated_at"), +]; + +const SKILLS_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("id"), + text_col!("name"), + nullable_text_col!("description"), + text_col!("directory"), + nullable_text_col!("repo_owner"), + nullable_text_col!("repo_name"), + nullable_text_col!("repo_branch"), + nullable_text_col!("readme_url"), + integer_col!("enabled_claude"), + integer_col!("enabled_codex"), + integer_col!("enabled_gemini"), + integer_col!("enabled_grokbuild"), + integer_col!("enabled_opencode"), + integer_col!("enabled_hermes"), + integer_col!("enabled_pi"), + integer_col!("installed_at"), + nullable_text_col!("content_hash"), + integer_col!("updated_at"), +]; + +const SKILL_REPOS_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("owner"), + text_col!("name"), + text_col!("branch"), + integer_col!("enabled"), +]; + +const SETTINGS_RESTORE_COLUMNS: &[RestoreColumnSpec] = + &[text_col!("key"), nullable_text_col!("value")]; + +const PROXY_CONFIG_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("app_type"), + integer_col!("proxy_enabled"), + text_col!("listen_address"), + integer_col!("listen_port"), + integer_col!("enable_logging"), + integer_col!("enabled"), + integer_col!("auto_failover_enabled"), + integer_col!("max_retries"), + integer_col!("streaming_first_byte_timeout"), + integer_col!("streaming_idle_timeout"), + integer_col!("non_streaming_timeout"), + integer_col!("circuit_failure_threshold"), + integer_col!("circuit_success_threshold"), + integer_col!("circuit_timeout_seconds"), + real_col!("circuit_error_rate_threshold"), + integer_col!("circuit_min_requests"), + text_col!("default_cost_multiplier"), + text_col!("pricing_model_source"), + integer_col!("live_takeover_active"), + text_col!("created_at"), + text_col!("updated_at"), +]; + +const PROVIDER_HEALTH_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("provider_id"), + text_col!("app_type"), + integer_col!("is_healthy"), + integer_col!("consecutive_failures"), + nullable_text_col!("last_success_at"), + nullable_text_col!("last_failure_at"), + nullable_text_col!("last_error"), + text_col!("updated_at"), +]; + +const PROXY_LOG_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("request_id"), + text_col!("provider_id"), + text_col!("app_type"), + text_col!("model"), + nullable_text_col!("request_model"), + nullable_text_col!("pricing_model"), + integer_col!("input_tokens"), + integer_col!("output_tokens"), + integer_col!("cache_read_tokens"), + integer_col!("cache_creation_tokens"), + integer_col!("input_token_semantics"), + text_col!("input_cost_usd"), + text_col!("output_cost_usd"), + text_col!("cache_read_cost_usd"), + text_col!("cache_creation_cost_usd"), + text_col!("total_cost_usd"), + integer_col!("latency_ms"), + nullable_integer_col!("first_token_ms"), + nullable_integer_col!("duration_ms"), + integer_col!("status_code"), + nullable_text_col!("error_message"), + nullable_text_col!("session_id"), + nullable_text_col!("provider_type"), + integer_col!("is_streaming"), + text_col!("cost_multiplier"), + integer_col!("created_at"), + text_col!("data_source"), +]; + +const MODEL_PRICING_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("model_id"), + text_col!("display_name"), + text_col!("input_cost_per_million"), + text_col!("output_cost_per_million"), + text_col!("cache_read_cost_per_million"), + text_col!("cache_creation_cost_per_million"), +]; + +const STREAM_LOG_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + integer_col!("id"), + text_col!("provider_id"), + text_col!("provider_name"), + text_col!("app_type"), + text_col!("status"), + integer_col!("success"), + text_col!("message"), + nullable_integer_col!("response_time_ms"), + nullable_integer_col!("http_status"), + nullable_text_col!("model_used"), + nullable_integer_col!("retry_count"), + integer_col!("tested_at"), +]; + +const PROXY_LIVE_BACKUP_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("app_type"), + text_col!("original_config"), + text_col!("backed_up_at"), +]; + +const USAGE_ROLLUP_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("date"), + text_col!("app_type"), + text_col!("provider_id"), + text_col!("model"), + text_col!("request_model"), + text_col!("pricing_model"), + integer_col!("request_count"), + integer_col!("success_count"), + integer_col!("input_tokens"), + integer_col!("output_tokens"), + integer_col!("cache_read_tokens"), + integer_col!("cache_creation_tokens"), + integer_col!("input_token_semantics"), + text_col!("total_cost_usd"), + integer_col!("avg_latency_ms"), +]; + +const SESSION_SYNC_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("file_path"), + integer_col!("last_modified"), + integer_col!("last_line_offset"), + integer_col!("last_synced_at"), +]; + +const PROFILE_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("id"), + text_col!("name"), + text_col!("payload"), + nullable_integer_col!("sort_order"), + nullable_integer_col!("created_at"), + nullable_integer_col!("updated_at"), +]; + +const PI_PROJECTION_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("provider_id"), + text_col!("provider_key"), + integer_col!("created_at"), + integer_col!("updated_at"), +]; + +const SKILL_DEPLOYMENT_RESTORE_COLUMNS: &[RestoreColumnSpec] = &[ + text_col!("app_type"), + text_col!("skill_id"), + text_col!("destination"), + text_col!("destination_key"), + text_col!("method"), + text_col!("source_identity"), + nullable_text_col!("deployed_digest"), + integer_col!("created_at"), + integer_col!("updated_at"), +]; + +/// Parent-before-child order is also the canonical copy order. +const RESTORE_TABLE_SPECS: &[RestoreTableSpec] = &[ + RestoreTableSpec { + name: "providers", + policy: RestorePolicy::PortableIncoming, + columns: PROVIDERS_RESTORE_COLUMNS, + validator: RestoreRowValidator::Provider, + parents: &[], + }, + RestoreTableSpec { + name: "provider_endpoints", + policy: RestorePolicy::PortableIncoming, + columns: PROVIDER_ENDPOINTS_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &["providers"], + }, + RestoreTableSpec { + name: "mcp_servers", + policy: RestorePolicy::PortableIncoming, + columns: MCP_RESTORE_COLUMNS, + validator: RestoreRowValidator::Mcp, + parents: &[], + }, + RestoreTableSpec { + name: "prompts", + policy: RestorePolicy::PortableIncoming, + columns: PROMPTS_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, + RestoreTableSpec { + name: "skills", + policy: RestorePolicy::PortableIncoming, + columns: SKILLS_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, + RestoreTableSpec { + name: "skill_repos", + policy: RestorePolicy::PortableIncoming, + columns: SKILL_REPOS_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, + RestoreTableSpec { + name: "settings", + policy: RestorePolicy::PortableIncoming, + columns: SETTINGS_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, + RestoreTableSpec { + name: "proxy_config", + policy: RestorePolicy::PortableIncoming, + columns: PROXY_CONFIG_RESTORE_COLUMNS, + validator: RestoreRowValidator::DecimalColumns(&[16]), + parents: &[], + }, + RestoreTableSpec { + name: "provider_health", + policy: RestorePolicy::RebuildRuntime, + columns: PROVIDER_HEALTH_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &["providers"], + }, + RestoreTableSpec { + name: "proxy_request_logs", + policy: RestorePolicy::PortableIncoming, + columns: PROXY_LOG_RESTORE_COLUMNS, + validator: RestoreRowValidator::DecimalColumns(&[11, 12, 13, 14, 15, 24]), + parents: &[], + }, + RestoreTableSpec { + name: "model_pricing", + policy: RestorePolicy::PortableIncoming, + columns: MODEL_PRICING_RESTORE_COLUMNS, + validator: RestoreRowValidator::DecimalColumns(&[2, 3, 4, 5]), + parents: &[], + }, + RestoreTableSpec { + name: "stream_check_logs", + policy: RestorePolicy::PortableIncoming, + columns: STREAM_LOG_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, + RestoreTableSpec { + name: "proxy_live_backup", + policy: RestorePolicy::PortableIncoming, + columns: PROXY_LIVE_BACKUP_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, + RestoreTableSpec { + name: "usage_daily_rollups", + policy: RestorePolicy::PortableIncoming, + columns: USAGE_ROLLUP_RESTORE_COLUMNS, + validator: RestoreRowValidator::DecimalColumns(&[13]), + parents: &[], + }, + RestoreTableSpec { + name: "session_log_sync", + policy: RestorePolicy::PortableIncoming, + columns: SESSION_SYNC_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, + RestoreTableSpec { + name: "profiles", + policy: RestorePolicy::PortableIncoming, + columns: PROFILE_RESTORE_COLUMNS, + validator: RestoreRowValidator::Profile, + parents: &[], + }, + RestoreTableSpec { + name: "pi_provider_projections", + policy: RestorePolicy::PreserveLive, + columns: PI_PROJECTION_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, + RestoreTableSpec { + name: "skill_deployments", + policy: RestorePolicy::PreserveLive, + columns: SKILL_DEPLOYMENT_RESTORE_COLUMNS, + validator: RestoreRowValidator::OpaqueStorage, + parents: &[], + }, +]; + +const SYNC_LIVE_OVERLAY_TABLES: &[&str] = &[ "proxy_request_logs", "stream_check_logs", "proxy_live_backup", "usage_daily_rollups", ]; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum RestoreFlavor { + UserRestore, + Sync, +} + +/// An untrusted schema can only exist behind this private wrapper. There is no +/// conversion from it to CanonicalStage. +struct UntrustedScratch { + connection: Connection, + _file: NamedTempFile, + cancellation: Arc, +} + +#[cfg(unix)] +fn open_nofollow(path: &Path) -> std::io::Result { + use std::os::unix::fs::OpenOptionsExt; + + OpenOptions::new() + .read(true) + .custom_flags(libc::O_NOFOLLOW) + .open(path) +} + +#[cfg(not(unix))] +fn open_nofollow(path: &Path) -> std::io::Result { + OpenOptions::new().read(true).open(path) +} + +#[cfg(unix)] +fn same_file_identity(opened: &Metadata, current: &Metadata) -> bool { + use std::os::unix::fs::MetadataExt; + + opened.dev() == current.dev() && opened.ino() == current.ino() +} + +#[cfg(not(unix))] +fn same_file_identity(opened: &Metadata, current: &Metadata) -> bool { + opened.len() == current.len() && opened.modified().ok() == current.modified().ok() +} + +fn validate_regular_file(path: &Path, max_bytes: u64) -> Result { + let metadata = fs::symlink_metadata(path).map_err(|error| AppError::io(path, error))?; + if !metadata.file_type().is_file() { + return Err(AppError::InvalidInput(format!( + "restore source must be a regular non-symlink file: {}", + path.display() + ))); + } + if metadata.len() > max_bytes { + return Err(AppError::InvalidInput(format!( + "restore source exceeds {max_bytes} bytes: {}", + path.display() + ))); + } + Ok(metadata) +} + +fn read_restore_file(path: &Path, max_bytes: u64) -> Result, AppError> { + let initial = validate_regular_file(path, max_bytes)?; + let mut file = open_nofollow(path).map_err(|error| AppError::io(path, error))?; + let opened = file.metadata().map_err(|error| AppError::io(path, error))?; + if !opened.file_type().is_file() + || opened.len() > max_bytes + || !same_file_identity(&initial, &opened) + { + return Err(AppError::InvalidInput(format!( + "restore source changed before open: {}", + path.display() + ))); + } + let mut bytes = Vec::new(); + let mut limited: Take<&mut File> = file.by_ref().take(max_bytes + 1); + limited + .read_to_end(&mut bytes) + .map_err(|error| AppError::io(path, error))?; + if bytes.len() as u64 > max_bytes { + return Err(AppError::InvalidInput(format!( + "restore source exceeds {max_bytes} bytes: {}", + path.display() + ))); + } + let completed = file.metadata().map_err(|error| AppError::io(path, error))?; + let current = fs::symlink_metadata(path).map_err(|error| AppError::io(path, error))?; + if !current.file_type().is_file() + || !same_file_identity(&opened, ¤t) + || opened.len() != bytes.len() as u64 + || completed.len() != bytes.len() as u64 + || current.len() != bytes.len() as u64 + || opened.modified().ok() != completed.modified().ok() + { + return Err(AppError::InvalidInput(format!( + "restore source changed during read: {}", + path.display() + ))); + } + Ok(bytes) +} + +impl UntrustedScratch { + fn empty() -> Result { + let file = NamedTempFile::new().map_err(|error| AppError::IoContext { + context: "create untrusted restore scratch".to_string(), + source: error, + })?; + let connection = + Connection::open(file.path()).map_err(|error| AppError::Database(error.to_string()))?; + let cancellation = Arc::new(AtomicBool::new(false)); + let scratch = Self { + connection, + _file: file, + cancellation, + }; + scratch.configure_untrusted_execution()?; + Ok(scratch) + } + + fn configure_untrusted_execution(&self) -> Result<(), AppError> { + self.connection + .set_db_config(DbConfig::SQLITE_DBCONFIG_ENABLE_TRIGGER, false) + .map_err(|error| AppError::Database(error.to_string()))?; + self.connection + .set_db_config(DbConfig::SQLITE_DBCONFIG_TRUSTED_SCHEMA, false) + .map_err(|error| AppError::Database(error.to_string()))?; + self.connection + .set_db_config(DbConfig::SQLITE_DBCONFIG_DEFENSIVE, true) + .map_err(|error| AppError::Database(error.to_string()))?; + self.connection + .set_db_config(DbConfig::SQLITE_DBCONFIG_DQS_DDL, false) + .map_err(|error| AppError::Database(error.to_string()))?; + self.connection + .set_db_config(DbConfig::SQLITE_DBCONFIG_DQS_DML, false) + .map_err(|error| AppError::Database(error.to_string()))?; + self.connection.set_limit(Limit::SQLITE_LIMIT_ATTACHED, 0); + self.connection + .set_limit(Limit::SQLITE_LIMIT_LENGTH, MAX_SQL_VALUE_BYTES); + self.connection + .set_limit(Limit::SQLITE_LIMIT_SQL_LENGTH, MAX_SQL_IMPORT_BYTES as i32); + self.connection + .set_limit(Limit::SQLITE_LIMIT_VDBE_OP, 1_000_000); + self.connection + .set_limit(Limit::SQLITE_LIMIT_TRIGGER_DEPTH, 0); + self.connection + .execute_batch(&format!( + "PRAGMA trusted_schema = OFF; + PRAGMA foreign_keys = OFF; + PRAGMA max_page_count = {};", + max_page_count() + )) + .map_err(|error| AppError::Database(error.to_string()))?; + + let steps = Arc::new(AtomicU64::new(0)); + let cancellation = Arc::clone(&self.cancellation); + let max_steps = max_vm_steps(); + self.connection.progress_handler( + PROGRESS_GRANULARITY as i32, + Some(move || { + cancellation.load(Ordering::Relaxed) + || steps.fetch_add(PROGRESS_GRANULARITY, Ordering::Relaxed) >= max_steps + }), + ); + Ok(()) + } + + fn from_sql(sql: &str) -> Result { + if sql.len() as u64 > MAX_SQL_IMPORT_BYTES { + return Err(AppError::InvalidInput(format!( + "SQL import exceeds {MAX_SQL_IMPORT_BYTES} bytes" + ))); + } + let scratch = Self::empty()?; + scratch.connection.authorizer(Some(import_authorizer)); + let result = scratch.connection.execute_batch(sql); + scratch.connection.authorizer( + None::) -> rusqlite::hooks::Authorization>, + ); + result.map_err(|error| AppError::Database(format!("execute SQL import: {error}")))?; + scratch.finish_input() + } + + fn from_binary(path: &Path) -> Result { + let initial = validate_regular_file(path, MAX_BINARY_RESTORE_BYTES)?; + let guard = open_nofollow(path).map_err(|error| AppError::io(path, error))?; + let opened = guard + .metadata() + .map_err(|error| AppError::io(path, error))?; + if !same_file_identity(&initial, &opened) { + return Err(AppError::InvalidInput(format!( + "binary restore source changed before open: {}", + path.display() + ))); + } + let source = Connection::open_with_flags( + path, + OpenFlags::SQLITE_OPEN_READ_ONLY + | OpenFlags::SQLITE_OPEN_NOFOLLOW + | OpenFlags::SQLITE_OPEN_PRIVATE_CACHE, + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let mut scratch = Self::empty()?; + { + let backup = Backup::new(&source, &mut scratch.connection) + .map_err(|error| AppError::Database(error.to_string()))?; + backup + .step(-1) + .map_err(|error| AppError::Database(error.to_string()))?; + } + let current = fs::symlink_metadata(path).map_err(|error| AppError::io(path, error))?; + if !current.file_type().is_file() + || !same_file_identity(&opened, ¤t) + || opened.len() != current.len() + || opened.modified().ok() != current.modified().ok() + { + return Err(AppError::InvalidInput(format!( + "binary restore source changed during snapshot: {}", + path.display() + ))); + } + drop(source); + drop(guard); + scratch.finish_input() + } + + fn finish_input(self) -> Result { + self.enforce_scratch_size()?; + let version = Database::get_user_version(&self.connection)?; + if version > SCHEMA_VERSION { + return Err(AppError::InvalidInput(format!( + "restore schema version {version} is newer than supported {SCHEMA_VERSION}" + ))); + } + self.drop_untrusted_executable_objects()?; + self.connection + .set_db_config(DbConfig::SQLITE_DBCONFIG_ENABLE_TRIGGER, true) + .map_err(|error| AppError::Database(error.to_string()))?; + Database::create_tables_on_conn(&self.connection)?; + Database::apply_schema_migrations_on_conn(&self.connection)?; + self.connection + .execute_batch("PRAGMA foreign_keys = ON;") + .map_err(|error| AppError::Database(error.to_string()))?; + self.enforce_scratch_size()?; + // Keep the progress/cancellation handler installed while fixed-column + // rows are decoded and copied into the canonical stage. The source is + // still untrusted during those SELECTs; dropping the handler here would + // leave the data-transfer half of restore outside the VM budget. + Ok(self) + } + + fn drop_untrusted_executable_objects(&self) -> Result<(), AppError> { + for schema in ["sqlite_schema", "sqlite_temp_schema"] { + let mut stmt = self + .connection + .prepare(&format!( + "SELECT type, name FROM {schema} + WHERE type IN ('trigger', 'view', 'index') AND sql IS NOT NULL + ORDER BY CASE type WHEN 'trigger' THEN 0 WHEN 'view' THEN 1 ELSE 2 END, name" + )) + .map_err(|error| AppError::Database(error.to_string()))?; + let objects = stmt + .query_map([], |row| { + Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)) + }) + .map_err(|error| AppError::Database(error.to_string()))? + .collect::, _>>() + .map_err(|error| AppError::Database(error.to_string()))?; + drop(stmt); + for (kind, name) in objects { + let keyword = match kind.as_str() { + "trigger" => "TRIGGER", + "view" => "VIEW", + "index" => "INDEX", + _ => { + return Err(AppError::Database(format!( + "unsupported scratch object type '{kind}'" + ))) + } + }; + let escaped = name.replace('"', "\"\""); + self.connection + .execute(&format!("DROP {keyword} IF EXISTS \"{escaped}\""), []) + .map_err(|error| AppError::Database(error.to_string()))?; + } + } + Ok(()) + } + + fn enforce_scratch_size(&self) -> Result<(), AppError> { + let page_count: u64 = self + .connection + .query_row("PRAGMA page_count", [], |row| row.get(0)) + .map_err(|error| AppError::Database(error.to_string()))?; + let page_size: u64 = self + .connection + .query_row("PRAGMA page_size", [], |row| row.get(0)) + .map_err(|error| AppError::Database(error.to_string()))?; + let logical_size = page_count.saturating_mul(page_size); + let file_size = self + ._file + .as_file() + .metadata() + .map_err(|error| AppError::io(self._file.path(), error))? + .len(); + if logical_size > MAX_SCRATCH_BYTES || file_size > MAX_SCRATCH_BYTES { + return Err(AppError::InvalidInput(format!( + "restore scratch exceeds {MAX_SCRATCH_BYTES} bytes" + ))); + } + Ok(()) + } +} + +fn quoted_columns(spec: &RestoreTableSpec) -> String { + spec.columns + .iter() + .map(|column| format!("\"{}\"", column.name)) + .collect::>() + .join(", ") +} + +fn validate_storage( + table: &str, + column: &RestoreColumnSpec, + value: &Value, +) -> Result<(), AppError> { + let valid = match value { + Value::Null => column.nullable, + Value::Text(_) => column.storage == StorageKind::Text, + Value::Integer(_) => column.storage == StorageKind::Integer, + Value::Real(_) => column.storage == StorageKind::Real, + Value::Blob(_) => false, + }; + if valid { + Ok(()) + } else { + Err(AppError::InvalidInput(format!( + "restore row has invalid storage class for {table}.{}", + column.name + ))) + } +} + +fn text_value<'a>(table: &str, column: &str, value: &'a Value) -> Result<&'a str, AppError> { + match value { + Value::Text(value) => Ok(value), + _ => Err(AppError::InvalidInput(format!( + "restore row requires text at {table}.{column}" + ))), + } +} + +fn validate_json_text(table: &str, column: &str, value: &Value) -> Result<(), AppError> { + let text = text_value(table, column, value)?; + serde_json::from_str::(text) + .map(|_| ()) + .map_err(|error| { + AppError::InvalidInput(format!( + "restore row has invalid JSON at {table}.{column}: {error}" + )) + }) +} + +fn validate_restore_row(spec: &RestoreTableSpec, values: &[Value]) -> Result<(), AppError> { + if values.len() != spec.columns.len() { + return Err(AppError::Database(format!( + "restore decoder width mismatch for '{}'", + spec.name + ))); + } + for (column, value) in spec.columns.iter().zip(values) { + validate_storage(spec.name, column, value)?; + } + match spec.validator { + RestoreRowValidator::OpaqueStorage => {} + RestoreRowValidator::Provider => { + let id = text_value(spec.name, "id", &values[0])?; + let app_type = text_value(spec.name, "app_type", &values[1])?; + let settings = text_value(spec.name, "settings_config", &values[3])?; + let meta = text_value(spec.name, "meta", &values[11])?; + crate::database::dao::providers::validate_provider_storage_json( + app_type, id, settings, meta, + )?; + for index in [14_usize, 15, 16] { + if let Value::Text(value) = &values[index] { + value.parse::().map_err(|error| { + AppError::InvalidInput(format!( + "restore provider has invalid decimal in '{}': {error}", + spec.columns[index].name + )) + })?; + } + } + } + RestoreRowValidator::Mcp => { + validate_json_text(spec.name, "server_config", &values[2])?; + validate_json_text(spec.name, "tags", &values[6])?; + } + RestoreRowValidator::Profile => { + validate_json_text(spec.name, "payload", &values[2])?; + } + RestoreRowValidator::DecimalColumns(indices) => { + for index in indices { + if let Value::Text(value) = &values[*index] { + value.parse::().map_err(|error| { + AppError::InvalidInput(format!( + "restore row has invalid decimal at {}.{}: {error}", + spec.name, spec.columns[*index].name + )) + })?; + } + } + } + } + Ok(()) +} + +fn copy_fixed_table( + source: &Connection, + target: &Connection, + spec: &RestoreTableSpec, + clear_target: bool, +) -> Result<(), AppError> { + if clear_target { + target + .execute(&format!("DELETE FROM \"{}\"", spec.name), []) + .map_err(|error| { + AppError::Database(format!("clear canonical table '{}': {error}", spec.name)) + })?; + } + let columns = quoted_columns(spec); + let select = format!("SELECT {columns} FROM \"{}\"", spec.name); + let placeholders = (1..=spec.columns.len()) + .map(|index| format!("?{index}")) + .collect::>() + .join(", "); + let insert = format!( + "INSERT INTO \"{}\" ({columns}) VALUES ({placeholders})", + spec.name + ); + let mut statement = source.prepare(&select).map_err(|error| { + AppError::InvalidInput(format!( + "restore source is missing a fixed column in '{}': {error}", + spec.name + )) + })?; + let mut rows = statement + .query([]) + .map_err(|error| AppError::Database(error.to_string()))?; + while let Some(row) = rows + .next() + .map_err(|error| AppError::Database(error.to_string()))? + { + let values = (0..spec.columns.len()) + .map(|index| row.get::<_, Value>(index)) + .collect::, _>>() + .map_err(|error| AppError::InvalidInput(error.to_string()))?; + validate_restore_row(spec, &values)?; + target + .execute(&insert, rusqlite::params_from_iter(values.iter())) + .map_err(|error| { + AppError::InvalidInput(format!( + "canonical insert into '{}' failed: {error}", + spec.name + )) + })?; + } + Ok(()) +} + +fn assert_restore_policy_topology() -> Result<(), AppError> { + let mut seen = std::collections::BTreeSet::new(); + for spec in RESTORE_TABLE_SPECS { + if !seen.insert(spec.name) { + return Err(AppError::Database(format!( + "duplicate restore policy for '{}'", + spec.name + ))); + } + for parent in spec.parents { + if !seen.contains(parent) { + return Err(AppError::Database(format!( + "restore policy '{}' precedes parent '{}'", + spec.name, parent + ))); + } + } + } + Ok(()) +} + +fn canonical_user_tables( + conn: &Connection, +) -> Result, AppError> { + let mut statement = conn + .prepare( + "SELECT name FROM sqlite_schema + WHERE type = 'table' AND name NOT LIKE 'sqlite_%' + ORDER BY name", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let tables = statement + .query_map([], |row| row.get::<_, String>(0)) + .map_err(|error| AppError::Database(error.to_string()))? + .collect::, _>>() + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(tables) +} + +fn assert_restore_policy_coverage( + canonical_tables: &std::collections::BTreeSet, + policy_tables: &std::collections::BTreeSet, +) -> Result<(), AppError> { + if canonical_tables == policy_tables { + Ok(()) + } else { + let missing_policy = canonical_tables + .difference(policy_tables) + .cloned() + .collect::>(); + let stale_policy = policy_tables + .difference(canonical_tables) + .cloned() + .collect::>(); + Err(AppError::Database(format!( + "canonical user tables and restore policy manifest differ; \ + missing policy={missing_policy:?}, stale policy={stale_policy:?}" + ))) + } +} + +fn validate_stage_rows(conn: &Connection) -> Result<(), AppError> { + for spec in RESTORE_TABLE_SPECS { + let columns = quoted_columns(spec); + let mut statement = conn + .prepare(&format!("SELECT {columns} FROM \"{}\"", spec.name)) + .map_err(|error| AppError::Database(error.to_string()))?; + let mut rows = statement + .query([]) + .map_err(|error| AppError::Database(error.to_string()))?; + while let Some(row) = rows + .next() + .map_err(|error| AppError::Database(error.to_string()))? + { + let values = (0..spec.columns.len()) + .map(|index| row.get::<_, Value>(index)) + .collect::, _>>() + .map_err(|error| AppError::InvalidInput(error.to_string()))?; + validate_restore_row(spec, &values)?; + } + } + Ok(()) +} + +fn validate_canonical_behaviors(conn: &Connection) -> Result<(), AppError> { + let transaction = conn + .unchecked_transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + + let endpoint_max: i64 = transaction + .query_row( + "SELECT COALESCE(MAX(id), 0) FROM provider_endpoints", + [], + |row| row.get(0), + ) + .map_err(|error| AppError::Database(error.to_string()))?; + transaction + .execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES ('__restore_Aa', '__probe', 'probe', '{}', '{}')", + [], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + transaction + .execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES ('__restore_aa', '__probe', 'probe', '{}', '{}')", + [], + ) + .map_err(|error| { + AppError::Database(format!( + "canonical provider key is not BINARY case-sensitive: {error}" + )) + })?; + transaction + .execute( + "INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at, last_used) + VALUES ('__restore_Aa', '__probe', 'https://probe.invalid', NULL, NULL)", + [], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let endpoint_id = transaction.last_insert_rowid(); + if endpoint_id <= endpoint_max { + return Err(AppError::Database( + "provider_endpoints AUTOINCREMENT did not advance past explicit restored IDs" + .to_string(), + )); + } + if transaction + .execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES ('__restore_Aa', '__probe', 'replacement', '{}', '{}')", + [], + ) + .is_ok() + { + return Err(AppError::Database( + "canonical provider duplicate policy is not ABORT".to_string(), + )); + } + let endpoint_count: i64 = transaction + .query_row( + "SELECT COUNT(*) FROM provider_endpoints + WHERE provider_id = '__restore_Aa' AND app_type = '__probe'", + [], + |row| row.get(0), + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if endpoint_count != 1 { + return Err(AppError::Database( + "duplicate provider insert replaced or cascaded endpoint rows".to_string(), + )); + } + + let stream_max: i64 = transaction + .query_row( + "SELECT COALESCE(MAX(id), 0) FROM stream_check_logs", + [], + |row| row.get(0), + ) + .map_err(|error| AppError::Database(error.to_string()))?; + transaction + .execute( + "INSERT INTO stream_check_logs + (provider_id, provider_name, app_type, status, success, message, tested_at) + VALUES ('__probe', 'probe', '__probe', 'ok', 1, 'ok', 1)", + [], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if transaction.last_insert_rowid() <= stream_max { + return Err(AppError::Database( + "stream_check_logs AUTOINCREMENT did not advance past explicit restored IDs" + .to_string(), + )); + } + transaction + .rollback() + .map_err(|error| AppError::Database(error.to_string())) +} + +fn validate_canonical_stage(stage: &CanonicalStage) -> Result<(), AppError> { + let conn = stage.connection(); + if Database::get_user_version(conn)? != SCHEMA_VERSION { + return Err(AppError::Database( + "canonical stage has an unexpected schema version".to_string(), + )); + } + assert_restore_policy_topology()?; + let policy_tables = RESTORE_TABLE_SPECS + .iter() + .map(|spec| spec.name.to_string()) + .collect::>(); + assert_restore_policy_coverage(&canonical_user_tables(conn)?, &policy_tables)?; + validate_stage_rows(conn)?; + let integrity: String = conn + .query_row("PRAGMA integrity_check(1)", [], |row| row.get(0)) + .map_err(|error| AppError::Database(error.to_string()))?; + if integrity != "ok" { + return Err(AppError::Database(format!( + "canonical integrity_check failed: {integrity}" + ))); + } + let mut foreign_keys = conn + .prepare("PRAGMA foreign_key_check") + .map_err(|error| AppError::Database(error.to_string()))?; + if foreign_keys + .query([]) + .map_err(|error| AppError::Database(error.to_string()))? + .next() + .map_err(|error| AppError::Database(error.to_string()))? + .is_some() + { + return Err(AppError::Database( + "canonical foreign_key_check failed".to_string(), + )); + } + drop(foreign_keys); + validate_canonical_behaviors(conn) +} + /// A database backup entry for the UI #[derive(Debug, serde::Serialize)] #[serde(rename_all = "camelCase")] @@ -96,7 +1293,7 @@ impl Database { /// 导出为 SQLite 兼容的 SQL 文本(内存字符串,完整导出) pub fn export_sql_string(&self) -> Result { let snapshot = self.snapshot_to_memory()?; - Self::dump_sql(&snapshot, &[]) + Self::dump_sql(&snapshot, DEVICE_LOCAL_TABLES) } /// Export SQL for sync (WebDAV), skipping local-only tables' data @@ -118,87 +1315,41 @@ impl Database { /// 从 SQL 文件导入,返回生成的备份 ID(若无备份则为空字符串) pub fn import_sql(&self, source_path: &Path) -> Result { - if !source_path.exists() { - return Err(AppError::InvalidInput(format!( - "SQL 文件不存在: {}", + let bytes = read_restore_file(source_path, MAX_SQL_IMPORT_BYTES)?; + let sql = std::str::from_utf8(&bytes).map_err(|error| { + AppError::InvalidInput(format!( + "SQL restore source is not UTF-8 ({}): {error}", source_path.display() - ))); - } - - let sql_raw = fs::read_to_string(source_path).map_err(|e| AppError::io(source_path, e))?; - let sql_content = sql_raw.trim_start_matches('\u{feff}'); - self.import_sql_string(sql_content) + )) + })?; + self.import_sql_string(sql) } /// 从 SQL 字符串导入,返回生成的备份 ID(若无备份则为空字符串) pub fn import_sql_string(&self, sql_raw: &str) -> Result { - self.import_sql_string_inner(sql_raw, &[]) + self.import_sql_string_inner(sql_raw, RestoreFlavor::UserRestore) } /// Import SQL generated for sync, then restore local-only tables from the /// current device snapshot before replacing the main database. pub(crate) fn import_sql_string_for_sync(&self, sql_raw: &str) -> Result { - self.import_sql_string_inner(sql_raw, SYNC_PRESERVE_TABLES) + self.import_sql_string_inner(sql_raw, RestoreFlavor::Sync) } fn import_sql_string_inner( &self, sql_raw: &str, - preserve_tables: &[&str], + flavor: RestoreFlavor, ) -> Result { let sql_content = sql_raw.trim_start_matches('\u{feff}'); Self::validate_cc_switch_sql_export(sql_content)?; - - // 导入前备份现有数据库 + let scratch = UntrustedScratch::from_sql(sql_content)?; + let stage = Self::build_canonical_stage(&scratch)?; let backup_path = self.backup_database_file()?; - - let local_snapshot = if preserve_tables.is_empty() { - None - } else { - Some(self.snapshot_to_memory()?) - }; - - // 在临时数据库执行导入,确保失败不会污染主库 - let temp_file = NamedTempFile::new().map_err(|e| AppError::IoContext { - context: "创建临时数据库文件失败".to_string(), - source: e, - })?; - let temp_path = temp_file.path().to_path_buf(); - let temp_conn = - Connection::open(&temp_path).map_err(|e| AppError::Database(e.to_string()))?; - - // authorizer 只覆盖外部 SQL,执行完立刻摘掉:紧随其后的 - // `create_tables_on_conn` / `apply_schema_migrations_on_conn` 是本程序自己的 - // schema 维护语句,不属于需要设防的输入,没必要让它们也过一遍守卫。 - temp_conn.authorizer(Some(import_authorizer)); - let batch_result = temp_conn.execute_batch(sql_content); - temp_conn.authorizer( - None::) -> rusqlite::hooks::Authorization>, - ); - batch_result.map_err(|e| AppError::Database(format!("执行 SQL 导入失败: {e}")))?; - - // 补齐缺失表/索引并进行基础校验 - Self::create_tables_on_conn(&temp_conn)?; - Self::apply_schema_migrations_on_conn(&temp_conn)?; - Self::validate_basic_state(&temp_conn)?; - if let Some(local_snapshot) = local_snapshot.as_ref() { - Self::restore_tables(local_snapshot, &temp_conn, preserve_tables)?; - } - - // 使用 Backup 将临时库原子写回主库 - { - let mut main_conn = lock_conn!(self.conn); - let backup = Backup::new(&temp_conn, &mut main_conn) - .map_err(|e| AppError::Database(e.to_string()))?; - backup - .step(-1) - .map_err(|e| AppError::Database(e.to_string()))?; - } - + self.publish_canonical_stage(stage, flavor)?; let backup_id = backup_path .and_then(|p| p.file_stem().map(|s| s.to_string_lossy().to_string())) .unwrap_or_default(); - Ok(backup_id) } @@ -219,6 +1370,75 @@ impl Database { Ok(snapshot) } + fn build_canonical_stage(scratch: &UntrustedScratch) -> Result { + let mut stage = Self::current_canonical_stage()?; + stage + .connection() + .execute_batch("PRAGMA foreign_keys = ON;") + .map_err(|error| AppError::Database(error.to_string()))?; + let transaction = stage + .connection_mut() + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + + // The factory may create canonical seed rows. Clear every non-seed + // table child-first so incoming rows are always inserted with plain + // INSERT into a clean target. + for spec in RESTORE_TABLE_SPECS.iter().rev() { + if spec.policy != RestorePolicy::SeedCanonical { + transaction + .execute(&format!("DELETE FROM \"{}\"", spec.name), []) + .map_err(|error| AppError::Database(error.to_string()))?; + } + } + for spec in RESTORE_TABLE_SPECS { + if spec.policy == RestorePolicy::PortableIncoming { + copy_fixed_table(&scratch.connection, &transaction, spec, false)?; + } + } + transaction + .commit() + .map_err(|error| AppError::Database(error.to_string()))?; + Self::validate_basic_state(stage.connection())?; + validate_canonical_stage(&stage)?; + Ok(stage) + } + + fn publish_canonical_stage( + &self, + mut stage: CanonicalStage, + flavor: RestoreFlavor, + ) -> Result<(), AppError> { + // CanonicalStage is the only accepted type. UntrustedScratch has no + // conversion or field access that can satisfy this boundary. + validate_canonical_stage(&stage)?; + let mut main_conn = lock_conn!(self.conn); + { + let transaction = stage + .connection_mut() + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + for spec in RESTORE_TABLE_SPECS { + let preserve = spec.policy == RestorePolicy::PreserveLive + || (flavor == RestoreFlavor::Sync + && SYNC_LIVE_OVERLAY_TABLES.contains(&spec.name)); + if preserve { + copy_fixed_table(&main_conn, &transaction, spec, true)?; + } + } + transaction + .commit() + .map_err(|error| AppError::Database(error.to_string()))?; + } + validate_canonical_stage(&stage)?; + let backup = Backup::new(stage.connection(), &mut main_conn) + .map_err(|error| AppError::Database(error.to_string()))?; + backup + .step(-1) + .map(|_| ()) + .map_err(|error| AppError::Database(error.to_string())) + } + fn validate_cc_switch_sql_export(sql: &str) -> Result<(), AppError> { let trimmed = sql.trim_start(); if trimmed.starts_with(CC_SWITCH_SQL_EXPORT_HEADER) { @@ -232,62 +1452,6 @@ impl Database { )) } - fn restore_tables( - source_conn: &Connection, - target_conn: &Connection, - tables: &[&str], - ) -> Result<(), AppError> { - for table in tables { - if !Self::table_exists(source_conn, table)? || !Self::table_exists(target_conn, table)? - { - continue; - } - - let columns = Self::get_table_columns(source_conn, table)?; - if columns.is_empty() { - continue; - } - - target_conn - .execute(&format!("DELETE FROM \"{table}\""), []) - .map_err(|e| AppError::Database(format!("清空表 {table} 失败: {e}")))?; - - let placeholders = (1..=columns.len()) - .map(|idx| format!("?{idx}")) - .collect::>() - .join(", "); - let cols = columns - .iter() - .map(|column| format!("\"{column}\"")) - .collect::>() - .join(", "); - let insert_sql = format!("INSERT INTO \"{table}\" ({cols}) VALUES ({placeholders})"); - - let mut stmt = source_conn - .prepare(&format!("SELECT * FROM \"{table}\"")) - .map_err(|e| AppError::Database(format!("读取表 {table} 失败: {e}")))?; - let mut rows = stmt - .query([]) - .map_err(|e| AppError::Database(format!("查询表 {table} 数据失败: {e}")))?; - - while let Some(row) = rows.next().map_err(|e| AppError::Database(e.to_string()))? { - let mut values = Vec::with_capacity(columns.len()); - for idx in 0..columns.len() { - values.push( - row.get::<_, rusqlite::types::Value>(idx) - .map_err(|e| AppError::Database(e.to_string()))?, - ); - } - - target_conn - .execute(&insert_sql, rusqlite::params_from_iter(values.iter())) - .map_err(|e| AppError::Database(format!("恢复表 {table} 数据失败: {e}")))?; - } - } - - Ok(()) - } - /// Periodic backup: create a new backup if the latest one is older than the configured interval pub(crate) fn periodic_backup_if_needed(&self) -> Result<(), AppError> { let interval_hours = crate::settings::effective_backup_interval_hours(); @@ -620,35 +1784,16 @@ impl Database { let backup_dir = get_app_config_dir().join("backups"); let backup_path = backup_dir.join(filename); - if !backup_path.exists() { - return Err(AppError::InvalidInput(format!( - "Backup file not found: {filename}" - ))); - } + // Build a canonical data-only stage before touching the live database. + let scratch = UntrustedScratch::from_binary(&backup_path)?; + let stage = Self::build_canonical_stage(&scratch)?; - // Step 1: Create safety backup of current database + // Create the safety backup only after the source has passed staging. let safety_backup = self.backup_database_file()?; let safety_id = safety_backup .and_then(|p| p.file_stem().map(|s| s.to_string_lossy().to_string())) .unwrap_or_default(); - - // Step 2: Open the backup file and restore it to the main database - let source_conn = - Connection::open(&backup_path).map_err(|e| AppError::Database(e.to_string()))?; - - { - let mut main_conn = lock_conn!(self.conn); - let backup = Backup::new(&source_conn, &mut main_conn) - .map_err(|e| AppError::Database(e.to_string()))?; - backup - .step(-1) - .map_err(|e| AppError::Database(e.to_string()))?; - } - - // Step 3: Run schema migrations (backup may be from an older version) - self.create_tables()?; - self.apply_schema_migrations()?; - self.ensure_model_pricing_seeded()?; + self.publish_canonical_stage(stage, RestoreFlavor::UserRestore)?; log::info!("Database restored from backup: {filename}, safety backup: {safety_id}"); Ok(safety_id) @@ -745,10 +1890,884 @@ impl Database { #[cfg(test)] mod tests { - use super::Database; + use super::{ + assert_restore_policy_coverage, validate_canonical_behaviors, validate_regular_file, + validate_stage_rows, Database, RestoreFlavor, RestorePolicy, RestoreRowValidator, + StorageKind, UntrustedScratch, MAX_BINARY_RESTORE_BYTES, MAX_SCRATCH_BYTES, + MAX_SQL_IMPORT_BYTES, RESTORE_TABLE_SPECS, SCHEMA_VERSION, TEST_MAX_PAGE_COUNT, + TEST_MAX_VM_STEPS, + }; use crate::error::AppError; use crate::settings::{update_settings, AppSettings}; + use rusqlite::backup::Backup; + use rusqlite::Connection; use serial_test::serial; + use sha2::{Digest, Sha256}; + use std::ffi::OsString; + use std::fs::{self, File}; + + #[derive(Clone, Copy)] + enum RestoreEntryPoint { + Sql, + Binary, + } + + struct TestHomeGuard(Option); + + impl TestHomeGuard { + fn set(path: &std::path::Path) -> Self { + let previous = std::env::var_os("CC_SWITCH_TEST_HOME"); + std::env::set_var("CC_SWITCH_TEST_HOME", path); + Self(previous) + } + } + + impl Drop for TestHomeGuard { + fn drop(&mut self) { + match self.0.take() { + Some(previous) => std::env::set_var("CC_SWITCH_TEST_HOME", previous), + None => std::env::remove_var("CC_SWITCH_TEST_HOME"), + } + } + } + + struct RestoreLimitGuard { + previous_vm_steps: Option, + previous_page_count: Option, + } + + impl RestoreLimitGuard { + fn set(vm_steps: Option, page_count: Option) -> Self { + let previous_vm_steps = TEST_MAX_VM_STEPS.with(|current| current.replace(vm_steps)); + let previous_page_count = + TEST_MAX_PAGE_COUNT.with(|current| current.replace(page_count)); + Self { + previous_vm_steps, + previous_page_count, + } + } + } + + impl Drop for RestoreLimitGuard { + fn drop(&mut self) { + TEST_MAX_VM_STEPS.with(|current| current.set(self.previous_vm_steps)); + TEST_MAX_PAGE_COUNT.with(|current| current.set(self.previous_page_count)); + } + } + + fn restore_policy_snapshot() -> serde_json::Value { + let full_specs = RESTORE_TABLE_SPECS + .iter() + .map(|spec| { + let policy = match spec.policy { + RestorePolicy::PortableIncoming => "portable_incoming", + RestorePolicy::PreserveLive => "preserve_live", + RestorePolicy::RebuildRuntime => "rebuild_runtime", + RestorePolicy::SeedCanonical => "seed_canonical", + }; + let validator = match spec.validator { + RestoreRowValidator::OpaqueStorage => { + serde_json::json!("opaque_storage") + } + RestoreRowValidator::Provider => serde_json::json!("provider"), + RestoreRowValidator::Mcp => serde_json::json!("mcp"), + RestoreRowValidator::Profile => serde_json::json!("profile"), + RestoreRowValidator::DecimalColumns(indices) => { + serde_json::json!({"decimalColumns": indices}) + } + }; + serde_json::json!({ + "name": spec.name, + "policy": policy, + "columns": spec.columns.iter().map(|column| { + let storage = match column.storage { + StorageKind::Text => "text", + StorageKind::Integer => "integer", + StorageKind::Real => "real", + }; + serde_json::json!([column.name, storage, column.nullable]) + }).collect::>(), + "validator": validator, + "parents": spec.parents, + }) + }) + .collect::>(); + let digest = Sha256::digest( + serde_json::to_vec(&full_specs).expect("serialize restore policy authority"), + ); + let tables = full_specs + .iter() + .map(|spec| { + serde_json::json!({ + "name": spec["name"], + "policy": spec["policy"], + }) + }) + .collect::>(); + serde_json::json!({ + "manifestVersion": 1, + "codeAuthority": "src-tauri/src/database/backup.rs", + "schemaVersion": SCHEMA_VERSION, + "limits": { + "sqlImportBytes": MAX_SQL_IMPORT_BYTES, + "binaryRestoreBytes": MAX_BINARY_RESTORE_BYTES, + "scratchBytes": MAX_SCRATCH_BYTES, + }, + "specSha256": format!("{digest:x}"), + "tables": tables, + }) + } + + #[test] + fn restore_policy_snapshot_is_exhaustive_and_detects_missing_table_fixture( + ) -> Result<(), AppError> { + let expected: serde_json::Value = serde_json::from_str(include_str!( + "../../../tests/fixtures/pi/restore-policy-v1.json" + )) + .expect("parse restore policy snapshot"); + assert_eq!(restore_policy_snapshot(), expected); + + let canonical = Database::current_canonical_stage()?; + let canonical_tables = super::canonical_user_tables(canonical.connection())?; + let policy_tables = RESTORE_TABLE_SPECS + .iter() + .map(|spec| spec.name.to_string()) + .collect::>(); + assert_restore_policy_coverage(&canonical_tables, &policy_tables)?; + + let mut negative_fixture = canonical_tables; + negative_fixture.insert("future_table_without_policy".to_string()); + let error = assert_restore_policy_coverage(&negative_fixture, &policy_tables) + .expect_err("a new table without a policy must fail the scanner"); + assert!(error.to_string().contains("future_table_without_policy")); + Ok(()) + } + + fn canonical_restore_source() -> Result { + let source = Database::memory()?; + { + let conn = crate::database::lock_conn!(source.conn); + conn.execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES ('remote-provider', 'pi', 'Remote Provider', '{}', '{}')", + [], + )?; + } + source.snapshot_to_memory() + } + + fn weak_ledger_source() -> Result { + let source = canonical_restore_source()?; + source.execute_batch( + "PRAGMA foreign_keys = OFF; + DROP TABLE pi_provider_projections; + DROP TABLE skill_deployments; + CREATE TABLE pi_provider_projections ( + provider_id TEXT, + provider_key TEXT, + created_at INTEGER, + updated_at INTEGER + ); + CREATE UNIQUE INDEX remote_projection_alternate + ON pi_provider_projections(updated_at); + CREATE TABLE skill_deployments ( + app_type TEXT, + skill_id TEXT, + destination TEXT, + destination_key TEXT, + method TEXT, + source_identity TEXT, + deployed_digest TEXT, + created_at INTEGER, + updated_at INTEGER + ); + CREATE UNIQUE INDEX remote_skill_alternate + ON skill_deployments(skill_id); + INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES ('remote-provider', 'remote-key', 10, 10); + INSERT INTO skill_deployments ( + app_type, skill_id, destination, destination_key, method, + source_identity, created_at, updated_at + ) VALUES ( + 'codex', 'remote-skill', '/remote', '/remote', 'move', + 'remote-source', 10, 10 + ); + PRAGMA foreign_keys = ON;", + )?; + Ok(source) + } + + fn weak_endpoint_source() -> Result { + let source = canonical_restore_source()?; + source.execute_batch( + "PRAGMA foreign_keys = OFF; + DROP TABLE provider_endpoints; + 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, + last_used INTEGER + ); + INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at) + VALUES ('remote-provider', 'pi', 'https://weak.test', 1); + PRAGMA foreign_keys = ON;", + )?; + Ok(source) + } + + fn seed_local_ledgers(target: &Database) -> Result<(), AppError> { + let conn = crate::database::lock_conn!(target.conn); + conn.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES ('local-provider', 'local-key', 20, 20)", + [], + )?; + conn.execute( + "INSERT INTO skill_deployments ( + app_type, skill_id, destination, destination_key, method, + source_identity, created_at, updated_at + ) VALUES ( + 'pi', 'local-skill', '/local', '/local', 'copy', + 'local-source', 20, 20 + )", + [], + )?; + Ok(()) + } + + fn run_restore_entry( + target: &Database, + source: &Connection, + entry_point: RestoreEntryPoint, + filename: &str, + ) -> Result { + match entry_point { + RestoreEntryPoint::Sql => { + let sql = Database::dump_sql(source, &[])?; + target.import_sql_string(&sql) + } + RestoreEntryPoint::Binary => { + let backup_dir = crate::config::get_app_config_dir().join("backups"); + std::fs::create_dir_all(&backup_dir) + .map_err(|error| AppError::io(&backup_dir, error))?; + let backup_path = backup_dir.join(filename); + let mut destination = Connection::open(&backup_path)?; + { + let backup = Backup::new(source, &mut destination)?; + backup.step(-1)?; + } + drop(destination); + target.restore_from_backup(filename) + } + } + } + + fn logical_snapshot(target: &Database) -> Result { + let snapshot = target.snapshot_to_memory()?; + let dump = Database::dump_sql(&snapshot, &[])?; + Ok(dump + .lines() + .filter(|line| !line.starts_with("-- 生成时间:")) + .collect::>() + .join("\n")) + } + + const WEAK_PROVIDERS_NO_KEY: &str = "CREATE TABLE providers ( + id TEXT NOT NULL, + app_type TEXT NOT NULL, + name TEXT NOT NULL, + settings_config TEXT NOT NULL, + website_url TEXT, + category TEXT, + created_at INTEGER, + sort_index INTEGER, + notes TEXT, + icon TEXT, + icon_color TEXT, + meta TEXT NOT NULL DEFAULT '{}', + is_current BOOLEAN NOT NULL DEFAULT 0, + in_failover_queue BOOLEAN NOT NULL DEFAULT 0, + cost_multiplier TEXT NOT NULL DEFAULT '1.0', + limit_daily_usd TEXT, + limit_monthly_usd TEXT, + provider_type TEXT + )"; + + const HOSTILE_PROVIDERS_NOCASE_REPLACE: &str = "CREATE TABLE providers ( + id TEXT COLLATE NOCASE NOT NULL, + app_type TEXT NOT NULL, + name TEXT NOT NULL, + settings_config TEXT NOT NULL, + website_url TEXT, + category TEXT, + created_at INTEGER, + sort_index INTEGER, + notes TEXT, + icon TEXT, + icon_color TEXT, + meta TEXT NOT NULL DEFAULT '{}', + is_current BOOLEAN NOT NULL DEFAULT 0, + in_failover_queue BOOLEAN NOT NULL DEFAULT 0, + cost_multiplier TEXT NOT NULL DEFAULT '1.0', + limit_daily_usd TEXT, + limit_monthly_usd TEXT, + provider_type TEXT, + PRIMARY KEY (id, app_type) ON CONFLICT REPLACE + )"; + + const WEAK_PROVIDER_ENDPOINTS: &str = "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 + )"; + + fn replace_schema_object( + exported: &mut String, + source: &Connection, + object_type: &str, + name: &str, + replacement: &str, + ) -> Result<(), AppError> { + let original: String = source.query_row( + "SELECT sql FROM sqlite_schema WHERE type = ?1 AND name = ?2", + rusqlite::params![object_type, name], + |row| row.get(0), + )?; + let needle = format!("{original};"); + if !exported.contains(&needle) { + return Err(AppError::Database(format!( + "test export did not contain schema object {object_type}/{name}" + ))); + } + *exported = exported.replacen(&needle, &format!("{replacement};"), 1); + Ok(()) + } + + fn rewritten_source( + provider_ddl: Option<&str>, + endpoint_ddl: Option<&str>, + ) -> Result { + let base = canonical_restore_source()?; + let mut exported = Database::dump_sql(&base, &[])?; + if let Some(definition) = provider_ddl { + replace_schema_object(&mut exported, &base, "table", "providers", definition)?; + } + if let Some(definition) = endpoint_ddl { + replace_schema_object( + &mut exported, + &base, + "table", + "provider_endpoints", + definition, + )?; + } + let source = Connection::open_in_memory()?; + source.execute_batch(&exported)?; + Ok(source) + } + + fn seed_live_restore_state(target: &Database) -> Result<(), AppError> { + let conn = crate::database::lock_conn!(target.conn); + conn.execute_batch( + "INSERT INTO providers + (id, app_type, name, settings_config, meta, in_failover_queue) + VALUES ('live-provider', 'pi', 'Live Provider', '{}', '{}', 1); + INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at, last_used) + VALUES ('live-provider', 'pi', 'https://live.invalid', NULL, NULL); + INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES ('live-provider', 'live-key', 20, 20); + INSERT INTO skill_deployments ( + app_type, skill_id, destination, destination_key, method, + source_identity, created_at, updated_at + ) VALUES ( + 'pi', 'live-skill', '/live', '/live', 'copy', + 'live-source', 20, 20 + );", + )?; + Ok(()) + } + + #[derive(Debug, Clone, Copy)] + enum InvalidRestoreCase { + ProviderJson, + StorageClass, + DuplicateProvider, + DuplicateEndpoint, + ForeignKeyOrphan, + FutureVersion, + } + + impl InvalidRestoreCase { + const ALL: [Self; 6] = [ + Self::ProviderJson, + Self::StorageClass, + Self::DuplicateProvider, + Self::DuplicateEndpoint, + Self::ForeignKeyOrphan, + Self::FutureVersion, + ]; + + fn label(self) -> &'static str { + match self { + Self::ProviderJson => "provider-json", + Self::StorageClass => "storage-class", + Self::DuplicateProvider => "duplicate-provider", + Self::DuplicateEndpoint => "duplicate-endpoint", + Self::ForeignKeyOrphan => "foreign-key-orphan", + Self::FutureVersion => "future-version", + } + } + } + + fn invalid_restore_source(case: InvalidRestoreCase) -> Result { + let source = match case { + InvalidRestoreCase::DuplicateProvider => { + rewritten_source(Some(WEAK_PROVIDERS_NO_KEY), None)? + } + InvalidRestoreCase::DuplicateEndpoint => { + rewritten_source(None, Some(WEAK_PROVIDER_ENDPOINTS))? + } + _ => canonical_restore_source()?, + }; + match case { + InvalidRestoreCase::ProviderJson => { + source.execute( + "UPDATE providers SET settings_config = '{' WHERE id = 'remote-provider'", + [], + )?; + } + InvalidRestoreCase::StorageClass => { + source.execute( + "UPDATE providers SET created_at = X'00' WHERE id = 'remote-provider'", + [], + )?; + } + InvalidRestoreCase::DuplicateProvider => { + source.execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES ('remote-provider', 'pi', 'Duplicate', '{}', '{}')", + [], + )?; + } + InvalidRestoreCase::DuplicateEndpoint => { + source.execute_batch( + "INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at, last_used) + VALUES + ('remote-provider', 'pi', 'https://duplicate.invalid', NULL, NULL), + ('remote-provider', 'pi', 'https://duplicate.invalid', 1, 2);", + )?; + } + InvalidRestoreCase::ForeignKeyOrphan => { + source.execute_batch( + "PRAGMA foreign_keys = OFF; + INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at, last_used) + VALUES ('missing', 'pi', 'https://orphan.invalid', NULL, NULL); + PRAGMA foreign_keys = ON;", + )?; + } + InvalidRestoreCase::FutureVersion => { + Database::set_user_version(&source, SCHEMA_VERSION + 1)?; + } + } + Ok(source) + } + + fn assert_weak_ledgers_are_rebuilt( + entry_point: RestoreEntryPoint, + filename: &str, + ) -> Result<(), AppError> { + let source = weak_ledger_source()?; + let target = Database::memory()?; + seed_local_ledgers(&target)?; + run_restore_entry(&target, &source, entry_point, filename)?; + + let conn = crate::database::lock_conn!(target.conn); + validate_stage_rows(&conn)?; + validate_canonical_behaviors(&conn)?; + let counts: (i64, i64, i64, i64, i64) = conn.query_row( + "SELECT + (SELECT COUNT(*) FROM providers + WHERE id = 'remote-provider' AND app_type = 'pi'), + (SELECT COUNT(*) FROM pi_provider_projections + WHERE provider_id = 'local-provider' AND provider_key = 'local-key'), + (SELECT COUNT(*) FROM pi_provider_projections + WHERE provider_id = 'remote-provider' OR provider_key = 'remote-key'), + (SELECT COUNT(*) FROM skill_deployments + WHERE skill_id = 'local-skill' AND destination_key = '/local'), + (SELECT COUNT(*) FROM skill_deployments + WHERE skill_id = 'remote-skill')", + [], + |row| { + Ok(( + row.get(0)?, + row.get(1)?, + row.get(2)?, + row.get(3)?, + row.get(4)?, + )) + }, + )?; + assert_eq!(counts, (1, 1, 0, 1, 0)); + Ok(()) + } + + fn assert_invalid_live_copy_is_atomic( + entry_point: RestoreEntryPoint, + filename: &str, + ) -> Result<(), AppError> { + let source = canonical_restore_source()?; + let target = Database::memory()?; + { + let conn = crate::database::lock_conn!(target.conn); + conn.execute_batch( + "PRAGMA ignore_check_constraints = ON; + INSERT INTO skill_deployments ( + app_type, skill_id, destination, destination_key, method, + source_identity, created_at, updated_at + ) VALUES ( + 'codex', 'invalid-local', '/invalid', '/invalid', 'move', + 'invalid-source', 1, 1 + ); + PRAGMA ignore_check_constraints = OFF;", + )?; + } + let before = logical_snapshot(&target)?; + assert!(run_restore_entry(&target, &source, entry_point, filename).is_err()); + assert_eq!(logical_snapshot(&target)?, before); + Ok(()) + } + + fn assert_weak_endpoint_is_canonicalized( + entry_point: RestoreEntryPoint, + filename: &str, + ) -> Result<(), AppError> { + let source = weak_endpoint_source()?; + let target = Database::memory()?; + seed_local_ledgers(&target)?; + run_restore_entry(&target, &source, entry_point, filename)?; + let conn = crate::database::lock_conn!(target.conn); + validate_stage_rows(&conn)?; + validate_canonical_behaviors(&conn)?; + let endpoint_count: i64 = conn.query_row( + "SELECT COUNT(*) FROM provider_endpoints + WHERE provider_id = 'remote-provider' + AND app_type = 'pi' + AND url = 'https://weak.test'", + [], + |row| row.get(0), + )?; + assert_eq!(endpoint_count, 1); + Ok(()) + } + + #[test] + #[serial] + fn sql_and_binary_restore_share_canonical_prepublication_contracts() -> Result<(), AppError> { + let test_home = tempfile::tempdir().map_err(|error| AppError::IoContext { + context: "create canonical restore test home".to_string(), + source: error, + })?; + let _home_guard = TestHomeGuard::set(test_home.path()); + + assert_weak_ledgers_are_rebuilt(RestoreEntryPoint::Sql, "unused-ledger.db")?; + assert_invalid_live_copy_is_atomic(RestoreEntryPoint::Sql, "unused-invalid.db")?; + assert_weak_endpoint_is_canonicalized(RestoreEntryPoint::Sql, "unused-endpoint.db")?; + + assert_weak_ledgers_are_rebuilt(RestoreEntryPoint::Binary, "ledger.db")?; + assert_invalid_live_copy_is_atomic(RestoreEntryPoint::Binary, "invalid-copy.db")?; + assert_weak_endpoint_is_canonicalized(RestoreEntryPoint::Binary, "weak-endpoint.db")?; + Ok(()) + } + + #[test] + #[serial] + fn public_restore_entries_discard_hostile_schema_and_publish_only_canonical_objects( + ) -> Result<(), AppError> { + let test_home = tempfile::tempdir().map_err(|error| AppError::IoContext { + context: "create hostile schema restore home".to_string(), + source: error, + })?; + let _home_guard = TestHomeGuard::set(test_home.path()); + let source = rewritten_source( + Some(HOSTILE_PROVIDERS_NOCASE_REPLACE), + Some(WEAK_PROVIDER_ENDPOINTS), + )?; + source.execute_batch( + "CREATE INDEX source_leak_index ON providers(name); + CREATE VIEW source_leak_view AS SELECT id, app_type FROM providers; + CREATE TRIGGER source_leak_trigger + AFTER INSERT ON provider_endpoints + BEGIN + UPDATE settings SET value = NEW.url WHERE key = 'source-leak'; + END; + INSERT INTO provider_endpoints + (id, provider_id, app_type, url, added_at, last_used) + VALUES + (7001, 'remote-provider', 'pi', 'https://weak.invalid', NULL, 8);", + )?; + + for (index, entry_point) in [RestoreEntryPoint::Sql, RestoreEntryPoint::Binary] + .into_iter() + .enumerate() + { + let target = Database::memory()?; + seed_local_ledgers(&target)?; + run_restore_entry( + &target, + &source, + entry_point, + &format!("hostile-schema-{index}.db"), + )?; + + let conn = crate::database::lock_conn!(target.conn); + validate_stage_rows(&conn)?; + validate_canonical_behaviors(&conn)?; + let leaked_objects: i64 = conn.query_row( + "SELECT COUNT(*) FROM sqlite_schema + WHERE name IN ( + 'source_leak_index', + 'source_leak_view', + 'source_leak_trigger' + )", + [], + |row| row.get(0), + )?; + assert_eq!(leaked_objects, 0, "source executable DDL must not publish"); + + let provider_schema: String = conn.query_row( + "SELECT sql FROM sqlite_schema + WHERE type = 'table' AND name = 'providers'", + [], + |row| row.get(0), + )?; + assert!(!provider_schema.to_ascii_uppercase().contains("NOCASE")); + assert!(!provider_schema.to_ascii_uppercase().contains("REPLACE")); + + let endpoint_schema: String = conn.query_row( + "SELECT sql FROM sqlite_schema + WHERE type = 'table' AND name = 'provider_endpoints'", + [], + |row| row.get(0), + )?; + let endpoint_schema = endpoint_schema.to_ascii_uppercase(); + assert!(endpoint_schema.contains("FOREIGN KEY")); + assert!(endpoint_schema.contains("UNIQUE")); + let endpoint: (i64, Option, Option) = conn.query_row( + "SELECT id, added_at, last_used FROM provider_endpoints + WHERE provider_id = 'remote-provider' + AND app_type = 'pi' + AND url = 'https://weak.invalid'", + [], + |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)), + )?; + assert_eq!(endpoint, (7001, None, Some(8))); + } + Ok(()) + } + + #[test] + #[serial] + fn public_restore_entries_abort_invalid_rows_without_live_or_ledger_mutation( + ) -> Result<(), AppError> { + let test_home = tempfile::tempdir().map_err(|error| AppError::IoContext { + context: "create invalid restore matrix home".to_string(), + source: error, + })?; + let _home_guard = TestHomeGuard::set(test_home.path()); + + for case in InvalidRestoreCase::ALL { + for (entry_index, entry_point) in [RestoreEntryPoint::Sql, RestoreEntryPoint::Binary] + .into_iter() + .enumerate() + { + let source = invalid_restore_source(case)?; + let target = Database::memory()?; + seed_live_restore_state(&target)?; + let before = logical_snapshot(&target)?; + let result = run_restore_entry( + &target, + &source, + entry_point, + &format!("invalid-{}-{entry_index}.db", case.label()), + ); + assert!( + result.is_err(), + "{} via entry {entry_index} must fail closed", + case.label() + ); + assert_eq!( + logical_snapshot(&target)?, + before, + "{} via entry {entry_index} changed live state", + case.label() + ); + } + } + Ok(()) + } + + #[test] + #[serial] + fn public_restore_entries_preserve_nulls_unknown_json_and_explicit_ids() -> Result<(), AppError> + { + let test_home = tempfile::tempdir().map_err(|error| AppError::IoContext { + context: "create lossless restore matrix home".to_string(), + source: error, + })?; + let _home_guard = TestHomeGuard::set(test_home.path()); + let source = canonical_restore_source()?; + let settings = r#"{"api":"future-api","futureShape":{"nested":[1,2,3]}}"#; + let meta = r#"{"futureMeta":{"opaque":true}}"#; + source.execute( + "UPDATE providers + SET settings_config = ?1, meta = ?2 + WHERE id = 'remote-provider' AND app_type = 'pi'", + rusqlite::params![settings, meta], + )?; + source.execute( + "INSERT INTO provider_endpoints + (id, provider_id, app_type, url, added_at, last_used) + VALUES (9001, 'remote-provider', 'pi', 'https://null.invalid', NULL, NULL)", + [], + )?; + + for (index, entry_point) in [RestoreEntryPoint::Sql, RestoreEntryPoint::Binary] + .into_iter() + .enumerate() + { + let target = Database::memory()?; + run_restore_entry( + &target, + &source, + entry_point, + &format!("lossless-{index}.db"), + )?; + let conn = crate::database::lock_conn!(target.conn); + let restored: (i64, Option, Option, String, String) = conn.query_row( + "SELECT e.id, e.added_at, e.last_used, p.settings_config, p.meta + FROM provider_endpoints AS e + JOIN providers AS p + ON p.id = e.provider_id AND p.app_type = e.app_type + WHERE e.url = 'https://null.invalid'", + [], + |row| { + Ok(( + row.get(0)?, + row.get(1)?, + row.get(2)?, + row.get(3)?, + row.get(4)?, + )) + }, + )?; + assert_eq!( + restored, + (9001, None, None, settings.to_string(), meta.to_string()) + ); + conn.execute( + "INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at, last_used) + VALUES ( + 'remote-provider', + 'pi', + 'https://next-id.invalid', + NULL, + NULL + )", + [], + )?; + assert!( + conn.last_insert_rowid() > 9001, + "AUTOINCREMENT must advance without copying sqlite_sequence" + ); + } + Ok(()) + } + + #[test] + #[serial] + fn every_supported_user_version_has_a_public_migration_sentinel() -> Result<(), AppError> { + let test_home = tempfile::tempdir().map_err(|error| AppError::IoContext { + context: "create restore migration matrix home".to_string(), + source: error, + })?; + let _home_guard = TestHomeGuard::set(test_home.path()); + + for version in 0..=SCHEMA_VERSION { + let source_db = Database::memory()?; + { + let conn = crate::database::lock_conn!(source_db.conn); + conn.execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES (?1, 'pi', ?2, '{}', '{}')", + rusqlite::params![ + format!("migration-v{version}"), + format!("Migration v{version}") + ], + )?; + Database::set_user_version(&conn, version)?; + } + let source = source_db.snapshot_to_memory()?; + let target = Database::memory()?; + run_restore_entry( + &target, + &source, + RestoreEntryPoint::Sql, + &format!("unused-migration-v{version}.db"), + )?; + let conn = crate::database::lock_conn!(target.conn); + let restored: i64 = conn.query_row( + "SELECT COUNT(*) FROM providers WHERE id = ?1 AND app_type = 'pi'", + [format!("migration-v{version}")], + |row| row.get(0), + )?; + assert_eq!(restored, 1, "SQL migration sentinel v{version}"); + } + + for version in [0, SCHEMA_VERSION - 1, SCHEMA_VERSION] { + let source_db = Database::memory()?; + { + let conn = crate::database::lock_conn!(source_db.conn); + conn.execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES (?1, 'pi', ?2, '{}', '{}')", + rusqlite::params![ + format!("binary-migration-v{version}"), + format!("Binary migration v{version}") + ], + )?; + Database::set_user_version(&conn, version)?; + } + let source = source_db.snapshot_to_memory()?; + let target = Database::memory()?; + run_restore_entry( + &target, + &source, + RestoreEntryPoint::Binary, + &format!("migration-v{version}.db"), + )?; + let conn = crate::database::lock_conn!(target.conn); + let restored: i64 = conn.query_row( + "SELECT COUNT(*) FROM providers WHERE id = ?1 AND app_type = 'pi'", + [format!("binary-migration-v{version}")], + |row| row.get(0), + )?; + assert_eq!(restored, 1, "binary migration sentinel v{version}"); + } + Ok(()) + } #[test] fn import_rejects_cross_file_statements_and_leaves_no_file_behind() -> Result<(), AppError> { @@ -788,6 +2807,171 @@ mod tests { Ok(()) } + #[test] + #[serial] + fn public_file_restore_entries_reject_symlink_directory_and_fifo() -> Result<(), AppError> { + let test_home = tempfile::tempdir().map_err(|error| AppError::IoContext { + context: "create restore file-shape home".to_string(), + source: error, + })?; + let _home_guard = TestHomeGuard::set(test_home.path()); + let database = Database::memory()?; + let backup_dir = crate::config::get_app_config_dir().join("backups"); + fs::create_dir_all(&backup_dir).map_err(|error| AppError::io(&backup_dir, error))?; + + let sql_directory = test_home.path().join("sql-directory"); + fs::create_dir(&sql_directory).map_err(|error| AppError::io(&sql_directory, error))?; + assert!(database.import_sql(&sql_directory).is_err()); + let binary_directory = backup_dir.join("binary-directory.db"); + fs::create_dir(&binary_directory) + .map_err(|error| AppError::io(&binary_directory, error))?; + assert!(database.restore_from_backup("binary-directory.db").is_err()); + + #[cfg(unix)] + { + use std::ffi::CString; + use std::os::unix::ffi::OsStrExt; + use std::os::unix::fs::symlink; + + let regular = test_home.path().join("regular-source"); + fs::write(®ular, super::CC_SWITCH_SQL_EXPORT_HEADER) + .map_err(|error| AppError::io(®ular, error))?; + let sql_symlink = test_home.path().join("symlink.sql"); + symlink(®ular, &sql_symlink).map_err(|error| AppError::io(&sql_symlink, error))?; + assert!(database.import_sql(&sql_symlink).is_err()); + + let binary_symlink = backup_dir.join("symlink.db"); + symlink(®ular, &binary_symlink) + .map_err(|error| AppError::io(&binary_symlink, error))?; + assert!(database.restore_from_backup("symlink.db").is_err()); + + for fifo in [ + test_home.path().join("source-fifo.sql"), + backup_dir.join("source-fifo.db"), + ] { + let path = CString::new(fifo.as_os_str().as_bytes()) + .map_err(|error| AppError::InvalidInput(error.to_string()))?; + let result = unsafe { libc::mkfifo(path.as_ptr(), 0o600) }; + if result != 0 { + return Err(AppError::io(&fifo, std::io::Error::last_os_error())); + } + } + assert!(database + .import_sql(&test_home.path().join("source-fifo.sql")) + .is_err()); + assert!(database.restore_from_backup("source-fifo.db").is_err()); + } + + Ok(()) + } + + #[test] + #[serial] + fn restore_file_size_limits_accept_n_and_publicly_reject_n_plus_one() -> Result<(), AppError> { + let test_home = tempfile::tempdir().map_err(|error| AppError::IoContext { + context: "create restore size-boundary home".to_string(), + source: error, + })?; + let _home_guard = TestHomeGuard::set(test_home.path()); + let database = Database::memory()?; + let backup_dir = crate::config::get_app_config_dir().join("backups"); + fs::create_dir_all(&backup_dir).map_err(|error| AppError::io(&backup_dir, error))?; + + let sql_n = test_home.path().join("sql-n"); + File::create(&sql_n) + .and_then(|file| file.set_len(MAX_SQL_IMPORT_BYTES)) + .map_err(|error| AppError::io(&sql_n, error))?; + assert_eq!( + validate_regular_file(&sql_n, MAX_SQL_IMPORT_BYTES)?.len(), + MAX_SQL_IMPORT_BYTES + ); + let sql_n_plus_one = test_home.path().join("sql-n-plus-one"); + File::create(&sql_n_plus_one) + .and_then(|file| file.set_len(MAX_SQL_IMPORT_BYTES + 1)) + .map_err(|error| AppError::io(&sql_n_plus_one, error))?; + assert!(database.import_sql(&sql_n_plus_one).is_err()); + + let binary_n = backup_dir.join("binary-n.db"); + File::create(&binary_n) + .and_then(|file| file.set_len(MAX_BINARY_RESTORE_BYTES)) + .map_err(|error| AppError::io(&binary_n, error))?; + assert_eq!( + validate_regular_file(&binary_n, MAX_BINARY_RESTORE_BYTES)?.len(), + MAX_BINARY_RESTORE_BYTES + ); + let binary_n_plus_one = backup_dir.join("binary-n-plus-one.db"); + File::create(&binary_n_plus_one) + .and_then(|file| file.set_len(MAX_BINARY_RESTORE_BYTES + 1)) + .map_err(|error| AppError::io(&binary_n_plus_one, error))?; + assert!(database + .restore_from_backup("binary-n-plus-one.db") + .is_err()); + Ok(()) + } + + #[test] + #[serial] + fn public_restore_entries_enforce_vm_and_page_budgets() -> Result<(), AppError> { + let test_home = tempfile::tempdir().map_err(|error| AppError::IoContext { + context: "create restore budget home".to_string(), + source: error, + })?; + let _home_guard = TestHomeGuard::set(test_home.path()); + let database = Database::memory()?; + + { + let _limit_guard = RestoreLimitGuard::set(Some(1_000), None); + let expensive = format!( + "{}\n\ + WITH RECURSIVE counter(value) AS (\n\ + VALUES(0)\n\ + UNION ALL SELECT value + 1 FROM counter WHERE value < 100000\n\ + ) SELECT SUM(value) FROM counter;", + super::CC_SWITCH_SQL_EXPORT_HEADER + ); + let error = database + .import_sql_string(&expensive) + .expect_err("VM budget must interrupt untrusted SQL"); + assert!( + error.to_string().to_ascii_lowercase().contains("interrupt"), + "unexpected VM budget error: {error}" + ); + } + + { + let _limit_guard = RestoreLimitGuard::set(None, Some(8)); + let page_heavy = format!( + "{}\n\ + CREATE TABLE filler (payload BLOB);\n\ + INSERT INTO filler(payload) VALUES (zeroblob(1048576));", + super::CC_SWITCH_SQL_EXPORT_HEADER + ); + assert!( + database.import_sql_string(&page_heavy).is_err(), + "SQL page budget must reject oversized scratch growth" + ); + } + + let backup_dir = crate::config::get_app_config_dir().join("backups"); + fs::create_dir_all(&backup_dir).map_err(|error| AppError::io(&backup_dir, error))?; + let binary_path = backup_dir.join("page-budget.db"); + let source = canonical_restore_source()?; + let mut destination = Connection::open(&binary_path)?; + { + let backup = Backup::new(&source, &mut destination)?; + backup.step(-1)?; + } + drop(destination); + { + let _limit_guard = RestoreLimitGuard::set(None, Some(8)); + assert!( + database.restore_from_backup("page-budget.db").is_err(), + "binary page budget must reject oversized scratch growth" + ); + } + Ok(()) + } + #[test] fn import_still_accepts_a_genuine_export() -> Result<(), AppError> { // 白名单收得紧,必须有一条回归防线证明它没误伤自家导出格式—— @@ -816,6 +3000,166 @@ mod tests { Ok(()) } + #[test] + fn portable_export_and_import_preserve_device_local_pi_ledgers() -> Result<(), AppError> { + let source = Database::memory()?; + { + let conn = crate::database::lock_conn!(source.conn); + conn.execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES ('remote', 'pi', 'Remote', '{}', '{}')", + [], + )?; + conn.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES ('remote', 'remote-key', 1, 1)", + [], + )?; + conn.execute( + "INSERT INTO skill_deployments ( + app_type, skill_id, destination, destination_key, method, + source_identity, created_at, updated_at + ) VALUES ('pi', 'remote-skill', '/remote', '/remote', + 'copy', 'remote-source', 1, 1)", + [], + )?; + } + let exported = source.export_sql_string()?; + assert!(!exported.contains("INSERT INTO \"pi_provider_projections\"")); + assert!(!exported.contains("INSERT INTO \"skill_deployments\"")); + + let target = Database::memory()?; + { + let conn = crate::database::lock_conn!(target.conn); + conn.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES ('local', 'local-key', 2, 2)", + [], + )?; + conn.execute( + "INSERT INTO skill_deployments ( + app_type, skill_id, destination, destination_key, method, + source_identity, created_at, updated_at + ) VALUES ('pi', 'local-skill', '/local', '/local', + 'symlink', 'local-source', 2, 2)", + [], + )?; + } + target.import_sql_string(&exported)?; + + let conn = crate::database::lock_conn!(target.conn); + let counts: (i64, i64, i64) = conn.query_row( + "SELECT + (SELECT COUNT(*) FROM providers WHERE id = 'remote' AND app_type = 'pi'), + (SELECT COUNT(*) FROM pi_provider_projections + WHERE provider_id = 'local' AND provider_key = 'local-key'), + (SELECT COUNT(*) FROM skill_deployments + WHERE skill_id = 'local-skill' AND destination_key = '/local')", + [], + |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)), + )?; + assert_eq!(counts, (1, 1, 1)); + Ok(()) + } + + #[test] + fn publish_copies_device_local_ledgers_at_the_commit_boundary() -> Result<(), AppError> { + let remote = Database::memory()?; + { + let conn = crate::database::lock_conn!(remote.conn); + conn.execute( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES ('remote', 'pi', 'Remote', '{}', '{}')", + [], + )?; + } + let exported = remote.export_sql_string()?; + let scratch = UntrustedScratch::from_sql(&exported)?; + let stage = Database::build_canonical_stage(&scratch)?; + + let target = Database::memory()?; + { + let conn = crate::database::lock_conn!(target.conn); + conn.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES ('created-after-staging', 'local-key', 2, 2)", + [], + )?; + } + + target.publish_canonical_stage(stage, RestoreFlavor::UserRestore)?; + + let conn = crate::database::lock_conn!(target.conn); + let counts: (i64, i64) = conn.query_row( + "SELECT + (SELECT COUNT(*) FROM providers WHERE id = 'remote' AND app_type = 'pi'), + (SELECT COUNT(*) FROM pi_provider_projections + WHERE provider_id = 'created-after-staging' AND provider_key = 'local-key')", + [], + |row| Ok((row.get(0)?, row.get(1)?)), + )?; + assert_eq!(counts, (1, 1)); + Ok(()) + } + + #[test] + fn public_sql_restore_discards_input_trigger_before_local_copy() -> Result<(), AppError> { + let staged_db = Database::memory()?; + { + let conn = crate::database::lock_conn!(staged_db.conn); + conn.execute_batch( + "INSERT INTO providers (id, app_type, name, settings_config, meta) + VALUES ('remote', 'pi', 'Remote', '{}', '{}'); + CREATE TRIGGER leak_local_projection + AFTER INSERT ON pi_provider_projections + BEGIN + INSERT OR REPLACE INTO settings (key, value) + VALUES ('leaked-provider-key', NEW.provider_key); + END;", + )?; + } + let exported = staged_db.export_sql_string()?; + + let target = Database::memory()?; + { + let conn = crate::database::lock_conn!(target.conn); + conn.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES ('local', 'local-secret-key', 1, 1)", + [], + )?; + } + + target.import_sql_string(&exported)?; + + let conn = crate::database::lock_conn!(target.conn); + let local_rows: i64 = conn.query_row( + "SELECT COUNT(*) FROM pi_provider_projections + WHERE provider_id = 'local' AND provider_key = 'local-secret-key'", + [], + |row| row.get(0), + )?; + let leaked_rows: i64 = conn.query_row( + "SELECT COUNT(*) FROM settings WHERE key = 'leaked-provider-key'", + [], + |row| row.get(0), + )?; + let trigger_rows: i64 = conn.query_row( + "SELECT COUNT(*) FROM sqlite_schema + WHERE type = 'trigger' AND name = 'leak_local_projection'", + [], + |row| row.get(0), + )?; + assert_eq!(local_rows, 1); + assert_eq!(leaked_rows, 0); + assert_eq!(trigger_rows, 0); + Ok(()) + } + #[test] fn sync_import_preserves_local_only_tables() -> Result<(), AppError> { let remote_db = Database::memory()?; diff --git a/src-tauri/src/database/dao/mod.rs b/src-tauri/src/database/dao/mod.rs index bb78f3075..b1a7ea941 100644 --- a/src-tauri/src/database/dao/mod.rs +++ b/src-tauri/src/database/dao/mod.rs @@ -4,12 +4,15 @@ pub mod failover; pub mod mcp; +pub mod pi_projections; pub mod profiles; pub mod prompts; +pub mod provider_write; pub mod providers; pub mod providers_seed; pub mod proxy; pub mod settings; +pub mod skill_deployments; pub mod skills; pub mod stream_check; pub mod universal_providers; diff --git a/src-tauri/src/database/dao/pi_projections.rs b/src-tauri/src/database/dao/pi_projections.rs new file mode 100644 index 000000000..59007f286 --- /dev/null +++ b/src-tauri/src/database/dao/pi_projections.rs @@ -0,0 +1,204 @@ +//! Device-local ownership ledger for exact keys in Pi's shared models.json. + +// The projection writer is introduced in a later contract-ordered commit. +#![allow(dead_code)] + +use crate::database::{lock_conn, Database}; +use crate::error::AppError; +use indexmap::IndexMap; +use rusqlite::{params, OptionalExtension}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct PiProviderProjection { + pub provider_id: String, + pub provider_key: String, + pub created_at: i64, + pub updated_at: i64, +} + +fn decode_projection(row: &rusqlite::Row<'_>) -> rusqlite::Result { + Ok(PiProviderProjection { + provider_id: row.get(0)?, + provider_key: row.get(1)?, + created_at: row.get(2)?, + updated_at: row.get(3)?, + }) +} + +impl Database { + pub(crate) fn get_pi_projection( + &self, + provider_id: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections WHERE provider_id = ?1", + [provider_id], + decode_projection, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn get_pi_projection_for_key( + &self, + provider_key: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections WHERE provider_key = ?1", + [provider_key], + decode_projection, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn get_pi_projection_manifest( + &self, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + let mut stmt = conn + .prepare( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections ORDER BY provider_id", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let rows = stmt + .query_map([], decode_projection) + .map_err(|error| AppError::Database(error.to_string()))?; + let mut manifest = IndexMap::new(); + for row in rows { + let projection = row.map_err(|error| AppError::Database(error.to_string()))?; + manifest.insert(projection.provider_id.clone(), projection); + } + Ok(manifest) + } + + /// Claim an exact key. Existing exact claims are idempotent; either-side + /// collisions fail and are never rewritten. + pub(crate) fn claim_pi_projection_key( + &self, + provider_id: &str, + provider_key: &str, + ) -> Result { + if provider_id.trim().is_empty() || provider_key.trim().is_empty() { + return Err(AppError::Config( + "Pi projection provider id and key must be non-empty".to_string(), + )); + } + let mut conn = lock_conn!(self.conn); + let tx = conn + .transaction() + .map_err(|error| AppError::Database(error.to_string()))?; + let by_provider = tx + .query_row( + "SELECT provider_id, provider_key, created_at, updated_at + FROM pi_provider_projections WHERE provider_id = ?1", + [provider_id], + decode_projection, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))?; + if let Some(existing) = by_provider { + if existing.provider_key != provider_key { + return Err(AppError::Config(format!( + "Pi provider '{provider_id}' already owns key '{}', not '{provider_key}'", + existing.provider_key + ))); + } + tx.commit() + .map_err(|error| AppError::Database(error.to_string()))?; + return Ok(existing); + } + if let Some(existing_owner) = tx + .query_row( + "SELECT provider_id FROM pi_provider_projections WHERE provider_key = ?1", + [provider_key], + |row| row.get::<_, String>(0), + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))? + { + return Err(AppError::Config(format!( + "Pi key '{provider_key}' is already owned by provider '{existing_owner}'" + ))); + } + let now = chrono::Utc::now().timestamp_millis(); + tx.execute( + "INSERT INTO pi_provider_projections + (provider_id, provider_key, created_at, updated_at) + VALUES (?1, ?2, ?3, ?3)", + params![provider_id, provider_key, now], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + tx.commit() + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(PiProviderProjection { + provider_id: provider_id.to_string(), + provider_key: provider_key.to_string(), + created_at: now, + updated_at: now, + }) + } + + pub(crate) fn delete_pi_projection_key( + &self, + provider_id: &str, + expected_key: &str, + ) -> Result { + let conn = lock_conn!(self.conn); + let removed = conn + .execute( + "DELETE FROM pi_provider_projections + WHERE provider_id = ?1 AND provider_key = ?2", + params![provider_id, expected_key], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + if removed == 0 + && conn + .query_row( + "SELECT 1 FROM pi_provider_projections WHERE provider_id = ?1", + [provider_id], + |_| Ok(()), + ) + .optional() + .map_err(|error| AppError::Database(error.to_string()))? + .is_some() + { + return Err(AppError::Config(format!( + "refusing to delete Pi projection '{provider_id}': expected key changed" + ))); + } + Ok(removed == 1) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn projection_claims_are_exact_idempotent_and_collision_safe() -> Result<(), AppError> { + let db = Database::memory()?; + let first = db.claim_pi_projection_key("provider-a", "native-a")?; + let repeated = db.claim_pi_projection_key("provider-a", "native-a")?; + assert_eq!(first, repeated); + assert!(db + .claim_pi_projection_key("provider-a", "native-b") + .is_err()); + assert!(db + .claim_pi_projection_key("provider-b", "native-a") + .is_err()); + assert_eq!(db.get_pi_projection_manifest()?.len(), 1); + assert!(db.delete_pi_projection_key("provider-a", "wrong").is_err()); + assert!(db.get_pi_projection("provider-a")?.is_some()); + assert!(db.delete_pi_projection_key("provider-a", "native-a")?); + assert!(db.get_pi_projection_for_key("native-a")?.is_none()); + Ok(()) + } +} 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..6b2d0e21a --- /dev/null +++ b/src-tauri/src/database/dao/provider_write.rs @@ -0,0 +1,531 @@ +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; + +#[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, + created_at: 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(), + created_at: input.created_at, + notes: input.notes.clone(), + meta, + icon: input.icon.clone(), + icon_color: input.icon_color.clone(), + }) + } +} + +#[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: ProviderRowUpdate, + 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 { + if key != endpoint.url { + return Err(AppError::InvalidInput(format!( + "provider endpoint key '{key}' must match endpoint URL '{}'", + endpoint.url + ))); + } + if !seen.insert(endpoint.url.clone()) { + return Err(AppError::InvalidInput(format!( + "duplicate initial provider endpoint '{}'", + endpoint.url + ))); + } + initial_endpoints.push(endpoint.try_into()?); + } + let key = ProviderKey::new(app_type, input.id.clone())?; + let row = ProviderRowUpdate::from_input(&input)?; + 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(), + )); + } + Ok(Self { + source, + target_id: input.id.clone(), + row: ProviderRowUpdate::from_input(input)?, + }) + } +} + +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, + 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, + row.created_at, + sort_index, + row.notes, + row.icon, + row.icon_color, + meta, + is_current, + in_failover_queue, + ], + ) + .map_err(|error| 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)] +pub(super) fn restore_provider_aggregate_on_tx( + tx: &Transaction<'_>, + key: &ProviderKey, + row: &ProviderRowUpdate, + sort_index: Option, + is_current: bool, + in_failover_queue: bool, + endpoints: &[NewEndpoint], +) -> Result<(), AppError> { + let updated = update_row(tx, key, row)?; + if updated != 1 { + return Err(AppError::NotFound(format!( + "provider '{}/{}'", + key.app_type, key.id + ))); + } + 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, + created_at = ?5, + notes = ?6, + icon = ?7, + icon_color = ?8, + meta = ?9 + WHERE id = ?10 AND app_type = ?11", + params![ + row.name, + settings_config, + row.website_url, + row.category, + row.created_at, + 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, + 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 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 + 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)?, + )) + }, + ) + .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 target = ProviderKey::new(&input.source.app_type, &input.target_id)?; + insert_row( + &tx, + &target, + &input.row, + 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/providers.rs b/src-tauri/src/database/dao/providers.rs index f64154b06..e42d8d893 100644 --- a/src-tauri/src/database/dao/providers.rs +++ b/src-tauri/src/database/dao/providers.rs @@ -1,111 +1,223 @@ -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, -); +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 { + 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)?, + }) + } + + 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)) +} + +/// Restore uses the exact production hydration decoder but copies the original +/// JSON bytes unchanged after validation, preserving unknown fields. +pub(crate) fn validate_provider_storage_json( + app_type: &str, + provider_id: &str, + settings_config: &str, + meta: &str, +) -> Result<(), AppError> { + decode_provider_json(app_type, provider_id, settings_config, meta).map(|_| ()) +} + +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 +244,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> { @@ -330,36 +327,6 @@ impl Database { 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 +410,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 +519,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 +549,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 +619,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 +685,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 +760,343 @@ 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.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/dao/skill_deployments.rs b/src-tauri/src/database/dao/skill_deployments.rs new file mode 100644 index 000000000..a5890a52f --- /dev/null +++ b/src-tauri/src/database/dao/skill_deployments.rs @@ -0,0 +1,208 @@ +//! Device-local evidence for Pi Skill deployments. + +// Pi skill reconciliation consumes this ledger in a later contract-ordered commit. +#![allow(dead_code)] + +use crate::database::{lock_conn, Database}; +use crate::error::AppError; +use rusqlite::{params, OptionalExtension}; +use serde::{Deserialize, Serialize}; +use std::str::FromStr; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum SkillDeploymentMethod { + Symlink, + Copy, +} + +impl SkillDeploymentMethod { + pub(crate) const fn as_str(self) -> &'static str { + match self { + Self::Symlink => "symlink", + Self::Copy => "copy", + } + } +} + +impl FromStr for SkillDeploymentMethod { + type Err = AppError; + + fn from_str(value: &str) -> Result { + match value { + "symlink" => Ok(Self::Symlink), + "copy" => Ok(Self::Copy), + _ => Err(AppError::Database(format!( + "unknown Pi Skill deployment method '{value}'" + ))), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct SkillDeployment { + pub skill_id: String, + pub destination: String, + pub destination_key: String, + pub method: SkillDeploymentMethod, + pub source_identity: String, + pub deployed_digest: Option, + pub created_at: i64, + pub updated_at: i64, +} + +fn decode_deployment(row: &rusqlite::Row<'_>) -> rusqlite::Result { + let method: String = row.get(3)?; + let method = method.parse().map_err(|error: AppError| { + rusqlite::Error::FromSqlConversionFailure(3, rusqlite::types::Type::Text, Box::new(error)) + })?; + Ok(SkillDeployment { + skill_id: row.get(0)?, + destination: row.get(1)?, + destination_key: row.get(2)?, + method, + source_identity: row.get(4)?, + deployed_digest: row.get(5)?, + created_at: row.get(6)?, + updated_at: row.get(7)?, + }) +} + +impl Database { + pub(crate) fn get_pi_skill_deployment( + &self, + skill_id: &str, + destination_key: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + conn.query_row( + "SELECT skill_id, destination, destination_key, method, + source_identity, deployed_digest, created_at, updated_at + FROM skill_deployments + WHERE app_type = 'pi' AND skill_id = ?1 AND destination_key = ?2", + params![skill_id, destination_key], + decode_deployment, + ) + .optional() + .map_err(|error| AppError::Database(error.to_string())) + } + + pub(crate) fn get_pi_skill_deployments( + &self, + skill_id: &str, + ) -> Result, AppError> { + let conn = lock_conn!(self.conn); + let mut stmt = conn + .prepare( + "SELECT skill_id, destination, destination_key, method, + source_identity, deployed_digest, created_at, updated_at + FROM skill_deployments + WHERE app_type = 'pi' AND skill_id = ?1 + ORDER BY created_at, destination_key", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + let rows = stmt + .query_map([skill_id], decode_deployment) + .map_err(|error| AppError::Database(error.to_string()))?; + rows.map(|row| row.map_err(|error| AppError::Database(error.to_string()))) + .collect() + } + + pub(crate) fn save_pi_skill_deployment( + &self, + deployment: &SkillDeployment, + ) -> Result<(), AppError> { + if deployment.skill_id.trim().is_empty() + || deployment.destination.trim().is_empty() + || deployment.destination_key.trim().is_empty() + || deployment.source_identity.trim().is_empty() + { + return Err(AppError::Config( + "Pi Skill deployment identity fields must be non-empty".to_string(), + )); + } + let conn = lock_conn!(self.conn); + conn.execute( + "INSERT INTO skill_deployments ( + app_type, skill_id, destination, destination_key, method, + source_identity, deployed_digest, created_at, updated_at + ) VALUES ('pi', ?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8) + ON CONFLICT(app_type, skill_id, destination_key) DO UPDATE SET + destination = excluded.destination, + method = excluded.method, + source_identity = excluded.source_identity, + deployed_digest = excluded.deployed_digest, + updated_at = excluded.updated_at", + params![ + deployment.skill_id, + deployment.destination, + deployment.destination_key, + deployment.method.as_str(), + deployment.source_identity, + deployment.deployed_digest, + deployment.created_at, + deployment.updated_at, + ], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + Ok(()) + } + + pub(crate) fn delete_pi_skill_deployment( + &self, + skill_id: &str, + destination_key: &str, + ) -> Result { + let conn = lock_conn!(self.conn); + conn.execute( + "DELETE FROM skill_deployments + WHERE app_type = 'pi' AND skill_id = ?1 AND destination_key = ?2", + params![skill_id, destination_key], + ) + .map(|count| count == 1) + .map_err(|error| AppError::Database(error.to_string())) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn deployment(skill_id: &str, destination_key: &str) -> SkillDeployment { + SkillDeployment { + skill_id: skill_id.into(), + destination: format!("/tmp/{destination_key}"), + destination_key: destination_key.into(), + method: SkillDeploymentMethod::Copy, + source_identity: format!("source:{skill_id}"), + deployed_digest: Some("sha256:initial".into()), + created_at: 10, + updated_at: 10, + } + } + + #[test] + fn skill_ledger_preserves_created_at_and_rejects_destination_collision() -> Result<(), AppError> + { + let db = Database::memory()?; + db.save_pi_skill_deployment(&deployment("one", "destination"))?; + let mut updated = deployment("one", "destination"); + updated.updated_at = 20; + updated.deployed_digest = Some("sha256:updated".into()); + db.save_pi_skill_deployment(&updated)?; + let saved = db + .get_pi_skill_deployment("one", "destination")? + .expect("deployment"); + assert_eq!(saved.created_at, 10); + assert_eq!(saved.updated_at, 20); + assert_eq!(saved.deployed_digest.as_deref(), Some("sha256:updated")); + + assert!(db + .save_pi_skill_deployment(&deployment("two", "destination")) + .is_err()); + assert_eq!(db.get_pi_skill_deployments("one")?.len(), 1); + assert!(db.delete_pi_skill_deployment("one", "destination")?); + Ok(()) + } +} diff --git a/src-tauri/src/database/dao/skills.rs b/src-tauri/src/database/dao/skills.rs index 488fde29d..edb7afefa 100644 --- a/src-tauri/src/database/dao/skills.rs +++ b/src-tauri/src/database/dao/skills.rs @@ -109,11 +109,28 @@ impl Database { pub fn save_skill(&self, skill: &InstalledSkill) -> Result<(), AppError> { let conn = lock_conn!(self.conn); conn.execute( - "INSERT OR REPLACE INTO skills + "INSERT INTO skills (id, name, description, directory, repo_owner, repo_name, repo_branch, readme_url, enabled_claude, enabled_codex, enabled_gemini, enabled_grokbuild, enabled_opencode, enabled_hermes, installed_at, content_hash, updated_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17)", + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17) + ON CONFLICT(id) DO UPDATE SET + name = excluded.name, + description = excluded.description, + directory = excluded.directory, + repo_owner = excluded.repo_owner, + repo_name = excluded.repo_name, + repo_branch = excluded.repo_branch, + readme_url = excluded.readme_url, + enabled_claude = excluded.enabled_claude, + enabled_codex = excluded.enabled_codex, + enabled_gemini = excluded.enabled_gemini, + enabled_grokbuild = excluded.enabled_grokbuild, + enabled_opencode = excluded.enabled_opencode, + enabled_hermes = excluded.enabled_hermes, + installed_at = excluded.installed_at, + content_hash = excluded.content_hash, + updated_at = excluded.updated_at", params![ skill.id, skill.name, @@ -262,3 +279,52 @@ impl Database { Ok(count) } } + +#[cfg(test)] +mod tests { + use super::*; + + fn installed_skill() -> InstalledSkill { + InstalledSkill { + id: "owner/repo:skill".into(), + name: "Skill".into(), + description: Some("before".into()), + directory: "skill".into(), + repo_owner: Some("owner".into()), + repo_name: Some("repo".into()), + repo_branch: Some("main".into()), + readme_url: None, + apps: SkillApps::default(), + installed_at: 10, + content_hash: Some("sha256:before".into()), + updated_at: 11, + } + } + + #[test] + fn legacy_skill_save_preserves_pi_desired_state() -> Result<(), AppError> { + let db = Database::memory()?; + let mut skill = installed_skill(); + db.save_skill(&skill)?; + { + let conn = lock_conn!(db.conn); + conn.execute( + "UPDATE skills SET enabled_pi = 1 WHERE id = ?1", + [&skill.id], + )?; + } + + skill.name = "Updated".into(); + skill.content_hash = Some("sha256:after".into()); + db.save_skill(&skill)?; + + let conn = lock_conn!(db.conn); + let saved: (String, String, bool) = conn.query_row( + "SELECT name, content_hash, enabled_pi FROM skills WHERE id = ?1", + [&skill.id], + |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)), + )?; + assert_eq!(saved, ("Updated".into(), "sha256:after".into(), true)); + 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 0b3931283..5e542b8c7 100644 --- a/src-tauri/src/database/schema.rs +++ b/src-tauri/src/database/schema.rs @@ -6,6 +6,498 @@ use super::{lock_conn, Database, SCHEMA_VERSION}; use crate::error::AppError; use rusqlite::{params, Connection}; use serde::Serialize; +use tempfile::NamedTempFile; + +/// A disk-backed database whose schema was created from this binary's current +/// schema code. Its fields are intentionally private: untrusted connections +/// cannot be wrapped or converted into a publishable stage. +pub(super) struct CanonicalStage { + connection: Connection, + _file: NamedTempFile, +} + +impl CanonicalStage { + pub(super) fn connection(&self) -> &Connection { + &self.connection + } + + pub(super) fn connection_mut(&mut self) -> &mut Connection { + &mut self.connection + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum CanonicalRestoreClass { + MigrateAndValidate, + RebuildAndPreserveLocal, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct CanonicalColumnSpec { + pub name: &'static str, + pub data_type: &'static str, + pub not_null: bool, + pub default: Option<&'static str>, + pub pk_position: i64, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct CanonicalIndexedColumnSpec { + pub name: &'static str, + pub collation: &'static str, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct CanonicalUniqueSpec { + pub columns: &'static [CanonicalIndexedColumnSpec], +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct CanonicalForeignKeySpec { + pub from: &'static [&'static str], + pub table: &'static str, + pub to: &'static [&'static str], + pub on_update: &'static str, + pub on_delete: &'static str, + pub match_type: &'static str, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum SchemaInvariantKind { + EndpointIdentityUnique, + EndpointParentForeignKey, + EndpointDeleteCascade, + ProjectionProviderKeyUnique, + SkillAppTypePiOnly, + SkillMethodAllowed, + SkillDestinationOwnedOnce, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct SchemaInvariantSpec { + pub name: &'static str, + pub kind: SchemaInvariantKind, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct CanonicalTableSpec { + pub name: &'static str, + pub definition: &'static str, + pub restore_class: CanonicalRestoreClass, + pub columns: &'static [CanonicalColumnSpec], + pub unique_tuples: &'static [CanonicalUniqueSpec], + pub foreign_keys: &'static [CanonicalForeignKeySpec], + pub invariants: &'static [SchemaInvariantSpec], +} + +const PROVIDERS_COLUMNS: &[CanonicalColumnSpec] = &[ + CanonicalColumnSpec { + name: "id", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 1, + }, + CanonicalColumnSpec { + name: "app_type", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 2, + }, + CanonicalColumnSpec { + name: "name", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "settings_config", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "website_url", + data_type: "TEXT", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "category", + data_type: "TEXT", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "created_at", + data_type: "INTEGER", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "sort_index", + data_type: "INTEGER", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "notes", + data_type: "TEXT", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "icon", + data_type: "TEXT", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "icon_color", + data_type: "TEXT", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "meta", + data_type: "TEXT", + not_null: true, + default: Some("'{}'"), + pk_position: 0, + }, + CanonicalColumnSpec { + name: "is_current", + data_type: "BOOLEAN", + not_null: true, + default: Some("0"), + pk_position: 0, + }, + CanonicalColumnSpec { + name: "in_failover_queue", + data_type: "BOOLEAN", + not_null: true, + default: Some("0"), + pk_position: 0, + }, +]; + +const PROVIDER_ENDPOINT_COLUMNS: &[CanonicalColumnSpec] = &[ + CanonicalColumnSpec { + name: "id", + data_type: "INTEGER", + not_null: false, + default: None, + pk_position: 1, + }, + CanonicalColumnSpec { + name: "provider_id", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "app_type", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "url", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "added_at", + data_type: "INTEGER", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "last_used", + data_type: "INTEGER", + not_null: false, + default: None, + pk_position: 0, + }, +]; + +const PI_PROJECTION_COLUMNS: &[CanonicalColumnSpec] = &[ + CanonicalColumnSpec { + name: "provider_id", + data_type: "TEXT", + not_null: false, + default: None, + pk_position: 1, + }, + CanonicalColumnSpec { + name: "provider_key", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "created_at", + data_type: "INTEGER", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "updated_at", + data_type: "INTEGER", + not_null: true, + default: None, + pk_position: 0, + }, +]; + +const SKILL_DEPLOYMENT_COLUMNS: &[CanonicalColumnSpec] = &[ + CanonicalColumnSpec { + name: "app_type", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 1, + }, + CanonicalColumnSpec { + name: "skill_id", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 2, + }, + CanonicalColumnSpec { + name: "destination", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "destination_key", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 3, + }, + CanonicalColumnSpec { + name: "method", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "source_identity", + data_type: "TEXT", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "deployed_digest", + data_type: "TEXT", + not_null: false, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "created_at", + data_type: "INTEGER", + not_null: true, + default: None, + pk_position: 0, + }, + CanonicalColumnSpec { + name: "updated_at", + data_type: "INTEGER", + not_null: true, + default: None, + pk_position: 0, + }, +]; + +const ENDPOINT_IDENTITY_COLUMNS: &[CanonicalIndexedColumnSpec] = &[ + CanonicalIndexedColumnSpec { + name: "provider_id", + collation: "BINARY", + }, + CanonicalIndexedColumnSpec { + name: "app_type", + collation: "BINARY", + }, + CanonicalIndexedColumnSpec { + name: "url", + collation: "BINARY", + }, +]; +const PROJECTION_KEY_COLUMNS: &[CanonicalIndexedColumnSpec] = &[CanonicalIndexedColumnSpec { + name: "provider_key", + collation: "BINARY", +}]; +const SKILL_DESTINATION_COLUMNS: &[CanonicalIndexedColumnSpec] = &[ + CanonicalIndexedColumnSpec { + name: "app_type", + collation: "BINARY", + }, + CanonicalIndexedColumnSpec { + name: "destination_key", + collation: "BINARY", + }, +]; + +const ENDPOINT_FOREIGN_KEYS: &[CanonicalForeignKeySpec] = &[CanonicalForeignKeySpec { + from: &["provider_id", "app_type"], + table: "providers", + to: &["id", "app_type"], + on_update: "NO ACTION", + on_delete: "CASCADE", + match_type: "NONE", +}]; + +const ENDPOINT_INVARIANTS: &[SchemaInvariantSpec] = &[ + SchemaInvariantSpec { + name: "endpoint_identity_unique", + kind: SchemaInvariantKind::EndpointIdentityUnique, + }, + SchemaInvariantSpec { + name: "endpoint_parent_fk", + kind: SchemaInvariantKind::EndpointParentForeignKey, + }, + SchemaInvariantSpec { + name: "endpoint_delete_cascade", + kind: SchemaInvariantKind::EndpointDeleteCascade, + }, +]; +const PROJECTION_INVARIANTS: &[SchemaInvariantSpec] = &[SchemaInvariantSpec { + name: "provider_key_unique", + kind: SchemaInvariantKind::ProjectionProviderKeyUnique, +}]; +const SKILL_INVARIANTS: &[SchemaInvariantSpec] = &[ + SchemaInvariantSpec { + name: "app_type_pi_only", + kind: SchemaInvariantKind::SkillAppTypePiOnly, + }, + SchemaInvariantSpec { + name: "method_symlink_or_copy", + kind: SchemaInvariantKind::SkillMethodAllowed, + }, + SchemaInvariantSpec { + name: "destination_owned_once", + kind: SchemaInvariantKind::SkillDestinationOwnedOnce, + }, +]; + +pub(crate) const CANONICAL_TABLE_SPECS: &[CanonicalTableSpec] = &[ + CanonicalTableSpec { + name: "providers", + definition: "providers ( + id TEXT NOT NULL, + app_type TEXT NOT NULL, + name TEXT NOT NULL, + settings_config TEXT NOT NULL, + website_url TEXT, + category TEXT, + created_at INTEGER, + sort_index INTEGER, + notes TEXT, + icon TEXT, + icon_color TEXT, + meta TEXT NOT NULL DEFAULT '{}', + is_current BOOLEAN NOT NULL DEFAULT 0, + in_failover_queue BOOLEAN NOT NULL DEFAULT 0, + PRIMARY KEY (id, app_type) + )", + restore_class: CanonicalRestoreClass::MigrateAndValidate, + columns: PROVIDERS_COLUMNS, + unique_tuples: &[], + foreign_keys: &[], + invariants: &[], + }, + CanonicalTableSpec { + name: "provider_endpoints", + definition: "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) + )", + restore_class: CanonicalRestoreClass::MigrateAndValidate, + columns: PROVIDER_ENDPOINT_COLUMNS, + unique_tuples: &[CanonicalUniqueSpec { + columns: ENDPOINT_IDENTITY_COLUMNS, + }], + foreign_keys: ENDPOINT_FOREIGN_KEYS, + invariants: ENDPOINT_INVARIANTS, + }, + CanonicalTableSpec { + name: "pi_provider_projections", + definition: "pi_provider_projections ( + provider_id TEXT PRIMARY KEY, + provider_key TEXT NOT NULL UNIQUE, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + )", + restore_class: CanonicalRestoreClass::RebuildAndPreserveLocal, + columns: PI_PROJECTION_COLUMNS, + unique_tuples: &[CanonicalUniqueSpec { + columns: PROJECTION_KEY_COLUMNS, + }], + foreign_keys: &[], + invariants: PROJECTION_INVARIANTS, + }, + CanonicalTableSpec { + name: "skill_deployments", + definition: "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) + )", + restore_class: CanonicalRestoreClass::RebuildAndPreserveLocal, + columns: SKILL_DEPLOYMENT_COLUMNS, + unique_tuples: &[CanonicalUniqueSpec { + columns: SKILL_DESTINATION_COLUMNS, + }], + foreign_keys: &[], + invariants: SKILL_INVARIANTS, + }, +]; #[derive(Serialize)] struct LegacySkillMigrationRow { @@ -14,6 +506,77 @@ struct LegacySkillMigrationRow { } impl Database { + /// Construct a publish-capable stage from an empty disk file using only + /// this binary's schema and migration code. + pub(super) fn current_canonical_stage() -> Result { + let file = NamedTempFile::new().map_err(|error| AppError::IoContext { + context: "create canonical restore stage".to_string(), + source: error, + })?; + let connection = + Connection::open(file.path()).map_err(|error| AppError::Database(error.to_string()))?; + connection + .execute_batch( + "PRAGMA foreign_keys = ON; + PRAGMA trusted_schema = OFF;", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + + // Starting from version zero exercises the normal migration chain and + // yields the exact same current objects as a fresh production database. + Self::create_tables_on_conn(&connection)?; + Self::set_user_version(&connection, 0)?; + Self::apply_schema_migrations_on_conn(&connection)?; + if Self::get_user_version(&connection)? != SCHEMA_VERSION { + return Err(AppError::Database( + "canonical stage factory did not reach the current schema version".to_string(), + )); + } + Ok(CanonicalStage { + connection, + _file: file, + }) + } + + pub(crate) fn canonical_table_spec( + name: &str, + ) -> Result<&'static CanonicalTableSpec, AppError> { + CANONICAL_TABLE_SPECS + .iter() + .find(|spec| spec.name == name) + .ok_or_else(|| AppError::Config(format!("unknown canonical table '{name}'"))) + } + + fn create_canonical_table_on_conn( + conn: &Connection, + name: &str, + if_not_exists: bool, + ) -> Result<(), AppError> { + let spec = Self::canonical_table_spec(name)?; + Self::create_canonical_table_as_on_conn(conn, spec, name, if_not_exists) + } + + fn create_canonical_table_as_on_conn( + conn: &Connection, + spec: &CanonicalTableSpec, + target_name: &str, + if_not_exists: bool, + ) -> Result<(), AppError> { + let qualifier = if if_not_exists { " IF NOT EXISTS" } else { "" }; + let definition = if target_name == spec.name { + spec.definition.to_string() + } else { + spec.definition.replacen(spec.name, target_name, 1) + }; + conn.execute(&format!("CREATE TABLE{qualifier} {definition}"), []) + .map_err(|error| { + AppError::Database(format!( + "failed to create canonical table '{target_name}': {error}" + )) + })?; + Ok(()) + } + /// 创建所有数据库表 pub(crate) fn create_tables(&self) -> Result<(), AppError> { let conn = lock_conn!(self.conn); @@ -23,41 +586,10 @@ impl Database { /// 在指定连接上创建表(供迁移和测试使用) pub(crate) fn create_tables_on_conn(conn: &Connection) -> Result<(), AppError> { // 1. Providers 表 - conn.execute( - "CREATE TABLE IF NOT EXISTS providers ( - id TEXT NOT NULL, - app_type TEXT NOT NULL, - name TEXT NOT NULL, - settings_config TEXT NOT NULL, - website_url TEXT, - category TEXT, - created_at INTEGER, - sort_index INTEGER, - notes TEXT, - icon TEXT, - icon_color TEXT, - meta TEXT NOT NULL DEFAULT '{}', - is_current BOOLEAN NOT NULL DEFAULT 0, - in_failover_queue BOOLEAN NOT NULL DEFAULT 0, - PRIMARY KEY (id, app_type) - )", - [], - ) - .map_err(|e| AppError::Database(e.to_string()))?; + Self::create_canonical_table_on_conn(conn, "providers", true)?; // 2. Provider Endpoints 表 - conn.execute( - "CREATE TABLE IF NOT EXISTS provider_endpoints ( - id INTEGER PRIMARY KEY AUTOINCREMENT, - provider_id TEXT NOT NULL, - 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 - )", - [], - ) - .map_err(|e| AppError::Database(e.to_string()))?; + Self::create_canonical_table_on_conn(conn, "provider_endpoints", true)?; // 3. MCP Servers 表 conn.execute( @@ -97,6 +629,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 +638,14 @@ impl Database { ) .map_err(|e| AppError::Database(e.to_string()))?; + // Exact models.json ownership is device-local and must never be + // inferred from provider names, content, or prefixes. + Self::create_canonical_table_on_conn(conn, "pi_provider_projections", true)?; + + // Pi Skill deployment ownership is independent from desired enablement + // in `skills.enabled_pi` and from live discovery. + Self::create_canonical_table_on_conn(conn, "skill_deployments", true)?; + // 6. Skill Repos 表 conn.execute( "CREATE TABLE IF NOT EXISTS skill_repos ( @@ -511,6 +1052,13 @@ impl Database { Self::migrate_v15_to_v16(conn)?; Self::set_user_version(conn, 16)?; } + 16 => { + log::info!( + "迁移数据库从 v16 到 v17(添加 Pi aggregate 与设备本地 ledger)" + ); + Self::migrate_v16_to_v17(conn)?; + Self::set_user_version(conn, 17)?; + } _ => { return Err(AppError::Database(format!( "未知的数据库版本 {version},无法迁移到 {SCHEMA_VERSION}" @@ -1523,6 +2071,88 @@ impl Database { crate::services::session_usage_codex::reset_codex_usage_on_conn(conn, &codex_dir) } + /// v16 -> v17: add the Pi desired bit, lossless endpoint metadata, and + /// device-local ownership ledgers. No ownership is inferred during + /// migration; both ledgers intentionally start empty. + 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")?; + // Older builds allowed duplicate rows for one logical endpoint. + // Merge their timestamps before rebuilding from the canonical + // definition. The fixed-column copy makes FK/UNIQUE/collation + // semantics part of migration rather than an optional index patch. + 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 + );", + ) + .map_err(|error| AppError::Database(error.to_string()))?; + const REBUILT_ENDPOINTS: &str = "provider_endpoints_v17_canonical"; + conn.execute(&format!("DROP TABLE IF EXISTS \"{REBUILT_ENDPOINTS}\""), []) + .map_err(|error| AppError::Database(error.to_string()))?; + let spec = Self::canonical_table_spec("provider_endpoints")?; + Self::create_canonical_table_as_on_conn(conn, spec, REBUILT_ENDPOINTS, false)?; + conn.execute( + &format!( + "INSERT INTO \"{REBUILT_ENDPOINTS}\" + (id, provider_id, app_type, url, added_at, last_used) + SELECT id, provider_id, app_type, url, added_at, last_used + FROM provider_endpoints" + ), + [], + ) + .map_err(|error| { + AppError::Database(format!( + "failed to copy provider_endpoints into canonical v17 table: {error}" + )) + })?; + conn.execute("DROP TABLE provider_endpoints", []) + .map_err(|error| AppError::Database(error.to_string()))?; + conn.execute( + &format!("ALTER TABLE \"{REBUILT_ENDPOINTS}\" RENAME TO provider_endpoints"), + [], + ) + .map_err(|error| AppError::Database(error.to_string()))?; + } else { + Self::create_canonical_table_on_conn(conn, "provider_endpoints", false)?; + } + if Self::table_exists(conn, "skills")? { + Self::add_column_if_missing( + conn, + "skills", + "enabled_pi", + "BOOLEAN NOT NULL DEFAULT 0", + )?; + } + Self::create_canonical_table_on_conn(conn, "pi_provider_projections", true)?; + Self::create_canonical_table_on_conn(conn, "skill_deployments", true) + } + /// 插入默认模型定价数据 /// 格式: (model_id, display_name, input, output, cache_read, cache_creation) /// 注意: model_id 使用短横线格式(如 claude-haiku-4-5),与 API 返回的模型名称标准化后一致 @@ -2780,7 +3410,7 @@ impl Database { Self::ensure_model_pricing_seeded_on_conn(&conn) } - fn ensure_model_pricing_seeded_on_conn(conn: &Connection) -> Result<(), AppError> { + pub(crate) fn ensure_model_pricing_seeded_on_conn(conn: &Connection) -> Result<(), AppError> { // 每次启动都执行 INSERT OR IGNORE,增量追加新模型;仅修复仍等于旧内置值的定价。 Self::seed_model_pricing(conn)?; Self::repair_current_model_pricing(conn) @@ -2938,6 +3568,73 @@ impl Database { #[cfg(test)] mod tests { use super::*; + use serde_json::json; + + fn canonical_manifest_from_specs() -> serde_json::Value { + let tables = CANONICAL_TABLE_SPECS + .iter() + .map(|spec| { + json!({ + "name": spec.name, + "restoreClass": spec.restore_class, + "columns": spec.columns.iter().map(|column| json!([ + column.name, + column.data_type, + column.not_null, + column.default, + column.pk_position + ])).collect::>(), + "uniqueTuples": spec.unique_tuples.iter().map(|tuple| { + tuple.columns.iter().map(|column| { + json!([column.name, column.collation]) + }).collect::>() + }).collect::>(), + "foreignKeys": spec.foreign_keys.iter().map(|foreign_key| json!({ + "from": foreign_key.from, + "table": foreign_key.table, + "to": foreign_key.to, + "onUpdate": foreign_key.on_update, + "onDelete": foreign_key.on_delete, + "match": foreign_key.match_type + })).collect::>(), + "checks": spec.invariants.iter().map(|invariant| invariant.name).collect::>() + }) + }) + .collect::>(); + json!({ + "manifestVersion": 1, + "schemaVersion": SCHEMA_VERSION, + "codeAuthority": "src-tauri/src/database/schema.rs", + "comparison": "semantic", + "tables": tables, + "futureUsageActivation": { + "commit": 13, + "proxy_request_logs.input_token_semantics": { + "type": "INTEGER", + "notNull": true, + "default": null, + "allowed": [1, 2, 3, 4] + }, + "usage_daily_rollups.input_token_semantics": { + "type": "INTEGER", + "notNull": true, + "default": null, + "allowed": [2] + } + } + }) + } + + #[test] + fn canonical_schema_specs_match_review_manifest() -> Result<(), AppError> { + let expected: serde_json::Value = serde_json::from_str(include_str!( + "../../../tests/fixtures/pi/canonical-schema-manifest-v1.json" + )) + .expect("parse canonical schema manifest"); + assert_eq!(canonical_manifest_from_specs(), expected); + + Ok(()) + } #[test] fn migrate_v12_to_v13_adds_input_token_semantics_columns() -> Result<(), AppError> { @@ -3080,7 +3777,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'), @@ -3093,4 +3790,80 @@ mod tests { assert_eq!(counts, (0, 1, 0, 1)); Ok(()) } + + #[test] + fn migrate_v16_to_v17_adds_pi_ledgers_without_inferred_ownership() -> 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 desired: i64 = conn.query_row( + "SELECT enabled_pi FROM skills WHERE id = 'existing'", + [], + |row| row.get(0), + )?; + 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!(desired, 0); + assert_eq!(ledgers, (0, 0)); + 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))); + assert!(conn + .execute( + "INSERT INTO provider_endpoints + (provider_id, app_type, url, added_at) + VALUES ('provider', 'pi', 'https://duplicate.test', 30)", + [], + ) + .is_err()); + Ok(()) + } } 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..010d5bcae 100644 --- a/src-tauri/src/deeplink/tests.rs +++ b/src-tauri/src/deeplink/tests.rs @@ -3,7 +3,7 @@ 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 +952,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..4cf39e8ff 100644 --- a/src-tauri/src/error.rs +++ b/src-tauri/src/error.rs @@ -9,6 +9,8 @@ pub enum AppError { Config(String), #[error("无效输入: {0}")] InvalidInput(String), + #[error("未找到: {0}")] + NotFound(String), #[error("IO 错误: {path}: {source}")] Io { path: String, diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index a493e3338..def70df1f 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -44,7 +44,10 @@ pub use codex_config::{get_codex_auth_path, get_codex_config_path, write_codex_l 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; @@ -56,7 +59,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, diff --git a/src-tauri/src/provider.rs b/src-tauri/src/provider.rs index 6569304b8..32852f5ad 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( 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/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..381d271d1 100644 --- a/src-tauri/src/services/provider/live.rs +++ b/src-tauri/src/services/provider/live.rs @@ -19,7 +19,9 @@ 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_to_mutation_input, reconcile_provider_record, +}; /// 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. @@ -1564,7 +1566,11 @@ 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 mut provider = existing; provider.name = display_name; provider.settings_config = settings_config; - if let Err(e) = state.db.save_provider("opencode", &provider) { + if let Err(e) = reconcile_provider_record( + &state.db, + "opencode", + provider_to_mutation_input(provider), + ) { log::warn!( "Failed to update OpenCode provider '{id}' from live config: {e}" ); @@ -1767,7 +1778,9 @@ pub fn import_opencode_providers_from_live(state: &AppState) -> Result Result { + let existing = existing.provider; if existing.settings_config != settings_config { let mut provider = existing; provider.settings_config = settings_config; - if let Err(e) = state.db.save_provider("openclaw", &provider) { + if let Err(e) = reconcile_provider_record( + &state.db, + "openclaw", + provider_to_mutation_input(provider), + ) { log::warn!( "Failed to update OpenClaw provider '{id}' from live config: {e}" ); @@ -1855,7 +1873,9 @@ pub fn import_openclaw_providers_from_live(state: &AppState) -> Result Result { + let existing = existing.provider; if existing.settings_config != config { let mut provider = existing; provider.settings_config = config; - if let Err(e) = state.db.save_provider("hermes", &provider) { + if let Err(e) = reconcile_provider_record( + &state.db, + "hermes", + provider_to_mutation_input(provider), + ) { log::warn!( "Failed to update Hermes provider '{name}' from live config: {e}" ); @@ -1923,7 +1948,9 @@ 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) +} + +/// Reconciliation paths must state their intent explicitly: inspect first, +/// then perform either strict create or strict one-row update. +pub(crate) fn reconcile_provider_record( + db: &crate::database::Database, + app_type: &str, + input: ProviderMutationInput, +) -> Result<(), AppError> { + let key = ProviderKey::new(app_type, input.id.clone())?; + if db.get_provider_aggregate(app_type, key.id())?.is_some() { + let row = ProviderRowUpdate::from_input(&input)?; + db.update_provider(&key, &row) + } else { + db.create_provider(NewProviderAggregate::from_input(app_type, input)?) + } +} + /// Result of a provider switch operation, including any non-fatal warnings #[derive(Debug, serde::Serialize, Default)] #[serde(rename_all = "camelCase")] @@ -437,6 +496,318 @@ 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_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 add_clears_usage_credentials_that_match_provider_config() { @@ -450,7 +821,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 +859,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 +909,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 +933,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 +975,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 +989,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 +1026,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 +1066,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 +1219,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 +1232,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 +1880,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 +1938,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 +2020,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 +2084,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 +2105,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 +2161,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 +2248,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 +2263,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 +2287,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 +2318,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 +2340,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 +2365,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 +2394,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 +2427,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 +2455,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 +2505,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 +2579,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 +2616,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 +2656,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 +2667,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 +2689,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 +2697,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 +2729,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 +2743,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 +2784,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 +2801,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 +2831,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 +2865,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 @@ -2551,10 +3028,10 @@ impl ProviderService { pub fn add( state: &AppState, app_type: AppType, - provider: Provider, + input: ProviderMutationInput, add_to_live: bool, ) -> Result { - let mut provider = provider; + 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 +3041,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 +3083,12 @@ 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 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 @@ -2677,8 +3161,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 +3197,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 +3230,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 +3244,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 = @@ -2943,9 +3444,13 @@ impl ProviderService { } } - if let Some(mut provider) = state.db.get_provider_by_id(id, app_type.as_str())? { + if let Some(mut provider) = state + .db + .get_provider_aggregate(app_type.as_str(), id)? + .map(|aggregate| aggregate.provider) + { Self::set_provider_live_config_managed(&mut provider, false); - state.db.save_provider(app_type.as_str(), &provider)?; + update_provider_record(state, &app_type, &provider_to_mutation_input(provider))?; } Ok(()) @@ -3097,7 +3602,11 @@ 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) + { // 切走前先把 live 里的可共享改动(含用户直接在应用内 // 装插件/加 hook/改偏好)同步进通用配置片段,再做剥离回填。 // 详见 sync_common_config_snippet_from_live 的文档。 @@ -3116,9 +3625,11 @@ impl ProviderService { ¤t_provider, live_config, ); - if let Err(e) = - state.db.save_provider(app_type.as_str(), ¤t_provider) - { + if let Err(e) = update_provider_record( + state, + &app_type, + &provider_to_mutation_input(current_provider), + ) { log::warn!("Backfill failed: {e}"); result .warnings @@ -3169,9 +3680,16 @@ 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(); + let mut updated = state + .db + .get_provider_aggregate(app_type.as_str(), &provider.id)? + .ok_or_else(|| { + AppError::NotFound(format!("provider '{}/{}'", app_type.as_str(), provider.id)) + })? + .provider; Self::set_provider_live_config_managed(&mut updated, true); - if let Err(e) = state.db.save_provider(app_type.as_str(), &updated) { + let update = provider_to_mutation_input(updated); + if let Err(e) = update_provider_record(state, &app_type, &update) { let rollback_result = match app_type { AppType::OpenCode => remove_opencode_provider_from_live(&provider.id), AppType::OpenClaw => remove_openclaw_provider_from_live(&provider.id), @@ -3273,9 +3791,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() @@ -3310,9 +3829,11 @@ impl ProviderService { } } - state - .db - .save_provider(app_type.as_str(), &updated_provider)?; + update_provider_record( + state, + &app_type, + &provider_to_mutation_input(updated_provider), + )?; } Ok(()) @@ -3828,9 +4349,10 @@ impl ProviderService { .map_err(|e| AppError::Message(format!("Serialization failed: {e}")))?; // 1) 先算出各供应商清理后的配置,但**先不落库** - let providers = state.db.get_all_providers(app.as_str())?; + let providers = state.db.get_all_provider_aggregates(app.as_str())?; let mut pending: Vec<(String, Provider, Value)> = Vec::new(); - for (id, provider) in providers { + for (id, aggregate) in providers { + let provider = aggregate.provider; let cleaned = match live::remove_common_config_from_settings( &app, &provider.settings_config, @@ -3880,7 +4402,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 会拿这份残缺 // 记录盖掉第一轮那份完整的。 @@ -3892,7 +4414,8 @@ impl ProviderService { for (id, provider, cleaned) in pending { let mut updated = provider; updated.settings_config = cleaned; - state.db.save_provider(app.as_str(), &updated)?; + let update = provider_to_mutation_input(updated); + update_provider_record(state, &app, &update)?; log::info!("已从 Gemini 供应商 '{id}' 中清除泄漏的共享凭据"); } @@ -4074,13 +4597,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) @@ -4633,7 +5154,11 @@ impl ProviderService { Self::merge_json(&mut merged, &claude_provider.settings_config); claude_provider.settings_config = merged; } - state.db.save_provider("claude", &claude_provider)?; + reconcile_provider_record( + &state.db, + "claude", + provider_to_mutation_input(claude_provider), + )?; } else { // 如果禁用了 Claude,删除对应的子供应商 let claude_id = format!("universal-claude-{id}"); @@ -4648,7 +5173,11 @@ impl ProviderService { Self::merge_json(&mut merged, &codex_provider.settings_config); codex_provider.settings_config = merged; } - state.db.save_provider("codex", &codex_provider)?; + reconcile_provider_record( + &state.db, + "codex", + provider_to_mutation_input(codex_provider), + )?; } else { let codex_id = format!("universal-codex-{id}"); let _ = state.db.delete_provider("codex", &codex_id); @@ -4662,7 +5191,11 @@ impl ProviderService { Self::merge_json(&mut merged, &gemini_provider.settings_config); gemini_provider.settings_config = merged; } - state.db.save_provider("gemini", &gemini_provider)?; + reconcile_provider_record( + &state.db, + "gemini", + provider_to_mutation_input(gemini_provider), + )?; } 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 4efee1537..104bbdc59 100644 --- a/src-tauri/src/services/proxy.rs +++ b/src-tauri/src/services/proxy.rs @@ -3759,7 +3759,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"); @@ -3945,7 +3945,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"); @@ -4031,7 +4031,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"); @@ -4092,7 +4092,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( @@ -4111,7 +4111,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"); @@ -4260,7 +4260,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( @@ -4338,7 +4339,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"); @@ -4418,7 +4419,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"); @@ -4530,7 +4531,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"); @@ -4648,7 +4649,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"); @@ -4784,7 +4785,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"); @@ -5261,7 +5262,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"); @@ -5317,7 +5318,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"); @@ -5382,9 +5383,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"); @@ -5457,9 +5458,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"); @@ -5608,11 +5609,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"); @@ -5695,9 +5696,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"); @@ -5995,9 +5996,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"); @@ -6170,9 +6171,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"); @@ -6414,9 +6415,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"); @@ -6550,9 +6551,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"); @@ -6632,9 +6633,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"); @@ -6916,7 +6917,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"); @@ -7170,9 +7171,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"); @@ -7237,9 +7238,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 bbb11cf02..7e84c6fdc 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 { @@ -2922,10 +2922,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/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/canonical-schema-manifest-v1.json b/tests/fixtures/pi/canonical-schema-manifest-v1.json new file mode 100644 index 000000000..be6e16b1e --- /dev/null +++ b/tests/fixtures/pi/canonical-schema-manifest-v1.json @@ -0,0 +1,107 @@ +{ + "manifestVersion": 1, + "schemaVersion": 17, + "codeAuthority": "src-tauri/src/database/schema.rs", + "comparison": "semantic", + "tables": [ + { + "name": "providers", + "restoreClass": "migrate_and_validate", + "columns": [ + ["id", "TEXT", true, null, 1], + ["app_type", "TEXT", true, null, 2], + ["name", "TEXT", true, null, 0], + ["settings_config", "TEXT", true, null, 0], + ["website_url", "TEXT", false, null, 0], + ["category", "TEXT", false, null, 0], + ["created_at", "INTEGER", false, null, 0], + ["sort_index", "INTEGER", false, null, 0], + ["notes", "TEXT", false, null, 0], + ["icon", "TEXT", false, null, 0], + ["icon_color", "TEXT", false, null, 0], + ["meta", "TEXT", true, "'{}'", 0], + ["is_current", "BOOLEAN", true, "0", 0], + ["in_failover_queue", "BOOLEAN", true, "0", 0] + ], + "uniqueTuples": [], + "foreignKeys": [], + "checks": [] + }, + { + "name": "provider_endpoints", + "restoreClass": "migrate_and_validate", + "columns": [ + ["id", "INTEGER", false, null, 1], + ["provider_id", "TEXT", true, null, 0], + ["app_type", "TEXT", true, null, 0], + ["url", "TEXT", true, null, 0], + ["added_at", "INTEGER", false, null, 0], + ["last_used", "INTEGER", false, null, 0] + ], + "uniqueTuples": [ + [["provider_id", "BINARY"], ["app_type", "BINARY"], ["url", "BINARY"]] + ], + "foreignKeys": [ + { + "from": ["provider_id", "app_type"], + "table": "providers", + "to": ["id", "app_type"], + "onUpdate": "NO ACTION", + "onDelete": "CASCADE", + "match": "NONE" + } + ], + "checks": ["endpoint_identity_unique", "endpoint_parent_fk", "endpoint_delete_cascade"] + }, + { + "name": "pi_provider_projections", + "restoreClass": "rebuild_and_preserve_local", + "columns": [ + ["provider_id", "TEXT", false, null, 1], + ["provider_key", "TEXT", true, null, 0], + ["created_at", "INTEGER", true, null, 0], + ["updated_at", "INTEGER", true, null, 0] + ], + "uniqueTuples": [ + [["provider_key", "BINARY"]] + ], + "foreignKeys": [], + "checks": ["provider_key_unique"] + }, + { + "name": "skill_deployments", + "restoreClass": "rebuild_and_preserve_local", + "columns": [ + ["app_type", "TEXT", true, null, 1], + ["skill_id", "TEXT", true, null, 2], + ["destination", "TEXT", true, null, 0], + ["destination_key", "TEXT", true, null, 3], + ["method", "TEXT", true, null, 0], + ["source_identity", "TEXT", true, null, 0], + ["deployed_digest", "TEXT", false, null, 0], + ["created_at", "INTEGER", true, null, 0], + ["updated_at", "INTEGER", true, null, 0] + ], + "uniqueTuples": [ + [["app_type", "BINARY"], ["destination_key", "BINARY"]] + ], + "foreignKeys": [], + "checks": ["app_type_pi_only", "method_symlink_or_copy", "destination_owned_once"] + } + ], + "futureUsageActivation": { + "commit": 13, + "proxy_request_logs.input_token_semantics": { + "type": "INTEGER", + "notNull": true, + "default": null, + "allowed": [1, 2, 3, 4] + }, + "usage_daily_rollups.input_token_semantics": { + "type": "INTEGER", + "notNull": true, + "default": null, + "allowed": [2] + } + } +} 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..787b22e71 --- /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"], + "ProviderRowUpdate": [ + "name", + "settings_config", + "website_url", + "category", + "created_at", + "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"] + } +} diff --git a/tests/fixtures/pi/restore-policy-v1.json b/tests/fixtures/pi/restore-policy-v1.json new file mode 100644 index 000000000..7e8ff29ed --- /dev/null +++ b/tests/fixtures/pi/restore-policy-v1.json @@ -0,0 +1,31 @@ +{ + "manifestVersion": 1, + "codeAuthority": "src-tauri/src/database/backup.rs", + "schemaVersion": 17, + "limits": { + "sqlImportBytes": 268435456, + "binaryRestoreBytes": 2147483648, + "scratchBytes": 2147483648 + }, + "specSha256": "c0a680de97b4d3eb291114ebfdf50fdfd65ce7f63819a0bc84f97cec6466bf07", + "tables": [ + {"name": "providers", "policy": "portable_incoming"}, + {"name": "provider_endpoints", "policy": "portable_incoming"}, + {"name": "mcp_servers", "policy": "portable_incoming"}, + {"name": "prompts", "policy": "portable_incoming"}, + {"name": "skills", "policy": "portable_incoming"}, + {"name": "skill_repos", "policy": "portable_incoming"}, + {"name": "settings", "policy": "portable_incoming"}, + {"name": "proxy_config", "policy": "portable_incoming"}, + {"name": "provider_health", "policy": "rebuild_runtime"}, + {"name": "proxy_request_logs", "policy": "portable_incoming"}, + {"name": "model_pricing", "policy": "portable_incoming"}, + {"name": "stream_check_logs", "policy": "portable_incoming"}, + {"name": "proxy_live_backup", "policy": "portable_incoming"}, + {"name": "usage_daily_rollups", "policy": "portable_incoming"}, + {"name": "session_log_sync", "policy": "portable_incoming"}, + {"name": "profiles", "policy": "portable_incoming"}, + {"name": "pi_provider_projections", "policy": "preserve_live"}, + {"name": "skill_deployments", "policy": "preserve_live"} + ] +}