mirror of
https://github.com/open-policy-agent/opa.git
synced 2026-08-13 03:42:35 -06:00
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:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)()
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user