[model-gateway] Add WASM support for middleware (#12471)

Signed-off-by: Tony Lu <tonylu@linux.alibaba.com>
This commit is contained in:
Tony Lu
2025-12-06 01:29:07 +08:00
committed by GitHub
parent 38daa29466
commit 5a46fb153d
41 changed files with 4740 additions and 8 deletions

2
.github/CODEOWNERS vendored
View File

@@ -41,5 +41,7 @@
/sgl-router/src/routers @CatherineSue @key4ng @slin1237
/sgl-router/src/tokenizer @slin1237 @CatherineSue
/sgl-router/src/tool_parser @slin1237 @CatherineSue
/sgl-router/src/wasm @tonyluj
/sgl-router/examples/wasm @tonyluj
/test/srt/ascend @ping1jing2 @iforgetmyname
/test/srt/test_modelopt* @Edwardf0t1

View File

@@ -110,6 +110,12 @@ tokio-postgres = { version = "0.7.15", features = ["runtime","with-chrono-0_4","
deadpool-postgres = "0.14.1"
# wasm dependencies
sha2 = "0.10"
wasmtime = { version = "38.0", features = ["component-model", "async"] }
wasmtime-wasi = "38.0"
async-channel = "2.5"
[build-dependencies]
tonic-prost-build = "0.14.2"
prost-build = "0.14.1"
@@ -123,6 +129,7 @@ http-body-util = "0.1"
portpicker = "0.1"
tempfile = "3.8"
lazy_static = "1.4"
wasm-encoder = "0.242"
npyz = { version = "0.8", features = ["npz"] } # For reading numpy .npz files in golden tests
[[bench]]

27
sgl-router/examples/wasm/.gitignore vendored Normal file
View File

@@ -0,0 +1,27 @@
# Rust build artifacts
target/
**/target/
# Cargo lock files (examples don't need locked dependencies)
Cargo.lock
**/Cargo.lock
# Generated WASM files
*.wasm
*.component.wasm
**/*.wasm
**/*.component.wasm
# Build scripts output
build/
# IDE files
.idea/
.vscode/
*.swp
*.swo
*~
# OS files
.DS_Store
Thumbs.db

View File

@@ -0,0 +1,102 @@
# WASM Guest Examples for sgl-router
This directory contains example WASM middleware components demonstrating how to implement custom middleware for sgl-router using the WebAssembly Component Model.
## Examples Overview
### [wasm-guest-auth](./wasm-guest-auth/)
API key authentication middleware that validates API keys for requests to `/api` and `/v1` paths.
**Features:**
- Validates API keys from `Authorization` header or `x-api-key` header
- Returns `401 Unauthorized` for missing or invalid keys
- Attach point: `OnRequest` only
**Use case:** Protect API endpoints with API key authentication.
### [wasm-guest-logging](./wasm-guest-logging/)
Request tracking and status code conversion middleware.
**Features:**
- Adds tracking headers (`x-request-id`, `x-wasm-processed`, `x-processed-at`, `x-api-route`)
- Converts `500` errors to `503` for better client handling
- Attach points: `OnRequest` and `OnResponse`
**Use case:** Request tracing and error status code conversion.
### [wasm-guest-ratelimit](./wasm-guest-ratelimit/)
Rate limiting middleware with configurable limits.
**Features:**
- Rate limiting per identifier (API Key, IP, or Request ID)
- Default: 60 requests per minute
- Returns `429 Too Many Requests` when limit exceeded
- Attach point: `OnRequest` only
**Note:** This is a simplified demonstration with per-instance state. For production, use router-level rate limiting with shared state.
**Use case:** Protect against request flooding and abuse.
## Quick Start
Each example includes its own README with detailed build and deployment instructions. See individual example directories for:
- Build instructions
- Deployment configuration
- Customization options
- Testing examples
## Common Prerequisites
All examples require:
- Rust toolchain (latest stable)
- `wasm32-wasip2` target: `rustup target add wasm32-wasip2`
- `wasm-tools`: `cargo install wasm-tools`
- sgl-router running with WASM enabled (`--enable-wasm`)
## Building All Examples
```bash
cd examples/wasm
for example in wasm-guest-auth wasm-guest-logging wasm-guest-ratelimit; do
echo "Building $example..."
cd $example && ./build.sh && cd ..
done
```
## Deploying Multiple Modules
You can deploy all three modules together:
```bash
curl -X POST http://localhost:3000/wasm \
-H "Content-Type: application/json" \
-d '{
"modules": [
{
"name": "auth-middleware",
"file_path": "/path/to/wasm_guest_auth.component.wasm",
"module_type": "Middleware",
"attach_points": [{"Middleware": "OnRequest"}]
},
{
"name": "logging-middleware",
"file_path": "/path/to/wasm_guest_logging.component.wasm",
"module_type": "Middleware",
"attach_points": [{"Middleware": "OnRequest"}, {"Middleware": "OnResponse"}]
},
{
"name": "ratelimit-middleware",
"file_path": "/path/to/wasm_guest_ratelimit.component.wasm",
"module_type": "Middleware",
"attach_points": [{"Middleware": "OnRequest"}]
}
]
}'
```
Modules execute in the order they are deployed. If a module returns `Reject`, subsequent modules won't execute.

View File

@@ -0,0 +1,10 @@
[package]
name = "wasm-guest-auth"
version = "0.1.0"
edition = "2021"
[lib]
crate-type = ["cdylib"]
[dependencies]
wit-bindgen = { version = "0.21", features = ["macros"] }

View File

@@ -0,0 +1,62 @@
# WASM Auth Example for sgl-router
This example demonstrates API key authentication middleware for sgl-router using the WebAssembly Component Model.
## Overview
This middleware validates API keys for requests to `/api` and `/v1` paths:
- Supports `Authorization: Bearer <key>` header
- Supports `Authorization: ApiKey <key>` header
- Supports `x-api-key` header
- Returns `401 Unauthorized` for missing or invalid keys
**Default API Key**: `secret-api-key-12345`
## Quick Start
### Build and Deploy
```bash
# Build
cd examples/wasm-guest-auth
./build.sh
# Deploy (replace file_path with actual path)
curl -X POST http://localhost:3000/wasm \
-H "Content-Type: application/json" \
-d '{
"modules": [{
"name": "auth-middleware",
"file_path": "/absolute/path/to/wasm_guest_auth.component.wasm",
"module_type": "Middleware",
"attach_points": [{"Middleware": "OnRequest"}]
}]
}'
```
### Customization
Modify `EXPECTED_API_KEY` in `src/lib.rs`:
```rust
const EXPECTED_API_KEY: &str = "your-secret-key";
```
## Testing
```bash
# Test unauthorized (returns 401)
curl -v http://localhost:3000/api/test
# Test authorized (passes)
curl -v http://localhost:3000/api/test \
-H "Authorization: Bearer secret-api-key-12345"
```
## Troubleshooting
- Verify API key matches `EXPECTED_API_KEY` in code
- Check request header format and path (`/api` or `/v1`)
- Verify module is attached to `OnRequest` phase
- Check router logs for errors

View File

@@ -0,0 +1,77 @@
#!/bin/bash
# Build script for WASM guest auth example
# This script simplifies the build process for the WASM middleware component
set -e
echo "Building WASM guest auth example..."
# Check if we're in the right directory
if [ ! -f "Cargo.toml" ]; then
echo "Error: Cargo.toml not found. Please run this script from the wasm-guest-auth directory."
exit 1
fi
# Check for required tools
command -v cargo >/dev/null 2>&1 || { echo "Error: cargo is required but not installed. Aborting." >&2; exit 1; }
# Check and install wasm32-wasip2 target
echo "Checking for wasm32-wasip2 target..."
if ! rustup target list --installed | grep -q "wasm32-wasip2"; then
echo "wasm32-wasip2 target not found. Installing..."
rustup target add wasm32-wasip2
echo "✓ wasm32-wasip2 target installed"
else
echo "✓ wasm32-wasip2 target already installed"
fi
# Check for wasm-tools
if ! command -v wasm-tools >/dev/null 2>&1; then
echo "Error: wasm-tools is required but not installed."
echo "Install it with: cargo install wasm-tools"
exit 1
fi
# Build with cargo (wit-bindgen uses cargo, not wasm-pack)
echo "Running cargo build..."
cargo build --target wasm32-wasip2 --release
# Output locations
WASM_MODULE="target/wasm32-wasip2/release/wasm_guest_auth.wasm"
WASM_COMPONENT="target/wasm32-wasip2/release/wasm_guest_auth.component.wasm"
if [ ! -f "$WASM_MODULE" ]; then
echo "Error: Build failed - WASM module not found"
exit 1
fi
# Check if the file is already a component
echo "Checking WASM file format..."
if wasm-tools print "$WASM_MODULE" 2>/dev/null | grep -q "^(\s*component"; then
echo "✓ WASM file is already in component format"
# Copy to component path for consistency
cp "$WASM_MODULE" "$WASM_COMPONENT"
else
# Wrap the WASM module into a component format
echo "Wrapping WASM module into component format..."
wasm-tools component new "$WASM_MODULE" -o "$WASM_COMPONENT"
if [ ! -f "$WASM_COMPONENT" ]; then
echo "Error: Failed to create component file"
exit 1
fi
fi
if [ -f "$WASM_COMPONENT" ]; then
echo ""
echo "✓ Build successful!"
echo " WASM module: $WASM_MODULE"
echo " WASM component: $WASM_COMPONENT"
echo ""
echo "Next steps:"
echo "1. Use the component file ($WASM_COMPONENT) when adding the module"
echo "2. Prepare the module configuration (see README.md for JSON format)"
echo "3. Use the API endpoint to add the module (see README.md for details)"
else
echo "Error: Component file not found"
exit 1
fi

View File

@@ -0,0 +1,70 @@
//! WASM Guest Auth Example for sgl-router
//!
//! This example demonstrates API key authentication middleware
//! for sgl-router using the WebAssembly Component Model.
//!
//! Features:
//! - API Key authentication
wit_bindgen::generate!({
path: "../../../src/wasm/interface",
world: "sgl-router",
});
use exports::sgl::router::{
middleware_on_request::Guest as OnRequestGuest,
middleware_on_response::Guest as OnResponseGuest,
};
use sgl::router::middleware_types::{Action, Request, Response};
/// Expected API Key (in production, this should be passed as configuration)
const EXPECTED_API_KEY: &str = "secret-api-key-12345";
/// Main middleware implementation
struct Middleware;
// Helper function to find header value
fn find_header_value(
headers: &[sgl::router::middleware_types::Header],
name: &str,
) -> Option<String> {
headers
.iter()
.find(|h| h.name.eq_ignore_ascii_case(name))
.map(|h| h.value.clone())
}
// Implement on-request interface
impl OnRequestGuest for Middleware {
fn on_request(req: Request) -> Action {
// API Key Authentication
// Check for API key in Authorization header for /api routes
if req.path.starts_with("/api") || req.path.starts_with("/v1") {
let api_key = find_header_value(&req.headers, "authorization")
.and_then(|h| {
h.strip_prefix("Bearer ")
.or_else(|| h.strip_prefix("ApiKey "))
.map(|s| s.to_string())
})
.or_else(|| find_header_value(&req.headers, "x-api-key"));
// Reject if API key is missing or invalid
if api_key.as_deref() != Some(EXPECTED_API_KEY) {
return Action::Reject(401);
}
}
// Authentication passed, continue processing
Action::Continue
}
}
// Implement on-response interface (empty - not used for auth)
impl OnResponseGuest for Middleware {
fn on_response(_resp: Response) -> Action {
Action::Continue
}
}
// Export the component
export!(Middleware);

View File

@@ -0,0 +1,10 @@
[package]
name = "wasm-guest-logging"
version = "0.1.0"
edition = "2021"
[lib]
crate-type = ["cdylib"]
[dependencies]
wit-bindgen = { version = "0.21", features = ["macros"] }

View File

@@ -0,0 +1,53 @@
# WASM Logging Example for sgl-router
This example demonstrates logging and tracing middleware for sgl-router using the WebAssembly Component Model.
## Overview
This middleware provides:
- **Request Tracking** - Adds tracking headers (`x-request-id`, `x-wasm-processed`, `x-processed-at`, `x-api-route`)
- **Status Code Conversion** - Converts `500` errors to `503`
## Quick Start
### Build and Deploy
```bash
# Build
cd examples/wasm-guest-logging
./build.sh
# Deploy (replace file_path with actual path)
curl -X POST http://localhost:3000/wasm \
-H "Content-Type: application/json" \
-d '{
"modules": [{
"name": "logging-middleware",
"file_path": "/absolute/path/to/wasm_guest_logging.component.wasm",
"module_type": "Middleware",
"attach_points": [{"Middleware": "OnRequest"}, {"Middleware": "OnResponse"}]
}]
}'
```
### Customization
Modify `on_request` or `on_response` functions in `src/lib.rs` to add custom tracking headers or status code conversions.
## Testing
```bash
# Check tracking headers
curl -v http://localhost:3000/v1/models 2>&1 | \
grep -E "(x-request-id|x-wasm-processed|x-processed-at)"
# Test status code conversion (requires endpoint returning 500)
curl -v http://localhost:3000/some-endpoint 2>&1 | grep -E "(< HTTP|500|503)"
```
## Troubleshooting
- Verify module attached to both `OnRequest` and `OnResponse` phases
- Check router logs for execution errors
- Ensure module built successfully

View File

@@ -0,0 +1,77 @@
#!/bin/bash
# Build script for WASM guest logging example
# This script simplifies the build process for the WASM middleware component
set -e
echo "Building WASM guest logging example..."
# Check if we're in the right directory
if [ ! -f "Cargo.toml" ]; then
echo "Error: Cargo.toml not found. Please run this script from the wasm-guest-logging directory."
exit 1
fi
# Check for required tools
command -v cargo >/dev/null 2>&1 || { echo "Error: cargo is required but not installed. Aborting." >&2; exit 1; }
# Check and install wasm32-wasip2 target
echo "Checking for wasm32-wasip2 target..."
if ! rustup target list --installed | grep -q "wasm32-wasip2"; then
echo "wasm32-wasip2 target not found. Installing..."
rustup target add wasm32-wasip2
echo "✓ wasm32-wasip2 target installed"
else
echo "✓ wasm32-wasip2 target already installed"
fi
# Check for wasm-tools
if ! command -v wasm-tools >/dev/null 2>&1; then
echo "Error: wasm-tools is required but not installed."
echo "Install it with: cargo install wasm-tools"
exit 1
fi
# Build with cargo (wit-bindgen uses cargo, not wasm-pack)
echo "Running cargo build..."
cargo build --target wasm32-wasip2 --release
# Output locations
WASM_MODULE="target/wasm32-wasip2/release/wasm_guest_logging.wasm"
WASM_COMPONENT="target/wasm32-wasip2/release/wasm_guest_logging.component.wasm"
if [ ! -f "$WASM_MODULE" ]; then
echo "Error: Build failed - WASM module not found"
exit 1
fi
# Check if the file is already a component
echo "Checking WASM file format..."
if wasm-tools print "$WASM_MODULE" 2>/dev/null | grep -q "^(\s*component"; then
echo "✓ WASM file is already in component format"
# Copy to component path for consistency
cp "$WASM_MODULE" "$WASM_COMPONENT"
else
# Wrap the WASM module into a component format
echo "Wrapping WASM module into component format..."
wasm-tools component new "$WASM_MODULE" -o "$WASM_COMPONENT"
if [ ! -f "$WASM_COMPONENT" ]; then
echo "Error: Failed to create component file"
exit 1
fi
fi
if [ -f "$WASM_COMPONENT" ]; then
echo ""
echo "✓ Build successful!"
echo " WASM module: $WASM_MODULE"
echo " WASM component: $WASM_COMPONENT"
echo ""
echo "Next steps:"
echo "1. Use the component file ($WASM_COMPONENT) when adding the module"
echo "2. Prepare the module configuration (see README.md for JSON format)"
echo "3. Use the API endpoint to add the module (see README.md for details)"
else
echo "Error: Component file not found"
exit 1
fi

View File

@@ -0,0 +1,88 @@
//! WASM Guest Logging Example for sgl-router
//!
//! This example demonstrates logging and tracing middleware
//! for sgl-router using the WebAssembly Component Model.
//!
//! Features:
//! - Request tracking and tracing headers
//! - Response status code conversion
wit_bindgen::generate!({
path: "../../../src/wasm/interface",
world: "sgl-router",
});
use exports::sgl::router::{
middleware_on_request::Guest as OnRequestGuest,
middleware_on_response::Guest as OnResponseGuest,
};
use sgl::router::middleware_types::{Action, Header, ModifyAction, Request, Response};
/// Main middleware implementation
struct Middleware;
// Helper function to create header
fn create_header(name: &str, value: &str) -> Header {
Header {
name: name.to_string(),
value: value.to_string(),
}
}
// Implement on-request interface
impl OnRequestGuest for Middleware {
fn on_request(req: Request) -> Action {
let mut modify_action = ModifyAction {
status: None,
headers_set: vec![],
headers_add: vec![],
headers_remove: vec![],
body_replace: None,
};
// Request Logging and Tracing
// Add tracing headers with request ID
modify_action
.headers_add
.push(create_header("x-request-id", &req.request_id));
modify_action
.headers_add
.push(create_header("x-wasm-processed", "true"));
modify_action.headers_add.push(create_header(
"x-processed-at",
&req.now_epoch_ms.to_string(),
));
// Add custom header for API requests
if req.path.starts_with("/api") || req.path.starts_with("/v1") {
modify_action
.headers_add
.push(create_header("x-api-route", "true"));
}
Action::Modify(modify_action)
}
}
// Implement on-response interface
impl OnResponseGuest for Middleware {
fn on_response(resp: Response) -> Action {
// Status code conversion: Convert 500 to 503 for better client handling
if resp.status == 500 {
let modify_action = ModifyAction {
status: Some(503),
headers_set: vec![],
headers_add: vec![],
headers_remove: vec![],
body_replace: None,
};
Action::Modify(modify_action)
} else {
// No modification needed
Action::Continue
}
}
}
// Export the component
export!(Middleware);

View File

@@ -0,0 +1,10 @@
[package]
name = "wasm-guest-ratelimit"
version = "0.1.0"
edition = "2021"
[lib]
crate-type = ["cdylib"]
[dependencies]
wit-bindgen = { version = "0.21", features = ["macros"] }

View File

@@ -0,0 +1,68 @@
# WASM Rate Limit Example for sgl-router
This example demonstrates rate limiting middleware for sgl-router using the WebAssembly Component Model.
## Overview
This middleware provides rate limiting:
- **Default**: 60 requests per minute per identifier
- **Identifier Priority**: API Key > IP Address > Request ID
- **Response**: Returns `429 Too Many Requests` when limit exceeded
**Important**: This is a simplified demonstration. Since WASM components are stateless, each worker thread maintains its own counter. For production, implement rate limiting at the router/host level with shared state.
## Quick Start
### Build and Deploy
```bash
# Build
cd examples/wasm-guest-ratelimit
./build.sh
# Deploy (replace file_path with actual path)
curl -X POST http://localhost:3000/wasm \
-H "Content-Type: application/json" \
-d '{
"modules": [{
"name": "ratelimit-middleware",
"file_path": "/absolute/path/to/wasm_guest_ratelimit.component.wasm",
"module_type": "Middleware",
"attach_points": [{"Middleware": "OnRequest"}]
}]
}'
```
### Customization
Modify constants in `src/lib.rs`:
```rust
const RATE_LIMIT_REQUESTS: u64 = 100; // requests per window
const RATE_LIMIT_WINDOW_MS: u64 = 60_000; // time window in ms
```
## Testing
```bash
# Send multiple requests (first 60 succeed, then 429)
for i in {1..65}; do
curl -s -o /dev/null -w "%{http_code}\n" \
http://localhost:3000/v1/models \
-H "Authorization: Bearer secret-api-key-12345"
done
```
## Limitations
- Per-instance state (not shared across workers)
- No cross-process state sharing
- Memory growth with unique identifiers
- State lost on instance restart
## Troubleshooting
- Verify module attached to `OnRequest` phase
- Check identifier extraction logic matches request format
- Note: Each WASM worker has separate counter

View File

@@ -0,0 +1,77 @@
#!/bin/bash
# Build script for WASM guest rate limit example
# This script simplifies the build process for the WASM middleware component
set -e
echo "Building WASM guest rate limit example..."
# Check if we're in the right directory
if [ ! -f "Cargo.toml" ]; then
echo "Error: Cargo.toml not found. Please run this script from the wasm-guest-ratelimit directory."
exit 1
fi
# Check for required tools
command -v cargo >/dev/null 2>&1 || { echo "Error: cargo is required but not installed. Aborting." >&2; exit 1; }
# Check and install wasm32-wasip2 target
echo "Checking for wasm32-wasip2 target..."
if ! rustup target list --installed | grep -q "wasm32-wasip2"; then
echo "wasm32-wasip2 target not found. Installing..."
rustup target add wasm32-wasip2
echo "✓ wasm32-wasip2 target installed"
else
echo "✓ wasm32-wasip2 target already installed"
fi
# Check for wasm-tools
if ! command -v wasm-tools >/dev/null 2>&1; then
echo "Error: wasm-tools is required but not installed."
echo "Install it with: cargo install wasm-tools"
exit 1
fi
# Build with cargo (wit-bindgen uses cargo, not wasm-pack)
echo "Running cargo build..."
cargo build --target wasm32-wasip2 --release
# Output locations
WASM_MODULE="target/wasm32-wasip2/release/wasm_guest_ratelimit.wasm"
WASM_COMPONENT="target/wasm32-wasip2/release/wasm_guest_ratelimit.component.wasm"
if [ ! -f "$WASM_MODULE" ]; then
echo "Error: Build failed - WASM module not found"
exit 1
fi
# Check if the file is already a component
echo "Checking WASM file format..."
if wasm-tools print "$WASM_MODULE" 2>/dev/null | grep -q "^(\s*component"; then
echo "✓ WASM file is already in component format"
# Copy to component path for consistency
cp "$WASM_MODULE" "$WASM_COMPONENT"
else
# Wrap the WASM module into a component format
echo "Wrapping WASM module into component format..."
wasm-tools component new "$WASM_MODULE" -o "$WASM_COMPONENT"
if [ ! -f "$WASM_COMPONENT" ]; then
echo "Error: Failed to create component file"
exit 1
fi
fi
if [ -f "$WASM_COMPONENT" ]; then
echo ""
echo "✓ Build successful!"
echo " WASM module: $WASM_MODULE"
echo " WASM component: $WASM_COMPONENT"
echo ""
echo "Next steps:"
echo "1. Use the component file ($WASM_COMPONENT) when adding the module"
echo "2. Prepare the module configuration (see README.md for JSON format)"
echo "3. Use the API endpoint to add the module (see README.md for details)"
else
echo "Error: Component file not found"
exit 1
fi

View File

@@ -0,0 +1,155 @@
//! WASM Guest Rate Limit Example for sgl-router
//!
//! This example demonstrates rate limiting middleware
//! for sgl-router using the WebAssembly Component Model.
//!
//! Features:
//! - Rate limiting based on API Key or IP address
//! - Fixed time window (e.g., 60 requests per minute)
//! - Returns 429 Too Many Requests when limit exceeded
//!
//! Note: This is a simplified implementation. Since WASM components are stateless,
//! each instance maintains its own counters. For production use, consider
//! implementing rate limiting at the host/router level with shared state.
wit_bindgen::generate!({
path: "../../../src/wasm/interface",
world: "sgl-router",
});
use std::cell::RefCell;
use exports::sgl::router::{
middleware_on_request::Guest as OnRequestGuest,
middleware_on_response::Guest as OnResponseGuest,
};
use sgl::router::middleware_types::{Action, Request, Response};
/// Main middleware implementation
struct Middleware;
// Rate limit configuration
const RATE_LIMIT_REQUESTS: u64 = 60; // Maximum requests per window
const RATE_LIMIT_WINDOW_MS: u64 = 60_000; // Time window in milliseconds (1 minute)
// Simple in-memory counter (per WASM instance)
// In a real implementation, this would be shared across all instances
// This is a simplified example for demonstration purposes
struct RateLimitState {
requests: Vec<(String, u64)>, // (identifier, timestamp_ms)
}
impl RateLimitState {
fn new() -> Self {
Self {
requests: Vec::new(),
}
}
// Clean up old entries outside the time window
fn cleanup(&mut self, current_time_ms: u64) {
let cutoff = current_time_ms.saturating_sub(RATE_LIMIT_WINDOW_MS);
self.requests.retain(|(_, timestamp)| *timestamp > cutoff);
}
// Check if identifier has exceeded rate limit
fn check_limit(&mut self, identifier: &str, current_time_ms: u64) -> bool {
self.cleanup(current_time_ms);
// Count requests in current window for this identifier
let count = self
.requests
.iter()
.filter(|(id, timestamp)| {
id == identifier
&& *timestamp > current_time_ms.saturating_sub(RATE_LIMIT_WINDOW_MS)
})
.count() as u64;
if count >= RATE_LIMIT_REQUESTS {
return false; // Limit exceeded
}
// Add new request
self.requests
.push((identifier.to_string(), current_time_ms));
true // Within limit
}
}
// Thread-local state (per WASM instance thread)
// Using thread_local! is safer than static mut as it avoids unsafe blocks
// and provides separate state for each thread automatically
thread_local! {
static RATE_LIMIT_STATE: RefCell<RateLimitState> = RefCell::new(RateLimitState::new());
}
fn get_identifier(req: &Request) -> String {
// Helper function to find header value
let find_header_value =
|headers: &[sgl::router::middleware_types::Header], name: &str| -> Option<String> {
headers
.iter()
.find(|h| h.name.eq_ignore_ascii_case(name))
.map(|h| h.value.clone())
};
// Prefer API Key as identifier (more stable than IP)
if let Some(auth_header) = find_header_value(&req.headers, "authorization") {
if auth_header.starts_with("Bearer ") {
return format!("api_key:{}", &auth_header[7..]);
} else if auth_header.starts_with("ApiKey ") {
return format!("api_key:{}", &auth_header[7..]);
}
}
if let Some(api_key) = find_header_value(&req.headers, "x-api-key") {
return format!("api_key:{}", api_key);
}
// Fall back to IP address from forwarded headers
if let Some(forwarded_for) = find_header_value(&req.headers, "x-forwarded-for") {
// Take first IP from comma-separated list
let ip = forwarded_for.split(',').next().unwrap_or("").trim();
if !ip.is_empty() {
return format!("ip:{}", ip);
}
}
if let Some(real_ip) = find_header_value(&req.headers, "x-real-ip") {
return format!("ip:{}", real_ip);
}
// Last resort: use request ID (not ideal, but better than nothing)
format!("req_id:{}", req.request_id)
}
// Implement on-request interface
impl OnRequestGuest for Middleware {
fn on_request(req: Request) -> Action {
let identifier = get_identifier(&req);
let current_time_ms = req.now_epoch_ms;
// Access thread-local state safely without unsafe blocks
// Each thread gets its own RateLimitState instance
RATE_LIMIT_STATE.with(|state| {
let mut state = state.borrow_mut();
if !state.check_limit(&identifier, current_time_ms) {
// Rate limit exceeded
return Action::Reject(429);
}
// Within rate limit, continue processing
Action::Continue
})
}
}
// Implement on-response interface (empty - not used for rate limiting)
impl OnResponseGuest for Middleware {
fn on_response(_resp: Response) -> Action {
Action::Continue
}
}
// Export the component
export!(Middleware);

View File

@@ -23,6 +23,7 @@ use crate::{
traits::Tokenizer,
},
tool_parser::ParserFactory as ToolParserFactory,
wasm::{config::WasmRuntimeConfig, module_manager::WasmModuleManager},
};
/// Error type for AppContext builder
@@ -57,6 +58,7 @@ pub struct AppContext {
pub worker_job_queue: Arc<OnceLock<Arc<JobQueue>>>,
pub workflow_engine: Arc<OnceLock<Arc<WorkflowEngine>>>,
pub mcp_manager: Arc<OnceLock<Arc<McpManager>>>,
pub wasm_manager: Option<Arc<WasmModuleManager>>,
}
pub struct AppContextBuilder {
@@ -76,6 +78,7 @@ pub struct AppContextBuilder {
worker_job_queue: Option<Arc<OnceLock<Arc<JobQueue>>>>,
workflow_engine: Option<Arc<OnceLock<Arc<WorkflowEngine>>>>,
mcp_manager: Option<Arc<OnceLock<Arc<McpManager>>>>,
wasm_manager: Option<Arc<WasmModuleManager>>,
}
impl AppContext {
@@ -115,6 +118,7 @@ impl AppContextBuilder {
worker_job_queue: None,
workflow_engine: None,
mcp_manager: None,
wasm_manager: None,
}
}
@@ -207,6 +211,11 @@ impl AppContextBuilder {
self
}
pub fn wasm_manager(mut self, wasm_manager: Option<Arc<WasmModuleManager>>) -> Self {
self.wasm_manager = wasm_manager;
self
}
pub fn build(self) -> Result<AppContext, AppContextBuildError> {
let router_config = self
.router_config
@@ -249,6 +258,7 @@ impl AppContextBuilder {
mcp_manager: self
.mcp_manager
.ok_or(AppContextBuildError("mcp_manager"))?,
wasm_manager: self.wasm_manager,
})
}
@@ -272,6 +282,7 @@ impl AppContextBuilder {
.with_workflow_engine()
.with_mcp_manager(&router_config)
.await?
.with_wasm_manager(&router_config)?
.router_config(router_config))
}
@@ -505,6 +516,19 @@ impl AppContextBuilder {
self.mcp_manager = Some(mcp_manager_lock);
Ok(self)
}
/// Create wasm manager if enabled in config
fn with_wasm_manager(mut self, config: &RouterConfig) -> Result<Self, String> {
self.wasm_manager = if config.enable_wasm {
Some(Arc::new(
WasmModuleManager::new(WasmRuntimeConfig::default())
.map_err(|e| format!("Failed to initialize WASM module manager: {}", e))?,
))
} else {
None
};
Ok(self)
}
}
impl Default for AppContextBuilder {

View File

@@ -327,6 +327,13 @@ impl RouterConfigBuilder {
self
}
// ==================== WASM ====================
pub fn enable_wasm(mut self, enable: bool) -> Self {
self.config.enable_wasm = enable;
self
}
pub fn model_path<S: Into<String>>(mut self, path: S) -> Self {
self.config.model_path = Some(path.into());
self

View File

@@ -72,6 +72,9 @@ pub struct RouterConfig {
/// Loaded from mcp_config_path during config creation
#[serde(skip)]
pub mcp_config: Option<crate::mcp::McpConfig>,
/// Enable WASM support
#[serde(default)]
pub enable_wasm: bool,
}
/// Tokenizer cache configuration
@@ -502,6 +505,7 @@ impl Default for RouterConfig {
client_identity: None,
ca_certificates: vec![],
mcp_config: None,
enable_wasm: false,
}
}
}

View File

@@ -17,7 +17,10 @@ use crate::{
app_context::AppContext,
config::{RouterConfig, RoutingMode},
core::workflow::{
steps::{McpServerConfigRequest, WorkerRemovalRequest},
steps::{
McpServerConfigRequest, WasmModuleConfigRequest, WasmModuleRemovalRequest,
WorkerRemovalRequest,
},
WorkflowContext, WorkflowEngine, WorkflowId, WorkflowInstanceId, WorkflowStatus,
},
mcp::McpConfig,
@@ -28,11 +31,27 @@ use crate::{
/// Job types for control plane operations
#[derive(Debug, Clone)]
pub enum Job {
AddWorker { config: Box<WorkerConfigRequest> },
RemoveWorker { url: String },
InitializeWorkersFromConfig { router_config: Box<RouterConfig> },
InitializeMcpServers { mcp_config: Box<McpConfig> },
RegisterMcpServer { config: Box<McpServerConfigRequest> },
AddWorker {
config: Box<WorkerConfigRequest>,
},
RemoveWorker {
url: String,
},
InitializeWorkersFromConfig {
router_config: Box<RouterConfig>,
},
InitializeMcpServers {
mcp_config: Box<McpConfig>,
},
RegisterMcpServer {
config: Box<McpServerConfigRequest>,
},
AddWasmModule {
config: Box<WasmModuleConfigRequest>,
},
RemoveWasmModule {
request: Box<WasmModuleRemovalRequest>,
},
}
impl Job {
@@ -44,10 +63,12 @@ impl Job {
Job::InitializeWorkersFromConfig { .. } => "InitializeWorkersFromConfig",
Job::InitializeMcpServers { .. } => "InitializeMcpServers",
Job::RegisterMcpServer { .. } => "RegisterMcpServer",
Job::AddWasmModule { .. } => "AddWasmModule",
Job::RemoveWasmModule { .. } => "RemoveWasmModule",
}
}
/// Get worker URL or MCP server name for logging
/// Get worker URL, MCP server name, or WASM module identifier for logging and status tracking
pub fn worker_url(&self) -> &str {
match self {
Job::AddWorker { config } => &config.url,
@@ -55,6 +76,8 @@ impl Job {
Job::InitializeWorkersFromConfig { .. } => "startup",
Job::InitializeMcpServers { .. } => "startup",
Job::RegisterMcpServer { config } => &config.name,
Job::AddWasmModule { config } => &config.descriptor.name,
Job::RemoveWasmModule { request } => &request.uuid_string,
}
}
}
@@ -347,6 +370,77 @@ impl JobQueue {
result
}
Job::AddWasmModule { config } => {
let engine = context
.workflow_engine
.get()
.ok_or_else(|| "Workflow engine not initialized".to_string())?;
let mut workflow_context = WorkflowContext::new(WorkflowInstanceId::new());
// Convert Box to Arc for context storage
let config_arc: Arc<WasmModuleConfigRequest> = Arc::new(*config.clone());
workflow_context.set_arc("wasm_module_config", config_arc);
workflow_context.set_arc("app_context", Arc::clone(context));
let instance_id = engine
.start_workflow(
WorkflowId::new("wasm_module_registration"),
workflow_context,
)
.await
.map_err(|e| {
format!("Failed to start WASM module registration workflow: {:?}", e)
})?;
debug!(
"Started WASM module registration workflow for {} (instance: {})",
config.descriptor.name, instance_id
);
let timeout_duration = Duration::from_secs(300); // 5 minutes
Self::wait_for_workflow_completion(
engine,
instance_id,
&config.descriptor.name,
timeout_duration,
)
.await
}
Job::RemoveWasmModule { request } => {
let engine = context
.workflow_engine
.get()
.ok_or_else(|| "Workflow engine not initialized".to_string())?;
let mut workflow_context = WorkflowContext::new(WorkflowInstanceId::new());
// Convert Box to Arc for context storage
let request_arc: Arc<WasmModuleRemovalRequest> = Arc::new(*request.clone());
workflow_context.set_arc("wasm_module_removal_request", request_arc);
workflow_context.set_arc("app_context", Arc::clone(context));
let instance_id = engine
.start_workflow(WorkflowId::new("wasm_module_removal"), workflow_context)
.await
.map_err(|e| {
format!("Failed to start WASM module removal workflow: {:?}", e)
})?;
debug!(
"Started WASM module removal workflow for {} (instance: {})",
request.module_uuid, instance_id
);
let timeout_duration = Duration::from_secs(60); // 1 minute
Self::wait_for_workflow_completion(
engine,
instance_id,
&request.module_uuid.to_string(),
timeout_duration,
)
.await
}
Job::InitializeWorkersFromConfig { router_config } => {
let api_key = router_config.api_key.clone();
let mut worker_count = 0;

View File

@@ -16,6 +16,7 @@ pub use executor::{FunctionStep, StepExecutor};
pub use state::WorkflowStateStore;
pub use steps::{
create_external_worker_registration_workflow, create_mcp_registration_workflow,
create_wasm_module_registration_workflow, create_wasm_module_removal_workflow,
create_worker_registration_workflow, create_worker_removal_workflow,
};
pub use types::*;

View File

@@ -5,9 +5,13 @@
//! - External worker registration (OpenAI, xAI, Anthropic, etc. - HTTPS only)
//! - Worker removal
//! - MCP server registration
//! - WASM module registration and removal
//! - Future: Tokenizer fetching, LoRA updates, etc.
pub mod external_worker_registration;
pub mod mcp_registration;
pub mod wasm_module_registration;
pub mod wasm_module_removal;
pub mod worker_registration;
pub mod worker_removal;
@@ -21,6 +25,15 @@ pub use mcp_registration::{
create_mcp_registration_workflow, ConnectMcpServerStep, DiscoverMcpInventoryStep,
McpServerConfigRequest, RegisterMcpServerStep, ValidateRegistrationStep,
};
pub use wasm_module_registration::{
create_wasm_module_registration_workflow, CalculateHashStep, CheckDuplicateStep,
LoadWasmBytesStep, RegisterModuleStep, ValidateDescriptorStep, ValidateWasmComponentStep,
WasmModuleConfigRequest,
};
pub use wasm_module_removal::{
create_wasm_module_removal_workflow, FindModuleToRemoveStep, RemoveModuleStep,
WasmModuleRemovalRequest,
};
pub use worker_registration::{
create_worker_registration_workflow, ActivateWorkerStep, CreateWorkerStep,
DetectConnectionModeStep, DiscoverDPInfoStep, DiscoverMetadataStep, RegisterWorkerStep,

View File

@@ -0,0 +1,474 @@
//! WASM Module Registration Workflow Steps
//!
//! Each step is atomic and performs a single operation in the WASM module registration process.
//!
//! Workflow order:
//! 1. ValidateDescriptor - Validate module descriptor (name, file_path, file existence)
//! 2. CalculateHash - Calculate SHA256 hash of the module file
//! 3. CheckDuplicate - Check for duplicate SHA256 hash
//! 4. LoadWasmBytes - Load WASM bytes into memory
//! 5. ValidateWasmComponent - Validate WASM component format
//! 6. RegisterModule - Register module in WasmModuleManager
use std::{sync::Arc, time::Duration};
use async_trait::async_trait;
use sha2::{Digest, Sha256};
use tracing::{debug, info};
use uuid::Uuid;
use wasmtime::{component::Component, Config, Engine};
use crate::{
app_context::AppContext,
core::workflow::*,
wasm::module::{WasmModule, WasmModuleDescriptor, WasmModuleMeta},
};
/// WASM module registration request
#[derive(Debug, Clone)]
pub struct WasmModuleConfigRequest {
/// Module descriptor containing name, file_path, attach_points, etc.
pub descriptor: WasmModuleDescriptor,
}
/// Step 1: Validate module descriptor
///
/// Validates that the module descriptor has all required fields:
/// - Module name is not empty
/// - File path is not empty
/// - File exists and is readable
/// - File size is not zero
pub struct ValidateDescriptorStep;
#[async_trait]
impl StepExecutor for ValidateDescriptorStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let config_request: Arc<WasmModuleConfigRequest> = context
.get("wasm_module_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_module_config".to_string()))?;
let descriptor = &config_request.descriptor;
debug!("Validating WASM module descriptor: {}", descriptor.name);
// Validate name
if descriptor.name.is_empty() {
return Err(WorkflowError::StepFailed {
step_id: StepId::new("validate_descriptor"),
message: "Module name cannot be empty".to_string(),
});
}
// Validate file path
if descriptor.file_path.is_empty() {
return Err(WorkflowError::StepFailed {
step_id: StepId::new("validate_descriptor"),
message: "Module file path cannot be empty".to_string(),
});
}
// Check if file exists and get size
let metadata = tokio::fs::metadata(&descriptor.file_path)
.await
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("validate_descriptor"),
message: format!("Failed to access file {}: {}", descriptor.file_path, e),
})?;
if metadata.len() == 0 {
return Err(WorkflowError::StepFailed {
step_id: StepId::new("validate_descriptor"),
message: "Module file size cannot be 0".to_string(),
});
}
// Store file size in context for later steps
context.set("file_size_bytes", metadata.len());
info!(
"Descriptor validated successfully for module: {}",
descriptor.name
);
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
false // Validation errors are not retryable (invalid input)
}
}
/// Step 2: Calculate SHA256 hash of the module file
///
/// Reads the file and calculates its SHA256 hash for deduplication.
/// This step is I/O intensive and may take time for large files.
pub struct CalculateHashStep;
#[async_trait]
impl StepExecutor for CalculateHashStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let config_request: Arc<WasmModuleConfigRequest> = context
.get("wasm_module_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_module_config".to_string()))?;
let file_path = &config_request.descriptor.file_path;
debug!("Calculating SHA256 hash for: {}", file_path);
// Read file in chunks to handle large files efficiently
let mut file =
tokio::fs::File::open(file_path)
.await
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("calculate_hash"),
message: format!("Failed to open file {}: {}", file_path, e),
})?;
let mut hasher = Sha256::new();
let mut buffer = vec![0u8; 1024 * 1024]; // 1MB buffer
loop {
use tokio::io::AsyncReadExt;
let bytes_read =
file.read(&mut buffer)
.await
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("calculate_hash"),
message: format!("Failed to read file {}: {}", file_path, e),
})?;
if bytes_read == 0 {
break;
}
hasher.update(&buffer[..bytes_read]);
}
let hash: [u8; 32] = hasher.finalize().into();
// Store hash in context
context.set("sha256_hash", hash);
info!("SHA256 hash calculated for: {}", file_path);
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
true // File I/O errors are retryable (network filesystem, etc.)
}
}
/// Step 3: Check for duplicate SHA256 hash
///
/// Checks if a module with the same SHA256 hash already exists in the manager.
/// This prevents duplicate modules from being registered.
pub struct CheckDuplicateStep;
#[async_trait]
impl StepExecutor for CheckDuplicateStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let config_request: Arc<WasmModuleConfigRequest> = context
.get("wasm_module_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_module_config".to_string()))?;
let app_context: Arc<AppContext> = context
.get("app_context")
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let sha256_hash: Arc<[u8; 32]> = context
.get("sha256_hash")
.ok_or_else(|| WorkflowError::ContextValueNotFound("sha256_hash".to_string()))?;
debug!(
"Checking for duplicate SHA256 hash for module: {}",
config_request.descriptor.name
);
// Get WASM module manager from app context
let wasm_manager =
app_context
.wasm_manager
.as_ref()
.ok_or_else(|| WorkflowError::StepFailed {
step_id: StepId::new("check_duplicate"),
message: "WASM module manager not initialized".to_string(),
})?;
// Check for duplicate hash using manager's internal method
wasm_manager
.check_duplicate_sha256_hash(sha256_hash.as_ref())
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("check_duplicate"),
message: format!("Duplicate SHA256 hash detected: {}", e),
})?;
info!(
"No duplicate found for module: {}",
config_request.descriptor.name
);
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
false // Duplicate check failures are not retryable
}
}
/// Step 4: Load WASM bytes into memory
///
/// Reads the entire WASM file into memory for faster execution.
/// This is an I/O operation that may take time for large files.
pub struct LoadWasmBytesStep;
#[async_trait]
impl StepExecutor for LoadWasmBytesStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let config_request: Arc<WasmModuleConfigRequest> = context
.get("wasm_module_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_module_config".to_string()))?;
let file_path = &config_request.descriptor.file_path;
debug!("Loading WASM bytes from: {}", file_path);
let wasm_bytes =
tokio::fs::read(file_path)
.await
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("load_wasm_bytes"),
message: format!("Failed to read WASM file {}: {}", file_path, e),
})?;
// Store WASM bytes in context
context.set("wasm_bytes", wasm_bytes);
info!("WASM bytes loaded from: {}", file_path);
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
true // File read errors are retryable
}
}
/// Step 5: Validate WASM component format
///
/// Validates that the loaded WASM bytes represent a valid component.
/// This catches format errors early during registration rather than during execution.
pub struct ValidateWasmComponentStep;
#[async_trait]
impl StepExecutor for ValidateWasmComponentStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let config_request: Arc<WasmModuleConfigRequest> = context
.get("wasm_module_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_module_config".to_string()))?;
let wasm_bytes: Arc<Vec<u8>> = context
.get("wasm_bytes")
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_bytes".to_string()))?;
debug!(
"Validating WASM component format for module: {}",
config_request.descriptor.name
);
// Create a temporary engine to validate the component
let mut config = Config::new();
config.async_support(true);
config.wasm_component_model(true);
let engine = Engine::new(&config).map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("validate_wasm_component"),
message: format!("Failed to create WASM engine: {}", e),
})?;
// Attempt to compile the component to validate it
Component::new(&engine, wasm_bytes.as_ref())
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("validate_wasm_component"),
message: format!(
"Invalid WASM component: {}. \
Hint: The WASM file must be in component format. \
If you're using wit-bindgen, use 'wasm-tools component new' to wrap the WASM module into a component.",
e
),
})?;
info!(
"WASM component validated successfully for module: {}",
config_request.descriptor.name
);
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
false // Validation errors are not retryable (invalid format)
}
}
/// Step 6: Register module in WasmModuleManager
///
/// Creates the WasmModule object and registers it in the manager's module map.
/// This is the final step that makes the module available for execution.
pub struct RegisterModuleStep;
#[async_trait]
impl StepExecutor for RegisterModuleStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let config_request: Arc<WasmModuleConfigRequest> = context
.get("wasm_module_config")
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_module_config".to_string()))?;
let app_context: Arc<AppContext> = context
.get("app_context")
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let sha256_hash: Arc<[u8; 32]> = context
.get("sha256_hash")
.ok_or_else(|| WorkflowError::ContextValueNotFound("sha256_hash".to_string()))?;
let file_size_bytes: Arc<u64> = context
.get("file_size_bytes")
.ok_or_else(|| WorkflowError::ContextValueNotFound("file_size_bytes".to_string()))?;
let wasm_bytes: Arc<Vec<u8>> = context
.get("wasm_bytes")
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_bytes".to_string()))?;
debug!(
"Registering WASM module in manager: {}",
config_request.descriptor.name
);
// Get WASM module manager from app context
let wasm_manager =
app_context
.wasm_manager
.as_ref()
.ok_or_else(|| WorkflowError::StepFailed {
step_id: StepId::new("register_module"),
message: "WASM module manager not initialized".to_string(),
})?;
// Create module metadata
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_else(|_| Duration::from_nanos(0))
.as_nanos() as u64;
let module_uuid = Uuid::new_v4();
let module = WasmModule {
module_uuid,
module_meta: WasmModuleMeta {
name: config_request.descriptor.name.clone(),
file_path: config_request.descriptor.file_path.clone(),
sha256_hash: *sha256_hash.as_ref(),
size_bytes: *file_size_bytes.as_ref(),
created_at: now,
last_accessed_at: now,
access_count: 0,
attach_points: config_request.descriptor.attach_points.clone(),
wasm_bytes: wasm_bytes.as_ref().clone(),
},
};
// Register module in manager
wasm_manager
.register_module_internal(module)
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("register_module"),
message: format!("Failed to register module: {}", e),
})?;
// Store module UUID in context for return value
context.set("module_uuid", module_uuid);
info!(
"WASM module registered successfully: {} (UUID: {})",
config_request.descriptor.name, module_uuid
);
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
false // Registration is a simple operation, not retryable
}
}
/// Create WASM module registration workflow
///
/// This workflow handles the complete process of registering a WASM module:
/// - Validates the module descriptor
/// - Calculates SHA256 hash for deduplication
/// - Checks for duplicates
/// - Loads WASM bytes into memory
/// - Validates WASM component format
/// - Registers the module in the manager
///
/// Workflow configuration:
/// - ValidateDescriptor: No retry, 5s timeout (fast validation)
/// - CalculateHash: 3 retries, 60s timeout (I/O intensive, may need retry)
/// - CheckDuplicate: No retry, 5s timeout (fast check)
/// - LoadWasmBytes: 3 retries, 60s timeout (I/O intensive)
/// - ValidateWasmComponent: No retry, 30s timeout (CPU intensive validation)
/// - RegisterModule: No retry, 5s timeout (fast registration)
pub fn create_wasm_module_registration_workflow() -> WorkflowDefinition {
WorkflowDefinition::new("wasm_module_registration", "WASM Module Registration")
.add_step(
StepDefinition::new(
"validate_descriptor",
"Validate Descriptor",
Arc::new(ValidateDescriptorStep),
)
.with_timeout(Duration::from_secs(5))
.with_failure_action(FailureAction::FailWorkflow),
)
.add_step(
StepDefinition::new(
"calculate_hash",
"Calculate SHA256 Hash",
Arc::new(CalculateHashStep),
)
.with_retry(RetryPolicy {
max_attempts: 3,
backoff: BackoffStrategy::Fixed(Duration::from_secs(1)),
})
.with_timeout(Duration::from_secs(60))
.with_failure_action(FailureAction::FailWorkflow),
)
.add_step(
StepDefinition::new(
"check_duplicate",
"Check Duplicate Hash",
Arc::new(CheckDuplicateStep),
)
.with_timeout(Duration::from_secs(5))
.with_failure_action(FailureAction::FailWorkflow),
)
.add_step(
StepDefinition::new(
"load_wasm_bytes",
"Load WASM Bytes",
Arc::new(LoadWasmBytesStep),
)
.with_retry(RetryPolicy {
max_attempts: 3,
backoff: BackoffStrategy::Fixed(Duration::from_secs(1)),
})
.with_timeout(Duration::from_secs(60))
.with_failure_action(FailureAction::FailWorkflow),
)
.add_step(
StepDefinition::new(
"validate_wasm_component",
"Validate WASM Component",
Arc::new(ValidateWasmComponentStep),
)
.with_timeout(Duration::from_secs(30))
.with_failure_action(FailureAction::FailWorkflow),
)
.add_step(
StepDefinition::new(
"register_module",
"Register Module",
Arc::new(RegisterModuleStep),
)
.with_timeout(Duration::from_secs(5))
.with_failure_action(FailureAction::FailWorkflow),
)
}

View File

@@ -0,0 +1,160 @@
//! WASM Module Removal Workflow Steps
//!
//! Each step is atomic and performs a single operation in the WASM module removal process.
//!
//! Workflow order:
//! 1. FindModuleToRemove - Find the module to remove by UUID
//! 2. RemoveModule - Remove module from WasmModuleManager
use std::{sync::Arc, time::Duration};
use async_trait::async_trait;
use tracing::{debug, info};
use uuid::Uuid;
use crate::{app_context::AppContext, core::workflow::*};
/// WASM module removal request
#[derive(Debug, Clone)]
pub struct WasmModuleRemovalRequest {
/// Module UUID to remove
pub module_uuid: Uuid,
/// Cached UUID string for worker_url() method
pub(crate) uuid_string: String,
}
impl WasmModuleRemovalRequest {
pub fn new(module_uuid: Uuid) -> Self {
Self {
module_uuid,
uuid_string: module_uuid.to_string(),
}
}
}
/// Step 1: Find module to remove
///
/// Verifies that the module exists before attempting removal.
pub struct FindModuleToRemoveStep;
#[async_trait]
impl StepExecutor for FindModuleToRemoveStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let removal_request: Arc<WasmModuleRemovalRequest> =
context.get("wasm_module_removal_request").ok_or_else(|| {
WorkflowError::ContextValueNotFound("wasm_module_removal_request".to_string())
})?;
let app_context: Arc<AppContext> = context
.get("app_context")
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
debug!("Finding module to remove: {}", removal_request.module_uuid);
// Get WASM module manager from app context
let wasm_manager =
app_context
.wasm_manager
.as_ref()
.ok_or_else(|| WorkflowError::StepFailed {
step_id: StepId::new("find_module_to_remove"),
message: "WASM module manager not initialized".to_string(),
})?;
// Check if module exists
let module = wasm_manager
.get_module(removal_request.module_uuid)
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("find_module_to_remove"),
message: format!("Failed to get module: {}", e),
})?;
if module.is_none() {
return Err(WorkflowError::StepFailed {
step_id: StepId::new("find_module_to_remove"),
message: format!("Module with UUID {} not found", removal_request.module_uuid),
});
}
info!("Module found for removal: {}", removal_request.module_uuid);
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
false // Module not found is not retryable
}
}
/// Step 2: Remove module from WasmModuleManager
///
/// Removes the module from the manager's module map.
pub struct RemoveModuleStep;
#[async_trait]
impl StepExecutor for RemoveModuleStep {
async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult<StepResult> {
let removal_request: Arc<WasmModuleRemovalRequest> =
context.get("wasm_module_removal_request").ok_or_else(|| {
WorkflowError::ContextValueNotFound("wasm_module_removal_request".to_string())
})?;
let app_context: Arc<AppContext> = context
.get("app_context")
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
debug!("Removing WASM module: {}", removal_request.module_uuid);
// Get WASM module manager from app context
let wasm_manager =
app_context
.wasm_manager
.as_ref()
.ok_or_else(|| WorkflowError::StepFailed {
step_id: StepId::new("remove_module"),
message: "WASM module manager not initialized".to_string(),
})?;
// Remove module from manager
wasm_manager
.remove_module_internal(removal_request.module_uuid)
.map_err(|e| WorkflowError::StepFailed {
step_id: StepId::new("remove_module"),
message: format!("Failed to remove module: {}", e),
})?;
info!(
"WASM module removed successfully: {}",
removal_request.module_uuid
);
Ok(StepResult::Success)
}
fn is_retryable(&self, _error: &WorkflowError) -> bool {
false // Removal is not retryable
}
}
/// Create WASM module removal workflow
///
/// This workflow handles the process of removing a WASM module:
/// - Finds the module to remove
/// - Removes it from the manager
///
/// Workflow configuration:
/// - FindModuleToRemove: No retry, 5s timeout (fast lookup)
/// - RemoveModule: No retry, 5s timeout (fast removal)
pub fn create_wasm_module_removal_workflow() -> WorkflowDefinition {
WorkflowDefinition::new("wasm_module_removal", "WASM Module Removal")
.add_step(
StepDefinition::new(
"find_module_to_remove",
"Find Module to Remove",
Arc::new(FindModuleToRemoveStep),
)
.with_timeout(Duration::from_secs(5))
.with_failure_action(FailureAction::FailWorkflow),
)
.add_step(
StepDefinition::new("remove_module", "Remove Module", Arc::new(RemoveModuleStep))
.with_timeout(Duration::from_secs(5))
.with_failure_action(FailureAction::FailWorkflow),
)
}

View File

@@ -18,3 +18,4 @@ pub mod service_discovery;
pub mod tokenizer;
pub mod tool_parser;
pub mod version;
pub mod wasm;

View File

@@ -345,6 +345,9 @@ struct CliArgs {
#[arg(long)]
mcp_config_path: Option<String>,
#[arg(long, default_value_t = false)]
enable_wasm: bool,
}
enum OracleConnectSource {
@@ -637,6 +640,7 @@ impl CliArgs {
.dp_aware(self.dp_aware)
.retries(!self.disable_retries)
.circuit_breaker(!self.disable_circuit_breaker)
.enable_wasm(self.enable_wasm)
.igw(self.enable_igw);
builder.build()

View File

@@ -21,7 +21,20 @@ use tower_http::trace::{MakeSpan, OnRequest, OnResponse, TraceLayer};
use tracing::{debug, error, field::Empty, info, info_span, warn, Span};
pub use crate::core::token_bucket::TokenBucket;
use crate::{metrics::RouterMetrics, server::AppState};
use crate::{
metrics::RouterMetrics,
server::AppState,
wasm::{
module::{MiddlewareAttachPoint, WasmModuleAttachPoint},
spec::{
apply_modify_action_to_headers, build_wasm_headers_from_axum_headers,
sgl::router::middleware_types::{
Action, Request as WasmRequest, Response as WasmResponse,
},
},
types::WasmComponentInput,
},
};
#[derive(Clone)]
pub struct AuthConfig {
@@ -554,3 +567,219 @@ pub async fn concurrency_limit_middleware(
}
}
}
pub async fn wasm_middleware(
State(app_state): State<Arc<AppState>>,
request: Request<Body>,
next: Next,
) -> Result<Response, StatusCode> {
// Check if WASM is enabled
if !app_state.context.router_config.enable_wasm {
return Ok(next.run(request).await);
}
// Get WASM manager
let wasm_manager = match &app_state.context.wasm_manager {
Some(manager) => manager,
None => {
return Ok(next.run(request).await);
}
};
// Get request ID from extensions or generate one
let request_id = request
.extensions()
.get::<RequestId>()
.map(|r| r.0.clone())
.unwrap_or_else(|| generate_request_id(request.uri().path()));
// ===== OnRequest Phase =====
let on_request_attach_point =
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnRequest);
let modules_on_request =
match wasm_manager.get_modules_by_attach_point(on_request_attach_point.clone()) {
Ok(modules) => modules,
Err(e) => {
error!("Failed to get WASM modules for OnRequest: {}", e);
return Ok(next.run(request).await);
}
};
// Extract request body once before processing modules
let method = request.method().clone();
let uri = request.uri().clone();
let mut headers = request.headers().clone();
let body_bytes = match axum::body::to_bytes(request.into_body(), usize::MAX).await {
Ok(bytes) => bytes.to_vec(),
Err(e) => {
error!("Failed to read request body: {}", e);
// Create a minimal request with empty body for error recovery
let error_request = Request::builder()
.uri(uri)
.body(Body::empty())
.unwrap_or_else(|_| Request::new(Body::empty()));
return Ok(next.run(error_request).await);
}
};
// Process each OnRequest module
let mut modified_body = body_bytes;
for module in modules_on_request {
// Build WebAssembly request from collected data
let wasm_headers = build_wasm_headers_from_axum_headers(&headers);
let wasm_request = WasmRequest {
method: method.to_string(),
path: uri.path().to_string(),
query: uri.query().unwrap_or("").to_string(),
headers: wasm_headers,
body: modified_body.clone(),
request_id: request_id.clone(),
now_epoch_ms: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_else(|_| {
// Fallback to 0 if system time is before UNIX_EPOCH
// This should never happen in practice, but provides a safe fallback
Duration::from_millis(0)
})
.as_millis() as u64,
};
// Execute WASM component
let action = match wasm_manager
.execute_module_for_attach_point(
&module,
on_request_attach_point.clone(),
WasmComponentInput::MiddlewareRequest(wasm_request),
)
.await
{
Some(action) => action,
None => continue, // Continue to next module on error
};
// Process action
match action {
Action::Continue => {
// Continue to next module or request processing
}
Action::Reject(status) => {
// Immediately reject the request
return Err(StatusCode::from_u16(status).unwrap_or(StatusCode::BAD_REQUEST));
}
Action::Modify(modify) => {
// Apply modifications to headers and body
apply_modify_action_to_headers(&mut headers, &modify);
// Apply body_replace
if let Some(body_bytes) = modify.body_replace {
modified_body = body_bytes;
}
}
}
}
// Reconstruct request with modifications
let mut final_request = Request::builder()
.method(method)
.uri(uri)
.body(Body::from(modified_body))
.unwrap_or_else(|_| Request::new(Body::empty()));
*final_request.headers_mut() = headers;
// Continue with request processing
let response = next.run(final_request).await;
// ===== OnResponse Phase =====
let on_response_attach_point =
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnResponse);
let modules_on_response =
match wasm_manager.get_modules_by_attach_point(on_response_attach_point.clone()) {
Ok(modules) => modules,
Err(e) => {
error!("Failed to get WASM modules for OnResponse: {}", e);
return Ok(response);
}
};
// Extract response data once before processing modules
let mut status = response.status();
let mut headers = response.headers().clone();
let mut body_bytes = match axum::body::to_bytes(response.into_body(), usize::MAX).await {
Ok(bytes) => bytes.to_vec(),
Err(e) => {
error!("Failed to read response body: {}", e);
// Create a minimal response with empty body for error recovery
let error_response = Response::builder()
.status(status)
.body(Body::empty())
.unwrap_or_else(|_| Response::new(Body::empty()));
return Ok(error_response);
}
};
// Process each OnResponse module
for module in modules_on_response {
// Build WebAssembly response from collected data
let wasm_headers = build_wasm_headers_from_axum_headers(&headers);
let wasm_response = WasmResponse {
status: status.as_u16(),
headers: wasm_headers,
body: body_bytes.clone(),
};
// Execute WASM component
let action = match wasm_manager
.execute_module_for_attach_point(
&module,
on_response_attach_point.clone(),
WasmComponentInput::MiddlewareResponse(wasm_response),
)
.await
{
Some(action) => action,
None => continue, // Continue to next module on error
};
// Process action - apply modifications incrementally
match action {
Action::Continue => {
// Continue to next module
}
Action::Reject(status_code) => {
// Override response status
status = StatusCode::from_u16(status_code).unwrap_or(StatusCode::BAD_REQUEST);
// Return immediately with current state
let final_response = Response::builder()
.status(status)
.body(Body::from(body_bytes))
.unwrap_or_else(|_| Response::new(Body::empty()));
let mut final_response = final_response;
*final_response.headers_mut() = headers;
return Ok(final_response);
}
Action::Modify(modify) => {
// Apply status modification
if let Some(new_status) = modify.status {
status = StatusCode::from_u16(new_status).unwrap_or(status);
}
// Apply headers modifications
apply_modify_action_to_headers(&mut headers, &modify);
// Apply body_replace
if let Some(new_body) = modify.body_replace {
body_bytes = new_body;
}
}
}
}
// Reconstruct final response with all modifications
let final_response = Response::builder()
.status(status)
.body(Body::from(body_bytes))
.unwrap_or_else(|_| Response::new(Body::empty()));
let mut final_response = final_response;
*final_response.headers_mut() = headers;
Ok(final_response)
}

View File

@@ -26,6 +26,7 @@ use crate::{
worker_to_info,
workflow::{
create_external_worker_registration_workflow, create_mcp_registration_workflow,
create_wasm_module_registration_workflow, create_wasm_module_removal_workflow,
create_worker_registration_workflow, create_worker_removal_workflow, LoggingSubscriber,
WorkflowEngine,
},
@@ -47,6 +48,7 @@ use crate::{
},
routers::{conversations, router_manager::RouterManager, RouterTrait},
service_discovery::{start_service_discovery, ServiceDiscoveryConfig},
wasm::route::{add_wasm_module, list_wasm_modules, remove_wasm_module},
};
#[derive(Clone)]
@@ -645,6 +647,10 @@ pub fn build_app(
.route_layer(axum::middleware::from_fn_with_state(
auth_config.clone(),
middleware::auth_middleware,
))
.route_layer(axum::middleware::from_fn_with_state(
app_state.clone(),
middleware::wasm_middleware,
));
let public_routes = Router::new()
@@ -660,6 +666,9 @@ pub fn build_app(
let admin_routes = Router::new()
.route("/flush_cache", post(flush_cache))
.route("/get_loads", get(get_loads))
.route("/wasm", post(add_wasm_module))
.route("/wasm/{module_uuid}", delete(remove_wasm_module))
.route("/wasm", get(list_wasm_modules))
.route_layer(axum::middleware::from_fn_with_state(
auth_config.clone(),
middleware::auth_middleware,
@@ -753,6 +762,8 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box<dyn std::error::Err
engine.register_workflow(create_external_worker_registration_workflow());
engine.register_workflow(create_worker_removal_workflow());
engine.register_workflow(create_mcp_registration_workflow());
engine.register_workflow(create_wasm_module_registration_workflow());
engine.register_workflow(create_wasm_module_removal_workflow());
app_context
.workflow_engine
.set(engine)

View File

@@ -595,6 +595,7 @@ mod tests {
worker_job_queue: Arc::new(std::sync::OnceLock::new()),
workflow_engine: Arc::new(std::sync::OnceLock::new()),
mcp_manager: Arc::new(std::sync::OnceLock::new()),
wasm_manager: None,
})
}

View File

@@ -0,0 +1,227 @@
# WebAssembly (WASM) Extensibility for sgl-router
This module provides WebAssembly-based extensibility for sgl-router, enabling dynamic, safe, and portable middleware execution without requiring router restarts or recompilation.
## Overview
The WASM module allows you to extend sgl-router functionality by deploying WebAssembly components that can:
- **Intercept requests/responses** at various lifecycle points (OnRequest, OnResponse)
- **Modify HTTP headers and bodies** before/after processing
- **Reject requests** with custom status codes
- **Execute custom logic** in a sandboxed, isolated environment
## Architecture
### Components
The WASM module consists of several key components:
```
src/wasm/
├── module.rs # Data structures (metadata, types, attach points)
├── module_manager.rs # Module lifecycle management (add/remove/list)
├── runtime.rs # WASM execution engine and thread pool
├── route.rs # HTTP API endpoints for module management
├── spec.rs # WASM interface types bindings and type conversions
├── types.rs # Generic input/output types
├── errors.rs # Error definitions
├── config.rs # Runtime configuration
└── interface/ # WebAssembly Interface Types definitions
```
### Execution Flow
```
1. HTTP Request arrives at router
2. Middleware chain checks for WASM modules attached to OnRequest
3. For each module:
a. Module manager retrieves pre-loaded WASM bytes
b. Runtime executes component in isolated worker thread
c. Component processes request via WASM type interface
d. Returns Action (Continue/Reject/Modify)
4. If Continue: proceed to next middleware/upstream
If Reject: return error response immediately
If Modify: apply changes (headers, body, status)
5. After upstream response:
- Modules attached to OnResponse process response
- Apply modifications
6. Return final response to client
```
### WebAssembly Interface Types
The module uses the WebAssembly Component Model with WASM interface type for type-safe communication between host and WASM components:
- **Request Processing**: `middleware-on-request::on-request(req: Request) -> Action`
- **Response Processing**: `middleware-on-response::on-response(resp: Response) -> Action`
- **Actions**: `Continue`, `Reject(status)`, or `Modify(modify-action)`
See [`interface/`](./interface/) for the complete interface definition.
## Usage
### Prerequisites
- sgl-router compiled with WASM support
- Rust toolchain (for building WASM components)
- `wasm32-wasip2` target: `rustup target add wasm32-wasip2`
- `wasm-tools`: `cargo install wasm-tools`
### Starting the Router
Enable WASM support when starting the router:
```bash
./sgl-router --enable-wasm --worker-urls=http://0.0.0.0:30000 --port=3000
```
### Deploying a WASM Module
Use the `/wasm` POST endpoint to deploy modules:
```bash
curl -X POST http://localhost:3000/wasm \
-H "Content-Type: application/json" \
-d '{
"modules": [{
"name": "my-middleware",
"file_path": "/path/to/my-component.component.wasm",
"module_type": "Middleware",
"attach_points": [{"Middleware": "OnRequest"}]
}]
}'
```
### Managing Modules
**List all modules:**
```bash
curl http://localhost:3000/wasm
```
**Remove a module:**
```bash
curl -X DELETE http://localhost:3000/wasm/{module-uuid}
```
### Module Configuration
Each module requires:
- **name**: Unique identifier for the module
- **file_path**: Absolute path to the WASM component file
- **module_type**: Currently supports `"Middleware"`
- **attach_points**: List of attachment points, e.g., `[{"Middleware": "OnRequest"}]`
Supported attachment points:
- `{"Middleware": "OnRequest"}` - Execute before forwarding to upstream
- `{"Middleware": "OnResponse"}` - Execute after receiving upstream response
- `{"Middleware": "OnError"}` - Not yet implemented
## Examples
See [`examples/wasm/`](../../examples/wasm/) for complete examples:
1. **[wasm-guest-auth](../../examples/wasm/wasm-guest-auth/)** - API key authentication middleware
2. **[wasm-guest-logging](../../examples/wasm/wasm-guest-logging/)** - Request tracking and status code conversion
3. **[wasm-guest-ratelimit](../../examples/wasm/wasm-guest-ratelimit/)** - Rate limiting middleware
Each example includes:
- Complete source code
- Build instructions
- Deployment examples
- Testing guidelines
## Security and Resource Management
### Sandboxing
WASM modules run in isolated environments provided by wasmtime, preventing:
- Direct system access
- Memory corruption of the host process
- Unauthorized network access
- File system access (unless explicitly granted via WASI)
### Resource Limits
Runtime configuration allows setting limits:
```rust
WasmRuntimeConfig {
max_memory_pages: 1024, // 64MB limit
max_execution_time_ms: 1000, // 1 second timeout
max_stack_size: 1024 * 1024, // 1MB stack
thread_pool_size: 4, // Worker threads
module_cache_size: 10, // Cached modules per worker
}
```
### Error Handling
- Failed module executions are logged and don't crash the router
- Invalid WASM components are rejected during load time
- Metrics track execution success/failure rates
## Metrics
The module exposes execution metrics via the `/wasm` GET endpoint:
```json
{
"modules": [...],
"metrics": {
"total_executions": 1000,
"successful_executions": 995,
"failed_executions": 5,
"total_execution_time_ms": 50000,
"max_execution_time_ms": 150,
"average_execution_time_ms": 50.0
}
}
```
## Development
### Building WASM Components
WASM components must be built using the Component Model. For Rust:
```bash
# 1. Build as WASM module
cargo build --target wasm32-wasip2 --release
# 2. Wrap into component format
wasm-tools component new target/wasm32-wasip2/release/my_module.wasm \
-o my_module.component.wasm
```
### WASM Interface Type
Define your component using the WASM interface from `interface/spec.*`:
```rust
wit_bindgen::generate!({
path: "../../../src/wasm/interface",
world: "sgl-router",
});
use exports::sgl::router::middleware_on_request::Guest as OnRequestGuest;
use sgl::router::middleware_types::{Request, Action};
struct Middleware;
impl OnRequestGuest for Middleware {
fn on_request(req: Request) -> Action {
// Your logic here
Action::Continue
}
}
export!(Middleware);
```

View File

@@ -0,0 +1,297 @@
//! WASM Runtime Configuration
//!
//! Defines configuration parameters for the WASM runtime,
//! including memory limits, execution timeouts, and thread pool settings.
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct WasmRuntimeConfig {
/// Maximum memory size in pages (64KB per page)
pub max_memory_pages: u32,
/// Maximum execution time in milliseconds
pub max_execution_time_ms: u64,
/// Maximum stack size in bytes
pub max_stack_size: usize,
/// Number of worker threads in the pool
pub thread_pool_size: usize,
/// Maximum number of modules to cache per worker
pub module_cache_size: usize,
}
impl Default for WasmRuntimeConfig {
fn default() -> Self {
let default_thread_pool_size = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4)
.max(1);
Self {
max_memory_pages: 1024, // 64MB
max_execution_time_ms: 1000, // 1 seconds
max_stack_size: 1024 * 1024, // 1MB
thread_pool_size: default_thread_pool_size, // based on cpu count
module_cache_size: 10, // Cache up to 10 modules per worker
}
}
}
impl WasmRuntimeConfig {
/// Validate the configuration parameters
pub fn validate(&self) -> Result<(), String> {
// Validate max_memory_pages
if self.max_memory_pages == 0 {
return Err("max_memory_pages cannot be 0".to_string());
}
if self.max_memory_pages > 65536 {
return Err("max_memory_pages cannot exceed 65536 (4GB)".to_string());
}
// Validate max_execution_time_ms
if self.max_execution_time_ms == 0 {
return Err("max_execution_time_ms cannot be 0".to_string());
}
if self.max_execution_time_ms > 300000 {
return Err("max_execution_time_ms cannot exceed 300000ms (5 minutes)".to_string());
}
// Validate max_stack_size
if self.max_stack_size == 0 {
return Err("max_stack_size cannot be 0".to_string());
}
if self.max_stack_size < 64 * 1024 {
return Err("max_stack_size must be at least 64KB".to_string());
}
if self.max_stack_size > 16 * 1024 * 1024 {
return Err("max_stack_size cannot exceed 16MB".to_string());
}
// Validate thread_pool_size
if self.thread_pool_size == 0 {
return Err("thread_pool_size cannot be 0".to_string());
}
if self.thread_pool_size > 128 {
return Err("thread_pool_size cannot exceed 128".to_string());
}
// Validate module_cache_size
if self.module_cache_size == 0 {
return Err("module_cache_size cannot be 0".to_string());
}
if self.module_cache_size > 1000 {
return Err("module_cache_size cannot exceed 1000".to_string());
}
Ok(())
}
/// Create a new config with validation
pub fn new(
max_memory_pages: u32,
max_execution_time_ms: u64,
max_stack_size: usize,
thread_pool_size: usize,
module_cache_size: usize,
) -> Result<Self, String> {
let config = Self {
max_memory_pages,
max_execution_time_ms,
max_stack_size,
thread_pool_size,
module_cache_size,
};
config.validate()?;
Ok(config)
}
/// Get the total memory size in bytes
pub fn get_total_memory_bytes(&self) -> u64 {
self.max_memory_pages as u64 * 64 * 1024 // 64KB per page
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config_validation() {
let config = WasmRuntimeConfig::default();
assert!(config.validate().is_ok());
}
#[test]
fn test_config_new_with_validation() {
let config = WasmRuntimeConfig::new(1024, 1000, 1024 * 1024, 2, 10);
assert!(config.is_ok());
}
#[test]
fn test_validation_max_memory_pages_zero() {
let config = WasmRuntimeConfig {
max_memory_pages: 0,
max_execution_time_ms: 1000,
max_stack_size: 1024 * 1024,
thread_pool_size: 2,
module_cache_size: 10,
};
let result = config.validate();
assert!(result.is_err());
assert!(result.unwrap_err().contains("max_memory_pages cannot be 0"));
}
#[test]
fn test_validation_max_memory_pages_too_large() {
let config = WasmRuntimeConfig {
max_memory_pages: 65537, // Exceeds 4GB limit
max_execution_time_ms: 1000,
max_stack_size: 1024 * 1024,
thread_pool_size: 2,
module_cache_size: 10,
};
let result = config.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("max_memory_pages cannot exceed 65536"));
}
#[test]
fn test_validation_max_execution_time_zero() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 0,
max_stack_size: 1024 * 1024,
thread_pool_size: 2,
module_cache_size: 10,
};
let result = config.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("max_execution_time_ms cannot be 0"));
}
#[test]
fn test_validation_max_execution_time_too_large() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 300001, // Exceeds 5 minutes
max_stack_size: 1024 * 1024,
thread_pool_size: 2,
module_cache_size: 10,
};
let result = config.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("max_execution_time_ms cannot exceed 300000ms"));
}
#[test]
fn test_validation_max_stack_size_too_small() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 1000,
max_stack_size: 32 * 1024, // Less than 64KB
thread_pool_size: 2,
module_cache_size: 10,
};
let result = config.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("max_stack_size must be at least 64KB"));
}
#[test]
fn test_validation_max_stack_size_too_large() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 1000,
max_stack_size: 17 * 1024 * 1024, // Exceeds 16MB
thread_pool_size: 2,
module_cache_size: 10,
};
let result = config.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("max_stack_size cannot exceed 16MB"));
}
#[test]
fn test_validation_thread_pool_size_zero() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 1000,
max_stack_size: 1024 * 1024,
thread_pool_size: 0,
module_cache_size: 10,
};
let result = config.validate();
assert!(result.is_err());
assert!(result.unwrap_err().contains("thread_pool_size cannot be 0"));
}
#[test]
fn test_validation_thread_pool_size_too_large() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 1000,
max_stack_size: 1024 * 1024,
thread_pool_size: 129, // Exceeds 128
module_cache_size: 10,
};
let result = config.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("thread_pool_size cannot exceed 128"));
}
#[test]
fn test_validation_module_cache_size_zero() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 1000,
max_stack_size: 1024 * 1024,
thread_pool_size: 2,
module_cache_size: 0,
};
let result = config.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("module_cache_size cannot be 0"));
}
#[test]
fn test_validation_module_cache_size_too_large() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 1000,
max_stack_size: 1024 * 1024,
thread_pool_size: 2,
module_cache_size: 1001, // Exceeds 1000
};
let result = config.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("module_cache_size cannot exceed 1000"));
}
#[test]
fn test_get_total_memory_bytes() {
let config = WasmRuntimeConfig {
max_memory_pages: 1024,
max_execution_time_ms: 1000,
max_stack_size: 1024 * 1024,
thread_pool_size: 2,
module_cache_size: 10,
};
// 1024 pages * 64KB = 64MB
assert_eq!(config.get_total_memory_bytes(), 64 * 1024 * 1024);
}
}

View File

@@ -0,0 +1,117 @@
//! WASM Error Types
//!
//! Defines comprehensive error types for the WASM subsystem,
//! including module, manager, and runtime errors.
use std::fmt;
use thiserror::Error;
pub type Result<T> = std::result::Result<T, WasmError>;
/// SHA256 hash wrapper for display purposes
#[derive(Debug, Clone, Copy)]
pub struct Sha256Hash(pub [u8; 32]);
impl fmt::Display for Sha256Hash {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let hex_string: String = self.0.iter().map(|b| format!("{:02x}", b)).collect();
write!(f, "{}", hex_string)
}
}
impl From<[u8; 32]> for Sha256Hash {
fn from(hash: [u8; 32]) -> Self {
Sha256Hash(hash)
}
}
#[derive(Debug, Error)]
pub enum WasmError {
#[error(transparent)]
Module(#[from] WasmModuleError),
#[error(transparent)]
Manager(#[from] WasmManagerError),
#[error(transparent)]
Runtime(#[from] WasmRuntimeError),
#[error(transparent)]
Io(#[from] std::io::Error),
#[error("{0}")]
Other(String),
}
#[derive(Debug, Error)]
pub enum WasmModuleError {
#[error("invalid module descriptor: {0}")]
InvalidDescriptor(String),
#[error("module with same sha256 already exists: {0}")]
DuplicateSha256(Sha256Hash),
#[error("module not found: {0}")]
NotFound(uuid::Uuid),
#[error("failed to read module file: {0}")]
FileRead(String),
#[error("validation failed: {0}")]
ValidationFailed(String),
#[error("attach point missing: {0}")]
AttachPointMissing(String),
#[error("invalid function for attach point: {0}")]
AttachPointFunctionInvalid(String),
}
#[derive(Debug, Error)]
pub enum WasmManagerError {
#[error("failed to acquire lock: {0}")]
LockFailed(String),
#[error("module add failed: {0}")]
ModuleAddFailed(String),
#[error("module remove failed: {0}")]
ModuleRemoveFailed(String),
#[error("runtime unavailable")]
RuntimeUnavailable,
#[error("execution failed: {0}")]
ExecutionFailed(String),
#[error("module {0} not found")]
ModuleNotFound(uuid::Uuid),
}
#[derive(Debug, Error)]
pub enum WasmRuntimeError {
#[error("failed to create engine: {0}")]
EngineCreateFailed(String),
#[error("failed to compile module: {0}")]
CompileFailed(String),
#[error("failed to create instance: {0}")]
InstanceCreateFailed(String),
#[error("function not found: {0}")]
FunctionNotFound(String),
#[error("execution timeout")]
Timeout,
#[error("execution failed: {0}")]
CallFailed(String),
}
impl From<wasmtime::Error> for WasmError {
fn from(value: wasmtime::Error) -> Self {
WasmError::Runtime(WasmRuntimeError::CallFailed(value.to_string()))
}
}

View File

@@ -0,0 +1,54 @@
package sgl:router;
interface middleware-types {
record header { name: string, value: string }
// onRequest
record request {
method: string,
path: string,
query: string,
headers: list<header>,
body: list<u8>,
request-id: string,
now-epoch-ms: u64,
}
// onResponse
record response {
status: u16,
headers: list<header>,
body: list<u8>,
}
// modify action
record modify-action {
status: option<u16>,
headers-set: list<header>,
headers-add: list<header>,
headers-remove: list<string>,
body-replace: option<list<u8>>,
}
// return actions
variant action {
continue,
reject(u16), // status code
modify(modify-action),
}
}
interface middleware-on-request {
use middleware-types.{request, action};
on-request: func(req: request) -> action;
}
interface middleware-on-response {
use middleware-types.{response, action};
on-response: func(resp: response) -> action;
}
world sgl-router {
export middleware-on-request;
export middleware-on-response;
}

View File

@@ -0,0 +1,13 @@
//! WebAssembly (WASM) module support for sgl-router
//!
//! This module provides WASM component execution capabilities using the WebAssembly Component Model.
//! It supports middleware execution at various attach points (OnRequest, OnResponse) with async support.
pub mod config;
pub mod errors;
pub mod module;
pub mod module_manager;
pub mod route;
pub mod runtime;
pub mod spec;
pub mod types;

View File

@@ -0,0 +1,209 @@
//! WASM Module Data Structures and Types
//!
//! This module defines the core data structures for managing WebAssembly components:
//! - Module metadata (UUID, name, file path, hash, timestamps, metrics)
//! - Module types and attachment points (Middleware hooks: OnRequest, OnResponse, OnError)
//! - API request/response types for module management
//! - Execution metrics and statistics
//!
//! The module provides custom serialization for:
//! - SHA256 hashes (hex string representation)
//! - Timestamps (ISO 8601 format for JSON output)
use serde::{Deserialize, Serialize, Serializer};
use uuid::Uuid;
/// Serialize [u8; 32] as hex string
fn serialize_sha256_hash<S>(hash: &[u8; 32], serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let hex_string = hash
.iter()
.map(|b| format!("{:02x}", b))
.collect::<String>();
serializer.serialize_str(&hex_string)
}
/// Deserialize hex string to [u8; 32]
fn deserialize_sha256_hash<'de, D>(deserializer: D) -> Result<[u8; 32], D::Error>
where
D: serde::Deserializer<'de>,
{
use serde::Deserialize;
let hex_string = String::deserialize(deserializer)?;
// Parse hex string to bytes
if hex_string.len() != 64 {
return Err(serde::de::Error::custom(format!(
"Invalid SHA256 hash length: expected 64 hex characters, got {}",
hex_string.len()
)));
}
let mut hash = [0u8; 32];
for (i, chunk) in hex_string.as_bytes().chunks(2).enumerate() {
if chunk.len() != 2 {
return Err(serde::de::Error::custom("Invalid hex string format"));
}
let byte_str = std::str::from_utf8(chunk)
.map_err(|e| serde::de::Error::custom(format!("Invalid UTF-8: {}", e)))?;
hash[i] = u8::from_str_radix(byte_str, 16)
.map_err(|e| serde::de::Error::custom(format!("Invalid hex digit: {}", e)))?;
}
Ok(hash)
}
/// Serialize u64 timestamp (nanoseconds since epoch) as ISO 8601 string
fn serialize_timestamp<S>(timestamp: &u64, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
use chrono::{DateTime, Utc};
// Convert nanoseconds to seconds and remaining nanoseconds
let secs = (*timestamp / 1_000_000_000) as i64;
let nanos = (*timestamp % 1_000_000_000) as u32;
match DateTime::<Utc>::from_timestamp(secs, nanos) {
Some(dt) => {
let s = dt.to_rfc3339_opts(chrono::SecondsFormat::Nanos, true);
serializer.serialize_str(&s)
}
None => {
// Fallback: format manually if timestamp is out of range
let s = format!("{}", timestamp);
serializer.serialize_str(&s)
}
}
}
/// Deserialize ISO 8601 string to u64 timestamp (nanoseconds since epoch)
fn deserialize_timestamp<'de, D>(deserializer: D) -> Result<u64, D::Error>
where
D: serde::Deserializer<'de>,
{
use chrono::{DateTime, Utc};
use serde::Deserialize;
let timestamp_str = String::deserialize(deserializer)?;
// Try to parse as ISO 8601 datetime (RFC 3339)
match DateTime::parse_from_rfc3339(&timestamp_str) {
Ok(dt) => {
// Convert to UTC and then to nanoseconds since epoch
let dt_utc = dt.with_timezone(&Utc);
let secs = dt_utc.timestamp();
let nanos = dt_utc.timestamp_subsec_nanos();
Ok((secs as u64) * 1_000_000_000 + (nanos as u64))
}
Err(_) => {
// Fallback: try to parse as u64 directly
timestamp_str
.parse::<u64>()
.map_err(|e| serde::de::Error::custom(format!("Invalid timestamp format: {}", e)))
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WasmModule {
// unique identifier for the module
pub module_uuid: Uuid,
pub module_meta: WasmModuleMeta,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum WasmModuleAddResult {
Success(Uuid),
Error(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WasmModuleDescriptor {
pub name: String,
pub file_path: String,
pub module_type: WasmModuleType,
pub attach_points: Vec<WasmModuleAttachPoint>,
pub add_result: Option<WasmModuleAddResult>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WasmModuleMeta {
// module name provided by the user
pub name: String,
// path to the module file
pub file_path: String,
// sha256 hash of the module file
#[serde(
serialize_with = "serialize_sha256_hash",
deserialize_with = "deserialize_sha256_hash"
)]
pub sha256_hash: [u8; 32],
// size of the module file in bytes
pub size_bytes: u64,
// timestamp of when the module was created (nanoseconds since epoch)
#[serde(
serialize_with = "serialize_timestamp",
deserialize_with = "deserialize_timestamp"
)]
pub created_at: u64,
// timestamp of when the module was last accessed (nanoseconds since epoch)
#[serde(
serialize_with = "serialize_timestamp",
deserialize_with = "deserialize_timestamp"
)]
pub last_accessed_at: u64,
// number of times the module was accessed
pub access_count: u64,
// attach points for the module
pub attach_points: Vec<WasmModuleAttachPoint>,
// Pre-loaded WASM component bytes (loaded into memory for faster execution)
#[serde(skip)]
pub wasm_bytes: Vec<u8>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Eq, PartialEq, Hash)]
pub enum WasmModuleType {
Middleware,
}
#[derive(Debug, Clone, Serialize, Deserialize, Eq, PartialEq, Hash)]
pub enum MiddlewareAttachPoint {
OnRequest,
OnResponse,
OnError,
}
#[derive(Debug, Clone, Serialize, Deserialize, Eq, PartialEq, Hash)]
pub enum WasmModuleAttachPoint {
Middleware(MiddlewareAttachPoint),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WasmModuleAddRequest {
pub modules: Vec<WasmModuleDescriptor>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WasmModuleAddResponse {
pub modules: Vec<WasmModuleDescriptor>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WasmModuleListResponse {
pub modules: Vec<WasmModule>,
pub metrics: WasmMetrics,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WasmMetrics {
pub total_executions: u64,
pub successful_executions: u64,
pub failed_executions: u64,
pub total_execution_time_ms: u64,
pub max_execution_time_ms: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub average_execution_time_ms: Option<f64>,
}

View File

@@ -0,0 +1,259 @@
//! WASM Module Manager
use std::{
collections::HashMap,
sync::{
atomic::{AtomicU64, Ordering},
Arc, RwLock,
},
};
use uuid::Uuid;
use crate::wasm::{
config::WasmRuntimeConfig,
errors::{Result, WasmError, WasmManagerError, WasmModuleError},
module::{WasmModule, WasmModuleAttachPoint},
runtime::WasmRuntime,
types::{WasmComponentInput, WasmComponentOutput},
};
pub struct WasmModuleManager {
modules: Arc<RwLock<HashMap<Uuid, WasmModule>>>,
runtime: Arc<WasmRuntime>,
// Metrics
total_executions: AtomicU64,
successful_executions: AtomicU64,
failed_executions: AtomicU64,
total_execution_time_ms: AtomicU64,
max_execution_time_ms: AtomicU64,
}
impl WasmModuleManager {
pub fn new(config: WasmRuntimeConfig) -> Result<Self> {
let runtime = Arc::new(WasmRuntime::new(config)?);
Ok(Self {
modules: Arc::new(RwLock::new(HashMap::new())),
runtime,
total_executions: AtomicU64::new(0),
successful_executions: AtomicU64::new(0),
failed_executions: AtomicU64::new(0),
total_execution_time_ms: AtomicU64::new(0),
max_execution_time_ms: AtomicU64::new(0),
})
}
pub fn with_default_config() -> Result<Self> {
Self::new(WasmRuntimeConfig::default())
}
/// Register a module (for workflow steps)
pub(crate) fn register_module_internal(&self, module: WasmModule) -> Result<()> {
let mut modules = self
.modules
.write()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
modules.insert(module.module_uuid, module);
Ok(())
}
/// Remove a module (for workflow steps)
pub(crate) fn remove_module_internal(&self, module_uuid: Uuid) -> Result<()> {
let mut modules = self
.modules
.write()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
if !modules.contains_key(&module_uuid) {
return Err(WasmManagerError::ModuleNotFound(module_uuid).into());
}
modules.remove(&module_uuid);
Ok(())
}
pub(crate) fn check_duplicate_sha256_hash(&self, sha256_hash: &[u8; 32]) -> Result<()> {
let modules = self
.modules
.read()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
if modules
.values()
.any(|module: &WasmModule| module.module_meta.sha256_hash == *sha256_hash)
{
return Err(WasmModuleError::DuplicateSha256((*sha256_hash).into()).into());
}
Ok(())
}
pub fn get_all_modules(&self) -> Result<Vec<WasmModule>> {
let modules = self
.modules
.read()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
Ok(modules.values().cloned().collect())
}
pub fn get_module(&self, module_uuid: Uuid) -> Result<Option<WasmModule>> {
let modules = self
.modules
.read()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
Ok(modules.get(&module_uuid).cloned())
}
pub fn get_modules(&self) -> Result<Vec<WasmModule>> {
let modules = self
.modules
.read()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
Ok(modules.values().cloned().collect())
}
/// get modules by attach point
pub fn get_modules_by_attach_point(
&self,
attach_point: WasmModuleAttachPoint,
) -> Result<Vec<WasmModule>> {
let modules = self
.modules
.read()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
Ok(modules
.values()
.filter(|module| module.module_meta.attach_points.contains(&attach_point))
.cloned()
.collect())
}
pub fn get_runtime(&self) -> &Arc<WasmRuntime> {
&self.runtime
}
/// Execute WASM module using WebAssembly component model based on attach_point
pub async fn execute_module_interface(
&self,
module_uuid: Uuid,
attach_point: WasmModuleAttachPoint,
input: WasmComponentInput,
) -> Result<WasmComponentOutput> {
let start_time = std::time::Instant::now();
// First, get the WASM bytes with a read lock (faster)
let wasm_bytes = {
let modules = self
.modules
.read()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
let module = modules
.get(&module_uuid)
.ok_or_else(|| WasmError::from(WasmManagerError::ModuleNotFound(module_uuid)))?;
// Clone the pre-loaded WASM bytes (already in memory, no file I/O)
module.module_meta.wasm_bytes.clone()
};
{
let mut modules = self
.modules
.write()
.map_err(|e| WasmManagerError::LockFailed(e.to_string()))?;
if let Some(module) = modules.get_mut(&module_uuid) {
// SystemTime::duration_since only fails if the system time is before UNIX_EPOCH,
// which should never happen in normal operation. If it does, use current time as fallback.
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_else(|_| {
// Fallback to a reasonable timestamp if system time is invalid
// This should never occur in practice, but provides a safe fallback
std::time::Duration::from_nanos(0)
})
.as_nanos() as u64;
module.module_meta.last_accessed_at = now;
module.module_meta.access_count += 1;
}
}
let result = self
.runtime
.execute_component_async(wasm_bytes, attach_point, input)
.await;
// Record metrics
let execution_time_ms = start_time.elapsed().as_millis() as u64;
self.total_executions.fetch_add(1, Ordering::Relaxed);
self.total_execution_time_ms
.fetch_add(execution_time_ms, Ordering::Relaxed);
// Update max execution time
self.max_execution_time_ms
.fetch_max(execution_time_ms, Ordering::Relaxed);
if result.is_ok() {
self.successful_executions.fetch_add(1, Ordering::Relaxed);
} else {
self.failed_executions.fetch_add(1, Ordering::Relaxed);
}
result
}
/// Execute WASM module using WebAssembly component model (sync version)
pub fn execute_module_interface_sync(
&self,
module_uuid: Uuid,
attach_point: WasmModuleAttachPoint,
input: WasmComponentInput,
) -> Result<WasmComponentOutput> {
let handle = tokio::runtime::Handle::current();
handle.block_on(self.execute_module_interface(module_uuid, attach_point, input))
}
/// Get current metrics
pub fn get_metrics(&self) -> (u64, u64, u64, u64, u64) {
(
self.total_executions.load(Ordering::Relaxed),
self.successful_executions.load(Ordering::Relaxed),
self.failed_executions.load(Ordering::Relaxed),
self.total_execution_time_ms.load(Ordering::Relaxed),
self.max_execution_time_ms.load(Ordering::Relaxed),
)
}
/// Execute a WASM module for a given attach point
/// Returns the Action if successful, or None if execution failed
///
/// This is a convenience method that wraps execute_module_interface and handles
/// error logging automatically.
pub async fn execute_module_for_attach_point(
&self,
module: &WasmModule,
attach_point: WasmModuleAttachPoint,
input: WasmComponentInput,
) -> Option<crate::wasm::spec::sgl::router::middleware_types::Action> {
use tracing::error;
let action_result = self
.execute_module_interface(module.module_uuid, attach_point, input)
.await;
match action_result {
Ok(output) => match output {
WasmComponentOutput::MiddlewareAction(action) => Some(action),
},
Err(e) => {
error!(
"Failed to execute WASM module {}: {}",
module.module_meta.name, e
);
None
}
}
}
}
impl Default for WasmModuleManager {
fn default() -> Self {
// with_default_config() should always succeed with default configuration.
// If it fails, it indicates a critical system configuration error.
Self::with_default_config()
.expect("Failed to create WasmModuleManager with default config. This should never happen with valid default configuration.")
}
}

View File

@@ -0,0 +1,228 @@
//! WASM HTTP API Routes
//!
//! Provides REST API endpoints for managing WASM modules:
//! - POST /wasm - Add modules
//! - DELETE /wasm/:uuid - Remove a module
//! - GET /wasm - List all modules with metrics
use std::{sync::Arc, time::Duration};
use axum::{
extract::{Json, Path, State},
http::StatusCode,
response::{IntoResponse, Response},
};
use uuid::Uuid;
use crate::{
core::{job_queue::Job, workflow::steps::WasmModuleConfigRequest},
server::AppState,
wasm::module::{
WasmMetrics, WasmModuleAddRequest, WasmModuleAddResponse, WasmModuleAddResult,
WasmModuleListResponse,
},
};
/// Wait for job completion by polling job status
/// Returns the job result message if successful
async fn wait_for_job_completion(
job_queue: &crate::core::job_queue::JobQueue,
status_key: &str,
timeout_duration: Duration,
) -> Result<String, String> {
let start = std::time::Instant::now();
let mut poll_interval = Duration::from_millis(100);
let max_poll_interval = Duration::from_millis(2000);
let poll_backoff = Duration::from_millis(200);
loop {
if start.elapsed() > timeout_duration {
return Err(format!("Job timeout after {}s", timeout_duration.as_secs()));
}
if let Some(job_status) = job_queue.get_status(status_key) {
match job_status.status.as_str() {
"pending" | "processing" => {
tokio::time::sleep(poll_interval).await;
poll_interval = (poll_interval + poll_backoff).min(max_poll_interval);
continue;
}
"failed" => {
let error_msg = job_status
.message
.unwrap_or_else(|| "Unknown error".to_string());
job_queue.remove_status(status_key);
return Err(error_msg);
}
_ => {
// Should not happen, but handle gracefully
job_queue.remove_status(status_key);
return Err("Unexpected job status".to_string());
}
}
} else {
// Job completed successfully (status was removed by record_job_completion)
// We need to get the result from the job execution
// Since job queue removes status on success, we can't get the result here
// We'll need to query the wasm manager to find the module by name
// For now, return a success message and let caller extract UUID from manager
return Ok("Job completed successfully".to_string());
}
}
}
pub async fn add_wasm_module(
State(state): State<Arc<AppState>>,
Json(config): Json<WasmModuleAddRequest>,
) -> Response {
let Some(_) = state.context.wasm_manager.as_ref() else {
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
};
let Some(job_queue) = state.context.worker_job_queue.get() else {
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
};
let mut status = StatusCode::OK;
let mut modules = config.modules.clone();
for module in modules.iter_mut() {
let wasm_config = WasmModuleConfigRequest {
descriptor: module.clone(),
};
let job = Job::AddWasmModule {
config: Box::new(wasm_config),
};
let worker_url = job.worker_url().to_string();
// Submit job to queue
match job_queue.submit(job).await {
Ok(_) => {
// Wait for job completion (timeout: 5 minutes)
let timeout = Duration::from_secs(300);
match wait_for_job_completion(job_queue, &worker_url, timeout).await {
Ok(_) => {
// Job completed successfully, but we need to get the UUID
// Since job queue removes status on success, we need to query
// the workflow engine or wasm manager to get the UUID
// For now, let's try to get it from the wasm manager by name
if let Some(wasm_manager) = state.context.wasm_manager.as_ref() {
if let Ok(all_modules) = wasm_manager.get_modules() {
if let Some(registered_module) = all_modules
.iter()
.find(|m| m.module_meta.name == module.name)
{
module.add_result = Some(WasmModuleAddResult::Success(
registered_module.module_uuid,
));
} else {
module.add_result = Some(WasmModuleAddResult::Error(
"Module registered but UUID not found".to_string(),
));
status = StatusCode::BAD_REQUEST;
}
} else {
module.add_result = Some(WasmModuleAddResult::Error(
"Failed to query registered modules".to_string(),
));
status = StatusCode::BAD_REQUEST;
}
} else {
module.add_result = Some(WasmModuleAddResult::Error(
"WASM manager not available".to_string(),
));
status = StatusCode::BAD_REQUEST;
}
}
Err(e) => {
module.add_result = Some(WasmModuleAddResult::Error(e));
status = StatusCode::BAD_REQUEST;
}
}
}
Err(e) => {
module.add_result = Some(WasmModuleAddResult::Error(format!(
"Failed to submit job: {}",
e
)));
status = StatusCode::BAD_REQUEST;
}
}
}
let response = WasmModuleAddResponse { modules };
(status, Json(response)).into_response()
}
pub async fn remove_wasm_module(
State(state): State<Arc<AppState>>,
Path(module_uuid_str): Path<String>,
) -> Response {
let Ok(module_uuid) = Uuid::parse_str(&module_uuid_str) else {
return StatusCode::BAD_REQUEST.into_response();
};
let Some(_) = state.context.wasm_manager.as_ref() else {
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
};
let Some(job_queue) = state.context.worker_job_queue.get() else {
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
};
use crate::core::workflow::steps::WasmModuleRemovalRequest;
let removal_request = WasmModuleRemovalRequest::new(module_uuid);
let job = Job::RemoveWasmModule {
request: Box::new(removal_request),
};
let worker_url = job.worker_url().to_string();
// Submit job to queue
match job_queue.submit(job).await {
Ok(_) => {
// Wait for job completion (timeout: 1 minute)
let timeout = Duration::from_secs(60);
match wait_for_job_completion(job_queue, &worker_url, timeout).await {
Ok(_) => (StatusCode::OK, "Module removed successfully").into_response(),
Err(e) => (StatusCode::BAD_REQUEST, e).into_response(),
}
}
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to submit job: {}", e),
)
.into_response(),
}
}
pub async fn list_wasm_modules(State(state): State<Arc<AppState>>) -> Response {
let Some(wasm_manager) = state.context.wasm_manager.as_ref() else {
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
};
let modules = wasm_manager.get_modules();
if let Ok(modules) = modules {
let (total, success, failed, total_time_ms, max_time_ms) = wasm_manager.get_metrics();
let average_execution_time_ms = if total > 0 {
Some(total_time_ms as f64 / total as f64)
} else {
None
};
let metrics = WasmMetrics {
total_executions: total,
successful_executions: success,
failed_executions: failed,
total_execution_time_ms: total_time_ms,
max_execution_time_ms: max_time_ms,
average_execution_time_ms,
};
let response = WasmModuleListResponse { modules, metrics };
(StatusCode::OK, Json(response)).into_response()
} else {
StatusCode::INTERNAL_SERVER_ERROR.into_response()
}
}

View File

@@ -0,0 +1,430 @@
//! WASM Runtime
//!
//! Manages WASM component execution using wasmtime with async support.
//! Provides a thread pool for concurrent WASM execution and metrics tracking.
use std::sync::{
atomic::{AtomicU64, Ordering},
Arc,
};
use tokio::sync::oneshot;
use tracing::{debug, error, info};
use wasmtime::{
component::{Component, Linker, ResourceTable},
Config, Engine, Store,
};
use wasmtime_wasi::WasiCtx;
use crate::wasm::{
config::WasmRuntimeConfig,
errors::{Result, WasmError, WasmRuntimeError},
module::{MiddlewareAttachPoint, WasmModuleAttachPoint},
spec::SglRouter,
types::{WasiState, WasmComponentInput, WasmComponentOutput},
};
pub struct WasmRuntime {
config: WasmRuntimeConfig,
thread_pool: Arc<WasmThreadPool>,
// Metrics
total_executions: AtomicU64,
successful_executions: AtomicU64,
failed_executions: AtomicU64,
total_execution_time_ms: AtomicU64,
max_execution_time_ms: AtomicU64,
}
pub struct WasmThreadPool {
sender: async_channel::Sender<WasmTask>,
receiver: async_channel::Receiver<WasmTask>,
workers: Vec<std::thread::JoinHandle<()>>,
// Metrics
total_tasks: AtomicU64,
completed_tasks: AtomicU64,
failed_tasks: AtomicU64,
}
pub enum WasmTask {
ExecuteComponent {
wasm_bytes: Vec<u8>,
attach_point: WasmModuleAttachPoint,
input: WasmComponentInput,
response: oneshot::Sender<Result<WasmComponentOutput>>,
},
}
impl WasmRuntime {
pub fn new(config: WasmRuntimeConfig) -> Result<Self> {
let thread_pool = Arc::new(WasmThreadPool::new(config.clone())?);
Ok(Self {
config,
thread_pool,
total_executions: AtomicU64::new(0),
successful_executions: AtomicU64::new(0),
failed_executions: AtomicU64::new(0),
total_execution_time_ms: AtomicU64::new(0),
max_execution_time_ms: AtomicU64::new(0),
})
}
pub fn with_default_config() -> Result<Self> {
Self::new(WasmRuntimeConfig::default())
}
pub fn get_config(&self) -> &WasmRuntimeConfig {
&self.config
}
/// get available cpu count and max recommended cpu count
pub fn get_cpu_info() -> (usize, usize) {
let cpu_count = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4);
let max_recommended = cpu_count.max(1);
(cpu_count, max_recommended)
}
/// get current thread pool status
pub fn get_thread_pool_info(&self) -> (usize, usize) {
let (_cpu_count, max_recommended) = Self::get_cpu_info();
let current_workers = self.thread_pool.workers.len();
(current_workers, max_recommended)
}
/// Execute WASM component using WASM interface based on attach_point
pub async fn execute_component_async(
&self,
wasm_bytes: Vec<u8>,
attach_point: WasmModuleAttachPoint,
input: WasmComponentInput,
) -> Result<WasmComponentOutput> {
let start_time = std::time::Instant::now();
let (response_tx, response_rx) = oneshot::channel();
let task = WasmTask::ExecuteComponent {
wasm_bytes,
attach_point,
input,
response: response_tx,
};
self.thread_pool.sender.send(task).await.map_err(|e| {
WasmRuntimeError::CallFailed(format!("Failed to send task to thread pool: {}", e))
})?;
let result = response_rx.await.map_err(|e| {
WasmRuntimeError::CallFailed(format!(
"Failed to receive response from thread pool: {}",
e
))
})?;
let execution_time_ms = start_time.elapsed().as_millis() as u64;
self.total_executions.fetch_add(1, Ordering::Relaxed);
self.total_execution_time_ms
.fetch_add(execution_time_ms, Ordering::Relaxed);
// Update max execution time
self.max_execution_time_ms
.fetch_max(execution_time_ms, Ordering::Relaxed);
if result.is_ok() {
self.successful_executions.fetch_add(1, Ordering::Relaxed);
} else {
self.failed_executions.fetch_add(1, Ordering::Relaxed);
}
result
}
/// Get current metrics
pub fn get_metrics(&self) -> (u64, u64, u64, u64, u64) {
(
self.total_executions.load(Ordering::Relaxed),
self.successful_executions.load(Ordering::Relaxed),
self.failed_executions.load(Ordering::Relaxed),
self.total_execution_time_ms.load(Ordering::Relaxed),
self.max_execution_time_ms.load(Ordering::Relaxed),
)
}
}
impl WasmThreadPool {
pub fn new(config: WasmRuntimeConfig) -> Result<Self> {
let (sender, receiver) = async_channel::unbounded();
let mut workers = Vec::new();
// set thread pool size based on cpu count
let max_workers = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4)
.max(1);
let num_workers = config.thread_pool_size.clamp(1, max_workers);
info!(
target: "sglang_router_rs::wasm::runtime",
"Initializing WASM runtime with {} workers",
num_workers
);
for worker_id in 0..num_workers {
let receiver = receiver.clone();
let config = config.clone();
let worker = std::thread::spawn(move || {
// create independent tokio runtime for this thread
let rt = match tokio::runtime::Runtime::new() {
Ok(rt) => rt,
Err(e) => {
error!(
target: "sglang_router_rs::wasm::runtime",
worker_id = worker_id,
"Failed to create tokio runtime: {}",
e
);
return;
}
};
rt.block_on(async {
Self::worker_loop(worker_id, receiver, config).await;
});
});
workers.push(worker);
}
Ok(Self {
sender,
receiver,
workers,
total_tasks: AtomicU64::new(0),
completed_tasks: AtomicU64::new(0),
failed_tasks: AtomicU64::new(0),
})
}
/// Get current thread pool metrics
pub fn get_metrics(&self) -> (u64, u64, u64) {
(
self.total_tasks.load(Ordering::Relaxed),
self.completed_tasks.load(Ordering::Relaxed),
self.failed_tasks.load(Ordering::Relaxed),
)
}
async fn worker_loop(
worker_id: usize,
receiver: async_channel::Receiver<WasmTask>,
config: WasmRuntimeConfig,
) {
debug!(
target: "sglang_router_rs::wasm::runtime",
worker_id = worker_id,
thread_id = ?std::thread::current().id(),
"Worker started"
);
let mut wasmtime_config = Config::new();
wasmtime_config.async_stack_size(config.max_stack_size);
wasmtime_config.async_support(true);
wasmtime_config.wasm_component_model(true); // Enable component model
let engine = match Engine::new(&wasmtime_config) {
Ok(engine) => engine,
Err(e) => {
error!(
target: "sglang_router_rs::wasm::runtime",
worker_id = worker_id,
"Failed to create engine: {}",
e
);
return;
}
};
loop {
let task = match receiver.recv().await {
Ok(task) => task,
Err(_) => {
debug!(
target: "sglang_router_rs::wasm::runtime",
worker_id = worker_id,
"Worker shutting down"
);
break; // channel closed, exit loop
}
};
match task {
WasmTask::ExecuteComponent {
wasm_bytes,
attach_point,
input,
response,
} => {
let result = Self::execute_component_in_worker(
&engine,
wasm_bytes,
attach_point,
input,
&config,
)
.await;
let _ = response.send(result);
}
}
}
}
async fn execute_component_in_worker(
engine: &Engine,
wasm_bytes: Vec<u8>,
attach_point: WasmModuleAttachPoint,
input: WasmComponentInput,
_config: &WasmRuntimeConfig,
) -> Result<WasmComponentOutput> {
// Compile component from bytes
// Note: The WASM file must be in component format (not plain WASM module)
// Use `wasm-tools component new` to wrap a WASM module into a component if needed
let component = Component::new(engine, &wasm_bytes).map_err(|e| {
WasmRuntimeError::CompileFailed(format!(
"failed to parse WebAssembly component: {}. \
Hint: The WASM file must be in component format. \
If you're using wit-bindgen, use 'wasm-tools component new' to wrap the WASM module into a component.",
e
))
})?;
let mut linker = Linker::<WasiState>::new(engine);
wasmtime_wasi::p2::add_to_linker_async(&mut linker)?;
let mut builder = WasiCtx::builder();
let mut store = Store::new(
engine,
WasiState {
ctx: builder.build(),
table: ResourceTable::new(),
},
);
let output = match attach_point {
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnRequest) => {
let request = match input {
WasmComponentInput::MiddlewareRequest(req) => req,
_ => {
return Err(WasmError::from(WasmRuntimeError::CallFailed(
"Expected MiddlewareRequest input for OnRequest attach point"
.to_string(),
)));
}
};
// Instantiate component (must use async instantiation when async support is enabled)
let bindings = SglRouter::instantiate_async(&mut store, &component, &linker)
.await
.map_err(|e| {
WasmError::from(WasmRuntimeError::InstanceCreateFailed(e.to_string()))
})?;
// Call on-request (async call when async support is enabled)
let action_result = bindings
.sgl_router_middleware_on_request()
.call_on_request(&mut store, &request)
.await
.map_err(|e| WasmError::from(WasmRuntimeError::CallFailed(e.to_string())))?;
WasmComponentOutput::MiddlewareAction(action_result)
}
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnResponse) => {
// Extract Response input
let response = match input {
WasmComponentInput::MiddlewareResponse(resp) => resp,
_ => {
return Err(WasmError::from(WasmRuntimeError::CallFailed(
"Expected MiddlewareResponse input for OnResponse attach point"
.to_string(),
)));
}
};
// Instantiate component (must use async instantiation when async support is enabled)
let bindings = SglRouter::instantiate_async(&mut store, &component, &linker)
.await
.map_err(|e| {
WasmError::from(WasmRuntimeError::InstanceCreateFailed(e.to_string()))
})?;
// Call on-response (async call when async support is enabled)
let action_result = bindings
.sgl_router_middleware_on_response()
.call_on_response(&mut store, &response)
.await
.map_err(|e| WasmError::from(WasmRuntimeError::CallFailed(e.to_string())))?;
WasmComponentOutput::MiddlewareAction(action_result)
}
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnError) => {
return Err(WasmError::from(WasmRuntimeError::CallFailed(
"OnError attach point not yet implemented".to_string(),
)));
}
};
Ok(output)
}
}
impl Drop for WasmThreadPool {
fn drop(&mut self) {
// close sender and receiver
self.sender.close();
self.receiver.close();
// wait for all workers to complete
for worker in self.workers.drain(..) {
let _ = worker.join();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::wasm::config::WasmRuntimeConfig;
#[test]
fn test_get_cpu_info() {
let (cpu_count, max_recommended) = WasmRuntime::get_cpu_info();
assert!(cpu_count > 0);
assert!(max_recommended > 0);
assert!(max_recommended >= cpu_count);
}
#[test]
fn test_config_default_values() {
let config = WasmRuntimeConfig::default();
assert_eq!(config.max_memory_pages, 1024);
assert_eq!(config.max_execution_time_ms, 1000);
assert_eq!(config.max_stack_size, 1024 * 1024);
assert!(config.thread_pool_size > 0);
assert_eq!(config.module_cache_size, 10);
}
#[test]
fn test_config_clone() {
let config = WasmRuntimeConfig::default();
let cloned_config = config.clone();
assert_eq!(config.max_memory_pages, cloned_config.max_memory_pages);
assert_eq!(
config.max_execution_time_ms,
cloned_config.max_execution_time_ms
);
assert_eq!(config.max_stack_size, cloned_config.max_stack_size);
assert_eq!(config.thread_pool_size, cloned_config.thread_pool_size);
assert_eq!(config.module_cache_size, cloned_config.module_cache_size);
}
}

View File

@@ -0,0 +1,60 @@
//! WebAssembly Interface Bindings and Type Conversions
//!
//! Contains wasmtime component bindings generated from interface definitions,
//! and helper functions to convert between Axum HTTP types and interface types.
use axum::http::{header, HeaderMap, HeaderValue};
wasmtime::component::bindgen!({
path: "src/wasm/interface",
world: "sgl-router",
imports: { default: async | trappable },
exports: { default: async },
});
/// Build WebAssembly headers from Axum HeaderMap
pub fn build_wasm_headers_from_axum_headers(
headers: &HeaderMap,
) -> Vec<sgl::router::middleware_types::Header> {
let mut wasm_headers = Vec::new();
for (name, value) in headers.iter() {
if let Ok(value_str) = value.to_str() {
wasm_headers.push(sgl::router::middleware_types::Header {
name: name.as_str().to_string(),
value: value_str.to_string(),
});
}
}
wasm_headers
}
/// Apply ModifyAction header modifications to Axum HeaderMap
pub fn apply_modify_action_to_headers(
headers: &mut HeaderMap,
modify: &sgl::router::middleware_types::ModifyAction,
) {
// Apply headers_set
for header_mod in &modify.headers_set {
if let (Ok(name), Ok(value)) = (
header_mod.name.parse::<header::HeaderName>(),
header_mod.value.parse::<HeaderValue>(),
) {
headers.insert(name, value);
}
}
// Apply headers_add
for header_mod in &modify.headers_add {
if let (Ok(name), Ok(value)) = (
header_mod.name.parse::<header::HeaderName>(),
header_mod.value.parse::<HeaderValue>(),
) {
headers.append(name, value);
}
}
// Apply headers_remove
for name_str in &modify.headers_remove {
if let Ok(name) = name_str.parse::<header::HeaderName>() {
headers.remove(name);
}
}
}

View File

@@ -0,0 +1,101 @@
//! WASM Component Type System
//!
//! Provides generic input/output types for WASM component execution
//! based on attach points.
use wasmtime::component::ResourceTable;
use wasmtime_wasi::{WasiCtx, WasiCtxView, WasiView};
use crate::wasm::{
module::{MiddlewareAttachPoint, WasmModuleAttachPoint},
spec::sgl::router::middleware_types,
};
/// Generic input type for WASM component execution
///
/// This enum represents all possible input types that can be passed
/// to a WASM component, determined by the attach_point.
#[derive(Debug, Clone)]
pub enum WasmComponentInput {
/// Middleware OnRequest input
MiddlewareRequest(middleware_types::Request),
/// Middleware OnResponse input
MiddlewareResponse(middleware_types::Response),
}
/// Generic output type from WASM component execution
///
/// This enum represents all possible output types that can be returned
/// from a WASM component, determined by the attach_point.
#[derive(Debug, Clone)]
pub enum WasmComponentOutput {
/// Middleware Action output
MiddlewareAction(middleware_types::Action),
}
impl WasmComponentInput {
/// Create input based on attach_point and raw data
///
/// This helper function validates that the attach_point matches
/// the expected input type.
pub fn from_attach_point(attach_point: &WasmModuleAttachPoint) -> Result<Self, String> {
match attach_point {
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnRequest) => {
// OnRequest expects a Request type, but we can't construct it here
// The caller should use MiddlewareRequest variant directly
Err("OnRequest requires MiddlewareRequest input. Use WasmComponentInput::MiddlewareRequest directly.".to_string())
}
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnResponse) => {
// OnResponse expects a Response type
Err("OnResponse requires MiddlewareResponse input. Use WasmComponentInput::MiddlewareResponse directly.".to_string())
}
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnError) => {
Err("OnError attach point not yet implemented".to_string())
}
}
}
/// Get the expected attach_point for this input type
pub fn expected_attach_point(&self) -> WasmModuleAttachPoint {
match self {
WasmComponentInput::MiddlewareRequest(_) => {
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnRequest)
}
WasmComponentInput::MiddlewareResponse(_) => {
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnResponse)
}
}
}
}
impl WasmComponentOutput {
/// Get the attach_point that produced this output type
pub fn from_attach_point(attach_point: &WasmModuleAttachPoint) -> Result<Self, String> {
match attach_point {
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnRequest) => {
// This would be set after execution
Err("Cannot create output before execution".to_string())
}
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnResponse) => {
Err("Cannot create output before execution".to_string())
}
WasmModuleAttachPoint::Middleware(MiddlewareAttachPoint::OnError) => {
Err("OnError attach point not yet implemented".to_string())
}
}
}
}
pub struct WasiState {
pub ctx: WasiCtx,
pub table: ResourceTable,
}
impl WasiView for WasiState {
fn ctx(&mut self) -> WasiCtxView<'_> {
WasiCtxView {
ctx: &mut self.ctx,
table: &mut self.table,
}
}
}

View File

@@ -0,0 +1,819 @@
//! WASM Module Integration Tests
//!
//! This test suite validates the complete WASM module management functionality:
//! - API endpoints (add, remove, list)
//! - Workflow integration
//! - Module execution
//! - Error handling
mod common;
use std::{sync::Arc, time::Duration};
use axum::{
body::{to_bytes, Body},
extract::Request,
http::{header::CONTENT_TYPE, StatusCode},
};
use sgl_model_gateway::{
app_context::AppContext,
config::RouterConfig,
core::workflow::{
create_wasm_module_registration_workflow, create_wasm_module_removal_workflow,
},
routers::RouterFactory,
server::{build_app, AppState},
wasm::{
module::{
WasmModuleAddRequest, WasmModuleAddResponse, WasmModuleAttachPoint,
WasmModuleDescriptor, WasmModuleListResponse, WasmModuleType,
},
module_manager::WasmModuleManager,
},
};
use tempfile::TempDir;
use tokio::fs;
use tower::ServiceExt;
use uuid::Uuid;
/// Create a test AppContext with WASM manager initialized
async fn create_test_context_with_wasm() -> Arc<AppContext> {
let config = RouterConfig::default();
// Initialize WASM manager first
let wasm_manager = Arc::new(
WasmModuleManager::with_default_config().expect("Failed to create WASM module manager"),
);
// Create AppContext with wasm_manager from the start
let client = reqwest::Client::new();
// Initialize registries
use sgl_model_gateway::{
core::{LoadMonitor, WorkerRegistry},
data_connector::{
MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage,
},
policies::PolicyRegistry,
};
let worker_registry = Arc::new(WorkerRegistry::new());
let policy_registry = Arc::new(PolicyRegistry::new(config.policy.clone()));
// Initialize storage backends
let response_storage = Arc::new(MemoryResponseStorage::new());
let conversation_storage = Arc::new(MemoryConversationStorage::new());
let conversation_item_storage = Arc::new(MemoryConversationItemStorage::new());
// Initialize load monitor
let load_monitor = Some(Arc::new(LoadMonitor::new(
worker_registry.clone(),
policy_registry.clone(),
client.clone(),
config.worker_startup_check_interval_secs,
)));
// Create empty OnceLock for worker job queue, workflow engine, and mcp manager
use std::sync::OnceLock;
let worker_job_queue = Arc::new(OnceLock::new());
let workflow_engine = Arc::new(OnceLock::new());
let mcp_manager_lock = Arc::new(OnceLock::new());
let app_context = Arc::new(
AppContext::builder()
.router_config(config.clone())
.client(client)
.rate_limiter(None)
.tokenizer(None)
.reasoning_parser_factory(None)
.tool_parser_factory(None)
.worker_registry(worker_registry)
.policy_registry(policy_registry)
.response_storage(response_storage)
.conversation_storage(conversation_storage)
.conversation_item_storage(conversation_item_storage)
.load_monitor(load_monitor)
.worker_job_queue(worker_job_queue)
.workflow_engine(workflow_engine)
.mcp_manager(mcp_manager_lock)
.wasm_manager(Some(wasm_manager))
.build()
.expect("Failed to build AppContext with WASM manager"),
);
// Initialize JobQueue after AppContext is created
let weak_context = Arc::downgrade(&app_context);
let job_queue = sgl_model_gateway::core::JobQueue::new(
sgl_model_gateway::core::JobQueueConfig::default(),
weak_context,
);
app_context
.worker_job_queue
.set(job_queue)
.expect("JobQueue should only be initialized once");
// Initialize WorkflowEngine and register workflows
use sgl_model_gateway::core::workflow::{
create_worker_registration_workflow, create_worker_removal_workflow, WorkflowEngine,
};
let engine = Arc::new(WorkflowEngine::new());
engine.register_workflow(create_worker_registration_workflow(&config));
engine.register_workflow(create_worker_removal_workflow());
engine.register_workflow(create_wasm_module_registration_workflow());
engine.register_workflow(create_wasm_module_removal_workflow());
app_context
.workflow_engine
.set(engine)
.expect("WorkflowEngine should only be initialized once");
// Initialize MCP manager with empty config
use sgl_model_gateway::mcp::{McpConfig, McpManager};
let empty_config = McpConfig {
servers: vec![],
pool: Default::default(),
proxy: None,
warmup: vec![],
inventory: Default::default(),
};
let mcp_manager = McpManager::with_defaults(empty_config)
.await
.expect("Failed to create MCP manager");
app_context
.mcp_manager
.set(Arc::new(mcp_manager))
.ok()
.expect("McpManager should only be initialized once");
app_context
}
/// Create a test WASM component file
/// Dynamically generates a valid WASM component programmatically without external tools
/// This ensures tests work in new environments without requiring pre-built files or external tools
async fn create_test_wasm_component(temp_dir: &TempDir) -> String {
use wasm_encoder::{Component, Module};
// Create a minimal valid WASM module first
// A minimal module needs at least a type section
let mut module = Module::new();
// Add an empty type section (0 types) - this is valid
let type_section = wasm_encoder::TypeSection::new();
module.section(&type_section);
let mut component = Component::new();
component.section(&wasm_encoder::ModuleSection(&module));
let component_bytes = component.as_slice().to_vec();
let component_path = temp_dir.path().join("test_module.component.wasm");
fs::write(&component_path, component_bytes)
.await
.expect("Failed to write WASM component file");
// Return absolute path
component_path
.canonicalize()
.expect("Failed to canonicalize path")
.to_str()
.unwrap()
.to_string()
}
/// Create a test app with WASM support
async fn create_test_app_with_wasm() -> (axum::Router, Arc<AppContext>, TempDir) {
let temp_dir = TempDir::new().expect("Failed to create temp directory");
let app_context = create_test_context_with_wasm().await;
// Create a dummy router (we only need the app for WASM endpoints)
let router = RouterFactory::create_router(&app_context)
.await
.expect("Failed to create router");
let router = Arc::from(router);
let app_state = Arc::new(AppState {
router,
context: app_context.clone(),
concurrency_queue_tx: None,
router_manager: None,
});
let request_id_headers = vec!["x-request-id".to_string(), "x-correlation-id".to_string()];
let app = build_app(
app_state,
sgl_model_gateway::middleware::AuthConfig { api_key: None },
256 * 1024 * 1024,
request_id_headers,
vec![], // cors_allowed_origins
);
(app, app_context, temp_dir)
}
// ============================================================================
// API Endpoint Tests
// ============================================================================
#[tokio::test]
async fn test_wasm_api_add_module() {
let (app, app_context, temp_dir) = create_test_app_with_wasm().await;
let wasm_file_path = create_test_wasm_component(&temp_dir).await;
let add_request = WasmModuleAddRequest {
modules: vec![WasmModuleDescriptor {
name: "test_module".to_string(),
file_path: wasm_file_path.clone(),
module_type: WasmModuleType::Middleware,
attach_points: vec![WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
)],
add_result: None,
}],
};
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/wasm")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&add_request).unwrap()))
.unwrap(),
)
.await
.unwrap();
let status = response.status();
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
let response_json: WasmModuleAddResponse = serde_json::from_slice(&body).unwrap();
assert_eq!(response_json.modules.len(), 1);
let module_result = &response_json.modules[0].add_result;
// Print error for debugging
if let Some(sgl_model_gateway::wasm::module::WasmModuleAddResult::Error(err)) = module_result {
eprintln!("Module registration failed: {}", err);
}
// If status is not OK, check the error message
if status != StatusCode::OK {
eprintln!("Response status: {:?}", status);
eprintln!("Response body: {}", String::from_utf8_lossy(&body));
panic!(
"Expected OK status but got {:?}. Error: {:?}",
status, module_result
);
}
assert!(module_result.is_some());
// Verify module is registered in wasm_manager
if let Some(wasm_manager) = app_context.wasm_manager.as_ref() {
let modules = wasm_manager.get_modules().expect("Failed to get modules");
assert!(!modules.is_empty(), "Module should be registered");
if let Some(sgl_model_gateway::wasm::module::WasmModuleAddResult::Success(uuid)) =
module_result
{
let module = wasm_manager
.get_module(*uuid)
.expect("Failed to get module");
assert!(module.is_some(), "Module should exist in manager");
}
}
}
#[tokio::test]
async fn test_wasm_api_add_module_invalid_file() {
let (app, _app_context, _temp_dir) = create_test_app_with_wasm().await;
let add_request = WasmModuleAddRequest {
modules: vec![WasmModuleDescriptor {
name: "test_module".to_string(),
file_path: "/nonexistent/path/to/module.component.wasm".to_string(),
module_type: WasmModuleType::Middleware,
attach_points: vec![WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
)],
add_result: None,
}],
};
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/wasm")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&add_request).unwrap()))
.unwrap(),
)
.await
.unwrap();
// Should return error status
assert!(
response.status() == StatusCode::BAD_REQUEST
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
);
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
let response_json: WasmModuleAddResponse = serde_json::from_slice(&body).unwrap();
assert_eq!(response_json.modules.len(), 1);
let module_result = &response_json.modules[0].add_result;
assert!(module_result.is_some());
// Verify it's an error result
if let Some(sgl_model_gateway::wasm::module::WasmModuleAddResult::Error(_)) = module_result {
// Expected error
} else {
panic!("Expected error result for invalid file path");
}
}
#[tokio::test]
async fn test_wasm_api_add_module_invalid_wasm() {
let (app, _app_context, temp_dir) = create_test_app_with_wasm().await;
// Create an invalid WASM file (just random bytes)
let invalid_wasm_path = temp_dir.path().join("invalid.component.wasm");
fs::write(&invalid_wasm_path, b"not a valid wasm file")
.await
.expect("Failed to write invalid WASM file");
let add_request = WasmModuleAddRequest {
modules: vec![WasmModuleDescriptor {
name: "invalid_module".to_string(),
file_path: invalid_wasm_path.to_str().unwrap().to_string(),
module_type: WasmModuleType::Middleware,
attach_points: vec![WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
)],
add_result: None,
}],
};
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/wasm")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&add_request).unwrap()))
.unwrap(),
)
.await
.unwrap();
// Should return error status
assert!(
response.status() == StatusCode::BAD_REQUEST
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
);
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
let response_json: WasmModuleAddResponse = serde_json::from_slice(&body).unwrap();
assert_eq!(response_json.modules.len(), 1);
let module_result = &response_json.modules[0].add_result;
assert!(module_result.is_some());
// Verify it's an error result
if let Some(sgl_model_gateway::wasm::module::WasmModuleAddResult::Error(_)) = module_result {
// Expected error
} else {
panic!("Expected error result for invalid WASM file");
}
}
#[tokio::test]
async fn test_wasm_api_list_modules() {
let (app, _app_context, temp_dir) = create_test_app_with_wasm().await;
let wasm_file_path = create_test_wasm_component(&temp_dir).await;
// First, add a module
let add_request = WasmModuleAddRequest {
modules: vec![WasmModuleDescriptor {
name: "test_module_list".to_string(),
file_path: wasm_file_path.clone(),
module_type: WasmModuleType::Middleware,
attach_points: vec![WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
)],
add_result: None,
}],
};
let add_response = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/wasm")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&add_request).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(add_response.status(), StatusCode::OK);
// Wait a bit for the job to complete
tokio::time::sleep(Duration::from_millis(500)).await;
// Now list modules
let list_response = app
.oneshot(
Request::builder()
.method("GET")
.uri("/wasm")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(list_response.status(), StatusCode::OK);
let body = to_bytes(list_response.into_body(), usize::MAX)
.await
.unwrap();
let response_json: WasmModuleListResponse = serde_json::from_slice(&body).unwrap();
assert!(
!response_json.modules.is_empty(),
"Should have at least one module"
);
assert!(response_json
.modules
.iter()
.any(|m| m.module_meta.name == "test_module_list"));
// Verify metrics are present (total_executions is u64, so always >= 0)
let _ = response_json.metrics.total_executions;
}
#[tokio::test]
async fn test_wasm_api_remove_module() {
let (app, app_context, temp_dir) = create_test_app_with_wasm().await;
let wasm_file_path = create_test_wasm_component(&temp_dir).await;
// First, add a module
let add_request = WasmModuleAddRequest {
modules: vec![WasmModuleDescriptor {
name: "test_module_remove".to_string(),
file_path: wasm_file_path.clone(),
module_type: WasmModuleType::Middleware,
attach_points: vec![WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
)],
add_result: None,
}],
};
let add_response = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/wasm")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&add_request).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(add_response.status(), StatusCode::OK);
let body = to_bytes(add_response.into_body(), usize::MAX)
.await
.unwrap();
let response_json: WasmModuleAddResponse = serde_json::from_slice(&body).unwrap();
// Wait for job to complete
tokio::time::sleep(Duration::from_millis(500)).await;
// Get the module UUID
let module_uuid =
if let Some(sgl_model_gateway::wasm::module::WasmModuleAddResult::Success(uuid)) =
&response_json.modules[0].add_result
{
*uuid
} else {
// If we can't get UUID from response, try to find it from manager
if let Some(wasm_manager) = app_context.wasm_manager.as_ref() {
let modules = wasm_manager.get_modules().expect("Failed to get modules");
modules
.iter()
.find(|m| m.module_meta.name == "test_module_remove")
.map(|m| m.module_uuid)
.expect("Module should be registered")
} else {
panic!("WASM manager not available");
}
};
// Now remove the module
let remove_response = app
.oneshot(
Request::builder()
.method("DELETE")
.uri(format!("/wasm/{}", module_uuid))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let remove_status = remove_response.status();
if remove_status != StatusCode::OK {
let body = to_bytes(remove_response.into_body(), usize::MAX)
.await
.unwrap();
eprintln!(
"Remove module failed with status {:?}: {}",
remove_status,
String::from_utf8_lossy(&body)
);
panic!("Expected OK status but got {:?}", remove_status);
}
// Wait for removal to complete
tokio::time::sleep(Duration::from_millis(500)).await;
// Verify module is removed
if let Some(wasm_manager) = app_context.wasm_manager.as_ref() {
let module = wasm_manager
.get_module(module_uuid)
.expect("Failed to get module");
assert!(module.is_none(), "Module should be removed");
}
}
#[tokio::test]
async fn test_wasm_api_remove_module_not_found() {
let (app, _app_context, _temp_dir) = create_test_app_with_wasm().await;
// Try to remove a non-existent module
let fake_uuid = Uuid::new_v4();
let remove_response = app
.oneshot(
Request::builder()
.method("DELETE")
.uri(format!("/wasm/{}", fake_uuid))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
// Should return error status
assert!(
remove_response.status() == StatusCode::BAD_REQUEST
|| remove_response.status() == StatusCode::NOT_FOUND
);
}
// ============================================================================
// WASM Functionality Tests
// ============================================================================
#[tokio::test]
async fn test_wasm_module_duplicate_sha256() {
let (app, _app_context, temp_dir) = create_test_app_with_wasm().await;
let wasm_file_path = create_test_wasm_component(&temp_dir).await;
// Add first module
let add_request1 = WasmModuleAddRequest {
modules: vec![WasmModuleDescriptor {
name: "test_module_dup1".to_string(),
file_path: wasm_file_path.clone(),
module_type: WasmModuleType::Middleware,
attach_points: vec![WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
)],
add_result: None,
}],
};
let response1 = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/wasm")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&add_request1).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response1.status(), StatusCode::OK);
// Wait for first job to complete
tokio::time::sleep(Duration::from_millis(500)).await;
// Try to add the same file again (should fail due to duplicate SHA256)
let add_request2 = WasmModuleAddRequest {
modules: vec![WasmModuleDescriptor {
name: "test_module_dup2".to_string(),
file_path: wasm_file_path.clone(), // Same file
module_type: WasmModuleType::Middleware,
attach_points: vec![WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
)],
add_result: None,
}],
};
let response2 = app
.oneshot(
Request::builder()
.method("POST")
.uri("/wasm")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&add_request2).unwrap()))
.unwrap(),
)
.await
.unwrap();
// Should return error status for duplicate
assert!(
response2.status() == StatusCode::BAD_REQUEST
|| response2.status() == StatusCode::INTERNAL_SERVER_ERROR
);
let body = to_bytes(response2.into_body(), usize::MAX).await.unwrap();
let response_json: WasmModuleAddResponse = serde_json::from_slice(&body).unwrap();
assert_eq!(response_json.modules.len(), 1);
let module_result = &response_json.modules[0].add_result;
assert!(module_result.is_some());
// Verify it's an error result (duplicate)
if let Some(sgl_model_gateway::wasm::module::WasmModuleAddResult::Error(err_msg)) =
module_result
{
assert!(
err_msg.contains("duplicate")
|| err_msg.contains("Duplicate")
|| err_msg.contains("SHA256")
);
} else {
panic!("Expected error result for duplicate SHA256");
}
}
#[tokio::test]
async fn test_wasm_module_execution() {
let (_app, app_context, temp_dir) = create_test_app_with_wasm().await;
let wasm_file_path = create_test_wasm_component(&temp_dir).await;
// First, add a module using the workflow directly
let wasm_manager = app_context
.wasm_manager
.as_ref()
.expect("WASM manager should be initialized");
let engine = app_context
.workflow_engine
.get()
.expect("Workflow engine should be initialized");
// Create workflow context for registration
use sgl_model_gateway::core::workflow::{
steps::WasmModuleConfigRequest, WorkflowContext, WorkflowId, WorkflowInstanceId,
};
let descriptor = WasmModuleDescriptor {
name: "test_execution_module".to_string(),
file_path: wasm_file_path.clone(),
module_type: WasmModuleType::Middleware,
attach_points: vec![WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
)],
add_result: None,
};
let config_request = WasmModuleConfigRequest { descriptor };
let mut workflow_context = WorkflowContext::new(WorkflowInstanceId::new());
workflow_context.set_arc("wasm_module_config", Arc::new(config_request));
workflow_context.set_arc("app_context", app_context.clone());
// Start workflow
let instance_id = engine
.start_workflow(
WorkflowId::new("wasm_module_registration"),
workflow_context,
)
.await
.expect("Failed to start workflow");
// Wait for workflow to complete
let timeout = Duration::from_secs(30);
let start = std::time::Instant::now();
let mut module_uuid: Option<Uuid> = None;
loop {
if start.elapsed() > timeout {
panic!("Workflow timeout");
}
let state = engine
.get_status(instance_id)
.expect("Failed to get workflow status");
match state.status {
sgl_model_gateway::core::workflow::WorkflowStatus::Completed => {
// Extract module UUID from context
if let Some(uuid_arc) = state.context.get::<Uuid>("module_uuid") {
module_uuid = Some(*uuid_arc.as_ref());
}
break;
}
sgl_model_gateway::core::workflow::WorkflowStatus::Failed => {
panic!("Workflow failed: {:?}", state);
}
_ => {
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
}
let module_uuid = module_uuid.expect("Module UUID should be in context");
// Verify module is registered
let module = wasm_manager
.get_module(module_uuid)
.expect("Failed to get module");
assert!(module.is_some(), "Module should be registered");
// Get initial metrics
let (initial_total, initial_success, initial_failed, _, _) = wasm_manager.get_metrics();
// Execute the module
use sgl_model_gateway::wasm::{
spec::sgl::router::middleware_types,
types::{WasmComponentInput, WasmComponentOutput},
};
let request = middleware_types::Request {
method: "GET".to_string(),
path: "/test".to_string(),
query: "".to_string(),
headers: vec![],
body: vec![],
request_id: "test-request-id".to_string(),
now_epoch_ms: 1000,
};
let input = WasmComponentInput::MiddlewareRequest(request);
let attach_point = WasmModuleAttachPoint::Middleware(
sgl_model_gateway::wasm::module::MiddlewareAttachPoint::OnRequest,
);
// Execute the module
let result = wasm_manager
.execute_module_interface(module_uuid, attach_point, input)
.await;
// Verify execution result
match result {
Ok(WasmComponentOutput::MiddlewareAction(action)) => {
// Verify action is valid (should be Continue, Reject, or Modify)
match action {
middleware_types::Action::Continue => {
// Expected for a simple middleware
}
middleware_types::Action::Reject(_) => {
// Also valid
}
middleware_types::Action::Modify(_) => {
// Also valid
}
}
}
Err(e) => {
// Execution might fail if the WASM component is not properly built
// This is acceptable for testing - we're testing the execution path, not the component itself
eprintln!(
"Module execution failed (expected if component is not properly built): {:?}",
e
);
}
}
// Verify metrics were updated
let (final_total, final_success, final_failed, _, _) = wasm_manager.get_metrics();
// Metrics should have increased (either success or failed)
assert!(
final_total > initial_total
|| final_failed > initial_failed
|| final_success > initial_success,
"Metrics should be updated after execution"
);
}