426 lines
9.5 KiB
Go
426 lines
9.5 KiB
Go
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()
|
||
}
|