From a8eba901172e070d595b88a34b13b11655061445 Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Thu, 16 Jul 2026 14:09:14 -0700 Subject: [PATCH 01/11] Allow GlobalState to have multiple AutoIncrementTrackers --- go/go.mod | 4 +++ .../doltcore/sqle/cluster/controller.go | 6 +--- go/libraries/doltcore/sqle/database.go | 6 ++-- .../doltcore/sqle/dsess/globalstate.go | 33 ++++++++++++++----- go/libraries/doltcore/sqle/dsess/session.go | 9 ++--- .../globalstate/auto_increment_tracker.go | 1 + .../doltcore/sqle/globalstate/global_state.go | 9 +++-- go/libraries/doltcore/sqle/tables.go | 8 ++--- 8 files changed, 46 insertions(+), 30 deletions(-) diff --git a/go/go.mod b/go/go.mod index b6ace51f237..6211543a5e5 100644 --- a/go/go.mod +++ b/go/go.mod @@ -209,3 +209,7 @@ require ( ) go 1.26.2 + +replace ( + github.com/dolthub/go-mysql-server => ../../go-mysql-server +) diff --git a/go/libraries/doltcore/sqle/cluster/controller.go b/go/libraries/doltcore/sqle/cluster/controller.go index 218c354da0d..05502049cf8 100644 --- a/go/libraries/doltcore/sqle/cluster/controller.go +++ b/go/libraries/doltcore/sqle/cluster/controller.go @@ -976,10 +976,6 @@ func (c *Controller) refreshAutoIncrementTrackersForSessionDatabases() error { // Non-versioned DBs don't participate in AUTO_INCREMENT global state continue } - ai, err := gsp.GetGlobalState().AutoIncrementTracker(sqlCtx) - if err != nil { - return fmt.Errorf("cluster/controller: auto-inc refresh: %s: tracker: %w", name, err) - } // Get working set roots only state, ok, err := sess.LookupDbState(sqlCtx, name) @@ -990,7 +986,7 @@ func (c *Controller) refreshAutoIncrementTrackersForSessionDatabases() error { // Not loaded in session; defer to lazy initialization on first use continue } - if err := ai.InitWithRoots(sqlCtx, state.WorkingSet()); err != nil { + if err := gsp.GetGlobalState().InitWithRoots(sqlCtx, state.WorkingSet()); err != nil { return fmt.Errorf("cluster/controller: auto-inc refresh: %s: init: %w", name, err) } } diff --git a/go/libraries/doltcore/sqle/database.go b/go/libraries/doltcore/sqle/database.go index a287c78926d..6889565b259 100644 --- a/go/libraries/doltcore/sqle/database.go +++ b/go/libraries/doltcore/sqle/database.go @@ -1924,7 +1924,7 @@ func (db Database) removeTableFromAutoIncrementTracker( wses = append(wses, ws) } - ait, err := db.gs.AutoIncrementTracker(ctx) + ait, err := db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) if err != nil { return err } @@ -2084,7 +2084,7 @@ func (db Database) createSqlTable(ctx *sql.Context, table string, schemaName str // Prevent any tables that use BINARY, CHAR, VARBINARY, VARCHAR prefixes if schema.HasAutoIncrement(doltSch) { - ait, err := db.gs.AutoIncrementTracker(ctx) + ait, err := db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) if err != nil { return err } @@ -2144,7 +2144,7 @@ func (db Database) createIndexedSqlTable(ctx *sql.Context, table string, schemaN } if schema.HasAutoIncrement(doltSch) { - ait, err := db.gs.AutoIncrementTracker(ctx) + ait, err := db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) if err != nil { return err } diff --git a/go/libraries/doltcore/sqle/dsess/globalstate.go b/go/libraries/doltcore/sqle/dsess/globalstate.go index b28c8e6e666..5482dc27b05 100644 --- a/go/libraries/doltcore/sqle/dsess/globalstate.go +++ b/go/libraries/doltcore/sqle/dsess/globalstate.go @@ -16,8 +16,6 @@ package dsess import ( "context" - "sync" - "github.com/dolthub/go-mysql-server/sql" "golang.org/x/sync/errgroup" @@ -26,6 +24,8 @@ import ( "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate" ) +var AutoIncrementTrackerKey = struct{}{} + func NewGlobalStateStoreForDb(ctx context.Context, dbName string, db *doltdb.DoltDB) (GlobalStateImpl, error) { branches, err := db.GetBranches(ctx) if err != nil { @@ -94,22 +94,37 @@ func NewGlobalStateStoreForDb(ctx context.Context, dbName string, db *doltdb.Dol } return GlobalStateImpl{ - aiTracker: tracker, - mu: &sync.Mutex{}, + aiTrackers: map[interface{}]globalstate.AutoIncrementTracker{AutoIncrementTrackerKey: tracker}, }, nil } type GlobalStateImpl struct { - aiTracker *AutoIncrementTracker - mu *sync.Mutex + aiTrackers map[interface{}]globalstate.AutoIncrementTracker } var _ globalstate.GlobalState = GlobalStateImpl{} -func (g GlobalStateImpl) AutoIncrementTracker(ctx *sql.Context) (globalstate.AutoIncrementTracker, error) { - return g.aiTracker, nil +func (g GlobalStateImpl) AutoIncrementTracker(ctx *sql.Context, key interface{}) (globalstate.AutoIncrementTracker, error) { + return g.aiTrackers[key], nil +} + +func (g GlobalStateImpl) AddAutoIncrementTracker(ctx *sql.Context, key interface{}, tracker globalstate.AutoIncrementTracker) error { + g.aiTrackers[key] = tracker + return nil } func (g GlobalStateImpl) Close() { - g.aiTracker.Close() + for _, tracker := range g.aiTrackers { + tracker.Close() + } +} + +func (g GlobalStateImpl) InitWithRoots(ctx *sql.Context, roots ...doltdb.Rootish) error { + for _, tracker := range g.aiTrackers { + err := tracker.InitWithRoots(ctx, roots...) + if err != nil { + return err + } + } + return nil } diff --git a/go/libraries/doltcore/sqle/dsess/session.go b/go/libraries/doltcore/sqle/dsess/session.go index 5d4894ac4b8..e25483ae1a9 100644 --- a/go/libraries/doltcore/sqle/dsess/session.go +++ b/go/libraries/doltcore/sqle/dsess/session.go @@ -1144,12 +1144,7 @@ func (d *DoltSession) ResetGlobals(ctx *sql.Context, dbName string, root doltdb. return err } - tracker, err := sessionState.dbState.globalState.AutoIncrementTracker(ctx) - if err != nil { - return err - } - - err = tracker.InitWithRoots(ctx, root) + err = sessionState.dbState.globalState.InitWithRoots(ctx, root) if err != nil { return err } @@ -1412,7 +1407,7 @@ func (d *DoltSession) addDB(ctx *sql.Context, db SqlDatabase) error { } sessionState.globalState = stateProvider.GetGlobalState() - tracker, err := sessionState.globalState.AutoIncrementTracker(ctx) + tracker, err := sessionState.globalState.AutoIncrementTracker(ctx, AutoIncrementTrackerKey) if err != nil { return err } diff --git a/go/libraries/doltcore/sqle/globalstate/auto_increment_tracker.go b/go/libraries/doltcore/sqle/globalstate/auto_increment_tracker.go index f8aed7a5498..e9ff60631c2 100644 --- a/go/libraries/doltcore/sqle/globalstate/auto_increment_tracker.go +++ b/go/libraries/doltcore/sqle/globalstate/auto_increment_tracker.go @@ -46,4 +46,5 @@ type AutoIncrementTracker interface { AcquireTableLock(ctx *sql.Context, tableName string) (func(), error) // InitWithRoots fills the AutoIncrementTracker with values pulled from each root in order. InitWithRoots(ctx context.Context, roots ...doltdb.Rootish) error + Close() } diff --git a/go/libraries/doltcore/sqle/globalstate/global_state.go b/go/libraries/doltcore/sqle/globalstate/global_state.go index 720e6e901b0..49a5dd37e86 100644 --- a/go/libraries/doltcore/sqle/globalstate/global_state.go +++ b/go/libraries/doltcore/sqle/globalstate/global_state.go @@ -14,13 +14,18 @@ package globalstate -import "github.com/dolthub/go-mysql-server/sql" +import ( + "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" + "github.com/dolthub/go-mysql-server/sql" +) // GlobalState is just a holding interface for pieces of global state, of which the auto increment tracking info is // the only example at the moment. type GlobalState interface { // AutoIncrementTracker returns the auto increment tracker for this global state. - AutoIncrementTracker(ctx *sql.Context) (AutoIncrementTracker, error) + AutoIncrementTracker(ctx *sql.Context, key interface{}) (AutoIncrementTracker, error) + AddAutoIncrementTracker(ctx *sql.Context, key interface{}, value AutoIncrementTracker) error + InitWithRoots(ctx *sql.Context, roots ...doltdb.Rootish) error } // GlobalStateProvider is an optional interface for databases that provide global state tracking diff --git a/go/libraries/doltcore/sqle/tables.go b/go/libraries/doltcore/sqle/tables.go index d4d0902d62a..e705b7fe254 100644 --- a/go/libraries/doltcore/sqle/tables.go +++ b/go/libraries/doltcore/sqle/tables.go @@ -1570,7 +1570,7 @@ func (t *AlterableDoltTable) AddColumn(ctx *sql.Context, column *sql.Column, ord } if column.AutoIncrement { - ait, err := t.db.gs.AutoIncrementTracker(ctx) + ait, err := t.db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) if err != nil { return err } @@ -1797,7 +1797,7 @@ func (t *AlterableDoltTable) RewriteInserter( // TODO: figure out locking. Other DBs automatically lock a table during this kind of operation, we should probably // do the same. We're messing with global auto-increment values here and it's not safe. - ait, err := t.db.gs.AutoIncrementTracker(ctx) + ait, err := t.db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) if err != nil { return nil, err } @@ -1854,7 +1854,7 @@ func fullTextRewriteEditor( // TODO: figure out locking. Other DBs automatically lock a table during this kind of operation, we should probably // do the same. We're messing with global auto-increment values here and it's not safe. - ait, err := t.db.gs.AutoIncrementTracker(ctx) + ait, err := t.db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) if err != nil { return nil, err } @@ -2276,7 +2276,7 @@ func (t *AlterableDoltTable) ModifyColumn(ctx *sql.Context, columnName string, c return err } - ait, err := t.db.gs.AutoIncrementTracker(ctx) + ait, err := t.db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) if err != nil { return err } From debe9fe1b854d89bad4e37d758a331f67724ff1d Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Mon, 20 Jul 2026 20:58:24 -0700 Subject: [PATCH 02/11] Add type-safe synchronized Map implementation. --- go/libraries/doltcore/sqle/dsess/sync_map.go | 51 ++++++++++++++++++++ 1 file changed, 51 insertions(+) create mode 100644 go/libraries/doltcore/sqle/dsess/sync_map.go diff --git a/go/libraries/doltcore/sqle/dsess/sync_map.go b/go/libraries/doltcore/sqle/dsess/sync_map.go new file mode 100644 index 00000000000..791c0fb151b --- /dev/null +++ b/go/libraries/doltcore/sqle/dsess/sync_map.go @@ -0,0 +1,51 @@ +// Copyright 2026 Dolthub, Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package dsess + +import "sync" + +// SyncMap is a simple generic wrapper around sync.Map, designed for type safety around callsites. +// Using this type instead of sync.Map removes the need for casts. +type SyncMap[Key any, Value any] struct { + m sync.Map +} + +func (a *SyncMap[Key, Value]) Store(key Key, val Value) { + a.m.Store(key, val) +} + +func (a *SyncMap[Key, Value]) Load(key Key) (val Value, ok bool) { + v, ok := a.m.Load(key) + if !ok { + return val, ok + } + return v.(Value), ok +} + +func (a *SyncMap[Key, Value]) Delete(key Key) { + a.m.Delete(key) +} + +func (a *SyncMap[Key, Value]) LoadOrStore(key Key, val Value) (actual Value, loaded bool) { + act, loaded := a.m.LoadOrStore(key, val) + if !loaded { + return actual, loaded + } + return act.(Value), loaded +} + +func (a *SyncMap[Key, Value]) CompareAndSwap(key Key, old Value, new Value) (swapped bool) { + return a.m.CompareAndSwap(key, old, new) +} From 754c6bf88377c174cc9007135af6d3a8aac6dfc3 Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Mon, 20 Jul 2026 20:59:30 -0700 Subject: [PATCH 03/11] Add globalstate/sequences package of interfaces --- .../sqle/globalstate/sequences/state.go | 66 +++++++++++++++++++ 1 file changed, 66 insertions(+) create mode 100644 go/libraries/doltcore/sqle/globalstate/sequences/state.go diff --git a/go/libraries/doltcore/sqle/globalstate/sequences/state.go b/go/libraries/doltcore/sqle/globalstate/sequences/state.go new file mode 100644 index 00000000000..fe04b991acd --- /dev/null +++ b/go/libraries/doltcore/sqle/globalstate/sequences/state.go @@ -0,0 +1,66 @@ +// Copyright 2026 Dolthub, Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sequences + +import ( + "context" + "github.com/dolthub/go-mysql-server/sql" +) + +// This package is a collection of interfaces that describe state machines that are controlled by the global state. +// It is a separate package from the parent package globalstate so that doltdb can depend on it. + +// A SequenceState is an incrementing state that must be shared across all branches and transactions. +// It corresponds to a table or root object in the database. +// |Self| should always be the same type as the implementation. +// It produces a sequence of |ValueType| +type SequenceState[Self any, ValueType comparable] interface { + // Next advances the state of the sequence, producing a SQL value and the next SequenceState. + Next() (sqlVal ValueType, hasNext bool, nextState Self, err error) + // CurrentValue returns the next SQL value in the sequence without advancing it. + CurrentValue() (sqlVal ValueType) + // WithValue sets the state such that the next SQL value will be the one provided, then returns the new SequenceState. + WithValue(sqlVal ValueType) Self + // WithSQLValue coerces the input into the required type, then sets the state such that the next SQL value will be the one provided, then returns the new SequenceState. + WithSQLValue(ctx *sql.Context, v interface{}) (Self, error) + // GreaterThan compares two states and determines which one has advanced further. + GreaterThan(other Self) bool + // AtEnd returns whether the sequence is at its end. If true, then subsequent calls to SequenceState will + // no longer advance the sequence and may return an error. + AtEnd() bool + // Merge combines two states, usually returning the one that's further along. If the states can't be combined, + // |ok| is false. + Merge(other Self) (merged Self, ok bool) +} + +// A SequencedRelation is a table or root object with a state that must be shared across all branches +// and transactions in order to ensure global uniqueness of created rows. It wraps a SequenceState. +// Because some implementations are immutable (doltdb.Table), the methods that set the state must return a new +// value instead of self-mutating. +type SequencedRelation[Self any, ValueType comparable, StateType SequenceState[StateType, ValueType]] interface { + // GetSequenceState returns the current SequenceState of the object. + GetSequenceState(ctx context.Context) (StateType, error) + // HasSequenceState returns whether the relation wraps a sequence. + // (This may be false, for instance, for tables that do not have an AUTO INCREMENT column) + HasSequenceState(ctx context.Context) (bool, error) + // GetAutoIncrementSqlType returns the SQL type generated by the sequence. + GetAutoIncrementSqlType(ctx context.Context) (sql.Type, bool, error) + // SetSequenceState unconditionally sets the SequenceState for the object. + SetSequenceState(ctx context.Context, val StateType) (Self, error) + // TrySetSequenceState attempts to set the SequenceState for the object, but may fail if the provided state is + // incompatible with the other state of the object. Returns a bool indicating success. + // (For example, setting a table's AUTO INCREMENT value to be less than a value currently written to the table will fail.) + TrySetSequenceState(ctx *sql.Context, val StateType) (Self, bool, error) +} From 3fc7f984ca699f1992bf05b6c5c45550fb1b3da6 Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Mon, 20 Jul 2026 22:49:53 -0700 Subject: [PATCH 04/11] Have AutoIncrementTracker implement SequenceTracker --- go/go.mod | 4 - go/libraries/doltcore/doltdb/table.go | 192 +++++++- .../doltcore/merge/merge_prolly_rows.go | 4 +- go/libraries/doltcore/sqle/database.go | 13 +- .../sqle/dsess/auto_increment_tracker.go | 56 +++ .../sqle/dsess/auto_increment_tracker_test.go | 7 +- .../doltcore/sqle/dsess/globalstate.go | 30 +- ...crement_tracker.go => sequence_tracker.go} | 440 ++++++++---------- go/libraries/doltcore/sqle/dsess/session.go | 4 +- .../globalstate/auto_increment_tracker.go | 50 -- .../doltcore/sqle/globalstate/global_state.go | 8 +- .../sqle/globalstate/sequence_tracker.go | 59 +++ .../sqle/globalstate/sequences/state.go | 6 +- go/libraries/doltcore/sqle/tables.go | 30 +- go/libraries/doltcore/sqle/temp_table.go | 6 +- .../sqle/writer/prolly_table_writer.go | 17 +- .../sqle/writer/prolly_write_session.go | 5 +- 17 files changed, 575 insertions(+), 356 deletions(-) create mode 100644 go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go rename go/libraries/doltcore/sqle/dsess/{autoincrement_tracker.go => sequence_tracker.go} (53%) delete mode 100644 go/libraries/doltcore/sqle/globalstate/auto_increment_tracker.go create mode 100644 go/libraries/doltcore/sqle/globalstate/sequence_tracker.go diff --git a/go/go.mod b/go/go.mod index 6211543a5e5..b6ace51f237 100644 --- a/go/go.mod +++ b/go/go.mod @@ -209,7 +209,3 @@ require ( ) go 1.26.2 - -replace ( - github.com/dolthub/go-mysql-server => ../../go-mysql-server -) diff --git a/go/libraries/doltcore/doltdb/table.go b/go/libraries/doltcore/doltdb/table.go index f2b147cc448..8f89a23e528 100644 --- a/go/libraries/doltcore/doltdb/table.go +++ b/go/libraries/doltcore/doltdb/table.go @@ -18,6 +18,12 @@ import ( "context" "errors" "fmt" + "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate/sequences" + + "github.com/dolthub/go-mysql-server/sql" + gmstypes "github.com/dolthub/go-mysql-server/sql/types" + "io" + "math" "unicode" "github.com/dolthub/dolt/go/libraries/doltcore/doltdb/durable" @@ -63,6 +69,8 @@ type Table struct { overriddenSchema schema.Schema } +var _ sequences.SequencedRelation[*Table, uint64, AutoIncrementState] = &Table{} + // NewTable creates a durable object which stores row data, index data, and schema. func NewTable(ctx context.Context, vrw types.ValueReadWriter, ns tree.NodeStore, sch schema.Schema, rows durable.Index, indexes durable.IndexSet, autoIncVal types.Value) (*Table, error) { dt, err := durable.NewTable(ctx, vrw, ns, sch, rows, indexes, autoIncVal) @@ -406,22 +414,198 @@ func (t *Table) RenameIndexRowData(ctx context.Context, oldIndexName, newIndexNa return t.SetIndexSet(ctx, indexes) } +// HasSequenceState returns whether the table has an AUTO_INCREMENT value. +func (t *Table) HasSequenceState(ctx context.Context) (bool, error) { + sch, err := t.GetSchema(ctx) + if err != nil { + return false, err + } + return schema.HasAutoIncrement(sch), nil +} + +// GetSequenceState implements SequencedRelation +func (t *Table) GetSequenceState(ctx context.Context) (AutoIncrementState, error) { + return t.GetAutoIncrementValue(ctx) +} + // GetAutoIncrementValue returns the current AUTO_INCREMENT value for this table. -func (t *Table) GetAutoIncrementValue(ctx context.Context) (uint64, error) { - return t.table.GetAutoIncrement(ctx) +func (t *Table) GetAutoIncrementValue(ctx context.Context) (AutoIncrementState, error) { + currentValue, err := t.table.GetAutoIncrement(ctx) + if err != nil { + return 0, err + } + return AutoIncrementState(currentValue), nil +} + +func (t *Table) GetSequenceSqlType(ctx context.Context) (sql.Type, bool, error) { + sch, err := t.table.GetSchema(ctx) + if err != nil { + return nil, false, err + } + + aiCol, ok := schema.GetAutoIncrementColumn(sch) + if !ok { + return nil, false, nil + } + + return aiCol.TypeInfo.ToSqlType(), true, nil } // SetAutoIncrementValue sets the current AUTO_INCREMENT value for this table. This method does not verify that the // value given is greater than current table values. Setting it lower than current table values will result in // incorrect key generation on future inserts, causing duplicate key errors. -func (t *Table) SetAutoIncrementValue(ctx context.Context, val uint64) (*Table, error) { - table, err := t.table.SetAutoIncrement(ctx, val) +func (t *Table) SetAutoIncrementValue(ctx context.Context, val AutoIncrementState) (*Table, error) { + table, err := t.table.SetAutoIncrement(ctx, val.CurrentValue()) if err != nil { return nil, err } return &Table{table: table}, nil } +// SetSequenceState implements sequences.SequencedRelation +func (t *Table) SetSequenceState(ctx context.Context, val AutoIncrementState) (*Table, error) { + return t.SetAutoIncrementValue(ctx, val) +} + +// SetSequenceState implements sequences.SequencedRelation +func (t *Table) TrySetSequenceState(ctx *sql.Context, newAutoIncVal AutoIncrementState) (*Table, bool, error) { + currentMax, err := getMaxAutoIncrementValue(ctx, t) + if err != nil { + return nil, false, err + } + currentMaxVal := AutoIncrementState(currentMax) + + if !newAutoIncVal.GreaterThan(currentMaxVal) { + return t, false, nil + } + + newTable, err := t.SetSequenceState(ctx, newAutoIncVal) + return newTable, true, err +} + +// CoerceAutoIncrementValue converts |val| into an AUTO_INCREMENT sequence value +func CoerceAutoIncrementValue(ctx *sql.Context, val interface{}) (uint64, error) { + switch typ := val.(type) { + case float32: + val = math.Round(float64(typ)) + case float64: + val = math.Round(typ) + } + + var err error + val, _, err = gmstypes.Uint64.Convert(ctx, val) + if err != nil { + return 0, err + } + if val == nil || val == uint64(0) { + return 0, nil + } + return val.(uint64), nil +} + +type AutoIncrementState uint64 + +func (s AutoIncrementState) Next() (aiVal uint64, ok bool, nextState AutoIncrementState, err error) { + if s == math.MaxUint64 { + return uint64(math.MaxUint64), false, s, nil + } + return uint64(s), true, s + 1, nil +} + +func (s AutoIncrementState) CurrentValue() uint64 { + return uint64(s) +} + +func (s AutoIncrementState) WithValue(v uint64) AutoIncrementState { + return AutoIncrementState(v) +} + +func (s AutoIncrementState) WithSQLValue(ctx *sql.Context, v interface{}) (AutoIncrementState, error) { + given, err := CoerceAutoIncrementValue(ctx, v) + if err != nil { + return s, err + } + return AutoIncrementState(given), nil +} + +func (s AutoIncrementState) GreaterThan(other AutoIncrementState) bool { + return s > other +} + +func (s AutoIncrementState) Merge(other AutoIncrementState) (AutoIncrementState, bool) { + if s > other { + return s, true + } + return other, true +} + +func (s AutoIncrementState) AtEnd() bool { + return s == math.MaxUint64 +} + +// getMaxAutoIncrementValue gets the highest value in a table's AUTO INCREMENT column +func getMaxAutoIncrementValue(ctx *sql.Context, table *Table) (uint64, error) { + // First, establish whether to update this table based on the given value and its current max value. + sch, err := table.GetSchema(ctx) + if err != nil { + return 0, err + } + + aiCol, ok := schema.GetAutoIncrementColumn(sch) + if !ok { + return 0, nil + } + + var indexData durable.Index + aiIndex, ok := sch.Indexes().GetIndexByColumnNames(aiCol.Name) + if ok { + indexes, err := table.GetIndexSet(ctx) + if err != nil { + return 0, err + } + + indexData, err = indexes.GetIndex(ctx, sch, nil, aiIndex.Name()) + if err != nil { + return 0, err + } + } else { + indexData, err = table.GetRowData(ctx) + if err != nil { + return 0, err + } + } + + maxValue, err := getMaxIndexValue(ctx, indexData) + if err != nil { + return 0, err + } + return CoerceAutoIncrementValue(ctx, maxValue) +} + +// getMaxIndexValue reads the highest value for the first column in an index. +func getMaxIndexValue(ctx *sql.Context, indexData durable.Index) (interface{}, error) { + idx, err := durable.ProllyMapFromIndex(indexData) + if err != nil { + return 0, err + } + + iter, err := idx.IterAllReverse(ctx) + if err != nil { + return 0, err + } + + kd, _ := idx.Descriptors() + k, _, err := iter.Next(ctx) + if err == io.EOF { + return 0, nil + } else if err != nil { + return 0, err + } + + // TODO: is the auto-inc column always the first column in the index? + return tree.GetField(ctx, kd, 0, k, idx.NodeStore()) +} + // AddColumnToRows adds the column named to row data as necessary and returns the resulting table. func (t *Table) AddColumnToRows(ctx context.Context, newCol string, newSchema schema.Schema) (*Table, error) { idx, err := t.table.GetTableRows(ctx) diff --git a/go/libraries/doltcore/merge/merge_prolly_rows.go b/go/libraries/doltcore/merge/merge_prolly_rows.go index c26bdc84dc6..109212b2619 100644 --- a/go/libraries/doltcore/merge/merge_prolly_rows.go +++ b/go/libraries/doltcore/merge/merge_prolly_rows.go @@ -140,8 +140,8 @@ func mergeAutoIncrementValues(ctx context.Context, tbl, otherTbl, resultTbl *dol if err != nil { return nil, err } - if autoVal < mergeAutoVal { - autoVal = mergeAutoVal + if mergeAutoVal.GreaterThan(autoVal) { + return resultTbl.SetAutoIncrementValue(ctx, mergeAutoVal) } return resultTbl.SetAutoIncrementValue(ctx, autoVal) } diff --git a/go/libraries/doltcore/sqle/database.go b/go/libraries/doltcore/sqle/database.go index 6889565b259..b53d9a57541 100644 --- a/go/libraries/doltcore/sqle/database.go +++ b/go/libraries/doltcore/sqle/database.go @@ -1924,7 +1924,7 @@ func (db Database) removeTableFromAutoIncrementTracker( wses = append(wses, ws) } - ait, err := db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) + ait, err := dsess.GetAutoIncrementTracker(ctx, db.gs) if err != nil { return err } @@ -2084,11 +2084,14 @@ func (db Database) createSqlTable(ctx *sql.Context, table string, schemaName str // Prevent any tables that use BINARY, CHAR, VARBINARY, VARCHAR prefixes if schema.HasAutoIncrement(doltSch) { - ait, err := db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) + ait, err := dsess.GetAutoIncrementTracker(ctx, db.gs) + if err != nil { + return err + } + err = ait.AddNewTable(tableName.Name, doltdb.AutoIncrementState(1)) if err != nil { return err } - ait.AddNewTable(tableName.Name) } return db.createDoltTable(ctx, tableName.Name, tableName.Schema, root, doltSch) @@ -2144,11 +2147,11 @@ func (db Database) createIndexedSqlTable(ctx *sql.Context, table string, schemaN } if schema.HasAutoIncrement(doltSch) { - ait, err := db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) + ait, err := dsess.GetAutoIncrementTracker(ctx, db.gs) if err != nil { return err } - ait.AddNewTable(tableName.Name) + ait.AddNewTable(tableName.Name, doltdb.AutoIncrementState(1)) } return db.createDoltTable(ctx, tableName.Name, tableName.Schema, root, doltSch) diff --git a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go new file mode 100644 index 00000000000..abd4e25b399 --- /dev/null +++ b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go @@ -0,0 +1,56 @@ +// Copyright 2023 Dolthub, Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package dsess + +import ( + "context" + + "github.com/dolthub/go-mysql-server/sql" + + "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" + "github.com/dolthub/dolt/go/libraries/doltcore/schema" + "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate" +) + +type DoltDBRelationSource struct{} + +func (s DoltDBRelationSource) GetRelation(ctx context.Context, root doltdb.RootValue, tName doltdb.TableName) (relation *doltdb.Table, resolvedName string, found bool, err error) { + return doltdb.GetTableInsensitive(ctx, root, tName) +} + +func (s DoltDBRelationSource) GetRelations(ctx context.Context, root doltdb.RootValue, cb func(doltdb.TableName, *doltdb.Table) (bool, error)) error { + return root.IterTables(ctx, func(name doltdb.TableName, table *doltdb.Table, sch schema.Schema) (stop bool, err error) { + return cb(name, table) + }) +} + +var _ RelationSource[*doltdb.Table, doltdb.AutoIncrementState, uint64] = (*DoltDBRelationSource)(nil) + +type AutoIncrementTracker = SequenceTracker[*doltdb.Table, doltdb.AutoIncrementState, uint64] + +// NewAutoIncrementTracker returns a new autoincrement tracker for the roots given. All roots sets must be +// considered because the auto increment value for a table is tracked globally, across all branches. +// Roots provided should be the working sets when available, or the branches when they are not (e.g. for remote +// branches that don't have a local working set) +func NewAutoIncrementTracker(ctx context.Context, dbName string, roots ...doltdb.Rootish) (*AutoIncrementTracker, error) { + return NewAutoIncrementTrackerI(ctx, dbName, DoltDBRelationSource{}, roots...) +} + +func GetAutoIncrementTracker(ctx *sql.Context, gs globalstate.GlobalState) (*AutoIncrementTracker, error) { + return GetSequenceTracker(ctx, gs, autoIncrementTrackerKey) +} + +// autoIncrementTrackerKey is the key used to store the AutoIncrmenetTracker in the GlobalState's map of SequenceTrackers +var autoIncrementTrackerKey = TrackerKey[*AutoIncrementTracker]{} diff --git a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker_test.go b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker_test.go index 03f03b87b55..a4f460bdc3c 100644 --- a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker_test.go +++ b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker_test.go @@ -17,7 +17,6 @@ package dsess import ( "context" "fmt" - "sync" "testing" "github.com/dolthub/go-mysql-server/sql" @@ -68,7 +67,7 @@ func TestCoerceAutoIncrementValue(t *testing.T) { for _, test := range tests { name := fmt.Sprintf("Coerce %v to %v", test.val, test.exp) t.Run(name, func(t *testing.T) { - act, err := CoerceAutoIncrementValue(ctx, test.val) + act, err := doltdb.CoerceAutoIncrementValue(ctx, test.val) if test.err { assert.Error(t, err) } else { @@ -83,7 +82,7 @@ func TestInitWithRoots(t *testing.T) { t.Run("EmptyRoots", func(t *testing.T) { ait := AutoIncrementTracker{ dbName: "test_database", - sequences: &sync.Map{}, + sequences: &SyncMap[string, doltdb.AutoIncrementState]{}, mm: mutexmap.NewMutexMap(), init: make(chan struct{}), cancelInit: make(chan struct{}), @@ -94,7 +93,7 @@ func TestInitWithRoots(t *testing.T) { t.Run("CloseCancelsInit", func(t *testing.T) { ait := AutoIncrementTracker{ dbName: "test_database", - sequences: &sync.Map{}, + sequences: &SyncMap[string, doltdb.AutoIncrementState]{}, mm: mutexmap.NewMutexMap(), init: make(chan struct{}), cancelInit: make(chan struct{}), diff --git a/go/libraries/doltcore/sqle/dsess/globalstate.go b/go/libraries/doltcore/sqle/dsess/globalstate.go index 5482dc27b05..6001ef89e25 100644 --- a/go/libraries/doltcore/sqle/dsess/globalstate.go +++ b/go/libraries/doltcore/sqle/dsess/globalstate.go @@ -24,8 +24,10 @@ import ( "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate" ) -var AutoIncrementTrackerKey = struct{}{} +// TrackerKey is a type used as a key into the GlobalStateImpl's map of SequenceTrackers +type TrackerKey[TrackerType globalstate.SequenceTrackerBase] struct{} +// NewGlobalStateStoreForDb creates a new GlobalState. It initially contains a single SequenceTracker: the AutoIncrementTracker func NewGlobalStateStoreForDb(ctx context.Context, dbName string, db *doltdb.DoltDB) (GlobalStateImpl, error) { branches, err := db.GetBranches(ctx) if err != nil { @@ -94,33 +96,33 @@ func NewGlobalStateStoreForDb(ctx context.Context, dbName string, db *doltdb.Dol } return GlobalStateImpl{ - aiTrackers: map[interface{}]globalstate.AutoIncrementTracker{AutoIncrementTrackerKey: tracker}, + sequenceTrackers: map[interface{}]globalstate.SequenceTrackerBase{autoIncrementTrackerKey: tracker}, }, nil } type GlobalStateImpl struct { - aiTrackers map[interface{}]globalstate.AutoIncrementTracker + sequenceTrackers map[interface{}]globalstate.SequenceTrackerBase } var _ globalstate.GlobalState = GlobalStateImpl{} -func (g GlobalStateImpl) AutoIncrementTracker(ctx *sql.Context, key interface{}) (globalstate.AutoIncrementTracker, error) { - return g.aiTrackers[key], nil +func (g GlobalStateImpl) GetSequenceTracker(ctx *sql.Context, key interface{}) (globalstate.SequenceTrackerBase, error) { + return g.sequenceTrackers[key], nil } -func (g GlobalStateImpl) AddAutoIncrementTracker(ctx *sql.Context, key interface{}, tracker globalstate.AutoIncrementTracker) error { - g.aiTrackers[key] = tracker +func (g GlobalStateImpl) AddSequenceTracker(ctx *sql.Context, key interface{}, tracker globalstate.SequenceTrackerBase) error { + g.sequenceTrackers[key] = tracker return nil } func (g GlobalStateImpl) Close() { - for _, tracker := range g.aiTrackers { + for _, tracker := range g.sequenceTrackers { tracker.Close() } } func (g GlobalStateImpl) InitWithRoots(ctx *sql.Context, roots ...doltdb.Rootish) error { - for _, tracker := range g.aiTrackers { + for _, tracker := range g.sequenceTrackers { err := tracker.InitWithRoots(ctx, roots...) if err != nil { return err @@ -128,3 +130,13 @@ func (g GlobalStateImpl) InitWithRoots(ctx *sql.Context, roots ...doltdb.Rootish } return nil } + +// GetSequenceTracker returns a SequenceTracker held by the globalstate.GlobalState, keyed by the provided key. +// This function performs the necessary cast so that the caller doesn't have to cast the result. +func GetSequenceTracker[T globalstate.SequenceTrackerBase](ctx *sql.Context, gs globalstate.GlobalState, key TrackerKey[T]) (result T, err error) { + aiti, err := gs.GetSequenceTracker(ctx, key) + if err != nil { + return result, err + } + return aiti.(T), nil +} diff --git a/go/libraries/doltcore/sqle/dsess/autoincrement_tracker.go b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go similarity index 53% rename from go/libraries/doltcore/sqle/dsess/autoincrement_tracker.go rename to go/libraries/doltcore/sqle/dsess/sequence_tracker.go index 60b55ef63f9..4b031a33584 100644 --- a/go/libraries/doltcore/sqle/dsess/autoincrement_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go @@ -18,24 +18,18 @@ import ( "context" "errors" "fmt" - "io" - "math" "strings" - "sync" "time" "github.com/dolthub/go-mysql-server/sql" - gmstypes "github.com/dolthub/go-mysql-server/sql/types" "golang.org/x/sync/errgroup" "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" - "github.com/dolthub/dolt/go/libraries/doltcore/doltdb/durable" "github.com/dolthub/dolt/go/libraries/doltcore/doltdb/gcctx" "github.com/dolthub/dolt/go/libraries/doltcore/ref" - "github.com/dolthub/dolt/go/libraries/doltcore/schema" "github.com/dolthub/dolt/go/libraries/doltcore/sqle/dsess/mutexmap" "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate" - "github.com/dolthub/dolt/go/store/prolly/tree" + "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate/sequences" ) type LockMode int64 @@ -46,11 +40,25 @@ var ( LockMode_Interleaved LockMode = 2 ) -type AutoIncrementTracker struct { +// A RelationSource maps table names to relations (which may be tables or root objects) at a supplied RootValue +type RelationSource[ + RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], + StateType sequences.SequenceState[StateType, ValueType], + ValueType comparable, +] interface { + GetRelation(ctx context.Context, root doltdb.RootValue, tName doltdb.TableName) (relation RelationType, resolvedName string, found bool, err error) + GetRelations(ctx context.Context, root doltdb.RootValue, cb func(doltdb.TableName, RelationType) (bool, error)) error +} + +type SequenceTracker[ + RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], + StateType sequences.SequenceState[StateType, ValueType], + ValueType comparable, +] struct { initErr error - sequences *sync.Map // map[string]uint64 + sequences *SyncMap[string, StateType] mm *mutexmap.MutexMap - // AutoIncrementTracker is lazily initialized by loading + // SequenceTracker is lazily initialized by loading // tracker state for every given |root|. On first access, we // block on initialization being completed and we terminally // return |initErr| if there was any error initializing. @@ -60,9 +68,12 @@ type AutoIncrementTracker struct { // async initialization and block on the process completing. cancelInit chan struct{} dbName string - // lockMode is the effective @@innodb_autoinc_lock_mode at the time of AutoIncrementTracker initialization. + // lockMode is the effective @@innodb_autoinc_lock_mode at the time of SequenceTracker initialization. // This value can only be set by config and cannot be changed in a running server. lockMode LockMode + // relationSource is how the tracker reads objects from a RootValue. + // It may read tables or RootObjects. + relationSource RelationSource[RelationType, StateType, ValueType] } // currentLockMode returns the effective @@innodb_autoinc_lock_mode stored in global server vars @@ -74,19 +85,22 @@ func currentLockMode() LockMode { return LockMode_Interleaved } -var _ globalstate.AutoIncrementTracker = &AutoIncrementTracker{} +func (a *SequenceTracker[RelationType, StateType, ValueType]) staticAssertTypes() { + var _ globalstate.SequenceTracker[RelationType, StateType, ValueType] = a +} -// NewAutoIncrementTracker returns a new autoincrement tracker for the roots given. All roots sets must be -// considered because the auto increment value for a table is tracked globally, across all branches. -// Roots provided should be the working sets when available, or the branches when they are not (e.g. for remote -// branches that don't have a local working set) -func NewAutoIncrementTracker(ctx context.Context, dbName string, roots ...doltdb.Rootish) (*AutoIncrementTracker, error) { - ait := AutoIncrementTracker{ - dbName: dbName, - sequences: &sync.Map{}, - mm: mutexmap.NewMutexMap(), - init: make(chan struct{}), - cancelInit: make(chan struct{}), +func NewAutoIncrementTrackerI[ + RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], + StateType sequences.SequenceState[StateType, ValueType], + ValueType comparable, +](ctx context.Context, dbName string, relationSource RelationSource[RelationType, StateType, ValueType], roots ...doltdb.Rootish) (*SequenceTracker[RelationType, StateType, ValueType], error) { + ait := SequenceTracker[RelationType, StateType, ValueType]{ + dbName: dbName, + sequences: &SyncMap[string, StateType]{}, + mm: mutexmap.NewMutexMap(), + init: make(chan struct{}), + cancelInit: make(chan struct{}), + relationSource: relationSource, } gcSafepointController := getGCSafepointController(ctx) ctx = context.Background() @@ -111,91 +125,90 @@ func getGCSafepointController(ctx context.Context) *gcctx.GCSafepointController return gcctx.GetGCSafepointController(ctx) } -func loadAutoIncValue(sequences *sync.Map, tableName string) (current uint64, hasCurrent bool) { +func loadSequenceState[StateType sequences.SequenceState[StateType, ValueType], ValueType comparable](sequences *SyncMap[string, StateType], tableName string) (current StateType, hasCurrent bool) { tableName = strings.ToLower(tableName) - stored, hasCurrent := sequences.Load(tableName) - if !hasCurrent { - return 0, false - } - return stored.(uint64), true + return sequences.Load(tableName) } -func (a *AutoIncrementTracker) initializeTableAutoIncrement(ctx *sql.Context, tableName string) (uint64, bool, error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAutoIncrement(ctx *sql.Context, tableName string, initialValue interface{}) (state StateType, hasState bool, err error) { sess := DSessFromSess(ctx.Session) ws, err := sess.WorkingSet(ctx, a.dbName) if err != nil { - return 0, false, err + return state, false, err } - table, _, ok, err := doltdb.GetTableInsensitive(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tableName}) + table, _, ok, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tableName}) if err != nil || !ok { - return 0, false, err + return state, false, err } - sch, err := table.GetSchema(ctx) + hasAutoIncrement, err := table.HasSequenceState(ctx) if err != nil { - return 0, false, err - } - if !schema.HasAutoIncrement(sch) { - return 0, false, nil + return state, false, err } - seq, err := table.GetAutoIncrementValue(ctx) - if err != nil { - return 0, false, err + var seq StateType + if !hasAutoIncrement { + // Create a new state based on the provided value + // TODO: This could cause problems when we need to create the more + // complicated sequence types for Doltgres + seq, err = state.WithSQLValue(ctx, initialValue) + if err != nil { + return state, false, err + } + } else { + seq, err = table.GetSequenceState(ctx) + if err != nil { + return state, false, err + } } - table, err = a.deepSet(ctx, tableName, table, ws.Ref(), seq) + relation, err := a.deepSet(ctx, tableName, table, ws.Ref(), seq) if err != nil { - return 0, false, err + return state, false, err } - seq, ok = loadAutoIncValue(a.sequences, tableName) + seq, ok = loadSequenceState(a.sequences, tableName) if ok { - return seq, true, nil + return state, true, nil } - seq, err = table.GetAutoIncrementValue(ctx) + seq, err = relation.GetSequenceState(ctx) if err != nil { - return 0, false, err + return state, false, err } a.sequences.Store(strings.ToLower(tableName), seq) - return seq, true, nil + return state, true, nil } -func (a *AutoIncrementTracker) Close() { +func (a *SequenceTracker[RelationType, StateType, ValueType]) Close() { close(a.cancelInit) <-a.init } // Current returns the next value to be generated in the auto increment sequence for |tableName|. -func (a *AutoIncrementTracker) Current(tableName string) (uint64, error) { - err := a.waitForInit() +func (a *SequenceTracker[RelationType, StateType, ValueType]) Current(relation string) (current StateType, err error) { + err = a.waitForInit() if err != nil { - return 0, err + return current, err } - seq, ok := loadAutoIncValue(a.sequences, tableName) + seq, ok := loadSequenceState(a.sequences, relation) if !ok { - return 0, nil + return current, nil } return seq, nil } // Next returns the next auto increment value for |tbl| using |insertVal| from an insert. If |insertVal| is // null or 0, it is generated from the sequence. -func (a *AutoIncrementTracker) Next(ctx *sql.Context, tbl string, insertVal interface{}) (uint64, error) { - err := a.waitForInit() +func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Context, tbl string, insertVal interface{}) (nextValue ValueType, err error) { + err = a.waitForInit() if err != nil { - return 0, err + return nextValue, err } tbl = strings.ToLower(tbl) - given, err := CoerceAutoIncrementValue(ctx, insertVal) - if err != nil { - return 0, err - } - // The read-modify-write of the sequence below must be atomic across concurrent inserters. In // interleaved lock mode (the default) the engine holds no statement-level lock, so we take a // short per-table lock here. @@ -206,7 +219,7 @@ func (a *AutoIncrementTracker) Next(ctx *sql.Context, tbl string, insertVal inte locked = true } - curr, ok := loadAutoIncValue(a.sequences, tbl) + currState, ok := loadSequenceState(a.sequences, tbl) if !ok { // Missing tracker state after initialization can happen when a running sql-server discovers a database // restored after startup, so initialize it here. @@ -217,80 +230,63 @@ func (a *AutoIncrementTracker) Next(ctx *sql.Context, tbl string, insertVal inte locked = true } - curr, ok = loadAutoIncValue(a.sequences, tbl) + currState, ok = loadSequenceState(a.sequences, tbl) } if !ok { - curr, ok, err = a.initializeTableAutoIncrement(ctx, tbl) + currState, ok, err = a.initializeTableAutoIncrement(ctx, tbl, insertVal) if err != nil { - return 0, err + return nextValue, err } if !ok { - return 0, fmt.Errorf("autoIncrementTracker: unable to find sequence for table %s", tbl) + return nextValue, fmt.Errorf("autoIncrementTracker: unable to find sequence for table %s", tbl) } } } - if given == 0 { + if insertVal == nil { // |given| is 0 or NULL - a.sequences.Store(tbl, curr+1) - return curr, nil + currentVal, _, nextState, err := currState.Next() + if err != nil { + return nextValue, err + } + a.sequences.Store(tbl, nextState) + return currentVal, nil } - if given >= curr { + givenState, err := currState.WithSQLValue(ctx, insertVal) + if err != nil { + return nextValue, err + } + given := givenState.CurrentValue() + + if !currState.GreaterThan(givenState) { // Check if the given value is valid for this column type - if !a.validateAutoIncrementBounds(ctx, tbl, given, false) { - return given, nil // Out of bounds, don't update sequence + if !a.validateAutoIncrementBounds(ctx, tbl, givenState, false) { + return givenState.CurrentValue(), nil // Out of bounds, don't update sequence } // Value is valid, determine next sequence value - nextVal := given - if a.validateAutoIncrementBounds(ctx, tbl, given, true) { - nextVal++ + if a.validateAutoIncrementBounds(ctx, tbl, givenState, true) { + _, _, givenState, err = givenState.Next() + if err != nil { + return nextValue, err + } } - a.sequences.Store(tbl, nextVal) + a.sequences.Store(tbl, givenState) return given, nil } - // |given| < curr return given, nil } -func (a *AutoIncrementTracker) CoerceAutoIncrementValue(ctx *sql.Context, val interface{}) (uint64, error) { - err := a.waitForInit() - if err != nil { - return 0, err - } - return CoerceAutoIncrementValue(ctx, val) -} - -// CoerceAutoIncrementValue converts |val| into an AUTO_INCREMENT sequence value -func CoerceAutoIncrementValue(ctx *sql.Context, val interface{}) (uint64, error) { - switch typ := val.(type) { - case float32: - val = math.Round(float64(typ)) - case float64: - val = math.Round(typ) - } - - var err error - val, _, err = gmstypes.Uint64.Convert(ctx, val) - if err != nil { - return 0, err - } - if val == nil || val == uint64(0) { - return 0, nil - } - return val.(uint64), nil -} - // Set sets the auto increment value for the table named, if it's greater than the one already registered for this // table. Otherwise, the update is silently disregarded. So far this matches the MySQL behavior, but Dolt uses the // maximum value for this table across all branches. -func (a *AutoIncrementTracker) Set(ctx *sql.Context, tableName string, table *doltdb.Table, ws ref.WorkingSetRef, newAutoIncVal uint64) (*doltdb.Table, error) { - err := a.waitForInit() +func (a *SequenceTracker[RelationType, StateType, ValueType]) Set(ctx *sql.Context, tableName string, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (newRelation RelationType, err error) { + err = a.waitForInit() if err != nil { - return nil, err + return newRelation, err } tableName = strings.ToLower(tableName) @@ -298,24 +294,26 @@ func (a *AutoIncrementTracker) Set(ctx *sql.Context, tableName string, table *do release := a.mm.Lock(tableName) defer release() - existing, ok := loadAutoIncValue(a.sequences, tableName) + existing, ok := loadSequenceState(a.sequences, tableName) if !ok { - existing = 0 - } - if newAutoIncVal > existing && a.validateAutoIncrementBounds(ctx, tableName, newAutoIncVal, true) { - a.sequences.Store(tableName, newAutoIncVal) - return table.SetAutoIncrementValue(ctx, newAutoIncVal) - } else if newAutoIncVal > existing { + a.sequences.Store(tableName, newSequenceState) + return table.SetSequenceState(ctx, newSequenceState) + } + gt := newSequenceState.GreaterThan(existing) + if gt && a.validateAutoIncrementBounds(ctx, tableName, newSequenceState, true) { + a.sequences.Store(tableName, newSequenceState) + return table.SetSequenceState(ctx, newSequenceState) + } else if gt { // Value is greater but out of bounds, don't update return table, nil } // Value is not greater than current, do deep check across branches - return a.deepSet(ctx, tableName, table, ws, newAutoIncVal) + return a.deepSet(ctx, tableName, table, ws, newSequenceState) } -// deepSet sets the auto increment value for the table named, if it's greater than the one on any branch head for this +// deepSet sets the sequence state for the table named, if it's greater than the one on any branch head for this // database, ignoring the current in-memory tracker value -func (a *AutoIncrementTracker) deepSet(ctx *sql.Context, tableName string, table *doltdb.Table, ws ref.WorkingSetRef, newAutoIncVal uint64) (*doltdb.Table, error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.Context, tableName string, table RelationType, ws ref.WorkingSetRef, newAutoIncVal StateType) (newRelation RelationType, err error) { sess := DSessFromSess(ctx.Session) db, ok := sess.Provider().BaseDatabase(ctx, a.dbName) @@ -324,63 +322,26 @@ func (a *AutoIncrementTracker) deepSet(ctx *sql.Context, tableName string, table return table, nil } - // First, establish whether to update this table based on the given value and its current max value. - sch, err := table.GetSchema(ctx) + table, success, err := table.TrySetSequenceState(ctx, newAutoIncVal) if err != nil { - return nil, err - } - - aiCol, ok := schema.GetAutoIncrementColumn(sch) - if !ok { - return nil, nil + return newRelation, err } - - var indexData durable.Index - aiIndex, ok := sch.Indexes().GetIndexByColumnNames(aiCol.Name) - if ok { - indexes, err := table.GetIndexSet(ctx) - if err != nil { - return nil, err - } - - indexData, err = indexes.GetIndex(ctx, sch, nil, aiIndex.Name()) - if err != nil { - return nil, err - } - } else { - indexData, err = table.GetRowData(ctx) - if err != nil { - return nil, err - } - } - - currentMax, err := getMaxIndexValue(ctx, indexData) - if err != nil { - return nil, err - } - - // If the given value is less than the current one, the operation is a no-op, bail out early - if newAutoIncVal <= currentMax { + if !success { return table, nil } - table, err = table.SetAutoIncrementValue(ctx, newAutoIncVal) - if err != nil { - return nil, err - } - // Now that we have established the current max for this table, reset the global max accordingly maxAutoInc := newAutoIncVal doltdbs := db.DoltDatabases() for _, db := range doltdbs { branches, err := db.GetBranches(ctx) if err != nil { - return nil, err + return newRelation, err } remotes, err := db.GetRemoteRefs(ctx) if err != nil { - return nil, err + return newRelation, err } rootRefs := make([]ref.DoltRef, 0, len(branches)+len(remotes)) @@ -393,7 +354,7 @@ func (a *AutoIncrementTracker) deepSet(ctx *sql.Context, tableName string, table case ref.BranchRefType: wsRef, err := ref.WorkingSetRefForHead(b) if err != nil { - return nil, err + return newRelation, err } if wsRef == ws { @@ -406,52 +367,55 @@ func (a *AutoIncrementTracker) deepSet(ctx *sql.Context, tableName string, table // use the branch head if there isn't a working set for it cm, err := db.ResolveCommitRef(ctx, b) if err != nil { - return nil, err + return newRelation, err } rootish = cm } else if err != nil { - return nil, err + return newRelation, err } else { rootish = ws } case ref.RemoteRefType: cm, err := db.ResolveCommitRef(ctx, b) if err != nil { - return nil, err + return newRelation, err } rootish = cm } root, err := rootish.ResolveRootValue(ctx) if err != nil { - return nil, err + return newRelation, err } - table, _, ok, err := doltdb.GetTableInsensitive(ctx, root, doltdb.TableName{Name: tableName}) + table, _, ok, err := a.relationSource.GetRelation(ctx, root, doltdb.TableName{Name: tableName}) if err != nil { - return nil, err + return newRelation, err } if !ok { continue } - sch, err := table.GetSchema(ctx) + hasAutoIncrement, err := table.HasSequenceState(ctx) if err != nil { - return nil, err + return newRelation, err } - if !schema.HasAutoIncrement(sch) { + if !hasAutoIncrement { continue } tableName = strings.ToLower(tableName) - seq, err := table.GetAutoIncrementValue(ctx) + seq, err := table.GetSequenceState(ctx) if err != nil { - return nil, err + return newRelation, err } - if seq > maxAutoInc { - maxAutoInc = seq + var mergeOk bool + maxAutoInc, mergeOk = maxAutoInc.Merge(seq) + if !mergeOk { + // TODO: This can't happen with AUTO INCREMENT but needs to be + // handled for sequences. } } } @@ -462,41 +426,8 @@ func (a *AutoIncrementTracker) deepSet(ctx *sql.Context, tableName string, table return table, nil } -func getMaxIndexValue(ctx *sql.Context, indexData durable.Index) (uint64, error) { - idx, err := durable.ProllyMapFromIndex(indexData) - if err != nil { - return 0, err - } - - iter, err := idx.IterAllReverse(ctx) - if err != nil { - return 0, err - } - - kd, _ := idx.Descriptors() - k, _, err := iter.Next(ctx) - if err == io.EOF { - return 0, nil - } else if err != nil { - return 0, err - } - - // TODO: is the auto-inc column always the first column in the index? - field, err := tree.GetField(ctx, kd, 0, k, idx.NodeStore()) - if err != nil { - return 0, err - } - - maxVal, err := CoerceAutoIncrementValue(ctx, field) - if err != nil { - return 0, err - } - - return maxVal, nil -} - // AddNewTable initializes a new table with an auto increment column to the tracker, as necessary -func (a *AutoIncrementTracker) AddNewTable(tableName string) error { +func (a *SequenceTracker[RelationType, StateType, ValueType]) AddNewTable(tableName string, initialState StateType) error { err := a.waitForInit() if err != nil { return err @@ -504,14 +435,14 @@ func (a *AutoIncrementTracker) AddNewTable(tableName string) error { tableName = strings.ToLower(tableName) // only initialize the sequence for this table if no other branch has such a table - a.sequences.LoadOrStore(tableName, uint64(1)) + a.sequences.LoadOrStore(tableName, initialState) return nil } // DropTable drops the table with the name given. // To establish the new auto increment value, callers must also pass all other working sets in scope that may include // a table with the same name, omitting the working set that just deleted the table named. -func (a *AutoIncrementTracker) DropTable(ctx *sql.Context, tableName string, wses ...*doltdb.WorkingSet) error { +func (a *SequenceTracker[RelationType, StateType, ValueType]) DropTable(ctx *sql.Context, tableName string, wses ...*doltdb.WorkingSet) error { err := a.waitForInit() if err != nil { return err @@ -522,11 +453,11 @@ func (a *AutoIncrementTracker) DropTable(ctx *sql.Context, tableName string, wse release := a.mm.Lock(tableName) defer release() - newHighestValue := uint64(1) + var newHighestValue *StateType // Get the new highest value from all tables in the working sets given for _, ws := range wses { - table, _, exists, err := doltdb.GetTableInsensitive(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tableName}) + table, _, exists, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tableName}) if err != nil { return err } @@ -535,29 +466,39 @@ func (a *AutoIncrementTracker) DropTable(ctx *sql.Context, tableName string, wse continue } - sch, err := table.GetSchema(ctx) + hasAutoIncrement, err := table.HasSequenceState(ctx) if err != nil { return err } - - if schema.HasAutoIncrement(sch) { - seq, err := table.GetAutoIncrementValue(ctx) + if hasAutoIncrement { + seq, err := table.GetSequenceState(ctx) if err != nil { return err } - - if seq > newHighestValue { - newHighestValue = seq + if newHighestValue == nil { + newHighestValue = &seq + } else { + var ok bool + *newHighestValue, ok = (*newHighestValue).Merge(seq) + if !ok { + // TODO: This can't happen with AUTO INCREMENT but needs to be + // handled for sequences. + } } + } } - a.sequences.Store(tableName, newHighestValue) + if newHighestValue != nil { + a.sequences.Store(tableName, *newHighestValue) + } else { + a.sequences.Delete(tableName) + } return nil } -func (a *AutoIncrementTracker) AcquireTableLock(ctx *sql.Context, tableName string) (func(), error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) AcquireTableLock(ctx *sql.Context, tableName string) (func(), error) { err := a.waitForInit() if err != nil { return nil, err @@ -570,7 +511,7 @@ func (a *AutoIncrementTracker) AcquireTableLock(ctx *sql.Context, tableName stri return a.mm.Lock(tableName), nil } -func (a *AutoIncrementTracker) waitForInit() error { +func (a *SequenceTracker[RelationType, StateType, ValueType]) waitForInit() error { select { case <-a.init: return a.initErr @@ -579,7 +520,7 @@ func (a *AutoIncrementTracker) waitForInit() error { } } -// This method will initialize the AutoIncrementTracker state with all +// This method will initialize the SequenceTracker state with all // data from the tables found in |roots|. This method closes the // |a.init| channel when it completes. It is meant to be run in a // goroutine, as in `go a.initWithRoots(...)`. When running this method, @@ -589,7 +530,7 @@ func (a *AutoIncrementTracker) waitForInit() error { // |initWithRoots| is called with appropriately outlives the end of // the method and that it participates in GC lifecycle callbacks // appropriately, if that is necessary. -func (a *AutoIncrementTracker) initWithRoots(ctx context.Context, roots ...doltdb.Rootish) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) initWithRoots(ctx context.Context, roots ...doltdb.Rootish) { defer close(a.init) // Cancel the parent context so that the errgroup work will @@ -619,28 +560,30 @@ func (a *AutoIncrementTracker) initWithRoots(ctx context.Context, roots ...doltd return err } - return r.IterTables(ctx, func(tableName doltdb.TableName, table *doltdb.Table, sch schema.Schema) (bool, error) { - if !schema.HasAutoIncrement(sch) { + init := func(tableName doltdb.TableName, relation RelationType) (bool, error) { + hasSequenceState, err := relation.HasSequenceState(ctx) + if err != nil { + return true, err + } + if !hasSequenceState { return false, nil } - - seq, err := table.GetAutoIncrementValue(ctx) + seq, err := relation.GetSequenceState(ctx) if err != nil { return true, err } tableNameStr := tableName.ToLower().Name if oldValue, loaded := a.sequences.LoadOrStore(tableNameStr, seq); loaded { - old := oldValue.(uint64) - for seq > old && !a.sequences.CompareAndSwap(tableNameStr, old, seq) { + for seq.GreaterThan(oldValue) && !a.sequences.CompareAndSwap(tableNameStr, oldValue, seq) { oldValue, _ = a.sequences.Load(tableNameStr) - old = oldValue.(uint64) - } } return false, nil - }) + } + + return a.relationSource.GetRelations(ctx, r, init) }) } @@ -649,7 +592,7 @@ func (a *AutoIncrementTracker) initWithRoots(ctx context.Context, roots ...doltd } // validateAutoIncrementBounds checks if a value (or value+1 if checkIncrement) is valid for the auto-increment column type -func (a *AutoIncrementTracker) validateAutoIncrementBounds(ctx *sql.Context, tbl string, val uint64, checkIncrement bool) bool { +func (a *SequenceTracker[RelationType, StateType, ValueType]) validateAutoIncrementBounds(ctx *sql.Context, tbl string, val StateType, checkIncrement bool) bool { sess := DSessFromSess(ctx.Session) db, ok := sess.Provider().BaseDatabase(ctx, a.dbName) if !ok || !db.Versioned() { @@ -661,38 +604,39 @@ func (a *AutoIncrementTracker) validateAutoIncrementBounds(ctx *sql.Context, tbl return true } - table, _, ok, err := doltdb.GetTableInsensitive(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tbl}) + table, _, ok, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tbl}) if err != nil || !ok { return true } - sch, err := table.GetSchema(ctx) - if err != nil { + hasSequenceState, err := table.HasSequenceState(ctx) + if !hasSequenceState { + // fail-open because the table writer could be in the process of adding auto-increment to a column return true } - aiCol, ok := schema.GetAutoIncrementColumn(sch) - if !ok { + sqlType, ok, err := table.GetSequenceSqlType(ctx) + if err != nil || !ok { return true } - sqlType := aiCol.TypeInfo.ToSqlType() - testVal := val if checkIncrement { - // Check if incrementing would overflow - nextVal := val + 1 - if nextVal < val { - return false // uint64 overflow + // TODO: Remove error parameter? + _, hasNext, nextVal, _ := val.Next() + // SequenceState can only error if there is no next value. + // Consider changing this to a separate |ok| return value. + if !hasNext { + return false } testVal = nextVal } - _, inRange, err := sqlType.Convert(ctx, testVal) + _, inRange, err := sqlType.Convert(ctx, testVal.CurrentValue()) return err == nil && inRange == sql.InRange } -func (a *AutoIncrementTracker) InitWithRoots(ctx context.Context, roots ...doltdb.Rootish) error { +func (a *SequenceTracker[RelationType, StateType, ValueType]) InitWithRoots(ctx context.Context, roots ...doltdb.Rootish) error { err := a.waitForInit() if err != nil { return err diff --git a/go/libraries/doltcore/sqle/dsess/session.go b/go/libraries/doltcore/sqle/dsess/session.go index e25483ae1a9..9adea55325d 100644 --- a/go/libraries/doltcore/sqle/dsess/session.go +++ b/go/libraries/doltcore/sqle/dsess/session.go @@ -1407,7 +1407,7 @@ func (d *DoltSession) addDB(ctx *sql.Context, db SqlDatabase) error { } sessionState.globalState = stateProvider.GetGlobalState() - tracker, err := sessionState.globalState.AutoIncrementTracker(ctx, AutoIncrementTrackerKey) + tracker, err := GetAutoIncrementTracker(ctx, sessionState.globalState) if err != nil { return err } @@ -2075,4 +2075,4 @@ func DefaultHead(ctx *sql.Context, baseName string, db SqlDatabase) (string, err // WriteSessFunc is a constructor that session builders use to // create fresh table editors. // The indirection avoids a writer/dsess package import cycle. -type WriteSessFunc func(dbName string, ws *doltdb.WorkingSet, aiTracker globalstate.AutoIncrementTracker, setter SessionRootSetter, opts editor.Options) WriteSession +type WriteSessFunc func(dbName string, ws *doltdb.WorkingSet, aiTracker *AutoIncrementTracker, setter SessionRootSetter, opts editor.Options) WriteSession diff --git a/go/libraries/doltcore/sqle/globalstate/auto_increment_tracker.go b/go/libraries/doltcore/sqle/globalstate/auto_increment_tracker.go deleted file mode 100644 index e9ff60631c2..00000000000 --- a/go/libraries/doltcore/sqle/globalstate/auto_increment_tracker.go +++ /dev/null @@ -1,50 +0,0 @@ -// Copyright 2021 Dolthub, Inc. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -package globalstate - -import ( - "context" - - "github.com/dolthub/go-mysql-server/sql" - - "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" - "github.com/dolthub/dolt/go/libraries/doltcore/ref" -) - -// AutoIncrementTracker knows how to get and set the current auto increment value for a table. It's defined as an -// interface here because implementations need to reach into session state, requiring a dependency on this package. -type AutoIncrementTracker interface { - // Current returns the current auto increment value for the given table. - Current(tableName string) (uint64, error) - // Next returns the next auto increment value for the given table, and increments the current value. - Next(ctx *sql.Context, tbl string, insertVal interface{}) (uint64, error) - // AddNewTable adds a new table to the tracker, initializing the auto increment value to 1. - AddNewTable(tableName string) error - // DropTable removes a table from the tracker. - DropTable(ctx *sql.Context, tableName string, wses ...*doltdb.WorkingSet) error - // CoerceAutoIncrementValue coerces the given value to a uint64, returning an error if it can't be done. - CoerceAutoIncrementValue(ctx *sql.Context, val interface{}) (uint64, error) - // Set sets the auto increment value for the given table. This operation may silently do nothing if this value is - // below the current value for this table. The table in the provided working set is assumed to already have the value - // given, so the new global maximum is computed without regard for its value in that working set. - Set(ctx *sql.Context, tableName string, table *doltdb.Table, ws ref.WorkingSetRef, newAutoIncVal uint64) (*doltdb.Table, error) - // AcquireTableLock acquires the auto increment lock on a table, and returns a callback function to release the lock. - // Depending on the value of the `innodb_autoinc_lock_mode` system variable, the engine may need to acquire and hold - // the lock for the duration of an insert statement. - AcquireTableLock(ctx *sql.Context, tableName string) (func(), error) - // InitWithRoots fills the AutoIncrementTracker with values pulled from each root in order. - InitWithRoots(ctx context.Context, roots ...doltdb.Rootish) error - Close() -} diff --git a/go/libraries/doltcore/sqle/globalstate/global_state.go b/go/libraries/doltcore/sqle/globalstate/global_state.go index 49a5dd37e86..5205c558da5 100644 --- a/go/libraries/doltcore/sqle/globalstate/global_state.go +++ b/go/libraries/doltcore/sqle/globalstate/global_state.go @@ -22,9 +22,11 @@ import ( // GlobalState is just a holding interface for pieces of global state, of which the auto increment tracking info is // the only example at the moment. type GlobalState interface { - // AutoIncrementTracker returns the auto increment tracker for this global state. - AutoIncrementTracker(ctx *sql.Context, key interface{}) (AutoIncrementTracker, error) - AddAutoIncrementTracker(ctx *sql.Context, key interface{}, value AutoIncrementTracker) error + // GetSequenceTracker returns the auto increment tracker for this global state. + GetSequenceTracker(ctx *sql.Context, key interface{}) (SequenceTrackerBase, error) + // AddSequenceTracker adds a new SequenceTracker to the GlobalState, accessible by the provided key. + AddSequenceTracker(ctx *sql.Context, key interface{}, value SequenceTrackerBase) error + // InitWithRoots initializes all of the state's SequenceTrackers InitWithRoots(ctx *sql.Context, roots ...doltdb.Rootish) error } diff --git a/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go b/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go new file mode 100644 index 00000000000..72c8ebf3755 --- /dev/null +++ b/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go @@ -0,0 +1,59 @@ +// Copyright 2021 Dolthub, Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package globalstate + +import ( + "context" + "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate/sequences" + + "github.com/dolthub/go-mysql-server/sql" + + "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" + "github.com/dolthub/dolt/go/libraries/doltcore/ref" +) + +// SequenceTrackerBase is the non-generic base interface for SequenceTracker +// It is useful for establishing an upper bound on type parameters in some circumstances. +type SequenceTrackerBase interface { + // AcquireTableLock acquires the auto increment lock on a relation, and returns a callback function to release the lock. + // Depending on the value of the `innodb_autoinc_lock_mode` system variable, the engine may need to acquire and hold + // the lock for the duration of an insert statement. + AcquireTableLock(ctx *sql.Context, tableName string) (func(), error) + // DropTable removes a relation from the tracker. + DropTable(ctx *sql.Context, relation string, wses ...*doltdb.WorkingSet) error + // InitWithRoots fills the SequenceTracker with values pulled from each root in order. + InitWithRoots(ctx context.Context, roots ...doltdb.Rootish) error + Close() +} + +// SequenceTracker knows how to get and set the current auto increment value for a relation (a table or a root object). +// It's defined as an interface here because implementations need to reach into session state, requiring a dependency on this package. +type SequenceTracker[ + RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], + StateType sequences.SequenceState[StateType, ValueType], + ValueType comparable, +] interface { + SequenceTrackerBase + // Current returns the current auto increment state for the given relation. + Current(relation string) (StateType, error) + // Next returns the next SQL value produced by the given relation, and advances that relation's state. + Next(ctx *sql.Context, relation string, insertVal interface{}) (ValueType, error) + // AddNewTable adds a new table to the tracker, initializing the auto increment value to the provided |initialState|. + AddNewTable(relation string, initialState StateType) error + // Set sets the auto increment value for the given relation. This operation may silently do nothing if this value is + // below the current value for this relation. The relation in the provided working set is assumed to already have the value + // given, so the new global maximum is computed without regard for its value in that working set. + Set(ctx *sql.Context, tableName string, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (RelationType, error) +} diff --git a/go/libraries/doltcore/sqle/globalstate/sequences/state.go b/go/libraries/doltcore/sqle/globalstate/sequences/state.go index fe04b991acd..caefe6dc656 100644 --- a/go/libraries/doltcore/sqle/globalstate/sequences/state.go +++ b/go/libraries/doltcore/sqle/globalstate/sequences/state.go @@ -25,7 +25,7 @@ import ( // A SequenceState is an incrementing state that must be shared across all branches and transactions. // It corresponds to a table or root object in the database. // |Self| should always be the same type as the implementation. -// It produces a sequence of |ValueType| +// It produces a sequence of |ValueType| each time Next() is called. type SequenceState[Self any, ValueType comparable] interface { // Next advances the state of the sequence, producing a SQL value and the next SequenceState. Next() (sqlVal ValueType, hasNext bool, nextState Self, err error) @@ -55,8 +55,8 @@ type SequencedRelation[Self any, ValueType comparable, StateType SequenceState[S // HasSequenceState returns whether the relation wraps a sequence. // (This may be false, for instance, for tables that do not have an AUTO INCREMENT column) HasSequenceState(ctx context.Context) (bool, error) - // GetAutoIncrementSqlType returns the SQL type generated by the sequence. - GetAutoIncrementSqlType(ctx context.Context) (sql.Type, bool, error) + // GetSequenceSqlType returns the SQL type generated by the sequence. + GetSequenceSqlType(ctx context.Context) (sql.Type, bool, error) // SetSequenceState unconditionally sets the SequenceState for the object. SetSequenceState(ctx context.Context, val StateType) (Self, error) // TrySetSequenceState attempts to set the SequenceState for the object, but may fail if the provided state is diff --git a/go/libraries/doltcore/sqle/tables.go b/go/libraries/doltcore/sqle/tables.go index e705b7fe254..c4a04a220d7 100644 --- a/go/libraries/doltcore/sqle/tables.go +++ b/go/libraries/doltcore/sqle/tables.go @@ -455,7 +455,11 @@ func (t *DoltTable) PeekNextAutoIncrementValue(ctx *sql.Context) (uint64, error) if err != nil { return 0, err } - return table.GetAutoIncrementValue(ctx) + seq, err := table.GetAutoIncrementValue(ctx) + if err != nil { + return 0, err + } + return uint64(seq), err } // Name returns the name of the table. @@ -1570,11 +1574,14 @@ func (t *AlterableDoltTable) AddColumn(ctx *sql.Context, column *sql.Column, ord } if column.AutoIncrement { - ait, err := t.db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) + ait, err := dsess.GetAutoIncrementTracker(ctx, t.db.gs) + if err != nil { + return err + } + err = ait.AddNewTable(t.tableName, doltdb.AutoIncrementState(1)) if err != nil { return err } - ait.AddNewTable(t.tableName) } newRoot, err := root.PutTable(ctx, t.TableName(), updatedTable) @@ -1797,7 +1804,7 @@ func (t *AlterableDoltTable) RewriteInserter( // TODO: figure out locking. Other DBs automatically lock a table during this kind of operation, we should probably // do the same. We're messing with global auto-increment values here and it's not safe. - ait, err := t.db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) + ait, err := dsess.GetAutoIncrementTracker(ctx, t.db.gs) if err != nil { return nil, err } @@ -1854,7 +1861,7 @@ func fullTextRewriteEditor( // TODO: figure out locking. Other DBs automatically lock a table during this kind of operation, we should probably // do the same. We're messing with global auto-increment values here and it's not safe. - ait, err := t.db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) + ait, err := dsess.GetAutoIncrementTracker(ctx, t.db.gs) if err != nil { return nil, err } @@ -2271,21 +2278,24 @@ func (t *AlterableDoltTable) ModifyColumn(ctx *sql.Context, columnName string, c return err } - updatedTable, err = updatedTable.SetAutoIncrementValue(ctx, seq) + updatedTable, err = updatedTable.SetAutoIncrementValue(ctx, doltdb.AutoIncrementState(seq)) if err != nil { return err } - ait, err := t.db.gs.AutoIncrementTracker(ctx, dsess.AutoIncrementTrackerKey) + ait, err := dsess.GetAutoIncrementTracker(ctx, t.db.gs) if err != nil { return err } // TODO: this isn't transactional, and it should be (but none of the auto increment tracking is) - ait.AddNewTable(t.tableName) + err = ait.AddNewTable(t.tableName, doltdb.AutoIncrementState(1)) + if err != nil { + return err + } // Since this is a new auto increment table, we don't need to exclude the current working set from consideration // when computing its new sequence value, hence the empty ref - _, err = ait.Set(ctx, t.tableName, updatedTable, ref.WorkingSetRef{}, seq) + _, err = ait.Set(ctx, t.tableName, updatedTable, ref.WorkingSetRef{}, doltdb.AutoIncrementState(seq)) if err != nil { return err } @@ -2361,7 +2371,7 @@ func (t *AlterableDoltTable) getFirstAutoIncrementValue( } } - seq, err := dsess.CoerceAutoIncrementValue(ctx, initialValue) + seq, err := doltdb.CoerceAutoIncrementValue(ctx, initialValue) if err != nil { return 0, err } diff --git a/go/libraries/doltcore/sqle/temp_table.go b/go/libraries/doltcore/sqle/temp_table.go index 3e96e281ba2..6dc48d676e2 100644 --- a/go/libraries/doltcore/sqle/temp_table.go +++ b/go/libraries/doltcore/sqle/temp_table.go @@ -484,7 +484,11 @@ func temporaryDoltSchema(ctx context.Context, pkSch sql.PrimaryKeySchema, tags [ } func (t *TempTable) PeekNextAutoIncrementValue(ctx *sql.Context) (uint64, error) { - return t.table.GetAutoIncrementValue(ctx) + autoIncState, err := t.table.GetAutoIncrementValue(ctx) + if err != nil { + return 0, err + } + return autoIncState.CurrentValue(), nil } func (t *TempTable) GetNextAutoIncrementValue(ctx *sql.Context, insertVal interface{}) (uint64, error) { diff --git a/go/libraries/doltcore/sqle/writer/prolly_table_writer.go b/go/libraries/doltcore/sqle/writer/prolly_table_writer.go index cf3853ded0d..b66cc0d588a 100644 --- a/go/libraries/doltcore/sqle/writer/prolly_table_writer.go +++ b/go/libraries/doltcore/sqle/writer/prolly_table_writer.go @@ -23,7 +23,6 @@ import ( "github.com/dolthub/dolt/go/libraries/doltcore/doltdb/durable" "github.com/dolthub/dolt/go/libraries/doltcore/schema" "github.com/dolthub/dolt/go/libraries/doltcore/sqle/dsess" - "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate" "github.com/dolthub/dolt/go/libraries/doltcore/sqle/index" "github.com/dolthub/dolt/go/store/hash" "github.com/dolthub/dolt/go/store/pool" @@ -52,7 +51,7 @@ type prollyTableWriter struct { writeSess dsess.WriteSession aiCol schema.Column - aiTracker globalstate.AutoIncrementTracker + aiTracker *dsess.AutoIncrementTracker aiAlterVal uint64 aiAltered bool // True when an ALTER TABLE affects the auto increment value aiSet bool // True when an INSERT/UPDATE affects the auto increment value @@ -176,7 +175,6 @@ func (w *prollyTableWriter) Insert(ctx *sql.Context, sqlRow sql.Row) (err error) // TODO: need schema name in ai tracker w.aiSet = true - w.aiTracker.Next(ctx, w.tblName.Name, sqlRow) return nil } @@ -273,12 +271,16 @@ func (w *prollyTableWriter) PreciseMatch() bool { // GetNextAutoIncrementValue implements TableWriter. func (w *prollyTableWriter) GetNextAutoIncrementValue(ctx *sql.Context, insertVal interface{}) (uint64, error) { - return w.aiTracker.Next(ctx, w.tblName.Name, insertVal) + v, err := w.aiTracker.Next(ctx, w.tblName.Name, insertVal) + if err != nil { + return 0, err + } + return v, nil } // SetAutoIncrementValue implements AutoIncrementSetter. func (w *prollyTableWriter) SetAutoIncrementValue(ctx *sql.Context, val uint64) error { - seq, err := w.aiTracker.CoerceAutoIncrementValue(ctx, val) + seq, err := doltdb.CoerceAutoIncrementValue(ctx, val) if err != nil { return err } @@ -396,13 +398,12 @@ func (w *prollyTableWriter) table(ctx *sql.Context) (tbl *doltdb.Table, err erro if w.aiCol.AutoIncrement { if w.aiAltered { - tbl, err = w.aiTracker.Set(ctx, w.tblName.Name, tbl, w.writeSess.GetWorkingSet().Ref(), w.aiAlterVal) + tbl, err = w.aiTracker.Set(ctx, w.tblName.Name, tbl, w.writeSess.GetWorkingSet().Ref(), doltdb.AutoIncrementState(w.aiAlterVal)) if err != nil { return nil, err } } else if w.aiSet { - var aiVal uint64 - aiVal, err = w.aiTracker.Current(w.tblName.Name) + aiVal, err := w.aiTracker.Current(w.tblName.Name) if err != nil { return nil, err } diff --git a/go/libraries/doltcore/sqle/writer/prolly_write_session.go b/go/libraries/doltcore/sqle/writer/prolly_write_session.go index 6afd83c2802..579521447c1 100644 --- a/go/libraries/doltcore/sqle/writer/prolly_write_session.go +++ b/go/libraries/doltcore/sqle/writer/prolly_write_session.go @@ -23,7 +23,6 @@ import ( "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" "github.com/dolthub/dolt/go/libraries/doltcore/sqle/dsess" - "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate" "github.com/dolthub/dolt/go/libraries/doltcore/table/editor" "github.com/dolthub/dolt/go/store/hash" ) @@ -31,7 +30,7 @@ import ( // NewWriteSession creates and returns a WriteSession. Inserting a nil root is not an error, as there are // locations that do not have a root at the time of this call. However, a root must be set through SetWorkingRoot before any // table editors are returned. -func NewWriteSession(dbName string, ws *doltdb.WorkingSet, aiTracker globalstate.AutoIncrementTracker, setter dsess.SessionRootSetter, opts editor.Options) dsess.WriteSession { +func NewWriteSession(dbName string, ws *doltdb.WorkingSet, aiTracker *dsess.AutoIncrementTracker, setter dsess.SessionRootSetter, opts editor.Options) dsess.WriteSession { return &prollyWriteSession{ dbName: dbName, tables: make(map[doltdb.TableName]*prollyTableWriter), @@ -47,7 +46,7 @@ func NewWriteSession(dbName string, ws *doltdb.WorkingSet, aiTracker globalstate type prollyWriteSession struct { dbName string tables map[doltdb.TableName]*prollyTableWriter - aiTracker globalstate.AutoIncrementTracker + aiTracker *dsess.AutoIncrementTracker workingSet *doltdb.WorkingSet setter dsess.SessionRootSetter targetStaging bool From d419eb854bdc032480a9701461cace61f8294ac5 Mon Sep 17 00:00:00 2001 From: nicktobey Date: Wed, 22 Jul 2026 06:25:35 +0000 Subject: [PATCH 05/11] [ga-format-pr] Run go/utils/repofmt/format_repo.sh and go/Godeps/update.sh --- go/libraries/doltcore/doltdb/table.go | 8 ++++---- go/libraries/doltcore/sqle/dsess/globalstate.go | 1 + go/libraries/doltcore/sqle/globalstate/global_state.go | 3 ++- .../doltcore/sqle/globalstate/sequence_tracker.go | 2 +- go/libraries/doltcore/sqle/globalstate/sequences/state.go | 1 + 5 files changed, 9 insertions(+), 6 deletions(-) diff --git a/go/libraries/doltcore/doltdb/table.go b/go/libraries/doltcore/doltdb/table.go index 8f89a23e528..f76565bec93 100644 --- a/go/libraries/doltcore/doltdb/table.go +++ b/go/libraries/doltcore/doltdb/table.go @@ -18,16 +18,16 @@ import ( "context" "errors" "fmt" - "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate/sequences" - - "github.com/dolthub/go-mysql-server/sql" - gmstypes "github.com/dolthub/go-mysql-server/sql/types" "io" "math" "unicode" + "github.com/dolthub/go-mysql-server/sql" + gmstypes "github.com/dolthub/go-mysql-server/sql/types" + "github.com/dolthub/dolt/go/libraries/doltcore/doltdb/durable" "github.com/dolthub/dolt/go/libraries/doltcore/schema" + "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate/sequences" "github.com/dolthub/dolt/go/store/hash" "github.com/dolthub/dolt/go/store/prolly/tree" "github.com/dolthub/dolt/go/store/types" diff --git a/go/libraries/doltcore/sqle/dsess/globalstate.go b/go/libraries/doltcore/sqle/dsess/globalstate.go index 6001ef89e25..9ef51b4d555 100644 --- a/go/libraries/doltcore/sqle/dsess/globalstate.go +++ b/go/libraries/doltcore/sqle/dsess/globalstate.go @@ -16,6 +16,7 @@ package dsess import ( "context" + "github.com/dolthub/go-mysql-server/sql" "golang.org/x/sync/errgroup" diff --git a/go/libraries/doltcore/sqle/globalstate/global_state.go b/go/libraries/doltcore/sqle/globalstate/global_state.go index 5205c558da5..cf08580ef2d 100644 --- a/go/libraries/doltcore/sqle/globalstate/global_state.go +++ b/go/libraries/doltcore/sqle/globalstate/global_state.go @@ -15,8 +15,9 @@ package globalstate import ( - "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" "github.com/dolthub/go-mysql-server/sql" + + "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" ) // GlobalState is just a holding interface for pieces of global state, of which the auto increment tracking info is diff --git a/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go b/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go index 72c8ebf3755..790a891793c 100644 --- a/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go @@ -16,12 +16,12 @@ package globalstate import ( "context" - "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate/sequences" "github.com/dolthub/go-mysql-server/sql" "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" "github.com/dolthub/dolt/go/libraries/doltcore/ref" + "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate/sequences" ) // SequenceTrackerBase is the non-generic base interface for SequenceTracker diff --git a/go/libraries/doltcore/sqle/globalstate/sequences/state.go b/go/libraries/doltcore/sqle/globalstate/sequences/state.go index caefe6dc656..69d4e2d303a 100644 --- a/go/libraries/doltcore/sqle/globalstate/sequences/state.go +++ b/go/libraries/doltcore/sqle/globalstate/sequences/state.go @@ -16,6 +16,7 @@ package sequences import ( "context" + "github.com/dolthub/go-mysql-server/sql" ) From bcb1aafc0a0fb6c6cf62bc3ade86dc9e81daa05f Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Thu, 23 Jul 2026 12:24:15 -0700 Subject: [PATCH 06/11] Rename function to NewSequenceTracker --- go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go | 2 +- go/libraries/doltcore/sqle/dsess/sequence_tracker.go | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go index abd4e25b399..e91ded91467 100644 --- a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go @@ -45,7 +45,7 @@ type AutoIncrementTracker = SequenceTracker[*doltdb.Table, doltdb.AutoIncrementS // Roots provided should be the working sets when available, or the branches when they are not (e.g. for remote // branches that don't have a local working set) func NewAutoIncrementTracker(ctx context.Context, dbName string, roots ...doltdb.Rootish) (*AutoIncrementTracker, error) { - return NewAutoIncrementTrackerI(ctx, dbName, DoltDBRelationSource{}, roots...) + return NewSequenceTracker(ctx, dbName, DoltDBRelationSource{}, roots...) } func GetAutoIncrementTracker(ctx *sql.Context, gs globalstate.GlobalState) (*AutoIncrementTracker, error) { diff --git a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go index 4b031a33584..492c76950a3 100644 --- a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go @@ -89,7 +89,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) staticAssertTypes( var _ globalstate.SequenceTracker[RelationType, StateType, ValueType] = a } -func NewAutoIncrementTrackerI[ +func NewSequenceTracker[ RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], StateType sequences.SequenceState[StateType, ValueType], ValueType comparable, From ddaf2b84cadc554fbdf49e5ae92671f0320b8f83 Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Tue, 28 Jul 2026 17:12:23 -0700 Subject: [PATCH 07/11] Rename and refactor code for managing global state. --- .../sqle/dsess/auto_increment_tracker.go | 2 +- .../doltcore/sqle/dsess/globalstate.go | 23 ++++++++++++------- .../doltcore/sqle/dsess/sequence_tracker.go | 12 +++++----- 3 files changed, 22 insertions(+), 15 deletions(-) diff --git a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go index e91ded91467..d7f9fdfa271 100644 --- a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go @@ -45,7 +45,7 @@ type AutoIncrementTracker = SequenceTracker[*doltdb.Table, doltdb.AutoIncrementS // Roots provided should be the working sets when available, or the branches when they are not (e.g. for remote // branches that don't have a local working set) func NewAutoIncrementTracker(ctx context.Context, dbName string, roots ...doltdb.Rootish) (*AutoIncrementTracker, error) { - return NewSequenceTracker(ctx, dbName, DoltDBRelationSource{}, roots...) + return NewSequenceTrackerFromRoots(ctx, dbName, DoltDBRelationSource{}, roots...) } func GetAutoIncrementTracker(ctx *sql.Context, gs globalstate.GlobalState) (*AutoIncrementTracker, error) { diff --git a/go/libraries/doltcore/sqle/dsess/globalstate.go b/go/libraries/doltcore/sqle/dsess/globalstate.go index 9ef51b4d555..a2f153bd7c5 100644 --- a/go/libraries/doltcore/sqle/dsess/globalstate.go +++ b/go/libraries/doltcore/sqle/dsess/globalstate.go @@ -23,21 +23,25 @@ import ( "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" "github.com/dolthub/dolt/go/libraries/doltcore/ref" "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate" + "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate/sequences" ) // TrackerKey is a type used as a key into the GlobalStateImpl's map of SequenceTrackers type TrackerKey[TrackerType globalstate.SequenceTrackerBase] struct{} -// NewGlobalStateStoreForDb creates a new GlobalState. It initially contains a single SequenceTracker: the AutoIncrementTracker -func NewGlobalStateStoreForDb(ctx context.Context, dbName string, db *doltdb.DoltDB) (GlobalStateImpl, error) { +func NewSequenceTracker[ + RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], + StateType sequences.SequenceState[StateType, ValueType], + ValueType comparable, +](ctx context.Context, dbName string, db *doltdb.DoltDB, relationSource RelationSource[RelationType, StateType, ValueType]) (*SequenceTracker[RelationType, StateType, ValueType], error) { branches, err := db.GetBranches(ctx) if err != nil { - return GlobalStateImpl{}, err + return nil, err } remotes, err := db.GetRemoteRefs(ctx) if err != nil { - return GlobalStateImpl{}, err + return nil, err } rootRefs := make([]ref.DoltRef, 0, len(branches)+len(remotes)) @@ -88,16 +92,19 @@ func NewGlobalStateStoreForDb(ctx context.Context, dbName string, db *doltdb.Dol err = eg.Wait() if err != nil { - return GlobalStateImpl{}, err + return nil, err } - tracker, err := NewAutoIncrementTracker(ctx, dbName, roots...) + return NewSequenceTrackerFromRoots(ctx, dbName, relationSource, roots...) +} + +func NewGlobalStateStoreForDb(ctx context.Context, dbName string, db *doltdb.DoltDB) (GlobalStateImpl, error) { + autoIncrementTracker, err := NewSequenceTracker(ctx, dbName, db, DoltDBRelationSource{}) if err != nil { return GlobalStateImpl{}, err } - return GlobalStateImpl{ - sequenceTrackers: map[interface{}]globalstate.SequenceTrackerBase{autoIncrementTrackerKey: tracker}, + sequenceTrackers: map[interface{}]globalstate.SequenceTrackerBase{autoIncrementTrackerKey: autoIncrementTracker}, }, nil } diff --git a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go index 492c76950a3..53d6cbdfd4b 100644 --- a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go @@ -89,7 +89,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) staticAssertTypes( var _ globalstate.SequenceTracker[RelationType, StateType, ValueType] = a } -func NewSequenceTracker[ +func NewSequenceTrackerFromRoots[ RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], StateType sequences.SequenceState[StateType, ValueType], ValueType comparable, @@ -262,12 +262,12 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Cont if !currState.GreaterThan(givenState) { // Check if the given value is valid for this column type - if !a.validateAutoIncrementBounds(ctx, tbl, givenState, false) { + if !a.validateBounds(ctx, tbl, givenState, false) { return givenState.CurrentValue(), nil // Out of bounds, don't update sequence } // Value is valid, determine next sequence value - if a.validateAutoIncrementBounds(ctx, tbl, givenState, true) { + if a.validateBounds(ctx, tbl, givenState, true) { _, _, givenState, err = givenState.Next() if err != nil { return nextValue, err @@ -300,7 +300,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Set(ctx *sql.Conte return table.SetSequenceState(ctx, newSequenceState) } gt := newSequenceState.GreaterThan(existing) - if gt && a.validateAutoIncrementBounds(ctx, tableName, newSequenceState, true) { + if gt && a.validateBounds(ctx, tableName, newSequenceState, false) { a.sequences.Store(tableName, newSequenceState) return table.SetSequenceState(ctx, newSequenceState) } else if gt { @@ -420,7 +420,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.C } } - if a.validateAutoIncrementBounds(ctx, tableName, maxAutoInc, true) { + if a.validateBounds(ctx, tableName, maxAutoInc, false) { a.sequences.Store(tableName, maxAutoInc) } return table, nil @@ -592,7 +592,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initWithRoots(ctx } // validateAutoIncrementBounds checks if a value (or value+1 if checkIncrement) is valid for the auto-increment column type -func (a *SequenceTracker[RelationType, StateType, ValueType]) validateAutoIncrementBounds(ctx *sql.Context, tbl string, val StateType, checkIncrement bool) bool { +func (a *SequenceTracker[RelationType, StateType, ValueType]) validateBounds(ctx *sql.Context, tbl string, val StateType, checkIncrement bool) bool { sess := DSessFromSess(ctx.Session) db, ok := sess.Provider().BaseDatabase(ctx, a.dbName) if !ok || !db.Versioned() { From 2eee36f0d6ea9f0030cfbb1145a19944530861b5 Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Tue, 28 Jul 2026 19:47:06 -0700 Subject: [PATCH 08/11] Fix returning loaded sequence state. --- go/libraries/doltcore/sqle/dsess/sequence_tracker.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go index 53d6cbdfd4b..c3cbf59b4cd 100644 --- a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go @@ -168,7 +168,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAut return state, false, err } - seq, ok = loadSequenceState(a.sequences, tableName) + state, ok = loadSequenceState(a.sequences, tableName) if ok { return state, true, nil } From ef151f6c30118d1b71f8a354abdb9acef211fe44 Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Thu, 30 Jul 2026 14:33:38 -0700 Subject: [PATCH 09/11] Rename types and add documentation based on PR feedback. --- go/libraries/doltcore/sqle/database.go | 6 +- .../sqle/dsess/auto_increment_tracker.go | 4 +- .../doltcore/sqle/dsess/sequence_tracker.go | 88 ++++++++++--------- .../sqle/globalstate/sequence_tracker.go | 20 ++--- .../sqle/globalstate/sequences/doc.go | 18 ++++ .../sqle/globalstate/sequences/state.go | 3 - go/libraries/doltcore/sqle/tables.go | 4 +- .../sqle/writer/prolly_table_writer.go | 2 +- 8 files changed, 83 insertions(+), 62 deletions(-) create mode 100644 go/libraries/doltcore/sqle/globalstate/sequences/doc.go diff --git a/go/libraries/doltcore/sqle/database.go b/go/libraries/doltcore/sqle/database.go index b53d9a57541..bbaae3b165c 100644 --- a/go/libraries/doltcore/sqle/database.go +++ b/go/libraries/doltcore/sqle/database.go @@ -1929,7 +1929,7 @@ func (db Database) removeTableFromAutoIncrementTracker( return err } - err = ait.DropTable(ctx, tableName, wses...) + err = ait.DropRelation(ctx, tableName, wses...) if err != nil { return err } @@ -2088,7 +2088,7 @@ func (db Database) createSqlTable(ctx *sql.Context, table string, schemaName str if err != nil { return err } - err = ait.AddNewTable(tableName.Name, doltdb.AutoIncrementState(1)) + err = ait.AddNewRelation(tableName.Name, doltdb.AutoIncrementState(1)) if err != nil { return err } @@ -2151,7 +2151,7 @@ func (db Database) createIndexedSqlTable(ctx *sql.Context, table string, schemaN if err != nil { return err } - ait.AddNewTable(tableName.Name, doltdb.AutoIncrementState(1)) + ait.AddNewRelation(tableName.Name, doltdb.AutoIncrementState(1)) } return db.createDoltTable(ctx, tableName.Name, tableName.Schema, root, doltSch) diff --git a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go index d7f9fdfa271..94d3e14111d 100644 --- a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go @@ -1,4 +1,4 @@ -// Copyright 2023 Dolthub, Inc. +// Copyright 2026 Dolthub, Inc. // // Licensed under the Apache License, Version 2.0 (the "License"); // you may not use this file except in compliance with the License. @@ -52,5 +52,5 @@ func GetAutoIncrementTracker(ctx *sql.Context, gs globalstate.GlobalState) (*Aut return GetSequenceTracker(ctx, gs, autoIncrementTrackerKey) } -// autoIncrementTrackerKey is the key used to store the AutoIncrmenetTracker in the GlobalState's map of SequenceTrackers +// autoIncrementTrackerKey is the key used to store the AutoIncrementTracker in the GlobalState's map of SequenceTrackers var autoIncrementTrackerKey = TrackerKey[*AutoIncrementTracker]{} diff --git a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go index c3cbf59b4cd..af9e6378915 100644 --- a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go @@ -1,4 +1,4 @@ -// Copyright 2023 Dolthub, Inc. +// Copyright 2026 Dolthub, Inc. // // Licensed under the Apache License, Version 2.0 (the "License"); // you may not use this file except in compliance with the License. @@ -46,6 +46,7 @@ type RelationSource[ StateType sequences.SequenceState[StateType, ValueType], ValueType comparable, ] interface { + // GetRelation gets a relation at a specific doltdb.RootValue GetRelation(ctx context.Context, root doltdb.RootValue, tName doltdb.TableName) (relation RelationType, resolvedName string, found bool, err error) GetRelations(ctx context.Context, root doltdb.RootValue, cb func(doltdb.TableName, RelationType) (bool, error)) error } @@ -85,10 +86,15 @@ func currentLockMode() LockMode { return LockMode_Interleaved } +// staticAssertTypes contains compile-time assertions that SequenceTracker implements interfaces. +// It does not need to be called. func (a *SequenceTracker[RelationType, StateType, ValueType]) staticAssertTypes() { var _ globalstate.SequenceTracker[RelationType, StateType, ValueType] = a } +// NewSequenceTrackerFromRoots creates and initializes a new SequenceTracker by querying |relationSource| for +// objects that need to be globally tracked, and computing a single tracked global state for each object by merging +// the state at every root in |roots| func NewSequenceTrackerFromRoots[ RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], StateType sequences.SequenceState[StateType, ValueType], @@ -125,19 +131,19 @@ func getGCSafepointController(ctx context.Context) *gcctx.GCSafepointController return gcctx.GetGCSafepointController(ctx) } -func loadSequenceState[StateType sequences.SequenceState[StateType, ValueType], ValueType comparable](sequences *SyncMap[string, StateType], tableName string) (current StateType, hasCurrent bool) { - tableName = strings.ToLower(tableName) - return sequences.Load(tableName) +func loadSequenceState[StateType sequences.SequenceState[StateType, ValueType], ValueType comparable](sequences *SyncMap[string, StateType], relationName string) (current StateType, hasCurrent bool) { + relationName = strings.ToLower(relationName) + return sequences.Load(relationName) } -func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAutoIncrement(ctx *sql.Context, tableName string, initialValue interface{}) (state StateType, hasState bool, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAutoIncrement(ctx *sql.Context, relationName string, initialValue interface{}) (state StateType, hasState bool, err error) { sess := DSessFromSess(ctx.Session) ws, err := sess.WorkingSet(ctx, a.dbName) if err != nil { return state, false, err } - table, _, ok, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tableName}) + table, _, ok, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: relationName}) if err != nil || !ok { return state, false, err } @@ -163,12 +169,12 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAut } } - relation, err := a.deepSet(ctx, tableName, table, ws.Ref(), seq) + relation, err := a.deepSet(ctx, relationName, table, ws.Ref(), seq) if err != nil { return state, false, err } - state, ok = loadSequenceState(a.sequences, tableName) + state, ok = loadSequenceState(a.sequences, relationName) if ok { return state, true, nil } @@ -177,7 +183,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAut if err != nil { return state, false, err } - a.sequences.Store(strings.ToLower(tableName), seq) + a.sequences.Store(strings.ToLower(relationName), seq) return state, true, nil } @@ -186,7 +192,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Close() { <-a.init } -// Current returns the next value to be generated in the auto increment sequence for |tableName|. +// Current returns the next value to be generated in the auto increment sequence for |relationName|. func (a *SequenceTracker[RelationType, StateType, ValueType]) Current(relation string) (current StateType, err error) { err = a.waitForInit() if err != nil { @@ -283,37 +289,37 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Cont // Set sets the auto increment value for the table named, if it's greater than the one already registered for this // table. Otherwise, the update is silently disregarded. So far this matches the MySQL behavior, but Dolt uses the // maximum value for this table across all branches. -func (a *SequenceTracker[RelationType, StateType, ValueType]) Set(ctx *sql.Context, tableName string, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (newRelation RelationType, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) Set(ctx *sql.Context, relationName string, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (newRelation RelationType, err error) { err = a.waitForInit() if err != nil { return newRelation, err } - tableName = strings.ToLower(tableName) + relationName = strings.ToLower(relationName) - release := a.mm.Lock(tableName) + release := a.mm.Lock(relationName) defer release() - existing, ok := loadSequenceState(a.sequences, tableName) + existing, ok := loadSequenceState(a.sequences, relationName) if !ok { - a.sequences.Store(tableName, newSequenceState) + a.sequences.Store(relationName, newSequenceState) return table.SetSequenceState(ctx, newSequenceState) } gt := newSequenceState.GreaterThan(existing) - if gt && a.validateBounds(ctx, tableName, newSequenceState, false) { - a.sequences.Store(tableName, newSequenceState) + if gt && a.validateBounds(ctx, relationName, newSequenceState, false) { + a.sequences.Store(relationName, newSequenceState) return table.SetSequenceState(ctx, newSequenceState) } else if gt { // Value is greater but out of bounds, don't update return table, nil } // Value is not greater than current, do deep check across branches - return a.deepSet(ctx, tableName, table, ws, newSequenceState) + return a.deepSet(ctx, relationName, table, ws, newSequenceState) } // deepSet sets the sequence state for the table named, if it's greater than the one on any branch head for this // database, ignoring the current in-memory tracker value -func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.Context, tableName string, table RelationType, ws ref.WorkingSetRef, newAutoIncVal StateType) (newRelation RelationType, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.Context, relationName string, table RelationType, ws ref.WorkingSetRef, newAutoIncVal StateType) (newRelation RelationType, err error) { sess := DSessFromSess(ctx.Session) db, ok := sess.Provider().BaseDatabase(ctx, a.dbName) @@ -388,7 +394,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.C return newRelation, err } - table, _, ok, err := a.relationSource.GetRelation(ctx, root, doltdb.TableName{Name: tableName}) + table, _, ok, err := a.relationSource.GetRelation(ctx, root, doltdb.TableName{Name: relationName}) if err != nil { return newRelation, err } @@ -405,7 +411,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.C continue } - tableName = strings.ToLower(tableName) + relationName = strings.ToLower(relationName) seq, err := table.GetSequenceState(ctx) if err != nil { return newRelation, err @@ -420,44 +426,44 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.C } } - if a.validateBounds(ctx, tableName, maxAutoInc, false) { - a.sequences.Store(tableName, maxAutoInc) + if a.validateBounds(ctx, relationName, maxAutoInc, false) { + a.sequences.Store(relationName, maxAutoInc) } return table, nil } -// AddNewTable initializes a new table with an auto increment column to the tracker, as necessary -func (a *SequenceTracker[RelationType, StateType, ValueType]) AddNewTable(tableName string, initialState StateType) error { +// AddNewRelation initializes a new table with an auto increment column to the tracker, as necessary +func (a *SequenceTracker[RelationType, StateType, ValueType]) AddNewRelation(relationName string, initialState StateType) error { err := a.waitForInit() if err != nil { return err } - tableName = strings.ToLower(tableName) + relationName = strings.ToLower(relationName) // only initialize the sequence for this table if no other branch has such a table - a.sequences.LoadOrStore(tableName, initialState) + a.sequences.LoadOrStore(relationName, initialState) return nil } -// DropTable drops the table with the name given. +// DropRelation drops the table with the name given. // To establish the new auto increment value, callers must also pass all other working sets in scope that may include // a table with the same name, omitting the working set that just deleted the table named. -func (a *SequenceTracker[RelationType, StateType, ValueType]) DropTable(ctx *sql.Context, tableName string, wses ...*doltdb.WorkingSet) error { +func (a *SequenceTracker[RelationType, StateType, ValueType]) DropRelation(ctx *sql.Context, relationName string, wses ...*doltdb.WorkingSet) error { err := a.waitForInit() if err != nil { return err } - tableName = strings.ToLower(tableName) + relationName = strings.ToLower(relationName) - release := a.mm.Lock(tableName) + release := a.mm.Lock(relationName) defer release() var newHighestValue *StateType // Get the new highest value from all tables in the working sets given for _, ws := range wses { - table, _, exists, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tableName}) + table, _, exists, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: relationName}) if err != nil { return err } @@ -490,15 +496,15 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) DropTable(ctx *sql } if newHighestValue != nil { - a.sequences.Store(tableName, *newHighestValue) + a.sequences.Store(relationName, *newHighestValue) } else { - a.sequences.Delete(tableName) + a.sequences.Delete(relationName) } return nil } -func (a *SequenceTracker[RelationType, StateType, ValueType]) AcquireTableLock(ctx *sql.Context, tableName string) (func(), error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) AcquireLock(ctx *sql.Context, relationName string) (func(), error) { err := a.waitForInit() if err != nil { return nil, err @@ -508,7 +514,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) AcquireTableLock(c // This shouldn't be possible, it's a serious programming error if it happens panic("Attempted to acquire AutoInc lock for entire insert operation, but lock mode was set to Interleaved") } - return a.mm.Lock(tableName), nil + return a.mm.Lock(relationName), nil } func (a *SequenceTracker[RelationType, StateType, ValueType]) waitForInit() error { @@ -560,7 +566,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initWithRoots(ctx return err } - init := func(tableName doltdb.TableName, relation RelationType) (bool, error) { + init := func(relationName doltdb.TableName, relation RelationType) (bool, error) { hasSequenceState, err := relation.HasSequenceState(ctx) if err != nil { return true, err @@ -573,10 +579,10 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initWithRoots(ctx return true, err } - tableNameStr := tableName.ToLower().Name - if oldValue, loaded := a.sequences.LoadOrStore(tableNameStr, seq); loaded { - for seq.GreaterThan(oldValue) && !a.sequences.CompareAndSwap(tableNameStr, oldValue, seq) { - oldValue, _ = a.sequences.Load(tableNameStr) + relationNameStr := relationName.ToLower().Name + if oldValue, loaded := a.sequences.LoadOrStore(relationNameStr, seq); loaded { + for seq.GreaterThan(oldValue) && !a.sequences.CompareAndSwap(relationNameStr, oldValue, seq) { + oldValue, _ = a.sequences.Load(relationNameStr) } } diff --git a/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go b/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go index 790a891793c..294e31de095 100644 --- a/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go @@ -1,4 +1,4 @@ -// Copyright 2021 Dolthub, Inc. +// Copyright 2026 Dolthub, Inc. // // Licensed under the Apache License, Version 2.0 (the "License"); // you may not use this file except in compliance with the License. @@ -27,18 +27,18 @@ import ( // SequenceTrackerBase is the non-generic base interface for SequenceTracker // It is useful for establishing an upper bound on type parameters in some circumstances. type SequenceTrackerBase interface { - // AcquireTableLock acquires the auto increment lock on a relation, and returns a callback function to release the lock. + // AcquireLock acquires the lock on the global state for a specific relation, and returns a callback function to release the lock. // Depending on the value of the `innodb_autoinc_lock_mode` system variable, the engine may need to acquire and hold // the lock for the duration of an insert statement. - AcquireTableLock(ctx *sql.Context, tableName string) (func(), error) - // DropTable removes a relation from the tracker. - DropTable(ctx *sql.Context, relation string, wses ...*doltdb.WorkingSet) error + AcquireLock(ctx *sql.Context, tableName string) (func(), error) + // DropRelation removes a relation from the tracker. + DropRelation(ctx *sql.Context, relation string, wses ...*doltdb.WorkingSet) error // InitWithRoots fills the SequenceTracker with values pulled from each root in order. InitWithRoots(ctx context.Context, roots ...doltdb.Rootish) error Close() } -// SequenceTracker knows how to get and set the current auto increment value for a relation (a table or a root object). +// SequenceTracker knows how to get and set the current sequence state for a relation (a table or a root object). // It's defined as an interface here because implementations need to reach into session state, requiring a dependency on this package. type SequenceTracker[ RelationType sequences.SequencedRelation[RelationType, ValueType, StateType], @@ -46,13 +46,13 @@ type SequenceTracker[ ValueType comparable, ] interface { SequenceTrackerBase - // Current returns the current auto increment state for the given relation. + // Current returns the current sequence state for the given relation. Current(relation string) (StateType, error) // Next returns the next SQL value produced by the given relation, and advances that relation's state. Next(ctx *sql.Context, relation string, insertVal interface{}) (ValueType, error) - // AddNewTable adds a new table to the tracker, initializing the auto increment value to the provided |initialState|. - AddNewTable(relation string, initialState StateType) error - // Set sets the auto increment value for the given relation. This operation may silently do nothing if this value is + // AddNewRelation adds a new table to the tracker, initializing the sequence state to the provided |initialState|. + AddNewRelation(relation string, initialState StateType) error + // Set sets the sequence state for the given relation. This operation may silently do nothing if this value is // below the current value for this relation. The relation in the provided working set is assumed to already have the value // given, so the new global maximum is computed without regard for its value in that working set. Set(ctx *sql.Context, tableName string, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (RelationType, error) diff --git a/go/libraries/doltcore/sqle/globalstate/sequences/doc.go b/go/libraries/doltcore/sqle/globalstate/sequences/doc.go new file mode 100644 index 00000000000..9cf7290992b --- /dev/null +++ b/go/libraries/doltcore/sqle/globalstate/sequences/doc.go @@ -0,0 +1,18 @@ +// Copyright 2026 Dolthub, Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// This package is a collection of interfaces that describe state machines that are controlled by the global state. +// It is a separate package from the parent package globalstate so that doltdb can depend on it. + +package sequences diff --git a/go/libraries/doltcore/sqle/globalstate/sequences/state.go b/go/libraries/doltcore/sqle/globalstate/sequences/state.go index 69d4e2d303a..026475bcda7 100644 --- a/go/libraries/doltcore/sqle/globalstate/sequences/state.go +++ b/go/libraries/doltcore/sqle/globalstate/sequences/state.go @@ -20,9 +20,6 @@ import ( "github.com/dolthub/go-mysql-server/sql" ) -// This package is a collection of interfaces that describe state machines that are controlled by the global state. -// It is a separate package from the parent package globalstate so that doltdb can depend on it. - // A SequenceState is an incrementing state that must be shared across all branches and transactions. // It corresponds to a table or root object in the database. // |Self| should always be the same type as the implementation. diff --git a/go/libraries/doltcore/sqle/tables.go b/go/libraries/doltcore/sqle/tables.go index c4a04a220d7..5fd6ea3bb78 100644 --- a/go/libraries/doltcore/sqle/tables.go +++ b/go/libraries/doltcore/sqle/tables.go @@ -1578,7 +1578,7 @@ func (t *AlterableDoltTable) AddColumn(ctx *sql.Context, column *sql.Column, ord if err != nil { return err } - err = ait.AddNewTable(t.tableName, doltdb.AutoIncrementState(1)) + err = ait.AddNewRelation(t.tableName, doltdb.AutoIncrementState(1)) if err != nil { return err } @@ -2289,7 +2289,7 @@ func (t *AlterableDoltTable) ModifyColumn(ctx *sql.Context, columnName string, c } // TODO: this isn't transactional, and it should be (but none of the auto increment tracking is) - err = ait.AddNewTable(t.tableName, doltdb.AutoIncrementState(1)) + err = ait.AddNewRelation(t.tableName, doltdb.AutoIncrementState(1)) if err != nil { return err } diff --git a/go/libraries/doltcore/sqle/writer/prolly_table_writer.go b/go/libraries/doltcore/sqle/writer/prolly_table_writer.go index b66cc0d588a..7cc16c5ccc5 100644 --- a/go/libraries/doltcore/sqle/writer/prolly_table_writer.go +++ b/go/libraries/doltcore/sqle/writer/prolly_table_writer.go @@ -294,7 +294,7 @@ func (w *prollyTableWriter) SetAutoIncrementValue(ctx *sql.Context, val uint64) // AcquireAutoIncrementLock implements AutoIncrementSetter. func (w *prollyTableWriter) AcquireAutoIncrementLock(ctx *sql.Context) (func(), error) { - return w.aiTracker.AcquireTableLock(ctx, w.tblName.Name) + return w.aiTracker.AcquireLock(ctx, w.tblName.Name) } // Close implements Closer From 4f7465e034d24ea8b0df7efe8672cb9fe6c016ac Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Thu, 30 Jul 2026 15:45:11 -0700 Subject: [PATCH 10/11] Change SequenceTracker to take doltdb.TableName instead of string parameters. --- go/libraries/doltcore/sqle/database.go | 11 ++- .../doltcore/sqle/dsess/sequence_tracker.go | 80 +++++++++---------- .../sqle/globalstate/sequence_tracker.go | 12 +-- go/libraries/doltcore/sqle/tables.go | 10 +-- .../sqle/writer/prolly_table_writer.go | 8 +- 5 files changed, 60 insertions(+), 61 deletions(-) diff --git a/go/libraries/doltcore/sqle/database.go b/go/libraries/doltcore/sqle/database.go index bbaae3b165c..194dc85ff90 100644 --- a/go/libraries/doltcore/sqle/database.go +++ b/go/libraries/doltcore/sqle/database.go @@ -1877,7 +1877,7 @@ func (db Database) dropTable(ctx *sql.Context, tableName string) error { if schema.HasAutoIncrement(sch) { ddb, _ := ds.GetDoltDB(ctx, db.RevisionQualifiedName()) - err = db.removeTableFromAutoIncrementTracker(ctx, tableName, ddb, ws.Ref()) + err = db.removeTableFromAutoIncrementTracker(ctx, tblName, ddb, ws.Ref()) if err != nil { return err } @@ -1892,7 +1892,7 @@ func (db Database) dropTable(ctx *sql.Context, tableName string) error { // otherwise. This operation is expensive if the func (db Database) removeTableFromAutoIncrementTracker( ctx *sql.Context, - tableName string, + tableName doltdb.TableName, ddb *doltdb.DoltDB, ws ref.WorkingSetRef, ) error { @@ -2088,7 +2088,7 @@ func (db Database) createSqlTable(ctx *sql.Context, table string, schemaName str if err != nil { return err } - err = ait.AddNewRelation(tableName.Name, doltdb.AutoIncrementState(1)) + err = ait.AddNewRelation(tableName, doltdb.AutoIncrementState(1)) if err != nil { return err } @@ -2151,7 +2151,10 @@ func (db Database) createIndexedSqlTable(ctx *sql.Context, table string, schemaN if err != nil { return err } - ait.AddNewRelation(tableName.Name, doltdb.AutoIncrementState(1)) + err = ait.AddNewRelation(tableName, doltdb.AutoIncrementState(1)) + if err != nil { + return err + } } return db.createDoltTable(ctx, tableName.Name, tableName.Schema, root, doltSch) diff --git a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go index af9e6378915..eb2d4c28748 100644 --- a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go @@ -18,7 +18,6 @@ import ( "context" "errors" "fmt" - "strings" "time" "github.com/dolthub/go-mysql-server/sql" @@ -57,7 +56,7 @@ type SequenceTracker[ ValueType comparable, ] struct { initErr error - sequences *SyncMap[string, StateType] + sequences *SyncMap[doltdb.TableName, StateType] mm *mutexmap.MutexMap // SequenceTracker is lazily initialized by loading // tracker state for every given |root|. On first access, we @@ -102,7 +101,7 @@ func NewSequenceTrackerFromRoots[ ](ctx context.Context, dbName string, relationSource RelationSource[RelationType, StateType, ValueType], roots ...doltdb.Rootish) (*SequenceTracker[RelationType, StateType, ValueType], error) { ait := SequenceTracker[RelationType, StateType, ValueType]{ dbName: dbName, - sequences: &SyncMap[string, StateType]{}, + sequences: &SyncMap[doltdb.TableName, StateType]{}, mm: mutexmap.NewMutexMap(), init: make(chan struct{}), cancelInit: make(chan struct{}), @@ -131,19 +130,18 @@ func getGCSafepointController(ctx context.Context) *gcctx.GCSafepointController return gcctx.GetGCSafepointController(ctx) } -func loadSequenceState[StateType sequences.SequenceState[StateType, ValueType], ValueType comparable](sequences *SyncMap[string, StateType], relationName string) (current StateType, hasCurrent bool) { - relationName = strings.ToLower(relationName) - return sequences.Load(relationName) +func loadSequenceState[StateType sequences.SequenceState[StateType, ValueType], ValueType comparable](sequences *SyncMap[doltdb.TableName, StateType], relationName doltdb.TableName) (current StateType, hasCurrent bool) { + return sequences.Load(relationName.ToLower()) } -func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAutoIncrement(ctx *sql.Context, relationName string, initialValue interface{}) (state StateType, hasState bool, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAutoIncrement(ctx *sql.Context, relationName doltdb.TableName, initialValue interface{}) (state StateType, hasState bool, err error) { sess := DSessFromSess(ctx.Session) ws, err := sess.WorkingSet(ctx, a.dbName) if err != nil { return state, false, err } - table, _, ok, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: relationName}) + table, _, ok, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), relationName) if err != nil || !ok { return state, false, err } @@ -183,7 +181,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAut if err != nil { return state, false, err } - a.sequences.Store(strings.ToLower(relationName), seq) + a.sequences.Store(relationName.ToLower(), seq) return state, true, nil } @@ -193,59 +191,59 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Close() { } // Current returns the next value to be generated in the auto increment sequence for |relationName|. -func (a *SequenceTracker[RelationType, StateType, ValueType]) Current(relation string) (current StateType, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) Current(relationName doltdb.TableName) (current StateType, err error) { err = a.waitForInit() if err != nil { return current, err } - seq, ok := loadSequenceState(a.sequences, relation) + seq, ok := loadSequenceState(a.sequences, relationName) if !ok { return current, nil } return seq, nil } -// Next returns the next auto increment value for |tbl| using |insertVal| from an insert. If |insertVal| is +// Next returns the next auto increment value for |relationName| using |insertVal| from an insert. If |insertVal| is // null or 0, it is generated from the sequence. -func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Context, tbl string, insertVal interface{}) (nextValue ValueType, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Context, relationName doltdb.TableName, insertVal interface{}) (nextValue ValueType, err error) { err = a.waitForInit() if err != nil { return nextValue, err } - tbl = strings.ToLower(tbl) + relationName = relationName.ToLower() // The read-modify-write of the sequence below must be atomic across concurrent inserters. In // interleaved lock mode (the default) the engine holds no statement-level lock, so we take a // short per-table lock here. locked := false if a.lockMode == LockMode_Interleaved { - release := a.mm.Lock(tbl) + release := a.mm.Lock(relationName) defer release() locked = true } - currState, ok := loadSequenceState(a.sequences, tbl) + currState, ok := loadSequenceState(a.sequences, relationName) if !ok { // Missing tracker state after initialization can happen when a running sql-server discovers a database // restored after startup, so initialize it here. if !locked { if a.lockMode == LockMode_Interleaved { - release := a.mm.Lock(tbl) + release := a.mm.Lock(relationName) defer release() locked = true } - currState, ok = loadSequenceState(a.sequences, tbl) + currState, ok = loadSequenceState(a.sequences, relationName) } if !ok { - currState, ok, err = a.initializeTableAutoIncrement(ctx, tbl, insertVal) + currState, ok, err = a.initializeTableAutoIncrement(ctx, relationName, insertVal) if err != nil { return nextValue, err } if !ok { - return nextValue, fmt.Errorf("autoIncrementTracker: unable to find sequence for table %s", tbl) + return nextValue, fmt.Errorf("autoIncrementTracker: unable to find sequence for table %s", relationName.Name) } } } @@ -256,7 +254,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Cont if err != nil { return nextValue, err } - a.sequences.Store(tbl, nextState) + a.sequences.Store(relationName, nextState) return currentVal, nil } @@ -268,18 +266,18 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Cont if !currState.GreaterThan(givenState) { // Check if the given value is valid for this column type - if !a.validateBounds(ctx, tbl, givenState, false) { + if !a.validateBounds(ctx, relationName, givenState, false) { return givenState.CurrentValue(), nil // Out of bounds, don't update sequence } // Value is valid, determine next sequence value - if a.validateBounds(ctx, tbl, givenState, true) { + if a.validateBounds(ctx, relationName, givenState, true) { _, _, givenState, err = givenState.Next() if err != nil { return nextValue, err } } - a.sequences.Store(tbl, givenState) + a.sequences.Store(relationName, givenState) return given, nil } @@ -289,13 +287,13 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Cont // Set sets the auto increment value for the table named, if it's greater than the one already registered for this // table. Otherwise, the update is silently disregarded. So far this matches the MySQL behavior, but Dolt uses the // maximum value for this table across all branches. -func (a *SequenceTracker[RelationType, StateType, ValueType]) Set(ctx *sql.Context, relationName string, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (newRelation RelationType, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) Set(ctx *sql.Context, relationName doltdb.TableName, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (newRelation RelationType, err error) { err = a.waitForInit() if err != nil { return newRelation, err } - relationName = strings.ToLower(relationName) + relationName = relationName.ToLower() release := a.mm.Lock(relationName) defer release() @@ -319,7 +317,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Set(ctx *sql.Conte // deepSet sets the sequence state for the table named, if it's greater than the one on any branch head for this // database, ignoring the current in-memory tracker value -func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.Context, relationName string, table RelationType, ws ref.WorkingSetRef, newAutoIncVal StateType) (newRelation RelationType, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.Context, relationName doltdb.TableName, table RelationType, ws ref.WorkingSetRef, newAutoIncVal StateType) (newRelation RelationType, err error) { sess := DSessFromSess(ctx.Session) db, ok := sess.Provider().BaseDatabase(ctx, a.dbName) @@ -394,7 +392,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.C return newRelation, err } - table, _, ok, err := a.relationSource.GetRelation(ctx, root, doltdb.TableName{Name: relationName}) + table, _, ok, err := a.relationSource.GetRelation(ctx, root, relationName) if err != nil { return newRelation, err } @@ -411,7 +409,6 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.C continue } - relationName = strings.ToLower(relationName) seq, err := table.GetSequenceState(ctx) if err != nil { return newRelation, err @@ -433,28 +430,27 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) deepSet(ctx *sql.C } // AddNewRelation initializes a new table with an auto increment column to the tracker, as necessary -func (a *SequenceTracker[RelationType, StateType, ValueType]) AddNewRelation(relationName string, initialState StateType) error { +func (a *SequenceTracker[RelationType, StateType, ValueType]) AddNewRelation(relationName doltdb.TableName, initialState StateType) error { err := a.waitForInit() if err != nil { return err } - relationName = strings.ToLower(relationName) // only initialize the sequence for this table if no other branch has such a table - a.sequences.LoadOrStore(relationName, initialState) + a.sequences.LoadOrStore(relationName.ToLower(), initialState) return nil } // DropRelation drops the table with the name given. // To establish the new auto increment value, callers must also pass all other working sets in scope that may include // a table with the same name, omitting the working set that just deleted the table named. -func (a *SequenceTracker[RelationType, StateType, ValueType]) DropRelation(ctx *sql.Context, relationName string, wses ...*doltdb.WorkingSet) error { +func (a *SequenceTracker[RelationType, StateType, ValueType]) DropRelation(ctx *sql.Context, relationName doltdb.TableName, wses ...*doltdb.WorkingSet) error { err := a.waitForInit() if err != nil { return err } - relationName = strings.ToLower(relationName) + relationName = relationName.ToLower() release := a.mm.Lock(relationName) defer release() @@ -463,7 +459,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) DropRelation(ctx * // Get the new highest value from all tables in the working sets given for _, ws := range wses { - table, _, exists, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: relationName}) + table, _, exists, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), relationName) if err != nil { return err } @@ -504,7 +500,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) DropRelation(ctx * return nil } -func (a *SequenceTracker[RelationType, StateType, ValueType]) AcquireLock(ctx *sql.Context, relationName string) (func(), error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) AcquireLock(ctx *sql.Context, relationName doltdb.TableName) (func(), error) { err := a.waitForInit() if err != nil { return nil, err @@ -579,10 +575,10 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initWithRoots(ctx return true, err } - relationNameStr := relationName.ToLower().Name - if oldValue, loaded := a.sequences.LoadOrStore(relationNameStr, seq); loaded { - for seq.GreaterThan(oldValue) && !a.sequences.CompareAndSwap(relationNameStr, oldValue, seq) { - oldValue, _ = a.sequences.Load(relationNameStr) + key := relationName.ToLower() + if oldValue, loaded := a.sequences.LoadOrStore(key, seq); loaded { + for seq.GreaterThan(oldValue) && !a.sequences.CompareAndSwap(key, oldValue, seq) { + oldValue, _ = a.sequences.Load(key) } } @@ -598,7 +594,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initWithRoots(ctx } // validateAutoIncrementBounds checks if a value (or value+1 if checkIncrement) is valid for the auto-increment column type -func (a *SequenceTracker[RelationType, StateType, ValueType]) validateBounds(ctx *sql.Context, tbl string, val StateType, checkIncrement bool) bool { +func (a *SequenceTracker[RelationType, StateType, ValueType]) validateBounds(ctx *sql.Context, relationName doltdb.TableName, val StateType, checkIncrement bool) bool { sess := DSessFromSess(ctx.Session) db, ok := sess.Provider().BaseDatabase(ctx, a.dbName) if !ok || !db.Versioned() { @@ -610,7 +606,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) validateBounds(ctx return true } - table, _, ok, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), doltdb.TableName{Name: tbl}) + table, _, ok, err := a.relationSource.GetRelation(ctx, ws.WorkingRoot(), relationName) if err != nil || !ok { return true } diff --git a/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go b/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go index 294e31de095..454c61b9e04 100644 --- a/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/globalstate/sequence_tracker.go @@ -30,9 +30,9 @@ type SequenceTrackerBase interface { // AcquireLock acquires the lock on the global state for a specific relation, and returns a callback function to release the lock. // Depending on the value of the `innodb_autoinc_lock_mode` system variable, the engine may need to acquire and hold // the lock for the duration of an insert statement. - AcquireLock(ctx *sql.Context, tableName string) (func(), error) + AcquireLock(ctx *sql.Context, tableName doltdb.TableName) (func(), error) // DropRelation removes a relation from the tracker. - DropRelation(ctx *sql.Context, relation string, wses ...*doltdb.WorkingSet) error + DropRelation(ctx *sql.Context, tableName doltdb.TableName, wses ...*doltdb.WorkingSet) error // InitWithRoots fills the SequenceTracker with values pulled from each root in order. InitWithRoots(ctx context.Context, roots ...doltdb.Rootish) error Close() @@ -47,13 +47,13 @@ type SequenceTracker[ ] interface { SequenceTrackerBase // Current returns the current sequence state for the given relation. - Current(relation string) (StateType, error) + Current(tableName doltdb.TableName) (StateType, error) // Next returns the next SQL value produced by the given relation, and advances that relation's state. - Next(ctx *sql.Context, relation string, insertVal interface{}) (ValueType, error) + Next(ctx *sql.Context, tableName doltdb.TableName, insertVal interface{}) (ValueType, error) // AddNewRelation adds a new table to the tracker, initializing the sequence state to the provided |initialState|. - AddNewRelation(relation string, initialState StateType) error + AddNewRelation(tableName doltdb.TableName, initialState StateType) error // Set sets the sequence state for the given relation. This operation may silently do nothing if this value is // below the current value for this relation. The relation in the provided working set is assumed to already have the value // given, so the new global maximum is computed without regard for its value in that working set. - Set(ctx *sql.Context, tableName string, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (RelationType, error) + Set(ctx *sql.Context, tableName doltdb.TableName, table RelationType, ws ref.WorkingSetRef, newSequenceState StateType) (RelationType, error) } diff --git a/go/libraries/doltcore/sqle/tables.go b/go/libraries/doltcore/sqle/tables.go index 5fd6ea3bb78..365fae4f32f 100644 --- a/go/libraries/doltcore/sqle/tables.go +++ b/go/libraries/doltcore/sqle/tables.go @@ -1073,7 +1073,7 @@ func (t *WritableDoltTable) truncate( if schema.HasAutoIncrement(sch) { ddb, _ := sess.GetDoltDB(ctx, t.db.RevisionQualifiedName()) - err = t.db.removeTableFromAutoIncrementTracker(ctx, t.Name(), ddb, ws.Ref()) + err = t.db.removeTableFromAutoIncrementTracker(ctx, t.TableName(), ddb, ws.Ref()) if err != nil { return nil, err } @@ -1578,7 +1578,7 @@ func (t *AlterableDoltTable) AddColumn(ctx *sql.Context, column *sql.Column, ord if err != nil { return err } - err = ait.AddNewRelation(t.tableName, doltdb.AutoIncrementState(1)) + err = ait.AddNewRelation(t.TableName(), doltdb.AutoIncrementState(1)) if err != nil { return err } @@ -2289,13 +2289,13 @@ func (t *AlterableDoltTable) ModifyColumn(ctx *sql.Context, columnName string, c } // TODO: this isn't transactional, and it should be (but none of the auto increment tracking is) - err = ait.AddNewRelation(t.tableName, doltdb.AutoIncrementState(1)) + err = ait.AddNewRelation(t.TableName(), doltdb.AutoIncrementState(1)) if err != nil { return err } // Since this is a new auto increment table, we don't need to exclude the current working set from consideration // when computing its new sequence value, hence the empty ref - _, err = ait.Set(ctx, t.tableName, updatedTable, ref.WorkingSetRef{}, doltdb.AutoIncrementState(seq)) + _, err = ait.Set(ctx, t.TableName(), updatedTable, ref.WorkingSetRef{}, doltdb.AutoIncrementState(seq)) if err != nil { return err } @@ -2306,7 +2306,7 @@ func (t *AlterableDoltTable) ModifyColumn(ctx *sql.Context, columnName string, c // TODO: this isn't transactional, and it should be sess := dsess.DSessFromSess(ctx.Session) ddb, _ := sess.GetDoltDB(ctx, t.db.RevisionQualifiedName()) - err = t.db.removeTableFromAutoIncrementTracker(ctx, t.Name(), ddb, ws.Ref()) + err = t.db.removeTableFromAutoIncrementTracker(ctx, t.TableName(), ddb, ws.Ref()) if err != nil { return err } diff --git a/go/libraries/doltcore/sqle/writer/prolly_table_writer.go b/go/libraries/doltcore/sqle/writer/prolly_table_writer.go index 7cc16c5ccc5..e633ce6e98b 100644 --- a/go/libraries/doltcore/sqle/writer/prolly_table_writer.go +++ b/go/libraries/doltcore/sqle/writer/prolly_table_writer.go @@ -271,7 +271,7 @@ func (w *prollyTableWriter) PreciseMatch() bool { // GetNextAutoIncrementValue implements TableWriter. func (w *prollyTableWriter) GetNextAutoIncrementValue(ctx *sql.Context, insertVal interface{}) (uint64, error) { - v, err := w.aiTracker.Next(ctx, w.tblName.Name, insertVal) + v, err := w.aiTracker.Next(ctx, w.tblName, insertVal) if err != nil { return 0, err } @@ -294,7 +294,7 @@ func (w *prollyTableWriter) SetAutoIncrementValue(ctx *sql.Context, val uint64) // AcquireAutoIncrementLock implements AutoIncrementSetter. func (w *prollyTableWriter) AcquireAutoIncrementLock(ctx *sql.Context) (func(), error) { - return w.aiTracker.AcquireLock(ctx, w.tblName.Name) + return w.aiTracker.AcquireLock(ctx, w.tblName) } // Close implements Closer @@ -398,12 +398,12 @@ func (w *prollyTableWriter) table(ctx *sql.Context) (tbl *doltdb.Table, err erro if w.aiCol.AutoIncrement { if w.aiAltered { - tbl, err = w.aiTracker.Set(ctx, w.tblName.Name, tbl, w.writeSess.GetWorkingSet().Ref(), doltdb.AutoIncrementState(w.aiAlterVal)) + tbl, err = w.aiTracker.Set(ctx, w.tblName, tbl, w.writeSess.GetWorkingSet().Ref(), doltdb.AutoIncrementState(w.aiAlterVal)) if err != nil { return nil, err } } else if w.aiSet { - aiVal, err := w.aiTracker.Current(w.tblName.Name) + aiVal, err := w.aiTracker.Current(w.tblName) if err != nil { return nil, err } From 97264f415fcbb6a5a77c44f2197c3eb2071394c3 Mon Sep 17 00:00:00 2001 From: Nick Tobey Date: Thu, 30 Jul 2026 17:33:29 -0700 Subject: [PATCH 11/11] Rename SequenceTracker.GetRelations to IterRelations and have it return an iter.Seq --- .../sqle/dsess/auto_increment_tracker.go | 22 ++++++++++++++----- .../doltcore/sqle/dsess/sequence_tracker.go | 20 ++++++++--------- 2 files changed, 26 insertions(+), 16 deletions(-) diff --git a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go index 94d3e14111d..9d0dd53415e 100644 --- a/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/auto_increment_tracker.go @@ -16,24 +16,35 @@ package dsess import ( "context" + "github.com/dolthub/dolt/go/libraries/doltcore/schema" + "iter" "github.com/dolthub/go-mysql-server/sql" "github.com/dolthub/dolt/go/libraries/doltcore/doltdb" - "github.com/dolthub/dolt/go/libraries/doltcore/schema" "github.com/dolthub/dolt/go/libraries/doltcore/sqle/globalstate" ) +// DoltDBRelationSource implements RelationSource +// Specializations of SequenceTracker (such as AutoIncrementTracker) require an interface to read relations of a specified +// type out of a doltdb.RootValue. DoltDBRelationSource provides the ability to read values of type doltdb.Table. type DoltDBRelationSource struct{} +// GetRelation implements RelationSource func (s DoltDBRelationSource) GetRelation(ctx context.Context, root doltdb.RootValue, tName doltdb.TableName) (relation *doltdb.Table, resolvedName string, found bool, err error) { return doltdb.GetTableInsensitive(ctx, root, tName) } -func (s DoltDBRelationSource) GetRelations(ctx context.Context, root doltdb.RootValue, cb func(doltdb.TableName, *doltdb.Table) (bool, error)) error { - return root.IterTables(ctx, func(name doltdb.TableName, table *doltdb.Table, sch schema.Schema) (stop bool, err error) { - return cb(name, table) - }) +// IterRelations implements RelationSource +func (s DoltDBRelationSource) IterRelations(ctx context.Context, root doltdb.RootValue) iter.Seq2[doltdb.TableName, *doltdb.Table] { + return func(yield func(doltdb.TableName, *doltdb.Table) bool) { + _ = root.IterTables(ctx, func(name doltdb.TableName, table *doltdb.Table, sch schema.Schema) (stop bool, err error) { + if !yield(name, table) { + return true, nil + } + return false, nil + }) + } } var _ RelationSource[*doltdb.Table, doltdb.AutoIncrementState, uint64] = (*DoltDBRelationSource)(nil) @@ -48,6 +59,7 @@ func NewAutoIncrementTracker(ctx context.Context, dbName string, roots ...doltdb return NewSequenceTrackerFromRoots(ctx, dbName, DoltDBRelationSource{}, roots...) } +// GetAutoIncrementTracker returns the AutoIncrementTracker stored within the global state. func GetAutoIncrementTracker(ctx *sql.Context, gs globalstate.GlobalState) (*AutoIncrementTracker, error) { return GetSequenceTracker(ctx, gs, autoIncrementTrackerKey) } diff --git a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go index eb2d4c28748..4f21313fc81 100644 --- a/go/libraries/doltcore/sqle/dsess/sequence_tracker.go +++ b/go/libraries/doltcore/sqle/dsess/sequence_tracker.go @@ -18,6 +18,7 @@ import ( "context" "errors" "fmt" + "iter" "time" "github.com/dolthub/go-mysql-server/sql" @@ -47,7 +48,7 @@ type RelationSource[ ] interface { // GetRelation gets a relation at a specific doltdb.RootValue GetRelation(ctx context.Context, root doltdb.RootValue, tName doltdb.TableName) (relation RelationType, resolvedName string, found bool, err error) - GetRelations(ctx context.Context, root doltdb.RootValue, cb func(doltdb.TableName, RelationType) (bool, error)) error + IterRelations(ctx context.Context, root doltdb.RootValue) iter.Seq2[doltdb.TableName, RelationType] } type SequenceTracker[ @@ -134,7 +135,7 @@ func loadSequenceState[StateType sequences.SequenceState[StateType, ValueType], return sequences.Load(relationName.ToLower()) } -func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeTableAutoIncrement(ctx *sql.Context, relationName doltdb.TableName, initialValue interface{}) (state StateType, hasState bool, err error) { +func (a *SequenceTracker[RelationType, StateType, ValueType]) initializeSequenceState(ctx *sql.Context, relationName doltdb.TableName, initialValue interface{}) (state StateType, hasState bool, err error) { sess := DSessFromSess(ctx.Session) ws, err := sess.WorkingSet(ctx, a.dbName) if err != nil { @@ -238,7 +239,7 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) Next(ctx *sql.Cont } if !ok { - currState, ok, err = a.initializeTableAutoIncrement(ctx, relationName, insertVal) + currState, ok, err = a.initializeSequenceState(ctx, relationName, insertVal) if err != nil { return nextValue, err } @@ -562,17 +563,17 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initWithRoots(ctx return err } - init := func(relationName doltdb.TableName, relation RelationType) (bool, error) { + for relationName, relation := range a.relationSource.IterRelations(ctx, r) { hasSequenceState, err := relation.HasSequenceState(ctx) if err != nil { - return true, err + return err } if !hasSequenceState { - return false, nil + continue } seq, err := relation.GetSequenceState(ctx) if err != nil { - return true, err + return err } key := relationName.ToLower() @@ -581,11 +582,8 @@ func (a *SequenceTracker[RelationType, StateType, ValueType]) initWithRoots(ctx oldValue, _ = a.sequences.Load(key) } } - - return false, nil } - - return a.relationSource.GetRelations(ctx, r, init) + return nil }) }