diff --git a/.buildkite/pipeline.yml b/.buildkite/pipeline.yml index 1cf6c10e..6fcaf5b2 100644 --- a/.buildkite/pipeline.yml +++ b/.buildkite/pipeline.yml @@ -88,6 +88,8 @@ steps: - -c - | apt-get update && apt-get install -y pkg-config libssl-dev protobuf-compiler + # wasm_middleware unit tests load the example component artifact + bash examples/wasm_middleware/build.sh cargo test --lib --bins agents: queue: "cpu_queue_premerge" @@ -126,6 +128,9 @@ steps: - -c - | apt-get update && apt-get install -y pkg-config libssl-dev protobuf-compiler python3 python3-pip python3-dev python3-venv + # WASM middleware pytest needs the example component + router binary + bash examples/wasm_middleware/build.sh + cargo build --release --bin vllm-router python3 -m venv /tmp/venv source /tmp/venv/bin/activate pip install -U pip setuptools wheel setuptools-rust diff --git a/.github/workflows/codespell.yml b/.github/workflows/codespell.yml index 83e58376..3c1c94a3 100644 --- a/.github/workflows/codespell.yml +++ b/.github/workflows/codespell.yml @@ -13,4 +13,4 @@ jobs: with: check_filenames: true skip: ./.git,./.github/workflows/codespell.yml,.git,*.png,*.jpg,*.svg,*.sum - ignore_words_list: "aks,te,hel" + ignore_words_list: "aks,te,hel,wit,wast" diff --git a/Cargo.lock b/Cargo.lock index 6b9ecec5..54919476 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,15 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "addr2line" +version = "0.25.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1b5d307320b3181d6d7954e663bd7c774a838b8220fe0593c86d9fb09f498b4b" +dependencies = [ + "gimli", +] + [[package]] name = "adler2" version = "2.0.1" @@ -37,6 +46,12 @@ version = "0.2.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" +[[package]] +name = "ambient-authority" +version = "0.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e9d4ee0d472d1cd2e28c97dfa124b3d8d992e10eb0a035f33f5d12e3a177ba3b" + [[package]] name = "android_system_properties" version = "0.1.5" @@ -108,6 +123,12 @@ version = "1.0.102" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" +[[package]] +name = "arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1" + [[package]] name = "async-broadcast" version = "0.7.2" @@ -197,7 +218,7 @@ dependencies = [ "futures-lite", "parking", "polling", - "rustix", + "rustix 1.1.4", "slab", "windows-sys 0.61.2", ] @@ -228,7 +249,7 @@ dependencies = [ "cfg-if", "event-listener 5.4.1", "futures-lite", - "rustix", + "rustix 1.1.4", ] [[package]] @@ -243,7 +264,7 @@ dependencies = [ "cfg-if", "futures-core", "futures-io", - "rustix", + "rustix 1.1.4", "signal-hook-registry", "slab", "windows-sys 0.61.2", @@ -565,6 +586,9 @@ name = "bumpalo" version = "3.20.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5d20789868f4b01b2f2caec9f5c4e0213b41e3e5702a50157d699ae31ced2fcb" +dependencies = [ + "allocator-api2", +] [[package]] name = "byteorder" @@ -578,6 +602,84 @@ version = "1.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1e748733b7cbc798e1434b6ac524f0c1ff2ab456fe201501e6497c8417a4fc33" +[[package]] +name = "cap-fs-ext" +version = "3.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "476f0d0003a760918ed4b1e039a59e11769030416f79c8222551d22785f7f70d" +dependencies = [ + "cap-primitives", + "cap-std", + "io-lifetimes", + "windows-sys 0.59.0", +] + +[[package]] +name = "cap-net-ext" +version = "3.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "150941cefd3df4de2fea24604ba4949371576f62e527410298333f7d431a1bc6" +dependencies = [ + "cap-primitives", + "cap-std", + "rustix 1.1.4", + "smallvec", +] + +[[package]] +name = "cap-primitives" +version = "3.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e0bf07d379916947be6c4a07f43684153d710a2896c31f9e97781362895596c" +dependencies = [ + "ambient-authority", + "fs-set-times", + "io-extras", + "io-lifetimes", + "ipnet", + "maybe-owned", + "rustix 1.1.4", + "rustix-linux-procfs", + "windows-sys 0.59.0", + "winx", +] + +[[package]] +name = "cap-rand" +version = "3.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ec6a5b75f54547c579a6b117c6fdd5f04f4ab7598de747b9f440a53592b3a4a" +dependencies = [ + "ambient-authority", + "rand 0.8.5", +] + +[[package]] +name = "cap-std" +version = "3.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a59e59fa26472d29680ece6a9f8ee8b0551a719a33df2f5240bde065ecbddfd7" +dependencies = [ + "cap-primitives", + "io-extras", + "io-lifetimes", + "rustix 1.1.4", +] + +[[package]] +name = "cap-time-ext" +version = "3.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b54c289326c70f1c697ebf0a31842a480932e5942b5fac92fcc46e87286b48e2" +dependencies = [ + "ambient-authority", + "cap-primitives", + "iana-time-zone", + "once_cell", + "rustix 1.1.4", + "winx", +] + [[package]] name = "cast" version = "0.3.0" @@ -723,6 +825,15 @@ dependencies = [ "cc", ] +[[package]] +name = "cobs" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0fa961b519f0b462e3a3b4a34b64d119eeaca1d59af726fe450bbba07a9fc0a1" +dependencies = [ + "thiserror 2.0.18", +] + [[package]] name = "colorchoice" version = "1.0.5" @@ -840,6 +951,144 @@ dependencies = [ "libc", ] +[[package]] +name = "cranelift-assembler-x64" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6835dba958b2ab7ab523e7e99296e0524317f60430a00cf5850562ef78ea7001" +dependencies = [ + "cranelift-assembler-x64-meta", +] + +[[package]] +name = "cranelift-assembler-x64-meta" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b6e4ce8ee6d899381fbdd9e6561336c651189d46cecaeee09b29e8d80aa786e" +dependencies = [ + "cranelift-srcgen", +] + +[[package]] +name = "cranelift-bforest" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cb6d37015df7ea4b60450c1229ad5f5819a1fb27434b063f8e6216dfbd0c42a" +dependencies = [ + "cranelift-entity", +] + +[[package]] +name = "cranelift-bitset" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "986bea0b0858b55192782120032ce9c15943fa073f186f6e479653c59e62c329" +dependencies = [ + "serde", + "serde_derive", +] + +[[package]] +name = "cranelift-codegen" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f30aeb2de7f97d6f26b4a1642615834daad58e2e4d7c027810010a3a32f22be" +dependencies = [ + "bumpalo", + "cranelift-assembler-x64", + "cranelift-bforest", + "cranelift-bitset", + "cranelift-codegen-meta", + "cranelift-codegen-shared", + "cranelift-control", + "cranelift-entity", + "cranelift-isle", + "gimli", + "hashbrown 0.15.5", + "log", + "pulley-interpreter", + "regalloc2", + "rustc-hash 2.1.1", + "serde", + "smallvec", + "target-lexicon 0.13.5", + "wasmtime-internal-math", +] + +[[package]] +name = "cranelift-codegen-meta" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cd5dd137fcdedef33b6fd40edf1ced024460d764ceb75833e8198a843395945c" +dependencies = [ + "cranelift-assembler-x64-meta", + "cranelift-codegen-shared", + "cranelift-srcgen", + "heck", + "pulley-interpreter", +] + +[[package]] +name = "cranelift-codegen-shared" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab54b260ef23a8f0f536679b9fc3b3b3e05353e8d1448f3ab83df02078e8be9b" + +[[package]] +name = "cranelift-control" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f3e569779ad70537f34a670d444ee3d75ae583b2023913f4682814b0979f7e8" +dependencies = [ + "arbitrary", +] + +[[package]] +name = "cranelift-entity" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ff53acc85f5c5f7d9315ff133a6671d329a0f04aa2d1a8a2e81d59709ccddcb" +dependencies = [ + "cranelift-bitset", + "serde", + "serde_derive", +] + +[[package]] +name = "cranelift-frontend" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab5976c0ff5bfadf61cd8bda81fea78ee5a07018b9cd03e66c0952c56684928b" +dependencies = [ + "cranelift-codegen", + "log", + "smallvec", + "target-lexicon 0.13.5", +] + +[[package]] +name = "cranelift-isle" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77b4f73d2288e9480fd2d1d9ab576394dce4805443d6148c6d819dbf78865ce4" + +[[package]] +name = "cranelift-native" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fe9650c2baf22fa1e2542a5bdd8152616ec2023d929c4cbb450ff677ad8d9c21" +dependencies = [ + "cranelift-codegen", + "libc", + "target-lexicon 0.13.5", +] + +[[package]] +name = "cranelift-srcgen" +version = "0.123.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ad4f61ae701d73c326d3df08c366b29ad10f1ba06c245092f217b8d2306746b" + [[package]] name = "crc32fast" version = "1.5.0" @@ -1192,6 +1441,18 @@ version = "1.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" +[[package]] +name = "embedded-io" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ef1a6892d9eef45c8fa6b9e0086428a2cca8491aca8f787c534a3d6d0bcb3ced" + +[[package]] +name = "embedded-io" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "edd0f118536f44f5ccd48bcb8b111bdc3de888b58c74639dfb034a357d0f206d" + [[package]] name = "encode_unicode" version = "1.0.0" @@ -1279,6 +1540,12 @@ dependencies = [ "pin-project-lite", ] +[[package]] +name = "fallible-iterator" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2acce4a10f12dc2fb14a218589d4f1f62ef011b2d0cc4b3cb1bba8e94da14649" + [[package]] name = "fancy-regex" version = "0.13.0" @@ -1296,6 +1563,17 @@ version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be" +[[package]] +name = "fd-lock" +version = "4.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ce92ff622d6dadf7349484f42c93271a0d49b7cc4d466a936405bacbe10aa78" +dependencies = [ + "cfg-if", + "rustix 1.1.4", + "windows-sys 0.59.0", +] + [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -1354,6 +1632,17 @@ dependencies = [ "percent-encoding", ] +[[package]] +name = "fs-set-times" +version = "0.20.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94e7099f6313ecacbe1256e8ff9d617b75d1bcb16a6fddef94866d225a01a14a" +dependencies = [ + "io-lifetimes", + "rustix 1.1.4", + "windows-sys 0.59.0", +] + [[package]] name = "fs_extra" version = "1.3.0" @@ -1511,6 +1800,17 @@ dependencies = [ "wasip3", ] +[[package]] +name = "gimli" +version = "0.32.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e629b9b98ef3dd8afe6ca2bd0f89306cec16d43d907889945bc5d6687f2f13c7" +dependencies = [ + "fallible-iterator", + "indexmap 2.13.0", + "stable_deref_trait", +] + [[package]] name = "glob" version = "0.3.3" @@ -1580,6 +1880,7 @@ dependencies = [ "allocator-api2", "equivalent", "foldhash 0.1.5", + "serde", ] [[package]] @@ -2045,6 +2346,22 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "io-extras" +version = "0.18.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2285ddfe3054097ef4b2fe909ef8c3bcd1ea52a8f0d274416caebeef39f04a65" +dependencies = [ + "io-lifetimes", + "windows-sys 0.59.0", +] + +[[package]] +name = "io-lifetimes" +version = "2.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06432fb54d3be7964ecd3649233cddf80db2832f47fec34c01f65b3d9d774983" + [[package]] name = "ipnet" version = "2.12.0" @@ -2326,6 +2643,12 @@ version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" +[[package]] +name = "leb128" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c83bff1d572d6b9aeef67ddfc8448e4a3737909cb28e81f97c791b9018703e52" + [[package]] name = "leb128fmt" version = "0.1.0" @@ -2338,6 +2661,12 @@ version = "0.2.183" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b5b646652bf6661599e1da8901b3b9522896f01e736bad5f723fe7a3a27f899d" +[[package]] +name = "libm" +version = "0.2.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981" + [[package]] name = "libredox" version = "0.1.14" @@ -2347,6 +2676,12 @@ dependencies = [ "libc", ] +[[package]] +name = "linux-raw-sys" +version = "0.4.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d26c52dbd32dccf2d10cac7725f8eae5296885fb5703b261f7d0a0739ec807ab" + [[package]] name = "linux-raw-sys" version = "0.12.1" @@ -2383,6 +2718,15 @@ version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" +[[package]] +name = "mach2" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d640282b302c0bb0a2a8e0233ead9035e3bed871f0b7e81fe4a1ec829765db44" +dependencies = [ + "libc", +] + [[package]] name = "macro_rules_attribute" version = "0.2.2" @@ -2420,12 +2764,27 @@ version = "0.8.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" +[[package]] +name = "maybe-owned" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4facc753ae494aeb6e3c22f839b158aebd4f9270f55cd3c79906c45476c47ab4" + [[package]] name = "memchr" version = "2.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" +[[package]] +name = "memfd" +version = "0.6.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57804b2c9b69967f1536a56f86297e367a33b19e98852ed624b84551cdbc0d90" +dependencies = [ + "rustix 1.1.4", +] + [[package]] name = "memo-map" version = "0.3.3" @@ -2620,6 +2979,18 @@ version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "830b246a0e5f20af87141b25c173cd1b609bd7779a4617d6ec582abaf90870f3" +[[package]] +name = "object" +version = "0.37.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff76201f031d8863c38aa7f905eca4f53abbfa15f609db4277d44cd8938f33fe" +dependencies = [ + "crc32fast", + "hashbrown 0.15.5", + "indexmap 2.13.0", + "memchr", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -2985,7 +3356,7 @@ dependencies = [ "concurrent-queue", "hermit-abi", "pin-project-lite", - "rustix", + "rustix 1.1.4", "windows-sys 0.61.2", ] @@ -3004,6 +3375,18 @@ dependencies = [ "rand 0.8.5", ] +[[package]] +name = "postcard" +version = "1.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6764c3b5dd454e283a30e6dfe78e9b31096d9e32036b5d1eaac7a6119ccb9a24" +dependencies = [ + "cobs", + "embedded-io 0.4.0", + "embedded-io 0.6.1", + "serde", +] + [[package]] name = "potential_utf" version = "0.1.4" @@ -3070,6 +3453,29 @@ dependencies = [ "syn", ] +[[package]] +name = "pulley-interpreter" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eb0a4b56042e461cc64456650182938e2d1ede98fa0c8a975027416a2809c414" +dependencies = [ + "cranelift-bitset", + "log", + "pulley-macros", + "wasmtime-internal-math", +] + +[[package]] +name = "pulley-macros" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "244667bea2e214273442a71f26adb12b88a41f66718fb2c6eea47c00f0dc325f" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "pyo3" version = "0.26.0" @@ -3351,6 +3757,20 @@ dependencies = [ "thiserror 2.0.18", ] +[[package]] +name = "regalloc2" +version = "0.12.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5216b1837de2149f8bc8e6d5f88a9326b63b8c836ed58ce4a0a29ec736a59734" +dependencies = [ + "allocator-api2", + "bumpalo", + "hashbrown 0.15.5", + "log", + "rustc-hash 2.1.1", + "smallvec", +] + [[package]] name = "regex" version = "1.12.3" @@ -3521,6 +3941,19 @@ dependencies = [ "semver", ] +[[package]] +name = "rustix" +version = "0.38.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" +dependencies = [ + "bitflags 2.11.0", + "errno", + "libc", + "linux-raw-sys 0.4.15", + "windows-sys 0.59.0", +] + [[package]] name = "rustix" version = "1.1.4" @@ -3530,13 +3963,23 @@ dependencies = [ "bitflags 2.11.0", "errno", "libc", - "linux-raw-sys", + "linux-raw-sys 0.12.1", "windows-sys 0.61.2", ] [[package]] -name = "rustls" -version = "0.23.37" +name = "rustix-linux-procfs" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2fc84bf7e9aa16c4f2c758f27412dc9841341e16aa682d9c7ac308fe3ee12056" +dependencies = [ + "once_cell", + "rustix 1.1.4", +] + +[[package]] +name = "rustls" +version = "0.23.37" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "758025cb5fccfd3bc2fd74708fd4682be41d99e5dff73c377c0646c6012c73a4" dependencies = [ @@ -3743,6 +4186,10 @@ name = "semver" version = "1.0.27" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d767eb0aabc880b29956c35734170f26ed551a859dbd361d140cdbeca61ab1e2" +dependencies = [ + "serde", + "serde_core", +] [[package]] name = "serde" @@ -3923,6 +4370,9 @@ name = "smallvec" version = "1.15.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" +dependencies = [ + "serde", +] [[package]] name = "socket2" @@ -4078,6 +4528,22 @@ dependencies = [ "version-compare", ] +[[package]] +name = "system-interface" +version = "0.27.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc4592f674ce18521c2a81483873a49596655b179f71c5e05d10c1fe66c78745" +dependencies = [ + "bitflags 2.11.0", + "cap-fs-ext", + "cap-std", + "fd-lock", + "io-lifetimes", + "rustix 0.38.44", + "windows-sys 0.59.0", + "winx", +] + [[package]] name = "target-lexicon" version = "0.12.16" @@ -4099,10 +4565,19 @@ dependencies = [ "fastrand", "getrandom 0.4.2", "once_cell", - "rustix", + "rustix 1.1.4", "windows-sys 0.61.2", ] +[[package]] +name = "termcolor" +version = "1.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06794f8f6c5c898b3275aebefa6b8a1cb24cd2c6c79397ab15774837a0bc5755" +dependencies = [ + "winapi-util", +] + [[package]] name = "thiserror" version = "1.0.69" @@ -4807,6 +5282,7 @@ name = "vllm_router_rs" version = "0.1.15" dependencies = [ "anyhow", + "async-channel 2.5.0", "async-trait", "axum 0.8.8", "backoff", @@ -4843,6 +5319,7 @@ dependencies = [ "rustls", "serde", "serde_json", + "sha2", "strum", "tempfile", "thiserror 2.0.18", @@ -4860,6 +5337,8 @@ dependencies = [ "ulid", "url", "uuid", + "wasmtime", + "wasmtime-wasi", "zmq", ] @@ -4965,6 +5444,16 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "wasm-encoder" +version = "0.236.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "724fccfd4f3c24b7e589d333fc0429c68042897a7e8a5f8694f31792471841e7" +dependencies = [ + "leb128fmt", + "wasmparser 0.236.1", +] + [[package]] name = "wasm-encoder" version = "0.244.0" @@ -4972,7 +5461,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "990065f2fe63003fe337b932cfb5e3b80e0b4d0f5ff650e6985b1048f62c8319" dependencies = [ "leb128fmt", - "wasmparser", + "wasmparser 0.244.0", ] [[package]] @@ -4983,8 +5472,8 @@ checksum = "bb0e353e6a2fbdc176932bbaab493762eb1255a7900fe0fea1a2f96c296cc909" dependencies = [ "anyhow", "indexmap 2.13.0", - "wasm-encoder", - "wasmparser", + "wasm-encoder 0.244.0", + "wasmparser 0.244.0", ] [[package]] @@ -5013,6 +5502,19 @@ dependencies = [ "web-sys", ] +[[package]] +name = "wasmparser" +version = "0.236.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a9b1e81f3eb254cf7404a82cee6926a4a3ccc5aad80cc3d43608a070c67aa1d7" +dependencies = [ + "bitflags 2.11.0", + "hashbrown 0.15.5", + "indexmap 2.13.0", + "semver", + "serde", +] + [[package]] name = "wasmparser" version = "0.244.0" @@ -5025,6 +5527,307 @@ dependencies = [ "semver", ] +[[package]] +name = "wasmprinter" +version = "0.236.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2df225df06a6df15b46e3f73ca066ff92c2e023670969f7d50ce7d5e695abbb1" +dependencies = [ + "anyhow", + "termcolor", + "wasmparser 0.236.1", +] + +[[package]] +name = "wasmtime" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7d05c745dc0978e589ef295958f3130122afc33d96af6bad3f0f06dbe7ac43a8" +dependencies = [ + "addr2line", + "anyhow", + "async-trait", + "bitflags 2.11.0", + "bumpalo", + "cc", + "cfg-if", + "encoding_rs", + "hashbrown 0.15.5", + "indexmap 2.13.0", + "libc", + "log", + "mach2", + "memfd", + "object", + "once_cell", + "postcard", + "pulley-interpreter", + "rayon", + "rustix 1.1.4", + "semver", + "serde", + "serde_derive", + "smallvec", + "target-lexicon 0.13.5", + "wasmparser 0.236.1", + "wasmtime-environ", + "wasmtime-internal-asm-macros", + "wasmtime-internal-component-macro", + "wasmtime-internal-component-util", + "wasmtime-internal-cranelift", + "wasmtime-internal-fiber", + "wasmtime-internal-jit-debug", + "wasmtime-internal-jit-icache-coherence", + "wasmtime-internal-math", + "wasmtime-internal-slab", + "wasmtime-internal-unwinder", + "wasmtime-internal-versioned-export-macros", + "wasmtime-internal-winch", + "windows-sys 0.60.2", +] + +[[package]] +name = "wasmtime-environ" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9fd1d43cfaa1a0859d2f4fccc15e7e571e2a88b357e81bc88ba6c501b83d925d" +dependencies = [ + "anyhow", + "cranelift-bitset", + "cranelift-entity", + "gimli", + "indexmap 2.13.0", + "log", + "object", + "postcard", + "semver", + "serde", + "serde_derive", + "smallvec", + "target-lexicon 0.13.5", + "wasm-encoder 0.236.1", + "wasmparser 0.236.1", + "wasmprinter", + "wasmtime-internal-component-util", +] + +[[package]] +name = "wasmtime-internal-asm-macros" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "515dd7158bf1719b41290cd2e6a2a46ec944484146816992f195af3720e49b3f" +dependencies = [ + "cfg-if", +] + +[[package]] +name = "wasmtime-internal-component-macro" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dfca017b7daa80ff217c66f105ee20e674d87b7c11dca85bbb9d0146e9f443fb" +dependencies = [ + "anyhow", + "proc-macro2", + "quote", + "syn", + "wasmtime-internal-component-util", + "wasmtime-internal-wit-bindgen", + "wit-parser 0.236.1", +] + +[[package]] +name = "wasmtime-internal-component-util" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1c3e218b51d2ef9eb181499e42c691512c1bc11dd6dc7746807a9b2b9290369c" + +[[package]] +name = "wasmtime-internal-cranelift" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ba1736927b58e50e741e407da7c037c0250f3e213833a09c89dcd8f73ae2eac" +dependencies = [ + "anyhow", + "cfg-if", + "cranelift-codegen", + "cranelift-control", + "cranelift-entity", + "cranelift-frontend", + "cranelift-native", + "gimli", + "itertools 0.14.0", + "log", + "object", + "pulley-interpreter", + "smallvec", + "target-lexicon 0.13.5", + "thiserror 2.0.18", + "wasmparser 0.236.1", + "wasmtime-environ", + "wasmtime-internal-math", + "wasmtime-internal-versioned-export-macros", +] + +[[package]] +name = "wasmtime-internal-fiber" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b238e4c20bddb900ec0cb380252d63e8d0644fd94de001119574f5921e895d9" +dependencies = [ + "anyhow", + "cc", + "cfg-if", + "libc", + "rustix 1.1.4", + "wasmtime-internal-asm-macros", + "wasmtime-internal-versioned-export-macros", + "windows-sys 0.60.2", +] + +[[package]] +name = "wasmtime-internal-jit-debug" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f259b13685ad51e3dcf58cb69031279ed0d79c25bc3ccc8b50e7160ed04fbfe" +dependencies = [ + "cc", + "wasmtime-internal-versioned-export-macros", +] + +[[package]] +name = "wasmtime-internal-jit-icache-coherence" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41fed85537936b16460bac352ad149052c025db50467c7bc539dd47b31439374" +dependencies = [ + "anyhow", + "cfg-if", + "libc", + "windows-sys 0.60.2", +] + +[[package]] +name = "wasmtime-internal-math" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "82fff10da41d0d15d90ebba70946a0aa16ed0957ae7b77e0b6d2a46e8221e555" +dependencies = [ + "libm", +] + +[[package]] +name = "wasmtime-internal-slab" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e44a8c097bab08d349d57dce1ab818859fefbe261ab3632b38fe127b1b551108" + +[[package]] +name = "wasmtime-internal-unwinder" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f40a57d5e7c221ce56391d7dca0a918ba17ea00185462c7facbf534d7745184" +dependencies = [ + "anyhow", + "cfg-if", + "cranelift-codegen", + "log", + "object", +] + +[[package]] +name = "wasmtime-internal-versioned-export-macros" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e085bfce1cb2089dbeef6e280a5d598666923d3dcd308712fe429fe43c9d19f5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "wasmtime-internal-winch" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4916cd526e1ce294984cc5b70264cfc0ca103b41ca665b58728bde250b6b82f" +dependencies = [ + "anyhow", + "cranelift-codegen", + "gimli", + "object", + "target-lexicon 0.13.5", + "wasmparser 0.236.1", + "wasmtime-environ", + "wasmtime-internal-cranelift", + "winch-codegen", +] + +[[package]] +name = "wasmtime-internal-wit-bindgen" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39ad2f9d9c3baa70ee4d157b55b4c08e5dc00f1d80aad3692b1e0806393914ea" +dependencies = [ + "anyhow", + "bitflags 2.11.0", + "heck", + "indexmap 2.13.0", + "wit-parser 0.236.1", +] + +[[package]] +name = "wasmtime-wasi" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "841c11707aaeaf09895677757614d83a23c45d2b2689df908d0005aa191a933b" +dependencies = [ + "anyhow", + "async-trait", + "bitflags 2.11.0", + "bytes", + "cap-fs-ext", + "cap-net-ext", + "cap-rand", + "cap-std", + "cap-time-ext", + "fs-set-times", + "futures", + "io-extras", + "io-lifetimes", + "rustix 1.1.4", + "system-interface", + "thiserror 2.0.18", + "tokio", + "tracing", + "url", + "wasmtime", + "wasmtime-wasi-io", + "wiggle", + "windows-sys 0.60.2", +] + +[[package]] +name = "wasmtime-wasi-io" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b76b03de2cba3c036f81074c17bb1c7b77662578424887ebbc090b4c23ba33e" +dependencies = [ + "anyhow", + "async-trait", + "bytes", + "futures", + "wasmtime", +] + +[[package]] +name = "wast" +version = "35.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ef140f1b49946586078353a453a1d28ba90adfc54dde75710bc1931de204d68" +dependencies = [ + "leb128", +] + [[package]] name = "web-sys" version = "0.3.91" @@ -5072,6 +5875,47 @@ dependencies = [ "rustls-pki-types", ] +[[package]] +name = "wiggle" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a333010f4b8b770a500e181e0abda564dde49a51db21aed4307fde800746ff97" +dependencies = [ + "anyhow", + "async-trait", + "bitflags 2.11.0", + "thiserror 2.0.18", + "tracing", + "wasmtime", + "wiggle-macro", +] + +[[package]] +name = "wiggle-generate" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e909aa247f90d2ba77b6860b36288cfeeb559fe01b140165f077b250cb1ba56" +dependencies = [ + "anyhow", + "heck", + "proc-macro2", + "quote", + "syn", + "witx", +] + +[[package]] +name = "wiggle-macro" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a84de50b5bf6530a15bbb228043561d525c8e21c8929693c3078f7cad2051ad3" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "wiggle-generate", +] + [[package]] name = "winapi" version = "0.3.9" @@ -5103,6 +5947,26 @@ version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" +[[package]] +name = "winch-codegen" +version = "36.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e826c012c68403725e77adf6b904c2ea809e5d464aaf25aa6eda14559300b3df" +dependencies = [ + "anyhow", + "cranelift-assembler-x64", + "cranelift-codegen", + "gimli", + "regalloc2", + "smallvec", + "target-lexicon 0.13.5", + "thiserror 2.0.18", + "wasmparser 0.236.1", + "wasmtime-environ", + "wasmtime-internal-cranelift", + "wasmtime-internal-math", +] + [[package]] name = "windows-core" version = "0.62.2" @@ -5413,6 +6277,16 @@ dependencies = [ "memchr", ] +[[package]] +name = "winx" +version = "0.36.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f3fd376f71958b862e7afb20cfe5a22830e1963462f3a17f49d82a6c1d1f42d" +dependencies = [ + "bitflags 2.11.0", + "windows-sys 0.59.0", +] + [[package]] name = "wit-bindgen" version = "0.51.0" @@ -5430,7 +6304,7 @@ checksum = "ea61de684c3ea68cb082b7a88508a8b27fcc8b797d738bfc99a82facf1d752dc" dependencies = [ "anyhow", "heck", - "wit-parser", + "wit-parser 0.244.0", ] [[package]] @@ -5477,10 +6351,28 @@ dependencies = [ "serde", "serde_derive", "serde_json", - "wasm-encoder", + "wasm-encoder 0.244.0", "wasm-metadata", - "wasmparser", - "wit-parser", + "wasmparser 0.244.0", + "wit-parser 0.244.0", +] + +[[package]] +name = "wit-parser" +version = "0.236.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "16e4833a20cd6e85d6abfea0e63a399472d6f88c6262957c17f546879a80ba15" +dependencies = [ + "anyhow", + "id-arena", + "indexmap 2.13.0", + "log", + "semver", + "serde", + "serde_derive", + "serde_json", + "unicode-xid", + "wasmparser 0.236.1", ] [[package]] @@ -5498,7 +6390,19 @@ dependencies = [ "serde_derive", "serde_json", "unicode-xid", - "wasmparser", + "wasmparser 0.244.0", +] + +[[package]] +name = "witx" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e366f27a5cabcddb2706a78296a40b8fcc451e1a6aba2fc1d94b4a01bdaaef4b" +dependencies = [ + "anyhow", + "log", + "thiserror 1.0.69", + "wast", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 9489359c..373408a3 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -80,6 +80,16 @@ strum = { version = "0.26", features = ["derive"] } once_cell = "1.21.3" zmq = "0.10.0" rmp-serde = "1.3" +wasmtime = { version = "36", default-features = false, features = [ + "component-model", + "cranelift", + "parallel-compilation", + "pooling-allocator", + "runtime", +] } +wasmtime-wasi = "36" +sha2 = "0.10" +async-channel = "2" [dev-dependencies] criterion = { version = "0.5", features = ["html_reports"] } diff --git a/Dockerfile.router b/Dockerfile.router index 43401a29..cb41b4ca 100644 --- a/Dockerfile.router +++ b/Dockerfile.router @@ -11,9 +11,10 @@ RUN apt-get update && apt-get install -y \ # Set working directory WORKDIR /app -# Copy Rust source files +# Copy Rust source files (wit/ is required by wasmtime bindgen in wasm_middleware) COPY Cargo.toml Cargo.lock ./ COPY src ./src +COPY wit ./wit # Build the Rust binary RUN cargo build --release diff --git a/MANIFEST.in b/MANIFEST.in index e1d6e7a9..ea30e3d3 100644 --- a/MANIFEST.in +++ b/MANIFEST.in @@ -1,3 +1,6 @@ # Must include: include Cargo.toml # Rust project configuration +include Cargo.lock # Locked Rust dependency graph for reproducible builds recursive-include src *.rs # Rust source files +# Required by wasmtime::component::bindgen!({ path: "wit", ... }) in src/wasm_middleware.rs +recursive-include wit * diff --git a/README.md b/README.md index ce4a7349..5bab02cd 100644 --- a/README.md +++ b/README.md @@ -80,6 +80,29 @@ vllm-router \ --intra-node-data-parallel-size 8 ``` +#### Optional WASM OnRequest middleware + +Load an independently built WASM Component plugin (see `examples/wasm_middleware/` and [RFC #236](https://github.com/vllm-project/router/issues/236)). By default it attaches only to `POST /v1/chat/completions` and fails closed on plugin errors: + +```bash +./examples/wasm_middleware/build.sh + +./target/release/vllm-router \ + --worker-urls http://localhost:8000 \ + --wasm-middleware ./examples/wasm_middleware/wasm_middleware_example.component.wasm \ + --wasm-middleware-route /v1/chat/completions +``` + +Additional paths can be attached with repeated `--wasm-middleware-route` flags (must be one of the protected inference routes). Without `--wasm-middleware`, the Router does not initialize Wasmtime. + +v0.1 resource / fail-closed defaults on attached routes (not configurable via CLI yet): + +- **Input body cap**: `min(10 MiB, --max-payload-size)`. Requests larger than this get **413** before the plugin runs, even if the plugin would only `Continue`. This is intentionally tighter than the Router's default 512 MiB payload limit. +- **Execution deadline**: **100 ms** per invocation (Wasmtime epoch interruption). Deadline / trap failures fail closed with **500**. +- **Queue full**: when the bounded worker queue is saturated, matching requests get **503**. + +Prometheus metrics for the WASM runtime are deferred to a later revision. + #### Prefill-Decode Disaggregation ```bash # When vLLM runs the NIXL connector, prefill/decode URLs are required. diff --git a/examples/wasm_middleware/.gitignore b/examples/wasm_middleware/.gitignore new file mode 100644 index 00000000..caa67449 --- /dev/null +++ b/examples/wasm_middleware/.gitignore @@ -0,0 +1,3 @@ +/target/ +*.wasm +Cargo.lock diff --git a/examples/wasm_middleware/Cargo.toml b/examples/wasm_middleware/Cargo.toml new file mode 100644 index 00000000..d6101c03 --- /dev/null +++ b/examples/wasm_middleware/Cargo.toml @@ -0,0 +1,12 @@ +[package] +name = "wasm_middleware_example" +version = "0.1.0" +edition = "2021" +publish = false +description = "Example OnRequest WASM middleware for vLLM Router" + +[lib] +crate-type = ["cdylib"] + +[dependencies] +wit-bindgen = "0.46" diff --git a/examples/wasm_middleware/README.md b/examples/wasm_middleware/README.md new file mode 100644 index 00000000..98105fa5 --- /dev/null +++ b/examples/wasm_middleware/README.md @@ -0,0 +1,30 @@ +# Example WASM OnRequest middleware + +Small demo guest for the vLLM Router pluggable middleware runtime +([RFC #236](https://github.com/vllm-project/router/issues/236)). + +## Build + +```bash +./build.sh +``` + +Produces `wasm_middleware_example.component.wasm` in this directory. + +## Behavior + +- `Reject(N)` when the body contains `__wasm_reject_N__` (e.g. `__wasm_reject_403__`) +- `Reject(400)` when the body contains `__wasm_reject__` +- Infinite loop when the body contains `__wasm_loop__` (for host timeout tests only) +- Otherwise `Modify` with header `x-wasm-middleware: example` (body unchanged) + +## Run with Router + +```bash +cargo run --release -- \ + --worker-urls http://localhost:8000 \ + --wasm-middleware ./examples/wasm_middleware/wasm_middleware_example.component.wasm +``` + +By default the host only invokes the plugin for `POST /v1/chat/completions`. +Use repeated `--wasm-middleware-route` flags to attach additional paths. diff --git a/examples/wasm_middleware/build.sh b/examples/wasm_middleware/build.sh new file mode 100755 index 00000000..12438f48 --- /dev/null +++ b/examples/wasm_middleware/build.sh @@ -0,0 +1,13 @@ +#!/usr/bin/env bash +set -euo pipefail + +cd "$(dirname "$0")" + +rustup target add wasm32-wasip2 >/dev/null +cargo build --release --target wasm32-wasip2 + +# wit-bindgen + wasm32-wasip2 already emits a component. +cp -f target/wasm32-wasip2/release/wasm_middleware_example.wasm \ + ./wasm_middleware_example.component.wasm + +echo "Built $(pwd)/wasm_middleware_example.component.wasm" diff --git a/examples/wasm_middleware/src/lib.rs b/examples/wasm_middleware/src/lib.rs new file mode 100644 index 00000000..82dec01b --- /dev/null +++ b/examples/wasm_middleware/src/lib.rs @@ -0,0 +1,80 @@ +//! Example OnRequest middleware for vLLM Router. +//! +//! Demo behavior (not product logic): +//! - Reject(N) when the body contains `__wasm_reject_N__` (e.g. `__wasm_reject_403__`) +//! - Reject(400) when the body contains `__wasm_reject__` +//! - Infinite loop when the body contains `__wasm_loop__` (for host timeout tests) +//! - Otherwise Modify: set `x-wasm-middleware: example` and leave the body unchanged + +#![allow(clippy::missing_safety_doc)] + +wit_bindgen::generate!({ + path: "../../wit", + world: "middleware", +}); + +use exports::vllm::router_middleware::on_request::{Guest, Request}; +use vllm::router_middleware::types::{Action, Header, ModifyAction}; + +struct Component; + +impl Guest for Component { + fn handle(req: Request) -> Action { + if body_contains_marker(&req.body, b"__wasm_loop__") { + // Intentional spin for host epoch-deadline tests. Do not use in production plugins. + loop {} + } + + if let Some(status) = parse_reject_status(&req.body) { + return Action::Reject(status); + } + + Action::Modify(ModifyAction { + headers_set: vec![Header { + name: "x-wasm-middleware".to_string(), + value: b"example".to_vec(), + }], + headers_add: Vec::new(), + headers_remove: Vec::new(), + body_replace: None, + }) + } +} + +fn body_contains_marker(body: &[u8], marker: &[u8]) -> bool { + body.windows(marker.len()).any(|window| window == marker) +} + +/// Parse `__wasm_reject__` (defaults to 400) or `__wasm_reject_NNN__`. +fn parse_reject_status(body: &[u8]) -> Option { + const PREFIX: &[u8] = b"__wasm_reject_"; + const BARE: &[u8] = b"__wasm_reject__"; + + if let Some(start) = find_subslice(body, PREFIX) { + let rest = &body[start + PREFIX.len()..]; + let digits: Vec = rest.iter().copied().take_while(|b| b.is_ascii_digit()).collect(); + if rest.get(digits.len()..) + .is_some_and(|tail| tail.starts_with(b"__")) + && !digits.is_empty() + { + if let Ok(text) = std::str::from_utf8(&digits) { + if let Ok(status) = text.parse::() { + return Some(status); + } + } + } + } + + if body_contains_marker(body, BARE) { + return Some(400); + } + None +} + +fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option { + haystack + .windows(needle.len()) + .position(|window| window == needle) +} + +export!(Component); diff --git a/py_src/vllm_router/router_args.py b/py_src/vllm_router/router_args.py index 5771bf49..9fe1f62e 100644 --- a/py_src/vllm_router/router_args.py +++ b/py_src/vllm_router/router_args.py @@ -33,6 +33,9 @@ class RouterArgs: eviction_interval_secs: int = 120 max_tree_size: int = 2**26 max_payload_size: int = 512 * 1024 * 1024 # 512MB default for large batches + wasm_middleware: Optional[str] = None + wasm_middleware_sha256: Optional[str] = None + wasm_middleware_routes: List[str] = dataclasses.field(default_factory=list) intra_node_data_parallel_size: int = ( 1 # Intra-node data parallel size (DP-aware routing automatically enabled when > 1) ) @@ -244,6 +247,31 @@ def add_cli_args( default=RouterArgs.max_payload_size, help="Maximum payload size in bytes", ) + parser.add_argument( + f"--{prefix}wasm-middleware", + type=str, + default=None, + help=( + "Path to a WASM Component Model OnRequest middleware artifact. " + "Fail-closed: plugin errors reject the request." + ), + ) + parser.add_argument( + f"--{prefix}wasm-middleware-sha256", + type=str, + default=None, + help="Optional SHA-256 hex digest that must match --wasm-middleware.", + ) + parser.add_argument( + f"--{prefix}wasm-middleware-route", + action="append", + dest="wasm_middleware_routes", + default=[], + help=( + "HTTP path that invokes the WASM middleware. Repeatable. " + "Defaults to /v1/chat/completions when --wasm-middleware is set." + ), + ) parser.add_argument( f"--{prefix}intra-node-data-parallel-size", type=int, @@ -518,6 +546,9 @@ def from_cli_args( return cls(**args_dict) def _validate_router_args(self): + if self.wasm_middleware_sha256 and not self.wasm_middleware: + raise ValueError("wasm_middleware_sha256 requires wasm_middleware") + # Validate configuration based on mode if self.vllm_pd_disaggregation: # Validate PD configuration - skip URL requirements if using service discovery diff --git a/py_test/fixtures/mock_worker.py b/py_test/fixtures/mock_worker.py index 08fe63bf..583a93f9 100644 --- a/py_test/fixtures/mock_worker.py +++ b/py_test/fixtures/mock_worker.py @@ -196,6 +196,13 @@ async def handle_text_request(request: Request): ], "worker_id": worker_id, "echo": data, + # Helpful for middleware e2e assertions (e.g. WASM header injection). + "request_headers": { + k: v + for k, v in request.headers.items() + if k.lower().startswith("x-wasm") + or k.lower() in {"content-type", "authorization"} + }, } return make_json_response(ret, status_code=200) diff --git a/py_test/fixtures/router_manager.py b/py_test/fixtures/router_manager.py index 98183218..163d8b37 100644 --- a/py_test/fixtures/router_manager.py +++ b/py_test/fixtures/router_manager.py @@ -32,20 +32,34 @@ def start_router( decode_urls: Optional[List[str]] = None, prefill_policy: Optional[str] = None, decode_policy: Optional[str] = None, + # Prefer the Rust binary when testing features that live in the host crate + # (e.g. WASM middleware) without requiring a rebuilt Python wheel. + router_bin: Optional[str] = None, ) -> ProcHandle: worker_urls = worker_urls or [] port = port or find_free_port() - cmd = [ - "python3", - "-m", - "vllm_router.launch_router", - "--host", - "127.0.0.1", - "--port", - str(port), - "--policy", - policy, - ] + if router_bin: + cmd = [ + router_bin, + "--host", + "127.0.0.1", + "--port", + str(port), + "--policy", + policy, + ] + else: + cmd = [ + "python3", + "-m", + "vllm_router.launch_router", + "--host", + "127.0.0.1", + "--port", + str(port), + "--policy", + policy, + ] # Avoid Prometheus port collisions by assigning a free port per router prom_port = find_free_port() cmd.extend( @@ -79,6 +93,7 @@ def start_router( "api_key": "--api-key", # Health/monitoring "worker_startup_check_interval": "--worker-startup-check-interval", + "worker_startup_timeout_secs": "--worker-startup-timeout-secs", # Cache-aware tuning "cache_threshold": "--cache-threshold", "balance_abs_threshold": "--balance-abs-threshold", @@ -101,10 +116,17 @@ def start_router( "queue_size": "--queue-size", "queue_timeout_secs": "--queue-timeout-secs", "rate_limit_tokens_per_second": "--rate-limit-tokens-per-second", + # WASM OnRequest middleware + "wasm_middleware": "--wasm-middleware", + "wasm_middleware_sha256": "--wasm-middleware-sha256", } for k, v in extra.items(): if v is None: continue + if k == "wasm_middleware_routes": + for route in v: + cmd.extend(["--wasm-middleware-route", str(route)]) + continue flag = flag_map.get(k) if not flag: continue diff --git a/py_test/integration/test_wasm_middleware.py b/py_test/integration/test_wasm_middleware.py new file mode 100644 index 00000000..894944a1 --- /dev/null +++ b/py_test/integration/test_wasm_middleware.py @@ -0,0 +1,96 @@ +"""Integration tests for WASM OnRequest middleware (real Router + mock worker).""" + +from __future__ import annotations + +import subprocess +from pathlib import Path + +import pytest +import requests + +REPO_ROOT = Path(__file__).resolve().parents[2] +EXAMPLE_DIR = REPO_ROOT / "examples" / "wasm_middleware" +WASM_COMPONENT = EXAMPLE_DIR / "wasm_middleware_example.component.wasm" + + +def _router_bin() -> Path: + for candidate in ( + REPO_ROOT / "target" / "release" / "vllm-router", + REPO_ROOT / "target" / "debug" / "vllm-router", + ): + if candidate.is_file(): + return candidate + pytest.skip( + "vllm-router binary not found; build with " + "`cargo build --release --bin vllm-router` first" + ) + + +def _ensure_wasm_component() -> Path: + if WASM_COMPONENT.is_file(): + return WASM_COMPONENT + build_sh = EXAMPLE_DIR / "build.sh" + if not build_sh.is_file(): + pytest.skip(f"missing example plugin build script at {build_sh}") + subprocess.check_call(["bash", str(build_sh)], cwd=EXAMPLE_DIR) + if not WASM_COMPONENT.is_file(): + pytest.skip(f"failed to build WASM component at {WASM_COMPONENT}") + return WASM_COMPONENT + + +def _header_map(payload: dict) -> dict: + headers = payload.get("request_headers") or {} + return {str(k).lower(): v for k, v in headers.items()} + + +@pytest.mark.integration +def test_wasm_middleware_modify_reject_and_path_isolation(router_manager, mock_workers): + """Example plugin: set header on chat, reject marker body, skip other routes.""" + wasm_path = _ensure_wasm_component() + router_bin = _router_bin() + _, urls, _ = mock_workers(n=1) + + rh = router_manager.start_router( + worker_urls=urls, + policy="round_robin", + router_bin=str(router_bin), + extra={ + "worker_startup_timeout_secs": 30, + "wasm_middleware": str(wasm_path), + }, + ) + + # Modify: chat should forward with the example header. + chat = requests.post( + f"{rh.url}/v1/chat/completions", + json={ + "model": "mock", + "messages": [{"role": "user", "content": "hi"}], + }, + timeout=30, + ) + assert chat.status_code == 200, chat.text + chat_headers = _header_map(chat.json()) + assert chat_headers.get("x-wasm-middleware") == "example", chat_headers + + # Reject: fail closed at the Router before the worker. + rejected = requests.post( + f"{rh.url}/v1/chat/completions", + json={ + "model": "mock", + "messages": [{"role": "user", "content": "nope"}], + "note": "__wasm_reject__", + }, + timeout=30, + ) + assert rejected.status_code == 400, rejected.text + + # Path isolation: completions is not attached by default. + completions = requests.post( + f"{rh.url}/v1/completions", + json={"model": "mock", "prompt": "hi"}, + timeout=30, + ) + assert completions.status_code == 200, completions.text + completion_headers = _header_map(completions.json()) + assert "x-wasm-middleware" not in completion_headers, completion_headers diff --git a/src/lib.rs b/src/lib.rs index 5c4c4066..68628b57 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -16,6 +16,7 @@ pub mod server; pub mod service_discovery; pub mod tokenizer; pub mod tree; +pub mod wasm_middleware; use crate::metrics::PrometheusConfig; #[pyclass(eq)] @@ -43,6 +44,9 @@ struct Router { eviction_interval_secs: u64, max_tree_size: usize, max_payload_size: usize, + wasm_middleware: Option, + wasm_middleware_sha256: Option, + wasm_middleware_routes: Vec, intra_node_data_parallel_size: usize, api_key: Option, api_key_validation_urls: Vec, @@ -310,6 +314,9 @@ impl Router { otlp_traces_endpoint = None, // KV connector default (PD disaggregation) kv_connector = String::from("nixl"), + wasm_middleware = None, + wasm_middleware_sha256 = None, + wasm_middleware_routes = vec![], ))] #[allow(clippy::too_many_arguments)] fn new( @@ -372,7 +379,25 @@ impl Router { enable_trace: bool, otlp_traces_endpoint: Option, kv_connector: String, + wasm_middleware: Option, + wasm_middleware_sha256: Option, + wasm_middleware_routes: Vec, ) -> PyResult { + if wasm_middleware_sha256 + .as_deref() + .map(str::trim) + .filter(|v| !v.is_empty()) + .is_some() + && wasm_middleware + .as_deref() + .map(str::trim) + .filter(|v| !v.is_empty()) + .is_none() + { + return Err(pyo3::exceptions::PyValueError::new_err( + "wasm_middleware_sha256 requires wasm_middleware", + )); + } Ok(Router { host, port, @@ -386,6 +411,9 @@ impl Router { eviction_interval_secs, max_tree_size, max_payload_size, + wasm_middleware, + wasm_middleware_sha256, + wasm_middleware_routes, intra_node_data_parallel_size, api_key, api_key_validation_urls, @@ -488,6 +516,9 @@ impl Router { port: self.port, router_config, max_payload_size: self.max_payload_size, + wasm_middleware: self.wasm_middleware.clone(), + wasm_middleware_sha256: self.wasm_middleware_sha256.clone(), + wasm_middleware_routes: self.wasm_middleware_routes.clone(), log_dir: self.log_dir.clone(), log_level: self.log_level.clone(), service_discovery_config, diff --git a/src/main.rs b/src/main.rs index 80526a42..27b8479f 100644 --- a/src/main.rs +++ b/src/main.rs @@ -170,6 +170,19 @@ struct CliArgs { #[arg(long, default_value_t = 536870912)] // 512MB max_payload_size: usize, + /// Path to a WASM Component Model OnRequest middleware artifact + #[arg(long)] + wasm_middleware: Option, + + /// Optional SHA-256 hex digest that must match --wasm-middleware + #[arg(long)] + wasm_middleware_sha256: Option, + + /// HTTP paths that invoke the WASM middleware (repeatable). + /// Defaults to /v1/chat/completions when --wasm-middleware is set. + #[arg(long = "wasm-middleware-route", action = ArgAction::Append)] + wasm_middleware_routes: Vec, + /// Intra-node data parallel size (number of DP replicas per worker URL). When > 1, the router will create multiple worker instances per URL, one for each DP rank. #[arg(long, default_value_t = 1)] intra_node_data_parallel_size: usize, @@ -598,6 +611,9 @@ impl CliArgs { port: self.port, router_config, max_payload_size: self.max_payload_size, + wasm_middleware: self.wasm_middleware.clone(), + wasm_middleware_sha256: self.wasm_middleware_sha256.clone(), + wasm_middleware_routes: self.wasm_middleware_routes.clone(), log_dir: self.log_dir.clone(), log_level: Some(self.log_level.clone()), service_discovery_config, @@ -722,7 +738,31 @@ Provide --worker-urls or PD flags as usual.", #[cfg(test)] mod tests { - use super::split_prefill_args_from_others; + use super::*; + + #[test] + fn parses_wasm_middleware_options() { + let args = CliArgs::try_parse_from([ + "vllm-router", + "--wasm-middleware", + "/tmp/example.component.wasm", + "--wasm-middleware-sha256", + "0000000000000000000000000000000000000000000000000000000000000000", + "--wasm-middleware-route", + "/v1/chat/completions", + "--wasm-middleware-route", + "/v1/completions", + ]) + .unwrap(); + assert_eq!( + args.wasm_middleware.as_deref(), + Some("/tmp/example.component.wasm") + ); + assert_eq!( + args.wasm_middleware_routes, + vec!["/v1/chat/completions", "/v1/completions"] + ); + } #[test] fn splits_prefill_args_and_preserves_other_args() { diff --git a/src/server.rs b/src/server.rs index d7c16df6..cb2d5f4f 100644 --- a/src/server.rs +++ b/src/server.rs @@ -680,6 +680,9 @@ pub struct ServerConfig { pub port: u16, pub router_config: RouterConfig, pub max_payload_size: usize, + pub wasm_middleware: Option, + pub wasm_middleware_sha256: Option, + pub wasm_middleware_routes: Vec, pub log_dir: Option, pub log_level: Option, pub service_discovery_config: Option, @@ -698,13 +701,14 @@ pub fn build_app( cors_allowed_origins: Vec, enable_transparent_proxy: bool, ) -> Router { - build_app_with_request_tracing( + build_app_with_wasm_middleware( app_state, max_payload_size, request_id_headers, cors_allowed_origins, enable_transparent_proxy, crate::otel_trace::is_otel_enabled(), + None, ) } @@ -717,8 +721,30 @@ pub fn build_app_with_request_tracing( enable_transparent_proxy: bool, enable_request_tracing: bool, ) -> Router { - // Create routes - let protected_routes = Router::new() + build_app_with_wasm_middleware( + app_state, + max_payload_size, + request_id_headers, + cors_allowed_origins, + enable_transparent_proxy, + enable_request_tracing, + None, + ) +} + +/// Build the Axum application with optional WASM OnRequest middleware. +pub fn build_app_with_wasm_middleware( + app_state: Arc, + max_payload_size: usize, + request_id_headers: Vec, + cors_allowed_origins: Vec, + enable_transparent_proxy: bool, + enable_request_tracing: bool, + wasm_runtime: Option>, +) -> Router { + // Layer order on the request path (outer → inner): + // concurrency_limit → optional WASM OnRequest → handler + let mut protected_routes = Router::new() .route("/generate", post(generate)) .route("/inference/v1/generate", post(inference_generate)) .route("/v1/chat/completions", post(v1_chat_completions)) @@ -736,11 +762,22 @@ pub fn build_app_with_request_tracing( .route( "/v1/responses/{response_id}/input", get(v1_responses_list_input_items), - ) - .route_layer(axum::middleware::from_fn_with_state( - app_state.clone(), - middleware::concurrency_limit_middleware, + ); + + if let Some(runtime) = wasm_runtime { + protected_routes = protected_routes.route_layer(axum::middleware::from_fn_with_state( + crate::wasm_middleware::WasmRouteMiddlewareState { + runtime, + max_payload_size, + }, + crate::wasm_middleware::wasm_on_request_middleware, )); + } + + let protected_routes = protected_routes.route_layer(axum::middleware::from_fn_with_state( + app_state.clone(), + middleware::concurrency_limit_middleware, + )); let public_routes = Router::new() .route("/liveness", get(liveness)) @@ -1023,13 +1060,62 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box", + ))); + } + None + }; + + let app = build_app_with_wasm_middleware( app_state, config.max_payload_size, request_id_headers, config.router_config.cors_allowed_origins.clone(), enable_transparent_proxy, crate::otel_trace::is_otel_enabled(), + wasm_runtime, ); let addr = format!("{}:{}", config.host, config.port); diff --git a/src/wasm_middleware.rs b/src/wasm_middleware.rs new file mode 100644 index 00000000..cbfa4a92 --- /dev/null +++ b/src/wasm_middleware.rs @@ -0,0 +1,1117 @@ +//! Startup-loaded WASM OnRequest middleware runtime for vLLM Router. +//! +//! Loads one Component Model artifact, optionally verifies a SHA-256 digest, +//! and executes it under bounded workers with a fresh Store per request. +//! The default attach point is `POST /v1/chat/completions`; additional paths +//! can be configured. Failures are fail-closed (no silent bypass). + +use std::collections::HashSet; +use std::num::NonZeroUsize; +use std::path::{Path, PathBuf}; +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use axum::{ + extract::{Request, State}, + http::{header, HeaderName, HeaderValue, StatusCode}, + middleware::Next, + response::{IntoResponse, Response}, +}; +use sha2::{Digest, Sha256}; +use thiserror::Error; +use tokio::sync::{oneshot, Semaphore}; +use tracing::{debug, warn}; +use wasmtime::component::{Component, Linker, ResourceTable}; +use wasmtime::{ + Config, Engine, InstanceAllocationStrategy, PoolingAllocationConfig, Store, StoreLimits, + StoreLimitsBuilder, +}; +use wasmtime_wasi::{WasiCtx, WasiCtxBuilder, WasiCtxView, WasiView}; + +wasmtime::component::bindgen!({ + path: "wit", + world: "middleware", +}); + +use crate::wasm_middleware::vllm::router_middleware::types::{ + Action as WitAction, Header as WitHeader, ModifyAction as WitModifyAction, + Request as WitRequest, +}; + +pub const DEFAULT_MAX_INPUT_BYTES: usize = 10 * 1024 * 1024; +pub const DEFAULT_MAX_OUTPUT_BYTES: usize = 10 * 1024 * 1024; +pub const DEFAULT_MAX_MEMORY_BYTES: usize = 64 * 1024 * 1024; +pub const DEFAULT_EXECUTION_DEADLINE: Duration = Duration::from_millis(100); +pub const DEFAULT_ROUTE: &str = "/v1/chat/completions"; +const EPOCH_INTERVAL: Duration = Duration::from_millis(10); + +/// Headers plugins are never allowed to mutate on the outbound request. +const IMMUTABLE_HEADERS: &[&str] = &[ + "authorization", + "proxy-authorization", + "cookie", + "set-cookie", + "host", + "content-length", + "transfer-encoding", + "connection", + "keep-alive", + "upgrade", + "te", + "trailer", + "x-forwarded-for", + "x-forwarded-host", + "x-forwarded-proto", + "forwarded", +]; + +/// Credential / session headers never copied into the guest envelope. +/// v0.1 uses a denylist; a configurable allowlist can replace this later. +const SENSITIVE_GUEST_HEADERS: &[&str] = &[ + "authorization", + "proxy-authorization", + "cookie", + "cookie2", + "set-cookie", + "set-cookie2", +]; + +/// Paths where the WASM OnRequest layer is mounted (`protected_routes`). +/// `--wasm-middleware-route` values outside this set are rejected at startup. +pub const SUPPORTED_WASM_ROUTES: &[&str] = &[ + "/generate", + "/inference/v1/generate", + "/v1/chat/completions", + "/v1/completions", + "/rerank", + "/v1/rerank", + "/v1/responses", + "/v1/embeddings", +]; + +#[derive(Debug, Clone)] +pub struct WasmMiddlewareConfig { + pub component_path: PathBuf, + pub sha256_hex: Option, + pub routes: Vec, + pub max_input_bytes: usize, + pub max_output_bytes: usize, + pub max_memory_bytes: usize, + pub execution_deadline: Duration, + pub worker_count: usize, + pub queue_capacity: usize, +} + +impl WasmMiddlewareConfig { + pub fn from_path(path: impl Into) -> Self { + let worker_count = std::thread::available_parallelism() + .map(NonZeroUsize::get) + .unwrap_or(1) + .clamp(1, 4); + Self { + component_path: path.into(), + sha256_hex: None, + routes: vec![DEFAULT_ROUTE.to_string()], + max_input_bytes: DEFAULT_MAX_INPUT_BYTES, + max_output_bytes: DEFAULT_MAX_OUTPUT_BYTES, + max_memory_bytes: DEFAULT_MAX_MEMORY_BYTES, + execution_deadline: DEFAULT_EXECUTION_DEADLINE, + worker_count, + queue_capacity: worker_count.saturating_mul(2).max(1), + } + } + + pub fn with_sha256_hex(mut self, digest: impl Into) -> Self { + self.sha256_hex = Some(normalize_digest(&digest.into())); + self + } + + pub fn with_routes(mut self, routes: Vec) -> Self { + self.routes = routes + .into_iter() + .map(|r| r.trim().to_string()) + .filter(|r| !r.is_empty()) + .collect(); + if self.routes.is_empty() { + self.routes.push(DEFAULT_ROUTE.to_string()); + } + self + } + + pub fn validate(&self) -> Result<(), WasmMiddlewareError> { + if self.max_input_bytes == 0 || self.max_output_bytes == 0 { + return Err(WasmMiddlewareError::InvalidConfig( + "wasm body limits must be greater than zero".into(), + )); + } + if self.max_memory_bytes < 64 * 1024 { + return Err(WasmMiddlewareError::InvalidConfig( + "wasm memory limit must be at least 64KiB".into(), + )); + } + if self.execution_deadline.is_zero() { + return Err(WasmMiddlewareError::InvalidConfig( + "wasm execution deadline must be greater than zero".into(), + )); + } + if self.worker_count == 0 || self.queue_capacity == 0 { + return Err(WasmMiddlewareError::InvalidConfig( + "wasm worker_count and queue_capacity must be greater than zero".into(), + )); + } + if self.routes.is_empty() { + return Err(WasmMiddlewareError::InvalidConfig( + "wasm middleware requires at least one route".into(), + )); + } + for route in &self.routes { + if !SUPPORTED_WASM_ROUTES + .iter() + .any(|supported| *supported == route) + { + return Err(WasmMiddlewareError::InvalidConfig(format!( + "unsupported wasm middleware route `{route}`; \ + supported routes: {}", + SUPPORTED_WASM_ROUTES.join(", ") + ))); + } + } + if let Some(digest) = &self.sha256_hex { + if digest.len() != 64 || !digest.chars().all(|c| c.is_ascii_hexdigit()) { + return Err(WasmMiddlewareError::InvalidConfig( + "wasm sha256 digest must be 64 hex characters".into(), + )); + } + } + Ok(()) + } +} + +fn normalize_digest(digest: &str) -> String { + digest.trim().to_ascii_lowercase() +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum WasmAction { + Continue, + Modify { + headers_set: Vec<(String, Vec)>, + headers_add: Vec<(String, Vec)>, + headers_remove: Vec, + body: Option>, + }, + Reject { + status: u16, + }, +} + +#[derive(Debug, Error)] +pub enum WasmMiddlewareError { + #[error("invalid wasm middleware config: {0}")] + InvalidConfig(String), + #[error("failed to read wasm component {path}: {source}")] + Read { + path: PathBuf, + #[source] + source: std::io::Error, + }, + #[error("wasm component digest mismatch: expected {expected}, got {actual}")] + DigestMismatch { expected: String, actual: String }, + #[error("failed to configure wasmtime engine: {0}")] + Engine(String), + #[error("failed to compile wasm component: {0}")] + Compile(String), + #[error("failed to instantiate wasm component: {0}")] + Instantiate(String), + #[error("wasm input body exceeds limit ({limit} bytes)")] + InputTooLarge { limit: usize }, + #[error("wasm output body exceeds limit ({limit} bytes)")] + OutputTooLarge { limit: usize }, + #[error("wasm execution queue is full")] + QueueFull, + #[error("wasm execution timed out after {0:?}")] + Timeout(Duration), + #[error("wasm component trap or call failure: {0}")] + Trap(String), + #[error("invalid wasm action: {0}")] + InvalidAction(String), + #[error("wasm worker pool failed: {0}")] + Worker(String), +} + +struct HostState { + ctx: WasiCtx, + table: ResourceTable, + limits: StoreLimits, +} + +impl WasiView for HostState { + fn ctx(&mut self) -> WasiCtxView<'_> { + WasiCtxView { + ctx: &mut self.ctx, + table: &mut self.table, + } + } +} + +struct InvokeRequest { + request: WitRequest, + response: oneshot::Sender>, +} + +struct RuntimeInner { + config: WasmMiddlewareConfig, + route_set: HashSet, + queue_tx: async_channel::Sender, + _workers: Vec>, + queue_permits: Arc, +} + +#[derive(Clone)] +pub struct WasmMiddlewareRuntime { + inner: Arc, +} + +impl WasmMiddlewareRuntime { + pub fn load(config: WasmMiddlewareConfig) -> Result { + config.validate()?; + let bytes = + std::fs::read(&config.component_path).map_err(|source| WasmMiddlewareError::Read { + path: config.component_path.clone(), + source, + })?; + let actual = hex_digest(&bytes); + if let Some(expected) = config.sha256_hex.as_ref().map(|d| normalize_digest(d)) { + if expected != actual { + return Err(WasmMiddlewareError::DigestMismatch { expected, actual }); + } + } + + let mut engine_config = Config::new(); + engine_config.wasm_component_model(true); + engine_config.epoch_interruption(true); + let mut pooling = PoolingAllocationConfig::default(); + pooling.total_component_instances(32); + pooling.total_core_instances(128); + pooling.max_core_instances_per_component(8); + pooling.total_memories(64); + pooling.max_memories_per_component(2); + pooling.total_tables(64); + pooling.max_tables_per_component(4); + engine_config.allocation_strategy(InstanceAllocationStrategy::Pooling(pooling)); + + let engine = Engine::new(&engine_config) + .map_err(|err| WasmMiddlewareError::Engine(err.to_string()))?; + let component = Component::new(&engine, &bytes) + .map_err(|err| WasmMiddlewareError::Compile(err.to_string()))?; + + let mut linker = Linker::new(&engine); + wasmtime_wasi::p2::add_to_linker_sync(&mut linker) + .map_err(|err| WasmMiddlewareError::Engine(err.to_string()))?; + + { + let mut store = make_store(&engine, &config)?; + let _ = Middleware::instantiate(&mut store, &component, &linker) + .map_err(|err| WasmMiddlewareError::Instantiate(err.to_string()))?; + } + + let engine_for_epoch = engine.clone(); + std::thread::Builder::new() + .name("wasm-middleware-epoch".into()) + .spawn(move || loop { + std::thread::sleep(EPOCH_INTERVAL); + engine_for_epoch.increment_epoch(); + }) + .map_err(|err| WasmMiddlewareError::Worker(err.to_string()))?; + + let (queue_tx, queue_rx) = async_channel::bounded::(config.queue_capacity); + let mut workers = Vec::with_capacity(config.worker_count); + for worker_id in 0..config.worker_count { + let queue_rx = queue_rx.clone(); + let engine = engine.clone(); + let component = component.clone(); + let linker = linker.clone(); + let worker_config = config.clone(); + let handle = std::thread::Builder::new() + .name(format!("wasm-middleware-worker-{worker_id}")) + .spawn(move || { + worker_loop( + worker_id, + queue_rx, + engine, + component, + linker, + worker_config, + ) + }) + .map_err(|err| WasmMiddlewareError::Worker(err.to_string()))?; + workers.push(handle); + } + + let route_set = config.routes.iter().cloned().collect(); + let queue_permits = Arc::new(Semaphore::new(config.queue_capacity)); + Ok(Self { + inner: Arc::new(RuntimeInner { + config, + route_set, + queue_tx, + _workers: workers, + queue_permits, + }), + }) + } + + pub fn component_path(&self) -> &Path { + &self.inner.config.component_path + } + + pub fn config(&self) -> &WasmMiddlewareConfig { + &self.inner.config + } + + pub fn matches_route(&self, path: &str) -> bool { + self.inner.route_set.contains(path) + } + + pub async fn handle_request( + &self, + request: WitRequest, + ) -> Result { + if request.body.len() > self.inner.config.max_input_bytes { + return Err(WasmMiddlewareError::InputTooLarge { + limit: self.inner.config.max_input_bytes, + }); + } + + let permit = self + .inner + .queue_permits + .clone() + .try_acquire_owned() + .map_err(|_| WasmMiddlewareError::QueueFull)?; + + let (response_tx, response_rx) = oneshot::channel(); + self.inner + .queue_tx + .try_send(InvokeRequest { + request, + response: response_tx, + }) + .map_err(|_| WasmMiddlewareError::QueueFull)?; + + let result = response_rx + .await + .map_err(|err| WasmMiddlewareError::Worker(err.to_string()))?; + drop(permit); + result + } +} + +fn hex_digest(bytes: &[u8]) -> String { + let digest = Sha256::digest(bytes); + digest.iter().map(|b| format!("{b:02x}")).collect() +} + +fn make_store( + engine: &Engine, + config: &WasmMiddlewareConfig, +) -> Result, WasmMiddlewareError> { + let ctx = WasiCtxBuilder::new().build(); + let limits = StoreLimitsBuilder::new() + .memory_size(config.max_memory_bytes) + .trap_on_grow_failure(true) + .build(); + let mut store = Store::new( + engine, + HostState { + ctx, + table: ResourceTable::new(), + limits, + }, + ); + store.limiter(|state| &mut state.limits); + let ticks = + (config.execution_deadline.as_millis() / EPOCH_INTERVAL.as_millis().max(1)).max(1) as u64; + store.set_epoch_deadline(ticks); + Ok(store) +} + +fn worker_loop( + worker_id: usize, + queue_rx: async_channel::Receiver, + engine: Engine, + component: Component, + linker: Linker, + config: WasmMiddlewareConfig, +) { + debug!(worker_id, "wasm middleware worker started"); + while let Ok(request) = queue_rx.recv_blocking() { + let started = Instant::now(); + let result = execute_once(&engine, &component, &linker, &config, request.request); + if let Err(err) = &result { + warn!( + worker_id, + elapsed_ms = started.elapsed().as_millis() as u64, + "wasm middleware invocation failed: {err}" + ); + } + let _ = request.response.send(result); + } + debug!(worker_id, "wasm middleware worker stopped"); +} + +fn execute_once( + engine: &Engine, + component: &Component, + linker: &Linker, + config: &WasmMiddlewareConfig, + request: WitRequest, +) -> Result { + let mut store = make_store(engine, config)?; + let bindings = Middleware::instantiate(&mut store, component, linker) + .map_err(|err| WasmMiddlewareError::Instantiate(err.to_string()))?; + + let action = bindings + .vllm_router_middleware_on_request() + .call_handle(&mut store, &request) + .map_err(|err| map_call_error(err, config.execution_deadline))?; + + map_action(action, config.max_output_bytes) +} + +fn map_call_error(err: wasmtime::Error, deadline: Duration) -> WasmMiddlewareError { + let message = err.to_string(); + if message.contains("epoch") + || message.contains("interrupt") + || message.contains("deadline") + || message.contains("execution time limit") + { + WasmMiddlewareError::Timeout(deadline) + } else { + WasmMiddlewareError::Trap(message) + } +} + +fn map_action( + action: WitAction, + max_output_bytes: usize, +) -> Result { + match action { + WitAction::Continue => Ok(WasmAction::Continue), + WitAction::Reject(status) => { + let Ok(code) = StatusCode::from_u16(status) else { + return Err(WasmMiddlewareError::InvalidAction(format!( + "reject status must be a valid HTTP status code, got {status}" + ))); + }; + if !(code.is_client_error() || code.is_server_error()) { + return Err(WasmMiddlewareError::InvalidAction(format!( + "reject status must be 4xx or 5xx, got {status}" + ))); + } + Ok(WasmAction::Reject { status }) + } + WitAction::Modify(WitModifyAction { + headers_set, + headers_add, + headers_remove, + body_replace, + }) => { + let body = match body_replace { + Some(body) => { + if body.len() > max_output_bytes { + return Err(WasmMiddlewareError::OutputTooLarge { + limit: max_output_bytes, + }); + } + Some(body) + } + None => None, + }; + Ok(WasmAction::Modify { + headers_set: headers_set.into_iter().map(|h| (h.name, h.value)).collect(), + headers_add: headers_add.into_iter().map(|h| (h.name, h.value)).collect(), + headers_remove, + body, + }) + } + } +} + +fn is_immutable_header(name: &str) -> bool { + IMMUTABLE_HEADERS + .iter() + .any(|forbidden| name.eq_ignore_ascii_case(forbidden)) +} + +fn header_bytes_to_value(value: &[u8]) -> Result { + HeaderValue::from_bytes(value).map_err(|err| { + WasmMiddlewareError::InvalidAction(format!("invalid header value bytes: {err}")) + }) +} + +fn apply_header_mutations( + headers: &mut http::HeaderMap, + headers_set: &[(String, Vec)], + headers_add: &[(String, Vec)], + headers_remove: &[String], +) -> Result<(), WasmMiddlewareError> { + for name in headers_remove { + if is_immutable_header(name) { + warn!(header = %name, "ignoring wasm attempt to remove immutable header"); + continue; + } + let header_name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| { + WasmMiddlewareError::InvalidAction(format!("invalid header name {name}: {err}")) + })?; + headers.remove(header_name); + } + for (name, value) in headers_set { + if is_immutable_header(name) { + warn!(header = %name, "ignoring wasm attempt to set immutable header"); + continue; + } + let header_name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| { + WasmMiddlewareError::InvalidAction(format!("invalid header name {name}: {err}")) + })?; + headers.insert(header_name, header_bytes_to_value(value)?); + } + for (name, value) in headers_add { + if is_immutable_header(name) { + warn!(header = %name, "ignoring wasm attempt to add immutable header"); + continue; + } + let header_name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| { + WasmMiddlewareError::InvalidAction(format!("invalid header name {name}: {err}")) + })?; + headers.append(header_name, header_bytes_to_value(value)?); + } + Ok(()) +} + +fn request_id_from_headers(headers: &http::HeaderMap) -> String { + for name in [ + "x-request-id", + "x-correlation-id", + "x-trace-id", + "request-id", + ] { + if let Some(value) = headers.get(name).and_then(|v| v.to_str().ok()) { + if !value.is_empty() { + return value.to_string(); + } + } + } + String::new() +} + +fn is_sensitive_guest_header(name: &str) -> bool { + SENSITIVE_GUEST_HEADERS + .iter() + .any(|forbidden| name.eq_ignore_ascii_case(forbidden)) +} + +fn collect_wit_headers(headers: &http::HeaderMap) -> Vec { + headers + .iter() + .filter(|(name, _)| !is_sensitive_guest_header(name.as_str())) + .map(|(name, value)| WitHeader { + name: name.as_str().to_string(), + value: value.as_bytes().to_vec(), + }) + .collect() +} + +/// Axum state for the WASM OnRequest middleware layer. +#[derive(Clone)] +pub struct WasmRouteMiddlewareState { + pub runtime: Arc, + pub max_payload_size: usize, +} + +/// Fail-closed WASM OnRequest adapter. Non-configured paths pass through unchanged. +pub async fn wasm_on_request_middleware( + State(state): State, + request: Request, + next: Next, +) -> Response { + let path = request.uri().path().to_string(); + if !state.runtime.matches_route(&path) { + return next.run(request).await; + } + + let method = request.method().as_str().to_string(); + let query = request.uri().query().unwrap_or_default().to_string(); + let request_id = request + .extensions() + .get::() + .map(|id| id.0.clone()) + .filter(|id| !id.is_empty()) + .unwrap_or_else(|| request_id_from_headers(request.headers())); + let wit_headers = collect_wit_headers(request.headers()); + + let (mut parts, body) = request.into_parts(); + // Bound the body read by the WASM input limit (and never above the server payload cap). + let input_limit = state + .runtime + .config() + .max_input_bytes + .min(state.max_payload_size); + let bytes = match axum::body::to_bytes(body, input_limit).await { + Ok(bytes) => bytes, + Err(err) => { + let message = err.to_string(); + if message.contains("length limit exceeded") { + warn!( + "wasm middleware rejected oversized body (limit {} bytes)", + input_limit + ); + return StatusCode::PAYLOAD_TOO_LARGE.into_response(); + } + warn!("Failed to read request body for wasm middleware: {err}"); + return StatusCode::BAD_REQUEST.into_response(); + } + }; + + let wit_request = WitRequest { + method, + path, + query, + headers: wit_headers, + body: bytes.to_vec(), + request_id, + }; + + match state.runtime.handle_request(wit_request).await { + Ok(WasmAction::Continue) => { + let request = Request::from_parts(parts, axum::body::Body::from(bytes)); + next.run(request).await + } + Ok(WasmAction::Modify { + headers_set, + headers_add, + headers_remove, + body, + }) => { + if let Err(err) = apply_header_mutations( + &mut parts.headers, + &headers_set, + &headers_add, + &headers_remove, + ) { + warn!("wasm middleware returned invalid header mutation: {err}"); + return StatusCode::INTERNAL_SERVER_ERROR.into_response(); + } + let body = body.unwrap_or_else(|| bytes.to_vec()); + parts.headers.remove(header::CONTENT_LENGTH); + let request = Request::from_parts(parts, axum::body::Body::from(body)); + next.run(request).await + } + Ok(WasmAction::Reject { status }) => { + warn!(status, "wasm middleware rejected request"); + StatusCode::from_u16(status) + .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR) + .into_response() + } + Err(WasmMiddlewareError::InputTooLarge { limit }) => { + warn!(limit, "wasm middleware input exceeds configured limit"); + StatusCode::PAYLOAD_TOO_LARGE.into_response() + } + Err(WasmMiddlewareError::QueueFull) => { + warn!("wasm middleware execution queue is full"); + StatusCode::SERVICE_UNAVAILABLE.into_response() + } + Err(err) => { + warn!("wasm middleware failed closed: {err}"); + StatusCode::INTERNAL_SERVER_ERROR.into_response() + } + } +} + +/// Resolve the example component path for tests (if already built). +pub fn example_component_artifact_path() -> Option { + let example_dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("examples/wasm_middleware"); + [ + example_dir.join("wasm_middleware_example.component.wasm"), + example_dir.join("target/wasm32-wasip2/release/wasm_middleware_example.wasm"), + ] + .into_iter() + .find(|p| p.is_file()) +} + +/// Build the example guest via `examples/wasm_middleware/build.sh` once per process. +/// +/// Tests should call this instead of requiring contributors to build the artifact +/// manually before `cargo test`. +pub fn ensure_example_component_artifact() -> Option { + use std::sync::OnceLock; + static ARTIFACT: OnceLock> = OnceLock::new(); + ARTIFACT + .get_or_init(|| { + let example_dir = + PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("examples/wasm_middleware"); + let build_sh = example_dir.join("build.sh"); + if !build_sh.is_file() { + eprintln!( + "wasm middleware tests: missing build script at {}", + build_sh.display() + ); + return None; + } + let status = std::process::Command::new("bash") + .arg(&build_sh) + .current_dir(&example_dir) + .status(); + match status { + Ok(code) if code.success() => example_component_artifact_path(), + Ok(code) => { + eprintln!( + "wasm middleware tests: build.sh exited with {code}; \ + install the wasm32-wasip2 target or check the example crate" + ); + example_component_artifact_path() + } + Err(err) => { + eprintln!("wasm middleware tests: failed to run build.sh: {err}"); + example_component_artifact_path() + } + } + }) + .clone() +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::{body::Body, routing::post, Json, Router}; + use http_body_util::BodyExt; + use serde_json::{json, Value}; + use std::time::Instant; + use tower::ServiceExt; + + fn require_component_path() -> PathBuf { + ensure_example_component_artifact().unwrap_or_else(|| { + panic!( + "failed to build/find the example WASM component via \ + `examples/wasm_middleware/build.sh` (wasm32-wasip2 target required)" + ) + }) + } + + fn wit_request(body: &[u8]) -> WitRequest { + WitRequest { + method: "POST".into(), + path: "/v1/chat/completions".into(), + query: String::new(), + headers: Vec::new(), + body: body.to_vec(), + request_id: "test".into(), + } + } + + #[tokio::test] + async fn example_plugin_sets_marker_header() { + let path = require_component_path(); + let runtime = WasmMiddlewareRuntime::load(WasmMiddlewareConfig::from_path(path)) + .expect("load wasm runtime"); + let action = runtime + .handle_request(wit_request(br#"{"messages":[]}"#)) + .await + .expect("invoke"); + match action { + WasmAction::Modify { + headers_set, body, .. + } => { + assert!(body.is_none()); + assert!(headers_set.iter().any(|(name, value)| { + name.eq_ignore_ascii_case("x-wasm-middleware") && value == b"example" + })); + } + other => panic!("expected Modify, got {other:?}"), + } + } + + #[tokio::test] + async fn example_plugin_rejects_marker_body() { + let path = require_component_path(); + let runtime = WasmMiddlewareRuntime::load(WasmMiddlewareConfig::from_path(path)) + .expect("load wasm runtime"); + let action = runtime + .handle_request(wit_request(br#"{"note":"__wasm_reject__"}"#)) + .await + .expect("invoke"); + assert_eq!(action, WasmAction::Reject { status: 400 }); + } + + #[test] + fn digest_mismatch_fails_startup() { + let path = require_component_path(); + match WasmMiddlewareRuntime::load( + WasmMiddlewareConfig::from_path(path).with_sha256_hex( + "0000000000000000000000000000000000000000000000000000000000000000", + ), + ) { + Err(WasmMiddlewareError::DigestMismatch { .. }) => {} + Ok(_) => panic!("expected digest mismatch"), + Err(other) => panic!("expected digest mismatch, got {other}"), + } + } + + #[test] + fn input_too_large_is_rejected_before_queue() { + let path = require_component_path(); + let mut config = WasmMiddlewareConfig::from_path(path); + config.max_input_bytes = 8; + let runtime = WasmMiddlewareRuntime::load(config).expect("load"); + let err = tokio::runtime::Runtime::new() + .unwrap() + .block_on(runtime.handle_request(wit_request(b"0123456789"))) + .expect_err("too large"); + assert!(matches!(err, WasmMiddlewareError::InputTooLarge { .. })); + } + + async fn echo_headers(request: Request) -> Json { + let marker = request + .headers() + .get("x-wasm-middleware") + .and_then(|v| v.to_str().ok()) + .unwrap_or("") + .to_string(); + Json(json!({ "x-wasm-middleware": marker })) + } + + fn test_app(runtime: Arc, max_payload_size: usize) -> Router { + Router::new() + .route("/v1/chat/completions", post(echo_headers)) + .route("/v1/completions", post(echo_headers)) + .layer(axum::middleware::from_fn_with_state( + WasmRouteMiddlewareState { + runtime, + max_payload_size, + }, + wasm_on_request_middleware, + )) + } + + async fn post_json(app: Router, uri: &str, body: &str) -> Response { + app.oneshot( + Request::builder() + .method("POST") + .uri(uri) + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from(body.to_string())) + .unwrap(), + ) + .await + .unwrap() + } + + #[tokio::test] + async fn axum_layer_applies_on_configured_route() { + let path = require_component_path(); + let runtime = Arc::new( + WasmMiddlewareRuntime::load(WasmMiddlewareConfig::from_path(path)).expect("load"), + ); + let app = test_app(runtime, 1024 * 1024); + + let chat = post_json(app.clone(), "/v1/chat/completions", r#"{"messages":[]}"#).await; + assert_eq!(chat.status(), StatusCode::OK); + let chat_body = chat.into_body().collect().await.unwrap().to_bytes(); + let chat_json: Value = serde_json::from_slice(&chat_body).unwrap(); + assert_eq!(chat_json["x-wasm-middleware"], "example"); + + let other = post_json(app, "/v1/completions", r#"{"prompt":"hi"}"#).await; + assert_eq!(other.status(), StatusCode::OK); + let other_body = other.into_body().collect().await.unwrap().to_bytes(); + let other_json: Value = serde_json::from_slice(&other_body).unwrap(); + assert_eq!(other_json["x-wasm-middleware"], ""); + } + + #[tokio::test] + async fn concurrent_requests_are_handled() { + let path = require_component_path(); + let mut config = WasmMiddlewareConfig::from_path(path); + config.worker_count = 4; + config.queue_capacity = 64; + let runtime = Arc::new(WasmMiddlewareRuntime::load(config).expect("load")); + + let mut tasks = Vec::new(); + for i in 0..32 { + let runtime = runtime.clone(); + tasks.push(tokio::spawn(async move { + let body = format!(r#"{{"messages":[],"n":{i}}}"#); + runtime.handle_request(wit_request(body.as_bytes())).await + })); + } + + for task in tasks { + let action = task.await.expect("join").expect("invoke"); + match action { + WasmAction::Modify { headers_set, .. } => { + assert!(headers_set.iter().any(|(name, value)| { + name.eq_ignore_ascii_case("x-wasm-middleware") && value == b"example" + })); + } + other => panic!("expected Modify, got {other:?}"), + } + } + } + + #[tokio::test] + async fn infinite_loop_plugin_times_out_without_blocking_later_requests() { + let path = require_component_path(); + let mut config = WasmMiddlewareConfig::from_path(path); + config.execution_deadline = Duration::from_millis(100); + config.worker_count = 1; + config.queue_capacity = 2; + let runtime = WasmMiddlewareRuntime::load(config).expect("load"); + + let started = Instant::now(); + let err = runtime + .handle_request(wit_request(br#"{"note":"__wasm_loop__"}"#)) + .await + .expect_err("loop should hit the execution deadline"); + // Epoch interruption may surface as Timeout or a generic Trap depending on Wasmtime. + assert!( + matches!( + err, + WasmMiddlewareError::Timeout(_) | WasmMiddlewareError::Trap(_) + ), + "expected Timeout or Trap, got {err}" + ); + assert!( + started.elapsed() < Duration::from_secs(2), + "epoch timeout took too long: {:?}", + started.elapsed() + ); + + // WASM worker must still accept subsequent requests after the trap/timeout. + let action = runtime + .handle_request(wit_request(br#"{"messages":[]}"#)) + .await + .expect("follow-up invoke"); + match action { + WasmAction::Modify { headers_set, .. } => { + assert!(headers_set.iter().any(|(name, value)| { + name.eq_ignore_ascii_case("x-wasm-middleware") && value == b"example" + })); + } + other => panic!("expected Modify after timeout recovery, got {other:?}"), + } + } + + #[tokio::test] + async fn middleware_maps_each_response_status_code() { + let path = require_component_path(); + + // Reject statuses returned by the guest. + for status in [400_u16, 401, 403, 404, 429, 500, 503] { + let runtime = Arc::new( + WasmMiddlewareRuntime::load(WasmMiddlewareConfig::from_path(&path)).expect("load"), + ); + let app = test_app(runtime, 1024 * 1024); + let body = format!(r#"{{"note":"__wasm_reject_{status}__"}}"#); + let response = post_json(app, "/v1/chat/completions", &body).await; + assert_eq!( + response.status(), + StatusCode::from_u16(status).unwrap(), + "reject marker should map to HTTP {status}" + ); + } + + // Input too large → 413 + { + let mut config = WasmMiddlewareConfig::from_path(&path); + config.max_input_bytes = 32; + let runtime = Arc::new(WasmMiddlewareRuntime::load(config).expect("load")); + let app = test_app(runtime, 1024 * 1024); + let oversized = format!(r#"{{"pad":"{}"}}"#, "x".repeat(64)); + let response = post_json(app, "/v1/chat/completions", &oversized).await; + assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE); + } + + // Execution timeout / trap → 500 + { + let mut config = WasmMiddlewareConfig::from_path(&path); + config.execution_deadline = Duration::from_millis(100); + let runtime = Arc::new(WasmMiddlewareRuntime::load(config).expect("load")); + let app = test_app(runtime, 1024 * 1024); + let response = + post_json(app, "/v1/chat/completions", r#"{"note":"__wasm_loop__"}"#).await; + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + } + + // Saturated queue → 503 + { + let mut config = WasmMiddlewareConfig::from_path(&path); + config.worker_count = 1; + config.queue_capacity = 1; + config.execution_deadline = Duration::from_millis(400); + let runtime = Arc::new(WasmMiddlewareRuntime::load(config).expect("load")); + let app = test_app(runtime, 1024 * 1024); + + let busy = tokio::spawn(post_json( + app.clone(), + "/v1/chat/completions", + r#"{"note":"__wasm_loop__"}"#, + )); + // Let the looping request take the only queue/worker slot. + tokio::time::sleep(Duration::from_millis(50)).await; + let rejected = post_json(app, "/v1/chat/completions", r#"{"messages":[]}"#).await; + assert_eq!(rejected.status(), StatusCode::SERVICE_UNAVAILABLE); + let busy_status = busy.await.expect("join").status(); + assert_eq!(busy_status, StatusCode::INTERNAL_SERVER_ERROR); + } + } + + #[test] + fn reject_status_must_be_valid_http_error() { + assert!(matches!( + map_action(WitAction::Reject(400), 1024), + Ok(WasmAction::Reject { status: 400 }) + )); + assert!(matches!( + map_action(WitAction::Reject(200), 1024), + Err(WasmMiddlewareError::InvalidAction(_)) + )); + assert!(matches!( + map_action(WitAction::Reject(65535), 1024), + Err(WasmMiddlewareError::InvalidAction(_)) + )); + } + + #[test] + fn headers_remove_invalid_name_fails_closed() { + let mut headers = http::HeaderMap::new(); + headers.insert("x-test", HeaderValue::from_static("1")); + let err = apply_header_mutations(&mut headers, &[], &[], &["bad name".into()]) + .expect_err("invalid remove name"); + assert!(matches!(err, WasmMiddlewareError::InvalidAction(_))); + } + + #[test] + fn sensitive_headers_are_not_copied_to_guest() { + let mut headers = http::HeaderMap::new(); + headers.insert(header::AUTHORIZATION, HeaderValue::from_static("secret")); + headers.insert( + header::CONTENT_TYPE, + HeaderValue::from_static("application/json"), + ); + let wit = collect_wit_headers(&headers); + assert!(wit + .iter() + .any(|h| h.name.eq_ignore_ascii_case("content-type"))); + assert!(!wit + .iter() + .any(|h| h.name.eq_ignore_ascii_case("authorization"))); + } + + #[test] + fn unsupported_route_fails_validation() { + let err = WasmMiddlewareConfig::from_path("/tmp/unused.component.wasm") + .with_routes(vec!["/health".into()]) + .validate() + .expect_err("unsupported route"); + assert!(matches!(err, WasmMiddlewareError::InvalidConfig(_))); + } +} diff --git a/tests/common/test_app.rs b/tests/common/test_app.rs index d88998ca..4ecffcea 100644 --- a/tests/common/test_app.rs +++ b/tests/common/test_app.rs @@ -5,7 +5,8 @@ use vllm_router_rs::{ config::RouterConfig, otel_trace, routers::RouterTrait, - server::{build_app_with_request_tracing, AppContext, AppState}, + server::{build_app_with_wasm_middleware, AppContext, AppState}, + wasm_middleware::WasmMiddlewareRuntime, }; /// Create a test Axum application using the actual server's build_app function @@ -25,6 +26,18 @@ pub fn create_test_app_with_tracing( client: Client, router_config: &RouterConfig, enable_request_tracing: bool, +) -> Router { + create_test_app_with_wasm(router, client, router_config, enable_request_tracing, None) +} + +/// Create a test Axum application with optional WASM OnRequest middleware. +#[allow(dead_code)] +pub fn create_test_app_with_wasm( + router: Arc, + client: Client, + router_config: &RouterConfig, + enable_request_tracing: bool, + wasm_runtime: Option>, ) -> Router { // Create AppContext let app_context = Arc::new( @@ -56,13 +69,14 @@ pub fn create_test_app_with_tracing( ] }); - // Use the actual server's build_app function - build_app_with_request_tracing( + // Use the actual server's build_app function (with optional WASM layer) + build_app_with_wasm_middleware( app_state, router_config.max_payload_size, request_id_headers, router_config.cors_allowed_origins.clone(), true, // enable_transparent_proxy enable_request_tracing, + wasm_runtime, ) } diff --git a/tests/test_wasm_middleware.rs b/tests/test_wasm_middleware.rs new file mode 100644 index 00000000..7a78dcbf --- /dev/null +++ b/tests/test_wasm_middleware.rs @@ -0,0 +1,184 @@ +//! Integration tests for WASM OnRequest middleware with mock HTTP workers. +//! +//! Builds the example component via `examples/wasm_middleware/build.sh` on demand. + +mod common; + +use axum::{ + body::Body, + extract::Request, + http::{header::CONTENT_TYPE, StatusCode}, +}; +use common::mock_worker::{self, HealthStatus, MockWorker, MockWorkerConfig, WorkerType}; +use common::test_app::create_test_app_with_wasm; +use http_body_util::BodyExt; +use reqwest::Client; +use serde_json::json; +use std::path::PathBuf; +use std::sync::Arc; +use tower::ServiceExt; +use vllm_router_rs::config::{RouterConfig, RoutingMode}; +use vllm_router_rs::routers::RouterFactory; +use vllm_router_rs::wasm_middleware::{ + ensure_example_component_artifact, WasmMiddlewareConfig, WasmMiddlewareRuntime, +}; + +fn require_component_path() -> PathBuf { + ensure_example_component_artifact().unwrap_or_else(|| { + panic!( + "failed to build/find the example WASM component via \ + `examples/wasm_middleware/build.sh` (wasm32-wasip2 target required)" + ) + }) +} + +fn header_values<'a>( + captured: &'a mock_worker::CapturedRequest, + name: &str, +) -> Option<&'a Vec> { + captured + .headers + .iter() + .find(|(k, _)| k.eq_ignore_ascii_case(name)) + .map(|(_, v)| v) +} + +#[tokio::test] +async fn wasm_middleware_modify_reject_and_path_isolation() { + let component = require_component_path(); + let wasm_runtime = Arc::new( + WasmMiddlewareRuntime::load(WasmMiddlewareConfig::from_path(component)) + .expect("load wasm runtime"), + ); + + let mut worker = MockWorker::new(MockWorkerConfig { + port: 0, + worker_type: WorkerType::Regular, + health_status: HealthStatus::Healthy, + response_delay_ms: 0, + fail_rate: 0.0, + }); + let worker_url = worker.start().await.expect("start mock worker"); + let worker_port: u16 = worker_url + .rsplit(':') + .next() + .unwrap() + .parse() + .expect("parse worker port"); + mock_worker::clear_captured_requests(worker_port); + + let config = RouterConfig { + mode: RoutingMode::Regular { + worker_urls: vec![worker_url], + }, + worker_startup_timeout_secs: 1, + worker_startup_check_interval_secs: 1, + ..Default::default() + }; + let app_context = common::create_test_context(config.clone()); + let router = RouterFactory::create_router(&app_context) + .await + .expect("create router"); + tokio::time::sleep(tokio::time::Duration::from_millis(300)).await; + + let app = create_test_app_with_wasm( + Arc::from(router), + Client::new(), + &config, + false, + Some(wasm_runtime), + ); + + // Modify: chat should forward with the example header. + let chat = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header(CONTENT_TYPE, "application/json") + .body(Body::from( + json!({ + "model": "mock-model", + "messages": [{"role": "user", "content": "hi"}] + }) + .to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(chat.status(), StatusCode::OK); + let _ = chat.into_body().collect().await.unwrap().to_bytes(); + + let captured = mock_worker::get_captured_requests(worker_port); + let chat_req = captured + .iter() + .find(|r| r.path == "/v1/chat/completions") + .expect("mock worker should receive chat request"); + let marker = header_values(chat_req, "x-wasm-middleware") + .and_then(|v| v.first()) + .map(String::as_str); + assert_eq!(marker, Some("example")); + + // Reject: fail closed at the Router before the worker. + mock_worker::clear_captured_requests(worker_port); + let rejected = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header(CONTENT_TYPE, "application/json") + .body(Body::from( + json!({ + "model": "mock-model", + "messages": [{"role": "user", "content": "nope"}], + "note": "__wasm_reject__" + }) + .to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(rejected.status(), StatusCode::BAD_REQUEST); + assert!( + mock_worker::get_captured_requests(worker_port).is_empty(), + "rejected requests must not reach the worker" + ); + + // Path isolation: completions is not attached by default. + mock_worker::clear_captured_requests(worker_port); + let completions = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/completions") + .header(CONTENT_TYPE, "application/json") + .body(Body::from( + json!({ + "model": "mock-model", + "prompt": "hi" + }) + .to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(completions.status(), StatusCode::OK); + let _ = completions.into_body().collect().await.unwrap().to_bytes(); + + let captured = mock_worker::get_captured_requests(worker_port); + let completion_req = captured + .iter() + .find(|r| r.path == "/v1/completions") + .expect("mock worker should receive completions request"); + assert!( + header_values(completion_req, "x-wasm-middleware").is_none(), + "default WASM attach point must not modify /v1/completions" + ); + + worker.stop().await; +} diff --git a/wit/router-middleware.wit b/wit/router-middleware.wit new file mode 100644 index 00000000..38c0a7c3 --- /dev/null +++ b/wit/router-middleware.wit @@ -0,0 +1,39 @@ +package vllm:router-middleware@0.1.0; + +interface types { + record header { + name: string, + value: list, + } + + record request { + method: string, + path: string, + query: string, + headers: list
, + body: list, + request-id: string, + } + + record modify-action { + headers-set: list
, + headers-add: list
, + headers-remove: list, + body-replace: option>, + } + + variant action { + continue, + reject(u16), + modify(modify-action), + } +} + +interface on-request { + use types.{request, action}; + handle: func(req: request) -> action; +} + +world middleware { + export on-request; +}