package kenro
import (
"context"
"fmt"
"math"
"strings"
"sync"
"github.com/tetratelabs/wazero"
"github.com/tetratelabs/wazero/api"
"github.com/tetratelabs/wazero/imports/wasi_snapshot_preview1"
)
const (
statusOK = 0
statusErr = 1
statusNull = 2
)
const (
aggUnion = 0
aggMVT = 1
aggExtent = 2
)
type runtime struct {
wz wazero.Runtime
compiled wazero.CompiledModule
cfg wazero.ModuleConfig
mu sync.Mutex
free []*instance
n int }
func newRuntime(ctx context.Context, wasm []byte) (*runtime, error) {
wz := wazero.NewRuntime(ctx)
if _, err := wasi_snapshot_preview1.Instantiate(ctx, wz); err != nil {
wz.Close(ctx)
return nil, fmt.Errorf("kenro: instantiating WASI: %w", err)
}
compiled, err := wz.CompileModule(ctx, wasm)
if err != nil {
wz.Close(ctx)
return nil, fmt.Errorf("kenro: compiling the kenro wasm module: %w", err)
}
return &runtime{
wz: wz,
compiled: compiled,
cfg: wazero.NewModuleConfig().WithStartFunctions().WithName(""),
}, nil
}
func (r *runtime) close(ctx context.Context) error {
r.mu.Lock()
r.free = nil
r.mu.Unlock()
return r.wz.Close(ctx)
}
func (r *runtime) acquire(ctx context.Context) (*instance, error) {
r.mu.Lock()
if n := len(r.free); n > 0 {
in := r.free[n-1]
r.free = r.free[:n-1]
r.mu.Unlock()
return in, nil
}
r.n++
r.mu.Unlock()
mod, err := r.wz.InstantiateModule(ctx, r.compiled, r.cfg)
if err != nil {
return nil, fmt.Errorf("kenro: instantiating the wasm module: %w", err)
}
return newInstance(mod)
}
func (r *runtime) release(ctx context.Context, in *instance) {
if in == nil {
return
}
if in.poisoned {
_ = in.mod.Close(ctx)
return
}
r.mu.Lock()
r.free = append(r.free, in)
r.mu.Unlock()
}
type instance struct {
mod api.Module
mem api.Memory
alloc api.Function
free api.Function
outPtr api.Function
outLen api.Function
retI64 api.Function
retF64 api.Function
aggNew api.Function
aggFinish api.Function
aggDrop api.Function
unionStep api.Function
extentStep api.Function
mvtStep api.Function
fns map[string]api.Function
scratch uint32
scratchLen uint32
stack []uint64
poisoned bool
}
func newInstance(mod api.Module) (*instance, error) {
in := &instance{
mod: mod,
mem: mod.Memory(),
fns: map[string]api.Function{},
stack: make([]uint64, 0, 12),
}
required := map[string]*api.Function{
"kenro_alloc": &in.alloc,
"kenro_free": &in.free,
"kenro_out_ptr": &in.outPtr,
"kenro_out_len": &in.outLen,
"kenro_ret_i64": &in.retI64,
"kenro_ret_f64": &in.retF64,
"k_agg_new": &in.aggNew,
"k_agg_finish": &in.aggFinish,
"k_agg_drop": &in.aggDrop,
}
for name, dst := range required {
f := mod.ExportedFunction(name)
if f == nil {
return nil, fmt.Errorf("kenro: wasm module is missing the %q export", name)
}
*dst = f
}
in.unionStep = mod.ExportedFunction("k_agg_union_step")
in.extentStep = mod.ExportedFunction("k_agg_extent_step")
in.mvtStep = mod.ExportedFunction("k_agg_mvt_step")
return in, nil
}
func (in *instance) fn(name string) (api.Function, error) {
if f, ok := in.fns[name]; ok {
return f, nil
}
f := in.mod.ExportedFunction(name)
if f == nil {
return nil, fmt.Errorf("kenro: wasm module is missing the %q export", name)
}
in.fns[name] = f
return f, nil
}
func (in *instance) call(ctx context.Context, f api.Function, params ...uint64) (int32, error) {
in.stack = append(in.stack[:0], params...)
if len(in.stack) == 0 {
in.stack = append(in.stack, 0)
}
if err := f.CallWithStack(ctx, in.stack); err != nil {
in.poisoned = true
return 0, fmt.Errorf("kenro: wasm trap (this is a bug in kenro): %w", err)
}
return api.DecodeI32(in.stack[0]), nil
}
func (in *instance) reserve(ctx context.Context, n uint32) (uint32, error) {
if n <= in.scratchLen {
return in.scratch, nil
}
size := uint32(1024)
for size < n {
size *= 2
}
if in.scratch != 0 {
if _, err := in.call(ctx, in.free, api.EncodeI32(int32(in.scratch)), api.EncodeU32(in.scratchLen)); err != nil {
return 0, err
}
}
in.stack = append(in.stack[:0], api.EncodeU32(size))
if err := in.alloc.CallWithStack(ctx, in.stack); err != nil {
in.poisoned = true
return 0, fmt.Errorf("kenro: wasm trap in kenro_alloc: %w", err)
}
ptr := api.DecodeU32(in.stack[0])
if ptr == 0 {
return 0, fmt.Errorf("kenro: out of wasm memory reserving %d bytes", size)
}
in.scratch, in.scratchLen = ptr, size
return ptr, nil
}
func (in *instance) out(ctx context.Context) ([]byte, error) {
if _, err := in.call(ctx, in.outPtr); err != nil {
return nil, err
}
ptr := api.DecodeU32(in.stack[0])
if _, err := in.call(ctx, in.outLen); err != nil {
return nil, err
}
n := api.DecodeU32(in.stack[0])
if n == 0 {
return nil, nil
}
b, ok := in.mem.Read(ptr, n)
if !ok {
return nil, fmt.Errorf("kenro: OUT buffer (%d bytes at %#x) is outside wasm memory", n, ptr)
}
return append([]byte(nil), b...), nil
}
func (in *instance) errFromOut(ctx context.Context) error {
msg, err := in.out(ctx)
if err != nil {
return err
}
if len(msg) == 0 {
return fmt.Errorf("kenro: unknown error")
}
return errString(msg)
}
type errString []byte
func (e errString) Error() string { return string(e) }
func (in *instance) i64(ctx context.Context) (int64, error) {
if _, err := in.call(ctx, in.retI64); err != nil {
return 0, err
}
return int64(in.stack[0]), nil
}
func (in *instance) f64(ctx context.Context) (float64, error) {
if _, err := in.call(ctx, in.retF64); err != nil {
return 0, err
}
return math.Float64frombits(in.stack[0]), nil
}
func (in *instance) writeBytes(base, off uint32, b []byte) (uint32, error) {
if len(b) == 0 {
return off, nil
}
if !in.mem.Write(base+off, b) {
return 0, fmt.Errorf("kenro: writing %d argument bytes past the end of wasm memory", len(b))
}
return off + uint32(len(b)), nil
}
func describeExports(in *instance, names []string) string {
missing := make([]string, 0, len(names))
for _, n := range names {
if in.mod.ExportedFunction(n) == nil {
missing = append(missing, n)
}
}
return strings.Join(missing, ", ")
}