diff --git a/sql/expression/function/aggregation/count_test.go b/sql/expression/function/aggregation/count_test.go index 10c769853b..87aa278ba6 100644 --- a/sql/expression/function/aggregation/count_test.go +++ b/sql/expression/function/aggregation/count_test.go @@ -137,6 +137,21 @@ func TestCountDistinctEvalString(t *testing.T) { require.Equal(int64(2), evalBuffer(t, b)) } +func TestCountDistinctEvalMultiColumn(t *testing.T) { + require := require.New(t) + ctx := sql.NewEmptyContext() + + c := NewCountDistinct( + expression.NewGetField(0, types.Text, "", true), + expression.NewGetField(1, types.Text, "", true), + ) + b, _ := c.NewBuffer(ctx) + + require.NoError(b.Update(ctx, sql.NewRow("a,", "b"))) + require.NoError(b.Update(ctx, sql.NewRow("a", ",b"))) + require.Equal(int64(2), evalBuffer(t, b)) +} + func TestCountDistinctEvalExtendedType(t *testing.T) { require := require.New(t) ctx := sql.NewEmptyContext() diff --git a/sql/expression/function/aggregation/unary_agg_buffers.go b/sql/expression/function/aggregation/unary_agg_buffers.go index f3e123b286..4e65caf0f1 100644 --- a/sql/expression/function/aggregation/unary_agg_buffers.go +++ b/sql/expression/function/aggregation/unary_agg_buffers.go @@ -486,7 +486,7 @@ func (c *countDistinctBuffer) Update(ctx *sql.Context, row sql.Row) error { } } - var str string + hash := xxhash.New() for _, val := range value.(sql.Row) { // skip nil values if val == nil { @@ -500,16 +500,11 @@ func (c *countDistinctBuffer) Update(ctx *sql.Context, row sql.Row) error { if !ok { return fmt.Errorf("count distinct unable to hash value: %s", err) } - str += vv + "," - } - - hash := xxhash.New() - _, err := hash.WriteString(str) - if err != nil { - return err + if _, err = fmt.Fprintf(hash, "%d:%s", len(vv), vv); err != nil { + return err + } } - h := hash.Sum64() - c.seen[h] = struct{}{} + c.seen[hash.Sum64()] = struct{}{} return nil }