diff --git a/executor.go b/executor.go index b3feb40aa..97d336b32 100644 --- a/executor.go +++ b/executor.go @@ -222,7 +222,7 @@ func (e *executor) Execute(ctx context.Context, index string, q *pql.Query, shar } // Can't do NewTx() this high up, because we need a specific shard. - // So start a ccx with a TxGroup and pass it down. + // So start a qcx with a TxGroup and pass it down. qcx := idx.holder.txf.NewQcx() qcx.write = needWriteTxn defer qcx.Abort() @@ -474,12 +474,17 @@ func (e *executor) execute(ctx context.Context, qcx *Qcx, index string, q *pql.Q colTranslations, rowTranslations = cols, rows } - // Don't bother calculating shards for query types that don't require it. - needsShards := needsShards(q.Calls) + needShards := false + if len(shards) == 0 { + for _, call := range q.Calls { + if needsShards(call) { + needShards = true + break + } + } + } - // If shards are specified, then use that value for shards. If shards aren't - // specified, then include all of them. - if len(shards) == 0 && needsShards { + if needShards { // Round up the number of shards. idx := e.Holder.Index(index) if idx == nil { @@ -491,9 +496,23 @@ func (e *executor) execute(ctx context.Context, qcx *Qcx, index string, q *pql.Q } } + lastWasWrite := false // Execute each call serially. results := make([]interface{}, 0, len(q.Calls)) for i, call := range q.Calls { + if lastWasWrite && needsShards(call) && needShards { + // Round up the number of shards. + idx := e.Holder.Index(index) + if idx == nil { + return nil, newNotFoundError(ErrIndexNotFound, index) + } + shards = idx.AvailableShards(includeRemote).Slice() + if len(shards) == 0 { + shards = []uint64{0} + } + } + + lastWasWrite = call.IsWrite() if err := validateQueryContext(ctx); err != nil { return nil, err @@ -655,7 +674,7 @@ func (e *executor) executeCall(ctx context.Context, qcx *Qcx, index string, c *p // If shards are specified, then use that value for shards. If shards aren't // specified, then include all of them. - if shards == nil && needsShards([]*pql.Call{c}) { + if shards == nil && needsShards(c) { // Round up the number of shards. idx := e.Holder.Index(index) if idx == nil { @@ -7717,22 +7736,18 @@ type execOptions struct { MaxMemory int64 } -func needsShards(calls []*pql.Call) bool { - if len(calls) == 0 { +func needsShards(call *pql.Call) bool { + if call == nil { return false } - for _, call := range calls { - switch call.Name { - case "Clear", "Set": - continue - case "Count", "TopN", "Rows": - return true - // default catches Bitmap calls - default: - return true - } + switch call.Name { + case "Clear", "Set": + return false + case "Count", "TopN", "Rows": + return true } - return false + // default catches Bitmap calls + return true } // SignedRow represents a signed *Row with two (neg/pos) *Rows. diff --git a/executor_test.go b/executor_test.go index bd592cce3..e641d8ef8 100644 --- a/executor_test.go +++ b/executor_test.go @@ -7169,14 +7169,12 @@ func TestMissingKeyRegression(t *testing.T) { query: `Difference(All(), Row(f="garbage"))`, expected: []interface{}{[]string{"a"}}, }, - /*{ - // Key translation works here, but it seems the actual count query is processing stale data. - // Uncomment it when the bug has been fixed. + { name: "SetAndCount", - query: `Set("a", f="example")` + "\n" + - `Count(Row(f="example"))`, + query: `Set("b", f="boo")` + "\n" + + `Count(Row(f="boo"))`, expected: []interface{}{true, uint64(1)}, - },*/ + }, { name: "CountNothing", query: `Count(Row(f="garbage"))`, diff --git a/pql/ast.go b/pql/ast.go index 3480f4def..04b754667 100644 --- a/pql/ast.go +++ b/pql/ast.go @@ -378,6 +378,18 @@ func (c *Call) HasCall(name string) bool { return false } +// IsWrite returns whether the call is a mutating call. +func (c *Call) IsWrite() bool { + if c == nil { + return false + } + switch c.Name { + case "Set", "Clear", "ClearRow", "Store", "SetBit": + return true + } + return false +} + // callInfo defines the arguments allowed for a particular PQL call, and // possibly things about its semantics. If allowUnknown is true, unfamiliar // non-reserved names are allowed on the assumption that they're field names.