diff --git a/v1/ast/check.go b/v1/ast/check.go index 7506719bba..a639fadf39 100644 --- a/v1/ast/check.go +++ b/v1/ast/check.go @@ -781,8 +781,8 @@ func (rc *refChecker) checkRef(curr *TypeEnv, node *typeTreeNode, ref Ref, idx i case RootDocumentNames.Contains(ref[0]): if idx != 0 { - node.Children().Iter(func(_, child util.T) bool { - _ = rc.checkRef(curr, child.(*typeTreeNode), ref, idx+1) // ignore error + node.Children().Iter(func(_ Value, child *typeTreeNode) bool { + _ = rc.checkRef(curr, child, ref, idx+1) // ignore error return false }) return nil @@ -1127,8 +1127,8 @@ func newArgError(loc *Location, builtinName Ref, msg string, have []types.Type, } func getOneOfForNode(node *typeTreeNode) (result []Value) { - node.Children().Iter(func(k, _ util.T) bool { - result = append(result, k.(Value)) + node.Children().Iter(func(k Value, _ *typeTreeNode) bool { + result = append(result, k) return false }) diff --git a/v1/ast/compare.go b/v1/ast/compare.go index 77343f9c86..6a2524815d 100644 --- a/v1/ast/compare.go +++ b/v1/ast/compare.go @@ -402,6 +402,10 @@ func TermValueCompare(a, b *Term) int { return a.Value.Compare(b.Value) } +func TermValueEqual(a, b *Term) bool { + return ValueEqual(a.Value, b.Value) +} + func ValueEqual(a, b Value) bool { // TODO(ae): why doesn't this work the same? // diff --git a/v1/ast/compile.go b/v1/ast/compile.go index 75f864f514..29dae29881 100644 --- a/v1/ast/compile.go +++ b/v1/ast/compile.go @@ -124,7 +124,7 @@ type Compiler struct { localvargen *localVarGenerator moduleLoader ModuleLoader - ruleIndices *util.HashMap + ruleIndices *util.HasherMap[Ref, RuleIndex] stages []stage maxErrs int sorted []string // list of sorted module names @@ -303,15 +303,10 @@ type stage struct { func NewCompiler() *Compiler { c := &Compiler{ - Modules: map[string]*Module{}, - RewrittenVars: map[Var]Var{}, - Required: &Capabilities{}, - ruleIndices: util.NewHashMap(func(a, b util.T) bool { - r1, r2 := a.(Ref), b.(Ref) - return r1.Equal(r2) - }, func(x util.T) int { - return x.(Ref).Hash() - }), + Modules: map[string]*Module{}, + RewrittenVars: map[Var]Var{}, + Required: &Capabilities{}, + ruleIndices: util.NewHasherMap[Ref, RuleIndex](RefEqual), maxErrs: CompileErrorLimitDefault, after: map[string][]CompilerStageDefinition{}, unsafeBuiltinsMap: map[string]struct{}{}, @@ -825,7 +820,7 @@ func (c *Compiler) RuleIndex(path Ref) RuleIndex { if !ok { return nil } - return r.(RuleIndex) + return r } // PassesTypeCheck determines whether the given body passes type checking @@ -1738,13 +1733,9 @@ func (c *Compiler) err(err *Error) { c.Errors = append(c.Errors, err) } -func (c *Compiler) getExports() *util.HashMap { +func (c *Compiler) getExports() *util.HasherMap[Ref, []Ref] { - rules := util.NewHashMap(func(a, b util.T) bool { - return a.(Ref).Equal(b.(Ref)) - }, func(v util.T) int { - return v.(Ref).Hash() - }) + rules := util.NewHasherMap[Ref, []Ref](RefEqual) for _, name := range c.sorted { mod := c.Modules[name] @@ -1757,18 +1748,30 @@ func (c *Compiler) getExports() *util.HashMap { return rules } -func hashMapAdd(rules *util.HashMap, pkg, rule Ref) { +func refSliceEqual(a, b []Ref) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if !a[i].Equal(b[i]) { + return false + } + } + return true +} + +func hashMapAdd(rules *util.HasherMap[Ref, []Ref], pkg, rule Ref) { prev, ok := rules.Get(pkg) if !ok { rules.Put(pkg, []Ref{rule}) return } - for _, p := range prev.([]Ref) { + for _, p := range prev { if p.Equal(rule) { return } } - rules.Put(pkg, append(prev.([]Ref), rule)) + rules.Put(pkg, append(prev, rule)) } func (c *Compiler) GetAnnotationSet() *AnnotationSet { @@ -1867,7 +1870,7 @@ func (c *Compiler) resolveAllRefs() { var ruleExports []Ref if x, ok := rules.Get(mod.Package.Path); ok { - ruleExports = x.([]Ref) + ruleExports = x } globals := getGlobals(mod.Package, ruleExports, mod.Imports) @@ -3014,7 +3017,7 @@ func (qc *queryCompiler) resolveRefs(qctx *QueryContext, body Body) (Body, error var ruleExports []Ref rules := qc.compiler.getExports() if exist, ok := rules.Get(pkg.Path); ok { - ruleExports = exist.([]Ref) + ruleExports = exist } globals = getGlobals(qctx.Package, ruleExports, qctx.Imports) diff --git a/v1/ast/compile_test.go b/v1/ast/compile_test.go index 9820ab0505..b045872568 100644 --- a/v1/ast/compile_test.go +++ b/v1/ast/compile_test.go @@ -410,28 +410,8 @@ func TestCompilerGetExports(t *testing.T) { // TODO(sr): add multi-val rule, and ref-with-var single-value rule. } - hashMap := func(ms map[string][]string) *util.HashMap { - rules := util.NewHashMap(func(a, b util.T) bool { - switch a := a.(type) { - case Ref: - return a.Equal(b.(Ref)) - case []Ref: - b := b.([]Ref) - if len(b) != len(a) { - return false - } - for i := range a { - if !a[i].Equal(b[i]) { - return false - } - } - return true - default: - panic("unreachable") - } - }, func(v util.T) int { - return v.(Ref).Hash() - }) + hashMap := func(ms map[string][]string) *util.HasherMap[Ref, []Ref] { + rules := util.NewHasherMap[Ref, []Ref](RefEqual) for r, rs := range ms { refs := make([]Ref, len(rs)) for i := range rs { @@ -449,13 +429,29 @@ func TestCompilerGetExports(t *testing.T) { c.Modules[strconv.Itoa(i)] = m c.sorted = append(c.sorted, strconv.Itoa(i)) } - if exp, act := hashMap(tc.exports), c.getExports(); !exp.Equal(act) { + if exp, act := hashMap(tc.exports), c.getExports(); !refMapEqual(exp, act) { t.Errorf("expected %v, got %v", exp, act) } }) } } +func refMapEqual(a, b *util.HasherMap[Ref, []Ref]) bool { + if a.Len() != b.Len() { + return false + } + return !a.Iter(func(k Ref, v []Ref) bool { + v2, ok := b.Get(k) + if !ok { + return true + } + if !refSliceEqual(v, v2) { + return true + } + return false + }) +} + func toRef(s string) Ref { switch t := MustParseTerm(s).Value.(type) { case Var: diff --git a/v1/ast/env.go b/v1/ast/env.go index fb374b1739..d7d72e1422 100644 --- a/v1/ast/env.go +++ b/v1/ast/env.go @@ -200,10 +200,7 @@ func (env *TypeEnv) getRefRecExtent(node *typeTreeNode) types.Type { children := []*types.StaticProperty{} - node.Children().Iter(func(k, v util.T) bool { - key := k.(Value) - child := v.(*typeTreeNode) - + node.Children().Iter(func(key Value, child *typeTreeNode) bool { tpe := env.getRefRecExtent(child) // NOTE(sr): Converting to Golang-native types here is an extension of what we did @@ -237,14 +234,14 @@ func (env *TypeEnv) wrap() *TypeEnv { type typeTreeNode struct { key Value value types.Type - children *util.HashMap + children *util.HasherMap[Value, *typeTreeNode] } func newTypeTree() *typeTreeNode { return &typeTreeNode{ key: nil, value: nil, - children: util.NewHashMap(valueEq, valueHash), + children: util.NewHasherMap[Value, *typeTreeNode](ValueEqual), } } @@ -253,10 +250,10 @@ func (n *typeTreeNode) Child(key Value) *typeTreeNode { if !ok { return nil } - return value.(*typeTreeNode) + return value } -func (n *typeTreeNode) Children() *util.HashMap { +func (n *typeTreeNode) Children() *util.HasherMap[Value, *typeTreeNode] { return n.children } @@ -267,7 +264,7 @@ func (n *typeTreeNode) Get(path Ref) types.Type { if !ok { return nil } - curr = child.(*typeTreeNode) + curr = child } return curr.Value() } @@ -285,7 +282,7 @@ func (n *typeTreeNode) PutOne(key Value, tpe types.Type) { child.key = key n.children.Put(key, child) } else { - child = c.(*typeTreeNode) + child = c } child.value = tpe @@ -302,7 +299,7 @@ func (n *typeTreeNode) Put(path Ref, tpe types.Type) { child.key = term.Value curr.children.Put(child.key, child) } else { - child = c.(*typeTreeNode) + child = c } curr = child @@ -324,8 +321,7 @@ func (n *typeTreeNode) Insert(path Ref, tpe types.Type, env *TypeEnv) { child.key = term.Value curr.children.Put(child.key, child) } else { - child = c.(*typeTreeNode) - + child = c if child.value != nil && i+1 < len(path) { // If child has an object value, merge the new value into it. if o, ok := child.value.(*types.Object); ok { @@ -426,13 +422,12 @@ func (n *typeTreeNode) String() string { b.WriteString(v.String()) } - n.children.Iter(func(_, v util.T) bool { - if child, ok := v.(*typeTreeNode); ok { - b.WriteString("\n\t+ ") - s := child.String() - s = strings.ReplaceAll(s, "\n", "\n\t") - b.WriteString(s) - } + n.children.Iter(func(_ Value, child *typeTreeNode) bool { + b.WriteString("\n\t+ ") + s := child.String() + s = strings.ReplaceAll(s, "\n", "\n\t") + b.WriteString(s) + return false }) @@ -472,8 +467,8 @@ func insertIntoObject(o *types.Object, path Ref, tpe types.Type, env *TypeEnv) ( func (n *typeTreeNode) Leafs() map[*Ref]types.Type { leafs := map[*Ref]types.Type{} - n.children.Iter(func(_, v util.T) bool { - collectLeafs(v.(*typeTreeNode), nil, leafs) + n.children.Iter(func(_ Value, v *typeTreeNode) bool { + collectLeafs(v, nil, leafs) return false }) return leafs @@ -485,8 +480,8 @@ func collectLeafs(n *typeTreeNode, path Ref, leafs map[*Ref]types.Type) { leafs[&nPath] = n.Value() return } - n.children.Iter(func(_, v util.T) bool { - collectLeafs(v.(*typeTreeNode), nPath, leafs) + n.children.Iter(func(_ Value, v *typeTreeNode) bool { + collectLeafs(v, nPath, leafs) return false }) } diff --git a/v1/ast/index.go b/v1/ast/index.go index 84f1b9b863..fd77e57e6b 100644 --- a/v1/ast/index.go +++ b/v1/ast/index.go @@ -235,7 +235,7 @@ type refindex struct { type refindices struct { isVirtual func(Ref) bool rules map[*Rule][]*refindex - frequency *util.HashMap + frequency *util.HasherMap[Ref, int] sorted []Ref } @@ -243,12 +243,7 @@ func newrefindices(isVirtual func(Ref) bool) *refindices { return &refindices{ isVirtual: isVirtual, rules: map[*Rule][]*refindex{}, - frequency: util.NewHashMap(func(a, b util.T) bool { - r1, r2 := a.(Ref), b.(Ref) - return r1.Equal(r2) - }, func(x util.T) int { - return x.(Ref).Hash() - }), + frequency: util.NewHasherMap[Ref, int](RefEqual), } } @@ -296,9 +291,9 @@ func (i *refindices) Sorted() []Ref { counts := make([]int, 0, i.frequency.Len()) i.sorted = make([]Ref, 0, i.frequency.Len()) - i.frequency.Iter(func(k, v util.T) bool { - counts = append(counts, v.(int)) - i.sorted = append(i.sorted, k.(Ref)) + i.frequency.Iter(func(k Ref, v int) bool { + counts = append(counts, v) + i.sorted = append(i.sorted, k) return false }) @@ -399,7 +394,7 @@ func (i *refindices) insert(rule *Rule, index *refindex) { count = 0 } - i.frequency.Put(index.Ref, count.(int)+1) + i.frequency.Put(index.Ref, count+1) for pos, other := range i.rules[rule] { if other.Ref.Equal(index.Ref) { @@ -467,7 +462,7 @@ type trieNode struct { next *trieNode any *trieNode undefined *trieNode - scalars *util.HashMap + scalars *util.HasherMap[Value, *trieNode] array *trieNode rules []*ruleNode } @@ -492,9 +487,7 @@ func (node *trieNode) String() string { } if node.scalars.Len() > 0 { buf := make([]string, 0, node.scalars.Len()) - node.scalars.Iter(func(k, v util.T) bool { - key := k.(Value) - val := v.(*trieNode) + node.scalars.Iter(func(key Value, val *trieNode) bool { buf = append(buf, fmt.Sprintf("scalar(%v):%p", key, val)) return false }) @@ -535,7 +528,7 @@ type ruleNode struct { func newTrieNodeImpl() *trieNode { return &trieNode{ - scalars: util.NewHashMap(valueEq, valueHash), + scalars: util.NewHasherMap[Value, *trieNode](ValueEqual), } } @@ -551,8 +544,7 @@ func (node *trieNode) Do(walker trieWalker) { node.undefined.Do(next) } - node.scalars.Iter(func(_, v util.T) bool { - child := v.(*trieNode) + node.scalars.Iter(func(_ Value, child *trieNode) bool { child.Do(next) return false }) @@ -618,7 +610,7 @@ func (node *trieNode) insertValue(value Value) *trieNode { child = newTrieNodeImpl() node.scalars.Put(value, child) } - return child.(*trieNode) + return child case *Array: if node.array == nil { node.array = newTrieNodeImpl() @@ -647,7 +639,7 @@ func (node *trieNode) insertArray(arr *Array) *trieNode { child = newTrieNodeImpl() node.scalars.Put(head, child) } - return child.(*trieNode).insertArray(arr.Slice(1, -1)) + return child.insertArray(arr.Slice(1, -1)) } panic("illegal value") @@ -712,7 +704,7 @@ func (node *trieNode) traverseValue(resolver ValueResolver, tr *trieTraversalRes if !ok { return nil } - return child.(*trieNode).Traverse(resolver, tr) + return child.Traverse(resolver, tr) } return nil @@ -741,7 +733,7 @@ func (node *trieNode) traverseArray(resolver ValueResolver, tr *trieTraversalRes if !ok { return nil } - return child.(*trieNode).traverseArray(resolver, tr, arr.Slice(1, -1)) + return child.traverseArray(resolver, tr, arr.Slice(1, -1)) } func (node *trieNode) traverseUnknown(resolver ValueResolver, tr *trieTraversalResult) error { @@ -767,12 +759,8 @@ func (node *trieNode) traverseUnknown(resolver ValueResolver, tr *trieTraversalR } var iterErr error - node.scalars.Iter(func(_, v util.T) bool { - child := v.(*trieNode) - if iterErr = child.traverseUnknown(resolver, tr); iterErr != nil { - return true - } - return false + node.scalars.Iter(func(_ Value, child *trieNode) bool { + return child.traverseUnknown(resolver, tr) != nil }) return iterErr diff --git a/v1/ast/map.go b/v1/ast/map.go index 5a64f32505..d0aa43755f 100644 --- a/v1/ast/map.go +++ b/v1/ast/map.go @@ -13,15 +13,14 @@ import ( // ValueMap represents a key/value map between AST term values. Any type of term // can be used as a key in the map. type ValueMap struct { - hashMap *util.HashMap + hashMap *util.TypedHashMap[Value, Value] } // NewValueMap returns a new ValueMap. func NewValueMap() *ValueMap { - vs := &ValueMap{ - hashMap: util.NewHashMap(valueEq, valueHash), + return &ValueMap{ + hashMap: util.NewTypedHashMap(ValueEqual, ValueEqual, Value.Hash, Value.Hash, nil), } - return vs } // MarshalJSON provides a custom marshaller for the ValueMap which @@ -39,16 +38,6 @@ func (vs *ValueMap) MarshalJSON() ([]byte, error) { return json.Marshal(tmp) } -// Copy returns a shallow copy of the ValueMap. -func (vs *ValueMap) Copy() *ValueMap { - if vs == nil { - return nil - } - cpy := NewValueMap() - cpy.hashMap = vs.hashMap.Copy() - return cpy -} - // Equal returns true if this ValueMap equals the other. func (vs *ValueMap) Equal(other *ValueMap) bool { if vs == nil { @@ -72,7 +61,7 @@ func (vs *ValueMap) Len() int { func (vs *ValueMap) Get(k Value) Value { if vs != nil { if v, ok := vs.hashMap.Get(k); ok { - return v.(Value) + return v } } return nil @@ -92,11 +81,7 @@ func (vs *ValueMap) Iter(iter func(Value, Value) bool) bool { if vs == nil { return false } - return vs.hashMap.Iter(func(kt, vt util.T) bool { - k := kt.(Value) - v := vt.(Value) - return iter(k, v) - }) + return vs.hashMap.Iter(iter) } // Put inserts a key k into the map with value v. @@ -121,13 +106,3 @@ func (vs *ValueMap) String() string { } return vs.hashMap.String() } - -func valueHash(v util.T) int { - return v.(Value).Hash() -} - -func valueEq(a, b util.T) bool { - av := a.(Value) - bv := b.(Value) - return av.Compare(bv) == 0 -} diff --git a/v1/ast/map_test.go b/v1/ast/map_test.go index 08f0fd2fb4..29a58bd2ca 100644 --- a/v1/ast/map_test.go +++ b/v1/ast/map_test.go @@ -38,22 +38,15 @@ func TestValueMapIter(t *testing.T) { } } -func TestValueMapCopy(t *testing.T) { - a := NewValueMap() - a.Put(String("x"), String("foo")) - a.Put(String("y"), String("bar")) - b := a.Copy() - b.Delete(String("y")) - if a.Get(String("y")) != String("bar") { - t.Fatalf("Unexpected a['y'] value: %v", a.Get(String("y"))) - } -} - func TestValueMapEqual(t *testing.T) { a := NewValueMap() a.Put(String("x"), String("foo")) a.Put(String("y"), String("bar")) - b := a.Copy() + + b := NewValueMap() + b.Put(String("x"), String("foo")) + b.Put(String("y"), String("bar")) + if !a.Equal(b) { t.Fatalf("Expected a == b but not for: %v / %v", a, b) } @@ -89,9 +82,6 @@ func TestValueMapString(t *testing.T) { func TestValueMapNil(t *testing.T) { var a *ValueMap - if a.Copy() != nil { - t.Fatalf("Expected nil map copy to be nil") - } a.Delete(String("foo")) var b *ValueMap if !a.Equal(b) { diff --git a/v1/ast/schema.go b/v1/ast/schema.go index e84a147a4a..3f9e2001d5 100644 --- a/v1/ast/schema.go +++ b/v1/ast/schema.go @@ -13,41 +13,32 @@ import ( // SchemaSet holds a map from a path to a schema. type SchemaSet struct { - m *util.HashMap + m *util.HasherMap[Ref, any] } // NewSchemaSet returns an empty SchemaSet. func NewSchemaSet() *SchemaSet { - - eqFunc := func(a, b util.T) bool { - return a.(Ref).Equal(b.(Ref)) - } - - hashFunc := func(x util.T) int { return x.(Ref).Hash() } - return &SchemaSet{ - m: util.NewHashMap(eqFunc, hashFunc), + m: util.NewHasherMap[Ref, any](RefEqual), } } // Put inserts a raw schema into the set. -func (ss *SchemaSet) Put(path Ref, raw interface{}) { +func (ss *SchemaSet) Put(path Ref, raw any) { ss.m.Put(path, raw) } // Get returns the raw schema identified by the path. -func (ss *SchemaSet) Get(path Ref) interface{} { - if ss == nil { - return nil +func (ss *SchemaSet) Get(path Ref) any { + if ss != nil { + if x, ok := ss.m.Get(path); ok { + return x + } } - x, ok := ss.m.Get(path) - if !ok { - return nil - } - return x + return nil } -func loadSchema(raw interface{}, allowNet []string) (types.Type, error) { +func loadSchema(raw any, allowNet []string) (types.Type, error) { jsonSchema, err := compileSchema(raw, allowNet) if err != nil { diff --git a/v1/dependencies/deps.go b/v1/dependencies/deps.go index 217516a0a9..ac8a1ecae7 100644 --- a/v1/dependencies/deps.go +++ b/v1/dependencies/deps.go @@ -6,7 +6,7 @@ package dependencies import ( "fmt" - "sort" + "slices" "github.com/open-policy-agent/opa/v1/ast" "github.com/open-policy-agent/opa/v1/util" @@ -168,25 +168,27 @@ func virtual(compiler *ast.Compiler, x interface{}, virtualRefs *dependencies) e } type dependencies struct { - refs *util.HashMap - visitedRules *util.HashMap + refs *util.HasherMap[ast.Ref, ast.Ref] + visitedRules *util.TypedHashMap[*ast.Rule, *ast.Rule] } func newRefSet() *dependencies { return &dependencies{ - refs: util.NewHashMap(func(a, b util.T) bool { - return a.(ast.Ref).Equal(b.(ast.Ref)) - }, func(a util.T) int { - return a.(ast.Ref).Hash() - }), - visitedRules: util.NewHashMap(func(a, b util.T) bool { - return a.(*ast.Rule).Equal(b.(*ast.Rule)) - }, func(a util.T) int { - return a.(*ast.Rule).Ref().Hash() - }), + refs: util.NewHasherMap[ast.Ref, ast.Ref](ast.RefEqual), + visitedRules: util.NewTypedHashMap[*ast.Rule, *ast.Rule]( + (*ast.Rule).Equal, + nil, + ruleHash, + nil, + nil, + ), } } +func ruleHash(r *ast.Rule) int { + return r.Ref().Hash() +} + func (rs *dependencies) add(r ast.Ref) { rs.refs.Put(r, r) } @@ -201,22 +203,18 @@ func (rs *dependencies) visited(rule *ast.Rule) bool { } func (rs *dependencies) toSlice() []ast.Ref { - var result []ast.Ref - rs.refs.Iter(func(k, _ util.T) bool { - result = append(result, k.(ast.Ref)) + result := make([]ast.Ref, 0, rs.refs.Len()) + rs.refs.Iter(func(k, _ ast.Ref) bool { + result = append(result, k) return false }) return result } func dedup(refs []ast.Ref) []ast.Ref { - sort.Slice(refs, func(i, j int) bool { - return refs[i].Compare(refs[j]) < 0 - }) + slices.SortFunc(refs, ast.RefCompare) - return filter(refs, func(a, b ast.Ref) bool { - return a.Compare(b) == 0 - }) + return slices.CompactFunc(refs, ast.RefEqual) } // filter removes all items from the list that cause pred to return true. It is diff --git a/v1/rego/rego.go b/v1/rego/rego.go index cdf0df7f03..ef638c95b9 100644 --- a/v1/rego/rego.go +++ b/v1/rego/rego.go @@ -2855,17 +2855,26 @@ func parseStringsToRefs(s []string) ([]ast.Ref, error) { func finishFunction(name string, bctx topdown.BuiltinContext, result *ast.Term, err error, iter func(*ast.Term) error) error { if err != nil { var e *HaltError + sb := strings.Builder{} if errors.As(err, &e) { + sb.Grow(len(name) + len(e.Error()) + 2) + sb.WriteString(name) + sb.WriteString(": ") + sb.WriteString(e.Error()) tdErr := &topdown.Error{ Code: topdown.BuiltinErr, - Message: fmt.Sprintf("%v: %v", name, e.Error()), + Message: sb.String(), Location: bctx.Location, } return topdown.Halt{Err: tdErr.Wrap(e)} } + sb.Grow(len(name) + len(err.Error()) + 2) + sb.WriteString(name) + sb.WriteString(": ") + sb.WriteString(err.Error()) tdErr := &topdown.Error{ Code: topdown.BuiltinErr, - Message: fmt.Sprintf("%v: %v", name, err.Error()), + Message: sb.String(), Location: bctx.Location, } return tdErr.Wrap(err) diff --git a/v1/topdown/cache.go b/v1/topdown/cache.go index 9d2416b7bd..42fb6ad3f5 100644 --- a/v1/topdown/cache.go +++ b/v1/topdown/cache.go @@ -44,7 +44,7 @@ type virtualCache struct { type virtualCacheElem struct { value *ast.Term - children *util.HashMap + children *util.HasherMap[*ast.Term, *virtualCacheElem] undefined bool } @@ -76,7 +76,7 @@ func (c *virtualCache) Get(ref ast.Ref) (*ast.Term, bool) { if !ok { return nil, false } - node = x.(*virtualCacheElem) + node = x } if node.undefined { return nil, true @@ -92,7 +92,7 @@ func (c *virtualCache) Put(ref ast.Ref, value *ast.Term) { for i := range ref { x, ok := node.children.Get(ref[i]) if ok { - node = x.(*virtualCacheElem) + node = x } else { next := newVirtualCacheElem() node.children.Put(ref[i], next) @@ -113,13 +113,13 @@ func (c *virtualCache) Keys() []ast.Ref { func keysRecursive(root ast.Ref, node *virtualCacheElem) []ast.Ref { var keys []ast.Ref - node.children.Iter(func(k, v util.T) bool { - ref := root.Append(k.(*ast.Term)) - if v.(*virtualCacheElem).value != nil { + node.children.Iter(func(k *ast.Term, v *virtualCacheElem) bool { + ref := root.Append(k) + if v.value != nil { keys = append(keys, ref) } - if v.(*virtualCacheElem).children.Len() > 0 { - keys = append(keys, keysRecursive(ref, v.(*virtualCacheElem))...) + if v.children.Len() > 0 { + keys = append(keys, keysRecursive(ref, v)...) } return false }) @@ -130,12 +130,8 @@ func newVirtualCacheElem() *virtualCacheElem { return &virtualCacheElem{children: newVirtualCacheHashMap()} } -func newVirtualCacheHashMap() *util.HashMap { - return util.NewHashMap(func(a, b util.T) bool { - return a.(*ast.Term).Equal(b.(*ast.Term)) - }, func(x util.T) int { - return x.(*ast.Term).Hash() - }) +func newVirtualCacheHashMap() *util.HasherMap[*ast.Term, *virtualCacheElem] { + return util.NewHasherMap[*ast.Term, *virtualCacheElem](ast.TermValueEqual) } // baseCache implements a trie structure to cache base documents read out of @@ -244,7 +240,7 @@ type comprehensionCache struct { type comprehensionCacheElem struct { value *ast.Term - children *util.HashMap + children *util.HasherMap[*ast.Term, *comprehensionCacheElem] } func newComprehensionCache() *comprehensionCache { @@ -281,7 +277,7 @@ func (c *comprehensionCacheElem) Get(key []*ast.Term) *ast.Term { if !ok { return nil } - node = x.(*comprehensionCacheElem) + node = x } return node.value } @@ -291,7 +287,7 @@ func (c *comprehensionCacheElem) Put(key []*ast.Term, value *ast.Term) { for i := range key { x, ok := node.children.Get(key[i]) if ok { - node = x.(*comprehensionCacheElem) + node = x } else { next := newComprehensionCacheElem() node.children.Put(key[i], next) @@ -301,12 +297,8 @@ func (c *comprehensionCacheElem) Put(key []*ast.Term, value *ast.Term) { node.value = value } -func newComprehensionCacheHashMap() *util.HashMap { - return util.NewHashMap(func(a, b util.T) bool { - return a.(*ast.Term).Equal(b.(*ast.Term)) - }, func(x util.T) int { - return x.(*ast.Term).Hash() - }) +func newComprehensionCacheHashMap() *util.HasherMap[*ast.Term, *comprehensionCacheElem] { + return util.NewHasherMap[*ast.Term, *comprehensionCacheElem](ast.TermValueEqual) } type functionMocksStack struct { diff --git a/v1/topdown/cache/cache.go b/v1/topdown/cache/cache.go index a2b80c0a77..064e9c4adc 100644 --- a/v1/topdown/cache/cache.go +++ b/v1/topdown/cache/cache.go @@ -372,19 +372,13 @@ type InterQueryValueCacheBucket interface { } type interQueryValueCacheBucket struct { - items util.TypedHashMap[ast.Value, any] + items util.HasherMap[ast.Value, any] config *NamedValueCacheConfig mtx sync.RWMutex } -func newItemsMap() *util.TypedHashMap[ast.Value, any] { - return util.NewTypedHashMap[ast.Value, any]( - func(a, b ast.Value) bool { return a.Compare(b) == 0 }, - func(any, any) bool { return false }, // map equality not supported - func(a ast.Value) int { return a.Hash() }, - func(any) int { return 0 }, // map equality not supported - nil, - ) +func newItemsMap() *util.HasherMap[ast.Value, any] { + return util.NewHasherMap[ast.Value, any](ast.ValueEqual) } func (c *interQueryValueCacheBucket) Get(k ast.Value) (any, bool) { diff --git a/v1/topdown/copypropagation/unionfind.go b/v1/topdown/copypropagation/unionfind.go index 679464250e..528c83a0f4 100644 --- a/v1/topdown/copypropagation/unionfind.go +++ b/v1/topdown/copypropagation/unionfind.go @@ -14,18 +14,14 @@ import ( type rankFunc func(*unionFindRoot, *unionFindRoot) (*unionFindRoot, *unionFindRoot) type unionFind struct { - roots *util.HashMap + roots *util.HasherMap[ast.Value, *unionFindRoot] parents *ast.ValueMap rank rankFunc } func newUnionFind(rank rankFunc) *unionFind { return &unionFind{ - roots: util.NewHashMap(func(a util.T, b util.T) bool { - return a.(ast.Value).Compare(b.(ast.Value)) == 0 - }, func(v util.T) int { - return v.(ast.Value).Hash() - }), + roots: util.NewHasherMap[ast.Value, *unionFindRoot](ast.ValueEqual), parents: ast.NewValueMap(), rank: rank, } @@ -53,7 +49,7 @@ func (uf *unionFind) Find(v ast.Value) (*unionFindRoot, bool) { if parent.Compare(v) == 0 { r, ok := uf.roots.Get(v) - return r.(*unionFindRoot), ok + return r, ok } return uf.Find(parent) @@ -93,13 +89,13 @@ func (uf *unionFind) String() string { map[string]ast.Value{}, } - uf.roots.Iter(func(k util.T, v util.T) bool { - o.Roots[k.(ast.Value).String()] = struct { + uf.roots.Iter(func(k ast.Value, v *unionFindRoot) bool { + o.Roots[k.String()] = struct { Constant *ast.Term Key ast.Value }{ - v.(*unionFindRoot).constant, - v.(*unionFindRoot).key, + v.constant, + v.key, } return true }) diff --git a/v1/topdown/http.go b/v1/topdown/http.go index 92218da982..58ffd2c122 100644 --- a/v1/topdown/http.go +++ b/v1/topdown/http.go @@ -781,28 +781,17 @@ type httpSendCacheEntry struct { // The httpSendCache is used for intra-query caching of http.send results. type httpSendCache struct { - entries *util.HashMap + entries *util.HasherMap[ast.Value, httpSendCacheEntry] } func newHTTPSendCache() *httpSendCache { return &httpSendCache{ - entries: util.NewHashMap(valueEq, valueHash), + entries: util.NewHasherMap[ast.Value, httpSendCacheEntry](ast.ValueEqual), } } -func valueHash(v util.T) int { - return ast.StringTerm(v.(ast.Value).String()).Hash() -} - -func valueEq(a, b util.T) bool { - av := a.(ast.Value) - bv := b.(ast.Value) - return av.String() == bv.String() -} - func (cache *httpSendCache) get(k ast.Value) *httpSendCacheEntry { if v, ok := cache.entries.Get(k); ok { - v := v.(httpSendCacheEntry) return &v } return nil diff --git a/v1/util/hashmap.go b/v1/util/hashmap.go index a6c584c570..cf6a385f41 100644 --- a/v1/util/hashmap.go +++ b/v1/util/hashmap.go @@ -12,6 +12,10 @@ import ( // T is a concise way to refer to T. type T interface{} +type Hasher interface { + Hash() int +} + type hashEntry[K any, V any] struct { k K v V @@ -177,3 +181,91 @@ func (h *TypedHashMap[K, V]) Update(other *TypedHashMap[K, V]) *TypedHashMap[K, }) return updated } + +type hasherEntry[K Hasher, V any] struct { + k K + v V + next *hasherEntry[K, V] +} + +// HasherMap represents a simpler version of TypedHashMap that uses Hasher's +// for keys, and requires only an equality function for keys. Ideally we'd have +// and Equal method for all key types too, and we could get rid of that requirement. +type HasherMap[K Hasher, V any] struct { + keq func(K, K) bool + table map[int]*hasherEntry[K, V] + size int +} + +// NewHasherMap returns a new empty HasherMap. +func NewHasherMap[K Hasher, V any](keq func(K, K) bool) *HasherMap[K, V] { + return &HasherMap[K, V]{ + keq: keq, + table: make(map[int]*hasherEntry[K, V]), + size: 0, + } +} + +// Get returns the value for k. +func (h *HasherMap[K, V]) Get(k K) (V, bool) { + for entry := h.table[k.Hash()]; entry != nil; entry = entry.next { + if h.keq(entry.k, k) { + return entry.v, true + } + } + var zero V + return zero, false +} + +// Put inserts a key/value pair into this HashMap. If the key is already present, the existing +// value is overwritten. +func (h *HasherMap[K, V]) Put(k K, v V) { + hash := k.Hash() + head := h.table[hash] + for entry := head; entry != nil; entry = entry.next { + if h.keq(entry.k, k) { + entry.v = v + return + } + } + h.table[hash] = &hasherEntry[K, V]{k: k, v: v, next: head} + h.size++ +} + +// Delete removes the key k. +func (h *HasherMap[K, V]) Delete(k K) { + hash := k.Hash() + var prev *hasherEntry[K, V] + for entry := h.table[hash]; entry != nil; entry = entry.next { + if h.keq(entry.k, k) { + if prev != nil { + prev.next = entry.next + } else { + h.table[hash] = entry.next + } + h.size-- + return + } + prev = entry + } +} + +// Iter invokes the iter function for each element in the HasherMap. +// If the iter function returns true, iteration stops and the return value is true. +// If the iter function never returns true, iteration proceeds through all elements +// and the return value is false. +func (h *HasherMap[K, V]) Iter(iter func(K, V) bool) bool { + for _, entry := range h.table { + for ; entry != nil; entry = entry.next { + if iter(entry.k, entry.v) { + return true + } + } + } + return false +} + +// Len returns the current size of this HashMap. +func (h *HasherMap[K, V]) Len() int { + return h.size +}