166 lines
3.8 KiB
Go
166 lines
3.8 KiB
Go
package store
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestOpenEmptyDirCreatesTables(t *testing.T) {
|
|
t.Parallel()
|
|
dir := t.TempDir()
|
|
|
|
db, err := Open(dir, "FULL")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = db.Close() }()
|
|
|
|
assertMigrationCount(t, db.Write, 2)
|
|
for _, table := range []string{
|
|
"endpoints", "settings", "talk_grants", "groups", "group_members",
|
|
"messages", "message_bodies", "deliveries", "receipts", "send_keys",
|
|
"admin_sessions", "api_tokens", "schema_migrations",
|
|
} {
|
|
assertTableExists(t, db.Write, table)
|
|
}
|
|
// 空目录首次启动不应产生迁移备份。
|
|
if entries, readErr := os.ReadDir(filepath.Join(dir, "backup")); readErr == nil && len(entries) > 0 {
|
|
t.Fatalf("unexpected backup files on empty start: %d", len(entries))
|
|
}
|
|
}
|
|
|
|
func TestOpenIdempotentNoRemigrate(t *testing.T) {
|
|
t.Parallel()
|
|
dir := t.TempDir()
|
|
|
|
db1, err := Open(dir, "FULL")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
assertMigrationCount(t, db1.Write, 2)
|
|
_ = db1.Close()
|
|
|
|
db2, err := Open(dir, "FULL")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = db2.Close() }()
|
|
assertMigrationCount(t, db2.Write, 2)
|
|
|
|
if entries, readErr := os.ReadDir(filepath.Join(dir, "backup")); readErr == nil && len(entries) > 0 {
|
|
t.Fatalf("idempotent reopen should not backup: %v", entries)
|
|
}
|
|
}
|
|
|
|
func TestMigrateBackupWhenNewVersion(t *testing.T) {
|
|
t.Parallel()
|
|
dir := t.TempDir()
|
|
|
|
// 模拟仅应用了 0001 的旧库。
|
|
w, err := OpenWriter(dir, "FULL")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, execErr := w.Exec(`
|
|
CREATE TABLE schema_migrations (
|
|
version INTEGER PRIMARY KEY,
|
|
applied_at INTEGER NOT NULL
|
|
)`); execErr != nil {
|
|
t.Fatal(execErr)
|
|
}
|
|
if _, execErr := w.Exec(`INSERT INTO schema_migrations(version, applied_at) VALUES(1, ?)`, time.Now().UnixMilli()); execErr != nil {
|
|
t.Fatal(execErr)
|
|
}
|
|
_ = w.Close()
|
|
|
|
db, err := Open(dir, "FULL")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = db.Close() }()
|
|
|
|
assertMigrationCount(t, db.Write, 2)
|
|
assertTableExists(t, db.Write, "endpoints")
|
|
assertTableExists(t, db.Write, "api_tokens")
|
|
|
|
backupDir := filepath.Join(dir, "backup")
|
|
entries, err := os.ReadDir(backupDir)
|
|
if err != nil {
|
|
t.Fatalf("backup dir missing: %v", err)
|
|
}
|
|
found := false
|
|
for _, e := range entries {
|
|
if strings.HasPrefix(e.Name(), "pre-migrate-") && strings.HasSuffix(e.Name(), ".db") {
|
|
found = true
|
|
info, statErr := e.Info()
|
|
if statErr != nil {
|
|
t.Fatal(statErr)
|
|
}
|
|
if info.Size() == 0 {
|
|
t.Fatalf("backup file empty: %s", e.Name())
|
|
}
|
|
}
|
|
}
|
|
if !found {
|
|
t.Fatal("expected pre-migrate-*.db backup")
|
|
}
|
|
}
|
|
|
|
func TestWriteQueueOneOpOneTx(t *testing.T) {
|
|
t.Parallel()
|
|
dir := t.TempDir()
|
|
db, err := Open(dir, "FULL")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer func() { _ = db.Close() }()
|
|
|
|
ctx := context.Background()
|
|
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
|
_, execErr := tx.Exec(
|
|
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
|
|
"registration_enabled", "0", time.Now().UnixMilli(),
|
|
)
|
|
return execErr
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
var value string
|
|
if scanErr := db.Read.QueryRow(`SELECT value FROM settings WHERE key = ?`, "registration_enabled").Scan(&value); scanErr != nil {
|
|
t.Fatal(scanErr)
|
|
}
|
|
if value != "0" {
|
|
t.Fatalf("want 0, got %q", value)
|
|
}
|
|
}
|
|
|
|
func assertMigrationCount(t *testing.T, db *sql.DB, want int) {
|
|
t.Helper()
|
|
var n int
|
|
if err := db.QueryRow(`SELECT COUNT(*) FROM schema_migrations`).Scan(&n); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if n != want {
|
|
t.Fatalf("migration rows: want %d got %d", want, n)
|
|
}
|
|
}
|
|
|
|
func assertTableExists(t *testing.T, db *sql.DB, name string) {
|
|
t.Helper()
|
|
var got string
|
|
err := db.QueryRow(
|
|
`SELECT name FROM sqlite_master WHERE type='table' AND name=?`,
|
|
name,
|
|
).Scan(&got)
|
|
if err != nil {
|
|
t.Fatalf("table %s missing: %v", name, err)
|
|
}
|
|
}
|