Merge pull request #52 from tsandall/nested-refs

Nested References
This commit is contained in:
Torin Sandall
2016-06-28 14:39:01 -07:00
committed by GitHub
10 changed files with 414 additions and 167 deletions
+21 -30
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -252,7 +252,7 @@ RefDot <- "." val:Var {
return str, nil
}
RefBracket <- "[" val:(Scalar / Var) "]" {
RefBracket <- "[" val:(Ref / Scalar / Var) "]" {
return val, nil
}
+51 -2
View File
@@ -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 {
+32
View File
@@ -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)
+33
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 {