mirror of
https://github.com/open-policy-agent/opa.git
synced 2026-08-12 19:32:48 -06:00
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:
+20
-20
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
+2
-2
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
Vendored
+1
-1
@@ -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
@@ -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
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user