fix: 修复写队列 busy 恢复、关闭安全、备份与配置构建问题

This commit is contained in:
Nixevol
2026-09-30 16:21:05 +08:00
parent 4059a1576b
commit 57c2f70678
23 changed files with 825 additions and 101 deletions
+8 -33
View File
@@ -12,6 +12,9 @@ import (
// MaxBodyBytesCap 是 max_body_bytes 的硬上限(DEVELOPMENT 11.1)。
const MaxBodyBytesCap = 262144
// MaxFrameBytesCap 是 max_frame_bytes 的硬上限,与 broker MaximumPacketSize 一致。
const MaxFrameBytesCap = 786432
// Config 对应 DEVELOPMENT 第 11.1 节的配置文件。
type Config struct {
Listen string `yaml:"listen"`
@@ -124,39 +127,8 @@ func applyEmptyDefaults(cfg *Config) {
if cfg.Log.Level == "" {
cfg.Log.Level = def.Log.Level
}
if cfg.Limits.MaxBodyBytes == 0 {
cfg.Limits.MaxBodyBytes = def.Limits.MaxBodyBytes
}
if cfg.Limits.MaxMetaBytes == 0 {
cfg.Limits.MaxMetaBytes = def.Limits.MaxMetaBytes
}
if cfg.Limits.MaxFrameBytes == 0 {
cfg.Limits.MaxFrameBytes = def.Limits.MaxFrameBytes
}
if cfg.Limits.MaxTTLSeconds == 0 {
cfg.Limits.MaxTTLSeconds = def.Limits.MaxTTLSeconds
}
if cfg.Limits.MaxScheduleSeconds == 0 {
cfg.Limits.MaxScheduleSeconds = def.Limits.MaxScheduleSeconds
}
if cfg.Limits.MaxGroupMembers == 0 {
cfg.Limits.MaxGroupMembers = def.Limits.MaxGroupMembers
}
if cfg.Limits.GraceSeconds == 0 {
cfg.Limits.GraceSeconds = def.Limits.GraceSeconds
}
if cfg.Limits.AckTimeoutSeconds == 0 {
cfg.Limits.AckTimeoutSeconds = def.Limits.AckTimeoutSeconds
}
if cfg.Limits.DeliveryWindow == 0 {
cfg.Limits.DeliveryWindow = def.Limits.DeliveryWindow
}
if cfg.Limits.ReceiptWindow == 0 {
cfg.Limits.ReceiptWindow = def.Limits.ReceiptWindow
}
if cfg.Limits.RequestsPerSecond == 0 {
cfg.Limits.RequestsPerSecond = def.Limits.RequestsPerSecond
}
// 数值字段不在这里把 0 改回默认:Load 已先填 Default 再解析 YAML,
// 未写的字段保留默认值;显式写 0 对 grace 等字段有意义,其余由 Validate 拒绝。
}
// Validate 校验 DEVELOPMENT 11.1 全部字段;拒绝 max_body_bytes > 262144。
@@ -189,6 +161,9 @@ func (c Config) Validate() error {
if c.Limits.MaxFrameBytes <= 0 {
errs = append(errs, "limits.max_frame_bytes must be > 0")
}
if c.Limits.MaxFrameBytes > MaxFrameBytesCap {
errs = append(errs, fmt.Sprintf("limits.max_frame_bytes must be <= %d", MaxFrameBytesCap))
}
if c.Limits.MaxFrameBytes < c.Limits.MaxBodyBytes {
errs = append(errs, "limits.max_frame_bytes must be >= limits.max_body_bytes")
}
+56
View File
@@ -48,6 +48,62 @@ func TestLoadAndValidate(t *testing.T) {
}
}
func TestLoadExplicitZeroGraceSeconds(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, "config.yaml")
body := "listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dir) + "\"\nlimits:\n grace_seconds: 0\n"
if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
t.Fatal(err)
}
cfg, err := Load(path)
if err != nil {
t.Fatal(err)
}
if err := cfg.Validate(); err != nil {
t.Fatal(err)
}
if cfg.Limits.GraceSeconds != 0 {
t.Fatalf("grace_seconds=%d, want 0", cfg.Limits.GraceSeconds)
}
if cfg.Limits.MaxBodyBytes != MaxBodyBytesCap {
t.Fatalf("omitted max_body_bytes=%d", cfg.Limits.MaxBodyBytes)
}
}
func TestValidateRejectsMaxFrameBytesTooLarge(t *testing.T) {
t.Parallel()
cfg := Default()
cfg.Limits.MaxFrameBytes = MaxFrameBytesCap + 1
err := cfg.Validate()
if err == nil {
t.Fatal("expected error")
}
if !strings.Contains(err.Error(), "max_frame_bytes") {
t.Fatalf("unexpected error: %v", err)
}
}
func TestLoadExplicitZeroMaxBodyBytesRejected(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, "config.yaml")
body := "listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dir) + "\"\nlimits:\n max_body_bytes: 0\n"
if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
t.Fatal(err)
}
cfg, err := Load(path)
if err != nil {
t.Fatal(err)
}
if cfg.Limits.MaxBodyBytes != 0 {
t.Fatalf("max_body_bytes=%d, want explicit 0", cfg.Limits.MaxBodyBytes)
}
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "max_body_bytes") {
t.Fatalf("want max_body_bytes error, got %v", err)
}
}
func TestValidateTrustedProxies(t *testing.T) {
t.Parallel()
cfg := Default()
+28 -5
View File
@@ -8,7 +8,13 @@ import (
"time"
)
const settingAdminPasswordHash = "admin_password_hash"
const (
settingAdminPasswordHash = "admin_password_hash"
adminPasswordSQL = `
INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at
`
)
// HasAdminPassword 检查 settings 中是否已有管理员密码哈希。
func HasAdminPassword(ctx context.Context, db *sql.DB) (bool, error) {
@@ -26,13 +32,30 @@ func HasAdminPassword(ctx context.Context, db *sql.DB) (bool, error) {
// SetAdminPasswordHash 写入或覆盖管理员密码哈希。
func SetAdminPasswordHash(ctx context.Context, db *sql.DB, phc string) error {
now := time.Now().UnixMilli()
_, err := db.ExecContext(ctx, `
INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at
`, settingAdminPasswordHash, phc, now)
_, err := db.ExecContext(ctx, adminPasswordSQL, settingAdminPasswordHash, phc, now)
return err
}
func setAdminPasswordHashTx(tx *sql.Tx, phc string) error {
now := time.Now().UnixMilli()
_, err := tx.Exec(adminPasswordSQL, settingAdminPasswordHash, phc, now)
return err
}
// ResetAdminPassword 更新管理员密码哈希并清空全部管理会话。
func (d *DB) ResetAdminPassword(ctx context.Context, phc string) error {
if d == nil || d.Queue == nil {
return fmt.Errorf("store: not open")
}
return d.Queue.Do(ctx, func(tx *sql.Tx) error {
if err := setAdminPasswordHashTx(tx, phc); err != nil {
return err
}
_, err := tx.Exec(`DELETE FROM admin_sessions`)
return err
})
}
// VacuumInto 对打开的写连接执行 VACUUM INTO(可用于运行中备份)。
func VacuumInto(ctx context.Context, db *sql.DB, outPath string) error {
if outPath == "" {
+38 -5
View File
@@ -1,6 +1,7 @@
package store
import (
"context"
"database/sql"
"fmt"
"os"
@@ -10,7 +11,8 @@ import (
_ "modernc.org/sqlite"
)
const dbFileName = "nixmsg.db"
// DBFileName 是数据目录下的主库文件名。
const DBFileName = "nixmsg.db"
// DB 持有读写连接与写入队列。
type DB struct {
@@ -25,7 +27,7 @@ func Open(dataDir, synchronous string) (*DB, error) {
if err := os.MkdirAll(dataDir, 0o755); err != nil {
return nil, fmt.Errorf("mkdir data_dir: %w", err)
}
dbPath := filepath.Join(dataDir, dbFileName)
dbPath := filepath.Join(dataDir, DBFileName)
existed, err := fileExists(dbPath)
if err != nil {
return nil, err
@@ -69,10 +71,41 @@ func (d *DB) Close() error {
return first
}
// Checkpoint 在写 goroutine 上执行 wal_checkpoint(TRUNCATE)。
func (d *DB) Checkpoint(ctx context.Context) error {
if d == nil || d.Queue == nil {
return fmt.Errorf("store: not open")
}
return d.Queue.Checkpoint(ctx)
}
// OpenWriter 按 DEVELOPMENT 7.7 打开写连接(带 _txlock=immediate,MaxOpenConns=1)。
func OpenWriter(dataDir, synchronous string) (*sql.DB, error) {
if err := os.MkdirAll(dataDir, 0o755); err != nil {
return nil, fmt.Errorf("mkdir data_dir: %w", err)
return openWriter(dataDir, synchronous, true)
}
// OpenExistingWriter 打开已有库的写连接,不创建数据目录;库文件不存在时返回错误。
func OpenExistingWriter(dataDir, synchronous string) (*sql.DB, error) {
dbPath := filepath.Join(dataDir, DBFileName)
ok, err := fileExists(dbPath)
if err != nil {
return nil, err
}
if !ok {
abs, absErr := filepath.Abs(dbPath)
if absErr != nil {
abs = dbPath
}
return nil, fmt.Errorf("database not found: %s", abs)
}
return openWriter(dataDir, synchronous, false)
}
func openWriter(dataDir, synchronous string, mkdir bool) (*sql.DB, error) {
if mkdir {
if err := os.MkdirAll(dataDir, 0o755); err != nil {
return nil, fmt.Errorf("mkdir data_dir: %w", err)
}
}
dsn, err := buildDSN(dataDir, synchronous, true)
if err != nil {
@@ -116,7 +149,7 @@ func buildDSN(dataDir, synchronous string, writer bool) (string, error) {
if sync != "FULL" && sync != "NORMAL" {
return "", fmt.Errorf("invalid sqlite_synchronous: %s", synchronous)
}
dbPath := filepath.ToSlash(filepath.Join(dataDir, dbFileName))
dbPath := filepath.ToSlash(filepath.Join(dataDir, DBFileName))
dsn := fmt.Sprintf(
"file:%s?_pragma=journal_mode(WAL)&_pragma=busy_timeout(5000)&_pragma=synchronous(%s)&_pragma=foreign_keys(ON)&_pragma=secure_delete(ON)",
dbPath,
+4 -3
View File
@@ -5,7 +5,6 @@ import (
"database/sql"
"os"
"path/filepath"
"strings"
"testing"
"time"
)
@@ -94,8 +93,10 @@ CREATE TABLE schema_migrations (
t.Fatalf("backup dir missing: %v", err)
}
found := false
var names []string
for _, e := range entries {
if strings.HasPrefix(e.Name(), "pre-migrate-") && strings.HasSuffix(e.Name(), ".db") {
names = append(names, e.Name())
if e.Name() == "pre-migrate-v1-to-v2.db" {
found = true
info, statErr := e.Info()
if statErr != nil {
@@ -107,7 +108,7 @@ CREATE TABLE schema_migrations (
}
}
if !found {
t.Fatal("expected pre-migrate-*.db backup")
t.Fatalf("expected pre-migrate-v1-to-v2.db, got %v", names)
}
}
+9
View File
@@ -0,0 +1,9 @@
//go:build !windows && !unix
package store
import "fmt"
func availableBytes(_ string) (uint64, error) {
return 0, fmt.Errorf("disk space probe unsupported on this platform")
}
+13
View File
@@ -0,0 +1,13 @@
//go:build unix
package store
import "golang.org/x/sys/unix"
func availableBytes(path string) (uint64, error) {
var st unix.Statfs_t
if err := unix.Statfs(path, &st); err != nil {
return 0, err
}
return uint64(st.Bavail) * uint64(st.Bsize), nil
}
+17
View File
@@ -0,0 +1,17 @@
//go:build windows
package store
import "golang.org/x/sys/windows"
func availableBytes(path string) (uint64, error) {
p, err := windows.UTF16PtrFromString(path)
if err != nil {
return 0, err
}
var free, total, totalFree uint64
if err := windows.GetDiskFreeSpaceEx(p, &free, &total, &totalFree); err != nil {
return 0, err
}
return free, nil
}
+74 -9
View File
@@ -18,9 +18,22 @@ var migrationFS embed.FS
// Migrate 应用尚未执行的嵌入迁移。
// 若 dbExisted 为 true(调用 Open/Migrate 前已有 nixmsg.db)且存在未应用版本,
// 先 VACUUM INTO <data_dir>/backup/pre-migrate-<时间>.db,再迁移。
// 先 VACUUM INTO <data_dir>/backup/pre-migrate-v{当前}-to-v{目标}.db,再迁移。
// 备份文件已存在则复用,避免迁移稳定失败时反复全量复制。
// 任一版本失败则返回错误,调用方不得继续带半新半旧库提供服务。
func Migrate(db *sql.DB, dataDir string, dbExisted bool) error {
if err := ensureMigrationsTable(db); err != nil {
return err
}
pending, err := pendingMigrations(db)
if err != nil {
return err
}
return applyPending(db, dataDir, dbExisted, pending)
}
func ensureMigrationsTable(db *sql.DB) error {
if _, err := db.Exec(`
CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
@@ -28,17 +41,21 @@ CREATE TABLE IF NOT EXISTS schema_migrations (
)`); err != nil {
return fmt.Errorf("ensure schema_migrations: %w", err)
}
return nil
}
pending, err := pendingMigrations(db)
if err != nil {
return err
}
func applyPending(db *sql.DB, dataDir string, dbExisted bool, pending []migrationFile) error {
if len(pending) == 0 {
return nil
}
if dbExisted {
if err := backupBeforeMigrate(db, dataDir); err != nil {
from, err := currentSchemaVersion(db)
if err != nil {
return err
}
to := pending[len(pending)-1].version
if err := backupBeforeMigrate(db, dataDir, from, to); err != nil {
return err
}
}
@@ -98,13 +115,33 @@ func pendingMigrations(db *sql.DB) ([]migrationFile, error) {
return pending, nil
}
func backupBeforeMigrate(db *sql.DB, dataDir string) error {
func currentSchemaVersion(db *sql.DB) (int, error) {
var v sql.NullInt64
if err := db.QueryRow(`SELECT MAX(version) FROM schema_migrations`).Scan(&v); err != nil {
return 0, fmt.Errorf("current schema version: %w", err)
}
if !v.Valid {
return 0, nil
}
return int(v.Int64), nil
}
func backupBeforeMigrate(db *sql.DB, dataDir string, fromVer, toVer int) error {
backupDir := filepath.Join(dataDir, "backup")
if err := os.MkdirAll(backupDir, 0o755); err != nil {
return fmt.Errorf("mkdir backup: %w", err)
}
stamp := time.Now().UTC().Format("20060102T150405")
backupPath := filepath.Join(backupDir, "pre-migrate-"+stamp+".db")
backupPath := filepath.Join(backupDir, fmt.Sprintf("pre-migrate-v%d-to-v%d.db", fromVer, toVer))
exists, err := fileExists(backupPath)
if err != nil {
return err
}
if exists {
return nil
}
if err := ensureDiskSpace(backupDir, backupNeedBytes(dataDir)); err != nil {
return err
}
// SQLite VACUUM INTO 需要字面量路径;统一用斜杠,并对单引号转义。
quoted := strings.ReplaceAll(filepath.ToSlash(backupPath), "'", "''")
if _, err := db.Exec("VACUUM INTO '" + quoted + "'"); err != nil {
@@ -113,6 +150,34 @@ func backupBeforeMigrate(db *sql.DB, dataDir string) error {
return nil
}
func backupNeedBytes(dataDir string) int64 {
var n int64
for _, name := range []string{DBFileName, DBFileName + "-wal", DBFileName + "-shm"} {
st, err := os.Stat(filepath.Join(dataDir, name))
if err == nil {
n += st.Size()
}
}
return n
}
func ensureDiskSpace(dir string, need int64) error {
if need < 0 {
need = 0
}
avail, err := availableBytes(dir)
if err != nil {
// 探测失败不阻断迁移,VACUUM INTO 自身会因空间不足报错。
return nil
}
const margin = 1 << 20
want := uint64(need) + margin
if avail < want {
return fmt.Errorf("insufficient disk space for migrate backup: have %d bytes, need ~%d", avail, want)
}
return nil
}
func applyMigration(db *sql.DB, m migrationFile) error {
tx, err := db.Begin()
if err != nil {
+73
View File
@@ -0,0 +1,73 @@
package store
import (
"os"
"path/filepath"
"testing"
"time"
)
func TestFailingMigrationReusesBackup(t *testing.T) {
t.Parallel()
dir := t.TempDir()
db, err := Open(dir, "FULL")
if err != nil {
t.Fatal(err)
}
_ = db.Close()
w, err := OpenWriter(dir, "FULL")
if err != nil {
t.Fatal(err)
}
defer func() { _ = w.Close() }()
pending := []migrationFile{{
version: 9999,
name: "9999_fail.sql",
body: "THIS IS NOT VALID SQL",
}}
if err := applyPending(w, dir, true, pending); err == nil {
t.Fatal("expected first failing migration to error")
}
backupDir := filepath.Join(dir, "backup")
first, err := os.ReadDir(backupDir)
if err != nil {
t.Fatal(err)
}
if len(first) != 1 || first[0].Name() != "pre-migrate-v2-to-v9999.db" {
t.Fatalf("first backups=%v", dirNames(first))
}
if err := applyPending(w, dir, true, pending); err == nil {
t.Fatal("expected second failing migration to error")
}
second, err := os.ReadDir(backupDir)
if err != nil {
t.Fatal(err)
}
if len(second) != 1 || second[0].Name() != first[0].Name() {
t.Fatalf("second backups=%v first=%v", dirNames(second), dirNames(first))
}
info1, err := first[0].Info()
if err != nil {
t.Fatal(err)
}
info2, err := second[0].Info()
if err != nil {
t.Fatal(err)
}
if info2.ModTime().After(info1.ModTime().Add(time.Second)) && info2.Size() != info1.Size() {
// 复用同一文件即可;时钟精度下允许 mtime 相同。
t.Logf("mtime first=%s second=%s", info1.ModTime(), info2.ModTime())
}
}
func dirNames(entries []os.DirEntry) []string {
out := make([]string, 0, len(entries))
for _, e := range entries {
out = append(out, e.Name())
}
return out
}
+113 -27
View File
@@ -28,6 +28,7 @@ type WriteFunc func(tx *sql.Tx) error
type writeJob struct {
ctx context.Context
fn WriteFunc
raw func(*sql.DB) error
res chan error
}
@@ -38,6 +39,7 @@ type Queue struct {
ch chan writeJob
done chan struct{}
closed atomic.Bool
sendMu sync.RWMutex
mu sync.Mutex
ready bool
@@ -65,34 +67,71 @@ func (q *Queue) Do(ctx context.Context, fn WriteFunc) error {
if fn == nil {
return errors.New("store: nil write func")
}
if q.closed.Load() {
return ErrQueueClosed
}
if err := ctx.Err(); err != nil {
return err
}
job := writeJob{ctx: ctx, fn: fn, res: make(chan error, 1)}
q.mu.Lock()
q.pending++
q.mu.Unlock()
select {
case q.ch <- job:
case <-ctx.Done():
q.mu.Lock()
q.pending--
q.mu.Unlock()
return ctx.Err()
case <-q.done:
q.mu.Lock()
q.pending--
q.mu.Unlock()
if err := q.enqueue(job); err != nil {
return err
}
return q.waitResult(ctx, job)
}
// ExecOnWriter 在写 goroutine 上、事务外执行 fn(如 PRAGMA wal_checkpoint)。
func (q *Queue) ExecOnWriter(ctx context.Context, fn func(*sql.DB) error) error {
if fn == nil {
return errors.New("store: nil exec func")
}
if err := ctx.Err(); err != nil {
return err
}
job := writeJob{ctx: ctx, raw: fn, res: make(chan error, 1)}
if err := q.enqueue(job); err != nil {
return err
}
return q.waitResult(ctx, job)
}
// Checkpoint 在写连接上执行 PRAGMA wal_checkpoint(TRUNCATE)。
func (q *Queue) Checkpoint(ctx context.Context) error {
return q.ExecOnWriter(ctx, func(db *sql.DB) error {
_, err := db.ExecContext(ctx, `PRAGMA wal_checkpoint(TRUNCATE)`)
return err
})
}
// Optimize 在写连接上执行 PRAGMA optimize。
func (q *Queue) Optimize(ctx context.Context) error {
return q.ExecOnWriter(ctx, func(db *sql.DB) error {
_, err := db.ExecContext(ctx, `PRAGMA optimize`)
return err
})
}
func (q *Queue) enqueue(job writeJob) error {
q.sendMu.RLock()
if q.closed.Load() {
q.sendMu.RUnlock()
return ErrQueueClosed
}
q.addPending(1)
select {
case q.ch <- job:
q.sendMu.RUnlock()
return nil
case <-job.ctx.Done():
q.addPending(-1)
q.sendMu.RUnlock()
return job.ctx.Err()
}
}
func (q *Queue) waitResult(ctx context.Context, job writeJob) error {
select {
case err := <-job.res:
return err
case <-ctx.Done():
// 操作可能仍在队列中执行;结果通道仍会被写端关闭式填入。
// 操作可能仍在队列中执行;结果通道仍会被写端填入。
select {
case err := <-job.res:
if err != nil {
@@ -112,8 +151,13 @@ func (q *Queue) loop() {
if !ok {
return
}
if job.raw != nil {
q.runRaw(job)
continue
}
batch := []writeJob{job}
timer := time.NewTimer(batchWait)
ranRaw := false
collect:
for len(batch) < maxBatchOps {
select {
@@ -121,25 +165,47 @@ func (q *Queue) loop() {
if !ok {
break collect
}
if j.raw != nil {
stopTimer(timer)
q.runBatch(batch)
q.runRaw(j)
ranRaw = true
break collect
}
batch = append(batch, j)
case <-timer.C:
break collect
}
}
timer.Stop()
q.runBatch(batch)
if !ranRaw {
stopTimer(timer)
q.runBatch(batch)
}
}
}
func stopTimer(timer *time.Timer) {
if !timer.Stop() {
select {
case <-timer.C:
default:
}
}
}
func (q *Queue) runRaw(job writeJob) {
defer q.addPending(-1)
if err := job.ctx.Err(); err != nil {
job.res <- err
return
}
job.res <- job.raw(q.db)
}
func (q *Queue) runBatch(batch []writeJob) {
started := time.Now()
defer func() {
q.mu.Lock()
q.pending -= len(batch)
if q.pending < 0 {
q.pending = 0
}
q.mu.Unlock()
q.addPending(-len(batch))
}()
// 过滤已取消的任务。
@@ -233,6 +299,7 @@ func (q *Queue) runBatch(batch []writeJob) {
}
return
}
q.markReady()
if q.OnBatchCommit != nil {
q.OnBatchCommit(time.Since(started))
}
@@ -252,7 +319,23 @@ func (q *Queue) markBusy(err error) {
q.lastWriteErr = err
}
// IsReady 写库是否仍可用(写失败后为 false)。
func (q *Queue) markReady() {
q.mu.Lock()
defer q.mu.Unlock()
q.ready = true
q.lastWriteErr = nil
}
func (q *Queue) addPending(delta int) {
q.mu.Lock()
defer q.mu.Unlock()
q.pending += delta
if q.pending < 0 {
q.pending = 0
}
}
// IsReady 写库是否仍可用(写失败后为 false;随后一次成功提交会恢复)。
func (q *Queue) IsReady() bool {
q.mu.Lock()
defer q.mu.Unlock()
@@ -293,11 +376,14 @@ func (q *Queue) Drain(ctx context.Context) error {
}
// Close 关闭队列:不再接受新任务,并等待写 goroutine 处理完已入队任务后退出。
// 在写锁内关闭数据通道,避免并发 Do 向已关闭 channel 发送而 panic。
func (q *Queue) Close() error {
if q.closed.Swap(true) {
return nil
}
q.sendMu.Lock()
close(q.ch)
q.sendMu.Unlock()
<-q.done
return nil
}
+116
View File
@@ -5,6 +5,9 @@ import (
"database/sql"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
@@ -148,6 +151,119 @@ func TestWriteQueueNoBacklogAt200PerSec(t *testing.T) {
}
}
func TestWriteQueueRecoversReadyAfterBusy(t *testing.T) {
t.Parallel()
dir := t.TempDir()
db, err := Open(dir, "FULL")
if err != nil {
t.Fatal(err)
}
defer func() { _ = db.Close() }()
db.Queue.markBusy(errors.New("injected busy"))
if db.Queue.IsReady() {
t.Fatal("expected not ready after markBusy")
}
ctx := context.Background()
if err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
"after_busy", "1", time.Now().UnixMilli(),
)
return e
}); err != nil {
t.Fatal(err)
}
if !db.Queue.IsReady() {
t.Fatalf("expected ready after successful Do, last=%v", db.Queue.LastWriteError())
}
}
func TestWriteQueueCloseConcurrentDoNoPanic(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()
const n = 1000
var wg sync.WaitGroup
wg.Add(n)
for i := 0; i < n; i++ {
i := i
go func() {
defer wg.Done()
_ = db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`,
fmt.Sprintf("close_%d", i%50), "1", time.Now().UnixMilli(),
)
return e
})
}()
}
time.Sleep(2 * time.Millisecond)
if err := db.Queue.Close(); err != nil {
t.Fatal(err)
}
wg.Wait()
err = db.Queue.Do(ctx, func(tx *sql.Tx) error { return nil })
if !errors.Is(err, ErrQueueClosed) {
t.Fatalf("want ErrQueueClosed, got %v", err)
}
}
func TestQueueCheckpointTruncatesWAL(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()
payload := strings.Repeat("x", 4096)
for i := 0; i < 300; i++ {
i := i
if err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
fmt.Sprintf("wal_%d", i), payload, time.Now().UnixMilli(),
)
return e
}); err != nil {
t.Fatal(err)
}
}
walPath := filepath.Join(dir, DBFileName+"-wal")
before, statErr := os.Stat(walPath)
if err := db.Checkpoint(ctx); err != nil {
t.Fatal(err)
}
if statErr != nil {
return
}
after, err := os.Stat(walPath)
if err != nil {
// TRUNCATE 后 WAL 可能被删掉,视为回落成功。
if os.IsNotExist(err) {
return
}
t.Fatal(err)
}
if before.Size() > 0 && after.Size() >= before.Size() {
t.Fatalf("WAL did not shrink: before=%d after=%d", before.Size(), after.Size())
}
}
func BenchmarkWriteQueue200PerSec(b *testing.B) {
dir := b.TempDir()
db, err := Open(dir, "FULL")