Refactor set and object types in ast package

Previously, sets and objects were not interfaces and as a result,
callers were relying on the underlying structure for operations such as
iteration.

These changes refactor the ast package to expose sets and objects as
interfaces so that we can change the underlying data structures without
affecting callers.
This commit is contained in:
Torin Sandall
2017-12-14 15:27:27 -08:00
parent 10182050d6
commit 61823a2920
27 changed files with 1242 additions and 1095 deletions
+20 -20
View File
@@ -304,7 +304,7 @@ func unify2(env *TypeEnv, a *Term, typeA types.Type, b *Term, typeB types.Type)
switch bv := b.Value.(type) {
case Object:
cv := av.Intersect(bv)
if len(av) == len(bv) && len(bv) == len(cv) {
if av.Len() == bv.Len() && bv.Len() == len(cv) {
for i := range cv {
if !unify2(env, cv[i][1], env.Get(cv[i][1]), cv[i][2], env.Get(cv[i][2])) {
return false
@@ -350,9 +350,9 @@ func unify1(env *TypeEnv, term *Term, tpe types.Type) bool {
return unify1Object(env, v, tpe)
case types.Any:
if types.Compare(tpe, types.A) == 0 {
for i := range v {
unify1(env, v[i][1], types.A)
}
v.Foreach(func(_, value *Term) {
unify1(env, value, types.A)
})
return true
}
unifies := false
@@ -362,15 +362,14 @@ func unify1(env *TypeEnv, term *Term, tpe types.Type) bool {
return unifies
}
return false
case *Set:
case Set:
switch tpe := tpe.(type) {
case *types.Set:
return unify1Set(env, v, tpe)
case types.Any:
if types.Compare(tpe, types.A) == 0 {
v.Iter(func(elem *Term) bool {
v.Foreach(func(elem *Term) {
unify1(env, elem, types.A)
return true
})
return true
}
@@ -410,32 +409,33 @@ func unify1Array(env *TypeEnv, val Array, tpe *types.Array) bool {
}
func unify1Object(env *TypeEnv, val Object, tpe *types.Object) bool {
if len(val) != len(tpe.Keys()) && tpe.DynamicValue() == nil {
if val.Len() != len(tpe.Keys()) && tpe.DynamicValue() == nil {
return false
}
for i := range val {
if IsConstant(val[i][0].Value) {
if child := selectConstant(tpe, val[i][0]); child != nil {
if !unify1(env, val[i][1], child) {
return false
stop := val.Until(func(k, v *Term) bool {
if IsConstant(k.Value) {
if child := selectConstant(tpe, k); child != nil {
if !unify1(env, v, child) {
return true
}
} else {
return false
return true
}
} else {
// Inferring type of value under dynamic key would involve unioning
// with all property values of tpe whose keys unify. For now, type
// these values as Any. We can investigate stricter inference in
// the future.
unify1(env, val[i][1], types.A)
unify1(env, v, types.A)
}
}
return true
return false
})
return !stop
}
func unify1Set(env *TypeEnv, val *Set, tpe *types.Set) bool {
func unify1Set(env *TypeEnv, val Set, tpe *types.Set) bool {
of := types.Values(tpe)
return !val.Iter(func(elem *Term) bool {
return !val.Until(func(elem *Term) bool {
return !unify1(env, elem, of)
})
}
@@ -604,7 +604,7 @@ func (rc *refChecker) checkRefLeaf(tpe types.Type, ref Ref, idx int) *Error {
}
}
case Array, Object, *Set:
case Array, Object, Set:
// Composite references operands may only be used with a set.
if !unifies(tpe, types.NewSet(types.A)) {
return newRefErrInvalid(ref[0].Location, ref, idx, tpe, types.NewSet(types.A), nil)
+5 -37
View File
@@ -7,7 +7,6 @@ package ast
import (
"encoding/json"
"fmt"
"sort"
"github.com/open-policy-agent/opa/util"
)
@@ -111,41 +110,10 @@ func Compare(a, b interface{}) int {
return termSliceCompare(a, b)
case Object:
b := b.(Object)
keysA := a.Keys()
keysB := b.Keys()
sort.Sort(termSlice(keysA))
sort.Sort(termSlice(keysB))
minLen := len(a)
if len(b) < len(a) {
minLen = len(b)
}
for i := 0; i < minLen; i++ {
keysCmp := Compare(keysA[i], keysB[i])
if keysCmp < 0 {
return -1
}
if keysCmp > 0 {
return 1
}
valA := a.Get(keysA[i])
valB := b.Get(keysB[i])
valCmp := Compare(valA, valB)
if valCmp != 0 {
return valCmp
}
}
if len(a) < len(b) {
return -1
}
if len(b) < len(a) {
return 1
}
return 0
case *Set:
b := b.(*Set)
sort.Sort(termSlice(*a))
sort.Sort(termSlice(*b))
return termSliceCompare(*a, *b)
return a.Compare(b)
case Set:
b := b.(Set)
return a.Compare(b)
case *ArrayComprehension:
b := b.(*ArrayComprehension)
if cmp := Compare(a.Term, b.Term); cmp != 0 {
@@ -222,7 +190,7 @@ func sortOrder(x interface{}) int {
return 6
case Object:
return 7
case *Set:
case Set:
return 8
case *ArrayComprehension:
return 9
+16 -19
View File
@@ -1540,7 +1540,7 @@ func resolveRef(globals map[Var]Ref, ref Ref) Ref {
} else {
r = append(r, x)
}
case Ref, Array, Object, *Set:
case Ref, Array, Object, Set:
r = append(r, resolveRefsInTerm(globals, x))
default:
r = append(r, x)
@@ -1607,14 +1607,12 @@ func resolveRefsInTerm(globals map[Var]Ref, term *Term) *Term {
cpy.Value = fqn
return &cpy
case Object:
o := Object{}
for _, i := range v {
k := resolveRefsInTerm(globals, i[0])
v := resolveRefsInTerm(globals, i[1])
o = append(o, Item(k, v))
}
cpy := *term
cpy.Value = o
cpy.Value, _ = v.Map(func(k, v *Term) (*Term, *Term, error) {
k = resolveRefsInTerm(globals, k)
v = resolveRefsInTerm(globals, v)
return k, v, nil
})
return &cpy
case Array:
a := Array{}
@@ -1625,12 +1623,10 @@ func resolveRefsInTerm(globals map[Var]Ref, term *Term) *Term {
cpy := *term
cpy.Value = a
return &cpy
case *Set:
s := &Set{}
for _, e := range *v {
x := resolveRefsInTerm(globals, e)
s.Add(x)
}
case Set:
s, _ := v.Map(func(e *Term) (*Term, error) {
return resolveRefsInTerm(globals, e), nil
})
cpy := *term
cpy.Value = s
return &cpy
@@ -1796,12 +1792,13 @@ func rewriteDynamicsOne(original *Expr, f *equalityFactory, term *Term, extras [
}
return extras, term
case Object:
for i := 0; i < len(v); i++ {
extras, v[i][0] = rewriteDynamicsOne(original, f, v[i][0], extras)
extras, v[i][1] = rewriteDynamicsOne(original, f, v[i][1], extras)
}
term.Value, _ = v.Map(func(k, v *Term) (*Term, *Term, error) {
extras, k = rewriteDynamicsOne(original, f, k, extras)
extras, v = rewriteDynamicsOne(original, f, v, extras)
return k, v, nil
})
return extras, term
case *Set:
case Set:
v, _ = v.Map(func(term *Term) (*Term, error) {
extras, term = rewriteDynamicsOne(original, f, term, extras)
return term, nil
+2 -2
View File
@@ -530,9 +530,9 @@ func TestCompilerCheckSafetyBodyErrors(t *testing.T) {
// Build slice of expected error messages.
expected := []string{}
MustParseTerm(tc.expected).Value.(*Set).Iter(func(x *Term) bool {
MustParseTerm(tc.expected).Value.(Set).Iter(func(x *Term) error {
expected = append(expected, makeErrMsg(string(x.Value.(Var))))
return false
return nil
})
sort.Strings(expected)
+11 -12
View File
@@ -63,20 +63,20 @@ func (env *TypeEnv) Get(x interface{}) types.Type {
static := []*types.StaticProperty{}
var dynamic *types.DynamicProperty
for _, pair := range x {
if IsConstant(pair[0].Value) {
k, err := JSON(pair[0].Value)
x.Foreach(func(k, v *Term) {
if IsConstant(k.Value) {
kjson, err := JSON(k.Value)
if err != nil {
panic("unreachable")
}
tpe := env.Get(pair[1].Value)
static = append(static, types.NewStaticProperty(k, tpe))
tpe := env.Get(v)
static = append(static, types.NewStaticProperty(kjson, tpe))
} else {
typeK := env.Get(pair[0].Value)
typeV := env.Get(pair[1].Value)
typeK := env.Get(k.Value)
typeV := env.Get(v.Value)
dynamic = types.NewDynamicProperty(typeK, typeV)
}
}
})
if len(static) == 0 && dynamic == nil {
dynamic = types.NewDynamicProperty(types.A, types.A)
@@ -84,12 +84,11 @@ func (env *TypeEnv) Get(x interface{}) types.Type {
return types.NewObject(static, dynamic)
case *Set:
case Set:
var tpe types.Type
x.Iter(func(elem *Term) bool {
x.Foreach(func(elem *Term) {
other := env.Get(elem.Value)
tpe = types.Or(tpe, other)
return false
})
if tpe == nil {
tpe = types.A
@@ -319,7 +318,7 @@ func selectRef(tpe types.Type, ref Ref) types.Type {
head, tail := ref[0], ref[1:]
switch head.Value.(type) {
case Var, Ref, Array, Object, *Set:
case Var, Ref, Array, Object, Set:
return selectRef(types.Values(tpe), tail)
default:
return selectRef(selectConstant(tpe, head), tail)
+1 -1
View File
@@ -237,7 +237,7 @@ func (i *baseDocEqIndex) getRefAndValueFromTerms(a, b *Term) (Ref, Value, bool)
}
switch x.(type) {
// No nested structures or values that require evaluation (other than var).
case Array, Object, *Set, *ArrayComprehension, *ObjectComprehension, *SetComprehension, Ref:
case Array, Object, Set, *ArrayComprehension, *ObjectComprehension, *SetComprehension, Ref:
stop = true
}
return stop
+708 -709
View File
File diff suppressed because it is too large Load Diff
+2 -2
View File
@@ -23,7 +23,7 @@ var InputRootDocument = VarTerm("input")
// RootDocumentNames contains the names of top-level documents that can be
// referred to in modules and queries.
var RootDocumentNames = &Set{
var RootDocumentNames Set = &set{
DefaultRootDocument,
InputRootDocument,
}
@@ -40,7 +40,7 @@ var InputRootRef = Ref{InputRootDocument}
// RootDocumentRefs contains the prefixes of top-level documents that all
// non-local references start with.
var RootDocumentRefs = &Set{
var RootDocumentRefs Set = &set{
NewTerm(DefaultRootRef),
NewTerm(InputRootRef),
}
+8 -9
View File
@@ -45,27 +45,26 @@ func ifacesToBody(i interface{}, a ...interface{}) Body {
}
func makeObject(head interface{}, tail interface{}, loc *Location) (*Term, error) {
obj := ObjectTerm()
obj.Location = loc
term := ObjectTerm()
term.Location = loc
// Empty object.
if head == nil {
return obj, nil
return term, nil
}
// Object definition above describes the "head" structure. We only care about the "Key" and "Term" elements.
headSlice := head.([]interface{})
obj.Value = append(obj.Value.(Object), Item(headSlice[0].(*Term), headSlice[len(headSlice) - 1].(*Term)))
obj := term.Value.(Object)
obj.Insert(headSlice[0].(*Term), headSlice[len(headSlice)-1].(*Term))
// Non-empty object, remaining key/value pairs.
tailSlice := tail.([]interface{})
for _, v := range tailSlice {
s := v.([]interface{})
// Object definition above describes the "tail" structure. We only care about the "Key" and "Term" elements.
obj.Value = append(obj.Value.(Object), Item(s[3].(*Term), s[len(s) - 1].(*Term)))
obj.Insert(s[3].(*Term), s[len(s)-1].(*Term))
}
return obj, nil
return term, nil
}
func makeArray(head interface{}, tail interface{}, loc *Location) (*Term, error) {
@@ -572,7 +571,7 @@ SetNonEmpty <- '{' _ head:Term tail:(_ ',' _ Term)* _ ','? _ '}' {
set := SetTerm()
set.Location = currentLocation(c)
val := set.Value.(*Set)
val := set.Value.(Set)
val.Add(head.(*Term))
tailSlice := tail.([]interface{})
+317 -137
View File
@@ -10,6 +10,7 @@ import (
"fmt"
"math/rand"
"regexp"
"sort"
"strconv"
"strings"
"time"
@@ -117,7 +118,7 @@ func InterfaceToValue(x interface{}) (Value, error) {
}
return r, nil
case map[string]interface{}:
r := Object{}
r := &object{}
for k, v := range x {
k, err := InterfaceToValue(k)
if err != nil {
@@ -127,11 +128,11 @@ func InterfaceToValue(x interface{}) (Value, error) {
if err != nil {
return nil, err
}
r = append(r, Item(&Term{Value: k}, &Term{Value: v}))
r.Insert(NewTerm(k), NewTerm(v))
}
return r, nil
case map[string]string:
r := Object{}
r := &object{}
for k, v := range x {
k, err := InterfaceToValue(k)
if err != nil {
@@ -141,7 +142,7 @@ func InterfaceToValue(x interface{}) (Value, error) {
if err != nil {
return nil, err
}
r = append(r, Item(&Term{Value: k}, &Term{Value: v}))
r.Insert(NewTerm(k), NewTerm(v))
}
return r, nil
default:
@@ -190,30 +191,38 @@ func ValueToInterface(v Value, resolver Resolver) (interface{}, error) {
return buf, nil
case Object:
buf := map[string]interface{}{}
for _, x := range v {
k, err := ValueToInterface(x[0].Value, resolver)
err := v.Iter(func(k, v *Term) error {
ki, err := ValueToInterface(k.Value, resolver)
if err != nil {
return nil, err
return err
}
asStr, stringKey := k.(string)
asStr, stringKey := ki.(string)
if !stringKey {
return nil, fmt.Errorf("object value has non-string key (%T)", k)
return fmt.Errorf("object value has non-string key (%T)", ki)
}
v, err := ValueToInterface(x[1].Value, resolver)
vi, err := ValueToInterface(v.Value, resolver)
if err != nil {
return nil, err
return err
}
buf[asStr] = v
buf[asStr] = vi
return nil
})
if err != nil {
return nil, err
}
return buf, nil
case *Set:
case Set:
buf := []interface{}{}
for _, x := range *v {
err := v.Iter(func(x *Term) error {
x1, err := ValueToInterface(x.Value, resolver)
if err != nil {
return nil, err
return err
}
buf = append(buf, x1)
return nil
})
if err != nil {
return nil, err
}
return buf, nil
case Ref:
@@ -275,7 +284,7 @@ func (term *Term) Copy() *Term {
cpy.Value = v.Copy()
case Array:
cpy.Value = v.Copy()
case *Set:
case Set:
cpy.Value = v.Copy()
case Object:
cpy.Value = v.Copy()
@@ -978,11 +987,36 @@ func (arr Array) String() string {
}
// Set represents a set as defined by the language.
type Set []*Term
type Set interface {
Value
Len() int
Copy() Set
Diff(Set) Set
Intersect(Set) Set
Union(Set) Set
Add(*Term)
Iter(func(*Term) error) error
Until(func(*Term) bool) bool
Foreach(func(*Term))
Contains(*Term) bool
Map(func(*Term) (*Term, error)) (Set, error)
Reduce(*Term, func(*Term, *Term) (*Term, error)) (*Term, error)
}
// NewSet returns a new Set containing t.
func NewSet(t ...*Term) Set {
s := &set{}
for i := range t {
s.Add(t[i])
}
return s
}
type set []*Term
// SetTerm returns a new Term representing a set containing terms t.
func SetTerm(t ...*Term) *Term {
s := &Set{}
s := &set{}
for i := range t {
s.Add(t[i])
}
@@ -992,22 +1026,22 @@ func SetTerm(t ...*Term) *Term {
}
// Copy returns a deep copy of s.
func (s *Set) Copy() *Set {
cpy := Set(termSliceCopy(*s))
return &cpy
func (s *set) Copy() Set {
cpy := set(termSliceCopy(*s))
return Set(&cpy)
}
// IsGround returns true if all terms in s are ground.
func (s *Set) IsGround() bool {
func (s *set) IsGround() bool {
return termSliceIsGround(*s)
}
// Hash returns a hash code for s.
func (s *Set) Hash() int {
func (s *set) Hash() int {
return termSliceHash(*s)
}
func (s *Set) String() string {
func (s *set) String() string {
sl := *s
@@ -1024,19 +1058,24 @@ func (s *Set) String() string {
return "{" + strings.Join(buf, ", ") + "}"
}
// Equal returns true if s is equal to v.
func (s *Set) Equal(v Value) bool {
return Compare(s, v) == 0
}
// Compare compares s to other, return <0, 0, or >0 if it is less than, equal to,
// or greater than other.
func (s *Set) Compare(other Value) int {
return Compare(s, other)
func (s *set) Compare(other Value) int {
o1 := sortOrder(s)
o2 := sortOrder(other)
if o1 < o2 {
return -1
} else if o1 > o2 {
return 1
}
t := other.(*set)
sort.Sort(termSlice(*s))
sort.Sort(termSlice(*t))
return termSliceCompare(*s, *t)
}
// Find returns the current value or a not found error.
func (s *Set) Find(path Ref) (Value, error) {
func (s *set) Find(path Ref) (Value, error) {
if len(path) == 0 {
return s, nil
}
@@ -1044,8 +1083,8 @@ func (s *Set) Find(path Ref) (Value, error) {
}
// Diff returns elements in s that are not in other.
func (s *Set) Diff(other *Set) *Set {
r := &Set{}
func (s *set) Diff(other Set) Set {
r := &set{}
for _, x := range *s {
if !other.Contains(x) {
r.Add(x)
@@ -1055,8 +1094,8 @@ func (s *Set) Diff(other *Set) *Set {
}
// Intersect returns the set containing elements in both s and other.
func (s *Set) Intersect(other *Set) *Set {
r := &Set{}
func (s *set) Intersect(other Set) Set {
r := &set{}
for _, x := range *s {
if other.Contains(x) {
r.Add(x)
@@ -1066,41 +1105,64 @@ func (s *Set) Intersect(other *Set) *Set {
}
// Union returns the set containing all elements of s and other.
func (s *Set) Union(other *Set) *Set {
r := &Set{}
for _, x := range *s {
func (s *set) Union(other Set) Set {
r := &set{}
s.Iter(func(x *Term) error {
r.Add(x)
}
for _, x := range *other {
return nil
})
other.Iter(func(x *Term) error {
r.Add(x)
}
return nil
})
return r
}
// Add updates s to include t.
func (s *Set) Add(t *Term) {
func (s *set) Add(t *Term) {
if s.Contains(t) {
return
}
*s = append(*s, t)
}
// Iter calls f on each element in s. If f returns true, iteration stops and the
// return value is true.
func (s *Set) Iter(f func(*Term) bool) (stop bool) {
// Iter calls f on each element in s. If f returns an error, iteration stops
// and the return value is the error.
func (s *set) Iter(f func(*Term) error) error {
sl := *s
for i := range sl {
if f(sl[i]) {
return true
if err := f(sl[i]); err != nil {
return err
}
}
return false
return nil
}
var errStop = errors.New("stop")
// Until calls f on each element in s. If f returns true, iteration stops.
func (s *set) Until(f func(*Term) bool) bool {
err := s.Iter(func(t *Term) error {
if f(t) {
return errStop
}
return nil
})
return err != nil
}
// Foreach calls f on each element in s.
func (s *set) Foreach(f func(*Term)) {
s.Iter(func(t *Term) error {
f(t)
return nil
})
}
// Map returns a new Set obtained by applying f to each value in s.
func (s *Set) Map(f func(*Term) (*Term, error)) (*Set, error) {
func (s *set) Map(f func(*Term) (*Term, error)) (Set, error) {
sl := *s
set := &Set{}
set := &set{}
for i := range sl {
term, err := f(sl[i])
if err != nil {
@@ -1114,7 +1176,7 @@ func (s *Set) Map(f func(*Term) (*Term, error)) (*Set, error) {
// Reduce returns a Term produced by applying f to each value in s. The first
// argument to f is the reduced value (starting with i) and the second argument
// to f is the element in s.
func (s *Set) Reduce(i *Term, f func(*Term, *Term) (*Term, error)) (*Term, error) {
func (s *set) Reduce(i *Term, f func(*Term, *Term) (*Term, error)) (*Term, error) {
sl := *s
for _, elem := range sl {
var err error
@@ -1127,19 +1189,54 @@ func (s *Set) Reduce(i *Term, f func(*Term, *Term) (*Term, error)) (*Term, error
}
// Contains returns true if t is in s.
func (s Set) Contains(t *Term) bool {
for i := range s {
if s[i].Equal(t) {
func (s *set) Contains(t *Term) bool {
sl := *s
for i := range sl {
if sl[i].Equal(t) {
return true
}
}
return false
}
// Object represents an object as defined by the language. Objects are similar to
// the same types as defined by JSON with the exception that they can contain
// Vars and References.
type Object [][2]*Term
// Len returns the number of elements in the set.
func (s *set) Len() int {
return len(*s)
}
// Object represents an object as defined by the language.
type Object interface {
Value
Len() int
Get(*Term) *Term
Copy() Object
Insert(*Term, *Term)
Iter(func(*Term, *Term) error) error
Until(func(*Term, *Term) bool) bool
Foreach(func(*Term, *Term))
Map(func(*Term, *Term) (*Term, *Term, error)) (Object, error)
Diff(other Object) Object
Intersect(other Object) [][3]*Term
Merge(other Object) (Object, bool)
Keys() []*Term
}
// NewObject creates a new Object with t.
func NewObject(t ...[2]*Term) Object {
obj := &object{}
for i := range t {
obj.Insert(t[i][0], t[i][1])
}
return obj
}
// ObjectTerm creates a new Term with an Object value.
func ObjectTerm(o ...[2]*Term) *Term {
obj := object(o)
return &Term{Value: &obj}
}
type object [][2]*Term
// Item is a helper for constructing an tuple containing two Terms
// representing a key/value pair in an Object.
@@ -1147,19 +1244,52 @@ func Item(key, value *Term) [2]*Term {
return [2]*Term{key, value}
}
// Equal returns true if obj is equal to other.
func (obj Object) Equal(other Value) bool {
return Compare(obj, other) == 0
}
// Compare compares obj to other, return <0, 0, or >0 if it is less than, equal to,
// or greater than other.
func (obj Object) Compare(other Value) int {
return Compare(obj, other)
func (obj *object) Compare(other Value) int {
o1 := sortOrder(obj)
o2 := sortOrder(other)
if o1 < o2 {
return -1
} else if o2 < o1 {
return 1
}
a := obj
b := other.(*object)
keysA := a.Keys()
keysB := b.Keys()
sort.Sort(termSlice(keysA))
sort.Sort(termSlice(keysB))
minLen := a.Len()
if b.Len() < a.Len() {
minLen = b.Len()
}
for i := 0; i < minLen; i++ {
keysCmp := Compare(keysA[i], keysB[i])
if keysCmp < 0 {
return -1
}
if keysCmp > 0 {
return 1
}
valA := a.Get(keysA[i])
valB := b.Get(keysB[i])
valCmp := Compare(valA, valB)
if valCmp != 0 {
return valCmp
}
}
if a.Len() < b.Len() {
return -1
}
if b.Len() < a.Len() {
return 1
}
return 0
}
// Find returns the value at the key or undefined.
func (obj Object) Find(path Ref) (Value, error) {
func (obj *object) Find(path Ref) (Value, error) {
if len(path) == 0 {
return obj, nil
}
@@ -1170,9 +1300,19 @@ func (obj Object) Find(path Ref) (Value, error) {
return value.Value.Find(path[1:])
}
func (obj *object) Insert(k, v *Term) {
for _, pair := range *obj {
if Compare(pair[0], k) == 0 {
pair[1] = v
return
}
}
*obj = append(*obj, [2]*Term{k, v})
}
// Get returns the value of k in obj if k exists, otherwise nil.
func (obj Object) Get(k *Term) *Term {
for _, pair := range obj {
func (obj *object) Get(k *Term) *Term {
for _, pair := range *obj {
if pair[0].Equal(k) {
return pair[1]
}
@@ -1181,90 +1321,120 @@ func (obj Object) Get(k *Term) *Term {
}
// Hash returns the hash code for the Value.
func (obj Object) Hash() int {
func (obj *object) Hash() int {
var hash int
for i := range obj {
hash += obj[i][0].Value.Hash()
hash += obj[i][1].Value.Hash()
for _, pair := range *obj {
hash += pair[0].Value.Hash()
hash += pair[1].Value.Hash()
}
return hash
}
// IsGround returns true if all of the Object key/value pairs are ground.
func (obj Object) IsGround() bool {
for i := range obj {
if !obj[i][0].IsGround() {
func (obj *object) IsGround() bool {
for _, pair := range *obj {
if !pair[0].IsGround() {
return false
}
if !obj[i][1].IsGround() {
if !pair[1].IsGround() {
return false
}
}
return true
}
// ObjectTerm creates a new Term with an Object value.
func ObjectTerm(o ...[2]*Term) *Term {
return &Term{Value: Object(o)}
}
// Copy returns a deep copy of obj.
func (obj Object) Copy() Object {
cpy := make(Object, len(obj))
for i := range obj {
cpy[i] = Item(obj[i][0].Copy(), obj[i][1].Copy())
}
func (obj *object) Copy() Object {
cpy, _ := obj.Map(func(k, v *Term) (*Term, *Term, error) {
return k.Copy(), v.Copy(), nil
})
return cpy
}
// Diff returns a new Object that contains only the key/value pairs that exist in obj.
func (obj Object) Diff(other Object) Object {
r := Object{}
for _, i := range obj {
found := false
for _, j := range other {
if j[0].Equal(i[0]) {
found = true
break
}
func (obj *object) Diff(other Object) Object {
r := &object{}
obj.Foreach(func(k, v *Term) {
if other.Get(k) == nil {
r.Insert(k, v)
}
if !found {
r = append(r, i)
}
}
})
return r
}
// Intersect returns a slice of term triplets that represent the intersection of keys
// between obj and other. For each intersecting key, the values from obj and other are included
// as the last two terms in the triplet (respectively).
func (obj Object) Intersect(other Object) [][3]*Term {
func (obj *object) Intersect(other Object) [][3]*Term {
r := [][3]*Term{}
for _, i := range obj {
for _, j := range other {
if i[0].Equal(j[0]) {
r = append(r, [...]*Term{&Term{Value: i[0].Value}, i[1], j[1]})
}
obj.Foreach(func(k, v *Term) {
if v2 := other.Get(k); v2 != nil {
r = append(r, [3]*Term{k, v, v2})
}
}
})
return r
}
// Iter calls the function f for each key-value pair in the object. If f
// returns an error, iteration stops and the error is returned.
func (obj *object) Iter(f func(*Term, *Term) error) error {
for _, pair := range *obj {
if err := f(pair[0], pair[1]); err != nil {
return err
}
}
return nil
}
// Until calls f for each key-value pair in the object. If f returns true,
// iteration stops.
func (obj *object) Until(f func(*Term, *Term) bool) bool {
err := obj.Iter(func(k, v *Term) error {
if f(k, v) {
return errStop
}
return nil
})
return err != nil
}
// Foreach calls f for each key-value pair in the object.
func (obj *object) Foreach(f func(*Term, *Term)) {
obj.Iter(func(k, v *Term) error {
f(k, v)
return nil
})
}
// Map returns a new Object constructed by mapping each element in the object
// using the function f.
func (obj *object) Map(f func(*Term, *Term) (*Term, *Term, error)) (Object, error) {
cpy := make(object, obj.Len())
for i, pair := range *obj {
k, v, err := f(pair[0], pair[1])
if err != nil {
return nil, err
}
cpy[i] = [2]*Term{k, v}
}
return &cpy, nil
}
// Keys returns the keys of obj.
func (obj Object) Keys() []*Term {
keys := make([]*Term, len(obj))
for i, pair := range obj {
func (obj *object) Keys() []*Term {
keys := make([]*Term, obj.Len())
for i, pair := range *obj {
keys[i] = pair[0]
}
return keys
}
// MarshalJSON returns JSON encoded bytes representing obj.
func (obj Object) MarshalJSON() ([]byte, error) {
if len(obj) == 0 {
func (obj *object) MarshalJSON() ([]byte, error) {
if obj.Len() == 0 {
return json.Marshal([]interface{}{})
}
sl := [][2]*Term(obj)
sl := [][2]*Term(*obj)
return json.Marshal(sl)
}
@@ -1272,32 +1442,42 @@ func (obj Object) MarshalJSON() ([]byte, error) {
// overlapping keys between obj and other, the values of associated with the keys are merged. Only
// objects can be merged with other objects. If the values cannot be merged, the second turn value
// will be false.
func (obj Object) Merge(other Object) (Object, bool) {
r := Object{}
r = append(r, obj.Diff(other)...)
r = append(r, other.Diff(obj)...)
for _, vs := range obj.Intersect(other) {
var merged Value
switch v1 := vs[1].Value.(type) {
case Object:
switch v2 := vs[2].Value.(type) {
case Object:
m, ok := v1.Merge(v2)
if !ok {
return nil, false
}
merged = m
func (obj object) Merge(other Object) (result Object, ok bool) {
result = &object{}
stop := obj.Until(func(k, v *Term) bool {
if v2 := other.Get(k); v2 == nil {
result.Insert(k, v)
} else {
obj1, ok1 := v.Value.(Object)
obj2, ok2 := v2.Value.(Object)
if !ok1 || !ok2 {
return true
}
obj3, ok := obj1.Merge(obj2)
if !ok {
return true
}
result.Insert(k, NewTerm(obj3))
}
if merged == nil {
return nil, false
}
r = append(r, [2]*Term{vs[0], &Term{Value: merged}})
return false
})
if stop {
return nil, false
}
return r, true
other.Foreach(func(k, v *Term) {
if v2 := obj.Get(k); v2 == nil {
result.Insert(k, v)
}
})
return result, true
}
func (obj Object) String() string {
// Len returns the number of elements in the object.
func (obj object) Len() int {
return len(obj)
}
func (obj object) String() string {
var buf []string
for _, p := range obj {
buf = append(buf, fmt.Sprintf("%s: %s", p[0], p[1]))
@@ -1687,7 +1867,7 @@ func unmarshalValue(d map[string]interface{}) (Value, error) {
}
case "set":
if s, err := unmarshalTermSliceValue(d); err == nil {
set := &Set{}
set := &set{}
for _, x := range s {
set.Add(x)
}
@@ -1695,12 +1875,12 @@ func unmarshalValue(d map[string]interface{}) (Value, error) {
}
case "object":
if s, ok := v.([]interface{}); ok {
buf := Object{}
buf := &object{}
for _, x := range s {
if i, ok := x.([]interface{}); ok && len(i) == 2 {
p, err := unmarshalTermSlice(i)
if err == nil {
buf = append(buf, Item(p[0], p[1]))
buf.Insert(p[0], p[1])
continue
}
}
+9 -9
View File
@@ -90,7 +90,7 @@ func TestObjectSetOperations(t *testing.T) {
b := MustParseTerm(`{"c": "q", "d": "e"}`).Value.(Object)
r1 := a.Diff(b)
if !r1.Equal(MustParseTerm(`{"a": "b"}`).Value) {
if r1.Compare(MustParseTerm(`{"a": "b"}`).Value) != 0 {
t.Errorf(`Expected a.Diff(b) to equal {"a": "b"} but got: %v`, r1)
}
@@ -108,7 +108,7 @@ func TestObjectSetOperations(t *testing.T) {
r3, ok := c.Merge(d)
expected := MustParseTerm(`{"a": {"b": [1], "x": [3], "c": {"d": 2, "y": 4}}}`).Value.(Object)
if !ok || !r3.Equal(expected) {
if !ok || r3.Compare(expected) != 0 {
t.Errorf("Expected c.Merge(d) to equal %v but got: %v", expected, r3)
}
}
@@ -403,7 +403,7 @@ func TestSetEqual(t *testing.T) {
func TestSetMap(t *testing.T) {
set := MustParseTerm(`{"foo", "bar", "baz", "qux"}`).Value.(*Set)
set := MustParseTerm(`{"foo", "bar", "baz", "qux"}`).Value.(Set)
result, err := set.Map(func(term *Term) (*Term, error) {
s := string(term.Value.(String))
@@ -419,7 +419,7 @@ func TestSetMap(t *testing.T) {
expected := MustParseTerm(`{"foo", "BAR", "BAZ", "qux"}`).Value
if !result.Equal(expected) {
if result.Compare(expected) != 0 {
t.Fatalf("Expected map result to be %v but got: %v", expected, result)
}
@@ -450,10 +450,10 @@ func TestSetOperations(t *testing.T) {
}
for _, tc := range tests {
s1 := MustParseTerm(tc.a).Value.(*Set)
s2 := MustParseTerm(tc.b).Value.(*Set)
s3 := MustParseTerm(tc.c).Value.(*Set)
var result *Set
s1 := MustParseTerm(tc.a).Value.(Set)
s2 := MustParseTerm(tc.b).Value.(Set)
s3 := MustParseTerm(tc.c).Value.(Set)
var result Set
if tc.op == "-" {
result = s1.Diff(s2)
} else if tc.op == "&" {
@@ -463,7 +463,7 @@ func TestSetOperations(t *testing.T) {
} else {
panic("bad operation")
}
if !result.Equal(s3) {
if result.Compare(s3) != 0 {
t.Errorf("Expected %v for %v %v %v but got: %v", s3, tc.a, tc.op, tc.b, result)
}
}
+8 -9
View File
@@ -179,18 +179,17 @@ func Transform(t Transformer, x interface{}) (interface{}, error) {
}
return y, nil
case Object:
for i, elem := range y {
k, err := transformTerm(t, elem[0])
return y.Map(func(k, v *Term) (*Term, *Term, error) {
k, err := transformTerm(t, k)
if err != nil {
return nil, err
return nil, nil, err
}
v, err := transformTerm(t, elem[1])
v, err = transformTerm(t, v)
if err != nil {
return nil, err
return nil, nil, err
}
y[i] = Item(k, v)
}
return y, nil
return k, v, nil
})
case Array:
for i := range y {
if y[i], err = transformTerm(t, y[i]); err != nil {
@@ -198,7 +197,7 @@ func Transform(t Transformer, x interface{}) (interface{}, error) {
}
}
return y, nil
case *Set:
case Set:
y, err = y.Map(func(term *Term) (*Term, error) {
return transformTerm(t, term)
})
+7 -4
View File
@@ -103,10 +103,13 @@ func (u *unifier) unify(a *Term, b *Term) {
case Ref:
u.markAllSafe(a)
case Object:
if len(a) == len(b) {
for i := range a {
u.unify(a[i][1], b[i][1])
}
if a.Len() == b.Len() {
a.Iter(func(k, v *Term) error {
if v2 := b.Get(k); v2 != nil {
u.unify(v, v2)
}
return nil
})
}
}
+11 -11
View File
@@ -81,18 +81,18 @@ func Walk(v Visitor, x interface{}) {
Walk(w, t)
}
case Object:
for _, t := range x {
Walk(w, t[0])
Walk(w, t[1])
}
x.Foreach(func(k, v *Term) {
Walk(w, k)
Walk(w, v)
})
case Array:
for _, t := range x {
Walk(w, t)
}
case *Set:
for _, t := range *x {
case Set:
x.Foreach(func(t *Term) {
Walk(w, t)
}
})
case *ArrayComprehension:
Walk(w, x.Term)
Walk(w, x.Body)
@@ -253,9 +253,9 @@ func (vis *VarVisitor) Vars() VarSet {
func (vis *VarVisitor) Visit(v interface{}) Visitor {
if vis.params.SkipObjectKeys {
if o, ok := v.(Object); ok {
for _, i := range o {
Walk(vis, i[1])
}
o.Foreach(func(_, v *Term) {
Walk(vis, v)
})
return nil
}
}
@@ -280,7 +280,7 @@ func (vis *VarVisitor) Visit(v interface{}) Visitor {
}
}
if vis.params.SkipSets {
if _, ok := v.(*Set); ok {
if _, ok := v.(Set); ok {
return nil
}
}
+1 -1
View File
@@ -16,7 +16,7 @@ import (
func All(x interface{}) (resolved []ast.Ref, err error) {
var rawResolved []ast.Ref
switch x := x.(type) {
case *ast.Module, *ast.Package, *ast.Import, *ast.Rule, *ast.Head, ast.Body, *ast.Expr, *ast.With, *ast.Term, ast.Ref, ast.Object, ast.Array, *ast.Set, *ast.ArrayComprehension:
case *ast.Module, *ast.Package, *ast.Import, *ast.Rule, *ast.Head, ast.Body, *ast.Expr, *ast.With, *ast.Term, ast.Ref, ast.Object, ast.Array, ast.Set, *ast.ArrayComprehension:
default:
return nil, fmt.Errorf("not an ast element: %v", x)
}
+7 -7
View File
@@ -394,7 +394,7 @@ func (w *writer) writeTerm(term *ast.Term, comments []*ast.Comment) []*ast.Comme
comments = w.writeObject(x, term.Location, comments)
case ast.Array:
comments = w.writeArray(x, term.Location, comments)
case *ast.Set:
case ast.Set:
comments = w.writeSet(x, term.Location, comments)
case *ast.ArrayComprehension:
comments = w.writeArrayComprehension(x, term.Location, comments)
@@ -426,9 +426,9 @@ func (w *writer) writeObject(obj ast.Object, loc *ast.Location, comments []*ast.
defer w.write("}")
var s []interface{}
for _, t := range obj {
s = append(s, t)
}
obj.Foreach(func(k, v *ast.Term) {
s = append(s, ast.Item(k, v))
})
comments = w.writeIterable(s, loc, comments, w.objectWriter())
return w.insertComments(comments, closingLoc(0, 0, '{', '}', loc))
}
@@ -445,14 +445,14 @@ func (w *writer) writeArray(arr ast.Array, loc *ast.Location, comments []*ast.Co
return w.insertComments(comments, closingLoc(0, 0, '[', ']', loc))
}
func (w *writer) writeSet(set *ast.Set, loc *ast.Location, comments []*ast.Comment) []*ast.Comment {
func (w *writer) writeSet(set ast.Set, loc *ast.Location, comments []*ast.Comment) []*ast.Comment {
w.write("{")
defer w.write("}")
var s []interface{}
for _, t := range *set {
set.Foreach(func(t *ast.Term) {
s = append(s, t)
}
})
comments = w.writeIterable(s, loc, comments, w.listWriter())
return w.insertComments(comments, closingLoc(0, 0, '{', '}', loc))
}
+1 -1
View File
@@ -52,7 +52,7 @@ func NewPathForRef(ref ast.Ref) (path Path, err error) {
Code: NotFoundErr,
Message: fmt.Sprintf("%v: does not exist", ref),
}
case ast.Array, ast.Object, *ast.Set:
case ast.Array, ast.Object, ast.Set:
return nil, fmt.Errorf("composites cannot be base document keys: %v", ref)
default:
return nil, fmt.Errorf("unresolved reference (indicates error in caller): %v", ref)
+19 -17
View File
@@ -16,9 +16,9 @@ func builtinCount(a ast.Value) (ast.Value, error) {
case ast.Array:
return ast.IntNumberTerm(len(a)).Value, nil
case ast.Object:
return ast.IntNumberTerm(len(a)).Value, nil
case *ast.Set:
return ast.IntNumberTerm(len(*a)).Value, nil
return ast.IntNumberTerm(a.Len()).Value, nil
case ast.Set:
return ast.IntNumberTerm(a.Len()).Value, nil
case ast.String:
return ast.IntNumberTerm(len(a)).Value, nil
}
@@ -37,16 +37,17 @@ func builtinSum(a ast.Value) (ast.Value, error) {
sum = new(big.Float).Add(sum, builtins.NumberToFloat(n))
}
return builtins.FloatToNumber(sum), nil
case *ast.Set:
case ast.Set:
sum := big.NewFloat(0)
for _, x := range *a {
err := a.Iter(func(x *ast.Term) error {
n, ok := x.Value.(ast.Number)
if !ok {
return nil, builtins.NewOperandElementErr(1, a, x.Value, ast.NumberTypeName)
return builtins.NewOperandElementErr(1, a, x.Value, ast.NumberTypeName)
}
sum = new(big.Float).Add(sum, builtins.NumberToFloat(n))
}
return builtins.FloatToNumber(sum), nil
return nil
})
return builtins.FloatToNumber(sum), err
}
return nil, builtins.NewOperandTypeErr(1, a, ast.SetTypeName, ast.ArrayTypeName)
}
@@ -63,16 +64,17 @@ func builtinProduct(a ast.Value) (ast.Value, error) {
product = new(big.Float).Mul(product, builtins.NumberToFloat(n))
}
return builtins.FloatToNumber(product), nil
case *ast.Set:
case ast.Set:
product := big.NewFloat(1)
for _, x := range *a {
err := a.Iter(func(x *ast.Term) error {
n, ok := x.Value.(ast.Number)
if !ok {
return nil, builtins.NewOperandElementErr(1, a, x.Value, ast.NumberTypeName)
return builtins.NewOperandElementErr(1, a, x.Value, ast.NumberTypeName)
}
product = new(big.Float).Mul(product, builtins.NumberToFloat(n))
}
return builtins.FloatToNumber(product), nil
return nil
})
return builtins.FloatToNumber(product), err
}
return nil, builtins.NewOperandTypeErr(1, a, ast.SetTypeName, ast.ArrayTypeName)
}
@@ -90,8 +92,8 @@ func builtinMax(a ast.Value) (ast.Value, error) {
}
}
return max, nil
case *ast.Set:
if len(*a) == 0 {
case ast.Set:
if a.Len() == 0 {
return nil, BuiltinEmpty{}
}
max, err := a.Reduce(ast.NullTerm(), func(max *ast.Term, elem *ast.Term) (*ast.Term, error) {
@@ -119,8 +121,8 @@ func builtinMin(a ast.Value) (ast.Value, error) {
}
}
return min, nil
case *ast.Set:
if len(*a) == 0 {
case ast.Set:
if a.Len() == 0 {
return nil, BuiltinEmpty{}
}
min, err := a.Reduce(ast.NullTerm(), func(min *ast.Term, elem *ast.Term) (*ast.Term, error) {
+2 -2
View File
@@ -97,8 +97,8 @@ func builtinMinus(a, b ast.Value) (ast.Value, error) {
return builtins.FloatToNumber(f), nil
}
s1, ok3 := a.(*ast.Set)
s2, ok4 := b.(*ast.Set)
s1, ok3 := a.(ast.Set)
s2, ok4 := b.(ast.Set)
if ok3 && ok4 {
return s1.Diff(s2), nil
+4 -6
View File
@@ -91,13 +91,11 @@ func (u *bindings) Plug(a *ast.Term) *ast.Term {
return &cpy
case ast.Object:
cpy := *a
obj := make(ast.Object, len(v))
for i := 0; i < len(obj); i++ {
obj[i] = ast.Item(u.Plug(v[i][0]), u.Plug(v[i][1]))
}
cpy.Value = obj
cpy.Value, _ = v.Map(func(k, v *ast.Term) (*ast.Term, *ast.Term, error) {
return u.Plug(k), u.Plug(v), nil
})
return &cpy
case *ast.Set:
case ast.Set:
cpy := *a
cpy.Value, _ = v.Map(func(x *ast.Term) (*ast.Term, error) {
return u.Plug(x), nil
+2 -2
View File
@@ -104,8 +104,8 @@ func NumberOperand(x ast.Value, pos int) (ast.Number, error) {
// SetOperand converts x to a set. If the cast fails, a descriptive error is
// returned.
func SetOperand(x ast.Value, pos int) (*ast.Set, error) {
s, ok := x.(*ast.Set)
func SetOperand(x ast.Value, pos int) (ast.Set, error) {
s, ok := x.(ast.Set)
if !ok {
return nil, NewOperandTypeErr(pos, x, ast.SetTypeName)
}
+50 -42
View File
@@ -330,7 +330,7 @@ func (e *eval) biunify(a, b *ast.Term, b1, b2 *bindings, iter unifyIterator) err
case ast.Object:
return e.biunifyObjects(vA, vB, b1, b2, iter)
}
case *ast.Set:
case ast.Set:
return e.biunifyValues(a, b, b1, b2, iter)
}
return nil
@@ -353,23 +353,22 @@ func (e *eval) biunifyArraysRec(a, b ast.Array, b1, b2 *bindings, iter unifyIter
}
func (e *eval) biunifyObjects(a, b ast.Object, b1, b2 *bindings, iter unifyIterator) error {
if len(a) != len(b) {
if a.Len() != b.Len() {
return nil
}
return e.biunifyObjectsRec(a, b, b1, b2, iter, 0)
return e.biunifyObjectsRec(a, b, b1, b2, iter, a.Keys(), 0)
}
func (e *eval) biunifyObjectsRec(a, b ast.Object, b1, b2 *bindings, iter unifyIterator, idx int) error {
if idx == len(a) {
func (e *eval) biunifyObjectsRec(a, b ast.Object, b1, b2 *bindings, iter unifyIterator, keys []*ast.Term, idx int) error {
if idx == len(keys) {
return iter()
}
item := a[idx]
other := b.Get(item[0])
if other == nil {
v2 := b.Get(keys[idx])
if v2 == nil {
return nil
}
return e.biunify(item[1], other, b1, b2, func() error {
return e.biunifyObjectsRec(a, b, b1, b2, iter, idx+1)
return e.biunify(a.Get(keys[idx]), v2, b1, b2, func() error {
return e.biunifyObjectsRec(a, b, b1, b2, iter, keys, idx+1)
})
}
@@ -514,7 +513,7 @@ func (e *eval) biunifyComprehensionArray(x *ast.ArrayComprehension, b *ast.Term,
}
func (e *eval) biunifyComprehensionSet(x *ast.SetComprehension, b *ast.Term, b1, b2 *bindings, iter unifyIterator) error {
result := &ast.Set{}
result := ast.NewSet()
child := e.closure(x.Body)
err := child.Run(func(child *eval) error {
result.Add(child.bindings.Plug(x.Term))
@@ -527,7 +526,7 @@ func (e *eval) biunifyComprehensionSet(x *ast.SetComprehension, b *ast.Term, b1,
}
func (e *eval) biunifyComprehensionObject(x *ast.ObjectComprehension, b *ast.Term, b1, b2 *bindings, iter unifyIterator) error {
result := ast.Object{}
result := ast.NewObject()
child := e.closure(x.Body)
err := child.Run(func(child *eval) error {
key := child.bindings.Plug(x.Key)
@@ -536,7 +535,7 @@ func (e *eval) biunifyComprehensionObject(x *ast.ObjectComprehension, b *ast.Ter
if exist != nil && !exist.Equal(value) {
return objectDocKeyConflictErr(x.Key.Location)
}
result = append(result, ast.Item(key, value))
result.Insert(key, value)
return nil
})
if err != nil {
@@ -868,13 +867,13 @@ func (e evalTree) enumerate(iter unifyIterator) error {
}
}
case ast.Object:
for _, pair := range doc {
err := e.e.biunify(pair[0], e.ref[e.pos], e.bindings, e.bindings, func() error {
return e.next(iter, pair[0])
err := doc.Iter(func(k, _ *ast.Term) error {
return e.e.biunify(k, e.ref[e.pos], e.bindings, e.bindings, func() error {
return e.next(iter, k)
})
if err != nil {
return err
}
})
if err != nil {
return err
}
}
}
@@ -926,7 +925,7 @@ func (e evalTree) leaves(plugged ast.Ref, node *ast.TreeNode) (ast.Object, error
return nil, nil
}
result := ast.Object{}
result := ast.NewObject()
for _, child := range node.Children {
if child.Hide {
@@ -953,7 +952,8 @@ func (e evalTree) leaves(plugged ast.Ref, node *ast.TreeNode) (ast.Object, error
}
if save != nil {
result, _ = result.Merge(ast.Object{ast.Item(plugged[len(plugged)-1], ast.NewTerm(save))})
v := ast.NewObject([2]*ast.Term{plugged[len(plugged)-1], ast.NewTerm(save)})
result, _ = result.Merge(v)
}
plugged = plugged[:len(plugged)-1]
@@ -1157,7 +1157,7 @@ func (e evalVirtualPartial) evalTerm(iter unifyIterator, term *ast.Term, termbin
func (e evalVirtualPartial) reduce(head *ast.Head, b *bindings, result *ast.Term) (*ast.Term, error) {
switch v := result.Value.(type) {
case *ast.Set:
case ast.Set:
v.Add(b.Plug(head.Key))
case ast.Object:
key := b.Plug(head.Key)
@@ -1166,7 +1166,7 @@ func (e evalVirtualPartial) reduce(head *ast.Head, b *bindings, result *ast.Term
if exist != nil && !exist.Equal(value) {
return nil, objectDocKeyConflictErr(head.Location)
}
v = append(v, ast.Item(key, value))
v.Insert(key, value)
result.Value = v
}
@@ -1343,23 +1343,17 @@ func (e evalTerm) enumerate(iter unifyIterator) error {
}
}
case ast.Object:
for _, pair := range v {
err := e.e.biunify(pair[0], e.ref[e.pos], e.termbindings, e.bindings, func() error {
return v.Iter(func(k, _ *ast.Term) error {
return e.e.biunify(k, e.ref[e.pos], e.termbindings, e.bindings, func() error {
return e.next(iter, e.bindings.Plug(e.ref[e.pos]))
})
if err != nil {
return err
}
}
case *ast.Set:
for _, elem := range *v {
err := e.e.biunify(elem, e.ref[e.pos], e.termbindings, e.bindings, func() error {
})
case ast.Set:
return v.Iter(func(elem *ast.Term) error {
return e.e.biunify(elem, e.ref[e.pos], e.termbindings, e.bindings, func() error {
return e.next(iter, e.bindings.Plug(e.ref[e.pos]))
})
if err != nil {
return err
}
}
})
}
return nil
@@ -1367,16 +1361,23 @@ func (e evalTerm) enumerate(iter unifyIterator) error {
func (e evalTerm) get(plugged *ast.Term) (*ast.Term, *bindings) {
switch v := e.term.Value.(type) {
case *ast.Set:
case ast.Set:
if v.IsGround() {
if v.Contains(plugged) {
return e.termbindings.apply(plugged)
}
} else {
for _, elem := range *v {
var t *ast.Term
var b *bindings
stop := v.Until(func(elem *ast.Term) bool {
if e.termbindings.Plug(elem).Equal(plugged) {
return e.termbindings.apply(plugged)
t, b = e.termbindings.apply(plugged)
return true
}
return false
})
if stop {
return t, b
}
}
case ast.Object:
@@ -1386,10 +1387,17 @@ func (e evalTerm) get(plugged *ast.Term) (*ast.Term, *bindings) {
return e.termbindings.apply(term)
}
} else {
for i := range v {
if e.termbindings.Plug(v[i][0]).Equal(plugged) {
return e.termbindings.apply(v[i][1])
var t *ast.Term
var b *bindings
stop := v.Until(func(k, v *ast.Term) bool {
if e.termbindings.Plug(k).Equal(plugged) {
t, b = e.termbindings.apply(v)
return true
}
return false
})
if stop {
return t, b
}
}
case ast.Array:
+2 -3
View File
@@ -66,10 +66,9 @@ func makeInput(pairs [][2]*ast.Term) (ast.Value, error) {
func makeTree(k ast.Ref, v *ast.Term) ast.Object {
var obj ast.Object
for i := len(k) - 1; i >= 1; i-- {
obj = ast.Object{ast.Item(k[i], v)}
obj = ast.NewObject(ast.Item(k[i], v))
v = &ast.Term{Value: obj}
obj = ast.Object{}
}
obj = ast.Object{ast.Item(k[0], v)}
obj = ast.NewObject(ast.Item(k[0], v))
return obj
}
+3 -3
View File
@@ -265,12 +265,12 @@ func prepareTest(ctx context.Context, t *testing.T, params fixtureParams, f func
}
func toTerm(qrs QueryResultSet) *ast.Term {
set := &ast.Set{}
set := ast.NewSet()
for _, qr := range qrs {
obj := ast.Object{}
obj := ast.NewObject()
for k, v := range qr {
if !k.IsWildcard() {
obj = append(obj, ast.Item(ast.NewTerm(k), v))
obj.Insert(ast.NewTerm(k), v)
}
}
set.Add(ast.NewTerm(obj))
+5 -7
View File
@@ -62,18 +62,16 @@ func builtinConcat(a, b ast.Value) (ast.Value, error) {
}
strs = append(strs, string(s))
}
case *ast.Set:
var err error
stopped := b.Iter(func(x *ast.Term) bool {
case ast.Set:
err := b.Iter(func(x *ast.Term) error {
s, ok := x.Value.(ast.String)
if !ok {
err = builtins.NewOperandElementErr(2, b, x.Value, ast.StringTypeName)
return true
return builtins.NewOperandElementErr(2, b, x.Value, ast.StringTypeName)
}
strs = append(strs, string(s))
return false
return nil
})
if stopped {
if err != nil {
return nil, err
}
default:
+11 -9
View File
@@ -2234,10 +2234,10 @@ func parseBindings(s string) *ast.ValueMap {
return nil
}
r := ast.NewValueMap()
for _, pair := range obj {
k, v := pair[0], pair[1]
obj.Iter(func(k, v *ast.Term) error {
r.Put(k.Value, v.Value)
}
return nil
})
return r
}
@@ -2248,13 +2248,15 @@ func parseVars(s string) map[ast.Var]ast.Value {
return nil
}
r := map[ast.Var]ast.Value{}
for _, pair := range obj {
k, v := pair[0].Value, pair[1].Value
if k, ok := k.(ast.Var); ok {
r[k] = v
} else {
return nil
stop := obj.Until(func(k, v *ast.Term) bool {
if asVar, ok := k.Value.(ast.Var); ok {
r[asVar] = v.Value
return false
}
return true
})
if stop {
return nil
}
return r
}
+10 -14
View File
@@ -30,27 +30,23 @@ func walk(input *ast.Term, path ast.Array, iter func(*ast.Term) error) error {
path = path[:len(path)-1]
}
case ast.Object:
for _, pair := range v {
path = append(path, pair[0])
if err := walk(pair[1], path, iter); err != nil {
return v.Iter(func(k, v *ast.Term) error {
path = append(path, k)
if err := walk(v, path, iter); err != nil {
return err
}
path = path[:len(path)-1]
}
case *ast.Set:
var err error
v.Iter(func(elem *ast.Term) bool {
if err != nil {
return true
}
return nil
})
case ast.Set:
return v.Iter(func(elem *ast.Term) error {
path = append(path, elem)
if err = walk(elem, path, iter); err != nil {
return true
if err := walk(elem, path, iter); err != nil {
return err
}
path = path[:len(path)-1]
return false
return nil
})
return err
}
return nil