wasm_sdk: redo interrupt handling, pass server ctx (#3317)

Now, when an interrupt happens, we'll clean up after ourselves: we keep calling
a cheap function to ensure that the trap has been trapped on.

To get there, we'll move the "defer-recover" further down the call stack.

Also, this changes the cancellation error returned by the topdown builtin. It no
longer is a builtinError with a message indicating that the context was cancelled
(or its deadline reached), but return a CancelErr, the same thing that happens via
the other cancellation mechanisms involved topdown's Cancel (like in cidr.expand).

Compared with master, this isn't worse:

    name                       old time/op    new time/op    delta
    RESTAuthzForbidAuthn-16       542µs ±15%     525µs ±12%    ~     (p=0.841 n=5+5)
    RESTAuthzForbidPath-16        818µs ± 2%     827µs ± 2%    ~     (p=0.556 n=4+5)
    RESTAuthzForbidMethod-16      864µs ± 2%     851µs ± 4%    ~     (p=0.310 n=5+5)
    RESTAuthzAllow10Paths-16      896µs ±20%     855µs ± 5%    ~     (p=1.000 n=5+5)
    RESTAuthzAllow100Paths-16    4.28ms ± 3%    4.07ms ± 3%  -4.97%  (p=0.008 n=5+5)

    name                       old alloc/op   new alloc/op   delta
    RESTAuthzForbidAuthn-16      68.8kB ± 1%    67.2kB ± 1%  -2.28%  (p=0.008 n=5+5)
    RESTAuthzForbidPath-16       68.5kB ± 0%    66.9kB ± 0%  -2.33%  (p=0.008 n=5+5)
    RESTAuthzForbidMethod-16     68.5kB ± 0%    66.9kB ± 0%  -2.33%  (p=0.008 n=5+5)
    RESTAuthzAllow10Paths-16     68.5kB ± 0%    66.9kB ± 0%  -2.33%  (p=0.008 n=5+5)
    RESTAuthzAllow100Paths-16    69.1kB ± 0%    67.5kB ± 0%  -2.31%  (p=0.008 n=5+5)

    name                       old allocs/op  new allocs/op  delta
    RESTAuthzForbidAuthn-16       1.73k ± 1%     1.64k ± 1%  -5.17%  (p=0.008 n=5+5)
    RESTAuthzForbidPath-16        1.72k ± 0%     1.63k ± 0%    ~     (p=0.079 n=4+5)
    RESTAuthzForbidMethod-16      1.72k ± 0%     1.63k ± 0%  -5.13%  (p=0.008 n=5+5)
    RESTAuthzAllow10Paths-16      1.72k ± 0%     1.63k ± 0%  -5.13%  (p=0.008 n=5+5)
    RESTAuthzAllow100Paths-16     1.72k ± 0%     1.63k ± 0%  -5.16%  (p=0.008 n=5+5)

Signed-off-by: Stephan Renatus <stephan.renatus@gmail.com>
This commit is contained in:
Stephan Renatus
2021-03-30 13:48:45 +02:00
committed by GitHub
parent afe33a5d7c
commit 683aed99ba
17 changed files with 237 additions and 233 deletions
+1 -1
View File
@@ -92,7 +92,7 @@ func main() {
// Evaluate the new policy.
if err := rego.SetPolicy(policy); err != nil {
if err := rego.SetPolicy(ctx, policy); err != nil {
fmt.Printf("error: %v\n", err)
return
}
+8 -8
View File
@@ -153,7 +153,7 @@ func (p *Pool) Release(vm *VM, metrics metrics.Metrics) {
// are constructed in advance before touching the pool. Returns
// either ErrNotReady, ErrInvalidPolicy or ErrInternal if an error
// occurs.
func (p *Pool) SetPolicyData(policy []byte, data []byte) error {
func (p *Pool) SetPolicyData(ctx context.Context, policy []byte, data []byte) error {
p.dataMtx.Lock()
defer p.dataMtx.Unlock()
@@ -196,7 +196,7 @@ func (p *Pool) SetPolicyData(policy []byte, data []byte) error {
return nil
}
err := p.setPolicyData(policy, data)
err := p.setPolicyData(ctx, policy, data)
if err != nil {
return fmt.Errorf("%v: %w", err, errors.ErrInternal)
}
@@ -207,31 +207,31 @@ func (p *Pool) SetPolicyData(policy []byte, data []byte) error {
// SetDataPath will update the current data on the VMs by setting the value at the
// specified path. If an error occurs the instance is still in a valid state, however
// the data will not have been modified.
func (p *Pool) SetDataPath(path []string, value interface{}) error {
func (p *Pool) SetDataPath(ctx context.Context, path []string, value interface{}) error {
p.dataMtx.Lock()
defer p.dataMtx.Unlock()
return p.updateVMs(func(vm *VM, opts vmOpts) error {
return vm.SetDataPath(path, value)
return vm.SetDataPath(ctx, path, value)
})
}
// RemoveDataPath will update the current data on the VMs by removing the value at the
// specified path. If an error occurs the instance is still in a valid state, however
// the data will not have been modified.
func (p *Pool) RemoveDataPath(path []string) error {
func (p *Pool) RemoveDataPath(ctx context.Context, path []string) error {
p.dataMtx.Lock()
defer p.dataMtx.Unlock()
return p.updateVMs(func(vm *VM, _ vmOpts) error {
return vm.RemoveDataPath(path)
return vm.RemoveDataPath(ctx, path)
})
}
// setPolicyData reinitializes the VMs one at a time.
func (p *Pool) setPolicyData(policy []byte, data []byte) error {
func (p *Pool) setPolicyData(ctx context.Context, policy []byte, data []byte) error {
return p.updateVMs(func(vm *VM, opts vmOpts) error {
opts.policy = policy
opts.data = data
return vm.SetPolicyData(opts)
return vm.SetPolicyData(ctx, opts)
})
}
+17 -12
View File
@@ -20,6 +20,7 @@ import (
)
func TestPoolCopyParsedDataOnInit(t *testing.T) {
ctx := context.Background()
module := `package test
p = data.a
@@ -45,10 +46,12 @@ func TestPoolCopyParsedDataOnInit(t *testing.T) {
poolSize := 4
testPool := initPoolWithData(t, uint32(poolSize), module, "test/p", data)
expected := `{{"result":{"b":[1,2,3,{"d":{"e":{"f":123}},"c":4}]}}}`
ensurePoolResults(t, testPool, poolSize, expected)
ensurePoolResults(t, ctx, testPool, poolSize, expected)
}
func TestPoolCopyParsedDataUpdateFull(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
module := `package test
p = data.a
@@ -59,24 +62,24 @@ func TestPoolCopyParsedDataUpdateFull(t *testing.T) {
testPool := initPoolWithData(t, uint32(poolSize), module, "test/p", data)
updated := []byte(`{"a": {"x": 123, "y": "bar"}}`)
err := testPool.SetPolicyData(testPool.Policy(), updated)
err := testPool.SetPolicyData(ctx, testPool.Policy(), updated)
if err != nil {
t.Fatalf("Unexpected error: %s", err)
}
expected := `{{"result":{"y":"bar","x":123}}}`
ensurePoolResults(t, testPool, poolSize, expected)
ensurePoolResults(t, ctx, testPool, poolSize, expected)
// Change it one more time, now that all VM's in the pool have been
// initialized and exercised at least once.
updated = []byte(`{"a": [1, 2, 3]}`)
err = testPool.SetPolicyData(testPool.Policy(), updated)
err = testPool.SetPolicyData(ctx, testPool.Policy(), updated)
if err != nil {
t.Fatalf("Unexpected error: %s", err)
}
expected = `{{"result":[1,2,3]}}`
ensurePoolResults(t, testPool, poolSize, expected)
ensurePoolResults(t, ctx, testPool, poolSize, expected)
}
func TestPoolCopyParsedDataUpdatePartial(t *testing.T) {
@@ -123,34 +126,36 @@ func TestPoolCopyParsedDataUpdatePartial(t *testing.T) {
}
for _, tc := range cases {
t.Run(tc.note, func(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
var err error
if tc.remove {
err = testPool.RemoveDataPath(tc.path)
err = testPool.RemoveDataPath(ctx, tc.path)
} else {
err = testPool.SetDataPath(tc.path, tc.update)
err = testPool.SetDataPath(ctx, tc.path, tc.update)
}
if err != nil {
t.Fatalf("Unexpected error: %s", err)
}
ensurePoolResults(t, testPool, poolSize, tc.expected)
ensurePoolResults(t, ctx, testPool, poolSize, tc.expected)
})
}
}
func ensurePoolResults(t *testing.T, testPool *wasm.Pool, poolSize int, expected string) {
func ensurePoolResults(t *testing.T, ctx context.Context, testPool *wasm.Pool, poolSize int, expected string) {
t.Helper()
var toRelease []*wasm.VM
for i := 0; i < poolSize; i++ {
vm, err := testPool.Acquire(context.Background(), metrics.New())
vm, err := testPool.Acquire(ctx, metrics.New())
if err != nil {
t.Fatalf("Unexpected error: %s", err)
}
toRelease = append(toRelease, vm)
result, err := vm.Eval(context.Background(), 0, nil, metrics.New(), time.Now())
result, err := vm.Eval(ctx, 0, nil, metrics.New(), time.Now())
if err != nil {
t.Fatalf("Unexpected error: %s", err)
}
@@ -189,7 +194,7 @@ func initPoolWithData(t *testing.T, size uint32, module string, entrypoint strin
testPool := wasm.NewPool(size, 16, 0)
err = testPool.SetPolicyData(compiler.Bundle().WasmModules[0].Raw, data)
err = testPool.SetPolicyData(ctx, compiler.Bundle().WasmModules[0].Raw, data)
if err != nil {
t.Fatalf("Unexpected error: %s", err)
}
+151 -156
View File
@@ -39,22 +39,22 @@ type VM struct {
baseHeapPtr int32
dataAddr int32
evalHeapPtr int32
eval func(int32) error
evalCtxGetResult func(int32) (int32, error)
evalCtxNew func() (int32, error)
evalCtxSetData func(int32, int32) error
evalCtxSetInput func(int32, int32) error
evalCtxSetEntrypoint func(int32, int32) error
heapPtrGet func() (int32, error)
heapPtrSet func(int32) error
jsonDump func(int32) (int32, error)
jsonParse func(int32, int32) (int32, error)
valueDump func(int32) (int32, error)
valueParse func(int32, int32) (int32, error)
malloc func(int32) (int32, error)
free func(int32) error
valueAddPath func(int32, int32, int32) (int32, error)
valueRemovePath func(int32, int32) (int32, error)
eval func(context.Context, int32) error
evalCtxGetResult func(context.Context, int32) (int32, error)
evalCtxNew func(context.Context) (int32, error)
evalCtxSetData func(context.Context, int32, int32) error
evalCtxSetInput func(context.Context, int32, int32) error
evalCtxSetEntrypoint func(context.Context, int32, int32) error
heapPtrGet func(context.Context) (int32, error)
heapPtrSet func(context.Context, int32) error
jsonDump func(context.Context, int32) (int32, error)
jsonParse func(context.Context, int32, int32) (int32, error)
valueDump func(context.Context, int32) (int32, error)
valueParse func(context.Context, int32, int32) (int32, error)
malloc func(context.Context, int32) (int32, error)
free func(context.Context, int32) error
valueAddPath func(context.Context, int32, int32, int32) (int32, error)
valueRemovePath func(context.Context, int32, int32) (int32, error)
}
type vmOpts struct {
@@ -67,6 +67,7 @@ type vmOpts struct {
}
func newVM(opts vmOpts) (*VM, error) {
ctx := context.Background()
v := &VM{}
cfg := wasmtime.NewConfig()
cfg.SetInterruptable(true)
@@ -101,30 +102,44 @@ func newVM(opts vmOpts) (*VM, error) {
v.memoryMax = opts.memoryMax
v.entrypointIDs = make(map[string]int32)
v.dataAddr = 0
v.eval = func(a int32) error { return callVoid(v, "eval", a) }
v.evalCtxGetResult = func(a int32) (int32, error) { return call(v, "opa_eval_ctx_get_result", a) }
v.evalCtxNew = func() (int32, error) { return call(v, "opa_eval_ctx_new") }
v.evalCtxSetData = func(a int32, b int32) error { return callVoid(v, "opa_eval_ctx_set_data", a, b) }
v.evalCtxSetInput = func(a int32, b int32) error { return callVoid(v, "opa_eval_ctx_set_input", a, b) }
v.evalCtxSetEntrypoint = func(a int32, b int32) error { return callVoid(v, "opa_eval_ctx_set_entrypoint", a, b) }
v.free = func(a int32) error { return callVoid(v, "opa_free", a) }
v.heapPtrGet = func() (int32, error) { return call(v, "opa_heap_ptr_get") }
v.heapPtrSet = func(a int32) error { return callVoid(v, "opa_heap_ptr_set", a) }
v.jsonDump = func(a int32) (int32, error) { return call(v, "opa_json_dump", a) }
v.jsonParse = func(a int32, b int32) (int32, error) { return call(v, "opa_json_parse", a, b) }
v.valueDump = func(a int32) (int32, error) { return call(v, "opa_value_dump", a) }
v.valueParse = func(a int32, b int32) (int32, error) { return call(v, "opa_value_parse", a, b) }
v.malloc = func(a int32) (int32, error) { return call(v, "opa_malloc", a) }
v.valueAddPath = func(a int32, b int32, c int32) (int32, error) { return call(v, "opa_value_add_path", a, b, c) }
v.valueRemovePath = func(a int32, b int32) (int32, error) { return call(v, "opa_value_remove_path", a, b) }
v.eval = func(ctx context.Context, a int32) error { return callVoid(ctx, v, "eval", a) }
v.evalCtxGetResult = func(ctx context.Context, a int32) (int32, error) { return call(ctx, v, "opa_eval_ctx_get_result", a) }
v.evalCtxNew = func(ctx context.Context) (int32, error) { return call(ctx, v, "opa_eval_ctx_new") }
v.evalCtxSetData = func(ctx context.Context, a int32, b int32) error {
return callVoid(ctx, v, "opa_eval_ctx_set_data", a, b)
}
v.evalCtxSetInput = func(ctx context.Context, a int32, b int32) error {
return callVoid(ctx, v, "opa_eval_ctx_set_input", a, b)
}
v.evalCtxSetEntrypoint = func(ctx context.Context, a int32, b int32) error {
return callVoid(ctx, v, "opa_eval_ctx_set_entrypoint", a, b)
}
v.free = func(ctx context.Context, a int32) error { return callVoid(ctx, v, "opa_free", a) }
v.heapPtrGet = func(ctx context.Context) (int32, error) { return call(ctx, v, "opa_heap_ptr_get") }
v.heapPtrSet = func(ctx context.Context, a int32) error { return callVoid(ctx, v, "opa_heap_ptr_set", a) }
v.jsonDump = func(ctx context.Context, a int32) (int32, error) { return call(ctx, v, "opa_json_dump", a) }
v.jsonParse = func(ctx context.Context, a int32, b int32) (int32, error) {
return call(ctx, v, "opa_json_parse", a, b)
}
v.valueDump = func(ctx context.Context, a int32) (int32, error) { return call(ctx, v, "opa_value_dump", a) }
v.valueParse = func(ctx context.Context, a int32, b int32) (int32, error) {
return call(ctx, v, "opa_value_parse", a, b)
}
v.malloc = func(ctx context.Context, a int32) (int32, error) { return call(ctx, v, "opa_malloc", a) }
v.valueAddPath = func(ctx context.Context, a int32, b int32, c int32) (int32, error) {
return call(ctx, v, "opa_value_add_path", a, b, c)
}
v.valueRemovePath = func(ctx context.Context, a int32, b int32) (int32, error) {
return call(ctx, v, "opa_value_remove_path", a, b)
}
// Initialize the heap.
if _, err := v.malloc(0); err != nil {
if _, err := v.malloc(ctx, 0); err != nil {
return nil, err
}
if v.baseHeapPtr, err = v.getHeapState(); err != nil {
if v.baseHeapPtr, err = v.getHeapState(ctx); err != nil {
return nil, err
}
@@ -143,17 +158,17 @@ func newVM(opts vmOpts) (*VM, error) {
}
v.dataAddr = opts.parsedDataAddr
v.evalHeapPtr = v.baseHeapPtr + int32(len(opts.parsedData))
err := v.setHeapState(v.evalHeapPtr)
err := v.setHeapState(ctx, v.evalHeapPtr)
if err != nil {
return nil, err
}
} else if opts.data != nil {
if v.dataAddr, err = v.toRegoJSON(opts.data, true); err != nil {
if v.dataAddr, err = v.toRegoJSON(ctx, opts.data, true); err != nil {
return nil, err
}
}
if v.evalHeapPtr, err = v.getHeapState(); err != nil {
if v.evalHeapPtr, err = v.getHeapState(ctx); err != nil {
return nil, err
}
@@ -164,7 +179,7 @@ func newVM(opts vmOpts) (*VM, error) {
return nil, err
}
builtins, err := v.fromRegoJSON(val.(int32), true)
builtins, err := v.fromRegoJSON(ctx, val.(int32), true)
if err != nil {
return nil, err
}
@@ -193,7 +208,7 @@ func newVM(opts vmOpts) (*VM, error) {
return nil, err
}
epMap, err := v.fromRegoJSON(val.(int32), true)
epMap, err := v.fromRegoJSON(ctx, val.(int32), true)
if err != nil {
return nil, err
}
@@ -215,10 +230,6 @@ func (i *VM) Eval(ctx context.Context, entrypoint int32, input *interface{}, met
metrics.Timer("wasm_vm_eval").Start()
defer metrics.Timer("wasm_vm_eval").Stop()
if err := i.clearInterrupts(ctx); err != nil {
return nil, fmt.Errorf("clear interrupts: %w", err)
}
metrics.Timer("wasm_vm_eval_prepare_input").Start()
// Setting the ctx here ensures that it'll be available to builtins that
@@ -227,45 +238,34 @@ func (i *VM) Eval(ctx context.Context, entrypoint int32, input *interface{}, met
// cancelled.
i.dispatcher.Reset(ctx, ns)
// Interrupt the VM if the context is cancelled.
done := make(chan struct{})
defer close(done)
go func() {
select {
case <-done:
case <-ctx.Done():
i.intHandle.Interrupt()
}
}()
err := i.setHeapState(i.evalHeapPtr)
err := i.setHeapState(ctx, i.evalHeapPtr)
if err != nil {
return nil, err
}
// Parse the input JSON and activate it with the data.
ctxAddr, err := i.evalCtxNew()
ctxAddr, err := i.evalCtxNew(ctx)
if err != nil {
return nil, err
}
if i.dataAddr != 0 {
if err := i.evalCtxSetData(ctxAddr, i.dataAddr); err != nil {
if err := i.evalCtxSetData(ctx, ctxAddr, i.dataAddr); err != nil {
return nil, err
}
}
if err := i.evalCtxSetEntrypoint(ctxAddr, int32(entrypoint)); err != nil {
if err := i.evalCtxSetEntrypoint(ctx, ctxAddr, int32(entrypoint)); err != nil {
return nil, err
}
if input != nil {
inputAddr, err := i.toRegoJSON(*input, false)
inputAddr, err := i.toRegoJSON(ctx, *input, false)
if err != nil {
return nil, err
}
if err := i.evalCtxSetInput(ctxAddr, inputAddr); err != nil {
if err := i.evalCtxSetInput(ctx, ctxAddr, inputAddr); err != nil {
return nil, err
}
}
@@ -273,41 +273,19 @@ func (i *VM) Eval(ctx context.Context, entrypoint int32, input *interface{}, met
// Evaluate the policy.
metrics.Timer("wasm_vm_eval_execute").Start()
func() {
defer func() {
if e := recover(); e != nil {
switch e := e.(type) {
case abortError:
err = fmt.Errorf(e.message)
case cancelledError:
err = errors.ErrCancelled
case builtinError:
err = e.err
if _, ok := err.(topdown.Halt); !ok {
err = nil
}
default:
panic(e)
}
}
}()
err = i.eval(ctxAddr)
}()
err = i.eval(ctx, ctxAddr)
metrics.Timer("wasm_vm_eval_execute").Stop()
if err != nil {
return nil, err
}
metrics.Timer("wasm_vm_eval_prepare_result").Start()
resultAddr, err := i.evalCtxGetResult(ctxAddr)
resultAddr, err := i.evalCtxGetResult(ctx, ctxAddr)
if err != nil {
return nil, err
}
serialized, err := i.valueDump(resultAddr)
serialized, err := i.valueDump(ctx, resultAddr)
if err != nil {
return nil, err
}
@@ -322,15 +300,12 @@ func (i *VM) Eval(ctx context.Context, entrypoint int32, input *interface{}, met
// Skip free'ing input and result JSON as the heap will be reset next round anyway.
return data[0:n], err
return data[0:n], nil
}
// SetPolicyData Will either update the VM's data or, if the policy changed,
// re-initialize the VM.
func (i *VM) SetPolicyData(opts vmOpts) error {
if err := i.clearInterrupts(context.TODO()); err != nil {
return fmt.Errorf("clear interrupts: %w", err)
}
func (i *VM) SetPolicyData(ctx context.Context, opts vmOpts) error {
if !bytes.Equal(opts.policy, i.policy) {
// Swap the instance to a new one, with new policy.
@@ -346,7 +321,7 @@ func (i *VM) SetPolicyData(opts vmOpts) error {
i.dataAddr = 0
var err error
if err = i.setHeapState(i.baseHeapPtr); err != nil {
if err = i.setHeapState(ctx, i.baseHeapPtr); err != nil {
return err
}
@@ -361,17 +336,17 @@ func (i *VM) SetPolicyData(opts vmOpts) error {
}
i.dataAddr = opts.parsedDataAddr
i.evalHeapPtr = i.baseHeapPtr + int32(len(opts.parsedData))
err := i.setHeapState(i.evalHeapPtr)
err := i.setHeapState(ctx, i.evalHeapPtr)
if err != nil {
return err
}
} else if opts.data != nil {
if i.dataAddr, err = i.toRegoJSON(opts.data, true); err != nil {
if i.dataAddr, err = i.toRegoJSON(ctx, opts.data, true); err != nil {
return err
}
}
if i.evalHeapPtr, err = i.getHeapState(); err != nil {
if i.evalHeapPtr, err = i.getHeapState(ctx); err != nil {
return err
}
@@ -409,29 +384,25 @@ func (i *VM) Entrypoints() map[string]int32 {
// SetDataPath will update the current data on the VM by setting the value at the
// specified path. If an error occurs the instance is still in a valid state, however
// the data will not have been modified.
func (i *VM) SetDataPath(path []string, value interface{}) error {
if err := i.clearInterrupts(context.TODO()); err != nil {
return fmt.Errorf("clear interrupts: %w", err)
}
func (i *VM) SetDataPath(ctx context.Context, path []string, value interface{}) error {
// Reset the heap ptr before patching the vm to try and keep any
// new allocations safe from subsequent heap resets on eval.
err := i.setHeapState(i.evalHeapPtr)
err := i.setHeapState(ctx, i.evalHeapPtr)
if err != nil {
return err
}
valueAddr, err := i.toRegoJSON(value, true)
valueAddr, err := i.toRegoJSON(ctx, value, true)
if err != nil {
return err
}
pathAddr, err := i.toRegoJSON(path, true)
pathAddr, err := i.toRegoJSON(ctx, path, true)
if err != nil {
return err
}
result, err := i.valueAddPath(i.dataAddr, pathAddr, valueAddr)
result, err := i.valueAddPath(ctx, i.dataAddr, pathAddr, valueAddr)
if err != nil {
return err
}
@@ -440,13 +411,13 @@ func (i *VM) SetDataPath(path []string, value interface{}) error {
// overall data object now.
// We do need to free the path
if err := i.free(pathAddr); err != nil {
if err := i.free(ctx, pathAddr); err != nil {
return err
}
// Update the eval heap pointer to accommodate for any new allocations done
// while patching.
i.evalHeapPtr, err = i.getHeapState()
i.evalHeapPtr, err = i.getHeapState(ctx)
if err != nil {
return err
}
@@ -462,22 +433,18 @@ func (i *VM) SetDataPath(path []string, value interface{}) error {
// RemoveDataPath will update the current data on the VM by removing the value at the
// specified path. If an error occurs the instance is still in a valid state, however
// the data will not have been modified.
func (i *VM) RemoveDataPath(path []string) error {
if err := i.clearInterrupts(context.TODO()); err != nil {
return fmt.Errorf("clear interrupts: %w", err)
}
pathAddr, err := i.toRegoJSON(path, true)
func (i *VM) RemoveDataPath(ctx context.Context, path []string) error {
pathAddr, err := i.toRegoJSON(ctx, path, true)
if err != nil {
return err
}
errc, err := i.valueRemovePath(i.dataAddr, pathAddr)
errc, err := i.valueRemovePath(ctx, i.dataAddr, pathAddr)
if err != nil {
return err
}
if err := i.free(pathAddr); err != nil {
if err := i.free(ctx, pathAddr); err != nil {
return err
}
@@ -490,8 +457,8 @@ func (i *VM) RemoveDataPath(path []string) error {
// fromRegoJSON parses serialized JSON from the Wasm memory buffer into
// native go types.
func (i *VM) fromRegoJSON(addr int32, free bool) (interface{}, error) {
serialized, err := i.jsonDump(addr)
func (i *VM) fromRegoJSON(ctx context.Context, addr int32, free bool) (interface{}, error) {
serialized, err := i.jsonDump(ctx, addr)
if err != nil {
return nil, err
}
@@ -513,7 +480,7 @@ func (i *VM) fromRegoJSON(addr int32, free bool) (interface{}, error) {
}
if free {
if err := i.free(serialized); err != nil {
if err := i.free(ctx, serialized); err != nil {
return nil, err
}
}
@@ -523,7 +490,7 @@ func (i *VM) fromRegoJSON(addr int32, free bool) (interface{}, error) {
// toRegoJSON converts go native JSON to Rego JSON. If the value is
// an AST type it will be dumped using its stringer.
func (i *VM) toRegoJSON(v interface{}, free bool) (int32, error) {
func (i *VM) toRegoJSON(ctx context.Context, v interface{}, free bool) (int32, error) {
var raw []byte
switch v := v.(type) {
case []byte:
@@ -541,20 +508,20 @@ func (i *VM) toRegoJSON(v interface{}, free bool) (int32, error) {
}
n := int32(len(raw))
p, err := i.malloc(n)
p, err := i.malloc(ctx, n)
if err != nil {
return 0, err
}
copy(i.memory.UnsafeData()[p:p+n], raw)
addr, err := i.valueParse(p, n)
addr, err := i.valueParse(ctx, p, n)
if err != nil {
return 0, err
}
if free {
if err := i.free(p); err != nil {
if err := i.free(ctx, p); err != nil {
return 0, err
}
}
@@ -564,8 +531,8 @@ func (i *VM) toRegoJSON(v interface{}, free bool) (int32, error) {
// fromRegoValue parses serialized opa values from the Wasm memory buffer into
// Rego AST types.
func (i *VM) fromRegoValue(addr int32, free bool) (*ast.Term, error) {
serialized, err := i.valueDump(addr)
func (i *VM) fromRegoValue(ctx context.Context, addr int32, free bool) (*ast.Term, error) {
serialized, err := i.valueDump(ctx, addr)
if err != nil {
return nil, err
}
@@ -580,7 +547,7 @@ func (i *VM) fromRegoValue(addr int32, free bool) (*ast.Term, error) {
result, err := ast.ParseTerm(string(data[0:n]))
if free {
if err := i.free(serialized); err != nil {
if err := i.free(ctx, serialized); err != nil {
return nil, err
}
}
@@ -588,12 +555,12 @@ func (i *VM) fromRegoValue(addr int32, free bool) (*ast.Term, error) {
return result, err
}
func (i *VM) getHeapState() (int32, error) {
return i.heapPtrGet()
func (i *VM) getHeapState(ctx context.Context) (int32, error) {
return i.heapPtrGet(ctx)
}
func (i *VM) setHeapState(ptr int32) error {
return i.heapPtrSet(ptr)
func (i *VM) setHeapState(ctx context.Context, ptr int32) error {
return i.heapPtrSet(ctx, ptr)
}
func (i *VM) cloneDataSegment() (int32, []byte) {
@@ -605,45 +572,73 @@ func (i *VM) cloneDataSegment() (int32, []byte) {
return i.dataAddr, patchedData
}
func call(vm *VM, name string, args ...int32) (int32, error) {
sl := make([]interface{}, len(args))
for i := range sl {
sl[i] = args[i]
}
x, err := vm.instance.GetExport(name).Func().Call(sl...)
func call(ctx context.Context, vm *VM, name string, args ...int32) (int32, error) {
res, err := callOrCancel(ctx, vm, name, args...)
if err != nil {
return 0, err
}
return x.(int32), nil
return res.(int32), nil
}
func callVoid(vm *VM, name string, args ...int32) error {
func callVoid(ctx context.Context, vm *VM, name string, args ...int32) error {
_, err := callOrCancel(ctx, vm, name, args...)
return err
}
func callOrCancel(ctx context.Context, vm *VM, name string, args ...int32) (interface{}, error) {
sl := make([]interface{}, len(args))
for i := range sl {
sl[i] = args[i]
}
_, err := vm.instance.GetExport(name).Func().Call(sl...)
return err
var res interface{}
errc := make(chan error) // unbuffered, we'll read it once at least
go func() {
var err error
// If this call into the VM ends up calling host functions (builtins not
// implemented in Wasm), and those panic, wasmtime will re-throw them,
// and this is where we deal with that:
defer func() {
if e := recover(); e != nil {
switch e := e.(type) {
case abortError:
errc <- fmt.Errorf(e.message)
case cancelledError:
errc <- errors.ErrCancelled
case builtinError:
if _, ok := e.err.(topdown.Halt); !ok {
errc <- nil
return
}
errc <- e.err
default:
panic(e)
}
}
}()
res, err = vm.instance.GetExport(name).Func().Call(sl...)
errc <- err
}()
select {
case <-ctx.Done():
// interrupt, wait for trap
vm.intHandle.Interrupt()
case err := <-errc:
if err != nil {
return 0, err
}
return res, nil
}
// clear trap: reaching this, we have definitely been interrupted
err := <-errc
for err != errors.ErrCancelled && !trap(err) {
_, err = vm.instance.GetExport("opa_heap_ptr_get").Func().Call() // cheap call
}
return 0, errors.ErrCancelled
}
func (i *VM) clearInterrupts(ctx context.Context) error {
// NOTE: It doesn't matter which exported function of the wasm module we call,
// any set traps will trigger. Let's call a cheap one.
_, err := i.heapPtrGet()
if err == nil {
return nil
}
if t, ok := err.(*wasmtime.Trap); ok && strings.HasPrefix(t.Message(), "wasm trap: interrupt") {
// Check if OUR ctx is done. If it wasn't, we've triggered the
// trap of a previous evaluation.
select {
case <-ctx.Done(): // chan closed
return errors.ErrCancelled
default: // don't block
}
return nil
}
return err
func trap(err error) bool {
t, ok := err.(*wasmtime.Trap)
return ok && strings.HasPrefix(t.Message(), "wasm trap: interrupt")
}
+2 -2
View File
@@ -38,7 +38,7 @@ type Loader struct {
// policyData captures the functions used in setting the policy and data.
type policyData interface {
SetPolicyData(policy []byte, data *interface{}) error
SetPolicyData(ctx context.Context, policy []byte, data *interface{}) error
}
// New constructs a new file loader periodically reloading the bundle
@@ -143,7 +143,7 @@ func (l *Loader) Load(ctx context.Context) error {
data = &v
}
return l.pd.SetPolicyData(b.WasmModules[0].Raw, data)
return l.pd.SetPolicyData(ctx, b.WasmModules[0].Raw, data)
}
// poller periodically downloads the bundle.
@@ -79,7 +79,7 @@ type testPolicyData struct {
updated chan struct{}
}
func (pd *testPolicyData) SetPolicyData(policy []byte, data *interface{}) error {
func (pd *testPolicyData) SetPolicyData(_ context.Context, policy []byte, data *interface{}) error {
pd.Lock()
defer pd.Unlock()
+3 -4
View File
@@ -15,9 +15,8 @@ import (
"time"
"github.com/open-policy-agent/opa/bundle"
"github.com/open-policy-agent/opa/internal/wasm/sdk/opa/errors"
"github.com/open-policy-agent/opa/internal/wasm/sdk/opa"
"github.com/open-policy-agent/opa/internal/wasm/sdk/opa/errors"
)
const (
@@ -54,7 +53,7 @@ type Loader struct {
// policyData captures the functions used in setting the policy and data.
type policyData interface {
SetPolicyData(policy []byte, data *interface{}) error
SetPolicyData(ctx context.Context, policy []byte, data *interface{}) error
}
// New constructs a new HTTP loader periodically downloading a bundle
@@ -199,7 +198,7 @@ func (l *Loader) Load(ctx context.Context) error {
data = &v
}
return l.pd.SetPolicyData(bundle.WasmModules[0].Raw, data)
return l.pd.SetPolicyData(ctx, bundle.WasmModules[0].Raw, data)
}
// get executes HTTP GET.
@@ -93,7 +93,7 @@ type testPolicyData struct {
updated chan struct{}
}
func (pd *testPolicyData) SetPolicyData(policy []byte, data *interface{}) error {
func (pd *testPolicyData) SetPolicyData(_ context.Context, policy []byte, data *interface{}) error {
pd.Lock()
defer pd.Unlock()
+14 -13
View File
@@ -58,6 +58,7 @@ func New() *OPA {
// configuration. If the configuration is invalid, it returns
// ErrInvalidConfig.
func (o *OPA) Init() (*OPA, error) {
ctx := context.Background()
if o.configErr != nil {
return nil, o.configErr
}
@@ -65,7 +66,7 @@ func (o *OPA) Init() (*OPA, error) {
o.pool = wasm.NewPool(o.poolSize, o.memoryMinPages, o.memoryMaxPages)
if len(o.policy) != 0 {
if err := o.pool.SetPolicyData(o.policy, o.data); err != nil {
if err := o.pool.SetPolicyData(ctx, o.policy, o.data); err != nil {
return nil, err
}
}
@@ -76,7 +77,7 @@ func (o *OPA) Init() (*OPA, error) {
// SetData updates the data for the subsequent Eval calls. Returns
// either ErrNotReady, ErrInvalidPolicyOrData, or ErrInternal if an
// error occurs.
func (o *OPA) SetData(v interface{}) error {
func (o *OPA) SetData(ctx context.Context, v interface{}) error {
if o.pool == nil {
return errors.ErrNotReady
}
@@ -89,27 +90,27 @@ func (o *OPA) SetData(v interface{}) error {
o.mutex.Lock()
defer o.mutex.Unlock()
return o.setPolicyData(o.policy, raw)
return o.setPolicyData(ctx, o.policy, raw)
}
// SetDataPath will update the current data on the VMs by setting the value at the
// specified path. If an error occurs the instance is still in a valid state, however
// the data will not have been modified.
func (o *OPA) SetDataPath(path []string, value interface{}) error {
return o.pool.SetDataPath(path, value)
func (o *OPA) SetDataPath(ctx context.Context, path []string, value interface{}) error {
return o.pool.SetDataPath(ctx, path, value)
}
// RemoveDataPath will update the current data on the VMs by removing the value at the
// specified path. If an error occurs the instance is still in a valid state, however
// the data will not have been modified.
func (o *OPA) RemoveDataPath(path []string) error {
return o.pool.RemoveDataPath(path)
func (o *OPA) RemoveDataPath(ctx context.Context, path []string) error {
return o.pool.RemoveDataPath(ctx, path)
}
// SetPolicy updates the policy for the subsequent Eval calls.
// Returns either ErrNotReady, ErrInvalidPolicy or ErrInternal if an
// error occurs.
func (o *OPA) SetPolicy(p []byte) error {
func (o *OPA) SetPolicy(ctx context.Context, p []byte) error {
if o.pool == nil {
return errors.ErrNotReady
}
@@ -117,13 +118,13 @@ func (o *OPA) SetPolicy(p []byte) error {
o.mutex.Lock()
defer o.mutex.Unlock()
return o.setPolicyData(p, o.data)
return o.setPolicyData(ctx, p, o.data)
}
// SetPolicyData updates both the policy and data for the subsequent
// Eval calls. Returns either ErrNotReady, ErrInvalidPolicyOrData, or
// ErrInternal if an error occurs.
func (o *OPA) SetPolicyData(policy []byte, data *interface{}) error {
func (o *OPA) SetPolicyData(ctx context.Context, policy []byte, data *interface{}) error {
if o.pool == nil {
return errors.ErrNotReady
}
@@ -140,11 +141,11 @@ func (o *OPA) SetPolicyData(policy []byte, data *interface{}) error {
o.mutex.Lock()
defer o.mutex.Unlock()
return o.setPolicyData(policy, raw)
return o.setPolicyData(ctx, policy, raw)
}
func (o *OPA) setPolicyData(policy []byte, data []byte) error {
if err := o.pool.SetPolicyData(policy, data); err != nil {
func (o *OPA) setPolicyData(ctx context.Context, policy []byte, data []byte) error {
if err := o.pool.SetPolicyData(ctx, policy, data); err != nil {
return err
}
+6 -4
View File
@@ -236,6 +236,8 @@ a = "c" { input > 2 }`,
for _, test := range tests {
t.Run(test.Description, func(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
policy := compileRegoToWasm(test.Policy, test.Query, dump)
data := []byte(test.Data)
if len(data) == 0 {
@@ -257,24 +259,24 @@ a = "c" { input > 2 }`,
case eval.NewPolicy != "" && eval.NewData != "":
policy := compileRegoToWasm(eval.NewPolicy, test.Query, dump)
data := parseJSON(eval.NewData)
if err := instance.SetPolicyData(policy, data); err != nil {
if err := instance.SetPolicyData(ctx, policy, data); err != nil {
t.Errorf(err.Error())
}
case eval.NewPolicy != "":
policy := compileRegoToWasm(eval.NewPolicy, test.Query, dump)
if err := instance.SetPolicy(policy); err != nil {
if err := instance.SetPolicy(ctx, policy); err != nil {
t.Errorf(err.Error())
}
case eval.NewData != "":
data := parseJSON(eval.NewData)
if err := instance.SetData(*data); err != nil {
if err := instance.SetData(ctx, *data); err != nil {
t.Errorf(err.Error())
}
}
r, err := instance.Eval(context.Background(), opa.EvalOpts{Input: parseJSON(eval.Input)})
r, err := instance.Eval(ctx, opa.EvalOpts{Input: parseJSON(eval.Input)})
if err != nil {
if test.WantErr == "" { // no error desired
t.Fatal(err.Error())
+4 -4
View File
@@ -650,7 +650,7 @@ func (m *Manager) onCommit(ctx context.Context, txn storage.Transaction, event s
}
m.setWasmResolvers(resolvers)
} else {
err := m.updateWasmResolversData(event)
err := m.updateWasmResolversData(ctx, event)
if err != nil {
panic(err)
}
@@ -694,7 +694,7 @@ func requiresWasmResolverReload(event storage.TriggerEvent) bool {
return false
}
func (m *Manager) updateWasmResolversData(event storage.TriggerEvent) error {
func (m *Manager) updateWasmResolversData(ctx context.Context, event storage.TriggerEvent) error {
m.wasmResolversMtx.Lock()
defer m.wasmResolversMtx.Unlock()
@@ -706,9 +706,9 @@ func (m *Manager) updateWasmResolversData(event storage.TriggerEvent) error {
for _, dataEvent := range event.Data {
var err error
if dataEvent.Removed {
err = resolver.RemoveDataPath(dataEvent.Path)
err = resolver.RemoveDataPath(ctx, dataEvent.Path)
} else {
err = resolver.SetDataPath(dataEvent.Path, dataEvent.Data)
err = resolver.SetDataPath(ctx, dataEvent.Path, dataEvent.Data)
}
if err != nil {
return fmt.Errorf("failed to update wasm runtime data: %s", err)
+6 -9
View File
@@ -150,7 +150,6 @@ allow {
http.send({"method": "get", "url": "%s", "raise_error": true})
}`,
ts.URL)
httpSendTopdownError := fmt.Sprintf(`http.send: Get "%s": context deadline exceeded`, ts.URL)
// This is a natively-implemented (for the wasm target) function that
// takes long.
@@ -170,13 +169,10 @@ allow {
errorCheck: topdown.IsCancel,
},
{
note: "http.send",
target: "rego",
policy: httpSend,
errorCheck: func(err error) bool {
var te *topdown.Error
return errors.As(err, &te) && te.Message == httpSendTopdownError
},
note: "http.send",
target: "rego",
policy: httpSend,
errorCheck: topdown.IsCancel,
},
{
note: "numbers.range",
@@ -206,7 +202,8 @@ allow {
policy: numbersRange,
errorCheck: func(err error) bool {
return errors.Is(err, sdk_errors.ErrCancelled)
}},
},
},
} {
t.Run(tc.target+"/"+tc.note, func(t *testing.T) {
defer leaktest.Check(t)()
+5 -5
View File
@@ -29,27 +29,27 @@ func (r *Resolver) Close() {
}
// Eval unimplemented.
func (r *Resolver) Eval(ctx context.Context, input resolver.Input) (resolver.Result, error) {
func (r *Resolver) Eval(context.Context, resolver.Input) (resolver.Result, error) {
panic("unreachable")
}
// SetData unimplemented.
func (r *Resolver) SetData(data interface{}) error {
func (r *Resolver) SetData(context.Context, interface{}) error {
panic("unreachable")
}
// SetDataPath unimplemented.
func (r *Resolver) SetDataPath(path []string, data interface{}) error {
func (r *Resolver) SetDataPath(context.Context, []string, interface{}) error {
panic("unreachable")
}
// RemoveDataPath unimplemented.
func (r *Resolver) RemoveDataPath(path []string) error {
func (r *Resolver) RemoveDataPath(context.Context, []string) error {
panic("unreachable")
}
// New unimplemented. Will always return an error.
func New(entrypoints []ast.Ref, policy []byte, data interface{}) (*Resolver, error) {
func New([]ast.Ref, []byte, interface{}) (*Resolver, error) {
return nil, errors.New("WebAssembly runtime not supported in this build")
}
+6 -6
View File
@@ -120,18 +120,18 @@ func (r *Resolver) Eval(ctx context.Context, input resolver.Input) (resolver.Res
}
// SetData will update the external data for the Wasm instance.
func (r *Resolver) SetData(data interface{}) error {
return r.o.SetData(data)
func (r *Resolver) SetData(ctx context.Context, data interface{}) error {
return r.o.SetData(ctx, data)
}
// SetDataPath will set the provided data on the wasm instance at the specified path.
func (r *Resolver) SetDataPath(path []string, data interface{}) error {
return r.o.SetDataPath(path, data)
func (r *Resolver) SetDataPath(ctx context.Context, path []string, data interface{}) error {
return r.o.SetDataPath(ctx, path, data)
}
// RemoveDataPath will remove any data at the specified path.
func (r *Resolver) RemoveDataPath(path []string) error {
return r.o.RemoveDataPath(path)
func (r *Resolver) RemoveDataPath(ctx context.Context, path []string) error {
return r.o.RemoveDataPath(ctx, path)
}
func getResult(evalResult *opa.Result) (ast.Value, error) {
+1
View File
@@ -1502,6 +1502,7 @@ func (s *Server) v1DataPost(w http.ResponseWriter, r *http.Request) {
result.Explanation, err = types.NewTraceV1(*buf, pretty)
if err != nil {
writer.ErrorAuto(w, err)
return
}
}
err = logger.Log(ctx, txn, decisionID, r.RemoteAddr, urlPath, "", goInput, input, nil, nil, m)
+8
View File
@@ -158,6 +158,14 @@ func handleHTTPSendErr(bctx BuiltinContext, err error) error {
if urlErr, ok := err.(*url.Error); ok && urlErr.Timeout() && bctx.Context.Err() == nil {
err = fmt.Errorf("%s %s: request timed out", urlErr.Op, urlErr.URL)
}
if err := bctx.Context.Err(); err != nil {
return Halt{
Err: &Error{
Code: CancelErr,
Message: fmt.Sprintf("http.send: timed out (%s)", err.Error()),
},
}
}
return handleBuiltinErr(ast.HTTPSend.Name, bctx.Location, err)
}
+3 -7
View File
@@ -62,7 +62,7 @@ func TestHTTPSendTimeout(t *testing.T) {
evalTimeout: 500 * time.Millisecond,
serverDelay: 5 * time.Second,
defaultTimeout: 1 * time.Minute,
expected: &Error{Code: BuiltinErr, Message: "context deadline exceeded"},
expected: &Error{Code: CancelErr, Message: "timed out (context deadline exceeded)"},
},
{
note: "param timeout less than default",
@@ -86,7 +86,7 @@ func TestHTTPSendTimeout(t *testing.T) {
evalTimeout: 500 * time.Millisecond,
serverDelay: 5 * time.Second,
defaultTimeout: 1 * time.Minute,
expected: &Error{Code: BuiltinErr, Message: "context deadline exceeded"},
expected: &Error{Code: CancelErr, Message: "timed out (context deadline exceeded)"},
},
}
@@ -96,9 +96,8 @@ func TestHTTPSendTimeout(t *testing.T) {
tsMtx.Unlock()
ctx := context.Background()
var cancel context.CancelFunc
if tc.evalTimeout > 0 {
ctx, cancel = context.WithTimeout(ctx, tc.evalTimeout)
ctx, _ = context.WithTimeout(ctx, tc.evalTimeout)
}
// TODO(patrick-east): Remove this along with the environment variable so that the "default" can't change
@@ -116,8 +115,5 @@ func TestHTTPSendTimeout(t *testing.T) {
// Put back the default (may not have changed)
defaultHTTPRequestTimeout = originalDefaultTimeout
if cancel != nil {
cancel()
}
}
}