Spanner: Add trivial sequence number generator. (#101777)
* Add trivial sequence number generator.
This commit is contained in:
@@ -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()
|
||||
}
|
||||
|
||||
|
||||
@@ -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.
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user