diff --git a/pkg/util/xorm/engine.go b/pkg/util/xorm/engine.go index 88135ef7a7c..c64b7b29f4f 100644 --- a/pkg/util/xorm/engine.go +++ b/pkg/util/xorm/engine.go @@ -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() } diff --git a/pkg/util/xorm/sequence.go b/pkg/util/xorm/sequence.go new file mode 100644 index 00000000000..bb74e2064b6 --- /dev/null +++ b/pkg/util/xorm/sequence.go @@ -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. +} diff --git a/pkg/util/xorm/sequence_test.go b/pkg/util/xorm/sequence_test.go new file mode 100644 index 00000000000..289e65e4034 --- /dev/null +++ b/pkg/util/xorm/sequence_test.go @@ -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) +} diff --git a/pkg/util/xorm/session_insert.go b/pkg/util/xorm/session_insert.go index 502a55e769e..09016f17255 100644 --- a/pkg/util/xorm/session_insert.go +++ b/pkg/util/xorm/session_insert.go @@ -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)) diff --git a/pkg/util/xorm/xorm.go b/pkg/util/xorm/xorm.go index b3178773518..40c1552cec5 100644 --- a/pkg/util/xorm/xorm.go +++ b/pkg/util/xorm/xorm.go @@ -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 }