Files

304 lines
6.0 KiB
Go

package store
import (
"context"
"database/sql"
"errors"
"fmt"
"sync"
"sync/atomic"
"time"
)
// ErrQueueClosed 表示写入队列已关闭。
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
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
ready bool
lastWriteErr error
pending int
// OnBatchCommit 可选;每次合并提交成功后回调耗时(秒级指标用)。
OnBatchCommit func(d time.Duration)
}
// NewQueue 创建合并写入队列并启动写 goroutine。
func NewQueue(db *sql.DB) *Queue {
q := &Queue{
db: db,
ch: make(chan writeJob, queueBuffSize),
done: make(chan struct{}),
ready: true,
}
go q.loop()
return q
}
// Do 提交写操作并等待提交结果。
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()
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) {
started := time.Now()
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)
for _, j := range active {
j.res <- errors.Join(ErrBusy, err)
}
return
}
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)
for _, o := range outcomes {
if o.success {
o.job.res <- errors.Join(ErrBusy, err)
} else {
o.job.res <- o.opErr
}
}
return
}
if q.OnBatchCommit != nil {
q.OnBatchCommit(time.Since(started))
}
for _, o := range outcomes {
if o.success {
o.job.res <- nil
} else {
o.job.res <- o.opErr
}
}
}
func (q *Queue) markBusy(err error) {
q.mu.Lock()
defer q.mu.Unlock()
q.ready = false
q.lastWriteErr = err
}
// IsReady 写库是否仍可用(写失败后为 false)。
func (q *Queue) IsReady() bool {
q.mu.Lock()
defer q.mu.Unlock()
return q.ready
}
// LastWriteError 返回最近一次基础设施写失败。
func (q *Queue) LastWriteError() error {
q.mu.Lock()
defer q.mu.Unlock()
return q.lastWriteErr
}
// Len 返回排队与批次中的写操作数。
func (q *Queue) Len() int {
q.mu.Lock()
defer q.mu.Unlock()
return q.pending
}
// 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.pending
q.mu.Unlock()
if n == 0 {
return nil
}
select {
case <-ctx.Done():
return ctx.Err()
case <-ticker.C:
}
}
}
// Close 关闭队列:不再接受新任务,并等待写 goroutine 处理完已入队任务后退出。
func (q *Queue) Close() error {
if q.closed.Swap(true) {
return nil
}
close(q.ch)
<-q.done
return nil
}