diff --git a/pkg/util/xorm/sequence.go b/pkg/util/xorm/sequence.go index 4e11c0fbab1..19192db96f3 100644 --- a/pkg/util/xorm/sequence.go +++ b/pkg/util/xorm/sequence.go @@ -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() { diff --git a/pkg/util/xorm/sequence_test.go b/pkg/util/xorm/sequence_test.go index 289e65e4034..b49bb568c8b 100644 --- a/pkg/util/xorm/sequence_test.go +++ b/pkg/util/xorm/sequence_test.go @@ -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) + } +}