Spanner: Add trivial sequence number generator. (#101777)

* Add trivial sequence number generator.
This commit is contained in:
Peter Štibraný
2025-03-10 17:37:44 +01:00
committed by GitHub
parent 8142aef64d
commit 9858e40a02
5 changed files with 128 additions and 15 deletions
+5 -1
View File
@@ -42,7 +42,8 @@ type Engine struct {
tagHandlers map[string]tagHandler
defaultContext context.Context
defaultContext context.Context
sequenceGenerator *sequenceGenerator // If not nil, this generator is used to generate auto-increment values for inserts.
}
// CondDeleted returns the conditions whether a record is soft deleted.
@@ -237,6 +238,9 @@ func (engine *Engine) NewSession() *Session {
// Close the engine
func (engine *Engine) Close() error {
if engine.sequenceGenerator != nil {
engine.sequenceGenerator.close()
}
return engine.db.Close()
}
+69
View File
@@ -0,0 +1,69 @@
package xorm
import (
"context"
"database/sql"
"errors"
"fmt"
)
type sequenceGenerator struct {
db *sql.DB
sequencesTable string
}
func newSequenceGenerator(db *sql.DB) *sequenceGenerator {
return &sequenceGenerator{
db: db,
sequencesTable: "autoincrement_sequences",
}
}
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)
tx, err := sg.db.BeginTx(ctx, &sql.TxOptions{Isolation: sql.LevelSerializable})
if err != nil {
return 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())
}
val := int64(1)
_, err := tx.ExecContext(ctx, "INSERT INTO "+sg.sequencesTable+" (name, next_value) VALUES(?, ?)", key, val+1)
if err != nil {
return 0, errors.Join(err, tx.Rollback())
}
return val, tx.Commit()
}
var val int64
if err := r.Scan(&val); err != nil {
return 0, errors.Join(err, tx.Rollback())
}
_, err = tx.ExecContext(ctx, "UPDATE "+sg.sequencesTable+" SET next_value = ? WHERE name = ?", val+1, key)
if err != nil {
return 0, errors.Join(err, tx.Rollback())
}
return val, tx.Commit()
}
func (sg *sequenceGenerator) close() {
// Nothing to do just yet.
}
+31
View File
@@ -0,0 +1,31 @@
package xorm
import (
"context"
"testing"
"github.com/stretchr/testify/require"
)
func TestSequenceGenerator(t *testing.T) {
eng, err := NewEngine("sqlite3", ":memory:")
require.NoError(t, err)
require.NotNil(t, eng)
require.Equal(t, "sqlite3", eng.DriverName())
_, 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)
val, err := sg.Next(context.Background(), "test", "test")
require.NoError(t, err)
require.Equal(t, int64(1), val)
val, err = sg.Next(context.Background(), "test", "different")
require.NoError(t, err)
require.Equal(t, int64(1), val)
val, err = sg.Next(context.Background(), "test", "different")
require.NoError(t, err)
require.Equal(t, int64(2), val)
}
+19 -14
View File
@@ -345,20 +345,25 @@ func (session *Session) innerInsert(bean any) (int64, error) {
return 0, err
}
//// XXX: hack to handle autoincrement in spanner
//if len(table.AutoIncrement) > 0 && session.engine.dialect.DBType() == "spanner" {
// var found bool
// for _, col := range colNames {
// if col == table.AutoIncrement {
// found = true
// break
// }
// }
// if !found {
// colNames = append(colNames, table.AutoIncrement)
// args = append(args, rand.Int63n(9e15))
// }
//}
// If engine has a sequence number generator, use it to produce values for auto-increment columns.
if len(table.AutoIncrement) > 0 && session.engine.sequenceGenerator != nil {
var found bool
for _, col := range colNames {
if col == table.AutoIncrement {
found = true
break
}
}
if !found {
seq, err := session.engine.sequenceGenerator.Next(session.ctx, table.Name, table.AutoIncrement)
if err != nil {
return 0, fmt.Errorf("failed to generate next value for auto_increment columns: %v", err)
}
colNames = append(colNames, table.AutoIncrement)
args = append(args, seq)
}
}
exprs := session.statement.exprColumns
colPlaces := strings.Repeat("?, ", len(colNames))
+4
View File
@@ -106,5 +106,9 @@ func NewEngine(driverName string, dataSourceName string) (*Engine, error) {
runtime.SetFinalizer(engine, close)
if dialect.DBType() == "spanner" {
engine.sequenceGenerator = newSequenceGenerator(db.DB)
}
return engine, nil
}