[model-gateway] Add WASM support for middleware (#12471)
Signed-off-by: Tony Lu <tonylu@linux.alibaba.com>
This commit is contained in:
2
.github/CODEOWNERS
vendored
2
.github/CODEOWNERS
vendored
@@ -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
|
||||
|
||||
@@ -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
27
sgl-router/examples/wasm/.gitignore
vendored
Normal 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
|
||||
102
sgl-router/examples/wasm/README.md
Normal file
102
sgl-router/examples/wasm/README.md
Normal 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.
|
||||
10
sgl-router/examples/wasm/wasm-guest-auth/Cargo.toml
Normal file
10
sgl-router/examples/wasm/wasm-guest-auth/Cargo.toml
Normal 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"] }
|
||||
62
sgl-router/examples/wasm/wasm-guest-auth/README.md
Normal file
62
sgl-router/examples/wasm/wasm-guest-auth/README.md
Normal 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
|
||||
77
sgl-router/examples/wasm/wasm-guest-auth/build.sh
Executable file
77
sgl-router/examples/wasm/wasm-guest-auth/build.sh
Executable 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
|
||||
70
sgl-router/examples/wasm/wasm-guest-auth/src/lib.rs
Normal file
70
sgl-router/examples/wasm/wasm-guest-auth/src/lib.rs
Normal 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);
|
||||
10
sgl-router/examples/wasm/wasm-guest-logging/Cargo.toml
Normal file
10
sgl-router/examples/wasm/wasm-guest-logging/Cargo.toml
Normal 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"] }
|
||||
53
sgl-router/examples/wasm/wasm-guest-logging/README.md
Normal file
53
sgl-router/examples/wasm/wasm-guest-logging/README.md
Normal 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
|
||||
77
sgl-router/examples/wasm/wasm-guest-logging/build.sh
Executable file
77
sgl-router/examples/wasm/wasm-guest-logging/build.sh
Executable 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
|
||||
88
sgl-router/examples/wasm/wasm-guest-logging/src/lib.rs
Normal file
88
sgl-router/examples/wasm/wasm-guest-logging/src/lib.rs
Normal 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);
|
||||
10
sgl-router/examples/wasm/wasm-guest-ratelimit/Cargo.toml
Normal file
10
sgl-router/examples/wasm/wasm-guest-ratelimit/Cargo.toml
Normal 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"] }
|
||||
68
sgl-router/examples/wasm/wasm-guest-ratelimit/README.md
Normal file
68
sgl-router/examples/wasm/wasm-guest-ratelimit/README.md
Normal 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
|
||||
77
sgl-router/examples/wasm/wasm-guest-ratelimit/build.sh
Executable file
77
sgl-router/examples/wasm/wasm-guest-ratelimit/build.sh
Executable 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
|
||||
155
sgl-router/examples/wasm/wasm-guest-ratelimit/src/lib.rs
Normal file
155
sgl-router/examples/wasm/wasm-guest-ratelimit/src/lib.rs
Normal 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);
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
@@ -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,
|
||||
|
||||
474
sgl-router/src/core/workflow/steps/wasm_module_registration.rs
Normal file
474
sgl-router/src/core/workflow/steps/wasm_module_registration.rs
Normal 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),
|
||||
)
|
||||
}
|
||||
160
sgl-router/src/core/workflow/steps/wasm_module_removal.rs
Normal file
160
sgl-router/src/core/workflow/steps/wasm_module_removal.rs
Normal 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),
|
||||
)
|
||||
}
|
||||
@@ -18,3 +18,4 @@ pub mod service_discovery;
|
||||
pub mod tokenizer;
|
||||
pub mod tool_parser;
|
||||
pub mod version;
|
||||
pub mod wasm;
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
227
sgl-router/src/wasm/README.md
Normal file
227
sgl-router/src/wasm/README.md
Normal 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);
|
||||
```
|
||||
297
sgl-router/src/wasm/config.rs
Normal file
297
sgl-router/src/wasm/config.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
117
sgl-router/src/wasm/errors.rs
Normal file
117
sgl-router/src/wasm/errors.rs
Normal 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()))
|
||||
}
|
||||
}
|
||||
54
sgl-router/src/wasm/interface/spec.wit
Normal file
54
sgl-router/src/wasm/interface/spec.wit
Normal 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;
|
||||
}
|
||||
13
sgl-router/src/wasm/mod.rs
Normal file
13
sgl-router/src/wasm/mod.rs
Normal 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;
|
||||
209
sgl-router/src/wasm/module.rs
Normal file
209
sgl-router/src/wasm/module.rs
Normal 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(×tamp_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>,
|
||||
}
|
||||
259
sgl-router/src/wasm/module_manager.rs
Normal file
259
sgl-router/src/wasm/module_manager.rs
Normal 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.")
|
||||
}
|
||||
}
|
||||
228
sgl-router/src/wasm/route.rs
Normal file
228
sgl-router/src/wasm/route.rs
Normal 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()
|
||||
}
|
||||
}
|
||||
430
sgl-router/src/wasm/runtime.rs
Normal file
430
sgl-router/src/wasm/runtime.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
60
sgl-router/src/wasm/spec.rs
Normal file
60
sgl-router/src/wasm/spec.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
}
|
||||
101
sgl-router/src/wasm/types.rs
Normal file
101
sgl-router/src/wasm/types.rs
Normal 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,
|
||||
}
|
||||
}
|
||||
}
|
||||
819
sgl-router/tests/wasm_test.rs
Normal file
819
sgl-router/tests/wasm_test.rs
Normal 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"
|
||||
);
|
||||
}
|
||||
Reference in New Issue
Block a user