Files
NixMsg/internal/store/queue_test.go
T

426 lines
9.5 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package store
import (
"context"
"database/sql"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
)
func TestWriteQueueSavepointIsolatesFailure(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()
var wg sync.WaitGroup
wg.Add(2)
errOK := make(chan error, 1)
errBad := make(chan error, 1)
go func() {
defer wg.Done()
errOK <- db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
"k_ok", "1", time.Now().UnixMilli(),
)
return e
})
}()
go func() {
defer wg.Done()
// 稍等,尽量与成功操作进同一批。
time.Sleep(500 * time.Microsecond)
errBad <- db.Queue.Do(ctx, func(tx *sql.Tx) error {
return errors.New("forced op failure")
})
}()
wg.Wait()
if e := <-errOK; e != nil {
t.Fatalf("ok op: %v", e)
}
if e := <-errBad; e == nil || e.Error() != "forced op failure" {
t.Fatalf("bad op: %v", e)
}
var value string
if scanErr := db.Read.QueryRow(`SELECT value FROM settings WHERE key = ?`, "k_ok").Scan(&value); scanErr != nil {
t.Fatalf("ok row missing: %v", scanErr)
}
if value != "1" {
t.Fatalf("value=%q", value)
}
if !db.Queue.IsReady() {
t.Fatal("op failure should not mark queue busy")
}
}
func TestWriteQueueBatchCommit(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 = 32
var wg sync.WaitGroup
wg.Add(n)
errs := make([]error, n)
for i := 0; i < n; i++ {
i := i
go func() {
defer wg.Done()
key := fmt.Sprintf("batch_%d", i)
errs[i] = db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
key, "1", time.Now().UnixMilli(),
)
return e
})
}()
}
wg.Wait()
for i, e := range errs {
if e != nil {
t.Fatalf("op %d: %v", i, e)
}
}
var count int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM settings WHERE key LIKE 'batch_%'`).Scan(&count); err != nil {
t.Fatal(err)
}
if count != n {
t.Fatalf("count=%d want %d", count, n)
}
}
func TestWriteQueueNoBacklogAt200PerSec(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 total = 200
start := time.Now()
var wg sync.WaitGroup
wg.Add(total)
for i := 0; i < total; i++ {
i := i
go func() {
defer wg.Done()
key := fmt.Sprintf("load_%d", i)
if e := db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
key, "1", time.Now().UnixMilli(),
)
return e
}); e != nil {
t.Errorf("op %d: %v", i, e)
}
}()
}
wg.Wait()
elapsed := time.Since(start)
if db.Queue.Len() != 0 {
t.Fatalf("queue backlog len=%d", db.Queue.Len())
}
// 200 条并发写入应在数秒内完成(合并提交);过长则合并未生效。
if elapsed > 5*time.Second {
t.Fatalf("200 writes took %s, queue likely not merging well", elapsed)
}
}
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 closeErr := db.Queue.Close(); closeErr != nil {
t.Fatal(closeErr)
}
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 doErr := 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
}); doErr != nil {
t.Fatal(doErr)
}
}
walPath := filepath.Join(dir, DBFileName+"-wal")
before, statErr := os.Stat(walPath)
if cpErr := db.Checkpoint(ctx); cpErr != nil {
t.Fatal(cpErr)
}
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")
if err != nil {
b.Fatal(err)
}
defer func() { _ = db.Close() }()
ctx := context.Background()
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
key := fmt.Sprintf("bench_%d", i)
if err := 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`,
key, "1", time.Now().UnixMilli(),
)
return e
}); err != nil {
b.Fatal(err)
}
}
}
// openSmallQueue 用小缓冲队列替换默认队列,便于测满通道时的 Close/Drain。
func openSmallQueue(t *testing.T, buf int) *DB {
t.Helper()
dir := t.TempDir()
db, err := Open(dir, "FULL")
if err != nil {
t.Fatal(err)
}
if err := db.Queue.Close(); err != nil {
t.Fatal(err)
}
db.Queue = newQueue(db.Write, buf)
return db
}
func TestQueueCloseUnblocksFullChannel(t *testing.T) {
t.Parallel()
const buf = 4
db := openSmallQueue(t, buf)
defer func() { _ = db.Close() }()
q := db.Queue
hold := make(chan struct{})
blockerStarted := make(chan struct{})
blockerErr := make(chan error, 1)
go func() {
blockerErr <- q.Do(context.Background(), func(tx *sql.Tx) error {
close(blockerStarted)
<-hold
return nil
})
}()
<-blockerStarted
ctx := context.Background()
var fillWG sync.WaitGroup
for i := 0; i < buf; i++ {
fillWG.Add(1)
go func() {
defer fillWG.Done()
_ = q.Do(ctx, func(tx *sql.Tx) error { return nil })
}()
}
deadline := time.Now().Add(2 * time.Second)
for q.Len() < buf+1 && time.Now().Before(deadline) {
time.Sleep(2 * time.Millisecond)
}
if q.Len() < buf+1 {
t.Fatalf("channel not full: len=%d", q.Len())
}
blockedErr := make(chan error, 1)
go func() {
blockedErr <- q.Do(ctx, func(tx *sql.Tx) error { return nil })
}()
// 等额外 Do 堵在 enqueue(pending 超过通道容量+正在执行的一条)。
deadline = time.Now().Add(2 * time.Second)
for q.Len() < buf+2 && time.Now().Before(deadline) {
time.Sleep(2 * time.Millisecond)
}
closeDone := make(chan error, 1)
go func() { closeDone <- q.Close() }()
select {
case err := <-blockedErr:
if !errors.Is(err, ErrQueueClosed) {
t.Fatalf("blocked Do: %v", err)
}
case <-time.After(2 * time.Second):
t.Fatal("Close did not unblock full-channel enqueue")
}
close(hold)
select {
case err := <-closeDone:
if err != nil {
t.Fatal(err)
}
case <-time.After(2 * time.Second):
t.Fatal("Close hung after writer released")
}
<-blockerErr
fillWG.Wait()
}
func TestQueueDrainTimeoutWhileWriterBlocked(t *testing.T) {
t.Parallel()
const buf = 4
db := openSmallQueue(t, buf)
defer func() { _ = db.Close() }()
q := db.Queue
hold := make(chan struct{})
started := make(chan struct{})
go func() {
_ = q.Do(context.Background(), func(tx *sql.Tx) error {
close(started)
<-hold
return nil
})
}()
<-started
ctx := context.Background()
var wg sync.WaitGroup
for i := 0; i < buf; i++ {
wg.Add(1)
go func() {
defer wg.Done()
_ = q.Do(ctx, func(tx *sql.Tx) error { return nil })
}()
}
deadline := time.Now().Add(2 * time.Second)
for q.Len() < buf+1 && time.Now().Before(deadline) {
time.Sleep(2 * time.Millisecond)
}
drainCtx, cancel := context.WithTimeout(context.Background(), 80*time.Millisecond)
defer cancel()
start := time.Now()
err := q.Drain(drainCtx)
elapsed := time.Since(start)
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("Drain err=%v want deadline", err)
}
if elapsed > 500*time.Millisecond {
t.Fatalf("Drain took %s, should return on timeout", elapsed)
}
close(hold)
wg.Wait()
}