Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
168 changes: 140 additions & 28 deletions pkg/frontend/data_branch_hashdiff.go
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,30 @@ type lcaProbeLayout struct {
enumValues []string
}

func sortDataBranchBatchByPrimaryKey(bat *batch.Batch, pkColIdx int, mp *mpool.MPool) error {
pkVec := bat.Vecs[pkColIdx]
if !isDataBranchFloatType(*pkVec.GetType()) {
return mergeutil.SortColumnsByIndex(bat.Vecs, pkColIdx, mp)
}

identityVec := vector.NewVec(types.T_uint64.ToType())
defer identityVec.Free(mp)
for row := range pkVec.Length() {
identity, isNull, err := dataBranchFloatPKIdentityAt(pkVec, row)
if err != nil {
return err
}
if err = vector.AppendFixed(identityVec, identity, isNull, mp); err != nil {
return err
}
}

cols := make([]*vector.Vector, 0, len(bat.Vecs)+1)
cols = append(cols, bat.Vecs...)
cols = append(cols, identityVec)
return mergeutil.SortColumnsByIndex(cols, len(cols)-1, mp)
}

func (layout lcaProbeLayout) columnNameForTargetIndex(targetIdx int) (string, bool) {
for i, candidateIdx := range layout.targetIdxes {
if candidateIdx == targetIdx {
Expand Down Expand Up @@ -289,8 +313,8 @@ func handleDelsOnLCA(

valsBuf.WriteString(fmt.Sprintf("row(%d,", i))
for j := range tuple {
if err = formatValIntoString(
ses, tuple[j], colTypes[expandedPKColIdxes[j]], valsBuf,
if err = formatValIntoStringWithFloatCast(
ses, tuple[j], colTypes[expandedPKColIdxes[j]], valsBuf, true,
); err != nil {
return nil, err
}
Expand Down Expand Up @@ -320,7 +344,7 @@ func handleDelsOnLCA(
valsBuf.WriteString(fmt.Sprintf("row(%d,", i))
b := tBat.Vecs[0].GetRawBytesAt(i)
val := types.DecodeValue(b, tBat.Vecs[0].GetType().Oid)
if err = formatValIntoString(ses, val, pkType, valsBuf); err != nil {
if err = formatValIntoStringWithFloatCast(ses, val, pkType, valsBuf, true); err != nil {
return nil, err
}
valsBuf.WriteString(")")
Expand Down Expand Up @@ -357,12 +381,14 @@ func handleDelsOnLCA(
)

for i := range quotedPKNames {
sqlBuf.WriteString(fmt.Sprintf("lca.%s = ", quotedPKNames[i]))
left := fmt.Sprintf("lca.%s", quotedPKNames[i])
right := fmt.Sprintf("pks.%s", quotedPKValueAliases[i])
if castType, ok := lcaProbeJoinCastType(colTypes[expandedPKColIdxes[i]]); ok {
sqlBuf.WriteString(fmt.Sprintf("cast(pks.%s as %s)", quotedPKValueAliases[i], castType))
} else {
sqlBuf.WriteString(fmt.Sprintf("pks.%s", quotedPKValueAliases[i]))
right = fmt.Sprintf("cast(%s as %s)", right, castType)
}
sqlBuf.WriteString(dataBranchSQLKeyEqual(
left, right, colTypes[expandedPKColIdxes[i]],
))
if i != len(quotedPKNames)-1 {
sqlBuf.WriteString(" AND ")
}
Expand Down Expand Up @@ -1053,14 +1079,15 @@ func hashDiffIfHasLCA(
wg sync.WaitGroup
atomicErr atomic.Value

baseDeleteBatches []batchWithKind
baseUpdateBatches []batchWithKind
baseDeleteBatches []batchWithKind
baseUpdateBatches []batchWithKind
restoreMissingKeys = make(map[string]struct{})
)

handleBaseDeleteAndUpdates := func(wrapped batchWithKind) error {
wrapped.side = diffSideBase
if err2 := mergeutil.SortColumnsByIndex(
wrapped.batch.Vecs, tblStuff.def.pkColIdx, ses.proc.Mp(),
if err2 := sortDataBranchBatchByPrimaryKey(
wrapped.batch, tblStuff.def.pkColIdx, ses.proc.Mp(),
); err2 != nil {
return err2
}
Expand All @@ -1076,6 +1103,50 @@ func hashDiffIfHasLCA(
handleTarDeleteAndUpdates := func(wrapped batchWithKind) (err2 error) {
wrapped.side = diffSideTarget
var pickConflictBat *batch.Batch
if wrapped.kind == diffInsert && wrapped.fromUpdate && len(restoreMissingKeys) > 0 {
var keep []int64
restoreBat := tblStuff.retPool.acquireRetBatch(tblStuff, false)
for rowIdx := range wrapped.batch.RowCount() {
key, keyErr := extractPKAsString(ses, tblStuff, wrapped.batch, rowIdx)
if keyErr != nil {
tblStuff.retPool.releaseRetBatch(restoreBat, false)
return keyErr
}
if _, restore := restoreMissingKeys[key]; !restore {
keep = append(keep, int64(rowIdx))
continue
}
if err2 = restoreBat.UnionOne(wrapped.batch, int64(rowIdx), ses.proc.Mp()); err2 != nil {
tblStuff.retPool.releaseRetBatch(restoreBat, false)
return err2
}
delete(restoreMissingKeys, key)
}
if restoreBat.Vecs[0].Length() > 0 {
restoreBat.SetRowCount(restoreBat.Vecs[0].Length())
if stop, e := emitBatch(emit, batchWithKind{
batch: restoreBat,
kind: diffInsert,
name: wrapped.name,
side: wrapped.side,
fromUpdate: true,
restoreMissing: true,
}, false, tblStuff.retPool); e != nil {
tblStuff.retPool.releaseRetBatch(wrapped.batch, false)
return e
} else if stop {
tblStuff.retPool.releaseRetBatch(wrapped.batch, false)
return nil
}
} else {
tblStuff.retPool.releaseRetBatch(restoreBat, false)
}
if len(keep) == 0 {
tblStuff.retPool.releaseRetBatch(wrapped.batch, false)
return nil
}
wrapped.batch.Shrink(keep, true)
}
if len(baseUpdateBatches) == 0 && len(baseDeleteBatches) == 0 {
// no need to check conflict
if stop, e := emitBatch(emit, wrapped, false, tblStuff.retPool); e != nil {
Expand All @@ -1086,8 +1157,8 @@ func hashDiffIfHasLCA(
return nil
}

if err2 = mergeutil.SortColumnsByIndex(
wrapped.batch.Vecs, tblStuff.def.pkColIdx, ses.proc.Mp(),
if err2 = sortDataBranchBatchByPrimaryKey(
wrapped.batch, tblStuff.def.pkColIdx, ses.proc.Mp(),
); err2 != nil {
return err2
}
Expand All @@ -1112,7 +1183,7 @@ func hashDiffIfHasLCA(

i, j := 0, 0
for i < tarVec.Length() && j < baseVec.Length() {
if cmp, err3 = compareSingleValInVector(
if cmp, err3 = compareDataBranchPrimaryKeyInVectors(
ctx, ses, i, j, tarVec, baseVec,
); err3 != nil {
return
Expand All @@ -1139,6 +1210,15 @@ func hashDiffIfHasLCA(
i++
j++
} else if copt.conflictOpt.Opt == tree.CONFLICT_ACCEPT {
if tarWrapped.kind == diffDelete && tarWrapped.fromUpdate &&
baseWrapped.kind == diffDelete && !baseWrapped.fromUpdate {
key, keyErr := extractPKAsString(ses, tblStuff, tarWrapped.batch, i)
if keyErr != nil {
err3 = keyErr
return
}
restoreMissingKeys[key] = struct{}{}
}
if tarWrapped.kind == diffDelete &&
baseWrapped.kind == diffDelete &&
!tarWrapped.fromUpdate && !baseWrapped.fromUpdate {
Expand Down Expand Up @@ -1254,11 +1334,6 @@ func hashDiffIfHasLCA(
return false
})

if wrapped.batch.RowCount() == 0 {
tblStuff.retPool.releaseRetBatch(wrapped.batch, false)
return
}

if pickConflictBat != nil {
if stop, e := emitBatch(emit, batchWithKind{
batch: pickConflictBat,
Expand All @@ -1271,6 +1346,10 @@ func hashDiffIfHasLCA(
return nil
}
}
if wrapped.batch.RowCount() == 0 {
tblStuff.retPool.releaseRetBatch(wrapped.batch, false)
return nil
}

stop, e := emitBatch(emit, wrapped, false, tblStuff.retPool)
if e != nil {
Expand Down Expand Up @@ -1390,6 +1469,12 @@ func hashDiffIfHasLCA(
if err = stepHandler(false); err != nil {
return
}
if len(restoreMissingKeys) != 0 {
return moerr.NewInternalErrorNoCtxf(
"data branch source update is missing %d replacement row(s)",
len(restoreMissingKeys),
)
}

// what can I do with these left base updates/inserts ?
if copt.conflictOpt == nil {
Expand Down Expand Up @@ -1721,10 +1806,11 @@ func findDeleteAndUpdateBat(
return err2
}
if err2 = send(batchWithKind{
name: tblName,
side: side,
batch: updateBat,
kind: diffInsert,
name: tblName,
side: side,
batch: updateBat,
kind: diffInsert,
fromUpdate: tblStuff.def.pkKind != fakeKind,
}); err2 != nil {
return err2
}
Expand Down Expand Up @@ -2092,13 +2178,19 @@ func diffDataHelper(
tarBat *batch.Batch
baseBat *batch.Batch
baseDeleteBat *batch.Batch
tarUpdateBat *batch.Batch
tarTuple types.Tuple
baseTuple types.Tuple
checkRet databranchutils.GetResult
)

tarBat = tblStuff.retPool.acquireRetBatch(tblStuff, false)
baseBat = tblStuff.retPool.acquireRetBatch(tblStuff, false)
defer func() {
if tarUpdateBat != nil {
tblStuff.retPool.releaseRetBatch(tarUpdateBat, false)
}
}()

if err2 = cursor.ForEach(func(key []byte, row []byte) error {
select {
Expand Down Expand Up @@ -2174,10 +2266,13 @@ func diffDataHelper(
if baseDeleteBat == nil {
baseDeleteBat = tblStuff.retPool.acquireRetBatch(tblStuff, false)
}
if tarUpdateBat == nil {
tarUpdateBat = tblStuff.retPool.acquireRetBatch(tblStuff, false)
}
if err2 = appendTupleToBat(ses, baseDeleteBat, baseTuple, tblStuff); err2 != nil {
return err2
}
if err2 = appendTupleToBat(ses, tarBat, tarTuple, tblStuff); err2 != nil {
if err2 = appendTupleToBat(ses, tarUpdateBat, tarTuple, tblStuff); err2 != nil {
return err2
}
} else {
Expand All @@ -2204,17 +2299,34 @@ func diffDataHelper(

if baseDeleteBat != nil {
if stop, err3 := emitBatch(emit, batchWithKind{
batch: baseDeleteBat,
kind: diffDelete,
name: tblStuff.baseRel.GetTableName(),
side: diffSideBase,
batch: baseDeleteBat,
kind: diffDelete,
name: tblStuff.baseRel.GetTableName(),
side: diffSideBase,
fromUpdate: true,
}, false, tblStuff.retPool); err3 != nil {
return err3
} else if stop {
return nil
}
}

if tarUpdateBat != nil {
stop, err3 := emitBatch(emit, batchWithKind{
batch: tarUpdateBat,
kind: diffInsert,
name: tblStuff.tarRel.GetTableName(),
side: diffSideTarget,
fromUpdate: true,
}, false, tblStuff.retPool)
tarUpdateBat = nil
if err3 != nil {
return err3
} else if stop {
return nil
}
}

if stop, err3 := emitBatch(emit, batchWithKind{
batch: tarBat,
kind: diffInsert,
Expand Down
56 changes: 52 additions & 4 deletions pkg/frontend/data_branch_hashdiff_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ package frontend
import (
"context"
"fmt"
"math"
"strings"
"sync"
"sync/atomic"
Expand Down Expand Up @@ -908,9 +909,10 @@ func (h *closeTrackingBranchHashmap) Close() error {
}

type capturedBatch struct {
kind string
side int
rows [][]any
kind string
side int
rows [][]any
fromUpdate bool
}

func TestRunLCAProbeWithReaderFallback_EarlyReturns(t *testing.T) {
Expand Down Expand Up @@ -1514,6 +1516,48 @@ func TestHandleDelsOnLCA_SQLPaths(t *testing.T) {
require.ErrorIs(t, err, wantErr)
})

t.Run("float primary key join matches NaN explicitly", func(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()

tblStuff := newTestBranchTableStuff(ctrl)
tblStuff.lcaRel = mock_frontend.NewMockRelation(ctrl)
tblStuff.def.colTypes[0] = types.T_float64.ToType()
targetDef := tblStuff.tarRel.GetTableDef(context.Background())
targetDef.Cols[0].Typ = plan.Type{Id: int32(types.T_float64)}
baseDef := tblStuff.baseRel.GetTableDef(context.Background())
baseDef.Cols[0].Typ = plan.Type{Id: int32(types.T_float64)}
lcaDef := newTestBranchTableDef("lca_tbl", "name")
lcaDef.Cols[0].Typ = plan.Type{Id: int32(types.T_float64)}
tblStuff.lcaRel.(*mock_frontend.MockRelation).EXPECT().GetTableDef(gomock.Any()).Return(lcaDef).AnyTimes()
tblStuff.lcaRel.(*mock_frontend.MockRelation).EXPECT().GetTableID(gomock.Any()).Return(uint64(76)).AnyTimes()

wantErr := moerr.NewInternalErrorNoCtx("stop after sql capture")
bh := mock_frontend.NewMockBackgroundExec(ctrl)
bh.EXPECT().Exec(gomock.Any(), gomock.Any()).
DoAndReturn(func(_ context.Context, sql string) error {
require.Contains(t, sql,
"values row(0,bit_cast(unhex('010000000000f87f') as double)),row(1,cast(1.25 as double))")
right := "cast(pks.`__mo_data_branch_pk_0` as DOUBLE)"
require.Contains(t, sql, dataBranchSQLKeyEqual("lca.`id`", right, types.T_float64.ToType()))
return wantErr
}).
Times(1)

tBat := batch.NewWithSize(1)
tBat.Vecs[0] = vector.NewVec(types.T_float64.ToType())
require.NoError(t, vector.AppendFixed(tBat.Vecs[0], math.NaN(), false, ses.proc.Mp()))
require.NoError(t, vector.AppendFixed(tBat.Vecs[0], 1.25, false, ses.proc.Mp()))
tBat.SetRowCount(2)
defer tBat.Clean(ses.proc.Mp())

_, err := handleDelsOnLCA(
context.Background(), ses, bh, tBat, tblStuff,
types.BuildTS(10, 0).ToTimestamp(),
)
require.ErrorIs(t, err, wantErr)
})

t.Run("internal aliases do not collide with user primary key", func(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
Expand Down Expand Up @@ -1984,7 +2028,9 @@ func TestHashDiff_NoLCABoundedUpdateKeepsLatestRow(t *testing.T) {
rows := decodeCapturedRows(t, w.batch, tblStuff.def.colTypes)
mu.Lock()
if len(rows) > 0 {
got = append(got, capturedBatch{kind: w.kind, side: w.side, rows: rows})
got = append(got, capturedBatch{
kind: w.kind, side: w.side, rows: rows, fromUpdate: w.fromUpdate,
})
}
mu.Unlock()
tblStuff.retPool.releaseRetBatch(w.batch, false)
Expand All @@ -2002,9 +2048,11 @@ func TestHashDiff_NoLCABoundedUpdateKeepsLatestRow(t *testing.T) {
require.Len(t, got, 2)
require.Equal(t, diffDelete, got[0].kind)
require.Equal(t, diffSideBase, got[0].side)
require.True(t, got[0].fromUpdate)
require.Equal(t, [][]any{{int64(1), "destination", "h1"}}, got[0].rows)
require.Equal(t, diffInsert, got[1].kind)
require.Equal(t, diffSideTarget, got[1].side)
require.True(t, got[1].fromUpdate)
require.Equal(t, [][]any{{int64(1), "bounded", "h1"}}, got[1].rows)
}

Expand Down
Loading
Loading