From a0c9eec410d4c07892dd6808dcf39a89d06c1fe5 Mon Sep 17 00:00:00 2001 From: reesporte Date: Wed, 1 Jun 2022 15:48:54 -0500 Subject: [PATCH] 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} ``` --- executor.go | 55 +++++++++++++++++++++++++++++------------------- executor_test.go | 10 ++++----- 2 files changed, 38 insertions(+), 27 deletions(-) diff --git a/executor.go b/executor.go index d21cfbb5b..34eaffe66 100644 --- a/executor.go +++ b/executor.go @@ -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 { diff --git a/executor_test.go b/executor_test.go index 51214aa3e..e2ce939fb 100644 --- a/executor_test.go +++ b/executor_test.go @@ -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,