From 45ea3ce57fee4a7c5bf4367b779c02e7c4d73a09 Mon Sep 17 00:00:00 2001 From: Torin Sandall Date: Mon, 15 Apr 2019 15:02:40 -0700 Subject: [PATCH] Add support for with keyword stacking (data) This is a follow-on to the previous commit. Note, with this change, input and data replacement are handled using the same code (whereas before the data replacement used the baseCache which could lead to slightly different behaviour). The only difference is that the data replacement has to keep track of prefixes so that base/virtual documents can be shadowed. Also, rename helper that merges with values into term to indicate purpose. Signed-off-by: Torin Sandall --- topdown/cache.go | 33 ++++++++++++++++++++ topdown/eval.go | 68 ++++++++++++++++++++++++++--------------- topdown/input.go | 12 ++++---- topdown/input_test.go | 10 +++--- topdown/query.go | 4 +-- topdown/topdown_test.go | 30 +++++++++++++++++- 6 files changed, 118 insertions(+), 39 deletions(-) diff --git a/topdown/cache.go b/topdown/cache.go index 29534d7d8b..4eeae1d59d 100644 --- a/topdown/cache.go +++ b/topdown/cache.go @@ -171,3 +171,36 @@ func (e *baseCacheElem) set(value ast.Value) { e.value = value e.children = map[ast.Value]*baseCacheElem{} } + +type refStack struct { + sl []refStackElem +} + +type refStackElem struct { + refs []ast.Ref +} + +func newRefStack() *refStack { + return &refStack{} +} + +func (s *refStack) Push(refs []ast.Ref) { + s.sl = append(s.sl, refStackElem{refs: refs}) +} + +func (s *refStack) Pop() { + s.sl = s.sl[:len(s.sl)-1] +} + +func (s *refStack) Prefixed(ref ast.Ref) bool { + if s != nil { + for i := len(s.sl) - 1; i >= 0; i-- { + for j := range s.sl[i].refs { + if ref.HasPrefix(s.sl[i].refs[j]) { + return true + } + } + } + } + return false +} diff --git a/topdown/eval.go b/topdown/eval.go index 8df370f980..bb9bb848cd 100644 --- a/topdown/eval.go +++ b/topdown/eval.go @@ -36,10 +36,11 @@ type eval struct { bindings *bindings store storage.Store baseCache *baseCache - withCache *baseCache txn storage.Transaction compiler *ast.Compiler input *ast.Term + data *ast.Term + targetStack *refStack tracers []Tracer instr *Instrumentation builtinCache builtins.Cache @@ -280,6 +281,7 @@ func (e *eval) evalWith(iter evalIterator) error { pairsInput := [][2]*ast.Term{} pairsData := [][2]*ast.Term{} + targets := []ast.Ref{} for i := range expr.With { plugged := e.bindings.Plug(expr.With[i].Value) @@ -288,9 +290,10 @@ func (e *eval) evalWith(iter evalIterator) error { } else if isDataRef(expr.With[i].Target) { pairsData = append(pairsData, [...]*ast.Term{expr.With[i].Target, plugged}) } + targets = append(targets, expr.With[i].Target.Value.(ast.Ref)) } - input, err := makeInput(e.input, pairsInput) + input, err := mergeTermWithValues(e.input, pairsInput) if err != nil { return &Error{ @@ -300,48 +303,55 @@ func (e *eval) evalWith(iter evalIterator) error { } } - old := e.evalWithPush(input, pairsData) + data, err := mergeTermWithValues(e.data, pairsData) + if err != nil { + return &Error{ + Code: ConflictErr, + Location: expr.Location, + Message: err.Error(), + } + } + + oldInput, oldData := e.evalWithPush(input, data, targets) err = e.evalStep(func(e *eval) error { - e.evalWithPop(old, pairsData) + e.evalWithPop(oldInput, oldData) err := e.next(iter) - old = e.evalWithPush(input, pairsData) + oldInput, oldData = e.evalWithPush(input, data, targets) return err }) - e.evalWithPop(old, pairsData) + e.evalWithPop(oldInput, oldData) return err } -func (e *eval) evalWithPush(input *ast.Term, data [][2]*ast.Term) *ast.Term { +func (e *eval) evalWithPush(input *ast.Term, data *ast.Term, targets []ast.Ref) (*ast.Term, *ast.Term) { - var old *ast.Term + var oldInput *ast.Term if input != nil { - old = e.input + oldInput = e.input e.input = input } - for _, pair := range data { - ref := pair[0].Value.(ast.Ref) - e.withCache.Put(ref, pair[1].Value) + var oldData *ast.Term + + if data != nil { + oldData = e.data + e.data = data } e.virtualCache.Push() + e.targetStack.Push(targets) - return old + return oldInput, oldData } -func (e *eval) evalWithPop(input *ast.Term, data [][2]*ast.Term) { - +func (e *eval) evalWithPop(input *ast.Term, data *ast.Term) { + e.targetStack.Pop() e.virtualCache.Pop() - - for _, pair := range data { - ref := pair[0].Value.(ast.Ref) - e.withCache.Remove(ref) - } - + e.data = data e.input = input } @@ -1001,8 +1011,17 @@ func (e *eval) Resolve(ref ast.Ref) (ast.Value, error) { if ref[0].Equal(ast.DefaultRootDocument) { - repValue, complete := e.withCache.Get(ref) - if complete { + var repValue ast.Value + + if e.data != nil { + if v, err := e.data.Value.Find(ref[1:]); err == nil { + repValue = v + } else { + repValue = nil + } + } + + if e.targetStack.Prefixed(ref) { return repValue, nil } @@ -1282,9 +1301,8 @@ func (e evalTree) next(iter unifyIterator, plugged *ast.Term) error { cpy := e cpy.plugged[e.pos] = plugged cpy.pos++ - _, complete := e.e.withCache.Get(cpy.plugged[:cpy.pos]) - if !complete { + if !e.e.targetStack.Prefixed(cpy.plugged[:cpy.pos]) { if e.node != nil { node = e.node.Child(plugged.Value) if node != nil && len(node.Values) > 0 { diff --git a/topdown/input.go b/topdown/input.go index 7b7a4fcf9e..a74ef199f6 100644 --- a/topdown/input.go +++ b/topdown/input.go @@ -10,10 +10,10 @@ import ( "github.com/open-policy-agent/opa/ast" ) -var errConflictingInputDoc = fmt.Errorf("conflicting input documents") -var errBadInputPath = fmt.Errorf("bad input document path") +var errConflictingDoc = fmt.Errorf("conflicting documents") +var errBadPath = fmt.Errorf("bad document path") -func makeInput(exist *ast.Term, pairs [][2]*ast.Term) (*ast.Term, error) { +func mergeTermWithValues(exist *ast.Term, pairs [][2]*ast.Term) (*ast.Term, error) { var result *ast.Term @@ -24,7 +24,7 @@ func makeInput(exist *ast.Term, pairs [][2]*ast.Term) (*ast.Term, error) { for _, pair := range pairs { if err := ast.IsValidImportPath(pair[0].Value); err != nil { - return nil, errBadInputPath + return nil, errBadPath } target := pair[0].Value.(ast.Ref) @@ -40,7 +40,7 @@ func makeInput(exist *ast.Term, pairs [][2]*ast.Term) (*ast.Term, error) { if child := node.Get(target[i]); child == nil { obj, ok := node.Value.(ast.Object) if !ok { - return nil, errConflictingInputDoc + return nil, errConflictingDoc } obj.Insert(target[i], ast.NewTerm(makeTree(target[i+1:], pair[1]))) done = true @@ -51,7 +51,7 @@ func makeInput(exist *ast.Term, pairs [][2]*ast.Term) (*ast.Term, error) { if !done { obj, ok := node.Value.(ast.Object) if !ok { - return nil, errConflictingInputDoc + return nil, errConflictingDoc } obj.Insert(target[len(target)-1], pair[1]) } diff --git a/topdown/input_test.go b/topdown/input_test.go index 83c1405f7c..36fdada46f 100644 --- a/topdown/input_test.go +++ b/topdown/input_test.go @@ -10,7 +10,7 @@ import ( "github.com/open-policy-agent/opa/ast" ) -func TestMakeInput(t *testing.T) { +func TestMergeTermWithValues(t *testing.T) { tests := []struct { note string @@ -46,12 +46,12 @@ func TestMakeInput(t *testing.T) { { note: "conflicting value", input: [][2]string{{"input", "[1,2,3]"}, {"input.a.b", "true"}}, - expected: errConflictingInputDoc, + expected: errConflictingDoc, }, { note: "conflicting merge", input: [][2]string{{`input.a.b`, `"c"`}, {`input.a.b.d`, `"d"`}}, - expected: errConflictingInputDoc, + expected: errConflictingDoc, }, { note: "ordered roots", @@ -61,7 +61,7 @@ func TestMakeInput(t *testing.T) { { note: "bad import path", input: [][2]string{{`input.a[1]`, `1`}}, - expected: errBadInputPath, + expected: errBadPath, }, { note: "existing merge", @@ -96,7 +96,7 @@ func TestMakeInput(t *testing.T) { exist = ast.MustParseTerm(tc.exist) } - input, err := makeInput(exist, pairs) + input, err := mergeTermWithValues(exist, pairs) switch e := tc.expected.(type) { case error: diff --git a/topdown/query.go b/topdown/query.go index fad5b08f59..edd85ee98c 100644 --- a/topdown/query.go +++ b/topdown/query.go @@ -141,7 +141,7 @@ func (q *Query) PartialRun(ctx context.Context) (partials []ast.Body, support [] compiler: q.compiler, store: q.store, baseCache: newBaseCache(), - withCache: newBaseCache(), + targetStack: newRefStack(), txn: q.txn, input: q.input, tracers: q.tracers, @@ -228,7 +228,7 @@ func (q *Query) Iter(ctx context.Context, iter func(QueryResult) error) error { compiler: q.compiler, store: q.store, baseCache: newBaseCache(), - withCache: newBaseCache(), + targetStack: newRefStack(), txn: q.txn, input: q.input, tracers: q.tracers, diff --git a/topdown/topdown_test.go b/topdown/topdown_test.go index e547c3e58c..2d0c3c4dea 100644 --- a/topdown/topdown_test.go +++ b/topdown/topdown_test.go @@ -2632,7 +2632,7 @@ func TestTopDownWithKeyword(t *testing.T) { }, { note: "with conflict", - exp: fmt.Errorf("conflicting input documents"), + exp: fmt.Errorf("conflicting documents"), modules: []string{`package ex loopback = __local0__ { true; __local0__ = input }`}, rules: []string{`p = true { data.ex.loopback with input.foo as "x" with input.foo.bar as "y" }`}, @@ -2647,6 +2647,21 @@ func TestTopDownWithKeyword(t *testing.T) { `p = x { q = x with input.a.b as 1 }`, }, }, + { + note: "with stack (data)", + exp: `{"a": {"b": 1, "c": 2, "d": 3}, "e": 4}`, + modules: []string{ + `package test.a + d = 3`, + `package test + e = 4`, + }, + rules: []string{ + `r = data.test { true }`, + `q = x { r = x with data.test.a.c as 2 }`, + `p = x { q = x with data.test.a.b as 1 }`, + }, + }, { note: "with stack overwrites", input: `{"a": {"b": 1, "c": 2}}`, @@ -2656,6 +2671,19 @@ func TestTopDownWithKeyword(t *testing.T) { `p = x { q = x with input.a as {"d": 3} }`, }, }, + { + note: "with stack overwrites (data)", + exp: `{"a": {"d": 3}}`, + modules: []string{ + `package test + + a = {"b": 1, "c": 2}`, + }, + rules: []string{ + `q = data.test { true }`, + `p = x { q = x with data.test.a as {"d": 3} }`, + }, + }, { note: "with invalidate", exp: `[2,3,4]`,