Compare commits

...

2 Commits

2 changed files with 170 additions and 0 deletions
@@ -123,6 +123,11 @@ func Open(path string) (*Repository, error) {
return nil, err return nil, err
} }
if err := migrateSchema(db); err != nil {
_ = db.Close()
return nil, err
}
return &Repository{db: db}, nil return &Repository{db: db}, nil
} }
@@ -1176,6 +1181,54 @@ func bootstrapSchema(db *sql.DB) error {
return nil return nil
} }
func migrateSchema(db *sql.DB) error {
if db == nil {
return errors.New("nil db")
}
ensureColumn := func(table, col, typ string) error {
var dummy interface{}
err := db.QueryRow(fmt.Sprintf("SELECT %s FROM %s LIMIT 1", col, table)).Scan(&dummy)
if err == nil || errors.Is(err, sql.ErrNoRows) {
return nil
}
if !strings.Contains(err.Error(), "no such column") {
return nil
}
if _, err := db.Exec(fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, col, typ)); err != nil {
return fmt.Errorf("add %s.%s: %w", table, col, err)
}
return nil
}
columnsByTable := map[string]map[string]string{
"node": {
"server_ip_v4": "VARCHAR(100)",
"server_ip_v6": "VARCHAR(100)",
"inx": "INTEGER NOT NULL DEFAULT 0",
},
"tunnel": {
"inx": "INTEGER NOT NULL DEFAULT 0",
},
"forward": {
"inx": "INTEGER NOT NULL DEFAULT 0",
},
"chain_tunnel": {
"inx": "INTEGER",
},
}
for table, cols := range columnsByTable {
for col, typ := range cols {
if err := ensureColumn(table, col, typ); err != nil {
return err
}
}
}
return nil
}
var osMkdirAll = func(path string) error { var osMkdirAll = func(path string) error {
return os.MkdirAll(path, 0o755) return os.MkdirAll(path, 0o755)
} }
@@ -2,6 +2,7 @@ package contract_test
import ( import (
"bytes" "bytes"
"database/sql"
"encoding/json" "encoding/json"
"io" "io"
"net/http" "net/http"
@@ -17,6 +18,8 @@ import (
"go-backend/internal/http/handler" "go-backend/internal/http/handler"
"go-backend/internal/http/response" "go-backend/internal/http/response"
"go-backend/internal/store/sqlite" "go-backend/internal/store/sqlite"
_ "modernc.org/sqlite"
) )
func TestCaptchaVerifyLoginContract(t *testing.T) { func TestCaptchaVerifyLoginContract(t *testing.T) {
@@ -213,3 +216,117 @@ func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *sqlite.
h := handler.New(repo, jwtSecret) h := handler.New(repo, jwtSecret)
return httpserver.NewRouter(h, jwtSecret), repo return httpserver.NewRouter(h, jwtSecret), repo
} }
func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "legacy-2.0.7-beta.db")
legacyDB, err := sql.Open("sqlite", dbPath)
if err != nil {
t.Fatalf("open legacy sqlite: %v", err)
}
t.Cleanup(func() {
_ = legacyDB.Close()
})
if _, err := legacyDB.Exec(`
CREATE TABLE IF NOT EXISTS node (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name VARCHAR(100) NOT NULL,
secret VARCHAR(100) NOT NULL,
server_ip VARCHAR(100) NOT NULL,
port TEXT NOT NULL,
interface_name VARCHAR(200),
version VARCHAR(100),
http INTEGER NOT NULL DEFAULT 0,
tls INTEGER NOT NULL DEFAULT 0,
socks INTEGER NOT NULL DEFAULT 0,
created_time INTEGER NOT NULL,
updated_time INTEGER,
status INTEGER NOT NULL,
tcp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]',
udp_listen_addr VARCHAR(100) NOT NULL DEFAULT '[::]'
)
`); err != nil {
t.Fatalf("create legacy node table: %v", err)
}
if _, err := legacyDB.Exec(`
CREATE TABLE IF NOT EXISTS tunnel (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name VARCHAR(100) NOT NULL,
traffic_ratio REAL NOT NULL DEFAULT 1.0,
type INTEGER NOT NULL,
protocol VARCHAR(10) NOT NULL DEFAULT 'tls',
flow INTEGER NOT NULL,
created_time INTEGER NOT NULL,
updated_time INTEGER NOT NULL,
status INTEGER NOT NULL,
in_ip TEXT
)
`); err != nil {
t.Fatalf("create legacy tunnel table: %v", err)
}
now := time.Now().UnixMilli()
if _, err := legacyDB.Exec(`
INSERT INTO node(name, secret, server_ip, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "legacy-node", "legacy-secret", "10.10.0.1", "10000-10010", "eth0", "v-old", 1, 1, 1, now, now, 1, "[::]", "[::]"); err != nil {
t.Fatalf("seed legacy node row: %v", err)
}
repo, err := sqlite.Open(dbPath)
if err != nil {
t.Fatalf("open migrated sqlite: %v", err)
}
t.Cleanup(func() {
_ = repo.Close()
})
nodes, err := repo.ListNodes()
if err != nil {
t.Fatalf("list nodes after migration: %v", err)
}
if len(nodes) != 1 {
t.Fatalf("expected 1 node after migration, got %d", len(nodes))
}
columns := readTableColumns(t, repo.DB(), "node")
for _, required := range []string{"server_ip_v4", "server_ip_v6", "inx"} {
if !columns[required] {
t.Fatalf("expected node column %q to exist after migration", required)
}
}
tunnelColumns := readTableColumns(t, repo.DB(), "tunnel")
if !tunnelColumns["inx"] {
t.Fatalf("expected tunnel column %q to exist after migration", "inx")
}
}
func readTableColumns(t *testing.T, db *sql.DB, table string) map[string]bool {
t.Helper()
rows, err := db.Query("PRAGMA table_info(" + table + ")")
if err != nil {
t.Fatalf("inspect %s columns: %v", table, err)
}
defer rows.Close()
columns := map[string]bool{}
for rows.Next() {
var cid, notNull, pk int
var name, typ string
var defaultValue sql.NullString
if err := rows.Scan(&cid, &name, &typ, &notNull, &defaultValue, &pk); err != nil {
t.Fatalf("scan %s pragma row: %v", table, err)
}
columns[name] = true
}
if err := rows.Err(); err != nil {
t.Fatalf("iterate %s pragma rows: %v", table, err)
}
return columns
}