mirror of
https://github.com/open-policy-agent/opa.git
synced 2026-08-13 03:42:35 -06:00
Performance improvements in formatter (#7967)
- Use new `util.SlicePool` to avoid cost of temporary slices in `formatTerm`
- Add `SkipDefensiveCopying` option and enable it for all `Source`* functions
```
676179 ns/op 995204 B/op 8850 allocs/op // regression in 0fb7526
513570 ns/op 378787 B/op 8775 allocs/op // addressed regression with sync.Pool
481681 ns/op 352130 B/op 7954 allocs/op // new util.NewSlicePool using only pointers
365116 ns/op 160528 B/op 2098 allocs/op // new SkipDefensiveCopying option
```
Signed-off-by: Anders Eknert <anders@eknert.com>
This commit is contained in:
+66
-30
@@ -18,6 +18,20 @@ import (
|
||||
"github.com/open-policy-agent/opa/internal/future"
|
||||
"github.com/open-policy-agent/opa/v1/ast"
|
||||
"github.com/open-policy-agent/opa/v1/types"
|
||||
"github.com/open-policy-agent/opa/v1/util"
|
||||
)
|
||||
|
||||
// defaultLocationFile is the file name used in `Ast()` for terms
|
||||
// without a location, as could happen when pretty-printing the
|
||||
// results of partial eval.
|
||||
const defaultLocationFile = "__format_default__"
|
||||
|
||||
var (
|
||||
elseVar ast.Value = ast.Var("else")
|
||||
|
||||
expandedConst = ast.NewBody(ast.NewExpr(ast.InternedTerm(true)))
|
||||
commentsSlicePool = util.NewSlicePool[*ast.Comment](50)
|
||||
varRegexp = regexp.MustCompile("^[[:alpha:]_][[:alpha:][:digit:]_]*$")
|
||||
)
|
||||
|
||||
// Opts lets you control the code formatting via `AstWithOpts()`.
|
||||
@@ -38,6 +52,11 @@ type Opts struct {
|
||||
// Imports are only removed if [Opts.RegoVersion] makes them redundant.
|
||||
DropV0Imports bool
|
||||
|
||||
// SkipDefensiveCopying, if true, will avoid deep-copying the AST before formatting it.
|
||||
// This is true by default for all Source* functions, but false by default for Ast* functions,
|
||||
// as some formatting operations may otherwise mutate the AST.
|
||||
SkipDefensiveCopying bool
|
||||
|
||||
Capabilities *ast.Capabilities
|
||||
}
|
||||
|
||||
@@ -48,16 +67,11 @@ func (o Opts) effectiveRegoVersion() ast.RegoVersion {
|
||||
return o.RegoVersion
|
||||
}
|
||||
|
||||
// defaultLocationFile is the file name used in `Ast()` for terms
|
||||
// without a location, as could happen when pretty-printing the
|
||||
// results of partial eval.
|
||||
const defaultLocationFile = "__format_default__"
|
||||
|
||||
// Source formats a Rego source file. The bytes provided must describe a complete
|
||||
// Rego module. If they don't, Source will return an error resulting from the attempt
|
||||
// to parse the bytes.
|
||||
func Source(filename string, src []byte) ([]byte, error) {
|
||||
return SourceWithOpts(filename, src, Opts{})
|
||||
return SourceWithOpts(filename, src, Opts{SkipDefensiveCopying: true})
|
||||
}
|
||||
|
||||
func SourceWithOpts(filename string, src []byte, opts Opts) ([]byte, error) {
|
||||
@@ -72,6 +86,9 @@ func SourceWithOpts(filename string, src []byte, opts Opts) ([]byte, error) {
|
||||
parserOpts.RegoVersion = ast.RegoV1
|
||||
}
|
||||
|
||||
// Copying the node does not make sense when both input and output are bytes.
|
||||
opts.SkipDefensiveCopying = true
|
||||
|
||||
if parserOpts.RegoVersion == ast.RegoUndefined {
|
||||
parserOpts.RegoVersion = ast.DefaultRegoVersion
|
||||
}
|
||||
@@ -166,7 +183,9 @@ func AstWithOpts(x any, opts Opts) ([]byte, error) {
|
||||
// The node has to be deep copied because it may be mutated below. Alternatively,
|
||||
// we could avoid the copy by checking if mutation will occur first. For now,
|
||||
// since format is not latency sensitive, just deep copy in all cases.
|
||||
x = ast.Copy(x)
|
||||
if !opts.SkipDefensiveCopying {
|
||||
x = ast.Copy(x)
|
||||
}
|
||||
|
||||
wildcards := map[ast.Var]*ast.Term{}
|
||||
|
||||
@@ -233,10 +252,11 @@ func AstWithOpts(x any, opts Opts) ([]byte, error) {
|
||||
}
|
||||
|
||||
case *ast.Rule:
|
||||
if len(n.Head.Ref()) > 2 {
|
||||
headLen := len(n.Head.Ref())
|
||||
if headLen > 2 {
|
||||
o.refHeads = true
|
||||
}
|
||||
if len(n.Head.Ref()) == 2 && n.Head.Key != nil && n.Head.Value == nil { // p.q contains "x"
|
||||
if headLen == 2 && n.Head.Key != nil && n.Head.Value == nil { // p.q contains "x"
|
||||
o.refHeads = true
|
||||
}
|
||||
}
|
||||
@@ -339,6 +359,7 @@ func AstWithOpts(x any, opts Opts) ([]byte, error) {
|
||||
if len(w.errs) > 0 {
|
||||
return nil, w.errs
|
||||
}
|
||||
|
||||
return squashTrailingNewlines(w.buf.Bytes()), nil
|
||||
}
|
||||
|
||||
@@ -545,8 +566,6 @@ func (w *writer) writeRules(rules []*ast.Rule, comments []*ast.Comment) ([]*ast.
|
||||
return comments, nil
|
||||
}
|
||||
|
||||
var expandedConst = ast.NewBody(ast.NewExpr(ast.InternedTerm(true)))
|
||||
|
||||
func (w *writer) groupableOneLiner(rule *ast.Rule) bool {
|
||||
// Location required to determine if two rules are adjacent in the policy.
|
||||
// If not, we respect line breaks between rules.
|
||||
@@ -667,8 +686,6 @@ func (w *writer) writeRule(rule *ast.Rule, isElse bool, comments []*ast.Comment)
|
||||
return comments, nil
|
||||
}
|
||||
|
||||
var elseVar ast.Value = ast.Var("else")
|
||||
|
||||
func (w *writer) writeElse(rule *ast.Rule, comments []*ast.Comment) ([]*ast.Comment, error) {
|
||||
// If there was nothing else on the line before the "else" starts
|
||||
// then preserve this style of else block, otherwise it will be
|
||||
@@ -1127,18 +1144,33 @@ func (w *writer) writeWith(with *ast.With, comments []*ast.Comment, indented boo
|
||||
return comments, nil
|
||||
}
|
||||
|
||||
// saveComments saves a copy of the comments slice in a pooled slice to and returns it.
|
||||
// This is to avoid having to create a new slice every time we need to save comments.
|
||||
// The caller is responsible for putting the slice back in the pool when done.
|
||||
func saveComments(comments []*ast.Comment) *[]*ast.Comment {
|
||||
cmlen := len(comments)
|
||||
saved := commentsSlicePool.Get(cmlen)
|
||||
|
||||
copy(*saved, comments)
|
||||
|
||||
return saved
|
||||
}
|
||||
|
||||
func (w *writer) writeTerm(term *ast.Term, comments []*ast.Comment) ([]*ast.Comment, error) {
|
||||
currentComments := make([]*ast.Comment, len(comments))
|
||||
copy(currentComments, comments)
|
||||
if len(comments) == 0 {
|
||||
return w.writeTermParens(false, term, comments)
|
||||
}
|
||||
|
||||
currentLen := w.buf.Len()
|
||||
currentComments := saveComments(comments)
|
||||
defer commentsSlicePool.Put(currentComments)
|
||||
|
||||
comments, err := w.writeTermParens(false, term, comments)
|
||||
if err != nil {
|
||||
if errors.As(err, &unexpectedCommentError{}) {
|
||||
w.buf.Truncate(currentLen)
|
||||
|
||||
comments, uErr := w.writeUnformatted(term.Location, currentComments)
|
||||
comments, uErr := w.writeUnformatted(term.Location, *currentComments)
|
||||
if uErr != nil {
|
||||
return nil, uErr
|
||||
}
|
||||
@@ -1156,16 +1188,16 @@ func (w *writer) writeUnformatted(location *ast.Location, currentComments []*ast
|
||||
return nil, errors.New("original unformatted text is empty")
|
||||
}
|
||||
|
||||
rawRule := string(location.Text)
|
||||
rowNum := len(strings.Split(rawRule, "\n"))
|
||||
rowNum := bytes.Count(location.Text, []byte{'\n'}) + 1
|
||||
|
||||
w.write(string(location.Text))
|
||||
w.writeBytes(location.Text)
|
||||
|
||||
comments := make([]*ast.Comment, 0, len(currentComments))
|
||||
for _, c := range currentComments {
|
||||
// if there is a body then wait to write the last comment
|
||||
if w.writeCommentOnFinalLine && c.Location.Row == location.Row+rowNum-1 {
|
||||
w.write(" " + string(c.Location.Text))
|
||||
w.write(" ")
|
||||
w.writeBytes(c.Location.Text)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -1227,19 +1259,19 @@ func (w *writer) writeTermParens(parens bool, term *ast.Term, comments []*ast.Co
|
||||
case ast.String:
|
||||
if term.Location.Text[0] == '`' {
|
||||
// To preserve raw strings, we need to output the original text,
|
||||
w.write(string(term.Location.Text))
|
||||
w.writeBytes(term.Location.Text)
|
||||
} else {
|
||||
// x.String() cannot be used by default because it can change the input string "\u0000" to "\x00"
|
||||
var after, quote string
|
||||
var after, quote []byte
|
||||
var found bool
|
||||
// term.Location.Text could contain the prefix `else :=`, remove it
|
||||
switch term.Location.Text[len(term.Location.Text)-1] {
|
||||
case '"':
|
||||
quote = "\""
|
||||
_, after, found = strings.Cut(string(term.Location.Text), quote)
|
||||
quote = []byte{'"'}
|
||||
_, after, found = bytes.Cut(term.Location.Text, quote)
|
||||
case '`':
|
||||
quote = "`"
|
||||
_, after, found = strings.Cut(string(term.Location.Text), quote)
|
||||
quote = []byte{'`'}
|
||||
_, after, found = bytes.Cut(term.Location.Text, quote)
|
||||
}
|
||||
|
||||
if !found {
|
||||
@@ -1247,7 +1279,8 @@ func (w *writer) writeTermParens(parens bool, term *ast.Term, comments []*ast.Co
|
||||
// e.g. partial_set.y to partial_set["y"]
|
||||
w.write(x.String())
|
||||
} else {
|
||||
w.write(quote + after)
|
||||
w.writeBytes(quote)
|
||||
w.writeBytes(after)
|
||||
}
|
||||
|
||||
}
|
||||
@@ -1310,8 +1343,6 @@ func (w *writer) writeBracketed(str string) {
|
||||
w.write("[" + str + "]")
|
||||
}
|
||||
|
||||
var varRegexp = regexp.MustCompile("^[[:alpha:]_][[:alpha:][:digit:]_]*$")
|
||||
|
||||
func (w *writer) writeRefStringPath(s ast.String, l *ast.Location) {
|
||||
str := string(s)
|
||||
if w.shouldBracketRefTerm(str, l) {
|
||||
@@ -2130,11 +2161,16 @@ func (w *writer) blankLine() {
|
||||
w.write("\n")
|
||||
}
|
||||
|
||||
// write the input string and writes it to the buffer.
|
||||
// write writes string s to the buffer.
|
||||
func (w *writer) write(s string) {
|
||||
w.buf.WriteString(s)
|
||||
}
|
||||
|
||||
// writeBytes writes []byte b to the buffer.
|
||||
func (w *writer) writeBytes(b []byte) {
|
||||
w.buf.Write(b)
|
||||
}
|
||||
|
||||
// writeLine writes the string on a newly started line, then terminate the line.
|
||||
func (w *writer) writeLine(s string) {
|
||||
if !w.inline {
|
||||
|
||||
+13
-11
@@ -1004,24 +1004,26 @@ func TestFormatKeywordsInRefs(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// 382 3064960 ns/op 4573131 B/op 26266 allocs/op // no optimizations
|
||||
// 685 1737719 ns/op 1972193 B/op 14160 allocs/op // pre-allocate partitionComments
|
||||
// 708 1674343 ns/op 1916700 B/op 11556 allocs/op // static memberRef & memberWithKeyRef
|
||||
// 746 1594546 ns/op 1882652 B/op 10644 allocs/op // various minor fixes
|
||||
// 1250 853508 ns/op 441730 B/op 8895 allocs/op // partitionComments early return if unchanged
|
||||
// 1396 812859 ns/op 362651 B/op 8811 allocs/op // partitionComments reuse backing array
|
||||
// 3064960 ns/op 4573131 B/op 26266 allocs/op // no optimizations
|
||||
// 1737719 ns/op 1972193 B/op 14160 allocs/op // pre-allocate partitionComments
|
||||
// 1674343 ns/op 1916700 B/op 11556 allocs/op // static memberRef & memberWithKeyRef
|
||||
// 1594546 ns/op 1882652 B/op 10644 allocs/op // various minor fixes
|
||||
// _853508 ns/op 441730 B/op 8895 allocs/op // partitionComments early return if unchanged
|
||||
// _812859 ns/op 362651 B/op 8811 allocs/op // partitionComments reuse backing array
|
||||
// _676179 ns/op 995204 B/op 8850 allocs/op // regression in 0fb7526
|
||||
// _513570 ns/op 378787 B/op 8775 allocs/op // addressed regression with sync.Pool
|
||||
// _481681 ns/op 352130 B/op 7954 allocs/op // new util.NewSlicePool using only pointers
|
||||
// _365116 ns/op 160528 B/op 2098 allocs/op // new SkipDefensiveCopying option
|
||||
func BenchmarkFormatLargePolicy(b *testing.B) {
|
||||
contents, err := os.ReadFile("testdata/bench.rego")
|
||||
if err != nil {
|
||||
b.Fatalf("Failed to read rego source: %v", err)
|
||||
}
|
||||
module := ast.MustParseModule(string(contents))
|
||||
opts := Opts{RegoVersion: ast.RegoV1, SkipDefensiveCopying: true}
|
||||
|
||||
b.ResetTimer()
|
||||
|
||||
for range b.N {
|
||||
_, err := AstWithOpts(module, Opts{RegoVersion: ast.RegoV1})
|
||||
if err != nil {
|
||||
for b.Loop() {
|
||||
if _, err := AstWithOpts(module, opts); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"math"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
@@ -91,3 +92,49 @@ func SplitMap[T any](text string, delim string, fn func(string) T) []T {
|
||||
|
||||
return sl
|
||||
}
|
||||
|
||||
// SlicePool is a pool for (pointers to) slices of type T.
|
||||
// It uses sync.Pool to pool the slices, and grows them as needed.
|
||||
type SlicePool[T any] struct {
|
||||
pool sync.Pool
|
||||
}
|
||||
|
||||
// NewSlicePool creates a new SlicePool for slices of type T with the given initial length.
|
||||
// This number is only a hint, as the slices will grow as needed. For best performance, store
|
||||
// slices of similar lengths in the same pool.
|
||||
func NewSlicePool[T any](length int) *SlicePool[T] {
|
||||
return &SlicePool[T]{
|
||||
pool: sync.Pool{
|
||||
New: func() any {
|
||||
s := make([]T, length)
|
||||
return &s
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// Get returns a pointer to a slice of type T with the given length
|
||||
// from the pool. The slice capacity will grow as needed to accommodate
|
||||
// the requested length. The returned slice will have all its elements
|
||||
// set to the zero value of T. Returns a pointer to avoid allocating.
|
||||
func (sp *SlicePool[T]) Get(length int) *[]T {
|
||||
s := sp.pool.Get().(*[]T)
|
||||
d := *s
|
||||
|
||||
if cap(d) < length {
|
||||
d = slices.Grow(d, length)
|
||||
}
|
||||
|
||||
d = d[:length] // reslice to requested length, while keeping capacity
|
||||
|
||||
clear(d)
|
||||
|
||||
*s = d
|
||||
|
||||
return s
|
||||
}
|
||||
|
||||
// Put returns a pointer to a slice of type T to the pool.
|
||||
func (sp *SlicePool[T]) Put(s *[]T) {
|
||||
sp.pool.Put(s)
|
||||
}
|
||||
|
||||
@@ -59,6 +59,15 @@ func BenchmarkSplitMap(b *testing.B) {
|
||||
})
|
||||
}
|
||||
|
||||
// Zero allocations
|
||||
func BenchmarkSlicePoolGetPut(b *testing.B) {
|
||||
sp := NewSlicePool[int](4)
|
||||
for b.Loop() {
|
||||
s := sp.Get(4)
|
||||
sp.Put(s)
|
||||
}
|
||||
}
|
||||
|
||||
func mustAtoi(s string) int {
|
||||
v, _ := strconv.Atoi(s)
|
||||
return v
|
||||
|
||||
Reference in New Issue
Block a user