fix: 写队列满时 Close 不再与 enqueue 死锁
This commit is contained in:
@@ -289,3 +289,137 @@ ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.upd
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 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()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user