mirror of
https://github.com/open-policy-agent/opa.git
synced 2026-08-12 19:32:48 -06:00
+21
-30
@@ -241,41 +241,32 @@ func (c *Compiler) resolveAllRefs() {
|
||||
|
||||
func (c *Compiler) resolveRef(globals map[Var]Value, ref Ref) Ref {
|
||||
|
||||
global := globals[ref[0].Value.(Var)]
|
||||
if global == nil {
|
||||
return ref
|
||||
}
|
||||
fqn := Ref{}
|
||||
switch global := global.(type) {
|
||||
case Ref:
|
||||
fqn = append(fqn, global...)
|
||||
for _, p := range ref[1:] {
|
||||
switch v := p.Value.(type) {
|
||||
case Var:
|
||||
global := globals[v]
|
||||
if global != nil {
|
||||
_, isRef := global.(Ref)
|
||||
if isRef {
|
||||
c.err("nested references in %v: %v => %v", ref, v, global)
|
||||
return ref
|
||||
r := Ref{}
|
||||
for i, x := range ref {
|
||||
switch v := x.Value.(type) {
|
||||
case Var:
|
||||
if g, ok := globals[v]; ok {
|
||||
switch g := g.(type) {
|
||||
case Ref:
|
||||
if i == 0 {
|
||||
r = append(r, g...)
|
||||
} else {
|
||||
r = append(r, &Term{Location: x.Location, Value: g[:]})
|
||||
}
|
||||
fqn = append(fqn, &Term{Location: p.Location, Value: global})
|
||||
} else {
|
||||
fqn = append(fqn, p)
|
||||
case Var:
|
||||
r = append(r, &Term{Value: g})
|
||||
}
|
||||
default:
|
||||
fqn = append(fqn, p)
|
||||
} else {
|
||||
r = append(r, x)
|
||||
}
|
||||
case Ref:
|
||||
r = append(r, c.resolveRefsInTerm(globals, x))
|
||||
default:
|
||||
r = append(r, x)
|
||||
}
|
||||
case Var:
|
||||
fqn = append(fqn, &Term{Value: global})
|
||||
fqn = append(fqn, ref[1:]...)
|
||||
default:
|
||||
c.err("unexpected %T: %v", global, global)
|
||||
return ref
|
||||
}
|
||||
|
||||
return fqn
|
||||
return r
|
||||
}
|
||||
|
||||
func (c *Compiler) resolveRefsInBody(globals map[Var]Value, body Body) Body {
|
||||
@@ -515,7 +506,7 @@ func (vis *ruleGraphBuilder) Visit(v interface{}) Visitor {
|
||||
for _, v := range findRules(vis.moduleTree, ref) {
|
||||
vis.edges[v] = struct{}{}
|
||||
}
|
||||
return nil
|
||||
return vis
|
||||
}
|
||||
|
||||
type ruleGraphTraveral struct {
|
||||
|
||||
+48
-4
@@ -16,9 +16,10 @@ func TestModuleTree(t *testing.T) {
|
||||
|
||||
mods := getCompilerTestModules()
|
||||
tree := NewModuleTree(mods)
|
||||
expectedSize := 6
|
||||
|
||||
if tree.Size() != 5 {
|
||||
t.Errorf("Expected size of 4 in module tree but got: %v", tree.Size())
|
||||
if tree.Size() != expectedSize {
|
||||
t.Errorf("Expected size of %v in module tree but got: %v", expectedSize, tree.Size())
|
||||
}
|
||||
|
||||
if r1 := findRules(tree, MustParseRef("data.a.b.c")); len(r1) != 0 {
|
||||
@@ -146,6 +147,7 @@ func TestCompilerCheckSafetyBodyReordering(t *testing.T) {
|
||||
// trivial cases
|
||||
{"noop", "x = 1, x != 0", "x = 1, x != 0"},
|
||||
{"var/ref", "a[i] = x, a = [1,2,3,4]", "a = [1,2,3,4], a[i] = x"},
|
||||
{"var/ref (nested)", "a = [1,2,3,4], a[b[i]] = x, b = [0,0,0,0]", "a = [1,2,3,4], b = [0,0,0,0], a[b[i]] = x"},
|
||||
{"negation",
|
||||
"a = [true, false], b = [true, false], not a[i], b[i]",
|
||||
"a = [true, false], b = [true, false], b[i], not a[i]"},
|
||||
@@ -302,6 +304,8 @@ func TestCompilerCheckSafetyBodyErrors(t *testing.T) {
|
||||
unsafeClosure1 :- x = [x | x = 1]
|
||||
unsafeClosure2 :- x = y, x = [y | y = 1]
|
||||
|
||||
unsafeNestedHead :- count(baz[i].attr[bar[dead.beef]], n)
|
||||
|
||||
negatedImport1 = true :- not foo
|
||||
negatedImport2 = true :- not bar
|
||||
negatedImport3 = true :- not baz
|
||||
@@ -327,6 +331,7 @@ func TestCompilerCheckSafetyBodyErrors(t *testing.T) {
|
||||
fmt.Errorf("unsafe variables in unboundArrayComprMixed1: [x z]"),
|
||||
fmt.Errorf("unsafe variables in unsafeClosure1: [x]"),
|
||||
fmt.Errorf("unsafe variables in unsafeClosure2: [y]"),
|
||||
fmt.Errorf("unsafe variables in unsafeNestedHead: [dead]"),
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(expected, c.Errors) {
|
||||
@@ -350,6 +355,7 @@ func TestCompilerResolveAllRefs(t *testing.T) {
|
||||
|
||||
assertNotFailed(t, c)
|
||||
|
||||
// Basic test cases.
|
||||
mod1 := c.Modules["mod1"]
|
||||
p := mod1.Rules[0]
|
||||
expr1 := p.Body[0]
|
||||
@@ -403,6 +409,17 @@ func TestCompilerResolveAllRefs(t *testing.T) {
|
||||
acTerm6 := ac(mod5.Rules[5])
|
||||
assertTermEqual(t, acTerm6.Body[0].Terms.([]*Term)[2].Value.(*ArrayComprehension).Body[0].Terms.([]*Term)[1], MustParseTerm("a.b.c.q[i]"))
|
||||
|
||||
// Nested references.
|
||||
mod6 := c.Modules["mod6"]
|
||||
nested1 := mod6.Rules[0].Body[0].Terms.(*Term)
|
||||
assertTermEqual(t, nested1, MustParseTerm("data.x[x[i].a[data.z.b[j]]]"))
|
||||
|
||||
nested2 := mod6.Rules[1].Body[1].Terms.(*Term)
|
||||
assertTermEqual(t, nested2, MustParseTerm("v[x[i]]"))
|
||||
|
||||
nested3 := mod6.Rules[3].Body[0].Terms.(*Term)
|
||||
assertTermEqual(t, nested3, MustParseTerm("data.x[data.a.b.nested.r]"))
|
||||
|
||||
}
|
||||
|
||||
func TestCompilerSetRuleGraph(t *testing.T) {
|
||||
@@ -464,6 +481,11 @@ func TestCompilerCheckRecursion(t *testing.T) {
|
||||
acp[x] :- acq[x]
|
||||
acq[x] :- a = [x | acp[x]], a[i] = x
|
||||
`),
|
||||
"newMod7": MustParseModule(`
|
||||
package rec6
|
||||
np[x] = y :- data.a[data.b.c[nq[x]]] = y
|
||||
nq[x] = y :- data.d[data.e[x].f[np[y]]]
|
||||
`),
|
||||
}
|
||||
|
||||
compileStages(c, "", "checkRecursion")
|
||||
@@ -479,10 +501,12 @@ func TestCompilerCheckRecursion(t *testing.T) {
|
||||
fmt.Errorf("recursion found in q: q, p, q"),
|
||||
fmt.Errorf("recursion found in acq: acq, acp, acq"),
|
||||
fmt.Errorf("recursion found in acp: acp, acq, acp"),
|
||||
fmt.Errorf("recursion found in np: np, nq, np"),
|
||||
fmt.Errorf("recursion found in nq: nq, np, nq"),
|
||||
}
|
||||
|
||||
if len(c.Errors) != len(expected) {
|
||||
t.Errorf("Expected exactly %v errors but got: %v", len(expected), c.Errors)
|
||||
t.Errorf("Expected exactly %v errors but got %v: %v", len(expected), len(c.Errors), c.Errors)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -672,5 +696,25 @@ func getCompilerTestModules() map[string]*Module {
|
||||
v :- [true | _ = [ true | q[i] = 1]]
|
||||
`)
|
||||
|
||||
return map[string]*Module{"mod2": mod2, "mod3": mod3, "mod1": mod1, "mod4": mod4, "mod5": mod5}
|
||||
mod6 := MustParseModule(`
|
||||
package a.b.nested
|
||||
|
||||
import data.x
|
||||
import x as y
|
||||
import data.z
|
||||
|
||||
p :- x[y[i].a[z.b[j]]]
|
||||
q :- x = v, v[y[i]]
|
||||
r = 1 :- true
|
||||
s :- x[r]
|
||||
`)
|
||||
|
||||
return map[string]*Module{
|
||||
"mod1": mod1,
|
||||
"mod2": mod2,
|
||||
"mod3": mod3,
|
||||
"mod4": mod4,
|
||||
"mod5": mod5,
|
||||
"mod6": mod6,
|
||||
}
|
||||
}
|
||||
|
||||
+9
-2
@@ -72,12 +72,19 @@ func TestRefTerms(t *testing.T) {
|
||||
assertParseOneTerm(t, "constants 2", "foo.bar[0].baz", RefTerm(VarTerm("foo"), StringTerm("bar"), NumberTerm(0), StringTerm("baz")))
|
||||
assertParseOneTerm(t, "variables", "foo.bar[0].baz[i]", RefTerm(VarTerm("foo"), StringTerm("bar"), NumberTerm(0), StringTerm("baz"), VarTerm("i")))
|
||||
assertParseOneTerm(t, "spaces", "foo[\"white space\"].bar", RefTerm(VarTerm("foo"), StringTerm("white space"), StringTerm("bar")))
|
||||
assertParseOneTerm(t, "nested", "foo[baz[1][borge[i]]].bar", RefTerm(
|
||||
VarTerm("foo"),
|
||||
RefTerm(
|
||||
VarTerm("baz"), NumberTerm(float64(1)), RefTerm(
|
||||
VarTerm("borge"), VarTerm("i"),
|
||||
),
|
||||
),
|
||||
StringTerm("bar"),
|
||||
))
|
||||
assertParseError(t, "missing component 1", "foo.")
|
||||
assertParseError(t, "missing component 2", "foo[].bar")
|
||||
assertParseError(t, "composite operand 1", "foo[[1,2,3]].bar")
|
||||
assertParseError(t, "composite operand 2", "foo[{1: 2}].bar")
|
||||
// TODO(tsandall): this may be allowed some day
|
||||
assertParseError(t, "nested refs", "foo[baz.qux].bar")
|
||||
}
|
||||
|
||||
func TestObjectWithScalars(t *testing.T) {
|
||||
|
||||
+1
-1
@@ -252,7 +252,7 @@ RefDot <- "." val:Var {
|
||||
return str, nil
|
||||
}
|
||||
|
||||
RefBracket <- "[" val:(Scalar / Var) "]" {
|
||||
RefBracket <- "[" val:(Ref / Scalar / Var) "]" {
|
||||
return val, nil
|
||||
}
|
||||
|
||||
|
||||
+51
-2
@@ -7,13 +7,12 @@ package ast
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"hash/fnv"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
import "hash/fnv"
|
||||
|
||||
// Location records a position in source code
|
||||
type Location struct {
|
||||
Text []byte // The original text fragment from the source.
|
||||
@@ -50,6 +49,46 @@ type Value interface {
|
||||
Hash() int
|
||||
}
|
||||
|
||||
// InterfaceToValue converts a native Go value x to a Value.
|
||||
func InterfaceToValue(x interface{}) (Value, error) {
|
||||
switch x := x.(type) {
|
||||
case nil:
|
||||
return Null{}, nil
|
||||
case bool:
|
||||
return Boolean(x), nil
|
||||
case float64:
|
||||
return Number(x), nil
|
||||
case string:
|
||||
return String(x), nil
|
||||
case []interface{}:
|
||||
r := Array{}
|
||||
for _, e := range x {
|
||||
e, err := InterfaceToValue(e)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r = append(r, &Term{Value: e})
|
||||
}
|
||||
return r, nil
|
||||
case map[string]interface{}:
|
||||
r := Object{}
|
||||
for k, v := range x {
|
||||
k, err := InterfaceToValue(k)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
v, err := InterfaceToValue(v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r = append(r, Item(&Term{Value: k}, &Term{Value: v}))
|
||||
}
|
||||
return r, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("illegal value: %v", x)
|
||||
}
|
||||
}
|
||||
|
||||
// Term is an argument to a function.
|
||||
type Term struct {
|
||||
Value Value // the value of the Term as represented in Go
|
||||
@@ -343,6 +382,16 @@ func (ref Ref) IsGround() bool {
|
||||
return termSliceIsGround(ref[1:])
|
||||
}
|
||||
|
||||
// IsNested returns true if this ref contains other Refs.
|
||||
func (ref Ref) IsNested() bool {
|
||||
for _, x := range ref {
|
||||
if _, ok := x.Value.(Ref); ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
var varRegexp = regexp.MustCompile("^[[:alpha:]_][[:alpha:][:digit:]_]*$")
|
||||
|
||||
func (ref Ref) String() string {
|
||||
|
||||
@@ -12,6 +12,38 @@ import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestInterfaceToValue(t *testing.T) {
|
||||
input := `
|
||||
{
|
||||
"x": [
|
||||
1,
|
||||
true,
|
||||
false,
|
||||
null,
|
||||
"hello",
|
||||
["goodbye", 1],
|
||||
{"y": 3.1}
|
||||
]
|
||||
}
|
||||
`
|
||||
var x interface{}
|
||||
if err := json.Unmarshal([]byte(input), &x); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
expected := MustParseTerm(input).Value
|
||||
|
||||
v, err := InterfaceToValue(x)
|
||||
if err != nil {
|
||||
t.Errorf("Unexpected error converting interface{} to ast.Value: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if !v.Equal(expected) {
|
||||
t.Errorf("Expected ast.Value to equal:\n%v\nBut got:\n%v", expected, v)
|
||||
}
|
||||
}
|
||||
|
||||
func TestObjectSetOperations(t *testing.T) {
|
||||
|
||||
a := MustParseTerm(`{"a": "b", "c": "d"}`).Value.(Object)
|
||||
|
||||
@@ -157,6 +157,39 @@ Steps
|
||||
| {"id":"s4","name":"dev","ports":["p1","p2"],"protocols":["http"]} |
|
||||
+-------------------------------------------------------------------------------+
|
||||
|
||||
One powerful thing about Rego and the REPL is that you can run queries using the same syntax that you would use to lookup values.
|
||||
|
||||
For example if `i` has value 0 then `data.servers[i]` returns the first value in the `data.servers` array:
|
||||
|
||||
> i = 0
|
||||
> data.servers[i]
|
||||
{
|
||||
"id": "s1",
|
||||
"name": "app",
|
||||
"ports": [
|
||||
"p1",
|
||||
"p2",
|
||||
"p3"
|
||||
],
|
||||
"protocols": [
|
||||
"https",
|
||||
"ssh"
|
||||
]
|
||||
}
|
||||
|
||||
That same expression `data.servers[i]` when `i` has no value defines a query that returns all the values of `i` and `data.servers[i]`:
|
||||
|
||||
> unset i
|
||||
> data.servers[i]
|
||||
+---+-------------------------------------------------------------------------------+
|
||||
| i | data.servers[i] |
|
||||
+---+-------------------------------------------------------------------------------+
|
||||
| 0 | {"id":"s1","name":"app","ports":["p1","p2","p3"],"protocols":["https","ssh"]} |
|
||||
| 1 | {"id":"s2","name":"db","ports":["p3"],"protocols":["mysql"]} |
|
||||
| 2 | {"id":"s3","name":"cache","ports":["p3"],"protocols":["memcache"]} |
|
||||
| 3 | {"id":"s4","name":"dev","ports":["p1","p2"],"protocols":["http"]} |
|
||||
+---+-------------------------------------------------------------------------------+
|
||||
|
||||
1. The REPL also understands the [Import and Package](/docs/lang.html#modules) directives.
|
||||
|
||||
> import data.servers
|
||||
|
||||
+20
-3
@@ -118,9 +118,26 @@ func (ds *DataStore) GetRef(ref ast.Ref) (interface{}, error) {
|
||||
if !ref[0].Equal(ast.DefaultRootDocument) {
|
||||
return nil, fmt.Errorf("illegal root %v: %v", ref[0], ref)
|
||||
}
|
||||
path, err := ref[1:].Underlying()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
path := []interface{}{}
|
||||
for _, x := range ref[1:] {
|
||||
switch v := x.Value.(type) {
|
||||
case ast.Ref:
|
||||
n, err := ds.GetRef(v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
path = append(path, n)
|
||||
case ast.String:
|
||||
path = append(path, string(v))
|
||||
case ast.Number:
|
||||
path = append(path, float64(v))
|
||||
case ast.Boolean:
|
||||
path = append(path, bool(v))
|
||||
case ast.Null:
|
||||
path = append(path, nil)
|
||||
default:
|
||||
return nil, fmt.Errorf("illegal reference element: %v", x)
|
||||
}
|
||||
}
|
||||
return ds.Get(path)
|
||||
}
|
||||
|
||||
+94
-99
@@ -278,9 +278,6 @@ func PlugValue(v ast.Value, ctx *Context) ast.Value {
|
||||
if b := ctx.Binding(v); b != nil {
|
||||
return b
|
||||
}
|
||||
if v.IsGround() {
|
||||
return v
|
||||
}
|
||||
var buf ast.Ref
|
||||
buf = append(buf, v[0])
|
||||
for _, p := range v[1:] {
|
||||
@@ -564,6 +561,7 @@ func evalExpr(ctx *Context, iter Iterator) error {
|
||||
v := tt.Value
|
||||
if !v.Equal(ast.Boolean(false)) {
|
||||
if v.IsGround() {
|
||||
ctx.traceSuccess(expr)
|
||||
return iter(ctx)
|
||||
}
|
||||
}
|
||||
@@ -573,34 +571,96 @@ func evalExpr(ctx *Context, iter Iterator) error {
|
||||
}
|
||||
}
|
||||
|
||||
func evalRef(ctx *Context, ref ast.Ref, iter Iterator) error {
|
||||
// If this reference refers to a local variable, evaluate against the binding.
|
||||
// Otherwise, evaluate against the database.
|
||||
if !ref[0].Equal(ast.DefaultRootDocument) {
|
||||
v := ctx.Binding(ref[0].Value)
|
||||
if v == nil {
|
||||
return unboundGlobalVarErr(ref)
|
||||
// evalRef evaluates the ast.Ref ref and calls the Iterator iter once for each
|
||||
// instance of ref that would be defined. If an error occurs during the evaluation
|
||||
// process, the return value is non-nil. Also, if iter returns an error, the return
|
||||
// value is non-nil.
|
||||
func evalRef(ctx *Context, ref, path ast.Ref, iter Iterator) error {
|
||||
|
||||
if len(ref) == 0 {
|
||||
// If this reference refers to a local variable, evaluate against the binding.
|
||||
// Otherwise, evaluate against the database.
|
||||
if !path[0].Equal(ast.DefaultRootDocument) {
|
||||
v := ctx.Binding(path[0].Value)
|
||||
if v == nil {
|
||||
return unboundGlobalVarErr(path)
|
||||
}
|
||||
return evalRefRuleResult(ctx, path, path[1:], v, iter)
|
||||
}
|
||||
return evalRefRuleResult(ctx, ref, ref[1:], v, iter)
|
||||
return evalRefRec(ctx, ast.Ref{path[0]}, path[1:], iter)
|
||||
}
|
||||
|
||||
return evalRefRec(ctx, ast.Ref{ref[0]}, ref[1:], iter)
|
||||
head, tail := ref[0], ref[1:]
|
||||
n, ok := head.Value.(ast.Ref)
|
||||
if !ok {
|
||||
path = append(path, head)
|
||||
return evalRef(ctx, tail, path, iter)
|
||||
}
|
||||
|
||||
return evalRef(ctx, n, ast.Ref{}, func(ctx *Context) error {
|
||||
if b := ctx.Binding(n); b == nil {
|
||||
p := PlugValue(n, ctx).(ast.Ref)
|
||||
v, err := lookupValue(ctx.DataStore, p)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ctx = ctx.BindValue(n, v)
|
||||
}
|
||||
tmp := append(path, head)
|
||||
return evalRef(ctx, tail, tmp, iter)
|
||||
})
|
||||
}
|
||||
|
||||
func evalRefRec(ctx *Context, path, tail ast.Ref, iter Iterator) error {
|
||||
|
||||
if len(tail) == 0 {
|
||||
return evalRefRecFinish(ctx, path, iter)
|
||||
}
|
||||
|
||||
if tail[0].IsGround() {
|
||||
return evalRefRecGround(ctx, path, tail, iter)
|
||||
}
|
||||
|
||||
return evalRefRecNonGround(ctx, path, tail, iter)
|
||||
}
|
||||
|
||||
func evalRefRecEnumColl(ctx *Context, path, tail ast.Ref, iter Iterator) error {
|
||||
func evalRefRecFinish(ctx *Context, path ast.Ref, iter Iterator) error {
|
||||
ok, err := lookupExists(ctx.DataStore, path)
|
||||
if err == nil && ok {
|
||||
return iter(ctx)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func evalRefRecGround(ctx *Context, path, tail ast.Ref, iter Iterator) error {
|
||||
// Check if the node exists. If the node does not exist, stop.
|
||||
// If the node exists and is a rule, evaluate the rule to produce a virtual doc.
|
||||
// Otherwise, process the rest of the reference.
|
||||
path = append(path, PlugTerm(tail[0], ctx))
|
||||
rules, err := lookupRule(ctx.DataStore, path)
|
||||
if err != nil {
|
||||
if storage.IsNotFound(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
if rules != nil {
|
||||
ref := append(path, tail[1:]...)
|
||||
return evalRefRule(ctx, ref, path, rules, iter)
|
||||
}
|
||||
return evalRefRec(ctx, path, tail[1:], iter)
|
||||
}
|
||||
|
||||
func evalRefRecNonGround(ctx *Context, path, tail ast.Ref, iter Iterator) error {
|
||||
// Check if the variable has a binding.
|
||||
// If there is a binding, process the rest of the reference normally.
|
||||
// If there is no binding, enumerate the collection referred to by the path.
|
||||
plugged := PlugTerm(tail[0], ctx)
|
||||
if plugged.IsGround() {
|
||||
path = append(path, plugged)
|
||||
return evalRefRec(ctx, path, tail[1:], iter)
|
||||
}
|
||||
return evalRefRecWalkColl(ctx, path, tail, iter)
|
||||
}
|
||||
|
||||
func evalRefRecWalkColl(ctx *Context, path, tail ast.Ref, iter Iterator) error {
|
||||
|
||||
node, err := ctx.DataStore.GetRef(path)
|
||||
if err != nil {
|
||||
@@ -641,45 +701,6 @@ func evalRefRecEnumColl(ctx *Context, path, tail ast.Ref, iter Iterator) error {
|
||||
}
|
||||
}
|
||||
|
||||
func evalRefRecFinish(ctx *Context, path ast.Ref, iter Iterator) error {
|
||||
ok, err := lookupExists(ctx.DataStore, path)
|
||||
if err == nil && ok {
|
||||
return iter(ctx)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func evalRefRecGround(ctx *Context, path, tail ast.Ref, iter Iterator) error {
|
||||
// Check if the node exists. If the node does not exist, stop.
|
||||
// If the node exists and is a rule, evaluate the rule to produce a virtual doc.
|
||||
// Otherwise, process the rest of the reference.
|
||||
path = append(path, tail[0])
|
||||
rules, err := lookupRule(ctx.DataStore, path)
|
||||
if err != nil {
|
||||
if storage.IsNotFound(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
if rules != nil {
|
||||
ref := append(path, tail[1:]...)
|
||||
return evalRefRule(ctx, ref, path, rules, iter)
|
||||
}
|
||||
return evalRefRec(ctx, path, tail[1:], iter)
|
||||
}
|
||||
|
||||
func evalRefRecNonGround(ctx *Context, path, tail ast.Ref, iter Iterator) error {
|
||||
// Check if the variable has a binding.
|
||||
// If there is a binding, process the rest of the reference normally.
|
||||
// If there is no binding, enumerate the collection referred to by the path.
|
||||
plugged := PlugTerm(tail[0], ctx)
|
||||
if plugged.IsGround() {
|
||||
path = append(path, plugged)
|
||||
return evalRefRec(ctx, path, tail[1:], iter)
|
||||
}
|
||||
return evalRefRecEnumColl(ctx, path, tail, iter)
|
||||
}
|
||||
|
||||
func evalRefRule(ctx *Context, ref ast.Ref, path ast.Ref, rules []*ast.Rule, iter Iterator) error {
|
||||
|
||||
suffix := ref[len(path):]
|
||||
@@ -1049,7 +1070,7 @@ func evalTermsRec(ctx *Context, iter Iterator, ts []*ast.Term) error {
|
||||
|
||||
switch head := head.Value.(type) {
|
||||
case ast.Ref:
|
||||
return evalRef(ctx, head, func(ctx *Context) error {
|
||||
return evalRef(ctx, head, ast.Ref{}, func(ctx *Context) error {
|
||||
return evalTermsRec(ctx, iter, tail)
|
||||
})
|
||||
case ast.Array:
|
||||
@@ -1075,7 +1096,7 @@ func evalTermsRecArray(ctx *Context, arr ast.Array, idx int, iter Iterator) erro
|
||||
}
|
||||
switch v := arr[idx].Value.(type) {
|
||||
case ast.Ref:
|
||||
return evalRef(ctx, v, func(ctx *Context) error {
|
||||
return evalRef(ctx, v, ast.Ref{}, func(ctx *Context) error {
|
||||
return evalTermsRecArray(ctx, arr, idx+1, iter)
|
||||
})
|
||||
case ast.Array:
|
||||
@@ -1101,10 +1122,10 @@ func evalTermsRecObject(ctx *Context, obj ast.Object, idx int, iter Iterator) er
|
||||
}
|
||||
switch k := obj[idx][0].Value.(type) {
|
||||
case ast.Ref:
|
||||
return evalRef(ctx, k, func(ctx *Context) error {
|
||||
return evalRef(ctx, k, ast.Ref{}, func(ctx *Context) error {
|
||||
switch v := obj[idx][1].Value.(type) {
|
||||
case ast.Ref:
|
||||
return evalRef(ctx, v, func(ctx *Context) error {
|
||||
return evalRef(ctx, v, ast.Ref{}, func(ctx *Context) error {
|
||||
return evalTermsRecObject(ctx, obj, idx+1, iter)
|
||||
})
|
||||
case ast.Array:
|
||||
@@ -1126,7 +1147,7 @@ func evalTermsRecObject(ctx *Context, obj ast.Object, idx int, iter Iterator) er
|
||||
default:
|
||||
switch v := obj[idx][1].Value.(type) {
|
||||
case ast.Ref:
|
||||
return evalRef(ctx, v, func(ctx *Context) error {
|
||||
return evalRef(ctx, v, ast.Ref{}, func(ctx *Context) error {
|
||||
return evalTermsRecObject(ctx, obj, idx+1, iter)
|
||||
})
|
||||
case ast.Array:
|
||||
@@ -1166,14 +1187,14 @@ func indexAvailable(ctx *Context, expr *ast.Expr) bool {
|
||||
a := ts[1].Value
|
||||
b := ts[2].Value
|
||||
|
||||
_, isRefA := a.(ast.Ref)
|
||||
_, isRefB := b.(ast.Ref)
|
||||
aRef, isRefA := a.(ast.Ref)
|
||||
bRef, isRefB := b.(ast.Ref)
|
||||
|
||||
if isRefA && !a.IsGround() {
|
||||
if isRefA && !a.IsGround() && !aRef.IsNested() {
|
||||
return b.IsGround() || isRefB
|
||||
}
|
||||
|
||||
if isRefB && !b.IsGround() {
|
||||
if isRefB && !b.IsGround() && !bRef.IsNested() {
|
||||
return a.IsGround() || isRefA
|
||||
}
|
||||
|
||||
@@ -1185,7 +1206,7 @@ func indexAvailable(ctx *Context, expr *ast.Expr) bool {
|
||||
// built on the fly.
|
||||
func indexBuildLazy(ctx *Context, ref ast.Ref) (bool, error) {
|
||||
|
||||
if ref.IsGround() {
|
||||
if ref.IsGround() || ref.IsNested() {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
@@ -1250,6 +1271,14 @@ func lookupRule(ds *storage.DataStore, ref ast.Ref) ([]*ast.Rule, error) {
|
||||
}
|
||||
}
|
||||
|
||||
func lookupValue(ds *storage.DataStore, ref ast.Ref) (ast.Value, error) {
|
||||
r, err := ds.GetRef(ref)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return ast.InterfaceToValue(r)
|
||||
}
|
||||
|
||||
func topDownQueryCompleteDoc(params *QueryParams, rules []*ast.Rule) (interface{}, error) {
|
||||
|
||||
var result ast.Value
|
||||
@@ -1354,37 +1383,3 @@ func topDownQueryPartialSetDoc(params *QueryParams, rules []*ast.Rule) (interfac
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// walkValue invokes the iterator for each AST value contained inside the supplied AST value.
|
||||
// If walkValue is called with a scalar, the iterator is invoked exactly once.
|
||||
// If walkValue is called with a reference, the iterator is invoked for each element in the reference.
|
||||
func walkValue(value ast.Value, iter func(ast.Value) bool) bool {
|
||||
switch value := value.(type) {
|
||||
case ast.Ref:
|
||||
for _, x := range value {
|
||||
if walkValue(x.Value, iter) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
case ast.Array:
|
||||
for _, x := range value {
|
||||
if walkValue(x.Value, iter) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
case ast.Object:
|
||||
for _, i := range value {
|
||||
if walkValue(i[0].Value, iter) {
|
||||
return true
|
||||
}
|
||||
if walkValue(i[1].Value, iter) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
default:
|
||||
return iter(value)
|
||||
}
|
||||
}
|
||||
|
||||
+105
-26
@@ -23,32 +23,32 @@ func TestEvalRef(t *testing.T) {
|
||||
expected interface{}
|
||||
}{
|
||||
{"data.c[i][j]", `[
|
||||
{"i": 0, "j": "x"},
|
||||
{"i": 0, "j": "y"},
|
||||
{"i": 0, "j": "z"}
|
||||
]`},
|
||||
{"i": 0, "j": "x"},
|
||||
{"i": 0, "j": "y"},
|
||||
{"i": 0, "j": "z"}
|
||||
]`},
|
||||
{"data.c[i][j][k]", `[
|
||||
{"i": 0, "j": "x", "k": 0},
|
||||
{"i": 0, "j": "x", "k": 1},
|
||||
{"i": 0, "j": "x", "k": 2},
|
||||
{"i": 0, "j": "y", "k": 0},
|
||||
{"i": 0, "j": "y", "k": 1},
|
||||
{"i": 0, "j": "z", "k": "p"},
|
||||
{"i": 0, "j": "z", "k": "q"}
|
||||
]`},
|
||||
{"i": 0, "j": "x", "k": 0},
|
||||
{"i": 0, "j": "x", "k": 1},
|
||||
{"i": 0, "j": "x", "k": 2},
|
||||
{"i": 0, "j": "y", "k": 0},
|
||||
{"i": 0, "j": "y", "k": 1},
|
||||
{"i": 0, "j": "z", "k": "p"},
|
||||
{"i": 0, "j": "z", "k": "q"}
|
||||
]`},
|
||||
{"data.d[x][y]", `[
|
||||
{"x": "e", "y": 0},
|
||||
{"x": "e", "y": 1}
|
||||
]`},
|
||||
{"x": "e", "y": 0},
|
||||
{"x": "e", "y": 1}
|
||||
]`},
|
||||
{`data.c[i]["x"][k]`, `[
|
||||
{"i": 0, "k": 0},
|
||||
{"i": 0, "k": 1},
|
||||
{"i": 0, "k": 2}
|
||||
]`},
|
||||
{"i": 0, "k": 0},
|
||||
{"i": 0, "k": 1},
|
||||
{"i": 0, "k": 2}
|
||||
]`},
|
||||
{"data.c[i][j][i]", `[
|
||||
{"i": 0, "j": "x"},
|
||||
{"i": 0, "j": "y"}
|
||||
]`},
|
||||
{"i": 0, "j": "x"},
|
||||
{"i": 0, "j": "y"}
|
||||
]`},
|
||||
{`data.c[i]["deadbeef"][k]`, nil},
|
||||
{`data.c[999]`, nil},
|
||||
}
|
||||
@@ -66,7 +66,7 @@ func TestEvalRef(t *testing.T) {
|
||||
switch e := tc.expected.(type) {
|
||||
case nil:
|
||||
var tmp *Context
|
||||
err := evalRef(ctx, ast.MustParseRef(tc.ref), func(ctx *Context) error {
|
||||
err := evalRef(ctx, ast.MustParseRef(tc.ref), ast.Ref{}, func(ctx *Context) error {
|
||||
tmp = ctx
|
||||
return nil
|
||||
})
|
||||
@@ -79,7 +79,7 @@ func TestEvalRef(t *testing.T) {
|
||||
}
|
||||
case string:
|
||||
expected := loadExpectedBindings(e)
|
||||
err := evalRef(ctx, ast.MustParseRef(tc.ref), func(ctx *Context) error {
|
||||
err := evalRef(ctx, ast.MustParseRef(tc.ref), ast.Ref{}, func(ctx *Context) error {
|
||||
if len(expected) > 0 {
|
||||
for j, exp := range expected {
|
||||
if exp.Equal(ctx.Locals) {
|
||||
@@ -135,6 +135,11 @@ func TestEvalTerms(t *testing.T) {
|
||||
{"x": "e", "y": 1}
|
||||
]`},
|
||||
{"data.d[x][y] = data.z[i]", `[]`},
|
||||
{"data.a[data.a[i]] = 3", `[
|
||||
{"i": 0, "data.a[i]": 1},
|
||||
{"i": 1, "data.a[i]": 2},
|
||||
{"i": 2, "data.a[i]": 3}
|
||||
]`},
|
||||
}
|
||||
|
||||
data := loadSmallTestData()
|
||||
@@ -212,6 +217,20 @@ func TestPlugValue(t *testing.T) {
|
||||
if !expected.Equal(r2) {
|
||||
t.Errorf("Expected %v but got %v", expected, r2)
|
||||
}
|
||||
|
||||
n := ast.MustParseTerm("a.b[x.y[i]]").Value
|
||||
|
||||
ctx3 := &Context{Locals: storage.NewBindings(), Globals: storage.NewBindings()}
|
||||
ctx3 = ctx3.BindVar(ast.Var("i"), ast.Number(1))
|
||||
ctx3 = ctx3.BindValue(ast.MustParseTerm("x.y[i]").Value, ast.Number(1))
|
||||
|
||||
expected = ast.MustParseTerm("a.b[1]").Value
|
||||
|
||||
r3 := PlugValue(n, ctx3)
|
||||
|
||||
if !expected.Equal(r3) {
|
||||
t.Errorf("Expected %v but got: %v", expected, r3)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTopDownCompleteDoc(t *testing.T) {
|
||||
@@ -512,6 +531,60 @@ func TestTopDownVirtualDocs(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestTopDownNestedReferences(t *testing.T) {
|
||||
tests := []struct {
|
||||
note string
|
||||
rules []string
|
||||
expected interface{}
|
||||
}{
|
||||
// nested base document references
|
||||
{"ground ref", []string{"p :- a[h[0][0]] = 2"}, "true"},
|
||||
{"non-ground ref", []string{"p[x] :- x = a[h[i][j]]"}, "[2,3,4,3,4]"},
|
||||
{"two deep", []string{"p[x] :- x = a[a[a[i]]]"}, "[3,4]"},
|
||||
{"two deep", []string{"p[x] :- x = a[h[i][a[j]]]"}, "[3,4,4]"},
|
||||
{"two deep repeated var", []string{"p[x] :- x = a[h[i][a[i]]]"}, "[3]"},
|
||||
{"no suffix", []string{"p :- 4 = a[three]"}, "true"},
|
||||
{"var ref", []string{"p[y] :- x = [1,2,3], y = a[x[_]]"}, "[2,3,4]"},
|
||||
{"undefined", []string{"p :- a[three.deadbeef] = x"}, ""},
|
||||
|
||||
// nested virtual document references
|
||||
{"vdoc ref: complete", []string{"p[x] :- x = a[q[_]]", "q = [2,3] :- true"}, "[3,4]"},
|
||||
{"vdoc ref: complete: ground", []string{"p[x] :- x = a[q[1]]", "q = [2,3] :- true"}, "[4]"},
|
||||
{"vdoc ref: complete: no suffix", []string{"p :- 2 = a[q]", "q = 1 :- true"}, "true"},
|
||||
{"vdoc ref: partial object", []string{
|
||||
"p[x] :- x = a[q[_]]",
|
||||
`q[k] = v :- o = {"a": 2, "b": 3, "c": 100}, o[k] = v`},
|
||||
"[3,4]"},
|
||||
{"vdoc ref: partial object: ground", []string{
|
||||
`p[x] :- x = a[q["b"]]`,
|
||||
`q[k] = v :- o = {"a": 2, "b": 3, "c": 100}, o[k] = v`},
|
||||
"[4]"},
|
||||
|
||||
// mixed cases
|
||||
{"vdoc ref: complete: nested bdoc ref", []string{
|
||||
"p[x] :- x = a[q[b[_]]]",
|
||||
`q = {"hello": 1, "goodbye": 3, "deadbeef": 1000} :- true`}, "[2,4]"},
|
||||
{"vdoc ref: partial object: nested bdoc ref", []string{
|
||||
"p[x] :- x = a[q[b[_]]]",
|
||||
// bind to value
|
||||
`q[k] = v :- o = {"hello": 1, "goodbye": 3, "deadbeef": 1000}, o[k] = v`}, "[2,4]"},
|
||||
{"vdoc ref: partial object: nested bdoc ref-2", []string{
|
||||
"p[x] :- x = a[q[d.e[_]]]",
|
||||
// bind to reference
|
||||
`q[k] = v :- strings[k] = v`}, "[3,4]"},
|
||||
{"vdoc ref: multiple", []string{
|
||||
"p[x] :- x = q[a[_]].v[r[a[_]]]",
|
||||
`q = [{"v": {}}, {"v": [0,0,1,2]}, {"v": [0,0,3,4]}, {"v": [0,0]}, {}] :- true`,
|
||||
"r = [1,2,3,4] :- true"}, "[1,2,3,4]"},
|
||||
}
|
||||
|
||||
data := loadSmallTestData()
|
||||
|
||||
for i, tc := range tests {
|
||||
runTopDownTestCase(t, data, i, tc.note, tc.rules, tc.expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTopDownVarReferences(t *testing.T) {
|
||||
|
||||
tests := []struct {
|
||||
@@ -885,9 +958,9 @@ func loadExpectedBindings(input string) []*storage.Bindings {
|
||||
for k, v := range bindings {
|
||||
switch v := v.(type) {
|
||||
case string:
|
||||
buf.Put(ast.Var(k), ast.String(v))
|
||||
buf.Put(ast.MustParseTerm(k).Value, ast.String(v))
|
||||
case float64:
|
||||
buf.Put(ast.Var(k), ast.Number(v))
|
||||
buf.Put(ast.MustParseTerm(k).Value, ast.Number(v))
|
||||
default:
|
||||
panic("unreachable")
|
||||
}
|
||||
@@ -968,6 +1041,12 @@ func loadSmallTestData() map[string]interface{} {
|
||||
"d": null
|
||||
}
|
||||
],
|
||||
"strings": {
|
||||
"foo": 1,
|
||||
"bar": 2,
|
||||
"baz": 3
|
||||
},
|
||||
"three": 3,
|
||||
"m": []
|
||||
}`), &data)
|
||||
if err != nil {
|
||||
|
||||
Reference in New Issue
Block a user