spanner: update the sequence generator to allocate sequences in batch (#102435)
* spanner: update the sequence generator to allocate sequences in batch * lock per sequence * handle error scenario * rollback on error * mutex-hat * implement sequent generator
This commit is contained in:
+110
-26
@@ -3,69 +3,153 @@ package xorm
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// batchState represents the state of a sequence batch
|
||||
type batchState struct {
|
||||
mu sync.Mutex
|
||||
nextValue int64
|
||||
lastValueInBatch int64
|
||||
}
|
||||
|
||||
type sequenceGenerator struct {
|
||||
db *sql.DB
|
||||
sequencesTable string
|
||||
batchSize int64
|
||||
|
||||
mu sync.Mutex
|
||||
batchStates map[string]*batchState // Track sequence batches per key (table:column)
|
||||
}
|
||||
|
||||
func newSequenceGenerator(db *sql.DB) *sequenceGenerator {
|
||||
return &sequenceGenerator{
|
||||
db: db,
|
||||
sequencesTable: "autoincrement_sequences",
|
||||
batchSize: 100, // Default batch size
|
||||
batchStates: make(map[string]*batchState),
|
||||
}
|
||||
}
|
||||
|
||||
func (sg *sequenceGenerator) Reset() {
|
||||
// Nothing to do. This generator always uses state from DB.
|
||||
sg.mu.Lock()
|
||||
defer sg.mu.Unlock()
|
||||
sg.batchStates = make(map[string]*batchState)
|
||||
}
|
||||
|
||||
func (sg *sequenceGenerator) Next(ctx context.Context, table, column string) (int64, error) {
|
||||
// Current implementation fetches new value for each Next call.
|
||||
key := fmt.Sprintf("%s:%s", table, column)
|
||||
|
||||
// First get or create the state with a global lock (only for map access)
|
||||
sg.mu.Lock()
|
||||
state, exists := sg.batchStates[key]
|
||||
if !exists {
|
||||
state = &batchState{
|
||||
nextValue: 0,
|
||||
lastValueInBatch: -1,
|
||||
}
|
||||
sg.batchStates[key] = state
|
||||
}
|
||||
sg.mu.Unlock() // Release global lock as soon as possible
|
||||
|
||||
// Now lock only the specific sequence state
|
||||
state.mu.Lock()
|
||||
defer state.mu.Unlock()
|
||||
|
||||
// If we've used all values in the current batch, get a new batch
|
||||
if state.nextValue > state.lastValueInBatch {
|
||||
start, end, err := sg.allocateNewBatch(ctx, key)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
state.nextValue = start
|
||||
state.lastValueInBatch = end
|
||||
}
|
||||
|
||||
// Return the next value from the batch
|
||||
val := state.nextValue
|
||||
state.nextValue++
|
||||
return val, nil
|
||||
}
|
||||
|
||||
// allocateNewBatch retrieves a new batch of sequence values from the database.
|
||||
// It returns the start and end values of the new batch on success.
|
||||
func (sg *sequenceGenerator) allocateNewBatch(ctx context.Context, key string) (start, end int64, retErr error) {
|
||||
tx, err := sg.db.BeginTx(ctx, &sql.TxOptions{Isolation: sql.LevelSerializable})
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return 0, 0, err
|
||||
}
|
||||
|
||||
// TODO: "FOR UPDATE". (Somehow this doesn't seem to be supported in Spanner emulator?)
|
||||
r, err := tx.QueryContext(ctx, "SELECT next_value FROM "+sg.sequencesTable+" WHERE name = ?", key)
|
||||
if err != nil {
|
||||
err2 := tx.Rollback()
|
||||
return 0, errors.Join(err, err2)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
// Sequence doesn't exist yet. Return 1, and put 2 into the table.
|
||||
if !r.Next() {
|
||||
if err := r.Err(); err != nil {
|
||||
return 0, errors.Join(err, tx.Rollback())
|
||||
defer func() {
|
||||
if retErr != nil {
|
||||
tx.Rollback()
|
||||
}
|
||||
val := int64(1)
|
||||
}()
|
||||
|
||||
_, err := tx.ExecContext(ctx, "INSERT INTO "+sg.sequencesTable+" (name, next_value) VALUES(?, ?)", key, val+1)
|
||||
// Query the current sequence value
|
||||
rows, err := tx.QueryContext(ctx, "SELECT next_value FROM "+sg.sequencesTable+" WHERE name = ?", key)
|
||||
if err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
// Handle case where sequence doesn't exist yet
|
||||
if !rows.Next() {
|
||||
if err = rows.Err(); err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
|
||||
// This is a new sequence - start from 1 and allocate a batch
|
||||
batchEnd := sg.batchSize
|
||||
nextBatchStart := batchEnd + 1
|
||||
|
||||
// Insert the next batch start value
|
||||
_, err = tx.ExecContext(ctx,
|
||||
"INSERT INTO "+sg.sequencesTable+" (name, next_value) VALUES(?, ?)",
|
||||
key, nextBatchStart)
|
||||
if err != nil {
|
||||
return 0, errors.Join(err, tx.Rollback())
|
||||
return 0, 0, err
|
||||
}
|
||||
|
||||
return val, tx.Commit()
|
||||
// Commit the transaction
|
||||
if err = tx.Commit(); err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
|
||||
return 1, batchEnd, nil
|
||||
}
|
||||
|
||||
var val int64
|
||||
if err := r.Scan(&val); err != nil {
|
||||
return 0, errors.Join(err, tx.Rollback())
|
||||
// Sequence exists - read current value and allocate next batch
|
||||
var batchStart int64
|
||||
if err = rows.Scan(&batchStart); err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
|
||||
_, err = tx.ExecContext(ctx, "UPDATE "+sg.sequencesTable+" SET next_value = ? WHERE name = ?", val+1, key)
|
||||
batchEnd := batchStart + sg.batchSize - 1
|
||||
nextBatchStart := batchEnd + 1
|
||||
|
||||
// Update the next batch start value
|
||||
_, err = tx.ExecContext(ctx,
|
||||
"UPDATE "+sg.sequencesTable+" SET next_value = ? WHERE name = ?",
|
||||
nextBatchStart, key)
|
||||
if err != nil {
|
||||
return 0, errors.Join(err, tx.Rollback())
|
||||
return 0, 0, err
|
||||
}
|
||||
|
||||
return val, tx.Commit()
|
||||
// Commit the transaction
|
||||
if err = tx.Commit(); err != nil {
|
||||
return 0, 0, err
|
||||
}
|
||||
|
||||
return batchStart, batchEnd, nil
|
||||
}
|
||||
|
||||
// SetBatchSize allows changing the batch size
|
||||
func (sg *sequenceGenerator) SetBatchSize(size int64) {
|
||||
sg.mu.Lock()
|
||||
defer sg.mu.Unlock()
|
||||
sg.batchSize = size
|
||||
}
|
||||
|
||||
func (sg *sequenceGenerator) close() {
|
||||
|
||||
@@ -2,6 +2,7 @@ package xorm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -29,3 +30,97 @@ func TestSequenceGenerator(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(2), val)
|
||||
}
|
||||
|
||||
func TestBatchSequenceAllocation(t *testing.T) {
|
||||
eng, err := NewEngine("sqlite3", ":memory:")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = eng.Exec("CREATE TABLE `autoincrement_sequences` (`name` STRING(128) NOT NULL PRIMARY KEY, `next_value` INT64 NOT NULL)")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create sequence generator with small batch size for testing
|
||||
sg := newSequenceGenerator(eng.db.DB)
|
||||
sg.SetBatchSize(10)
|
||||
|
||||
// First batch (1-10)
|
||||
for i := 1; i <= 10; i++ {
|
||||
val, err := sg.Next(context.Background(), "test", "batch")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(i), val)
|
||||
}
|
||||
|
||||
// Next value should trigger a new batch (11-20)
|
||||
val, err := sg.Next(context.Background(), "test", "batch")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(11), val)
|
||||
|
||||
// Check database value is now set to next batch
|
||||
var nextVal int64
|
||||
err = eng.db.QueryRow("SELECT next_value FROM autoincrement_sequences WHERE name = 'test:batch'").Scan(&nextVal)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(21), nextVal, "Database should store the start of the next batch")
|
||||
|
||||
// Continue getting values from the second batch
|
||||
for i := 12; i <= 20; i++ {
|
||||
val, err := sg.Next(context.Background(), "test", "batch")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(i), val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentSequenceAccess(t *testing.T) {
|
||||
eng, err := NewEngine("sqlite3", ":memory:")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = eng.Exec("CREATE TABLE `autoincrement_sequences` (`name` STRING(128) NOT NULL PRIMARY KEY, `next_value` INT64 NOT NULL)")
|
||||
require.NoError(t, err)
|
||||
|
||||
sg := newSequenceGenerator(eng.db.DB)
|
||||
sg.SetBatchSize(100)
|
||||
|
||||
// Launch multiple goroutines to get sequence values
|
||||
const numRoutines = 50
|
||||
const valuesPerRoutine = 20
|
||||
results := make([]int64, numRoutines*valuesPerRoutine)
|
||||
|
||||
var wg sync.WaitGroup
|
||||
var mu sync.Mutex // To protect results slice
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
for i := 0; i < numRoutines; i++ {
|
||||
wg.Add(1)
|
||||
go func(routineID int) {
|
||||
defer wg.Done()
|
||||
|
||||
for j := 0; j < valuesPerRoutine; j++ {
|
||||
val, err := sg.Next(ctx, "test", "concurrent")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Store the result in our results array
|
||||
mu.Lock()
|
||||
results[routineID*valuesPerRoutine+j] = val
|
||||
mu.Unlock()
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
|
||||
// Check that we have the expected number of values
|
||||
require.Equal(t, numRoutines*valuesPerRoutine, len(results))
|
||||
|
||||
// Create a map to check for duplicates
|
||||
seen := make(map[int64]bool)
|
||||
for _, val := range results {
|
||||
// Verify we haven't seen this value before
|
||||
require.False(t, seen[val], "Found duplicate sequence value: %d", val)
|
||||
seen[val] = true
|
||||
}
|
||||
|
||||
// Verify the range of values is correct (all values should be between 1 and numRoutines*valuesPerRoutine)
|
||||
require.Equal(t, numRoutines*valuesPerRoutine, len(seen), "Should have exactly the right number of unique values")
|
||||
for i := int64(1); i <= int64(numRoutines*valuesPerRoutine); i++ {
|
||||
require.True(t, seen[i], "Missing sequence value: %d", i)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user