Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -259,6 +259,9 @@ https://服务器IP:18888/随机入口路径
# 安装脚本
curl -fsSL https://raw.githubusercontent.com/luoye663/nxpanel/main/install.sh | sudo bash

# 升级脚本
curl -fsSL https://raw.githubusercontent.com/luoye663/nxpanel/main/upgrade.sh | sudo bash

# 如需卸载,执行
curl -fsSL https://raw.githubusercontent.com/luoye663/nxpanel/main/uninstall.sh | sudo bash
```
Expand Down
191 changes: 0 additions & 191 deletions internal/db/db.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,8 @@
package db

import (
"crypto/rand"
"database/sql"
"embed"
"encoding/hex"
"fmt"
"log/slog"
"sort"
Expand Down Expand Up @@ -155,198 +153,9 @@ func RunMigrations(db *sql.DB) error {
slog.Info("数据库迁移完成", "version", m.Version)
}

if err := ensureAuthAccountSchema(db); err != nil {
return fmt.Errorf("修复访问账户数据库结构失败: %w", err)
}

return nil
}

func ensureAuthAccountSchema(db *sql.DB) error {
if err := ensureAuthAccountTables(db); err != nil {
return err
}
if err := migrateLegacyAuthRules(db); err != nil {
return err
}
_, _ = db.Exec("INSERT OR IGNORE INTO schema_migrations (version) VALUES (2)")
return nil
}

func ensureAuthAccountTables(db *sql.DB) error {
statements := []string{
`CREATE TABLE IF NOT EXISTS auth_accounts (
id TEXT PRIMARY KEY,
scope TEXT NOT NULL CHECK (scope IN ('global','site')),
site_id TEXT REFERENCES sites(id) ON DELETE CASCADE,
username TEXT NOT NULL UNIQUE,
password_hash TEXT NOT NULL,
enabled INTEGER NOT NULL DEFAULT 1,
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
CHECK ((scope = 'global' AND site_id IS NULL) OR (scope = 'site' AND site_id IS NOT NULL))
)`,
`CREATE INDEX IF NOT EXISTS idx_auth_accounts_scope_site ON auth_accounts(scope, site_id)`,
`CREATE TABLE IF NOT EXISTS site_auth_rule_accounts (
rule_id TEXT NOT NULL REFERENCES site_auth_rules(id) ON DELETE CASCADE,
account_id TEXT NOT NULL REFERENCES auth_accounts(id) ON DELETE RESTRICT,
PRIMARY KEY (rule_id, account_id)
)`,
`CREATE INDEX IF NOT EXISTS idx_site_auth_rule_accounts_account ON site_auth_rule_accounts(account_id)`,
`CREATE TABLE IF NOT EXISTS site_proxy_auth_accounts (
proxy_id TEXT NOT NULL REFERENCES site_proxy(id) ON DELETE CASCADE,
account_id TEXT NOT NULL REFERENCES auth_accounts(id) ON DELETE RESTRICT,
PRIMARY KEY (proxy_id, account_id)
)`,
`CREATE INDEX IF NOT EXISTS idx_site_proxy_auth_accounts_account ON site_proxy_auth_accounts(account_id)`,
}
for _, stmt := range statements {
if _, err := db.Exec(stmt); err != nil {
return err
}
}
if ok, err := columnExists(db, "site_proxy", "auth_enabled"); err != nil {
return err
} else if !ok {
if _, err := db.Exec(`ALTER TABLE site_proxy ADD COLUMN auth_enabled INTEGER NOT NULL DEFAULT 0`); err != nil {
return err
}
}
if ok, err := columnExists(db, "site_proxy", "auth_htpasswd_path"); err != nil {
return err
} else if !ok {
if _, err := db.Exec(`ALTER TABLE site_proxy ADD COLUMN auth_htpasswd_path TEXT NOT NULL DEFAULT ''`); err != nil {
return err
}
}
return nil
}

func migrateLegacyAuthRules(db *sql.DB) error {
rows, err := db.Query(`SELECT r.id, r.site_id, r.username, r.password_hash, r.created_at, r.updated_at
FROM site_auth_rules r
LEFT JOIN site_auth_rule_accounts ra ON ra.rule_id = r.id
WHERE ra.rule_id IS NULL AND r.username IS NOT NULL AND r.username != ''`)
if err != nil {
return err
}
defer rows.Close()

type legacyRule struct {
ruleID, siteID, username, passwordHash, createdAt, updatedAt string
}
var rules []legacyRule
for rows.Next() {
var rule legacyRule
if err := rows.Scan(&rule.ruleID, &rule.siteID, &rule.username, &rule.passwordHash, &rule.createdAt, &rule.updatedAt); err != nil {
return err
}
rules = append(rules, rule)
}
if err := rows.Err(); err != nil {
return err
}

for _, rule := range rules {
accountID, ok, err := compatibleAccountID(db, rule.username, rule.passwordHash)
if err != nil {
return err
}
if !ok {
username, err := uniqueAuthUsername(db, rule.username, rule.ruleID)
if err != nil {
return err
}
accountID = newMigrationID()
if _, err := db.Exec(`INSERT INTO auth_accounts (id, scope, site_id, username, password_hash, enabled, created_at, updated_at)
VALUES (?, 'site', ?, ?, ?, 1, ?, ?)`, accountID, rule.siteID, username, rule.passwordHash, rule.createdAt, rule.updatedAt); err != nil {
return err
}
}
if _, err := db.Exec(`INSERT OR IGNORE INTO site_auth_rule_accounts (rule_id, account_id) VALUES (?, ?)`, rule.ruleID, accountID); err != nil {
return err
}
}
return nil
}

func compatibleAccountID(db *sql.DB, username, passwordHash string) (string, bool, error) {
var id, existingHash string
err := db.QueryRow(`SELECT id, password_hash FROM auth_accounts WHERE username = ?`, username).Scan(&id, &existingHash)
if err == sql.ErrNoRows {
return "", false, nil
}
if err != nil {
return "", false, err
}
if existingHash == passwordHash {
return id, true, nil
}
return "", false, nil
}

func uniqueAuthUsername(db *sql.DB, username, ruleID string) (string, error) {
exists, err := authUsernameExists(db, username)
if err != nil {
return "", err
}
if !exists {
return username, nil
}
base := username + "_" + ruleID
for i := 0; ; i++ {
candidate := base
if i > 0 {
candidate = fmt.Sprintf("%s_%d", base, i)
}
exists, err := authUsernameExists(db, candidate)
if err != nil {
return "", err
}
if !exists {
return candidate, nil
}
}
}

func authUsernameExists(db *sql.DB, username string) (bool, error) {
var count int
if err := db.QueryRow(`SELECT COUNT(*) FROM auth_accounts WHERE username = ?`, username).Scan(&count); err != nil {
return false, err
}
return count > 0, nil
}

func columnExists(db *sql.DB, tableName, columnName string) (bool, error) {
rows, err := db.Query("PRAGMA table_info(" + tableName + ")")
if err != nil {
return false, err
}
defer rows.Close()
for rows.Next() {
var cid int
var name, typ string
var notNull int
var defaultValue any
var pk int
if err := rows.Scan(&cid, &name, &typ, &notNull, &defaultValue, &pk); err != nil {
return false, err
}
if name == columnName {
return true, nil
}
}
return false, rows.Err()
}

func newMigrationID() string {
b := make([]byte, 8)
if _, err := rand.Read(b); err != nil {
panic("crypto/rand.Read failed: " + err.Error())
}
return "aa_" + hex.EncodeToString(b)
}

// loadMigrations 从 embed.FS 加载所有迁移文件并按版本号排序
func loadMigrations() ([]Migration, error) {
entries, err := migrationsFS.ReadDir("migrations")
Expand Down
Loading