feat: 实现写入队列合并事务与 SAVEPOINT 隔离
This commit is contained in:
+201
-34
@@ -4,7 +4,9 @@ import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -14,25 +16,45 @@ var ErrQueueClosed = errors.New("store: write queue closed")
|
||||
// ErrBusy 表示写库基础设施失败(可映射为协议 busy;/readyz 应失败)。
|
||||
var ErrBusy = errors.New("store: busy")
|
||||
|
||||
const (
|
||||
maxBatchOps = 256
|
||||
batchWait = 2 * time.Millisecond
|
||||
queueBuffSize = 1024
|
||||
)
|
||||
|
||||
// WriteFunc 在单个写事务中执行的操作。
|
||||
type WriteFunc func(tx *sql.Tx) error
|
||||
|
||||
// Queue 写入队列:提交一个写操作并拿到结果。
|
||||
//
|
||||
// P1 仍为一操作一事务;P2 换成合并提交(最多 256 / 2ms + SAVEPOINT)。
|
||||
type writeJob struct {
|
||||
ctx context.Context
|
||||
fn WriteFunc
|
||||
res chan error
|
||||
}
|
||||
|
||||
// Queue 写入队列:写 goroutine 合并提交(最多 256 个或凑满 2ms),每操作用 SAVEPOINT 隔离。
|
||||
type Queue struct {
|
||||
db *sql.DB
|
||||
|
||||
ch chan writeJob
|
||||
done chan struct{}
|
||||
closed atomic.Bool
|
||||
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
inflight int
|
||||
ready bool
|
||||
lastWriteErr error
|
||||
pending int
|
||||
}
|
||||
|
||||
// NewQueue 创建简单写入队列(一操作一事务)。
|
||||
// NewQueue 创建合并写入队列并启动写 goroutine。
|
||||
func NewQueue(db *sql.DB) *Queue {
|
||||
return &Queue{db: db, ready: true}
|
||||
q := &Queue{
|
||||
db: db,
|
||||
ch: make(chan writeJob, queueBuffSize),
|
||||
done: make(chan struct{}),
|
||||
ready: true,
|
||||
}
|
||||
go q.loop()
|
||||
return q
|
||||
}
|
||||
|
||||
// Do 提交写操作并等待提交结果。
|
||||
@@ -40,37 +62,180 @@ func (q *Queue) Do(ctx context.Context, fn WriteFunc) error {
|
||||
if fn == nil {
|
||||
return errors.New("store: nil write func")
|
||||
}
|
||||
q.mu.Lock()
|
||||
if q.closed {
|
||||
q.mu.Unlock()
|
||||
if q.closed.Load() {
|
||||
return ErrQueueClosed
|
||||
}
|
||||
q.inflight++
|
||||
q.mu.Unlock()
|
||||
|
||||
defer func() {
|
||||
q.mu.Lock()
|
||||
q.inflight--
|
||||
q.mu.Unlock()
|
||||
}()
|
||||
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
tx, err := q.db.BeginTx(ctx, nil)
|
||||
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()
|
||||
return ErrQueueClosed
|
||||
}
|
||||
select {
|
||||
case err := <-job.res:
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
// 操作可能仍在队列中执行;结果通道仍会被写端关闭式填入。
|
||||
select {
|
||||
case err := <-job.res:
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return ctx.Err()
|
||||
case <-time.After(30 * time.Second):
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (q *Queue) loop() {
|
||||
defer close(q.done)
|
||||
for {
|
||||
job, ok := <-q.ch
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
batch := []writeJob{job}
|
||||
timer := time.NewTimer(batchWait)
|
||||
collect:
|
||||
for len(batch) < maxBatchOps {
|
||||
select {
|
||||
case j, ok := <-q.ch:
|
||||
if !ok {
|
||||
break collect
|
||||
}
|
||||
batch = append(batch, j)
|
||||
case <-timer.C:
|
||||
break collect
|
||||
}
|
||||
}
|
||||
timer.Stop()
|
||||
q.runBatch(batch)
|
||||
}
|
||||
}
|
||||
|
||||
func (q *Queue) runBatch(batch []writeJob) {
|
||||
defer func() {
|
||||
q.mu.Lock()
|
||||
q.pending -= len(batch)
|
||||
if q.pending < 0 {
|
||||
q.pending = 0
|
||||
}
|
||||
q.mu.Unlock()
|
||||
}()
|
||||
|
||||
// 过滤已取消的任务。
|
||||
active := make([]writeJob, 0, len(batch))
|
||||
for _, j := range batch {
|
||||
if err := j.ctx.Err(); err != nil {
|
||||
j.res <- err
|
||||
continue
|
||||
}
|
||||
active = append(active, j)
|
||||
}
|
||||
if len(active) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
tx, err := q.db.BeginTx(context.Background(), nil)
|
||||
if err != nil {
|
||||
q.markBusy(err)
|
||||
return errors.Join(ErrBusy, err)
|
||||
for _, j := range active {
|
||||
j.res <- errors.Join(ErrBusy, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err := fn(tx); err != nil {
|
||||
_ = tx.Rollback()
|
||||
return err
|
||||
|
||||
type outcome struct {
|
||||
job writeJob
|
||||
opErr error
|
||||
success bool
|
||||
}
|
||||
outcomes := make([]outcome, 0, len(active))
|
||||
for i, j := range active {
|
||||
sp := fmt.Sprintf("sp_%d", i)
|
||||
if _, err := tx.Exec("SAVEPOINT " + sp); err != nil {
|
||||
_ = tx.Rollback()
|
||||
q.markBusy(err)
|
||||
for _, o := range outcomes {
|
||||
if o.success {
|
||||
o.job.res <- errors.Join(ErrBusy, err)
|
||||
}
|
||||
}
|
||||
for _, rest := range active[i:] {
|
||||
rest.res <- errors.Join(ErrBusy, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
opErr := j.fn(tx)
|
||||
if opErr != nil {
|
||||
if _, rbErr := tx.Exec("ROLLBACK TO " + sp); rbErr != nil {
|
||||
_ = tx.Rollback()
|
||||
q.markBusy(rbErr)
|
||||
for _, o := range outcomes {
|
||||
if o.success {
|
||||
o.job.res <- errors.Join(ErrBusy, rbErr)
|
||||
}
|
||||
}
|
||||
j.res <- opErr
|
||||
for _, rest := range active[i+1:] {
|
||||
rest.res <- errors.Join(ErrBusy, rbErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
_, _ = tx.Exec("RELEASE " + sp)
|
||||
outcomes = append(outcomes, outcome{job: j, opErr: opErr, success: false})
|
||||
continue
|
||||
}
|
||||
if _, err := tx.Exec("RELEASE " + sp); err != nil {
|
||||
_ = tx.Rollback()
|
||||
q.markBusy(err)
|
||||
for _, o := range outcomes {
|
||||
if o.success {
|
||||
o.job.res <- errors.Join(ErrBusy, err)
|
||||
}
|
||||
}
|
||||
j.res <- errors.Join(ErrBusy, err)
|
||||
for _, rest := range active[i+1:] {
|
||||
rest.res <- errors.Join(ErrBusy, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
outcomes = append(outcomes, outcome{job: j, success: true})
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
q.markBusy(err)
|
||||
return errors.Join(ErrBusy, err)
|
||||
for _, o := range outcomes {
|
||||
if o.success {
|
||||
o.job.res <- errors.Join(ErrBusy, err)
|
||||
} else {
|
||||
o.job.res <- o.opErr
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
for _, o := range outcomes {
|
||||
if o.success {
|
||||
o.job.res <- nil
|
||||
} else {
|
||||
o.job.res <- o.opErr
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (q *Queue) markBusy(err error) {
|
||||
@@ -94,20 +259,20 @@ func (q *Queue) LastWriteError() error {
|
||||
return q.lastWriteErr
|
||||
}
|
||||
|
||||
// Len 返回进行中的写操作数。
|
||||
// Len 返回排队与批次中的写操作数。
|
||||
func (q *Queue) Len() int {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
return q.inflight
|
||||
return q.pending
|
||||
}
|
||||
|
||||
// Drain 等待已进行中的写操作完成,或 ctx 取消。
|
||||
// Drain 等待已入队操作完成,或 ctx 取消。关闭后也会等到 loop 退出。
|
||||
func (q *Queue) Drain(ctx context.Context) error {
|
||||
ticker := time.NewTicker(5 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
q.mu.Lock()
|
||||
n := q.inflight
|
||||
n := q.pending
|
||||
q.mu.Unlock()
|
||||
if n == 0 {
|
||||
return nil
|
||||
@@ -120,10 +285,12 @@ func (q *Queue) Drain(ctx context.Context) error {
|
||||
}
|
||||
}
|
||||
|
||||
// Close 关闭队列,之后 Do 返回 ErrQueueClosed。
|
||||
// Close 关闭队列:不再接受新任务,并等待写 goroutine 处理完已入队任务后退出。
|
||||
func (q *Queue) Close() error {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
q.closed = true
|
||||
if q.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
close(q.ch)
|
||||
<-q.done
|
||||
return nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user