Files
DoGaMa-serv/internal/module/runtime.go
T
codex 35c11aff6d
CI / validate (pull_request) Successful in 26m22s
feat(vrising): support template runtime requirements
2026-08-25 22:32:15 +02:00

475 lines
17 KiB
Go

// Package module executes game adapters in a capability-limited WebAssembly sandbox.
package module
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"os"
"path"
"strings"
"sync"
"time"
"github.com/tetratelabs/wazero"
"github.com/tetratelabs/wazero/api"
"github.com/tetratelabs/wazero/imports/wasi_snapshot_preview1"
)
const (
ABI = "dogama:game-module@1.0.0"
maxRequestBytes = 64 << 10
maxHostCallsPerCall = 32
)
var capabilityExport = map[string]string{
"server_info": "get_server_info", "metrics": "get_metrics", "player_list": "list_players",
"online_save": "save_world", "graceful_shutdown": "shutdown", "announcement": "send_announcement",
"kick": "kick_player", "ban": "ban_player", "unban": "unban_player", "list_bans": "list_bans",
}
type Limits struct {
MemoryMB uint32
Timeout time.Duration
MaxResponseBytes int
MaxConcurrentCall int
}
type Binding struct {
InstanceID string
ContainerPort int
AllowedMethods map[string]bool
Configuration map[string]string
Secrets map[string]string
}
// TCPRequest is a single, bounded exchange with the template-declared TCP
// integration endpoint. The destination is never supplied by a module.
type TCPRequest struct {
Body []byte `json:"body"`
}
type HTTPRequest struct {
Method string `json:"method"`
Path string `json:"path"`
Header map[string]string `json:"header,omitempty"`
Body []byte `json:"body,omitempty"`
}
type HTTPResponse struct {
Status int `json:"status"`
Header map[string][]string `json:"header,omitempty"`
Body []byte `json:"body,omitempty"`
}
type Runtime struct {
wasm []byte
capabilities []string
limits Limits
binding Binding
client *http.Client
dialContext func(context.Context, string, string) (net.Conn, error)
sem chan struct{}
mu sync.Mutex
failures int
openUntil time.Time
}
func New(ctx context.Context, wasm []byte, checksum string, capabilities []string, limits Limits, binding Binding) (*Runtime, error) {
transport, dialContext, err := pinnedTransport(ctx, binding)
if err != nil {
return nil, err
}
return newWithTransportAndDial(wasm, checksum, capabilities, limits, binding, transport, dialContext)
}
func newWithTransport(wasm []byte, checksum string, capabilities []string, limits Limits, binding Binding, transport http.RoundTripper) (*Runtime, error) {
return newWithTransportAndDial(wasm, checksum, capabilities, limits, binding, transport, func(context.Context, string, string) (net.Conn, error) {
return nil, errors.New("TCP exchange is unavailable in this test runtime")
})
}
func newWithTransportAndDial(wasm []byte, checksum string, capabilities []string, limits Limits, binding Binding, transport http.RoundTripper, dialContext func(context.Context, string, string) (net.Conn, error)) (*Runtime, error) {
if len(wasm) == 0 || limits.MemoryMB == 0 || limits.MemoryMB > 256 || limits.Timeout <= 0 || limits.MaxResponseBytes < 1 || limits.MaxResponseBytes > 8<<20 || limits.MaxConcurrentCall < 1 || limits.MaxConcurrentCall > 16 {
return nil, errors.New("invalid module runtime configuration")
}
digest := sha256.Sum256(wasm)
if hex.EncodeToString(digest[:]) != checksum {
return nil, errors.New("module checksum mismatch")
}
if binding.InstanceID == "" || binding.ContainerPort < 1 || binding.ContainerPort > 65535 || transport == nil || dialContext == nil {
return nil, errors.New("invalid instance API binding")
}
seen := map[string]bool{}
for _, capability := range capabilities {
if _, ok := capabilityExport[capability]; !ok || seen[capability] {
return nil, errors.New("invalid module capability")
}
seen[capability] = true
}
r := &Runtime{wasm: append([]byte(nil), wasm...), capabilities: append([]string(nil), capabilities...), limits: limits, binding: binding, sem: make(chan struct{}, limits.MaxConcurrentCall), dialContext: dialContext}
r.client = &http.Client{Transport: transport, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}
return r, nil
}
func pinnedTransport(ctx context.Context, binding Binding) (http.RoundTripper, func(context.Context, string, string) (net.Conn, error), error) {
hostname := "dogama-" + strings.ToLower(binding.InstanceID)
lookupCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
defer cancel()
addresses, err := net.DefaultResolver.LookupIPAddr(lookupCtx, hostname)
if err != nil || len(addresses) == 0 {
return nil, nil, errors.New("resolve bound instance API")
}
ip := addresses[0].IP
if ip == nil || ip.IsUnspecified() || ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsMulticast() {
return nil, nil, errors.New("unsafe instance API address")
}
pinned := net.JoinHostPort(ip.String(), fmt.Sprint(binding.ContainerPort))
dialer := &net.Dialer{Timeout: 2 * time.Second, KeepAlive: 30 * time.Second}
dialContext := func(ctx context.Context, network, _ string) (net.Conn, error) {
if network != "tcp" && network != "tcp4" && network != "tcp6" {
return nil, errors.New("unsupported instance API network")
}
return dialer.DialContext(ctx, "tcp", pinned)
}
return &http.Transport{
DisableCompression: true,
Proxy: nil,
DialContext: dialContext,
}, dialContext, nil
}
func (r *Runtime) Call(ctx context.Context, operation string, request any, response any) error {
if !r.allowedOperation(operation) {
return errors.New("unsupported")
}
r.mu.Lock()
open := time.Now().Before(r.openUntil)
r.mu.Unlock()
if open {
return errors.New("module circuit breaker is open")
}
select {
case r.sem <- struct{}{}:
defer func() { <-r.sem }()
case <-ctx.Done():
return ctx.Err()
}
callCtx, cancel := context.WithTimeout(ctx, r.limits.Timeout)
defer cancel()
input, err := json.Marshal(request)
if err != nil || len(input) > maxRequestBytes {
return errors.New("invalid module request")
}
result, err := r.invoke(callCtx, operation, input)
if err != nil {
r.recordFailure()
return err
}
r.recordSuccess()
decoder := json.NewDecoder(bytes.NewReader(result))
decoder.DisallowUnknownFields()
if err := decoder.Decode(response); err != nil || decoder.Decode(&struct{}{}) != io.EOF {
return errors.New("invalid module response")
}
return nil
}
// CallData unwraps the normalized module envelope for typed service callers.
// Call remains available for diagnostics and callers that need the envelope.
func (r *Runtime) CallData(ctx context.Context, operation string, request any, response any) error {
var envelope struct {
OK bool `json:"ok"`
Data json.RawMessage `json:"data"`
Error *struct {
Message string `json:"message"`
} `json:"error"`
}
if err := r.Call(ctx, operation, request, &envelope); err != nil {
return err
}
if !envelope.OK {
if envelope.Error != nil && envelope.Error.Message != "" {
return errors.New(envelope.Error.Message)
}
return errors.New("module operation failed")
}
decoder := json.NewDecoder(bytes.NewReader(envelope.Data))
decoder.DisallowUnknownFields()
if err := decoder.Decode(response); err != nil || decoder.Decode(&struct{}{}) != io.EOF {
return errors.New("invalid module response")
}
return nil
}
func (r *Runtime) allowedOperation(operation string) bool {
if operation == "initialize" || operation == "test_connection" || operation == "get_server_status" {
return true
}
for _, capability := range r.capabilities {
if capabilityExport[capability] == operation {
return true
}
}
return false
}
type callState struct{ hostCalls int }
type stateKey struct{}
func (r *Runtime) invoke(ctx context.Context, operation string, input []byte) ([]byte, error) {
return r.invokeWithConfig(ctx, operation, input, wazero.NewRuntimeConfigInterpreter())
}
func (r *Runtime) invokeWithConfig(ctx context.Context, operation string, input []byte, baseConfig wazero.RuntimeConfig) ([]byte, error) {
ctx = context.WithValue(ctx, stateKey{}, &callState{})
pages := (r.limits.MemoryMB*1024*1024 + 65535) / 65536
config := baseConfig.WithMemoryLimitPages(pages).WithCloseOnContextDone(true).WithDebugInfoEnabled(false)
runtime := wazero.NewRuntimeWithConfig(ctx, config)
defer func() { _ = runtime.Close(ctx) }()
if _, err := wasi_snapshot_preview1.Instantiate(ctx, runtime); err != nil {
return nil, fmt.Errorf("instantiate restricted WASI: %w", err)
}
host := runtime.NewHostModuleBuilder("dogama_host")
host.NewFunctionBuilder().WithFunc(r.httpRequest).Export("http_request")
host.NewFunctionBuilder().WithFunc(r.tcpExchange).Export("tcp_exchange")
host.NewFunctionBuilder().WithFunc(r.getSecret).Export("get_secret")
host.NewFunctionBuilder().WithFunc(r.getConfig).Export("get_config")
if _, err := host.Instantiate(ctx); err != nil {
return nil, fmt.Errorf("instantiate module host: %w", err)
}
compiled, err := runtime.CompileModule(ctx, r.wasm)
if err != nil {
return nil, errors.New("compile module")
}
if err := validateABI(compiled, r.capabilities); err != nil {
return nil, err
}
instance, err := runtime.InstantiateModule(ctx, compiled, wazero.NewModuleConfig().WithName("").WithStartFunctions("_initialize").WithStdin(bytes.NewReader(nil)).WithStdout(io.Discard).WithStderr(io.Discard))
if err != nil {
return nil, errors.New("instantiate module")
}
alloc := instance.ExportedFunction("dogama_alloc")
fn := instance.ExportedFunction(operation)
if alloc == nil || fn == nil {
return nil, errors.New("missing module export")
}
allocated, err := alloc.Call(ctx, uint64(len(input)))
if err != nil || len(allocated) != 1 || allocated[0] > 1<<32-1 || !instance.Memory().Write(uint32(allocated[0]), input) {
return nil, errors.New("write module request")
}
out, err := alloc.Call(ctx, uint64(r.limits.MaxResponseBytes))
if err != nil || len(out) != 1 || out[0] > 1<<32-1 {
return nil, errors.New("allocate module response")
}
result, err := fn.Call(ctx, allocated[0], uint64(len(input)), out[0], uint64(r.limits.MaxResponseBytes))
if err != nil || len(result) != 1 || int32(result[0]) < 0 || result[0] > uint64(r.limits.MaxResponseBytes) {
return nil, errors.New("module execution failed")
}
body, ok := instance.Memory().Read(uint32(out[0]), uint32(result[0]))
if !ok {
return nil, errors.New("read module response")
}
return append([]byte(nil), body...), nil
}
// tcpExchange provides no socket handle to the module. It writes one bounded
// request to the instance-pinned integration endpoint then returns bytes read
// until a short idle period, EOF, or the call deadline. This permits framed
// protocols such as Source RCON while keeping every operation bounded.
func (r *Runtime) tcpExchange(ctx context.Context, mod api.Module, requestPtr, requestLen, responsePtr, responseCap uint32) int32 {
state, _ := ctx.Value(stateKey{}).(*callState)
if state == nil || state.hostCalls >= maxHostCallsPerCall || requestLen > maxRequestBytes || responseCap > uint32(r.limits.MaxResponseBytes) {
return -1
}
state.hostCalls++
raw, ok := mod.Memory().Read(requestPtr, requestLen)
if !ok {
return -1
}
var request TCPRequest
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.DisallowUnknownFields()
if decoder.Decode(&request) != nil || len(request.Body) == 0 || len(request.Body) > maxRequestBytes {
return -2
}
conn, err := r.dialContext(ctx, "tcp", "")
if err != nil {
if errors.Is(err, context.DeadlineExceeded) || isTimeout(err) {
return -4
}
return -3
}
defer func() { _ = conn.Close() }()
deadline := time.Now().Add(r.limits.Timeout)
if value, exists := ctx.Deadline(); exists && value.Before(deadline) {
deadline = value
}
if err := conn.SetDeadline(deadline); err != nil {
return -3
}
if _, err := conn.Write(request.Body); err != nil {
if isTimeout(err) {
return -4
}
return -3
}
result := make([]byte, 0, min(int(responseCap), 4096))
buffer := make([]byte, min(4096, int(responseCap)))
for {
if len(result) > 0 {
_ = conn.SetReadDeadline(time.Now().Add(50 * time.Millisecond))
}
n, readErr := conn.Read(buffer)
if n > 0 {
if len(result)+n > int(responseCap) {
return -5
}
result = append(result, buffer[:n]...)
}
if readErr != nil {
if errors.Is(readErr, io.EOF) || (len(result) > 0 && isTimeout(readErr)) {
break
}
if isTimeout(readErr) {
return -4
}
return -3
}
}
if len(result) == 0 || !mod.Memory().Write(responsePtr, result) {
return -5
}
return int32(len(result))
}
func isTimeout(err error) bool {
var netErr net.Error
return errors.Is(err, os.ErrDeadlineExceeded) || (errors.As(err, &netErr) && netErr.Timeout())
}
func min(a, b int) int {
if a < b {
return a
}
return b
}
func validateABI(compiled wazero.CompiledModule, capabilities []string) error {
exports := compiled.ExportedFunctions()
for _, name := range []string{"dogama_alloc", "initialize", "test_connection", "get_server_status"} {
if exports[name] == nil {
return fmt.Errorf("required module export %s is missing", name)
}
}
for _, capability := range capabilities {
if exports[capabilityExport[capability]] == nil {
return fmt.Errorf("capability export %s is missing", capability)
}
}
for _, imported := range compiled.ImportedFunctions() {
moduleName, name, importedFunction := imported.Import()
if importedFunction && moduleName != "dogama_host" && moduleName != wasi_snapshot_preview1.ModuleName {
return fmt.Errorf("forbidden import %s.%s", moduleName, name)
}
}
return nil
}
func (r *Runtime) httpRequest(ctx context.Context, mod api.Module, requestPtr, requestLen, responsePtr, responseCap uint32) int32 {
state, _ := ctx.Value(stateKey{}).(*callState)
if state == nil || state.hostCalls >= maxHostCallsPerCall || requestLen > maxRequestBytes || responseCap > uint32(r.limits.MaxResponseBytes) {
return -1
}
state.hostCalls++
raw, ok := mod.Memory().Read(requestPtr, requestLen)
if !ok {
return -1
}
var request HTTPRequest
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.DisallowUnknownFields()
if decoder.Decode(&request) != nil || !r.binding.AllowedMethods[request.Method] || len(request.Body) > maxRequestBytes || !validRelativeAPIPath(request.Path) {
return -2
}
endpoint := fmt.Sprintf("http://dogama-%s:%d%s", strings.ToLower(r.binding.InstanceID), r.binding.ContainerPort, request.Path)
httpRequest, err := http.NewRequestWithContext(ctx, request.Method, endpoint, bytes.NewReader(request.Body))
if err != nil {
return -2
}
for name, value := range request.Header {
canonical := http.CanonicalHeaderKey(name)
if canonical != "Accept" && canonical != "Content-Type" && canonical != "Authorization" {
return -2
}
httpRequest.Header.Set(canonical, value)
}
response, err := r.client.Do(httpRequest)
if err != nil {
return -3
}
defer func() { _ = response.Body.Close() }()
body, err := io.ReadAll(io.LimitReader(response.Body, int64(r.limits.MaxResponseBytes)+1))
if err != nil || len(body) > r.limits.MaxResponseBytes {
return -4
}
encoded, err := json.Marshal(HTTPResponse{Status: response.StatusCode, Header: map[string][]string{"Content-Type": response.Header.Values("Content-Type")}, Body: body})
if err != nil || len(encoded) > int(responseCap) || !mod.Memory().Write(responsePtr, encoded) {
return -4
}
return int32(len(encoded))
}
func (r *Runtime) getSecret(ctx context.Context, mod api.Module, keyPtr, keyLen, valuePtr, valueCap uint32) int32 {
return r.readBoundValue(ctx, mod, keyPtr, keyLen, valuePtr, valueCap, r.binding.Secrets)
}
func (r *Runtime) getConfig(ctx context.Context, mod api.Module, keyPtr, keyLen, valuePtr, valueCap uint32) int32 {
return r.readBoundValue(ctx, mod, keyPtr, keyLen, valuePtr, valueCap, r.binding.Configuration)
}
func (r *Runtime) readBoundValue(ctx context.Context, mod api.Module, keyPtr, keyLen, valuePtr, valueCap uint32, values map[string]string) int32 {
state, _ := ctx.Value(stateKey{}).(*callState)
if state == nil || state.hostCalls >= maxHostCallsPerCall || keyLen > 128 {
return -1
}
state.hostCalls++
raw, ok := mod.Memory().Read(keyPtr, keyLen)
if !ok {
return -1
}
value, exists := values[string(raw)]
if !exists || len(value) > int(valueCap) || !mod.Memory().Write(valuePtr, []byte(value)) {
return -2
}
return int32(len(value))
}
func validRelativeAPIPath(value string) bool {
parsed, err := url.Parse(value)
return err == nil && !parsed.IsAbs() && parsed.Host == "" && parsed.RawQuery == "" && parsed.Fragment == "" && strings.HasPrefix(value, "/v1/api/") && path.Clean(value) == value && !strings.Contains(value, "\\")
}
func (r *Runtime) recordFailure() {
r.mu.Lock()
defer r.mu.Unlock()
r.failures++
if r.failures >= 3 {
r.openUntil = time.Now().Add(30 * time.Second)
}
}
func (r *Runtime) recordSuccess() {
r.mu.Lock()
r.failures, r.openUntil = 0, time.Time{}
r.mu.Unlock()
}