Compare commits

...
Author SHA1 Message Date
dependabot[bot]andGitHub 95dcefaaa5 Bump actions/setup-python from 6.1.0 to 6.3.0
Bumps [actions/setup-python](https://github.com/actions/setup-python) from 6.1.0 to 6.3.0.
- [Release notes](https://github.com/actions/setup-python/releases)
- [Commits](https://github.com/actions/setup-python/compare/83679a892e2d95755f2dac6acb0bfd1e9ac5d548...ece7cb06caefa5fff74198d8649806c4678c61a1)

---
updated-dependencies:
- dependency-name: actions/setup-python
  dependency-version: 6.3.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
...

Signed-off-by: dependabot[bot] <support@github.com>
2026-06-30 12:18:49 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
536047755e Bump actions/checkout from 6.0.1 to 7.0.0 (#33057)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-06-30 13:16:20 +01:00
1907d3854a [Bugfix] Reject negative values for max_logprobs and long_prefill_token_threshold (#44002)
Signed-off-by: jwzheng96 <jianweizheng@pku.edu.cn>
Signed-off-by: JianweiZheng <32029023+jwzheng96@users.noreply.github.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Wentao Ye <44945378+yewentao256@users.noreply.github.com>
Co-authored-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-06-30 13:01:03 +01:00
Chaojun ZhangandGitHub ea9ddf59fc [XPU][CI] Enable shared loader test (#45977)
Signed-off-by: Chaojun Zhang <chaojun.zhang@intel.com>
2026-06-30 11:20:33 +00:00
8cf7c4d8ad [Attention Backend] add HPC-Ops Attention backend (#46020)
Signed-off-by: chengvjiang <chengvjiang@tencent.com>
Co-authored-by: chengvjiang <chengvjiang@tencent.com>
Co-authored-by: Andreas Karatzas <akaratza@amd.com>
2026-06-30 18:17:43 +08:00
8e9d70fdd5 [Kernel][XPU] Adjust kernel unit tests for XPU (#45140)
Signed-off-by: Dobrzyniewicz, Agata <agata.dobrzyniewicz@intel.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-06-30 09:57:27 +00:00
Juan Pérez de AlgabaandGitHub 364ee36af1 fix(security): prevent image decompression bomb OOM denial of service (#47010)
Signed-off-by: jperezde <jperezde@redhat.com>
2026-06-30 09:39:22 +00:00
Nicolò LucchesiandGitHub 06fae69114 [Misc] Mistral label alert (#47132)
Signed-off-by: NickLucche <nicolo.lucchesi@mistral.ai>
2026-06-30 09:02:07 +00:00
14f8660a18 [CI/Build] Add CPU test dependency pre-commit hooks (#47032)
Signed-off-by: jiang1.li <jiang1.li@intel.com>
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
Co-authored-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-06-30 07:59:13 +00:00
aed541def4 [Bugfix][Responses] Set completed status for Harmony function calls (#46945)
Signed-off-by: amanambak <aman.paswan@ambak.com>
Co-authored-by: amanambak <aman.paswan@ambak.com>
Co-authored-by: Chauncey <chaunceyjiang@gmail.com>
2026-06-30 07:55:14 +00:00
ChaunceyGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2bc20e8aba [Frontend] Add Streaming Parser Engine and new Kimi k2.5/k2.6/k2.7 Parser (#46610)
Signed-off-by: chaunceyjiang <chaunceyjiang@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-30 07:53:17 +00:00
Chaojun ZhangandGitHub 8cc242335d [XPU] Optimize XPU worker shutdown logic to prevent resource leak (#46433)
Signed-off-by: Chaojun Zhang <chaojun.zhang@intel.com>
2026-06-30 15:27:21 +08:00
Andreas KaratzasandGitHub ba22cb6765 [ROCm][Ray][CI] Keep assigned GPU visible for weight transfer (#47000)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-06-30 14:59:18 +08:00
Uros MarkovicandGitHub 81bcced482 [Bugfix][ROCm] Preserve MoE weight padding for unquantized Triton path (#46381)
Signed-off-by: Uros Markovic <umarkovi@amd.com>
2026-06-30 14:47:57 +08:00
Kunshang JiandGitHub fb42e5219e [Platform] Replace torch.cuda.mem_get_info with torch.accelerator.get_memory_info (#44825)
Signed-off-by: Kunshang Ji <kunshang.ji@intel.com>
Signed-off-by: Kunshang Ji <jikunshang95@gmail.com>
2026-06-30 14:39:52 +08:00
Dakai AnandGitHub 0feca7ffa8 PD disagg with Mooncake Connector: GDN support (Qwen3.5) and MLA support (Deepseek-V4-Flash) (#46807) 2026-06-29 23:29:04 -07:00
97b5ce5c39 [Bugfix] Raise VLLMValidationError for non-integer logit_bias keys (#46612)
Signed-off-by: muhammadfawaz1 <135441198+muhammadfawaz1@users.noreply.github.com>
Co-authored-by: Mahad Durrani <114791389+mahadrehmann@users.noreply.github.com>
2026-06-30 06:18:59 +00:00
Andreas KaratzasandGitHub 4236514098 [ROCm][CI][Multimodal] Use ROCm-aware FA availability check for Unlimited-OCR (#47004)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-06-30 14:03:13 +08:00
Blas Rodriguez IrizarandGitHub e45c8a9f4b [Rust Frontend] Start current wave for a stale DP FirstRequest (#46833)
Signed-off-by: Blas Rodriguez Irizar <rodrigblas@gmail.com>
2026-06-30 05:13:09 +00:00
Wei ZhaoandGitHub b153dd3f28 [Bugfix] Use larger workspace size for Flashinfer MLA LSE (#47074)
Signed-off-by: wzhao18 <wzhao18.sz@gmail.com>
2026-06-29 22:11:03 -07:00
ReidandGitHub 930f8dc0a1 [Bugfix][Rust Frontend] Reject prompt_logprobs for streaming generate (#46839)
Signed-off-by: reidliu41 <reid201711@gmail.com>
2026-06-30 05:10:07 +00:00
ReidandGitHub a16dbd5b85 [Rust Frontend] Avoid LoRA registry scans without active LoRA requests (#47040)
Signed-off-by: reidliu41 <reid201711@gmail.com>
2026-06-30 04:58:19 +00:00
bec232a914 Secondary tier implementation for PD disaggregation (#42285)
Signed-off-by: Liran Schour <lirans@il.ibm.com>
Signed-off-by: liranschour <liranschour@users.noreply.github.com>
Co-authored-by: Or Ozeri <or@ozery.com>
Co-authored-by: Or Ozeri <oro@il.ibm.com>
2026-06-30 07:51:44 +03:00
b5c9e1ac33 [LoRA] Add language-backbone LoRA support for MiniCPM-V 4.6 (#46740)
Signed-off-by: linitra24 <Joy25810@foxmail.com>
Co-authored-by: Jee Jee Li <pandaleefree@gmail.com>
2026-06-30 04:19:31 +00:00
ae2c4f3db7 [XPU][UT]Fix xpu pass_config.fuse_norm_quant assert issue (#46804)
Signed-off-by: Lai, Yejing <yejing.lai@intel.com>
Co-authored-by: Kunshang Ji <kunshang.ji@intel.com>
2026-06-29 21:13:44 -07:00
ganeshandGitHub fca432e60a [Bugfix] Propagate default stop_token_ids to per-request SamplingParams (#35076)
Signed-off-by: sriganesh123 <arjulasriganesh@gmail.com>
2026-06-30 12:10:09 +08:00
hclGitHubmergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
af1ee8c475 fix(config): reject negative max_logprobs (except -1) and long_prefill_token_threshold (#44070)
Signed-off-by: Chenglun Hu <chenglunhu@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-30 04:02:36 +00:00
5b4cb69523 [Bugfix][MLA] Fix LSE log-base mismatch in DCP + FlashInfer MLA decode (#47079)
Signed-off-by: girasoley <girasoleyang@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-06-29 19:15:02 -07:00
9fc0c08026 [ROCm][CI] Make tests/v1/shutdown an importable package (#47085)
Signed-off-by: pei.zhang <pei.zhang@amd.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-06-29 21:01:27 -05:00
f2b5fabb23 [ROCm][CI] Move LM Eval Large Models (8 GPUs) to mi300 pool (#47094)
Signed-off-by: pei.zhang <pei.zhang@amd.com>
Co-authored-by: Andreas Karatzas <akaratza@amd.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-06-29 20:59:08 -05:00
Tahsin TunanGitHubBugen Zhaomergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
b8cb75b149 [Rust Frontend] Add static HTTPS and mTLS support for HTTP and gRPC (#45890)
Co-authored-by: Bugen Zhao <i@bugenzhao.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Signed-off-by: Tahsin Tunan <tahsintunan@gmail.com>
Signed-off-by: Bugen Zhao <i@bugenzhao.com>
2026-06-30 01:45:59 +00:00
Thien TranandGitHub 43916891b2 [GDN] Improve kkt kernel of CuteDSL prefill backend (#46346)
Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>
2026-06-29 18:34:18 -07:00
cda05ee8c4 [Bugfix][Reasoning] Fix thinking_token_budget not enforced on re-entry after forced end (#43757)
Signed-off-by: Ashwin Giridharan <girida@amazon.com>
Signed-off-by: Cursor Agent <cursor-agent@cursor.com>
Co-authored-by: Cursor Agent <cursor-agent@cursor.com>
Co-authored-by: Simon Mo <simon.mo@hey.com>
2026-06-30 01:04:25 +00:00
weishuandGitHub 77654d080c [KVTransfer] MultiConnector: merge kv_transfer_params dicts across connectors (#46777)
Signed-off-by: deng451e <838677410@qq.com>
2026-06-30 00:25:05 +00:00
Wentao YeandGitHub 75698e60b3 [Bug] Fix sparse attention issue for GLM5.2 non-torch compile path (#47083)
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-06-29 15:45:53 -07:00
Andreas KaratzasandGitHub 8632c884dc [ROCm][CI] Use spawn around the threaded OTLP test (#47003)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
2026-06-29 16:34:05 -05:00
c3734e8334 [CI][Bugfix] Add cohere_melody to ROCm test requirements (#47072)
Signed-off-by: pei.zhang <pei.zhang@amd.com>
Co-authored-by: Claude <noreply@anthropic.com>
2026-06-29 16:29:47 -05:00
53f7553f09 [ROCm][DeepEP] Stabilize high-throughput DBO for DP+EP (#46990)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
Signed-off-by: Andreas Karatzas <Andreas.Karatzas@amd.com>
Co-authored-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
2026-06-29 14:28:02 -07:00
4eb227992a [ROCm][CI] Make memory sampling less racy in tests and sleep mode (#45490)
Signed-off-by: Andreas Karatzas <akaratza@amd.com>
Signed-off-by: Codex <codex@example.invalid>
Co-authored-by: Codex <codex@example.invalid>
2026-06-29 14:26:41 -07:00
Micah WilliamsonandGitHub ebcf511ec3 [ROCm][CI] Soft Fail Spec Decode Ngram + Suffix and Entrypoints Integration (LLM) AMD Mirrors (#47067)
Signed-off-by: Micah Williamson <micah.williamson@amd.com>
2026-06-29 16:24:08 -05:00
Matthew BonanniandGitHub 8fc1b2d046 Fix FA4 dynamic_causal for full attention layers (#46659)
Signed-off-by: Matthew Bonanni <mbonanni@redhat.com>
2026-06-29 14:23:34 -07:00
Harry MellorandGitHub 5316638a5e Fix transient dependency issues caused by requirements/common.txt (#47015)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
2026-06-29 14:20:33 -07:00
zhrrrandGitHub 61ab70ec3b [Model Runner V2] support mamba hybrid models align prefix cache (#42406)
Signed-off-by: zhuhaoran <zhuhaoran.zhr@alibaba-inc.com>
2026-06-29 14:09:16 -07:00
155 changed files with 16334 additions and 1748 deletions
@@ -23,4 +23,5 @@ steps:
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
export VLLM_WORKER_MULTIPROC_METHOD=spawn &&
pytest -v -s basic_correctness/test_cpu_offload.py &&
pytest -v -s basic_correctness/test_mem.py::test_end_to_end'
+1 -1
View File
@@ -128,10 +128,10 @@ steps:
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
export VLLM_WORKER_MULTIPROC_METHOD=spawn &&
(pytest -v -s lora/test_mixtral.py --deselect="tests/lora/test_mixtral.py::test_mixtral_lora[4]" || true) &&
pytest -v -s lora/test_quant_model.py --deselect="tests/lora/test_quant_model.py::test_quant_model_lora[model0]" --deselect="tests/lora/test_quant_model.py::test_quant_model_lora[model1]" --deselect="tests/lora/test_quant_model.py::test_quant_model_tp_equality[model0]" &&
pytest -v -s lora/test_transformers_model.py &&
pytest -v -s lora/test_chatglm3_tp.py &&
pytest -v -s lora/test_llama_tp.py::test_llama_lora &&
pytest -s -v lora/test_minicpmv_tp.py'
- label: LoRA Multimodal
@@ -0,0 +1,27 @@
group: Models - Distributed
depends_on:
- image-build-xpu
steps:
- label: Distributed Model Tests (2 GPUs)
key: distributed-model-tests-2-gpus
timeout_in_minutes: 50
device: intel_gpu
agent_tags:
label: production
gpu: 2+
mem: 24+
no_plugin: true
working_dir: "."
env:
REGISTRY: "public.ecr.aws/q9t5s3a7"
REPO: "vllm-ci-test-repo"
VLLM_TEST_DEVICE: "xpu"
source_file_dependencies:
- vllm/model_executor/model_loader/sharded_state_loader.py
- vllm/model_executor/models/
- tests/model_executor/model_loader/test_sharded_state_loader.py
commands:
- >-
bash .buildkite/scripts/hardware_ci/run-intel-test.sh
'cd tests &&
pytest -v -s model_executor/model_loader/test_sharded_state_loader.py -m "not slow_test"'
+21 -21
View File
@@ -1198,6 +1198,27 @@ steps:
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
- pytest -s -v evals/gsm8k/test_gsm8k_correctness.py --config-list-file=configs/models-mi3xx.txt
- label: ROCm LM Eval Large Models (8 GPUs) # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi300]
agent_pool: mi300_8
optional: true
num_gpus: 8
working_dir: "/vllm-workspace/.buildkite/lm-eval-harness"
source_file_dependencies:
- vllm/model_executor/models/
- vllm/model_executor/model_loader/
- vllm/model_executor/layers/quantization/
- vllm/v1/attention/backends/
- vllm/v1/attention/selector.py
- vllm/model_executor/layers/layernorm.py
- csrc/
- vllm/_aiter_ops.py
- vllm/platforms/rocm.py
commands:
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-large-rocm.txt --tp-size=8
#--------------------------------------------------------- mi300 · examples ----------------------------------------------------------#
- label: Examples # TBD
@@ -2392,27 +2413,6 @@ steps:
- export VLLM_USE_DEEP_GEMM=0
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-large-rocm-fp8.txt --tp-size=4
- label: ROCm LM Eval Large Models (8 GPUs) # TBD
timeout_in_minutes: 180
mirror_hardwares: [amdexperimental, amdproduction, amdgfx942nightly, amdmi325]
agent_pool: mi325_8
optional: true
num_gpus: 8
working_dir: "/vllm-workspace/.buildkite/lm-eval-harness"
source_file_dependencies:
- vllm/model_executor/models/
- vllm/model_executor/model_loader/
- vllm/model_executor/layers/quantization/
- vllm/v1/attention/backends/
- vllm/v1/attention/selector.py
- vllm/model_executor/layers/layernorm.py
- csrc/
- vllm/_aiter_ops.py
- vllm/platforms/rocm.py
commands:
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
- pytest -s -v test_lm_eval_correctness.py --config-list-file=configs/models-large-rocm.txt --tp-size=8
#----------------------------------------------------- mi325 · models / language -----------------------------------------------------#
- label: Language Models Test (Extended Generation) # TBD
+2
View File
@@ -29,6 +29,8 @@ steps:
mirror:
amd:
device: mi325_1
# TODO(akaratza): Test after Torch >= 2.12 bump
soft_fail: true
depends_on:
- image-build-amd
+2
View File
@@ -94,6 +94,8 @@ steps:
amd:
device: mi325_1
timeout_in_minutes: 65
# TODO(akaratza): Test after Torch >= 2.12 bump
soft_fail: true
depends_on:
- image-build-amd
source_file_dependencies:
+1 -1
View File
@@ -327,7 +327,7 @@ jobs:
message: 'CC {users} for ROCm-related issue',
},
mistral: {
users: ['patrickvonplaten', 'juliendenize', 'andylolu2'],
users: ['patrickvonplaten', 'juliendenize', 'andylolu2', 'NickLucche'],
message: 'CC {users} for Mistral-related issue',
},
// Add more label -> user mappings here
+1 -1
View File
@@ -27,7 +27,7 @@ jobs:
timeout-minutes: 30
steps:
- uses: actions/checkout@8e8c483db84b4bee98b60c0593521ed34d9990e8 # v6.0.1
- uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
- uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
with:
+2 -2
View File
@@ -48,8 +48,8 @@ jobs:
if: always() && (needs.pre-run-check.result == 'success' || needs.pre-run-check.result == 'skipped')
runs-on: [self-hosted, linux, x64, vllm-runners]
steps:
- uses: actions/checkout@8e8c483db84b4bee98b60c0593521ed34d9990e8 # v6.0.1
- uses: actions/setup-python@83679a892e2d95755f2dac6acb0bfd1e9ac5d548 # v6.1.0
- uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
- uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0
with:
python-version: "3.12"
# Provide shellcheck on PATH so tools/pre_commit/shellcheck.sh skips its
+13
View File
@@ -131,6 +131,19 @@ repos:
--python-version, "3.12",
]
files: ^requirements/(common|xpu|test/xpu)\.(in|txt)$
- id: pip-compile
alias: pip-compile-cpu
name: pip-compile-cpu
args: [
requirements/test/cuda.in,
-o, requirements/test/cpu.txt,
--index-strategy, unsafe-best-match,
--torch-backend, cpu,
--python-platform, x86_64-manylinux_2_28,
--python-version, "3.12",
]
files: ^requirements/(common|cpu|test/(cuda|cpu))\.(in|txt)$
exclude: ^requirements/test/cuda\.txt$
- id: pip-compile
alias: pip-compile-docs
name: pip-compile-docs
+9 -19
View File
@@ -193,26 +193,16 @@ FROM base AS vllm-test-deps
WORKDIR /vllm-workspace
# Copy test requirements
COPY requirements/test/cuda.in requirements/test/cpu.in
# Test requirements are compiled from requirements/test/cuda.in into
# requirements/test/cpu.txt by the pip-compile-cpu pre-commit hook, which
# resolves CPU wheels via uv's --torch-backend cpu.
COPY requirements/test/cpu.txt requirements/test/cpu.txt
RUN \
sed -i '/mamba_ssm/d' requirements/test/cpu.in && \
remove_packages_not_supported_on_aarch64() { \
case "$(uname -m)" in \
aarch64|arm64) \
sed -i '/decord/d' requirements/test/cpu.in; \
sed -i '/terratorch/d' requirements/test/cpu.in; \
;; \
esac; \
}; \
remove_packages_not_supported_on_aarch64 && \
sed -i 's/^torch==.*/torch==2.11.0/g' requirements/test/cpu.in && \
sed -i 's/torchaudio.*/torchaudio/g' requirements/test/cpu.in && \
sed -i 's/torchvision.*/torchvision/g' requirements/test/cpu.in && \
# Related issue: https://github.com/vllm-project/vllm/pull/38800#issuecomment-4228314305
sed -i 's/^sentence-transformers.*/sentence-transformers==5.3.0/g' requirements/test/cpu.in && \
uv pip compile requirements/test/cpu.in -o requirements/test/cpu.txt --index-strategy unsafe-best-match --torch-backend cpu
# cpu.txt is compiled for x86_64, so platform markers are resolved away. Drop
# packages unavailable on aarch64 (decord, terratorch) for arm builds.
RUN case "$(uname -m)" in \
aarch64|arm64) sed -i '/^decord==/d; /^terratorch==/d' requirements/test/cpu.txt ;; \
esac
RUN --mount=type=cache,target=/root/.cache/uv \
uv pip install -r requirements/test/cpu.txt
+1
View File
@@ -167,6 +167,7 @@ Priority is **1 = highest** (tried first).
| `FLASH_ATTN` | FA4* | fp16, bf16 | `auto`, `float16`, `bfloat16` | %16 | Any | ✅ | ✅ | ❌ | ✅ | All | ≥10.0 |
| `FLASH_ATTN_DIFFKV` | | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ✅ | Decoder | Any |
| `FLEX_ATTENTION` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16` | %16 | Any | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder Only | Any |
| `HPC_ATTN` | | fp16, bf16 | `auto`, `fp8_e4m3` | 64 | 128 | ❌ | ❌ | ❌ | ❌ | Decoder | ≥9.0 |
| `ROCM_AITER_FA` | | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | 16, 32 | 64, 128, 256 | ✅ | ✅ | ❌ | ❌ | Decoder | N/A |
| `ROCM_AITER_UNIFIED_ATTN` | | bf16 | `auto`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ✅ | ❌ | ✅ | ❌ | All | N/A |
| `ROCM_ATTN` | | fp16, bf16, fp32 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3`, `fp8_e5m2` | %16 | 32, 64, 80, 96, 128, 160, 192, 224, 256 | ❌ | ✅ | ✅ | ❌ | Decoder, Encoder, Encoder Only | N/A |
+14
View File
@@ -120,6 +120,20 @@ To enable KV cache sharing between multiple vLLM instances using the same `root_
PYTHONHASHSEED=0 vllm serve ...
```
### P2P (Including P/D)
The P2P tier (`type: "p2p"`) shares completed KV blocks between vLLM instances over RDMA via NIXL. Each instance binds a control socket on `host:port` and exchanges blocks directly with peers — no shared filesystem required.
| Key | Required | Default | Notes |
| --- | --- | --- | --- |
| `type` | yes | — | Must be `p2p`. |
| `host` | no | `0.0.0.0` | Address the control socket binds to. |
| `port` | no | `7777` | Port for the control socket. Must be reachable from peers. |
| `backends` | no | `["UCX"]` | NIXL transport backends. See [NixlConnector Usage Guide](nixl_connector_usage.md#selecting-a-nixl-transport-backend-plugin) for available backends and selection guidance. |
| `num_threads` | no | `4` | NIXL agent worker threads. Only used when `backends` is UCX-only; ignored when any non-UCX backend is requested. |
The `backends` and `num_threads` options mirror the conditional logic used by [`NixlConnector`](nixl_connector_usage.md#selecting-a-nixl-transport-backend-plugin): when any non-UCX backend is configured, NIXL is initialised with `backends=...`; otherwise it falls back to a UCX-only agent with the configured `num_threads`. This lets the P2P tier use a different transport (e.g. `MOONCAKE`, `GDS_MT`, `LIBFABRIC`) than the main `NixlConnector` running in the same process.
## Tuning Tips
- `cpu_bytes_to_use`: a bigger CPU tier means fewer trips to slower secondary tiers and a higher hit rate. The value is total across all workers, not per-worker. Leave headroom for the rest of the host workload.
+1 -1
View File
@@ -586,7 +586,7 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
| `MiDashengLMModel` | MiDashengLM | T + A<sup>+</sup> | `mispeech/midashenglm-7b` | | ✅︎ |
| `MiMoV2OmniForCausalLM` | MiMo-V2.5-Omni | T + I<sup>E+</sup> + V<sup>E+</sup> + A<sup>+</sup> | `XiaomiMiMo/MiMo-V2.5-Omni` | | ✅︎ |
| `MiniCPMO` | MiniCPM-O | T + I<sup>E+</sup> + V<sup>E+</sup> + A<sup>E+</sup> | `openbmb/MiniCPM-o-2_6`, etc. | ✅︎ | ✅︎ |
| `MiniCPMV` | MiniCPM-V | T + I<sup>E+</sup> + V<sup>E+</sup> | `openbmb/MiniCPM-V-2` (see note), `openbmb/MiniCPM-Llama3-V-2_5`, `openbmb/MiniCPM-V-2_6`, `openbmb/MiniCPM-V-4`, `openbmb/MiniCPM-V-4_5`, etc. | ✅︎ | |
| `MiniCPMV` | MiniCPM-V | T + I<sup>E+</sup> + V<sup>E+</sup> | `openbmb/MiniCPM-V-2` (see note), `openbmb/MiniCPM-Llama3-V-2_5`, `openbmb/MiniCPM-V-2_6`, `openbmb/MiniCPM-V-4`, `openbmb/MiniCPM-V-4_5`, `openbmb/MiniCPM-V-4_6`, etc. | ✅︎ | |
| `MiniMaxM3SparseForConditionalGeneration` | MiniMax-M3 | T + I<sup>+</sup> + V<sup>+</sup> | `MiniMaxAI/MiniMax-M3`, `MiniMaxAI/MiniMax-M3-MXFP8`, etc. | | ✅︎ |
| `MiniMaxVL01ForConditionalGeneration` | MiniMax-VL | T + I<sup>E+</sup> | `MiniMaxAI/MiniMax-VL-01`, etc. | | ✅︎ |
| `Mistral3ForConditionalGeneration` | Mistral3 (HF Transformers) | T + I<sup>+</sup> | `mistralai/Mistral-Small-3.1-24B-Instruct-2503`, etc. | ✅︎ | ✅︎ |
+15
View File
@@ -85,6 +85,21 @@ significantly reduce the attack surface for these types of abuse.
Also, consider setting `VLLM_MEDIA_URL_ALLOW_REDIRECTS=0` to prevent HTTP
redirects from being followed to bypass domain restrictions.
### 5. **Restrict Media Decode Sizes:**
Compressed media files can expand into gigabytes of memory during decoding. vLLM
enforces decode-size limits to prevent out-of-memory denial of service:
| Environment Variable | Default | Description |
| --- | --- | --- |
| `VLLM_MAX_IMAGE_PIXELS` | `178956970` (~179M pixels) | Maximum decoded image size in pixels. Images exceeding this are rejected before raster memory is allocated. Default matches PIL's built-in 2x decompression-bomb threshold (~680 MB for RGB). |
| `VLLM_MAX_AUDIO_CLIP_FILESIZE_MB` | `25` | Maximum filesize in MB for a single audio file. |
| `VLLM_MAX_AUDIO_DECODE_DURATION_S` | `600` | Maximum decoded audio duration in seconds. Prevents compressed audio from expanding into gigabytes of float32 PCM. |
Setting any of these to `0` disables the corresponding limit. This is **not
recommended** for deployments exposed to untrusted users, as it removes the
protection against resource-exhaustion attacks.
## Security and Firewalls: Protecting Exposed vLLM Systems
While vLLM is designed to allow unsafe network services to be isolated to
File diff suppressed because it is too large Load Diff
+2 -4
View File
@@ -1,3 +1,5 @@
-r ../common.txt
# testing
pytest
tensorizer==2.10.1
@@ -13,7 +15,6 @@ albumentations # required for Nemotron Parse in test_common.py
av # required for audio_in_video tests
backoff # required for phi4mm test
blobfile # required for kimi-vl test
einops # required for MPT, qwen-vl
httpx
librosa # required for audio tests
vector_quantize_pytorch # required for minicpmo_26 test
@@ -34,7 +35,6 @@ matplotlib # required for qwen-vl test
mistral_common[image,audio] >= 1.11.5 # required for voxtral test
num2words # required for smolvlm test
open_clip_torch==2.32.0 # Required for nemotron_vl test, Nemotron Parse in test_common.py
opencv-python-headless >= 4.13.0 # required for video test
datamodel_code_generator # required for minicpm3 test
lm-eval[api]>=0.4.12 # required for model evaluation test
mteb[bm25s]>=2, <3 # required for mteb test
@@ -55,11 +55,9 @@ grpcio-reflection==1.78.0
arctic-inference == 0.1.1; platform_machine == "x86_64" # Required for suffix decoding test
numba == 0.65.0 # Required for N-gram speculative decoding
numpy
runai-model-streamer[s3,gcs,azure]==0.15.7
fastsafetensors>=0.3.2
instanttensor>=0.1.5; platform_machine == "x86_64"
pydantic>=2.12 # 2.11 leads to error on python 3.13
decord==0.6.0; platform_machine == "x86_64"
# terratorch is temporarily disabled while PyPI has the `lightning` package
# in `quarantined` status (every published terratorch version transitively
+288 -18
View File
@@ -9,6 +9,7 @@ aiohappyeyeballs==2.6.1
aiohttp==3.13.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# aiohttp-cors
# datasets
# fsspec
@@ -24,17 +25,34 @@ albumentations==1.4.6
alembic==1.16.4
# via optuna
annotated-doc==0.0.4
# via fastapi
# via
# fastapi
# typer
annotated-types==0.7.0
# via pydantic
anyio==4.6.2.post1
anthropic==0.112.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
anyio==4.14.1
# via
# anthropic
# httpx
# mcp
# openai
# sse-starlette
# starlette
# watchfiles
apache-tvm-ffi==0.1.9
# via
# -c requirements/cuda.txt
# xgrammar
arctic-inference==0.1.1
# via -r requirements/test/cuda.in
argcomplete==3.5.1
# via datamodel-code-generator
astor==0.8.1
# via depyf
attrs==24.2.0
# via
# aiohttp
@@ -59,6 +77,8 @@ bitsandbytes==0.49.2
# via -r requirements/test/cuda.in
black==24.10.0
# via datamodel-code-generator
blake3==1.0.9
# via -r requirements/test/../common.txt
blobfile==3.0.0
# via -r requirements/test/cuda.in
bm25s==0.2.13
@@ -76,12 +96,17 @@ bounded-pool-executor==0.0.3
buildkite-test-collector==0.1.9
# via -r requirements/test/cuda.in
cachetools==5.5.2
# via google-auth
# via
# -r requirements/test/../common.txt
# google-auth
cbor2==6.1.2
# via -r requirements/test/../common.txt
certifi==2024.8.30
# via
# httpcore
# httpx
# requests
# sentry-sdk
cffi==2.0.0
# via
# cryptography
@@ -98,9 +123,11 @@ click==8.1.7
# jiwer
# nltk
# ray
# rich-toolkit
# schemathesis
# typer
# uvicorn
cloudpickle==3.1.2
# via -r requirements/test/../common.txt
cohere-melody==0.9.0
# via -r requirements/test/cuda.in
colorama==0.4.6
@@ -111,6 +138,10 @@ colorful==0.5.6
# via ray
colorlog==6.10.1
# via optuna
compressed-tensors==0.17.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
contourpy==1.3.0
# via matplotlib
coverage==7.10.6
@@ -149,30 +180,49 @@ decorator==5.1.1
# via librosa
decord==0.6.0
# via -r requirements/test/cuda.in
depyf==0.20.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
detect-installer==0.1.0
# via fastapi-cloud-cli
dill==0.3.8
# via
# datasets
# depyf
# evaluate
# lm-eval
# multiprocess
diskcache==5.6.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
distlib==0.3.9
# via virtualenv
distro==1.9.0
# via
# anthropic
# openai
dnspython==2.7.0
# via email-validator
docker==7.1.0
# via gpt-oss
docopt==0.6.2
# via num2words
docstring-parser==0.18.0
# via anthropic
einops==0.8.1
# via
# -r requirements/test/cuda.in
# -r requirements/test/../common.txt
# encodec
# vector-quantize-pytorch
# vocos
einx==0.3.0
# via vector-quantize-pytorch
email-validator==2.2.0
# via pydantic
# via
# fastapi
# pydantic
encodec==0.1.1
# via vocos
et-xmlfile==2.0.0
@@ -182,7 +232,17 @@ evaluate==0.4.3
fastapi==0.136.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# gpt-oss
# model-hosting-container-standards
fastapi-cli==0.0.27
# via fastapi
fastapi-cloud-cli==0.21.0
# via fastapi-cli
fastar==0.11.0
# via
# fastapi
# fastapi-cloud-cli
fastparquet==2024.11.0
# via genai-perf
fastrlock==0.8.2
@@ -194,6 +254,7 @@ fastsafetensors==0.3.2
filelock==3.16.1
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# blobfile
# datasets
# huggingface-hub
@@ -243,7 +304,10 @@ google-crc32c==1.7.1
google-resumable-media==2.7.2
# via google-cloud-storage
googleapis-common-protos==1.70.0
# via google-api-core
# via
# google-api-core
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
gpt-oss==0.0.8
# via -r requirements/test/cuda.in
graphql-core==3.2.6
@@ -254,6 +318,7 @@ grpcio==1.78.0
# via
# -r requirements/test/cuda.in
# grpcio-reflection
# opentelemetry-exporter-otlp-proto-grpc
# ray
grpcio-reflection==1.78.0
# via -r requirements/test/cuda.in
@@ -275,12 +340,22 @@ html2text==2025.4.15
# via gpt-oss
httpcore==1.0.6
# via httpx
httptools==0.8.0
# via uvicorn
httpx==0.27.2
# via
# -r requirements/test/cuda.in
# anthropic
# fastapi
# fastapi-cloud-cli
# huggingface-hub
# mcp
# model-hosting-container-standards
# openai
# perceptron
# schemathesis
httpx-sse==0.4.3
# via mcp
huggingface-hub==1.10.2
# via
# accelerate
@@ -314,6 +389,8 @@ idna==3.10
# httpx
# requests
# yarl
ijson==3.5.0
# via -r requirements/test/../common.txt
imagehash==4.3.2
# via -r requirements/test/cuda.in
imageio==2.37.0
@@ -326,6 +403,8 @@ iniconfig==2.0.0
# via pytest
instanttensor==0.1.5
# via -r requirements/test/cuda.in
interegular==0.3.3
# via lm-format-enforcer
isodate==0.7.2
# via azure-storage-blob
isort==5.13.2
@@ -333,15 +412,21 @@ isort==5.13.2
jinja2==3.1.6
# via
# datamodel-code-generator
# fastapi
# genai-perf
# lm-eval
# torch
jiter==0.15.0
# via
# anthropic
# openai
jiwer==3.0.5
# via -r requirements/test/cuda.in
jmespath==1.0.1
# via
# boto3
# botocore
# model-hosting-container-standards
joblib==1.4.2
# via
# librosa
@@ -350,7 +435,9 @@ joblib==1.4.2
jsonschema==4.23.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# hypothesis-jsonschema
# mcp
# mistral-common
# ray
jsonschema-rs==0.46.5
@@ -365,6 +452,10 @@ kaleido==0.2.1
# via genai-perf
kiwisolver==1.4.7
# via matplotlib
lark==1.2.2
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
lazy-loader==0.4
# via
# librosa
@@ -373,10 +464,20 @@ libnacl==2.1.0
# via tensorizer
librosa==0.10.2.post1
# via -r requirements/test/cuda.in
llguidance==1.7.6
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
llvmlite==0.47.0
# via numba
lm-eval==0.4.12
# via -r requirements/test/cuda.in
lm-format-enforcer==0.11.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
loguru==0.7.3
# via compressed-tensors
lxml==5.3.0
# via
# blobfile
@@ -398,12 +499,19 @@ mbstrdecoder==1.1.3
# dataproperty
# pytablewriter
# typepy
mcp==1.28.1
# via -r requirements/test/../common.txt
mdurl==0.1.2
# via markdown-it-py
mistral-common==1.11.5
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# -r requirements/test/cuda.in
model-hosting-container-standards==0.1.16
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
more-itertools==10.5.0
# via lm-eval
mpmath==1.3.0
@@ -418,6 +526,8 @@ msgpack==1.1.0
# via
# librosa
# ray
msgspec==0.21.1
# via -r requirements/test/../common.txt
mteb==2.8.3
# via -r requirements/test/cuda.in
multidict==6.1.0
@@ -434,6 +544,8 @@ networkx==3.2.1
# via
# scikit-image
# torch
ninja==1.13.0
# via -r requirements/test/../common.txt
nltk==3.9.1
# via rouge-score
num2words==0.5.14
@@ -445,7 +557,7 @@ numba==0.65.0
# librosa
numpy==2.2.6
# via
# -r requirements/test/cuda.in
# -r requirements/test/../common.txt
# accelerate
# albumentations
# bitsandbytes
@@ -489,6 +601,7 @@ numpy==2.2.6
# transformers
# tritonclient
# vocos
# xgrammar
nvidia-cublas==13.1.0.3
# via
# cuda-toolkit
@@ -530,9 +643,14 @@ nvidia-nvtx==13.0.85
# via cuda-toolkit
open-clip-torch==2.32.0
# via -r requirements/test/cuda.in
openai==2.44.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
openai-harmony==0.0.4
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# gpt-oss
opencensus==0.11.4
# via ray
@@ -541,7 +659,7 @@ opencensus-context==0.1.3
opencv-python-headless==4.13.0.90
# via
# -c requirements/common.txt
# -r requirements/test/cuda.in
# -r requirements/test/../common.txt
# albumentations
# mistral-common
openpyxl==3.1.5
@@ -549,24 +667,54 @@ openpyxl==3.1.5
opentelemetry-api==1.35.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
# opentelemetry-exporter-prometheus
# opentelemetry-sdk
# opentelemetry-semantic-conventions
opentelemetry-exporter-otlp==1.35.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
opentelemetry-exporter-otlp-proto-common==1.35.0
# via
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
opentelemetry-exporter-otlp-proto-grpc==1.35.0
# via opentelemetry-exporter-otlp
opentelemetry-exporter-otlp-proto-http==1.35.0
# via opentelemetry-exporter-otlp
opentelemetry-exporter-prometheus==0.56b0
# via ray
opentelemetry-proto==1.35.0
# via ray
# via
# opentelemetry-exporter-otlp-proto-common
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
# ray
opentelemetry-sdk==1.35.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
# opentelemetry-exporter-prometheus
# ray
opentelemetry-semantic-conventions==0.56b0
# via opentelemetry-sdk
opentelemetry-semantic-conventions-ai==0.4.13
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
optuna==3.6.1
# via genai-perf
orjson==3.11.5
# via genai-perf
outlines-core==0.2.14
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
packaging==24.2
# via
# accelerate
@@ -578,6 +726,7 @@ packaging==24.2
# fastparquet
# huggingface-hub
# lazy-loader
# lm-format-enforcer
# matplotlib
# optuna
# peft
@@ -597,6 +746,8 @@ pandas==2.2.3
# fastparquet
# genai-perf
# statsmodels
partial-json-parser==0.2.1.1.post7
# via -r requirements/test/../common.txt
pathspec==0.12.1
# via black
pathvalidate==3.2.1
@@ -611,6 +762,7 @@ perf-analyzer==0.1.0
# via genai-perf
pillow==10.4.0
# via
# -r requirements/test/../common.txt
# genai-perf
# imagehash
# imageio
@@ -644,8 +796,14 @@ pqdm==0.2.0
prometheus-client==0.22.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# opentelemetry-exporter-prometheus
# prometheus-fastapi-instrumentator
# ray
prometheus-fastapi-instrumentator==8.0.2
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
propcache==0.2.0
# via
# aiohttp
@@ -655,6 +813,7 @@ proto-plus==1.26.1
protobuf==6.33.6
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# google-api-core
# googleapis-common-protos
# grpcio-reflection
@@ -664,11 +823,14 @@ protobuf==6.33.6
# tensorizer
psutil==6.1.0
# via
# -r requirements/test/../common.txt
# accelerate
# peft
# tensorizer
py==1.11.0
# via pytest-forked
py-cpuinfo==9.0.0
# via -r requirements/test/../common.txt
py-spy==0.4.0
# via ray
pyarrow==23.0.0
@@ -681,6 +843,8 @@ pyasn1==0.6.1
# rsa
pyasn1-modules==0.4.2
# via google-auth
pybase64==1.4.3
# via -r requirements/test/../common.txt
pycountry==24.6.1
# via pydantic-extra-types
pycparser==2.22
@@ -690,26 +854,43 @@ pycryptodomex==3.22.0
pydantic==2.12.0
# via
# -c requirements/common.txt
# -r requirements/test/cuda.in
# -r requirements/test/../common.txt
# albumentations
# anthropic
# compressed-tensors
# datamodel-code-generator
# fastapi
# fastapi-cloud-cli
# gpt-oss
# lm-format-enforcer
# mcp
# mistral-common
# model-hosting-container-standards
# mteb
# openai
# openai-harmony
# pydantic-extra-types
# pydantic-settings
# ray
# xgrammar
pydantic-core==2.41.1
# via pydantic
pydantic-extra-types==2.10.5
# via mistral-common
# via
# fastapi
# mistral-common
pydantic-settings==2.14.2
# via
# fastapi
# mcp
pygments==2.18.0
# via
# pytest
# rich
pyjwt==2.11.0
# via msal
# via
# mcp
# msal
pyparsing==3.2.0
# via matplotlib
pyrate-limiter==4.4.0
@@ -751,6 +932,16 @@ python-dateutil==2.9.0.post0
# matplotlib
# pandas
# typepy
python-dotenv==1.2.2
# via
# pydantic-settings
# uvicorn
python-json-logger==4.1.0
# via -r requirements/test/../common.txt
python-multipart==0.0.32
# via
# fastapi
# mcp
python-rapidjson==1.20
# via tritonclient
pytrec-eval-terrier==0.5.7
@@ -763,12 +954,14 @@ pywavelets==1.9.0
# via imagehash
pyyaml==6.0.2
# via
# -r requirements/test/../common.txt
# accelerate
# albumentations
# datamodel-code-generator
# datasets
# genai-perf
# huggingface-hub
# lm-format-enforcer
# optuna
# peft
# ray
@@ -776,7 +969,12 @@ pyyaml==6.0.2
# schemathesis
# timm
# transformers
# uvicorn
# vocos
pyzmq==27.1.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
rapidfuzz==3.12.1
# via jiwer
ray==2.48.0
@@ -789,6 +987,7 @@ referencing==0.35.1
# jsonschema-specifications
regex==2026.2.28
# via
# -r requirements/test/../common.txt
# nltk
# open-clip-torch
# sacrebleu
@@ -797,6 +996,7 @@ regex==2026.2.28
requests==2.32.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# azure-core
# buildkite-test-collector
# datasets
@@ -809,6 +1009,7 @@ requests==2.32.3
# mistral-common
# msal
# mteb
# opentelemetry-exporter-otlp-proto-http
# pooch
# ray
# responses
@@ -822,8 +1023,15 @@ rich==13.9.4
# genai-perf
# mteb
# perceptron
# rich-toolkit
# schemathesis
# typer
rich-toolkit==0.20.1
# via
# fastapi-cli
# fastapi-cloud-cli
rignore==0.7.6
# via fastapi-cloud-cli
rouge-score==0.1.2
# via lm-eval
rpds-py==0.20.1
@@ -847,6 +1055,7 @@ sacrebleu==2.4.3
safetensors==0.7.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# accelerate
# open-clip-torch
# peft
@@ -882,9 +1091,17 @@ sentence-transformers==5.2.0
# via
# -r requirements/test/cuda.in
# mteb
sentencepiece==0.2.1
# via -r requirements/test/../common.txt
sentry-sdk==2.63.0
# via fastapi-cloud-cli
setproctitle==1.3.7
# via -r requirements/test/../common.txt
setuptools==77.0.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# model-hosting-container-standards
# pytablewriter
# torch
shellingham==1.5.4
@@ -894,6 +1111,7 @@ shellingham==1.5.4
six==1.16.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# junit-xml
# opencensus
# python-dateutil
@@ -902,8 +1120,9 @@ smart-open==7.1.0
# via ray
sniffio==1.3.1
# via
# anyio
# anthropic
# httpx
# openai
sortedcontainers==2.4.0
# via hypothesis
soundfile==0.12.1
@@ -922,10 +1141,17 @@ sqlalchemy==2.0.41
# optuna
sqlitedict==2.1.0
# via lm-eval
sse-starlette==3.4.5
# via mcp
starlette==1.3.1
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# fastapi
# mcp
# model-hosting-container-standards
# prometheus-fastapi-instrumentator
# sse-starlette
# starlette-testclient
starlette-testclient==0.4.1
# via schemathesis
@@ -933,6 +1159,8 @@ statsmodels==0.14.4
# via genai-perf
structlog==25.4.0
# via gpt-oss
supervisor==4.3.0
# via model-hosting-container-standards
sympy==1.13.3
# via
# einx
@@ -962,6 +1190,7 @@ tifffile==2025.3.30
tiktoken==0.12.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# gpt-oss
# lm-eval
# mistral-common
@@ -973,6 +1202,7 @@ timm==1.0.17
tokenizers==0.22.2
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# -r requirements/test/cuda.in
# transformers
torch==2.11.0+cu130
@@ -981,6 +1211,7 @@ torch==2.11.0+cu130
# -r requirements/test/cuda.in
# accelerate
# bitsandbytes
# compressed-tensors
# encodec
# instanttensor
# mteb
@@ -994,6 +1225,7 @@ torch==2.11.0+cu130
# torchvision
# vector-quantize-pytorch
# vocos
# xgrammar
torchaudio==2.11.0+cu130
# via
# -c requirements/cuda.txt
@@ -1009,6 +1241,7 @@ torchvision==0.26.0+cu130
# timm
tqdm==4.67.3
# via
# -r requirements/test/../common.txt
# datasets
# evaluate
# huggingface-hub
@@ -1016,6 +1249,7 @@ tqdm==4.67.3
# mteb
# nltk
# open-clip-torch
# openai
# optuna
# peft
# pqdm
@@ -1025,15 +1259,20 @@ tqdm==4.67.3
transformers==5.5.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# -r requirements/test/cuda.in
# compressed-tensors
# genai-perf
# peft
# sentence-transformers
# transformers-stream-generator
# xgrammar
transformers-stream-generator==0.0.5
# via -r requirements/test/cuda.in
triton==3.6.0
# via torch
# via
# torch
# xgrammar
tritonclient==2.64.0
# via -r requirements/test/cuda.in
typepy==1.3.2
@@ -1041,8 +1280,10 @@ typepy==1.3.2
# dataproperty
# pytablewriter
# tabledata
typer==0.15.2
typer==0.26.8
# via
# fastapi-cli
# fastapi-cloud-cli
# fastsafetensors
# huggingface-hub
# perceptron
@@ -1050,9 +1291,13 @@ typer==0.15.2
typing-extensions==4.15.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# aiosignal
# albumentations
# alembic
# anthropic
# anyio
# apache-tvm-ffi
# azure-core
# azure-identity
# azure-storage-blob
@@ -1062,9 +1307,13 @@ typing-extensions==4.15.0
# huggingface-hub
# librosa
# lm-eval
# mcp
# mistral-common
# mteb
# openai
# opentelemetry-api
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
# opentelemetry-sdk
# opentelemetry-semantic-conventions
# pqdm
@@ -1072,17 +1321,20 @@ typing-extensions==4.15.0
# pydantic-core
# pydantic-extra-types
# pytest-asyncio
# rich-toolkit
# schemathesis
# sentence-transformers
# sqlalchemy
# starlette
# torch
# typer
# typing-inspection
# xgrammar
typing-inspection==0.4.2
# via
# fastapi
# mcp
# pydantic
# pydantic-settings
tzdata==2024.2
# via pandas
urllib3==2.2.3
@@ -1092,23 +1344,41 @@ urllib3==2.2.3
# docker
# requests
# responses
# sentry-sdk
# tritonclient
uvicorn==0.35.0
# via gpt-oss
# via
# fastapi
# fastapi-cli
# fastapi-cloud-cli
# gpt-oss
# mcp
uvloop==0.22.1
# via uvicorn
vector-quantize-pytorch==1.21.2
# via -r requirements/test/cuda.in
virtualenv==20.31.2
# via ray
vocos==0.1.0
# via -r requirements/test/cuda.in
watchfiles==1.2.0
# via
# -r requirements/test/../common.txt
# uvicorn
wcwidth==0.2.13
# via ftfy
websockets==16.0
# via uvicorn
werkzeug==3.1.3
# via schemathesis
word2number==1.1
# via lm-eval
wrapt==1.17.2
# via smart-open
xgrammar==0.2.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
xxhash==3.5.0
# via
# datasets
+1 -4
View File
@@ -15,7 +15,6 @@ albumentations # required for Nemotron Parse in test_common.py
av # required for audio_in_video tests
backoff # required for phi4mm test
blobfile # required for kimi-vl test
einops # required for MPT, qwen-vl
httpx
librosa # required for audio tests
vector_quantize_pytorch # required for minicpmo_26 test
@@ -33,7 +32,6 @@ matplotlib # required for qwen-vl test
mistral_common[image,audio]>=1.11.5 # required for voxtral test
num2words # required for smolvlm test
open_clip_torch==2.32.0 # Required for nemotron_vl test, Nemotron Parse in test_common.py
opencv-python-headless>=4.13.0 # required for video test
datamodel_code_generator # required for minicpm3 test
lm-eval[api]>=0.4.12 # required for model evaluation test
mteb[bm25s]>=2, <3 # required for mteb test
@@ -54,11 +52,9 @@ grpcio-reflection==1.78.0
arctic-inference==0.1.1 # Required for suffix decoding test
numba==0.65.0 # Required for N-gram speculative decoding
numpy
runai-model-streamer[s3,gcs,azure]==0.15.7
fastsafetensors>=0.3.2
instanttensor>=0.1.5
pydantic>=2.12 # 2.11 leads to error on python 3.13
decord==0.6.0
# Prithvi tests
@@ -74,6 +70,7 @@ gpt-oss>=0.0.7; python_version > '3.11'
perceptron # required for isaac test
kaldi-native-fbank>=1.18.7 # required for fireredasr2 test
cohere_melody>=0.9.0 # required for cohere command reasoning parser test
# Newer versions of datasets require torchcoded, that makes the tests fail in CI because of a missing library.
# Older versions are in conflict with terratorch requirements.
+2 -4
View File
@@ -130,6 +130,8 @@ cloudpickle==3.1.2
# via
# -r requirements/test/../common.txt
# tilelang
cohere-melody==0.9.0
# via -r requirements/test/rocm.in
colorama==0.4.6
# via
# perceptron
@@ -205,7 +207,6 @@ docstring-parser==0.17.0
einops==0.8.2
# via
# -r requirements/test/../common.txt
# -r requirements/test/rocm.in
# encodec
# vector-quantize-pytorch
# vocos
@@ -561,7 +562,6 @@ numba==0.65.0
numpy==2.2.6
# via
# -r requirements/test/../common.txt
# -r requirements/test/rocm.in
# accelerate
# albumentations
# bitsandbytes
@@ -630,7 +630,6 @@ opencv-python-headless==4.13.0.92
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# -r requirements/test/rocm.in
# albumentations
# mistral-common
openpyxl==3.1.5
@@ -834,7 +833,6 @@ pydantic==2.12.5
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# -r requirements/test/rocm.in
# albumentations
# anthropic
# compressed-tensors
+2
View File
@@ -1,3 +1,5 @@
-r ../common.txt
# --- Test Infrastructure ---
tblib
pytest
+316 -4
View File
@@ -11,6 +11,7 @@ aiohappyeyeballs==2.6.1
aiohttp==3.13.4
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# fsspec
# gpt-oss
# lm-eval
@@ -24,12 +25,25 @@ annotated-doc==0.0.4
# typer
annotated-types==0.7.0
# via pydantic
anthropic==0.112.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
anyio==4.13.0
# via
# anthropic
# httpx
# mcp
# openai
# sse-starlette
# starlette
# watchfiles
apache-tvm-ffi==0.1.12
# via xgrammar
arctic-inference==0.1.1
# via -r requirements/test/xpu.in
astor==0.8.1
# via depyf
attrs==26.1.0
# via
# aiohttp
@@ -39,6 +53,8 @@ audioread==3.0.1
# via
# -r requirements/test/xpu.in
# librosa
blake3==1.0.9
# via -r requirements/test/../common.txt
blobfile==3.0.0
# via -r requirements/test/xpu.in
bm25s==0.2.13
@@ -47,13 +63,20 @@ bm25s==0.2.13
# mteb
bounded-pool-executor==0.0.3
# via pqdm
cachetools==7.1.4
# via -r requirements/test/../common.txt
cbor2==6.1.2
# via -r requirements/test/../common.txt
certifi==2026.2.25
# via
# httpcore
# httpx
# requests
# sentry-sdk
cffi==2.0.0
# via soundfile
# via
# cryptography
# soundfile
chardet==5.2.0
# via mbstrdecoder
charset-normalizer==3.4.6
@@ -64,13 +87,22 @@ click==8.3.1
# via
# jiwer
# nltk
# rich-toolkit
# schemathesis
# typer
# uvicorn
cloudpickle==3.1.2
# via -r requirements/test/../common.txt
colorama==0.4.6
# via sacrebleu
compressed-tensors==0.17.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
coverage==7.13.5
# via pytest-cov
cryptography==49.0.0
# via pyjwt
dataproperty==1.1.0
# via
# pytablewriter
@@ -82,16 +114,35 @@ datasets==4.8.4
# mteb
decorator==5.2.1
# via librosa
depyf==0.20.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
detect-installer==0.1.0
# via fastapi-cloud-cli
dill==0.4.1
# via
# datasets
# depyf
# evaluate
# lm-eval
# multiprocess
diskcache==5.6.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
distro==1.9.0
# via
# anthropic
# openai
dnspython==2.8.0
# via email-validator
docker==7.1.0
# via gpt-oss
docopt==0.6.2
# via num2words
docstring-parser==0.18.0
# via anthropic
dpcpp-cpp-rt==2025.3.2
# via
# onemkl-sycl-blas
@@ -100,15 +151,30 @@ dpcpp-cpp-rt==2025.3.2
# onemkl-sycl-rng
# onemkl-sycl-sparse
# torch
einops==0.8.2
# via -r requirements/test/../common.txt
email-validator==2.3.0
# via
# fastapi
# pydantic
evaluate==0.4.6
# via lm-eval
fastapi==0.135.2
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# gpt-oss
# model-hosting-container-standards
fastapi-cli==0.0.27
# via fastapi
fastapi-cloud-cli==0.21.0
# via fastapi-cli
fastar==0.11.0
# via fastapi-cloud-cli
filelock==3.25.2
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# blobfile
# datasets
# huggingface-hub
@@ -124,10 +190,16 @@ fsspec==2026.2.0
# evaluate
# huggingface-hub
# torch
googleapis-common-protos==1.75.0
# via
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
gpt-oss==0.0.8
# via -r requirements/test/xpu.in
graphql-core==3.2.8
# via hypothesis-graphql
grpcio==1.81.1
# via opentelemetry-exporter-otlp-proto-grpc
h11==0.16.0
# via
# httpcore
@@ -140,11 +212,21 @@ html2text==2025.4.15
# via gpt-oss
httpcore==1.0.9
# via httpx
httptools==0.8.0
# via uvicorn
httpx==0.28.1
# via
# anthropic
# datasets
# fastapi
# fastapi-cloud-cli
# huggingface-hub
# mcp
# model-hosting-container-standards
# openai
# schemathesis
httpx-sse==0.4.3
# via mcp
huggingface-hub==1.10.2
# via
# accelerate
@@ -166,9 +248,12 @@ hypothesis-jsonschema==0.23.1
idna==3.11
# via
# anyio
# email-validator
# httpx
# requests
# yarl
ijson==3.5.0
# via -r requirements/test/../common.txt
imageio==2.37.3
# via scikit-image
impi-rt==2021.17.2
@@ -212,13 +297,22 @@ intel-sycl-rt==2025.3.2
# dpcpp-cpp-rt
# oneccl
# torch
interegular==0.3.3
# via lm-format-enforcer
jinja2==3.1.6
# via
# -c requirements/xpu.txt
# fastapi
# lm-eval
# torch
jiter==0.15.0
# via
# anthropic
# openai
jiwer==4.0.0
# via -r requirements/test/xpu.in
jmespath==1.1.0
# via model-hosting-container-standards
joblib==1.5.3
# via
# librosa
@@ -227,7 +321,9 @@ joblib==1.5.3
jsonschema==4.26.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# hypothesis-jsonschema
# mcp
# mistral-common
# schemathesis
jsonschema-rs==0.45.0
@@ -236,16 +332,30 @@ jsonschema-specifications==2025.9.1
# via jsonschema
junit-xml==1.9
# via schemathesis
lark==1.2.2
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
lazy-loader==0.5
# via
# librosa
# scikit-image
librosa==0.10.2.post1
# via -r requirements/test/xpu.in
llguidance==1.7.6
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
llvmlite==0.47.0
# via numba
lm-eval==0.4.12
# via -r requirements/test/xpu.in
lm-format-enforcer==0.11.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
loguru==0.7.3
# via compressed-tensors
lxml==6.0.2
# via
# blobfile
@@ -262,11 +372,14 @@ mbstrdecoder==1.1.4
# dataproperty
# pytablewriter
# typepy
mcp==1.28.1
# via -r requirements/test/../common.txt
mdurl==0.1.2
# via markdown-it-py
mistral-common==1.11.5
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# -r requirements/test/xpu.in
mkl==2025.3.1
# via
@@ -276,6 +389,10 @@ mkl==2025.3.1
# onemkl-sycl-rng
# onemkl-sycl-sparse
# torch
model-hosting-container-standards==0.1.16
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
modelscope==1.35.3
# via -r requirements/test/xpu.in
more-itertools==10.8.0
@@ -284,6 +401,8 @@ mpmath==1.3.0
# via sympy
msgpack==1.1.2
# via librosa
msgspec==0.21.1
# via -r requirements/test/../common.txt
mteb==2.12.7
# via -r requirements/test/xpu.in
multidict==6.7.1
@@ -298,6 +417,8 @@ networkx==3.6.1
# via
# scikit-image
# torch
ninja==1.13.0
# via -r requirements/test/../common.txt
nltk==3.9.4
# via rouge-score
num2words==0.5.14
@@ -308,6 +429,7 @@ numba==0.65.0
# librosa
numpy==2.2.6
# via
# -r requirements/test/../common.txt
# accelerate
# albumentations
# bm25s
@@ -333,6 +455,7 @@ numpy==2.2.6
# tifffile
# torchvision
# transformers
# xgrammar
oneccl==2021.17.2
# via
# oneccl-devel
@@ -356,15 +479,65 @@ onemkl-sycl-rng==2025.3.1
# via torch
onemkl-sycl-sparse==2025.3.1
# via torch
openai==2.44.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
openai-harmony==0.0.8
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# gpt-oss
opencv-python-headless==4.13.0.92
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# albumentations
# mistral-common
opentelemetry-api==1.43.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
# opentelemetry-sdk
# opentelemetry-semantic-conventions
opentelemetry-exporter-otlp==1.43.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
opentelemetry-exporter-otlp-proto-common==1.43.0
# via
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
opentelemetry-exporter-otlp-proto-grpc==1.43.0
# via opentelemetry-exporter-otlp
opentelemetry-exporter-otlp-proto-http==1.43.0
# via opentelemetry-exporter-otlp
opentelemetry-proto==1.43.0
# via
# opentelemetry-exporter-otlp-proto-common
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
opentelemetry-sdk==1.43.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
# opentelemetry-semantic-conventions-ai
opentelemetry-semantic-conventions==0.64b0
# via
# opentelemetry-sdk
# opentelemetry-semantic-conventions-ai
opentelemetry-semantic-conventions-ai==0.5.1
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
outlines-core==0.2.14
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
packaging==26.0
# via
# -c requirements/xpu.txt
@@ -373,6 +546,7 @@ packaging==26.0
# evaluate
# huggingface-hub
# lazy-loader
# lm-format-enforcer
# modelscope
# pooch
# pytest
@@ -384,10 +558,13 @@ pandas==3.0.1
# via
# datasets
# evaluate
partial-json-parser==0.2.1.1.post7
# via -r requirements/test/../common.txt
pathvalidate==3.3.1
# via pytablewriter
pillow==12.1.1
# via
# -r requirements/test/../common.txt
# imageio
# mistral-common
# scikit-image
@@ -410,16 +587,37 @@ portalocker==3.2.0
# via sacrebleu
pqdm==0.2.0
# via -r requirements/test/xpu.in
prometheus-client==0.25.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# prometheus-fastapi-instrumentator
prometheus-fastapi-instrumentator==8.0.2
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
propcache==0.4.1
# via
# aiohttp
# yarl
protobuf==7.35.1
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# googleapis-common-protos
# opentelemetry-proto
psutil==7.2.2
# via accelerate
# via
# -r requirements/test/../common.txt
# accelerate
py==1.11.0
# via pytest-forked
py-cpuinfo==9.0.0
# via -r requirements/test/../common.txt
pyarrow==23.0.1
# via datasets
pybase64==1.4.3
# via -r requirements/test/../common.txt
pycountry==26.2.16
# via pydantic-extra-types
pycparser==3.0
@@ -429,23 +627,41 @@ pycryptodomex==3.23.0
pydantic==2.12.5
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# albumentations
# anthropic
# compressed-tensors
# fastapi
# fastapi-cloud-cli
# gpt-oss
# lm-format-enforcer
# mcp
# mistral-common
# model-hosting-container-standards
# mteb
# openai
# openai-harmony
# pydantic-extra-types
# pydantic-settings
# xgrammar
pydantic-core==2.41.5
# via pydantic
pydantic-extra-types==2.11.1
# via mistral-common
# via
# fastapi
# mistral-common
pydantic-settings==2.14.2
# via
# fastapi
# mcp
pyelftools==0.32
# via triton-xpu
pygments==2.20.0
# via
# pytest
# rich
pyjwt==2.13.0
# via mcp
pyrate-limiter==4.1.0
# via schemathesis
pystemmer==3.0.0
@@ -480,19 +696,36 @@ python-dateutil==2.9.0.post0
# via
# pandas
# typepy
python-dotenv==1.2.2
# via
# pydantic-settings
# uvicorn
python-json-logger==4.1.0
# via -r requirements/test/../common.txt
python-multipart==0.0.32
# via
# fastapi
# mcp
pytrec-eval-terrier==0.5.10
# via mteb
pytz==2026.1.post1
# via typepy
pyyaml==6.0.3
# via
# -r requirements/test/../common.txt
# accelerate
# albumentations
# datasets
# huggingface-hub
# lm-format-enforcer
# schemathesis
# timm
# transformers
# uvicorn
pyzmq==27.1.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
rapidfuzz==3.12.1
# via
# -r requirements/test/xpu.in
@@ -503,6 +736,7 @@ referencing==0.37.0
# jsonschema-specifications
regex==2026.3.32
# via
# -r requirements/test/../common.txt
# nltk
# sacrebleu
# tiktoken
@@ -510,6 +744,7 @@ regex==2026.3.32
requests==2.33.1
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# datasets
# docker
# evaluate
@@ -518,6 +753,7 @@ requests==2.33.1
# mistral-common
# modelscope
# mteb
# opentelemetry-exporter-otlp-proto-http
# pooch
# schemathesis
# starlette-testclient
@@ -525,8 +761,15 @@ requests==2.33.1
rich==14.3.3
# via
# mteb
# rich-toolkit
# schemathesis
# typer
rich-toolkit==0.20.1
# via
# fastapi-cli
# fastapi-cloud-cli
rignore==0.7.6
# via fastapi-cloud-cli
rouge-score==0.1.2
# via lm-eval
rpds-py==0.30.0
@@ -538,6 +781,7 @@ sacrebleu==2.6.0
safetensors==0.7.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# accelerate
# timm
# transformers
@@ -564,10 +808,18 @@ scipy==1.17.1
# sentence-transformers
sentence-transformers==5.3.0
# via mteb
sentencepiece==0.2.1
# via -r requirements/test/../common.txt
sentry-sdk==2.63.0
# via fastapi-cloud-cli
setproctitle==1.3.7
# via -r requirements/test/../common.txt
setuptools==80.10.2
# via
# -c requirements/common.txt
# -c requirements/xpu.txt
# -r requirements/test/../common.txt
# model-hosting-container-standards
# modelscope
# pytablewriter
# torch
@@ -576,9 +828,14 @@ shellingham==1.5.4
six==1.17.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# junit-xml
# python-dateutil
# rouge-score
sniffio==1.3.1
# via
# anthropic
# openai
sortedcontainers==2.4.0
# via hypothesis
soundfile==0.13.1
@@ -593,15 +850,24 @@ soxr==0.5.0.post1
# mistral-common
sqlitedict==2.1.0
# via lm-eval
sse-starlette==3.4.5
# via mcp
starlette==1.3.1
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# fastapi
# mcp
# model-hosting-container-standards
# prometheus-fastapi-instrumentator
# sse-starlette
# starlette-testclient
starlette-testclient==0.4.1
# via schemathesis
structlog==25.5.0
# via gpt-oss
supervisor==4.3.0
# via model-hosting-container-standards
sympy==1.14.0
# via torch
tabledata==1.3.4
@@ -636,6 +902,7 @@ tifffile==2026.3.3
tiktoken==0.12.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# gpt-oss
# lm-eval
# mistral-common
@@ -644,19 +911,23 @@ timm==1.0.17
tokenizers==0.22.2
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# transformers
torch==2.12.0+xpu
# via
# -c requirements/xpu.txt
# accelerate
# compressed-tensors
# mteb
# sentence-transformers
# timm
# torchvision
# xgrammar
torchvision==0.27.0+xpu
# via timm
tqdm==4.67.3
# via
# -r requirements/test/../common.txt
# datasets
# evaluate
# huggingface-hub
@@ -664,13 +935,19 @@ tqdm==4.67.3
# modelscope
# mteb
# nltk
# openai
# pqdm
# sentence-transformers
# transformers
transformers==5.5.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# compressed-tensors
# sentence-transformers
# xgrammar
triton==3.7.1
# via xgrammar
triton-xpu==3.7.1
# via torch
typepy==1.3.4
@@ -680,36 +957,53 @@ typepy==1.3.4
# tabledata
typer==0.24.1
# via
# fastapi-cli
# fastapi-cloud-cli
# huggingface-hub
# transformers
typing-extensions==4.15.0
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
# aiosignal
# albumentations
# anthropic
# anyio
# apache-tvm-ffi
# chz
# fastapi
# grpcio
# huggingface-hub
# librosa
# lm-eval
# mcp
# mistral-common
# mteb
# openai
# opentelemetry-api
# opentelemetry-exporter-otlp-proto-grpc
# opentelemetry-exporter-otlp-proto-http
# opentelemetry-sdk
# opentelemetry-semantic-conventions
# pqdm
# pydantic
# pydantic-core
# pydantic-extra-types
# pytest-asyncio
# referencing
# rich-toolkit
# schemathesis
# sentence-transformers
# starlette
# torch
# typing-inspection
# xgrammar
typing-inspection==0.4.2
# via
# fastapi
# mcp
# pydantic
# pydantic-settings
umf==1.0.3
# via
# intel-cmplr-lib-ur
@@ -720,12 +1014,30 @@ urllib3==2.6.3
# docker
# modelscope
# requests
# sentry-sdk
uvicorn==0.42.0
# via gpt-oss
# via
# fastapi
# fastapi-cli
# fastapi-cloud-cli
# gpt-oss
# mcp
uvloop==0.22.1
# via uvicorn
watchfiles==1.2.0
# via
# -r requirements/test/../common.txt
# uvicorn
websockets==16.0
# via uvicorn
werkzeug==3.1.7
# via schemathesis
word2number==1.1
# via lm-eval
xgrammar==0.2.3
# via
# -c requirements/common.txt
# -r requirements/test/../common.txt
xxhash==3.6.0
# via
# datasets
+69 -20
View File
@@ -272,6 +272,18 @@ version = "1.1.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0"
[[package]]
name = "auto_enums"
version = "0.8.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2e4487600931c9a89f8db7ffbdf3fbdd45bb7bd85e26861f659a463cd0dff966"
dependencies = [
"derive_utils",
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]]
name = "auto_impl"
version = "1.3.0"
@@ -938,6 +950,17 @@ dependencies = [
"unicode-xid",
]
[[package]]
name = "derive_utils"
version = "0.15.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "362f47930db19fe7735f527e6595e4900316b893ebf6d48ad3d31be928d57dd6"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]]
name = "digest"
version = "0.10.7"
@@ -1478,9 +1501,9 @@ dependencies = [
[[package]]
name = "h2"
version = "0.4.13"
version = "0.4.15"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2f44da3a8150a6703ed5d34e164b875fd14c2cdab9af1252a9a1020bde2bdc54"
checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155"
dependencies = [
"atomic-waker",
"bytes",
@@ -1638,9 +1661,9 @@ checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9"
[[package]]
name = "hyper"
version = "1.8.1"
version = "1.10.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2ab2d4f250c3d7b1c9fcdff1cece94ea4e2dfbec68614f7b87cb205f24ca9d11"
checksum = "55281c53a1894c864990125767da440a4e630446785086f52523b20033b74498"
dependencies = [
"atomic-waker",
"bytes",
@@ -1653,7 +1676,6 @@ dependencies = [
"httpdate",
"itoa",
"pin-project-lite",
"pin-utils",
"smallvec",
"tokio",
"want",
@@ -2569,15 +2591,14 @@ dependencies = [
[[package]]
name = "openssl"
version = "0.10.76"
version = "0.10.81"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "951c002c75e16ea2c65b8c7e4d3d51d5530d8dfa7d060b4776828c88cfb18ecf"
checksum = "77823a27f0babb03091cb9ed9ef80af3b39dbc82f97e8fa530374b7dafd87a45"
dependencies = [
"bitflags",
"cfg-if",
"foreign-types",
"libc",
"once_cell",
"openssl-macros",
"openssl-sys",
]
@@ -2610,9 +2631,9 @@ dependencies = [
[[package]]
name = "openssl-sys"
version = "0.9.112"
version = "0.9.117"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "57d55af3b3e226502be1526dfdba67ab0e9c96fc293004e79576b2b9edb0dbdb"
checksum = "b47e7e6bb2c38cd930d25a23b40fa52e068c10e85f3e03a7f5ba5aaca5713695"
dependencies = [
"cc",
"libc",
@@ -2783,12 +2804,6 @@ version = "0.2.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd"
[[package]]
name = "pin-utils"
version = "0.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184"
[[package]]
name = "pkg-config"
version = "0.3.32"
@@ -2988,7 +3003,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "343d3bd7056eda839b03204e68deff7d1b13aba7af2b2fd16890697274262ee7"
dependencies = [
"heck",
"itertools 0.10.5",
"itertools 0.14.0",
"log",
"multimap",
"petgraph",
@@ -3009,7 +3024,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "27c6023962132f4b30eb4c172c91ce92d933da334c59c23cddee82358ddafb0b"
dependencies = [
"anyhow",
"itertools 0.10.5",
"itertools 0.14.0",
"proc-macro2",
"quote",
"syn 2.0.117",
@@ -3503,9 +3518,9 @@ dependencies = [
[[package]]
name = "rustls-pki-types"
version = "1.14.0"
version = "1.14.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "be040f8b0a225e40375822a563fa9524378b9d63112f53e19ffff34df5d33fdd"
checksum = "30a7197ae7eb376e574fe940d068c30fe0462554a3ddbe4eca7838e049c937a9"
dependencies = [
"zeroize",
]
@@ -4385,6 +4400,22 @@ dependencies = [
"serde_json",
]
[[package]]
name = "tls-listener"
version = "0.11.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1461056cc1ef47003f7ee16e4cef3741068d4c7f6b627bfce49b7c00c120a530"
dependencies = [
"axum",
"futures-util",
"openssl",
"pin-project-lite",
"thiserror 2.0.18",
"tokio",
"tokio-openssl",
"tracing",
]
[[package]]
name = "tokenizers"
version = "0.22.2"
@@ -4457,6 +4488,17 @@ dependencies = [
"tokio",
]
[[package]]
name = "tokio-openssl"
version = "0.6.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "59df6849caa43bb7567f9a36f863c447d95a11d5903c9cc334ba32576a27eadd"
dependencies = [
"openssl",
"openssl-sys",
"tokio",
]
[[package]]
name = "tokio-rustls"
version = "0.26.4"
@@ -5220,6 +5262,7 @@ dependencies = [
"anyhow",
"async-openai",
"asynk-strim-attr",
"auto_enums",
"axum",
"bytes",
"clap",
@@ -5227,10 +5270,13 @@ dependencies = [
"expect-test",
"futures",
"http-body",
"hyper",
"hyper-util",
"indexmap 2.13.0",
"itertools 0.14.0",
"libc",
"llm-multimodal",
"openssl",
"prost",
"prost-types",
"rmp-serde",
@@ -5242,8 +5288,11 @@ dependencies = [
"sha2",
"socket2",
"subtle",
"tempfile",
"thiserror-ext",
"tls-listener",
"tokio",
"tokio-openssl",
"tokio-stream",
"tokio-util",
"tonic",
+10
View File
@@ -26,6 +26,7 @@ arc-swap = "1.9.0"
async-openai = { version = "0.33.1", default-features = false, features = ["native-tls"] }
async-trait = "0.1.89"
asynk-strim-attr = "0.1.0"
auto_enums = { version = "0.8.9", features = ["tokio1"] }
axum = "0.8.8"
base64 = "0.22.1"
bytemuck = { version = "1.25.0", features = ["extern_crate_alloc"] }
@@ -43,6 +44,12 @@ half = { version = "2.7.1", features = ["bytemuck"] }
hex = "0.4.3"
hf-hub = { version = "0.5.0", default-features = false, features = ["tokio"] }
http-body = "1.0.1"
hyper = { version = "1.10.1", features = ["http1", "server"] }
hyper-util = { version = "0.1.20", features = [
"server-graceful",
"service",
"tokio",
] }
indexmap = "2.13.0"
itertools = "0.14.0"
libc = "0.2.177"
@@ -54,6 +61,7 @@ native-tls-vendored = { package = "native-tls", version = "0.2.18", features = [
ndarray = { version = "0.16.1", features = ["serde"] }
openai-harmony = { package = "oss-harmony", git = "https://github.com/oss-harmony/harmony", tag = "v0.0.11", default-features = false }
openai-protocol = "1.6.0"
openssl = "0.10"
parking_lot = "0.12.5"
paste = "1.0.15"
prometheus-client = "0.24.0"
@@ -89,6 +97,7 @@ thiserror = "2.0.16"
thiserror-ext = "0.3.0"
tiktoken-rs = "0.9.1"
time = { version = "0.3.47", features = ["formatting", "local-offset", "macros"] }
tls-listener = { version = "0.11.2", default-features = false, features = ["openssl", "tokio-net", "axum"] }
tokenizers = "0.22.0"
tokio = { version = "1.47.1", features = [
"macros",
@@ -97,6 +106,7 @@ tokio = { version = "1.47.1", features = [
"sync",
"time",
] }
tokio-openssl = "0.6"
tokio-stream = "0.1"
tokio-util = { version = "0.7.18", features = ["rt"] }
tonic = "0.14.5"
+66 -1
View File
@@ -25,7 +25,7 @@ use vllm_managed_engine::ManagedEngineConfig;
use vllm_managed_engine::cli::{ManagedEngineArgs, repartition_managed_engine_args};
use vllm_server::{
ApiServerOptions, ChatTemplateContentFormatOption, Config, CoordinatorMode, CorsConfig,
HttpListenerMode, ParserSelection, RendererSelection,
DEFAULT_KEEP_ALIVE_TIMEOUT, HttpListenerMode, ParserSelection, RendererSelection, TlsConfig,
};
use crate::cli::unsupported::UnsupportedArgs;
@@ -154,6 +154,11 @@ pub struct SharedRuntimeArgs {
#[arg(long, default_value_t = 0)]
#[serde(default)]
pub shutdown_timeout: u64,
/// Maximum idle time (seconds) on a keep-alive HTTP connection before the
/// server closes it (default 5).
#[arg(long = "http-timeout-keep-alive", env = "VLLM_HTTP_TIMEOUT_KEEP_ALIVE")]
#[serde(default)]
pub http_timeout_keep_alive: Option<u64>,
/// The file path to the chat template, or the template in single-line form
/// for the specified model.
@@ -257,6 +262,34 @@ pub struct SharedRuntimeArgs {
#[serde(default)]
pub allow_credentials: bool,
/// The file path to the SSL key file. When omitted, the key is read from
/// `--ssl-certfile` (combined PEM).
#[arg(long)]
#[serde(default)]
pub ssl_keyfile: Option<String>,
/// The file path to the SSL cert file. Enables TLS when set.
#[arg(long)]
#[serde(default)]
pub ssl_certfile: Option<String>,
/// The CA certificates file used to verify client certificates (mTLS).
#[arg(long)]
#[serde(default)]
pub ssl_ca_certs: Option<String>,
/// Whether a client certificate is required: 0 = none, 1 = optional,
/// 2 = required (mirrors Python's `ssl.CERT_*`).
#[arg(long, default_value_t = 0, value_parser = clap::value_parser!(i32).range(0..=2))]
#[serde(default)]
pub ssl_cert_reqs: i32,
/// OpenSSL cipher string for HTTPS (TLS 1.2 and below).
/// When unset, the linked OpenSSL's default suites are used.
#[arg(long)]
#[serde(default)]
pub ssl_ciphers: Option<String>,
/// Unsupported Python vLLM frontend arguments recognized but not yet
/// implemented in Rust.
#[educe(Debug(ignore))]
@@ -277,6 +310,13 @@ impl SharedRuntimeArgs {
Duration::from_secs(self.shutdown_timeout)
}
/// Maximum idle time on a keep-alive HTTP connection before the server
/// closes it.
pub fn keep_alive_timeout(&self) -> Duration {
self.http_timeout_keep_alive
.map_or(DEFAULT_KEEP_ALIVE_TIMEOUT, Duration::from_secs)
}
/// Apply fallback logic for API key configuration from env variables.
fn apply_env_api_key_fallback(&mut self) {
if self.api_key.is_empty()
@@ -301,8 +341,10 @@ impl SharedRuntimeArgs {
) -> Config {
let ready_timeout = self.ready_timeout();
let shutdown_timeout = self.shutdown_timeout();
let keep_alive_timeout = self.keep_alive_timeout();
let api_server_options = self.api_server_options();
let cors = self.cors_config();
let tls = self.tls_config();
Config {
transport_mode: TransportMode::Bootstrapped {
@@ -329,10 +371,12 @@ impl SharedRuntimeArgs {
max_logprobs: self.max_logprobs,
api_server_options,
cors,
tls,
api_keys: self.api_key,
disable_log_stats: self.disable_log_stats,
grpc_port: self.grpc_port,
shutdown_timeout,
keep_alive_timeout,
}
}
@@ -349,8 +393,10 @@ impl SharedRuntimeArgs {
) -> Config {
let ready_timeout = self.ready_timeout();
let shutdown_timeout = self.shutdown_timeout();
let keep_alive_timeout = self.keep_alive_timeout();
let api_server_options = self.api_server_options();
let cors = self.cors_config();
let tls = self.tls_config();
Config {
transport_mode: TransportMode::HandshakeOwner {
@@ -375,10 +421,12 @@ impl SharedRuntimeArgs {
max_logprobs: self.max_logprobs,
api_server_options,
cors,
tls,
api_keys: self.api_key,
disable_log_stats: self.disable_log_stats,
grpc_port: self.grpc_port,
shutdown_timeout,
keep_alive_timeout,
}
}
@@ -398,6 +446,23 @@ impl SharedRuntimeArgs {
allow_credentials: self.allow_credentials,
}
}
/// Build the TLS config: `Some` when any `ssl_*` argument is set, else
/// `None` (plaintext). The combination is validated in [`Config::validate`].
fn tls_config(&self) -> Option<TlsConfig> {
let tls_requested = self.ssl_certfile.is_some()
|| self.ssl_keyfile.is_some()
|| self.ssl_ca_certs.is_some()
|| self.ssl_cert_reqs != 0
|| self.ssl_ciphers.is_some();
tls_requested.then(|| TlsConfig {
cert_file: self.ssl_certfile.clone(),
key_file: self.ssl_keyfile.clone(),
ca_certs: self.ssl_ca_certs.clone(),
cert_reqs: self.ssl_cert_reqs,
ciphers: self.ssl_ciphers.clone(),
})
}
}
fn default_engine_ready_timeout_secs() -> u64 {
+167 -9
View File
@@ -41,6 +41,7 @@ fn serve_args_forward_python_flags_with_separator() {
max_logprobs: None,
grpc_port: None,
shutdown_timeout: 0,
http_timeout_keep_alive: None,
chat_template: None,
default_chat_template_kwargs: None,
chat_template_content_format: Auto,
@@ -65,6 +66,11 @@ fn serve_args_forward_python_flags_with_separator() {
],
),
allow_credentials: false,
ssl_keyfile: None,
ssl_certfile: None,
ssl_ca_certs: None,
ssl_cert_reqs: 0,
ssl_ciphers: None,
},
managed_engine: ManagedEngineArgs {
python: "../vllm/.venv/bin/python",
@@ -363,6 +369,140 @@ fn serve_passes_enable_prompt_tokens_details_into_config() {
assert!(config.api_server_options.enable_prompt_tokens_details);
}
#[test]
fn serve_passes_tls_into_config() {
let cli = Cli::try_parse_from([
"vllm-rs",
"serve",
"Qwen/Qwen3-0.6B",
"--ssl-certfile",
"/tmp/cert.pem",
"--ssl-keyfile",
"/tmp/key.pem",
"--ssl-ca-certs",
"/tmp/ca.pem",
"--ssl-cert-reqs",
"2",
])
.unwrap();
let Command::Serve(args) = cli.command else {
panic!("expected serve args");
};
let config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
let tls = config.tls.expect("tls configured");
assert_eq!(tls.cert_file.as_deref(), Some("/tmp/cert.pem"));
assert_eq!(tls.key_file.as_deref(), Some("/tmp/key.pem"));
assert_eq!(tls.ca_certs.as_deref(), Some("/tmp/ca.pem"));
assert_eq!(tls.cert_reqs, 2);
}
#[test]
fn serve_without_ssl_flags_has_no_tls() {
let cli = Cli::try_parse_from(["vllm-rs", "serve", "Qwen/Qwen3-0.6B"]).unwrap();
let Command::Serve(args) = cli.command else {
panic!("expected serve args");
};
let config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
assert!(config.tls.is_none());
}
#[test]
fn serve_ssl_keyfile_without_certfile_fails_validation() {
let cli = Cli::try_parse_from([
"vllm-rs",
"serve",
"Qwen/Qwen3-0.6B",
"--ssl-keyfile",
"/tmp/key.pem",
])
.unwrap();
let Command::Serve(args) = cli.command else {
panic!("expected serve args");
};
let config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
// TLS is requested (a key was given) but there is no certificate, so
// validation fails loud rather than silently serving plaintext.
assert_eq!(config.tls.as_ref().expect("tls requested").cert_file, None);
let err = config.validate().unwrap_err().to_string();
assert!(err.contains("--ssl-certfile is required"), "{err}");
}
#[test]
fn serve_mtls_without_ca_certs_fails_validation() {
let cli = Cli::try_parse_from([
"vllm-rs",
"serve",
"Qwen/Qwen3-0.6B",
"--ssl-certfile",
"/tmp/cert.pem",
"--ssl-cert-reqs",
"2",
])
.unwrap();
let Command::Serve(args) = cli.command else {
panic!("expected serve args");
};
let config = args.to_frontend_config("tcp://127.0.0.1:62100".to_string());
// Client-cert verification without a CA bundle has nothing to verify
// against, so it fails loud at startup.
let err = config.validate().unwrap_err().to_string();
assert!(err.contains("--ssl-ca-certs is required"), "{err}");
}
#[test]
fn frontend_args_json_passes_tls_into_config() {
let cli = Cli::try_parse_from([
"vllm-rs",
"frontend",
"--listen-fd",
"3",
"--input-address",
"ipc:///tmp/input.sock",
"--output-address",
"ipc:///tmp/output.sock",
"--args-json",
r#"{"model_tag":"Qwen/Qwen3-0.6B","ssl_certfile":"/tmp/cert.pem","ssl_keyfile":"/tmp/key.pem"}"#,
])
.unwrap();
let Command::Frontend(args) = cli.command else {
panic!("expected frontend args");
};
let config = args.into_config();
let tls = config.tls.expect("tls configured");
assert_eq!(tls.cert_file.as_deref(), Some("/tmp/cert.pem"));
assert_eq!(tls.key_file.as_deref(), Some("/tmp/key.pem"));
}
#[test]
fn frontend_args_json_rejects_out_of_range_cert_reqs() {
let cli = Cli::try_parse_from([
"vllm-rs",
"frontend",
"--listen-fd",
"3",
"--input-address",
"ipc:///tmp/input.sock",
"--output-address",
"ipc:///tmp/output.sock",
"--args-json",
r#"{"model_tag":"Qwen/Qwen3-0.6B","ssl_certfile":"/tmp/cert.pem","ssl_cert_reqs":5}"#,
])
.unwrap();
let Command::Frontend(args) = cli.command else {
panic!("expected frontend args");
};
// The JSON path bypasses clap's range check, so validate() is the only guard.
let config = args.into_config();
let err = config.validate().unwrap_err().to_string();
assert!(err.contains("--ssl-cert-reqs"), "{err}");
}
#[test]
fn frontend_args_json_passes_enable_request_id_headers_into_config() {
let cli = Cli::try_parse_from([
@@ -481,13 +621,13 @@ fn serve_args_reject_unsupported_flag_arg() {
"vllm-rs",
"serve",
"Qwen/Qwen3-0.6B",
"--ssl-keyfile",
"/tmp/key.pem",
"--root-path",
"/prefix",
])
.unwrap_err();
expect![[r#"
error: invalid value '/tmp/key.pem' for '--ssl-keyfile <SSL_KEYFILE>': argument is not implemented in Rust frontend yet
error: invalid value '/prefix' for '--root-path <ROOT_PATH>': argument is not implemented in Rust frontend yet
Remove this unsupported argument to continue.
@@ -562,6 +702,7 @@ fn frontend_args_accept_json() {
max_logprobs: None,
grpc_port: None,
shutdown_timeout: 0,
http_timeout_keep_alive: None,
chat_template: None,
default_chat_template_kwargs: None,
chat_template_content_format: Auto,
@@ -586,6 +727,11 @@ fn frontend_args_accept_json() {
],
),
allow_credentials: false,
ssl_keyfile: None,
ssl_certfile: None,
ssl_ca_certs: None,
ssl_cert_reqs: 0,
ssl_ciphers: None,
},
},
),
@@ -798,14 +944,14 @@ fn frontend_args_json_rejects_unsupported_fields() {
"--output-address",
"ipc:///tmp/output.sock",
"--args-json",
r#"{"model_tag":"Qwen/Qwen3-0.6B","ssl_keyfile":"/tmp/key.pem"}"#,
r#"{"model_tag":"Qwen/Qwen3-0.6B","root_path":"/prefix"}"#,
])
.unwrap_err();
expect![[r#"
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","ssl_keyfile":"/tmp/key.pem"}' for '--args-json <JSON>':
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","root_path":"/prefix"}' for '--args-json <JSON>':
The following arguments are not implemented in Rust frontend yet:
- ssl_keyfile
- root_path
Remove these arguments to continue.
@@ -825,16 +971,16 @@ fn frontend_args_json_aggregates_multiple_unsupported_fields() {
"--output-address",
"ipc:///tmp/output.sock",
"--args-json",
r#"{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","ssl_keyfile":"/tmp/key.pem"}"#,
r#"{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","root_path":"/prefix"}"#,
])
.unwrap_err();
let actual = error.to_string().replace(": \n", ":\n");
expect![[r#"
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","ssl_keyfile":"/tmp/key.pem"}' for '--args-json <JSON>':
error: invalid value '{"model_tag":"Qwen/Qwen3-0.6B","response_role":"assistant","root_path":"/prefix"}' for '--args-json <JSON>':
The following arguments are not implemented in Rust frontend yet:
- response_role
- ssl_keyfile
- root_path
Remove these arguments to continue.
@@ -1077,6 +1223,7 @@ fn serve_args_accept_handshake_aliases() {
max_logprobs: None,
grpc_port: None,
shutdown_timeout: 0,
http_timeout_keep_alive: None,
chat_template: None,
default_chat_template_kwargs: None,
chat_template_content_format: Auto,
@@ -1101,6 +1248,11 @@ fn serve_args_accept_handshake_aliases() {
],
),
allow_credentials: false,
ssl_keyfile: None,
ssl_certfile: None,
ssl_ca_certs: None,
ssl_cert_reqs: 0,
ssl_ciphers: None,
},
managed_engine: ManagedEngineArgs {
python: "python3",
@@ -1234,10 +1386,12 @@ fn serve_frontend_config_uses_dp_address_as_advertised_host() {
],
allow_credentials: false,
},
tls: None,
api_keys: [],
disable_log_stats: false,
grpc_port: None,
shutdown_timeout: 0ns,
keep_alive_timeout: 5s,
}
"#]]
.assert_debug_eq(&Config {
@@ -1315,10 +1469,12 @@ fn serve_frontend_config_keeps_tcp_transport_for_non_local_only_topology() {
],
allow_credentials: false,
},
tls: None,
api_keys: [],
disable_log_stats: false,
grpc_port: None,
shutdown_timeout: 0ns,
keep_alive_timeout: 5s,
}
"#]]
.assert_debug_eq(&config);
@@ -1414,10 +1570,12 @@ fn frontend_config_uses_external_coordinator_when_coordinator_address_is_present
],
allow_credentials: false,
},
tls: None,
api_keys: [],
disable_log_stats: false,
grpc_port: None,
shutdown_timeout: 0ns,
keep_alive_timeout: 5s,
}
"#]]
.assert_debug_eq(&config);
-21
View File
@@ -526,18 +526,6 @@ pub struct ServerUnsupportedArgs {
#[arg(long)]
pub disable_access_log_for_endpoints: Option<Noop>,
/// The file path to the SSL key file.
#[arg(long)]
pub ssl_keyfile: Option<Unsupported>,
/// The file path to the SSL cert file.
#[arg(long)]
pub ssl_certfile: Option<Unsupported>,
/// The CA certificates file.
#[arg(long)]
pub ssl_ca_certs: Option<Unsupported>,
/// Refresh SSL Context when SSL certificate files change
#[arg(
long,
@@ -547,15 +535,6 @@ pub struct ServerUnsupportedArgs {
)]
pub enable_ssl_refresh: Option<Unsupported>,
/// Whether client certificate is required (see stdlib ssl module's).
#[arg(long)]
pub ssl_cert_reqs: Option<Unsupported>,
/// SSL cipher suites for HTTPS (TLS 1.2 and below only).
/// Example: 'ECDHE-RSA-AES256-GCM-SHA384:ECDHE-RSA-CHACHA20-POLY1305'
#[arg(long)]
pub ssl_ciphers: Option<Unsupported>,
/// FastAPI root_path when app is behind a path based routing proxy.
#[arg(long)]
pub root_path: Option<Unsupported>,
@@ -100,6 +100,7 @@ impl EngineRoutingState {
pub struct RequestRegistry {
closed: bool,
requests: HashMap<String, TrackedRequest>,
active_lora_requests: usize,
routing_per_engine: BTreeMap<EngineId, EngineRoutingState>,
}
@@ -108,6 +109,7 @@ impl RequestRegistry {
Self {
closed: false,
requests: HashMap::default(),
active_lora_requests: 0,
routing_per_engine: engines
.iter()
.map(|engine| (engine.engine_id.clone(), EngineRoutingState::default()))
@@ -133,15 +135,19 @@ impl RequestRegistry {
let engine_id = self.choose_engine_for_request(data_parallel_rank)?;
let (tx, rx) = mpsc::unbounded_channel();
let lora = lora_name.map(|adapter_name| LoraRequestState {
adapter_name,
phase: LoraPhase::Waiting,
});
if lora.is_some() {
self.active_lora_requests += 1;
}
self.requests.insert(
request_id,
TrackedRequest {
sender: tx,
engine_id: engine_id.clone(),
lora: lora_name.map(|adapter_name| LoraRequestState {
adapter_name,
phase: LoraPhase::Waiting,
}),
lora,
},
);
@@ -230,6 +236,10 @@ impl RequestRegistry {
/// Snapshot the adapter names of tracked LoRA requests as
/// (running, waiting) sets. Feeds the `vllm:lora_requests_info` gauge.
pub fn lora_adapter_states(&self) -> (BTreeSet<String>, BTreeSet<String>) {
if self.active_lora_requests == 0 {
return (BTreeSet::new(), BTreeSet::new());
}
let mut running = BTreeSet::new();
let mut waiting = BTreeSet::new();
for lora in self.requests.values().filter_map(|tracked| tracked.lora.as_ref()) {
@@ -283,6 +293,7 @@ impl RequestRegistry {
}
self.closed = true;
self.active_lora_requests = 0;
std::mem::take(&mut self.requests)
.into_values()
.map(|tracked| tracked.sender)
@@ -322,6 +333,9 @@ impl RequestRegistry {
#[must_use]
pub fn remove(&mut self, request_id: &str) -> Option<(OutputSender, EngineId)> {
let tracked = self.requests.remove(request_id)?;
if tracked.lora.is_some() {
self.active_lora_requests -= 1;
}
self.routing_per_engine
.get_mut(&tracked.engine_id)
.expect("request registry must track all known engines")
@@ -359,6 +373,11 @@ impl RequestRegistry {
pub fn is_closed(&self) -> bool {
self.closed
}
#[cfg(test)]
fn active_lora_requests(&self) -> usize {
self.active_lora_requests
}
}
/// Internal registry for tracking active utility calls and their waiting
@@ -574,6 +593,63 @@ mod tests {
);
}
#[test]
fn registry_counts_only_active_lora_requests() {
let mut registry = RequestRegistry::new(&[connected_engine(EngineId::from(b"engine-0"))]);
registry.register("req-plain".to_string(), None, None).unwrap();
assert_eq!(registry.active_lora_requests(), 0);
assert_eq!(
registry.lora_adapter_states(),
(adapter_names(&[]), adapter_names(&[]))
);
registry
.register(
"req-lora-a".to_string(),
Some("adapter-a".to_string()),
None,
)
.unwrap();
registry
.register(
"req-lora-b".to_string(),
Some("adapter-b".to_string()),
None,
)
.unwrap();
assert_eq!(registry.active_lora_requests(), 2);
drop(registry.remove("req-plain"));
assert_eq!(registry.active_lora_requests(), 2);
drop(registry.finish_many(&["req-lora-a".to_string()]));
assert_eq!(registry.active_lora_requests(), 1);
drop(registry.abort_many(&["req-lora-b".to_string()], 0.0));
assert_eq!(registry.active_lora_requests(), 0);
assert_eq!(
registry.lora_adapter_states(),
(adapter_names(&[]), adapter_names(&[]))
);
}
#[test]
fn registry_clears_lora_count_on_close() {
let mut registry = RequestRegistry::new(&[connected_engine(EngineId::from(b"engine-0"))]);
registry
.register("req-lora".to_string(), Some("adapter-a".to_string()), None)
.unwrap();
assert_eq!(registry.active_lora_requests(), 1);
drop(registry.close());
assert_eq!(registry.active_lora_requests(), 0);
assert_eq!(
registry.lora_adapter_states(),
(adapter_names(&[]), adapter_names(&[]))
);
}
#[test]
fn registry_drops_lora_tracking_on_abort() {
let mut registry = RequestRegistry::new(&[connected_engine(EngineId::from(b"engine-0"))]);
@@ -20,6 +20,27 @@ pub(crate) struct CoordinatorStateSnapshot {
pub engines_running: bool,
}
impl CoordinatorStateSnapshot {
/// Resume the engines for a `FirstRequest` and return the wave to broadcast
/// and the engine to exclude from the wakeup.
///
/// The request may have been stamped with a `request_wave` older than
/// `current_wave` if a `WaveComplete` advanced it after the command was
/// enqueued. Such a request still needs serving, so the current wave is
/// broadcast to every engine (`exclude = None`); the wave is never rewound.
/// A non-stale request excludes the engine that already received it. Mirrors
/// the Python coordinator's front-end path.
pub(crate) fn start_wave_for_first_request(
&mut self,
request_wave: u32,
target_engine_index: u32,
) -> (u32, Option<u32>) {
self.engines_running = true;
let exclude = (request_wave >= self.current_wave).then_some(target_engine_index);
(self.current_wave, exclude)
}
}
/// Shared in-process coordinator state.
pub(crate) type CoordinatorState = Mutex<CoordinatorStateSnapshot>;
@@ -27,9 +27,10 @@ use crate::protocol::{
struct StartDpWaveMessage {
/// DP wave number that all engines should start processing.
wave: u32,
/// Engine index that already received the triggering request and should not
/// receive an extra wakeup notification.
exclude_engine_index: u32,
/// Engine index that already received the triggering request and so does not
/// need an extra wakeup. `None` wakes every engine (used when the triggering
/// request was for a stale wave).
exclude_engine_index: Option<u32>,
}
/// Background half of the in-process coordinator.
@@ -57,7 +58,11 @@ impl InProcCoordinatorRunner {
}
/// Broadcast Python-compatible `START_DP_WAVE` to all connected engines.
async fn broadcast_start_wave(&mut self, wave: u32, exclude_engine_index: u32) -> Result<()> {
async fn broadcast_start_wave(
&mut self,
wave: u32,
exclude_engine_index: Option<u32>,
) -> Result<()> {
let payload = encode_msgpack(&StartDpWaveMessage {
wave,
exclude_engine_index,
@@ -86,13 +91,17 @@ impl InProcCoordinatorRunner {
engine_id: target_engine_id.to_vec(),
}
})?;
self.state.lock().current_wave = wave;
let (current_wave, exclude) = {
let mut state = self.state.lock();
state.start_wave_for_first_request(wave, target_engine_index)
};
debug!(
wave,
exclude_engine_index = target_engine_index,
current_wave,
request_wave = wave,
?exclude,
"starting DP wave after first request while engines were paused"
);
self.broadcast_start_wave(wave, target_engine_index).await?;
self.broadcast_start_wave(current_wave, exclude).await?;
}
}
Ok(())
@@ -150,7 +159,7 @@ impl InProcCoordinatorRunner {
exclude_engine_index = engine_index,
"starting DP wave after stale-wave notification from engine"
);
self.broadcast_start_wave(wave, engine_index).await?;
self.broadcast_start_wave(wave, Some(engine_index)).await?;
}
}
},
@@ -202,3 +211,48 @@ impl InProcCoordinatorRunner {
inner.close_registries(Arc::new(error));
}
}
#[cfg(test)]
mod tests {
use crate::coordinator::handle::CoordinatorStateSnapshot;
/// A `FirstRequest` for the current wave starts that wave and excludes the
/// engine that already received the triggering request.
#[test]
fn first_request_for_current_wave_excludes_target() {
let mut state = CoordinatorStateSnapshot {
current_wave: 3,
engines_running: false,
};
let (wave, exclude) = state.start_wave_for_first_request(3, 2);
assert_eq!(wave, 3);
assert_eq!(exclude, Some(2));
assert!(state.engines_running);
assert_eq!(state.current_wave, 3);
}
/// A `FirstRequest` whose wave was superseded by a racing `WaveComplete`
/// (`request_wave < current_wave`) must still start the request's wave: it
/// broadcasts the current wave and wakes every engine (`exclude = None`)
/// rather than rewinding the wave or dropping the request.
#[test]
fn stale_first_request_starts_current_wave_for_all_engines() {
let mut state = CoordinatorStateSnapshot {
current_wave: 4,
engines_running: false,
};
// Request stamped with wave 3 while the coordinator already advanced to 4.
let (wave, exclude) = state.start_wave_for_first_request(3, 2);
assert_eq!(
wave, 4,
"must broadcast the current wave, not the stale one"
);
assert_eq!(exclude, None, "a stale request must wake every engine");
assert!(state.engines_running);
assert_eq!(state.current_wave, 4, "wave must not be rewound");
}
}
@@ -1285,18 +1285,24 @@ async fn dropping_multiple_live_streams_aborts_all_in_a_burst() {
)
.await;
let abort =
timeout(Duration::from_secs(1), recv_engine_message(dealer)).await.unwrap();
assert_eq!(abort[0].as_ref(), &[0x01]);
let ids: Vec<String> = rmp_serde::from_slice(&abort[1]).unwrap();
// Aborts may coalesce into one burst or split across several.
let mut aborted = BTreeSet::new();
while aborted.len() < 3 {
let abort =
timeout(Duration::from_secs(1), recv_engine_message(dealer)).await.unwrap();
assert_eq!(abort[0].as_ref(), &[0x01]);
let ids: Vec<String> = rmp_serde::from_slice(&abort[1]).unwrap();
aborted.extend(ids);
}
assert_eq!(
ids,
vec![
aborted,
BTreeSet::from([
"req-1".to_string(),
"req-2".to_string(),
"req-3".to_string()
]
])
);
// No spurious extra aborts.
assert!(
timeout(Duration::from_millis(100), recv_engine_message(dealer)).await.is_err()
);
+7
View File
@@ -7,14 +7,18 @@ license.workspace = true
[dependencies]
anyhow.workspace = true
asynk-strim-attr.workspace = true
auto_enums.workspace = true
axum.workspace = true
educe.workspace = true
futures.workspace = true
http-body.workspace = true
hyper.workspace = true
hyper-util.workspace = true
indexmap.workspace = true
itertools.workspace = true
libc.workspace = true
llm-multimodal.workspace = true
openssl.workspace = true
prost.workspace = true
prost-types.workspace = true
rmpv.workspace = true
@@ -25,7 +29,9 @@ sha2.workspace = true
socket2.workspace = true
subtle.workspace = true
thiserror-ext.workspace = true
tls-listener.workspace = true
tokio.workspace = true
tokio-openssl.workspace = true
tokio-stream.workspace = true
tokio-util.workspace = true
tonic.workspace = true
@@ -54,6 +60,7 @@ clap.workspace = true
expect-test.workspace = true
rmp-serde.workspace = true
serial_test.workspace = true
tempfile.workspace = true
tower.workspace = true
vllm-engine-core-client = { workspace = true, features = ["test-util"] }
zeromq.workspace = true
@@ -71,10 +71,12 @@ async fn main() -> Result<()> {
max_logprobs: None,
api_server_options: ApiServerOptions::default(),
cors: CorsConfig::default(),
tls: None,
api_keys: Vec::new(),
disable_log_stats: false,
grpc_port: None,
shutdown_timeout: Duration::ZERO,
keep_alive_timeout: Duration::from_secs(5),
};
let bind_address = format!("127.0.0.1:{port}");
+61
View File
@@ -10,6 +10,10 @@ use serde_json::Value;
use vllm_chat::{ChatTemplateContentFormatOption, ParserSelection, RendererSelection};
use vllm_engine_core_client::{CoordinatorMode as EngineCoreCoordinatorMode, TransportMode};
/// Default keep-alive idle timeout (seconds); also the head-read bound
/// when keep-alive is disabled (`0`).
pub const DEFAULT_KEEP_ALIVE_TIMEOUT: Duration = Duration::from_secs(5);
/// How the HTTP server obtains its listening socket.
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub enum HttpListenerMode {
@@ -99,6 +103,54 @@ impl CorsConfig {
}
}
/// TLS settings mirroring Python's uvicorn `ssl_*` arguments.
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct TlsConfig {
/// PEM certificate chain file. Required when TLS is configured; may also
/// hold the private key (combined PEM) when `key_file` is unset.
pub cert_file: Option<String>,
/// PEM private key file. When `None`, the key is read from `cert_file`
/// (combined PEM).
pub key_file: Option<String>,
/// PEM CA bundle used to verify client certificates (mTLS). Required when
/// `cert_reqs` is non-zero.
pub ca_certs: Option<String>,
/// Client-certificate requirement, mirroring Python's `ssl.CERT_*`:
/// 0 = none, 1 = optional, 2 = required.
pub cert_reqs: i32,
/// OpenSSL cipher string for TLS 1.2 and below, mirroring Python's
/// `ssl.set_ciphers`. `None` keeps the forward-secret AEAD default.
pub ciphers: Option<String>,
}
impl TlsConfig {
/// Structurally validate the TLS arguments; the cert/key material is parsed
/// later, when the OpenSSL context is built.
pub fn validate(&self) -> Result<()> {
if self.cert_file.is_none() {
bail!(
"--ssl-certfile is required to enable TLS; \
--ssl-keyfile/--ssl-ca-certs/--ssl-cert-reqs/--ssl-ciphers \
cannot be used without it"
);
}
if !matches!(self.cert_reqs, 0..=2) {
bail!(
"--ssl-cert-reqs must be 0 (none), 1 (optional), or 2 (required), got {}",
self.cert_reqs
);
}
if self.cert_reqs != 0 && self.ca_certs.is_none() {
bail!(
"--ssl-ca-certs is required when --ssl-cert-reqs is {} \
(client certificate verification)",
self.cert_reqs
);
}
Ok(())
}
}
/// Normalized runtime configuration for the minimal OpenAI-compatible server.
#[derive(Educe, Clone, PartialEq, Eq, Serialize)]
#[educe(Debug)]
@@ -138,6 +190,9 @@ pub struct Config {
pub api_server_options: ApiServerOptions,
/// CORS settings applied to every HTTP response.
pub cors: CorsConfig,
/// TLS settings. `None` serves plaintext HTTP; `Some` terminates TLS at the
/// listener.
pub tls: Option<TlsConfig>,
/// API keys accepted as bearer tokens for guarded routes.
#[serde(skip_serializing)]
#[educe(Debug(method(fmt_redacted_api_keys)))]
@@ -150,6 +205,9 @@ pub struct Config {
pub grpc_port: Option<u16>,
/// Maximum time to wait for active HTTP/gRPC requests to drain on shutdown.
pub shutdown_timeout: Duration,
/// Maximum idle time on a keep-alive HTTP connection before the server
/// closes it (`VLLM_HTTP_TIMEOUT_KEEP_ALIVE`, default 5s).
pub keep_alive_timeout: Duration,
}
impl Config {
@@ -158,6 +216,9 @@ impl Config {
pub fn validate(&self) -> Result<()> {
vllm_chat::validate_parser_overrides(&self.tool_call_parser, &self.reasoning_parser)?;
self.cors.validate()?;
if let Some(tls) = &self.tls {
tls.validate()?;
}
if let Some(max_logprobs) = self.max_logprobs
&& max_logprobs < -1
{
+78 -1
View File
@@ -4,16 +4,21 @@ mod convert;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use futures::{Stream, StreamExt as _};
use futures::{Stream, StreamExt as _, stream};
use thiserror_ext::AsReport as _;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::sync::mpsc;
use tokio_openssl::SslStream;
use tokio_stream::wrappers::ReceiverStream;
use tonic::transport::server::{Connected, TcpConnectInfo};
use tonic::{Request, Response, Status};
use tracing::info;
use vllm_text::{DecodedTextEvent, TextOutputStreamExt as _};
use self::convert::ResponseOpts;
use crate::listener::{Listener, ListenerIo};
use crate::state::AppState;
/// Generated protobuf/gRPC types for the `vllm` package.
@@ -26,6 +31,78 @@ pub use pb::generate_server::GenerateServer;
#[cfg(test)]
mod tests;
/// Newtype over `tokio-openssl`'s `SslStream` so we can implement tonic's
/// [`Connected`] on it (the orphan rule blocks doing so on the foreign type).
pub(crate) struct GrpcTlsStream {
inner: SslStream<ListenerIo>,
}
impl GrpcTlsStream {
pub(crate) fn new(inner: SslStream<ListenerIo>) -> Self {
Self { inner }
}
}
impl AsyncRead for GrpcTlsStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().inner).poll_read(cx, buf)
}
}
impl AsyncWrite for GrpcTlsStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.get_mut().inner).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().inner).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().inner).poll_shutdown(cx)
}
}
impl Connected for GrpcTlsStream {
type ConnectInfo = TcpConnectInfo;
fn connect_info(&self) -> TcpConnectInfo {
self.inner.get_ref().connect_info()
}
}
/// Adapt the shared server listener into tonic's incoming stream shape.
pub(crate) fn incoming(listener: Listener) -> impl Stream<Item = std::io::Result<ListenerIo>> {
stream::unfold(listener, |mut listener| async move {
let (io, _) = axum::serve::Listener::accept(&mut listener).await;
Some((Ok(io), listener))
})
}
/// Wrap the gRPC listener so each accepted connection completes a TLS handshake
/// before tonic serves it.
pub(crate) fn tls_incoming(
listener: Listener,
context: openssl::ssl::SslContext,
handshake_timeout: std::time::Duration,
) -> impl Stream<Item = std::io::Result<GrpcTlsStream>> {
tls_listener::builder(context)
.handshake_timeout(handshake_timeout)
.listen(listener)
.map(|res| {
res.map(|(inner, _addr)| GrpcTlsStream::new(inner))
.map_err(std::io::Error::other)
})
}
/// gRPC Generate service implementation backed by the shared application state.
pub struct GenerateServiceImpl {
state: Arc<AppState>,
+357 -15
View File
@@ -1,11 +1,19 @@
use std::future::Future;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use futures::StreamExt as _;
use hyper_util::rt::TokioIo;
use openssl::ssl::{SslConnector, SslFiletype, SslMethod};
use serial_test::serial;
use tonic::transport::Server as TonicServer;
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
use tokio::net::TcpStream;
use tokio_openssl::SslStream;
use tonic::transport::{Channel, Endpoint, Server as TonicServer, Uri};
use tower::service_fn;
use vllm_chat::{
ChatBackend, ChatLlm, ChatRenderer, ChatRequest, ChatTextBackend, DefaultChatOutputProcessor,
DynChatOutputProcessor, DynChatRenderer, NewChatOutputProcessorOptions, RenderedPrompt,
@@ -22,8 +30,11 @@ use zeromq::prelude::{SocketRecv, SocketSend};
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
use super::pb::generate_client::GenerateClient;
use super::{GenerateServer, GenerateServiceImpl, pb};
use super::{GenerateServer, GenerateServiceImpl, incoming, pb, tls_incoming};
use crate::listener::Listener;
use crate::state::AppState;
use crate::tls;
use crate::tls_tests::{TestCerts, server_tls};
// ========================================================================================
// Helpers (mirrors the patterns in routes/tests.rs)
@@ -211,17 +222,12 @@ impl ChatRenderer for FakeTextBackend {
}
}
/// Spin up a gRPC server backed by a mock engine that serves a single request
/// with the given output specs. Returns the client, the gRPC server task, and
/// the mock engine task.
async fn grpc_test_server(
/// Build the gRPC service + mock engine that serves a single request with the
/// given output specs. Shared by the plaintext and TLS server fixtures.
async fn setup_grpc_service(
engine_id: impl Into<EngineId>,
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
) -> (
GenerateClient<tonic::transport::Channel>,
tokio::task::JoinHandle<()>,
MockEngineTask,
) {
) -> (GenerateServer<GenerateServiceImpl>, MockEngineTask) {
let ipc = IpcNamespace::new().expect("create ipc namespace");
let handshake_address = ipc.handshake_endpoint();
let engine_id = engine_id.into();
@@ -259,14 +265,29 @@ async fn grpc_test_server(
Arc::new(FakeTextBackend) as Arc<dyn ChatTextBackend>,
);
let state = Arc::new(AppState::new(vec!["test-model".to_string()], chat));
let svc = GenerateServer::new(GenerateServiceImpl::new(state));
(
GenerateServer::new(GenerateServiceImpl::new(state)),
engine_task,
)
}
/// Spin up a plaintext gRPC server backed by a mock engine. Returns the client,
/// the gRPC server task, and the mock engine task.
async fn grpc_test_server(
engine_id: impl Into<EngineId>,
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
) -> (
GenerateClient<tonic::transport::Channel>,
tokio::task::JoinHandle<()>,
MockEngineTask,
) {
let (svc, engine_task) = setup_grpc_service(engine_id, output_specs).await;
// Bind to an OS-assigned port.
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
let addr = listener.local_addr().expect("local addr");
let server_task = tokio::spawn(async move {
let incoming = tokio_stream::wrappers::TcpListenerStream::new(listener);
let incoming = incoming(Listener::Tcp(listener));
TonicServer::builder()
.add_service(svc)
.serve_with_incoming(incoming)
@@ -274,7 +295,6 @@ async fn grpc_test_server(
.expect("grpc server");
});
// Connect the client.
let grpc_client = GenerateClient::connect(format!("http://{addr}"))
.await
.expect("connect grpc client");
@@ -282,6 +302,158 @@ async fn grpc_test_server(
(grpc_client, server_task, engine_task)
}
/// Spin up a TLS gRPC server (server cert from `certs`, `cert_reqs` mTLS mode).
/// Returns the address, the server task, and the mock engine task.
async fn grpc_tls_test_server(
engine_id: impl Into<EngineId>,
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
certs: &TestCerts,
cert_reqs: i32,
) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) {
let (svc, engine_task) = setup_grpc_service(engine_id, output_specs).await;
let context = tls::build_grpc_server_config(&server_tls(certs, cert_reqs))
.expect("build grpc tls config");
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
let addr = listener.local_addr().expect("local addr").to_string();
let server_task = tokio::spawn(async move {
let incoming = tls_incoming(Listener::Tcp(listener), context, tls::TLS_HANDSHAKE_TIMEOUT);
TonicServer::builder()
.add_service(svc)
.serve_with_incoming(incoming)
.await
.expect("grpc tls server");
});
(addr, server_task, engine_task)
}
/// Build a tonic `Generate` client over a tokio-openssl connector, optionally
/// with a client identity for mTLS. Hand-rolled because tonic 0.14 ships no
/// OpenSSL transport.
async fn grpc_tls_client(
certs: &TestCerts,
addr: &str,
identity: Option<&str>,
) -> Result<GenerateClient<Channel>, tonic::transport::Error> {
let ca = certs.path("ca.pem");
let identity = identity.map(|name| {
(
certs.path(&format!("{name}.pem")),
certs.path(&format!("{name}.key")),
)
});
let target = addr.to_string();
let connector = service_fn(move |_: Uri| {
let ca = ca.clone();
let identity = identity.clone();
let target = target.clone();
async move {
let tcp = TcpStream::connect(&target).await?;
let mut builder =
SslConnector::builder(SslMethod::tls_client()).map_err(io::Error::other)?;
builder.set_ca_file(&ca).map_err(io::Error::other)?;
if let Some((cert, key)) = &identity {
builder.set_certificate_chain_file(cert).map_err(io::Error::other)?;
builder.set_private_key_file(key, SslFiletype::PEM).map_err(io::Error::other)?;
}
let mut config = builder.build().configure().map_err(io::Error::other)?;
config.set_verify_hostname(false);
config.set_alpn_protos(b"\x02h2").map_err(io::Error::other)?;
let ssl = config.into_ssl("127.0.0.1").map_err(io::Error::other)?;
let mut stream = SslStream::new(ssl, tcp).map_err(io::Error::other)?;
Pin::new(&mut stream).connect().await.map_err(io::Error::other)?;
Ok::<_, io::Error>(TokioIo::new(stream))
}
});
let channel = Endpoint::from_shared(format!("https://{addr}"))
.expect("grpc endpoint")
.connect_with_connector(connector)
.await?;
Ok(GenerateClient::new(channel))
}
/// Complete a raw TLS handshake against the gRPC port (offering ALPN `h2`) for
/// the ALPN-negotiation assertion.
async fn grpc_tls_handshake(
certs: &TestCerts,
addr: &str,
) -> io::Result<Pin<Box<SslStream<TcpStream>>>> {
let tcp = TcpStream::connect(addr).await?;
let mut builder = SslConnector::builder(SslMethod::tls_client()).map_err(io::Error::other)?;
builder.set_ca_file(certs.path("ca.pem")).map_err(io::Error::other)?;
let mut config = builder.build().configure().map_err(io::Error::other)?;
config.set_verify_hostname(false);
config.set_alpn_protos(b"\x02h2").map_err(io::Error::other)?;
let ssl = config.into_ssl("127.0.0.1").map_err(io::Error::other)?;
let mut stream = Box::pin(SslStream::new(ssl, tcp).map_err(io::Error::other)?);
stream.as_mut().connect().await.map_err(io::Error::other)?;
Ok(stream)
}
/// Spin up a plaintext gRPC server, optionally with HTTP/2 keepalive set to
/// `keepalive` for both the PING interval and the unanswered-PING timeout.
async fn grpc_server_with_keepalive(
engine_id: impl Into<EngineId>,
keepalive: Option<Duration>,
) -> (String, tokio::task::JoinHandle<()>, MockEngineTask) {
let (svc, engine_task) = setup_grpc_service(engine_id, default_stream_output_specs()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind grpc listener");
let addr = listener.local_addr().expect("local addr").to_string();
let mut builder = TonicServer::builder();
if let Some(interval) = keepalive {
builder = builder
.http2_keepalive_interval(Some(interval))
.http2_keepalive_timeout(Some(interval));
}
let server_task = tokio::spawn(async move {
let incoming = incoming(Listener::Tcp(listener));
builder
.add_service(svc)
.serve_with_incoming(incoming)
.await
.expect("grpc server");
});
(addr, server_task, engine_task)
}
/// Establish an HTTP/2 connection (preface + SETTINGS exchange) then go silent,
/// ACKing the server's SETTINGS but never its keepalive PINGs. Returns whether
/// the SERVER closes the connection within `wait`. A minimal hand-rolled h2 peer
/// because a real client auto-ACKs PINGs and so can never be kept-alive-evicted.
async fn h2_unresponsive_peer_closed_within(addr: &str, wait: Duration) -> bool {
let mut tcp = TcpStream::connect(addr).await.expect("connect");
tcp.write_all(b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n").await.expect("preface");
tcp.write_all(&[0, 0, 0, 0x4, 0, 0, 0, 0, 0]).await.expect("client settings");
let closed = tokio::time::timeout(wait, async {
let mut header = [0u8; 9];
while tcp.read_exact(&mut header).await.is_ok() {
let len = u32::from_be_bytes([0, header[0], header[1], header[2]]) as usize;
let frame_type = header[3];
let flags = header[4];
let mut payload = vec![0u8; len];
if tcp.read_exact(&mut payload).await.is_err() {
return;
}
// ACK the server's SETTINGS so the only thing left unanswered is PINGs.
if frame_type == 0x4 && flags & 0x1 == 0 {
let _ = tcp.write_all(&[0, 0, 0, 0x4, 0x1, 0, 0, 0, 0]).await;
}
}
})
.await;
closed.is_ok()
}
// ========================================================================================
// Tests
// ========================================================================================
@@ -720,3 +892,173 @@ async fn unary_generate_output_text_defaults_to_true() {
engine_task.await.expect("mock engine task");
server_task.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn grpc_generate_succeeds_over_tls() {
let certs = TestCerts::generate();
let (addr, server_task, engine_task) = grpc_tls_test_server(
b"engine-grpc-tls-unary",
default_stream_output_specs(),
&certs,
0,
)
.await;
let mut client = grpc_tls_client(&certs, &addr, None).await.expect("tls client");
let response = client
.generate(pb::GenerateRequest {
request_id: "test-tls-unary".to_string(),
model: "test-model".to_string(),
prompt: Some(pb::generate_request::Prompt::Text("hello".to_string())),
stopping: Some(pb::StoppingCriteria {
max_new_tokens: 10,
..Default::default()
}),
response: Some(pb::ResponseOptions {
output_text: Some(true),
..Default::default()
}),
..Default::default()
})
.await
.expect("unary generate over tls")
.into_inner();
assert_eq!(response.outputs.expect("outputs present").text, "hi");
engine_task.await.expect("mock engine task");
server_task.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn grpc_tls_negotiates_h2_alpn() {
let certs = TestCerts::generate();
let (addr, server_task, _engine_task) = grpc_tls_test_server(
b"engine-grpc-tls-alpn",
default_stream_output_specs(),
&certs,
0,
)
.await;
let stream = grpc_tls_handshake(&certs, &addr).await.expect("handshake");
assert_eq!(
stream.ssl().selected_alpn_protocol(),
Some(&b"h2"[..]),
"server must negotiate h2 ALPN"
);
server_task.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn grpc_mtls_required_rejects_client_without_certificate() {
let certs = TestCerts::generate();
let (addr, server_task, _engine_task) = grpc_tls_test_server(
b"engine-grpc-tls-mtls-reject",
default_stream_output_specs(),
&certs,
2,
)
.await;
// With TLS 1.3 the missing-client-cert rejection surfaces on first use, not
// at the handshake, so drive an RPC and assert the call fails.
let outcome = match grpc_tls_client(&certs, &addr, None).await {
Err(_) => Err(()),
Ok(mut client) => client
.generate(pb::GenerateRequest {
request_id: "test-tls-mtls-reject".to_string(),
model: "test-model".to_string(),
prompt: Some(pb::generate_request::Prompt::Text("hello".to_string())),
stopping: Some(pb::StoppingCriteria {
max_new_tokens: 10,
..Default::default()
}),
..Default::default()
})
.await
.map(|_| ())
.map_err(|_| ()),
};
assert!(
outcome.is_err(),
"mTLS-required gRPC must reject a client without a certificate"
);
server_task.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn grpc_mtls_required_accepts_valid_client_certificate() {
let certs = TestCerts::generate();
let (addr, server_task, engine_task) = grpc_tls_test_server(
b"engine-grpc-tls-mtls-accept",
default_stream_output_specs(),
&certs,
2,
)
.await;
let mut client = grpc_tls_client(&certs, &addr, Some("client")).await.expect("mtls client");
let response = client
.generate(pb::GenerateRequest {
request_id: "test-tls-mtls".to_string(),
model: "test-model".to_string(),
prompt: Some(pb::generate_request::Prompt::Text("hello".to_string())),
stopping: Some(pb::StoppingCriteria {
max_new_tokens: 10,
..Default::default()
}),
response: Some(pb::ResponseOptions {
output_text: Some(true),
..Default::default()
}),
..Default::default()
})
.await
.expect("mtls generate over tls")
.into_inner();
assert_eq!(response.outputs.expect("outputs present").text, "hi");
engine_task.await.expect("mock engine task");
server_task.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn grpc_keepalive_closes_unresponsive_connection() {
let (addr, server_task, _engine_task) =
grpc_server_with_keepalive(b"engine-grpc-keepalive", Some(Duration::from_millis(150)))
.await;
let closed = h2_unresponsive_peer_closed_within(&addr, Duration::from_secs(5)).await;
assert!(
closed,
"keepalive must close a peer that stops answering PINGs"
);
server_task.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn grpc_without_keepalive_keeps_unresponsive_connection_open() {
// Without keepalive the same unresponsive peer is NOT
// closed, proving the close above is attributable to keepalive.
let (addr, server_task, _engine_task) =
grpc_server_with_keepalive(b"engine-grpc-no-keepalive", None).await;
let closed = h2_unresponsive_peer_closed_within(&addr, Duration::from_secs(1)).await;
assert!(
!closed,
"without keepalive an idle h2 connection must stay open"
);
server_task.abort();
}
+205 -35
View File
@@ -10,20 +10,34 @@ mod routes;
mod runtime;
mod server_info;
mod state;
mod tls;
#[cfg(test)]
mod tls_tests;
mod utils;
use std::future::Future;
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use anyhow::{Context as _, Result};
use axum::Router;
use axum::serve::ListenerExt as _;
pub use config::{ApiServerOptions, Config, CoordinatorMode, CorsConfig, HttpListenerMode};
use axum::body::Body;
use axum::http::Request;
pub use config::{
ApiServerOptions, Config, CoordinatorMode, CorsConfig, DEFAULT_KEEP_ALIVE_TIMEOUT,
HttpListenerMode, TlsConfig,
};
use futures::FutureExt as _;
use hyper::body::Incoming;
use hyper::server::conn::http1;
use hyper_util::rt::{TokioIo, TokioTimer};
use hyper_util::server::graceful::GracefulShutdown;
use hyper_util::service::TowerToHyperService;
use tokio::net::TcpListener;
use tokio::time::{Instant, sleep_until};
use tokio_stream::wrappers::TcpListenerStream;
use tokio_util::either::Either;
use tokio_util::sync::CancellationToken;
use tonic::transport::Server as TonicServer;
use tower::ServiceExt as _;
use tracing::{info, trace, warn};
use vllm_chat::{ChatLlm, LoadModelBackendsOptions, load_model_backends};
pub use vllm_chat::{ChatTemplateContentFormatOption, ParserSelection, RendererSelection};
@@ -36,6 +50,13 @@ use crate::routes::build_router;
use crate::server_info::ServerInfoSnapshot;
use crate::state::AppState;
/// How often the server PINGs an idle gRPC connection to reap a dead peer;
/// tonic enables no keepalive by default. 2h matches the gRPC-core default.
const GRPC_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(7200);
/// How long the server waits for a keepalive PING reply before dropping the gRPC
/// connection. 20s matches the gRPC-core default.
const GRPC_KEEPALIVE_TIMEOUT: Duration = Duration::from_secs(20);
/// Resolve the public model names accepted by the frontend.
fn effective_served_model_names(model: &str, served_model_name: &[String]) -> Vec<String> {
if served_model_name.is_empty() {
@@ -45,6 +66,17 @@ fn effective_served_model_names(model: &str, served_model_name: &[String]) -> Ve
}
}
/// Choose the gRPC listener host. It follows the HTTP TCP host when there is
/// one; otherwise (unix socket or inherited fd) it defaults to IPv4 loopback
/// rather than all interfaces, so the side-car is never accidentally
/// network-exposed.
fn grpc_bind_host(listener_mode: &HttpListenerMode) -> &str {
match listener_mode {
HttpListenerMode::BindTcp { host, .. } => host.as_str(),
HttpListenerMode::BindUnix { .. } | HttpListenerMode::InheritedFd { .. } => "127.0.0.1",
}
}
/// Build the shared application state for one configured model and one engine
/// client.
async fn build_state(config: &Config) -> Result<Arc<AppState>> {
@@ -130,6 +162,15 @@ where
{
config.validate().context("invalid OpenAI frontend configuration")?;
// Build the TLS server config once, up front, so a bad cert/key fails fast
// before the (potentially long) engine handshake.
let tls_config = config
.tls
.as_ref()
.map(tls::build_server_config)
.transpose()
.context("invalid TLS configuration")?;
// Also check shutdown during the (potentially long) startup handshake.
let state = tokio::select! {
result = build_state(&config) => result?,
@@ -144,40 +185,39 @@ where
// Optionally bind the gRPC Generate server on a separate port. Bind
// synchronously here so bind errors (port in use, permission denied, ...)
// surface before we start serving, rather than being deferred until
// shutdown. The gRPC listener follows the same host as the HTTP listener so
// that enabling --grpc-port does not accidentally expose the service on all
// interfaces when HTTP is intentionally local-only.
// surface before serving rather than being deferred until shutdown.
let grpc_setup = if let Some(grpc_port) = config.grpc_port {
let grpc_host = match &config.listener_mode {
HttpListenerMode::BindTcp { host, .. } => host.as_str(),
HttpListenerMode::BindUnix { .. } | HttpListenerMode::InheritedFd { .. } => "0.0.0.0",
};
let grpc_host = grpc_bind_host(&config.listener_mode);
let grpc_listener = TcpListener::bind((grpc_host, grpc_port))
.await
.with_context(|| format!("failed to bind gRPC listener on {grpc_host}:{grpc_port}"))?;
let addr = grpc_listener.local_addr()?;
let grpc_listener = Listener::Tcp(grpc_listener);
// gRPC reuses the HTTP TLS config (same SslContext) plus ALPN h2.
let grpc_tls = config
.tls
.as_ref()
.map(tls::build_grpc_server_config)
.transpose()
.context("invalid gRPC TLS configuration")?;
let svc = grpc::GenerateServer::new(grpc::GenerateServiceImpl::new(state.clone()));
let svc = TonicServer::builder()
.http2_keepalive_interval(Some(GRPC_KEEPALIVE_INTERVAL))
.http2_keepalive_timeout(Some(GRPC_KEEPALIVE_TIMEOUT))
.layer(middleware::request_runtime_layer(state.clone()))
.add_service(svc);
info!(%addr, "starting gRPC server");
Some((grpc_listener, svc))
info!(%addr, tls = grpc_tls.is_some(), "starting gRPC server");
Some((grpc_listener, svc, grpc_tls))
} else {
None
};
info!(%bind_address, %model, "starting OpenAI server");
// Set TCP_NODELAY on accepted connections to reduce latency.
// By `tap_io` we will do this on every accepted connection.
let listener = listener.tap_io(|io| {
if let Either::Left(tcp_stream) = io
&& let Err(err) = tcp_stream.set_nodelay(true)
{
trace!(error = %err, "failed to enable TCP_NODELAY on accepted HTTP connection");
}
});
let scheme = if tls_config.is_some() {
"https"
} else {
"http"
};
info!(%bind_address, %scheme, %model, "starting OpenAI server");
// Run HTTP and gRPC concurrently under a child token of the caller's shutdown
// token. Caller cancellation propagates into both protocols; if either
@@ -208,17 +248,27 @@ where
}
});
// 0 disables keep-alive but still bounds the head read (default), so a
// silent client cannot hold the connection open.
let keep_alive_timeout = config.keep_alive_timeout;
let timeouts = ConnectionTimeouts {
handshake: tls::TLS_HANDSHAKE_TIMEOUT,
header_read: if keep_alive_timeout.is_zero() {
DEFAULT_KEEP_ALIVE_TIMEOUT
} else {
keep_alive_timeout
},
keep_alive_enabled: !keep_alive_timeout.is_zero(),
};
let http_fut = {
let shutdown = server_shutdown.child_token();
let server_shutdown = server_shutdown.clone();
let force_shutdown = force_shutdown.clone();
async move {
let server =
axum::serve(listener, app).with_graceful_shutdown(shutdown.cancelled_owned());
let result = tokio::select! {
result = server => {
result.context("HTTP server failed")
result = serve_listener(listener, tls_config, app, shutdown.cancelled_owned(), timeouts) => {
result
}
_ = force_shutdown.cancelled() => {
warn!("HTTP graceful shutdown deadline elapsed; aborting server");
@@ -236,16 +286,24 @@ where
let server_shutdown = server_shutdown.clone();
let force_shutdown = force_shutdown.clone();
async move {
let Some((grpc_listener, svc)) = grpc_setup else {
let Some((grpc_listener, svc, grpc_tls)) = grpc_setup else {
// No gRPC configured: just wait for shutdown so we do not race the
// join! by resolving early and tripping the cancellation token.
shutdown.cancelled().await;
return Ok(());
};
let server = svc.serve_with_incoming_shutdown(
TcpListenerStream::new(grpc_listener),
shutdown.cancelled_owned(),
);
// Box to unify the TLS and plaintext arms' different stream types.
let server = match grpc_tls {
Some(context) => {
let incoming =
grpc::tls_incoming(grpc_listener, context, tls::TLS_HANDSHAKE_TIMEOUT);
svc.serve_with_incoming_shutdown(incoming, shutdown.cancelled_owned()).boxed()
}
None => {
let incoming = grpc::incoming(grpc_listener);
svc.serve_with_incoming_shutdown(incoming, shutdown.cancelled_owned()).boxed()
}
};
let result = tokio::select! {
result = server => {
@@ -272,6 +330,99 @@ where
state.shutdown(shutdown_deadline).await
}
/// Per-connection timeouts applied while serving HTTP/HTTPS.
#[derive(Clone, Copy)]
pub(crate) struct ConnectionTimeouts {
/// Max time for a client to complete the TLS handshake (TLS path only).
pub(crate) handshake: Duration,
/// HTTP/1 header-read timeout (bounds idle keep-alive and the head read).
pub(crate) header_read: Duration,
/// Whether HTTP/1 keep-alive is enabled; `false` closes after each response.
pub(crate) keep_alive_enabled: bool,
}
/// Apply optional TLS termination and per-connection HTTP timeouts, then serve
/// `app`. Shared by [`serve_with_router_extension`] and the TLS tests.
async fn serve_listener(
listener: Listener,
tls: Option<openssl::ssl::SslContext>,
app: Router,
shutdown: impl Future<Output = ()> + Send + 'static,
timeouts: ConnectionTimeouts,
) -> Result<()> {
match tls {
Some(context) => {
// tls-listener terminates TLS (handshake + timeout); serve_connections
// owns the HTTP keep-alive/idle bound that axum::serve cannot express.
// Failed handshakes (incl. timeouts) log at ERROR via tls-listener.
let listener = tls_listener::builder(context)
.handshake_timeout(timeouts.handshake)
.listen(listener);
serve_connections(
listener,
app,
shutdown,
timeouts.header_read,
timeouts.keep_alive_enabled,
)
.await
.context("HTTPS server failed")
}
None => serve_connections(
listener,
app,
shutdown,
timeouts.header_read,
timeouts.keep_alive_enabled,
)
.await
.context("HTTP server failed"),
}
}
/// Serve `app` per connection (HTTP/1) with a keep-alive idle timeout and
/// graceful drain. Hand-rolled on hyper because [`axum::serve()`] takes no config.
async fn serve_connections<L>(
mut listener: L,
app: Router,
shutdown: impl Future<Output = ()> + Send,
header_read: Duration,
keep_alive_enabled: bool,
) -> Result<()>
where
L: axum::serve::Listener,
{
let graceful = GracefulShutdown::new();
let mut shutdown = std::pin::pin!(shutdown);
loop {
let (io, _addr) = tokio::select! {
conn = listener.accept() => conn,
() = &mut shutdown => break,
};
let service = TowerToHyperService::new(
app.clone().map_request(|req: Request<Incoming>| req.map(Body::new)),
);
let mut builder = http1::Builder::new();
builder.timer(TokioTimer::new()).header_read_timeout(header_read);
if !keep_alive_enabled {
builder.keep_alive(false);
}
let connection = builder.serve_connection(TokioIo::new(io), service);
let connection = graceful.watch(connection);
tokio::spawn(async move {
if let Err(err) = connection.await {
trace!(error = %err, "failed to serve connection");
}
});
}
drop(listener);
graceful.shutdown().await;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
@@ -293,4 +444,23 @@ mod tests {
served_names
);
}
#[test]
fn grpc_bind_host_follows_http_tcp_host() {
let mode = HttpListenerMode::BindTcp {
host: "0.0.0.0".to_string(),
port: 8000,
};
assert_eq!(grpc_bind_host(&mode), "0.0.0.0");
}
#[test]
fn grpc_bind_host_defaults_to_loopback_without_tcp_host() {
let unix = HttpListenerMode::BindUnix {
path: "/tmp/vllm.sock".to_string(),
};
let inherited = HttpListenerMode::InheritedFd { fd: 3 };
assert_eq!(grpc_bind_host(&unix), "127.0.0.1");
assert_eq!(grpc_bind_host(&inherited), "127.0.0.1");
}
}
+96 -14
View File
@@ -1,28 +1,49 @@
//! Unified HTTP listener wrapper for the Rust frontend.
//! Unified listener wrapper for the Rust frontend.
//!
//! This module hides the difference between TCP and Unix-domain listeners so
//! the rest of the server can bind or inherit one socket and pass it to
//! `axum::serve(...)` through a single type.
use std::io::Result;
use std::net::TcpListener as StdTcpListener;
use std::net::{SocketAddr, TcpListener as StdTcpListener};
use std::os::fd::{FromRawFd, IntoRawFd, OwnedFd};
use std::os::unix::net::UnixListener as StdUnixListener;
use std::pin::Pin;
use std::task::{Context, Poll, ready};
use auto_enums::enum_derive;
use socket2::Socket;
use tls_listener::{AsyncAccept, AsyncListener};
use tokio::net::{TcpListener, TcpStream, UnixListener, UnixStream};
use tokio_util::either::Either;
use tonic::transport::server::{Connected, TcpConnectInfo};
use tracing::trace;
use crate::HttpListenerMode;
/// Runtime listener type used by the OpenAI-compatible HTTP server, which is
/// either a TCP listener or a Unix-domain listener.
/// Runtime listener type used by the OpenAI-compatible HTTP or gRPC server,
/// which is either a TCP listener or a Unix-domain listener.
#[derive(Debug)]
pub enum Listener {
Tcp(TcpListener),
Unix(UnixListener),
}
/// Runtime listener I/O type which is either a TCP stream or a Unix-domain stream.
#[derive(Debug)]
#[enum_derive(tokio1::AsyncRead, tokio1::AsyncWrite)]
pub enum ListenerIo {
Tcp(TcpStream),
Unix(UnixStream),
}
/// Runtime listener address type which is either a TCP address or a Unix-domain address.
#[derive(Debug)]
#[allow(dead_code)]
pub enum ListenerAddr {
Tcp(SocketAddr),
Unix(tokio::net::unix::SocketAddr),
}
impl Listener {
/// Bind or adopt the listener described by the frontend configuration.
///
@@ -70,34 +91,95 @@ impl Listener {
Ok(Self::Tcp(TcpListener::from_std(std_listener)?))
}
}
fn listener_addr(&self) -> Result<ListenerAddr> {
match self {
Self::Tcp(listener) => listener.local_addr().map(ListenerAddr::Tcp),
Self::Unix(listener) => listener.local_addr().map(ListenerAddr::Unix),
}
}
}
impl Connected for ListenerIo {
type ConnectInfo = TcpConnectInfo;
fn connect_info(&self) -> TcpConnectInfo {
match self {
Self::Tcp(stream) => stream.connect_info(),
Self::Unix(_) => TcpConnectInfo {
local_addr: None,
remote_addr: None,
},
}
}
}
/// Attempt to set `TCP_NODELAY` on the accepted TCP stream.
fn enable_tcp_nodelay(stream: TcpStream) -> TcpStream {
if let Err(err) = stream.set_nodelay(true) {
trace!(error = %err, "failed to enable TCP_NODELAY on accepted TCP connection");
}
stream
}
/// Allow the unified listener to plug directly into `axum::serve(...)`.
impl axum::serve::Listener for Listener {
type Addr = Either<std::net::SocketAddr, tokio::net::unix::SocketAddr>;
type Io = Either<TcpStream, UnixStream>;
type Addr = ListenerAddr;
type Io = ListenerIo;
async fn accept(&mut self) -> (Self::Io, Self::Addr) {
match self {
Self::Tcp(listener) => {
let (io, addr) = listener.accept().await;
(Either::Left(io), Either::Left(addr))
let (io, addr) = axum::serve::Listener::accept(listener).await;
(
ListenerIo::Tcp(enable_tcp_nodelay(io)),
ListenerAddr::Tcp(addr),
)
}
Self::Unix(listener) => {
let (io, addr) = listener.accept().await;
(Either::Right(io), Either::Right(addr))
let (io, addr) = axum::serve::Listener::accept(listener).await;
(ListenerIo::Unix(io), ListenerAddr::Unix(addr))
}
}
}
fn local_addr(&self) -> Result<Self::Addr> {
match self {
Self::Tcp(listener) => listener.local_addr().map(Either::Left),
Self::Unix(listener) => listener.local_addr().map(Either::Right),
self.listener_addr()
}
}
/// Allow the unified listener to be adaptable to `tls_listener`.
impl AsyncAccept for Listener {
type Connection = ListenerIo;
type Address = ListenerAddr;
type Error = std::io::Error;
fn poll_accept(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(Self::Connection, Self::Address)>> {
match self.get_mut() {
Self::Tcp(listener) => {
let (io, addr) = ready!(listener.poll_accept(cx))?;
Poll::Ready(Ok((
ListenerIo::Tcp(enable_tcp_nodelay(io)),
ListenerAddr::Tcp(addr),
)))
}
Self::Unix(listener) => {
let (io, addr) = ready!(listener.poll_accept(cx))?;
Poll::Ready(Ok((ListenerIo::Unix(io), ListenerAddr::Unix(addr))))
}
}
}
}
impl AsyncListener for Listener {
fn local_addr(&self) -> Result<Self::Address> {
self.listener_addr()
}
}
#[cfg(test)]
mod tests {
use std::net::{Ipv4Addr, SocketAddrV4};
@@ -34,14 +34,20 @@ pub(super) fn validate_request_compat(
);
}
if let Some(prompt_logprobs) = request.sampling_params.prompt_logprobs
&& prompt_logprobs < 0
&& prompt_logprobs != -1
{
bail_invalid_request!(
param = "sampling_params",
"`prompt_logprobs` must be a non-negative value or -1."
);
if let Some(prompt_logprobs) = request.sampling_params.prompt_logprobs {
if prompt_logprobs < 0 && prompt_logprobs != -1 {
bail_invalid_request!(
param = "sampling_params",
"`prompt_logprobs` must be a non-negative value or -1."
);
}
if request.stream {
bail_invalid_request!(
param = "sampling_params",
"`prompt_logprobs` are not available when `stream=true`."
);
}
}
Ok(())
@@ -97,4 +103,54 @@ mod tests {
};
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
}
#[test]
fn validate_request_compat_rejects_streaming_prompt_logprobs() {
let request: GenerateRequest = serde_json::from_value(json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
"token_ids": [11, 22],
"stream": true,
"sampling_params": {
"prompt_logprobs": 0
}
}))
.expect("parse request");
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
let request: GenerateRequest = serde_json::from_value(json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
"token_ids": [11, 22],
"stream": true,
"sampling_params": {
"prompt_logprobs": 1
}
}))
.expect("parse request");
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
let request: GenerateRequest = serde_json::from_value(json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
"token_ids": [11, 22],
"stream": true,
"sampling_params": {
"prompt_logprobs": -1
}
}))
.expect("parse request");
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_err());
}
#[test]
fn validate_request_compat_accepts_non_stream_prompt_logprobs() {
let request: GenerateRequest = serde_json::from_value(json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
"token_ids": [11, 22],
"stream": false,
"sampling_params": {
"prompt_logprobs": 1
}
}))
.expect("parse request");
assert!(validate_request_compat(&request, &served(&["Qwen/Qwen1.5-0.5B-Chat"])).is_ok());
}
}
+39
View File
@@ -4129,6 +4129,45 @@ async fn raw_generate_rejects_empty_token_ids() {
assert_eq!(json["error"]["param"], "token_ids");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn raw_generate_rejects_streaming_prompt_logprobs() {
let mut app = test_app().await;
for prompt_logprobs in [0, 1] {
let response = app
.call(
Request::builder()
.method("POST")
.uri("/inference/v1/generate")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "Qwen/Qwen1.5-0.5B-Chat",
"token_ids": [11, 22],
"stream": true,
"sampling_params": {
"prompt_logprobs": prompt_logprobs
}
})
.to_string(),
))
.expect("build request"),
)
.await
.expect("call app");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = to_bytes(response.into_body(), usize::MAX).await.expect("read body");
let json: serde_json::Value = serde_json::from_slice(&body).expect("decode json");
assert_eq!(json["error"]["param"], "sampling_params");
assert_eq!(
json["error"]["message"],
"`prompt_logprobs` are not available when `stream=true`."
);
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn raw_generate_rejects_wrong_model() {
+119
View File
@@ -0,0 +1,119 @@
//! OpenSSL server-config construction for TLS termination.
//!
//! Builds an OpenSSL [`SslContext`] from the uvicorn-style `ssl_*` arguments
//! (certificate chain, private key, mTLS client verifier, optional cipher list).
//! The `tls-listener` crate drives the handshake on each accepted connection.
//!
//! Crypto runs through whichever OpenSSL the binary links (system by default,
//! vendored when built with that feature).
use std::path::Path;
use std::time::Duration;
use anyhow::{Context as _, Result};
use openssl::ssl::{
AlpnError, SslAcceptor, SslAcceptorBuilder, SslContext, SslContextBuilder, SslFiletype,
SslMethod, SslOptions, SslVerifyMode, select_next_proto,
};
use crate::config::TlsConfig;
/// Time a client has to complete the TLS handshake before the connection is dropped.
pub(crate) const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(60);
/// ALPN wire bytes for HTTP/2 (length-prefixed).
const ALPN_H2: &[u8] = b"\x02h2";
/// Build the shared OpenSSL acceptor from validated [`TlsConfig`]: the full
/// certificate chain, the private key (`key_file`, or the certificate file when
/// unset), the mTLS client verifier, and an optional cipher list.
///
/// Starts from the Mozilla intermediate baseline (forward-secret AEAD suites,
/// TLS 1.2 floor, server cipher preference, no compression), a slightly
/// stricter subset of the Python frontend's default suites; `--ssl-ciphers`
/// overrides it.
fn build_server_builder(tls: &TlsConfig) -> Result<SslAcceptorBuilder> {
let cert_file = tls.cert_file.as_deref().context("--ssl-certfile is required to enable TLS")?;
let mut builder = SslAcceptor::mozilla_intermediate_v5(SslMethod::tls_server())
.context("failed to initialize TLS")?;
builder.set_options(SslOptions::CIPHER_SERVER_PREFERENCE);
// Load the whole chain (leaf + intermediates), not just the leaf, so
// deployments behind an intermediate CA serve a complete chain.
ensure_exists(cert_file, "--ssl-certfile")?;
builder.set_certificate_chain_file(cert_file).with_context(|| {
format!("failed to parse certificate chain in --ssl-certfile {cert_file:?}")
})?;
// When `key_file` is unset the key is read from the certificate file
// (combined PEM).
let key_file = tls.key_file.as_deref().unwrap_or(cert_file);
ensure_exists(key_file, "private key file")?;
builder
.set_private_key_file(key_file, SslFiletype::PEM)
.with_context(|| format!("failed to parse private key in {key_file:?}"))?;
builder
.check_private_key()
.context("the certificate and private key do not match")?;
configure_client_auth(&mut builder, tls)?;
if let Some(ciphers) = tls.ciphers.as_deref().filter(|c| !c.is_empty()) {
builder
.set_cipher_list(ciphers)
.with_context(|| format!("invalid --ssl-ciphers {ciphers:?}"))?;
}
Ok(builder)
}
/// Build the HTTP [`SslContext`] (HTTP/1.1; no ALPN, matching uvicorn).
pub(crate) fn build_server_config(tls: &TlsConfig) -> Result<SslContext> {
Ok(build_server_builder(tls)?.build().into_context())
}
/// Build the gRPC [`SslContext`]: identical to [`build_server_config`] but
/// negotiates ALPN `h2`, which HTTP/2 over TLS requires.
pub(crate) fn build_grpc_server_config(tls: &TlsConfig) -> Result<SslContext> {
let mut builder = build_server_builder(tls)?;
builder.set_alpn_select_callback(|_ssl, client| {
select_next_proto(ALPN_H2, client).ok_or(AlpnError::NOACK)
});
Ok(builder.build().into_context())
}
/// Fail loudly with a flag-named message when a configured file is missing,
/// distinguishing it from a malformed-PEM error raised later by OpenSSL (whose
/// `ErrorStack` does not name the offending file).
fn ensure_exists(path: &str, what: &str) -> Result<()> {
std::fs::metadata(Path::new(path))
.map(drop)
.with_context(|| format!("failed to read {what} {path:?}"))
}
/// Apply the `cert_reqs` client-certificate policy: 0 = none, 1 = optional
/// (verify if presented, allow anonymous), 2 = required. `PEER` without a custom
/// verify callback still rejects a presented-but-untrusted certificate.
fn configure_client_auth(builder: &mut SslContextBuilder, tls: &TlsConfig) -> Result<()> {
if tls.cert_reqs == 0 {
builder.set_verify(SslVerifyMode::NONE);
return Ok(());
}
let ca_file = tls
.ca_certs
.as_deref()
.context("--ssl-ca-certs is required for client certificate verification")?;
ensure_exists(ca_file, "--ssl-ca-certs")?;
builder
.set_ca_file(ca_file)
.with_context(|| format!("failed to parse --ssl-ca-certs {ca_file:?}"))?;
let mut mode = SslVerifyMode::PEER;
if tls.cert_reqs == 2 {
mode |= SslVerifyMode::FAIL_IF_NO_PEER_CERT;
}
builder.set_verify(mode);
Ok(())
}
+688
View File
@@ -0,0 +1,688 @@
//! TLS tests: `build_server_config` unit checks plus end-to-end OpenSSL handshakes
//! through the production `serve_listener` path, with a trivial router since TLS
//! terminates below the app.
use std::pin::Pin;
use std::time::Duration;
use axum::Router;
use axum::routing::get;
use openssl::asn1::Asn1Time;
use openssl::bn::{BigNum, MsbOption};
use openssl::ec::{EcGroup, EcKey};
use openssl::hash::MessageDigest;
use openssl::nid::Nid;
use openssl::pkey::{PKey, Private};
use openssl::ssl::{SslConnector, SslFiletype, SslMethod, SslVersion};
use openssl::x509::extension::{BasicConstraints, KeyUsage, SubjectAlternativeName};
use openssl::x509::{X509, X509NameBuilder};
use tempfile::TempDir;
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
use tokio::net::TcpStream;
use tokio_openssl::SslStream;
use tokio_util::sync::CancellationToken;
use crate::config::{HttpListenerMode, TlsConfig};
use crate::listener::Listener;
use crate::{ConnectionTimeouts, serve_listener, tls};
// ============================================================================
// Test infrastructure
// ============================================================================
/// A throwaway CA + server/client/untrusted/chain cert set as PEM files in a
/// temp dir; dropping it deletes them.
pub(crate) struct TestCerts {
dir: TempDir,
}
impl TestCerts {
pub(crate) fn generate() -> Self {
let dir = tempfile::tempdir().expect("tempdir");
let (ca, ca_key) = build_ca();
let (server, server_key) = build_leaf("server", &["127.0.0.1", "localhost"], &ca, &ca_key);
let (client, client_key) = build_leaf("client", &[], &ca, &ca_key);
let (untrusted, untrusted_key) = build_self_signed("untrusted client");
// Leaf signed by an intermediate (itself signed by the root); the cert
// file holds leaf + intermediate, for the chain-serving test.
let (intermediate, intermediate_key) = build_intermediate(&ca, &ca_key);
let (chain_leaf, chain_leaf_key) = build_leaf(
"chain",
&["127.0.0.1", "localhost"],
&intermediate,
&intermediate_key,
);
let server_pem = pem(&server);
let server_key_pem = key_pem(&server_key);
let files = [
("ca.pem", pem(&ca)),
("server.pem", server_pem.clone()),
("server.key", server_key_pem.clone()),
("client.pem", pem(&client)),
("client.key", key_pem(&client_key)),
("untrusted_client.pem", pem(&untrusted)),
("untrusted_client.key", key_pem(&untrusted_key)),
(
"server_combined.pem",
format!("{server_pem}{server_key_pem}"),
),
(
"server_chain.pem",
format!("{}{}", pem(&chain_leaf), pem(&intermediate)),
),
("server_chain.key", key_pem(&chain_leaf_key)),
];
for (name, contents) in files {
std::fs::write(dir.path().join(name), contents).expect("write fixture");
}
Self { dir }
}
/// Absolute path to a fixture by name; the file need not exist.
pub(crate) fn path(&self, name: &str) -> String {
self.dir.path().join(name).to_str().expect("utf-8 path").to_string()
}
}
fn gen_key() -> PKey<Private> {
let group = EcGroup::from_curve_name(Nid::X9_62_PRIME256V1).expect("ec group");
let ec = EcKey::generate(&group).expect("ec key");
PKey::from_ec_key(ec).expect("pkey")
}
fn serial() -> openssl::asn1::Asn1Integer {
let mut bn = BigNum::new().expect("bignum");
bn.rand(159, MsbOption::MAYBE_ZERO, false).expect("rand serial");
bn.to_asn1_integer().expect("asn1 serial")
}
fn x509_name(cn: &str) -> openssl::x509::X509Name {
let mut builder = X509NameBuilder::new().expect("name builder");
builder.append_entry_by_text("CN", cn).expect("cn");
builder.build()
}
fn pem(cert: &X509) -> String {
String::from_utf8(cert.to_pem().expect("cert pem")).expect("utf-8 cert")
}
fn key_pem(key: &PKey<Private>) -> String {
String::from_utf8(key.private_key_to_pem_pkcs8().expect("key pem")).expect("utf-8 key")
}
/// A self-signed CA used to sign the server/client leaf certs.
fn build_ca() -> (X509, PKey<Private>) {
let key = gen_key();
let name = x509_name("vLLM Test CA");
let mut builder = X509::builder().expect("x509 builder");
builder.set_version(2).expect("version");
builder.set_serial_number(&serial()).expect("serial");
builder.set_subject_name(&name).expect("subject");
builder.set_issuer_name(&name).expect("issuer");
builder.set_pubkey(&key).expect("pubkey");
builder
.set_not_before(&Asn1Time::days_from_now(0).expect("nb"))
.expect("set nb");
builder
.set_not_after(&Asn1Time::days_from_now(3650).expect("na"))
.expect("set na");
builder
.append_extension(BasicConstraints::new().critical().ca().build().expect("bc"))
.expect("ext bc");
builder
.append_extension(
KeyUsage::new().critical().key_cert_sign().crl_sign().build().expect("ku"),
)
.expect("ext ku");
builder.sign(&key, MessageDigest::sha256()).expect("sign ca");
(builder.build(), key)
}
/// A CA-signed leaf cert with optional subject-alternative names (IP or DNS).
fn build_leaf(cn: &str, sans: &[&str], ca: &X509, ca_key: &PKey<Private>) -> (X509, PKey<Private>) {
let key = gen_key();
let mut builder = X509::builder().expect("x509 builder");
builder.set_version(2).expect("version");
builder.set_serial_number(&serial()).expect("serial");
builder.set_subject_name(&x509_name(cn)).expect("subject");
builder.set_issuer_name(ca.subject_name()).expect("issuer");
builder.set_pubkey(&key).expect("pubkey");
builder
.set_not_before(&Asn1Time::days_from_now(0).expect("nb"))
.expect("set nb");
builder
.set_not_after(&Asn1Time::days_from_now(3650).expect("na"))
.expect("set na");
builder
.append_extension(BasicConstraints::new().build().expect("bc"))
.expect("ext bc");
if !sans.is_empty() {
let mut san = SubjectAlternativeName::new();
for entry in sans {
if entry.parse::<std::net::IpAddr>().is_ok() {
san.ip(entry);
} else {
san.dns(entry);
}
}
let ext = san.build(&builder.x509v3_context(Some(ca), None)).expect("san");
builder.append_extension(ext).expect("ext san");
}
builder.sign(ca_key, MessageDigest::sha256()).expect("sign leaf");
(builder.build(), key)
}
/// A self-signed leaf not chained to the CA, for the untrusted-client test.
fn build_self_signed(cn: &str) -> (X509, PKey<Private>) {
let key = gen_key();
let name = x509_name(cn);
let mut builder = X509::builder().expect("x509 builder");
builder.set_version(2).expect("version");
builder.set_serial_number(&serial()).expect("serial");
builder.set_subject_name(&name).expect("subject");
builder.set_issuer_name(&name).expect("issuer");
builder.set_pubkey(&key).expect("pubkey");
builder
.set_not_before(&Asn1Time::days_from_now(0).expect("nb"))
.expect("set nb");
builder
.set_not_after(&Asn1Time::days_from_now(3650).expect("na"))
.expect("set na");
builder
.append_extension(BasicConstraints::new().build().expect("bc"))
.expect("ext bc");
builder.sign(&key, MessageDigest::sha256()).expect("sign self");
(builder.build(), key)
}
/// A CA-capable intermediate signed by the root, for the full-chain test.
fn build_intermediate(ca: &X509, ca_key: &PKey<Private>) -> (X509, PKey<Private>) {
let key = gen_key();
let mut builder = X509::builder().expect("x509 builder");
builder.set_version(2).expect("version");
builder.set_serial_number(&serial()).expect("serial");
builder
.set_subject_name(&x509_name("vLLM Test Intermediate CA"))
.expect("subject");
builder.set_issuer_name(ca.subject_name()).expect("issuer");
builder.set_pubkey(&key).expect("pubkey");
builder
.set_not_before(&Asn1Time::days_from_now(0).expect("nb"))
.expect("set nb");
builder
.set_not_after(&Asn1Time::days_from_now(3650).expect("na"))
.expect("set na");
builder
.append_extension(BasicConstraints::new().critical().ca().build().expect("bc"))
.expect("ext bc");
builder
.append_extension(
KeyUsage::new().critical().key_cert_sign().crl_sign().build().expect("ku"),
)
.expect("ext ku");
builder.sign(ca_key, MessageDigest::sha256()).expect("sign intermediate");
(builder.build(), key)
}
pub(crate) fn server_tls(certs: &TestCerts, cert_reqs: i32) -> TlsConfig {
TlsConfig {
cert_file: Some(certs.path("server.pem")),
key_file: Some(certs.path("server.key")),
ca_certs: (cert_reqs != 0).then(|| certs.path("ca.pem")),
cert_reqs,
ciphers: None,
}
}
/// A plaintext-listener TLS config for `build_server_config` checks (`cert_reqs`
/// 0, no client auth), with the cert/key files chosen by the caller.
fn build_tls(certs: &TestCerts, cert: &str, key: Option<&str>) -> TlsConfig {
TlsConfig {
cert_file: Some(certs.path(cert)),
key_file: key.map(|k| certs.path(k)),
ca_certs: None,
cert_reqs: 0,
ciphers: None,
}
}
/// Generous per-connection timeouts that never fire during the fast tests.
const TEST_TIMEOUTS: ConnectionTimeouts = ConnectionTimeouts {
handshake: Duration::from_secs(60),
header_read: Duration::from_secs(5),
keep_alive_enabled: true,
};
async fn spawn_server(tls_config: Option<TlsConfig>) -> (String, CancellationToken) {
spawn_server_with_timeouts(tls_config, TEST_TIMEOUTS).await
}
/// Bind an ephemeral listener and serve a trivial router via the production
/// `serve_listener`, optionally with TLS. The listener is bound (and thus
/// accepting into the backlog) before returning, so a client may connect
/// immediately without a sleep.
async fn spawn_server_with_timeouts(
tls_config: Option<TlsConfig>,
timeouts: ConnectionTimeouts,
) -> (String, CancellationToken) {
let listener = Listener::bind(&HttpListenerMode::BindTcp {
host: "127.0.0.1".to_string(),
port: 0,
})
.await
.expect("bind listener");
let addr = listener.local_addr().expect("local addr");
let server_config =
tls_config.map(|cfg| tls::build_server_config(&cfg).expect("build server config"));
let app = Router::new().route("/health", get(|| async { "ok" }));
let shutdown = CancellationToken::new();
let server_shutdown = shutdown.clone();
tokio::spawn(async move {
let _ = serve_listener(
listener,
server_config,
app,
server_shutdown.cancelled_owned(),
timeouts,
)
.await;
});
(addr, shutdown)
}
/// Open a TLS connection trusting the test CA and finish the handshake,
/// optionally presenting a client identity (`<name>.pem` + `<name>.key`) for
/// mTLS. Hostname verification is disabled (the IP-SAN match is not under test);
/// chain verification stays on, so an untrusted server cert is still rejected.
async fn connect_tls(
certs: &TestCerts,
addr: &str,
identity: Option<&str>,
) -> std::io::Result<Pin<Box<SslStream<TcpStream>>>> {
let tcp = TcpStream::connect(addr).await?;
let mut builder = SslConnector::builder(SslMethod::tls_client()).expect("connector builder");
builder.set_ca_file(certs.path("ca.pem")).expect("trust ca");
if let Some(name) = identity {
builder
.set_certificate_chain_file(certs.path(&format!("{name}.pem")))
.expect("client cert");
builder
.set_private_key_file(certs.path(&format!("{name}.key")), SslFiletype::PEM)
.expect("client key");
}
let connector = builder.build();
let mut config = connector.configure().expect("configure");
config.set_verify_hostname(false);
let ssl = config.into_ssl("127.0.0.1").expect("ssl");
let mut stream = Box::pin(SslStream::new(ssl, tcp).expect("client ssl stream"));
stream.as_mut().connect().await.map_err(std::io::Error::other)?;
Ok(stream)
}
/// Issue an HTTPS GET (with `Connection: close`), optionally with an mTLS identity.
async fn https_get(
certs: &TestCerts,
addr: &str,
identity: Option<&str>,
) -> std::io::Result<String> {
let mut stream = connect_tls(certs, addr, identity).await?;
stream
.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n")
.await?;
let mut response = String::new();
stream.read_to_string(&mut response).await?;
Ok(response)
}
/// Attempt a handshake offering only a legacy CBC+SHA1 suite over TLS 1.2,
/// capping the version so TLS 1.3 cannot rescue the negotiation.
async fn legacy_suite_handshake(certs: &TestCerts, addr: &str) -> std::io::Result<()> {
let tcp = TcpStream::connect(addr).await?;
let mut builder = SslConnector::builder(SslMethod::tls_client()).expect("connector builder");
builder.set_ca_file(certs.path("ca.pem")).expect("trust ca");
builder.set_max_proto_version(Some(SslVersion::TLS1_2)).expect("cap tls1.2");
builder
.set_cipher_list("ECDHE-ECDSA-AES256-SHA:@SECLEVEL=0")
.expect("legacy cipher");
let connector = builder.build();
let mut config = connector.configure().expect("configure");
config.set_verify_hostname(false);
let ssl = config.into_ssl("127.0.0.1").expect("ssl");
let stream = SslStream::new(ssl, tcp).expect("client ssl stream");
tokio::pin!(stream);
stream.as_mut().connect().await.map_err(std::io::Error::other)
}
async fn plain_get(addr: &str) -> std::io::Result<String> {
let mut tcp = TcpStream::connect(addr).await?;
tcp.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n")
.await?;
let mut response = String::new();
tcp.read_to_string(&mut response).await?;
Ok(response)
}
// ============================================================================
// Tests
// ============================================================================
#[test]
fn builds_from_combined_pem() {
// Key omitted: it is read from the combined cert+key file.
let certs = TestCerts::generate();
assert!(tls::build_server_config(&build_tls(&certs, "server_combined.pem", None)).is_ok());
}
#[test]
fn rejects_missing_cert_file() {
let certs = TestCerts::generate();
assert!(tls::build_server_config(&build_tls(&certs, "does_not_exist.pem", None)).is_err());
}
#[test]
fn accepts_valid_cipher_list() {
let certs = TestCerts::generate();
let mut cfg = build_tls(&certs, "server.pem", Some("server.key"));
cfg.ciphers = Some("ECDHE-ECDSA-AES256-GCM-SHA384".to_string());
assert!(tls::build_server_config(&cfg).is_ok());
}
#[test]
fn rejects_invalid_cipher_list() {
let certs = TestCerts::generate();
let mut cfg = build_tls(&certs, "server.pem", Some("server.key"));
cfg.ciphers = Some("THIS-IS-NOT-A-CIPHER".to_string());
assert!(tls::build_server_config(&cfg).is_err());
}
#[test]
fn rejects_mismatched_cert_and_key() {
// check_private_key must reject a key that does not match the certificate.
let certs = TestCerts::generate();
let tls = build_tls(&certs, "client.pem", Some("server.key"));
assert!(tls::build_server_config(&tls).is_err());
}
#[tokio::test]
async fn https_request_succeeds_over_tls() {
let certs = TestCerts::generate();
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 0))).await;
let response = https_get(&certs, &addr, None).await.expect("https request");
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
shutdown.cancel();
}
#[tokio::test]
async fn serves_full_certificate_chain() {
// Cert file holds leaf + intermediate; a client trusting only the root can
// verify only if the server sends the intermediate, guarding against a
// leaf-only load.
let certs = TestCerts::generate();
let tls = TlsConfig {
cert_file: Some(certs.path("server_chain.pem")),
key_file: Some(certs.path("server_chain.key")),
ca_certs: None,
cert_reqs: 0,
ciphers: None,
};
let (addr, shutdown) = spawn_server(Some(tls)).await;
let response = https_get(&certs, &addr, None).await.expect("chained https request");
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
shutdown.cancel();
}
#[tokio::test]
async fn rejects_legacy_cipher_only_client() {
let certs = TestCerts::generate();
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 0))).await;
let result = legacy_suite_handshake(&certs, &addr).await;
assert!(result.is_err(), "legacy-only client must be rejected");
shutdown.cancel();
}
#[tokio::test]
async fn ssl_ciphers_override_widens_past_preset() {
// Counterpart to rejects_legacy_cipher_only_client: --ssl-ciphers set to that
// same legacy suite lets the client through, proving the override beats the preset.
let certs = TestCerts::generate();
let mut tls = server_tls(&certs, 0);
tls.ciphers = Some("ECDHE-ECDSA-AES256-SHA:@SECLEVEL=0".to_string());
let (addr, shutdown) = spawn_server(Some(tls)).await;
let result = legacy_suite_handshake(&certs, &addr).await;
assert!(
result.is_ok(),
"override must allow the legacy suite: {result:?}"
);
shutdown.cancel();
}
#[tokio::test]
async fn mtls_required_rejects_client_without_certificate() {
let certs = TestCerts::generate();
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 2))).await;
let result = https_get(&certs, &addr, None).await;
assert!(
result.is_err(),
"handshake must fail without a client certificate"
);
shutdown.cancel();
}
#[tokio::test]
async fn mtls_required_accepts_valid_client_certificate() {
let certs = TestCerts::generate();
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 2))).await;
let response = https_get(&certs, &addr, Some("client")).await.expect("mtls request");
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
shutdown.cancel();
}
#[tokio::test]
async fn mtls_optional_allows_anonymous_and_authenticated() {
let certs = TestCerts::generate();
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 1))).await;
let anonymous = https_get(&certs, &addr, None).await.expect("anonymous request");
assert!(anonymous.starts_with("HTTP/1.1 200"), "{anonymous}");
let authenticated =
https_get(&certs, &addr, Some("client")).await.expect("authenticated request");
assert!(authenticated.starts_with("HTTP/1.1 200"), "{authenticated}");
shutdown.cancel();
}
#[tokio::test]
async fn mtls_rejects_untrusted_client_certificate() {
// Optional (1) still verifies a presented cert, so a self-signed cert not
// chained to the CA is rejected in both modes, not just required (2).
let certs = TestCerts::generate();
for cert_reqs in [1, 2] {
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, cert_reqs))).await;
let result = https_get(&certs, &addr, Some("untrusted_client")).await;
assert!(
result.is_err(),
"cert_reqs={cert_reqs}: untrusted client cert must be rejected"
);
shutdown.cancel();
}
}
#[tokio::test]
async fn plain_http_serves_when_tls_is_disabled() {
let (addr, shutdown) = spawn_server(None).await;
let response = plain_get(&addr).await.expect("http request");
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
shutdown.cancel();
}
#[tokio::test]
async fn tls_handshake_timeout_drops_silent_client() {
// Silent client (no ClientHello) must be dropped at the handshake deadline.
let certs = TestCerts::generate();
let timeouts = ConnectionTimeouts {
handshake: Duration::from_millis(150),
header_read: Duration::from_secs(5),
keep_alive_enabled: true,
};
let (addr, shutdown) = spawn_server_with_timeouts(Some(server_tls(&certs, 0)), timeouts).await;
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
let mut buf = [0u8; 1];
let read = tokio::time::timeout(Duration::from_secs(5), tcp.read(&mut buf)).await;
assert!(
matches!(read, Ok(Ok(0)) | Ok(Err(_))),
"server must drop a stalled TLS handshake (expected close, got {read:?})"
);
shutdown.cancel();
}
#[tokio::test]
async fn keep_alive_timeout_closes_idle_connection() {
// Idle keep-alive connection must be closed at the deadline.
let timeouts = ConnectionTimeouts {
handshake: Duration::from_secs(60),
header_read: Duration::from_millis(150),
keep_alive_enabled: true,
};
let (addr, shutdown) = spawn_server_with_timeouts(None, timeouts).await;
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
// No `Connection: close`, so it stays alive until the idle deadline.
tcp.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n")
.await
.expect("write request");
let drained = tokio::time::timeout(Duration::from_secs(5), async {
let mut buf = [0u8; 1024];
loop {
match tcp.read(&mut buf).await {
Ok(0) => return Ok(()),
Ok(_) => continue,
Err(err) => return Err(err),
}
}
})
.await;
assert!(
matches!(drained, Ok(Ok(()))),
"server must close an idle keep-alive connection (got {drained:?})"
);
shutdown.cancel();
}
#[tokio::test]
async fn keep_alive_timeout_closes_idle_tls_connection() {
// The keep-alive idle bound lives in serve_connections, below TLS; assert it
// still fires through tls-listener's post-handshake SslStream, not just plaintext.
let certs = TestCerts::generate();
let timeouts = ConnectionTimeouts {
handshake: Duration::from_secs(60),
header_read: Duration::from_millis(150),
keep_alive_enabled: true,
};
let (addr, shutdown) = spawn_server_with_timeouts(Some(server_tls(&certs, 0)), timeouts).await;
let mut stream = connect_tls(&certs, &addr, None).await.expect("handshake");
// No `Connection: close`, so the connection stays alive until the idle deadline.
stream
.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n")
.await
.expect("write request");
let closed = tokio::time::timeout(Duration::from_secs(5), async {
let mut buf = [0u8; 1024];
loop {
// A clean close_notify (Ok(0)) or an abrupt TLS EOF both mean the
// server closed; only the outer timeout (still open) is a failure.
match stream.read(&mut buf).await {
Ok(0) | Err(_) => break,
Ok(_) => continue,
}
}
})
.await;
assert!(
closed.is_ok(),
"server must close an idle keep-alive TLS connection at the deadline"
);
shutdown.cancel();
}
#[tokio::test]
async fn idle_timeout_closes_silent_client() {
// Silent client closed by the header-read timeout (http1-only arms it from byte 0).
let timeouts = ConnectionTimeouts {
handshake: Duration::from_secs(60),
header_read: Duration::from_millis(150),
keep_alive_enabled: true,
};
let (addr, shutdown) = spawn_server_with_timeouts(None, timeouts).await;
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
let mut buf = [0u8; 1];
let read = tokio::time::timeout(Duration::from_secs(5), tcp.read(&mut buf)).await;
assert!(
matches!(read, Ok(Ok(0)) | Ok(Err(_))),
"server must close a silent client (expected close, got {read:?})"
);
shutdown.cancel();
}
#[tokio::test]
async fn keep_alive_zero_disables_keep_alive() {
// 0 disables keep-alive (serve, then close), like uvicorn's timeout_keep_alive=0.
let timeouts = ConnectionTimeouts {
handshake: Duration::from_secs(60),
header_read: Duration::from_secs(5),
keep_alive_enabled: false,
};
let (addr, shutdown) = spawn_server_with_timeouts(None, timeouts).await;
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
tcp.write_all(b"GET /health HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n")
.await
.expect("write request");
let mut response = String::new();
let read =
tokio::time::timeout(Duration::from_secs(5), tcp.read_to_string(&mut response)).await;
assert!(
read.is_ok(),
"server must close after one response, not hang"
);
assert!(response.starts_with("HTTP/1.1 200"), "{response}");
// Assert `Connection: close`, not just 200: a 0 header-read timeout would also
// serve an immediate request, so 200 alone wouldn't prove keep-alive is off.
assert!(
response.to_ascii_lowercase().contains("connection: close"),
"keep-alive must be disabled (expected Connection: close): {response}"
);
shutdown.cancel();
}
#[tokio::test]
async fn disabled_keep_alive_still_closes_silent_client() {
// Even with keep-alive off, the head read stays bounded, so a silent client
// is dropped rather than held open.
let timeouts = ConnectionTimeouts {
handshake: Duration::from_secs(60),
header_read: Duration::from_millis(150),
keep_alive_enabled: false,
};
let (addr, shutdown) = spawn_server_with_timeouts(None, timeouts).await;
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
let mut buf = [0u8; 1];
let read = tokio::time::timeout(Duration::from_secs(5), tcp.read(&mut buf)).await;
assert!(
matches!(read, Ok(Ok(0)) | Ok(Err(_))),
"disabled keep-alive must still close a silent client (got {read:?})"
);
shutdown.cancel();
}
+13 -13
View File
@@ -23,7 +23,7 @@ def test_python_error():
error happening from the C++ side.
"""
allocator = get_mem_allocator_instance()
total_bytes = current_platform.mem_get_info()[1]
total_bytes = torch.accelerator.get_memory_info()[1]
alloc_bytes = int(total_bytes * 0.7)
tensors = []
with allocator.use_memory_pool():
@@ -64,9 +64,9 @@ def test_basic_cumem():
output = x + y + z
assert torch.allclose(output, torch.ones_like(output) * 3)
free_bytes = current_platform.mem_get_info()[0]
free_bytes = torch.accelerator.get_memory_info()[0]
allocator.sleep()
free_bytes_after_sleep = current_platform.mem_get_info()[0]
free_bytes_after_sleep = torch.accelerator.get_memory_info()[0]
assert free_bytes_after_sleep > free_bytes
allocator.wake_up()
@@ -99,9 +99,9 @@ def test_cumem_with_cudagraph():
with torch.cuda.graph(model_graph):
y = model(x)
free_bytes = current_platform.mem_get_info()[0]
free_bytes = torch.accelerator.get_memory_info()[0]
allocator.sleep()
free_bytes_after_sleep = current_platform.mem_get_info()[0]
free_bytes_after_sleep = torch.accelerator.get_memory_info()[0]
assert free_bytes_after_sleep > free_bytes
allocator.wake_up()
@@ -132,7 +132,7 @@ def test_cumem_with_cudagraph():
],
)
def test_end_to_end(model: str):
free, total = current_platform.mem_get_info()
free, total = torch.accelerator.get_memory_info()
used_bytes_baseline = total - free # in case other process is running
llm = LLM(model, enable_sleep_mode=True)
prompt = "How are you?"
@@ -144,7 +144,7 @@ def test_end_to_end(model: str):
# test sleep level 1 here.
llm.sleep(level=1)
free_gpu_bytes_after_sleep, total = current_platform.mem_get_info()
free_gpu_bytes_after_sleep, total = torch.accelerator.get_memory_info()
used_bytes = total - free_gpu_bytes_after_sleep - used_bytes_baseline
# now the memory usage is mostly cudagraph memory pool,
# and it should be less than the model weights (1B model, 2GiB weights)
@@ -164,7 +164,7 @@ def test_end_to_end(model: str):
llm.sleep(level=1)
llm.wake_up(tags=["weights"])
free_gpu_bytes_wake_up_w, total = current_platform.mem_get_info()
free_gpu_bytes_wake_up_w, total = torch.accelerator.get_memory_info()
used_bytes = total - free_gpu_bytes_wake_up_w - used_bytes_baseline
# should just reallocate memory for weights (1B model, ~2GiB weights)
@@ -181,7 +181,7 @@ def test_end_to_end(model: str):
@create_new_process_for_each_test()
def test_deep_sleep():
model = "hmellor/tiny-random-LlamaForCausalLM"
free, total = current_platform.mem_get_info()
free, total = torch.accelerator.get_memory_info()
used_bytes_baseline = total - free # in case other process is running
llm = LLM(model, enable_sleep_mode=True)
prompt = "How are you?"
@@ -191,13 +191,13 @@ def test_deep_sleep():
# Put the engine to deep sleep
llm.sleep(level=2)
free_gpu_bytes_after_sleep, total = current_platform.mem_get_info()
free_gpu_bytes_after_sleep, total = torch.accelerator.get_memory_info()
used_bytes = total - free_gpu_bytes_after_sleep - used_bytes_baseline
assert used_bytes < 3 * GiB_bytes
llm.wake_up(tags=["weights"])
llm.collective_rpc("reload_weights")
free_gpu_bytes_wake_up_w, total = current_platform.mem_get_info()
free_gpu_bytes_wake_up_w, total = torch.accelerator.get_memory_info()
used_bytes = total - free_gpu_bytes_wake_up_w - used_bytes_baseline
assert used_bytes < 4 * GiB_bytes
@@ -213,7 +213,7 @@ def test_deep_sleep():
def test_deep_sleep_async():
async def test():
model = "hmellor/tiny-random-LlamaForCausalLM"
free, total = current_platform.mem_get_info()
free, total = torch.accelerator.get_memory_info()
used_bytes_baseline = total - free # in case other process is running
engine_args = AsyncEngineArgs(
model=model,
@@ -232,7 +232,7 @@ def test_deep_sleep_async():
await llm.wake_up(tags=["weights"])
await llm.collective_rpc("reload_weights")
free_gpu_bytes_wake_up_w, total = current_platform.mem_get_info()
free_gpu_bytes_wake_up_w, total = torch.accelerator.get_memory_info()
used_bytes = total - free_gpu_bytes_wake_up_w - used_bytes_baseline
assert used_bytes < 4 * GiB_bytes
+59 -14
View File
@@ -29,9 +29,43 @@ from vllm.distributed.weight_transfer.nccl_engine import (
NCCLWeightTransferInitInfo,
NCCLWeightTransferUpdateInfo,
)
from vllm.platforms import current_platform
from vllm.utils.network_utils import get_open_port
def _weight_transfer_ray_env_vars() -> dict[str, str]:
if not current_platform.is_rocm():
return {}
return {
"RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES": "1",
"RAY_EXPERIMENTAL_NOSET_HIP_VISIBLE_DEVICES": "1",
"RAY_EXPERIMENTAL_NOSET_ROCR_VISIBLE_DEVICES": "1",
}
def _init_ray_for_weight_transfer() -> None:
if ray.is_initialized():
return
ray.init(
ignore_reinit_error=True,
runtime_env={"env_vars": _weight_transfer_ray_env_vars()},
)
def _get_ray_assigned_device() -> torch.device:
gpu_ids = ray.get_gpu_ids()
if not gpu_ids:
return torch.device("cuda:0")
return torch.device(f"cuda:{int(gpu_ids[0])}")
def _set_ray_assigned_device() -> torch.device:
device = _get_ray_assigned_device()
torch.accelerator.set_device(device)
return device
def create_mock_parallel_config(
rank: int = 0,
world_size: int = 1,
@@ -321,6 +355,8 @@ def trainer_broadcast_tensor(
"""Trainer task that broadcasts a tensor via NCCL."""
import torch
device = _set_ray_assigned_device()
from vllm.distributed.device_communicators.pynccl import PyNcclCommunicator
from vllm.distributed.utils import StatelessProcessGroup
@@ -331,12 +367,11 @@ def trainer_broadcast_tensor(
rank=0,
world_size=world_size,
)
# Ray sets CUDA_VISIBLE_DEVICES, so device 0 is the assigned GPU
comm = PyNcclCommunicator(pg, device=0)
comm = PyNcclCommunicator(pg, device=device.index)
# Create and broadcast the tensor
dtype = getattr(torch, tensor_dtype)
tensor_to_send = torch.ones(tensor_shape, dtype=dtype, device="cuda:0")
tensor_to_send = torch.ones(tensor_shape, dtype=dtype, device=device)
comm.broadcast(tensor_to_send, src=0, stream=torch.cuda.current_stream())
torch.accelerator.synchronize()
@@ -356,6 +391,8 @@ def inference_receive_tensor(
import torch
_set_ray_assigned_device()
from vllm.config.parallel import ParallelConfig
from vllm.config.weight_transfer import WeightTransferConfig
from vllm.distributed.weight_transfer.nccl_engine import (
@@ -435,7 +472,7 @@ def test_nccl_weight_transfer_between_processes():
This test verifies that the NCCLWeightTransferEngine can receive
tensors broadcast by a trainer process via NCCL.
"""
ray.init(ignore_reinit_error=True)
_init_ray_for_weight_transfer()
master_address = "127.0.0.1"
master_port = get_open_port()
@@ -473,6 +510,8 @@ def trainer_broadcast_sparse_tensor(
"""Trainer task that broadcasts sparse patches via NCCL."""
import torch
device = _set_ray_assigned_device()
from vllm.distributed.device_communicators.pynccl import PyNcclCommunicator
from vllm.distributed.utils import StatelessProcessGroup
from vllm.distributed.weight_transfer.base import SparseWeightPatch
@@ -487,12 +526,12 @@ def trainer_broadcast_sparse_tensor(
rank=0,
world_size=world_size,
)
comm = PyNcclCommunicator(pg, device=0)
comm = PyNcclCommunicator(pg, device=device.index)
patch = SparseWeightPatch(
name="test.weight",
indices=torch.tensor([1, 7, 25], dtype=torch.int32, device="cuda:0"),
values=torch.tensor([10.0, 20.0, 30.0], dtype=torch.float32, device="cuda:0"),
indices=torch.tensor([1, 7, 25], dtype=torch.int32, device=device),
values=torch.tensor([10.0, 20.0, 30.0], dtype=torch.float32, device=device),
)
NCCLWeightTransferEngine.trainer_send_sparse_weights(
iter([patch]),
@@ -513,6 +552,8 @@ def inference_receive_sparse_tensor(
import torch
device = _set_ray_assigned_device()
from vllm.config.parallel import ParallelConfig
from vllm.config.weight_transfer import WeightTransferConfig
from vllm.distributed.weight_transfer.nccl_engine import (
@@ -540,7 +581,7 @@ def inference_receive_sparse_tensor(
)
)
target = torch.zeros(30, dtype=torch.float32, device="cuda")
target = torch.zeros(30, dtype=torch.float32, device=device)
def apply_sparse_patches(patches: list[SparseWeightPatch]):
for patch in patches:
@@ -556,9 +597,9 @@ def inference_receive_sparse_tensor(
engine.receive_sparse_weights(update_info, apply_sparse_patches)
torch.accelerator.synchronize()
expected = torch.zeros(30, dtype=torch.float32, device="cuda")
expected = torch.zeros(30, dtype=torch.float32, device=device)
expected[[1, 7, 25]] = torch.tensor(
[10.0, 20.0, 30.0], dtype=torch.float32, device="cuda"
[10.0, 20.0, 30.0], dtype=torch.float32, device=device
)
success = torch.equal(target, expected)
engine.shutdown()
@@ -574,7 +615,7 @@ def inference_receive_sparse_tensor(
)
def test_nccl_sparse_weight_transfer_between_processes():
"""Test NCCL sparse weight transfer from trainer to inference process."""
ray.init(ignore_reinit_error=True)
_init_ray_for_weight_transfer()
master_address = "127.0.0.1"
master_port = get_open_port()
@@ -933,16 +974,18 @@ class TrainerActor:
"""Trainer actor that creates and holds CUDA IPC handles."""
def __init__(self, tensor_shape: list[int], tensor_dtype: str):
device = _set_ray_assigned_device()
# Create tensor on GPU and keep it alive
dtype = getattr(torch, tensor_dtype)
self.tensor = torch.ones(tensor_shape, dtype=dtype, device="cuda:0")
self.tensor = torch.ones(tensor_shape, dtype=dtype, device=device)
self.tensor.fill_(42.0) # Fill with 42 to verify correct transfer
# Create IPC handle (tensor must stay alive for IPC to work)
# reduce_tensor returns (rebuild_func, args); we only send args
# since the receiver imports rebuild_cuda_tensor directly.
_, ipc_args = reduce_tensor(self.tensor)
gpu_uuid = get_physical_gpu_id(0)
gpu_uuid = get_physical_gpu_id(device.index)
torch.accelerator.synchronize()
@@ -974,6 +1017,8 @@ def inference_receive_ipc_tensor(
import torch
_set_ray_assigned_device()
from vllm.config.parallel import ParallelConfig
from vllm.config.weight_transfer import WeightTransferConfig
from vllm.distributed.weight_transfer.ipc_engine import (
@@ -1072,7 +1117,7 @@ def test_ipc_weight_transfer_between_processes(mode: str):
from ray.util.placement_group import placement_group
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
ray.init(ignore_reinit_error=True)
_init_ray_for_weight_transfer()
# Create a placement group to ensure both processes are on the same GPU
# Use fractional GPUs so both tasks can share the same GPU bundle
@@ -76,3 +76,60 @@ async def test_chat_logit_bias_invalid(client):
assert error.status_code == 400
assert str(invalid_token_id) in error_message
assert str(vocab_size) in error_message
@pytest.mark.asyncio
async def test_chat_logit_bias_non_integer_key(client):
"""Test that a non-integer logit_bias key is rejected with a clean,
informative error instead of a raw 'invalid literal for int()' message."""
with pytest.raises(openai.BadRequestError) as excinfo:
await client.chat.completions.create(
model=MODEL_NAME,
messages=[{"role": "user", "content": "Testing invalid logit bias key"}],
max_tokens=5,
logit_bias={"not_a_token_id": 50},
)
error = excinfo.value
error_message = str(error)
assert error.status_code == 400
assert "not_a_token_id" in error_message
assert "logit_bias" in error_message
@pytest.mark.asyncio
async def test_chat_logit_bias_non_numeric_value(client):
"""Test that a non-numeric logit_bias value is rejected with a message
that names the specific offending token, not just a generic TypeError."""
with pytest.raises(openai.BadRequestError) as excinfo:
await client.chat.completions.create(
model=MODEL_NAME,
messages=[{"role": "user", "content": "Testing invalid logit bias value"}],
max_tokens=5,
logit_bias={"1": "not_a_number"},
)
error = excinfo.value
error_message = str(error)
assert error.status_code == 400
assert "logit_bias" in error_message
@pytest.mark.asyncio
async def test_chat_logit_bias_multiple_non_integer_keys(client):
"""Test that ALL invalid logit_bias keys are reported together,
not just the first one encountered."""
with pytest.raises(openai.BadRequestError) as excinfo:
await client.chat.completions.create(
model=MODEL_NAME,
messages=[{"role": "user", "content": "Testing multiple bad keys"}],
max_tokens=5,
logit_bias={"bad1": 50.0, "bad2": 20.0},
)
error_message = str(excinfo.value)
assert excinfo.value.status_code == 400
assert "bad1" in error_message
assert "bad2" in error_message
@@ -142,6 +142,7 @@ class TestHarmonyToResponseOutput:
)
assert output_items[0].call_id.startswith("call_")
assert output_items[0].id.startswith("fc_")
assert output_items[0].status == "completed"
def test_commentary_with_python_recipient_creates_reasoning(self):
"""Test that commentary with recipient='python' creates reasoning items."""
@@ -0,0 +1,160 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Unit tests for stop_token_ids propagation from default_sampling_params
to SamplingParams in ChatCompletionRequest and CompletionRequest.
Regression test for https://github.com/vllm-project/vllm/issues/22519
where gpt-oss model stop tokens (e.g., </call> = 200012) were loaded into
default_sampling_params at server startup but silently discarded on every
request because to_sampling_params() never fell back to defaults.
"""
import pytest
from vllm.entrypoints.openai.chat_completion.protocol import (
ChatCompletionRequest,
)
from vllm.entrypoints.openai.completion.protocol import (
CompletionRequest,
)
class TestChatCompletionStopTokenIds:
"""Test stop_token_ids merging in ChatCompletionRequest.to_sampling_params()."""
@pytest.fixture
def minimal_chat_request(self):
return ChatCompletionRequest(
model="test-model",
messages=[{"role": "user", "content": "hello"}],
)
def test_default_stop_token_ids_applied(self, minimal_chat_request):
"""Server-default stop_token_ids are applied when client sends none."""
default_sampling_params = {
"stop_token_ids": [200012, 200002],
}
sampling_params = minimal_chat_request.to_sampling_params(
max_tokens=100,
default_sampling_params=default_sampling_params,
)
assert set(sampling_params.stop_token_ids) == {200012, 200002}
def test_client_stop_token_ids_merged_with_defaults(self):
"""Client-specified stop_token_ids are merged with server defaults."""
request = ChatCompletionRequest(
model="test-model",
messages=[{"role": "user", "content": "hello"}],
stop_token_ids=[99999],
)
default_sampling_params = {
"stop_token_ids": [200012, 200002],
}
sampling_params = request.to_sampling_params(
max_tokens=100,
default_sampling_params=default_sampling_params,
)
assert set(sampling_params.stop_token_ids) == {200012, 200002, 99999}
assert sampling_params.stop_token_ids == [99999, 200012, 200002]
def test_no_stop_token_ids_anywhere(self, minimal_chat_request):
"""When neither client nor server specifies stop_token_ids, result is empty."""
sampling_params = minimal_chat_request.to_sampling_params(
max_tokens=100,
default_sampling_params={},
)
assert not sampling_params.stop_token_ids
def test_only_client_stop_token_ids(self):
"""Client stop_token_ids work when no server defaults exist."""
request = ChatCompletionRequest(
model="test-model",
messages=[{"role": "user", "content": "hello"}],
stop_token_ids=[42, 43],
)
sampling_params = request.to_sampling_params(
max_tokens=100,
default_sampling_params={},
)
assert set(sampling_params.stop_token_ids) == {42, 43}
def test_duplicate_stop_token_ids_deduplicated(self):
"""Overlapping stop_token_ids between client and server are deduplicated."""
request = ChatCompletionRequest(
model="test-model",
messages=[{"role": "user", "content": "hello"}],
stop_token_ids=[200012, 55555],
)
default_sampling_params = {
"stop_token_ids": [200012, 200002],
}
sampling_params = request.to_sampling_params(
max_tokens=100,
default_sampling_params=default_sampling_params,
)
assert set(sampling_params.stop_token_ids) == {200012, 200002, 55555}
assert sampling_params.stop_token_ids == [200012, 55555, 200002]
assert len(sampling_params.stop_token_ids) == 3
class TestCompletionStopTokenIds:
"""Test stop_token_ids merging in CompletionRequest.to_sampling_params()."""
@pytest.fixture
def minimal_completion_request(self):
return CompletionRequest(
model="test-model",
prompt="hello",
)
def test_default_stop_token_ids_applied(self, minimal_completion_request):
"""Server-default stop_token_ids are applied when client sends none."""
default_sampling_params = {
"stop_token_ids": [200012, 200002],
}
sampling_params = minimal_completion_request.to_sampling_params(
max_tokens=100,
default_sampling_params=default_sampling_params,
)
assert set(sampling_params.stop_token_ids) == {200012, 200002}
def test_client_stop_token_ids_merged_with_defaults(self):
"""Client-specified stop_token_ids are merged with server defaults."""
request = CompletionRequest(
model="test-model",
prompt="hello",
stop_token_ids=[99999],
)
default_sampling_params = {
"stop_token_ids": [200012, 200002],
}
sampling_params = request.to_sampling_params(
max_tokens=100,
default_sampling_params=default_sampling_params,
)
assert set(sampling_params.stop_token_ids) == {200012, 200002, 99999}
assert sampling_params.stop_token_ids == [99999, 200012, 200002]
def test_no_stop_token_ids_anywhere(self, minimal_completion_request):
"""When neither client nor server specifies stop_token_ids, result is empty."""
sampling_params = minimal_completion_request.to_sampling_params(
max_tokens=100,
default_sampling_params={},
)
assert not sampling_params.stop_token_ids
+5 -5
View File
@@ -347,7 +347,7 @@ def test_selective_state_update(dim, dstate, has_z, itype):
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
if itype == torch.bfloat16:
rtol, atol = 1e-2, 5e-2
if torch.version.hip:
if current_platform.is_rocm() or current_platform.is_xpu():
atol *= 2
# set seed
set_random_seed(0)
@@ -437,7 +437,7 @@ def test_selective_state_update_varlen(dim, dstate, has_z, itype, max_seq_len):
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
if itype == torch.bfloat16:
rtol, atol = 5e-2, 1.5e-1
if torch.version.hip:
if current_platform.is_rocm() or current_platform.is_xpu():
atol *= 2
# set seed
set_random_seed(0)
@@ -700,7 +700,7 @@ def test_selective_state_update_with_batch_indices(
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
if itype == torch.bfloat16:
rtol, atol = 1e-1, 1e-1
if torch.version.hip:
if current_platform.is_rocm() or current_platform.is_xpu():
atol *= 2
# set seed
torch.random.manual_seed(0)
@@ -865,7 +865,7 @@ def test_selective_state_update_with_num_accepted_tokens(
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
if itype == torch.bfloat16:
rtol, atol = 5e-2, 1.5e-1
if torch.version.hip:
if current_platform.is_rocm() or current_platform.is_xpu():
atol *= 2
set_random_seed(0)
@@ -991,7 +991,7 @@ def test_selective_state_update_varlen_with_num_accepted(
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (5e-3, 1e-2)
if itype == torch.bfloat16:
rtol, atol = 5e-2, 1.5e-1
if torch.version.hip:
if current_platform.is_rocm() or current_platform.is_xpu():
atol *= 2
set_random_seed(0)
@@ -0,0 +1,180 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Equivalence test for ``precopy_mamba_align_fused_kernel``.
The V2 "align" pre-copy must migrate mamba state across block boundaries with
byte-identical semantics to the V1 copy specs (``get_conv_copy_spec`` /
``get_temporal_copy_spec``):
* conv state (SD layout, conv_width > 0): shift the sliding window by
``token_bias`` tokens -- ``state[bt[src_col], token_bias:]`` ->
``state[bt[dst_col], :conv_width - token_bias]``.
* temporal state (conv_width == 0): ``token_bias`` selects the accepted
speculative column -- ``state[bt[src_col + token_bias]]`` ->
``state[bt[dst_col]]``.
The kernel must also no-op when ``src_col < 0`` (fresh request) or
``src_col == dst_col`` (no boundary crossed).
"""
from __future__ import annotations
import torch
from vllm.platforms import current_platform
from vllm.v1.worker.mamba_utils import precopy_mamba_align_fused_kernel
try:
import pytest
pytestmark = pytest.mark.skipif(
not current_platform.is_cuda(),
reason="precopy_mamba_align_fused_kernel needs CUDA/Triton",
)
_parametrize = pytest.mark.parametrize
except ModuleNotFoundError: # allow running directly as ``python <thisfile>``
pytest = None
def _parametrize(_name, _values):
def _deco(fn):
return fn
return _deco
NUM_LAYERS = 3
CONV_WIDTH = 4 # conv_kernel - 1 + num_spec
CONV_DIM = 96
SSM_SHAPE = (4, 16, 16)
MAX_COLS = 8
def _build_state(num_blocks, device):
"""Per-layer (conv SD [nb, width, dim] bf16, ssm [nb, *shape] fp32) pools."""
convs, ssms = [], []
for _ in range(NUM_LAYERS):
convs.append(
torch.randn(
num_blocks, CONV_WIDTH, CONV_DIM, dtype=torch.bfloat16, device=device
)
)
ssms.append(
torch.randn(num_blocks, *SSM_SHAPE, dtype=torch.float32, device=device)
)
return convs, ssms
def _build_meta(convs, ssms, device):
"""Flattened per-(layer, state-type) metadata, ordered conv, ssm per layer."""
n = NUM_LAYERS * 2
base = torch.zeros(n, dtype=torch.int64, device=device)
blk_stride = torch.zeros(n, dtype=torch.int64, device=device)
elem = torch.zeros(n, dtype=torch.int32, device=device)
inner = torch.zeros(n, dtype=torch.int64, device=device)
width = torch.zeros(n, dtype=torch.int32, device=device)
group = torch.zeros(n, dtype=torch.int32, device=device)
drc = torch.zeros(n, dtype=torch.int32, device=device) # DS rows (unused, SD)
drs = torch.zeros(n, dtype=torch.int64, device=device)
i = 0
for layer in range(NUM_LAYERS):
conv, ssm = convs[layer], ssms[layer]
# conv (SD): width = size(1), inner = stride(1)
base[i] = conv.data_ptr()
blk_stride[i] = conv.stride(0) * conv.element_size()
elem[i] = conv.element_size()
width[i] = conv.size(1)
inner[i] = conv.stride(1)
i += 1
# ssm (temporal): width = 0, inner = elems per block
base[i] = ssm.data_ptr()
blk_stride[i] = ssm.stride(0) * ssm.element_size()
elem[i] = ssm.element_size()
width[i] = 0
inner[i] = ssm[0].numel()
i += 1
return base, blk_stride, elem, inner, width, group, drc, drs
def _reference(convs, ssms, bt, src_col, dst_col, bias, num_reqs):
"""Apply the V1 copy semantics on clones, reading from the pre-copy state."""
conv_pre = [c.clone() for c in convs]
ssm_pre = [s.clone() for s in ssms]
conv_ref = [c.clone() for c in convs]
ssm_ref = [s.clone() for s in ssms]
for r in range(num_reqs):
sc, dc, tb = int(src_col[r]), int(dst_col[r]), int(bias[r])
if sc < 0 or sc == dc:
continue
sblk, dblk = int(bt[r, sc]), int(bt[r, dc])
tblk = int(bt[r, sc + tb]) # temporal src column shifted by bias
for layer in range(NUM_LAYERS):
conv_ref[layer][dblk, : CONV_WIDTH - tb] = conv_pre[layer][sblk, tb:]
ssm_ref[layer][dblk] = ssm_pre[layer][tblk]
return conv_ref, ssm_ref
@_parametrize("num_reqs", [1, 4, 16])
@_parametrize("token_bias", [0, 1, 2])
def test_precopy_matches_v1_copy_specs(num_reqs, token_bias):
device = torch.device("cuda")
torch.manual_seed(0)
# Distinct physical block per (req, col) so copies never alias.
num_blocks = num_reqs * MAX_COLS + 1
bt = torch.empty(num_reqs, MAX_COLS, dtype=torch.int32, device=device)
for r in range(num_reqs):
bt[r] = torch.arange(
1 + r * MAX_COLS, 1 + (r + 1) * MAX_COLS, dtype=torch.int32, device=device
)
# Per-req columns: req 0 fresh (src=-1, skip), req 1 same block (skip),
# the rest cross from col 1 -> col 0 with the given spec token bias.
src_col = torch.full((num_reqs,), 1, dtype=torch.int32, device=device)
dst_col = torch.zeros(num_reqs, dtype=torch.int32, device=device)
bias = torch.full((num_reqs,), token_bias, dtype=torch.int32, device=device)
if num_reqs >= 1:
src_col[0] = -1 # fresh -> no copy
if num_reqs >= 2:
dst_col[1] = 1 # src_col == dst_col -> no copy
convs, ssms = _build_state(num_blocks, device)
conv_ref, ssm_ref = _reference(
convs, ssms, bt.cpu(), src_col.cpu(), dst_col.cpu(), bias.cpu(), num_reqs
)
base, blk_stride, elem, inner, width, group, drc, drs = _build_meta(
convs, ssms, device
)
bt_ptrs = torch.tensor([bt.data_ptr()], dtype=torch.int64, device=device)
idx_mapping = torch.arange(num_reqs, dtype=torch.int32, device=device)
grid = (num_reqs, NUM_LAYERS * 2)
precopy_mamba_align_fused_kernel[grid](
dst_col,
src_col,
bias,
bt_ptrs,
bt.stride(0),
base,
blk_stride,
elem,
inner,
width,
group,
drc,
drs,
idx_mapping,
num_reqs,
COPY_BLOCK_SIZE=1024,
CONV_STATE_DIM_FIRST=False,
)
torch.accelerator.synchronize()
for layer in range(NUM_LAYERS):
torch.testing.assert_close(convs[layer], conv_ref[layer], rtol=0, atol=0)
torch.testing.assert_close(ssms[layer], ssm_ref[layer], rtol=0, atol=0)
if __name__ == "__main__":
for nr in (1, 4, 16):
for tb in (0, 1, 2):
test_precopy_matches_v1_copy_specs(nr, tb)
print(f"OK num_reqs={nr} token_bias={tb}")
+1 -1
View File
@@ -406,7 +406,7 @@ def test_fused_moe_int64_overflow(workspace_init):
Reproduces the scenario from PR #34279.
"""
# ~12 GB GPU memory needed for intermediate caches
free_mem = torch.cuda.mem_get_info()[0]
free_mem = torch.accelerator.get_memory_info()[0]
if free_mem < 12 * 1024**3:
pytest.skip("Insufficient GPU memory for overflow test")
@@ -97,7 +97,7 @@ def test_sharded_state_loader(
ctx = mp.get_context("spawn")
platform_args = {}
if current_platform.is_rocm():
if current_platform.is_rocm() or current_platform.is_xpu():
platform_args["max_num_seqs"] = 1
# Run in separate processes for memory & CUDA isolation
@@ -83,7 +83,7 @@ def _ru_maxrss_bytes() -> int | None:
def _gpu_used_bytes() -> int:
torch.accelerator.synchronize()
free_bytes, total_bytes = current_platform.mem_get_info()
free_bytes, total_bytes = torch.accelerator.get_memory_info()
return int(total_bytes - free_bytes)
+54
View File
@@ -205,3 +205,57 @@ def test_image_media_io_load_file(tmp_path):
with pytest.raises(ValueError, match="Failed to load image"):
image_io.load_file(truncated_real_file)
def test_image_pixel_limit_respected():
"""A small image within the pixel limit loads successfully."""
import vllm.envs as envs
image = Image.new("RGB", (100, 100), (255, 0, 0))
from io import BytesIO
buf = BytesIO()
image.save(buf, format="PNG")
data = buf.getvalue()
assert envs.VLLM_MAX_IMAGE_PIXELS >= 100 * 100
image_io = ImageMediaIO()
result = image_io.load_bytes(data)
assert result.media.size == (100, 100)
def test_image_pixel_limit_rejected(monkeypatch):
"""An image exceeding the pixel limit is rejected before raster decode."""
import vllm.envs as envs
monkeypatch.setattr(envs, "VLLM_MAX_IMAGE_PIXELS", 100)
image = Image.new("RGB", (20, 20), (0, 255, 0))
from io import BytesIO
buf = BytesIO()
image.save(buf, format="PNG")
data = buf.getvalue()
image_io = ImageMediaIO()
with pytest.raises(ValueError, match="exceed"):
image_io.load_bytes(data)
def test_image_pixel_limit_disabled(monkeypatch):
"""Setting VLLM_MAX_IMAGE_PIXELS=0 disables the pixel limit."""
import vllm.envs as envs
monkeypatch.setattr(envs, "VLLM_MAX_IMAGE_PIXELS", 0)
image = Image.new("RGB", (1000, 1000), (0, 0, 255))
from io import BytesIO
buf = BytesIO()
image.save(buf, format="PNG")
data = buf.getvalue()
image_io = ImageMediaIO()
result = image_io.load_bytes(data)
assert result.media.size == (1000, 1000)
+94 -1
View File
@@ -31,6 +31,7 @@ from vllm.entrypoints.openai.chat_completion.protocol import (
from vllm.parser.engine.registered_adapters import (
Gemma4Parser,
Glm47MoeParser,
KimiK2Parser,
MinimaxM2Parser,
NemotronV3Parser,
Qwen3Parser,
@@ -717,6 +718,96 @@ def _build_glm47_moe(scenario: Scenario, validate: bool = True) -> Sample:
return sample
# ── Kimi K2 (native tool-call section, starts in REASONING) ──────────
_KIMI_K2_VOCAB: dict[str, int] = {
"<think>": 50,
"</think>": 51,
"<|tool_calls_section_begin|>": 60,
"<|tool_calls_section_end|>": 61,
"<|tool_call_begin|>": 62,
"<|tool_call_end|>": 63,
"<|tool_call_argument_begin|>": 64,
}
def _kimi_k2_tool_segments(
tool_calls: list[ToolCallSpec],
) -> list[tuple[str, bool]]:
segs: list[tuple[str, bool]] = [("<|tool_calls_section_begin|>", True)]
for index, tc in enumerate(tool_calls):
args = json.dumps(tc.arguments, ensure_ascii=False, separators=(",", ":"))
segs.extend(
[
("<|tool_call_begin|>", True),
(f"functions.{tc.name}:{index}\n", False),
("<|tool_call_argument_begin|>", True),
(args, False),
("<|tool_call_end|>", True),
]
)
segs.append(("<|tool_calls_section_end|>", True))
return segs
def _kimi_k2_segments(scenario: Scenario) -> list[tuple[str, bool]]:
segs: list[tuple[str, bool]] = []
if scenario.reasoning is not None:
segs.append(("<think>", True))
segs.append((scenario.reasoning, False))
if scenario.content is not None or scenario.tool_calls is not None:
segs.append(("</think>", True))
if scenario.content is not None:
segs.append((scenario.content, False))
if scenario.tool_calls is not None:
segs.extend(_kimi_k2_tool_segments(scenario.tool_calls))
return segs
def _build_kimi_k2(
scenario: Scenario,
validate: bool = True,
thinking: bool = True,
) -> Sample:
expected_reasoning = (
scenario.reasoning.rstrip()
if (thinking and scenario.reasoning is not None)
else None
)
if thinking and scenario.reasoning is None:
expected_reasoning = ""
sample = _make_sample(
sample_id=f"kimi_k2-{scenario.id}",
description=scenario.description,
vocab=_KIMI_K2_VOCAB,
segments=_kimi_k2_segments(scenario),
expected_reasoning=expected_reasoning,
expected_content=_qwen3_expected_content(scenario),
expected_tool_calls=_expected_tc(scenario),
tools=_expected_tools(scenario),
chat_template_kwargs=None if thinking else {"thinking": False},
)
if validate:
_validate_sample(
sample,
KimiK2Parser,
chat_template_kwargs=sample.chat_template_kwargs,
)
return sample
_KIMI_K2_SCENARIOS = [
*SCENARIOS,
Scenario(
id="trailing-reasoning-whitespace",
description="Reasoning trailing whitespace is stripped",
reasoning="Reasoning with trailing whitespace. \n\t",
content="Done.",
),
]
# ── Registry and public API ──────────────────────────────────────────
_BUILDERS: dict[str, Any] = {
@@ -726,6 +817,7 @@ _BUILDERS: dict[str, Any] = {
"nemotron_v3": _build_nemotron_v3,
"seed_oss": _build_seed_oss,
"glm47_moe": _build_glm47_moe,
"kimi_k2": _build_kimi_k2,
}
@@ -733,7 +825,8 @@ _BUILDERS: dict[str, Any] = {
def build_samples(model: str) -> tuple[Sample, ...]:
"""Build all scenario samples for a model, self-validated."""
builder = _BUILDERS[model]
return tuple(builder(s) for s in SCENARIOS)
scenarios = _KIMI_K2_SCENARIOS if model == "kimi_k2" else SCENARIOS
return tuple(builder(s) for s in scenarios)
def build_sample(model: str, scenario: Scenario) -> Sample:
@@ -7,7 +7,6 @@ import pytest
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
from vllm.entrypoints.openai.engine.protocol import DeltaMessage
from vllm.reasoning.identity_reasoning_parser import IdentityReasoningParser
from vllm.reasoning.kimi_k2_reasoning_parser import KimiK2ReasoningParser
from vllm.tokenizers import get_tokenizer
@@ -33,20 +32,6 @@ def kimi_k2_tokenizer():
return get_tokenizer(tokenizer_name=REASONING_MODEL_NAME, trust_remote_code=True)
def test_parser_selection_thinking_enabled(kimi_k2_tokenizer):
parser = KimiK2ReasoningParser(
kimi_k2_tokenizer, chat_template_kwargs={"thinking": True}
)
assert parser._identity_parser is None
def test_parser_selection_thinking_disabled(kimi_k2_tokenizer):
parser = KimiK2ReasoningParser(
kimi_k2_tokenizer, chat_template_kwargs={"thinking": False}
)
assert isinstance(parser._identity_parser, IdentityReasoningParser)
def test_extract_reasoning_with_think_tags(kimi_k2_tokenizer):
parser = KimiK2ReasoningParser(kimi_k2_tokenizer)
request = ChatCompletionRequest(model="test-model", messages=[], temperature=1.0)
@@ -65,7 +50,7 @@ def test_extract_reasoning_empty_thinking(kimi_k2_tokenizer):
reasoning, content = parser.extract_reasoning(
"<think></think>final answer", request
)
assert reasoning == ""
assert reasoning is None
assert content == "final answer"
@@ -96,8 +81,8 @@ def test_streaming_reasoning_then_content(kimi_k2_tokenizer):
"""Token-by-token streaming: reasoning tokens then content after </think>."""
parser = KimiK2ReasoningParser(kimi_k2_tokenizer)
think_id = parser._start_token_id
end_think_id = parser._end_token_id
think_id = parser._parser_engine._start_token_id
end_think_id = parser._parser_engine._end_token_id
# Use a real token ID from the tokenizer for regular content
regular_id = kimi_k2_tokenizer.encode("hello", add_special_tokens=False)[0]
@@ -154,8 +139,8 @@ def test_streaming_tool_section_ends_reasoning(kimi_k2_tokenizer):
"""<|tool_calls_section_begin|> in delta ends reasoning during streaming."""
parser = KimiK2ReasoningParser(kimi_k2_tokenizer)
think_id = parser._start_token_id
tool_begin_id = parser._tool_section_start_token_id
think_id = parser._parser_engine._start_token_id
tool_begin_id = parser._parser_engine._tool_section_start_token_id
regular_id = kimi_k2_tokenizer.encode("hello", add_special_tokens=False)[0]
# Tool section token arrives — should transition from reasoning to content
@@ -169,50 +154,3 @@ def test_streaming_tool_section_ends_reasoning(kimi_k2_tokenizer):
)
assert isinstance(result, DeltaMessage)
assert result.content == "<|tool_calls_section_begin|>"
def test_streaming_end_token_id_buffered(mock_kimi_k2_tokenizer):
"""When stop sequences buffer text, </think> ID arrives before its text.
The token ID is present in delta_token_ids but the actual string is not
yet in delta_text (still buffered). The parser must return None to wait
for the next delta, instead of calling find() which returns -1 and
silently corrupting the text split.
"""
parser = KimiK2ReasoningParser(mock_kimi_k2_tokenizer)
think_id = parser._start_token_id
end_think_id = parser._end_token_id
# Simulate: </think> ID arrived but text not yet flushed.
# Two token IDs in delta to bypass the single-special-token guard.
result = parser.extract_reasoning_streaming(
previous_text="some reasoning",
current_text="some reasoning extra",
delta_text="extra", # </think> text not yet flushed
previous_token_ids=[think_id],
current_token_ids=[think_id, end_think_id, 999],
delta_token_ids=[end_think_id, 999],
)
assert result is None
def test_streaming_tool_section_id_buffered(mock_kimi_k2_tokenizer):
"""When stop sequences buffer text, tool section start ID arrives before its text.
Same buffering scenario as above but for <|tool_calls_section_begin|>.
Without the guard, find() returns -1 and delta_text[:tool_index] silently
drops the last character of reasoning.
"""
parser = KimiK2ReasoningParser(mock_kimi_k2_tokenizer)
think_id = parser._start_token_id
tool_begin_id = parser._tool_section_start_token_id
result = parser.extract_reasoning_streaming(
previous_text="some reasoning",
current_text="some reasoning extra",
delta_text="extra", # tool section text not yet flushed
previous_token_ids=[think_id],
current_token_ids=[think_id, tool_begin_id, 999],
delta_token_ids=[tool_begin_id, 999],
)
assert result is None
+3 -1
View File
@@ -1292,7 +1292,9 @@ def test_vllm_config_explicit_overrides():
compilation_config=compilation_config,
)
assert config.compilation_config.cudagraph_mode == CUDAGraphMode.NONE
assert config.compilation_config.pass_config.enable_qk_norm_rope_fusion is True
assert config.compilation_config.pass_config.enable_qk_norm_rope_fusion is (
current_platform.is_cuda_alike() or current_platform.is_xpu()
)
# Mode should still use default for O2
assert config.compilation_config.mode == CompilationMode.VLLM_COMPILE
@@ -102,45 +102,36 @@ def test_get_model_structural_tag_supports_vllm_hermes(
)
assert isinstance(tag, StructuralTag)
assert tag.model_dump() == {
"type": "structural_tag",
"format": {
"type": "tags_with_separator",
"tags": [
{
"type": "tag",
"begin": '<tool_call>\n{"name": "get_weather", "arguments": ',
"content": {
"type": "json_schema",
"json_schema": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
"style": "json",
},
"end": "}\n</tool_call>",
},
{
"type": "tag",
"begin": '<tool_call>{"name": "get_weather", "arguments": ',
"content": {
"type": "json_schema",
"json_schema": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
"style": "json",
},
"end": "}</tool_call>",
},
],
"separator": "",
"at_least_one": True,
"stop_after_first": False,
},
# Assert the semantically meaningful structure rather than the full
# model_dump(), which gains version-specific keys across xgrammar releases
# (e.g. "any_order" was added to json_schema content in 0.2.3).
dump = tag.model_dump()
assert dump["type"] == "structural_tag"
fmt = dump["format"]
assert fmt["type"] == "tags_with_separator"
assert fmt["separator"] == ""
assert fmt["at_least_one"] is True
assert fmt["stop_after_first"] is False
expected_schema = {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
}
expected_tags = [
('<tool_call>\n{"name": "get_weather", "arguments": ', "}\n</tool_call>"),
('<tool_call>{"name": "get_weather", "arguments": ', "}</tool_call>"),
]
assert len(fmt["tags"]) == len(expected_tags)
for tag_dump, (begin, end) in zip(fmt["tags"], expected_tags):
assert tag_dump["type"] == "tag"
assert tag_dump["begin"] == begin
assert tag_dump["end"] == end
content = tag_dump["content"]
assert content["type"] == "json_schema"
assert content["json_schema"] == expected_schema
def test_hermes_required_tool_calls_use_empty_separator():
+53 -4
View File
@@ -604,6 +604,13 @@ class RemoteVLLMServer:
mem_info = nvmlDeviceGetMemoryInfo(handle)
total_used += mem_info.used
return total_used
elif current_platform.is_xpu():
total_used = 0
device_count = current_platform.device_count()
for i in range(device_count):
free, total = torch.xpu.mem_get_info(i)
total_used += total - free
return total_used
except Exception as e:
print(f"[RemoteOpenAIServer] Could not query GPU memory: {e}")
return None
@@ -1501,6 +1508,9 @@ def wait_for_gpu_memory_to_clear(
threshold_bytes: int | dict[int, int] | None = None,
threshold_ratio: float | dict[int, float] | None = None,
timeout_s: float = 120,
stable_duration_s: float = 0,
stable_tolerance_bytes: int = 512 * 1024**2,
poll_interval_s: float = 5,
) -> None:
assert threshold_bytes is not None or threshold_ratio is not None
devices = get_physical_device_indices(devices)
@@ -1528,8 +1538,13 @@ def wait_for_gpu_memory_to_clear(
# Use nvml instead of pytorch to reduce measurement error from torch cuda
# context.
start_time = time.time()
stable_since: float | None = None
stable_used_bytes: dict[int, int] | None = None
while True:
output_raw = record_gpu_memory_usage_stats(devices=devices)
used_bytes_by_device = {
device: int(gb_used * 2**30) for device, (gb_used, _) in output_raw.items()
}
output = {
device: f"{gb_used:.02f}/{gb_total:.02f}"
for device, (gb_used, gb_total) in output_raw.items()
@@ -1577,15 +1592,45 @@ def wait_for_gpu_memory_to_clear(
dur_s = time.time() - start_time
if all_free:
print(f"Done waiting for free GPU memory on ({threshold=}) {dur_s=:.02f}")
break
if stable_duration_s <= 0:
print(
f"Done waiting for free GPU memory on devices {devices=} "
f"({threshold=}) {dur_s=:.02f}"
)
break
now = time.time()
if stable_used_bytes is None:
stable_since = now
stable_used_bytes = used_bytes_by_device
else:
memory_changed = any(
abs(used_bytes_by_device[device] - stable_used_bytes[device])
> stable_tolerance_bytes
for device in devices
)
if memory_changed:
stable_since = now
stable_used_bytes = used_bytes_by_device
elif (
stable_since is not None and now - stable_since >= stable_duration_s
):
print(
f"Done waiting for stable free GPU memory on devices "
f"{devices=} ({threshold=}) {dur_s=:.02f}"
)
break
else:
stable_since = None
stable_used_bytes = None
if dur_s >= timeout_s:
raise ValueError(
f"Memory of devices not free after {dur_s=:.02f} ({threshold=})"
f"Memory of devices {devices=} not free after "
f"{dur_s=:.02f} ({threshold=})"
)
time.sleep(5)
time.sleep(poll_interval_s)
def wait_for_rocm_memory_to_settle(
@@ -1606,11 +1651,15 @@ def wait_for_rocm_memory_to_settle(
num_gpus = current_platform.device_count()
if num_gpus == 0:
return
if threshold_ratio is None:
threshold_ratio = 0.1
wait_for_gpu_memory_to_clear(
devices=list(range(num_gpus)),
threshold_ratio=threshold_ratio,
timeout_s=timeout_s,
stable_duration_s=2.0,
poll_interval_s=1.0,
)
+11 -9
View File
@@ -36,7 +36,7 @@ def test_memory_profiling():
weights_memory = 128 * 1024 * 1024 * 4 # 512 MiB
def measure_current_non_torch():
free, total = torch.cuda.mem_get_info()
free, total = torch.accelerator.get_memory_info()
current_used = total - free
current_torch = torch.accelerator.memory_reserved()
current_non_torch = current_used - current_torch
@@ -81,8 +81,9 @@ def test_memory_snapshot_uses_psutil_on_integrated_gpu():
with (
patch("vllm.utils.mem_utils.current_platform") as mock_platform,
patch("vllm.utils.mem_utils.psutil") as mock_psutil,
patch("torch.accelerator") as mock_accelerator,
):
mock_platform.mem_get_info.return_value = (
mock_accelerator.get_memory_info.return_value = (
mock_cuda_free,
mock_cuda_total,
)
@@ -90,8 +91,8 @@ def test_memory_snapshot_uses_psutil_on_integrated_gpu():
mock_platform.memory_stats.return_value = {
"allocated_bytes.all.peak": 0,
}
mock_platform.memory_reserved.return_value = 0
mock_platform.current_device = lambda: "cuda:0"
mock_accelerator.memory_reserved.return_value = 0
mock_accelerator.current_device = lambda: "cuda:0"
mock_vmem = MagicMock()
mock_vmem.available = mock_psutil_available
@@ -105,24 +106,25 @@ def test_memory_snapshot_uses_psutil_on_integrated_gpu():
def test_memory_snapshot_uses_cuda_on_discrete_gpu():
"""On discrete GPUs, free_memory should come from CUDA mem_get_info."""
"""On discrete GPUs, free_memory should come from accelerator get_memory_info."""
mock_cuda_free = 70 * 1024**3
mock_cuda_total = 80 * 1024**3
with (
patch("vllm.utils.mem_utils.current_platform") as mock_platform,
patch("vllm.utils.mem_utils.psutil") as mock_psutil,
patch("torch.accelerator") as mock_accelerator,
):
mock_platform.mem_get_info.return_value = (
mock_accelerator.get_memory_info.return_value = (
mock_cuda_free,
mock_cuda_total,
)
mock_platform.is_integrated_gpu.return_value = False
mock_platform.memory_stats.return_value = {
mock_accelerator.memory_stats.return_value = {
"allocated_bytes.all.peak": 0,
}
mock_platform.memory_reserved.return_value = 0
mock_platform.current_device = lambda: "cuda:0"
mock_accelerator.memory_reserved.return_value = 0
mock_accelerator.current_device = lambda: "cuda:0"
snapshot = MemorySnapshot(device="cuda:0")
+280 -16
View File
@@ -11,6 +11,7 @@ import datasets
import pytest
import torch
import vllm.envs as envs
from tests.utils import create_new_process_for_each_test
from vllm import LLM, SamplingParams, TokensPrompt
from vllm.config import CacheConfig
@@ -494,12 +495,7 @@ def apply_patch(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(mamba_utils, "do_mamba_copy_block", fake_copy_fn)
@create_new_process_for_each_test()
def test_mamba_prefix_cache(monkeypatch: pytest.MonkeyPatch):
run_ref_mamba_state_in_subprocess()
apply_patch(monkeypatch)
prompt_dataset = datasets.load_dataset("heheda/a_long_article")
full_prompt = prompt_dataset["train"][0]["text"]
def get_mamba_prefix_cache_step_configs() -> dict[str, TestConfig]:
tests = {
"accept_1": TestConfig(
num_prompt_tokens=554,
@@ -731,6 +727,27 @@ def test_mamba_prefix_cache(monkeypatch: pytest.MonkeyPatch):
),
}
return tests
def fill_following_kv_cache_block_ids(test_config: TestConfig) -> None:
for step_action_prev, step_action_next in zip(
test_config.step_actions[:-1], test_config.step_actions[1:]
):
if len(step_action_next.kv_cache_block_ids) == 0:
step_action_next.kv_cache_block_ids = (
step_action_prev.kv_cache_block_ids.copy()
)
@create_new_process_for_each_test()
def test_mamba_prefix_cache_mrv1(monkeypatch: pytest.MonkeyPatch):
run_ref_mamba_state_in_subprocess()
apply_patch(monkeypatch)
prompt_dataset = datasets.load_dataset("heheda/a_long_article")
full_prompt = prompt_dataset["train"][0]["text"]
tests = get_mamba_prefix_cache_step_configs()
engine = LLM(
model=MODEL,
enable_prefix_caching=True,
@@ -758,16 +775,7 @@ def test_mamba_prefix_cache(monkeypatch: pytest.MonkeyPatch):
)
global cur_step_action_idx
cur_step_action_idx = 0
for step_action_prev, step_action_next in zip(
test_config.step_actions[:-1], test_config.step_actions[1:]
):
if (
step_action_next.kv_cache_block_ids is not None
and len(step_action_next.kv_cache_block_ids) == 0
):
prev_block_ids = step_action_prev.kv_cache_block_ids
if prev_block_ids is not None:
step_action_next.kv_cache_block_ids = prev_block_ids.copy()
fill_following_kv_cache_block_ids(test_config)
global step_actions
step_actions = test_config.step_actions
_ = engine.generate(
@@ -787,3 +795,259 @@ def test_mamba_prefix_cache(monkeypatch: pytest.MonkeyPatch):
del engine
torch.accelerator.empty_cache()
cleanup_dist_env_and_memory()
@create_new_process_for_each_test()
def test_mamba_prefix_cache_mrv2(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "1")
envs.disable_envs_cache()
from vllm.v1.worker.gpu.model_runner import GPUModelRunner as MRV2GPUModelRunner
from vllm.v1.worker.gpu.model_states.mamba_hybrid import (
MambaHybridModelState,
)
from vllm.v1.worker.gpu.sample.output import SamplerOutput as MRV2SamplerOutput
events: list[int] = []
original_execute_model = MRV2GPUModelRunner.execute_model
original_sample = MRV2GPUModelRunner.sample
original_preprocess_state = MambaHybridModelState.preprocess_state
original_postprocess_state = MambaHybridModelState.postprocess_state
original_step_action_fn = InprocClient.get_output
original_allocate_slots = KVCacheManager.allocate_slots
captured: dict[str, Any] = {}
def temporal_states(model_state, block_tables, kv_cache_config):
# Qwen3-Next keeps the temporal (ssm) state as the last Mamba cache.
forward_context = (
model_state.vllm_config.compilation_config.static_forward_context
)
group_ids, _ = get_mamba_groups(kv_cache_config)
for group_id in group_ids:
block_table = block_tables[group_id]
for layer_name in kv_cache_config.kv_cache_groups[group_id].layer_names:
yield forward_context[layer_name].kv_cache[-1], block_table
def temporal_block(temporal_state, block_table, col):
return temporal_state[int(block_table[0, col].item())]
def wrapped_preprocess_state(
self: MambaHybridModelState,
input_batch: Any,
block_tables: tuple[torch.Tensor, ...],
kv_cache_config: KVCacheConfig,
num_computed_tokens: torch.Tensor,
) -> None:
captured["block_tables"] = block_tables
captured["kv_cache_config"] = kv_cache_config
expected = (
None if cur_step_action is None else cur_step_action.preprocess_copy_idx
)
snapshots = []
if expected is not None and expected != (-1, -1):
for temporal, bt in temporal_states(self, block_tables, kv_cache_config):
snapshots.append(
(temporal, bt, temporal_block(temporal, bt, expected[0]).clone())
)
ret = original_preprocess_state(
self, input_batch, block_tables, kv_cache_config, num_computed_tokens
)
if cur_step_action is not None:
req_idx = int(input_batch.idx_mapping[0].item())
src_col = int(self._mamba_src_col_gpu[req_idx].item())
off = int(self._mamba_src_off_gpu[req_idx].item())
dst = int(self._mamba_state_idx_gpu[req_idx].item())
actual = (-1, -1) if src_col < 0 or src_col == dst else (src_col + off, dst)
assert actual == expected, (
f"V2 align preprocess copy: expected={expected}, "
f"actual={actual}, {cur_step_action=}"
)
for temporal, bt, src_state in snapshots:
torch.testing.assert_close(
temporal_block(temporal, bt, expected[1]), src_state
)
return ret
def wrapped_postprocess_state(
self: MambaHybridModelState,
idx_mapping: torch.Tensor,
num_sampled: torch.Tensor | int,
num_computed_tokens: torch.Tensor | None = None,
) -> None:
action = cur_step_action
block_tables = captured.get("block_tables")
kv_cache_config = captured.get("kv_cache_config")
# The postprocess kernel does not expose its indices, so only the copy
# case is checked, by effect: snapshot the src block, expect dst == src.
if (
action is None
or num_computed_tokens is None
or block_tables is None
or action.postprocess_copy_idx == (-1, -1)
):
return original_postprocess_state(
self, idx_mapping, num_sampled, num_computed_tokens
)
expected = action.postprocess_copy_idx
snapshots = [
(temporal, bt, temporal_block(temporal, bt, expected[0]).clone())
for temporal, bt in temporal_states(self, block_tables, kv_cache_config)
]
ret = original_postprocess_state(
self, idx_mapping, num_sampled, num_computed_tokens
)
for temporal, bt, src_state in snapshots:
torch.testing.assert_close(
temporal_block(temporal, bt, expected[1]), src_state
)
return ret
def wrapped_execute_model(
self: MRV2GPUModelRunner,
scheduler_output: SchedulerOutput,
*args: Any,
**kwargs: Any,
):
events.extend(
req.num_computed_tokens for req in scheduler_output.scheduled_new_reqs
)
events.extend(scheduler_output.scheduled_cached_reqs.num_computed_tokens)
if cur_step_action is not None:
num_scheduled_tokens = next(
iter(scheduler_output.num_scheduled_tokens.values())
)
assert num_scheduled_tokens == cur_step_action.num_scheduled_tokens
ret = original_execute_model(self, scheduler_output, *args, **kwargs)
if cur_step_action is not None and self.execute_model_state is not None:
input_batch = self.execute_model_state.input_batch
assert (
cur_step_action.num_computed_tokens_start
== input_batch.positions[input_batch.query_start_loc[0]].item()
)
return ret
def fake_sample(
self: MRV2GPUModelRunner,
hidden_states: torch.Tensor,
input_batch: Any,
grammar_output: Any,
):
if cur_step_action is None:
return original_sample(self, hidden_states, input_batch, grammar_output)
num_reqs = input_batch.num_reqs
sampled_token_ids = torch.ones(
(num_reqs, self.num_speculative_steps + 1),
device=hidden_states.device,
dtype=torch.int64,
)
num_logits = torch.tensor(
input_batch.cu_num_logits_np[1 : num_reqs + 1]
- input_batch.cu_num_logits_np[:num_reqs],
device=hidden_states.device,
dtype=torch.int32,
)
accepted = torch.full_like(num_logits, num_accepted_tokens)
num_sampled = torch.minimum(accepted, num_logits)
prefill_lens = self.req_states.prefill_len.gpu[input_batch.idx_mapping]
is_chunked_prefill = input_batch.seq_lens[:num_reqs] < prefill_lens
num_sampled = torch.where(is_chunked_prefill, 0, num_sampled)
num_rejected = torch.where(is_chunked_prefill, 0, num_logits - num_sampled)
sampler_output = MRV2SamplerOutput(
sampled_token_ids=sampled_token_ids,
logprobs_tensors=None,
num_nans=None,
num_sampled=num_sampled,
)
return sampler_output, num_sampled, num_rejected
monkeypatch.setattr(
InprocClient,
"get_output",
get_fake_step_action_fn(original_step_action_fn),
)
monkeypatch.setattr(
KVCacheManager,
"allocate_slots",
get_fake_allocate_slots_fn(original_allocate_slots),
)
monkeypatch.setattr(MRV2GPUModelRunner, "execute_model", wrapped_execute_model)
monkeypatch.setattr(MRV2GPUModelRunner, "sample", fake_sample)
monkeypatch.setattr(
MambaHybridModelState, "preprocess_state", wrapped_preprocess_state
)
monkeypatch.setattr(
MambaHybridModelState, "postprocess_state", wrapped_postprocess_state
)
engine = LLM(
model=MODEL,
load_format="dummy",
enforce_eager=True,
skip_tokenizer_init=True,
enable_prefix_caching=True,
block_size=BLOCK_SIZE,
mamba_cache_mode="align",
speculative_config={
"method": "qwen3_next_mtp",
"num_speculative_tokens": num_speculative_tokens,
},
max_num_batched_tokens=3072,
max_model_len=BLOCK_SIZE * 12,
hf_overrides={"num_hidden_layers": NUM_HIDDEN_LAYERS},
seed=42,
)
try:
tests = get_mamba_prefix_cache_step_configs()
global step_actions
global cur_step_action_idx
global num_accepted_tokens
for test_name, test_config in tests.items():
num_accepted_tokens = test_config.num_accepted_tokens
cur_step_action_idx = 0
fill_following_kv_cache_block_ids(test_config)
step_actions = test_config.step_actions
sampling_params = SamplingParams(
temperature=0.0,
max_tokens=test_config.num_generated_tokens,
ignore_eos=True,
)
_ = engine.generate(
[TokensPrompt(prompt_token_ids=[1] * test_config.num_prompt_tokens)],
sampling_params=sampling_params,
)
assert cur_step_action_idx == len(test_config.step_actions), test_name
assert (
engine.llm_engine.engine_core.engine_core.scheduler.reset_prefix_cache()
)
step_actions = []
cur_step_action_idx = 0
num_accepted_tokens = 1
prompt = TokensPrompt(prompt_token_ids=[1] * (BLOCK_SIZE * 2))
sampling_params = SamplingParams(
temperature=0.0,
max_tokens=1,
ignore_eos=True,
)
_ = engine.generate([prompt], sampling_params=sampling_params)
first_event_count = len(events)
_ = engine.generate([prompt], sampling_params=sampling_params)
second_events = events[first_event_count:]
prefix_hits = [
num_computed_tokens
for num_computed_tokens in second_events
if num_computed_tokens >= BLOCK_SIZE
]
assert prefix_hits, (
"Expected the second identical prompt to hit prefix cache, "
f"got events={second_events!r}"
)
assert engine.llm_engine.engine_core.engine_core.scheduler.reset_prefix_cache()
finally:
del engine
torch.accelerator.empty_cache()
cleanup_dist_env_and_memory()
@@ -33,14 +33,37 @@ from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_utils import
)
from vllm.utils.network_utils import get_open_port
from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheConfig,
KVCacheGroupSpec,
)
from vllm.v1.request import RequestStatus
from .utils import create_request, create_scheduler, create_vllm_config
def _make_test_kv_cache_config() -> KVCacheConfig:
return KVCacheConfig(num_blocks=0, kv_cache_tensors=[], kv_cache_groups=[])
return KVCacheConfig(
num_blocks=0,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
[
"model.layers.0.self_attn",
"model.layers.1.self_attn",
"model.layers.0.mla_attn",
"model.layers.1.eagle_attn",
],
FullAttentionSpec(
block_size=16,
num_kv_heads=4,
head_size=64,
dtype=torch.float16,
),
)
],
)
class FakeMooncakeWrapper:
@@ -126,6 +149,8 @@ async def test_build_transfer_params_separates_prefill_pp_layers():
worker.is_kv_producer = True
worker.tp_rank = 0
worker.tp_size = 1
worker.kv_cache_config = _make_test_kv_cache_config()
worker._physical_blocks_per_logical_kv_block = 1
worker.transfer_topo = SimpleNamespace(local_replicates_kv_cache=False)
block_len = 256
@@ -206,6 +231,7 @@ async def test_build_transfer_params_separates_prefill_pp_layers():
req_blocks={"d-req-pp": (transfer_id, [[20, 21]])},
kv_caches_base_addr=[region.base_addr for region in remote_regions],
block_lens=[region.block_len for region in remote_regions],
kv_block_lens=[region.kv_block_len for region in remote_regions],
registered_layer_names=[region.layer_name for region in remote_regions],
registered_layer_indices=[region.layer_index for region in remote_regions],
)
@@ -266,6 +292,7 @@ async def test_send_kv_to_decode_aligns_consumer_regions_by_layer_metadata(
kv_half = block_len // 2
prefill_worker.kv_caches_base_addr = [0x1000]
prefill_worker.block_len_per_layer = [block_len]
prefill_worker.kv_block_len_per_layer = [kv_half]
prefill_worker.registered_layer_names = ["model.layers.1.self_attn"]
prefill_worker.registered_layer_indices = [1]
@@ -294,6 +321,7 @@ async def test_send_kv_to_decode_aligns_consumer_regions_by_layer_metadata(
req_blocks={"d-req-layer-align": (transfer_id, [[20]])},
kv_caches_base_addr=[0xA000, 0xB000],
block_lens=[block_len, block_len],
kv_block_lens=[kv_half, kv_half],
registered_layer_names=[
"model.layers.0.self_attn",
"model.layers.1.self_attn",
@@ -804,7 +832,9 @@ async def test_kv_producer(monkeypatch):
prefill_worker = prefill_connector.connector_worker
prefill_worker.kv_caches_base_addr = [0x1000]
block_len = 4096
kv_half = block_len // 2
prefill_worker.block_len_per_layer = [block_len]
prefill_worker.kv_block_len_per_layer = [kv_half]
prefill_worker.registered_layer_names = ["model.layers.0.self_attn"]
prefill_worker.registered_layer_indices = [0]
@@ -832,6 +862,7 @@ async def test_kv_producer(monkeypatch):
req_blocks={"d-req-1": (transfer_id, [[20, 21]])},
kv_caches_base_addr=[0x2000],
block_lens=[block_len],
kv_block_lens=[kv_half],
registered_layer_names=["model.layers.0.self_attn"],
registered_layer_indices=[0],
)
@@ -845,8 +876,6 @@ async def test_kv_producer(monkeypatch):
) as mock_send_blocks:
# With blocks-first layout, each block is virtually split
# into K and V halves, producing non-coalesced transfers.
kv_half = block_len // 2
def expected_split_transfers(src_base, dst_base, src_blocks, dst_blocks):
"""Build expected (src_ptrs, dst_ptrs, lengths) for
virtual-split K/V transfers."""
@@ -981,6 +1010,7 @@ async def test_kv_consumuer(monkeypatch):
decode_worker = decode_connector.connector_worker
decode_worker.kv_caches_base_addr = [0x1000]
decode_worker.block_len_per_layer = [4096]
decode_worker.kv_block_len_per_layer = [4096]
decode_worker.registered_layer_names = ["model.layers.0.self_attn"]
decode_worker.registered_layer_indices = [0]
decode_worker.rpc_port = 54321
@@ -1236,6 +1266,7 @@ async def test_kv_producer_heterogeneous_tp(monkeypatch, d_tp_size):
prefill_worker.kv_caches_base_addr = [0x1000]
prefill_worker.block_len_per_layer = [local_block_len]
prefill_worker.kv_block_len_per_layer = [local_block_len // 2]
prefill_worker.registered_layer_names = ["model.layers.0.self_attn"]
prefill_worker.registered_layer_indices = [0]
@@ -1283,6 +1314,7 @@ async def test_kv_producer_heterogeneous_tp(monkeypatch, d_tp_size):
},
kv_caches_base_addr=[0x2000],
block_lens=[remote_block_len],
kv_block_lens=[remote_block_len // 2],
registered_layer_names=["model.layers.0.self_attn"],
registered_layer_indices=[0],
)
@@ -257,6 +257,7 @@ async def test_build_transfer_params_multi_group_trimming(monkeypatch):
},
kv_caches_base_addr=[0x2000],
block_lens=[block_len],
kv_block_lens=[block_len],
)
local_regions = [
@@ -348,6 +349,7 @@ async def test_build_transfer_params_group_count_mismatch(monkeypatch):
},
kv_caches_base_addr=[0x2000],
block_lens=[block_len],
kv_block_lens=[block_len],
)
local_regions = [
@@ -0,0 +1,386 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for MooncakeConnector hybrid FA + GDN support.
GDN is represented as a MambaSpec in vLLM, so these tests exercise the
Mooncake MambaSpec path with mamba_type=GDN_ATTN. Mamba2 is intentionally not
validated by this test module.
"""
import asyncio
from types import SimpleNamespace
from unittest.mock import patch
import pytest
import torch
from vllm.config import set_current_vllm_config
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_connector import (
KVConnectorRole,
MooncakeConnector,
MooncakeConnectorScheduler,
MooncakeConnectorWorker,
MooncakeXferMetadata,
SendBlockMeta,
TransferRegion,
)
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheConfig,
KVCacheGroupSpec,
MambaSpec,
)
from .test_mooncake_connector import patch_worker_dependencies
from .utils import create_request, create_vllm_config
def noop_shutdown():
pass
def make_hybrid_gdn_kv_cache_config(block_size: int) -> KVCacheConfig:
gdn_spec = MambaSpec(
block_size=block_size,
shapes=((6, 3), (1, 2, 2)),
dtypes=(torch.float16, torch.float16),
mamba_type=MambaAttentionBackendEnum.GDN_ATTN,
)
assert gdn_spec.mamba_type == MambaAttentionBackendEnum.GDN_ATTN
return KVCacheConfig(
num_blocks=16,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["model.layers.0.self_attn"],
FullAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float16,
),
),
KVCacheGroupSpec(
["model.layers.1.linear_attn"],
gdn_spec,
),
],
)
def make_hybrid_gdn_scheduler(kv_role: str) -> MooncakeConnectorScheduler:
vllm_config = create_vllm_config(
kv_connector="MooncakeConnector",
kv_role=kv_role,
)
vllm_config.scheduler_config.disable_hybrid_kv_cache_manager = False
return MooncakeConnectorScheduler(
vllm_config=vllm_config,
engine_id="test-engine",
kv_cache_config=make_hybrid_gdn_kv_cache_config(
vllm_config.cache_config.block_size
),
)
@pytest.mark.cpu_test
def test_hybrid_gdn_remote_prefill_uses_mamba_n_minus_one():
scheduler = make_hybrid_gdn_scheduler(kv_role="kv_consumer")
request = create_request(num_tokens=10, do_remote_prefill=True)
num_new_tokens, is_async = scheduler.get_num_new_matched_tokens(
request, num_computed_tokens=0
)
assert num_new_tokens == request.num_prompt_tokens - 1
assert is_async is True
@pytest.mark.cpu_test
def test_hybrid_gdn_remote_decode_truncates_prefill_once():
scheduler = make_hybrid_gdn_scheduler(kv_role="kv_producer")
request = create_request(num_tokens=10, do_remote_decode=True)
original_tokens = list(request.prompt_token_ids)
num_new_tokens, is_async = scheduler.get_num_new_matched_tokens(
request, num_computed_tokens=0
)
assert num_new_tokens == 0
assert is_async is False
assert request.prompt_token_ids == original_tokens[:-1]
assert request._all_token_ids == original_tokens[:-1]
assert request.num_prompt_tokens == len(original_tokens) - 1
assert request.max_tokens == 1
assert request.kv_transfer_params["_p_side_truncated"] is True
scheduler.get_num_new_matched_tokens(request, num_computed_tokens=0)
assert request.prompt_token_ids == original_tokens[:-1]
def test_register_kv_caches_emits_fa_and_gdn_regions(monkeypatch):
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
vllm_config = create_vllm_config(
kv_connector="MooncakeConnector",
kv_role="kv_consumer",
)
kv_cache_config = make_hybrid_gdn_kv_cache_config(
vllm_config.cache_config.block_size
)
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
connector = MooncakeConnector(
vllm_config,
KVConnectorRole.WORKER,
kv_cache_config,
)
worker = connector.connector_worker
fa_cache = torch.empty((2, 2, 11), dtype=torch.float16)
gdn_conv_state = torch.empty((2, 22), dtype=torch.float16)
gdn_ssm_state = torch.empty((2, 4), dtype=torch.float16)
worker.register_kv_caches(
{
"model.layers.0.self_attn": fa_cache,
"model.layers.1.linear_attn": (gdn_conv_state, gdn_ssm_state),
}
)
assert worker.transfer_topo.is_mamba is True
assert worker.registered_layer_names == [
"model.layers.0.self_attn",
"model.layers.1.linear_attn",
]
assert worker.registered_group_indices == [0, 1]
assert worker.kv_caches_base_addr == [
fa_cache.data_ptr(),
gdn_conv_state.data_ptr(),
]
worker.shutdown()
worker.shutdown = noop_shutdown
connector.connector_worker = None
def test_register_kv_caches_deduplicates_shared_backing_memory(monkeypatch):
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
vllm_config = create_vllm_config(
kv_connector="MooncakeConnector",
kv_role="kv_consumer",
)
kv_cache_config = make_hybrid_gdn_kv_cache_config(
vllm_config.cache_config.block_size
)
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
connector = MooncakeConnector(
vllm_config,
KVConnectorRole.WORKER,
kv_cache_config,
)
worker = connector.connector_worker
backing = torch.empty((4, 64), dtype=torch.float16)
fa_cache = backing[:2, :16]
gdn_conv_state = backing[:3]
gdn_ssm_state = torch.empty((3, 4), dtype=torch.float16)
with patch.object(
worker.engine, "batch_register_memory", return_value=0
) as batch_register_memory:
worker.register_kv_caches(
{
"model.layers.0.self_attn": fa_cache,
"model.layers.1.linear_attn": (gdn_conv_state, gdn_ssm_state),
}
)
assert worker.kv_caches_base_addr == [
fa_cache.data_ptr(),
gdn_conv_state.data_ptr(),
]
batch_register_memory.assert_called_once()
registered_ptrs, registered_lens = batch_register_memory.call_args[0]
assert registered_ptrs == [backing.data_ptr()]
assert registered_lens == [backing.untyped_storage().nbytes()]
worker.shutdown()
worker.shutdown = noop_shutdown
connector.connector_worker = None
def test_hybrid_gdn_transfer_params_preserve_group_identity(monkeypatch):
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
vllm_config = create_vllm_config(
kv_connector="MooncakeConnector",
kv_role="kv_producer",
)
kv_cache_config = make_hybrid_gdn_kv_cache_config(
vllm_config.cache_config.block_size
)
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
connector = MooncakeConnector(
vllm_config,
KVConnectorRole.WORKER,
kv_cache_config,
)
worker = connector.connector_worker
block_len = 0x100
transfer_id = "xfer-hybrid-gdn"
async def build_transfer_params():
send_meta = SendBlockMeta(
p_req_id="p-hybrid-gdn",
transfer_id=transfer_id,
local_block_ids=[
[10, 11],
[NULL_BLOCK_ID, 4],
],
ready=asyncio.Event(),
)
return await worker._build_transfer_params(
[("d-hybrid-gdn", send_meta)],
xfer_meta,
local_regions,
remote_regions,
)
xfer_meta = MooncakeXferMetadata(
remote_hostname="consumer-host",
remote_port=54321,
remote_tp_size=1,
remote_tp_rank=0,
req_blocks={
"d-hybrid-gdn": (
transfer_id,
[
[30, 31],
[NULL_BLOCK_ID, 7],
],
)
},
kv_caches_base_addr=[],
block_lens=[],
kv_block_lens=[],
)
local_regions = [
TransferRegion(
layer_name="model.layers.1.linear_attn",
layer_index=1,
base_addr=0x5000,
block_len=block_len,
kv_block_len=block_len,
group_index=1,
),
TransferRegion(
layer_name="model.layers.0.self_attn",
layer_index=0,
base_addr=0x1000,
block_len=block_len,
kv_block_len=block_len,
group_index=0,
),
]
remote_regions = [
TransferRegion(
layer_name="model.layers.1.linear_attn",
layer_index=1,
base_addr=0x6000,
block_len=block_len,
kv_block_len=block_len,
group_index=1,
),
TransferRegion(
layer_name="model.layers.0.self_attn",
layer_index=0,
base_addr=0x2000,
block_len=block_len,
kv_block_len=block_len,
group_index=0,
),
]
src_ptrs, dst_ptrs, lengths, err_reqs, err_msg = asyncio.run(
build_transfer_params()
)
assert err_reqs == []
assert err_msg is None
assert src_ptrs == [
0x5000 + 4 * block_len,
0x1000 + 10 * block_len,
]
assert dst_ptrs == [
0x6000 + 7 * block_len,
0x2000 + 30 * block_len,
]
assert lengths == [block_len, 2 * block_len]
worker.shutdown()
worker.shutdown = noop_shutdown
connector.connector_worker = None
def test_logical_to_kernel_block_ids_expands_fa_not_gdn():
worker = object.__new__(MooncakeConnectorWorker)
worker.shutdown = noop_shutdown
worker._physical_blocks_per_logical_kv_block = 17
worker.kv_cache_config = make_hybrid_gdn_kv_cache_config(block_size=544)
block_ids = [[2], [2]]
kernel_block_ids = worker._logical_to_kernel_block_ids(block_ids)
assert kernel_block_ids == [list(range(34, 51)), [2]]
def test_hybrid_gdn_splits_fa_regions_but_keeps_gdn_state_whole(
monkeypatch,
):
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
vllm_config = create_vllm_config(
kv_connector="MooncakeConnector",
kv_role="kv_producer",
)
kv_cache_config = make_hybrid_gdn_kv_cache_config(
vllm_config.cache_config.block_size
)
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
connector = MooncakeConnector(
vllm_config,
KVConnectorRole.WORKER,
kv_cache_config,
)
worker = connector.connector_worker
worker.transfer_topo = SimpleNamespace(virtually_split_kv_in_blocks=True)
regions = worker._get_transfer_regions(
base_addrs=[0x1000, 0x2000],
block_lens=[0x100, 0x100],
kv_block_lens=[0x40, 0x100],
layer_names=[
"model.layers.0.self_attn",
"model.layers.1.linear_attn",
],
layer_indices=[0, 1],
group_indices=[0, 1],
)
assert [
(region.group_index, region.base_addr, region.kv_block_len)
for region in regions
] == [
(0, 0x1000, 0x40),
(0, 0x1040, 0x40),
(1, 0x2000, 0x100),
]
worker.shutdown()
worker.shutdown = noop_shutdown
connector.connector_worker = None
@@ -29,9 +29,8 @@ def _gpu_snapshot(tag: str, prev_alloc: float = 0.0) -> dict:
torch.accelerator.synchronize()
alloc = torch.accelerator.memory_allocated()
reserved = torch.accelerator.memory_reserved()
# mem_get_info is not available on torch.accelerator
try:
drv_free, drv_total = torch.cuda.mem_get_info()
drv_free, drv_total = torch.accelerator.get_memory_info()
drv_used = drv_total - drv_free
drv_pct = drv_used / drv_total * 100
except Exception:
@@ -0,0 +1,331 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
P2PConnector proxy server for OffloadingConnector + TieringOffloadingSpec.
Unlike NixlConnector (which returns remote_host/remote_port in the prefill
response), OffloadingConnector does not embed connector coordinates in its
response. This proxy injects the prefiller's P2PConnector address into
kv_transfer_params before forwarding the decode request so the decoder knows
where to pull KV blocks from.
Usage:
.venv/bin/python p2p_connector_proxy.py \
--port 8192 \
--prefiller-host 127.0.0.1 --prefiller-port 8100 \
--decoder-host 127.0.0.1 --decoder-port 8200 \
--p2p-connector-host 127.0.0.1 --p2p-connector-port 7777
"""
import argparse
import asyncio
import itertools
import logging
import os
import uuid
from contextlib import asynccontextmanager
import httpx
from fastapi import FastAPI, Request
from fastapi.responses import StreamingResponse
logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG)
@asynccontextmanager
async def lifespan(app: FastAPI):
app.state.prefill_clients = []
app.state.decode_clients = []
for i, (host, port) in enumerate(global_args.prefiller_instances):
app.state.prefill_clients.append(
{
"client": httpx.AsyncClient(
timeout=None,
base_url=f"http://{host}:{port}/v1",
limits=httpx.Limits(
max_connections=None, max_keepalive_connections=None
),
),
"host": host,
"port": port,
"id": i,
}
)
for i, (host, port) in enumerate(global_args.decoder_instances):
app.state.decode_clients.append(
{
"client": httpx.AsyncClient(
timeout=None,
base_url=f"http://{host}:{port}/v1",
limits=httpx.Limits(
max_connections=None, max_keepalive_connections=None
),
),
"host": host,
"port": port,
"id": i,
}
)
app.state.prefill_iterator = itertools.cycle(range(len(app.state.prefill_clients)))
app.state.decode_iterator = itertools.cycle(range(len(app.state.decode_clients)))
mode = "decoder-first" if global_args.decoder_first else "prefiller-first"
pd_host = global_args.p2p_connector_host
pd_port = global_args.p2p_connector_port
print(
f"Proxy ready [{mode}]: "
f"{len(app.state.prefill_clients)} prefiller(s), "
f"{len(app.state.decode_clients)} decoder(s). "
f"P2PConnector at {pd_host}:{pd_port}"
)
yield
for ci in app.state.prefill_clients:
await ci["client"].aclose()
for ci in app.state.decode_clients:
await ci["client"].aclose()
app = FastAPI(lifespan=lifespan)
def parse_args():
p = argparse.ArgumentParser()
p.add_argument("--port", type=int, default=8192)
p.add_argument("--host", type=str, default="127.0.0.1")
p.add_argument("--prefiller-hosts", type=str, nargs="+", default=["127.0.0.1"])
p.add_argument("--prefiller-ports", type=int, nargs="+", default=[8100])
p.add_argument("--decoder-hosts", type=str, nargs="+", default=["127.0.0.1"])
p.add_argument("--decoder-ports", type=int, nargs="+", default=[8200])
# P2PConnector coordinates of the prefiller — injected into decode requests.
p.add_argument(
"--p2p-connector-host",
type=str,
default="127.0.0.1",
help="Host of the prefiller's P2PConnector ZMQ socket",
)
p.add_argument(
"--p2p-connector-port",
type=int,
default=7777,
help="Port of the prefiller's P2PConnector ZMQ socket",
)
# P2PConnector coordinates of the decoder — injected into prefill requests
# so the prefiller's submit_store can resolve the peer to push KV to.
p.add_argument(
"--decoder-p2p-connector-host",
type=str,
default="127.0.0.1",
help="Host of the decoder's P2PConnector ZMQ socket",
)
p.add_argument(
"--decoder-p2p-connector-port",
type=int,
default=7778,
help="Port of the decoder's P2PConnector ZMQ socket",
)
p.add_argument(
"--decoder-first",
action="store_true",
help="Send decode request before prefill so decoder is already "
"waiting when KV blocks arrive (decoder-first mode)",
)
args = p.parse_args()
if len(args.prefiller_hosts) != len(args.prefiller_ports):
raise ValueError("Prefiller host/port count mismatch")
if len(args.decoder_hosts) != len(args.decoder_ports):
raise ValueError("Decoder host/port count mismatch")
args.prefiller_instances = list(zip(args.prefiller_hosts, args.prefiller_ports))
args.decoder_instances = list(zip(args.decoder_hosts, args.decoder_ports))
return args
def _get_next(app, service: str):
if service == "prefill":
return app.state.prefill_clients[next(app.state.prefill_iterator)]
return app.state.decode_clients[next(app.state.decode_iterator)]
def _auth_headers(request_id: str) -> dict:
headers: dict = {"X-Request-Id": request_id}
api_key = os.environ.get("OPENAI_API_KEY", "")
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
return headers
async def _prefill(client_info, endpoint, req_data, request_id):
"""Send a prefill-only request (max_tokens=1) to the prefiller."""
data = req_data.copy()
data["kv_transfer_params"] = {
"decode": {
"kv_request_id": request_id,
},
}
data["stream"] = False
data["max_tokens"] = 1
data.pop("max_completion_tokens", None)
data.pop("stream_options", None)
data.pop("min_tokens", None)
data.pop("min_completion_tokens", None)
headers = _auth_headers(request_id)
resp = await client_info["client"].post(endpoint, json=data, headers=headers)
resp.raise_for_status()
await resp.aread()
return resp
async def _stream_decode(client_info, endpoint, req_data, request_id):
headers = _auth_headers(request_id)
async with client_info["client"].stream(
"POST", endpoint, json=req_data, headers=headers
) as resp:
resp.raise_for_status()
async for chunk in resp.aiter_bytes():
yield chunk
async def _handle_completions(api: str, request: Request):
try:
req_data = await request.json()
request_id = str(uuid.uuid4())
prefill_client = _get_next(request.app, "prefill")
await _prefill(prefill_client, api, req_data, request_id)
# Inject the prefiller's P2PConnector address so the decoder can pull
# KV blocks from it via the P2PConnector transport.
req_data["kv_transfer_params"] = {
"prefill": {
"kv_request_id": request_id,
"remote_host": global_args.p2p_connector_host,
"remote_port": global_args.p2p_connector_port,
},
}
decode_client = _get_next(request.app, "decode")
logger.debug("prefill=%s decode=%s", prefill_client, decode_client)
async def generate():
async for chunk in _stream_decode(decode_client, api, req_data, request_id):
yield chunk
return StreamingResponse(generate(), media_type="application/json")
except Exception as e:
import sys
import traceback
print(f"Proxy error on {api}: {e}")
print("".join(traceback.format_exception(*sys.exc_info())))
raise
async def _handle_completions_decoder_first(api: str, request: Request):
"""Decoder-first mode: send decode request before prefill.
The decoder establishes its request and starts polling for KV blocks
immediately. The prefill is then sent so the prefiller computes and
pushes blocks to the already-waiting decoder.
"""
try:
req_data = await request.json()
request_id = str(uuid.uuid4())
prefill_client = _get_next(request.app, "prefill")
decode_client = _get_next(request.app, "decode")
decode_data = req_data.copy()
decode_data["kv_transfer_params"] = {
"prefill": {
"kv_request_id": request_id,
"remote_host": global_args.p2p_connector_host,
"remote_port": global_args.p2p_connector_port,
},
}
async def generate():
queue: asyncio.Queue = asyncio.Queue()
async def _run_decode():
try:
async for chunk in _stream_decode(
decode_client, api, decode_data, request_id
):
await queue.put(("data", chunk))
except Exception as exc:
await queue.put(("error", exc))
finally:
await queue.put(("done", None))
# 1. Start decode request — decoder is now waiting for KV blocks
asyncio.create_task(_run_decode())
# 2. Send prefill — blocks are computed and pushed to the decoder
try:
await _prefill(prefill_client, api, req_data, request_id)
except Exception as exc:
logger.warning("decoder-first: prefill failed: %s", exc)
logger.debug(
"decoder-first: prefill done, streaming decode prefill=%s decode=%s",
prefill_client,
decode_client,
)
# 3. Stream the decode response
while True:
kind, value = await queue.get()
if kind == "done":
break
if kind == "error":
raise value # type: ignore[misc]
yield value
return StreamingResponse(generate(), media_type="application/json")
except Exception as e:
import sys
import traceback
print(f"Proxy error on {api}: {e}")
print("".join(traceback.format_exception(*sys.exc_info())))
raise
def _route_handler(api: str):
if global_args.decoder_first:
return lambda req: _handle_completions_decoder_first(api, req)
return lambda req: _handle_completions(api, req)
@app.post("/v1/completions")
async def completions(request: Request):
return await _route_handler("/completions")(request)
@app.post("/v1/chat/completions")
async def chat_completions(request: Request):
return await _route_handler("/chat/completions")(request)
@app.get("/healthcheck")
async def healthcheck():
return {
"status": "ok",
"prefill_instances": len(app.state.prefill_clients),
"decode_instances": len(app.state.decode_clients),
}
if __name__ == "__main__":
global global_args
global_args = parse_args()
import uvicorn
uvicorn.run(app, host=global_args.host, port=global_args.port)
+312
View File
@@ -0,0 +1,312 @@
#!/bin/bash
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#
# Accuracy test driver for the p2p connector
# (OffloadingConnector + TieringOffloadingSpec + p2p tier).
#
# Mirrors tests/v1/kv_connector/nixl_integration/run_accuracy_test.sh:
# brings up N prefillers + M decoders on the local host, fronts them
# with p2p_connector_proxy.py, then runs the connector-agnostic
# test_accuracy.py (gsm8k via lm_eval) against the proxy.
#
# Knobs (env vars unless flagged otherwise):
# MODEL_NAMES space-separated model list (default: Llama-3.2-1B-Instruct)
# NUM_PREFILL_INSTANCES default 1
# NUM_DECODE_INSTANCES default 1
# PREFILLER_TP_SIZE default 1
# DECODER_TP_SIZE default 1
# GPU_MEMORY_UTILIZATION default 0.45
# MAX_MODEL_LEN default 512
# PREFILL_BLOCK_SIZE default 128
# DECODE_BLOCK_SIZE default 128
# CPU_BYTES default 209715200 (200 MB)
# VLLM_SERVE_EXTRA_ARGS comma-separated extra args for vllm serve
# --decoder-first toggle decoder-first proxy mode
#
# Examples:
# bash tests/v1/kv_offload/tiering/p2p/run_accuracy_test.sh
# NUM_PREFILL_INSTANCES=2 NUM_DECODE_INSTANCES=2 \
# bash tests/v1/kv_offload/tiering/p2p/run_accuracy_test.sh
# bash tests/v1/kv_offload/tiering/p2p/run_accuracy_test.sh --decoder-first
set -xe
# ---------------------------------------------------------------------------
# Args
# ---------------------------------------------------------------------------
DECODER_FIRST="false"
while [[ $# -gt 0 ]]; do
case $1 in
--decoder-first)
DECODER_FIRST="true"
shift 1
;;
*)
echo "Unknown option $1"
echo "Usage: $0 [--decoder-first]"
exit 1
;;
esac
done
# ---------------------------------------------------------------------------
# Models
# ---------------------------------------------------------------------------
MODEL_NAMES=${MODEL_NAMES:-}
if [[ -n "$MODEL_NAMES" ]]; then
# shellcheck disable=SC2206
MODELS=($MODEL_NAMES)
else
MODELS=(
"meta-llama/Llama-3.2-1B-Instruct"
)
fi
# ---------------------------------------------------------------------------
# Defaults
# ---------------------------------------------------------------------------
NUM_PREFILL_INSTANCES=${NUM_PREFILL_INSTANCES:-1}
NUM_DECODE_INSTANCES=${NUM_DECODE_INSTANCES:-1}
PREFILLER_TP_SIZE=${PREFILLER_TP_SIZE:-1}
DECODER_TP_SIZE=${DECODER_TP_SIZE:-1}
GPU_MEMORY_UTILIZATION=${GPU_MEMORY_UTILIZATION:-0.45}
MAX_MODEL_LEN=${MAX_MODEL_LEN:-512}
PREFILL_BLOCK_SIZE=${PREFILL_BLOCK_SIZE:-128}
DECODE_BLOCK_SIZE=${DECODE_BLOCK_SIZE:-128}
CPU_BYTES=${CPU_BYTES:-209715200}
VLLM_SERVE_EXTRA_ARGS=${VLLM_SERVE_EXTRA_ARGS:-}
# Base ports — per-instance offsets layered on top.
PREFILL_HTTP_BASE=8100
DECODE_HTTP_BASE=8200
PREFILL_PD_BASE=7777
DECODE_PD_BASE=$((PREFILL_PD_BASE + NUM_PREFILL_INSTANCES))
PROXY_PORT=8192
P2P_HOST=127.0.0.1
# ---------------------------------------------------------------------------
# Resolve repo root + venv (works in .venv and /workspace/venv pods)
# ---------------------------------------------------------------------------
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../../.." && pwd -P)}"
if [[ -z "${VLLM_BIN:-}" ]]; then
if [[ -x "${GIT_ROOT}/.venv/bin/vllm" ]]; then
VLLM_BIN="${GIT_ROOT}/.venv/bin/vllm"
elif [[ -x "/workspace/venv/bin/vllm" ]]; then
VLLM_BIN="/workspace/venv/bin/vllm"
else
VLLM_BIN="$(command -v vllm)"
fi
fi
if [[ -z "${PYTHON_BIN:-}" ]]; then
if [[ -x "${GIT_ROOT}/.venv/bin/python" ]]; then
PYTHON_BIN="${GIT_ROOT}/.venv/bin/python"
elif [[ -x "/workspace/venv/bin/python" ]]; then
PYTHON_BIN="/workspace/venv/bin/python"
else
PYTHON_BIN="$(command -v python3 || command -v python)"
fi
fi
echo "Using vllm: ${VLLM_BIN}"
echo "Using python: ${PYTHON_BIN}"
SMI_BIN=$(command -v nvidia-smi || command -v rocm-smi || echo "")
# Trap SIGINT/SIGTERM/EXIT to kill background jobs.
trap 'kill $(jobs -pr) 2>/dev/null || true' SIGINT SIGTERM EXIT
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
wait_for_server() {
local port=$1
timeout 1200 bash -c "
until curl -s localhost:${port}/v1/completions > /dev/null; do
sleep 1
done" && return 0 || return 1
}
cleanup_instances() {
echo "Cleaning up any running vLLM / proxy instances..."
pkill -f "vllm serve" || true
pkill -f "p2p_connector_proxy.py" || true
sleep 2
}
get_num_gpus() {
if [[ "$SMI_BIN" == *"nvidia"* ]]; then
$SMI_BIN --query-gpu=name --format=csv,noheader | wc -l
elif [[ "$SMI_BIN" == *"rocm"* ]]; then
$SMI_BIN -l | grep -c GPU
else
echo "1"
fi
}
# Build the OffloadingConnector kv-transfer-config for a given PD port.
# Mirrors deploy_local.sh:131.
build_kv_config() {
local pd_port=$1
printf '{"kv_connector":"OffloadingConnector","kv_role":"kv_both","kv_connector_extra_config":{"spec_name":"TieringOffloadingSpec","cpu_bytes_to_use":%s,"secondary_tiers":[{"type":"p2p","host":"%s","port":%s}]}}' \
"${CPU_BYTES}" "${P2P_HOST}" "${pd_port}"
}
# ---------------------------------------------------------------------------
# Per-model run
# ---------------------------------------------------------------------------
run_tests_for_model() {
local model_name=$1
echo "================================"
echo "Testing model: $model_name"
echo " prefillers=${NUM_PREFILL_INSTANCES} (tp=${PREFILLER_TP_SIZE})"
echo " decoders=${NUM_DECODE_INSTANCES} (tp=${DECODER_TP_SIZE})"
echo " decoder_first=${DECODER_FIRST}"
echo "================================"
PREFILL_HOSTS=()
PREFILL_PORTS=()
PREFILL_PD_PORTS=()
DECODE_HOSTS=()
DECODE_PORTS=()
DECODE_PD_PORTS=()
local num_gpus
num_gpus=$(get_num_gpus)
local next_gpu=0
# ---- Prefillers ----
for i in $(seq 0 $((NUM_PREFILL_INSTANCES-1))); do
local gpu_id=$((i * PREFILLER_TP_SIZE % num_gpus))
local cuda_devs="${gpu_id}"
for (( j=1; j < PREFILLER_TP_SIZE; j++ )); do
cuda_devs="${cuda_devs},$(((gpu_id + j) % num_gpus))"
done
next_gpu=$(((gpu_id + PREFILLER_TP_SIZE) % num_gpus))
local http_port=$((PREFILL_HTTP_BASE + i))
local pd_port=$((PREFILL_PD_BASE + i))
local kv_cfg
kv_cfg=$(build_kv_config "${pd_port}")
echo "Prefiller $i: gpu=[${cuda_devs}] http=${http_port} pd=${pd_port}"
BASE_CMD="CUDA_VISIBLE_DEVICES=${cuda_devs} \
PYTHONHASHSEED=42 \
${VLLM_BIN} serve ${model_name} \
--port ${http_port} \
--enforce-eager \
--block-size ${PREFILL_BLOCK_SIZE} \
--gpu-memory-utilization ${GPU_MEMORY_UTILIZATION} \
--max-model-len ${MAX_MODEL_LEN} \
--tensor-parallel-size ${PREFILLER_TP_SIZE} \
--kv-transfer-config '${kv_cfg}'"
if [[ -n "$VLLM_SERVE_EXTRA_ARGS" ]]; then
IFS=',' read -r -a extra_args <<< "$VLLM_SERVE_EXTRA_ARGS"
for arg in "${extra_args[@]}"; do
BASE_CMD="${BASE_CMD} $arg"
done
fi
eval "${BASE_CMD} &"
PREFILL_HOSTS+=("${P2P_HOST}")
PREFILL_PORTS+=("${http_port}")
PREFILL_PD_PORTS+=("${pd_port}")
done
# ---- Decoders ----
for i in $(seq 0 $((NUM_DECODE_INSTANCES-1))); do
local gpu_id=$(((next_gpu + i * DECODER_TP_SIZE) % num_gpus))
local cuda_devs="${gpu_id}"
for (( j=1; j < DECODER_TP_SIZE; j++ )); do
cuda_devs="${cuda_devs},$(((gpu_id + j) % num_gpus))"
done
local http_port=$((DECODE_HTTP_BASE + i))
local pd_port=$((DECODE_PD_BASE + i))
local kv_cfg
kv_cfg=$(build_kv_config "${pd_port}")
echo "Decoder $i: gpu=[${cuda_devs}] http=${http_port} pd=${pd_port}"
BASE_CMD="CUDA_VISIBLE_DEVICES=${cuda_devs} \
PYTHONHASHSEED=42 \
${VLLM_BIN} serve ${model_name} \
--port ${http_port} \
--enforce-eager \
--block-size ${DECODE_BLOCK_SIZE} \
--gpu-memory-utilization ${GPU_MEMORY_UTILIZATION} \
--max-model-len ${MAX_MODEL_LEN} \
--tensor-parallel-size ${DECODER_TP_SIZE} \
--kv-transfer-config '${kv_cfg}'"
if [[ -n "$VLLM_SERVE_EXTRA_ARGS" ]]; then
IFS=',' read -r -a extra_args <<< "$VLLM_SERVE_EXTRA_ARGS"
for arg in "${extra_args[@]}"; do
BASE_CMD="${BASE_CMD} $arg"
done
fi
eval "${BASE_CMD} &"
DECODE_HOSTS+=("${P2P_HOST}")
DECODE_PORTS+=("${http_port}")
DECODE_PD_PORTS+=("${pd_port}")
done
# ---- Wait for HTTP readiness ----
for port in "${PREFILL_PORTS[@]}"; do
echo "Waiting for prefill instance on port $port to start..."
wait_for_server "$port"
done
for port in "${DECODE_PORTS[@]}"; do
echo "Waiting for decode instance on port $port to start..."
wait_for_server "$port"
done
# ---- Proxy ----
# The proxy currently advertises a single prefiller PD address to decoders.
# For the 1xM and matched NxM common cases the first prefiller's PD coords
# are the right pick; multi-prefiller PD round-robin is a follow-up.
PROXY_CMD="${PYTHON_BIN} ${SCRIPT_DIR}/p2p_connector_proxy.py \
--port ${PROXY_PORT} \
--host ${P2P_HOST} \
--prefiller-hosts ${PREFILL_HOSTS[*]} \
--prefiller-ports ${PREFILL_PORTS[*]} \
--decoder-hosts ${DECODE_HOSTS[*]} \
--decoder-ports ${DECODE_PORTS[*]} \
--p2p-connector-host ${P2P_HOST} \
--p2p-connector-port ${PREFILL_PD_PORTS[0]} \
--decoder-p2p-connector-host ${P2P_HOST} \
--decoder-p2p-connector-port ${DECODE_PD_PORTS[0]}"
if [[ "${DECODER_FIRST}" == "true" ]]; then
PROXY_CMD="${PROXY_CMD} --decoder-first"
fi
echo "Starting proxy: ${PROXY_CMD}"
eval "${PROXY_CMD} &"
sleep 5
# ---- Run accuracy test (reused from nixl_integration) ----
echo "Running tests for $model_name"
TEST_MODEL=$model_name "${PYTHON_BIN}" -m pytest -s -x \
"${GIT_ROOT}/tests/v1/kv_connector/nixl_integration/test_accuracy.py"
cleanup_instances
sleep 3
}
# ---------------------------------------------------------------------------
# Drive
# ---------------------------------------------------------------------------
for model in "${MODELS[@]}"; do
run_tests_for_model "$model"
done
echo "All tests completed!"
@@ -0,0 +1,330 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for DataTransport base class and NixlTransport."""
from __future__ import annotations
import ctypes
from unittest.mock import MagicMock, patch
import numpy as np
from vllm.v1.kv_offload.tiering.p2p.data.base import PollResult
from vllm.v1.kv_offload.tiering.p2p.data.nixl import NixlTransport
# ---------------------------------------------------------------------------
# DataTransport base class tests
# ---------------------------------------------------------------------------
class TestDataTransportBase:
"""Tests for the DataTransport abstract base properties."""
def _make_view(self, num_blocks: int = 8, block_len: int = 1024) -> memoryview:
"""Create a memoryview with the given shape."""
buf = np.zeros((num_blocks, block_len), dtype=np.uint8)
return memoryview(buf)
def test_properties(self):
"""base_addr, num_blocks, block_len are set from memoryview shape."""
view = self._make_view(num_blocks=4, block_len=2048)
# Use NixlTransport (concrete) with NIXL mocked away
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
transport = NixlTransport("test:1", view)
assert transport.num_blocks == 4
assert transport.block_len == 2048
assert transport.base_addr == ctypes.addressof(ctypes.c_char.from_buffer(view))
def test_config_fingerprint_empty_when_no_fields(self):
"""No config fields → empty fingerprint."""
view = self._make_view()
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
transport = NixlTransport("test:1", view, config_fields=None)
assert transport.config_fingerprint == ""
def test_config_fingerprint_deterministic(self):
"""Same config fields → same fingerprint."""
view = self._make_view()
fields = {"model": "llama", "dtype": "float16", "block_size_factor": 1}
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
t1 = NixlTransport("test:1", view, config_fields=fields)
t2 = NixlTransport("test:2", view, config_fields=fields)
assert t1.config_fingerprint == t2.config_fingerprint
assert len(t1.config_fingerprint) == 16
def test_config_fingerprint_differs_for_different_fields(self):
"""Different config fields → different fingerprint."""
view = self._make_view()
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
t1 = NixlTransport("test:1", view, config_fields={"model": "a"})
t2 = NixlTransport("test:2", view, config_fields={"model": "b"})
assert t1.config_fingerprint != t2.config_fingerprint
# ---------------------------------------------------------------------------
# NixlTransport tests (with mocked NIXL agent)
# ---------------------------------------------------------------------------
class TestNixlTransportWithMockedAgent:
"""Tests for NixlTransport logic with a mocked NIXL agent."""
def _make_transport(self) -> NixlTransport:
"""Create a NixlTransport with mocked NIXL internals."""
view = memoryview(np.zeros((8, 1024), dtype=np.uint8))
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
transport = NixlTransport("test:1", view)
# Manually set up a mock agent after construction
agent = MagicMock()
agent.add_remote_agent.return_value = "nixl-peer-name"
agent.get_xfer_descs.return_value = MagicMock()
agent.prep_xfer_dlist.return_value = MagicMock()
agent.make_prepped_xfer.return_value = MagicMock(name="handle")
agent.transfer.return_value = None
agent.check_xfer_state.return_value = "PROC"
agent.get_agent_metadata.return_value = b"test-metadata"
transport._agent = agent
transport._local_dlist = MagicMock()
return transport
def test_available_false_without_nixl(self):
"""Without NIXL installed, available is False."""
view = memoryview(np.zeros((4, 512), dtype=np.uint8))
with patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", None):
transport = NixlTransport("test:1", view)
assert transport.available is False
def test_available_true_with_agent(self):
transport = self._make_transport()
assert transport.available is True
def test_get_agent_metadata(self):
transport = self._make_transport()
assert transport.get_agent_metadata() == b"test-metadata"
def test_write_blocks_returns_none_for_unknown_peer(self):
"""write_blocks returns None if peer not registered."""
transport = self._make_transport()
result = transport.write_blocks("unknown:1", [0, 1], [2, 3])
assert result is None
def test_write_blocks_returns_transfer_id(self):
"""write_blocks returns an integer transfer_id on success."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
tid = transport.write_blocks("peer:1", [0, 1], [2, 3])
assert tid is not None
assert isinstance(tid, int)
def test_write_blocks_increments_transfer_id(self):
"""Each write_blocks call gets a unique transfer_id."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
tid1 = transport.write_blocks("peer:1", [0], [1])
tid2 = transport.write_blocks("peer:1", [2], [3])
assert tid1 != tid2
def test_poll_empty_when_no_inflight(self):
"""poll returns empty when nothing is inflight."""
transport = self._make_transport()
result = transport.poll()
assert result == PollResult(done=(), failed=())
def test_poll_returns_done_when_transfer_completes(self):
"""Completed transfer appears in poll().done."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
tid = transport.write_blocks("peer:1", [0], [1])
# Simulate completion
transport._agent.check_xfer_state.return_value = "DONE"
result = transport.poll()
assert tid in result.done
assert result.failed == ()
# Handle released
transport._agent.release_xfer_handle.assert_called()
def test_poll_returns_failed_for_error_state(self):
"""Transfer in error state appears in poll().failed."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
tid = transport.write_blocks("peer:1", [0], [1])
transport._agent.check_xfer_state.return_value = "ERR"
result = transport.poll()
assert result.done == ()
assert tid in result.failed
def test_poll_ignores_in_progress(self):
"""Transfers in PROC/PEND state stay inflight."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
transport.write_blocks("peer:1", [0], [1])
transport._agent.check_xfer_state.return_value = "PROC"
result = transport.poll()
assert result.done == ()
assert result.failed == ()
transport._agent.check_xfer_state.return_value = "PEND"
result = transport.poll()
assert result.done == ()
assert result.failed == ()
def test_cancel_removes_inflight(self):
"""cancel removes transfers and releases handles."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
tid = transport.write_blocks("peer:1", [0], [1])
assert tid in transport._inflight
result = transport.cancel([tid])
assert result == []
assert tid not in transport._inflight
transport._agent.release_xfer_handle.assert_called()
def test_cancel_ignores_unknown_ids(self):
"""cancel with unknown IDs doesn't crash."""
transport = self._make_transport()
assert transport.cancel([999, 1000]) == []
assert transport.cancel([999, 1000], mode="wait") == []
def test_cancel_wait_release_succeeds(self):
"""wait-mode cancel that succeeds pops the entry and returns []."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
tid = transport.write_blocks("peer:1", [0], [1])
assert tid in transport._inflight
result = transport.cancel([tid], mode="wait")
assert result == []
assert tid not in transport._inflight
transport._agent.release_xfer_handle.assert_called_once()
def test_cancel_wait_release_raises(self):
"""wait-mode cancel keeps the entry and returns the tid on raise."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
tid = transport.write_blocks("peer:1", [0], [1])
transport._agent.release_xfer_handle.side_effect = RuntimeError(
"NIXL_ERR_REPOST_ACTIVE"
)
result = transport.cancel([tid], mode="wait")
assert result == [tid]
assert tid in transport._inflight
def test_cancel_wait_then_poll_completes(self):
"""A wait-cancel that left a tid pending later completes via poll."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
tid = transport.write_blocks("peer:1", [0], [1])
transport._agent.release_xfer_handle.side_effect = RuntimeError("busy")
assert transport.cancel([tid], mode="wait") == [tid]
assert tid in transport._inflight
transport._agent.release_xfer_handle.side_effect = None
transport._agent.check_xfer_state.return_value = "DONE"
result = transport.poll()
assert tid in result.done
assert tid not in transport._inflight
def test_add_and_remove_remote_peer(self):
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
assert "peer:1" in transport._remote_dlists
transport.remove_remote_peer("peer:1")
assert "peer:1" not in transport._remote_dlists
transport._agent.release_dlist_handle.assert_called()
transport._agent.remove_remote_agent.assert_called()
def test_close_releases_everything(self):
"""close releases all handles and clears state."""
transport = self._make_transport()
transport.add_remote_peer("peer:1", b"meta", 0x1000, 8, 1024)
transport.write_blocks("peer:1", [0], [1])
transport.close()
assert transport._agent is None
assert transport._inflight == {}
assert transport._remote_dlists == {}
# ---------------------------------------------------------------------------
# NIXL agent-config selection
# ---------------------------------------------------------------------------
class TestNixlAgentConfigSelection:
"""Tests that backends/num_threads pick the right nixl_agent_config call.
Mirrors the conditional in
vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py:325-329.
"""
def _make_view(self) -> memoryview:
return memoryview(np.zeros((4, 512), dtype=np.uint8))
def test_non_ucx_backends_passes_backends_kwarg(self):
"""When any non-UCX backend is requested, pass backends + telemetry."""
agent_cls = MagicMock()
config_fn = MagicMock(return_value=MagicMock(name="cfg"))
with (
patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", agent_cls),
patch(
"vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgentConfig", config_fn
),
):
NixlTransport("test:1", self._make_view(), backends=["MOONCAKE"])
config_fn.assert_called_once_with(backends=["MOONCAKE"], capture_telemetry=True)
# num_threads must NOT be passed on the non-UCX branch.
assert "num_threads" not in config_fn.call_args.kwargs
def test_ucx_only_passes_num_threads(self):
"""UCX-only configuration passes num_threads + telemetry, no backends."""
agent_cls = MagicMock()
config_fn = MagicMock(return_value=MagicMock(name="cfg"))
with (
patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", agent_cls),
patch(
"vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgentConfig", config_fn
),
):
NixlTransport("test:1", self._make_view(), num_threads=8)
config_fn.assert_called_once_with(num_threads=8, capture_telemetry=True)
assert "backends" not in config_fn.call_args.kwargs
def test_default_backends_is_ucx_only(self):
"""No backends arg → defaults to UCX-only branch."""
agent_cls = MagicMock()
config_fn = MagicMock(return_value=MagicMock(name="cfg"))
with (
patch("vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgent", agent_cls),
patch(
"vllm.v1.kv_offload.tiering.p2p.data.nixl._NixlAgentConfig", config_fn
),
):
NixlTransport("test:1", self._make_view())
# Default num_threads=4, no backends kwarg.
config_fn.assert_called_once_with(num_threads=4, capture_telemetry=True)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,230 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for vllm.v1.kv_offload.tiering.p2p.control.zmq."""
from __future__ import annotations
import socket
import time
import pytest
import zmq
from vllm.v1.kv_offload.tiering.p2p.control.zmq import (
ZmqConnection,
ZmqTransport,
_Sockets,
)
def _free_port() -> int:
"""Find a free TCP port."""
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("127.0.0.1", 0))
return s.getsockname()[1]
def _make_transport(host: str = "127.0.0.1", attempts: int = 8):
"""Construct a ZmqTransport on a fresh port, retrying on bind collisions.
Why: _free_port() releases the probe socket before ZmqTransport binds the
same port a parallel test run can steal it in between. Retrying on
ZMQError/OSError closes that race without a production change.
"""
last_err: Exception | None = None
for _ in range(attempts):
port = _free_port()
try:
return ZmqTransport(f"{host}:{port}", host, port), port
except (zmq.ZMQError, OSError) as e:
last_err = e
assert last_err is not None
raise last_err
def _wait_for_inbound(transport: ZmqTransport, deadline: float = 2.0):
"""Poll until at least one new inbound connection is accepted, or fail."""
end = time.monotonic() + deadline
while time.monotonic() < end:
new = transport.poll()
if new:
return new
time.sleep(0.005)
raise AssertionError(f"no inbound connection within {deadline}s")
def _wait_for_messages(
transport: ZmqTransport,
conn: ZmqConnection,
n: int,
deadline: float = 2.0,
) -> list[dict]:
"""Poll until `conn` has received at least `n` messages, then return them."""
end = time.monotonic() + deadline
msgs: list[dict] = []
while time.monotonic() < end:
transport.poll()
msgs.extend(conn.recv())
if len(msgs) >= n:
return msgs
time.sleep(0.005)
raise AssertionError(f"got {len(msgs)}/{n} messages within {deadline}s")
def _make_mock_connection(peer_id: str = "test:1234") -> ZmqConnection:
"""Create a ZmqConnection with mock sockets for unit testing."""
from unittest.mock import MagicMock
sockets = _Sockets(dealer=MagicMock(), monitor=MagicMock())
return ZmqConnection(peer_id, sockets)
class TestZmqConnection:
"""Tests for ZmqConnection in isolation (no real sockets)."""
def test_enqueue_and_recv(self):
"""Messages enqueued are returned by recv() in order."""
conn = _make_mock_connection()
conn.enqueue({"type": "a"})
conn.enqueue({"type": "b"})
msgs = conn.recv()
assert list(msgs) == [{"type": "a"}, {"type": "b"}]
# Second recv is empty
assert not conn.recv()
def test_recv_returns_empty_initially(self):
conn = _make_mock_connection()
assert not conn.recv()
def test_alive_initially_true(self):
conn = _make_mock_connection()
assert conn.alive is True
def test_mark_dead(self):
conn = _make_mock_connection()
conn.mark_dead()
assert conn.alive is False
def test_send_raises_when_closed(self):
conn = _make_mock_connection()
conn.mark_dead()
with pytest.raises(RuntimeError, match="closed connection"):
conn.send({"type": "test"})
class TestZmqTransportConnectivity:
"""Integration tests for ZmqTransport with real ZMQ sockets."""
def test_connect_and_send_message(self):
"""Two transports can connect and exchange messages."""
transport_a, port_a = _make_transport()
transport_b, port_b = _make_transport()
try:
peer_a_id = f"127.0.0.1:{port_a}"
conn_b_to_a = transport_b.connect(peer_a_id)
conn_b_to_a.send({"type": "hello", "data": 42})
new_conns = _wait_for_inbound(transport_a)
assert len(new_conns) == 1
conn_a_from_b = new_conns[0]
assert conn_a_from_b.peer_id == f"127.0.0.1:{port_b}"
msgs = _wait_for_messages(transport_a, conn_a_from_b, 1)
assert msgs == [{"type": "hello", "data": 42}]
finally:
transport_a.close()
transport_b.close()
def test_bidirectional_messaging(self):
"""Both sides can send and receive after connection."""
transport_a, port_a = _make_transport()
transport_b, _ = _make_transport()
try:
conn_b = transport_b.connect(f"127.0.0.1:{port_a}")
conn_b.send({"type": "connect", "from": "b"})
new_conns = _wait_for_inbound(transport_a)
assert len(new_conns) == 1
conn_a = new_conns[0]
conn_a.send({"type": "reply", "from": "a"})
msgs = _wait_for_messages(transport_b, conn_b, 1)
assert msgs == [{"type": "reply", "from": "a"}]
finally:
transport_a.close()
transport_b.close()
def test_poll_returns_empty_when_no_connections(self):
transport, _ = _make_transport()
try:
assert not transport.poll()
finally:
transport.close()
def test_multiple_messages(self):
"""Multiple messages are buffered and returned together."""
transport_a, port_a = _make_transport()
transport_b, _ = _make_transport()
try:
conn_b = transport_b.connect(f"127.0.0.1:{port_a}")
conn_b.send({"seq": 1})
conn_b.send({"seq": 2})
conn_b.send({"seq": 3})
new_conns = _wait_for_inbound(transport_a)
assert len(new_conns) == 1
conn_a = new_conns[0]
msgs = _wait_for_messages(transport_a, conn_a, 3)
assert [m["seq"] for m in msgs] == [1, 2, 3]
finally:
transport_a.close()
transport_b.close()
def test_duplicate_connect_asserts(self):
"""Connecting to the same peer twice raises AssertionError."""
# port_a is never bound — we just need a syntactically-valid peer id.
port_a = _free_port()
transport_b, _ = _make_transport()
try:
transport_b.connect(f"127.0.0.1:{port_a}")
with pytest.raises(AssertionError, match="already exists"):
transport_b.connect(f"127.0.0.1:{port_a}")
finally:
transport_b.close()
def test_dead_connection_removed_on_poll(self):
"""Dead connections are cleaned up during poll."""
transport_a, port_a = _make_transport()
transport_b, _ = _make_transport()
try:
conn_b = transport_b.connect(f"127.0.0.1:{port_a}")
conn_b.send({"type": "hello"})
new_conns = _wait_for_inbound(transport_a)
assert len(new_conns) == 1
# Mark the inbound connection dead manually.
new_conns[0].mark_dead()
# Pruning is synchronous within poll().
transport_a.poll()
assert len(transport_a._connections) == 0
finally:
transport_a.close()
transport_b.close()
def test_close_is_idempotent(self):
"""Calling close() twice doesn't raise."""
transport, _ = _make_transport()
transport.close()
transport.close() # should not raise
@@ -1257,3 +1257,127 @@ def test_thinking_budget_long_thinking_section_end_marker_found_at_correct_index
assert h._state[0]["start_thinking"] == 0
assert h._state[0]["end_thinking"] == expected_end_idx
# --- Thinking budget re-entry tests (issue #43708) ---
# Regression tests: after budget forces end-of-thinking token sequence,
# the state machine must detect and enforce budget on subsequent blocks.
class TestThinkingBudgetReentry:
THINK_START = 100
THINK_END_SINGLE = [200]
THINK_END_MULTI = [200, 201, 202]
BUDGET = 5
CONTENT_TOKEN = 50
THINK_TOKEN = 60
@staticmethod
def _make_holder(end_token_ids: list[int]) -> ThinkingBudgetStateHolder:
class FakeReasoningConfig:
reasoning_start_token_ids = [TestThinkingBudgetReentry.THINK_START]
reasoning_end_token_ids: list[int] = []
enabled = True
cfg = FakeReasoningConfig()
cfg.reasoning_end_token_ids = end_token_ids
return ThinkingBudgetStateHolder(
reasoning_config=cfg,
max_num_seqs=8,
num_spec_tokens=0,
device=torch.device("cpu"),
is_pin_memory=False,
)
@staticmethod
def _sync_batch(holder: ThinkingBudgetStateHolder, budget: int) -> None:
holder.sync_batch(
BatchUpdate(
batch_size=1,
removed=(),
added=[(0, SamplingParams(thinking_token_budget=budget), None, [])],
moved=(),
)
)
@staticmethod
def _step(holder: ThinkingBudgetStateHolder, output_tok_ids: list[int]) -> None:
holder.update_state(
output_token_ids=[output_tok_ids],
spec_token_ids=None,
repeat_indices=None,
)
def _exhaust_budget(self, holder: ThinkingBudgetStateHolder) -> list[int]:
output = [self.THINK_START]
self._step(holder, list(output))
for _ in range(self.BUDGET):
output.append(self.THINK_TOKEN)
self._step(holder, list(output))
assert holder._state[0]["in_end"]
return output
def _accept_end_tokens(
self,
holder: ThinkingBudgetStateHolder,
output: list[int],
end_token_ids: list[int],
) -> None:
for tok in end_token_ids:
output.append(tok)
self._step(holder, list(output))
def test_single_token_end_reentry(self):
holder = self._make_holder(self.THINK_END_SINGLE)
self._sync_batch(holder, self.BUDGET)
output = self._exhaust_budget(holder)
self._accept_end_tokens(holder, output, self.THINK_END_SINGLE)
for _ in range(3):
output.append(self.CONTENT_TOKEN)
self._step(holder, list(output))
output.append(self.THINK_START)
self._step(holder, list(output))
for _ in range(self.BUDGET):
output.append(self.THINK_TOKEN)
self._step(holder, list(output))
assert holder._state[0]["in_end"], (
"Second thinking block must also be budget-enforced"
)
def test_multi_token_end_reentry(self):
holder = self._make_holder(self.THINK_END_MULTI)
self._sync_batch(holder, self.BUDGET)
output = self._exhaust_budget(holder)
self._accept_end_tokens(holder, output, self.THINK_END_MULTI)
assert not holder._state[0]["in_end"]
output.append(self.THINK_START)
self._step(holder, list(output))
for _ in range(self.BUDGET):
output.append(self.THINK_TOKEN)
self._step(holder, list(output))
assert holder._state[0]["in_end"], (
"Immediate re-entry after multi-token end must be enforced"
)
def test_single_block_not_broken(self):
holder = self._make_holder(self.THINK_END_SINGLE)
self._sync_batch(holder, self.BUDGET)
output = self._exhaust_budget(holder)
self._accept_end_tokens(holder, output, self.THINK_END_SINGLE)
for _ in range(20):
output.append(self.CONTENT_TOKEN)
self._step(holder, list(output))
assert not holder._state[0]["in_end"]
assert not holder._state[0]["in_think"]
+1 -1
View File
@@ -1285,7 +1285,7 @@ def test_token_logprobs_large_batch_int64_row_offset():
batch_size = 2**31 // vocab_size + 64 # batch_size * vocab_size > 2**31
# logits (the large input) plus small logprob/rank outputs; ~1 GB headroom.
required_bytes = batch_size * vocab_size * 4 + (1 << 30)
if torch.cuda.mem_get_info()[0] < required_bytes:
if torch.accelerator.get_memory_info()[0] < required_bytes:
pytest.skip(f"needs ~{required_bytes / 1e9:.0f} GB of free GPU memory")
logits = torch.randn(batch_size, vocab_size, device=device, dtype=torch.float32)
+1 -1
View File
@@ -426,7 +426,7 @@ class TestTritonTopkTopp:
# logits is modified in place; the only extra device memory is the
# per-SM scratch buffer (~num_sm * vocab), so allow ~1 GB of headroom.
required_bytes = batch_size * vocab_size * 4 + (1 << 30)
if torch.cuda.mem_get_info()[0] < required_bytes:
if torch.accelerator.get_memory_info()[0] < required_bytes:
pytest.skip(f"needs ~{required_bytes / 1e9:.0f} GB of free GPU memory")
logits = torch.randn(
View File
+70 -51
View File
@@ -7,6 +7,8 @@ import time
from opentelemetry.sdk.environment_variables import OTEL_EXPORTER_OTLP_TRACES_INSECURE
from vllm import LLM, SamplingParams
from vllm.distributed import cleanup_dist_env_and_memory
from vllm.platforms import current_platform
from vllm.tracing import SpanAttributes
# Import shared fixtures from the tracing conftest
@@ -23,6 +25,11 @@ def test_traces(
):
with monkeypatch.context() as m:
m.setenv(OTEL_EXPORTER_OTLP_TRACES_INSECURE, "true")
if current_platform.is_rocm():
# The fake OTLP server starts gRPC worker threads before the engine
# core is launched. On ROCm CI, forking while those threads are
# active can segfault in gRPC during engine startup or teardown.
m.setenv("VLLM_WORKER_MULTIPROC_METHOD", "spawn")
sampling_params = SamplingParams(
temperature=0.01,
@@ -30,58 +37,70 @@ def test_traces(
max_tokens=256,
)
model = "facebook/opt-125m"
llm = LLM(
model=model,
otlp_traces_endpoint=FAKE_TRACE_SERVER_ADDRESS,
gpu_memory_utilization=0.3,
disable_log_stats=False,
)
prompts = ["This is a short prompt"]
outputs = llm.generate(prompts, sampling_params=sampling_params)
print(f"test_traces outputs is : {outputs}")
llm = None
try:
llm = LLM(
model=model,
otlp_traces_endpoint=FAKE_TRACE_SERVER_ADDRESS,
gpu_memory_utilization=0.3,
disable_log_stats=False,
)
prompts = ["This is a short prompt"]
outputs = llm.generate(prompts, sampling_params=sampling_params)
print(f"test_traces outputs is : {outputs}")
# Wait for the "llm_request" span to be exported.
# The BatchSpanProcessor batches spans and exports them periodically,
# so we need to wait specifically for the llm_request span to appear.
timeout = 15
deadline = time.time() + timeout
llm_request_spans = []
while time.time() < deadline:
all_spans = trace_service.get_all_spans()
llm_request_spans = [s for s in all_spans if s["name"] == "llm_request"]
if llm_request_spans:
break
time.sleep(0.5)
# Wait for the "llm_request" span to be exported.
# The BatchSpanProcessor batches spans and exports them periodically,
# so we need to wait specifically for the llm_request span to appear.
timeout = 15
deadline = time.time() + timeout
llm_request_spans = []
while time.time() < deadline:
all_spans = trace_service.get_all_spans()
llm_request_spans = [s for s in all_spans if s["name"] == "llm_request"]
if llm_request_spans:
break
time.sleep(0.5)
assert len(llm_request_spans) == 1, (
f"Expected exactly 1 'llm_request' span, but got {len(llm_request_spans)}. "
f"All span names: {[s['name'] for s in all_spans]}"
)
assert len(llm_request_spans) == 1, (
f"Expected exactly 1 'llm_request' span, but got "
f"{len(llm_request_spans)}. "
f"All span names: {[s['name'] for s in all_spans]}"
)
attributes = llm_request_spans[0]["attributes"]
# assert attributes.get(SpanAttributes.GEN_AI_RESPONSE_MODEL) == model
assert attributes.get(SpanAttributes.GEN_AI_REQUEST_ID) == outputs[0].request_id
assert (
attributes.get(SpanAttributes.GEN_AI_REQUEST_TEMPERATURE)
== sampling_params.temperature
)
assert (
attributes.get(SpanAttributes.GEN_AI_REQUEST_TOP_P) == sampling_params.top_p
)
assert (
attributes.get(SpanAttributes.GEN_AI_REQUEST_MAX_TOKENS)
== sampling_params.max_tokens
)
assert attributes.get(SpanAttributes.GEN_AI_REQUEST_N) == sampling_params.n
assert attributes.get(SpanAttributes.GEN_AI_USAGE_PROMPT_TOKENS) == len(
outputs[0].prompt_token_ids
)
completion_tokens = sum(len(o.token_ids) for o in outputs[0].outputs)
assert (
attributes.get(SpanAttributes.GEN_AI_USAGE_COMPLETION_TOKENS)
== completion_tokens
)
attributes = llm_request_spans[0]["attributes"]
# assert attributes.get(SpanAttributes.GEN_AI_RESPONSE_MODEL) == model
assert (
attributes.get(SpanAttributes.GEN_AI_REQUEST_ID)
== outputs[0].request_id
)
assert (
attributes.get(SpanAttributes.GEN_AI_REQUEST_TEMPERATURE)
== sampling_params.temperature
)
assert (
attributes.get(SpanAttributes.GEN_AI_REQUEST_TOP_P)
== sampling_params.top_p
)
assert (
attributes.get(SpanAttributes.GEN_AI_REQUEST_MAX_TOKENS)
== sampling_params.max_tokens
)
assert attributes.get(SpanAttributes.GEN_AI_REQUEST_N) == sampling_params.n
assert attributes.get(SpanAttributes.GEN_AI_USAGE_PROMPT_TOKENS) == len(
outputs[0].prompt_token_ids
)
completion_tokens = sum(len(o.token_ids) for o in outputs[0].outputs)
assert (
attributes.get(SpanAttributes.GEN_AI_USAGE_COMPLETION_TOKENS)
== completion_tokens
)
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_TIME_IN_QUEUE) > 0
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_TIME_TO_FIRST_TOKEN) > 0
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_E2E) > 0
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_TIME_IN_QUEUE) > 0
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_TIME_TO_FIRST_TOKEN) > 0
assert attributes.get(SpanAttributes.GEN_AI_LATENCY_E2E) > 0
finally:
if llm is not None:
shutdown_timeout = 60.0 if current_platform.is_rocm() else 5.0
llm.llm_engine.engine_core.shutdown(timeout=shutdown_timeout)
cleanup_dist_env_and_memory()
+2 -1
View File
@@ -8,11 +8,12 @@ import regex as re
# Regex: match `torch.cuda.xxx` but allow `torch.accelerator.xxx`
# --------------------------------------------------------------------------- #
_TORCH_CUDA_PATTERNS = [
r"\btorch\.cuda\.(empty_cache|synchronize|device_count|current_device|memory_reserved|memory_allocated|max_memory_allocated|max_memory_reserved|reset_peak_memory_stats|memory_stats|set_device|device\()\b",
r"\btorch\.cuda\.(empty_cache|synchronize|device_count|current_device|memory_reserved|memory_allocated|max_memory_allocated|max_memory_reserved|reset_peak_memory_stats|memory_stats|mem_get_info|set_device|device\()\b",
r"\btorch\.cuda\.(manual_seed|manual_seed_all)\b",
r"\bwith\storch\.cuda\.device\b",
# Calls torch.cuda.{_is_compiled/_device_count_amdsmi/_device_count_nvml} internally
r"\bcuda_device_count_stateless\(\)\b",
r"\bcurrent_platform\.mem_get_info\(\)\b",
]
ALLOWED_FILES = {
+5 -2
View File
@@ -262,10 +262,12 @@ class PassConfig:
"Fusion enabled but reshape elimination disabled. "
"RMSNorm + padding fusion might not work"
)
if self.enable_qk_norm_rope_fusion and not current_platform.is_cuda_alike():
if self.enable_qk_norm_rope_fusion and not (
current_platform.is_cuda_alike() or current_platform.is_xpu()
):
logger.warning_once(
"QK Norm + RoPE fusion enabled but the current platform is not "
"CUDA or ROCm. The fusion will be disabled."
"CUDA, ROCm or XPU. The fusion will be disabled."
)
self.enable_qk_norm_rope_fusion = False
if self.fuse_act_padding and not current_platform.is_rocm():
@@ -757,6 +759,7 @@ class CompilationConfig:
"vllm::sparse_attn_indexer",
"vllm::rocm_aiter_sparse_attn_indexer",
"vllm::deepseek_v4_attention",
"vllm::hpc_rope_norm_forward",
]
def compute_hash(self) -> str:
+1 -1
View File
@@ -213,7 +213,7 @@ class ModelConfig:
flexibility."""
enable_return_routed_experts: bool = False
"""Whether to return routed experts."""
max_logprobs: int = 20
max_logprobs: int = Field(default=20, ge=-1)
"""Maximum number of log probabilities to return when `logprobs` is
specified in `SamplingParams`. The default value comes the default for the
OpenAI Chat Completions API. -1 means no cap, i.e. all (output_length *
+2 -2
View File
@@ -77,9 +77,9 @@ class SchedulerConfig:
this less than max_num_partial_prefills will allow shorter prompts to jump
the queue in front of longer prompts in some cases, improving latency."""
long_prefill_token_threshold: int = 0
long_prefill_token_threshold: int = Field(default=0, ge=0)
"""For chunked prefill, a request is considered long if the prompt is
longer than this number of tokens."""
longer than this number of tokens. 0 disables the cap (default)."""
enable_chunked_prefill: bool = True
"""If True, prefill requests can be chunked based
+19 -13
View File
@@ -931,16 +931,28 @@ class VllmConfig:
model_type,
)
from vllm.platforms import current_platform
from vllm.v1.executor.abstract import Executor
executor_backend = self.parallel_config.distributed_executor_backend
executor_class = Executor.get_class(self)
executor_supports_async_sched = executor_class.supports_async_scheduling()
uses_rocm_deepep_ht_dbo = (
current_platform.is_rocm()
and self.parallel_config.enable_dbo
and self.parallel_config.all2all_backend == "deepep_high_throughput"
)
if self.scheduler_config.async_scheduling:
# Async scheduling explicitly enabled, hard fail any incompatibilities.
# Currently, async scheduling only support eagle speculative
# decoding.
if uses_rocm_deepep_ht_dbo:
raise ValueError(
"Async scheduling is not compatible with ROCm DeepEP "
"high-throughput DBO. Please use --no-async-scheduling or "
"select a different all2all backend."
)
if self.speculative_config is not None:
if (
self.speculative_config.method not in get_args(EagleModelTypes)
@@ -1000,6 +1012,13 @@ class VllmConfig:
executor_backend,
)
self.scheduler_config.async_scheduling = False
elif uses_rocm_deepep_ht_dbo:
logger.warning_once(
"Async scheduling is disabled for ROCm DeepEP "
"high-throughput DBO because that combination can corrupt "
"DP+EP generation accuracy."
)
self.scheduler_config.async_scheduling = False
else:
self.scheduler_config.async_scheduling = True
@@ -1044,8 +1063,6 @@ class VllmConfig:
"VLLM_WORKER_MULTIPROC_METHOD set to spawn"
)
from vllm.platforms import current_platform
if (
self.model_config is not None
and self.scheduler_config.enable_chunked_prefill
@@ -1997,13 +2014,6 @@ class VllmConfig:
model_config = self.model_config
speculative_config = self.speculative_config
if (
model_config is not None
and model_config.has_inner_state
and self.cache_config.mamba_cache_mode == "align"
):
unsupported.append("hybrid/mamba models with align cache mode")
if self.parallel_config.prefill_context_parallel_size > 1:
unsupported.append("prefill context parallelism")
@@ -2152,10 +2162,6 @@ class VllmConfig:
"to schedule a multiple of block_size tokens even if they are "
"in the middle of a mm input"
)
# TODO: support align mamba cache mode for model runner v2
assert not envs.VLLM_USE_V2_MODEL_RUNNER, (
"Model Runner V2 has not yet supported mamba_cache_mode='align'. "
)
@model_validator(mode="after")
def validate_nvfp4_kv_cache_with_mla(self) -> "VllmConfig":
+35 -21
View File
@@ -95,40 +95,54 @@ def mma_bf16(
return cute.TensorSSA(vec, 4, Float32)
@dsl_user_op
def _bf16x2_abs(a: Uint32, *, loc=None, ip=None) -> Uint32:
def _bf16x2_unary(asm: str, a: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip)],
"abs.bf16x2 $0, $1;",
f"{asm}.bf16x2 $0, $1;",
"=r,r",
has_side_effects=False,
is_align_stack=False,
loc=loc,
ip=ip,
)
return Uint32(out)
def _bf16x2_binary(asm: str, a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
f"{asm}.bf16x2 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
loc=loc,
ip=ip,
)
return Uint32(out)
@dsl_user_op
def _bf16x2_abs(a: Uint32, *, loc=None, ip=None) -> Uint32:
return _bf16x2_unary("abs", a, loc=loc, ip=ip)
@dsl_user_op
def _bf16x2_neg(a: Uint32, *, loc=None, ip=None) -> Uint32:
return _bf16x2_unary("neg", a, loc=loc, ip=ip)
@dsl_user_op
def _bf16x2_max(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
"max.bf16x2 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
)
return Uint32(out)
return _bf16x2_binary("max", a, b, loc=loc, ip=ip)
@dsl_user_op
def _bf16x2_mul(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
"mul.rn.bf16x2 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
)
return Uint32(out)
return _bf16x2_binary("mul.rn", a, b, loc=loc, ip=ip)
@dsl_user_op
def _bf16x2_sub(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
return _bf16x2_binary("sub.rn", a, b, loc=loc, ip=ip)
+34 -30
View File
@@ -90,17 +90,18 @@ def mma_f16(
loc=None,
ip=None,
) -> None:
nvvm.tcgen05_mma(
nvvm.Tcgen05MMAKind.F16,
NVVM_CTA_GROUP_MAP[cta_group],
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
Uint64(a_desc).ir_value(loc=loc, ip=ip),
Uint64(b_desc).ir_value(loc=loc, ip=ip),
Int32(idesc).ir_value(loc=loc, ip=ip),
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
with cute.arch.elect_one():
nvvm.tcgen05_mma(
nvvm.Tcgen05MMAKind.F16,
NVVM_CTA_GROUP_MAP[cta_group],
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
Uint64(a_desc).ir_value(loc=loc, ip=ip),
Uint64(b_desc).ir_value(loc=loc, ip=ip),
Int32(idesc).ir_value(loc=loc, ip=ip),
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
@dsl_user_op
@@ -115,17 +116,18 @@ def mma_ts_f16(
loc=None,
ip=None,
) -> None:
nvvm.tcgen05_mma(
nvvm.Tcgen05MMAKind.F16,
NVVM_CTA_GROUP_MAP[cta_group],
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
_make_tmem_llvm_ptr(a_tmem, loc=loc, ip=ip),
Uint64(b_desc).ir_value(loc=loc, ip=ip),
Int32(idesc).ir_value(loc=loc, ip=ip),
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
with cute.arch.elect_one():
nvvm.tcgen05_mma(
nvvm.Tcgen05MMAKind.F16,
NVVM_CTA_GROUP_MAP[cta_group],
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
_make_tmem_llvm_ptr(a_tmem, loc=loc, ip=ip),
Uint64(b_desc).ir_value(loc=loc, ip=ip),
Int32(idesc).ir_value(loc=loc, ip=ip),
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
@dsl_user_op
@@ -133,15 +135,17 @@ def commit(mbar, cta_mask=None, cta_group: int = 1, *, loc=None, ip=None):
mbar_llvm = mbar.to_llvm_ptr(loc=loc, ip=ip)
group = NVVM_CTA_GROUP_MAP[cta_group]
if cutlass.const_expr(cta_mask is not None):
nvvm.tcgen05_commit_arrive(
mbar_llvm,
multicast_mask=cta_mask.ir_value(loc=loc, ip=ip),
group=group,
loc=loc,
ip=ip,
)
with cute.arch.elect_one():
nvvm.tcgen05_commit_arrive(
mbar_llvm,
multicast_mask=cta_mask.ir_value(loc=loc, ip=ip),
group=group,
loc=loc,
ip=ip,
)
else:
nvvm.tcgen05_commit_arrive(mbar_llvm, group=group, loc=loc, ip=ip)
with cute.arch.elect_one():
nvvm.tcgen05_commit_arrive(mbar_llvm, group=group, loc=loc, ip=ip)
@dsl_user_op
@@ -22,7 +22,6 @@ from vllm.config import VllmConfig
from vllm.distributed.kv_transfer.kv_connector.utils import (
EngineId,
TransferTopology,
get_current_attn_backend,
get_current_attn_backends,
)
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
@@ -51,10 +50,18 @@ from vllm.platforms import current_platform
from vllm.utils.math_utils import cdiv
from vllm.utils.network_utils import get_ip, make_zmq_path, make_zmq_socket
from vllm.v1.attention.backend import AttentionMetadata
from vllm.v1.attention.backends.utils import get_kv_cache_layout
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID, get_kv_cache_layout
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import FullAttentionSpec, SlidingWindowSpec
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheSpec,
MambaSpec,
MLAAttentionSpec,
SlidingWindowMLASpec,
SlidingWindowSpec,
)
from vllm.v1.request import RequestStatus
from vllm.v1.worker.block_table import BlockTable
from vllm.v1.worker.utils import select_common_block_size
logger = init_logger(__name__)
@@ -85,6 +92,7 @@ class TransferRegion:
base_addr: int
block_len: int
kv_block_len: int
group_index: int = 0
def _get_tp_ratio(local_tp_size: int, remote_tp_size: int) -> int:
@@ -111,24 +119,58 @@ def _get_tp_ratio(local_tp_size: int, remote_tp_size: int) -> int:
def _expand_transfer_regions(
base_addrs: list[int],
block_lens: list[int],
kv_block_lens: list[int],
layer_names: list[str],
layer_indices: list[int],
is_kv_layout_blocks_first: bool,
group_indices: list[int] | None = None,
split_kv_regions: list[bool] | None = None,
) -> list[TransferRegion]:
"""Expand registered KV tensors into the regions transferred by Mooncake."""
assert (
len(base_addrs) == len(block_lens) == len(layer_names) == len(layer_indices)
len(base_addrs)
== len(block_lens)
== len(kv_block_lens)
== len(layer_names)
== len(layer_indices)
), (
"Mooncake transfer regions require matching metadata lengths, got "
f"base_addrs={len(base_addrs)}, block_lens={len(block_lens)}, "
f"kv_block_lens={len(kv_block_lens)}, "
f"layer_names={len(layer_names)}, "
f"layer_indices={len(layer_indices)}."
)
if group_indices is None:
group_indices = [0] * len(layer_names)
assert len(group_indices) == len(layer_names), (
"Mooncake transfer regions require matching group metadata lengths, "
f"got group_indices={len(group_indices)}, layer_names={len(layer_names)}."
)
if split_kv_regions is None:
split_kv_regions = [is_kv_layout_blocks_first] * len(layer_names)
assert len(split_kv_regions) == len(layer_names), (
"Mooncake transfer regions require matching split metadata, "
f"got split_kv_regions={len(split_kv_regions)}, "
f"layer_names={len(layer_names)}."
)
regions: list[TransferRegion] = []
for base_addr, block_len, layer_name, layer_index in zip(
base_addrs, block_lens, layer_names, layer_indices
for (
base_addr,
block_len,
kv_block_len,
layer_name,
layer_index,
group_index,
split_kv_region,
) in zip(
base_addrs,
block_lens,
kv_block_lens,
layer_names,
layer_indices,
group_indices,
split_kv_regions,
):
kv_block_len = block_len // 2 if is_kv_layout_blocks_first else block_len
regions.append(
TransferRegion(
layer_name=layer_name,
@@ -136,9 +178,10 @@ def _expand_transfer_regions(
base_addr=base_addr,
block_len=block_len,
kv_block_len=kv_block_len,
group_index=group_index,
)
)
if is_kv_layout_blocks_first:
if split_kv_region:
regions.append(
TransferRegion(
layer_name=layer_name,
@@ -146,6 +189,7 @@ def _expand_transfer_regions(
base_addr=base_addr + kv_block_len,
block_len=block_len,
kv_block_len=kv_block_len,
group_index=group_index,
)
)
return regions
@@ -308,6 +352,17 @@ def _align_transfer_regions(
f"{remote_region.layer_index}."
),
)
if local_region.group_index != remote_region.group_index:
return (
[],
[],
(
"Mooncake registered group index mismatch for "
f"{local_region.layer_name}: producer="
f"{local_region.group_index}, consumer="
f"{remote_region.group_index}."
),
)
aligned_local.append(local_region)
aligned_remote.append(remote_region)
@@ -332,8 +387,10 @@ class MooncakeXferMetadata(
req_blocks: dict[ReqId, tuple[TransferId, list[list[int]]]]
kv_caches_base_addr: list[int]
block_lens: list[int]
kv_block_lens: list[int]
registered_layer_names: list[str] = msgspec.field(default_factory=list)
registered_layer_indices: list[int] = msgspec.field(default_factory=list)
registered_group_indices: list[int] = msgspec.field(default_factory=list)
class MooncakeXferResponseStatus(IntEnum):
@@ -581,6 +638,9 @@ class MooncakeConnectorScheduler:
for g in kv_cache_config.kv_cache_groups
)
)
# GDN is represented as a MambaSpec in vLLM. This Mooncake MambaSpec
# path is currently tested with GDN; Mamba2 is not validated yet.
self._has_mamba = kv_cache_config.has_mamba_layers
# Requests that need to start recv/send.
# New requests are added by update_state_after_alloc in
@@ -617,6 +677,38 @@ class MooncakeConnectorScheduler:
for i, blocks in enumerate(block_ids)
]
def _get_remote_prefill_token_count(self, num_prompt_tokens: int) -> int:
"""D-side only. Returns N-1 for Mamba models since the decoder
always recomputes the last token and must start from h(N-1)."""
if self._has_mamba and num_prompt_tokens > 1:
return num_prompt_tokens - 1
return num_prompt_tokens
def _truncate_mamba_request_for_prefill(self, request: "Request") -> None:
"""P-side only: drop the last prompt token so the prefiller computes
h(N-1) instead of h(N). The decoder recomputes the last token to
derive h(N) correctly.
Guarded by ``_p_side_truncated`` to avoid repeated truncation if the
request is preempted and rescheduled."""
params = request.kv_transfer_params
if (
params is not None
and not params.get("_p_side_truncated")
and request.num_prompt_tokens > 1
):
if request.prompt_token_ids is not None:
request.prompt_token_ids.pop()
elif request.prompt_embeds is not None:
request.prompt_embeds = request.prompt_embeds[:-1]
else:
return
request._all_token_ids.pop()
request.num_prompt_tokens -= 1
request.max_tokens = 1
params["_p_side_truncated"] = True
def get_num_new_matched_tokens(
self, request: "Request", num_computed_tokens: int
) -> tuple[int, bool]:
@@ -650,10 +742,15 @@ class MooncakeConnectorScheduler:
# Remote prefill: get all prompt blocks from remote.
assert not self.is_kv_producer
token_ids = request.prompt_token_ids or []
count = len(token_ids) - num_computed_tokens
count = self._get_remote_prefill_token_count(len(token_ids)) - (
num_computed_tokens
)
if count > 0:
return count, True
if params.get("do_remote_decode") and self._has_mamba:
self._truncate_mamba_request_for_prefill(request)
# No remote prefill for this request.
return 0, False
@@ -802,7 +899,7 @@ class MooncakeConnectorWorker:
self,
vllm_config: VllmConfig,
engine_id: str,
kv_cache_config: "KVCacheConfig | None" = None,
kv_cache_config: "KVCacheConfig",
):
if TransferEngine is None:
logger.error("Mooncake is not available")
@@ -831,10 +928,15 @@ class MooncakeConnectorWorker:
protocol = kv_transfer_config.kv_connector_extra_config.get( # type: ignore[union-attr]
"mooncake_protocol", "rdma"
)
device_name = kv_transfer_config.kv_connector_extra_config.get( # type: ignore[union-attr]
"device_name", ""
)
logger.info(
"The Mooncake Transfer Engine is using %s as its protocol.", protocol
)
ret_value = self.engine.initialize(self.hostname, "P2PHANDSHAKE", protocol, "")
ret_value = self.engine.initialize(
self.hostname, "P2PHANDSHAKE", protocol, device_name
)
if ret_value != 0:
raise RuntimeError("Mooncake Transfer Engine initialization failed.")
@@ -852,10 +954,11 @@ class MooncakeConnectorWorker:
self.engine_id: EngineId = engine_id
self.tp_rank = get_tensor_model_parallel_rank()
self.tp_size = get_tensor_model_parallel_world_size()
self.num_blocks = 0
self.block_len_per_layer: list[int] = []
self.kv_block_len_per_layer: list[int] = []
self.registered_layer_names: list[str] = []
self.registered_layer_indices: list[int] = []
self.registered_group_indices: list[int] = []
self.seen_base_addresses: list[int] = []
assert (parallel_config := vllm_config.parallel_config)
@@ -916,26 +1019,40 @@ class MooncakeConnectorWorker:
self.cache_config = vllm_config.cache_config
self.kv_cache_config = kv_cache_config
self.use_mla = self.model_config.use_mla
self._physical_blocks_per_logical_kv_block = 1
self._sync_block_size_with_kernel()
# Get the attention backend from the first layer
# NOTE (NickLucche) models with multiple backends are not supported yet
backend = get_current_attn_backend(vllm_config)
self.backend_name = backend.get_name()
self.attn_backends = get_current_attn_backends(vllm_config)
self.kv_cache_layout = get_kv_cache_layout()
logger.debug("Detected attention backend %s", self.backend_name)
logger.debug(
"Detected attention backends %s",
[backend.get_name() for backend in self.attn_backends],
)
logger.debug("Detected kv cache layout %s", self.kv_cache_layout)
self._tp_size: dict[EngineId, int] = {self.engine_id: self.tp_size}
self._layer_specs: dict[str, KVCacheSpec] = {}
for group in kv_cache_config.kv_cache_groups:
group_spec = group.kv_cache_spec
specs_by_layer = getattr(group_spec, "kv_cache_specs", {})
for layer_name in group.layer_names:
self._layer_specs[layer_name] = specs_by_layer.get(
layer_name, group_spec
)
self._layer_group_indices: dict[str, int] = {
layer: group_index
for group_index, group in enumerate(kv_cache_config.kv_cache_groups)
for layer in group.layer_names
}
self.transfer_topo = TransferTopology(
tp_rank=self.tp_rank,
tp_size=self.tp_size,
block_size=self.block_size,
engine_id=self.engine_id,
is_mla=self.use_mla,
is_mamba=False,
is_mamba=kv_cache_config.has_mamba_layers,
total_num_kv_heads=self.model_config.get_total_num_kv_heads(),
attn_backends=[backend],
attn_backends=self.attn_backends,
)
self.async_zmq_ctx = zmq.asyncio.Context()
@@ -958,6 +1075,9 @@ class MooncakeConnectorWorker:
kernel_block_size,
)
assert self.block_size > kernel_block_size
self._physical_blocks_per_logical_kv_block = (
self.block_size // kernel_block_size
)
self.block_size = kernel_block_size
def __del__(self):
@@ -1092,14 +1212,18 @@ class MooncakeConnectorWorker:
local_regions = self._get_transfer_regions(
self.kv_caches_base_addr,
self.block_len_per_layer,
self.kv_block_len_per_layer,
self.registered_layer_names,
self.registered_layer_indices,
self.registered_group_indices,
)
remote_regions = self._get_transfer_regions(
meta.kv_caches_base_addr,
meta.block_lens,
meta.kv_block_lens,
meta.registered_layer_names,
meta.registered_layer_indices,
meta.registered_group_indices,
)
local_regions, remote_regions, align_err = _align_transfer_regions(
local_regions, remote_regions
@@ -1271,6 +1395,32 @@ class MooncakeConnectorWorker:
remote_tp_ranks,
)
def _logical_to_kernel_block_ids(
self, block_ids: list[list[int]]
) -> list[list[int]]:
# For example, if a 544-token logical block is served by 32-token
# FA kernel blocks, FA block id k expands to [17k, ..., 17k + 16],
# while the matching Mamba/GDN state block remains k. Only attention
# groups need logical block ids expanded to kernel block ids; Mamba/GDN
# state block ids stay in the logical/page-id space.
if self._physical_blocks_per_logical_kv_block == 1:
return block_ids
block_arange = np.arange(self._physical_blocks_per_logical_kv_block).reshape(
1, -1
)
group_specs = self.kv_cache_config.kv_cache_groups
return [
BlockTable.map_to_kernel_blocks(
np.array(group),
self._physical_blocks_per_logical_kv_block,
block_arange,
).tolist()
if not isinstance(group_specs[i].kv_cache_spec, MambaSpec)
else group
for i, group in enumerate(block_ids)
]
async def _build_transfer_params(
self,
ready_reqs: list[tuple[ReqId, SendBlockMeta]],
@@ -1293,14 +1443,6 @@ class MooncakeConnectorWorker:
):
continue
# Per-group partial hit trimming, then flatten.
# With HMA, groups share the same KV tensor but use different
# block ranges. We trim and concatenate so the coalescer and
# address math see one flat block list — same as non-HMA, but
# now including blocks from every group.
local_block_ids: list[int] = []
remote_block_ids: list[int] = []
has_block_error = False
if len(send_meta.local_block_ids) != len(remote_block_ids_per_group):
logger.error(
"req %s: KV group count mismatch: local=%d, remote=%d",
@@ -1312,26 +1454,55 @@ class MooncakeConnectorWorker:
if err_msg is None:
err_msg = "KV group count mismatch"
continue
for local_group, remote_group in zip(
send_meta.local_block_ids, remote_block_ids_per_group
# Keep KV-cache group identity. Hybrid/HMA groups can carry
# different semantics (e.g. full-attention KV pages vs GDN/Mamba
# inner-state slots), so their block IDs must not be flattened and
# reused for every registered region.
local_block_ids_by_group: list[list[int]] = []
remote_block_ids_by_group: list[list[int]] = []
has_block_error = False
group_specs = self.kv_cache_config.kv_cache_groups
for group_index, (local_group, remote_group) in enumerate(
zip(send_meta.local_block_ids, remote_block_ids_per_group)
):
is_mamba_group = isinstance(
group_specs[group_index].kv_cache_spec,
MambaSpec,
)
if is_mamba_group:
# Mamba/GDN prefix caching can use null blocks only as
# align-mode placeholders. They do not carry transferable
# state, so skip them on both producer and consumer sides.
local_group = [
block_id
for block_id in local_group
if block_id != NULL_BLOCK_ID
]
remote_group = [
block_id
for block_id in remote_group
if block_id != NULL_BLOCK_ID
]
n_local = len(local_group)
n_remote = len(remote_group)
if n_local < n_remote:
logger.error(
"req %s: local blocks(%d) < remote blocks(%d) "
"in a KV cache group",
"in a KV cache group (is_mamba_group=%s)",
d_req_id,
n_local,
n_remote,
is_mamba_group,
)
has_block_error = True
break
if n_local > n_remote:
elif n_local > n_remote:
# Partial prefix cache hit: just read uncomputed blocks.
local_group = local_group[-n_remote:]
local_block_ids.extend(local_group)
remote_block_ids.extend(remote_group)
local_group = local_group[-n_remote:] if n_remote > 0 else []
local_block_ids_by_group.append(local_group)
remote_block_ids_by_group.append(remote_group)
if has_block_error:
err_reqs.append(d_req_id)
@@ -1339,22 +1510,44 @@ class MooncakeConnectorWorker:
err_msg = "P num blocks less than D"
continue
if not local_block_ids:
if not any(local_block_ids_by_group):
continue
# Group by indices
group_local_block_ids, group_remote_block_ids = group_concurrent_contiguous(
local_block_ids, remote_block_ids
local_block_ids_by_group = self._logical_to_kernel_block_ids(
local_block_ids_by_group
)
remote_block_ids_by_group = self._logical_to_kernel_block_ids(
remote_block_ids_by_group
)
for local_region, remote_region in zip(local_regions, remote_regions):
should_transfer, src_region_offset, dst_region_offset, transfer_len = (
self._get_sender_transfer_plan(
local_kv_block_len=local_region.kv_block_len,
remote_kv_block_len=remote_region.kv_block_len,
remote_tp_rank=agent_meta.remote_tp_rank,
remote_tp_size=agent_meta.remote_tp_size,
)
assert local_region.group_index == remote_region.group_index, (
"Aligned Mooncake transfer regions must belong to the same "
"KV group."
)
group_index = local_region.group_index
assert group_index < len(local_block_ids_by_group), (
"Transfer region references a missing KV group."
)
local_block_ids = local_block_ids_by_group[group_index]
remote_block_ids = remote_block_ids_by_group[group_index]
if not local_block_ids:
continue
# Group by indices within this region's KV-cache group only.
group_local_block_ids, group_remote_block_ids = (
group_concurrent_contiguous(local_block_ids, remote_block_ids)
)
(
should_transfer,
src_region_offset,
dst_region_offset,
transfer_len,
) = self._get_sender_transfer_plan(
local_kv_block_len=local_region.kv_block_len,
remote_kv_block_len=remote_region.kv_block_len,
remote_tp_rank=agent_meta.remote_tp_rank,
remote_tp_size=agent_meta.remote_tp_size,
)
if not should_transfer:
# Replicated KV cache: only one producer rank in the TP group
@@ -1368,7 +1561,7 @@ class MooncakeConnectorWorker:
"Computed source transfer region exceeds local KV block size."
)
assert dst_region_offset + transfer_len <= remote_region.kv_block_len, (
"Computed destination transfer region exceeds remote KV block size."
"Destination transfer region exceeds remote KV block size."
)
# Collapse one contiguous block group into a single larger
# transfer descriptor when the per-block copy is identical.
@@ -1411,28 +1604,10 @@ class MooncakeConnectorWorker:
)
lengths.append(transfer_len)
if local_region is local_regions[0]:
logger.debug(
"Mooncake transfer plan for request %s: local_tp=%d "
"remote_tp=%d remote_tp_rank=%d local_block_len=%d "
"remote_block_len=%d src_offset=%d dst_offset=%d "
"transfer_len=%d coalesce=%s",
d_req_id,
self.tp_size,
agent_meta.remote_tp_size,
agent_meta.remote_tp_rank,
local_region.block_len,
remote_region.block_len,
src_region_offset,
dst_region_offset,
transfer_len,
can_coalesce,
)
logger.debug(
"Sending kv_caches for request %s (%d blocks) to %s",
d_req_id,
len(local_block_ids),
sum(len(group) for group in local_block_ids_by_group),
remote_session,
)
@@ -1480,18 +1655,33 @@ class MooncakeConnectorWorker:
logger.info("Registering KV_Caches. use_mla: %s", self.use_mla)
kv_data_ptrs = []
kv_data_lens = []
seen_base_addresses = []
kv_data_ptrs: list[int] = []
kv_data_lens: list[int] = []
region_base_addresses: list[int] = []
seen_storage_ptrs: set[int] = set()
self.block_len_per_layer = []
self.kv_block_len_per_layer = []
self.registered_layer_names = []
self.registered_layer_indices = []
self.registered_group_indices = []
split_k_and_v = self.transfer_topo.split_k_and_v
tensor_size_bytes = None
for layer_name, cache_or_caches in kv_caches.items():
layer_index = extract_layer_index(layer_name)
cache_list = cache_or_caches if split_k_and_v else [cache_or_caches]
layer_spec = self._layer_specs.get(layer_name)
if layer_spec is None:
logger.debug(
"Skipping layer %s because no KV cache spec is present.",
layer_name,
)
continue
if isinstance(layer_spec, MambaSpec):
conv, _ = cache_or_caches
cache_list = [conv]
else:
cache_list = self.transfer_topo.get_transfer_cache_regions(
cache_or_caches, layer_spec
)
logger.debug(
"registering layer %s with %d cache tensor(s)",
layer_name,
@@ -1501,45 +1691,46 @@ class MooncakeConnectorWorker:
for cache in cache_list:
self._log_debug_cache_registration(layer_name, cache)
base_addr = cache.data_ptr()
if base_addr in seen_base_addresses:
continue
seen_base_addresses.append(base_addr)
if tensor_size_bytes is None:
tensor_size_bytes = cache.nbytes
self.num_blocks = cache.shape[0]
assert cache.shape[0] == self.num_blocks, (
"All kv cache tensors must have the same number of blocks"
)
# Use stride-based block length so RDMA reaches the last
# block's padding (e.g. DeepseekV4 MLA alignment). stride(0)
# reflects the actual byte distance between consecutive
# blocks in GPU memory, which matches or exceeds the
# shape-based size.
block_len = cache.stride(0) * cache.element_size()
region_base_addresses.append(base_addr)
if isinstance(layer_spec, (MLAAttentionSpec, SlidingWindowMLASpec)):
kv_block_len = layer_spec.page_size_bytes
elif self.transfer_topo.virtually_split_kv_in_blocks and not isinstance(
layer_spec, MambaSpec
):
kv_block_len = block_len // 2
else:
kv_block_len = block_len
self.block_len_per_layer.append(block_len)
self.kv_block_len_per_layer.append(kv_block_len)
self.registered_layer_names.append(layer_name)
self.registered_layer_indices.append(layer_index)
kv_data_ptrs.append(base_addr)
kv_data_lens.append(self.num_blocks * block_len)
self.registered_group_indices.append(
self._layer_group_indices[layer_name]
)
storage = cache.untyped_storage()
storage_addr = storage.data_ptr()
if storage_addr not in seen_storage_ptrs:
seen_storage_ptrs.add(storage_addr)
kv_data_ptrs.append(storage_addr)
kv_data_lens.append(storage.nbytes())
self.kv_caches_base_addr = seen_base_addresses
self.seen_base_addresses = seen_base_addresses
self.kv_caches_base_addr = region_base_addresses
self.seen_base_addresses = kv_data_ptrs
if not kv_data_ptrs:
raise RuntimeError("No KV cache tensors were registered with Mooncake.")
ret_value = self.engine.batch_register_memory(kv_data_ptrs, kv_data_lens)
if ret_value != 0:
raise RuntimeError("Mooncake batch memory registration failed.")
assert tensor_size_bytes is not None
assert self.num_blocks != 0
self.device_kv_caches = kv_caches
logger.debug(
"registered num_blocks=%d block_lens=%s",
self.num_blocks,
"registered block_lens=%s kv_block_lens=%s",
self.block_len_per_layer,
self.kv_block_len_per_layer,
)
# No need to launch server for D node.
@@ -1642,8 +1833,10 @@ class MooncakeConnectorWorker:
},
kv_caches_base_addr=self.kv_caches_base_addr,
block_lens=self.block_len_per_layer,
kv_block_lens=self.kv_block_len_per_layer,
registered_layer_names=self.registered_layer_names,
registered_layer_indices=self.registered_layer_indices,
registered_group_indices=self.registered_group_indices,
)
encoded_data = self._encoder.encode(metadata)
@@ -1852,15 +2045,34 @@ class MooncakeConnectorWorker:
self,
base_addrs: list[int],
block_lens: list[int],
kv_block_lens: list[int],
layer_names: list[str],
layer_indices: list[int],
group_indices: list[int] | None = None,
) -> list[TransferRegion]:
if not group_indices:
group_indices = [
self._layer_group_indices.get(layer_name, 0)
for layer_name in layer_names
]
split_kv_regions = None
if self.transfer_topo.virtually_split_kv_in_blocks:
split_kv_regions = [
not isinstance(
self._layer_specs[layer_name],
(MambaSpec, MLAAttentionSpec, SlidingWindowMLASpec),
)
for layer_name in layer_names
]
return _expand_transfer_regions(
base_addrs=base_addrs,
block_lens=block_lens,
kv_block_lens=kv_block_lens,
layer_names=layer_names,
layer_indices=layer_indices,
is_kv_layout_blocks_first=self.transfer_topo.virtually_split_kv_in_blocks,
group_indices=group_indices,
split_kv_regions=split_kv_regions,
)
def _get_sender_transfer_plan(
@@ -492,12 +492,15 @@ class MultiConnector(KVConnectorBase_V1, SupportsHMA):
async_saves += 1
if txfer_params is not None:
if kv_txfer_params is not None:
# TODO we can probably change this to merge the dicts here,
# checking for key clashes.
raise RuntimeError(
"Only one connector can produce KV transfer params"
)
kv_txfer_params = txfer_params
clashes = set(kv_txfer_params) & set(txfer_params)
if clashes:
raise RuntimeError(
"Key clash in kv_transfer_params from multiple "
f"connectors: {clashes}"
)
kv_txfer_params.update(txfer_params)
else:
kv_txfer_params = txfer_params
if async_saves > 1:
self._extra_async_saves[request.request_id] = async_saves - 1
@@ -320,7 +320,7 @@ class NixlBaseConnectorScheduler:
logger.warning("Connection listener got unexpected message %s", msg)
sock.send_multipart((identity, b"", encoded_data[target_tp_rank]))
def _mamba_prefill_token_count(self, num_prompt_tokens: int) -> int:
def _get_remote_prefill_token_count(self, num_prompt_tokens: int) -> int:
"""D-side only. Returns N-1 for Mamba models since the decoder
always recomputes the last token and must start from h(N-1)."""
if self._has_mamba and num_prompt_tokens > 1:
@@ -60,7 +60,7 @@ class NixlPullConnectorScheduler(NixlBaseConnectorScheduler):
if params is not None and params.get("do_remote_prefill"):
# Remote prefill: get all prompt blocks from remote.
token_ids = request.prompt_token_ids or []
actual = self._mamba_prefill_token_count(len(token_ids))
actual = self._get_remote_prefill_token_count(len(token_ids))
count = actual - num_computed_tokens
if count > 0:
return count, True
@@ -116,7 +116,7 @@ class NixlPushConnectorScheduler(NixlBaseConnectorScheduler):
if params is not None and params.get("do_remote_prefill"):
token_ids = request.prompt_token_ids or []
actual = self._mamba_prefill_token_count(len(token_ids))
actual = self._get_remote_prefill_token_count(len(token_ids))
count = actual - num_computed_tokens
if count > 0:
return count, True
@@ -610,6 +610,18 @@ class ChatCompletionRequest(OpenAIBaseModel):
"min_p", self._DEFAULT_SAMPLING_PARAMS["min_p"]
)
# Merge server-default stop_token_ids (e.g., model-specific tokens
# like </call> for gpt-oss) with any request-specified ones
stop_token_ids = self.stop_token_ids
default_stop_ids = default_sampling_params.get("stop_token_ids")
if default_stop_ids:
if not stop_token_ids:
stop_token_ids = list(default_stop_ids)
else:
stop_token_ids = list(
dict.fromkeys([*stop_token_ids, *default_stop_ids])
)
prompt_logprobs = self.prompt_logprobs
if prompt_logprobs is None and self.echo:
prompt_logprobs = self.top_logprobs
@@ -661,7 +673,7 @@ class ChatCompletionRequest(OpenAIBaseModel):
min_p=min_p,
seed=self.seed,
stop=self.stop,
stop_token_ids=self.stop_token_ids,
stop_token_ids=stop_token_ids,
logprobs=self.top_logprobs if self.logprobs else None,
prompt_logprobs=prompt_logprobs,
ignore_eos=self.ignore_eos,
+13 -1
View File
@@ -288,6 +288,18 @@ class CompletionRequest(OpenAIBaseModel):
"min_p", self._DEFAULT_SAMPLING_PARAMS["min_p"]
)
# Merge server-default stop_token_ids (e.g., model-specific tokens
# like </call> for gpt-oss) with any request-specified ones
stop_token_ids = self.stop_token_ids
default_stop_ids = default_sampling_params.get("stop_token_ids")
if default_stop_ids:
if not stop_token_ids:
stop_token_ids = list(default_stop_ids)
else:
stop_token_ids = list(
dict.fromkeys([*stop_token_ids, *default_stop_ids])
)
prompt_logprobs = self.prompt_logprobs
if prompt_logprobs is None and self.echo:
prompt_logprobs = self.logprobs
@@ -341,7 +353,7 @@ class CompletionRequest(OpenAIBaseModel):
min_p=min_p,
seed=self.seed,
stop=self.stop,
stop_token_ids=self.stop_token_ids,
stop_token_ids=stop_token_ids,
logprobs=self.logprobs,
ignore_eos=self.ignore_eos,
max_tokens=max_tokens if not echo_without_generation else 1,
@@ -318,6 +318,7 @@ def _parse_function_call(message: Message, recipient: str) -> list[ResponseOutpu
type="function_call",
name=function_name,
id=f"fc_{random_id}",
status="completed",
)
output_items.append(response_item)
return output_items
+9
View File
@@ -79,6 +79,7 @@ if TYPE_CHECKING:
VLLM_MAX_AUDIO_CLIP_FILESIZE_MB: int = 25
VLLM_MAX_AUDIO_DECODE_DURATION_S: int = 600
VLLM_MAX_AUDIO_PREPROCESS_WORKERS: int = max(1, min(os.cpu_count() or 1, 2))
VLLM_MAX_IMAGE_PIXELS: int = 178_956_970
VLLM_VIDEO_LOADER_BACKEND: str = "opencv"
VLLM_MEDIA_CONNECTOR: str = "http"
VLLM_MM_HASHER_ALGORITHM: str = "blake3"
@@ -954,6 +955,13 @@ environment_variables: dict[str, Callable[[], Any]] = {
str(max(1, min(os.cpu_count() or 1, 2))),
)
),
# Maximum decoded image size in pixels. Small compressed images can
# expand into gigabytes of raster memory. This limit is enforced before
# decoding so the memory is never allocated. Default matches PIL's
# built-in 2x decompression-bomb threshold (~179M pixels, ~680 MB RGB).
"VLLM_MAX_IMAGE_PIXELS": lambda: int(
os.getenv("VLLM_MAX_IMAGE_PIXELS", "178956970")
),
# Backend for Video IO — selects the frame-sampling algorithm.
# - "opencv": uniform sampling.
# - "opencv_dynamic": duration-aware dynamic sampling.
@@ -2083,6 +2091,7 @@ def compile_factors() -> dict[str, object]:
"VLLM_MAX_AUDIO_CLIP_FILESIZE_MB",
"VLLM_MAX_AUDIO_DECODE_DURATION_S",
"VLLM_MAX_AUDIO_PREPROCESS_WORKERS",
"VLLM_MAX_IMAGE_PIXELS",
"VLLM_VIDEO_LOADER_BACKEND",
"VLLM_MEDIA_CONNECTOR",
"VLLM_OBJECT_STORAGE_SHM_BUFFER_NAME",
@@ -458,6 +458,7 @@ class Attention(nn.Module, AttentionLayerBase):
# shape does not match the query shape, so we optionally let the model
# definition specify the output tensor shape.
output_shape: torch.Size | None = None,
output_dtype: torch.dtype | None = None,
) -> torch.Tensor:
"""
The KV cache is stored inside this class and is accessed via
@@ -472,7 +473,8 @@ class Attention(nn.Module, AttentionLayerBase):
torch.ops.vllm.maybe_calc_kv_scales(
query, key, value, _encode_layer_name(self.layer_name)
)
output_dtype = query.dtype
if output_dtype is None:
output_dtype = query.dtype
if self.query_quant is not None:
# quantizing with a simple torch operation enables
# torch.compile to fuse this into previous ops
@@ -818,14 +818,14 @@ class MLAAttention(nn.Module, AttentionLayerBase):
attn_out,
lse,
get_dcp_group(),
is_lse_base_on_e=True,
is_lse_base_on_e=self.impl.lse_base_on_e,
)
else:
attn_out = cp_lse_ag_out_rs(
attn_out,
lse,
get_dcp_group(),
is_lse_base_on_e=True,
is_lse_base_on_e=self.impl.lse_base_on_e,
)
# v_up projection
@@ -301,6 +301,13 @@ def convert_to_unquantized_kernel_format(
is_gated_act_gemm=is_act_and_mul,
)
if (
unquantized_backend == UnquantizedMoeBackend.TRITON
and current_platform.is_rocm()
and envs.VLLM_ROCM_MOE_PADDING
):
# Skip .contiguous(): it would undo the ROCm MoE weight padding.
return w13_weight, w2_weight
return w13_weight.contiguous(), w2_weight.contiguous()
@@ -12,6 +12,7 @@ from vllm.model_executor.layers.fused_moe.topk_weight_and_reduce import (
TopKWeightAndReduceDelegate,
)
from vllm.model_executor.layers.fused_moe.utils import moe_kernel_quantize_input
from vllm.platforms import current_platform
from vllm.utils.math_utils import round_up
from vllm.v1.worker.ubatching import (
dbo_current_ubatch_id,
@@ -59,6 +60,7 @@ class DeepEPHTPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
self.dp_size = dp_size
self.rank_expert_offset = rank_expert_offset
self.async_prepare = True
self.sync_dbo_comm = current_platform.is_rocm()
# The dispatch function returns a handle that the combine function
# requires. Under DBO microbatching we must track one handle per
@@ -68,6 +70,13 @@ class DeepEPHTPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
# From https://github.com/deepseek-ai/DeepEP/blob/9fe9021f29c9083cd1808ab36b740208524d9f63/deep_ep/buffer.py#L164
self.available_rank_configs = [2, 4, 8, 16, 24, 32, 64, 128, 144, 160]
def _sync_dbo_comm_if_needed(self) -> None:
if self.sync_dbo_comm and dbo_enabled():
# ROCm DeepEP HT dispatch/combine reuse Buffer-owned communication
# workspace. Do not let the next DBO ubatch reuse that workspace
# before this ubatch's HT kernel has completed.
torch.cuda.current_stream().synchronize()
def num_dispatchers(self) -> int:
return self.num_dispatchers_
@@ -161,6 +170,8 @@ class DeepEPHTPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
allocate_on_comm_stream=False,
)
self._sync_dbo_comm_if_needed()
# record the handle for this ubatch
a2a_idx = dbo_current_ubatch_id()
self.handles[a2a_idx] = handle
@@ -375,6 +386,8 @@ class DeepEPHTPrepareAndFinalize(mk.FusedMoEPrepareAndFinalizeModular):
allocate_on_comm_stream=False,
)
self._sync_dbo_comm_if_needed()
dbo_switch_to_compute()
if do_async:
@@ -0,0 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from vllm.model_executor.layers.hpc.hpc_module import HpcModule
from vllm.model_executor.layers.hpc.rope_norm import HpcRopeNorm, QkNormPolicy
__all__ = [
"HpcModule",
"HpcRopeNorm",
"QkNormPolicy",
]
@@ -0,0 +1,18 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch.nn as nn
class HpcModule(nn.Module):
def __init__(self):
super().__init__()
@classmethod
def support(cls, *args, **kwargs):
return True
def process_weights_after_loading(self, model):
pass
def forward(self, *args, **kwargs):
pass
+408
View File
@@ -0,0 +1,408 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""HPC fused RoPE + QK-Norm + KV-Cache-Write (+ optional FP8 Q quant).
Decoupled from HpcAttentionImpl; extra params are passed via layer attrs.
"""
from __future__ import annotations
from enum import IntEnum
from typing import Any
import torch
from vllm.config import get_current_vllm_config_or_none
from vllm.forward_context import ForwardContext, get_forward_context
from vllm.logger import init_logger
from vllm.model_executor.custom_op import CustomOp
from vllm.model_executor.layers.hpc.hpc_module import HpcModule
from vllm.utils.torch_utils import direct_register_custom_op
from vllm.v1.attention.backends.hpc_attn import HpcAttnMetadata
from vllm.v1.attention.backends.registry import AttentionBackendEnum
logger = init_logger(__name__)
_hpc_rope_norm_instances: dict[str, HpcRopeNorm] = {}
class QkNormPolicy(IntEnum):
"""Order of QK-RMSNorm relative to RoPE in the fused HPC rope_norm kernel.
The values are part of the HPC kernel ABI (passed through as ints), so they
must stay in sync with the kernel's expectations.
"""
# No QK-Norm: apply RoPE only.
NONE = 0
# Apply RoPE first, then QK-RMSNorm.
ROPE_THEN_NORM = 1
# Apply QK-RMSNorm first, then RoPE (e.g. HunYuan V3).
NORM_THEN_ROPE = 2
def hpc_rope_norm_forward(
qkv: torch.Tensor,
output: torch.Tensor,
layer_name: str,
) -> None:
"""Top-level custom op: RoPE + QK-Norm + KV-Cache-Write + FP8 Q quant.
Fully opaque to torch.compile (dynamo).
"""
forward_context: ForwardContext = get_forward_context()
attn_metadata: Any = forward_context.attn_metadata
if isinstance(attn_metadata, dict):
attn_metadata = attn_metadata[layer_name]
if attn_metadata is None:
output.zero_()
return
attn_layer = forward_context.no_compile_layers[layer_name]
# bind_kv_cache stores the per-layer KV cache as a single 5D tensor
# (num_blocks, 2, block_size, num_kv_heads, head_size), so use it directly.
kv_cache = attn_layer.kv_cache
if kv_cache.numel() == 0:
output.zero_()
return
assert kv_cache.dim() == 5, (
f"Expected kv_cache to have 5 dims, got {tuple(kv_cache.shape)}"
)
rope_norm = _hpc_rope_norm_instances[layer_name]
rope_norm._forward_impl(qkv, kv_cache, attn_metadata, attn_layer, output)
def hpc_rope_norm_forward_fake(
qkv: torch.Tensor,
output: torch.Tensor,
layer_name: str,
) -> None:
"""Fake impl for torch.compile trace; output is a mutated arg."""
return
direct_register_custom_op(
op_name="hpc_rope_norm_forward",
op_func=hpc_rope_norm_forward,
mutates_args=["output"],
fake_impl=hpc_rope_norm_forward_fake,
)
@CustomOp.register("hpc_rope_norm")
class HpcRopeNorm(CustomOp, HpcModule):
"""HPC fused RoPE + QK-Norm + KV-Cache-Write (+ optional FP8 Q quant).
Registered as a sub-module in model layers (e.g. HunYuanAttention).
Norm weights are extracted from fallback norm modules via
process_weights_after_loading() after all weights are loaded.
forward() is dispatched by CustomOp framework:
- In compiled mode: forward_cuda() calls torch.ops.vllm.hpc_rope_norm_forward
as a splitting point internal Python control flow is opaque
to torch.compile and not captured by CUDA Graph.
- In eager/native mode: forward_native() falls back to forward_cuda().
"""
def __init__(
self,
num_heads: int,
num_kv_heads: int,
head_dim: int,
cos_sin_cache: torch.Tensor,
use_qk_norm: bool,
fallback_qnorm: torch.nn.Module | None,
fallback_knorm: torch.nn.Module | None,
kv_cache_dtype: str,
layer_name: str,
qk_norm_policy: QkNormPolicy = QkNormPolicy.ROPE_THEN_NORM,
) -> None:
super().__init__()
self.num_heads = num_heads
self.num_kv_heads = num_kv_heads
self.head_dim = head_dim
self.use_qk_norm = use_qk_norm
self.q_size = num_heads * head_dim
self.kv_size = num_kv_heads * head_dim
# Register as a non-persistent buffer so it participates in sleep
# level-2 save/restore (CuMemAllocator) but is excluded from the
# checkpoint state_dict.
self.register_buffer("cos_sin_cache", cos_sin_cache.float(), persistent=False)
self.fallback_qnorm = fallback_qnorm
self.fallback_knorm = fallback_knorm
self.head_per_group = num_heads // num_kv_heads
# Pre-allocate norm weight tensors as Parameters so they are tracked by
# CuMemAllocator (for sleep/wake_up) and have stable addresses for CUDA
# Graph replay. process_weights_after_loading() updates them inplace via
# copy_() so refit does not invalidate captured graph tensor pointers.
# Shape is [head_dim] to match the HPC kernel's q/k_norm_weight layout.
if use_qk_norm and fallback_qnorm is not None:
self.qnorm_weight: torch.nn.Parameter | None = torch.nn.Parameter(
torch.empty(head_dim, dtype=torch.float32),
requires_grad=False,
)
else:
self.qnorm_weight = None
if use_qk_norm and fallback_knorm is not None:
self.knorm_weight: torch.nn.Parameter | None = torch.nn.Parameter(
torch.empty(head_dim, dtype=torch.float32),
requires_grad=False,
)
else:
self.knorm_weight = None
self.use_fp8 = "fp8" in kv_cache_dtype
# The RMSNorm/RoPE ordering is model dependent (e.g. HunYuan V3 applies
# QK-Norm before RoPE -> NORM_THEN_ROPE), so it is supplied by the
# caller. When QK-Norm is disabled the policy is forced to NONE.
self.qk_norm_policy = qk_norm_policy if use_qk_norm else QkNormPolicy.NONE
# Register layer_name + add self to the global instance registry so the
# module-level custom op (hpc_rope_norm_forward) can route back here.
self.layer_name: str | None = None
self.register_layer_name(layer_name)
@classmethod
def support(
cls,
num_heads: int,
num_kv_heads: int,
head_dim: int,
kv_cache_dtype: str,
) -> bool:
"""Check whether HpcRopeNorm is supported for the given config."""
# HpcRopeNorm is only enabled together with the HPC attention backend.
vllm_config = get_current_vllm_config_or_none()
if (
vllm_config is None
or vllm_config.attention_config.backend != AttentionBackendEnum.HPC_ATTN
):
return False
if kv_cache_dtype not in ("fp8_e4m3", "auto"):
logger.warning_once(
f"hpc rope_norm not support kv_cache_dtype:{kv_cache_dtype}, "
"only support fp8_e4m3, bfloat16"
)
return False
if head_dim not in (128,):
logger.warning_once("hpc rope_norm only support head_dim == 128.")
return False
head_per_group = num_heads // num_kv_heads
if head_per_group not in (4, 8):
logger.warning_once("hpc rope_norm only support head_per_group in [4, 8].")
return False
logger.info_once("enable hpc rope_norm")
return True
def process_weights_after_loading(self, model: torch.nn.Module = None) -> None:
"""Copy norm weights (float32) from fallback norm modules inplace.
Uses copy_() to preserve tensor addresses for CUDA Graph / refit
compatibility. Called by the model's load_weights() after all weights
are loaded (and generically from the model loader for DummyModelLoader
/ sleep-wake_up reload paths).
"""
if self.use_qk_norm:
if self.fallback_qnorm is not None and self.qnorm_weight is not None:
self.qnorm_weight.data.copy_(self.fallback_qnorm.weight.data.float())
if self.fallback_knorm is not None and self.knorm_weight is not None:
self.knorm_weight.data.copy_(self.fallback_knorm.weight.data.float())
def register_layer_name(self, layer_name: str) -> None:
"""Register layer_name and add self to the global registry.
The global registry is needed because the bottom-level torch op
(hpc_rope_norm_forward) is a module-level function and needs to
route back to the correct instance via layer_name.
"""
self.layer_name = layer_name
_hpc_rope_norm_instances[layer_name] = self
logger.debug(
"[rope_norm] registered HpcRopeNorm for layer: %s",
layer_name,
)
def forward_native(
self,
qkv: torch.Tensor,
layer_name: str,
) -> torch.Tensor:
"""Native fallback path: delegates to forward_cuda().
For now, the default native path will use CUDA backend path.
Other platforms may override via OOT registration.
"""
return self.forward_cuda(qkv, layer_name)
def forward_cuda(
self,
qkv: torch.Tensor,
layer_name: str,
) -> torch.Tensor:
"""CUDA path: invoke the torch custom op as a compile splitting point."""
num_tokens = qkv.shape[0]
output = torch.empty(
(num_tokens, self.num_heads, self.head_dim),
dtype=torch.float8_e4m3fn if self.use_fp8 else qkv.dtype,
device=qkv.device,
)
torch.ops.vllm.hpc_rope_norm_forward(qkv, output, layer_name)
return output
def _forward_impl(
self,
qkv: torch.Tensor,
kv_cache: torch.Tensor,
attn_metadata: HpcAttnMetadata,
attn_layer: torch.nn.Module,
output: torch.Tensor,
) -> None:
"""Actual forward logic called by the custom op.
Writes processed q into *output* and attaches extra params
(e.g. FP8 scales) to *attn_layer* as attributes.
"""
import hpc
num_actual_tokens = attn_metadata.num_actual_tokens
num_prefill_reqs = attn_metadata.num_prefills
num_decode_reqs = attn_metadata.num_decodes
num_decode_tokens = attn_metadata.num_decode_tokens
qkv = qkv[:num_actual_tokens]
num_prefill_tokens = num_actual_tokens - num_decode_tokens
# KV cache for the FP8 path is stored as uint8; view it as fp8 so the
# rope_norm_store_kv_fp8 kernel can write quantized K/V in-place.
if self.use_fp8:
kv_cache = kv_cache.view(torch.float8_e4m3fn)
# Per-tensor K/V scales (shape [1]) used by the FP8 kernel.
k_scale = attn_layer._k_scale.reshape(1)
v_scale = attn_layer._v_scale.reshape(1)
q_norm_weight = (
self.qnorm_weight if self.qk_norm_policy != QkNormPolicy.NONE else None
)
k_norm_weight = (
self.knorm_weight if self.qk_norm_policy != QkNormPolicy.NONE else None
)
# Dynamic per-token-per-head Q quant + per-tensor K/V (dqskv).
# rope_norm_store_kv_fp8 is registered as a torch op whose ``quant_policy``
# argument is typed as ``int``; pybind cannot cast the hpc.QuantType enum
# automatically, so pass its integer ``.value``.
QUANT_POLICY_DQSKV = hpc.QuantType.QPERTOKEN_PERHEAD_KPERTENSOR_VPERTENSOR.value
# --- Prefill ---
if num_prefill_reqs > 0:
seq_lens_prefill = attn_metadata.seq_lens[num_decode_reqs:]
cu_seqlens_prefill = attn_metadata.qo_indptr
max_seqlens = attn_metadata.max_query_len
block_table_prefill = attn_metadata.block_table_tensor[num_decode_reqs:]
qkv_prefill = qkv[num_decode_tokens:]
out_q_prefill = output[
num_decode_tokens : num_decode_tokens + num_prefill_tokens
]
if self.use_fp8:
_, q_scale, split_k_flag = hpc.rope_norm_store_kv_fp8(
key_cache=kv_cache[:, 0],
value_cache=kv_cache[:, 1],
qkv=qkv_prefill,
cos_sin=self.cos_sin_cache,
num_seqlen_per_req=seq_lens_prefill,
q_index=cu_seqlens_prefill,
kvcache_indices=block_table_prefill,
is_prefill=True,
k_scale=k_scale,
v_scale=v_scale,
quant_policy=QUANT_POLICY_DQSKV,
max_seqlens=max_seqlens,
q_norm_weight=q_norm_weight,
k_norm_weight=k_norm_weight,
qk_norm_policy=self.qk_norm_policy,
out_q=out_q_prefill,
)
attn_metadata.hpc_prefill_q_scale = q_scale
else:
hpc.rope_norm_store_kv(
kv_cache[:, 0],
kv_cache[:, 1],
qkv_prefill,
self.cos_sin_cache,
seq_lens_prefill,
cu_seqlens_prefill,
block_table_prefill,
True, # is_prefill
q_norm_weight=q_norm_weight,
k_norm_weight=k_norm_weight,
out_q=out_q_prefill,
qk_norm_policy=self.qk_norm_policy,
)
# --- Decode ---
if num_decode_reqs > 0:
num_seq_kvcache = attn_metadata.seq_lens[:num_decode_reqs]
block_table_decode = attn_metadata.block_table_tensor[:num_decode_reqs]
qkv_decode = qkv[:num_decode_tokens]
# Single-token decode: q_index is the per-request prefix sum
# [0, 1, ..., num_decode_reqs].
qo_indptr_decode = torch.arange(
num_decode_reqs + 1, dtype=torch.int32, device=qkv.device
)
out_q_decode = output[:num_decode_tokens]
if self.use_fp8:
_, q_scale, split_k_flag = hpc.rope_norm_store_kv_fp8(
key_cache=kv_cache[:, 0],
value_cache=kv_cache[:, 1],
qkv=qkv_decode,
cos_sin=self.cos_sin_cache,
num_seqlen_per_req=num_seq_kvcache,
q_index=qo_indptr_decode,
kvcache_indices=block_table_decode,
is_prefill=False,
k_scale=k_scale,
v_scale=v_scale,
quant_policy=QUANT_POLICY_DQSKV,
max_seqlens=1,
q_norm_weight=q_norm_weight,
k_norm_weight=k_norm_weight,
qk_norm_policy=self.qk_norm_policy,
out_q=out_q_decode,
)
attn_metadata.hpc_decode_q_scale = q_scale
if split_k_flag is not None:
attn_metadata.hpc_split_k_flag = split_k_flag
else:
hpc.rope_norm_store_kv(
kv_cache[:, 0],
kv_cache[:, 1],
qkv_decode,
self.cos_sin_cache,
num_seq_kvcache,
qo_indptr_decode,
block_table_decode,
False, # is_prefill
q_norm_weight=q_norm_weight,
k_norm_weight=k_norm_weight,
out_q=out_q_decode,
qk_norm_policy=self.qk_norm_policy,
)
@@ -174,31 +174,31 @@ def chunk_gated_delta_rule_cutedsl(
When ``core_attn_out`` is provided, ``output`` is an unsqueezed view of
that buffer.
"""
q_3d = q.squeeze(0)
k_3d = k.squeeze(0)
v_3d = v.squeeze(0)
g_2d = g.squeeze(0)
beta_2d = beta.squeeze(0)
q = q.squeeze(0)
k = k.squeeze(0)
v = v.squeeze(0)
g = g.squeeze(0)
beta = beta.squeeze(0)
_, _, head_k_dim = k_3d.shape
_, num_v_heads, head_v_dim = v_3d.shape
_, _, K_dim = k.shape
_, num_v_heads, V_dim = v.shape
chunk_size = 64
upper_bound_chunks = chunk_indices.shape[0]
pad_t = upper_bound_chunks * chunk_size
total_chunks_ptr = chunk_offsets[-1:]
g_cu = torch.empty_like(g_2d, dtype=torch.float32)
u = q_3d.new_empty(pad_t, num_v_heads, head_v_dim)
w = q_3d.new_empty(pad_t, num_v_heads, head_k_dim)
g_cu = torch.empty_like(g, dtype=torch.float32)
u = q.new_empty(pad_t, num_v_heads, V_dim)
w = q.new_empty(pad_t, num_v_heads, K_dim)
num_sms = torch.cuda.get_device_properties(q.device).multi_processor_count
kkt_inv_uw_cutedsl(
k_3d,
v_3d,
k,
v,
u,
w,
g_2d,
beta_2d,
g,
beta,
g_cu,
cu_seqlens,
chunk_indices,
@@ -206,16 +206,11 @@ def chunk_gated_delta_rule_cutedsl(
num_sms=num_sms,
)
h = k_3d.new_empty(
upper_bound_chunks,
num_v_heads,
head_v_dim,
head_k_dim,
)
v_new = q_3d.new_empty(pad_t, num_v_heads, head_v_dim)
h = k.new_empty(upper_bound_chunks, num_v_heads, V_dim, K_dim)
v_new = q.new_empty(pad_t, num_v_heads, V_dim)
final_state = torch.empty_like(initial_state)
h_cutedsl(
k_3d,
k,
u,
w,
v_new,
@@ -227,12 +222,12 @@ def chunk_gated_delta_rule_cutedsl(
chunk_offsets,
)
output = core_attn_out if core_attn_out is not None else torch.empty_like(v_3d)
scale = head_k_dim**-0.5
output = core_attn_out if core_attn_out is not None else torch.empty_like(v)
scale = K_dim**-0.5
o_cutedsl(
q_3d,
k_3d,
v_new.view(upper_bound_chunks, chunk_size, num_v_heads, head_v_dim),
q,
k,
v_new,
h,
g_cu,
output,

Some files were not shown because too many files have changed in this diff Show More