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
aggExtent3D = 3
)
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
extent3dStep 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.extent3dStep = mod.ExportedFunction("k_agg_extent3d_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, ", ")
}