diff --git a/pkg/frontend/computation_wrapper.go b/pkg/frontend/computation_wrapper.go index 3510b51bbaf45..53e2601a46f74 100644 --- a/pkg/frontend/computation_wrapper.go +++ b/pkg/frontend/computation_wrapper.go @@ -544,8 +544,12 @@ func initExecuteStmtParam(execCtx *ExecCtx, ses *Session, cwft *TxnComputationWr } } - // rebuild plan when schema changed - if change { + // FK-sensitive plans also depend on the current foreign_key_checks session + // value, which does not invalidate prepared statements. Rebuild them for + // every EXECUTE so both enabled->disabled and disabled->enabled transitions + // observe the current setting. + fkSensitive := shouldRebuildPreparePlan(false, preparePlan.Plan) + if change || fkSensitive { originPrepareStmt := &tree.PrepareStmt{ Name: tree.Identifier(prepareStmt.Name), Stmt: prepareStmt.PrepareStmt, @@ -568,7 +572,7 @@ func initExecuteStmtParam(execCtx *ExecCtx, ses *Session, cwft *TxnComputationWr // query); recompiling would fail with ErrCantCompileForPrepare on every // execution, so leave it to the regular compile path (isPrepare=false). // See: https://github.com/matrixorigin/matrixone/issues/25614 - if change && prepareStmt.compile != nil { + if (change || fkSensitive) && prepareStmt.compile != nil { prepareStmt.compile.FreeOperator() prepareStmt.compile.SetIsPrepare(false) prepareStmt.compile.Release() @@ -648,6 +652,14 @@ func shouldCachePrepareCompile(p *plan.Plan) bool { return !query.GetHasForeignKeyAction() } +func shouldRebuildPreparePlan(schemaChanged bool, p *plan.Plan) bool { + if schemaChanged || p == nil { + return schemaChanged + } + query := p.GetQuery() + return query != nil && query.GetHasForeignKeyAction() +} + func createCompile( execCtx *ExecCtx, ses FeSession, diff --git a/pkg/frontend/mysql_cmd_executor.go b/pkg/frontend/mysql_cmd_executor.go index f83aa47187f92..6a657bdb82337 100644 --- a/pkg/frontend/mysql_cmd_executor.go +++ b/pkg/frontend/mysql_cmd_executor.go @@ -4310,6 +4310,9 @@ func checkNodeCanCache(p *plan2.Plan) bool { return true } if q, ok := p.Plan.(*plan2.Plan_Query); ok { + if q.Query.GetHasForeignKeyAction() { + return false + } for _, node := range q.Query.Nodes { if node.NotCacheable { return false diff --git a/pkg/frontend/prepared_fk_cache_test.go b/pkg/frontend/prepared_fk_cache_test.go index 2b4b5775f1f98..bbbca18587352 100644 --- a/pkg/frontend/prepared_fk_cache_test.go +++ b/pkg/frontend/prepared_fk_cache_test.go @@ -40,6 +40,15 @@ func TestShouldCachePrepareCompileForeignKeyActions(t *testing.T) { require.False(t, shouldCachePrepareCompile(makePlan(plan.Query_UPDATE, true))) require.False(t, shouldCachePrepareCompile(makePlan(plan.Query_DELETE, true))) + require.False(t, shouldCachePrepareCompile(makePlan(plan.Query_INSERT, true))) + + require.True(t, checkNodeCanCache(makePlan(plan.Query_INSERT, false))) + require.False(t, checkNodeCanCache(makePlan(plan.Query_INSERT, true))) + + require.False(t, shouldRebuildPreparePlan(false, nil)) + require.False(t, shouldRebuildPreparePlan(false, makePlan(plan.Query_INSERT, false))) + require.True(t, shouldRebuildPreparePlan(false, makePlan(plan.Query_INSERT, true))) + require.True(t, shouldRebuildPreparePlan(true, makePlan(plan.Query_INSERT, false))) } func TestShouldCachePrepareCompileRejectsIcebergScan(t *testing.T) { diff --git a/pkg/sql/colexec/lockop/fetch.go b/pkg/sql/colexec/lockop/fetch.go index 042fd053f78df..55107d4faf683 100644 --- a/pkg/sql/colexec/lockop/fetch.go +++ b/pkg/sql/colexec/lockop/fetch.go @@ -49,57 +49,72 @@ var ( // GetFetchRowsFunc get FetchLockRowsFunc based on primary key type func GetFetchRowsFunc(t types.Type) FetchLockRowsFunc { + var fetcher FetchLockRowsFunc switch t.Oid { case types.T_bool: - return fetchBoolRows + fetcher = fetchBoolRows case types.T_bit: - return fetchUint64Rows + fetcher = fetchUint64Rows case types.T_int8: - return fetchInt8Rows + fetcher = fetchInt8Rows case types.T_int16: - return fetchInt16Rows + fetcher = fetchInt16Rows case types.T_int32: - return fetchInt32Rows + fetcher = fetchInt32Rows case types.T_int64: - return fetchInt64Rows + fetcher = fetchInt64Rows case types.T_uint8: - return fetchUint8Rows + fetcher = fetchUint8Rows case types.T_uint16: - return fetchUint16Rows + fetcher = fetchUint16Rows case types.T_uint32: - return fetchUint32Rows + fetcher = fetchUint32Rows case types.T_uint64: - return fetchUint64Rows + fetcher = fetchUint64Rows case types.T_float32: - return fetchFloat32Rows + fetcher = fetchFloat32Rows case types.T_float64: - return fetchFloat64Rows + fetcher = fetchFloat64Rows case types.T_date: - return fetchDateRows + fetcher = fetchDateRows case types.T_year: - return fetchYearRows + fetcher = fetchYearRows case types.T_time: - return fetchTimeRows + fetcher = fetchTimeRows case types.T_datetime: - return fetchDateTimeRows + fetcher = fetchDateTimeRows case types.T_timestamp: - return fetchTimestampRows + fetcher = fetchTimestampRows case types.T_decimal64: - return fetchDecimal64Rows + fetcher = fetchDecimal64Rows case types.T_decimal128: - return fetchDecimal128Rows + fetcher = fetchDecimal128Rows case types.T_decimal256: - return fetchDecimal256Rows + fetcher = fetchDecimal256Rows case types.T_uuid: - return fetchUUIDRows + fetcher = fetchUUIDRows case types.T_char, types.T_varchar, types.T_binary, types.T_varbinary: - return fetchVarlenaRows + fetcher = fetchVarlenaRows // T_json, T_blob, T_array_float32 etc. cannot be PK. case types.T_enum: - return fetchEnumRows + fetcher = fetchEnumRows default: panic(fmt.Sprintf("not support for %s", t.String())) } + return func( + vec *vector.Vector, + packer *types.Packer, + tp types.Type, + max int, + lockTable bool, + filter RowsFilter, + filterCols []int32, + ) (bool, [][]byte, lock.Granularity) { + if !lockTable && vec.IsConstNull() { + return false, nil, lock.Granularity_Row + } + return fetcher(vec, packer, tp, max, lockTable, filter, filterCols) + } } func fetchBoolRows( @@ -714,7 +729,7 @@ func fetchVarlenaRows( n := vec.Length() data, area := vector.MustVarlenaRawData(vec) if n == 1 { - if filter != nil && + if vec.GetNulls().Contains(0) || filter != nil && !filter(0, filterCols) { return false, nil, lock.Granularity_Row } @@ -727,7 +742,7 @@ func fetchVarlenaRows( initialized := false applied := 0 for i := 0; i < n; i++ { - if filter != nil && + if vec.GetNulls().Contains(uint64(i)) || filter != nil && !filter(i, filterCols) { continue } @@ -757,7 +772,7 @@ func fetchVarlenaRows( } rows := make([][]byte, 0, n) for idx := range data { - if filter != nil && + if vec.GetNulls().Contains(uint64(idx)) || filter != nil && !filter(idx, filterCols) { continue } @@ -799,7 +814,7 @@ func fetchFixedRowsWithCompare[T any]( n := vec.Length() values := vector.MustFixedColWithTypeCheck[T](vec) if n == 1 { - if filter != nil && !filter(0, filterCols) { + if vec.GetNulls().Contains(0) || filter != nil && !filter(0, filterCols) { return false, nil, lock.Granularity_Row } return true, [][]byte{fn(values[0])}, lock.Granularity_Row @@ -809,7 +824,7 @@ func fetchFixedRowsWithCompare[T any]( initialized := false applied := 0 for row, v := range values { - if filter != nil && + if vec.GetNulls().Contains(uint64(row)) || filter != nil && !filter(row, filterCols) { continue } @@ -838,7 +853,7 @@ func fetchFixedRowsWithCompare[T any]( } rows := make([][]byte, 0, n) for row, v := range values { - if filter != nil && + if vec.GetNulls().Contains(uint64(row)) || filter != nil && !filter(row, filterCols) { continue } diff --git a/pkg/sql/colexec/lockop/fetch_test.go b/pkg/sql/colexec/lockop/fetch_test.go index e657a33aaef51..e04be4e7826b7 100644 --- a/pkg/sql/colexec/lockop/fetch_test.go +++ b/pkg/sql/colexec/lockop/fetch_test.go @@ -1472,6 +1472,62 @@ func TestDecimal128(t *testing.T) { assert.True(t, bytes.Compare(minDecimal128, maxDecimal128) < 0) } +func TestFetchRowsSkipsConstNull(t *testing.T) { + mp := mpool.MustNew("test") + vec := vector.NewConstNull(types.T_uint64.ToType(), 1, mp) + defer vec.Free(mp) + + fetcher := GetFetchRowsFunc(types.T_uint64.ToType()) + ok, rows, granularity := fetcher(vec, types.NewPacker(), types.T_uint64.ToType(), 1, false, nil, nil) + require.False(t, ok) + require.Empty(t, rows) + require.Equal(t, lock.Granularity_Row, granularity) + + ok, rows, granularity = fetcher(vec, types.NewPacker(), types.T_uint64.ToType(), 1, true, nil, nil) + require.True(t, ok) + require.Len(t, rows, 2) + require.Equal(t, lock.Granularity_Range, granularity) +} + +func TestFetchRowsSkipsPartialNull(t *testing.T) { + mp := mpool.MustNew("test") + + fixed := vector.NewVec(types.T_uint64.ToType()) + require.NoError(t, vector.AppendFixed(fixed, uint64(0), true, mp)) + require.NoError(t, vector.AppendFixed(fixed, uint64(0), false, mp)) + require.NoError(t, vector.AppendFixed(fixed, uint64(1), false, mp)) + defer fixed.Free(mp) + fetcher := GetFetchRowsFunc(types.T_uint64.ToType()) + packer := types.NewPacker() + ok, rows, granularity := fetcher(fixed, packer, types.T_uint64.ToType(), 10, false, nil, nil) + require.True(t, ok) + require.Equal(t, lock.Granularity_Row, granularity) + require.Len(t, rows, 2) + packer.Reset() + packer.EncodeUint64(0) + require.Equal(t, packer.Bytes(), rows[0]) + packer.Reset() + packer.EncodeUint64(1) + require.Equal(t, packer.Bytes(), rows[1]) + + text := vector.NewVec(types.T_varchar.ToType()) + require.NoError(t, vector.AppendBytes(text, nil, true, mp)) + require.NoError(t, vector.AppendBytes(text, []byte{}, false, mp)) + require.NoError(t, vector.AppendBytes(text, []byte("x"), false, mp)) + defer text.Free(mp) + fetcher = GetFetchRowsFunc(types.T_varchar.ToType()) + ok, rows, granularity = fetcher(text, packer, types.T_varchar.ToType(), 10, false, nil, nil) + require.True(t, ok) + require.Equal(t, lock.Granularity_Row, granularity) + require.Len(t, rows, 2) + packer.Reset() + packer.EncodeStringType([]byte{}) + require.Equal(t, packer.Bytes(), rows[0]) + packer.Reset() + packer.EncodeStringType([]byte("x")) + require.Equal(t, packer.Bytes(), rows[1]) +} + func TestDecimal256(t *testing.T) { packer := types.NewPacker() decimal256Fn := func(v types.Decimal256) []byte { diff --git a/pkg/sql/colexec/lockop/lock_op.go b/pkg/sql/colexec/lockop/lock_op.go index 76ece2e2349dc..46646a8cafa9c 100644 --- a/pkg/sql/colexec/lockop/lock_op.go +++ b/pkg/sql/colexec/lockop/lock_op.go @@ -122,6 +122,8 @@ func (lockOp *LockOp) Prepare(proc *process.Process) error { } else { lockOp.ctr.parker.Reset() } + lockOp.ctr.materializeInput = lockOp.hasMergeableTargetGroup() + lockOp.ctr.bufferEmitted = false return nil } @@ -147,6 +149,9 @@ func callNonBlocking( proc *process.Process, lockOp *LockOp) (vm.CallResult, error) { analyzer := lockOp.OpAnalyzer + if lockOp.ctr.materializeInput { + return callWithMaterializedInput(proc, lockOp, analyzer) + } result, err := vm.ChildrenCall(lockOp.GetChildren(0), proc, analyzer) if err != nil { @@ -174,6 +179,72 @@ func callNonBlocking( return result, nil } +func callWithMaterializedInput( + proc *process.Process, + lockOp *LockOp, + analyzer process.Analyzer, +) (vm.CallResult, error) { + // A mergeable group represents several key columns in one physical lock + // namespace. Consume the complete input before taking its first lock so every + // transaction derives the same global key order regardless of batch boundaries. + if lockOp.ctr.bufferEmitted { + return vm.CallResult{}, lockOp.ctr.retryError + } + for { + result, err := vm.ChildrenCall(lockOp.GetChildren(0), proc, analyzer) + if err != nil { + return result, err + } + if result.Batch == nil { + if lockOp.ctr.bufferedInput == nil { + lockOp.ctr.bufferEmitted = true + if lockOp.ctr.retryError == nil { + err = lockTalbeIfLockCountIsZero(proc, lockOp) + } + return result, err + } + lockOp.ctr.lockCount += int64(lockOp.ctr.bufferedInput.RowCount()) + if err = performLock(lockOp.ctr.bufferedInput, proc, lockOp, analyzer, -1); err != nil { + return result, err + } + lockOp.ctr.bufferEmitted = true + result.Batch = lockOp.ctr.bufferedInput + return result, nil + } + if result.Batch.IsEmpty() { + continue + } + lockOp.ctr.bufferedInput, err = lockOp.ctr.bufferedInput.AppendWithCopy( + proc.Ctx, proc.Mp(), result.Batch) + if err != nil { + return result, err + } + } +} + +func (lockOp *LockOp) hasMergeableTargetGroup() bool { + for idx := 0; idx < len(lockOp.targets); { + end := idx + 1 + for end < len(lockOp.targets) && mergeableLockTargets(lockOp.targets[idx], lockOp.targets[end]) { + end++ + } + if end-idx > 1 { + return true + } + idx = end + } + return false +} + +func mergeableLockTargets(left, right lockTarget) bool { + return left.tableID == right.tableID && left.mode == right.mode && + left.primaryColumnType == right.primaryColumnType && left.filter == nil && right.filter == nil && + !left.lockTable && !right.lockTable && left.lockRows == nil && right.lockRows == nil && + left.changeDef == right.changeDef && + left.partitionColumnIndexInBatch == right.partitionColumnIndexInBatch && + left.refreshTimestampIndexInBatch == right.refreshTimestampIndexInBatch +} + // if input vec is not allnull and has null, return a copy vector without null value func getVec(proc *process.Process, vec *vector.Vector) (*vector.Vector, error) { if vec.HasNull() { @@ -204,10 +275,34 @@ func performLock( targetIdx int, ) error { needRetry := false + consumed := make([]bool, len(lockOp.targets)) for idx, target := range lockOp.targets { + if consumed[idx] { + continue + } if targetIdx != -1 && targetIdx != idx { continue } + group := []int{idx} + if targetIdx == -1 && target.filter == nil && !target.lockTable && target.lockRows == nil { + for next := idx + 1; next < len(lockOp.targets); next++ { + candidate := lockOp.targets[next] + if !mergeableLockTargets(target, candidate) { + break + } + group = append(group, next) + consumed[next] = true + } + } + primaryIdx := idx + for _, groupIdx := range group { + vec := resultVector(bat, lockOp.targets[groupIdx].primaryColumnIndexInBatch) + if vec != nil && !vec.AllNull() { + primaryIdx = groupIdx + break + } + } + target = lockOp.targets[primaryIdx] if proc.GetTxnOperator().LockSkipped(target.tableID, target.mode) { return nil } @@ -236,11 +331,26 @@ func performLock( } } } */ + fetchRows := lockOp.ctr.fetchers[primaryIdx] + if len(group) > 1 { + fetchRows = func( + _ *vector.Vector, + packer *types.Packer, + pkType types.Type, + maxCountPerLock int, + lockTable bool, + _ RowsFilter, + _ []int32, + ) (bool, [][]byte, lock.Granularity) { + return lockOp.fetchMergedLockRows( + bat, group, packer, pkType, maxCountPerLock, lockTable) + } + } locked, defChanged, refreshTS, err := doLock( proc.Ctx, lockOp.engine, analyzer, - lockOp.ctr.relations[idx], + lockOp.ctr.relations[primaryIdx], target.tableID, proc, bat, @@ -248,8 +358,8 @@ func performLock( target.primaryColumnType, target.partitionColumnIndexInBatch, DefaultLockOptions(lockOp.ctr.parker). - WithLockMode(lock.LockMode_Exclusive). - WithFetchLockRowsFunc(lockOp.ctr.fetchers[idx]). + WithLockMode(target.mode). + WithFetchLockRowsFunc(fetchRows). WithMaxBytesPerLock(int(proc.GetLockService().GetConfig().MaxLockRowCount)). WithFilterRows(target.filter, filterCols). WithLockTable(target.lockTable, target.changeDef). @@ -305,6 +415,44 @@ func performLock( return nil } +func (lockOp *LockOp) fetchMergedLockRows( + bat *batch.Batch, + group []int, + packer *types.Packer, + pkType types.Type, + maxCountPerLock int, + lockTable bool, +) (bool, [][]byte, lock.Granularity) { + // Disable per-target range conversion. A range chosen from only the first + // target can omit keys from later targets; convert only after the complete + // group has been sorted and deduplicated. + var rows [][]byte + for _, groupIdx := range group { + groupTarget := lockOp.targets[groupIdx] + has, targetRows, _ := lockOp.ctr.fetchers[groupIdx]( + resultVector(bat, groupTarget.primaryColumnIndexInBatch), + packer, pkType, math.MaxInt, lockTable, nil, nil) + if has { + rows = append(rows, targetRows...) + } + } + if len(rows) == 0 { + return false, nil, lock.Granularity_Row + } + rows = dedupLockRows(rows) + if len(rows) > maxCountPerLock { + return true, [][]byte{rows[0], rows[len(rows)-1]}, lock.Granularity_Range + } + return true, rows, lock.Granularity_Row +} + +func resultVector(bat *batch.Batch, idx int32) *vector.Vector { + if bat == nil { + return nil + } + return bat.GetVector(idx) +} + // LockTable lock table, all rows in the table will be locked, and wait current txn // closed. func LockTable( @@ -313,6 +461,17 @@ func LockTable( tableID uint64, pkType types.Type, changeDef bool) error { + return LockTableWithMode(eng, proc, tableID, pkType, lock.LockMode_Exclusive, changeDef) +} + +// LockTableWithMode locks all rows in a table with the specified lock mode. +func LockTableWithMode( + eng engine.Engine, + proc *process.Process, + tableID uint64, + pkType types.Type, + mode lock.LockMode, + changeDef bool) error { txnOp := proc.GetTxnOperator() if !txnOp.Txn().IsPessimistic() { return nil @@ -336,6 +495,7 @@ func LockTable( }() opts := DefaultLockOptions(parker). + WithLockMode(mode). WithLockTable(true, changeDef). WithFetchLockRowsFunc(GetFetchRowsFunc(pkType)) _, defChanged, refreshTS, err := doLock( @@ -1411,6 +1571,7 @@ func (lockOp *LockOp) AddLockTargetWithPartitionAndMode( } func (lockOp *LockOp) Reset(proc *process.Process, pipelineFailed bool, err error) { + lockOp.cleanBufferedInput(proc) lockOp.resetParker() lockOp.ctr.retryError = nil lockOp.ctr.defChanged = false @@ -1419,10 +1580,19 @@ func (lockOp *LockOp) Reset(proc *process.Process, pipelineFailed bool, err erro // Free free mem func (lockOp *LockOp) Free(proc *process.Process, pipelineFailed bool, err error) { + lockOp.cleanBufferedInput(proc) lockOp.cleanParker() lockOp.ctr.relations = nil } +func (lockOp *LockOp) cleanBufferedInput(proc *process.Process) { + if lockOp.ctr.bufferedInput != nil && proc != nil { + lockOp.ctr.bufferedInput.Clean(proc.Mp()) + } + lockOp.ctr.bufferedInput = nil + lockOp.ctr.bufferEmitted = false +} + func (lockOp *LockOp) ExecProjection(proc *process.Process, input *batch.Batch) (*batch.Batch, error) { return input, nil } @@ -1540,7 +1710,8 @@ func lockTalbeIfLockCountIsZero( if !target.lockTableAtTheEnd { continue } - err := LockTable(lockOp.engine, proc, target.tableID, target.primaryColumnType, false) + err := LockTableWithMode( + lockOp.engine, proc, target.tableID, target.primaryColumnType, target.mode, false) if err != nil { return err } diff --git a/pkg/sql/colexec/lockop/lock_op_test.go b/pkg/sql/colexec/lockop/lock_op_test.go index f14a5827acff0..500e0d60d33f9 100644 --- a/pkg/sql/colexec/lockop/lock_op_test.go +++ b/pkg/sql/colexec/lockop/lock_op_test.go @@ -1087,8 +1087,9 @@ func TestCallLockOpLocksTableAtEOFWhenNoRowsProduced(t *testing.T) { IsFirst: false, IsLast: false, } - arg.AddLockTarget(tableID, nil, 0, pkType, -1, -1, nil, true) - arg.LockTable(tableID, false) + arg.AddLockTargetWithMode( + tableID, nil, lock.LockMode_Shared, 0, pkType, -1, -1, nil, true) + arg.LockTableWithMode(tableID, lock.LockMode_Shared, false) resetChildren(arg, nil) defer arg.Free(proc, false, nil) @@ -1096,6 +1097,16 @@ func TestCallLockOpLocksTableAtEOFWhenNoRowsProduced(t *testing.T) { _, err := vm.Exec(arg, proc) require.NoError(t, err) require.True(t, proc.GetTxnOperator().HasLockTable(tableID)) + + sharedTxn, err := proc.Base.TxnClient.New(proc.Ctx, timestamp.Timestamp{}) + require.NoError(t, err) + defer func() { require.NoError(t, sharedTxn.Rollback(proc.Ctx)) }() + sharedProc := process.NewTopProcess(proc.Ctx, mpool.MustNewZero(), proc.Base.TxnClient, + sharedTxn, nil, proc.GetLockService(), nil, nil, nil, nil, nil) + require.NoError(t, LockTableWithMode( + nil, sharedProc, tableID, pkType, lock.LockMode_Shared, false)) + require.NoError(t, LockTable(nil, proc, tableID+1, pkType, false)) + require.True(t, proc.GetTxnOperator().HasLockTable(tableID+1)) }, ) } @@ -1691,3 +1702,65 @@ func TestDedupLockRows_Idempotent(t *testing.T) { twice := dedupLockRows(append([][]byte(nil), once...)) require.Equal(t, once, twice) } + +func TestFetchMergedLockRowsUsesAllTargetsForRange(t *testing.T) { + mp := mpool.MustNew("test") + pkType := types.T_int32.ToType() + bat := batch.NewWithSize(2) + bat.Vecs[0] = testutil.MakeInt32Vector([]int32{1, 100}, nil, mp) + bat.Vecs[1] = testutil.MakeInt32Vector([]int32{2, 200}, nil, mp) + bat.SetRowCount(2) + defer bat.Clean(mp) + + arg := NewArgument() + arg.AddLockTargetWithMode(1, nil, lock.LockMode_Shared, 0, pkType, -1, -1, nil, false) + arg.AddLockTargetWithMode(1, nil, lock.LockMode_Shared, 1, pkType, -1, -1, nil, false) + arg.ctr.fetchers = []FetchLockRowsFunc{GetFetchRowsFunc(pkType), GetFetchRowsFunc(pkType)} + packer := types.NewPacker() + defer packer.Close() + + ok, rows, granularity := arg.fetchMergedLockRows(bat, []int{0, 1}, packer, pkType, 2, false) + require.True(t, ok) + require.Equal(t, lock.Granularity_Range, granularity) + require.Len(t, rows, 2) + packer.Reset() + packer.EncodeInt32(1) + require.Equal(t, packer.Bytes(), rows[0]) + packer.Reset() + packer.EncodeInt32(200) + require.Equal(t, packer.Bytes(), rows[1]) +} + +func TestLockOpMaterializesMergeableTargetsAcrossBatches(t *testing.T) { + runLockOpTest(t, func(proc *process.Process) { + pkType := types.T_int32.ToType() + makeBatch := func(left, right int32) *batch.Batch { + bat := batch.NewWithSize(2) + bat.Vecs[0] = testutil.MakeInt32Vector([]int32{left}, nil, proc.Mp()) + bat.Vecs[1] = testutil.MakeInt32Vector([]int32{right}, nil, proc.Mp()) + bat.SetRowCount(1) + return bat + } + first := makeBatch(2, 1) + second := makeBatch(1, 2) + defer first.Clean(proc.Mp()) + defer second.Clean(proc.Mp()) + + arg := NewArgumentByEngine(nil) + arg.AddLockTargetWithMode(1, nil, lock.LockMode_Shared, 0, pkType, -1, -1, nil, false) + arg.AddLockTargetWithMode(1, nil, lock.LockMode_Shared, 1, pkType, -1, -1, nil, false) + arg.AppendChild(colexec.NewMockOperator().WithBatchs([]*batch.Batch{first, second})) + require.NoError(t, arg.Prepare(proc)) + arg.ctr.hasNewVersionInRange = testFunc + require.True(t, arg.ctr.materializeInput) + + result, err := vm.Exec(arg, proc) + require.NoError(t, err) + require.NotNil(t, result.Batch) + require.Equal(t, 2, result.Batch.RowCount()) + result, err = vm.Exec(arg, proc) + require.NoError(t, err) + require.Nil(t, result.Batch) + arg.Free(proc, false, nil) + }) +} diff --git a/pkg/sql/colexec/lockop/types.go b/pkg/sql/colexec/lockop/types.go index 41183d1b9ad97..d131b3753008a 100644 --- a/pkg/sql/colexec/lockop/types.go +++ b/pkg/sql/colexec/lockop/types.go @@ -141,4 +141,7 @@ type state struct { relations []engine.Relation hasNewVersionInRange hasNewVersionInRangeFunc lockCount int64 + materializeInput bool + bufferedInput *batch.Batch + bufferEmitted bool } diff --git a/pkg/sql/compile/compile.go b/pkg/sql/compile/compile.go index 4de7d0165288c..59775de637b49 100644 --- a/pkg/sql/compile/compile.go +++ b/pkg/sql/compile/compile.go @@ -527,12 +527,21 @@ func (c *Compile) runOnce() (err error) { c.proc.Base.StageCache.Clear() }() - // Pre-check: REPLACE parent→child FK RESTRICT constraints must be - // verified before the REPLACE execution modifies any rows. + // REPLACE parent checks and actions run before the main pipeline. query := c.pn.GetQuery() if query != nil && query.StmtType == plan.Query_INSERT && len(query.GetDetectSqls()) != 0 { + if err = validateReplaceParentTxnMode( + c.proc.Ctx, query, c.proc.GetTxnOperator().Txn().IsPessimistic()); err != nil { + return err + } for _, sql := range query.DetectSqls { - if strings.HasPrefix(sql, "REPLACE_PARENT_CHK:") { + if strings.HasPrefix(sql, "REPLACE_PARENT_PLAN:") { + continue + } else if strings.HasPrefix(sql, "REPLACE_PARENT_LOCK:") { + if err = c.runSql(strings.TrimPrefix(sql, "REPLACE_PARENT_LOCK:")); err != nil { + return err + } + } else if strings.HasPrefix(sql, "REPLACE_PARENT_CHK:") { if err = runDetectSql(c, strings.TrimPrefix(sql, "REPLACE_PARENT_CHK:")); err != nil { // Only translate the "check returned false" signal into the // parent-row-referenced error; pass through real execution @@ -543,6 +552,10 @@ func (c *Compile) runOnce() (err error) { } return err } + } else if strings.HasPrefix(sql, "REPLACE_PARENT_ACTION:") { + if err = c.runSql(strings.TrimPrefix(sql, "REPLACE_PARENT_ACTION:")); err != nil { + return err + } } } } @@ -636,12 +649,14 @@ func (c *Compile) runOnce() (err error) { query = c.pn.GetQuery() if query != nil && (query.StmtType == plan.Query_INSERT || query.StmtType == plan.Query_UPDATE) && len(query.GetDetectSqls()) != 0 { - // Filter out pre-check SQLs (already executed before the main operation). - // The modern INSERT path enforces child→parent existence in-plan now, so the - // remaining DetectSqls are self-referencing FK checks (plain 1452 message). + // Filter out REPLACE parent-side checks and actions already executed before + // the main operation. var postCheckSqls []string for _, sql := range query.DetectSqls { - if strings.HasPrefix(sql, "REPLACE_PARENT_CHK:") { + if strings.HasPrefix(sql, "REPLACE_PARENT_LOCK:") || + strings.HasPrefix(sql, "REPLACE_PARENT_PLAN:") || + strings.HasPrefix(sql, "REPLACE_PARENT_CHK:") || + strings.HasPrefix(sql, "REPLACE_PARENT_ACTION:") { continue } postCheckSqls = append(postCheckSqls, sql) @@ -658,6 +673,20 @@ func (c *Compile) runOnce() (err error) { return err } +func validateReplaceParentTxnMode(ctx context.Context, query *plan.Query, pessimistic bool) error { + if pessimistic || query == nil { + return nil + } + for _, sql := range query.DetectSqls { + if strings.HasPrefix(sql, "REPLACE_PARENT_LOCK:") || + strings.HasPrefix(sql, "REPLACE_PARENT_PLAN:") { + return moerr.NewNotSupported(ctx, + "REPLACE on a referenced parent table in optimistic transaction mode") + } + } + return nil +} + // add log to check if background sql return NeedRetry error when origin sql execute successfully func (c *Compile) debugLogFor19288(err error, bsql string) { if c.isRetryErr(err) { @@ -809,11 +838,12 @@ func (c *Compile) lockTable() error { for _, tableID := range tableIDs { tbl := c.lockTables[tableID] typ := plan2.MakeTypeByPlan2Type(tbl.PrimaryColTyp) - if err := lockop.LockTable( + if err := lockop.LockTableWithMode( c.e, c.proc, tbl.TableId, typ, + tbl.Mode, false); err != nil { return err } diff --git a/pkg/sql/compile/compile_test.go b/pkg/sql/compile/compile_test.go index 2b631267055ec..8a9823ea7e86f 100644 --- a/pkg/sql/compile/compile_test.go +++ b/pkg/sql/compile/compile_test.go @@ -32,8 +32,11 @@ import ( "github.com/matrixorigin/matrixone/pkg/container/vector" "github.com/matrixorigin/matrixone/pkg/defines" "github.com/matrixorigin/matrixone/pkg/lockservice" + lockpb "github.com/matrixorigin/matrixone/pkg/pb/lock" + "github.com/matrixorigin/matrixone/pkg/pb/plan" "github.com/matrixorigin/matrixone/pkg/pb/timestamp" "github.com/matrixorigin/matrixone/pkg/perfcounter" + "github.com/matrixorigin/matrixone/pkg/sql/colexec/lockop" "github.com/matrixorigin/matrixone/pkg/txn/client" "github.com/matrixorigin/matrixone/pkg/txn/rpc" @@ -45,7 +48,6 @@ import ( "github.com/matrixorigin/matrixone/pkg/common/buffer" "github.com/matrixorigin/matrixone/pkg/container/batch" mock_frontend "github.com/matrixorigin/matrixone/pkg/frontend/test" - "github.com/matrixorigin/matrixone/pkg/pb/plan" "github.com/matrixorigin/matrixone/pkg/pb/txn" "github.com/matrixorigin/matrixone/pkg/sql/colexec" "github.com/matrixorigin/matrixone/pkg/sql/colexec/group" @@ -213,6 +215,40 @@ func TestShouldPrePipelineLockTable(t *testing.T) { require.False(t, target.LockTableAtTheEnd) } +func TestConstructLockOpPreservesSharedTableMode(t *testing.T) { + for _, lockTable := range []bool{false, true} { + t.Run(fmt.Sprintf("table=%t", lockTable), func(t *testing.T) { + node := &plan.Node{LockTargets: []*plan.LockTarget{{ + TableId: 42, PrimaryColTyp: plan.Type{Id: int32(types.T_int64)}, + Mode: lockpb.LockMode_Shared, LockTable: lockTable, + }}} + + op, err := constructLockOp(node, nil) + require.NoError(t, err) + targets := op.CopyToPipelineTarget() + require.Len(t, targets, 1) + assert.Equal(t, lockTable, targets[0].LockTable) + assert.Equal(t, lockpb.LockMode_Shared, targets[0].Mode) + }) + } +} + +func TestValidateReplaceParentTxnMode(t *testing.T) { + ctx := context.Background() + query := &plan.Query{DetectSqls: []string{"REPLACE_PARENT_LOCK:select 1 for update"}} + + require.NoError(t, validateReplaceParentTxnMode(ctx, query, true)) + require.ErrorContains(t, validateReplaceParentTxnMode(ctx, query, false), + "optimistic transaction mode") + query.DetectSqls = []string{"REPLACE_PARENT_PLAN:"} + require.NoError(t, validateReplaceParentTxnMode(ctx, query, true)) + require.ErrorContains(t, validateReplaceParentTxnMode(ctx, query, false), + "optimistic transaction mode") + require.NoError(t, validateReplaceParentTxnMode(ctx, + &plan.Query{DetectSqls: []string{"select true"}}, false)) + require.NoError(t, validateReplaceParentTxnMode(ctx, nil, false)) +} + func TestLockTableLocksAllPrePipelineTargets(t *testing.T) { runtime.RunTest( "", @@ -258,7 +294,8 @@ func TestLockTableLocksAllPrePipelineTargets(t *testing.T) { c := &Compile{ proc: proc, lockTables: map[uint64]*plan.LockTarget{ - 10: {TableId: 10, PrimaryColTyp: plan.Type{Id: int32(types.T_int32)}}, + 10: {TableId: 10, PrimaryColTyp: plan.Type{Id: int32(types.T_int32)}, + Mode: lockpb.LockMode_Shared}, 11: {TableId: 11, PrimaryColTyp: plan.Type{Id: int32(types.T_int32)}}, }, } @@ -266,6 +303,14 @@ func TestLockTableLocksAllPrePipelineTargets(t *testing.T) { require.NoError(t, c.lockTable()) require.True(t, txnOp.HasLockTable(10)) require.True(t, txnOp.HasLockTable(11)) + + sharedTxn, err := txnClient.New(ctx, timestamp.Timestamp{}) + require.NoError(t, err) + defer func() { require.NoError(t, sharedTxn.Rollback(ctx)) }() + sharedProc := process.NewTopProcess(ctx, mpool.MustNewZero(), txnClient, sharedTxn, + nil, services[0], nil, nil, nil, nil, nil) + require.NoError(t, lockop.LockTableWithMode(nil, sharedProc, 10, + types.T_int32.ToType(), lockpb.LockMode_Shared, false)) }, nil, ) diff --git a/pkg/sql/compile/operator.go b/pkg/sql/compile/operator.go index 78ca821222587..88e02448c31d0 100644 --- a/pkg/sql/compile/operator.go +++ b/pkg/sql/compile/operator.go @@ -799,11 +799,11 @@ func constructLockOp(node *plan.Node, eng engine.Engine) (*lockop.LockOp, error) partitionColPos = target.PartitionColIdxInBat } typ := plan2.MakeTypeByPlan2Type(target.PrimaryColTyp) - arg.AddLockTarget(target.GetTableId(), target.GetObjRef(), target.GetPrimaryColIdxInBat(), typ, partitionColPos, target.GetRefreshTsIdxInBat(), target.GetLockRows(), target.GetLockTableAtTheEnd()) + arg.AddLockTargetWithMode(target.GetTableId(), target.GetObjRef(), target.GetMode(), target.GetPrimaryColIdxInBat(), typ, partitionColPos, target.GetRefreshTsIdxInBat(), target.GetLockRows(), target.GetLockTableAtTheEnd()) } for _, target := range node.LockTargets { if target.LockTable { - arg.LockTable(target.TableId, false) + arg.LockTableWithMode(target.TableId, target.Mode, false) } } return arg, nil diff --git a/pkg/sql/compile/remoterun.go b/pkg/sql/compile/remoterun.go index 43cf50ceb0e8c..0df6cc44181b3 100644 --- a/pkg/sql/compile/remoterun.go +++ b/pkg/sql/compile/remoterun.go @@ -992,11 +992,11 @@ func convertToVmOperator(opr *pipeline.Instruction, ctx *scopeContext, eng engin lockArg := lockop.NewArgumentByEngine(eng) for _, target := range t.Targets { typ := plan2.MakeTypeByPlan2Type(target.PrimaryColTyp) - lockArg.AddLockTarget(target.GetTableId(), target.GetObjRef(), target.GetPrimaryColIdxInBat(), typ, target.PartitionColIdxInBat, target.GetRefreshTsIdxInBat(), target.GetLockRows(), target.GetLockTableAtTheEnd()) + lockArg.AddLockTargetWithMode(target.GetTableId(), target.GetObjRef(), target.GetMode(), target.GetPrimaryColIdxInBat(), typ, target.PartitionColIdxInBat, target.GetRefreshTsIdxInBat(), target.GetLockRows(), target.GetLockTableAtTheEnd()) } for _, target := range t.Targets { if target.LockTable { - lockArg.LockTable(target.TableId, target.ChangeDef) + lockArg.LockTableWithMode(target.TableId, target.Mode, target.ChangeDef) } } op = lockArg diff --git a/pkg/sql/compile/remoterun_test.go b/pkg/sql/compile/remoterun_test.go index cba350c2d9f0f..94fed0fc7bb1c 100644 --- a/pkg/sql/compile/remoterun_test.go +++ b/pkg/sql/compile/remoterun_test.go @@ -36,6 +36,7 @@ import ( "github.com/matrixorigin/matrixone/pkg/container/vector" "github.com/matrixorigin/matrixone/pkg/defines" mock_frontend "github.com/matrixorigin/matrixone/pkg/frontend/test" + lockpb "github.com/matrixorigin/matrixone/pkg/pb/lock" "github.com/matrixorigin/matrixone/pkg/pb/pipeline" planpb "github.com/matrixorigin/matrixone/pkg/pb/plan" "github.com/matrixorigin/matrixone/pkg/pb/txn" @@ -359,6 +360,37 @@ func TestRemoteRunOperatorCodecRoundTrip(t *testing.T) { require.IsType(t, &intersectall.IntersectAll{}, restored) require.Equal(t, vm.IntersectAll, restored.OpType()) }) + + t.Run("SharedTableLock", func(t *testing.T) { + original := lockop.NewArgumentByEngine(nil) + original.AddLockTargetWithMode(42, nil, lockpb.LockMode_Shared, 0, + types.T_int64.ToType(), -1, -1, nil, false) + original.LockTableWithMode(42, lockpb.LockMode_Shared, false) + + restored := roundTrip(t, original) + defer restored.Release() + restoredLock, ok := restored.(*lockop.LockOp) + require.True(t, ok) + targets := restoredLock.CopyToPipelineTarget() + require.Len(t, targets, 1) + require.True(t, targets[0].LockTable) + require.Equal(t, lockpb.LockMode_Shared, targets[0].Mode) + }) + + t.Run("SharedRowLock", func(t *testing.T) { + original := lockop.NewArgumentByEngine(nil) + original.AddLockTargetWithMode(43, nil, lockpb.LockMode_Shared, 0, + types.T_int64.ToType(), -1, -1, nil, false) + + restored := roundTrip(t, original) + defer restored.Release() + restoredLock, ok := restored.(*lockop.LockOp) + require.True(t, ok) + targets := restoredLock.CopyToPipelineTarget() + require.Len(t, targets, 1) + require.False(t, targets[0].LockTable) + require.Equal(t, lockpb.LockMode_Shared, targets[0].Mode) + }) } func TestExternalScanParquetRowGroupShardsRoundtrip(t *testing.T) { diff --git a/pkg/sql/plan/bind_insert.go b/pkg/sql/plan/bind_insert.go index 042a5c5383543..023a0c799c404 100644 --- a/pkg/sql/plan/bind_insert.go +++ b/pkg/sql/plan/bind_insert.go @@ -17,6 +17,7 @@ package plan import ( "context" "fmt" + "slices" "strings" "github.com/google/uuid" @@ -24,6 +25,7 @@ import ( "github.com/matrixorigin/matrixone/pkg/common/moerr" "github.com/matrixorigin/matrixone/pkg/defines" "github.com/matrixorigin/matrixone/pkg/logutil" + lockpb "github.com/matrixorigin/matrixone/pkg/pb/lock" "github.com/matrixorigin/matrixone/pkg/pb/plan" "github.com/matrixorigin/matrixone/pkg/sql/parsers/tree" "github.com/matrixorigin/matrixone/pkg/sql/util" @@ -1034,24 +1036,276 @@ func (builder *QueryBuilder) appendModernChildFkMarkOks( childColPos func(colName string) int32, ) (int32, []*plan.Expr, error) { selectNode := builder.qry.Nodes[lastNodeID] + inputTypes := make([]plan.Type, len(selectNode.ProjectList)) + for i, expr := range selectNode.ProjectList { + inputTypes[i] = expr.Typ + } id2name := make(map[uint64]string, len(tableDef.Cols)) for _, col := range tableDef.Cols { id2name[col.ColId] = col.Name } - oks := make([]*plan.Expr, 0, len(tableDef.Fkeys)) + nonSelfFks := make([]*plan.ForeignKeyDef, 0, len(tableDef.Fkeys)) for _, fk := range tableDef.Fkeys { - if fk.ForeignTbl == 0 { - continue // self-referencing FK handled post-execution via DetectSql + if fk.ForeignTbl != 0 { + nonSelfFks = append(nonSelfFks, fk) + } + } + if len(nonSelfFks) == 0 { + return lastNodeID, nil, nil + } + lockForeignKeys := true + if proc := builder.compCtx.GetProcess(); proc != nil { + if txnOp := proc.GetTxnOperator(); txnOp != nil { + lockForeignKeys = txnOp.Txn().IsPessimistic() + } + } + + parentColNames := func(parent *plan.TableDef, colIDs []uint64) ([]string, error) { + idToName := make(map[uint64]string, len(parent.Cols)) + for _, col := range parent.Cols { + idToName[col.ColId] = col.Name + } + names := make([]string, len(colIDs)) + for i, id := range colIDs { + var ok bool + if names[i], ok = idToName[id]; !ok { + return nil, moerr.NewInternalErrorf(builder.GetContext(), + "foreign-key parent column %d not found", id) + } + } + return names, nil + } + partsEqual := func(parts, names []string) bool { + if len(parts) != len(names) { + return false + } + for i := range parts { + if catalog.ResolveAlias(parts[i]) != names[i] { + return false + } + } + return true + } + validationFks := nonSelfFks + type foreignKeyLock struct { + tableDef *plan.TableDef + objRef *plan.ObjectRef + expr *plan.Expr + typ plan.Type + lockTable bool + } + foreignKeyLocks := make([]foreignKeyLock, 0, len(nonSelfFks)) + if lockForeignKeys { + type orderedFK struct { + fk *plan.ForeignKeyDef + key string + } + ordered := make([]orderedFK, 0, len(nonSelfFks)) + for _, fk := range nonSelfFks { + _, parentTableDef, err := builder.compCtx.ResolveById(fk.ForeignTbl, bindCtx.snapshot) + if err != nil { + return 0, nil, err + } + if parentTableDef == nil { + return 0, nil, moerr.NewInternalErrorf(builder.GetContext(), "parent table %d not found", fk.ForeignTbl) + } + referencedNames, err := parentColNames(parentTableDef, fk.ForeignCols) + if err != nil { + return 0, nil, err + } + pkeyNames := []string(nil) + if parentTableDef.Pkey != nil { + pkeyNames = parentTableDef.Pkey.Names + if len(pkeyNames) == 0 && parentTableDef.Pkey.PkeyColName != "" { + pkeyNames = []string{parentTableDef.Pkey.PkeyColName} + } + } + // Base-table locks sort before hidden-index locks, matching the parent + // REPLACE pre-phase. Hidden targets then use the physical index-table + // name as the stable order shared by both sides. + targetKey := "0:" + if !partsEqual(pkeyNames, referencedNames) { + for _, idxDef := range parentTableDef.Indexes { + if idxDef.Unique && partsEqual(idxDef.Parts, referencedNames) { + targetKey = "1:" + idxDef.IndexTableName + break + } + } + } + ordered = append(ordered, orderedFK{ + fk: fk, + key: fmt.Sprintf("%020d:%s", fk.ForeignTbl, targetKey), + }) } + slices.SortStableFunc(ordered, func(left, right orderedFK) int { + return strings.Compare(left.key, right.key) + }) + nonSelfFks = make([]*plan.ForeignKeyDef, len(ordered)) + for i := range ordered { + nonSelfFks[i] = ordered[i].fk + } + } + + if lockForeignKeys { + for _, fk := range nonSelfFks { + parentObjRef, parentTableDef, err := builder.compCtx.ResolveById(fk.ForeignTbl, bindCtx.snapshot) + if err != nil { + return 0, nil, err + } + if parentTableDef == nil { + return 0, nil, moerr.NewInternalErrorf(builder.GetContext(), "parent table %d not found", fk.ForeignTbl) + } + referencedNames, err := parentColNames(parentTableDef, fk.ForeignCols) + if err != nil { + return 0, nil, err + } + childExprs := make([]*plan.Expr, len(fk.Cols)) + for i, childColID := range fk.Cols { + pos := childColPos(id2name[childColID]) + childExpr := &plan.Expr{Typ: inputTypes[pos], Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: selectTag, ColPos: int32(pos), + }}} + var parentCol *plan.ColDef + for _, col := range parentTableDef.Cols { + if col.ColId == fk.ForeignCols[i] { + parentCol = col + break + } + } + if parentCol == nil { + return 0, nil, moerr.NewInternalErrorf(builder.GetContext(), + "foreign-key parent column %s not found", referencedNames[i]) + } + childExprs[i], err = makePlan2AssignmentCastExpr( + builder.GetContext(), childExpr, parentCol.Typ) + if err != nil { + return 0, nil, err + } + } + + var lockExpr *plan.Expr + var lockTableDef *plan.TableDef + var lockObjRef *plan.ObjectRef + lockTable := false + var pkeyNames []string + if parentTableDef.Pkey != nil { + pkeyNames = parentTableDef.Pkey.Names + if len(pkeyNames) == 0 && parentTableDef.Pkey.PkeyColName != "" { + pkeyNames = []string{parentTableDef.Pkey.PkeyColName} + } + } + if partsEqual(pkeyNames, referencedNames) { + lockTableDef = parentTableDef + lockObjRef = parentObjRef + if len(childExprs) == 1 { + lockExpr = childExprs[0] + } else { + lockExpr, err = BindFuncExprImplByPlanExpr(builder.GetContext(), "serial", childExprs) + if err != nil { + return 0, nil, err + } + } + } else { + var matchedIndex *plan.IndexDef + for _, idxDef := range parentTableDef.Indexes { + if idxDef.Unique && partsEqual(idxDef.Parts, referencedNames) { + matchedIndex = idxDef + break + } + } + if matchedIndex == nil { + // Legacy schemas may reference a non-unique prefix of a composite key. + // Such a reference has no physical point-lock key, so serialize it with + // parent mutations before the validation scan using a shared table lock. + lockTableDef = parentTableDef + lockObjRef = parentObjRef + lockExpr = childExprs[0] + lockTable = true + } else { + lockObjRef, lockTableDef, err = builder.compCtx.ResolveIndexTableByRef( + parentObjRef, matchedIndex.IndexTableName, bindCtx.snapshot) + if err != nil { + return 0, nil, err + } + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(matchedIndex.IndexAlgoParams) + if err != nil { + return 0, nil, err + } + keyParts := make([]*plan.Expr, len(childExprs)) + for i, expr := range childExprs { + keyParts[i], err = builder.makeIndexPartExprFromInputExpr(expr, referencedNames[i], prefixLengths) + if err != nil { + return 0, nil, err + } + } + if indexTableStoresSerializedKey(matchedIndex) { + lockExpr, err = BindFuncExprImplByPlanExpr(builder.GetContext(), "serial", keyParts) + if err != nil { + return 0, nil, err + } + } else { + lockExpr = keyParts[0] + } + } + } + lockPkPos, lockTyp := getPkPos(lockTableDef, false) + if lockPkPos < 0 { + return 0, nil, moerr.NewInternalErrorf(builder.GetContext(), + "foreign-key lock table %s has no primary key", lockTableDef.Name) + } + foreignKeyLocks = append(foreignKeyLocks, foreignKeyLock{ + tableDef: lockTableDef, + objRef: lockObjRef, + expr: lockExpr, + typ: lockTyp, + lockTable: lockTable, + }) + } + } + + if lockForeignKeys { + rowProject := getProjectionByLastNodeWithTag(builder, lastNodeID, selectTag) + lockTag := builder.genNewBindTag() + lockProject := slices.Clone(rowProject) + lockTargets := make([]*plan.LockTarget, 0, len(foreignKeyLocks)) + for _, fkLock := range foreignKeyLocks { + lockProject = append(lockProject, fkLock.expr) + lockTargets = append(lockTargets, &plan.LockTarget{ + TableId: fkLock.tableDef.TblId, ObjRef: fkLock.objRef, + PrimaryColIdxInBat: int32(len(lockProject) - 1), PrimaryColRelPos: lockTag, + PrimaryColTyp: fkLock.typ, Mode: lockpb.LockMode_Shared, LockTable: fkLock.lockTable, + }) + } + lockInputID := builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, Children: []int32{lastNodeID}, + ProjectList: lockProject, BindingTags: []int32{lockTag}, + }, bindCtx) + lockOutput := getProjectionByLastNodeWithTag(builder, lockInputID, lockTag) + lockNodeID := builder.appendNode(&plan.Node{ + NodeType: plan.Node_LOCK_OP, Children: []int32{lockInputID}, + TableDef: foreignKeyLocks[0].tableDef, LockTargets: lockTargets, + }, bindCtx) + lastNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, Children: []int32{lockNodeID}, + ProjectList: slices.Clone(lockOutput[:len(rowProject)]), BindingTags: []int32{selectTag}, + }, bindCtx) + + // Materialize the row image only after every referenced key has been locked. + // The final validation/DML step consumes this single sink dependency, avoiding + // unsupported multi-hop sink chains while preserving lock-before-scan ordering. + lockSinkID := appendSinkNodeWithTag(builder, bindCtx, lastNodeID, selectTag) + lockStep := builder.appendStep(lockSinkID) + lastNodeID = builder.appendTaggedSinkScan(bindCtx, lockStep, selectTag) + selectNode = builder.qry.Nodes[lastNodeID] + } + oks := make([]*plan.Expr, 0, len(nonSelfFks)) + for _, fk := range validationFks { parentObjRef, parentTableDef, err := builder.compCtx.ResolveById(fk.ForeignTbl, bindCtx.snapshot) if err != nil { return 0, nil, err } - if parentTableDef == nil { - return 0, nil, moerr.NewInternalErrorf(builder.GetContext(), "parent table %d not found", fk.ForeignTbl) - } parentTag := builder.genNewBindTag() builder.addNameByColRef(parentTag, parentTableDef) diff --git a/pkg/sql/plan/bind_replace.go b/pkg/sql/plan/bind_replace.go index cbdfd82be227c..6a52ba156fda1 100644 --- a/pkg/sql/plan/bind_replace.go +++ b/pkg/sql/plan/bind_replace.go @@ -17,6 +17,7 @@ package plan import ( "context" "fmt" + "slices" "strings" "github.com/matrixorigin/matrixone/pkg/catalog" @@ -121,6 +122,7 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( break } } + needsOldIndexMaintenance := !isFakePK || hasUniqueIdx // get old columns from existing main table // @@ -186,7 +188,7 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( var err error for i, idxDef := range tableDef.Indexes { - if skipUniqueIdx[i] { + if skipUniqueIdx[i] && !needsOldIndexMaintenance { continue } idxObjRefs[i], idxTableDefs[i], err = builder.compCtx.ResolveIndexTableByRef(objRef, idxDef.IndexTableName, bindCtx.snapshot) @@ -242,11 +244,14 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }) } - var err error for i, idxDef := range tableDef.Indexes { - if skipUniqueIdx[i] { + if skipUniqueIdx[i] && !needsOldIndexMaintenance { continue } + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(idxDef.IndexAlgoParams) + if err != nil { + return 0, err + } idxObjRefs[i], idxTableDefs[i], err = builder.compCtx.ResolveIndexTableByRef(objRef, idxDef.IndexTableName, bindCtx.snapshot) if err != nil { return 0, err @@ -255,11 +260,31 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( oldColName2Idx[idxDef.IndexTableName+"."+catalog.IndexTablePrimaryColName] = oldColName2Idx[tableDef.Name+"."+tableDef.Pkey.PkeyColName] if !indexTableStoresSerializedKey(idxDef) { - oldColName2Idx[idxDef.IndexTableName+"."+catalog.IndexTableIndexColName] = oldColName2Idx[tableDef.Name+"."+indexPrimaryPartName(idxDef)] + partName := indexPrimaryPartName(idxDef) + if prefixLengths[partName] > 0 { + colIdx := tableDef.Name2ColIndex[partName] + partExpr := &plan.Expr{ + Typ: tableDef.Cols[colIdx].Typ, + Expr: &plan.Expr_Col{ + Col: &plan.ColRef{RelPos: oldScanTag, ColPos: colIdx}, + }, + } + idxExpr, err := builder.makeIndexPartExprFromInputExpr(partExpr, partName, prefixLengths) + if err != nil { + return 0, err + } + oldColName2Idx[idxDef.IndexTableName+"."+catalog.IndexTableIndexColName] = [2]int32{ + fullProjTag, int32(len(fullProjList)), + } + fullProjList = append(fullProjList, idxExpr) + } else { + oldColName2Idx[idxDef.IndexTableName+"."+catalog.IndexTableIndexColName] = oldColName2Idx[tableDef.Name+"."+partName] + } } else { args := make([]*plan.Expr, len(idxDef.Parts)) for j, part := range idxDef.Parts { - colIdx := tableDef.Name2ColIndex[catalog.ResolveAlias(part)] + partName := catalog.ResolveAlias(part) + colIdx := tableDef.Name2ColIndex[partName] args[j] = &plan.Expr{ Typ: tableDef.Cols[colIdx].Typ, Expr: &plan.Expr_Col{ @@ -269,6 +294,12 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }, }, } + if prefixLengths[partName] > 0 { + args[j], err = builder.makeIndexPartExprFromInputExpr(args[j], partName, prefixLengths) + if err != nil { + return 0, err + } + } } idxExpr := args[0] @@ -300,6 +331,10 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( if !idxDef.Unique || skipUniqueIdx[i] { continue } + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(idxDef.IndexAlgoParams) + if err != nil { + return 0, err + } var ukPartConds []*plan.Expr for _, part := range idxDef.Parts { colName := catalog.ResolveAlias(part) @@ -323,6 +358,16 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }, }, } + if prefixLengths[colName] > 0 { + lExpr, err = builder.makeIndexPartExprFromInputExpr(lExpr, colName, prefixLengths) + if err != nil { + return 0, err + } + rExpr, err = builder.makeIndexPartExprFromInputExpr(rExpr, colName, prefixLengths) + if err != nil { + return 0, err + } + } partCond, _ := BindFuncExprImplByPlanExpr(builder.GetContext(), "=", []*plan.Expr{lExpr, rExpr}) ukPartConds = append(ukPartConds, partCond) } @@ -363,6 +408,10 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( if !idxDef.Unique || skipUniqueIdx[i] { continue } + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(idxDef.IndexAlgoParams) + if err != nil { + return 0, err + } var ukPartConds []*plan.Expr for _, part := range idxDef.Parts { colName := catalog.ResolveAlias(part) @@ -386,6 +435,16 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }, }, } + if prefixLengths[colName] > 0 { + lExpr, err = builder.makeIndexPartExprFromInputExpr(lExpr, colName, prefixLengths) + if err != nil { + return 0, err + } + rExpr, err = builder.makeIndexPartExprFromInputExpr(rExpr, colName, prefixLengths) + if err != nil { + return 0, err + } + } partCond, _ := BindFuncExprImplByPlanExpr(builder.GetContext(), "=", []*plan.Expr{lExpr, rExpr}) ukPartConds = append(ukPartConds, partCond) } @@ -430,6 +489,14 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( oldMainRowIDPos := oldColName2Idx[tableDef.Name+"."+catalog.Row_ID] oldMainPKPos := oldColName2Idx[tableDef.Name+"."+tableDef.Pkey.PkeyColName] + buildParentFKActions := len(tableDef.RefChildTbls) > 0 + if buildParentFKActions { + enabled, err := IsForeignKeyChecksEnabled(builder.compCtx) + if err != nil { + return 0, err + } + buildParentFKActions = enabled + } replaceDedupOldColList := func(first [2]int32) []plan.ColRef { oldCols := make([]plan.ColRef, 0, 3+len(tableDef.Indexes)) seen := make(map[[2]int32]struct{}, 3+len(tableDef.Indexes)) @@ -446,6 +513,16 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( appendOldCol(first) appendOldCol(oldMainRowIDPos) appendOldCol(oldMainPKPos) + if buildParentFKActions { + // Parent-side FK actions consume the actual old row selected by the + // REPLACE conflict joins. Preserve every base column so FKs that + // reference any UNIQUE key can reuse the delete action planner. + for _, col := range tableDef.Cols { + if pos, ok := oldColName2Idx[tableDef.Name+"."+col.Name]; ok { + appendOldCol(pos) + } + } + } for i, idxDef := range tableDef.Indexes { if idxTableDefs[i] == nil { continue @@ -526,7 +603,7 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( requiredOldCols[catalog.Row_ID] = struct{}{} requiredOldCols[tableDef.Pkey.PkeyColName] = struct{}{} for i, idxDef := range tableDef.Indexes { - if skipUniqueIdx[i] { + if skipUniqueIdx[i] && !needsOldIndexMaintenance { continue } if !indexTableStoresSerializedKey(idxDef) { @@ -535,6 +612,9 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( } captureList := make([]plan.OldColCapture, 0, len(requiredOldCols)) for i, col := range tableDef.Cols { + if buildParentFKActions { + requiredOldCols[col.Name] = struct{}{} + } if _, needed := requiredOldCols[col.Name]; !needed { continue } @@ -659,8 +739,7 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( // get old RowID for index tables for i, idxDef := range tableDef.Indexes { - // Skipped unique index (statically-NULL key): not stored, so no old row to fetch. - if skipUniqueIdx[i] { + if skipUniqueIdx[i] && !needsOldIndexMaintenance { continue } idxTag := builder.genNewBindTag() @@ -729,6 +808,8 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( // unique key, so the deleted row's PK can differ from the inserted row's PK. var replaceOldPkPos int32 var replaceOldPkTyp plan.Type + var replaceOldParentPos []int32 + oldParentColFinalPos := make(map[string]int32) { insertCols := make([]plan.ColRef, len(tableDef.Cols)-1) @@ -788,6 +869,7 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }, }, }) + oldParentColFinalPos[catalog.Row_ID] = deleteCols[0].ColPos oldPkPos := oldColName2Idx[tableDef.Name+"."+tableDef.Pkey.PkeyColName] deleteCols[1].RelPos = finalProjTag @@ -821,6 +903,7 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }, }, }) + oldParentColFinalPos[tableDef.Pkey.PkeyColName] = replaceOldPkPos updateCtxList = append(updateCtxList, &plan.UpdateCtx{ ObjRef: objRef, TableDef: tableDef, @@ -832,18 +915,31 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }) } - for i, idxDef := range tableDef.Indexes { - // A unique index whose key is statically NULL for this statement is not stored - // (serial(...) is NULL), matching the INSERT path which skips index maintenance - // for a NULL key. Nothing to insert into or delete from its index table. - if skipUniqueIdx[i] { + orderedIndexPos := make([]int, len(tableDef.Indexes)) + for i := range orderedIndexPos { + orderedIndexPos[i] = i + } + slices.SortStableFunc(orderedIndexPos, func(left, right int) int { + return strings.Compare( + tableDef.Indexes[left].IndexTableName, + tableDef.Indexes[right].IndexTableName, + ) + }) + for _, i := range orderedIndexPos { + idxDef := tableDef.Indexes[i] + if skipUniqueIdx[i] && !needsOldIndexMaintenance { continue } insertCols := make([]plan.ColRef, 2) deleteCols := make([]plan.ColRef, 2) newIdxPos := colName2Idx[idxDef.IndexTableName+"."+catalog.IndexTableIndexColName] - if indexTableStoresSerializedKey(idxDef) { + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(idxDef.IndexAlgoParams) + if err != nil { + return 0, err + } + partName := indexPrimaryPartName(idxDef) + if indexTableStoresSerializedKey(idxDef) || prefixLengths[partName] > 0 { idxExpr := &plan.Expr{ Typ: fullProjList[newIdxPos].Typ, Expr: &plan.Expr_Col{ @@ -894,6 +990,9 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }, } finalProjList = append(finalProjList, idxExpr) + if len(idxDef.Parts) == 1 && !indexTableStoresSerializedKey(idxDef) && prefixLengths[partName] == 0 { + oldParentColFinalPos[catalog.ResolveAlias(idxDef.Parts[0])] = oldIdxPos + } insertCols[0].RelPos = finalProjTag insertCols[0].ColPos = int32(newIdxPos) @@ -913,13 +1012,16 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( }) if idxDef.Unique { + if !skipUniqueIdx[i] { + lockTargets = append(lockTargets, &plan.LockTarget{ + TableId: idxTableDefs[i].TblId, + ObjRef: idxObjRefs[i], + PrimaryColIdxInBat: int32(newIdxPos), + PrimaryColRelPos: finalProjTag, + PrimaryColTyp: finalProjList[newIdxPos].Typ, + }) + } lockTargets = append(lockTargets, &plan.LockTarget{ - TableId: idxTableDefs[i].TblId, - ObjRef: idxObjRefs[i], - PrimaryColIdxInBat: int32(newIdxPos), - PrimaryColRelPos: finalProjTag, - PrimaryColTyp: finalProjList[newIdxPos].Typ, - }, &plan.LockTarget{ TableId: idxTableDefs[i].TblId, ObjRef: idxObjRefs[i], PrimaryColIdxInBat: int32(oldIdxPos), @@ -929,6 +1031,31 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( } } + if buildParentFKActions { + // Append the auxiliary old-parent image after every existing DML/index + // column. Native index keys rely on the legacy prefix positions above. + replaceOldParentPos = make([]int32, len(tableDef.Cols)) + for i, col := range tableDef.Cols { + if finalPos, ok := oldParentColFinalPos[col.Name]; ok { + replaceOldParentPos[i] = finalPos + continue + } + oldPos, ok := oldColName2Idx[tableDef.Name+"."+col.Name] + if !ok { + return 0, moerr.NewInternalErrorf(builder.GetContext(), + "bind replace err, can not find old parent column %s", col.Name) + } + replaceOldParentPos[i] = int32(len(finalProjList)) + finalProjList = append(finalProjList, &plan.Expr{ + Typ: fullProjList[oldPos[1]].Typ, + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: fullProjTag, + ColPos: oldPos[1], + }}, + }) + } + } + lastNodeID = builder.appendNode(&plan.Node{ NodeType: plan.Node_PROJECT, Children: []int32{lastNodeID}, @@ -946,15 +1073,88 @@ func (builder *QueryBuilder) appendDedupAndMultiUpdateNodesForBindReplace( irregularIndexes, tableDef, objRef) } - if len(lockTargets) > 0 { + if len(lockTargets) > 0 && !buildParentFKActions { lastNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_LOCK_OP, + Children: []int32{lastNodeID}, + TableDef: tableDef, + // LOCK_OP is a pass-through node. Keep the projection tag so a + // following shared SINK can remap every requested column correctly. + BindingTags: []int32{finalProjTag}, + LockTargets: lockTargets, + }, bindCtx) + reCheckifNeedLockWholeTable(builder) + } + + if len(replaceOldParentPos) > 0 { + // Execute parent-side FK actions from the same evaluated and locked old-row + // set consumed by MULTI_UPDATE. This supports VALUES parameters/functions + // and REPLACE SELECT/TABLE without serializing their AST into background SQL. + evaluatedSinkID := appendSinkNode(builder, bindCtx, lastNodeID) + if builder.preserveSinkProjection == nil { + builder.preserveSinkProjection = make(map[int32]struct{}) + } + builder.preserveSinkProjection[evaluatedSinkID] = struct{}{} + evaluatedStep := builder.appendStep(evaluatedSinkID) + + lockedSourceID := appendSinkScanNode(builder, bindCtx, evaluatedStep) + builder.qry.Nodes[lockedSourceID].BindingTags = []int32{finalProjTag} + lockedSourceID = builder.appendNode(&plan.Node{ NodeType: plan.Node_LOCK_OP, - Children: []int32{lastNodeID}, + Children: []int32{lockedSourceID}, TableDef: tableDef, - BindingTags: []int32{builder.genNewBindTag()}, + BindingTags: []int32{finalProjTag}, LockTargets: lockTargets, }, bindCtx) + if builder.preserveLockProjection == nil { + builder.preserveLockProjection = make(map[int32]struct{}) + } + builder.preserveLockProjection[lockedSourceID] = struct{}{} reCheckifNeedLockWholeTable(builder) + + sharedSinkID := appendSinkNode(builder, bindCtx, lockedSourceID) + builder.preserveSinkProjection[sharedSinkID] = struct{}{} + sharedStep := builder.appendStep(sharedSinkID) + + actionSourceID := appendSinkScanNode(builder, bindCtx, sharedStep) + actionInputTag := builder.genNewBindTag() + builder.qry.Nodes[actionSourceID].BindingTags = []int32{actionInputTag} + actionTag := builder.genNewBindTag() + actionProjection := make([]*plan.Expr, len(tableDef.Cols)) + for i, col := range tableDef.Cols { + actionProjection[i] = &plan.Expr{ + Typ: col.Typ, + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: actionInputTag, + ColPos: replaceOldParentPos[i], + }}, + } + } + actionSourceID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, + Children: []int32{actionSourceID}, + ProjectList: actionProjection, + BindingTags: []int32{actionTag}, + }, bindCtx) + actionSinkID := appendSinkNode(builder, bindCtx, actionSourceID) + builder.preserveSinkProjection[actionSinkID] = struct{}{} + actionStep := builder.appendStep(actionSinkID) + + delCtx := getDmlPlanCtx() + delCtx.objRef = objRef + delCtx.tableDef = tableDef + delCtx.sourceStep = actionStep + delCtx.rowIdPos = int(tableDef.Name2ColIndex[catalog.Row_ID]) + delCtx.allDelTableIDs = map[uint64]struct{}{tableDef.TblId: {}} + delCtx.skipTargetDelete = true + err := buildDeletePlans(builder.compCtx, builder, bindCtx, delCtx) + putDmlPlanCtx(delCtx) + if err != nil { + return 0, err + } + + lastNodeID = appendSinkScanNode(builder, bindCtx, sharedStep) + builder.qry.Nodes[lastNodeID].BindingTags = []int32{finalProjTag} } // Self-referencing FK constraint checks are handled by DetectSqls (generated in @@ -1136,8 +1336,23 @@ func (builder *QueryBuilder) appendNodesForReplaceStmt( idxTableName := idxDef.IndexTableName colName2Idx[idxTableName+"."+catalog.IndexTablePrimaryColName] = pkPos + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(idxDef.IndexAlgoParams) + if err != nil { + return 0, nil, nil, err + } if !indexTableStoresSerializedKey(idxDef) { - colName2Idx[idxTableName+"."+catalog.IndexTableIndexColName] = colName2Idx[tableDef.Name+"."+indexPrimaryPartName(idxDef)] + partName := indexPrimaryPartName(idxDef) + partPos := colName2Idx[tableDef.Name+"."+partName] + if prefixLengths[partName] > 0 { + idxExpr, err := builder.makeIndexPartExprFromInputExpr(projList2[partPos], partName, prefixLengths) + if err != nil { + return 0, nil, nil, err + } + colName2Idx[idxTableName+"."+catalog.IndexTableIndexColName] = int32(len(projList2)) + projList2 = append(projList2, idxExpr) + } else { + colName2Idx[idxTableName+"."+catalog.IndexTableIndexColName] = partPos + } } else { argsLen := len(idxDef.Parts) args := make([]*plan.Expr, argsLen) @@ -1149,7 +1364,15 @@ func (builder *QueryBuilder) appendNodesForReplaceStmt( errMsg := fmt.Sprintf("bind insert err, can not find colName = %s", idxDef.Parts[k]) return 0, nil, nil, moerr.NewInternalError(builder.GetContext(), errMsg) } - args[k] = DeepCopyExpr(projList2[colPos]) + partName := catalog.ResolveAlias(idxDef.Parts[k]) + if prefixLengths[partName] > 0 { + args[k], err = builder.makeIndexPartExprFromInputExpr(projList2[colPos], partName, prefixLengths) + if err != nil { + return 0, nil, nil, err + } + } else { + args[k] = DeepCopyExpr(projList2[colPos]) + } } funcName := "serial" diff --git a/pkg/sql/plan/build.go b/pkg/sql/plan/build.go index f331072a214d1..f890ebdc36d57 100644 --- a/pkg/sql/plan/build.go +++ b/pkg/sql/plan/build.go @@ -183,7 +183,28 @@ func bindAndOptimizeReplaceQuery(ctx CompilerContext, stmt *tree.Replace, isPrep if err != nil { return nil, err } - if len(tblInfo.tableDefs) == 1 { + // FK checks/actions are all disabled when foreign_key_checks is off, the + // same way MySQL skips foreign-key enforcement. Gate every FK SQL below + // (self-referencing checks, the RESTRICT pre-check, and the non-self + // parent-side actions) under one guard so the behavior is consistent. + fkChecksEnabled, err := IsForeignKeyChecksEnabled(ctx) + if err != nil { + return nil, err + } + if len(tblInfo.tableDefs) == 1 && + (len(tblInfo.tableDefs[0].Fkeys) > 0 || len(tblInfo.tableDefs[0].RefChildTbls) > 0) { + // The presence or absence of DetectSqls depends on the session's + // foreign_key_checks value. Keep the plan FK-sensitive even when the + // variable is currently off, otherwise a cached plan built without the + // checks could survive after they are enabled. + query.HasForeignKeyAction = true + } + if fkChecksEnabled && len(tblInfo.tableDefs) == 1 { + if len(tblInfo.tableDefs[0].RefChildTbls) > 0 { + // Parent-side actions are part of the modern REPLACE plan. Keep a + // marker solely for the optimistic-transaction fail-closed guard. + query.DetectSqls = append(query.DetectSqls, "REPLACE_PARENT_PLAN:") + } sqls, err := genSqlsForCheckFKSelfRefer( ctx.GetContext(), tblInfo.objRef[0].SchemaName, @@ -194,7 +215,7 @@ func bindAndOptimizeReplaceQuery(ctx CompilerContext, stmt *tree.Replace, isPrep if err != nil { return nil, err } - query.DetectSqls = sqls + query.DetectSqls = append(query.DetectSqls, sqls...) // Generate pre-check SQLs for parent→child safety (RESTRICT). preCheckSqls, err := genPreCheckSqlsForReplaceFKSelfRefer( diff --git a/pkg/sql/plan/build_dml_util.go b/pkg/sql/plan/build_dml_util.go index 29833d36c8a72..c838a27e5283d 100644 --- a/pkg/sql/plan/build_dml_util.go +++ b/pkg/sql/plan/build_dml_util.go @@ -31,6 +31,7 @@ import ( "github.com/matrixorigin/matrixone/pkg/container/types" "github.com/matrixorigin/matrixone/pkg/defines" "github.com/matrixorigin/matrixone/pkg/logutil" + lockpb "github.com/matrixorigin/matrixone/pkg/pb/lock" "github.com/matrixorigin/matrixone/pkg/pb/plan" "github.com/matrixorigin/matrixone/pkg/sql/parsers/tree" "github.com/matrixorigin/matrixone/pkg/sql/plan/function" @@ -85,6 +86,7 @@ type dmlPlanCtx struct { tableDef *TableDef beginIdx int sourceStep int32 + sourceTag int32 isMulti bool // needAggFilter drives two behaviours: an any_value aggregation for dedup // and an isnotnull(row_id) filter for join-target NULL-row protection. After @@ -104,20 +106,26 @@ type dmlPlanCtx struct { updatePkCol bool //if update stmt will update the primary key or one of pks pkFilterExprs []*Expr isDeleteWithoutFilters bool + // skipTargetDelete reuses the parent-reference action planner for a row set + // whose base-table delete is owned by another operator (modern REPLACE). + // Recursive child actions still build their normal delete/update branches. + skipTargetDelete bool + preserveUpdateSourceProjection bool } // information of deleteNode, which is about the deleted table type deleteNodeInfo struct { - objRef *ObjectRef - tableDef *TableDef - IsClusterTable bool - deleteIndex int // The array index position of the rowid column - indexTableNames []string - foreignTbl []uint64 - addAffectedRows bool // for hidden table, should not update affect Rows, e.g. delete 1 row from table t with schema like a int, b unique key, c key, affact rows should be 1 instead of 3 - pkPos int - pkTyp plan.Type - lockTable bool + objRef *ObjectRef + tableDef *TableDef + IsClusterTable bool + deleteIndex int // The array index position of the rowid column + indexTableNames []string + foreignTbl []uint64 + addAffectedRows bool // for hidden table, should not update affect Rows, e.g. delete 1 row from table t with schema like a int, b unique key, c key, affact rows should be 1 instead of 3 + pkPos int + pkTyp plan.Type + lockTable bool + preserveProjection bool } // buildInsertPlans build insert plan. @@ -199,11 +207,29 @@ func buildUpdatePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC var err error // sink_scan -> project -> [agg] -> [filter] -> sink lastNodeId := appendSinkScanNode(builder, bindCtx, updatePlanCtx.sourceStep) + if updatePlanCtx.preserveUpdateSourceProjection { + if builder.preserveScanProjection == nil { + builder.preserveScanProjection = make(map[int32]struct{}) + } + builder.preserveScanProjection[lastNodeId] = struct{}{} + if builder.positionalSinkScans == nil { + builder.positionalSinkScans = make(map[int32]struct{}) + } + builder.positionalSinkScans[lastNodeId] = struct{}{} + } lastNodeId, err = makePreUpdateDeletePlan(ctx, builder, bindCtx, updatePlanCtx, lastNodeId) if err != nil { return err } - lastNodeId = appendSinkNode(builder, bindCtx, lastNodeId) + if updatePlanCtx.preserveUpdateSourceProjection { + lastNodeId = appendSinkNode(builder, bindCtx, lastNodeId) + if builder.preserveSinkProjection == nil { + builder.preserveSinkProjection = make(map[int32]struct{}) + } + builder.preserveSinkProjection[lastNodeId] = struct{}{} + } else { + lastNodeId = appendSinkNode(builder, bindCtx, lastNodeId) + } nextSourceStep := builder.appendStep(lastNodeId) updatePlanCtx.sourceStep = nextSourceStep @@ -215,6 +241,10 @@ func buildUpdatePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC // sink_scan -> project -> preinsert -> sink lastNodeId = appendSinkScanNode(builder, bindCtx, updatePlanCtx.sourceStep) + if updatePlanCtx.preserveUpdateSourceProjection { + builder.preserveScanProjection[lastNodeId] = struct{}{} + builder.positionalSinkScans[lastNodeId] = struct{}{} + } lastNode := builder.qry.Nodes[lastNodeId] newCols := make([]*ColDef, 0, len(updatePlanCtx.tableDef.Cols)) oldRowIdPos := len(updatePlanCtx.tableDef.Cols) - 1 @@ -261,9 +291,18 @@ func buildUpdatePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC lastNodeId = builder.appendNode(projectNode, bindCtx) //append preinsert node lastNodeId = appendPreInsertNode(builder, bindCtx, updatePlanCtx.objRef, updatePlanCtx.tableDef, lastNodeId, true) + if updatePlanCtx.preserveUpdateSourceProjection { + if builder.preservePreInsertProjection == nil { + builder.preservePreInsertProjection = make(map[int32]struct{}) + } + builder.preservePreInsertProjection[lastNodeId] = struct{}{} + } //append sink node lastNodeId = appendSinkNode(builder, bindCtx, lastNodeId) + if updatePlanCtx.preserveUpdateSourceProjection { + builder.preserveSinkProjection[lastNodeId] = struct{}{} + } sourceStep := builder.appendStep(lastNodeId) // build insert plan. @@ -307,6 +346,227 @@ func checkDeleteOptToTruncate(ctx CompilerContext) (bool, error) { } } +func appendRecursiveCascadeLockNode( + builder *QueryBuilder, + bindCtx *BindContext, + delCtx *dmlPlanCtx, + sourceNodeID int32, +) (int32, error) { + pkPos, pkTyp := getPkPos(delCtx.tableDef, false) + if pkPos < 0 { + return 0, moerr.NewInternalErrorf( + builder.GetContext(), "cascade lock key is unavailable for table %s", delCtx.tableDef.Name) + } + + lockTag := delCtx.sourceTag + var rowProject []*plan.Expr + if lockTag != 0 { + rowProject = getProjectionByLastNodeWithTag(builder, sourceNodeID, lockTag) + } else { + rowProject = getProjectionByLastNode(builder, sourceNodeID) + } + lockProject := slices.Clone(rowProject) + lockTargets := []*plan.LockTarget{{ + TableId: delCtx.tableDef.TblId, ObjRef: delCtx.objRef, + PrimaryColIdxInBat: int32(pkPos), PrimaryColRelPos: lockTag, + PrimaryColTyp: pkTyp, Mode: lockpb.LockMode_Exclusive, + }} + targetTables := map[uint64]struct{}{delCtx.tableDef.TblId: {}} + baseTableLock := false + + parentColName := make(map[uint64]string, len(delCtx.tableDef.Cols)) + parentColPos := make(map[string]int32, len(delCtx.tableDef.Cols)) + for i, col := range delCtx.tableDef.Cols { + parentColName[col.ColId] = col.Name + parentColPos[col.Name] = int32(i) + } + var pkeyNames []string + if delCtx.tableDef.Pkey != nil { + pkeyNames = delCtx.tableDef.Pkey.Names + if len(pkeyNames) == 0 && delCtx.tableDef.Pkey.PkeyColName != "" { + pkeyNames = []string{delCtx.tableDef.Pkey.PkeyColName} + } + } + partsEqual := func(parts, names []string) bool { + if len(parts) != len(names) { + return false + } + for i := range parts { + if catalog.ResolveAlias(parts[i]) != catalog.ResolveAlias(names[i]) { + return false + } + } + return true + } + + seenChild := make(map[uint64]struct{}, len(delCtx.tableDef.RefChildTbls)) + for _, childTableID := range delCtx.tableDef.RefChildTbls { + if _, ok := seenChild[childTableID]; ok { + continue + } + seenChild[childTableID] = struct{}{} + var childTableDef *plan.TableDef + var err error + if childTableID == 0 { + childTableDef = delCtx.tableDef + } else { + _, childTableDef, err = builder.compCtx.ResolveById(childTableID, bindCtx.snapshot) + if err != nil { + return 0, err + } + } + if childTableDef == nil { + return 0, moerr.NewInternalErrorf( + builder.GetContext(), "cascade child table %d is unavailable", childTableID) + } + for _, fk := range childTableDef.Fkeys { + selfRefer := fk.ForeignTbl == 0 && childTableDef.TblId == delCtx.tableDef.TblId + if fk.ForeignTbl != delCtx.tableDef.TblId && !selfRefer { + continue + } + referencedNames := make([]string, len(fk.ForeignCols)) + for i, colID := range fk.ForeignCols { + name, ok := parentColName[colID] + if !ok { + return 0, moerr.NewInternalErrorf(builder.GetContext(), + "foreign-key referenced column %d is unavailable in table %s", colID, delCtx.tableDef.Name) + } + referencedNames[i] = name + } + if partsEqual(pkeyNames, referencedNames) { + continue + } + + var matchedIndex *plan.IndexDef + for _, idxDef := range delCtx.tableDef.Indexes { + if idxDef.Unique && partsEqual(idxDef.Parts, referencedNames) { + matchedIndex = idxDef + break + } + } + if matchedIndex == nil { + if !baseTableLock { + lockTargets = append(lockTargets, &plan.LockTarget{ + TableId: delCtx.tableDef.TblId, ObjRef: delCtx.objRef, + PrimaryColIdxInBat: int32(pkPos), PrimaryColRelPos: lockTag, + PrimaryColTyp: pkTyp, Mode: lockpb.LockMode_Exclusive, LockTable: true, + }) + baseTableLock = true + } + continue + } + + indexObjRef, indexTableDef, err := builder.compCtx.ResolveIndexTableByRef( + delCtx.objRef, matchedIndex.IndexTableName, bindCtx.snapshot) + if err != nil { + return 0, err + } + if _, ok := targetTables[indexTableDef.TblId]; ok { + continue + } + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(matchedIndex.IndexAlgoParams) + if err != nil { + return 0, err + } + keyParts := make([]*plan.Expr, len(referencedNames)) + for i, name := range referencedNames { + pos, ok := parentColPos[name] + if !ok { + return 0, moerr.NewInternalErrorf(builder.GetContext(), + "foreign-key referenced column %s is unavailable in table %s", name, delCtx.tableDef.Name) + } + inputExpr := &plan.Expr{Typ: delCtx.tableDef.Cols[pos].Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: lockTag, ColPos: pos, Name: name, + }}} + keyParts[i], err = builder.makeIndexPartExprFromInputExpr(inputExpr, name, prefixLengths) + if err != nil { + return 0, err + } + } + var keyExpr *plan.Expr + if indexTableStoresSerializedKey(matchedIndex) { + keyExpr, err = BindFuncExprImplByPlanExpr(builder.GetContext(), "serial", keyParts) + if err != nil { + return 0, err + } + } else { + keyExpr = keyParts[0] + } + indexPkPos, indexPkTyp := getPkPos(indexTableDef, false) + if indexPkPos < 0 { + return 0, moerr.NewInternalErrorf(builder.GetContext(), + "cascade lock index table %s has no primary key", indexTableDef.Name) + } + lockProject = append(lockProject, keyExpr) + lockTargets = append(lockTargets, &plan.LockTarget{ + TableId: indexTableDef.TblId, ObjRef: indexObjRef, + PrimaryColIdxInBat: int32(len(lockProject) - 1), PrimaryColRelPos: lockTag, + PrimaryColTyp: indexPkTyp, Mode: lockpb.LockMode_Exclusive, + }) + targetTables[indexTableDef.TblId] = struct{}{} + } + } + + slices.SortStableFunc(lockTargets, func(left, right *plan.LockTarget) int { + if left.TableId < right.TableId { + return -1 + } + if left.TableId > right.TableId { + return 1 + } + if !left.LockTable && right.LockTable { + return -1 + } + if left.LockTable && !right.LockTable { + return 1 + } + return 0 + }) + if len(lockProject) > len(rowProject) { + sourceNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, Children: []int32{sourceNodeID}, + ProjectList: lockProject, BindingTags: func() []int32 { + if lockTag == 0 { + return nil + } + return []int32{lockTag} + }(), + }, bindCtx) + } + lockedNodeID := builder.appendNode(&plan.Node{ + NodeType: plan.Node_LOCK_OP, Children: []int32{sourceNodeID}, + TableDef: delCtx.tableDef, BindingTags: func() []int32 { + if lockTag == 0 { + return nil + } + return []int32{lockTag} + }(), + LockTargets: lockTargets, + }, bindCtx) + if builder.preserveLockProjection == nil { + builder.preserveLockProjection = make(map[int32]struct{}) + } + builder.preserveLockProjection[lockedNodeID] = struct{}{} + if len(lockProject) > len(rowProject) { + var lockedProject []*plan.Expr + if lockTag != 0 { + lockedProject = getProjectionByLastNodeWithTag(builder, lockedNodeID, lockTag) + } else { + lockedProject = getProjectionByLastNode(builder, lockedNodeID) + } + lockedNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, Children: []int32{lockedNodeID}, + ProjectList: slices.Clone(lockedProject[:len(rowProject)]), BindingTags: func() []int32 { + if lockTag == 0 { + return nil + } + return []int32{lockTag} + }(), + }, bindCtx) + } + return lockedNodeID, nil +} + // buildDeletePlans build preinsert plan. /* [o1]sink_scan -> join[u1] -> sink @@ -323,37 +583,82 @@ func checkDeleteOptToTruncate(ctx CompilerContext) (bool, error) { [o1]sink_scan -> join[f1 inner join c4 on f1.id = c4.fid, get c3.*, update cols] -> sink ...(like update) // update stmt: if have refChild table with cascade */ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindContext, delCtx *dmlPlanCtx) error { + appendDeleteSourceScan := func() int32 { + var nodeID int32 + if delCtx.sourceTag != 0 { + nodeID = appendSinkScanNodeWithTag(builder, bindCtx, delCtx.sourceStep, delCtx.sourceTag) + } else { + nodeID = appendSinkScanNode(builder, bindCtx, delCtx.sourceStep) + } + if delCtx.skipTargetDelete { + if builder.preserveScanProjection == nil { + builder.preserveScanProjection = make(map[int32]struct{}) + } + builder.preserveScanProjection[nodeID] = struct{}{} + if builder.positionalSinkScans == nil { + builder.positionalSinkScans = make(map[int32]struct{}) + } + builder.positionalSinkScans[nodeID] = struct{}{} + } + return nodeID + } + if delCtx.isFkRecursionCall && len(delCtx.tableDef.RefChildTbls) > 0 { + lockedSourceID := appendDeleteSourceScan() + lockTag := delCtx.sourceTag + var err error + lockedSourceID, err = appendRecursiveCascadeLockNode(builder, bindCtx, delCtx, lockedSourceID) + if err != nil { + return err + } + if builder.preserveLockProjection == nil { + builder.preserveLockProjection = make(map[int32]struct{}) + } + builder.preserveLockProjection[lockedSourceID] = struct{}{} + var lockedSinkID int32 + if lockTag != 0 { + lockedSinkID = appendSinkNodeWithTag(builder, bindCtx, lockedSourceID, lockTag) + } else { + lockedSinkID = appendSinkNode(builder, bindCtx, lockedSourceID) + } + if builder.preserveSinkProjection == nil { + builder.preserveSinkProjection = make(map[int32]struct{}) + } + builder.preserveSinkProjection[lockedSinkID] = struct{}{} + delCtx.sourceStep = builder.appendStep(lockedSinkID) + } // When the same child table is reached multiple times (e.g. two FKs pointing to the // same parent), we merge the delete sources with a UNION chain. `deleteNode[tblId]` // always stores the SINK node id of the current merged plan so every subsequent entry // can follow the same code path regardless of how many times we merge. - if sinkNodeId, ok := builder.deleteNode[delCtx.tableDef.TblId]; ok { - step := getStepByNodeId(builder, sinkNodeId) - if step == -1 || delCtx.sourceStep == -1 { - panic("steps should not be -1") - } - - oldDelPlanSinkScanNodeId := appendSinkScanNode(builder, bindCtx, int32(step)) - thisDelPlanSinkScanNodeId := appendSinkScanNode(builder, bindCtx, delCtx.sourceStep) - unionProjection := getProjectionByLastNode(builder, sinkNodeId) - unionNode := &plan.Node{ - NodeType: plan.Node_UNION, - Children: []int32{oldDelPlanSinkScanNodeId, thisDelPlanSinkScanNodeId}, - ProjectList: unionProjection, - } - unionNodeId := builder.appendNode(unionNode, bindCtx) - newSinkNodeId := appendSinkNode(builder, bindCtx, unionNodeId) - endStep := builder.appendStep(newSinkNodeId) - for i, n := range builder.qry.Nodes { - if n.NodeType == plan.Node_SINK_SCAN && n.SourceStep[0] == int32(step) && i != int(oldDelPlanSinkScanNodeId) { - n.SourceStep[0] = endStep + if !delCtx.skipTargetDelete { + if sinkNodeId, ok := builder.deleteNode[delCtx.tableDef.TblId]; ok { + step := getStepByNodeId(builder, sinkNodeId) + if step == -1 || delCtx.sourceStep == -1 { + panic("steps should not be -1") + } + + oldDelPlanSinkScanNodeId := appendSinkScanNode(builder, bindCtx, int32(step)) + thisDelPlanSinkScanNodeId := appendSinkScanNode(builder, bindCtx, delCtx.sourceStep) + unionProjection := getProjectionByLastNode(builder, sinkNodeId) + unionNode := &plan.Node{ + NodeType: plan.Node_UNION, + Children: []int32{oldDelPlanSinkScanNodeId, thisDelPlanSinkScanNodeId}, + ProjectList: unionProjection, } + unionNodeId := builder.appendNode(unionNode, bindCtx) + newSinkNodeId := appendSinkNode(builder, bindCtx, unionNodeId) + endStep := builder.appendStep(newSinkNodeId) + for i, n := range builder.qry.Nodes { + if n.NodeType == plan.Node_SINK_SCAN && n.SourceStep[0] == int32(step) && i != int(oldDelPlanSinkScanNodeId) { + n.SourceStep[0] = endStep + } + } + // Store the new SINK (not the UNION) so the next merge can find the step directly. + builder.deleteNode[delCtx.tableDef.TblId] = newSinkNodeId + return nil + } else { + builder.deleteNode[delCtx.tableDef.TblId] = builder.qry.Steps[delCtx.sourceStep] } - // Store the new SINK (not the UNION) so the next merge can find the step directly. - builder.deleteNode[delCtx.tableDef.TblId] = newSinkNodeId - return nil - } else { - builder.deleteNode[delCtx.tableDef.TblId] = builder.qry.Steps[delCtx.sourceStep] } isUpdate := delCtx.updateColLength > 0 @@ -363,45 +668,48 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC // both UK and SK. To handle SK case, we will have flags to indicate if it's UK or SK. canTruncate := delCtx.isDeleteWithoutFilters - accountId, err := ctx.GetAccountId() - if err != nil { - return err - } - enabled, err := IsForeignKeyChecksEnabled(ctx) if err != nil { return err } - deleteOptToTruncate, err := checkDeleteOptToTruncate(ctx) - if err != nil { - return err - } - - if enabled && len(delCtx.tableDef.RefChildTbls) > 0 || - delCtx.tableDef.ViewSql != nil || - (util.TableIsClusterTable(delCtx.tableDef.GetTableType()) && accountId != catalog.System_Account) || - delCtx.objRef.PubInfo != nil || !deleteOptToTruncate || - delCtx.tableDef.Partition != nil { - canTruncate = false + if !delCtx.skipTargetDelete { + accountId, err := ctx.GetAccountId() + if err != nil { + return err + } + deleteOptToTruncate, err := checkDeleteOptToTruncate(ctx) + if err != nil { + return err + } + if enabled && len(delCtx.tableDef.RefChildTbls) > 0 || + delCtx.tableDef.ViewSql != nil || + (util.TableIsClusterTable(delCtx.tableDef.GetTableType()) && accountId != catalog.System_Account) || + delCtx.objRef.PubInfo != nil || !deleteOptToTruncate || + delCtx.tableDef.Partition != nil { + canTruncate = false + } } - // create delete index plans - err = buildDeleteIndexPlans(ctx, builder, bindCtx, delCtx) - if err != nil { - return err - } + lastNodeId := appendDeleteSourceScan() + if !delCtx.skipTargetDelete { + // create delete index plans + err = buildDeleteIndexPlans(ctx, builder, bindCtx, delCtx) + if err != nil { + return err + } - // delete origin table - lastNodeId := appendSinkScanNode(builder, bindCtx, delCtx.sourceStep) - pkPos, pkTyp := getPkPos(delCtx.tableDef, false) - delNodeInfo := makeDeleteNodeInfo(ctx, delCtx.objRef, delCtx.tableDef, delCtx.rowIdPos, true, pkPos, pkTyp, delCtx.lockTable) - lastNodeId, err = makeOneDeletePlan(builder, bindCtx, lastNodeId, delNodeInfo, false, false, canTruncate) - putDeleteNodeInfo(delNodeInfo) - if err != nil { - return err + // delete origin table + pkPos, pkTyp := getPkPos(delCtx.tableDef, false) + delNodeInfo := makeDeleteNodeInfo(ctx, delCtx.objRef, delCtx.tableDef, delCtx.rowIdPos, true, pkPos, pkTyp, delCtx.lockTable) + delNodeInfo.preserveProjection = delCtx.sourceTag != 0 + lastNodeId, err = makeOneDeletePlan(builder, bindCtx, lastNodeId, delNodeInfo, false, false, canTruncate) + putDeleteNodeInfo(delNodeInfo) + if err != nil { + return err + } + builder.appendStep(lastNodeId) } - builder.appendStep(lastNodeId) // if some table references to this table if enabled && len(delCtx.tableDef.RefChildTbls) > 0 { @@ -444,6 +752,18 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC childId2name := make(map[uint64]string) childProjectList := make([]*Expr, len(childTableDef.Cols)) childForJoinProject := make([]*Expr, len(childTableDef.Cols)) + parentRelPos := int32(0) + childRelPos := int32(1) + var childScanTag int32 + var childBindingTags []int32 + if delCtx.skipTargetDelete { + childScanTag = builder.genNewBindTag() + childBindingTags = []int32{childScanTag} + childRelPos = childScanTag + } + if delCtx.sourceTag != 0 { + parentRelPos = delCtx.sourceTag + } childRowIdPos := -1 for idx, col := range childTableDef.Cols { childPosMap[col.Name] = int32(idx) @@ -462,7 +782,7 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC Typ: col.Typ, Expr: &plan.Expr_Col{ Col: &plan.ColRef{ - RelPos: 1, + RelPos: childRelPos, ColPos: int32(idx), Name: col.Name, }, @@ -472,8 +792,260 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC childRowIdPos = idx } } + childScanProject := childProjectList + if delCtx.skipTargetDelete { + // Column pruning reconstructs the physical scan projection from + // referenced child columns. Starting with the full logical list + // would leave stale original ordinals ahead of that compact list. + childScanProject = nil + } + + combinedSetNull := make(map[*plan.ForeignKeyDef]struct{}) + if !isUpdate { + setNullFks := make([]*plan.ForeignKeyDef, 0, len(childTableDef.Fkeys)) + for _, fk := range childTableDef.Fkeys { + fkSelfReferCond := fk.ForeignTbl == 0 && childTableDef.TblId == delCtx.tableDef.TblId + if (fk.ForeignTbl == delCtx.tableDef.TblId || fkSelfReferCond) && + fk.OnDelete == plan.ForeignKeyDef_SET_NULL { + setNullFks = append(setNullFks, fk) + } + } + if len(setNullFks) > 1 { + builder.qry.HasForeignKeyAction = true + for _, fk := range setNullFks { + combinedSetNull[fk] = struct{}{} + } + + childTag := builder.genNewBindTag() + parentTag := builder.genNewBindTag() + childNodeID := builder.appendNode(&plan.Node{ + NodeType: plan.Node_TABLE_SCAN, Stats: &plan.Stats{}, + ObjRef: childObjRef, TableDef: CloneTableDefForPlan(childTableDef, true), + ProjectList: childProjectList, BindingTags: []int32{childTag}, + }, bindCtx) + if builder.preserveScanProjection == nil { + builder.preserveScanProjection = make(map[int32]struct{}) + } + builder.preserveScanProjection[childNodeID] = struct{}{} + parentNodeID := appendDeleteSourceScan() + parentNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, Children: []int32{parentNodeID}, + ProjectList: DeepCopyExprList(builder.qry.Nodes[parentNodeID].ProjectList), + BindingTags: []int32{parentTag}, + }, bindCtx) + + fkMatches := make([]*Expr, len(setNullFks)) + markerByColumn := make(map[string][]int) + var anyMatch *Expr + for fkIdx, fk := range setNullFks { + for i, childColID := range fk.Cols { + childName := childId2name[childColID] + parentName := idNameMap[fk.ForeignCols[i]] + leftExpr := &Expr{Typ: *childTypMap[childName], Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: childTag, ColPos: childPosMap[childName], Name: childName, + }}} + rightExpr := &Expr{Typ: *nameTypMap[parentName], Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: parentTag, ColPos: nameIdxMap[parentName], Name: parentName, + }}} + partMatch, bindErr := BindFuncExprImplByPlanExpr( + builder.GetContext(), "=", []*Expr{leftExpr, rightExpr}) + if bindErr != nil { + return bindErr + } + if fkMatches[fkIdx] == nil { + fkMatches[fkIdx] = partMatch + } else { + fkMatches[fkIdx], err = BindFuncExprImplByPlanExpr( + builder.GetContext(), "and", []*Expr{fkMatches[fkIdx], partMatch}) + if err != nil { + return err + } + } + markerByColumn[childName] = append(markerByColumn[childName], fkIdx) + } + if anyMatch == nil { + anyMatch = DeepCopyExpr(fkMatches[fkIdx]) + } else { + anyMatch, err = BindFuncExprImplByPlanExpr( + builder.GetContext(), "or", []*Expr{anyMatch, DeepCopyExpr(fkMatches[fkIdx])}) + if err != nil { + return err + } + } + } + + joinProjection := make([]*Expr, 0, len(childTableDef.Cols)+len(fkMatches)) + for i, col := range childTableDef.Cols { + joinProjection = append(joinProjection, &Expr{ + Typ: col.Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: childTag, ColPos: int32(i), Name: col.Name, + }}, + }) + } + joinProjection = append(joinProjection, DeepCopyExprList(fkMatches)...) + combinedTag := builder.genNewBindTag() + combinedNodeID := builder.appendNode(&plan.Node{ + NodeType: plan.Node_JOIN, Children: []int32{childNodeID, parentNodeID}, + JoinType: plan.Node_INNER, OnList: []*Expr{anyMatch}, + }, bindCtx) + combinedNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, Children: []int32{combinedNodeID}, + ProjectList: joinProjection, BindingTags: []int32{combinedTag}, + }, bindCtx) + groupTag := builder.genNewBindTag() + aggTag := builder.genNewBindTag() + groupBy := make([]*Expr, 0, 2) + childGroupPos := make([]int32, len(childTableDef.Cols)) + childAggPos := make([]int32, len(childTableDef.Cols)) + aggList := make([]*Expr, 0, len(childTableDef.Cols)-2+len(fkMatches)) + for i, col := range childTableDef.Cols { + colExpr := &Expr{Typ: col.Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: combinedTag, ColPos: int32(i), Name: col.Name, + }}} + if col.Name == catalog.Row_ID || col.Name == catalog.FakePrimaryKeyColName { + childGroupPos[i] = int32(len(groupBy)) + childAggPos[i] = -1 + groupBy = append(groupBy, colExpr) + continue + } + childAggPos[i] = int32(len(aggList)) + colAgg, bindErr := BindFuncExprImplByPlanExpr( + builder.GetContext(), "any_value", []*Expr{colExpr}) + if bindErr != nil { + return bindErr + } + aggList = append(aggList, colAgg) + } + markerAggOffset := len(aggList) + for i := range fkMatches { + marker := &Expr{Typ: fkMatches[i].Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: combinedTag, ColPos: int32(len(childTableDef.Cols) + i), + }}} + markerAgg, bindErr := BindFuncExprImplByPlanExpr( + builder.GetContext(), "max", []*Expr{marker}) + if bindErr != nil { + return bindErr + } + aggList = append(aggList, markerAgg) + } + combinedNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_AGG, Children: []int32{combinedNodeID}, + GroupBy: groupBy, AggList: aggList, BindingTags: []int32{groupTag, aggTag}, + SpillMem: builder.aggSpillMem, + }, bindCtx) + actionTag := builder.genNewBindTag() + actionProjection := make([]*Expr, 0, len(childTableDef.Cols)+len(fkMatches)) + for i, col := range childTableDef.Cols { + relPos := aggTag + colPos := childAggPos[i] + if childAggPos[i] < 0 { + relPos = groupTag + colPos = childGroupPos[i] + } + actionProjection = append(actionProjection, &Expr{ + Typ: col.Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: relPos, ColPos: colPos, Name: col.Name, + }}, + }) + } + for i := range fkMatches { + expr := aggList[markerAggOffset+i] + actionProjection = append(actionProjection, &Expr{ + Typ: expr.Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: aggTag, ColPos: int32(markerAggOffset + i), + }}, + }) + } + combinedNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, Children: []int32{combinedNodeID}, + ProjectList: actionProjection, BindingTags: []int32{actionTag}, + }, bindCtx) + actionSinkID := appendSinkNodeWithTag(builder, bindCtx, combinedNodeID, actionTag) + if builder.preserveSinkProjection == nil { + builder.preserveSinkProjection = make(map[int32]struct{}) + } + builder.preserveSinkProjection[actionSinkID] = struct{}{} + actionStep := builder.appendStep(actionSinkID) + combinedNodeID = builder.appendTaggedSinkScan(bindCtx, actionStep, actionTag) + if builder.preserveScanProjection == nil { + builder.preserveScanProjection = make(map[int32]struct{}) + } + builder.preserveScanProjection[combinedNodeID] = struct{}{} + + updateMap := make(map[string]int) + insertColPos := make([]int, 0, len(childTableDef.Cols)-1) + projectList := make([]*Expr, len(childTableDef.Cols)) + for i, col := range childTableDef.Cols { + projectList[i] = &Expr{Typ: col.Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: actionTag, ColPos: int32(i), Name: col.Name, + }}} + } + for columnName, markers := range markerByColumn { + var matched *Expr + for _, markerIdx := range markers { + marker := &Expr{Typ: fkMatches[markerIdx].Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: actionTag, ColPos: int32(len(childTableDef.Cols) + markerIdx), + }}} + if matched == nil { + matched = marker + } else { + matched, err = BindFuncExprImplByPlanExpr( + builder.GetContext(), "or", []*Expr{matched, marker}) + if err != nil { + return err + } + } + } + colPos := childPosMap[columnName] + nullExpr := &Expr{Typ: *childTypMap[columnName], Expr: &plan.Expr_Lit{Lit: &Const{Isnull: true}}} + updated, bindErr := BindFuncExprImplByPlanExpr( + builder.GetContext(), "if", []*Expr{matched, nullExpr, DeepCopyExpr(projectList[colPos])}) + if bindErr != nil { + return bindErr + } + updateMap[columnName] = len(projectList) + projectList = append(projectList, updated) + } + for i, col := range childTableDef.Cols { + if col.Name == catalog.Row_ID || col.Name == catalog.CPrimaryKeyColName { + continue + } + if pos, ok := updateMap[col.Name]; ok { + insertColPos = append(insertColPos, pos) + } else { + insertColPos = append(insertColPos, i) + } + } + combinedNodeID = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, Children: []int32{combinedNodeID}, ProjectList: projectList, + }, bindCtx) + combinedNodeID = appendSinkNode(builder, bindCtx, combinedNodeID) + builder.preserveSinkProjection[combinedNodeID] = struct{}{} + combinedStep := builder.appendStep(combinedNodeID) + upPlanCtx := getDmlPlanCtx() + upPlanCtx.objRef = childObjRef + upPlanCtx.tableDef = childTableDef + upPlanCtx.updateColLength = len(updateMap) + upPlanCtx.rowIdPos = childRowIdPos + upPlanCtx.sourceStep = combinedStep + upPlanCtx.updateColPosMap = updateMap + upPlanCtx.allDelTableIDs = map[uint64]struct{}{} + upPlanCtx.insertColPos = insertColPos + upPlanCtx.isFkRecursionCall = true + upPlanCtx.updatePkCol = false + upPlanCtx.preserveUpdateSourceProjection = true + err = buildUpdatePlans(ctx, builder, bindCtx, upPlanCtx, false) + putDmlPlanCtx(upPlanCtx) + if err != nil { + return err + } + } + } for _, fk := range childTableDef.Fkeys { + if _, ok := combinedSetNull[fk]; ok { + continue + } //child table fk self refer //if the child table in the delete table list, something must be done fkSelfReferCond := fk.ForeignTbl == 0 && @@ -497,6 +1069,17 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC joinConds := make([]*Expr, len(fk.Cols)) rightConds := make([]*Expr, len(fk.Cols)) leftConds := make([]*Expr, len(fk.Cols)) + parentActionTag := parentRelPos + var parentActionProjection []*Expr + parentActionStep := int32(-1) + if delCtx.skipTargetDelete { + parentActionTag = builder.genNewBindTag() + projectionLen := len(fk.Cols) + if fkSelfReferCond && !isUpdate && fk.OnDelete == plan.ForeignKeyDef_CASCADE { + projectionLen++ + } + parentActionProjection = make([]*Expr, projectionLen) + } // use for join's projection & filter's condExpr var oneLeftCond *Expr var oneLeftCondName string @@ -515,12 +1098,29 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC childColumnName := col.Name originColumnName := idNameMap[fk.ForeignCols[i]] + leftColPos := nameIdxMap[originColumnName] + if delCtx.skipTargetDelete { + parentExpr := &Expr{ + Typ: *nameTypMap[originColumnName], + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: 0, + ColPos: leftColPos, + Name: originColumnName, + }}, + } + parentActionProjection[i], err = BindFuncExprImplByPlanExpr( + builder.GetContext(), "coalesce", []*Expr{parentExpr, DeepCopyExpr(parentExpr)}) + if err != nil { + return err + } + leftColPos = int32(i) + } leftExpr := &Expr{ Typ: *nameTypMap[originColumnName], Expr: &plan.Expr_Col{ Col: &plan.ColRef{ - RelPos: 0, - ColPos: nameIdxMap[originColumnName], + RelPos: parentActionTag, + ColPos: leftColPos, Name: originColumnName, }, }, @@ -543,7 +1143,7 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC Typ: *childTypMap[childColumnName], Expr: &plan.Expr_Col{ Col: &plan.ColRef{ - RelPos: 1, + RelPos: childRelPos, ColPos: childPosMap[childColumnName], Name: childColumnName, }, @@ -567,6 +1167,23 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC } } } + if len(parentActionProjection) > len(fk.Cols) { + rowIDPos, ok := nameIdxMap[catalog.Row_ID] + rowIDTyp, typOK := nameTypMap[catalog.Row_ID] + if !ok || !typOK { + return moerr.NewInternalErrorf( + builder.GetContext(), "self-referencing cascade root rowid is unavailable") + } + rowIDExpr := &Expr{ + Typ: *rowIDTyp, + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: 0, + ColPos: rowIDPos, + Name: catalog.Row_ID, + }}, + } + parentActionProjection[len(fk.Cols)] = rowIDExpr + } for idx, col := range childTableDef.Cols { if col.Name != catalog.Row_ID && col.Name != catalog.CPrimaryKeyColName { @@ -585,7 +1202,27 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC refAction = fk.OnDelete } - lastNodeId = appendSinkScanNode(builder, bindCtx, delCtx.sourceStep) + lastNodeId = appendDeleteSourceScan() + if delCtx.skipTargetDelete { + // This projection intentionally compacts the referenced parent + // keys to the leading positions; let the source scan prune and remap it. + delete(builder.preserveScanProjection, lastNodeId) + builder.qry.Nodes[lastNodeId].BindingTags = []int32{0} + lastNodeId = builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, + Children: []int32{lastNodeId}, + ProjectList: parentActionProjection, + BindingTags: []int32{parentActionTag}, + }, bindCtx) + parentSinkID := appendSinkNodeWithTag(builder, bindCtx, lastNodeId, parentActionTag) + if builder.preserveSinkProjection == nil { + builder.preserveSinkProjection = make(map[int32]struct{}) + } + builder.preserveSinkProjection[parentSinkID] = struct{}{} + parentActionStep = builder.appendStep(parentSinkID) + lastNodeId = appendSinkScanNodeWithTag(builder, bindCtx, parentActionStep, parentActionTag) + builder.positionalSinkScans[lastNodeId] = struct{}{} + } // deal with case: update t1 set a = a. then do not need to check constraint if isUpdate { var filterExpr, tmpExpr *Expr @@ -647,15 +1284,24 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC Stats: &plan.Stats{}, ObjRef: childObjRef, TableDef: copiedTableDef, - ProjectList: childProjectList, + ProjectList: childScanProject, + BindingTags: childBindingTags, }, bindCtx) + joinProjection := []*Expr{oneLeftCond} + joinType := plan.Node_SEMI + if delCtx.sourceTag != 0 { + joinType = plan.Node_INNER + } + if delCtx.skipTargetDelete { + joinProjection = []*Expr{oneLeftCond} + } lastNodeId = builder.appendNode(&plan.Node{ NodeType: plan.Node_JOIN, Children: []int32{lastNodeId, rightId}, - JoinType: plan.Node_SEMI, + JoinType: joinType, OnList: joinConds, - ProjectList: []*Expr{oneLeftCond}, + ProjectList: joinProjection, }, bindCtx) colExpr := &Expr{ @@ -666,8 +1312,16 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC }, }, } + if delCtx.skipTargetDelete { + colExpr = DeepCopyExpr(oneLeftCond) + } errExpr := makePlan2StringConstExprWithType("Cannot delete or update a parent row: a foreign key constraint fails") - isEmptyExpr, err := BindFuncExprImplByPlanExpr(builder.GetContext(), "isempty", []*Expr{colExpr}) + var isEmptyExpr *Expr + if delCtx.skipTargetDelete { + isEmptyExpr = makePlan2BoolConstExprWithType(false) + } else { + isEmptyExpr, err = BindFuncExprImplByPlanExpr(builder.GetContext(), "isempty", []*Expr{colExpr}) + } if err != nil { return err } @@ -675,11 +1329,22 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC if err != nil { return err } + if delCtx.skipTargetDelete { + lastNodeId = builder.appendNode(&Node{ + NodeType: plan.Node_PROJECT, + Children: []int32{lastNodeId}, + ProjectList: []*Expr{assertExpr}, + BindingTags: []int32{builder.genNewBindTag()}, + }, bindCtx) + builder.appendStep(lastNodeId) + break + } + filterProjection := getProjectionByLastNode(builder, lastNodeId) filterNode := &Node{ NodeType: plan.Node_FILTER, Children: []int32{lastNodeId}, FilterList: []*Expr{assertExpr}, - ProjectList: getProjectionByLastNode(builder, lastNodeId), + ProjectList: filterProjection, IsEnd: true, } lastNodeId = builder.appendNode(filterNode, bindCtx) @@ -693,7 +1358,8 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC Stats: &plan.Stats{}, ObjRef: childObjRef, TableDef: CloneTableDefForPlan(childTableDef, true), - ProjectList: childProjectList, + ProjectList: childScanProject, + BindingTags: childBindingTags, }, bindCtx) lastNodeId = builder.appendNode(&plan.Node{ NodeType: plan.Node_JOIN, @@ -704,6 +1370,9 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC }, bindCtx) // inner join cannot dealwith null expr in projectList. so we append a project node projectProjection := getProjectionByLastNode(builder, lastNodeId) + if delCtx.skipTargetDelete { + projectProjection = DeepCopyExprList(builder.qry.Nodes[lastNodeId].ProjectList) + } for _, e := range rightConds { projectProjection = append(projectProjection, &plan.Expr{ Typ: e.Typ, @@ -718,9 +1387,19 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC NodeType: plan.Node_PROJECT, Children: []int32{lastNodeId}, ProjectList: projectProjection, + BindingTags: []int32{builder.genNewBindTag()}, }, bindCtx) - lastNodeId = appendAggNodeForFkJoin(builder, bindCtx, lastNodeId) + if delCtx.skipTargetDelete { + projectTag := builder.qry.Nodes[lastNodeId].BindingTags[0] + lastNodeId = appendSinkNodeWithTag(builder, bindCtx, lastNodeId, projectTag) + if builder.preserveSinkProjection == nil { + builder.preserveSinkProjection = make(map[int32]struct{}) + } + builder.preserveSinkProjection[lastNodeId] = struct{}{} + } else { + lastNodeId = appendAggNodeForFkJoin(builder, bindCtx, lastNodeId) + } newSourceStep := builder.appendStep(lastNodeId) @@ -737,6 +1416,7 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC upPlanCtx.insertColPos = insertColPos upPlanCtx.isFkRecursionCall = true upPlanCtx.updatePkCol = updatePk + upPlanCtx.preserveUpdateSourceProjection = delCtx.skipTargetDelete err = buildUpdatePlans(ctx, builder, bindCtx, upPlanCtx, false) putDmlPlanCtx(upPlanCtx) @@ -750,11 +1430,14 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC Stats: &plan.Stats{}, ObjRef: childObjRef, TableDef: childTableDef, - ProjectList: childProjectList, + ProjectList: childScanProject, + BindingTags: childBindingTags, }, bindCtx) - //skip cascade for fk self refer - if !fkSelfReferCond { + // Legacy DELETE keeps the existing self-reference guard. Modern + // REPLACE owns the target-row delete in MULTI_UPDATE, so its old-row + // action path must explicitly collect self-referencing descendants. + if !fkSelfReferCond || delCtx.skipTargetDelete { builder.qry.HasForeignKeyAction = true if isUpdate { // update stmt get plan : sink_scan -> join[f1 inner join c1 on f1.id = c1.fid, get c1.* & update cols] -> sink then + updatePlans @@ -791,14 +1474,51 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC } } else { // delete stmt get plan : sink_scan -> join[f1 inner join c1 on f1.id = c1.fid, get c1.*] -> sink then + deletePlans + childActionTag := int32(0) + if delCtx.sourceTag != 0 || delCtx.skipTargetDelete { + // The join projects child-table columns unchanged, so retain + // their tag for sink remapping instead of inventing an alias + // that JOIN cannot map back to its inputs. + childActionTag = childScanTag + } lastNodeId = builder.appendNode(&plan.Node{ NodeType: plan.Node_JOIN, Children: []int32{lastNodeId, rightId}, JoinType: plan.Node_INNER, OnList: joinConds, ProjectList: childForJoinProject, + BindingTags: func() []int32 { + if childActionTag == 0 { + return nil + } + return []int32{childActionTag} + }(), }, bindCtx) - lastNodeId = appendSinkNode(builder, bindCtx, lastNodeId) + if fkSelfReferCond && delCtx.skipTargetDelete { + lastNodeId, err = appendExcludeReplaceOldRows( + builder, bindCtx, lastNodeId, + childScanTag, int32(childRowIdPos), + parentActionStep, parentActionTag, int32(len(fk.Cols))) + if err != nil { + return err + } + lastNodeId, err = appendSelfReferCascadeSource( + builder, bindCtx, lastNodeId, + childObjRef, childTableDef, fk, childPosMap, + parentActionStep, parentActionTag, int32(len(fk.Cols))) + if err != nil { + return err + } + childActionTag = 0 + } else if childActionTag != 0 { + lastNodeId = appendSinkNodeWithTag(builder, bindCtx, lastNodeId, childActionTag) + if builder.preserveSinkProjection == nil { + builder.preserveSinkProjection = make(map[int32]struct{}) + } + builder.preserveSinkProjection[lastNodeId] = struct{}{} + } else { + lastNodeId = appendSinkNode(builder, bindCtx, lastNodeId) + } newSourceStep := builder.appendStep(lastNodeId) //make deletePlans @@ -811,8 +1531,10 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC upPlanCtx.isMulti = false upPlanCtx.rowIdPos = childRowIdPos upPlanCtx.sourceStep = newSourceStep + upPlanCtx.sourceTag = childActionTag upPlanCtx.beginIdx = 0 upPlanCtx.allDelTableIDs = allDelTableIDs + upPlanCtx.isFkRecursionCall = true err := buildDeletePlans(ctx, builder, bindCtx, upPlanCtx) putDmlPlanCtx(upPlanCtx) @@ -831,6 +1553,80 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC return nil } +// appendExcludeReplaceOldRows keeps cascade DELETE ownership disjoint from the +// old rows deleted by the main REPLACE MULTI_UPDATE. +func appendExcludeReplaceOldRows( + builder *QueryBuilder, + bindCtx *BindContext, + inputNodeID int32, + candidateTag int32, + candidateRowIDPos int32, + oldRowsSourceStep int32, + oldRowsSourceTag int32, + oldRowIDPos int32, +) (int32, error) { + inputProject := builder.qry.Nodes[inputNodeID].ProjectList + if candidateRowIDPos < 0 || int(candidateRowIDPos) >= len(inputProject) || + oldRowsSourceStep < 0 || int(oldRowsSourceStep) >= len(builder.qry.Steps) { + return 0, moerr.NewInternalErrorf( + builder.GetContext(), "self-referencing cascade old-row source is incomplete") + } + inputNodeID = builder.appendNode(&Node{ + NodeType: plan.Node_PROJECT, + Children: []int32{inputNodeID}, + ProjectList: DeepCopyExprList(inputProject), + BindingTags: []int32{candidateTag}, + }, bindCtx) + oldRowsScanID := appendSinkScanNodeWithTag( + builder, bindCtx, oldRowsSourceStep, oldRowsSourceTag) + oldRowsTag := builder.genNewBindTag() + oldRowsScanID = builder.appendNode(&Node{ + NodeType: plan.Node_PROJECT, + Children: []int32{oldRowsScanID}, + ProjectList: DeepCopyExprList(builder.qry.Nodes[oldRowsScanID].ProjectList), + BindingTags: []int32{oldRowsTag}, + }, bindCtx) + if oldRowIDPos < 0 || int(oldRowIDPos) >= len(builder.qry.Nodes[oldRowsScanID].ProjectList) { + return 0, moerr.NewInternalErrorf( + builder.GetContext(), "self-referencing cascade old rowid is unavailable") + } + + candidateRowID := &Expr{ + Typ: builder.qry.Nodes[inputNodeID].ProjectList[candidateRowIDPos].Typ, + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: candidateTag, + ColPos: candidateRowIDPos, + Name: catalog.Row_ID, + }}, + } + oldRowID := &Expr{ + Typ: builder.qry.Nodes[oldRowsScanID].ProjectList[oldRowIDPos].Typ, + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: oldRowsTag, + ColPos: oldRowIDPos, + Name: catalog.Row_ID, + }}, + } + rowIDEqual, err := BindFuncExprImplByPlanExpr( + builder.GetContext(), "=", []*Expr{candidateRowID, oldRowID}) + if err != nil { + return 0, err + } + + markJoinID, notOwnedByReplace, err := builder.insertMarkJoin( + inputNodeID, oldRowsScanID, []*Expr{rowIDEqual}, nil, true, bindCtx) + if err != nil { + return 0, err + } + return builder.appendNode(&Node{ + NodeType: plan.Node_FILTER, + Children: []int32{markJoinID}, + FilterList: []*Expr{notOwnedByReplace}, + ProjectList: getProjectionByLastNodeWithTag(builder, inputNodeID, candidateTag), + BindingTags: []int32{candidateTag}, + }, bindCtx), nil +} + // appendAggNodeForFkJoin append agg node. to deal with these case: // create table f (a int, b int, primary key(a,b)); // insert into f values (1,1),(1,2),(1,3),(2,3); @@ -838,7 +1634,11 @@ func buildDeletePlans(ctx CompilerContext, builder *QueryBuilder, bindCtx *BindC // insert into c values (1,1),(2,1),(3,2); // update f set a = 10 where b=1; we need update c only once for 2 rows. not three times for 6 rows. func appendAggNodeForFkJoin(builder *QueryBuilder, bindCtx *BindContext, lastNodeId int32) int32 { + lastNode := builder.qry.Nodes[lastNodeId] groupByList := getProjectionByLastNode(builder, lastNodeId) + if len(lastNode.BindingTags) > 0 { + groupByList = getProjectionByLastNodeWithTag(builder, lastNodeId, lastNode.BindingTags[0]) + } aggProject := make([]*Expr, len(groupByList)) for i, e := range groupByList { aggProject[i] = &Expr{ @@ -856,6 +1656,7 @@ func appendAggNodeForFkJoin(builder *QueryBuilder, bindCtx *BindContext, lastNod GroupBy: groupByList, Children: []int32{lastNodeId}, ProjectList: aggProject, + BindingTags: []int32{builder.genNewBindTag(), builder.genNewBindTag()}, SpillMem: builder.aggSpillMem, }, bindCtx) lastNodeId = appendSinkNode(builder, bindCtx, lastNodeId) @@ -863,6 +1664,165 @@ func appendAggNodeForFkJoin(builder *QueryBuilder, bindCtx *BindContext, lastNod return lastNodeId } +// appendSelfReferCascadeSource expands the directly matched child rows into the +// complete descendant set for a self-referencing ON DELETE CASCADE. The +// Every recursion level excludes the complete main-REPLACE old-row set, which +// also prevents a valid FK cycle from revisiting a replaced row. The final +// aggregate removes duplicate row images produced by converging paths. +func appendSelfReferCascadeSource( + builder *QueryBuilder, + bindCtx *BindContext, + initialNodeID int32, + childObjRef *ObjectRef, + childTableDef *TableDef, + fk *ForeignKeyDef, + childPosMap map[string]int32, + oldRowsSourceStep int32, + oldRowsSourceTag int32, + oldRowIDPos int32, +) (int32, error) { + cteTag := builder.genNewBindTag() + initialNodeID = builder.appendNode(&Node{ + NodeType: plan.Node_PROJECT, + Children: []int32{initialNodeID}, + ProjectList: DeepCopyExprList(builder.qry.Nodes[initialNodeID].ProjectList), + BindingTags: []int32{cteTag}, + }, bindCtx) + initialSinkID := appendSinkNodeWithTag(builder, bindCtx, initialNodeID, cteTag) + if builder.preserveSinkProjection == nil { + builder.preserveSinkProjection = make(map[int32]struct{}) + } + builder.preserveSinkProjection[initialSinkID] = struct{}{} + initialSourceStep := builder.appendStep(initialSinkID) + + recursiveScanID := builder.appendNode(&Node{ + NodeType: plan.Node_RECURSIVE_SCAN, + SourceStep: []int32{initialSourceStep}, + ProjectList: getProjectionByLastNodeWithTag(builder, initialSinkID, cteTag), + BindingTags: []int32{cteTag}, + TableDef: CloneTableDefForPlan(childTableDef, true), + }, bindCtx) + + descendantTag := builder.genNewBindTag() + descendantJoinProject := make([]*Expr, len(childTableDef.Cols)) + for i, col := range childTableDef.Cols { + descendantJoinProject[i] = &Expr{ + Typ: col.Typ, + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: descendantTag, + ColPos: int32(i), + Name: col.Name, + }}, + } + } + descendantScanID := builder.appendNode(&Node{ + NodeType: plan.Node_TABLE_SCAN, + Stats: &plan.Stats{}, + ObjRef: childObjRef, + TableDef: CloneTableDefForPlan(childTableDef, true), + BindingTags: []int32{descendantTag}, + }, bindCtx) + + recursiveConds := make([]*Expr, len(fk.Cols)) + for i, childColID := range fk.Cols { + childName := "" + parentName := "" + for _, col := range childTableDef.Cols { + if col.ColId == childColID { + childName = col.Name + } + if col.ColId == fk.ForeignCols[i] { + parentName = col.Name + } + } + childPos, childOK := childPosMap[childName] + parentPos, parentOK := childPosMap[parentName] + if !childOK || !parentOK { + return 0, moerr.NewInternalErrorf( + builder.GetContext(), "self-referencing foreign key column mapping is incomplete") + } + leftExpr := &Expr{ + Typ: childTableDef.Cols[parentPos].Typ, + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: cteTag, + ColPos: parentPos, + Name: parentName, + }}, + } + rightExpr := &Expr{ + Typ: childTableDef.Cols[childPos].Typ, + Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: descendantTag, + ColPos: childPos, + Name: childName, + }}, + } + cond, err := BindFuncExprImplByPlanExpr(builder.GetContext(), "=", []*Expr{leftExpr, rightExpr}) + if err != nil { + return 0, err + } + recursiveConds[i] = cond + } + rowIDPos, rowIDOK := childPosMap[catalog.Row_ID] + if !rowIDOK { + return 0, moerr.NewInternalErrorf( + builder.GetContext(), "self-referencing cascade rowid is unavailable") + } + recursiveJoinID := builder.appendNode(&Node{ + NodeType: plan.Node_JOIN, + Children: []int32{recursiveScanID, descendantScanID}, + JoinType: plan.Node_INNER, + OnList: recursiveConds, + ProjectList: descendantJoinProject, + }, bindCtx) + var err error + recursiveJoinID, err = appendExcludeReplaceOldRows( + builder, bindCtx, recursiveJoinID, + descendantTag, rowIDPos, + oldRowsSourceStep, oldRowsSourceTag, oldRowIDPos) + if err != nil { + return 0, err + } + recursiveJoinID = builder.appendNode(&Node{ + NodeType: plan.Node_PROJECT, + Children: []int32{recursiveJoinID}, + ProjectList: DeepCopyExprList(descendantJoinProject), + BindingTags: []int32{cteTag}, + }, bindCtx) + recursiveSinkID := appendSinkNodeWithTag(builder, bindCtx, recursiveJoinID, cteTag) + builder.qry.Nodes[recursiveSinkID].RecursiveCte = true + builder.preserveSinkProjection[recursiveSinkID] = struct{}{} + recursiveSourceStep := builder.appendStep(recursiveSinkID) + + cteScanID := builder.appendNode(&Node{ + NodeType: plan.Node_RECURSIVE_CTE, + SourceStep: []int32{ + initialSourceStep, + recursiveSourceStep, + }, + ProjectList: getProjectionByLastNodeWithTag(builder, initialSinkID, cteTag), + BindingTags: []int32{cteTag}, + }, bindCtx) + cteSourceStep := int32(len(builder.qry.Steps)) + builder.qry.Nodes[recursiveScanID].SourceStep[0] = cteSourceStep + cteSinkID := appendSinkNodeWithTag(builder, bindCtx, cteScanID, cteTag) + builder.qry.Nodes[cteSinkID].RecursiveSink = true + builder.preserveSinkProjection[cteSinkID] = struct{}{} + cteSourceStep = builder.appendStep(cteSinkID) + + lastNodeID := appendSinkScanNodeWithTag(builder, bindCtx, cteSourceStep, cteTag) + if builder.preserveScanProjection == nil { + builder.preserveScanProjection = make(map[int32]struct{}) + } + builder.preserveScanProjection[lastNodeID] = struct{}{} + lastNodeID = builder.appendNode(&Node{ + NodeType: plan.Node_PROJECT, + Children: []int32{lastNodeID}, + ProjectList: DeepCopyExprList(builder.qry.Nodes[lastNodeID].ProjectList[:len(childTableDef.Cols)]), + }, bindCtx) + return appendAggNodeForFkJoin(builder, bindCtx, lastNodeID), nil +} + // buildInsertPlansWithRelatedHiddenTable build insert plan recursively for origin table func buildInsertPlansWithRelatedHiddenTable( stmt *tree.Insert, ctx CompilerContext, builder *QueryBuilder, bindCtx *BindContext, objRef *ObjectRef, @@ -1053,6 +2013,12 @@ func appendPureInsertBranch(ctx CompilerContext, builder *QueryBuilder, bindCtx ProjectList: insertProjection, } lastNodeId = builder.appendNode(insertNode, bindCtx) + if _, preserve := builder.preserveSinkProjection[builder.qry.Steps[sourceStep]]; preserve { + if builder.preserveInsertProjection == nil { + builder.preserveInsertProjection = make(map[int32]struct{}) + } + builder.preserveInsertProjection[lastNodeId] = struct{}{} + } builder.appendStep(lastNodeId) } @@ -1142,6 +2108,9 @@ func makeOneDeletePlan( TruncateTable: truncateTable, }, } + if delNodeInfo.preserveProjection { + deleteNode.ProjectList = getProjectionByLastNode(builder, lastNodeId) + } lastNodeId = builder.appendNode(deleteNode, bindCtx) return lastNodeId, nil @@ -3061,10 +4030,26 @@ func makePreUpdateDeletePlan( Children: []int32{lastNodeId}, LockTargets: []*plan.LockTarget{lockTarget}, } + if delCtx.preserveUpdateSourceProjection { + lockNode.ProjectList = getProjectionByLastNode(builder, lastNodeId) + } lastNodeId = builder.appendNode(lockNode, bindCtx) //lock new pk for update statement (if update pk) - if delCtx.updateColLength > 0 && delCtx.updatePkCol && delCtx.tableDef.Pkey != nil { + updatesPrimaryKey := false + if delCtx.tableDef.Pkey != nil { + if delCtx.tableDef.Pkey.PkeyColName == catalog.CPrimaryKeyColName { + for _, colName := range delCtx.tableDef.Pkey.Names { + if _, ok := delCtx.updateColPosMap[colName]; ok { + updatesPrimaryKey = true + break + } + } + } else { + _, updatesPrimaryKey = delCtx.updateColPosMap[delCtx.tableDef.Pkey.PkeyColName] + } + } + if delCtx.updateColLength > 0 && delCtx.updatePkCol && updatesPrimaryKey { newPkPos := int32(0) // for compound primary key, we need append hidden pk column to the project list diff --git a/pkg/sql/plan/build_test.go b/pkg/sql/plan/build_test.go index f791a3e2c0879..a7d31f5a8d2c0 100644 --- a/pkg/sql/plan/build_test.go +++ b/pkg/sql/plan/build_test.go @@ -18,9 +18,11 @@ import ( "bytes" "context" "encoding/json" + "fmt" "os" "strings" "testing" + "time" "github.com/golang/mock/gomock" "github.com/stretchr/testify/assert" @@ -30,14 +32,33 @@ import ( "github.com/matrixorigin/matrixone/pkg/common/moerr" moruntime "github.com/matrixorigin/matrixone/pkg/common/runtime" "github.com/matrixorigin/matrixone/pkg/container/types" + lockpb "github.com/matrixorigin/matrixone/pkg/pb/lock" "github.com/matrixorigin/matrixone/pkg/pb/plan" + txnpb "github.com/matrixorigin/matrixone/pkg/pb/txn" "github.com/matrixorigin/matrixone/pkg/sql/parsers/dialect" "github.com/matrixorigin/matrixone/pkg/sql/parsers/dialect/mysql" "github.com/matrixorigin/matrixone/pkg/sql/parsers/tree" "github.com/matrixorigin/matrixone/pkg/testutil" + "github.com/matrixorigin/matrixone/pkg/txn/client" "github.com/matrixorigin/matrixone/pkg/util/executor" + "github.com/matrixorigin/matrixone/pkg/vm/process" ) +type txnModeTestOperator struct { + client.TxnOperator + meta txnpb.TxnMeta +} + +func (o txnModeTestOperator) Txn() txnpb.TxnMeta { + return o.meta +} + +func setMockTxnMode(mock *MockOptimizer, mode txnpb.TxnMode) { + proc := testutil.NewProc(nil) + proc.Base.TxnOperator = txnModeTestOperator{meta: txnpb.TxnMeta{Mode: mode}} + mock.ctxt.GetProcessFunc = func() *process.Process { return proc } +} + type sqlModeMockCompilerContext struct { *MockCompilerContext sqlMode string @@ -3568,6 +3589,33 @@ func TestReplaceSelfRefCascade(t *testing.T) { assert.False(t, strings.HasPrefix(sql, "REPLACE_PARENT_CHK:"), "CASCADE self-ref FK should NOT generate parent-child pre-check, got: %s", sql) } + assert.True(t, queryDeletesTable(query, "self_ref_cascade"), + "CASCADE self-ref FK must build a descendant delete branch") + assert.True(t, queryHasNodeType(query, plan.Node_RECURSIVE_CTE), + "CASCADE self-ref FK must recursively collect the full descendant chain") + oldRowExclusions := 0 + for _, node := range query.Nodes { + if node.NodeType == plan.Node_JOIN && node.JoinType == plan.Node_ANTI { + oldRowExclusions++ + } + } + assert.GreaterOrEqual(t, oldRowExclusions, 2, + "initial and recursive cascade sources must exclude main REPLACE old rows") + cascadeLocks := 0 + for _, node := range query.Nodes { + if node.NodeType != plan.Node_LOCK_OP || len(node.Children) != 1 || + query.Nodes[node.Children[0]].NodeType != plan.Node_SINK_SCAN { + continue + } + for _, target := range node.LockTargets { + if target.TableId == mock.ctxt.tables["self_ref_cascade"].TblId && + target.Mode == lockpb.LockMode_Exclusive { + cascadeLocks++ + } + } + } + assert.GreaterOrEqual(t, cascadeLocks, 2, + "root and recursively cascaded rows must each lock a materialized source") } func TestReplaceDetectSqls(t *testing.T) { @@ -3583,6 +3631,7 @@ func TestReplaceDetectSqls(t *testing.T) { query := logicPlan.GetQuery() assert.NotNil(t, query) + assert.True(t, query.GetHasForeignKeyAction(), "FK-sensitive REPLACE must not be cached") var preCheck string for _, sql := range query.DetectSqls { @@ -3599,6 +3648,192 @@ func TestReplaceDetectSqls(t *testing.T) { assert.Contains(t, preCheck, "(1)", "pre-check SQL should embed the supplied PK value") } +func TestReplaceForeignKeyPlanRemainsSensitiveWhenChecksDisabled(t *testing.T) { + mock := NewMockOptimizer(true) + mock.ctxt.ResolveVariableFunc = func(name string, _, _ bool) (interface{}, error) { + switch name { + case "foreign_key_checks": + return int64(0), nil + case "sql_mode": + return "", nil + default: + return nil, moerr.NewInternalError(context.Background(), "unexpected variable") + } + } + logicPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_cp VALUES (1, 'new')") + require.NoError(t, err) + query := logicPlan.GetQuery() + require.True(t, query.GetHasForeignKeyAction()) + require.Empty(t, query.GetDetectSqls()) +} + +func TestReplaceParentSideFKAutoIncrementZero(t *testing.T) { + parseReplace := func(t *testing.T, literal string) *tree.Replace { + t.Helper() + stmt, err := mysql.ParseOne(context.Background(), + "REPLACE INTO replace_fk_cp VALUES ("+literal+", 'new')", 1) + require.NoError(t, err) + return stmt.(*tree.Replace) + } + + for _, tc := range []struct { + name string + sqlMode string + literal string + wantAction bool + predicate string + }{ + {name: "numeric zero allocates generated value", sqlMode: "", literal: "0", wantAction: false}, + {name: "string zero allocates generated value", sqlMode: "", literal: "'0'", wantAction: false}, + {name: "hex zero allocates generated value", sqlMode: "", literal: "0x0", wantAction: false}, + {name: "bit zero allocates generated value", sqlMode: "", literal: "b'0'", wantAction: false}, + {name: "hex nonzero is explicit key", sqlMode: "", literal: "0x1", wantAction: true, predicate: "`id` = cast(0x1 as INT)"}, + {name: "zero is explicit key", sqlMode: "NO_AUTO_VALUE_ON_ZERO", literal: "0", wantAction: true, predicate: "`id` = cast(0 as INT)"}, + } { + t.Run(tc.name, func(t *testing.T) { + mock := NewMockOptimizer(true) + parent := DeepCopyTableDef(mock.ctxt.tables["replace_fk_cp"], true) + parent.Cols[parent.Name2ColIndex["id"]].Typ.AutoIncr = true + mock.ctxt.ResolveVariableFunc = func(name string, _, _ bool) (interface{}, error) { + if name == "sql_mode" { + return tc.sqlMode, nil + } + if name == "foreign_key_checks" { + return int64(1), nil + } + return nil, moerr.NewInternalError(context.Background(), "unexpected variable") + } + + _, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, parseReplace(t, tc.literal)) + require.NoError(t, err) + if tc.wantAction { + require.Len(t, actions, 1) + assert.Contains(t, actions[0], tc.predicate) + } else { + assert.Empty(t, actions) + } + }) + } +} + +func TestChildInsertSkipsForeignKeyLockBarrierInOptimisticMode(t *testing.T) { + for _, tc := range []struct { + name string + sql string + }{ + {name: "insert", sql: "INSERT INTO replace_fk_c VALUES (10, 1), (11, 1)"}, + {name: "insert ignore", sql: "INSERT IGNORE INTO replace_fk_c VALUES (10, 1), (11, 1)"}, + {name: "on duplicate key update", sql: "INSERT INTO replace_fk_c VALUES (10, 1), (11, 1) ON DUPLICATE KEY UPDATE pid = VALUES(pid)"}, + {name: "replace", sql: "REPLACE INTO replace_fk_c VALUES (10, 1), (11, 1)"}, + } { + t.Run(tc.name, func(t *testing.T) { + mock := NewMockOptimizer(true) + setMockTxnMode(mock, txnpb.TxnMode_Optimistic) + + logicPlan, err := runOneStmt(mock, t, tc.sql) + require.NoError(t, err) + query := logicPlan.GetQuery() + for _, node := range query.Nodes { + for _, target := range node.LockTargets { + assert.NotEqual(t, lockpb.LockMode_Shared, target.Mode, + "optimistic child FK validation must not plan prerequisite shared locks") + } + } + assert.Len(t, query.Steps, 1, + "optimistic FK validation must remain in the streaming DML step") + }) + } + // The row count is deliberately much larger than the cases above. Plan shape + // must remain one streaming step; only the VALUE_SCAN payload may grow. + values := make([]string, 256) + for i := range values { + values[i] = fmt.Sprintf("(%d, 1)", i+100) + } + mock := NewMockOptimizer(true) + setMockTxnMode(mock, txnpb.TxnMode_Optimistic) + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES "+strings.Join(values, ",")) + require.NoError(t, err) + assert.Len(t, logicPlan.GetQuery().Steps, 1) + for _, node := range logicPlan.GetQuery().Nodes { + for _, target := range node.LockTargets { + assert.NotEqual(t, lockpb.LockMode_Shared, target.Mode) + } + } +} + +func TestChildInsertKeepsForeignKeyLockBarrierInPessimisticMode(t *testing.T) { + mock := NewMockOptimizer(true) + setMockTxnMode(mock, txnpb.TxnMode_Pessimistic) + + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES (10, 1), (11, 1)") + require.NoError(t, err) + query := logicPlan.GetQuery() + hasLock := false + hasSinkScan := false + for _, node := range query.Nodes { + hasLock = hasLock || node.NodeType == plan.Node_LOCK_OP + hasSinkScan = hasSinkScan || node.NodeType == plan.Node_SINK_SCAN + } + assert.True(t, hasLock) + assert.True(t, hasSinkScan) + assert.Greater(t, len(query.Steps), 1) +} + +func TestDeepCopyQueryKeepsReplaceDetectionSQLIndependent(t *testing.T) { + original := &plan.Query{DetectSqls: []string{ + "REPLACE_PARENT_LOCK:select 1 for update", + "REPLACE_PARENT_CHK:select true", + }} + copied := DeepCopyQuery(original) + require.Equal(t, original.DetectSqls, copied.DetectSqls) + copied.DetectSqls[0] = "changed" + assert.Equal(t, "REPLACE_PARENT_LOCK:select 1 for update", original.DetectSqls[0]) +} + +func TestReplaceParentSideFKMaterializesGeneratedUniqueKey(t *testing.T) { + mock := NewMockOptimizer(true) + parent := DeepCopyTableDef(mock.ctxt.tables["replace_fk_cp"], true) + if parent.Name2ColIndex == nil { + parent.Name2ColIndex = make(map[string]int32, len(parent.Cols)+1) + for i, col := range parent.Cols { + parent.Name2ColIndex[col.Name] = int32(i) + } + } + generatedPos := int32(len(parent.Cols)) + parent.Name2ColIndex["g"] = generatedPos + parent.Cols = append(parent.Cols, &plan.ColDef{ + Name: "g", + ColId: 999, + Typ: plan.Type{Id: int32(types.T_int32)}, + GeneratedCol: &plan.GeneratedCol{ + Expr: &plan.Expr{ + Typ: plan.Type{Id: int32(types.T_int32)}, + Expr: &plan.Expr_F{F: &plan.Function{ + Func: &plan.ObjectRef{ObjName: "+"}, + Args: []*plan.Expr{ + {Typ: plan.Type{Id: int32(types.T_int32)}, Expr: &plan.Expr_Col{Col: &plan.ColRef{ColPos: 0}}}, + makePlan2Int32ConstExprWithType(1), + }, + }}, + }, + OriginString: "`id` + 1", + }, + }) + parent.Indexes = append(parent.Indexes, &plan.IndexDef{Unique: true, Parts: []string{"g"}}) + stmt, err := mysql.ParseOne(context.Background(), "REPLACE INTO replace_fk_cp(id, v) VALUES (1, 'new')", 1) + require.NoError(t, err) + + lockSQL, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, stmt.(*tree.Replace)) + require.NoError(t, err) + require.Len(t, actions, 1) + assert.Contains(t, actions[0], "`__mo_replace_parent`.`g` = cast((select cast(`id` + 1 as INT)") + assert.Contains(t, actions[0], "cast(1 as INT) as `id`") + _, err = mysql.ParseOne(context.Background(), lockSQL, 1) + require.NoError(t, err, "generated parent lock SQL must be parseable") +} + func TestReplaceDetectSqlsExplicitColumnsCaseInsensitive(t *testing.T) { mock := NewMockOptimizer(true) @@ -3676,6 +3911,1267 @@ func TestReplaceDetectSqlsMultipleRows(t *testing.T) { assert.Contains(t, preCheck, "3", "pre-check IN list should contain row 3's PK") } +func assertReplaceParentPlanMarker(t *testing.T, query *plan.Query) { + t.Helper() + require.Contains(t, query.DetectSqls, "REPLACE_PARENT_PLAN:") +} + +func queryHasNodeType(query *plan.Query, typ plan.Node_NodeType) bool { + for _, node := range query.Nodes { + if node.NodeType == typ { + return true + } + } + return false +} + +func queryHasFKAssert(query *plan.Query) bool { + for _, node := range query.Nodes { + if node.NodeType == plan.Node_FILTER && node.IsEnd { + return true + } + for _, expr := range node.ProjectList { + if fn := expr.GetF(); fn != nil && fn.Func.ObjName == "assert" { + return true + } + } + } + return false +} + +func queryDeletesTable(query *plan.Query, table string) bool { + for _, node := range query.Nodes { + if node.NodeType == plan.Node_DELETE && node.DeleteCtx != nil && + node.DeleteCtx.TableDef != nil && node.DeleteCtx.TableDef.Name == table { + return true + } + } + return false +} + +func queryUpdatesTable(query *plan.Query, table string) bool { + for _, node := range query.Nodes { + for _, updateCtx := range node.UpdateCtxList { + if updateCtx.TableDef != nil && updateCtx.TableDef.Name == table { + return true + } + } + if node.NodeType == plan.Node_INSERT && node.InsertCtx != nil && + node.InsertCtx.TableDef != nil && node.InsertCtx.TableDef.Name == table { + return true + } + } + return false +} + +func assertLockTargetTypesMatchInput(t *testing.T, query *plan.Query) { + t.Helper() + for _, node := range query.Nodes { + if node.NodeType != plan.Node_LOCK_OP { + continue + } + require.Len(t, node.Children, 1) + input := query.Nodes[node.Children[0]] + for _, target := range node.LockTargets { + require.Less(t, int(target.PrimaryColIdxInBat), len(input.ProjectList)) + assert.Equal(t, target.PrimaryColTyp.Id, input.ProjectList[target.PrimaryColIdxInBat].Typ.Id) + } + } +} + +func TestReplaceParentSideFKRestrict(t *testing.T) { + mock := NewMockOptimizer(true) + + // REPLACE on a parent table whose PK is referenced by a child with + // ON DELETE RESTRICT must generate a REPLACE_PARENT_CHK: pre-check SQL + // against the child table (issue #24951, 3.2 RESTRICT case). + logicPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_p VALUES (1, 'p1_new')") + if err != nil { + t.Fatalf("%+v", err) + } + + query := logicPlan.GetQuery() + assert.NotNil(t, query) + assertReplaceParentPlanMarker(t, query) + assert.True(t, queryHasNodeType(query, plan.Node_LOCK_OP)) + assertLockTargetTypesMatchInput(t, query) + assert.True(t, queryHasFKAssert(query), "RESTRICT must assert that no child row references the locked old parent") +} + +func TestReplaceParentSideFKCascade(t *testing.T) { + mock := NewMockOptimizer(true) + + // REPLACE on a parent table whose PK is referenced by a child with + // ON DELETE CASCADE must generate a REPLACE_PARENT_ACTION: delete SQL + // against the child table (issue #24951, 3.2 CASCADE case). + logicPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_cp VALUES (1, 'p1_new')") + if err != nil { + t.Fatalf("%+v", err) + } + + query := logicPlan.GetQuery() + assert.NotNil(t, query) + assertReplaceParentPlanMarker(t, query) + assert.True(t, queryHasNodeType(query, plan.Node_LOCK_OP)) + assert.True(t, queryDeletesTable(query, "replace_fk_cc"), "CASCADE must build a child delete branch") +} + +func TestReplaceParentSideFKExplicitColumns(t *testing.T) { + mock := NewMockOptimizer(true) + + // Explicit column list (mixed case) must still resolve the PK position and + // generate the parent-side pre-check. + logicPlan, err := runOneStmt(mock, t, + "REPLACE INTO replace_fk_p (ID, V) VALUES (1, 'p1_new')") + if err != nil { + t.Fatalf("%+v", err) + } + + query := logicPlan.GetQuery() + assert.NotNil(t, query) + + assertReplaceParentPlanMarker(t, query) + assert.True(t, queryHasFKAssert(query)) +} + +func TestReplaceParentSideFKNoAction(t *testing.T) { + mock := NewMockOptimizer(true) + + // ON DELETE NO ACTION behaves like RESTRICT: it must generate a + // REPLACE_PARENT_CHK: pre-check, not a CASCADE/SET NULL action. + logicPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_np VALUES (1, 'p1_new')") + if err != nil { + t.Fatalf("%+v", err) + } + + query := logicPlan.GetQuery() + assert.NotNil(t, query) + + assertReplaceParentPlanMarker(t, query) + assert.True(t, queryHasFKAssert(query)) +} + +func TestReplaceParentSideFKSetDefault(t *testing.T) { + mock := NewMockOptimizer(true) + + logicPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_dp VALUES (1, 'p1_new')") + if err != nil { + t.Fatalf("%+v", err) + } + + query := logicPlan.GetQuery() + assert.NotNil(t, query) + + assertReplaceParentPlanMarker(t, query) + assert.True(t, queryHasFKAssert(query)) +} + +func TestReplaceParentSideFKMultiRow(t *testing.T) { + mock := NewMockOptimizer(true) + + // Multi-row REPLACE: every literal PK value must be embedded into the same + // parent-side action IN list (issue #24951 data-integrity case). + logicPlan, err := runOneStmt(mock, t, + "REPLACE INTO replace_fk_cp VALUES (1, 'a'), (2, 'b')") + if err != nil { + t.Fatalf("%+v", err) + } + + query := logicPlan.GetQuery() + assert.NotNil(t, query) + + assertReplaceParentPlanMarker(t, query) + assert.True(t, queryDeletesTable(query, "replace_fk_cc")) +} + +func TestReplaceParentSideFKMixedLiteralRows(t *testing.T) { + mock := NewMockOptimizer(true) + + // Mixed literal/function input is evaluated once by the main row-image plan. + logicPlan, err := runOneStmt(mock, t, + "REPLACE INTO replace_fk_cp VALUES (1, 'a'), (rand(), 'b')") + require.NoError(t, err) + assertReplaceParentPlanMarker(t, logicPlan.GetQuery()) +} + +func TestReplaceParentSideFKSetNull(t *testing.T) { + mock := NewMockOptimizer(true) + + // REPLACE on a parent table whose PK is referenced by a child with + // ON DELETE SET NULL must generate a REPLACE_PARENT_ACTION: update SQL + // that nulls the child FK column. + logicPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_sp VALUES (1, 'p1_new')") + if err != nil { + t.Fatalf("%+v", err) + } + + query := logicPlan.GetQuery() + assert.NotNil(t, query) + + assertReplaceParentPlanMarker(t, query) + assert.True(t, queryUpdatesTable(query, "replace_fk_sc"), "SET NULL must build a child update branch") +} + +func TestReplaceParentSideFKCombinesSetNullActions(t *testing.T) { + mock := NewMockOptimizer(true) + child := DeepCopyTableDef(mock.ctxt.tables["replace_fk_sc"], true) + mock.ctxt.tables["replace_fk_sc"] = child + if child.Name2ColIndex == nil { + child.Name2ColIndex = make(map[string]int32) + for i, col := range child.Cols { + child.Name2ColIndex[col.Name] = int32(i) + } + } + rowIDPos := len(child.Cols) - 1 + child.Cols = append(child.Cols, nil) + copy(child.Cols[rowIDPos+1:], child.Cols[rowIDPos:]) + child.Cols[rowIDPos] = &plan.ColDef{ + Name: "pid2", ColId: 10, Typ: plan.Type{Id: int32(types.T_int32), Width: 32}, + } + child.Name2ColIndex["pid2"] = int32(rowIDPos) + child.Name2ColIndex[catalog.Row_ID] = int32(rowIDPos + 1) + child.Fkeys = append(child.Fkeys, &plan.ForeignKeyDef{ + Name: "fk_replace_sc_2", Cols: []uint64{10}, ForeignTbl: 77005, ForeignCols: []uint64{0}, + OnDelete: plan.ForeignKeyDef_SET_NULL, OnUpdate: plan.ForeignKeyDef_SET_NULL, + }) + + logicPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_sp VALUES (1, 'p1_new')") + require.NoError(t, err) + query := logicPlan.GetQuery() + updates := 0 + foundPhysicalRowGrouping := false + for _, node := range query.Nodes { + if node.NodeType == plan.Node_AGG { + for _, groupExpr := range node.GroupBy { + if groupExpr.Typ.Id == int32(types.T_Rowid) { + foundPhysicalRowGrouping = true + } + } + } + if node.NodeType == plan.Node_INSERT && node.InsertCtx != nil && + node.InsertCtx.TableDef != nil && node.InsertCtx.TableDef.Name == "replace_fk_sc" { + updates++ + } + } + assert.True(t, foundPhysicalRowGrouping, + "combined SET NULL actions must group by Row_ID so physically distinct duplicate rows remain distinct") + require.Equal(t, 1, updates, + "all SET NULL columns for one child row must be emitted by one base-table update") +} + +func TestReplaceRecursiveCascadeLocksReferencedUniqueIndexKey(t *testing.T) { + mock := NewMockOptimizer(true) + cascadeChild := DeepCopyTableDef(mock.ctxt.tables["replace_fk_cc"], true) + mock.ctxt.tables["replace_fk_cc"] = cascadeChild + rootObj := mock.ctxt.objects["replace_fk_cp"] + + if cascadeChild.Name2ColIndex == nil { + cascadeChild.Name2ColIndex = make(map[string]int32, len(cascadeChild.Cols)+1) + } + rowIDPos := int32(-1) + for i, col := range cascadeChild.Cols { + cascadeChild.Name2ColIndex[col.Name] = int32(i) + if col.Name == catalog.Row_ID { + rowIDPos = int32(i) + } + } + require.GreaterOrEqual(t, rowIDPos, int32(0)) + cascadeChild.Cols = append(cascadeChild.Cols, nil) + copy(cascadeChild.Cols[rowIDPos+1:], cascadeChild.Cols[rowIDPos:]) + cascadeChild.Cols[rowIDPos] = &plan.ColDef{ + Name: "u", ColId: 10, Typ: plan.Type{Id: int32(types.T_varchar), Width: 20}, + } + for i, col := range cascadeChild.Cols { + cascadeChild.Name2ColIndex[col.Name] = int32(i) + } + const ( + indexTableID = uint64(77911) + grandchildID = uint64(77912) + ) + indexTableName := "__mo_index_replace_fk_cc_u" + cascadeChild.Indexes = append(cascadeChild.Indexes, &plan.IndexDef{ + IndexName: "uk_u", IndexTableName: indexTableName, Parts: []string{"u"}, + Unique: true, TableExist: true, IndexAlgo: catalog.MoIndexDefaultAlgo.ToString(), + }) + cascadeChild.RefChildTbls = []uint64{grandchildID} + + indexTable := &plan.TableDef{ + TblId: indexTableID, Name: indexTableName, + Cols: []*plan.ColDef{ + {Name: catalog.IndexTableIndexColName, ColId: 0, + Typ: plan.Type{Id: int32(types.T_varchar), Width: 20}}, + {Name: catalog.Row_ID, ColId: 1, Hidden: true, + Typ: plan.Type{Id: int32(types.T_Rowid), Width: 16}}, + }, + Pkey: &plan.PrimaryKeyDef{Names: []string{catalog.IndexTableIndexColName}, + PkeyColName: catalog.IndexTableIndexColName}, + Name2ColIndex: map[string]int32{catalog.IndexTableIndexColName: 0, catalog.Row_ID: 1}, + } + grandchild := &plan.TableDef{ + TblId: grandchildID, Name: "replace_fk_gc", + Cols: []*plan.ColDef{ + {Name: "id", ColId: 0, Typ: plan.Type{Id: int32(types.T_int32), Width: 32}}, + {Name: "cu", ColId: 1, Typ: plan.Type{Id: int32(types.T_varchar), Width: 20}}, + {Name: catalog.Row_ID, ColId: 2, Hidden: true, Typ: plan.Type{Id: int32(types.T_Rowid), Width: 16}}, + }, + Pkey: &plan.PrimaryKeyDef{Names: []string{"id"}, PkeyColName: "id"}, + Fkeys: []*plan.ForeignKeyDef{{ + Name: "fk_replace_gc", Cols: []uint64{1}, ForeignTbl: cascadeChild.TblId, + ForeignCols: []uint64{10}, OnDelete: plan.ForeignKeyDef_RESTRICT, + OnUpdate: plan.ForeignKeyDef_RESTRICT, + }}, + Name2ColIndex: map[string]int32{"id": 0, "cu": 1, catalog.Row_ID: 2}, + } + registerTable := func(tableDef *plan.TableDef) { + mock.ctxt.tables[tableDef.Name] = tableDef + mock.ctxt.objects[tableDef.Name] = &plan.ObjectRef{ + Obj: int64(tableDef.TblId), SchemaName: rootObj.SchemaName, ObjName: tableDef.Name, + } + mock.ctxt.id2name[tableDef.TblId] = tableDef.Name + } + registerTable(indexTable) + registerTable(grandchild) + + builder := NewQueryBuilder(plan.Query_DELETE, mock.CurrentContext(), false, true) + bindCtx := NewBindContext(builder, nil) + sourceTag := builder.genNewBindTag() + sourceProject := make([]*plan.Expr, len(cascadeChild.Cols)) + for i, col := range cascadeChild.Cols { + sourceProject[i] = &plan.Expr{Typ: col.Typ, Expr: &plan.Expr_Col{Col: &plan.ColRef{ + RelPos: sourceTag, ColPos: int32(i), Name: col.Name, + }}} + } + sourceNodeID := builder.appendNode(&plan.Node{ + NodeType: plan.Node_TABLE_SCAN, ObjRef: mock.ctxt.objects[cascadeChild.Name], + TableDef: cascadeChild, ProjectList: sourceProject, BindingTags: []int32{sourceTag}, + }, bindCtx) + delCtx := &dmlPlanCtx{ + objRef: mock.ctxt.objects[cascadeChild.Name], tableDef: cascadeChild, sourceTag: sourceTag, + } + outputNodeID, err := appendRecursiveCascadeLockNode(builder, bindCtx, delCtx, sourceNodeID) + require.NoError(t, err) + builder.appendStep(outputNodeID) + query, err := builder.createQuery() + require.NoError(t, err) + foundBaseLock := false + foundUniqueLock := false + for _, node := range query.Nodes { + if node.NodeType != plan.Node_LOCK_OP { + continue + } + for _, target := range node.LockTargets { + if target.Mode != lockpb.LockMode_Exclusive { + continue + } + if target.TableId == cascadeChild.TblId { + foundBaseLock = true + } + if target.TableId == indexTableID { + foundUniqueLock = true + require.Len(t, node.Children, 1) + lockInput := query.Nodes[node.Children[0]] + require.Less(t, int(target.PrimaryColIdxInBat), len(lockInput.ProjectList)) + assert.Equal(t, target.PrimaryColTyp.Id, + lockInput.ProjectList[target.PrimaryColIdxInBat].Typ.Id) + } + } + } + assert.True(t, foundBaseLock, "recursive cascade must lock the current table primary key") + assert.True(t, foundUniqueLock, + "recursive cascade must lock the hidden UNIQUE namespace referenced by the grandchild") +} + +func TestReplaceParentSideFKNonLiteralSkip(t *testing.T) { + mock := NewMockOptimizer(true) + + // Non-literal expressions are evaluated by the main REPLACE row image. + logicPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_p VALUES (rand(), 'x')") + require.NoError(t, err) + assertReplaceParentPlanMarker(t, logicPlan.GetQuery()) +} + +func TestReplaceParentSideFKUnsupportedSources(t *testing.T) { + mock := NewMockOptimizer(true) + preparedSQL := "REPLACE INTO replace_fk_p VALUES (?, 'x')" + stmts, err := mysql.Parse(mock.CurrentContext().GetContext(), preparedSQL, 1) + require.NoError(t, err) + logicPlan, err := BuildPlan(mock.CurrentContext(), stmts[0], true) + require.NoError(t, err) + assertReplaceParentPlanMarker(t, logicPlan.GetQuery()) + + selectSQL := "REPLACE INTO replace_fk_p SELECT deptno, dname FROM dept" + logicPlan, err = runOneStmt(mock, t, selectSQL) + require.NoError(t, err) + assertReplaceParentPlanMarker(t, logicPlan.GetQuery()) +} + +func TestReplaceParentSideFKUniquePrefixConflict(t *testing.T) { + mock := NewMockOptimizer(true) + parent := DeepCopyTableDef(mock.ctxt.tables["replace_fk_cp"], true) + parent.Indexes = append(parent.Indexes, &plan.IndexDef{ + Unique: true, + Parts: []string{"v"}, + IndexAlgoParams: `{"prefix_lengths":"v:4"}`, + }) + stmt, err := mysql.ParseOne(context.Background(), + "REPLACE INTO replace_fk_cp VALUES (2, 'abcdyyyy')", 1) + require.NoError(t, err) + + _, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, stmt.(*tree.Replace)) + require.NoError(t, err) + require.Len(t, actions, 1) + assert.Contains(t, actions[0], "substring(`__mo_replace_parent`.`v`, 1, 4)") + assert.Contains(t, actions[0], `substring(cast("abcdyyyy" as VARCHAR(20)), 1, 4)`) +} + +func TestReplaceParentSideFKAssignmentCastAndLock(t *testing.T) { + mock := NewMockOptimizer(true) + parent := DeepCopyTableDef(mock.ctxt.tables["replace_fk_cp"], true) + for _, col := range parent.Cols { + if col.Name == "v" { + col.Typ = plan.Type{Id: int32(types.T_decimal64), Width: 5, Scale: 2} + } + } + parent.Indexes = append(parent.Indexes, &plan.IndexDef{Unique: true, Parts: []string{"v"}}) + stmt, err := mysql.ParseOne(context.Background(), + "REPLACE INTO replace_fk_cp VALUES (2, 1.234)", 1) + require.NoError(t, err) + + lockSQL, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, stmt.(*tree.Replace)) + require.NoError(t, err) + require.Len(t, actions, 1) + assert.Contains(t, lockSQL, "`__mo_replace_parent`.`v` = cast(1.234 as DECIMAL(5,2))") + assert.Contains(t, lockSQL, "for update") + assert.Contains(t, actions[0], "`__mo_replace_parent`.`v` = cast(1.234 as DECIMAL(5,2))") +} + +func TestReplaceParentSideFKLocksReferencedUniqueIndexKey(t *testing.T) { + mock := NewMockOptimizer(true) + parent := mock.ctxt.tables["replace_fk_p"] + child := mock.ctxt.tables["replace_fk_c"] + child.Cols[1].Typ = plan.Type{Id: int32(types.T_varchar), Width: 20} + child.Fkeys[0].ForeignCols = []uint64{1} + parent.Indexes = append(parent.Indexes, &plan.IndexDef{ + IndexName: "uk_v", IndexTableName: "__mo_index_replace_fk_p_v", + Parts: []string{"v"}, Unique: true, TableExist: true, + }) + indexTable := &plan.TableDef{ + TblId: 77900, Name: "__mo_index_replace_fk_p_v", + Cols: []*plan.ColDef{{Name: catalog.IndexTableIndexColName, ColId: 0, Typ: parent.Cols[1].Typ}}, + Pkey: &plan.PrimaryKeyDef{Names: []string{catalog.IndexTableIndexColName}, + PkeyColName: catalog.IndexTableIndexColName}, + Name2ColIndex: map[string]int32{catalog.IndexTableIndexColName: 0}, + } + mock.ctxt.tables[indexTable.Name] = indexTable + mock.ctxt.objects[indexTable.Name] = &plan.ObjectRef{ + Obj: int64(indexTable.TblId), SchemaName: mock.ctxt.objects["replace_fk_p"].SchemaName, + ObjName: indexTable.Name, + } + modernPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_p VALUES (1, 'new')") + require.NoError(t, err) + assertLockTargetTypesMatchInput(t, modernPlan.GetQuery()) + + omittedPlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_p(id) VALUES (1)") + require.NoError(t, err) + oldIndexLockCount := 0 + oldIndexUpdateFound := false + for _, node := range omittedPlan.GetQuery().Nodes { + for _, target := range node.LockTargets { + if target.TableId == indexTable.TblId { + oldIndexLockCount++ + } + } + if node.NodeType == plan.Node_MULTI_UPDATE { + for _, updateCtx := range node.UpdateCtxList { + if updateCtx.TableDef.TblId == indexTable.TblId { + oldIndexUpdateFound = true + } + } + } + } + assert.Equal(t, 1, oldIndexLockCount, + "a NULL replacement key must skip only the new-key lock and retain the old-key lock") + assert.True(t, oldIndexUpdateFound, + "a NULL replacement key must retain old hidden-index deletion") + stmt, err := mysql.ParseOne(context.Background(), + "REPLACE INTO replace_fk_p VALUES (1, 'new')", 1) + require.NoError(t, err) + + lockSQL, checks, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_p"], parent, + stmt.(*tree.Replace)) + require.NoError(t, err) + require.Len(t, checks, 1) + assert.Empty(t, actions) + assert.Contains(t, lockSQL, fmt.Sprintf("from `%s`.`__mo_index_replace_fk_p_v`", + mock.ctxt.objects["replace_fk_p"].SchemaName)) + assert.Contains(t, lockSQL, + "where `__mo_replace_fk_idx_0`.`__mo_index_idx_col` = `__mo_replace_parent`.`v` for update") + assert.Contains(t, lockSQL, + "select `__mo_replace_parent`.`id`, (select `__mo_replace_fk_idx_0`.`__mo_index_idx_col`") + assert.Contains(t, lockSQL, "for update") + + lockPlan, err := runOneStmt(mock, t, lockSQL) + require.NoError(t, err) + for _, node := range lockPlan.GetQuery().Nodes { + if node.NodeType != plan.Node_LOCK_OP { + continue + } + require.Len(t, node.Children, 1) + lockInput := lockPlan.GetQuery().Nodes[node.Children[0]] + for _, target := range node.LockTargets { + require.Less(t, int(target.PrimaryColIdxInBat), len(lockInput.ProjectList)) + assert.Equal(t, target.PrimaryColTyp.Id, lockInput.ProjectList[target.PrimaryColIdxInBat].Typ.Id) + } + } +} + +func TestReplaceParentSideFKLocksCompositePrefixUniqueIndexOnce(t *testing.T) { + mock := NewMockOptimizer(true) + parent := mock.ctxt.tables["replace_fk_cp"] + child := mock.ctxt.tables["replace_fk_cc"] + parent.Cols[1].Typ = plan.Type{Id: int32(types.T_text), Width: types.MaxStringSize} + child.Cols[1].Typ = parent.Cols[1].Typ + child.Fkeys[0].Cols = []uint64{0, 1} + child.Fkeys[0].ForeignCols = []uint64{0, 1} + child.Fkeys = append(child.Fkeys, child.Fkeys[0]) + parent.Indexes = append(parent.Indexes, &plan.IndexDef{ + IndexName: "uk_id_v", IndexTableName: "__mo_index_replace_fk_cp_id_v", + Parts: []string{"id", "v"}, Unique: true, TableExist: true, + IndexAlgoParams: `{"prefix_lengths":"v:4"}`, + }) + stmt, err := mysql.ParseOne(context.Background(), + "REPLACE INTO replace_fk_cp VALUES (1, 'abcdefgh')", 1) + require.NoError(t, err) + + lockSQL, checks, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, + stmt.(*tree.Replace)) + require.NoError(t, err) + assert.Empty(t, checks) + require.Len(t, actions, 2) + assert.Equal(t, 1, strings.Count(lockSQL, + fmt.Sprintf("from `%s`.`__mo_index_replace_fk_cp_id_v`", + mock.ctxt.objects["replace_fk_cp"].SchemaName))) + assert.Contains(t, lockSQL, + "serial(`__mo_replace_parent`.`id`, cast(substring(`__mo_replace_parent`.`v`, 1, 4) as VARCHAR(65535)))") + assert.Contains(t, lockSQL, "for update") +} + +func TestReplaceParentSideFKGeneratedColumnValueMapping(t *testing.T) { + mock := NewMockOptimizer(true) + parent := DeepCopyTableDef(mock.ctxt.tables["replace_fk_cp"], true) + generated := &plan.ColDef{ + Name: "g", ColId: 99, Typ: plan.Type{Id: int32(types.T_int32), Width: 32}, + GeneratedCol: &plan.GeneratedCol{Expr: makePlan2Int64ConstExprWithType(1)}, + } + parent.Cols = append(parent.Cols[:1], append([]*plan.ColDef{generated}, parent.Cols[1:]...)...) + parent.Name2ColIndex = make(map[string]int32, len(parent.Cols)) + for i, col := range parent.Cols { + parent.Name2ColIndex[col.Name] = int32(i) + } + parent.Indexes = append(parent.Indexes, &plan.IndexDef{Unique: true, Parts: []string{"v"}}) + + parse := func(sql string) *tree.Replace { + stmt, err := mysql.ParseOne(context.Background(), sql, 1) + require.NoError(t, err) + return stmt.(*tree.Replace) + } + assertMappedValue := func(stmt *tree.Replace, value string) { + _, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, stmt) + require.NoError(t, err) + require.Len(t, actions, 1) + assert.Contains(t, actions[0], fmt.Sprintf("`__mo_replace_parent`.`v` = cast(\"%s\" as VARCHAR(20))", value)) + } + + assertMappedValue(parse("REPLACE INTO replace_fk_cp VALUES (2, 'implicit')"), "implicit") + explicit := parse("REPLACE INTO replace_fk_cp (id, g, v) VALUES (2, DEFAULT, 'explicit')") + values := explicit.Rows.Select.(*tree.ValuesClause) + values.Rows[0] = tree.Exprs{values.Rows[0][0], values.Rows[0][2]} + assertMappedValue(explicit, "explicit") +} + +func TestReplaceParentSideFKRejectsOverWidthConflictLiteral(t *testing.T) { + mock := NewMockOptimizer(true) + parent := DeepCopyTableDef(mock.ctxt.tables["replace_fk_cp"], true) + for _, col := range parent.Cols { + if col.Name == "v" { + col.Typ = plan.Type{Id: int32(types.T_varchar), Width: 3} + } + } + parent.Indexes = append(parent.Indexes, &plan.IndexDef{Unique: true, Parts: []string{"v"}}) + stmt, err := mysql.ParseOne(context.Background(), + "REPLACE INTO replace_fk_cp VALUES (2, 'abcd')", 1) + require.NoError(t, err) + + lockSQL, checks, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, stmt.(*tree.Replace)) + require.ErrorContains(t, err, "larger than Dest length") + assert.Empty(t, lockSQL) + assert.Empty(t, checks) + assert.Empty(t, actions) +} + +func TestChildInsertLocksForeignKeyParentShared(t *testing.T) { + mock := NewMockOptimizer(true) + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES (10, 1)") + require.NoError(t, err) + + parentID := mock.ctxt.tables["replace_fk_p"].TblId + query := logicPlan.GetQuery() + found := false + lockNodeID := int32(-1) + parentScanIDs := make([]int32, 0, 1) + for nodeID, node := range query.Nodes { + if node.NodeType == plan.Node_TABLE_SCAN && node.TableDef != nil && node.TableDef.TblId == parentID { + assert.Empty(t, node.LockTargets, "the raw parent scan must not carry a shared lock") + parentScanIDs = append(parentScanIDs, int32(nodeID)) + } + for _, target := range node.LockTargets { + if target.TableId == parentID && target.Mode == lockpb.LockMode_Shared { + found = true + lockNodeID = int32(nodeID) + assert.Equal(t, int32(0), target.PrimaryColRelPos) + require.Len(t, node.Children, 1) + lockInput := query.Nodes[node.Children[0]] + require.Less(t, int(target.PrimaryColIdxInBat), len(lockInput.ProjectList)) + assert.Equal(t, target.PrimaryColTyp.Id, lockInput.ProjectList[target.PrimaryColIdxInBat].Typ.Id) + } + } + } + assert.True(t, found, "child FK validation must hold a shared lock on its parent row") + require.NotEmpty(t, parentScanIDs) + stepContaining := func(target int32) int { + var contains func(int32) bool + contains = func(nodeID int32) bool { + if nodeID == target { + return true + } + for _, childID := range query.Nodes[nodeID].Children { + if contains(childID) { + return true + } + } + return false + } + for step, rootID := range query.Steps { + if contains(rootID) { + return step + } + } + return -1 + } + var stepDependsOn func(int, int, map[int]bool) bool + stepDependsOn = func(step, dependency int, visited map[int]bool) bool { + if step == dependency { + return true + } + if visited[step] { + return false + } + visited[step] = true + var nodeDependsOn func(int32) bool + nodeDependsOn = func(nodeID int32) bool { + node := query.Nodes[nodeID] + for _, sourceStep := range node.SourceStep { + if stepDependsOn(int(sourceStep), dependency, visited) { + return true + } + } + for _, childID := range node.Children { + if nodeDependsOn(childID) { + return true + } + } + return false + } + return nodeDependsOn(query.Steps[step]) + } + lockStep := stepContaining(lockNodeID) + require.GreaterOrEqual(t, lockStep, 0) + assert.Equal(t, plan.Node_SINK, query.Nodes[query.Steps[lockStep]].NodeType, + "a dependent SINK_SCAN must consume a materialized lock stage") + for _, scanID := range parentScanIDs { + parentStep := stepContaining(scanID) + require.GreaterOrEqual(t, parentStep, 0) + assert.True(t, stepDependsOn(parentStep, lockStep, make(map[int]bool)), + "the parent scan must consume the referenced-key lock step output") + } +} + +func TestChildInsertLockKeyUsesParentDecimalType(t *testing.T) { + mock := NewMockOptimizer(true) + parent := mock.ctxt.tables["replace_fk_p"] + child := mock.ctxt.tables["replace_fk_c"] + parent.Cols[0].Typ = plan.Type{Id: int32(types.T_decimal64), Width: 5, Scale: 2} + child.Cols[1].Typ = plan.Type{Id: int32(types.T_decimal64), Width: 5, Scale: 3} + + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES (10, 1.230)") + require.NoError(t, err) + for _, node := range logicPlan.GetQuery().Nodes { + for _, target := range node.LockTargets { + if target.TableId != parent.TblId || target.Mode != lockpb.LockMode_Shared { + continue + } + lockInput := logicPlan.GetQuery().Nodes[node.Children[0]] + require.Less(t, int(target.PrimaryColIdxInBat), len(lockInput.ProjectList)) + lockKey := lockInput.ProjectList[target.PrimaryColIdxInBat] + assert.Equal(t, int32(types.T_decimal64), lockKey.Typ.Id) + assert.Equal(t, int32(2), lockKey.Typ.Scale) + assert.Equal(t, target.PrimaryColTyp, lockKey.Typ) + require.NotNil(t, lockKey.GetF()) + assert.Equal(t, "cast", lockKey.GetF().Func.ObjName) + return + } + } + t.Fatal("decimal parent shared lock not found") +} + +func TestChildInsertChainsMultipleForeignKeyLocks(t *testing.T) { + mock := NewMockOptimizer(true) + child := mock.ctxt.tables["replace_fk_c"] + fkCopy := *child.Fkeys[0] + child.Fkeys = append(child.Fkeys, &fkCopy) + + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES (10, 1)") + require.NoError(t, err) + query := logicPlan.GetQuery() + lockIDs := make([]int32, 0, 2) + parentID := mock.ctxt.tables["replace_fk_p"].TblId + for nodeID, node := range query.Nodes { + for _, target := range node.LockTargets { + if target.TableId == parentID && target.Mode == lockpb.LockMode_Shared { + lockIDs = append(lockIDs, int32(nodeID)) + } + } + } + require.Len(t, lockIDs, 2) + + contains := func(root, target int32) bool { + var visit func(int32) bool + visit = func(nodeID int32) bool { + if nodeID == target { + return true + } + for _, childID := range query.Nodes[nodeID].Children { + if visit(childID) { + return true + } + } + return false + } + return visit(root) + } + lockStepRoot := int32(-1) + for _, stepRoot := range query.Steps { + if contains(stepRoot, lockIDs[0]) && contains(stepRoot, lockIDs[1]) { + lockStepRoot = stepRoot + break + } + } + require.NotEqual(t, int32(-1), lockStepRoot) + assert.Equal(t, plan.Node_SINK, query.Nodes[lockStepRoot].NodeType) + assert.True(t, contains(lockIDs[0], lockIDs[1]) || contains(lockIDs[1], lockIDs[0]), + "foreign-key lock stages must form one serial data pipeline") +} + +func TestChildInsertLocksCompositeParentPrimaryKey(t *testing.T) { + mock := NewMockOptimizer(true) + parent := mock.ctxt.tables["replace_fk_p"] + child := mock.ctxt.tables["replace_fk_c"] + parent.Cols = append(parent.Cols, + &plan.ColDef{Name: "k", ColId: 3, Typ: plan.Type{Id: int32(types.T_int32), Width: 32}}, + &plan.ColDef{Name: catalog.CPrimaryKeyColName, ColId: 4, Hidden: true, + Typ: plan.Type{Id: int32(types.T_varchar), Width: 65535}}, + ) + parent.Pkey = &plan.PrimaryKeyDef{Names: []string{"id", "k"}, PkeyColName: catalog.CPrimaryKeyColName} + if parent.Name2ColIndex == nil { + parent.Name2ColIndex = make(map[string]int32, len(parent.Cols)) + for i, col := range parent.Cols { + parent.Name2ColIndex[col.Name] = int32(i) + } + } + parent.Name2ColIndex["k"] = int32(len(parent.Cols) - 2) + parent.Name2ColIndex[catalog.CPrimaryKeyColName] = int32(len(parent.Cols) - 1) + child.Fkeys[0].Cols = []uint64{0, 1} + child.Fkeys[0].ForeignCols = []uint64{0, 3} + + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES (10, 1)") + require.NoError(t, err) + for _, node := range logicPlan.GetQuery().Nodes { + for _, target := range node.LockTargets { + if target.TableId != parent.TblId || target.Mode != lockpb.LockMode_Shared { + continue + } + lockInput := logicPlan.GetQuery().Nodes[node.Children[0]] + require.Less(t, int(target.PrimaryColIdxInBat), len(lockInput.ProjectList), + "lock input=%+v target=%+v", lockInput, target) + assert.Equal(t, target.PrimaryColTyp.Id, lockInput.ProjectList[target.PrimaryColIdxInBat].Typ.Id) + assert.Equal(t, int32(types.T_varchar), target.PrimaryColTyp.Id) + return + } + } + t.Fatal("composite parent primary key shared lock not found") +} + +func TestChildInsertLocksCompositeParentPrimaryKeyPrefixTable(t *testing.T) { + mock := NewMockOptimizer(true) + parent := mock.ctxt.tables["replace_fk_p"] + parent.Cols = append(parent.Cols, + &plan.ColDef{Name: "k", ColId: 3, Typ: plan.Type{Id: int32(types.T_int32), Width: 32}}, + &plan.ColDef{Name: catalog.CPrimaryKeyColName, ColId: 4, Hidden: true, + Typ: plan.Type{Id: int32(types.T_varchar), Width: 65535}}, + ) + parent.Pkey = &plan.PrimaryKeyDef{Names: []string{"id", "k"}, PkeyColName: catalog.CPrimaryKeyColName} + if parent.Name2ColIndex == nil { + parent.Name2ColIndex = make(map[string]int32, len(parent.Cols)) + for i, col := range parent.Cols { + parent.Name2ColIndex[col.Name] = int32(i) + } + } + parent.Name2ColIndex["k"] = int32(len(parent.Cols) - 2) + parent.Name2ColIndex[catalog.CPrimaryKeyColName] = int32(len(parent.Cols) - 1) + + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES (10, 1)") + require.NoError(t, err) + query := logicPlan.GetQuery() + stepContaining := func(target int32) int { + var contains func(int32) bool + contains = func(nodeID int32) bool { + if nodeID == target { + return true + } + for _, childID := range query.Nodes[nodeID].Children { + if contains(childID) { + return true + } + } + return false + } + for step, rootID := range query.Steps { + if contains(rootID) { + return step + } + } + return -1 + } + foundParentScan := false + for _, node := range query.Nodes { + if node.NodeType == plan.Node_TABLE_SCAN && node.TableDef != nil && node.TableDef.TblId == parent.TblId { + foundParentScan = true + } + } + for nodeID, node := range query.Nodes { + for _, target := range node.LockTargets { + if target.TableId != parent.TblId || target.Mode != lockpb.LockMode_Shared { + continue + } + assert.True(t, target.LockTable) + lockStep := stepContaining(int32(nodeID)) + require.GreaterOrEqual(t, lockStep, 0) + assert.Less(t, lockStep, len(query.Steps)-1) + assert.True(t, foundParentScan) + return + } + } + t.Fatal("composite parent primary-key prefix shared table lock not found") +} + +func TestChildInsertLocksReferencedUniqueIndexKey(t *testing.T) { + mock := NewMockOptimizer(true) + parent := mock.ctxt.tables["replace_fk_p"] + child := mock.ctxt.tables["replace_fk_c"] + child.Cols[1].Typ = plan.Type{Id: int32(types.T_varchar), Width: 20} + child.Fkeys[0].ForeignCols = []uint64{1} + indexName := "__mo_index_fk_parent_v" + indexID := uint64(77901) + parent.Indexes = append(parent.Indexes, &plan.IndexDef{ + IndexName: "uk_v", IndexTableName: indexName, Parts: []string{"v"}, + Unique: true, TableExist: true, IndexAlgo: catalog.MoIndexDefaultAlgo.ToString(), + }) + indexTable := &plan.TableDef{ + TblId: indexID, Name: indexName, + Cols: []*plan.ColDef{ + {Name: catalog.IndexTableIndexColName, ColId: 0, Typ: plan.Type{Id: int32(types.T_varchar), Width: 20}}, + {Name: catalog.Row_ID, ColId: 1, Hidden: true, Typ: plan.Type{Id: int32(types.T_Rowid)}}, + }, + Pkey: &plan.PrimaryKeyDef{Names: []string{catalog.IndexTableIndexColName}, + PkeyColName: catalog.IndexTableIndexColName}, + Name2ColIndex: map[string]int32{catalog.IndexTableIndexColName: 0, catalog.Row_ID: 1}, + } + mock.ctxt.tables[indexName] = indexTable + mock.ctxt.objects[indexName] = &plan.ObjectRef{ + Obj: int64(indexID), SchemaName: mock.ctxt.objects["replace_fk_p"].SchemaName, ObjName: indexName, + } + + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES (10, 'x')") + require.NoError(t, err) + for _, node := range logicPlan.GetQuery().Nodes { + for _, target := range node.LockTargets { + if target.TableId == indexID && target.Mode == lockpb.LockMode_Shared { + assert.Equal(t, int32(types.T_varchar), target.PrimaryColTyp.Id) + return + } + } + } + t.Fatal("referenced unique-index shared lock not found") +} + +func TestReplaceAndChildInsertUseCanonicalForeignKeyLockOrder(t *testing.T) { + mock := NewMockOptimizer(true) + parent := mock.ctxt.tables["replace_fk_p"] + child := mock.ctxt.tables["replace_fk_c"] + if parent.Name2ColIndex == nil { + parent.Name2ColIndex = make(map[string]int32) + for i, col := range parent.Cols { + parent.Name2ColIndex[col.Name] = int32(i) + } + } + if child.Name2ColIndex == nil { + child.Name2ColIndex = make(map[string]int32) + for i, col := range child.Cols { + child.Name2ColIndex[col.Name] = int32(i) + } + } + parentPos := len(parent.Cols) - 1 + parent.Cols = append(parent.Cols, nil) + copy(parent.Cols[parentPos+1:], parent.Cols[parentPos:]) + parent.Cols[parentPos] = &plan.ColDef{ + Name: "k", ColId: 3, Typ: plan.Type{Id: int32(types.T_varchar), Width: 20}, + } + parent.Name2ColIndex["k"] = int32(parentPos) + parent.Name2ColIndex[catalog.Row_ID] = int32(parentPos + 1) + child.Cols[1].Typ = plan.Type{Id: int32(types.T_varchar), Width: 20} + childPos := len(child.Cols) - 1 + child.Cols = append(child.Cols, nil) + copy(child.Cols[childPos+1:], child.Cols[childPos:]) + child.Cols[childPos] = &plan.ColDef{ + Name: "pid2", ColId: 2, Typ: plan.Type{Id: int32(types.T_varchar), Width: 20}, + } + child.Name2ColIndex["pid2"] = int32(childPos) + child.Name2ColIndex[catalog.Row_ID] = int32(childPos + 1) + child.Fkeys[0].ForeignCols = []uint64{1} + child.Fkeys = append(child.Fkeys, &plan.ForeignKeyDef{ + Cols: []uint64{2}, ForeignTbl: parent.TblId, ForeignCols: []uint64{3}, + }) + + addIndex := func(indexName, tableName string, tableID uint64, part string) { + parent.Indexes = append(parent.Indexes, &plan.IndexDef{ + IndexName: indexName, IndexTableName: tableName, Parts: []string{part}, + Unique: true, TableExist: true, IndexAlgo: catalog.MoIndexDefaultAlgo.ToString(), + }) + mock.ctxt.tables[tableName] = &plan.TableDef{ + TblId: tableID, Name: tableName, + Cols: []*plan.ColDef{ + {Name: catalog.IndexTableIndexColName, ColId: 0, Typ: plan.Type{Id: int32(types.T_varchar), Width: 20}}, + {Name: catalog.Row_ID, ColId: 1, Hidden: true, Typ: plan.Type{Id: int32(types.T_Rowid)}}, + }, + Pkey: &plan.PrimaryKeyDef{Names: []string{catalog.IndexTableIndexColName}, + PkeyColName: catalog.IndexTableIndexColName}, + Name2ColIndex: map[string]int32{catalog.IndexTableIndexColName: 0, catalog.Row_ID: 1}, + } + mock.ctxt.objects[tableName] = &plan.ObjectRef{ + Obj: int64(tableID), SchemaName: mock.ctxt.objects["replace_fk_p"].SchemaName, ObjName: tableName, + } + } + // Declaration order is z then a; physical lock order must be a then z. + addIndex("uk_v", "__mo_index_z", 77911, "v") + addIndex("uk_k", "__mo_index_a", 77912, "k") + + logicPlan, err := runOneStmt(mock, t, "INSERT INTO replace_fk_c VALUES (10, 'x', 'y')") + require.NoError(t, err) + query := logicPlan.GetQuery() + lockNode := make(map[uint64]int32) + for nodeID, node := range query.Nodes { + for _, target := range node.LockTargets { + if target.Mode == lockpb.LockMode_Shared { + lockNode[target.TableId] = int32(nodeID) + } + } + } + require.Contains(t, lockNode, uint64(77911)) + require.Contains(t, lockNode, uint64(77912)) + var contains func(int32, int32) bool + contains = func(root, target int32) bool { + if root == target { + return true + } + for _, childID := range query.Nodes[root].Children { + if contains(childID, target) { + return true + } + } + return false + } + assert.True(t, contains(lockNode[77911], lockNode[77912]), + "z lock must depend on the lexically earlier a lock regardless of FK declaration order") + + replacePlan, err := runOneStmt(mock, t, "REPLACE INTO replace_fk_p VALUES (1, 'x', 'y')") + require.NoError(t, err) + var replaceLockOrder []uint64 + for _, node := range replacePlan.GetQuery().Nodes { + if node.NodeType != plan.Node_LOCK_OP || len(node.LockTargets) == 0 { + continue + } + for _, target := range node.LockTargets { + replaceLockOrder = append(replaceLockOrder, target.TableId) + } + break + } + require.Equal(t, []uint64{parent.TblId, parent.TblId, 77912, 77912, 77911, 77911}, replaceLockOrder, + "REPLACE must lock the base table first and hidden unique indexes by physical table name") +} + +func TestDeepCopyPreservesSharedLockMode(t *testing.T) { + assert.Nil(t, DeepCopyLockTarget(nil)) + target := &plan.LockTarget{ + TableId: 42, + ObjRef: &plan.ObjectRef{Obj: 42, ObjName: "parent"}, + Mode: lockpb.LockMode_Shared, + PrimaryColRelPos: 11, + FilterColRelPos: 12, + PartitionColIdxInBat: 13, + HasPartitionCol: true, + LockRows: makePlan2Int64ConstExprWithType(7), + } + assertScalarFields := func(t *testing.T, copied *plan.LockTarget) { + t.Helper() + assert.Equal(t, lockpb.LockMode_Shared, copied.Mode) + assert.Equal(t, int32(11), copied.PrimaryColRelPos) + assert.Equal(t, int32(12), copied.FilterColRelPos) + assert.Equal(t, int32(13), copied.PartitionColIdxInBat) + assert.True(t, copied.HasPartitionCol) + } + + direct := DeepCopyLockTarget(target) + require.NotSame(t, target, direct) + assertScalarFields(t, direct) + require.NotSame(t, target.ObjRef, direct.ObjRef) + require.NotSame(t, target.LockRows, direct.LockRows) + + node := &plan.Node{NodeType: plan.Node_LOCK_OP, LockTargets: []*plan.LockTarget{target}} + nodeCopy := DeepCopyNode(node) + require.Len(t, nodeCopy.LockTargets, 1) + assertScalarFields(t, nodeCopy.LockTargets[0]) + require.NotSame(t, target, nodeCopy.LockTargets[0]) + + queryCopy := DeepCopyQuery(&plan.Query{Nodes: []*plan.Node{node}}) + require.Len(t, queryCopy.Nodes, 1) + require.Len(t, queryCopy.Nodes[0].LockTargets, 1) + assertScalarFields(t, queryCopy.Nodes[0].LockTargets[0]) + require.NotSame(t, target, queryCopy.Nodes[0].LockTargets[0]) +} + +func TestReplaceParentSideFKOmittedUniqueDefaults(t *testing.T) { + parseReplace := func(t *testing.T, sql string) *tree.Replace { + t.Helper() + stmt, err := mysql.ParseOne(context.Background(), sql, 1) + require.NoError(t, err) + return stmt.(*tree.Replace) + } + newParent := func(t *testing.T) (*MockOptimizer, *plan.TableDef) { + t.Helper() + mock := NewMockOptimizer(true) + parent := DeepCopyTableDef(mock.ctxt.tables["replace_fk_cp"], true) + parent.Indexes = append(parent.Indexes, &plan.IndexDef{ + Unique: true, + Parts: []string{"v"}, + }) + return mock, parent + } + parentCol := func(t *testing.T, parent *plan.TableDef, name string) *plan.ColDef { + t.Helper() + for _, col := range parent.Cols { + if col.Name == name { + return col + } + } + t.Fatalf("missing parent column %s", name) + return nil + } + + t.Run("nullable default cannot conflict", func(t *testing.T) { + mock, parent := newParent(t) + v := parentCol(t, parent, "v") + v.Default = &plan.Default{NullAbility: true} + + _, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, + parseReplace(t, "REPLACE INTO replace_fk_cp(id) VALUES (1)")) + require.NoError(t, err) + require.Len(t, actions, 1) + assert.Contains(t, actions[0], "`__mo_replace_parent`.`id` = cast(1 as INT)") + assert.NotContains(t, actions[0], "`__mo_replace_parent`.`v` =") + }) + + t.Run("constant prefix default participates", func(t *testing.T) { + mock, parent := newParent(t) + parent.Indexes[len(parent.Indexes)-1].IndexAlgoParams = `{"prefix_lengths":"v:4"}` + v := parentCol(t, parent, "v") + v.Default = &plan.Default{ + NullAbility: true, + Expr: makeStringConstExpr(v.Typ, "abcdyyyy"), + } + + _, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, + parseReplace(t, "REPLACE INTO replace_fk_cp(id) VALUES (1)")) + require.NoError(t, err) + require.Len(t, actions, 1) + assert.Contains(t, actions[0], "substring(`__mo_replace_parent`.`v`, 1, 4)") + assert.Contains(t, actions[0], `substring(cast("abcdyyyy" as VARCHAR(20)), 1, 4)`) + }) + + t.Run("dynamic default fails closed", func(t *testing.T) { + mock, parent := newParent(t) + v := parentCol(t, parent, "v") + v.Default = &plan.Default{ + NullAbility: true, + Expr: &plan.Expr{ + Typ: v.Typ, + Expr: &plan.Expr_Col{Col: &plan.ColRef{Name: "dynamic_default"}}, + }, + } + + _, _, _, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, + parseReplace(t, "REPLACE INTO replace_fk_cp(id) VALUES (1)")) + require.ErrorContains(t, err, "non-literal default conflict key") + }) + + t.Run("numeric literal defaults participate", func(t *testing.T) { + cases := []struct { + name string + typ plan.Type + literal *plan.Expr_Lit + expected string + }{ + {name: "int8", typ: plan.Type{Id: int32(types.T_int8)}, literal: makePlan2Int8ConstExpr(-8), expected: "cast(-8 as TINYINT)"}, + {name: "int16", typ: plan.Type{Id: int32(types.T_int16)}, literal: makePlan2Int16ConstExpr(-16), expected: "cast(-16 as SMALLINT)"}, + {name: "int32", typ: plan.Type{Id: int32(types.T_int32)}, literal: makePlan2Int32ConstExpr(-32), expected: "cast(-32 as INT)"}, + {name: "int64", typ: plan.Type{Id: int32(types.T_int64)}, literal: makePlan2Int64ConstExpr(-64), expected: "cast(-64 as BIGINT)"}, + {name: "uint8", typ: plan.Type{Id: int32(types.T_uint8)}, literal: makePlan2Uint8ConstExpr(8), expected: "cast(8 as TINYINT UNSIGNED)"}, + {name: "uint16", typ: plan.Type{Id: int32(types.T_uint16)}, literal: makePlan2Uint16ConstExpr(16), expected: "cast(16 as SMALLINT UNSIGNED)"}, + {name: "uint32", typ: plan.Type{Id: int32(types.T_uint32)}, literal: makePlan2Uint32ConstExpr(32), expected: "cast(32 as INT UNSIGNED)"}, + {name: "uint64", typ: plan.Type{Id: int32(types.T_uint64)}, literal: makePlan2Uint64ConstExpr(64), expected: "cast(64 as BIGINT UNSIGNED)"}, + {name: "float32", typ: plan.Type{Id: int32(types.T_float32)}, literal: makePlan2Float32ConstExpr(1.25), expected: "cast(1.25 as FLOAT)"}, + {name: "float64", typ: plan.Type{Id: int32(types.T_float64)}, literal: makePlan2Float64ConstExpr(2.5), expected: "cast(2.5 as DOUBLE)"}, + {name: "bool", typ: plan.Type{Id: int32(types.T_bool)}, literal: makePlan2BoolConstExpr(true), expected: "cast(true as BOOL)"}, + { + name: "enum", typ: plan.Type{Id: int32(types.T_enum), Enumvalues: "small,medium,large"}, + literal: &plan.Expr_Lit{Lit: &plan.Literal{Value: &plan.Literal_EnumVal{EnumVal: 2}}}, + expected: `cast(2 as ENUM("small","medium","large"))`, + }, + { + name: "decimal64", typ: plan.Type{Id: int32(types.T_decimal64), Width: 5, Scale: 2}, + literal: &plan.Expr_Lit{Lit: &plan.Literal{Value: &plan.Literal_Decimal64Val{ + Decimal64Val: &plan.Decimal64{A: 123}, + }}}, expected: "cast(1.23 as DECIMAL(5,2))", + }, + { + name: "decimal128", typ: plan.Type{Id: int32(types.T_decimal128), Width: 20, Scale: 2}, + literal: &plan.Expr_Lit{Lit: &plan.Literal{Value: &plan.Literal_Decimal128Val{ + Decimal128Val: &plan.Decimal128{A: 123}, + }}}, expected: "cast(1.23 as DECIMAL(20,2))", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + mock, parent := newParent(t) + v := parentCol(t, parent, "v") + v.Typ = tc.typ + v.Default = &plan.Default{NullAbility: false, Expr: &plan.Expr{Typ: tc.typ, Expr: tc.literal}} + + lockSQL, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, + parseReplace(t, "REPLACE INTO replace_fk_cp(id) VALUES (1)")) + require.NoError(t, err) + require.Len(t, actions, 1) + assert.Contains(t, lockSQL, tc.expected) + assert.Contains(t, actions[0], tc.expected) + _, err = mysql.ParseOne(context.Background(), lockSQL, 1) + require.NoError(t, err, "generated parent lock SQL must be parseable") + }) + } + }) + + t.Run("temporal defaults participate", func(t *testing.T) { + dateValue, err := types.ParseDateCast("2026-07-15") + require.NoError(t, err) + timeValue, err := types.ParseTime("12:34:56.123", 3) + require.NoError(t, err) + datetimeValue, err := types.ParseDatetime("2026-07-15 12:34:56.123", 3) + require.NoError(t, err) + + location := time.UTC + mockForLocation := NewMockOptimizer(true) + if sessionLocation := mockForLocation.ctxt.GetProcess().GetSessionInfo().TimeZone; sessionLocation != nil { + location = sessionLocation + } + timestampValue, err := types.ParseTimestamp(location, "2026-07-15 12:34:56.123", 3) + require.NoError(t, err) + + cases := []struct { + name string + typ plan.Type + literal *plan.Expr_Lit + expected string + }{ + { + name: "date", + typ: plan.Type{Id: int32(types.T_date)}, + literal: makePlan2DateConstExpr(int32(dateValue)), + expected: `"2026-07-15"`, + }, + { + name: "time", + typ: plan.Type{Id: int32(types.T_time), Scale: 3}, + literal: makePlan2TimeConstExpr(int64(timeValue)), + expected: `"12:34:56.123"`, + }, + { + name: "datetime", + typ: plan.Type{Id: int32(types.T_datetime), Scale: 3}, + literal: makePlan2DateTimeConstExpr(int64(datetimeValue)), + expected: `"2026-07-15 12:34:56.123"`, + }, + { + name: "timestamp", + typ: plan.Type{Id: int32(types.T_timestamp), Scale: 3}, + literal: makePlan2TimestampConstExpr(int64(timestampValue)), + expected: `"2026-07-15 12:34:56.123"`, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + mock, parent := newParent(t) + v := parentCol(t, parent, "v") + v.Typ = tc.typ + v.Default = &plan.Default{ + NullAbility: true, + Expr: &plan.Expr{Typ: tc.typ, Expr: tc.literal}, + } + + _, _, actions, err := genParentSideReplaceFKSqls( + &mock.ctxt, mock.ctxt.objects["replace_fk_cp"], parent, + parseReplace(t, "REPLACE INTO replace_fk_cp(id) VALUES (1)")) + require.NoError(t, err) + require.Len(t, actions, 1) + expectedType := strings.ToUpper(types.T(tc.typ.Id).String()) + if tc.typ.Scale > 0 && types.T(tc.typ.Id) != types.T_date { + expectedType += fmt.Sprintf("(%d)", tc.typ.Scale) + } + assert.Contains(t, actions[0], "`__mo_replace_parent`.`v` = cast("+tc.expected+" as "+expectedType+")") + }) + } + }) +} + func TestReplaceODKU(t *testing.T) { mock := NewMockOptimizer(true) // INSERT ON DUPLICATE KEY UPDATE should be rewritten to REPLACE path diff --git a/pkg/sql/plan/build_util.go b/pkg/sql/plan/build_util.go index 60e17167805a5..921e8194a311c 100644 --- a/pkg/sql/plan/build_util.go +++ b/pkg/sql/plan/build_util.go @@ -18,7 +18,11 @@ import ( "context" "fmt" "regexp" + "slices" + "strconv" "strings" + "time" + "unicode/utf8" "github.com/matrixorigin/matrixone/pkg/catalog" "github.com/matrixorigin/matrixone/pkg/common/moerr" @@ -1188,6 +1192,630 @@ func genPreCheckSqlsForReplaceFKSelfRefer( return ret, nil } +func genParentSideReplaceFKSqls( + ctx CompilerContext, + parentRef *plan.ObjectRef, + parent *plan.TableDef, + stmt *tree.Replace, +) (string, []string, []string, error) { + if stmt.Rows == nil || len(parent.RefChildTbls) == 0 { + return "", nil, nil, nil + } + hasNonSelfReference := false + for _, childID := range parent.RefChildTbls { + if childID != 0 { + hasNonSelfReference = true + break + } + } + if !hasNonSelfReference { + return "", nil, nil, nil + } + values, ok := stmt.Rows.Select.(*tree.ValuesClause) + if !ok { + return "", nil, nil, moerr.NewNotSupported(ctx.GetContext(), "REPLACE SELECT/TABLE on a referenced parent table") + } + positions := make(map[string]int) + if len(stmt.Columns) > 0 { + pos := 0 + for _, col := range stmt.Columns { + name := strings.ToLower(string(col)) + if colIdx, found := parent.Name2ColIndex[name]; found && + int(colIdx) < len(parent.Cols) && parent.Cols[colIdx].GeneratedCol != nil { + continue + } + positions[name] = pos + pos++ + } + } else { + pos := 0 + for _, col := range parent.Cols { + if !col.Hidden && col.GeneratedCol == nil { + positions[strings.ToLower(col.Name)] = pos + pos++ + } + } + } + quoteIdentifier := func(name string) string { + return "`" + strings.ReplaceAll(name, "`", "``") + "`" + } + qualifiedCol := func(alias, name string) string { + return quoteIdentifier(alias) + "." + quoteIdentifier(name) + } + findParentCol := func(name string) (*plan.ColDef, bool) { + if pos, found := parent.Name2ColIndex[name]; found && int(pos) < len(parent.Cols) { + return parent.Cols[pos], true + } + for _, col := range parent.Cols { + if col.Name == name { + return col, true + } + } + return nil, false + } + treatAutoIncrementZeroAsGenerated := true + if sqlMode, err := ctx.ResolveVariable("sql_mode", true, false); err != nil { + return "", nil, nil, err + } else if mode, ok := sqlMode.(string); ok { + treatAutoIncrementZeroAsGenerated = !strings.Contains(strings.ToUpper(mode), "NO_AUTO_VALUE_ON_ZERO") + } + isAssignmentConvertedZero := func(expr tree.Expr, col *plan.ColDef) (bool, error) { + var value *tree.NumVal + switch input := expr.(type) { + case *tree.NumVal: + value = input + case *tree.StrVal: + value = tree.NewNumVal(input.String(), input.String(), false, tree.P_char) + default: + return false, nil + } + if value.ValType == tree.P_null || value.ValType == tree.P_bool { + return false, nil + } + proc := ctx.GetProcess() + if proc == nil { + return false, moerr.NewInternalError(ctx.GetContext(), + "cannot materialize auto-increment value without a process") + } + converted, err := MakeInsertValueConstExpr(proc, value, &types.Type{ + Oid: types.T(col.Typ.Id), + Width: col.Typ.Width, + Scale: col.Typ.Scale, + }) + if err != nil { + return false, err + } + if converted == nil { + binder := NewDefaultBinder(ctx.GetContext(), nil, nil, plan.Type{}, nil) + converted, err = binder.BindExpr(expr, 0, true) + if err != nil { + return false, err + } + converted, err = forceAssignmentCastExpr(ctx.GetContext(), converted, col.Typ) + if err != nil { + return false, err + } + converted, err = ConstantFold(batch.EmptyForConstFoldBatch, converted, proc, false, true) + if err != nil { + return false, err + } + } + lit := converted.GetLit() + if lit == nil { + return false, nil + } + switch val := lit.Value.(type) { + case *plan.Literal_I8Val: + return val.I8Val == 0, nil + case *plan.Literal_I16Val: + return val.I16Val == 0, nil + case *plan.Literal_I32Val: + return val.I32Val == 0, nil + case *plan.Literal_I64Val: + return val.I64Val == 0, nil + case *plan.Literal_U8Val: + return val.U8Val == 0, nil + case *plan.Literal_U16Val: + return val.U16Val == 0, nil + case *plan.Literal_U32Val: + return val.U32Val == 0, nil + case *plan.Literal_U64Val: + return val.U64Val == 0, nil + default: + return false, nil + } + } + + type uniqueKey struct { + parts []string + prefixLengths map[string]int + } + uniqueKeys := make([]uniqueKey, 0, 1+len(parent.Indexes)) + if parent.Pkey != nil && len(parent.Pkey.Names) > 0 { + uniqueKeys = append(uniqueKeys, uniqueKey{parts: parent.Pkey.Names}) + } + for _, idx := range parent.Indexes { + if !idx.Unique { + continue + } + parts := make([]string, len(idx.Parts)) + for i, part := range idx.Parts { + parts[i] = catalog.ResolveAlias(part) + } + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(idx.IndexAlgoParams) + if err != nil { + return "", nil, nil, err + } + uniqueKeys = append(uniqueKeys, uniqueKey{parts: parts, prefixLengths: prefixLengths}) + } + + const parentAlias = "__mo_replace_parent" + literalFmt := tree.NewFmtCtx(dialect.MYSQL, tree.WithQuoteString(true)) + formatStringLiteral := func(value string) string { + tree.NewNumVal(value, value, false, tree.P_char).Format(literalFmt) + formatted := literalFmt.String() + literalFmt.Reset() + return formatted + } + formatLiteral := func(expr *plan.Expr) (string, bool, bool) { + lit := expr.GetLit() + if lit == nil { + return "", false, false + } + if lit.Isnull { + return "", true, true + } + switch val := lit.Value.(type) { + case *plan.Literal_I8Val: + return strconv.FormatInt(int64(val.I8Val), 10), false, true + case *plan.Literal_I16Val: + return strconv.FormatInt(int64(val.I16Val), 10), false, true + case *plan.Literal_I32Val: + return strconv.FormatInt(int64(val.I32Val), 10), false, true + case *plan.Literal_I64Val: + return strconv.FormatInt(val.I64Val, 10), false, true + case *plan.Literal_U8Val: + return strconv.FormatUint(uint64(val.U8Val), 10), false, true + case *plan.Literal_U16Val: + return strconv.FormatUint(uint64(val.U16Val), 10), false, true + case *plan.Literal_U32Val: + return strconv.FormatUint(uint64(val.U32Val), 10), false, true + case *plan.Literal_U64Val: + return strconv.FormatUint(val.U64Val, 10), false, true + case *plan.Literal_Fval: + return strconv.FormatFloat(float64(val.Fval), 'g', -1, 32), false, true + case *plan.Literal_Dval: + return strconv.FormatFloat(val.Dval, 'g', -1, 64), false, true + case *plan.Literal_Bval: + return strconv.FormatBool(val.Bval), false, true + case *plan.Literal_EnumVal: + return strconv.FormatUint(uint64(val.EnumVal), 10), false, true + case *plan.Literal_Decimal64Val: + return types.Decimal64(val.Decimal64Val.A).Format(expr.Typ.Scale), false, true + case *plan.Literal_Decimal128Val: + decimal := types.Decimal128{ + B0_63: uint64(val.Decimal128Val.A), + B64_127: uint64(val.Decimal128Val.B), + } + return decimal.Format(expr.Typ.Scale), false, true + case *plan.Literal_Dateval: + return formatStringLiteral(types.Date(val.Dateval).String()), false, true + case *plan.Literal_Timeval: + return formatStringLiteral(types.Time(val.Timeval).String2(expr.Typ.Scale)), false, true + case *plan.Literal_Datetimeval: + return formatStringLiteral(types.Datetime(val.Datetimeval).String2(expr.Typ.Scale)), false, true + case *plan.Literal_Timestampval: + location := time.UTC + if proc := ctx.GetProcess(); proc != nil && proc.GetSessionInfo().TimeZone != nil { + location = proc.GetSessionInfo().TimeZone + } + value := types.Timestamp(val.Timestampval).String2(location, expr.Typ.Scale) + return formatStringLiteral(value), false, true + case *plan.Literal_Sval: + return formatStringLiteral(val.Sval), false, true + default: + return "", false, false + } + } + materializeDefault := func(colName string) (string, bool, error) { + col, found := findParentCol(colName) + if !found { + return "", false, moerr.NewInternalErrorf(ctx.GetContext(), + "REPLACE conflict column %s not found", colName) + } + if col.GeneratedCol != nil { + return "", false, moerr.NewNotSupported(ctx.GetContext(), + "REPLACE with an omitted generated conflict key") + } + defaultExpr, err := getDefaultExpr(ctx.GetContext(), col) + if err != nil { + return "", false, err + } + formatted, isNull, ok := formatLiteral(defaultExpr) + if !ok { + return "", false, moerr.NewNotSupported(ctx.GetContext(), + "REPLACE with a non-literal default conflict key") + } + return formatted, isNull, nil + } + sqlTypeForColumn := func(col *plan.ColDef) string { + typ := types.T(col.Typ.Id) + switch typ { + case types.T_time, types.T_datetime, types.T_timestamp: + if col.Typ.Scale > 0 { + return fmt.Sprintf("%s(%d)", makeTypeByPlan2Type(col.Typ).String(), col.Typ.Scale) + } + case types.T_enum: + values := strings.Split(col.Typ.Enumvalues, ",") + for i := range values { + values[i] = formatStringLiteral(values[i]) + } + return "ENUM(" + strings.Join(values, ",") + ")" + } + if isSetPlanType(&col.Typ) { + values := strings.Split(col.Typ.Enumvalues, ",") + for i := range values { + values[i] = formatStringLiteral(values[i]) + } + return "SET(" + strings.Join(values, ",") + ")" + } + return makeTypeByPlan2Type(col.Typ).DescString() + } + castForColumn := func(value string, col *plan.ColDef) string { + return fmt.Sprintf("cast(%s as %s)", value, sqlTypeForColumn(col)) + } + formatInputLiteral := func(expr tree.Expr, col *plan.ColDef) (string, error) { + switch value := expr.(type) { + case *tree.NumVal: + if value.ValType == tree.P_char && + (types.T(col.Typ.Id) == types.T_char || types.T(col.Typ.Id) == types.T_varchar) && + col.Typ.Width > 0 && int32(utf8.RuneCountInString(value.String())) > col.Typ.Width { + return "", moerr.NewInvalidInputf(ctx.GetContext(), + "Src length %d is larger than Dest length %d", + utf8.RuneCountInString(value.String()), col.Typ.Width) + } + expr.Format(literalFmt) + case *tree.StrVal: + if (types.T(col.Typ.Id) == types.T_char || types.T(col.Typ.Id) == types.T_varchar) && + col.Typ.Width > 0 && int32(utf8.RuneCountInString(value.String())) > col.Typ.Width { + return "", moerr.NewInvalidInputf(ctx.GetContext(), + "Src length %d is larger than Dest length %d", + utf8.RuneCountInString(value.String()), col.Typ.Width) + } + expr.Format(literalFmt) + default: + return "", moerr.NewNotSupported(ctx.GetContext(), "REPLACE with a non-literal conflict key") + } + formatted := literalFmt.String() + literalFmt.Reset() + return formatted, nil + } + var materializeInputColumn func(tree.Exprs, int, map[int]bool) (string, error) + materializeInputColumn = func(row tree.Exprs, colIdx int, visiting map[int]bool) (string, error) { + if colIdx < 0 || colIdx >= len(parent.Cols) { + return "", moerr.NewInternalErrorf(ctx.GetContext(), "REPLACE conflict column position %d not found", colIdx) + } + col := parent.Cols[colIdx] + if col.GeneratedCol != nil { + if visiting[colIdx] { + return "", moerr.NewInternalErrorf(ctx.GetContext(), + "cyclic generated column dependency at %s", col.Name) + } + visiting[colIdx] = true + defer delete(visiting, colIdx) + if strings.TrimSpace(col.GeneratedCol.OriginString) == "" { + return "", moerr.NewNotSupportedf(ctx.GetContext(), + "REPLACE with generated conflict key %s lacking its source expression", col.Name) + } + refs := collectRefColPos(col.GeneratedCol.Expr) + slices.Sort(refs) + refs = slices.Compact(refs) + selectParts := make([]string, 0, len(refs)) + for _, refPos := range refs { + refExpr, err := materializeInputColumn(row, int(refPos), visiting) + if err != nil { + return "", err + } + selectParts = append(selectParts, fmt.Sprintf("%s as %s", refExpr, + quoteIdentifier(parent.Cols[refPos].Name))) + } + generatedExpr := castForColumn(col.GeneratedCol.OriginString, col) + if len(selectParts) == 0 { + return "(select " + generatedExpr + ")", nil + } + return fmt.Sprintf("(select %s from (select %s) as %s)", generatedExpr, + strings.Join(selectParts, ", "), quoteIdentifier("__mo_replace_input")), nil + } + + pos, supplied := positions[strings.ToLower(col.Name)] + if !supplied || pos >= len(row) { + if col.Typ.AutoIncr { + return "null", nil + } + value, isNull, err := materializeDefault(col.Name) + if err != nil { + return "", err + } + if isNull { + return "null", nil + } + return castForColumn(value, col), nil + } + if _, ok := row[pos].(*tree.DefaultVal); ok { + value, isNull, err := materializeDefault(col.Name) + if err != nil { + return "", err + } + if isNull { + return "null", nil + } + return castForColumn(value, col), nil + } + if col.Typ.AutoIncr && treatAutoIncrementZeroAsGenerated { + zero, err := isAssignmentConvertedZero(row[pos], col) + if err != nil { + return "", err + } + if zero { + return "null", nil + } + } + value, err := formatInputLiteral(row[pos], col) + if err != nil { + return "", err + } + return castForColumn(value, col), nil + } + conflictPredicates := make([]string, 0, len(values.Rows)*len(uniqueKeys)) + for _, row := range values.Rows { + for _, key := range uniqueKeys { + parts := make([]string, 0, len(key.parts)) + keyCannotConflict := false + for _, colName := range key.parts { + col, found := findParentCol(colName) + if !found { + return "", nil, nil, moerr.NewInternalErrorf(ctx.GetContext(), + "REPLACE conflict column %s not found", colName) + } + pos, supplied := positions[strings.ToLower(colName)] + var incomingExpr string + if !supplied || pos >= len(row) { + if col.Typ.AutoIncr { + keyCannotConflict = true + break + } + var isNull bool + var err error + if col.GeneratedCol != nil { + colIdx := -1 + for i, candidate := range parent.Cols { + if strings.EqualFold(candidate.Name, colName) { + colIdx = i + break + } + } + incomingExpr, err = materializeInputColumn(row, colIdx, make(map[int]bool)) + } else { + incomingExpr, isNull, err = materializeDefault(colName) + } + if err != nil { + return "", nil, nil, err + } + if isNull { + keyCannotConflict = true + break + } + } else { + zero := false + if col.Typ.AutoIncr && treatAutoIncrementZeroAsGenerated { + var zeroErr error + zero, zeroErr = isAssignmentConvertedZero(row[pos], col) + if zeroErr != nil { + return "", nil, nil, zeroErr + } + } + if zero { + // PRE_INSERT turns an explicit numeric zero into NULL and + // allocates a fresh auto-increment value unless + // NO_AUTO_VALUE_ON_ZERO is enabled. Such a value cannot + // identify an old parent row during this pre-phase. + keyCannotConflict = true + break + } + switch value := row[pos].(type) { + case *tree.NumVal: + if value.ValType == tree.P_char && + (types.T(col.Typ.Id) == types.T_char || types.T(col.Typ.Id) == types.T_varchar) && + col.Typ.Width > 0 && int32(utf8.RuneCountInString(value.String())) > col.Typ.Width { + return "", nil, nil, moerr.NewInvalidInputf(ctx.GetContext(), + "Src length %d is larger than Dest length %d", + utf8.RuneCountInString(value.String()), col.Typ.Width) + } + row[pos].Format(literalFmt) + incomingExpr = literalFmt.String() + literalFmt.Reset() + case *tree.StrVal: + if (types.T(col.Typ.Id) == types.T_char || types.T(col.Typ.Id) == types.T_varchar) && + col.Typ.Width > 0 && int32(utf8.RuneCountInString(value.String())) > col.Typ.Width { + return "", nil, nil, moerr.NewInvalidInputf(ctx.GetContext(), + "Src length %d is larger than Dest length %d", + utf8.RuneCountInString(value.String()), col.Typ.Width) + } + row[pos].Format(literalFmt) + incomingExpr = literalFmt.String() + literalFmt.Reset() + case *tree.DefaultVal: + var isNull bool + var err error + incomingExpr, isNull, err = materializeDefault(colName) + if err != nil { + return "", nil, nil, err + } + if isNull { + keyCannotConflict = true + break + } + default: + return "", nil, nil, moerr.NewNotSupported(ctx.GetContext(), "REPLACE with a non-literal conflict key") + } + } + if keyCannotConflict { + break + } + incomingExpr = castForColumn(incomingExpr, col) + parentExpr := qualifiedCol(parentAlias, colName) + if length := key.prefixLengths[colName]; length > 0 { + parentExpr = fmt.Sprintf("substring(%s, 1, %d)", parentExpr, length) + incomingExpr = fmt.Sprintf("substring(%s, 1, %d)", incomingExpr, length) + } + parts = append(parts, fmt.Sprintf("%s = %s", parentExpr, incomingExpr)) + } + if keyCannotConflict { + continue + } + conflictPredicates = append(conflictPredicates, "("+strings.Join(parts, " and ")+")") + } + } + if len(conflictPredicates) == 0 { + return "", nil, nil, nil + } + conflictPredicate := "(" + strings.Join(conflictPredicates, " or ") + ")" + parentTable := quoteIdentifier(parentRef.SchemaName) + "." + quoteIdentifier(parent.Name) + + var checks, actions []string + referencedIndexes := make(map[string]*plan.IndexDef) + partsEqual := func(parts, names []string) bool { + if len(parts) != len(names) { + return false + } + for i := range parts { + if catalog.ResolveAlias(parts[i]) != names[i] { + return false + } + } + return true + } + pkeyNames := []string(nil) + if parent.Pkey != nil { + pkeyNames = parent.Pkey.Names + if len(pkeyNames) == 0 && parent.Pkey.PkeyColName != "" { + pkeyNames = []string{parent.Pkey.PkeyColName} + } + } + seen := make(map[uint64]bool) + for _, childID := range parent.RefChildTbls { + if childID == 0 || seen[childID] { + continue + } + seen[childID] = true + childRef, child, err := ctx.ResolveById(childID, nil) + if err != nil { + return "", nil, nil, err + } + if child == nil { + return "", nil, nil, moerr.NewInternalError(ctx.GetContext(), fmt.Sprintf("referencing table %d not found", childID)) + } + for _, fk := range child.Fkeys { + if fk.ForeignTbl != parent.TblId { + continue + } + if len(fk.Cols) == 0 || len(fk.Cols) != len(fk.ForeignCols) { + return "", nil, nil, moerr.NewInternalError(ctx.GetContext(), "invalid parent foreign key definition") + } + parentCols, err := colIdsToNames(ctx.GetContext(), fk.ForeignCols, parent.Cols) + if err != nil { + return "", nil, nil, err + } + if !partsEqual(pkeyNames, parentCols) { + for _, idx := range parent.Indexes { + if idx.Unique && partsEqual(idx.Parts, parentCols) { + if idx.IndexTableName == "" { + return "", nil, nil, moerr.NewInternalErrorf(ctx.GetContext(), + "unique index %s has no index table", idx.IndexName) + } + referencedIndexes[idx.IndexTableName] = idx + break + } + } + } + childCols, err := colIdsToNames(ctx.GetContext(), fk.Cols, child.Cols) + if err != nil { + return "", nil, nil, err + } + childTable := quoteIdentifier(childRef.SchemaName) + "." + quoteIdentifier(child.Name) + joinParts := make([]string, len(childCols)) + for i := range childCols { + joinParts[i] = fmt.Sprintf("%s.%s = %s", + childTable, quoteIdentifier(childCols[i]), qualifiedCol(parentAlias, parentCols[i])) + } + exists := fmt.Sprintf("exists (select 1 from %s as %s where %s and %s)", + parentTable, quoteIdentifier(parentAlias), conflictPredicate, strings.Join(joinParts, " and ")) + switch fk.OnDelete { + case plan.ForeignKeyDef_RESTRICT, plan.ForeignKeyDef_NO_ACTION, plan.ForeignKeyDef_SET_DEFAULT: + checks = append(checks, fmt.Sprintf("select count(*) = 0 from %s where %s", childTable, exists)) + case plan.ForeignKeyDef_CASCADE: + actions = append(actions, fmt.Sprintf("delete from %s where %s", childTable, exists)) + case plan.ForeignKeyDef_SET_NULL: + setParts := make([]string, len(childCols)) + for i, col := range childCols { + setParts[i] = quoteIdentifier(col) + " = null" + } + actions = append(actions, fmt.Sprintf("update %s set %s where %s", + childTable, strings.Join(setParts, ", "), exists)) + } + } + } + + indexNames := make([]string, 0, len(referencedIndexes)) + for name := range referencedIndexes { + indexNames = append(indexNames, name) + } + slices.Sort(indexNames) + lockSelects := make([]string, 0, len(indexNames)+1) + if parent.Pkey != nil && parent.Pkey.PkeyColName != "" { + lockSelects = append(lockSelects, qualifiedCol(parentAlias, parent.Pkey.PkeyColName)) + } else { + lockSelects = append(lockSelects, "1") + } + for i, indexName := range indexNames { + idx := referencedIndexes[indexName] + prefixLengths, err := catalog.IndexPrefixLengthsFromParamsWithError(idx.IndexAlgoParams) + if err != nil { + return "", nil, nil, err + } + keyParts := make([]string, len(idx.Parts)) + for partPos, part := range idx.Parts { + colName := catalog.ResolveAlias(part) + col, found := findParentCol(colName) + if !found { + return "", nil, nil, moerr.NewInternalErrorf(ctx.GetContext(), + "REPLACE referenced index column %s not found", colName) + } + partExpr := qualifiedCol(parentAlias, colName) + if length := prefixLengths[colName]; length > 0 { + partExpr = fmt.Sprintf("substring(%s, 1, %d)", partExpr, length) + if prefixType, ok := indexTableKeyTypeForPrefix(col.Typ); ok { + partExpr = fmt.Sprintf("cast(%s as %s)", partExpr, makeTypeByPlan2Type(prefixType).DescString()) + } + } + keyParts[partPos] = partExpr + } + keyExpr := keyParts[0] + if indexTableStoresSerializedKey(idx) { + keyExpr = "serial(" + strings.Join(keyParts, ", ") + ")" + } + indexAlias := fmt.Sprintf("__mo_replace_fk_idx_%d", i) + indexTable := quoteIdentifier(parentRef.SchemaName) + "." + quoteIdentifier(indexName) + indexKey := qualifiedCol(indexAlias, catalog.IndexTableIndexColName) + lockSelects = append(lockSelects, fmt.Sprintf( + "(select %s from %s as %s where %s = %s for update)", + indexKey, indexTable, quoteIdentifier(indexAlias), indexKey, keyExpr)) + } + parentLock := fmt.Sprintf("select %s from %s as %s where %s for update", + strings.Join(lockSelects, ", "), parentTable, quoteIdentifier(parentAlias), conflictPredicate) + return parentLock, checks, actions, nil +} + func cleanHint(originSql string) string { re := regexp.MustCompile(`/\*[^!].*?\*/`) cleanSQL := re.ReplaceAllString(originSql, "") diff --git a/pkg/sql/plan/deepcopy.go b/pkg/sql/plan/deepcopy.go index 27bab76f99dd4..cae9c956b73ba 100644 --- a/pkg/sql/plan/deepcopy.go +++ b/pkg/sql/plan/deepcopy.go @@ -157,16 +157,21 @@ func DeepCopyLockTarget(target *plan.LockTarget) *plan.LockTarget { return nil } return &plan.LockTarget{ - TableId: target.TableId, - ObjRef: DeepCopyObjectRef(target.ObjRef), - PrimaryColIdxInBat: target.PrimaryColIdxInBat, - PrimaryColTyp: target.PrimaryColTyp, - RefreshTsIdxInBat: target.RefreshTsIdxInBat, - FilterColIdxInBat: target.FilterColIdxInBat, - LockTable: target.LockTable, - Block: target.Block, - LockRows: DeepCopyExpr(target.LockRows), - LockTableAtTheEnd: target.LockTableAtTheEnd, + TableId: target.TableId, + ObjRef: DeepCopyObjectRef(target.ObjRef), + PrimaryColIdxInBat: target.PrimaryColIdxInBat, + PrimaryColTyp: target.PrimaryColTyp, + RefreshTsIdxInBat: target.RefreshTsIdxInBat, + FilterColIdxInBat: target.FilterColIdxInBat, + LockTable: target.LockTable, + Block: target.Block, + Mode: target.Mode, + PrimaryColRelPos: target.PrimaryColRelPos, + FilterColRelPos: target.FilterColRelPos, + LockRows: DeepCopyExpr(target.LockRows), + LockTableAtTheEnd: target.LockTableAtTheEnd, + PartitionColIdxInBat: target.PartitionColIdxInBat, + HasPartitionCol: target.HasPartitionCol, } } @@ -609,6 +614,7 @@ func DeepCopyQuery(qry *plan.Query) *plan.Query { Params: DeepCopyExprList(qry.Params), Headings: qry.Headings, HasForeignKeyAction: qry.HasForeignKeyAction, + DetectSqls: slices.Clone(qry.DetectSqls), } for idx, node := range qry.Nodes { newQry.Nodes[idx] = DeepCopyNode(node) diff --git a/pkg/sql/plan/mock.go b/pkg/sql/plan/mock.go index 21f104cca5b22..32b02e0ee0955 100644 --- a/pkg/sql/plan/mock.go +++ b/pkg/sql/plan/mock.go @@ -53,6 +53,8 @@ type MockCompilerContext struct { GetDatabaseIdFunc func(string, *Snapshot) (uint64, error) ResolveAccountIdsFunc func([]string) ([]uint32, error) ResolveFunc func(string, string, *Snapshot) (*ObjectRef, *TableDef) + ResolveVariableFunc func(string, bool, bool) (interface{}, error) + GetProcessFunc func() *process.Process } func (m *MockCompilerContext) GetLowerCaseTableNames() int64 { @@ -102,6 +104,9 @@ func (m *MockCompilerContext) ResolveAccountIds(accountNames []string) ([]uint32 } func (m *MockCompilerContext) ResolveVariable(varName string, isSystemVar, isGlobalVar bool) (interface{}, error) { + if m.ResolveVariableFunc != nil { + return m.ResolveVariableFunc(varName, isSystemVar, isGlobalVar) + } vars := make(map[string]interface{}) vars["str_var"] = "str" vars["int_var"] = 20 @@ -155,15 +160,16 @@ func NewEmptyCompilerContext() *MockCompilerContext { } type Schema struct { - cols []col - pks []int - idxs []index - fks []*ForeignKeyDef - clusterby *ClusterByDef - outcnt float64 - tblId int64 - isView bool - viewCfg ViewCfg + cols []col + pks []int + idxs []index + fks []*ForeignKeyDef + refChildTbls []uint64 + clusterby *ClusterByDef + outcnt float64 + tblId int64 + isView bool + viewCfg ViewCfg // tableType overrides TableType when non-empty; used to mock index tables // carrying an algo-specific type (e.g. ivfflat "metadata"). tableType string @@ -1037,7 +1043,175 @@ func NewMockCompilerContext(isDml bool) *MockCompilerContext { OnUpdate: plan.ForeignKeyDef_CASCADE, }, }, - outcnt: 10, + refChildTbls: []uint64{0}, + outcnt: 10, + } + + /* + Parent-side FK action fixtures for REPLACE (issue #24951). + + create table replace_fk_p(id int primary key, v varchar(20)); + create table replace_fk_c(id int primary key, pid int, + foreign key(pid) references replace_fk_p(id) on delete restrict); + + create table replace_fk_cp(id int primary key, v varchar(20)); + create table replace_fk_cc(id int primary key, pid int, + foreign key(pid) references replace_fk_cp(id) on delete cascade); + */ + constraintTestSchema["replace_fk_p"] = &Schema{ + tblId: 77001, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"v", types.T_varchar, true, 20, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + refChildTbls: []uint64{77002}, + outcnt: 4, + } + constraintTestSchema["replace_fk_c"] = &Schema{ + tblId: 77002, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"pid", types.T_int32, true, 32, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + fks: []*plan.ForeignKeyDef{ + { + Name: "fk_replace_c", + Cols: []uint64{1}, // pid + ForeignTbl: 77001, + ForeignCols: []uint64{0}, // replace_fk_p.id + OnDelete: plan.ForeignKeyDef_RESTRICT, + OnUpdate: plan.ForeignKeyDef_RESTRICT, + }, + }, + outcnt: 4, + } + constraintTestSchema["replace_fk_cp"] = &Schema{ + tblId: 77003, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"v", types.T_varchar, true, 20, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + refChildTbls: []uint64{77004}, + outcnt: 4, + } + constraintTestSchema["replace_fk_cc"] = &Schema{ + tblId: 77004, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"pid", types.T_int32, true, 32, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + fks: []*plan.ForeignKeyDef{ + { + Name: "fk_replace_cc", + Cols: []uint64{1}, // pid + ForeignTbl: 77003, + ForeignCols: []uint64{0}, // replace_fk_cp.id + OnDelete: plan.ForeignKeyDef_CASCADE, + OnUpdate: plan.ForeignKeyDef_CASCADE, + }, + }, + outcnt: 4, + } + constraintTestSchema["replace_fk_sp"] = &Schema{ + tblId: 77005, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"v", types.T_varchar, true, 20, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + refChildTbls: []uint64{77006}, + outcnt: 4, + } + constraintTestSchema["replace_fk_sc"] = &Schema{ + tblId: 77006, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"pid", types.T_int32, true, 32, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + fks: []*plan.ForeignKeyDef{ + { + Name: "fk_replace_sc", + Cols: []uint64{1}, // pid + ForeignTbl: 77005, + ForeignCols: []uint64{0}, // replace_fk_sp.id + OnDelete: plan.ForeignKeyDef_SET_NULL, + OnUpdate: plan.ForeignKeyDef_SET_NULL, + }, + }, + outcnt: 4, + } + constraintTestSchema["replace_fk_np"] = &Schema{ + tblId: 77007, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"v", types.T_varchar, true, 20, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + refChildTbls: []uint64{77008}, + outcnt: 4, + } + constraintTestSchema["replace_fk_nc"] = &Schema{ + tblId: 77008, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"pid", types.T_int32, true, 32, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + fks: []*plan.ForeignKeyDef{ + { + Name: "fk_replace_nc", + Cols: []uint64{1}, // pid + ForeignTbl: 77007, + ForeignCols: []uint64{0}, // replace_fk_np.id + OnDelete: plan.ForeignKeyDef_NO_ACTION, + OnUpdate: plan.ForeignKeyDef_NO_ACTION, + }, + }, + outcnt: 4, + } + constraintTestSchema["replace_fk_dp"] = &Schema{ + tblId: 77009, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"v", types.T_varchar, true, 20, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + refChildTbls: []uint64{77010}, + outcnt: 4, + } + constraintTestSchema["replace_fk_dc"] = &Schema{ + tblId: 77010, + cols: []col{ + {"id", types.T_int32, true, 32, 0}, + {"pid", types.T_int32, true, 32, 0}, + {catalog.Row_ID, types.T_Rowid, false, 16, 0}, + }, + pks: []int{0}, + fks: []*plan.ForeignKeyDef{ + { + Name: "fk_replace_dc", + Cols: []uint64{1}, // pid + ForeignTbl: 77009, + ForeignCols: []uint64{0}, // replace_fk_dp.id + OnDelete: plan.ForeignKeyDef_SET_DEFAULT, + OnUpdate: plan.ForeignKeyDef_SET_DEFAULT, + }, + }, + outcnt: 4, } /* @@ -1513,6 +1687,10 @@ func NewMockCompilerContext(isDml bool) *MockCompilerContext { tableDef.Fkeys = table.fks } + if table.refChildTbls != nil { + tableDef.RefChildTbls = table.refChildTbls + } + if table.clusterby != nil { tableDef.ClusterBy = &plan.ClusterByDef{ Name: "__mo_cbkey_003pid005pname", @@ -1702,6 +1880,9 @@ func (m *MockCompilerContext) SetContext(ctx context.Context) { } func (m *MockCompilerContext) GetProcess() *process.Process { + if m.GetProcessFunc != nil { + return m.GetProcessFunc() + } proc := testutil.NewProc(nil) moruntime.ServiceRuntime(proc.GetService()).SetGlobalVariables( moruntime.InternalSQLExecutor, diff --git a/pkg/sql/plan/opt_misc.go b/pkg/sql/plan/opt_misc.go index 7665a52022a31..85151f2c37215 100644 --- a/pkg/sql/plan/opt_misc.go +++ b/pkg/sql/plan/opt_misc.go @@ -120,6 +120,15 @@ func (builder *QueryBuilder) removeSimpleProjections(nodeID int32, parentType pl } } + case plan.Node_LOCK_OP: + for i, childID := range node.Children { + newChildID, childProjMap := builder.removeSimpleProjections(childID, node.NodeType, true, colRefCnt) + node.Children[i] = newChildID + for ref, expr := range childProjMap { + projMap[ref] = expr + } + } + default: for i, childID := range node.Children { newChildID, childProjMap := builder.removeSimpleProjections(childID, node.NodeType, flag, colRefCnt) diff --git a/pkg/sql/plan/pushdown.go b/pkg/sql/plan/pushdown.go index 8ab2ca0ab7676..e0eb5223153d7 100644 --- a/pkg/sql/plan/pushdown.go +++ b/pkg/sql/plan/pushdown.go @@ -43,6 +43,11 @@ func (builder *QueryBuilder) pushdownFilters(nodeID int32, filters []*plan.Expr, switch node.NodeType { case plan.Node_AGG: + // Legacy positional aggregates have no global binding tags. Keep filters + // above them because tag-based replacement cannot address their outputs. + if len(node.BindingTags) < 2 { + return originalNodeID, filters + } groupTag := node.BindingTags[0] aggregateTag := node.BindingTags[1] @@ -157,6 +162,11 @@ func (builder *QueryBuilder) pushdownFilters(nodeID int32, filters []*plan.Expr, node.Children[0] = childID case plan.Node_FILTER: + // IsEnd filters are terminal assertions/action selectors. Moving their + // predicates below joins can change both assertion scope and marker layout. + if node.IsEnd { + return originalNodeID, filters + } canPushdown = filters if !node.RollupFilter { for _, filter := range node.FilterList { @@ -511,6 +521,9 @@ func (builder *QueryBuilder) pushdownFilters(nodeID int32, filters []*plan.Expr, break } + if len(node.BindingTags) == 0 { + node.BindingTags = []int32{0} + } projectTag := node.BindingTags[0] for _, filter := range filters { diff --git a/pkg/sql/plan/query_builder.go b/pkg/sql/plan/query_builder.go index ed80af0d71419..338a5e418b3e4 100644 --- a/pkg/sql/plan/query_builder.go +++ b/pkg/sql/plan/query_builder.go @@ -1136,6 +1136,15 @@ func (builder *QueryBuilder) remapAllColRefs(nodeID int32, step int32, colRefCnt } case plan.Node_AGG: + // Some DML duplicate-check aggregates intentionally expose positional + // output. They are already fully projected and cannot participate in the + // global-tag pruning protocol used by regular aggregates. + if len(node.BindingTags) < 2 { + for i := range node.ProjectList { + remapping.addColRef([2]int32{0, int32(i)}) + } + return remapping, nil + } groupTag := node.BindingTags[0] aggregateTag := node.BindingTags[1] groupSize := int32(len(node.GroupBy)) @@ -1944,8 +1953,16 @@ func (builder *QueryBuilder) remapAllColRefs(nodeID int32, step int32, colRefCnt } case plan.Node_SINK_SCAN, plan.Node_RECURSIVE_SCAN, plan.Node_RECURSIVE_CTE: + if len(node.BindingTags) == 0 { + node.BindingTags = []int32{0} + } tag := node.BindingTags[0] var newProjList []*plan.Expr + if _, preserve := builder.preserveScanProjection[nodeID]; preserve { + for i := range node.ProjectList { + colRefCnt[[2]int32{tag, int32(i)}] = 1 + } + } for i, expr := range node.ProjectList { globalRef := [2]int32{tag, int32(i)} @@ -1970,11 +1987,32 @@ func (builder *QueryBuilder) remapAllColRefs(nodeID int32, step int32, colRefCnt } } + if len(newProjList) == 0 && len(node.ProjectList) > 0 { + for i, expr := range node.ProjectList { + globalRef := [2]int32{tag, int32(i)} + newProjList = append(newProjList, &plan.Expr{ + Typ: expr.Typ, + Expr: &plan.Expr_Col{Col: &ColRef{ + RelPos: 0, + ColPos: int32(i), + }}, + }) + remapping.addColRef(globalRef) + } + } node.ProjectList = newProjList case plan.Node_SINK: childNode := builder.qry.Nodes[node.Children[0]] + if len(childNode.BindingTags) == 0 { + childNode.BindingTags = []int32{0} + } resultTag := childNode.BindingTags[0] + if _, preserve := builder.preserveSinkProjection[nodeID]; preserve { + for i := range node.ProjectList { + colRefBool[[2]int32{step, int32(i)}] = true + } + } for i := range childNode.ProjectList { if colRefBool[[2]int32{step, int32(i)}] { colRefCnt[[2]int32{resultTag, int32(i)}] = 1 @@ -2133,6 +2171,11 @@ func (builder *QueryBuilder) remapAllColRefs(nodeID int32, step int32, colRefCnt case plan.Node_LOCK_OP: preNode := builder.qry.Nodes[node.Children[0]] + if _, preserve := builder.preserveLockProjection[nodeID]; preserve && len(preNode.BindingTags) > 0 { + for i := range preNode.ProjectList { + colRefCnt[[2]int32{preNode.BindingTags[0], int32(i)}] = 1 + } + } var pkExprs []*plan.Expr var oldPkPos [][2]int32 @@ -2274,6 +2317,11 @@ func (builder *QueryBuilder) remapAllColRefs(nodeID int32, step int32, colRefCnt } case plan.Node_INSERT, plan.Node_DELETE: + if _, preserve := builder.preserveInsertProjection[nodeID]; preserve { + for _, expr := range builder.qry.Nodes[node.Children[0]].ProjectList { + increaseRefCnt(expr, 1, colRefCnt) + } + } childRemapping, err := builder.remapAllColRefs(node.Children[0], step, colRefCnt, colRefBool, sinkColRef) if err != nil { return nil, err @@ -2400,6 +2448,11 @@ func (builder *QueryBuilder) remapAllColRefs(nodeID int32, step int32, colRefCnt } case plan.Node_PRE_INSERT: + if _, preserve := builder.preservePreInsertProjection[nodeID]; preserve { + for _, expr := range builder.qry.Nodes[node.Children[0]].ProjectList { + increaseRefCnt(expr, 1, colRefCnt) + } + } childRemapping, err := builder.remapAllColRefs(node.Children[0], step, colRefCnt, colRefBool, sinkColRef) if err != nil { return nil, err @@ -2703,6 +2756,21 @@ func (builder *QueryBuilder) createQuery() (*Query, error) { // after createQuery can translate pre-prune positions into the materialized // sink's post-prune layout. builder.sinkColRef = sinkColRef + for nodeID := range builder.positionalSinkScans { + node := builder.qry.Nodes[nodeID] + if node.NodeType != plan.Node_SINK_SCAN || len(node.SourceStep) == 0 { + continue + } + for _, expr := range node.ProjectList { + col, ok := expr.Expr.(*plan.Expr_Col) + if !ok { + continue + } + if newPos, ok := sinkColRef[[2]int32{node.SourceStep[0], col.Col.ColPos}]; ok { + col.Col.ColPos = int32(newPos) + } + } + } err = builder.lockTableIfLockNoRowsAtTheEndForDelAndUpdate() if err != nil { diff --git a/pkg/sql/plan/types.go b/pkg/sql/plan/types.go index 3dc242f28530e..d1d461d9c2b34 100644 --- a/pkg/sql/plan/types.go +++ b/pkg/sql/plan/types.go @@ -180,12 +180,18 @@ type QueryBuilder struct { qry *plan.Query compCtx CompilerContext - ctxByNode []*BindContext - nameByColRef map[[2]int32]string - protectedScans map[int32]int - projectSpecialGuards map[int32]*specialIndexGuard - indexHintsByScan map[int32]*indexHintSet - indexHintOwnerByNode map[int32]int32 + ctxByNode []*BindContext + nameByColRef map[[2]int32]string + protectedScans map[int32]int + projectSpecialGuards map[int32]*specialIndexGuard + indexHintsByScan map[int32]*indexHintSet + indexHintOwnerByNode map[int32]int32 + preserveSinkProjection map[int32]struct{} + preserveLockProjection map[int32]struct{} + preservePreInsertProjection map[int32]struct{} + preserveInsertProjection map[int32]struct{} + preserveScanProjection map[int32]struct{} + positionalSinkScans map[int32]struct{} tag2Table map[int32]*TableDef tag2NodeID map[int32]int32 diff --git a/test/distributed/cases/dml/replace/replace.result b/test/distributed/cases/dml/replace/replace.result index f8651c6949111..f6d2e971f16e4 100644 --- a/test/distributed/cases/dml/replace/replace.result +++ b/test/distributed/cases/dml/replace/replace.result @@ -287,6 +287,294 @@ name email value a y@a.com 111 b z@a.com 222 drop table t_replace_multi_uk_batch; +drop table if exists fk_c; +drop table if exists fk_p; +create table fk_p(id int primary key, v varchar(20)); +create table fk_c(id int primary key, pid int, foreign key(pid) references fk_p(id) on delete restrict); +insert into fk_p values (1,'p1'); +insert into fk_c values (10,1); +replace into fk_p values (1,'p1_new'); +internal error: Cannot delete or update a parent row: a foreign key constraint fails +select * from fk_p order by id; +id v +1 p1 +select * from fk_c order by id; +id pid +10 1 +replace into fk_p values (2,'p2_new'); +select * from fk_p order by id; +id v +1 p1 +2 p2_new +drop table fk_c; +drop table fk_p; +drop table if exists fk_cc; +drop table if exists fk_cp; +create table fk_cp(id int primary key, v varchar(20)); +create table fk_cc(id int primary key, pid int, foreign key(pid) references fk_cp(id) on delete cascade); +insert into fk_cp values (1,'p1'); +insert into fk_cc values (10,1); +replace into fk_cp values (1,'p1_new'); +select * from fk_cp order by id; +id v +1 p1_new +select * from fk_cc order by id; +id pid +drop table fk_cc; +drop table fk_cp; +drop table if exists fk_sc; +drop table if exists fk_sp; +create table fk_sp(id int primary key, v varchar(20)); +create table fk_sc(id int primary key, pid int, foreign key(pid) references fk_sp(id) on delete set null); +insert into fk_sp values (1,'p1'); +insert into fk_sc values (10,1); +replace into fk_sp values (1,'p1_new'); +select * from fk_sp order by id; +id v +1 p1_new +select * from fk_sc order by id; +id pid +10 null +drop table fk_sc; +drop table fk_sp; +drop table if exists fk_dup_sc; +drop table if exists fk_dup_sp; +create table fk_dup_sp(id int primary key); +create table fk_dup_sc(pid1 int, pid2 int, note int, +foreign key(pid1) references fk_dup_sp(id) on delete set null, +foreign key(pid2) references fk_dup_sp(id) on delete set null); +insert into fk_dup_sp values (1); +insert into fk_dup_sc values (1,1,7),(1,1,7); +replace into fk_dup_sp values (1); +select count(*) from fk_dup_sc where pid1 is null and pid2 is null and note = 7; +count(*) +2 +select count(*) from fk_dup_sc where pid1 = 1 or pid2 = 1; +count(*) +0 +drop table fk_dup_sc; +drop table fk_dup_sp; +drop table if exists fk_dc; +drop table if exists fk_dp; +create table fk_dp(id int primary key, v varchar(20)); +create table fk_dc(id int primary key, pid int default 2, +foreign key(pid) references fk_dp(id) on delete set default); +insert into fk_dp values (1,'p1'),(2,'p2'); +insert into fk_dc values (10,1); +replace into fk_dp values (1,'p1_new'); +internal error: Cannot delete or update a parent row: a foreign key constraint fails +select * from fk_dp order by id; +id v +1 p1 +2 p2 +select * from fk_dc order by id; +id pid +10 1 +drop table fk_dc; +drop table fk_dp; +drop table if exists fk_mc; +drop table if exists fk_mp; +create table fk_mp(id int primary key, v varchar(20)); +create table fk_mc(id int primary key, pid int, foreign key(pid) references fk_mp(id) on delete cascade); +insert into fk_mp values (1,'p1'),(2,'p2'),(3,'p3'); +insert into fk_mc values (10,1),(20,2),(30,3); +replace into fk_mp values (1,'p1_new'),(2,'p2_new'); +select * from fk_mp order by id; +id v +1 p1_new +2 p2_new +3 p3 +select * from fk_mc order by id; +id pid +30 3 +drop table fk_mc; +drop table fk_mp; +drop table if exists fk_review_c; +drop table if exists fk_review_p; +create table fk_review_p(id int primary key, u int unique, v int); +create table fk_review_c(id int primary key, pid int, +foreign key(pid) references fk_review_p(id) on delete cascade); +insert into fk_review_p values (1, 10, 100); +insert into fk_review_c values (1, 1); +replace into fk_review_p values (2, 10, 200); +select * from fk_review_p order by id; +id u v +2 10 200 +select * from fk_review_c order by id; +id pid +create table fk_review_src(id int, u int, v int); +insert into fk_review_src values (1, 10, 400); +replace into fk_review_p select * from fk_review_src; +select * from fk_review_p order by id; +id u v +1 10 400 +select * from fk_review_c order by id; +id pid +drop table fk_review_src; +drop table fk_review_c; +drop table fk_review_p; +create table fk_auto_p(id int auto_increment primary key, u int unique, v int); +create table fk_auto_c(id int primary key, pid int, +foreign key(pid) references fk_auto_p(id) on delete cascade); +insert into fk_auto_p(u, v) values (10, 100); +insert into fk_auto_c values (1, 1); +replace into fk_auto_p(u, v) values (10, 200); +select u, v from fk_auto_p; +u v +10 200 +select * from fk_auto_c; +id pid +drop table fk_auto_c; +drop table fk_auto_p; +create table fk_nonpk_p(id int primary key, u int unique); +create table fk_nonpk_c(id int primary key, parent_u int, +foreign key(parent_u) references fk_nonpk_p(u) on delete cascade); +insert into fk_nonpk_p values (1, 10); +insert into fk_nonpk_c values (1, 10); +replace into fk_nonpk_p values (2, 10); +select * from fk_nonpk_p; +id u +2 10 +select * from fk_nonpk_c; +id parent_u +drop table fk_nonpk_c; +drop table fk_nonpk_p; +create table fk_fanout_p(id int primary key, u int unique, v int unique); +create table fk_fanout_c(id int primary key, pid int, +foreign key(pid) references fk_fanout_p(id) on delete cascade); +insert into fk_fanout_p values (1, 10, 100), (2, 20, 200); +insert into fk_fanout_c values (1, 1), (2, 2); +replace into fk_fanout_p values (3, 10, 200); +select * from fk_fanout_p; +id u v +3 10 200 +select * from fk_fanout_c; +id pid +drop table fk_fanout_c; +drop table fk_fanout_p; +create table fk_prefix_p(id int primary key, body varchar(64), unique key u(body(4))); +create table fk_prefix_c(id int primary key, pid int, +foreign key(pid) references fk_prefix_p(id) on delete cascade); +insert into fk_prefix_p values (1, 'abcdxxxx'); +insert into fk_prefix_c values (1, 1); +replace into fk_prefix_p values (2, 'abcdyyyy'); +select * from fk_prefix_p; +id body +2 abcdyyyy +select * from fk_prefix_c; +id pid +drop table fk_prefix_c; +drop table fk_prefix_p; +create table fk_omitted_uk_p( +id int primary key, +u varchar(20) unique, +v int +); +create table fk_omitted_uk_c( +id int primary key, +pid int, +foreign key(pid) references fk_omitted_uk_p(id) on delete cascade +); +insert into fk_omitted_uk_p values (1, 'x', 100); +insert into fk_omitted_uk_c values (1, 1); +replace into fk_omitted_uk_p(id) values (1); +select * from fk_omitted_uk_p; +id u v +1 NULL NULL +select * from fk_omitted_uk_c; +id pid +insert into fk_omitted_uk_p values (2, 'x', 200); +select * from fk_omitted_uk_p order by id; +id u v +1 NULL NULL +2 x 200 +drop table fk_omitted_uk_c; +drop table fk_omitted_uk_p; +create table fk_temporal_default_p( +id int primary key, +d date unique default '2026-07-15', +t time unique default '12:34:56', +dt datetime unique default '2026-07-15 12:34:56', +ts timestamp unique default '2026-07-15 12:34:56' +); +create table fk_temporal_default_c( +id int primary key, +pid int, +foreign key(pid) references fk_temporal_default_p(id) on delete cascade +); +insert into fk_temporal_default_p(id) values (1); +insert into fk_temporal_default_c values (1, 1); +replace into fk_temporal_default_p(id) values (2); +select * from fk_temporal_default_p; +id d t dt ts +2 2026-07-15 12:34:56 2026-07-15 12:34:56 2026-07-15 12:34:56 +select * from fk_temporal_default_c; +id pid +drop table fk_temporal_default_c; +drop table fk_temporal_default_p; +create table fk_decimal_cast_p(id int primary key, u decimal(5,2) unique); +create table fk_decimal_cast_c(id int primary key, pid int, +foreign key(pid) references fk_decimal_cast_p(id) on delete cascade); +insert into fk_decimal_cast_p values (1, 1.23); +insert into fk_decimal_cast_c values (1, 1); +replace into fk_decimal_cast_p values (2, 1.234); +select * from fk_decimal_cast_p; +id u +2 1.23 +select * from fk_decimal_cast_c; +id pid +drop table fk_decimal_cast_c; +drop table fk_decimal_cast_p; +create table fk_param_p(id int primary key, v int); +create table fk_param_c(id int primary key, pid int, +foreign key(pid) references fk_param_p(id) on delete restrict); +prepare fk_param_key_stmt from 'replace into fk_param_p values (?, ?)'; +set @fk_id = 1, @fk_v = 300; +execute fk_param_key_stmt using @fk_id, @fk_v; +deallocate prepare fk_param_key_stmt; +insert into fk_param_p values (1, 100); +Duplicate entry '1' for key 'id' +insert into fk_param_c values (1, 1); +prepare fk_review_stmt from 'replace into fk_param_p values (1, ?)'; +set @fk_v = 300; +execute fk_review_stmt using @fk_v; +internal error: Cannot delete or update a parent row: a foreign key constraint fails +select * from fk_param_p; +id v +1 300 +select * from fk_param_c; +id pid +1 1 +deallocate prepare fk_review_stmt; +drop table fk_param_c; +drop table fk_param_p; +drop table if exists fk_comp_c; +drop table if exists fk_comp_p; +create table fk_comp_p(a int, b int, primary key(a, b)); +create table fk_comp_c(id int primary key, a int, b int, +foreign key(a, b) references fk_comp_p(a, b) on delete cascade); +insert into fk_comp_p values (1, 1); +insert into fk_comp_c values (1, 1, 1); +replace into fk_comp_p values (1, 1); +select * from fk_comp_p; +a b +1 1 +select * from fk_comp_c; +id a b +drop table fk_comp_c; +drop table fk_comp_p; +drop table if exists `fk``tick_c`; +drop table if exists `fk``tick_p`; +create table `fk``tick_p`(`id``x` int primary key); +create table `fk``tick_c`(id int primary key, `pid``x` int, +foreign key(`pid``x`) references `fk``tick_p`(`id``x`) on delete cascade); +insert into `fk``tick_p` values (1); +insert into `fk``tick_c` values (1, 1); +replace into `fk``tick_p` values (1); +select * from `fk``tick_c`; +id pid`x +drop table `fk``tick_c`; +drop table `fk``tick_p`; drop table if exists replace_fk_c; drop table if exists replace_fk_p; create table replace_fk_p(id int primary key); @@ -421,3 +709,50 @@ null 1 null-a-2 1 1 one-replaced 2 2 two drop table t_replace_fakepk_comp_null; +drop table if exists t_replace_self_cascade; +create table t_replace_self_cascade ( +id int primary key, +parent_id int, +foreign key (parent_id) references t_replace_self_cascade(id) on delete cascade +); +insert into t_replace_self_cascade values (1, null), (2, 1), (3, 2), (4, 1); +replace into t_replace_self_cascade values (1, null); +select * from t_replace_self_cascade order by id; +id parent_id +1 null +insert into t_replace_self_cascade values (2, 1), (3, 2); +replace into t_replace_self_cascade values (10, null); +select * from t_replace_self_cascade order by id; +id parent_id +1 null +2 1 +3 2 +10 null +delete from t_replace_self_cascade; +insert into t_replace_self_cascade values (1, null); +update t_replace_self_cascade set parent_id = 1 where id = 1; +replace into t_replace_self_cascade values (1, 1); +select row_count(); +row_count() +2 +select * from t_replace_self_cascade order by id; +id parent_id +1 1 +delete from t_replace_self_cascade; +insert into t_replace_self_cascade values (1, null), (2, 1); +replace into t_replace_self_cascade values (1, null), (2, 1); +select row_count(); +row_count() +4 +select * from t_replace_self_cascade order by id; +id parent_id +1 null +2 1 +delete from t_replace_self_cascade; +insert into t_replace_self_cascade values (1, null), (2, 1); +update t_replace_self_cascade set parent_id = 2 where id = 1; +replace into t_replace_self_cascade values (1, null); +select * from t_replace_self_cascade order by id; +id parent_id +1 null +drop table t_replace_self_cascade; diff --git a/test/distributed/cases/dml/replace/replace.test b/test/distributed/cases/dml/replace/replace.test index bb46ebaa75672..c4b3bb60727fa 100644 --- a/test/distributed/cases/dml/replace/replace.test +++ b/test/distributed/cases/dml/replace/replace.test @@ -237,6 +237,249 @@ insert into t_replace_multi_uk_batch (name, email, value) values ('a', 'x@a.com' replace into t_replace_multi_uk_batch (name, email, value) values ('a', 'y@a.com', 111), ('b', 'z@a.com', 222); select name, email, value from t_replace_multi_uk_batch order by name; drop table t_replace_multi_uk_batch; +-- parent-side foreign key actions during REPLACE (delete-then-insert semantics) +-- ON DELETE RESTRICT: replacing a referenced parent row must fail and leave it unchanged +drop table if exists fk_c; +drop table if exists fk_p; +create table fk_p(id int primary key, v varchar(20)); +create table fk_c(id int primary key, pid int, foreign key(pid) references fk_p(id) on delete restrict); +insert into fk_p values (1,'p1'); +insert into fk_c values (10,1); +replace into fk_p values (1,'p1_new'); +select * from fk_p order by id; +select * from fk_c order by id; +-- replacing an unreferenced parent row is allowed +replace into fk_p values (2,'p2_new'); +select * from fk_p order by id; +drop table fk_c; +drop table fk_p; + +-- ON DELETE CASCADE: replacing a referenced parent row cascades the delete to children +drop table if exists fk_cc; +drop table if exists fk_cp; +create table fk_cp(id int primary key, v varchar(20)); +create table fk_cc(id int primary key, pid int, foreign key(pid) references fk_cp(id) on delete cascade); +insert into fk_cp values (1,'p1'); +insert into fk_cc values (10,1); +replace into fk_cp values (1,'p1_new'); +select * from fk_cp order by id; +select * from fk_cc order by id; +drop table fk_cc; +drop table fk_cp; + +-- ON DELETE SET NULL: replacing a referenced parent row nulls the child fk columns +drop table if exists fk_sc; +drop table if exists fk_sp; +create table fk_sp(id int primary key, v varchar(20)); +create table fk_sc(id int primary key, pid int, foreign key(pid) references fk_sp(id) on delete set null); +insert into fk_sp values (1,'p1'); +insert into fk_sc values (10,1); +replace into fk_sp values (1,'p1_new'); +select * from fk_sp order by id; +select * from fk_sc order by id; +drop table fk_sc; +drop table fk_sp; + +-- combined SET NULL keeps physically distinct rows with identical business columns +drop table if exists fk_dup_sc; +drop table if exists fk_dup_sp; +create table fk_dup_sp(id int primary key); +create table fk_dup_sc(pid1 int, pid2 int, note int, + foreign key(pid1) references fk_dup_sp(id) on delete set null, + foreign key(pid2) references fk_dup_sp(id) on delete set null); +insert into fk_dup_sp values (1); +insert into fk_dup_sc values (1,1,7),(1,1,7); +replace into fk_dup_sp values (1); +select count(*) from fk_dup_sc where pid1 is null and pid2 is null and note = 7; +select count(*) from fk_dup_sc where pid1 = 1 or pid2 = 1; +drop table fk_dup_sc; +drop table fk_dup_sp; + +-- ON DELETE SET DEFAULT follows DELETE semantics and restricts referenced parents +drop table if exists fk_dc; +drop table if exists fk_dp; +create table fk_dp(id int primary key, v varchar(20)); +create table fk_dc(id int primary key, pid int default 2, + foreign key(pid) references fk_dp(id) on delete set default); +insert into fk_dp values (1,'p1'),(2,'p2'); +insert into fk_dc values (10,1); +replace into fk_dp values (1,'p1_new'); +select * from fk_dp order by id; +select * from fk_dc order by id; +drop table fk_dc; +drop table fk_dp; + +-- multi-row REPLACE: every conflicting parent row applies its child action (CASCADE) +drop table if exists fk_mc; +drop table if exists fk_mp; +create table fk_mp(id int primary key, v varchar(20)); +create table fk_mc(id int primary key, pid int, foreign key(pid) references fk_mp(id) on delete cascade); +insert into fk_mp values (1,'p1'),(2,'p2'),(3,'p3'); +insert into fk_mc values (10,1),(20,2),(30,3); +replace into fk_mp values (1,'p1_new'),(2,'p2_new'); +select * from fk_mp order by id; +select * from fk_mc order by id; +drop table fk_mc; +drop table fk_mp; + +-- Unsupported parent-side FK shapes must fail instead of silently skipping actions. +drop table if exists fk_review_c; +drop table if exists fk_review_p; +create table fk_review_p(id int primary key, u int unique, v int); +create table fk_review_c(id int primary key, pid int, + foreign key(pid) references fk_review_p(id) on delete cascade); +insert into fk_review_p values (1, 10, 100); +insert into fk_review_c values (1, 1); +replace into fk_review_p values (2, 10, 200); +select * from fk_review_p order by id; +select * from fk_review_c order by id; +create table fk_review_src(id int, u int, v int); +insert into fk_review_src values (1, 10, 400); +replace into fk_review_p select * from fk_review_src; +select * from fk_review_p order by id; +select * from fk_review_c order by id; +drop table fk_review_src; +drop table fk_review_c; +drop table fk_review_p; + +create table fk_auto_p(id int auto_increment primary key, u int unique, v int); +create table fk_auto_c(id int primary key, pid int, + foreign key(pid) references fk_auto_p(id) on delete cascade); +insert into fk_auto_p(u, v) values (10, 100); +insert into fk_auto_c values (1, 1); +replace into fk_auto_p(u, v) values (10, 200); +select u, v from fk_auto_p; +select * from fk_auto_c; +drop table fk_auto_c; +drop table fk_auto_p; + +create table fk_nonpk_p(id int primary key, u int unique); +create table fk_nonpk_c(id int primary key, parent_u int, + foreign key(parent_u) references fk_nonpk_p(u) on delete cascade); +insert into fk_nonpk_p values (1, 10); +insert into fk_nonpk_c values (1, 10); +replace into fk_nonpk_p values (2, 10); +select * from fk_nonpk_p; +select * from fk_nonpk_c; +drop table fk_nonpk_c; +drop table fk_nonpk_p; + +create table fk_fanout_p(id int primary key, u int unique, v int unique); +create table fk_fanout_c(id int primary key, pid int, + foreign key(pid) references fk_fanout_p(id) on delete cascade); +insert into fk_fanout_p values (1, 10, 100), (2, 20, 200); +insert into fk_fanout_c values (1, 1), (2, 2); +replace into fk_fanout_p values (3, 10, 200); +select * from fk_fanout_p; +select * from fk_fanout_c; +drop table fk_fanout_c; +drop table fk_fanout_p; + +create table fk_prefix_p(id int primary key, body varchar(64), unique key u(body(4))); +create table fk_prefix_c(id int primary key, pid int, + foreign key(pid) references fk_prefix_p(id) on delete cascade); +insert into fk_prefix_p values (1, 'abcdxxxx'); +insert into fk_prefix_c values (1, 1); +replace into fk_prefix_p values (2, 'abcdyyyy'); +select * from fk_prefix_p; +select * from fk_prefix_c; +drop table fk_prefix_c; +drop table fk_prefix_p; + +create table fk_omitted_uk_p( + id int primary key, + u varchar(20) unique, + v int +); +create table fk_omitted_uk_c( + id int primary key, + pid int, + foreign key(pid) references fk_omitted_uk_p(id) on delete cascade +); +insert into fk_omitted_uk_p values (1, 'x', 100); +insert into fk_omitted_uk_c values (1, 1); +replace into fk_omitted_uk_p(id) values (1); +select * from fk_omitted_uk_p; +select * from fk_omitted_uk_c; +insert into fk_omitted_uk_p values (2, 'x', 200); +select * from fk_omitted_uk_p order by id; +drop table fk_omitted_uk_c; +drop table fk_omitted_uk_p; + +create table fk_temporal_default_p( + id int primary key, + d date unique default '2026-07-15', + t time unique default '12:34:56', + dt datetime unique default '2026-07-15 12:34:56', + ts timestamp unique default '2026-07-15 12:34:56' +); +create table fk_temporal_default_c( + id int primary key, + pid int, + foreign key(pid) references fk_temporal_default_p(id) on delete cascade +); +insert into fk_temporal_default_p(id) values (1); +insert into fk_temporal_default_c values (1, 1); +replace into fk_temporal_default_p(id) values (2); +select * from fk_temporal_default_p; +select * from fk_temporal_default_c; +drop table fk_temporal_default_c; +drop table fk_temporal_default_p; + +create table fk_decimal_cast_p(id int primary key, u decimal(5,2) unique); +create table fk_decimal_cast_c(id int primary key, pid int, + foreign key(pid) references fk_decimal_cast_p(id) on delete cascade); +insert into fk_decimal_cast_p values (1, 1.23); +insert into fk_decimal_cast_c values (1, 1); +replace into fk_decimal_cast_p values (2, 1.234); +select * from fk_decimal_cast_p; +select * from fk_decimal_cast_c; +drop table fk_decimal_cast_c; +drop table fk_decimal_cast_p; + +create table fk_param_p(id int primary key, v int); +create table fk_param_c(id int primary key, pid int, + foreign key(pid) references fk_param_p(id) on delete restrict); +prepare fk_param_key_stmt from 'replace into fk_param_p values (?, ?)'; +set @fk_id = 1, @fk_v = 300; +execute fk_param_key_stmt using @fk_id, @fk_v; +deallocate prepare fk_param_key_stmt; +insert into fk_param_p values (1, 100); +insert into fk_param_c values (1, 1); +prepare fk_review_stmt from 'replace into fk_param_p values (1, ?)'; +set @fk_v = 300; +execute fk_review_stmt using @fk_v; +select * from fk_param_p; +select * from fk_param_c; +deallocate prepare fk_review_stmt; +drop table fk_param_c; +drop table fk_param_p; + +drop table if exists fk_comp_c; +drop table if exists fk_comp_p; +create table fk_comp_p(a int, b int, primary key(a, b)); +create table fk_comp_c(id int primary key, a int, b int, + foreign key(a, b) references fk_comp_p(a, b) on delete cascade); +insert into fk_comp_p values (1, 1); +insert into fk_comp_c values (1, 1, 1); +replace into fk_comp_p values (1, 1); +select * from fk_comp_p; +select * from fk_comp_c; +drop table fk_comp_c; +drop table fk_comp_p; + +drop table if exists `fk``tick_c`; +drop table if exists `fk``tick_p`; +create table `fk``tick_p`(`id``x` int primary key); +create table `fk``tick_c`(id int primary key, `pid``x` int, + foreign key(`pid``x`) references `fk``tick_p`(`id``x`) on delete cascade); +insert into `fk``tick_p` values (1); +insert into `fk``tick_c` values (1, 1); +replace into `fk``tick_p` values (1); +select * from `fk``tick_c`; +drop table `fk``tick_c`; +drop table `fk``tick_p`; + -- child-side foreign key check during REPLACE: an inserted/replaced child row -- must reference an existing parent row, otherwise REPLACE fails and the -- previous conflicting row stays intact @@ -333,3 +576,34 @@ select a, b, c from t_replace_fakepk_comp_null order by a, b, c; replace into t_replace_fakepk_comp_null (a, b, c) values (1, 1, 'one-replaced'); select a, b, c from t_replace_fakepk_comp_null order by a, b, c; drop table t_replace_fakepk_comp_null; + +-- Self-referencing CASCADE uses the conflicting old row as the recursive root. +drop table if exists t_replace_self_cascade; +create table t_replace_self_cascade ( + id int primary key, + parent_id int, + foreign key (parent_id) references t_replace_self_cascade(id) on delete cascade +); +insert into t_replace_self_cascade values (1, null), (2, 1), (3, 2), (4, 1); +replace into t_replace_self_cascade values (1, null); +select * from t_replace_self_cascade order by id; +insert into t_replace_self_cascade values (2, 1), (3, 2); +replace into t_replace_self_cascade values (10, null); +select * from t_replace_self_cascade order by id; +delete from t_replace_self_cascade; +insert into t_replace_self_cascade values (1, null); +update t_replace_self_cascade set parent_id = 1 where id = 1; +replace into t_replace_self_cascade values (1, 1); +select row_count(); +select * from t_replace_self_cascade order by id; +delete from t_replace_self_cascade; +insert into t_replace_self_cascade values (1, null), (2, 1); +replace into t_replace_self_cascade values (1, null), (2, 1); +select row_count(); +select * from t_replace_self_cascade order by id; +delete from t_replace_self_cascade; +insert into t_replace_self_cascade values (1, null), (2, 1); +update t_replace_self_cascade set parent_id = 2 where id = 1; +replace into t_replace_self_cascade values (1, null); +select * from t_replace_self_cascade order by id; +drop table t_replace_self_cascade; diff --git a/test/distributed/cases/pessimistic_transaction/replace_fk_unique_lock.result b/test/distributed/cases/pessimistic_transaction/replace_fk_unique_lock.result new file mode 100644 index 0000000000000..bf4c0b14e2f7e --- /dev/null +++ b/test/distributed/cases/pessimistic_transaction/replace_fk_unique_lock.result @@ -0,0 +1,224 @@ +drop database if exists replace_fk_unique_lock; +create database replace_fk_unique_lock; +use replace_fk_unique_lock; +create table parent_restrict(id int primary key, u int unique); +create table child_restrict(id int primary key, parent_u int, +foreign key(parent_u) references parent_restrict(u) on delete restrict); +insert into parent_restrict values (1, 10); +begin; +replace into parent_restrict values (1, 20); +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +insert into child_restrict values (1, 10); +Lock wait timeout exceeded; try restarting transaction +rollback; +commit; +select * from parent_restrict; +➤ id[4,32,0] ¦ u[4,32,0] 𝄀 +1 ¦ 20 +select * from child_restrict; +➤ id[4,32,0] ¦ parent_u[4,32,0] +create table parent_cascade(id int primary key, u int unique); +create table child_cascade(id int primary key, parent_u int, +foreign key(parent_u) references parent_cascade(u) on delete cascade); +insert into parent_cascade values (1, 10); +insert into child_cascade values (1, 10); +begin; +replace into parent_cascade values (1, 20); +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +insert into child_cascade values (2, 10); +Lock wait timeout exceeded; try restarting transaction +rollback; +commit; +select * from parent_cascade; +➤ id[4,32,0] ¦ u[4,32,0] 𝄀 +1 ¦ 20 +select * from child_cascade; +➤ id[4,32,0] ¦ parent_u[4,32,0] +create table parent_decimal(id int primary key, u decimal(5,2) unique); +create table child_decimal(id int primary key, parent_u decimal(5,3), +foreign key(parent_u) references parent_decimal(u) on delete restrict); +insert into parent_decimal values (1, 1.23); +begin; +insert into child_decimal values (1, 1.230); +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +replace into parent_decimal values (1, 2.00); +Lock wait timeout exceeded; try restarting transaction +rollback; +rollback; +select * from parent_decimal; +➤ id[4,32,0] ¦ u[3,5,2] 𝄀 +1 ¦ 1.23 +select * from child_decimal; +➤ id[4,32,0] ¦ parent_u[3,5,3] +create table parent_generated( +id int primary key, +g int generated always as (id + 1), +u int unique +); +create table child_generated(id int primary key, parent_id int, +foreign key(parent_id) references parent_generated(id) on delete cascade); +insert into parent_generated(id, u) values (1, 10), (2, 20); +insert into child_generated values (1, 1), (2, 2); +replace into parent_generated values (3, 10); +replace into parent_generated(id, g, u) values (4, default, 20); +select id, g, u from parent_generated order by id; +➤ id[4,32,0] ¦ g[4,32,0] ¦ u[4,32,0] 𝄀 +3 ¦ 4 ¦ 10 𝄀 +4 ¦ 5 ¦ 20 +select * from child_generated; +➤ id[4,32,0] ¦ parent_id[4,32,0] +create table parent_generated_unique( +id int primary key, +a int, +g int generated always as (a + 1), +unique key uk_g(g) +); +create table child_generated_unique(id int primary key, parent_id int, +foreign key(parent_id) references parent_generated_unique(id) on delete cascade); +insert into parent_generated_unique(id, a) values (1, 10); +insert into child_generated_unique values (1, 1); +replace into parent_generated_unique(id, a) values (2, 10); +select id, a, g from parent_generated_unique; +➤ id[4,32,0] ¦ a[4,32,0] ¦ g[4,32,0] 𝄀 +2 ¦ 10 ¦ 11 +select * from child_generated_unique; +➤ id[4,32,0] ¦ parent_id[4,32,0] +create table parent_auto_zero(id int auto_increment primary key, v varchar(20)); +create table child_auto_zero(id int primary key, parent_id int, +foreign key(parent_id) references parent_auto_zero(id) on delete cascade); +set session sql_mode = 'NO_AUTO_VALUE_ON_ZERO'; +insert into parent_auto_zero values (0, 'old'); +insert into child_auto_zero values (1, 0); +set session sql_mode = ''; +replace into parent_auto_zero values (0, 'allocated'); +select * from parent_auto_zero order by id; +➤ id[4,32,0] ¦ v[12,-1,0] 𝄀 +0 ¦ old 𝄀 +1 ¦ allocated +select * from child_auto_zero; +➤ id[4,32,0] ¦ parent_id[4,32,0] 𝄀 +1 ¦ 0 +replace into parent_auto_zero values (0x0, 'allocated-hex'); +replace into parent_auto_zero values (b'0', 'allocated-bit'); +select * from parent_auto_zero order by id; +➤ id[4,32,0] ¦ v[12,-1,0] 𝄀 +0 ¦ old 𝄀 +1 ¦ allocated 𝄀 +2 ¦ allocated-hex 𝄀 +3 ¦ allocated-bit +select * from child_auto_zero; +➤ id[4,32,0] ¦ parent_id[4,32,0] 𝄀 +1 ¦ 0 +set session sql_mode = 'NO_AUTO_VALUE_ON_ZERO'; +replace into parent_auto_zero values (0, 'explicit-zero'); +select * from parent_auto_zero order by id; +➤ id[4,32,0] ¦ v[12,-1,0] 𝄀 +0 ¦ explicit-zero 𝄀 +1 ¦ allocated 𝄀 +2 ¦ allocated-hex 𝄀 +3 ¦ allocated-bit +select * from child_auto_zero; +➤ id[4,32,0] ¦ parent_id[4,32,0] +set session sql_mode = ''; +create table parent_cached(id int primary key, v int); +create table child_cached(id int primary key, parent_id int, +foreign key(parent_id) references parent_cached(id) on delete cascade); +insert into parent_cached values (1, 10); +insert into child_cached values (1, 1); +prepare replace_cached from 'replace into parent_cached values (1, 20)'; +set foreign_key_checks = 0; +execute replace_cached; +select * from child_cached; +➤ id[4,32,0] ¦ parent_id[4,32,0] 𝄀 +1 ¦ 1 +set foreign_key_checks = 1; +execute replace_cached; +select * from child_cached; +➤ id[4,32,0] ¦ parent_id[4,32,0] +deallocate prepare replace_cached; +insert into child_cached values (2, 1); +replace into parent_cached values (1, 30); +select * from child_cached; +➤ id[4,32,0] ¦ parent_id[4,32,0] +insert into child_cached values (3, 1); +set foreign_key_checks = 0; +replace into parent_cached values (1, 30); +select * from child_cached; +➤ id[4,32,0] ¦ parent_id[4,32,0] 𝄀 +3 ¦ 1 +set foreign_key_checks = 1; +replace into parent_cached values (1, 30); +select * from child_cached; +➤ id[4,32,0] ¦ parent_id[4,32,0] +create table parent_lock_order(id int primary key, a int unique, b int unique); +create table child_lock_order(id int primary key, parent_b int, parent_a int, +foreign key(parent_b) references parent_lock_order(b), +foreign key(parent_a) references parent_lock_order(a)); +insert into parent_lock_order values (1, 10, 20); +begin; +replace into parent_lock_order values (1, 11, 21); +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +insert into child_lock_order values (1, 20, 10); +Lock wait timeout exceeded; try restarting transaction +rollback; +commit; +select * from parent_lock_order; +➤ id[4,32,0] ¦ a[4,32,0] ¦ b[4,32,0] 𝄀 +1 ¦ 11 ¦ 21 +select * from child_lock_order; +➤ id[4,32,0] ¦ parent_b[4,32,0] ¦ parent_a[4,32,0] +create table parent_nullable_lock(id int primary key, u int unique); +create table child_nullable_lock(id int primary key, parent_u int, +foreign key(parent_u) references parent_nullable_lock(u)); +insert into parent_nullable_lock values (1, 10); +begin; +replace into parent_nullable_lock(id) values (1); +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +insert into child_nullable_lock values (1, 10); +Lock wait timeout exceeded; try restarting transaction +rollback; +commit; +insert into parent_nullable_lock values (2, 10); +select * from parent_nullable_lock order by id; +➤ id[4,32,0] ¦ u[4,32,0] 𝄀 +1 ¦ null 𝄀 +2 ¦ 10 +select * from child_nullable_lock; +➤ id[4,32,0] ¦ parent_u[4,32,0] +create table parent_dynamic(id int primary key, u int unique); +create table child_dynamic(id int primary key, parent_id int, +foreign key(parent_id) references parent_dynamic(id) on delete cascade); +create table replace_source(id int, u int); +insert into parent_dynamic values (1, 10), (2, 20); +insert into child_dynamic values (1, 1), (2, 2); +prepare replace_dynamic from 'replace into parent_dynamic values (?, ?)'; +set @replace_id = 1; +set @replace_u = 11; +execute replace_dynamic using @replace_id, @replace_u; +select * from parent_dynamic order by id; +➤ id[4,32,0] ¦ u[4,32,0] 𝄀 +1 ¦ 11 𝄀 +2 ¦ 20 +select * from child_dynamic order by id; +➤ id[4,32,0] ¦ parent_id[4,32,0] 𝄀 +2 ¦ 2 +deallocate prepare replace_dynamic; +insert into replace_source values (2, 22); +replace into parent_dynamic select id, u from replace_source; +select * from parent_dynamic order by id; +➤ id[4,32,0] ¦ u[4,32,0] 𝄀 +1 ¦ 11 𝄀 +2 ¦ 22 +select * from child_dynamic order by id; +➤ id[4,32,0] ¦ parent_id[4,32,0] +drop database replace_fk_unique_lock; diff --git a/test/distributed/cases/pessimistic_transaction/replace_fk_unique_lock.sql b/test/distributed/cases/pessimistic_transaction/replace_fk_unique_lock.sql new file mode 100644 index 0000000000000..312693671bb98 --- /dev/null +++ b/test/distributed/cases/pessimistic_transaction/replace_fk_unique_lock.sql @@ -0,0 +1,194 @@ +-- @suite + +-- @case +-- @desc:test REPLACE and child FK validation contend on non-PK unique keys +-- @label:bvt +drop database if exists replace_fk_unique_lock; +create database replace_fk_unique_lock; +use replace_fk_unique_lock; + +create table parent_restrict(id int primary key, u int unique); +create table child_restrict(id int primary key, parent_u int, + foreign key(parent_u) references parent_restrict(u) on delete restrict); +insert into parent_restrict values (1, 10); +begin; +replace into parent_restrict values (1, 20); +-- @session:id=1{ +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +-- @pattern +insert into child_restrict values (1, 10); +rollback; +-- @session} +commit; +select * from parent_restrict; +select * from child_restrict; + +create table parent_cascade(id int primary key, u int unique); +create table child_cascade(id int primary key, parent_u int, + foreign key(parent_u) references parent_cascade(u) on delete cascade); +insert into parent_cascade values (1, 10); +insert into child_cascade values (1, 10); +begin; +replace into parent_cascade values (1, 20); +-- @session:id=1{ +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +-- @pattern +insert into child_cascade values (2, 10); +rollback; +-- @session} +commit; +select * from parent_cascade; +select * from child_cascade; + +create table parent_decimal(id int primary key, u decimal(5,2) unique); +create table child_decimal(id int primary key, parent_u decimal(5,3), + foreign key(parent_u) references parent_decimal(u) on delete restrict); +insert into parent_decimal values (1, 1.23); +begin; +insert into child_decimal values (1, 1.230); +-- @session:id=1{ +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +-- @pattern +replace into parent_decimal values (1, 2.00); +rollback; +-- @session} +rollback; +select * from parent_decimal; +select * from child_decimal; + +create table parent_generated( + id int primary key, + g int generated always as (id + 1), + u int unique +); +create table child_generated(id int primary key, parent_id int, + foreign key(parent_id) references parent_generated(id) on delete cascade); +insert into parent_generated(id, u) values (1, 10), (2, 20); +insert into child_generated values (1, 1), (2, 2); +replace into parent_generated values (3, 10); +replace into parent_generated(id, g, u) values (4, default, 20); +select id, g, u from parent_generated order by id; +select * from child_generated; + +create table parent_generated_unique( + id int primary key, + a int, + g int generated always as (a + 1), + unique key uk_g(g) +); +create table child_generated_unique(id int primary key, parent_id int, + foreign key(parent_id) references parent_generated_unique(id) on delete cascade); +insert into parent_generated_unique(id, a) values (1, 10); +insert into child_generated_unique values (1, 1); +replace into parent_generated_unique(id, a) values (2, 10); +select id, a, g from parent_generated_unique; +select * from child_generated_unique; + +create table parent_auto_zero(id int auto_increment primary key, v varchar(20)); +create table child_auto_zero(id int primary key, parent_id int, + foreign key(parent_id) references parent_auto_zero(id) on delete cascade); +set session sql_mode = 'NO_AUTO_VALUE_ON_ZERO'; +insert into parent_auto_zero values (0, 'old'); +insert into child_auto_zero values (1, 0); +set session sql_mode = ''; +replace into parent_auto_zero values (0, 'allocated'); +select * from parent_auto_zero order by id; +select * from child_auto_zero; +replace into parent_auto_zero values (0x0, 'allocated-hex'); +replace into parent_auto_zero values (b'0', 'allocated-bit'); +select * from parent_auto_zero order by id; +select * from child_auto_zero; +set session sql_mode = 'NO_AUTO_VALUE_ON_ZERO'; +replace into parent_auto_zero values (0, 'explicit-zero'); +select * from parent_auto_zero order by id; +select * from child_auto_zero; +set session sql_mode = ''; + +create table parent_cached(id int primary key, v int); +create table child_cached(id int primary key, parent_id int, + foreign key(parent_id) references parent_cached(id) on delete cascade); +insert into parent_cached values (1, 10); +insert into child_cached values (1, 1); +prepare replace_cached from 'replace into parent_cached values (1, 20)'; +set foreign_key_checks = 0; +execute replace_cached; +select * from child_cached; +set foreign_key_checks = 1; +execute replace_cached; +select * from child_cached; +deallocate prepare replace_cached; + +insert into child_cached values (2, 1); +replace into parent_cached values (1, 30); +select * from child_cached; +insert into child_cached values (3, 1); +set foreign_key_checks = 0; +replace into parent_cached values (1, 30); +select * from child_cached; +set foreign_key_checks = 1; +replace into parent_cached values (1, 30); +select * from child_cached; + +create table parent_lock_order(id int primary key, a int unique, b int unique); +create table child_lock_order(id int primary key, parent_b int, parent_a int, + foreign key(parent_b) references parent_lock_order(b), + foreign key(parent_a) references parent_lock_order(a)); +insert into parent_lock_order values (1, 10, 20); +begin; +replace into parent_lock_order values (1, 11, 21); +-- @session:id=1{ +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +-- @pattern +insert into child_lock_order values (1, 20, 10); +rollback; +-- @session} +commit; +select * from parent_lock_order; +select * from child_lock_order; + +create table parent_nullable_lock(id int primary key, u int unique); +create table child_nullable_lock(id int primary key, parent_u int, + foreign key(parent_u) references parent_nullable_lock(u)); +insert into parent_nullable_lock values (1, 10); +begin; +replace into parent_nullable_lock(id) values (1); +-- @session:id=1{ +use replace_fk_unique_lock; +set session lock_wait_timeout = 1; +begin; +-- @pattern +insert into child_nullable_lock values (1, 10); +rollback; +-- @session} +commit; +insert into parent_nullable_lock values (2, 10); +select * from parent_nullable_lock order by id; +select * from child_nullable_lock; + +create table parent_dynamic(id int primary key, u int unique); +create table child_dynamic(id int primary key, parent_id int, + foreign key(parent_id) references parent_dynamic(id) on delete cascade); +create table replace_source(id int, u int); +insert into parent_dynamic values (1, 10), (2, 20); +insert into child_dynamic values (1, 1), (2, 2); +prepare replace_dynamic from 'replace into parent_dynamic values (?, ?)'; +set @replace_id = 1; +set @replace_u = 11; +execute replace_dynamic using @replace_id, @replace_u; +select * from parent_dynamic order by id; +select * from child_dynamic order by id; +deallocate prepare replace_dynamic; +insert into replace_source values (2, 22); +replace into parent_dynamic select id, u from replace_source; +select * from parent_dynamic order by id; +select * from child_dynamic order by id; + +drop database replace_fk_unique_lock;