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"} + ] +}