// Package store 负责 SQLite 持久化:站点设置、用户、会话、仓库与推送记录。 package store import ( "crypto/rand" "database/sql" "encoding/hex" "errors" "fmt" "os" "path/filepath" "strconv" "strings" "time" _ "modernc.org/sqlite" ) // ErrNotFound 表示记录不存在。 var ErrNotFound = errors.New("not found") const timeFormat = time.RFC3339Nano // Store 是数据库访问入口。 type Store struct { db *sql.DB dbPath string // 用于迁移前备份;为空表示不备份 } // Open 打开(必要时创建)SQLite 数据库并完成建表。 func Open(path string) (*Store, error) { dsn := "file:" + path + "?_pragma=busy_timeout(15000)&_pragma=journal_mode(WAL)&_pragma=synchronous(NORMAL)&_pragma=foreign_keys(1)" db, err := sql.Open("sqlite", dsn) if err != nil { return nil, err } // 小型工作室场景:单连接即可,彻底避免 SQLite 写锁竞争。 db.SetMaxOpenConns(1) if err := db.Ping(); err != nil { _ = db.Close() return nil, err } s := &Store{db: db, dbPath: path} if err := s.migrate(); err != nil { _ = db.Close() return nil, err } return s, nil } // Close 关闭数据库。 func (s *Store) Close() error { return s.db.Close() } // schema 是基线结构,按顺序执行。 // // 注意语句顺序:SQLite 不允许为尚不存在的表建索引,所以索引必须排在其 // 对应的 CREATE TABLE 之后。 var schema = []string{ `CREATE TABLE IF NOT EXISTS settings ( key TEXT PRIMARY KEY, value TEXT NOT NULL DEFAULT '' )`, `CREATE TABLE IF NOT EXISTS users ( id INTEGER PRIMARY KEY AUTOINCREMENT, username TEXT NOT NULL UNIQUE COLLATE NOCASE, display_name TEXT NOT NULL DEFAULT '', email TEXT NOT NULL DEFAULT '', bio TEXT NOT NULL DEFAULT '', password_hash TEXT NOT NULL, avatar TEXT NOT NULL DEFAULT '', is_admin INTEGER NOT NULL DEFAULT 0, is_disabled INTEGER NOT NULL DEFAULT 0, created_at TEXT NOT NULL, last_login_at TEXT NOT NULL DEFAULT '' )`, `CREATE TABLE IF NOT EXISTS sessions ( id TEXT PRIMARY KEY, user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, csrf TEXT NOT NULL DEFAULT '', ip TEXT NOT NULL DEFAULT '', user_agent TEXT NOT NULL DEFAULT '', created_at TEXT NOT NULL, expires_at TEXT NOT NULL )`, `CREATE TABLE IF NOT EXISTS repos ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL UNIQUE COLLATE NOCASE, description TEXT NOT NULL DEFAULT '', owner_id INTEGER REFERENCES users(id) ON DELETE SET NULL, default_branch TEXT NOT NULL DEFAULT 'main', is_archived INTEGER NOT NULL DEFAULT 0, ai_enabled INTEGER NOT NULL DEFAULT 1, created_at TEXT NOT NULL, updated_at TEXT NOT NULL )`, `CREATE TABLE IF NOT EXISTS pushes ( id INTEGER PRIMARY KEY AUTOINCREMENT, repo_id INTEGER NOT NULL REFERENCES repos(id) ON DELETE CASCADE, user_id INTEGER REFERENCES users(id) ON DELETE SET NULL, actor_username TEXT NOT NULL DEFAULT '', actor_display TEXT NOT NULL DEFAULT '', ref TEXT NOT NULL DEFAULT '', old_sha TEXT NOT NULL DEFAULT '', new_sha TEXT NOT NULL DEFAULT '', commit_count INTEGER NOT NULL DEFAULT 0, details TEXT NOT NULL DEFAULT '[]', whats_new TEXT NOT NULL DEFAULT '', ai_status TEXT NOT NULL DEFAULT 'pending', created_at TEXT NOT NULL )`, // 索引。原先只有三处,而 ListUsers / ListRepos 恰恰在 owner_id、 // user_id 上做相关子查询,列表页每次都要全表扫。 `CREATE INDEX IF NOT EXISTS idx_repos_owner ON repos(owner_id)`, `CREATE INDEX IF NOT EXISTS idx_repos_updated ON repos(updated_at DESC)`, `CREATE INDEX IF NOT EXISTS idx_pushes_repo ON pushes(repo_id, id DESC)`, `CREATE INDEX IF NOT EXISTS idx_pushes_created ON pushes(created_at DESC)`, `CREATE INDEX IF NOT EXISTS idx_pushes_user ON pushes(user_id, id DESC)`, `CREATE INDEX IF NOT EXISTS idx_sessions_expires ON sessions(expires_at)`, `CREATE INDEX IF NOT EXISTS idx_sessions_user ON sessions(user_id)`, } // ---------------------------------------------------------------- 迁移 // // schema 里每条语句都是 IF NOT EXISTS,本身幂等;schema_meta 记录的版本号 // 用来做两件事: // 1. 降级保护:数据库版本高于二进制版本时拒绝启动,避免旧代码把新结构写坏。 // 2. 增量升级:新增的、无法用 IF NOT EXISTS 表达的变更(加列、改类型、 // 填数据)放进 migrations 切片,按版本号依次执行。 // // 版本 1 = 最原始的 5 张表(无 schema_meta),用于识别旧版数据库。 // SchemaVersion 是当前二进制支持的最新数据库版本。 const SchemaVersion = 1 // execer 让迁移既能跑在事务里(*sql.Tx)也能直接跑在连接上(*sql.DB)。 type execer interface { Exec(query string, args ...any) (sql.Result, error) Query(query string, args ...any) (*sql.Rows, error) QueryRow(query string, args ...any) *sql.Row } // migration 是一条按序执行的结构变更。 type migration struct { version int name string run func(execer) error } // migrations 追加式列表:新增变更时在末尾追加,version 必须递增。 var migrations = []migration{ // 示例: // {version: 2, name: "repos 加 visibility", run: func(db *sql.DB) error { // _, err := db.Exec(`ALTER TABLE repos ADD COLUMN visibility TEXT NOT NULL DEFAULT 'public'`) // return err // }}, } const schemaMetaTable = `CREATE TABLE IF NOT EXISTS schema_meta ( key TEXT PRIMARY KEY, value TEXT NOT NULL DEFAULT '' )` // dbVersion 读取数据库当前版本;没有 schema_meta 的旧库返回 0。 func dbVersion(db *sql.DB) (int, error) { var name string err := db.QueryRow(`SELECT name FROM sqlite_master WHERE type='table' AND name='schema_meta'`).Scan(&name) if errors.Is(err, sql.ErrNoRows) { return 0, nil } if err != nil { return 0, err } var v string if err := db.QueryRow(`SELECT value FROM schema_meta WHERE key='schema_version'`).Scan(&v); err != nil { if errors.Is(err, sql.ErrNoRows) { return 0, nil } return 0, err } n, err := strconv.Atoi(strings.TrimSpace(v)) if err != nil { return 0, fmt.Errorf("无法解析 schema 版本号 %q", v) } return n, nil } func setDBVersion(db *sql.DB, v int) error { _, err := db.Exec(`INSERT INTO schema_meta(key, value) VALUES('schema_version', ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value`, strconv.Itoa(v)) return err } func (s *Store) migrate() error { // 是否需要在迁移前备份,必须在建表之前判断,而且不能只看文件大小: // 驱动一连上就会创建出带页头的文件,"全新安装"看起来也像是有内容的库。 // 精确判据是"users 表是否已存在"。 var existing int err := s.db.QueryRow(`SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name IN ('users','repos','pushes')`).Scan(&existing) if err != nil { return fmt.Errorf("migrate: %w", err) } hasLegacyData := existing > 0 for _, stmt := range schema { if _, err := s.db.Exec(stmt); err != nil { return fmt.Errorf("migrate: %w", err) } } if _, err := s.db.Exec(schemaMetaTable); err != nil { return fmt.Errorf("migrate: %w", err) } current, err := dbVersion(s.db) if err != nil { return fmt.Errorf("migrate: %w", err) } // 降级保护:用旧二进制打开被新版本迁移过的库,字段可能对不上, // 与其运行到一半报诡异错误,不如明确拒绝启动。 if current > SchemaVersion { return fmt.Errorf("数据库结构版本为 v%d,高于当前程序支持的 v%d;请先升级 gitcat 再启动", current, SchemaVersion) } if current == SchemaVersion { return nil } if hasLegacyData { s.backup() } for _, m := range migrations { if m.version <= current { continue } if err := s.applyMigration(m); err != nil { return fmt.Errorf("迁移 v%d(%s)失败: %w", m.version, m.name, err) } } return setDBVersion(s.db, SchemaVersion) } // applyMigration 在一个事务里执行单条迁移。 func (s *Store) applyMigration(m migration) error { tx, err := s.db.Begin() if err != nil { return err } defer func() { _ = tx.Rollback() }() if err := m.run(tx); err != nil { return err } if _, err := tx.Exec(`INSERT INTO schema_meta(key, value) VALUES('schema_version', ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value`, strconv.Itoa(m.version)); err != nil { return err } return tx.Commit() } // backupDir 是数据库备份目录名。 const backupDir = "backups" // backup 用 SQLite 的 VACUUM INTO 生成一个一致性快照。 // // 迁移前留一份副本,是为了在迁移逻辑有问题时还能把数据捞回来。 // 首次安装(空库)不需要备份。 func (s *Store) backup() { if s.dbPath == "" { return } st, err := os.Stat(s.dbPath) if err != nil || st.Size() < 4096 { return // 还没建过表的空库 } dir := filepath.Join(filepath.Dir(s.dbPath), backupDir) if err := os.MkdirAll(dir, 0o755); err != nil { fmt.Fprintf(os.Stderr, "创建备份目录失败: %v\n", err) return } // 文件名带随机后缀:同一秒内连续两次启动(例如服务自动重启)也不会撞名。 dst := filepath.Join(dir, fmt.Sprintf("gitcat-%s-%s.db", time.Now().Format("20060102-150405"), randomToken(4))) // VACUUM INTO 的目标路径不能是绑定参数(SQLite 语法限制),只能自己拼; // 时间与随机串都由本进程生成,不含外部输入。 if _, err := s.db.Exec(`VACUUM INTO ?`, dst); err != nil { fmt.Fprintf(os.Stderr, "迁移前备份数据库失败(不影响升级继续): %v\n", err) return } fmt.Fprintf(os.Stderr, "迁移前已备份数据库: %s\n", dst) } // SchemaVersion 返回数据库当前的版本号。 func (s *Store) SchemaVersion() int { v, err := dbVersion(s.db) if err != nil { return 0 } return v } func nowString() string { return time.Now().UTC().Format(timeFormat) } func parseTime(s string) time.Time { if s == "" { return time.Time{} } if t, err := time.Parse(timeFormat, s); err == nil { return t } if t, err := time.Parse(time.RFC3339, s); err == nil { return t } return time.Time{} } func isUniqueErr(err error) bool { return err != nil && strings.Contains(strings.ToLower(err.Error()), "unique constraint") } // randomToken 生成用于文件名等非安全场景的短随机串。 func randomToken(n int) string { b := make([]byte, n) if _, err := rand.Read(b); err != nil { return strconv.FormatInt(time.Now().UnixNano(), 36) } return hex.EncodeToString(b) } // ---------------------------------------------------------------- settings // GetSetting 读取配置项,不存在时返回空串。 func (s *Store) GetSetting(key string) string { var v string err := s.db.QueryRow(`SELECT value FROM settings WHERE key = ?`, key).Scan(&v) if err != nil { return "" } return v } // SetSetting 写入(或覆盖)配置项。 func (s *Store) SetSetting(key, value string) error { _, err := s.db.Exec(`INSERT INTO settings(key, value) VALUES(?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value`, key, value) return err } // Settings 返回全部配置项。 func (s *Store) Settings() map[string]string { out := map[string]string{} rows, err := s.db.Query(`SELECT key, value FROM settings`) if err != nil { return out } defer rows.Close() for rows.Next() { var k, v string if rows.Scan(&k, &v) == nil { out[k] = v } } return out } // Installed 判断是否已完成安装。 func (s *Store) Installed() bool { return s.GetSetting("installed") == "1" } // ---------------------------------------------------------------- users // User 表示一个站点账号。 type User struct { ID int64 Username string DisplayName string Email string Bio string PasswordHash string Avatar string IsAdmin bool IsDisabled bool CreatedAt time.Time LastLoginAt time.Time // 统计字段(按需填充) RepoCount int PushCount int } // Display 返回用于展示的名字。 func (u *User) Display() string { if u != nil && u.DisplayName != "" { return u.DisplayName } if u == nil { return "" } return u.Username } // Initial 返回头像占位字符。 func (u *User) Initial() string { name := u.Display() for _, r := range name { return strings.ToUpper(string(r)) } return "?" } func scanUser(row interface { Scan(dest ...any) error }) (*User, error) { var u User var created, lastLogin string var admin, disabled int err := row.Scan(&u.ID, &u.Username, &u.DisplayName, &u.Email, &u.Bio, &u.PasswordHash, &u.Avatar, &admin, &disabled, &created, &lastLogin) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, ErrNotFound } return nil, err } u.IsAdmin = admin == 1 u.IsDisabled = disabled == 1 u.CreatedAt = parseTime(created) u.LastLoginAt = parseTime(lastLogin) return &u, nil } const userCols = `id, username, display_name, email, bio, password_hash, avatar, is_admin, is_disabled, created_at, last_login_at` // CreateUser 新建账号。 func (s *Store) CreateUser(username, display, email, hash string, isAdmin bool) (*User, error) { ts := nowString() admin := 0 if isAdmin { admin = 1 } res, err := s.db.Exec(`INSERT INTO users(username, display_name, email, password_hash, is_admin, created_at) VALUES(?, ?, ?, ?, ?, ?)`, username, display, email, hash, admin, ts) if err != nil { if isUniqueErr(err) { return nil, fmt.Errorf("用户名 %q 已存在", username) } return nil, err } id, _ := res.LastInsertId() return s.UserByID(id) } // UserByID 按主键查询用户。 func (s *Store) UserByID(id int64) (*User, error) { return scanUser(s.db.QueryRow(`SELECT `+userCols+` FROM users WHERE id = ?`, id)) } // UserByUsername 按用户名查询(不区分大小写)。 func (s *Store) UserByUsername(name string) (*User, error) { return scanUser(s.db.QueryRow(`SELECT `+userCols+` FROM users WHERE username = ? COLLATE NOCASE`, name)) } // UserByEmail 按邮箱查询。 func (s *Store) UserByEmail(email string) (*User, error) { return scanUser(s.db.QueryRow(`SELECT `+userCols+` FROM users WHERE email = ? COLLATE NOCASE AND email <> ''`, email)) } // ListUsers 返回所有用户,附带仓库数与推送数。 func (s *Store) ListUsers() ([]*User, error) { rows, err := s.db.Query(`SELECT ` + userCols + `, (SELECT COUNT(*) FROM repos r WHERE r.owner_id = users.id), (SELECT COUNT(*) FROM pushes p WHERE p.user_id = users.id) FROM users ORDER BY is_admin DESC, username COLLATE NOCASE ASC`) if err != nil { return nil, err } defer rows.Close() var out []*User for rows.Next() { var u User var created, lastLogin string var admin, disabled int if err := rows.Scan(&u.ID, &u.Username, &u.DisplayName, &u.Email, &u.Bio, &u.PasswordHash, &u.Avatar, &admin, &disabled, &created, &lastLogin, &u.RepoCount, &u.PushCount); err != nil { return nil, err } u.IsAdmin = admin == 1 u.IsDisabled = disabled == 1 u.CreatedAt = parseTime(created) u.LastLoginAt = parseTime(lastLogin) out = append(out, &u) } return out, rows.Err() } // UpdateProfile 更新用户资料(不含密码)。 func (s *Store) UpdateProfile(id int64, display, email, bio, avatar string) error { _, err := s.db.Exec(`UPDATE users SET display_name = ?, email = ?, bio = ?, avatar = ? WHERE id = ?`, display, email, bio, avatar, id) return err } // SetPassword 更新密码哈希。 func (s *Store) SetPassword(id int64, hash string) error { _, err := s.db.Exec(`UPDATE users SET password_hash = ? WHERE id = ?`, hash, id) return err } // SetUserFlags 更新管理员 / 禁用状态。 func (s *Store) SetUserFlags(id int64, isAdmin, isDisabled bool) error { b2i := func(b bool) int { if b { return 1 } return 0 } _, err := s.db.Exec(`UPDATE users SET is_admin = ?, is_disabled = ? WHERE id = ?`, b2i(isAdmin), b2i(isDisabled), id) return err } // TouchLogin 记录最近登录时间。 func (s *Store) TouchLogin(id int64) error { _, err := s.db.Exec(`UPDATE users SET last_login_at = ? WHERE id = ?`, nowString(), id) return err } // DeleteUser 删除用户(仓库 owner 置空,推送记录保留人名快照)。 func (s *Store) DeleteUser(id int64) error { _, err := s.db.Exec(`DELETE FROM users WHERE id = ?`, id) return err } // Counts 返回用户 / 仓库 / 推送总数。 func (s *Store) Counts() (users, repos, pushes int, err error) { err = s.db.QueryRow(`SELECT (SELECT COUNT(*) FROM users), (SELECT COUNT(*) FROM repos), (SELECT COUNT(*) FROM pushes)`).Scan(&users, &repos, &pushes) return } // CountAdmins 返回管理员数量。 func (s *Store) CountAdmins() (int, error) { var n int err := s.db.QueryRow(`SELECT COUNT(*) FROM users WHERE is_admin = 1`).Scan(&n) return n, err } // ---------------------------------------------------------------- sessions // Session 表示一次登录会话。 type Session struct { ID string UserID int64 CSRF string CreatedAt time.Time ExpiresAt time.Time } // CreateSession 写入会话记录。 func (s *Store) CreateSession(id string, userID int64, csrf string, ttl time.Duration, ip, ua string) error { ts := time.Now().UTC() _, err := s.db.Exec(`INSERT INTO sessions(id, user_id, csrf, ip, user_agent, created_at, expires_at) VALUES(?, ?, ?, ?, ?, ?, ?)`, id, userID, csrf, ip, ua, ts.Format(timeFormat), ts.Add(ttl).Format(timeFormat)) return err } // SessionByID 查询有效会话,并返回对应用户。 func (s *Store) SessionByID(id string) (*Session, *User, error) { var sess Session var created, expires string err := s.db.QueryRow(`SELECT id, user_id, csrf, created_at, expires_at FROM sessions WHERE id = ?`, id). Scan(&sess.ID, &sess.UserID, &sess.CSRF, &created, &expires) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, nil, ErrNotFound } return nil, nil, err } sess.CreatedAt = parseTime(created) sess.ExpiresAt = parseTime(expires) if time.Now().After(sess.ExpiresAt) { _ = s.DeleteSession(id) return nil, nil, ErrNotFound } u, err := s.UserByID(sess.UserID) if err != nil { return nil, nil, err } if u.IsDisabled { _ = s.DeleteSession(id) return nil, nil, ErrNotFound } return &sess, u, nil } // DeleteSession 删除会话(退出登录)。 func (s *Store) DeleteSession(id string) error { _, err := s.db.Exec(`DELETE FROM sessions WHERE id = ?`, id) return err } // CleanupSessions 清理过期会话。 func (s *Store) CleanupSessions() error { _, err := s.db.Exec(`DELETE FROM sessions WHERE expires_at < ?`, nowString()) return err }