diff --git a/src-tauri/src/database.rs b/src-tauri/src/database.rs index 8db7d2731..1f8b6b23c 100644 --- a/src-tauri/src/database.rs +++ b/src-tauri/src/database.rs @@ -85,7 +85,7 @@ impl Database { notes TEXT, icon TEXT, icon_color TEXT, - meta TEXT, + meta TEXT NOT NULL DEFAULT '{}', is_current BOOLEAN NOT NULL DEFAULT 0, PRIMARY KEY (id, app_type) )", @@ -115,7 +115,7 @@ impl Database { description TEXT, homepage TEXT, docs TEXT, - tags TEXT, + tags TEXT NOT NULL DEFAULT '[]', enabled_claude BOOLEAN NOT NULL DEFAULT 0, enabled_codex BOOLEAN NOT NULL DEFAULT 0, enabled_gemini BOOLEAN NOT NULL DEFAULT 0 @@ -146,7 +146,7 @@ impl Database { "CREATE TABLE IF NOT EXISTS skills ( key TEXT PRIMARY KEY, installed BOOLEAN NOT NULL DEFAULT 0, - installed_at INTEGER + installed_at INTEGER NOT NULL DEFAULT 0 )", [], ) @@ -157,7 +157,7 @@ impl Database { "CREATE TABLE IF NOT EXISTS skill_repos ( owner TEXT NOT NULL, name TEXT NOT NULL, - branch TEXT NOT NULL, + branch TEXT NOT NULL DEFAULT 'main', enabled BOOLEAN NOT NULL DEFAULT 1, skills_path TEXT, PRIMARY KEY (owner, name) @@ -201,32 +201,207 @@ impl Database { Self::apply_schema_migrations_on_conn(&conn) } + fn validate_identifier(s: &str, kind: &str) -> Result<(), AppError> { + if s.is_empty() { + return Err(AppError::Database(format!("{kind} 不能为空"))); + } + if !s + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '_') + { + return Err(AppError::Database(format!( + "非法{kind}: {s},仅允许字母、数字和下划线" + ))); + } + Ok(()) + } + + fn table_exists(conn: &Connection, table: &str) -> Result { + Self::validate_identifier(table, "表名")?; + + let sql = "SELECT 1 FROM sqlite_master WHERE type='table' AND name = ?1 LIMIT 1;"; + match conn.query_row(sql, params![table], |row| row.get::<_, i64>(0)) { + Ok(_) => Ok(true), + Err(rusqlite::Error::QueryReturnedNoRows) => Ok(false), + Err(e) => Err(AppError::Database(format!( + "检查表是否存在失败: {e}" + ))), + } + } + + fn has_column(conn: &Connection, table: &str, column: &str) -> Result { + Self::validate_identifier(table, "表名")?; + Self::validate_identifier(column, "列名")?; + + let sql = format!("PRAGMA table_info(\"{table}\");"); + let mut stmt = conn + .prepare(&sql) + .map_err(|e| AppError::Database(format!("读取表结构失败: {e}")))?; + let mut rows = stmt + .query([]) + .map_err(|e| AppError::Database(format!("查询表结构失败: {e}")))?; + while let Some(row) = rows.next().map_err(|e| AppError::Database(e.to_string()))? { + let name: String = row + .get(1) + .map_err(|e| AppError::Database(format!("读取列名失败: {e}")))?; + if name.eq_ignore_ascii_case(column) { + return Ok(true); + } + } + Ok(false) + } + + fn add_column_if_missing( + conn: &Connection, + table: &str, + column: &str, + definition: &str, + ) -> Result { + Self::validate_identifier(table, "表名")?; + Self::validate_identifier(column, "列名")?; + + if !Self::table_exists(conn, table)? { + return Err(AppError::Database(format!( + "表 {table} 不存在,无法添加列 {column}" + ))); + } + if Self::has_column(conn, table, column)? { + return Ok(false); + } + + let sql = format!("ALTER TABLE \"{table}\" ADD COLUMN {definition};"); + conn.execute(&sql, []) + .map_err(|e| AppError::Database(format!("为表 {table} 添加列 {column} 失败: {e}")))?; + log::info!("已为表 {table} 添加缺失列 {column}"); + Ok(true) + } + fn apply_schema_migrations_on_conn(conn: &Connection) -> Result<(), AppError> { + conn.execute("SAVEPOINT schema_migration;", []) + .map_err(|e| AppError::Database(format!("开启迁移 savepoint 失败: {e}")))?; + let mut version = Self::get_user_version(conn)?; if version > SCHEMA_VERSION { + conn.execute("ROLLBACK TO schema_migration;", []).ok(); + conn.execute("RELEASE schema_migration;", []).ok(); return Err(AppError::Database(format!( "数据库版本过新({version}),当前应用仅支持 {SCHEMA_VERSION},请升级应用后再尝试。" ))); } - while version < SCHEMA_VERSION { - match version { - 0 => { - log::info!("检测到 user_version=0,设置为初始版本 {SCHEMA_VERSION}"); - Self::set_user_version(conn, SCHEMA_VERSION)?; - } - _ => { - return Err(AppError::Database(format!( - "未知的数据库版本 {version},无法迁移到 {SCHEMA_VERSION}" - ))); + let result = (|| { + while version < SCHEMA_VERSION { + match version { + 0 => { + log::info!("检测到 user_version=0,迁移到 1(补齐缺失列并设置版本)"); + Self::add_column_if_missing(conn, "providers", "category", "TEXT")?; + Self::add_column_if_missing(conn, "providers", "created_at", "INTEGER")?; + Self::add_column_if_missing(conn, "providers", "sort_index", "INTEGER")?; + Self::add_column_if_missing(conn, "providers", "notes", "TEXT")?; + Self::add_column_if_missing(conn, "providers", "icon", "TEXT")?; + Self::add_column_if_missing(conn, "providers", "icon_color", "TEXT")?; + Self::add_column_if_missing( + conn, + "providers", + "meta", + "TEXT NOT NULL DEFAULT '{}'", + )?; + Self::add_column_if_missing( + conn, + "providers", + "is_current", + "INTEGER NOT NULL DEFAULT 0", + )?; + + Self::add_column_if_missing( + conn, + "provider_endpoints", + "added_at", + "INTEGER", + )?; + + Self::add_column_if_missing(conn, "mcp_servers", "description", "TEXT")?; + Self::add_column_if_missing(conn, "mcp_servers", "homepage", "TEXT")?; + Self::add_column_if_missing(conn, "mcp_servers", "docs", "TEXT")?; + Self::add_column_if_missing( + conn, + "mcp_servers", + "tags", + "TEXT NOT NULL DEFAULT '[]'", + )?; + Self::add_column_if_missing( + conn, + "mcp_servers", + "enabled_codex", + "INTEGER NOT NULL DEFAULT 0", + )?; + Self::add_column_if_missing( + conn, + "mcp_servers", + "enabled_gemini", + "INTEGER NOT NULL DEFAULT 0", + )?; + + Self::add_column_if_missing(conn, "prompts", "description", "TEXT")?; + Self::add_column_if_missing( + conn, + "prompts", + "enabled", + "INTEGER NOT NULL DEFAULT 1", + )?; + Self::add_column_if_missing(conn, "prompts", "created_at", "INTEGER")?; + Self::add_column_if_missing(conn, "prompts", "updated_at", "INTEGER")?; + + Self::add_column_if_missing( + conn, + "skills", + "installed_at", + "INTEGER NOT NULL DEFAULT 0", + )?; + + Self::add_column_if_missing( + conn, + "skill_repos", + "branch", + "TEXT NOT NULL DEFAULT 'main'", + )?; + Self::add_column_if_missing( + conn, + "skill_repos", + "enabled", + "INTEGER NOT NULL DEFAULT 1", + )?; + Self::add_column_if_missing(conn, "skill_repos", "skills_path", "TEXT")?; + + Self::set_user_version(conn, SCHEMA_VERSION)?; + } + _ => { + return Err(AppError::Database(format!( + "未知的数据库版本 {version},无法迁移到 {SCHEMA_VERSION}" + ))); + } } + + version = Self::get_user_version(conn)?; } - version = Self::get_user_version(conn)?; - } + Ok(()) + })(); - Ok(()) + match result { + Ok(_) => { + conn.execute("RELEASE schema_migration;", []).map_err(|e| { + AppError::Database(format!("提交迁移 savepoint 失败: {e}")) + })?; + Ok(()) + } + Err(e) => { + conn.execute("ROLLBACK TO schema_migration;", []).ok(); + conn.execute("RELEASE schema_migration;", []).ok(); + Err(e) + } + } } /// 创建内存快照以避免长时间持有数据库锁 @@ -1312,4 +1487,96 @@ mod tests { "unexpected error: {err}" ); } + + #[test] + fn migration_adds_missing_columns_for_providers() { + let conn = Connection::open_in_memory().expect("open memory db"); + + // 创建旧版 providers 表,缺少新增列 + conn.execute_batch( + r#" + CREATE TABLE providers ( + id TEXT NOT NULL, + app_type TEXT NOT NULL, + name TEXT NOT NULL, + settings_config TEXT NOT NULL, + PRIMARY KEY (id, app_type) + ); + CREATE TABLE provider_endpoints ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + provider_id TEXT NOT NULL, + app_type TEXT NOT NULL, + url TEXT NOT NULL + ); + CREATE TABLE mcp_servers ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + server_config TEXT NOT NULL + ); + CREATE TABLE prompts ( + id TEXT NOT NULL, + app_type TEXT NOT NULL, + name TEXT NOT NULL, + content TEXT NOT NULL, + PRIMARY KEY (id, app_type) + ); + CREATE TABLE skills ( + key TEXT PRIMARY KEY, + installed BOOLEAN NOT NULL DEFAULT 0 + ); + CREATE TABLE skill_repos ( + owner TEXT NOT NULL, + name TEXT NOT NULL, + PRIMARY KEY (owner, name) + ); + CREATE TABLE settings ( + key TEXT PRIMARY KEY, + value TEXT + ); + "#, + ) + .expect("seed old schema"); + + Database::apply_schema_migrations_on_conn(&conn).expect("apply migrations"); + + // 验证关键新增列已补齐 + for (table, column) in [ + ("providers", "meta"), + ("providers", "is_current"), + ("provider_endpoints", "added_at"), + ("mcp_servers", "enabled_gemini"), + ("prompts", "updated_at"), + ("skills", "installed_at"), + ("skill_repos", "enabled"), + ] { + assert!( + Database::has_column(&conn, table, column).expect("check column"), + "{table}.{column} should exist after migration" + ); + } + + // 验证 meta 列约束保持一致 + let mut stmt = conn + .prepare("PRAGMA table_info(\"providers\");") + .expect("prepare pragma"); + let mut rows = stmt.query([]).expect("query pragma"); + let mut meta_not_null = None; + let mut meta_default = None; + while let Some(row) = rows.next().expect("read row") { + let name: String = row.get(1).expect("name"); + if name == "meta" { + meta_not_null = Some(row.get::<_, i64>(3).expect("notnull")); + meta_default = row.get::<_, Option>(4).ok().flatten(); + break; + } + } + assert_eq!(meta_not_null, Some(1), "meta should be NOT NULL"); + let normalized_default = meta_default.map(|s| s.trim_matches('\'').to_string()); + assert_eq!(normalized_default.as_deref(), Some("{}"), "meta default should be '{}'"); + + assert_eq!( + Database::get_user_version(&conn).expect("version after migration"), + SCHEMA_VERSION + ); + } }