297 lines
5.8 KiB
Go
297 lines
5.8 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
|
|
}
|
|
|
|
// 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) {
|
|
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
|
|
}
|
|
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
|
|
}
|