pql.Decimal for DecimalVal in ValCount&GroupCount

This way we can avoid annoying floating point rounding errors.

Check out FB-1359 for an example:
```
--- FAIL: TestExecutor_GroupByStrings (0.55s)
    --- FAIL: TestExecutor_GroupByStrings/3 (0.00s)
        executor_test.go:5433: unexpected result at 0:
             got:{Group:[generals.1.r1] Count:5 Agg:2775
DecimalAgg:27.749999999999996}
            want:{Group:[generals.1.r1] Count:5 Agg:2775 DecimalAgg:27.75}
```
This commit is contained in:
reesporte 2022-06-01 15:48:54 -05:00 • committed by reesporte
parent 50787fd37a
commit a0c9eec410
2 changed files with 38 additions and 27 deletions

View file

@ -1892,6 +1892,7 @@ func (e *executor) executeSumCountShard(ctx context.Context, qcx *Qcx, index str
}
if field.Type() == FieldTypeDecimal {
out.FloatVal = float64(int64(vsum)+(int64(vcount)*bsig.Base)) / math.Pow(10, float64(bsig.Scale))
out.DecimalVal = &pql.Decimal{Value: (int64(vsum) + (int64(vcount) * bsig.Base)), Scale: bsig.Scale}
}
return out, nil
}
@ -3000,7 +3001,11 @@ func (e *executor) executeGroupBy(ctx context.Context, qcx *Qcx, index string, c
if err := ctx.Err(); err != nil {
return err
}
return mergeGroupCounts(other, findGroupCounts(v), limit)
merged, err := mergeGroupCounts(other, findGroupCounts(v), limit)
if err != nil {
return err
}
return merged
}
// Get full result set.
other, err := e.mapReduce(ctx, index, shards, c, opt, mapFn, reduceFn)
@ -3115,7 +3120,7 @@ func (e *executor) executeGroupBy(ctx context.Context, qcx *Qcx, index string, c
}
}
for _, res := range results {
if res.DecimalAgg != 0 && aggType == "sum" {
if res.DecimalAgg != nil && aggType == "sum" {
aggType = "decimalSum"
break
}
@ -3359,31 +3364,31 @@ func (g *GroupCounts) MarshalJSON() ([]byte, error) {
// GroupCount represents a result item for a group by query.
type GroupCount struct {
Group []FieldRow `json:"group"`
Count uint64 `json:"count"`
Agg int64 `json:"-"`
DecimalAgg float64 `json:"-"`
Group []FieldRow `json:"group"`
Count uint64 `json:"count"`
Agg int64 `json:"-"`
DecimalAgg *pql.Decimal `json:"-"`
}
type groupCountSum struct {
Group []FieldRow `json:"group"`
Count uint64 `json:"count"`
Agg int64 `json:"sum"`
DecimalAgg float64 `json:"-"`
Group []FieldRow `json:"group"`
Count uint64 `json:"count"`
Agg int64 `json:"sum"`
DecimalAgg *pql.Decimal `json:"-"`
}
type groupCountAggregate struct {
Group []FieldRow `json:"group"`
Count uint64 `json:"count"`
Agg int64 `json:"aggregate"`
DecimalAgg float64 `json:"-"`
Group []FieldRow `json:"group"`
Count uint64 `json:"count"`
Agg int64 `json:"aggregate"`
DecimalAgg *pql.Decimal `json:"-"`
}
type groupCountDecimalSum struct {
Group []FieldRow `json:"group"`
Count uint64 `json:"count"`
Agg int64 `json:"-"`
DecimalAgg float64 `json:"sum"`
Group []FieldRow `json:"group"`
Count uint64 `json:"count"`
Agg int64 `json:"-"`
DecimalAgg *pql.Decimal `json:"sum"`
}
var _ GroupCount = GroupCount(groupCountSum{})
@ -3406,7 +3411,7 @@ func (g *GroupCount) Clone() (r *GroupCount) {
// mergeGroupCounts merges two slices of GroupCounts throwing away any that go
// beyond the limit. It assume that the two slices are sorted by the row ids in
// the fields of the group counts. It may modify its arguments.
func mergeGroupCounts(a, b []GroupCount, limit int) []GroupCount {
func mergeGroupCounts(a, b []GroupCount, limit int) ([]GroupCount, error) {
if limit > len(a)+len(b) {
limit = len(a) + len(b)
}
@ -3420,7 +3425,13 @@ func mergeGroupCounts(a, b []GroupCount, limit int) []GroupCount {
case 0:
a[i].Count += b[j].Count
a[i].Agg += b[j].Agg
a[i].DecimalAgg += b[j].DecimalAgg
if a[i].DecimalAgg != nil && b[j].DecimalAgg != nil {
sum, ok := pql.AddDecimal(*a[i].DecimalAgg, *b[j].DecimalAgg)
if !ok {
return nil, fmt.Errorf("cannot add %s and %s, decimal overflow", a[i].DecimalAgg, b[j].DecimalAgg)
}
a[i].DecimalAgg = &sum
}
ret = append(ret, a[i])
i++
j++
@ -3435,7 +3446,7 @@ func mergeGroupCounts(a, b []GroupCount, limit int) []GroupCount {
for ; j < len(b) && len(ret) < limit; j++ {
ret = append(ret, b[j])
}
return ret
return ret, nil
}
// Compare is used in ordering two GroupCount objects.
@ -8131,7 +8142,7 @@ func (gbi *groupByIterator) Next(ctx context.Context) (ret GroupCount, done bool
}
ret.Count = uint64(result.Count)
ret.Agg = result.Val
ret.DecimalAgg = result.FloatVal
ret.DecimalAgg = result.DecimalVal
}
}
if ret.Count == 0 {

View file

@ -5358,15 +5358,15 @@ func TestExecutor_GroupByStrings(t *testing.T) {
{
query: "GroupBy(Rows(generals), aggregate=Sum(field=dv))",
expected: []pilosa.GroupCount{
{Group: []pilosa.FieldRow{{Field: "generals", RowID: 1, RowKey: "r1"}}, Count: 5, Agg: 2775, DecimalAgg: 27.75},
{Group: []pilosa.FieldRow{{Field: "generals", RowID: 2, RowKey: "r2"}}, Count: 5, Agg: 3220, DecimalAgg: 32.20},
{Group: []pilosa.FieldRow{{Field: "generals", RowID: 1, RowKey: "r1"}}, Count: 5, Agg: 2775, DecimalAgg: &pql.Decimal{Value: 2775, Scale: 2}},
{Group: []pilosa.FieldRow{{Field: "generals", RowID: 2, RowKey: "r2"}}, Count: 5, Agg: 3220, DecimalAgg: &pql.Decimal{Value: 3220, Scale: 2}},
},
},
{
query: "GroupBy(Rows(generals), aggregate=Sum(field=ndv))",
expected: []pilosa.GroupCount{
{Group: []pilosa.FieldRow{{Field: "generals", RowID: 1, RowKey: "r1"}}, Count: 5, Agg: -2775, DecimalAgg: -277.5},
{Group: []pilosa.FieldRow{{Field: "generals", RowID: 2, RowKey: "r2"}}, Count: 5, Agg: -3220, DecimalAgg: -322.0},
{Group: []pilosa.FieldRow{{Field: "generals", RowID: 1, RowKey: "r1"}}, Count: 5, Agg: -2775, DecimalAgg: &pql.Decimal{Value: -2775, Scale: 1}},
{Group: []pilosa.FieldRow{{Field: "generals", RowID: 2, RowKey: "r2"}}, Count: 5, Agg: -3220, DecimalAgg: &pql.Decimal{Value: -3220, Scale: 1}},
},
},
{
@ -5509,7 +5509,7 @@ func TestExecutor_GroupByStrings(t *testing.T) {
}
for i, tst := range tests {
t.Run(fmt.Sprintf("%d", i), func(t *testing.T) {
t.Run(fmt.Sprintf("%s%d", tst.query, i), func(t *testing.T) {
r, err := c.GetNode(0).API.Query(context.Background(), &pilosa.QueryRequest{
Index: "istring",
Query: tst.query,