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
14 changes: 13 additions & 1 deletion pkg/common/hashmap/iterator.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,19 @@ func validateIteratorVectors(
return mpool.ErrAllocationAccountInvalid
}
for _, vec := range vecs {
if vec == nil || start > vec.Length() || count > vec.Length()-start {
if vec == nil {
return mpool.ErrAllocationAccountInvalid
}
// Const vectors physically store one value and broadcast it across the
// caller's logical row range. Hash encoders already handle that contract
// by reading row zero, so only require physical storage when rows are read.
if vec.IsConst() {
if count > 0 && vec.Length() == 0 {
return mpool.ErrAllocationAccountInvalid
}
continue
}
if start > vec.Length() || count > vec.Length()-start {
return mpool.ErrAllocationAccountInvalid
}
}
Expand Down
54 changes: 51 additions & 3 deletions pkg/common/hashmap/strhashmap_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -854,14 +854,14 @@ func TestHashMapIteratorsRejectMalformedRowShapes(t *testing.T) {
},
} {
m, iterator := makeIterator()
short, err := vector.NewConstFixed(types.T_int32.ToType(), int32(1), 1, mp)
require.NoError(t, err)
short := vector.NewVec(types.T_int32.ToType())
require.NoError(t, vector.AppendFixed(short, int32(1), false, mp))
for _, vecs := range [][]*vector.Vector{
nil,
{nil},
{short},
} {
_, _, err = iterator.Insert(0, 2, vecs)
_, _, err := iterator.Insert(0, 2, vecs)
require.ErrorIs(t, err, mpool.ErrAllocationAccountInvalid)
_, _, err = iterator.Find(1, 1, vecs)
require.ErrorIs(t, err, mpool.ErrAllocationAccountInvalid)
Expand All @@ -873,6 +873,54 @@ func TestHashMapIteratorsRejectMalformedRowShapes(t *testing.T) {
require.Zero(t, mp.CurrNB())
}

func TestHashMapIteratorsBroadcastConstVectors(t *testing.T) {
mp := mpool.MustNewZero()
for _, makeIterator := range []func() (HashMap, Iterator){
func() (HashMap, Iterator) {
m, err := NewIntHashMap(true, mp)
require.NoError(t, err)
return m, m.NewIterator()
},
func() (HashMap, Iterator) {
m, err := NewStrHashMap(true, mp)
require.NoError(t, err)
return m, m.NewIterator()
},
} {
m, iterator := makeIterator()
constant, err := vector.NewConstFixed(types.T_int32.ToType(), int32(7), 1, mp)
require.NoError(t, err)

values, zValues, err := iterator.Insert(UnitLimit, 2, []*vector.Vector{constant})
require.NoError(t, err)
require.Equal(t, []uint64{1, 1}, values)
require.Equal(t, []int64{1, 1}, zValues)
require.Equal(t, uint64(1), m.GroupCount())

values, zValues, err = iterator.Find(UnitLimit*2, 2, []*vector.Vector{constant})
require.NoError(t, err)
require.Equal(t, []uint64{1, 1}, values)
require.Equal(t, []int64{1, 1}, zValues)

constantNull := vector.NewConstNull(types.T_int32.ToType(), 1, mp)
values, zValues, err = iterator.Insert(UnitLimit*3, 2, []*vector.Vector{constantNull})
require.NoError(t, err)
require.Equal(t, []uint64{2, 2}, values)
require.Equal(t, []int64{1, 1}, zValues)
require.Equal(t, uint64(2), m.GroupCount())

emptyConstant := vector.NewConstNull(types.T_int32.ToType(), 0, mp)
_, _, err = iterator.Insert(UnitLimit*4, 1, []*vector.Vector{emptyConstant})
require.ErrorIs(t, err, mpool.ErrAllocationAccountInvalid)

emptyConstant.Free(mp)
constantNull.Free(mp)
constant.Free(mp)
m.Free()
}
require.Zero(t, mp.CurrNB())
}

func runFloatHashMapShape(
t *testing.T,
composite bool,
Expand Down
4 changes: 2 additions & 2 deletions pkg/frontend/mysql_cmd_executor.go
Original file line number Diff line number Diff line change
Expand Up @@ -762,8 +762,8 @@ func getDataFromPipeline(obj FeSession, execCtx *ExecCtx, bat *batch.Batch, crs
}
tTime := time.Since(begin)
n := 0
if !isPerformStatement(execCtx.stmt) && bat != nil && bat.Vecs[0] != nil {
n = bat.Vecs[0].Length()
if !isPerformStatement(execCtx.stmt) && bat != nil {
n = bat.RowCount()
ses.sentRows.Add(int64(n))
}

Expand Down
2 changes: 1 addition & 1 deletion pkg/frontend/mysql_protocol.go
Original file line number Diff line number Diff line change
Expand Up @@ -367,7 +367,7 @@ func (mp *MysqlProtocolImpl) GetBool(id PropertyID) bool {
}

func (mp *MysqlProtocolImpl) Write(execCtx *ExecCtx, crs *perfcounter.CounterSet, bat *batch.Batch) error {
n := bat.Vecs[0].Length()
n := bat.RowCount()
//TODO: remove this MRS here
//Create a new temporary result set per pipeline thread.
mrs := MysqlResultSet{}
Expand Down
107 changes: 107 additions & 0 deletions pkg/frontend/mysql_protocol_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3200,6 +3200,113 @@ func Test_resultset(t *testing.T) {
})
}

func TestMysqlProtocolWriteUsesLogicalBatchRowCount(t *testing.T) {
for _, test := range []struct {
name string
cmd CommandType
}{
{name: "text protocol", cmd: COM_QUERY},
{name: "binary protocol", cmd: COM_STMT_EXECUTE},
} {
t.Run(test.name, func(t *testing.T) {
sv, err := getSystemVariables("test/system_vars_config.toml")
require.NoError(t, err)
pu := config.NewParameterUnit(sv, nil, nil, nil)
pu.SV.SkipCheckUser = true
pu.SV.KillRountinesInterval = 0
setSessionAlloc("", NewLeakCheckAllocator())
setPu("", pu)

ioses, err := NewIOSession(&testConn{}, pu, "")
require.NoError(t, err)
proto := NewMysqlClientProtocol("", 0, ioses, 1024, sv)
t.Cleanup(proto.Close)

ses := NewSession(context.Background(), "", proto, nil)
ses.SetCmd(test.cmd)
proto.ses = ses
mrs := &MysqlResultSet{}
for _, name := range []string{"projection_value", "total"} {
column := &MysqlColumn{}
column.SetName(name)
column.SetColumnType(defines.MYSQL_TYPE_LONGLONG)
mrs.AddColumn(column)
}
ses.SetMysqlResultSet(mrs)

mp := mpool.MustNewZero()
constant, err := vector.NewConstFixed(
types.T_int64.ToType(), int64(7), 1, mp,
)
require.NoError(t, err)
totals := vector.NewVec(types.T_int64.ToType())
for _, value := range []int64{10, 20, 30} {
require.NoError(t, vector.AppendFixed(totals, value, false, mp))
}
bat := batch.NewWithSize(2)
bat.SetVector(0, constant)
bat.SetVector(1, totals)
bat.SetRowCount(3)
t.Cleanup(func() {
bat.Clean(mp)
require.Zero(t, mp.CurrNB())
})

execCtx := &ExecCtx{
reqCtx: context.Background(),
ses: ses,
}
err = getDataFromPipeline(ses, execCtx, bat, nil)
require.NoError(t, err)
require.Equal(t, bat.RowCount(), proto.tcpConn.packetInBuf)
require.Equal(t, int64(bat.RowCount()), ses.sentRows.Load())
})
}
}

func TestMysqlProtocolWriteSendsLogicalConstNullRow(t *testing.T) {
sv, err := getSystemVariables("test/system_vars_config.toml")
require.NoError(t, err)
pu := config.NewParameterUnit(sv, nil, nil, nil)
pu.SV.SkipCheckUser = true
pu.SV.KillRountinesInterval = 0
setSessionAlloc("", NewLeakCheckAllocator())
setPu("", pu)

ioses, err := NewIOSession(&testConn{}, pu, "")
require.NoError(t, err)
proto := NewMysqlClientProtocol("", 0, ioses, 1024, sv)
t.Cleanup(proto.Close)

ses := NewSession(context.Background(), "", proto, nil)
ses.SetCmd(COM_QUERY)
proto.ses = ses
mrs := &MysqlResultSet{}
column := &MysqlColumn{}
column.SetName("null_input")
column.SetColumnType(defines.MYSQL_TYPE_VARCHAR)
mrs.AddColumn(column)
ses.SetMysqlResultSet(mrs)

mp := mpool.MustNewZero()
bat := batch.NewWithSize(1)
bat.SetVector(0, vector.NewConstNull(types.T_varchar.ToType(), 0, mp))
bat.SetRowCount(1)
t.Cleanup(func() {
bat.Clean(mp)
require.Zero(t, mp.CurrNB())
})

execCtx := &ExecCtx{
reqCtx: context.Background(),
ses: ses,
}
err = getDataFromPipeline(ses, execCtx, bat, nil)
require.NoError(t, err)
require.Equal(t, 1, proto.tcpConn.packetInBuf)
require.Equal(t, int64(1), ses.sentRows.Load())
}

func TestSendResultSetPropagatesRowEncodingErrors(t *testing.T) {
sv, err := getSystemVariables("test/system_vars_config.toml")
require.NoError(t, err)
Expand Down
107 changes: 101 additions & 6 deletions pkg/frontend/query_result.go
Original file line number Diff line number Diff line change
Expand Up @@ -110,13 +110,18 @@ func initQueryResulConfig(ctx context.Context, ses *Session) error {
}

func saveBatch(ctx context.Context, ses *Session, bat *batch.Batch) error {
n := uint64(0)
if bat != nil && bat.Vecs[0] != nil {
n = uint64(bat.Vecs[0].Length())
}
n := uint64(bat.RowCount())
ses.queryRowCount += n

s := ses.curResultSize + float64(bat.Size())/(1024*1024)
writeBat, release, err := prepareQueryResultBatchForWrite(bat, ses.GetMemPool())
if err != nil {
return err
}
if release != nil {
defer release()
}

s := ses.curResultSize + float64(writeBat.Size())/(1024*1024)
if s > ses.limitResultSize {
ses.Debug(ctx, "open save query result", zap.Float64("current result size:", s))
return nil
Expand All @@ -129,7 +134,7 @@ func saveBatch(ctx context.Context, ses *Session, bat *batch.Batch) error {
if err != nil {
return err
}
_, err = writer.Write(bat)
_, err = writer.Write(writeBat)
if err != nil {
return err
}
Expand All @@ -146,6 +151,96 @@ func saveBatch(ctx context.Context, ses *Session, bat *batch.Batch) error {
return err
}

func prepareQueryResultBatchForWrite(
bat *batch.Batch,
mp *mpool.MPool,
) (*batch.Batch, func(), error) {
rowCount := bat.RowCount()
for i, vec := range bat.Vecs {
if vec == nil {
return nil, nil, moerr.NewInternalErrorNoCtxf(
"query result column %d is nil", i,
)
}
if vec.Length() != rowCount {
return normalizeQueryResultBatchForWrite(bat, mp, rowCount)
}
}
return bat, nil, nil
}

func normalizeQueryResultBatchForWrite(
bat *batch.Batch,
mp *mpool.MPool,
rowCount int,
) (*batch.Batch, func(), error) {
writeBat := batch.NewWithSize(len(bat.Vecs))
writeBat.Attrs = bat.Attrs
copy(writeBat.Vecs, bat.Vecs)
writeBat.SetRowCount(rowCount)
cloned := make([]*vector.Vector, 0, len(bat.Vecs))

for i, vec := range bat.Vecs {
if vec == nil {
freeQueryResultVectors(cloned, mp)
return nil, nil, moerr.NewInternalErrorNoCtxf(
"query result column %d is nil", i,
)
}
if vec.Length() == rowCount {
continue
}
if !vec.IsConst() && vec.Length() > rowCount {
freeQueryResultVectors(cloned, mp)
return nil, nil, moerr.NewInternalErrorNoCtxf(
"query result column %d has %d rows, batch has %d",
i, vec.Length(), rowCount,
)
}
if !vec.IsConst() {
for row := vec.Length(); row < rowCount; row++ {
if !vec.GetNulls().Contains(uint64(row)) {
freeQueryResultVectors(cloned, mp)
return nil, nil, moerr.NewInternalErrorNoCtxf(
"query result column %d is missing non-null row %d",
i, row,
)
}
}
}
dup, err := vec.Dup(mp)
if err != nil {
freeQueryResultVectors(cloned, mp)
return nil, nil, err
}
cloned = append(cloned, dup)
// Prepared-parameter provenance is execution-only and is not part of the
// stable vector encoding. Drop it before extending the persisted logical
// length so a heterogeneous sidecar cannot grow with the result batch.
dup.SetPrepareParamKinds(nil)
if vec.IsConst() {
dup.SetLength(rowCount)
} else {
for dup.Length() < rowCount {
if err := dup.UnionNull(mp); err != nil {
freeQueryResultVectors(cloned, mp)
return nil, nil, err
}
}
}
writeBat.Vecs[i] = dup
}
return writeBat, func() {
freeQueryResultVectors(cloned, mp)
}, nil
}

func freeQueryResultVectors(vecs []*vector.Vector, mp *mpool.MPool) {
for _, vec := range vecs {
vec.Free(mp)
}
}

func saveBatches(ctx context.Context, ses *Session, data []*batch.Batch) error {
for _, b := range data {
if err := saveBatch(ctx, ses, b); err != nil {
Expand Down
Loading
Loading