Add Partial function to Rego objects

Partial allows callers to partially evaluate a query and obtain the
resulting partially evaluated queries and support modules from the
evaluation engine.

This is a lower level interface compared to the other API (PartialEval)
that returns a result which can be evaluated normally (which takes
advantage of partial evaluation for optimization purposes.) This new
interface is targetted at use cases where callers are not interested in
a binary policy decision but instead need the conditions that should be
evaluated at a later time.

Deprecate the old PartialEval function and rename it to PartialResult to
avoid some (but certainly not all) confusion.

Signed-off-by: Torin Sandall <torinsandall@gmail.com>
This commit is contained in:
Torin Sandall
2018-06-26 20:20:21 -07:00
parent d3b9f53e1f
commit a55a542809
3 changed files with 163 additions and 46 deletions
+40 -1
View File
@@ -397,7 +397,7 @@ q = {1, 2, 3} { true }`,
// filename: example_error.rego
}
func ExampleRego_PartialEval() {
func ExampleRego_PartialResult() {
ctx := context.Background()
@@ -530,3 +530,42 @@ func ExampleRego_PartialEval() {
// input 2 allowed: false
// input 3 allowed: true
}
func ExampleRego_Partial() {
ctx := context.Background()
// Define a simple policy for example purposes.
module := `package test
allow {
input.method = read_methods[_]
input.path = ["reviews", user]
input.user = user
}
allow {
input.method = read_methods[_]
input.path = ["reviews", _]
input.is_admin
}
read_methods = ["GET"]
`
r := rego.New(rego.Query("data.test.allow == true"), rego.Module("example.rego", module))
pq, err := r.Partial(ctx)
if err != nil {
// Handle error.
}
// Inspect result.
for i := range pq.Queries {
fmt.Printf("Query #%d: %v\n", i+1, pq.Queries[i])
}
// Output:
//
// Query #1: "GET" = input.method; input.path = ["reviews", _]; neq(input.is_admin, false)
// Query #2: "GET" = input.method; input.path = ["reviews", user3]; user3 = input.user
}
+122 -44
View File
@@ -20,6 +20,13 @@ import (
const defaultPartialNamespace = "partial"
// PartialQueries contains the queries and support modules produced by partial
// evaluation.
type PartialQueries struct {
Queries []ast.Body `json:"queries,omitempty"`
Support []*ast.Module `json:"modules,omitempty"`
}
// PartialResult represents the result of partial evaluation. The result can be
// used to generate a new query that can be run when inputs are known.
type PartialResult struct {
@@ -118,6 +125,7 @@ type Rego struct {
rawInput *interface{}
input ast.Value
unknowns []string
parsedUnknowns []*ast.Term
partialNamespace string
modules []rawModule
compiler *ast.Compiler
@@ -199,6 +207,14 @@ func Unknowns(unknowns []string) func(r *Rego) {
}
}
// ParsedUnknowns returns an argument that sets the values to treat as unknown
// during partial evaluation.
func ParsedUnknowns(unknowns []*ast.Term) func(r *Rego) {
return func(r *Rego) {
r.parsedUnknowns = unknowns
}
}
// PartialNamespace returns an argument that sets the namespace to use for
// partial evaluation results. The namespace must be a valid package path
// component.
@@ -335,8 +351,13 @@ func (r *Rego) Eval(ctx context.Context) (ResultSet, error) {
return r.eval(ctx, qc, compiled, txn)
}
// PartialEval partially evaluates this Rego object and returns a PartialResult.
// PartialEval has been deprecated and renamed to PartialResult.
func (r *Rego) PartialEval(ctx context.Context) (PartialResult, error) {
return r.PartialResult(ctx)
}
// PartialResult partially evaluates this Rego object and returns a PartialResult.
func (r *Rego) PartialResult(ctx context.Context) (PartialResult, error) {
if len(r.query) == 0 && len(r.parsedQuery) == 0 {
return PartialResult{}, fmt.Errorf("cannot evaluate empty query")
@@ -373,7 +394,52 @@ func (r *Rego) PartialEval(ctx context.Context) (PartialResult, error) {
defer r.store.Abort(ctx, txn)
}
return r.partialEval(ctx, compiled, txn, ast.Wildcard)
partialNamespace := r.partialNamespace
if partialNamespace == "" {
partialNamespace = defaultPartialNamespace
}
return r.partialResult(ctx, compiled, txn, partialNamespace, ast.Wildcard)
}
// Partial runs partial evaluation on r and returns the result.
func (r *Rego) Partial(ctx context.Context) (*PartialQueries, error) {
if len(r.query) == 0 && len(r.parsedQuery) == 0 {
return nil, fmt.Errorf("cannot evaluate empty query")
}
parsed, query, err := r.parse()
if err != nil {
return nil, err
}
err = r.compileModules(parsed)
if err != nil {
return nil, err
}
_, compiled, err := r.compileQuery(nil, query)
if err != nil {
return nil, err
}
txn := r.txn
if txn == nil {
txn, err = r.store.NewTransaction(ctx)
if err != nil {
return nil, err
}
defer r.store.Abort(ctx, txn)
}
partialNamespace := r.partialNamespace
if partialNamespace == "" {
partialNamespace = defaultPartialNamespace
}
return r.partial(ctx, compiled, txn, partialNamespace)
}
func (r *Rego) parse() (map[string]*ast.Module, ast.Body, error) {
@@ -567,34 +633,69 @@ func (r *Rego) eval(ctx context.Context, qc ast.QueryCompiler, compiled ast.Body
return rs, nil
}
func (r *Rego) partialEval(ctx context.Context, compiled ast.Body, txn storage.Transaction, output *ast.Term) (PartialResult, error) {
func (r *Rego) partialResult(ctx context.Context, compiled ast.Body, txn storage.Transaction, partialNamespace string, output *ast.Term) (PartialResult, error) {
pq, err := r.partial(ctx, compiled, txn, partialNamespace)
if err != nil {
return PartialResult{}, err
}
// Construct module for queries.
module := ast.MustParseModule("package " + partialNamespace)
module.Rules = make([]*ast.Rule, len(pq.Queries))
for i, body := range pq.Queries {
module.Rules[i] = &ast.Rule{
Head: ast.NewHead(ast.Var("__result__"), nil, output),
Body: body,
Module: module,
}
}
// Update compiler with partial evaluation output.
r.compiler.Modules["__partialresult__"] = module
for i, module := range pq.Support {
r.compiler.Modules[fmt.Sprintf("__partialsupport%d__", i)] = module
}
r.compiler.Compile(r.compiler.Modules)
if r.compiler.Failed() {
return PartialResult{}, r.compiler.Errors
}
result := PartialResult{
compiler: r.compiler,
store: r.store,
body: ast.MustParseBody(fmt.Sprintf("data.%v.__result__", partialNamespace)),
}
return result, nil
}
func (r *Rego) partial(ctx context.Context, compiled ast.Body, txn storage.Transaction, partialNamespace string) (*PartialQueries, error) {
var unknowns []*ast.Term
// Use input document as unknown if caller has not specified any.
if r.unknowns == nil {
unknowns = []*ast.Term{ast.InputRootDocument}
} else {
if r.parsedUnknowns != nil {
unknowns = r.parsedUnknowns
} else if r.unknowns != nil {
unknowns = make([]*ast.Term, len(r.unknowns))
for i := range r.unknowns {
var err error
unknowns[i], err = ast.ParseTerm(r.unknowns[i])
if err != nil {
return PartialResult{}, err
return nil, err
}
}
}
partialNamespace := r.partialNamespace
if partialNamespace == "" {
partialNamespace = defaultPartialNamespace
} else {
// Use input document as unknown if caller has not specified any.
unknowns = []*ast.Term{ast.InputRootDocument}
}
// Check partial namespace to ensure it's valid.
if term, err := ast.ParseTerm(partialNamespace); err != nil {
return PartialResult{}, err
return nil, err
} else if _, ok := term.Value.(ast.Var); !ok {
return PartialResult{}, fmt.Errorf("bad partial namespace")
return nil, fmt.Errorf("bad partial namespace")
}
q := topdown.NewQuery(compiled).
@@ -623,40 +724,17 @@ func (r *Rego) partialEval(ctx context.Context, compiled ast.Body, txn storage.T
c.Cancel()
})
partials, support, err := q.PartialRun(ctx)
queries, support, err := q.PartialRun(ctx)
if err != nil {
return PartialResult{}, err
return nil, err
}
// Construct module for queries.
module := ast.MustParseModule("package " + partialNamespace)
module.Rules = make([]*ast.Rule, len(partials))
for i, body := range partials {
module.Rules[i] = &ast.Rule{
Head: ast.NewHead(ast.Var("__result__"), nil, output),
Body: body,
Module: module,
}
pq := &PartialQueries{
Queries: queries,
Support: support,
}
// Update compiler with partial evaluation output.
r.compiler.Modules["__partialresult__"] = module
for i, module := range support {
r.compiler.Modules[fmt.Sprintf("__partialsupport%d__", i)] = module
}
r.compiler.Compile(r.compiler.Modules)
if r.compiler.Failed() {
return PartialResult{}, r.compiler.Errors
}
result := PartialResult{
compiler: r.compiler,
store: r.store,
body: ast.MustParseBody(fmt.Sprintf("data.%v.__result__", partialNamespace)),
}
return result, nil
return pq, nil
}
func (r *Rego) rewriteQueryToCaptureValue(qc ast.QueryCompiler, query ast.Body) (ast.Body, error) {
+1 -1
View File
@@ -1390,7 +1390,7 @@ func (s *Server) makeRego(ctx context.Context, partial bool, txn storage.Transac
opts = append(opts, rego.Transaction(txn), rego.Query(path), rego.Metrics(m), rego.Instrument(instrument))
r := rego.New(opts...)
var err error
pr, err = r.PartialEval(ctx)
pr, err = r.PartialResult(ctx)
if err != nil {
return nil, err
}