forked from Karylab-cklius/vllm
Compare commits
87
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b31b57a391 | ||
|
|
d8f483dc30 | ||
|
|
dc148dc4d7 | ||
|
|
7cf7cbcd95 | ||
|
|
c231d1f290 | ||
|
|
db808b3961 | ||
|
|
00ebf19cca | ||
|
|
ded6676458 | ||
|
|
7a327f0b4f | ||
|
|
1ab9522935 | ||
|
|
0fc2512094 | ||
|
|
62c7d8009f | ||
|
|
ab80b3dff4 | ||
|
|
91055efd36 | ||
|
|
3675bcff67 | ||
|
|
bdbd7278b6 | ||
|
|
5dc36a4fa5 | ||
|
|
aab7af0bcb | ||
|
|
536047755e | ||
|
|
1907d3854a | ||
|
|
ea9ddf59fc | ||
|
|
8cf7c4d8ad | ||
|
|
8e9d70fdd5 | ||
|
|
364ee36af1 | ||
|
|
06fae69114 | ||
|
|
14f8660a18 | ||
|
|
aed541def4 | ||
|
|
2bc20e8aba | ||
|
|
8cc242335d | ||
|
|
ba22cb6765 | ||
|
|
81bcced482 | ||
|
|
fb42e5219e | ||
|
|
0feca7ffa8 | ||
|
|
97b5ce5c39 | ||
|
|
4236514098 | ||
|
|
e45c8a9f4b | ||
|
|
b153dd3f28 | ||
|
|
930f8dc0a1 | ||
|
|
a16dbd5b85 | ||
|
|
bec232a914 | ||
|
|
b5c9e1ac33 | ||
|
|
ae2c4f3db7 | ||
|
|
fca432e60a | ||
|
|
af1ee8c475 | ||
|
|
5b4cb69523 | ||
|
|
9fc0c08026 | ||
|
|
f2b5fabb23 | ||
|
|
b8cb75b149 | ||
|
|
43916891b2 | ||
|
|
cda05ee8c4 | ||
|
|
77654d080c | ||
|
|
75698e60b3 | ||
|
|
8632c884dc | ||
|
|
c3734e8334 | ||
|
|
53f7553f09 | ||
|
|
4eb227992a | ||
|
|
ebcf511ec3 | ||
|
|
8fc1b2d046 | ||
|
|
5316638a5e | ||
|
|
61ab70ec3b | ||
|
|
a309d4fe60 | ||
|
|
72f639927f | ||
|
|
8ad4a01825 | ||
|
|
7be582697b | ||
|
|
030c9523bd | ||
|
|
4708292d48 | ||
|
|
debec6440b | ||
|
|
c8fb2963bd | ||
|
|
379acd4e4f | ||
|
|
07d33e575b | ||
|
|
36bbecd643 | ||
|
|
6149187a4c | ||
|
|
49e28e8e91 | ||
|
|
0ca39c4f1f | ||
|
|
6185d73882 | ||
|
|
bc8481af09 | ||
|
|
59575da46d | ||
|
|
3483240b7e | ||
|
|
eddfd4cf21 | ||
|
|
a4e3cb40d0 | ||
|
|
ab132ee98b | ||
|
|
e186107870 | ||
|
|
0e207dac78 | ||
|
|
9e86352c60 | ||
|
|
5051698e41 | ||
|
|
db28ae2d07 | ||
|
|
f6bb8682ee |
@@ -0,0 +1,25 @@
|
||||
group: Zen CPU
|
||||
depends_on: []
|
||||
steps:
|
||||
- label: Zen-CPU-Kernel Tests
|
||||
depends_on: []
|
||||
soft_fail: false
|
||||
device: zen5
|
||||
no_plugin: true
|
||||
source_file_dependencies:
|
||||
- setup.py
|
||||
- docker/Dockerfile.cpu
|
||||
- requirements/cpu.txt
|
||||
- requirements/build/cpu.txt
|
||||
- csrc/cpu/
|
||||
- cmake/cpu_extension.cmake
|
||||
- CMakeLists.txt
|
||||
- vllm/model_executor/layers/utils.py
|
||||
- vllm/platforms/cpu.py
|
||||
- vllm/platforms/zen_cpu.py
|
||||
- vllm/platforms/__init__.py
|
||||
- tests/model_executor/test_cpu_unquantized_gemm_dispatch.py
|
||||
commands:
|
||||
- |
|
||||
bash .buildkite/scripts/hardware_ci/run-zen-cpu-test.sh 20m "
|
||||
pytest -x -v -s tests/model_executor/test_cpu_unquantized_gemm_dispatch.py"
|
||||
@@ -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'
|
||||
|
||||
@@ -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"'
|
||||
@@ -125,7 +125,5 @@ steps:
|
||||
pip install open-clip-torch --no-deps &&
|
||||
cd tests &&
|
||||
pytest -v -s models/multimodal/processing/test_tensor_schema.py
|
||||
--deselect "tests/models/multimodal/processing/test_tensor_schema.py::test_model_tensor_schema[mistralai/Mistral-Large-3-675B-Instruct-2512-NVFP4]"
|
||||
--deselect "tests/models/multimodal/processing/test_tensor_schema.py::test_model_tensor_schema[Qwen/Qwen2.5-Omni-7B-AWQ]"
|
||||
--num-shards=$$BUILDKITE_PARALLEL_JOB_COUNT --shard-id=$$BUILDKITE_PARALLEL_JOB'
|
||||
parallelism: 4
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
#!/bin/bash
|
||||
|
||||
# This script build the CPU docker image and run the offline inference inside the container.
|
||||
# It serves a sanity check for compilation and basic model usage.
|
||||
set -euox pipefail
|
||||
|
||||
# allow to bind to different cores
|
||||
CORE_RANGE=${CORE_RANGE:-48-95}
|
||||
NUMA_NODE=${NUMA_NODE:-1}
|
||||
IMAGE_NAME="zen-cpu-test-$NUMA_NODE"
|
||||
TIMEOUT_VAL=$1
|
||||
TEST_COMMAND=$2
|
||||
|
||||
# building the docker image
|
||||
echo "--- :docker: Building Docker image"
|
||||
docker build --progress plain --tag "$IMAGE_NAME" --target vllm-zen-test -f docker/Dockerfile.cpu .
|
||||
|
||||
# Run the image, setting --shm-size=4g for tensor parallel.
|
||||
docker run --rm --cpuset-cpus="$CORE_RANGE" --cpuset-mems="$NUMA_NODE" -v ~/.cache/huggingface:/root/.cache/huggingface --privileged=true -e HF_TOKEN -e VLLM_CPU_KVCACHE_SPACE=16 -e VLLM_CPU_CI_ENV=1 -e VLLM_CPU_SIM_MULTI_NUMA=1 --shm-size=4g "$IMAGE_NAME" \
|
||||
timeout "$TIMEOUT_VAL" bash -c "set -euox pipefail; echo \"--- Print packages\"; pip list; echo \"--- Running tests\"; ${TEST_COMMAND}"
|
||||
+26
-21
@@ -845,10 +845,12 @@ steps:
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/entrypoints/serve
|
||||
- tests/entrypoints/scale_out
|
||||
commands:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -v -s entrypoints/serve --ignore=entrypoints/serve/dev/rpc
|
||||
- PYTHONPATH=/vllm-workspace pytest -v -s entrypoints/serve/dev/rpc
|
||||
- pytest -v -s entrypoints/scale_out
|
||||
|
||||
- label: Entrypoints Integration (API Server OpenAI - Part 1) # TBD
|
||||
timeout_in_minutes: 180
|
||||
@@ -1196,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
|
||||
@@ -2390,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
|
||||
@@ -2559,10 +2561,12 @@ steps:
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/entrypoints/serve
|
||||
- tests/entrypoints/scale_out
|
||||
commands:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -v -s entrypoints/serve --ignore=entrypoints/serve/dev/rpc
|
||||
- PYTHONPATH=/vllm-workspace pytest -v -s entrypoints/serve/dev/rpc
|
||||
- pytest -v -s entrypoints/scale_out
|
||||
|
||||
- label: Entrypoints Integration (API Server OpenAI - Part 1) # TBD
|
||||
timeout_in_minutes: 180
|
||||
@@ -2824,6 +2828,7 @@ steps:
|
||||
- rocm-smi
|
||||
- python3 examples/basic/offline_inference/chat.py --attention-backend TRITON_ATTN
|
||||
- pytest -v -s tests/kernels/attention/test_attention_selector.py
|
||||
- pytest -v -s tests/kernels/attention/test_rocm_aiter_mla_decode_metadata.py
|
||||
|
||||
- label: Kernels Attention Test %N # TBD
|
||||
timeout_in_minutes: 60
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -40,10 +42,12 @@ steps:
|
||||
source_file_dependencies:
|
||||
- vllm/
|
||||
- tests/entrypoints/serve
|
||||
- tests/entrypoints/scale_out
|
||||
commands:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
- pytest -v -s entrypoints/serve --ignore=entrypoints/serve/dev/rpc
|
||||
- PYTHONPATH=/vllm-workspace pytest -v -s entrypoints/serve/dev/rpc
|
||||
- pytest -v -s entrypoints/scale_out
|
||||
mirror:
|
||||
amd:
|
||||
device: mi325_1
|
||||
|
||||
@@ -46,7 +46,7 @@ steps:
|
||||
- vllm/v1/engine/
|
||||
- tests/utils.py
|
||||
# - tests/entrypoints/serve/dev/rpc/test_collective_rpc.py
|
||||
- tests/entrypoints/serve/disagg/test_serving_tokens.py
|
||||
- tests/entrypoints/scale_out/token_in_token_out/test_serving_tokens.py
|
||||
- tests/entrypoints/serve/instrumentator/test_basic.py
|
||||
- tests/entrypoints/serve/instrumentator/test_metrics.py
|
||||
# - tests/entrypoints/serve/dev/test_sleep.py
|
||||
@@ -55,7 +55,7 @@ steps:
|
||||
- export VLLM_WORKER_MULTIPROC_METHOD=spawn
|
||||
# - pytest -v -s entrypoints/serve/dev/rpc/test_collective_rpc.py
|
||||
- pytest -v -s entrypoints/serve/instrumentator/test_basic.py -k "not show_version and not server_load"
|
||||
- pytest -v -s entrypoints/serve/disagg/test_serving_tokens.py -k "not stream and not lora and not test_generate_logprobs and not stop_string_workflow"
|
||||
- pytest -v -s entrypoints/scale_out/token_in_token_out/test_serving_tokens.py -k "not stream and not lora and not test_generate_logprobs and not stop_string_workflow"
|
||||
- pytest -v -s entrypoints/serve/instrumentator/test_metrics.py -k "text and not show and not run_batch and not test_metrics_counts and not test_metrics_exist"
|
||||
# - pytest -v -s entrypoints/serve/dev/test_sleep.py
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -48,7 +48,7 @@ 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/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||
- uses: actions/setup-python@83679a892e2d95755f2dac6acb0bfd1e9ac5d548 # v6.1.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -140,6 +140,21 @@ if(Python_VERSION VERSION_GREATER_EQUAL "3.11")
|
||||
WITH_SOABI)
|
||||
endif()
|
||||
|
||||
#
|
||||
# fs_io extension (pure CXX; must stay above the non-CUDA device branch
|
||||
# so CPU builds define the target before the early return).
|
||||
# GIL-releasing filesystem helpers for FileSystemTierManager.
|
||||
#
|
||||
if(Python_VERSION VERSION_GREATER_EQUAL "3.11")
|
||||
define_extension_target(
|
||||
fs_io_C
|
||||
DESTINATION vllm
|
||||
LANGUAGE CXX
|
||||
SOURCES csrc/fs_io.cpp
|
||||
USE_SABI 3.11
|
||||
WITH_SOABI)
|
||||
endif()
|
||||
|
||||
#
|
||||
# Forward the non-CUDA device extensions to external CMake scripts.
|
||||
#
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
#include <Python.h>
|
||||
|
||||
#include <unistd.h>
|
||||
|
||||
#include <vector>
|
||||
|
||||
extern "C" {
|
||||
|
||||
static void _batch_lookup(const std::vector<const char*>& paths,
|
||||
std::vector<int>& exists_flags) {
|
||||
for (size_t i = 0; i < paths.size(); i++) {
|
||||
exists_flags[i] = (access(paths[i], F_OK) == 0) ? 1 : 0;
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Check file existence for a batch of paths.
|
||||
/// @param paths list[str] – absolute paths to check.
|
||||
/// @return list[bool] – True if the corresponding path exists, False otherwise.
|
||||
/// @note Releases the GIL for the entire batch. File existence via access(2).
|
||||
static PyObject* batch_lookup(PyObject* /*self*/, PyObject* args) {
|
||||
PyObject* path_list;
|
||||
if (!PyArg_ParseTuple(args, "O!", &PyList_Type, &path_list)) {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
const Py_ssize_t n = PyList_Size(path_list);
|
||||
std::vector<const char*> paths(n);
|
||||
for (Py_ssize_t i = 0; i < n; i++) {
|
||||
paths[i] = PyUnicode_AsUTF8AndSize(PyList_GetItem(path_list, i), nullptr);
|
||||
if (paths[i] == nullptr) {
|
||||
return nullptr;
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<int> exists_flags(n);
|
||||
{
|
||||
Py_BEGIN_ALLOW_THREADS _batch_lookup(paths, exists_flags);
|
||||
Py_END_ALLOW_THREADS
|
||||
}
|
||||
|
||||
PyObject* result = PyList_New(n);
|
||||
if (result == nullptr) {
|
||||
return nullptr;
|
||||
}
|
||||
for (Py_ssize_t i = 0; i < n; i++) {
|
||||
PyList_SetItem(result, i, PyBool_FromLong(exists_flags[i]));
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
static PyMethodDef fs_io_C_methods[] = {
|
||||
{"batch_lookup", batch_lookup, METH_VARARGS,
|
||||
"batch_lookup(paths: list[str]) -> list[bool]\n"
|
||||
"\n"
|
||||
"Check file existence for a batch of paths."},
|
||||
{nullptr, nullptr, 0, nullptr},
|
||||
};
|
||||
|
||||
static struct PyModuleDef fs_io_C_module = {
|
||||
PyModuleDef_HEAD_INIT, "fs_io_C", "Filesystem helpers for KV offload", -1,
|
||||
fs_io_C_methods,
|
||||
};
|
||||
|
||||
PyMODINIT_FUNC PyInit_fs_io_C(void) { return PyModule_Create(&fs_io_C_module); }
|
||||
|
||||
} // extern "C"
|
||||
+1
-1
@@ -793,7 +793,7 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
# Install FlashInfer JIT cache (requires CUDA-version-specific index URL)
|
||||
# https://docs.flashinfer.ai/installation.html
|
||||
# From versions.json: .flashinfer.version
|
||||
ARG FLASHINFER_VERSION=0.6.12
|
||||
ARG FLASHINFER_VERSION=0.6.13
|
||||
RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
uv pip install --system flashinfer-jit-cache==${FLASHINFER_VERSION} \
|
||||
--index-url https://flashinfer.ai/whl/cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d '.')
|
||||
|
||||
+36
-19
@@ -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
|
||||
@@ -328,3 +318,30 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install "vllm[zen]"
|
||||
|
||||
ENTRYPOINT ["vllm", "serve"]
|
||||
|
||||
######################### ZEN CPU TEST IMAGE #########################
|
||||
FROM vllm-openai-zen AS vllm-zen-test
|
||||
|
||||
WORKDIR /vllm-workspace
|
||||
|
||||
COPY --from=vllm-test-deps /vllm-workspace/requirements/test/cpu.txt requirements/test/cpu.txt
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install -r requirements/test/cpu.txt
|
||||
|
||||
ADD ./tests/ ./tests/
|
||||
ADD ./examples/ ./examples/
|
||||
ADD ./benchmarks/ ./benchmarks/
|
||||
ADD ./vllm/collect_env.py .
|
||||
ADD ./docker/ ./docker/
|
||||
ADD ./.buildkite/ ./.buildkite/
|
||||
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv pip install -e tests/vllm_test_utils
|
||||
|
||||
# enable fast downloads from hf (for testing)
|
||||
ENV HF_XET_HIGH_PERFORMANCE=1
|
||||
# increase timeout for hf downloads (for testing)
|
||||
ENV HF_HUB_DOWNLOAD_TIMEOUT=60
|
||||
|
||||
ENTRYPOINT []
|
||||
|
||||
@@ -257,13 +257,13 @@ RUN pip install setuptools==75.6.0 packaging==23.2 ninja==1.11.1.3 build==1.2.2.
|
||||
|
||||
|
||||
# build flashinfer for torch nightly from source around 10 mins
|
||||
# release version: v0.6.12
|
||||
# release version: v0.6.13
|
||||
# todo(elainewy): cache flashinfer build result for faster build
|
||||
ENV CCACHE_DIR=/root/.cache/ccache
|
||||
RUN --mount=type=cache,target=/root/.cache/ccache \
|
||||
--mount=type=cache,target=/root/.cache/uv \
|
||||
echo "git clone flashinfer..." \
|
||||
&& git clone --depth 1 --branch v0.6.12 --recursive https://github.com/flashinfer-ai/flashinfer.git \
|
||||
&& git clone --depth 1 --branch v0.6.13 --recursive https://github.com/flashinfer-ai/flashinfer.git \
|
||||
&& cd flashinfer \
|
||||
&& git submodule update --init --recursive \
|
||||
&& echo "finish git clone flashinfer..." \
|
||||
|
||||
@@ -68,7 +68,7 @@
|
||||
"default": "true"
|
||||
},
|
||||
"FLASHINFER_VERSION": {
|
||||
"default": "0.6.12"
|
||||
"default": "0.6.13"
|
||||
},
|
||||
"GDRCOPY_CUDA_VERSION": {
|
||||
"default": "12.8"
|
||||
|
||||
@@ -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 |
|
||||
@@ -220,7 +221,7 @@ MLA decode backends are selected using the standard
|
||||
| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. |
|
||||
| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | ------ | --------- | --- | --------------- | ------------ |
|
||||
| `CUTLASS_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 128 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
|
||||
| `FLASHINFER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | 10.x |
|
||||
| `FLASHINFER_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 10.x |
|
||||
| `FLASHINFER_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | 10.x |
|
||||
| `FLASHINFER_MLA_SPARSE_SM120` | bf16 | `auto`, `fp8`, `fp8_e4m3`, `fp8_ds_mla` | 64, 256 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | 12.x |
|
||||
| `FLASHMLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | 64 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | 9.x-10.x |
|
||||
|
||||
@@ -89,7 +89,7 @@ To be used with a particular `FusedMoEPrepareAndFinalizeModular` subclass, MoE k
|
||||
| gpt oss triton | standard | N/A | N/A | <sup>5</sup> | Y | Y | [`triton_kernel_fused_experts`][vllm.model_executor.layers.fused_moe.experts.gpt_oss_triton_kernels_moe.triton_kernel_fused_experts],</br>[`OAITritonExperts`][vllm.model_executor.layers.fused_moe.experts.gpt_oss_triton_kernels_moe.OAITritonExperts] |
|
||||
| marlin | standard,</br>batched | <sup>3</sup> / N/A | <sup>3</sup> / N/A | silu,</br>swigluoai | Y | Y | [`fused_marlin_moe`][vllm.model_executor.layers.fused_moe.experts.marlin_moe.fused_marlin_moe],</br>[`MarlinExperts`][vllm.model_executor.layers.fused_moe.experts.marlin_moe.MarlinExperts],</br>[`BatchedMarlinExperts`][vllm.model_executor.layers.fused_moe.experts.marlin_moe.BatchedMarlinExperts] |
|
||||
| trtllm | standard | mxfp4,</br>nvfp4 | G(16),G(32) | <sup>5</sup> | N | Y | [`TrtLlmMxfp4ExpertsMonolithic`][vllm.model_executor.layers.fused_moe.experts.trtllm_mxfp4_moe.TrtLlmMxfp4ExpertsMonolithic],</br>[`TrtLlmMxfp4ExpertsModular`][vllm.model_executor.layers.fused_moe.experts.trtllm_mxfp4_moe.TrtLlmMxfp4ExpertsModular],</br>[`TrtLlmNvFp4ExpertsMonolithic`][vllm.model_executor.layers.fused_moe.experts.trtllm_nvfp4_moe.TrtLlmNvFp4ExpertsMonolithic],</br>[`TrtLlmNvfp4ExpertsModular`][vllm.model_executor.layers.fused_moe.experts.trtllm_nvfp4_moe.TrtLlmNvFp4ExpertsModular] |
|
||||
| hpc | standard | fp8 | G(128),T | silu | Y | Y | [`HPCExperts`][vllm.model_executor.layers.fused_moe.experts.hpc.HPCExperts] |
|
||||
| hpc | standard | fp8 | G(128),T | silu | Y | Y | [`HPCExperts`][vllm.model_executor.layers.fused_moe.hpc_moe.HPCExperts] |
|
||||
| rocm aiter moe | standard | mxfp4,</br>fp8 | G(32),G(128),A,T | silu, gelu,</br>swigluoai | Y | N | `rocm_aiter_fused_experts`,</br>`AiterExperts` |
|
||||
| cpu_fused_moe | standard | N/A | N/A | silu | N | N | [`CPUFusedMOE`][vllm.model_executor.layers.fused_moe.cpu_fused_moe.CPUFusedMOE] |
|
||||
| naive batched<sup>4</sup> | batched | int8,</br>fp8 | G,A,T | silu, gelu | <sup>6</sup> | Y | [`NaiveBatchedExperts`][vllm.model_executor.layers.fused_moe.experts.fused_batched_moe.NaiveBatchedExperts] |
|
||||
|
||||
@@ -13,5 +13,6 @@ vLLM's examples are organized into the following categories:
|
||||
- **[`rl/`](../../examples/rl)** – Reinforcement learning examples.
|
||||
- **[`deployment/`](../../examples/deployment)** – Examples for deploying vLLM in production.
|
||||
- **[`ray_serving/`](../../examples/ray_serving)** – Scalable serving using Ray.
|
||||
- **[`disaggregated/`](../../examples/disaggregated)** – Examples for disaggregated serving (separate prefill and decode), including various kv cache connectors (LMCache, Mooncake, FlexKV, P2P NCCL) and failure recovery.
|
||||
- **[`disaggregated/`](../../examples/disaggregated)** – Examples for Disaggregated P/D (Prefill/Decoding) inference, including various kv cache connectors (LMCache, Mooncake, FlexKV, P2P NCCL) and failure recovery.
|
||||
- **[`scale_out/`](../../examples/scale_out)** – Examples for Token In <> Token Out API Server.
|
||||
- **[`observability/`](../../examples/observability)** – Metrics, logging, tracing (OpenTelemetry), and dashboards (Grafana, Perses).
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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. | ✅︎ | ✅︎ |
|
||||
@@ -626,8 +626,6 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen
|
||||
| `Step3VLForConditionalGeneration` | Step3-VL | T + I<sup>+</sup> | `stepfun-ai/step3` | | ✅︎ |
|
||||
| `StepVLForConditionalGeneration` | Step3-VL-10B | T + I<sup>+</sup> | `stepfun-ai/Step3-VL-10B` | | ✅︎ |
|
||||
| `Step3p7ForConditionalGeneration` | Step-3.7-Flash | T + I<sup>+</sup> | `stepfun-ai/Step-3.7-Flash` | | ✅︎ |
|
||||
| `TarsierForConditionalGeneration` | Tarsier | T + I<sup>E+</sup> | `omni-search/Tarsier-7b`, `omni-search/Tarsier-34b` | | ✅︎ |
|
||||
| `Tarsier2ForConditionalGeneration`<sup>^</sup> | Tarsier2 | T + I<sup>E+</sup> + V<sup>E+</sup> | `omni-research/Tarsier2-Recap-7b`, `omni-research/Tarsier2-7b-0115` | | ✅︎ |
|
||||
| `UltravoxModel` | Ultravox | T + A<sup>E+</sup> | `fixie-ai/ultravox-v0_5-llama-3_2-1b` | ✅︎ | ✅︎ |
|
||||
| `UnlimitedOCRForCausalLM` | Unlimited-OCR | T + I<sup>+</sup> | `baidu/Unlimited-OCR`, etc. | ✅︎ | ✅︎ |
|
||||
|
||||
|
||||
@@ -119,9 +119,9 @@ For further details on profiling vLLM, please refer to [this page](../../contrib
|
||||
- `/ping` - SageMaker health check
|
||||
- `/invocations` - SageMaker-compatible endpoint (routes to the same inference functions as `/v1` endpoints)
|
||||
|
||||
## Disaggregated Everything
|
||||
## Scale-Out APIs
|
||||
|
||||
### Tokens IN <> Tokens OUT
|
||||
### Tokens IN <> Tokens OUT APIs
|
||||
|
||||
- `/inference/v1/generate` - Generate completions
|
||||
- `/abort_requests` - Abort in-flight requests (only when `--tokens-only` is also set)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1275,55 +1275,6 @@ def load_step_vl(question: str, image_urls: list[str]) -> ModelRequestData:
|
||||
)
|
||||
|
||||
|
||||
def load_tarsier(question: str, image_urls: list[str]) -> ModelRequestData:
|
||||
model_name = "omni-research/Tarsier-7b"
|
||||
|
||||
engine_args = EngineArgs(
|
||||
model=model_name,
|
||||
trust_remote_code=True,
|
||||
max_model_len=4096,
|
||||
limit_mm_per_prompt={"image": len(image_urls)},
|
||||
)
|
||||
|
||||
prompt = f"USER: {'<image>' * len(image_urls)}\n{question}\n ASSISTANT:"
|
||||
image_data = [fetch_image(url) for url in image_urls]
|
||||
|
||||
return ModelRequestData(
|
||||
engine_args=engine_args,
|
||||
prompt=prompt,
|
||||
image_data=image_data,
|
||||
)
|
||||
|
||||
|
||||
def load_tarsier2(question: str, image_urls: list[str]) -> ModelRequestData:
|
||||
model_name = "omni-research/Tarsier2-Recap-7b"
|
||||
|
||||
engine_args = EngineArgs(
|
||||
model=model_name,
|
||||
trust_remote_code=True,
|
||||
max_model_len=32768,
|
||||
limit_mm_per_prompt={"image": len(image_urls)},
|
||||
hf_overrides={
|
||||
"architectures": ["Tarsier2ForConditionalGeneration"],
|
||||
"model_type": "tarsier2",
|
||||
},
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n"
|
||||
f"<|im_start|>user\n<|vision_start|>{'<|image_pad|>' * len(image_urls)}"
|
||||
f"<|vision_end|>{question}<|im_end|>\n"
|
||||
"<|im_start|>assistant\n"
|
||||
)
|
||||
image_data = [fetch_image(url) for url in image_urls]
|
||||
|
||||
return ModelRequestData(
|
||||
engine_args=engine_args,
|
||||
prompt=prompt,
|
||||
image_data=image_data,
|
||||
)
|
||||
|
||||
|
||||
# GLM-4.1V
|
||||
def load_glm4_1v(question: str, image_urls: list[str]) -> ModelRequestData:
|
||||
model_name = "zai-org/GLM-4.1V-9B-Thinking"
|
||||
@@ -1507,8 +1458,6 @@ model_example_map = {
|
||||
"smolvlm": load_smolvlm,
|
||||
"step3": load_step3,
|
||||
"stepvl": load_step_vl,
|
||||
"tarsier": load_tarsier,
|
||||
"tarsier2": load_tarsier2,
|
||||
"glm4_1v": load_glm4_1v,
|
||||
"glm4_5v": load_glm4_5v,
|
||||
"glm4_5v_fp8": load_glm4_5v_fp8,
|
||||
|
||||
@@ -2347,65 +2347,6 @@ def run_step_vl(questions: list[str], modality: str) -> ModelRequestData:
|
||||
)
|
||||
|
||||
|
||||
# omni-research/Tarsier-7b
|
||||
def run_tarsier(questions: list[str], modality: str) -> ModelRequestData:
|
||||
assert modality == "image"
|
||||
model_name = "omni-research/Tarsier-7b"
|
||||
|
||||
engine_args = EngineArgs(
|
||||
model=model_name,
|
||||
trust_remote_code=True,
|
||||
max_model_len=4096,
|
||||
limit_mm_per_prompt={modality: 1},
|
||||
)
|
||||
prompts = [(f"USER: <image>\n{question} ASSISTANT:") for question in questions]
|
||||
|
||||
return ModelRequestData(
|
||||
engine_args=engine_args,
|
||||
prompts=prompts,
|
||||
)
|
||||
|
||||
|
||||
def run_tarsier2(questions: list[str], modality: str) -> ModelRequestData:
|
||||
model_name = "omni-research/Tarsier2-Recap-7b"
|
||||
|
||||
mm_limit = {"image": 1, "video": 1} if modality == "image+video" else {modality: 1}
|
||||
engine_args = EngineArgs(
|
||||
model=model_name,
|
||||
max_model_len=4096,
|
||||
hf_overrides={
|
||||
"architectures": ["Tarsier2ForConditionalGeneration"],
|
||||
"model_type": "tarsier2",
|
||||
},
|
||||
limit_mm_per_prompt=mm_limit,
|
||||
)
|
||||
|
||||
image_placeholder = "<|vision_start|><|image_pad|><|vision_end|>"
|
||||
video_placeholder = "<|vision_start|><|video_pad|><|vision_end|>"
|
||||
|
||||
if modality == "image":
|
||||
placeholder = image_placeholder
|
||||
elif modality == "video":
|
||||
placeholder = video_placeholder
|
||||
elif modality == "image+video":
|
||||
placeholder = image_placeholder + video_placeholder
|
||||
|
||||
prompts = [
|
||||
(
|
||||
"<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n"
|
||||
f"<|im_start|>user\n{placeholder}"
|
||||
f"{question}<|im_end|>\n"
|
||||
"<|im_start|>assistant\n"
|
||||
)
|
||||
for question in questions
|
||||
]
|
||||
|
||||
return ModelRequestData(
|
||||
engine_args=engine_args,
|
||||
prompts=prompts,
|
||||
)
|
||||
|
||||
|
||||
model_example_map = {
|
||||
"aria": run_aria,
|
||||
"aya_vision": run_aya_vision,
|
||||
@@ -2479,8 +2420,6 @@ model_example_map = {
|
||||
"smolvlm": run_smolvlm,
|
||||
"step3": run_step3,
|
||||
"stepvl": run_step_vl,
|
||||
"tarsier": run_tarsier,
|
||||
"tarsier2": run_tarsier2,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -8,10 +8,10 @@ torch==2.11.0
|
||||
torchaudio==2.11.0
|
||||
# These must be updated alongside torch
|
||||
torchvision==0.26.0 # Required for phi3v processor. See https://github.com/pytorch/vision?tab=readme-ov-file#installation for corresponding version
|
||||
PyNvVideoCodec==2.0.4
|
||||
PyNvVideoCodec==2.1.0
|
||||
# FlashInfer should be updated together with the Dockerfile
|
||||
flashinfer-python==0.6.12
|
||||
flashinfer-cubin==0.6.12
|
||||
flashinfer-python==0.6.13
|
||||
flashinfer-cubin==0.6.13
|
||||
apache-tvm-ffi==0.1.9
|
||||
tilelang==0.1.9
|
||||
nvidia-cudnn-frontend>=1.19.1
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
-r ../common.txt
|
||||
|
||||
# --- Test Infrastructure ---
|
||||
tblib
|
||||
pytest
|
||||
|
||||
+316
-4
@@ -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
|
||||
|
||||
Generated
+70
-20
@@ -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",
|
||||
@@ -5261,6 +5310,7 @@ dependencies = [
|
||||
"vllm-llm",
|
||||
"vllm-metrics",
|
||||
"vllm-text",
|
||||
"vllm-tokenizer",
|
||||
"zeromq",
|
||||
]
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -50,6 +50,7 @@ tokio.workspace = true
|
||||
tracing-subscriber.workspace = true
|
||||
uuid.workspace = true
|
||||
vllm-engine-core-client = { workspace = true, features = ["test-util"] }
|
||||
vllm-tokenizer = { workspace = true, features = ["test-utils"] }
|
||||
zeromq.workspace = true
|
||||
|
||||
[lints]
|
||||
|
||||
@@ -154,7 +154,8 @@ mod tests {
|
||||
use thiserror_ext::AsReport as _;
|
||||
use vllm_text::Prompt;
|
||||
use vllm_text::backend::hf::TokenizerSource;
|
||||
use vllm_text::tokenizer::{DynTokenizer, Tokenizer};
|
||||
use vllm_text::tokenizer::DynTokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::HfChatBackend;
|
||||
use crate::backend::{ChatBackend, LoadModelBackendsOptions, NewChatOutputProcessorOptions};
|
||||
@@ -196,32 +197,8 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
struct TestTokenizer;
|
||||
|
||||
impl Tokenizer for TestTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
_text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<Vec<u32>> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
_token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<String> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn test_tokenizer() -> DynTokenizer {
|
||||
Arc::new(TestTokenizer)
|
||||
Arc::new(TestTokenizer::new())
|
||||
}
|
||||
|
||||
fn backend_for_selection(
|
||||
|
||||
@@ -563,7 +563,7 @@ mod tests {
|
||||
|
||||
use llm_multimodal::TokenId;
|
||||
use vllm_engine_core_client::protocol::tensor::WireArrayData;
|
||||
use vllm_text::tokenizer::{IncrementalDecoder, Tokenizer, TokenizerError};
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::*;
|
||||
|
||||
@@ -574,60 +574,14 @@ mod tests {
|
||||
const LLAMA4_TILE_X_SEPARATOR_ID: u32 = 200093;
|
||||
const LLAMA4_TILE_Y_SEPARATOR_ID: u32 = 200094;
|
||||
|
||||
struct TestTokenizer;
|
||||
|
||||
impl Tokenizer for TestTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> std::result::Result<Vec<u32>, TokenizerError> {
|
||||
Ok(match text {
|
||||
"<|image|>" => vec![LLAMA4_IMAGE_ID],
|
||||
text => text.bytes().map(u32::from).collect(),
|
||||
})
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
_token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> std::result::Result<String, TokenizerError> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
match token {
|
||||
"<|image_start|>" => Some(LLAMA4_IMAGE_START_ID),
|
||||
"<|image_end|>" => Some(LLAMA4_IMAGE_END_ID),
|
||||
"<|image|>" => Some(LLAMA4_IMAGE_ID),
|
||||
"<|patch|>" => Some(LLAMA4_PATCH_ID),
|
||||
"<|tile_x_separator|>" => Some(LLAMA4_TILE_X_SEPARATOR_ID),
|
||||
"<|tile_y_separator|>" => Some(LLAMA4_TILE_Y_SEPARATOR_ID),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn id_to_token(&self, id: u32) -> Option<String> {
|
||||
match id {
|
||||
LLAMA4_IMAGE_START_ID => Some("<|image_start|>".to_string()),
|
||||
LLAMA4_IMAGE_END_ID => Some("<|image_end|>".to_string()),
|
||||
LLAMA4_IMAGE_ID => Some("<|image|>".to_string()),
|
||||
LLAMA4_PATCH_ID => Some("<|patch|>".to_string()),
|
||||
LLAMA4_TILE_X_SEPARATOR_ID => Some("<|tile_x_separator|>".to_string()),
|
||||
LLAMA4_TILE_Y_SEPARATOR_ID => Some("<|tile_y_separator|>".to_string()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn create_decode_stream(
|
||||
&self,
|
||||
_prompt_token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
_min_bytes_to_buffer: usize,
|
||||
) -> Box<dyn IncrementalDecoder + '_> {
|
||||
unreachable!("not used")
|
||||
}
|
||||
fn llama4_tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
.with_regular_token("<|image_start|>", LLAMA4_IMAGE_START_ID)
|
||||
.with_regular_token("<|image_end|>", LLAMA4_IMAGE_END_ID)
|
||||
.with_regular_token("<|image|>", LLAMA4_IMAGE_ID)
|
||||
.with_regular_token("<|patch|>", LLAMA4_PATCH_ID)
|
||||
.with_regular_token("<|tile_x_separator|>", LLAMA4_TILE_X_SEPARATOR_ID)
|
||||
.with_regular_token("<|tile_y_separator|>", LLAMA4_TILE_Y_SEPARATOR_ID)
|
||||
}
|
||||
|
||||
fn test_info(model_type: &str, config: serde_json::Value) -> MultimodalModelInfo {
|
||||
@@ -635,7 +589,7 @@ mod tests {
|
||||
model_id: format!("{model_type}-test"),
|
||||
model_type: Some(model_type.to_string()),
|
||||
config,
|
||||
tokenizer: TokenizerResolver(Arc::new(TestTokenizer)),
|
||||
tokenizer: TokenizerResolver(Arc::new(llama4_tokenizer())),
|
||||
};
|
||||
let spec = context
|
||||
.resolve_model_spec()
|
||||
|
||||
@@ -189,46 +189,19 @@ impl ChatOutputProcessor for DefaultChatOutputProcessor {
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use vllm_tokenizer::Tokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::DefaultChatOutputProcessor;
|
||||
use crate::Error;
|
||||
use crate::parser::ParserSelection;
|
||||
use crate::request::ChatRequest;
|
||||
|
||||
struct FakeTokenizer;
|
||||
|
||||
impl Tokenizer for FakeTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.chars().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(token_ids
|
||||
.iter()
|
||||
.map(|token_id| char::from_u32(*token_id).unwrap_or('\u{FFFD}'))
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
match token {
|
||||
"<|channel>" => Some(1),
|
||||
"<channel|>" => Some(2),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn tokenizer() -> Arc<FakeTokenizer> {
|
||||
Arc::new(FakeTokenizer)
|
||||
fn tokenizer() -> Arc<TestTokenizer> {
|
||||
Arc::new(
|
||||
TestTokenizer::new()
|
||||
.with_regular_token("<|channel>", 256)
|
||||
.with_regular_token("<channel|>", 257),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -1,14 +1,8 @@
|
||||
//! Harmony output tests share the upstream `openai-harmony` tiktoken cache.
|
||||
//!
|
||||
//! Use a file lock for tests that load the encoding so `cargo nextest` cannot
|
||||
//! start multiple processes that concurrently populate the same cache file.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use futures::executor::block_on;
|
||||
use futures::{TryStreamExt as _, stream};
|
||||
use openai_harmony::chat::{Message, Role};
|
||||
use serial_test::file_serial;
|
||||
use vllm_text::output::{DecodedLogprobs, DecodedPositionLogprobs, DecodedTextEvent, Finished};
|
||||
|
||||
use super::*;
|
||||
@@ -91,7 +85,6 @@ fn request_with_tools() -> ChatRequest {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[file_serial(harmony_tiktoken_cache)]
|
||||
fn interrupted_final_message_is_preserved() {
|
||||
let tokens = completion_tokens(&[text_message("final", "hello")]);
|
||||
let events = block_on(collect_events(
|
||||
@@ -127,7 +120,6 @@ fn interrupted_final_message_is_preserved() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[file_serial(harmony_tiktoken_cache)]
|
||||
fn eos_flush_preserves_trailing_replacement_text() {
|
||||
let mut tokens = completion_tokens(&[text_message("final", "Hi")]);
|
||||
tokens.pop();
|
||||
@@ -153,7 +145,6 @@ fn eos_flush_preserves_trailing_replacement_text() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[file_serial(harmony_tiktoken_cache)]
|
||||
fn interrupted_analysis_message_is_preserved() {
|
||||
let tokens = completion_tokens(&[text_message("analysis", "think")]);
|
||||
let events = block_on(collect_events(
|
||||
@@ -189,7 +180,6 @@ fn interrupted_analysis_message_is_preserved() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[file_serial(harmony_tiktoken_cache)]
|
||||
fn commentary_preamble_is_visible_but_commentary_tool_payload_is_not() {
|
||||
let tokens = completion_tokens(&[
|
||||
text_message("commentary", "Let me check."),
|
||||
@@ -217,7 +207,6 @@ fn commentary_preamble_is_visible_but_commentary_tool_payload_is_not() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[file_serial(harmony_tiktoken_cache)]
|
||||
fn multiple_messages_get_newline_separators() {
|
||||
let tokens = completion_tokens(&[
|
||||
text_message("analysis", "first think"),
|
||||
@@ -249,7 +238,6 @@ fn multiple_messages_get_newline_separators() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[file_serial(harmony_tiktoken_cache)]
|
||||
fn tool_calls_stream_arguments_and_finish_with_local_id_shape() {
|
||||
let tokens = completion_tokens(&[tool_message(
|
||||
"get_weather",
|
||||
@@ -302,7 +290,6 @@ fn tool_calls_stream_arguments_and_finish_with_local_id_shape() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[file_serial(harmony_tiktoken_cache)]
|
||||
fn semantic_events_precede_same_update_logprobs() {
|
||||
let tokens = completion_tokens(&[text_message("final", "hello")]);
|
||||
let events = block_on(collect_events(
|
||||
@@ -353,7 +340,6 @@ fn rejects_generic_parser_overrides() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[file_serial(harmony_tiktoken_cache)]
|
||||
fn allows_auto_auto_only() {
|
||||
validate_harmony_parser_overrides(&ParserSelection::Auto, &ParserSelection::Auto).unwrap();
|
||||
let _ = HarmonyChatOutputProcessor::new(&ChatRequest::for_test()).unwrap();
|
||||
|
||||
@@ -1,32 +1,9 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use vllm_tokenizer::Tokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::{ReasoningParserFactory, names};
|
||||
|
||||
struct FakeTokenizer;
|
||||
|
||||
impl Tokenizer for FakeTokenizer {
|
||||
fn encode(&self, text: &str, _add_special_tokens: bool) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.chars().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(token_ids
|
||||
.iter()
|
||||
.map(|token_id| char::from_u32(*token_id).unwrap_or('\u{FFFD}'))
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn factory_contains_and_lists_registered_parsers() {
|
||||
let factory = ReasoningParserFactory::new();
|
||||
@@ -107,7 +84,7 @@ fn factory_resolves_minimax_m3_before_generic_minimax() {
|
||||
|
||||
#[test]
|
||||
fn factory_rejects_unknown_parser_names() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(TestTokenizer::new());
|
||||
let factory = ReasoningParserFactory::new();
|
||||
let error = match factory.create("missing", tokenizer) {
|
||||
Ok(_) => panic!("expected parser lookup to fail"),
|
||||
|
||||
@@ -75,39 +75,14 @@ impl UnifiedParserFactory {
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use vllm_tokenizer::Tokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::{UnifiedParserFactory, names};
|
||||
|
||||
struct FakeTokenizer;
|
||||
|
||||
impl Tokenizer for FakeTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.chars().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(token_ids
|
||||
.iter()
|
||||
.map(|token_id| char::from_u32(*token_id).unwrap_or('\u{FFFD}'))
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
match token {
|
||||
"<|channel>" => Some(1),
|
||||
"<channel|>" => Some(2),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
fn tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
.with_regular_token("<|channel>", 256)
|
||||
.with_regular_token("<channel|>", 257)
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -119,6 +94,6 @@ mod tests {
|
||||
factory.resolve_name_for_model("google/gemma-4-27b-it"),
|
||||
Some(names::GEMMA4)
|
||||
);
|
||||
factory.create(names::GEMMA4, &[], Arc::new(FakeTokenizer)).unwrap();
|
||||
factory.create(names::GEMMA4, &[], Arc::new(tokenizer())).unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
+34
-103
@@ -21,15 +21,17 @@ use vllm_engine_core_client::protocol::{
|
||||
use vllm_engine_core_client::test_utils::{IpcNamespace, spawn_mock_engine_task};
|
||||
use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig};
|
||||
use vllm_llm::Llm;
|
||||
use vllm_text::tokenizer::{DynTokenizer, Tokenizer};
|
||||
use vllm_text::tokenizer::DynTokenizer;
|
||||
use vllm_text::{
|
||||
DecodedLogprobs, DecodedPositionLogprobs, DecodedPromptLogprobs, DecodedTokenLogprob, Prompt,
|
||||
TextBackend,
|
||||
};
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
use zeromq::prelude::{SocketRecv, SocketSend};
|
||||
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
|
||||
|
||||
const SPECIAL_STOP_TOKEN_ID: u32 = 256;
|
||||
const UNKNOWN_DECODE_TOKEN_ID: u32 = 10_000;
|
||||
|
||||
fn request_output(
|
||||
request_id: &str,
|
||||
@@ -158,45 +160,18 @@ async fn connect_chat_llm_with_ipc(
|
||||
struct FakeChatBackend {
|
||||
has_template: bool,
|
||||
model_id: String,
|
||||
tokenizer: DynTokenizer,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct FakeChatTokenizer;
|
||||
|
||||
impl Tokenizer for FakeChatTokenizer {
|
||||
fn encode(&self, text: &str, _add_special_tokens: bool) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.bytes().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
let bytes = token_ids
|
||||
.iter()
|
||||
.filter_map(|id| {
|
||||
if skip_special_tokens && *id == SPECIAL_STOP_TOKEN_ID {
|
||||
None
|
||||
} else {
|
||||
Some(*id as u8)
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
Ok(String::from_utf8_lossy(&bytes).into_owned())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
match token {
|
||||
"<think>" => Some(0xF001),
|
||||
"</think>" => Some(0xF002),
|
||||
"<|START_THINKING|>" => Some(0xF003),
|
||||
"<|END_THINKING|>" => Some(0xF004),
|
||||
"◁think▷" => Some(0xF005),
|
||||
"◁/think▷" => Some(0xF006),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
fn fake_chat_tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
.with_special_token("<stop>", SPECIAL_STOP_TOKEN_ID)
|
||||
.with_regular_token("<think>", 0xF001)
|
||||
.with_regular_token("</think>", 0xF002)
|
||||
.with_regular_token("<|START_THINKING|>", 0xF003)
|
||||
.with_regular_token("<|END_THINKING|>", 0xF004)
|
||||
.with_regular_token("◁think▷", 0xF005)
|
||||
.with_regular_token("◁/think▷", 0xF006)
|
||||
}
|
||||
|
||||
impl fmt::Debug for FakeChatBackend {
|
||||
@@ -210,6 +185,7 @@ impl FakeChatBackend {
|
||||
Self {
|
||||
has_template: true,
|
||||
model_id: "test-model".to_string(),
|
||||
tokenizer: Arc::new(fake_chat_tokenizer()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -217,6 +193,7 @@ impl FakeChatBackend {
|
||||
Self {
|
||||
has_template: false,
|
||||
model_id: "test-model".to_string(),
|
||||
tokenizer: Arc::new(fake_chat_tokenizer()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -224,13 +201,19 @@ impl FakeChatBackend {
|
||||
Self {
|
||||
has_template: true,
|
||||
model_id: model_id.into(),
|
||||
tokenizer: Arc::new(fake_chat_tokenizer()),
|
||||
}
|
||||
}
|
||||
|
||||
fn with_tokenizer(mut self, tokenizer: DynTokenizer) -> Self {
|
||||
self.tokenizer = tokenizer;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
impl TextBackend for FakeChatBackend {
|
||||
fn tokenizer(&self) -> DynTokenizer {
|
||||
Arc::new(FakeChatTokenizer)
|
||||
Arc::clone(&self.tokenizer)
|
||||
}
|
||||
|
||||
fn model_id(&self) -> &str {
|
||||
@@ -282,65 +265,6 @@ impl ChatRenderer for FakeChatBackend {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
struct FailingDecodeBackend {
|
||||
inner: FakeChatBackend,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct FailingDecodeTokenizer;
|
||||
|
||||
impl Tokenizer for FailingDecodeTokenizer {
|
||||
fn encode(&self, text: &str, add_special_tokens: bool) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
FakeChatTokenizer.encode(text, add_special_tokens)
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
if token_ids.contains(&(b'i' as u32)) {
|
||||
return Err(vllm_tokenizer::TokenizerError("decode failed".to_string()));
|
||||
}
|
||||
FakeChatTokenizer.decode(token_ids, skip_special_tokens)
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
FakeChatTokenizer.token_to_id(token)
|
||||
}
|
||||
}
|
||||
|
||||
impl TextBackend for FailingDecodeBackend {
|
||||
fn tokenizer(&self) -> DynTokenizer {
|
||||
Arc::new(FailingDecodeTokenizer)
|
||||
}
|
||||
|
||||
fn model_id(&self) -> &str {
|
||||
self.inner.model_id()
|
||||
}
|
||||
}
|
||||
|
||||
impl ChatBackend for FailingDecodeBackend {
|
||||
fn chat_renderer(&self) -> DynChatRenderer {
|
||||
Arc::new(self.clone())
|
||||
}
|
||||
|
||||
fn new_chat_output_processor(
|
||||
&self,
|
||||
_request: &mut ChatRequest,
|
||||
_options: NewChatOutputProcessorOptions<'_>,
|
||||
) -> vllm_chat::Result<DynChatOutputProcessor> {
|
||||
Ok(Box::new(DefaultChatOutputProcessor::plain_text_only()))
|
||||
}
|
||||
}
|
||||
|
||||
impl ChatRenderer for FailingDecodeBackend {
|
||||
fn render(&self, request: &ChatRequest) -> vllm_chat::Result<RenderedPrompt> {
|
||||
self.inner.render(request)
|
||||
}
|
||||
}
|
||||
|
||||
/// Skip `LogprobsDelta` events that carry only token_ids (no logprobs),
|
||||
/// returning the next semantically interesting event.
|
||||
async fn next_semantic<S>(stream: &mut S) -> Option<Result<ChatEvent, vllm_chat::Error>>
|
||||
@@ -738,7 +662,12 @@ async fn chat_stream_reports_decode_failure_as_error_event() {
|
||||
send_outputs(
|
||||
push,
|
||||
EngineCoreOutputs {
|
||||
outputs: vec![request_output("chat-4", vec![b'i' as u32], None, None)],
|
||||
outputs: vec![request_output(
|
||||
"chat-4",
|
||||
vec![UNKNOWN_DECODE_TOKEN_ID],
|
||||
None,
|
||||
None,
|
||||
)],
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
@@ -747,9 +676,8 @@ async fn chat_stream_reports_decode_failure_as_error_event() {
|
||||
},
|
||||
);
|
||||
|
||||
let backend: Arc<dyn ChatTextBackend> = Arc::new(FailingDecodeBackend {
|
||||
inner: FakeChatBackend::new(),
|
||||
});
|
||||
let backend: Arc<dyn ChatTextBackend> =
|
||||
Arc::new(FakeChatBackend::new().with_tokenizer(Arc::new(TestTokenizer::new())));
|
||||
let chat = connect_chat_llm_with_ipc(
|
||||
EngineCoreClientConfig::new_single(handshake_address),
|
||||
&ipc,
|
||||
@@ -769,7 +697,10 @@ async fn chat_stream_reports_decode_failure_as_error_event() {
|
||||
|
||||
match timeout(Duration::from_secs(2), stream.next()).await.unwrap() {
|
||||
Some(Err(vllm_chat::Error::Text(vllm_text::Error::Tokenizer(message)))) => {
|
||||
assert_eq!(message, "decode failed");
|
||||
assert_eq!(
|
||||
message,
|
||||
format!("test tokenizer cannot decode unknown token id {UNKNOWN_DECODE_TOKEN_ID}")
|
||||
);
|
||||
}
|
||||
other => panic!("unexpected event after close: {other:?}"),
|
||||
}
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
//! Text-level roundtrip tests for the real chat-template and output-processor pairing.
|
||||
//! Roundtrip tests for the real chat-template and output-processor pairing.
|
||||
//!
|
||||
//! The invariant under test is that a structured assistant message rendered as history can be
|
||||
//! parsed from the generated assistant completion and then rendered back to the exact same
|
||||
//! assistant-completion text.
|
||||
//! assistant completion.
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
@@ -18,6 +18,10 @@ use vllm_chat::{
|
||||
RendererSelection, load_model_backends,
|
||||
};
|
||||
use vllm_text::{DecodedTextEvent, Finished, Prompt};
|
||||
use vllm_tokenizer::Tokenizer;
|
||||
|
||||
const TEXT_COMPLETION_CHUNK_CHARS: usize = 7;
|
||||
const TOKEN_COMPLETION_CHUNK_TOKENS: usize = 1;
|
||||
|
||||
/// One model/parser configuration used to run the fixed roundtrip fixtures.
|
||||
#[derive(Clone)]
|
||||
@@ -191,14 +195,28 @@ impl RoundtripCase {
|
||||
sort_json_keys: false,
|
||||
}
|
||||
}
|
||||
|
||||
/// GPT-OSS Harmony token-id renderer and native Harmony output processor.
|
||||
fn gpt_oss() -> Self {
|
||||
Self {
|
||||
model_id: "openai/gpt-oss-20b",
|
||||
assistant_stop_suffix: "", // not applicable for token-id cases
|
||||
tool_call_parser: ParserSelection::Auto,
|
||||
reasoning_parser: ParserSelection::Auto,
|
||||
thinking_behavior: ThinkingBehavior::Always { value: true },
|
||||
json_fmt: compact_json_fmt(),
|
||||
sort_json_keys: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
macro_rules! roundtrip_tests {
|
||||
($($case:ident => [$($fixture:ident),* $(,)?]),+ $(,)?) => {
|
||||
($($case:ident => [$($(#[$fixture_attr:meta])* $fixture:ident),* $(,)?]),+ $(,)?) => {
|
||||
paste::paste! {
|
||||
$(
|
||||
$(
|
||||
#[tokio::test]
|
||||
$(#[$fixture_attr])*
|
||||
#[file_serial([<hf_ $case>])]
|
||||
async fn [<roundtrip_ $case _ $fixture>]() -> Result<()> {
|
||||
[<run_roundtrip_ $fixture>](RoundtripCase::$case()).await
|
||||
@@ -217,9 +235,9 @@ roundtrip_tests! {
|
||||
glm47 => [reasoning_and_content, tool_call_mix],
|
||||
seed_oss => [reasoning_and_content],
|
||||
step3p5 => [reasoning_and_content],
|
||||
|
||||
gemma4 => [tool_call_mix], // Gemma4 strips reasoning in history if there's no tool call
|
||||
kimi_k25 => [tool_call_mix], // Kimi K2.5 strips reasoning in history
|
||||
gpt_oss => [tool_call_mix], // Harmony strips reasoning in history if there's no tool call
|
||||
}
|
||||
|
||||
/// Run the fixed reasoning+content fixture for one model/parser case.
|
||||
@@ -421,10 +439,10 @@ struct RoundtripResult {
|
||||
parsed_message: AssistantMessage,
|
||||
/// Assistant-completion suffix cut from rendering the expected assistant as
|
||||
/// history.
|
||||
closed_completion: String,
|
||||
closed_completion: Prompt,
|
||||
/// Assistant-completion suffix cut after rendering the parsed assistant
|
||||
/// back as history.
|
||||
rerendered_closed_completion: String,
|
||||
rerendered_closed_completion: Prompt,
|
||||
}
|
||||
|
||||
/// Render, parse, and rerender one assistant turn through the production
|
||||
@@ -436,60 +454,59 @@ async fn run_roundtrip(
|
||||
assistant: AssistantMessage,
|
||||
) -> Result<RoundtripResult> {
|
||||
let renderer = backends.chat_backend.chat_renderer();
|
||||
let (prompt, closed_completion_text) =
|
||||
render_closed_completion(renderer.as_ref(), request, &assistant)?;
|
||||
let completion_body = closed_completion_text
|
||||
.strip_suffix(case.assistant_stop_suffix)
|
||||
.with_context(|| {
|
||||
format!(
|
||||
"closed assistant completion did not end with {:?}: {:?}",
|
||||
case.assistant_stop_suffix, closed_completion_text
|
||||
)
|
||||
})?;
|
||||
let rendered = render_closed_completion(renderer.as_ref(), request, &assistant)?;
|
||||
|
||||
let parsed_message =
|
||||
parse_completion(case, backends, request, &prompt, completion_body).await?;
|
||||
let (_, rerendered_closed_completion) =
|
||||
render_closed_completion(renderer.as_ref(), request, &parsed_message)?;
|
||||
let parsed_message = parse_completion(case, backends, request, &rendered).await?;
|
||||
let rerendered = render_closed_completion(renderer.as_ref(), request, &parsed_message)?;
|
||||
|
||||
Ok(RoundtripResult {
|
||||
parsed_message,
|
||||
closed_completion: closed_completion_text,
|
||||
rerendered_closed_completion,
|
||||
closed_completion: rendered.completion,
|
||||
rerendered_closed_completion: rerendered.completion,
|
||||
})
|
||||
}
|
||||
|
||||
/// Rendered prompt/completion artifacts at the renderer boundary.
|
||||
struct RenderedTurn {
|
||||
prompt: Prompt,
|
||||
completion: Prompt,
|
||||
}
|
||||
|
||||
/// Render `history` as a production prompt and `history + assistant` as closed
|
||||
/// history, then return the production prompt and assistant-completion suffix.
|
||||
fn render_closed_completion(
|
||||
renderer: &dyn vllm_chat::ChatRenderer,
|
||||
base_request: &ChatRequest,
|
||||
assistant: &AssistantMessage,
|
||||
) -> Result<(String, String)> {
|
||||
) -> Result<RenderedTurn> {
|
||||
let mut prompt_request = base_request.clone();
|
||||
prompt_request.chat_options.generation_prompt_mode = GenerationPromptMode::StartNewAssistant;
|
||||
let prompt = render_text(renderer, &prompt_request).context("failed to render prompt")?;
|
||||
let prompt = renderer.render(&prompt_request).context("failed to render prompt")?.prompt;
|
||||
|
||||
let mut full_request = base_request.clone();
|
||||
full_request.chat_options.generation_prompt_mode = GenerationPromptMode::NoGenerationPrompt;
|
||||
full_request.messages.push(ChatMessage::from(assistant.clone()));
|
||||
let full = render_text(renderer, &full_request).context("failed to render full prompt")?;
|
||||
let full = renderer.render(&full_request).context("failed to render full prompt")?.prompt;
|
||||
|
||||
ensure!(
|
||||
full.starts_with(&prompt),
|
||||
"full prompt must extend production prompt\nprompt: {prompt:?}\nfull: {full:?}"
|
||||
);
|
||||
let completion = full[prompt.len()..].to_string();
|
||||
let completion = match (&prompt, full) {
|
||||
(Prompt::Text(prompt), Prompt::Text(full)) => {
|
||||
ensure!(
|
||||
full.starts_with(prompt),
|
||||
"full prompt must extend production prompt\nprompt: {prompt:?}\nfull: {full:?}"
|
||||
);
|
||||
Prompt::Text(full[prompt.len()..].to_string())
|
||||
}
|
||||
(Prompt::TokenIds(prompt), Prompt::TokenIds(full)) => {
|
||||
ensure!(
|
||||
full.starts_with(prompt),
|
||||
"full prompt must extend production prompt\nprompt: {prompt:?}\nfull: {full:?}"
|
||||
);
|
||||
Prompt::TokenIds(full[prompt.len()..].to_vec())
|
||||
}
|
||||
(prompt, full) => bail!("prompt kind changed between renders: {prompt:?} vs {full:?}"),
|
||||
};
|
||||
|
||||
Ok((prompt, completion))
|
||||
}
|
||||
|
||||
/// Render one chat request and require a text prompt.
|
||||
fn render_text(renderer: &dyn vllm_chat::ChatRenderer, request: &ChatRequest) -> Result<String> {
|
||||
match renderer.render(request)?.prompt {
|
||||
Prompt::Text(text) => Ok(text),
|
||||
other => bail!("roundtrip tests expect text prompts, got {other:?}"),
|
||||
}
|
||||
Ok(RenderedTurn { prompt, completion })
|
||||
}
|
||||
|
||||
/// Feed one rendered assistant completion body into the real output processor
|
||||
@@ -498,13 +515,15 @@ async fn parse_completion(
|
||||
case: &RoundtripCase,
|
||||
backends: &vllm_chat::LoadedModelBackends,
|
||||
base_request: &ChatRequest,
|
||||
prompt: &str,
|
||||
completion_body: &str,
|
||||
rendered: &RenderedTurn,
|
||||
) -> Result<AssistantMessage> {
|
||||
let tokenizer = backends.text_backend.tokenizer();
|
||||
let prompt_token_ids = tokenizer
|
||||
.encode(prompt, base_request.add_special_tokens)
|
||||
.context("failed to encode rendered prompt")?;
|
||||
let prompt_token_ids = match &rendered.prompt {
|
||||
Prompt::Text(prompt) => tokenizer
|
||||
.encode(prompt, base_request.add_special_tokens)
|
||||
.context("failed to encode rendered prompt")?,
|
||||
Prompt::TokenIds(token_ids) => token_ids.clone(),
|
||||
};
|
||||
|
||||
let mut request = base_request.clone();
|
||||
let processor = backends.chat_backend.new_chat_output_processor(
|
||||
@@ -515,7 +534,12 @@ async fn parse_completion(
|
||||
},
|
||||
)?;
|
||||
|
||||
let decoded = decoded_completion_stream(prompt_token_ids, completion_body);
|
||||
let decoded = decoded_completion_stream(
|
||||
tokenizer.as_ref(),
|
||||
prompt_token_ids,
|
||||
&rendered.completion,
|
||||
case.assistant_stop_suffix,
|
||||
)?;
|
||||
let mut events = processor.process(decoded)?;
|
||||
|
||||
while let Some(event) = events.next().await {
|
||||
@@ -538,16 +562,46 @@ async fn parse_completion(
|
||||
/// split into small chunks to exercise streaming parser state across marker
|
||||
/// and JSON boundaries.
|
||||
fn decoded_completion_stream(
|
||||
tokenizer: &dyn Tokenizer,
|
||||
prompt_token_ids: Vec<u32>,
|
||||
completion_body: &str,
|
||||
) -> Pin<Box<dyn Stream<Item = vllm_chat::Result<DecodedTextEvent>> + Send>> {
|
||||
let prompt_token_count = prompt_token_ids.len();
|
||||
completion: &Prompt,
|
||||
assistant_stop_suffix: &str,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = vllm_chat::Result<DecodedTextEvent>> + Send>>> {
|
||||
let mut events = vec![DecodedTextEvent::Start {
|
||||
prompt_token_ids: Arc::from(prompt_token_ids.into_boxed_slice()),
|
||||
prompt_token_ids: Arc::from(prompt_token_ids.clone().into_boxed_slice()),
|
||||
prompt_logprobs: None,
|
||||
}];
|
||||
|
||||
let chunks = split_by_chars(completion_body, 7);
|
||||
let chunks = match completion {
|
||||
Prompt::Text(text) => {
|
||||
let body = text.strip_suffix(assistant_stop_suffix).with_context(|| {
|
||||
format!(
|
||||
"closed assistant completion did not end with {:?}: {:?}",
|
||||
assistant_stop_suffix, text
|
||||
)
|
||||
})?;
|
||||
split_by_chars(body, TEXT_COMPLETION_CHUNK_CHARS)
|
||||
.into_iter()
|
||||
.map(|delta| DecodedCompletionChunk {
|
||||
delta,
|
||||
token_ids: Vec::new(), // unused for text-level roundtrip cases
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
Prompt::TokenIds(token_ids) => {
|
||||
ensure!(
|
||||
assistant_stop_suffix.is_empty(),
|
||||
"token-id roundtrip cases do not support text stop suffixes"
|
||||
);
|
||||
incremental_decode_chunks(
|
||||
tokenizer,
|
||||
&prompt_token_ids,
|
||||
token_ids,
|
||||
TOKEN_COMPLETION_CHUNK_TOKENS,
|
||||
)?
|
||||
}
|
||||
};
|
||||
|
||||
if chunks.is_empty() {
|
||||
events.push({
|
||||
DecodedTextEvent::TextDelta {
|
||||
@@ -555,11 +609,7 @@ fn decoded_completion_stream(
|
||||
token_ids: Vec::new(),
|
||||
logprobs: None,
|
||||
finished: Some(Finished {
|
||||
usage: vllm_llm::TokenUsage {
|
||||
prompt_token_count: 0,
|
||||
output_token_count: 0,
|
||||
cached_token_count: 0,
|
||||
},
|
||||
usage: Default::default(),
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
}),
|
||||
@@ -569,24 +619,26 @@ fn decoded_completion_stream(
|
||||
let last_index = chunks.len() - 1;
|
||||
for (index, chunk) in chunks.into_iter().enumerate() {
|
||||
let finished = (index == last_index).then(|| Finished {
|
||||
usage: vllm_llm::TokenUsage {
|
||||
prompt_token_count,
|
||||
output_token_count: completion_body.chars().count(),
|
||||
cached_token_count: 0,
|
||||
},
|
||||
usage: Default::default(),
|
||||
finish_reason: FinishReason::stop_eos(),
|
||||
kv_transfer_params: None,
|
||||
});
|
||||
events.push(DecodedTextEvent::TextDelta {
|
||||
delta: chunk,
|
||||
token_ids: Vec::new(),
|
||||
delta: chunk.delta,
|
||||
token_ids: chunk.token_ids,
|
||||
logprobs: None,
|
||||
finished,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
stream::iter(events).map(Ok).boxed()
|
||||
Ok(stream::iter(events).map(Ok).boxed())
|
||||
}
|
||||
|
||||
/// One decoded completion chunk fed into the output processor.
|
||||
struct DecodedCompletionChunk {
|
||||
delta: String,
|
||||
token_ids: Vec<u32>,
|
||||
}
|
||||
|
||||
/// Split text into chunks containing at most `chunk_chars` Unicode scalar
|
||||
@@ -612,6 +664,49 @@ fn split_by_chars(text: &str, chunk_chars: usize) -> Vec<String> {
|
||||
chunks
|
||||
}
|
||||
|
||||
/// Split token ids into chunks containing at most `chunk_size` ids.
|
||||
fn split_by_count(token_ids: &[u32], chunk_size: usize) -> Vec<Vec<u32>> {
|
||||
token_ids.chunks(chunk_size).map(<[u32]>::to_vec).collect()
|
||||
}
|
||||
|
||||
/// Decode token ids incrementally using the production tokenizer stream.
|
||||
fn incremental_decode_chunks(
|
||||
tokenizer: &dyn Tokenizer,
|
||||
prompt_token_ids: &[u32],
|
||||
token_ids: &[u32],
|
||||
chunk_size: usize,
|
||||
) -> Result<Vec<DecodedCompletionChunk>> {
|
||||
let mut decoder = tokenizer.create_decode_stream(prompt_token_ids, false, 0);
|
||||
let mut chunks = Vec::new();
|
||||
for chunk_token_ids in split_by_count(token_ids, chunk_size) {
|
||||
let mut delta = String::new();
|
||||
for token_id in chunk_token_ids.iter().copied() {
|
||||
decoder.push_token(token_id)?;
|
||||
while let Some(chunk) = decoder.next_chunk() {
|
||||
delta.push_str(&chunk);
|
||||
}
|
||||
}
|
||||
chunks.push(DecodedCompletionChunk {
|
||||
delta,
|
||||
token_ids: chunk_token_ids,
|
||||
});
|
||||
}
|
||||
|
||||
let (last_chunk, _) = decoder.flush(None)?;
|
||||
if let Some(last_chunk) = last_chunk {
|
||||
if let Some(delta) = chunks.last_mut() {
|
||||
delta.delta.push_str(&last_chunk);
|
||||
} else {
|
||||
chunks.push(DecodedCompletionChunk {
|
||||
delta: last_chunk,
|
||||
token_ids: Vec::new(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
Ok(chunks)
|
||||
}
|
||||
|
||||
/// Build a chat request fixture with parser-enabling tool-choice semantics.
|
||||
fn roundtrip_request(
|
||||
request_id: impl Into<String>,
|
||||
|
||||
+66
-1
@@ -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 {
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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()
|
||||
);
|
||||
|
||||
@@ -23,6 +23,7 @@ expect-test.workspace = true
|
||||
futures.workspace = true
|
||||
openai-protocol.workspace = true
|
||||
tool-parser.workspace = true
|
||||
vllm-tokenizer = { workspace = true, features = ["test-utils"] }
|
||||
|
||||
[[bench]]
|
||||
name = "deepseek_v3"
|
||||
|
||||
@@ -27,6 +27,10 @@ impl Tokenizer for BenchTokenizer {
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
Some(u32::MAX)
|
||||
}
|
||||
|
||||
fn id_to_token(&self, _id: u32) -> Option<String> {
|
||||
Some("\u{FFFD}".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// Bench-only adapter that exposes a unified parser through the tool-parser
|
||||
|
||||
@@ -49,11 +49,14 @@ mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::SeedOssReasoningParser;
|
||||
use crate::reasoning::{ReasoningParser, tests::FakeTokenizer};
|
||||
use crate::reasoning::{
|
||||
ReasoningParser,
|
||||
tests::{SEED_THINK_END_ID, SEED_THINK_START_ID, fake_tokenizer},
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn without_prompt_markers_expects_start_token() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = SeedOssReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("implicit reasoning</seed:think>answer").unwrap();
|
||||
@@ -66,10 +69,10 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn picks_up_prompt_start_boundary() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = SeedOssReasoningParser::new(tokenizer).unwrap();
|
||||
// Prompt prefills `<seed:think>` (id 10), opening reasoning before the stream.
|
||||
parser.initialize(&[10]).unwrap();
|
||||
// Prompt prefills `<seed:think>`, opening reasoning before the stream.
|
||||
parser.initialize(&[SEED_THINK_START_ID]).unwrap();
|
||||
|
||||
let delta = parser.push("reason</seed:think>answer").unwrap();
|
||||
assert_eq!(delta.reasoning.as_deref(), Some("reason"));
|
||||
@@ -78,10 +81,10 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn respects_prompt_end_boundary() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = SeedOssReasoningParser::new(tokenizer).unwrap();
|
||||
// Prompt already closed reasoning with `</seed:think>` (id 11).
|
||||
parser.initialize(&[11]).unwrap();
|
||||
// Prompt already closed reasoning with `</seed:think>`.
|
||||
parser.initialize(&[SEED_THINK_END_ID]).unwrap();
|
||||
|
||||
let delta = parser.push("answer").unwrap();
|
||||
assert_eq!(delta.reasoning, None);
|
||||
@@ -91,7 +94,7 @@ mod tests {
|
||||
#[test]
|
||||
fn handles_explicit_start_token() {
|
||||
// An explicit start delimiter must not leak into reasoning text.
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = SeedOssReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("<seed:think>reason</seed:think>answer").unwrap();
|
||||
@@ -103,7 +106,7 @@ mod tests {
|
||||
fn streams_explicit_start_token_across_pushes() {
|
||||
// Start token, reasoning body, end token, and content arrive in separate
|
||||
// streaming deltas.
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = SeedOssReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let mut reasoning = String::new();
|
||||
@@ -131,9 +134,9 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn handles_partial_delimiters_across_pushes() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = SeedOssReasoningParser::new(tokenizer).unwrap();
|
||||
parser.initialize(&[10]).unwrap();
|
||||
parser.initialize(&[SEED_THINK_START_ID]).unwrap();
|
||||
|
||||
// Closing delimiter `</seed:think>` arrives in two halves.
|
||||
let first = parser.push("reason</seed:").unwrap();
|
||||
|
||||
@@ -127,14 +127,17 @@ mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::Step3p5ReasoningParser;
|
||||
use crate::reasoning::{ReasoningParser, tests::FakeTokenizer};
|
||||
use crate::reasoning::{
|
||||
ReasoningParser,
|
||||
tests::{THINK_START_ID, fake_tokenizer},
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn picks_up_prompt_start_boundary() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
// Prompt prefills `<think>` (id 1), opening reasoning before the stream.
|
||||
parser.initialize(&[1]).unwrap();
|
||||
// Prompt prefills `<think>`, opening reasoning before the stream.
|
||||
parser.initialize(&[THINK_START_ID]).unwrap();
|
||||
|
||||
let delta = parser.push("This is a reasoning section</think>This is the rest").unwrap();
|
||||
assert_eq!(
|
||||
@@ -146,7 +149,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn handles_unterminated_reasoning() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let pushed = parser.push("<think>reason without end").unwrap();
|
||||
@@ -159,7 +162,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn handles_empty_input() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let pushed = parser.push("").unwrap();
|
||||
@@ -172,9 +175,9 @@ mod tests {
|
||||
fn complex_newline_pattern_trims_only_single_framing_newline_each_side() {
|
||||
// Only the immediately-adjacent framing `\n` is dropped on each side of
|
||||
// `</think>`; surrounding newlines remain part of reasoning/content.
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
parser.initialize(&[1]).unwrap();
|
||||
parser.initialize(&[THINK_START_ID]).unwrap();
|
||||
|
||||
let delta = parser
|
||||
.push("\n This is a \n reasoning section\n\n\n</think>\n\nThis is the rest")
|
||||
@@ -188,7 +191,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn drops_framing_newlines_in_single_push() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("<think>reason\n</think>\nanswer").unwrap();
|
||||
@@ -198,7 +201,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn drops_framing_newlines_across_pushes() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
// The trailing `\n` from the first push is held until we know whether
|
||||
@@ -219,7 +222,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn replays_held_newline_when_more_reasoning_follows() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let first = parser.push("<think>reason\n").unwrap();
|
||||
@@ -232,7 +235,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn finish_flushes_held_newline_in_unterminated_stream() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let first = parser.push("<think>reason\n").unwrap();
|
||||
@@ -245,7 +248,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn preserves_inner_newlines_in_reasoning() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("<think>line1\nline2</think>tail").unwrap();
|
||||
@@ -257,7 +260,7 @@ mod tests {
|
||||
fn trims_only_one_trailing_reasoning_newline() {
|
||||
// Only the single framing newline immediately before `</think>` is
|
||||
// dropped; earlier newlines in the reasoning body are preserved.
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("<think>reason\n\n</think>answer").unwrap();
|
||||
@@ -269,7 +272,7 @@ mod tests {
|
||||
fn drops_only_first_content_newline_after_transition() {
|
||||
// The leading-`\n` drop applies only to the first content delta after
|
||||
// `</think>`; later deltas pass through untouched.
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let first = parser.push("<think>reason</think>").unwrap();
|
||||
@@ -288,7 +291,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn passes_through_clean_boundary_without_framing_newlines() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("<think>reason</think>tail").unwrap();
|
||||
@@ -298,7 +301,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn handles_empty_reasoning_section() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Step3p5ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("<think></think>answer").unwrap();
|
||||
|
||||
@@ -1,54 +1,42 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use vllm_tokenizer::Tokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::{
|
||||
DeepSeekR1ReasoningParser, DelimitedReasoningParser, MiniMaxM3ReasoningParser,
|
||||
Qwen3ReasoningParser, ReasoningParser,
|
||||
};
|
||||
|
||||
pub(crate) struct FakeTokenizer;
|
||||
pub(crate) const THINK_START_ID: u32 = 256;
|
||||
pub(crate) const THINK_END_ID: u32 = 257;
|
||||
pub(crate) const START_THINKING_ID: u32 = 258;
|
||||
pub(crate) const END_THINKING_ID: u32 = 259;
|
||||
pub(crate) const MINIMAX_THINK_START_ID: u32 = 260;
|
||||
pub(crate) const MINIMAX_THINK_END_ID: u32 = 261;
|
||||
pub(crate) const SPECIAL_BOUNDARY_ID: u32 = 262;
|
||||
pub(crate) const MM_THINK_START_ID: u32 = 263;
|
||||
pub(crate) const MM_THINK_END_ID: u32 = 264;
|
||||
pub(crate) const SEED_THINK_START_ID: u32 = 265;
|
||||
pub(crate) const SEED_THINK_END_ID: u32 = 266;
|
||||
|
||||
impl Tokenizer for FakeTokenizer {
|
||||
fn encode(&self, text: &str, _add_special_tokens: bool) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.chars().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(token_ids
|
||||
.iter()
|
||||
.map(|token_id| char::from_u32(*token_id).unwrap_or('\u{FFFD}'))
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
match token {
|
||||
"<think>" => Some(1),
|
||||
"</think>" => Some(2),
|
||||
"<|START_THINKING|>" => Some(3),
|
||||
"<|END_THINKING|>" => Some(4),
|
||||
"◁think▷" => Some(5),
|
||||
"◁/think▷" => Some(6),
|
||||
"<mm:think>" => Some(8),
|
||||
"</mm:think>" => Some(9),
|
||||
"<seed:think>" => Some(10),
|
||||
"</seed:think>" => Some(11),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_special_id(&self, token_id: u32) -> bool {
|
||||
token_id == 7
|
||||
}
|
||||
pub(crate) fn fake_tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
.with_regular_token("<think>", THINK_START_ID)
|
||||
.with_regular_token("</think>", THINK_END_ID)
|
||||
.with_regular_token("<|START_THINKING|>", START_THINKING_ID)
|
||||
.with_regular_token("<|END_THINKING|>", END_THINKING_ID)
|
||||
.with_regular_token("◁think▷", MINIMAX_THINK_START_ID)
|
||||
.with_regular_token("◁/think▷", MINIMAX_THINK_END_ID)
|
||||
.with_special_token("<special-boundary>", SPECIAL_BOUNDARY_ID)
|
||||
.with_regular_token("<mm:think>", MM_THINK_START_ID)
|
||||
.with_regular_token("</mm:think>", MM_THINK_END_ID)
|
||||
.with_regular_token("<seed:think>", SEED_THINK_START_ID)
|
||||
.with_regular_token("</seed:think>", SEED_THINK_END_ID)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn delimited_content_only_stream() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser =
|
||||
DelimitedReasoningParser::new(tokenizer, "<think>", "</think>", false).unwrap();
|
||||
|
||||
@@ -60,7 +48,7 @@ fn delimited_content_only_stream() {
|
||||
|
||||
#[test]
|
||||
fn delimited_single_chunk_with_reasoning_and_content() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser =
|
||||
DelimitedReasoningParser::new(tokenizer, "<think>", "</think>", false).unwrap();
|
||||
|
||||
@@ -71,7 +59,7 @@ fn delimited_single_chunk_with_reasoning_and_content() {
|
||||
|
||||
#[test]
|
||||
fn delimited_partial_tokens_across_chunks() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser =
|
||||
DelimitedReasoningParser::new(tokenizer, "<think>", "</think>", false).unwrap();
|
||||
|
||||
@@ -83,10 +71,10 @@ fn delimited_partial_tokens_across_chunks() {
|
||||
|
||||
#[test]
|
||||
fn delimited_finish_flushes_buffer() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser =
|
||||
DelimitedReasoningParser::new(tokenizer, "<think>", "</think>", false).unwrap();
|
||||
parser.initialize(&[1]);
|
||||
parser.initialize(&[THINK_START_ID]);
|
||||
|
||||
let delta = parser.push("unfinished</thi");
|
||||
assert_eq!(delta.reasoning.as_deref(), Some("unfinished"));
|
||||
@@ -96,7 +84,7 @@ fn delimited_finish_flushes_buffer() {
|
||||
|
||||
#[test]
|
||||
fn qwen3_without_prompt_markers_expects_start_token() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Qwen3ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("reason</think>answer").unwrap();
|
||||
@@ -106,9 +94,9 @@ fn qwen3_without_prompt_markers_expects_start_token() {
|
||||
|
||||
#[test]
|
||||
fn qwen3_prompt_end_marker_starts_in_content() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Qwen3ReasoningParser::new(tokenizer).unwrap();
|
||||
parser.initialize(&[2]).unwrap();
|
||||
parser.initialize(&[THINK_END_ID]).unwrap();
|
||||
|
||||
let delta = parser.push("answer").unwrap();
|
||||
assert_eq!(delta.reasoning, None);
|
||||
@@ -117,7 +105,7 @@ fn qwen3_prompt_end_marker_starts_in_content() {
|
||||
|
||||
#[test]
|
||||
fn qwen3_tolerates_old_and_new_formats() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
|
||||
let mut old_parser = Qwen3ReasoningParser::new(tokenizer.clone()).unwrap();
|
||||
let old = old_parser.push("<think>reason</think>answer").unwrap();
|
||||
@@ -125,7 +113,7 @@ fn qwen3_tolerates_old_and_new_formats() {
|
||||
assert_eq!(old.content.as_deref(), Some("answer"));
|
||||
|
||||
let mut new_parser = Qwen3ReasoningParser::new(tokenizer).unwrap();
|
||||
new_parser.initialize(&[1]).unwrap();
|
||||
new_parser.initialize(&[THINK_START_ID]).unwrap();
|
||||
let new = new_parser.push("reason</think>answer").unwrap();
|
||||
assert_eq!(new.reasoning.as_deref(), Some("reason"));
|
||||
assert_eq!(new.content.as_deref(), Some("answer"));
|
||||
@@ -133,10 +121,10 @@ fn qwen3_tolerates_old_and_new_formats() {
|
||||
|
||||
#[test]
|
||||
fn qwen3_stops_scanning_at_last_special_token() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = Qwen3ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
parser.initialize(&[1, 7]).unwrap();
|
||||
parser.initialize(&[THINK_START_ID, SPECIAL_BOUNDARY_ID]).unwrap();
|
||||
|
||||
let delta = parser.push("answer").unwrap();
|
||||
assert_eq!(delta.reasoning, None);
|
||||
@@ -145,7 +133,7 @@ fn qwen3_stops_scanning_at_last_special_token() {
|
||||
|
||||
#[test]
|
||||
fn deepseek_r1_defaults_to_reasoning_without_prompt_boundary() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = DeepSeekR1ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("reason</think>answer").unwrap();
|
||||
@@ -155,10 +143,10 @@ fn deepseek_r1_defaults_to_reasoning_without_prompt_boundary() {
|
||||
|
||||
#[test]
|
||||
fn deepseek_r1_stops_scanning_at_last_special_token() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = DeepSeekR1ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
parser.initialize(&[2, 7]).unwrap();
|
||||
parser.initialize(&[THINK_END_ID, SPECIAL_BOUNDARY_ID]).unwrap();
|
||||
|
||||
let delta = parser.push("reason</think>answer").unwrap();
|
||||
assert_eq!(delta.reasoning.as_deref(), Some("reason"));
|
||||
@@ -167,7 +155,7 @@ fn deepseek_r1_stops_scanning_at_last_special_token() {
|
||||
|
||||
#[test]
|
||||
fn minimax_m3_handles_explicit_think_delimiters() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = MiniMaxM3ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("<mm:think>reason</mm:think>answer").unwrap();
|
||||
@@ -177,7 +165,7 @@ fn minimax_m3_handles_explicit_think_delimiters() {
|
||||
|
||||
#[test]
|
||||
fn minimax_m3_drops_leading_end_marker() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = MiniMaxM3ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("</mm:think>answer").unwrap();
|
||||
@@ -187,7 +175,7 @@ fn minimax_m3_drops_leading_end_marker() {
|
||||
|
||||
#[test]
|
||||
fn minimax_m3_preserves_non_leading_end_marker() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = MiniMaxM3ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
let delta = parser.push("XXX</mm:think>YYY").unwrap();
|
||||
@@ -197,7 +185,7 @@ fn minimax_m3_preserves_non_leading_end_marker() {
|
||||
|
||||
#[test]
|
||||
fn minimax_m3_drops_split_leading_end_marker() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = MiniMaxM3ReasoningParser::new(tokenizer).unwrap();
|
||||
|
||||
assert!(parser.push("</mm").unwrap().is_empty());
|
||||
@@ -208,9 +196,9 @@ fn minimax_m3_drops_split_leading_end_marker() {
|
||||
|
||||
#[test]
|
||||
fn minimax_m3_uses_prompt_prefilled_start_marker() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = MiniMaxM3ReasoningParser::new(tokenizer).unwrap();
|
||||
parser.initialize(&[8]).unwrap();
|
||||
parser.initialize(&[MM_THINK_START_ID]).unwrap();
|
||||
|
||||
let delta = parser.push("reason</mm:think>answer").unwrap();
|
||||
assert_eq!(delta.reasoning.as_deref(), Some("reason"));
|
||||
@@ -219,9 +207,9 @@ fn minimax_m3_uses_prompt_prefilled_start_marker() {
|
||||
|
||||
#[test]
|
||||
fn minimax_m3_uses_prompt_prefilled_end_marker() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(fake_tokenizer());
|
||||
let mut parser = MiniMaxM3ReasoningParser::new(tokenizer).unwrap();
|
||||
parser.initialize(&[9]).unwrap();
|
||||
parser.initialize(&[MM_THINK_END_ID]).unwrap();
|
||||
|
||||
let delta = parser.push("answer").unwrap();
|
||||
assert_eq!(delta.reasoning, None);
|
||||
|
||||
@@ -154,7 +154,7 @@ mod tests {
|
||||
"flag": true,
|
||||
"payload": { "nested": true },
|
||||
"items": [1, 2],
|
||||
"empty": null,
|
||||
"empty": "null",
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
@@ -166,7 +166,13 @@ impl JsonParamType {
|
||||
|
||||
// Typically, these types are already handled by checking the "type" field, but
|
||||
// we can also infer them from their characteristic fields if "type" is missing.
|
||||
if schema.contains_key("enum") {
|
||||
if let Some(values) = schema.get("enum").and_then(Value::as_array) {
|
||||
// Enum values are treated as strings, except that a `null` member
|
||||
// makes the parameter nullable (mirrors Python's enum type
|
||||
// inference), so a literal "null" coerces to JSON null.
|
||||
if values.iter().any(Value::is_null) {
|
||||
return Some(Self::one_of(vec![Self::String, Self::Null]));
|
||||
}
|
||||
return Some(Self::String);
|
||||
}
|
||||
if schema.contains_key("items") {
|
||||
@@ -277,9 +283,12 @@ impl JsonParamType {
|
||||
|
||||
/// Convert one parameter input to a normalized JSON value.
|
||||
fn convert_with_optional_schema(param_type: Option<&JsonParamType>, input: &ParamInput) -> Value {
|
||||
// For literal `null`, always convert to JSON null value.
|
||||
// Coerce the literal text `null` to JSON null, except for `string`-typed
|
||||
// params, where it must stay the string "null": a model emitting the literal
|
||||
// text "null" for a string field means the string, not a missing value.
|
||||
if let ParamInput::Text(value) = input
|
||||
&& value.eq_ignore_ascii_case("null")
|
||||
&& param_type != Some(&JsonParamType::String)
|
||||
{
|
||||
return Value::Null;
|
||||
}
|
||||
@@ -685,21 +694,43 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn convert_params_preserves_null_for_known_param() {
|
||||
let schemas = ToolSchemas::from_tools(&[test_tool(
|
||||
"convert",
|
||||
json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"value": { "type": "string" }
|
||||
}
|
||||
}),
|
||||
)]);
|
||||
fn string_param_preserves_literal_null_text() {
|
||||
// A `string`-typed param whose value is the literal text "null"/"NULL"
|
||||
// must stay a string (the original case is preserved), rather than being
|
||||
// coerced to JSON null. Non-string types keep coercing "null" to null.
|
||||
let params = ToolSchema::from_schema(&json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": { "type": "string" },
|
||||
"count": { "type": "integer" },
|
||||
"anything": {}
|
||||
}
|
||||
}));
|
||||
|
||||
let converted = schemas
|
||||
.convert_params_with_schema("convert", vec![("value".to_string(), "NULL".to_string())]);
|
||||
assert_eq!(params.convert("name", text("null")), json!("null"));
|
||||
assert_eq!(params.convert("name", text("NULL")), json!("NULL"));
|
||||
// Non-string and schema-less params are unchanged: "null" -> null.
|
||||
assert_eq!(params.convert("count", text("null")), json!(null));
|
||||
assert_eq!(params.convert("anything", text("null")), json!(null));
|
||||
}
|
||||
|
||||
assert_eq!(converted.get("value"), Some(&json!(null)));
|
||||
#[test]
|
||||
fn nullable_enum_param_coerces_literal_null() {
|
||||
// An enum that includes `null` admits a null value, so a literal "null"
|
||||
// must coerce to JSON null (matching Python's `extract_types_from_schema`,
|
||||
// which infers `null` from the enum values), while a non-null enum keeps
|
||||
// "null" as a string.
|
||||
let params = ToolSchema::from_schema(&json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"mode": { "enum": [null, "auto"] },
|
||||
"color": { "enum": ["red", "green"] }
|
||||
}
|
||||
}));
|
||||
|
||||
assert_eq!(params.convert("mode", text("null")), json!(null));
|
||||
assert_eq!(params.convert("mode", text("auto")), json!("auto"));
|
||||
assert_eq!(params.convert("color", text("null")), json!("null"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -841,7 +872,7 @@ mod tests {
|
||||
"user_id": 42,
|
||||
"urgent": true,
|
||||
"note": "Please leave at front desk.",
|
||||
"nil": null,
|
||||
"nil": "NULL",
|
||||
"shipping": {
|
||||
"city": "Singapore",
|
||||
"zip": 18956
|
||||
|
||||
@@ -124,42 +124,17 @@ impl UnifiedParser for CombinedParser {
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use vllm_tokenizer::Tokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::CombinedParser;
|
||||
use crate::reasoning::{Qwen3ReasoningParser, ReasoningDelta, ReasoningParser};
|
||||
use crate::tool::{Qwen3XmlToolParser, Tool, ToolParser};
|
||||
use crate::unified::{UnifiedParser, UnifiedParserEvent, UnifiedParserOutput};
|
||||
|
||||
struct FakeTokenizer;
|
||||
|
||||
impl Tokenizer for FakeTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.chars().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(token_ids
|
||||
.iter()
|
||||
.map(|token_id| char::from_u32(*token_id).unwrap_or('\u{FFFD}'))
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
match token {
|
||||
"<think>" => Some(1),
|
||||
"</think>" => Some(2),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
fn tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
.with_regular_token("<think>", 256)
|
||||
.with_regular_token("</think>", 257)
|
||||
}
|
||||
|
||||
fn test_tools() -> Vec<Tool> {
|
||||
@@ -273,7 +248,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn combined_parser_emits_reasoning_and_text() {
|
||||
let tokenizer = Arc::new(FakeTokenizer);
|
||||
let tokenizer = Arc::new(tokenizer());
|
||||
let reasoning = Qwen3ReasoningParser::create(tokenizer).unwrap();
|
||||
let mut parser = CombinedParser::new(Some(reasoning), None);
|
||||
|
||||
|
||||
@@ -515,7 +515,7 @@ mod tests {
|
||||
|
||||
use serde_json::{Value, json};
|
||||
use thiserror_ext::AsReport;
|
||||
use vllm_tokenizer::Tokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
use winnow::combinator::{eof, terminated};
|
||||
use winnow::error::ErrMode;
|
||||
use winnow::prelude::*;
|
||||
@@ -527,66 +527,15 @@ mod tests {
|
||||
use crate::tool::Tool;
|
||||
use crate::unified::{UnifiedParserEvent, parsing_failed};
|
||||
|
||||
struct FakeTokenizer;
|
||||
const CHANNEL_START_ID: u32 = 256;
|
||||
const CHANNEL_END_ID: u32 = 257;
|
||||
const TURN_BOUNDARY_ID: u32 = 258;
|
||||
|
||||
impl Tokenizer for FakeTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.chars().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(token_ids
|
||||
.iter()
|
||||
.map(|token_id| char::from_u32(*token_id).unwrap_or('\u{FFFD}'))
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
match token {
|
||||
CHANNEL_START => Some(100),
|
||||
CHANNEL_END => Some(101),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_special_id(&self, token_id: u32) -> bool {
|
||||
matches!(token_id, 100..=105)
|
||||
}
|
||||
}
|
||||
|
||||
struct MissingTokenTokenizer;
|
||||
|
||||
impl Tokenizer for MissingTokenTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.chars().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(token_ids
|
||||
.iter()
|
||||
.map(|token_id| char::from_u32(*token_id).unwrap_or('\u{FFFD}'))
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
None
|
||||
}
|
||||
fn tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
.with_special_token(CHANNEL_START, CHANNEL_START_ID)
|
||||
.with_special_token(CHANNEL_END, CHANNEL_END_ID)
|
||||
.with_special_token("<turn-boundary>", TURN_BOUNDARY_ID)
|
||||
}
|
||||
|
||||
trait UnifiedParserTestExt {
|
||||
@@ -716,12 +665,12 @@ mod tests {
|
||||
}
|
||||
|
||||
fn test_parser() -> Gemma4UnifiedParser {
|
||||
Gemma4UnifiedParser::new(&test_tools(), Arc::new(FakeTokenizer)).unwrap()
|
||||
Gemma4UnifiedParser::new(&test_tools(), Arc::new(tokenizer())).unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gemma4_create_requires_channel_start_token() {
|
||||
let error = match Gemma4UnifiedParser::new(&test_tools(), Arc::new(MissingTokenTokenizer)) {
|
||||
let error = match Gemma4UnifiedParser::new(&test_tools(), Arc::new(TestTokenizer::new())) {
|
||||
Ok(_) => panic!("expected missing token error"),
|
||||
Err(error) => error,
|
||||
};
|
||||
@@ -1046,7 +995,7 @@ mod tests {
|
||||
#[test]
|
||||
fn gemma4_initialize_open_channel_prompt_starts_in_reasoning() {
|
||||
let mut parser = test_parser();
|
||||
parser.initialize(&[100, 3000, 3001]).unwrap();
|
||||
parser.initialize(&[CHANNEL_START_ID, 3000, 3001]).unwrap();
|
||||
|
||||
let output = parser.parse_complete("reason<channel|>answer").unwrap();
|
||||
|
||||
@@ -1057,7 +1006,7 @@ mod tests {
|
||||
#[test]
|
||||
fn gemma4_initialize_turn_prompt_starts_in_text() {
|
||||
let mut parser = test_parser();
|
||||
parser.initialize(&[104, 3000, 3001]).unwrap();
|
||||
parser.initialize(&[TURN_BOUNDARY_ID, 3000, 3001]).unwrap();
|
||||
|
||||
let output = parser.parse_complete("<|channel>thought\nreason<channel|>answer").unwrap();
|
||||
|
||||
@@ -1068,7 +1017,7 @@ mod tests {
|
||||
#[test]
|
||||
fn gemma4_initialize_special_token_caps_boundary_scan() {
|
||||
let mut parser = test_parser();
|
||||
parser.initialize(&[100, 3000, 104, 3001]).unwrap();
|
||||
parser.initialize(&[CHANNEL_START_ID, 3000, TURN_BOUNDARY_ID, 3001]).unwrap();
|
||||
|
||||
let output = parser.parse_complete("answer").unwrap();
|
||||
|
||||
@@ -1079,7 +1028,7 @@ mod tests {
|
||||
#[test]
|
||||
fn gemma4_initialize_closed_channel_prompt_starts_in_text() {
|
||||
let mut parser = test_parser();
|
||||
parser.initialize(&[100, 3000, 3001, 101]).unwrap();
|
||||
parser.initialize(&[CHANNEL_START_ID, 3000, 3001, CHANNEL_END_ID]).unwrap();
|
||||
|
||||
let output = parser.parse_complete("answer").unwrap();
|
||||
|
||||
|
||||
@@ -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,8 +60,11 @@ clap.workspace = true
|
||||
expect-test.workspace = true
|
||||
rmp-serde.workspace = true
|
||||
serial_test.workspace = true
|
||||
tempfile.workspace = true
|
||||
tokio = { workspace = true, features = ["test-util"] }
|
||||
tower.workspace = true
|
||||
vllm-engine-core-client = { workspace = true, features = ["test-util"] }
|
||||
vllm-tokenizer = { workspace = true, features = ["test-utils"] }
|
||||
zeromq.workspace = true
|
||||
|
||||
[lints]
|
||||
|
||||
@@ -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}");
|
||||
|
||||
@@ -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
|
||||
{
|
||||
|
||||
@@ -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,
|
||||
@@ -16,14 +24,18 @@ use vllm_engine_core_client::protocol::{
|
||||
use vllm_engine_core_client::test_utils::{IpcNamespace, spawn_mock_engine_task};
|
||||
use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig, EngineId};
|
||||
use vllm_llm::Llm;
|
||||
use vllm_text::tokenizer::{DynTokenizer, Tokenizer};
|
||||
use vllm_text::tokenizer::DynTokenizer;
|
||||
use vllm_text::{Prompt, TextBackend};
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
use zeromq::prelude::{SocketRecv, SocketSend};
|
||||
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
|
||||
|
||||
use super::pb::generate_client::GenerateClient;
|
||||
use super::{GenerateServer, GenerateServiceImpl, pb};
|
||||
use crate::listener::{Listener, MaybeTlsListener};
|
||||
use crate::state::AppState;
|
||||
use crate::tls;
|
||||
use crate::tls_tests::{TestCerts, server_tls};
|
||||
|
||||
// ========================================================================================
|
||||
// Helpers (mirrors the patterns in routes/tests.rs)
|
||||
@@ -144,37 +156,9 @@ fn test_llm(client: EngineCoreClient) -> Llm {
|
||||
#[derive(Clone, Debug)]
|
||||
struct FakeTextBackend;
|
||||
|
||||
#[derive(Debug)]
|
||||
struct FakeTokenizer;
|
||||
|
||||
impl Tokenizer for FakeTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.bytes().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<String> {
|
||||
Ok(
|
||||
String::from_utf8_lossy(&token_ids.iter().map(|id| *id as u8).collect::<Vec<_>>())
|
||||
.into_owned(),
|
||||
)
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
token.bytes().next().map(u32::from)
|
||||
}
|
||||
}
|
||||
|
||||
impl TextBackend for FakeTextBackend {
|
||||
fn tokenizer(&self) -> DynTokenizer {
|
||||
Arc::new(FakeTokenizer)
|
||||
Arc::new(TestTokenizer::new())
|
||||
}
|
||||
|
||||
fn model_id(&self) -> &str {
|
||||
@@ -211,17 +195,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 +238,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 = MaybeTlsListener::plain(Listener::Tcp(listener));
|
||||
TonicServer::builder()
|
||||
.add_service(svc)
|
||||
.serve_with_incoming(incoming)
|
||||
@@ -274,7 +268,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 +275,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 = MaybeTlsListener::tls(Listener::Tcp(listener), context);
|
||||
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 = MaybeTlsListener::plain(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 +865,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();
|
||||
}
|
||||
|
||||
+159
-34
@@ -10,20 +10,33 @@ 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 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};
|
||||
@@ -31,11 +44,18 @@ use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig};
|
||||
use vllm_llm::Llm;
|
||||
use vllm_text::TextLlm;
|
||||
|
||||
use crate::listener::Listener;
|
||||
use crate::listener::{Listener, MaybeTlsListener};
|
||||
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 +65,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 +161,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?,
|
||||
@@ -138,46 +178,45 @@ where
|
||||
let listener = Listener::bind(&config.listener_mode)
|
||||
.await
|
||||
.context("failed to bind listener for OpenAI server")?;
|
||||
let bind_address = listener.local_addr()?;
|
||||
let bind_address = listener.local_addr_display()?;
|
||||
let model = state.primary_model_name().to_owned();
|
||||
let app = extend_router(build_router(state.clone()));
|
||||
|
||||
// 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,13 +247,28 @@ 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 {
|
||||
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 listener = match tls_config {
|
||||
Some(context) => MaybeTlsListener::tls(listener, context),
|
||||
None => MaybeTlsListener::plain(listener),
|
||||
};
|
||||
let server = serve_connections(listener, app, shutdown.cancelled_owned(), timeouts);
|
||||
|
||||
let result = tokio::select! {
|
||||
result = server => {
|
||||
@@ -236,16 +290,17 @@ 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(),
|
||||
);
|
||||
let incoming = match grpc_tls {
|
||||
Some(context) => MaybeTlsListener::tls(grpc_listener, context),
|
||||
None => MaybeTlsListener::plain(grpc_listener),
|
||||
};
|
||||
let server = svc.serve_with_incoming_shutdown(incoming, shutdown.cancelled_owned());
|
||||
|
||||
let result = tokio::select! {
|
||||
result = server => {
|
||||
@@ -272,6 +327,57 @@ where
|
||||
state.shutdown(shutdown_deadline).await
|
||||
}
|
||||
|
||||
/// Per-connection timeouts applied while serving HTTP/HTTPS.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct ConnectionTimeouts {
|
||||
/// 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,
|
||||
}
|
||||
|
||||
/// 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,
|
||||
timeouts: ConnectionTimeouts,
|
||||
) -> 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(timeouts.header_read);
|
||||
if !timeouts.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 +399,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");
|
||||
}
|
||||
}
|
||||
|
||||
+187
-15
@@ -1,28 +1,50 @@
|
||||
//! 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 openssl::ssl::SslContext;
|
||||
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;
|
||||
use crate::{HttpListenerMode, tls};
|
||||
|
||||
/// 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.
|
||||
///
|
||||
@@ -40,7 +62,7 @@ impl Listener {
|
||||
|
||||
/// Return a log-friendly local address string for either TCP or Unix
|
||||
/// sockets.
|
||||
pub fn local_addr(&self) -> Result<String> {
|
||||
pub fn local_addr_display(&self) -> Result<String> {
|
||||
match self {
|
||||
Self::Tcp(listener) => Ok(listener.local_addr()?.to_string()),
|
||||
Self::Unix(listener) => Ok(match listener.local_addr()?.as_pathname() {
|
||||
@@ -70,30 +92,180 @@ impl Listener {
|
||||
Ok(Self::Tcp(TcpListener::from_std(std_listener)?))
|
||||
}
|
||||
}
|
||||
|
||||
fn local_addr(&self) -> Result<ListenerAddr> {
|
||||
match self {
|
||||
Self::Tcp(listener) => listener.local_addr().map(ListenerAddr::Tcp),
|
||||
Self::Unix(listener) => listener.local_addr().map(ListenerAddr::Unix),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Allow the unified listener to plug directly into tonic's gRPC server.
|
||||
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> {
|
||||
self.local_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.local_addr()
|
||||
}
|
||||
}
|
||||
|
||||
/// A listener that may be either a plain TCP/UDS listener or a TLS listener over it.
|
||||
pub enum MaybeTlsListener {
|
||||
Plain(Listener),
|
||||
Tls(tls_listener::TlsListener<Listener, SslContext>),
|
||||
}
|
||||
|
||||
impl MaybeTlsListener {
|
||||
/// Create a plain listener without TLS.
|
||||
pub fn plain(listener: Listener) -> Self {
|
||||
Self::Plain(listener)
|
||||
}
|
||||
|
||||
/// Create a TLS listener over the given plain listener.
|
||||
pub fn tls(listener: Listener, context: SslContext) -> Self {
|
||||
Self::Tls(
|
||||
tls_listener::builder(context)
|
||||
.handshake_timeout(tls::TLS_HANDSHAKE_TIMEOUT)
|
||||
.listen(listener),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// Listener I/O type that may be either a plain TCP/UDS stream or a TLS stream over it.
|
||||
#[derive(Debug)]
|
||||
#[enum_derive(tokio1::AsyncRead, tokio1::AsyncWrite)]
|
||||
pub enum MaybeTlsStream {
|
||||
Plain(ListenerIo),
|
||||
Tls(tokio_openssl::SslStream<ListenerIo>),
|
||||
}
|
||||
|
||||
/// Allow the maybe-TLS listener to plug directly into `axum::serve(...)`.
|
||||
impl axum::serve::Listener for MaybeTlsListener {
|
||||
type Addr = ListenerAddr;
|
||||
type Io = MaybeTlsStream;
|
||||
|
||||
async fn accept(&mut self) -> (Self::Io, Self::Addr) {
|
||||
match self {
|
||||
Self::Tcp(listener) => listener.local_addr().map(Either::Left),
|
||||
Self::Unix(listener) => listener.local_addr().map(Either::Right),
|
||||
Self::Plain(listener) => {
|
||||
let (io, addr) = axum::serve::Listener::accept(listener).await;
|
||||
(MaybeTlsStream::Plain(io), addr)
|
||||
}
|
||||
Self::Tls(tls_listener) => {
|
||||
let (io, addr) = axum::serve::Listener::accept(tls_listener).await;
|
||||
(MaybeTlsStream::Tls(io), addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn local_addr(&self) -> tokio::io::Result<Self::Addr> {
|
||||
match self {
|
||||
Self::Plain(listener) => listener.local_addr(),
|
||||
Self::Tls(tls_listener) => tls_listener.local_addr(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Allow the maybe-TLS listener to plug directly into tonic's gRPC server.
|
||||
impl Connected for MaybeTlsStream {
|
||||
type ConnectInfo = TcpConnectInfo;
|
||||
|
||||
fn connect_info(&self) -> TcpConnectInfo {
|
||||
match self {
|
||||
Self::Plain(stream) => stream.connect_info(),
|
||||
Self::Tls(stream) => stream.get_ref().connect_info(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Allow the maybe-TLS listener to be adaptable to tonic's incoming stream shape.
|
||||
impl futures::Stream for MaybeTlsListener {
|
||||
type Item = std::io::Result<MaybeTlsStream>;
|
||||
|
||||
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
match self.get_mut() {
|
||||
Self::Plain(listener) => {
|
||||
let listener = Pin::new(listener);
|
||||
let (io, _) = ready!(listener.poll_accept(cx))?;
|
||||
Poll::Ready(Some(Ok(MaybeTlsStream::Plain(io))))
|
||||
}
|
||||
Self::Tls(tls_listener) => {
|
||||
let tls_listener = Pin::new(tls_listener);
|
||||
let (io, _) =
|
||||
ready!(tls_listener.poll_accept(cx)).map_err(std::io::Error::other)?;
|
||||
Poll::Ready(Some(Ok(MaybeTlsStream::Tls(io))))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -24,8 +24,9 @@ use vllm_engine_core_client::protocol::{
|
||||
use vllm_engine_core_client::test_utils::{IpcNamespace, spawn_mock_engine_task};
|
||||
use vllm_engine_core_client::{EngineCoreClient, EngineCoreClientConfig, EngineId};
|
||||
use vllm_llm::Llm;
|
||||
use vllm_text::tokenizer::{DynTokenizer, Tokenizer};
|
||||
use vllm_text::tokenizer::DynTokenizer;
|
||||
use vllm_text::{Prompt, TextBackend};
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
use zeromq::prelude::{SocketRecv, SocketSend};
|
||||
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
|
||||
|
||||
@@ -151,37 +152,9 @@ fn test_llm(client: EngineCoreClient) -> Llm {
|
||||
#[derive(Clone, Debug)]
|
||||
struct FakeChatBackend;
|
||||
|
||||
#[derive(Debug)]
|
||||
struct FakeChatTokenizer;
|
||||
|
||||
impl Tokenizer for FakeChatTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.bytes().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<String> {
|
||||
Ok(
|
||||
String::from_utf8_lossy(&token_ids.iter().map(|id| *id as u8).collect::<Vec<_>>())
|
||||
.into_owned(),
|
||||
)
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
token.bytes().next().map(u32::from)
|
||||
}
|
||||
}
|
||||
|
||||
impl TextBackend for FakeChatBackend {
|
||||
fn tokenizer(&self) -> DynTokenizer {
|
||||
Arc::new(FakeChatTokenizer)
|
||||
Arc::new(TestTokenizer::new())
|
||||
}
|
||||
|
||||
fn model_id(&self) -> &str {
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -194,7 +194,7 @@ mod tests {
|
||||
use axum::http::HeaderMap;
|
||||
use serde_json::json;
|
||||
use vllm_text::Prompt;
|
||||
use vllm_text::tokenizer::Tokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::prepare_completion_request;
|
||||
use crate::lora::LoraModelResolution;
|
||||
@@ -212,32 +212,8 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct TestTokenizer;
|
||||
|
||||
impl Tokenizer for TestTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<Vec<u32>> {
|
||||
Ok(text.bytes().map(u32::from).collect())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<String> {
|
||||
Ok(
|
||||
String::from_utf8_lossy(&token_ids.iter().map(|id| *id as u8).collect::<Vec<_>>())
|
||||
.into_owned(),
|
||||
)
|
||||
}
|
||||
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
None
|
||||
}
|
||||
fn test_tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
}
|
||||
|
||||
fn base_request_json() -> serde_json::Value {
|
||||
@@ -297,7 +273,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
@@ -340,7 +316,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare")
|
||||
.text_request
|
||||
@@ -374,7 +350,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
@@ -399,7 +375,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
@@ -422,7 +398,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
@@ -446,7 +422,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
@@ -471,7 +447,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
@@ -498,7 +474,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
@@ -523,7 +499,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
|
||||
@@ -549,7 +525,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
assert_eq!(prepared.text_request.sampling_params.logprobs, Some(1));
|
||||
@@ -574,7 +550,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
request_context(&headers, None),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
assert_eq!(prepared.text_request.data_parallel_rank, Some(3));
|
||||
@@ -593,7 +569,7 @@ mod tests {
|
||||
request,
|
||||
&served(&["Qwen/Qwen1.5-0.5B-Chat"]),
|
||||
ResolvedRequestContext::default(),
|
||||
&TestTokenizer,
|
||||
&test_tokenizer(),
|
||||
)
|
||||
.expect("prepare");
|
||||
assert_eq!(prepared.text_request.data_parallel_rank, None);
|
||||
|
||||
@@ -40,8 +40,9 @@ use vllm_engine_core_client::{
|
||||
};
|
||||
use vllm_llm::Llm;
|
||||
use vllm_metrics::METRICS;
|
||||
use vllm_text::tokenizer::{DynTokenizer, Tokenizer};
|
||||
use vllm_text::tokenizer::DynTokenizer;
|
||||
use vllm_text::{Prompt, TextBackend};
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
use zeromq::prelude::{SocketRecv, SocketSend};
|
||||
use zeromq::{DealerSocket, PushSocket, ZmqMessage};
|
||||
|
||||
@@ -417,82 +418,20 @@ struct FakeChatBackend {
|
||||
}
|
||||
|
||||
/// Synthetic BOS id used when `add_special_tokens` is true in tests.
|
||||
const FAKE_BOS_TOKEN_ID: u32 = 1;
|
||||
const FAKE_BOS_TOKEN_ID: u32 = 256;
|
||||
const UNKNOWN_DECODE_TOKEN_ID: u32 = 10_000;
|
||||
|
||||
#[derive(Debug)]
|
||||
struct FakeChatTokenizer;
|
||||
|
||||
impl Tokenizer for FakeChatTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
add_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<Vec<u32>> {
|
||||
let mut token_ids = Vec::new();
|
||||
if add_special_tokens {
|
||||
token_ids.push(FAKE_BOS_TOKEN_ID);
|
||||
}
|
||||
let mut rest = text;
|
||||
while !rest.is_empty() {
|
||||
if let Some(stripped) = rest.strip_prefix("<image>") {
|
||||
token_ids.push(999);
|
||||
rest = stripped;
|
||||
continue;
|
||||
}
|
||||
if let Some(stripped) = rest.strip_prefix("<|image_pad|>") {
|
||||
token_ids.push(151655);
|
||||
rest = stripped;
|
||||
continue;
|
||||
}
|
||||
|
||||
let ch = rest.chars().next().expect("rest is not empty");
|
||||
let mut buf = [0; 4];
|
||||
token_ids.extend(ch.encode_utf8(&mut buf).bytes().map(u32::from));
|
||||
rest = &rest[ch.len_utf8()..];
|
||||
}
|
||||
Ok(token_ids)
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<String> {
|
||||
Ok(
|
||||
String::from_utf8_lossy(&token_ids.iter().map(|id| *id as u8).collect::<Vec<_>>())
|
||||
.into_owned(),
|
||||
)
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
match token {
|
||||
"<image>" => Some(999),
|
||||
"<|image_pad|>" => Some(151655),
|
||||
"<think>" => Some(0xF001),
|
||||
"</think>" => Some(0xF002),
|
||||
"<|START_THINKING|>" => Some(0xF003),
|
||||
"<|END_THINKING|>" => Some(0xF004),
|
||||
"◁think▷" => Some(0xF005),
|
||||
"◁/think▷" => Some(0xF006),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn id_to_token(&self, id: u32) -> Option<String> {
|
||||
match id {
|
||||
FAKE_BOS_TOKEN_ID => Some("<bos>".to_string()),
|
||||
999 => Some("<image>".to_string()),
|
||||
151655 => Some("<|image_pad|>".to_string()),
|
||||
0xF001 => Some("<think>".to_string()),
|
||||
0xF002 => Some("</think>".to_string()),
|
||||
0xF003 => Some("<|START_THINKING|>".to_string()),
|
||||
0xF004 => Some("<|END_THINKING|>".to_string()),
|
||||
0xF005 => Some("◁think▷".to_string()),
|
||||
0xF006 => Some("◁/think▷".to_string()),
|
||||
id if id < 128 => char::from_u32(id).map(|ch| ch.to_string()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
fn fake_chat_tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
.with_bos_token("<bos>", FAKE_BOS_TOKEN_ID)
|
||||
.with_regular_token("<image>", 999)
|
||||
.with_regular_token("<|image_pad|>", 151655)
|
||||
.with_regular_token("<think>", 0xF001)
|
||||
.with_regular_token("</think>", 0xF002)
|
||||
.with_regular_token("<|START_THINKING|>", 0xF003)
|
||||
.with_regular_token("<|END_THINKING|>", 0xF004)
|
||||
.with_regular_token("◁think▷", 0xF005)
|
||||
.with_regular_token("◁/think▷", 0xF006)
|
||||
}
|
||||
|
||||
impl FakeChatBackend {
|
||||
@@ -530,7 +469,7 @@ impl fmt::Debug for FakeChatBackend {
|
||||
|
||||
impl TextBackend for FakeChatBackend {
|
||||
fn tokenizer(&self) -> DynTokenizer {
|
||||
Arc::new(FakeChatTokenizer)
|
||||
Arc::new(fake_chat_tokenizer())
|
||||
}
|
||||
|
||||
fn model_id(&self) -> &str {
|
||||
@@ -630,7 +569,7 @@ fn qwen_multimodal_model_info() -> vllm_chat::multimodal::MultimodalModelInfo {
|
||||
Some("qwen2_vl".to_string()),
|
||||
Some(&config_path),
|
||||
None,
|
||||
Arc::new(FakeChatTokenizer),
|
||||
Arc::new(fake_chat_tokenizer()),
|
||||
)
|
||||
.expect("load multimodal info")
|
||||
.expect("qwen multimodal info is registered");
|
||||
@@ -638,70 +577,6 @@ fn qwen_multimodal_model_info() -> vllm_chat::multimodal::MultimodalModelInfo {
|
||||
info
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
struct FailingDecodeChatBackend;
|
||||
|
||||
#[derive(Debug)]
|
||||
struct FailingDecodeTokenizer;
|
||||
|
||||
impl Tokenizer for FailingDecodeTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
text: &str,
|
||||
add_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<Vec<u32>> {
|
||||
FakeChatTokenizer.encode(text, add_special_tokens)
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
skip_special_tokens: bool,
|
||||
) -> vllm_text::tokenizer::Result<String> {
|
||||
if token_ids.contains(&(b'i' as u32)) {
|
||||
return Err(vllm_text::tokenizer::TokenizerError(
|
||||
"forced decode failure for streaming test".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
FakeChatTokenizer.decode(token_ids, skip_special_tokens)
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
FakeChatTokenizer.token_to_id(token)
|
||||
}
|
||||
}
|
||||
|
||||
impl TextBackend for FailingDecodeChatBackend {
|
||||
fn tokenizer(&self) -> DynTokenizer {
|
||||
Arc::new(FailingDecodeTokenizer)
|
||||
}
|
||||
|
||||
fn model_id(&self) -> &str {
|
||||
"test-model"
|
||||
}
|
||||
}
|
||||
|
||||
impl ChatBackend for FailingDecodeChatBackend {
|
||||
fn chat_renderer(&self) -> DynChatRenderer {
|
||||
Arc::new(self.clone())
|
||||
}
|
||||
|
||||
fn new_chat_output_processor(
|
||||
&self,
|
||||
_request: &mut ChatRequest,
|
||||
_options: NewChatOutputProcessorOptions<'_>,
|
||||
) -> vllm_chat::Result<DynChatOutputProcessor> {
|
||||
Ok(Box::new(DefaultChatOutputProcessor::plain_text_only()))
|
||||
}
|
||||
}
|
||||
|
||||
impl ChatRenderer for FailingDecodeChatBackend {
|
||||
fn render(&self, request: &ChatRequest) -> vllm_chat::Result<vllm_chat::RenderedPrompt> {
|
||||
FakeChatBackend::new().render(request)
|
||||
}
|
||||
}
|
||||
|
||||
async fn test_models_with_engine_outputs_and_backend_inner(
|
||||
engine_id: impl Into<EngineId>,
|
||||
output_specs: Vec<(Vec<u32>, Option<EngineCoreFinishReason>)>,
|
||||
@@ -2755,8 +2630,8 @@ async fn load_endpoint_resets_when_stream_response_is_dropped() {
|
||||
#[serial]
|
||||
async fn stream_error_is_returned_as_openai_error_sse() {
|
||||
let (app, engine_task) = test_app_with_backend_and_stream_output_specs(
|
||||
Arc::new(FailingDecodeChatBackend),
|
||||
default_stream_output_specs(),
|
||||
Arc::new(FakeChatBackend::new()),
|
||||
vec![(vec![UNKNOWN_DECODE_TOKEN_ID], None)],
|
||||
)
|
||||
.await;
|
||||
let response = app
|
||||
@@ -2789,7 +2664,9 @@ async fn stream_error_is_returned_as_openai_error_sse() {
|
||||
assert!(text.contains("\"role\":\"assistant\""), "{text}");
|
||||
assert!(text.contains("\"type\":\"server_error\""), "{text}");
|
||||
assert!(
|
||||
text.contains("forced decode failure for streaming test"),
|
||||
text.contains(&format!(
|
||||
"test tokenizer cannot decode unknown token id {UNKNOWN_DECODE_TOKEN_ID}"
|
||||
)),
|
||||
"{text}"
|
||||
);
|
||||
assert!(!text.contains("\"usage\":"), "{text}");
|
||||
@@ -4129,6 +4006,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() {
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
@@ -0,0 +1,677 @@
|
||||
//! TLS tests: `build_server_config` unit checks plus end-to-end OpenSSL handshakes
|
||||
//! through the production listener/connection 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, MaybeTlsListener};
|
||||
use crate::{ConnectionTimeouts, serve_connections, 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 {
|
||||
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
|
||||
/// listener/connection path. 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_display().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 listener = match server_config {
|
||||
Some(context) => MaybeTlsListener::tls(listener, context),
|
||||
None => MaybeTlsListener::plain(listener),
|
||||
};
|
||||
let _ = serve_connections(listener, 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(start_paused = true)]
|
||||
async fn tls_handshake_timeout_drops_silent_client() {
|
||||
// Silent client (no ClientHello) must be dropped at the handshake deadline.
|
||||
let certs = TestCerts::generate();
|
||||
let (addr, shutdown) = spawn_server(Some(server_tls(&certs, 0))).await;
|
||||
|
||||
let mut tcp = TcpStream::connect(&addr).await.expect("connect");
|
||||
tokio::task::yield_now().await;
|
||||
tokio::time::advance(tls::TLS_HANDSHAKE_TIMEOUT + Duration::from_millis(1)).await;
|
||||
tokio::task::yield_now().await;
|
||||
|
||||
let mut buf = [0u8; 1];
|
||||
let read = tokio::time::timeout(Duration::from_secs(1), 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 {
|
||||
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 {
|
||||
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 {
|
||||
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 {
|
||||
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 {
|
||||
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();
|
||||
}
|
||||
@@ -31,6 +31,7 @@ serial_test.workspace = true
|
||||
tempfile.workspace = true
|
||||
tokio.workspace = true
|
||||
vllm-llm = { workspace = true, features = ["test-util"] }
|
||||
vllm-tokenizer = { workspace = true, features = ["test-utils"] }
|
||||
|
||||
[lints]
|
||||
workspace = true
|
||||
|
||||
@@ -275,6 +275,7 @@ mod tests {
|
||||
use std::collections::{BTreeSet, HashMap};
|
||||
|
||||
use serial_test::file_serial;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::*;
|
||||
use crate::backend::hf::HfTextBackend;
|
||||
@@ -282,60 +283,8 @@ mod tests {
|
||||
use crate::error::{LogprobsError, TokenIdsError};
|
||||
use crate::request::{Prompt, TextRequest};
|
||||
|
||||
/// Stub tokenizer that returns empty token IDs — sufficient for tests that
|
||||
/// don't exercise bad-words tokenization.
|
||||
struct StubTokenizer;
|
||||
|
||||
impl Tokenizer for StubTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
_text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(vec![])
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
_token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn stub_tokenizer() -> StubTokenizer {
|
||||
StubTokenizer
|
||||
}
|
||||
|
||||
struct FixedTokenizer {
|
||||
token_ids: Vec<u32>,
|
||||
}
|
||||
|
||||
impl Tokenizer for FixedTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
_text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
Ok(self.token_ids.clone())
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
_token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(String::new())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
None
|
||||
}
|
||||
fn stub_tokenizer() -> TestTokenizer {
|
||||
TestTokenizer::new()
|
||||
}
|
||||
|
||||
fn sample_request() -> TextRequest {
|
||||
@@ -952,9 +901,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn lower_sampling_params_rejects_out_of_vocab_bad_words() {
|
||||
let tokenizer = FixedTokenizer {
|
||||
token_ids: vec![1999, 2000],
|
||||
};
|
||||
let tokenizer = TestTokenizer::new().with_regular_token("blocked", 2000);
|
||||
let error = lower_sampling_params(
|
||||
SamplingParams {
|
||||
bad_words: Some(vec!["blocked".to_string()]),
|
||||
|
||||
@@ -323,37 +323,11 @@ mod tests {
|
||||
use futures::{Stream, stream};
|
||||
use vllm_engine_core_client::AbortCause;
|
||||
use vllm_llm::GenerateOutput;
|
||||
use vllm_tokenizer::Tokenizer;
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::*;
|
||||
use crate::output::TextOutputStreamExt as _;
|
||||
|
||||
/// Backend that treats each token ID as a raw byte, producing lossy UTF-8.
|
||||
struct ByteTokenizer;
|
||||
|
||||
impl Tokenizer for ByteTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
_text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
let bytes = token_ids.iter().map(|id| *id as u8).collect::<Vec<_>>();
|
||||
Ok(String::from_utf8_lossy(&bytes).into_owned())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
unreachable!()
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper: run `decoded_text_event_stream` to completion and return the
|
||||
/// collected output.
|
||||
async fn run_to_completion(
|
||||
@@ -366,7 +340,7 @@ mod tests {
|
||||
token_ids,
|
||||
Some(FinishReason::Length),
|
||||
))]);
|
||||
let tokenizer: DynTokenizer = Arc::new(ByteTokenizer);
|
||||
let tokenizer: DynTokenizer = Arc::new(TestTokenizer::new());
|
||||
decoded_text_event_stream("test".into(), tokenizer, raw_stream, decode_options, false)
|
||||
.collect_output()
|
||||
.await
|
||||
@@ -419,7 +393,7 @@ mod tests {
|
||||
))),
|
||||
dropped_cause: Arc::clone(&dropped_cause),
|
||||
};
|
||||
let tokenizer: DynTokenizer = Arc::new(ByteTokenizer);
|
||||
let tokenizer: DynTokenizer = Arc::new(TestTokenizer::new());
|
||||
|
||||
let output = decoded_text_event_stream(
|
||||
"test".into(),
|
||||
|
||||
@@ -129,40 +129,13 @@ fn decode_position_logprobs<T: Tokenizer + ?Sized>(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use vllm_llm::{Logprobs, PositionLogprobs, TokenLogprob};
|
||||
use vllm_tokenizer::test_utils::TestTokenizer;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ByteTokenizer;
|
||||
|
||||
impl vllm_tokenizer::Tokenizer for ByteTokenizer {
|
||||
fn encode(
|
||||
&self,
|
||||
_text: &str,
|
||||
_add_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<Vec<u32>> {
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
fn decode(
|
||||
&self,
|
||||
token_ids: &[u32],
|
||||
_skip_special_tokens: bool,
|
||||
) -> vllm_tokenizer::Result<String> {
|
||||
Ok(String::from_utf8_lossy(
|
||||
&token_ids.iter().map(|token_id| *token_id as u8).collect::<Vec<_>>(),
|
||||
)
|
||||
.into_owned())
|
||||
}
|
||||
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
unreachable!()
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decode_logprobs_decodes_every_candidate_token() {
|
||||
let tokenizer = ByteTokenizer;
|
||||
let tokenizer = TestTokenizer::new();
|
||||
let logprobs = Logprobs {
|
||||
positions: vec![PositionLogprobs {
|
||||
entries: vec![
|
||||
@@ -205,7 +178,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn decode_prompt_logprobs_separates_first_prompt_token() {
|
||||
let tokenizer = ByteTokenizer;
|
||||
let tokenizer = TestTokenizer::new();
|
||||
let logprobs = Logprobs {
|
||||
positions: vec![PositionLogprobs {
|
||||
entries: vec![TokenLogprob {
|
||||
|
||||
@@ -4,6 +4,9 @@ version.workspace = true
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[features]
|
||||
test-utils = []
|
||||
|
||||
[dependencies]
|
||||
base64.workspace = true
|
||||
fastokens.workspace = true
|
||||
|
||||
@@ -170,6 +170,10 @@ mod tests {
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
fn id_to_token(&self, _id: u32) -> Option<String> {
|
||||
unreachable!()
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -248,6 +252,10 @@ mod tests {
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
fn id_to_token(&self, _id: u32) -> Option<String> {
|
||||
unreachable!()
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -320,6 +328,10 @@ mod tests {
|
||||
fn token_to_id(&self, _token: &str) -> Option<u32> {
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
fn id_to_token(&self, _id: u32) -> Option<String> {
|
||||
unreachable!()
|
||||
}
|
||||
}
|
||||
|
||||
/// Without the char-boundary fix, this panics slicing mid-emoji.
|
||||
|
||||
@@ -8,6 +8,8 @@ mod error;
|
||||
mod hf;
|
||||
mod incremental;
|
||||
mod tekken;
|
||||
#[cfg(any(test, feature = "test-utils"))]
|
||||
pub mod test_utils;
|
||||
mod tiktoken;
|
||||
|
||||
pub use error::{Result, TokenizerError};
|
||||
@@ -28,11 +30,7 @@ pub trait Tokenizer: Send + Sync {
|
||||
fn token_to_id(&self, token: &str) -> Option<u32>;
|
||||
|
||||
/// Convert one token ID into the tokenizer's raw token string.
|
||||
fn id_to_token(&self, _id: u32) -> Option<String> {
|
||||
// TODO: remove default impl and require this to be implemented by all
|
||||
// tokenizers
|
||||
None
|
||||
}
|
||||
fn id_to_token(&self, id: u32) -> Option<String>;
|
||||
|
||||
/// Return the vocabulary size. Backends that cannot report it fall back to
|
||||
/// `usize::MAX`, an effectively unbounded value used only by test stubs.
|
||||
|
||||
@@ -0,0 +1,434 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use crate::{Result, Tokenizer, TokenizerError};
|
||||
|
||||
const FIRST_CONFIGURED_TOKEN_ID: u32 = 256;
|
||||
|
||||
/// Whether a configured test token should be treated as special.
|
||||
///
|
||||
/// Special tokens are skipped by [`Tokenizer::decode`] when
|
||||
/// `skip_special_tokens` is set. Regular configured tokens are always emitted.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum TestTokenKind {
|
||||
/// Token is skipped when `skip_special_tokens = true`.
|
||||
Special,
|
||||
/// Token is emitted regardless of `skip_special_tokens`.
|
||||
Regular,
|
||||
}
|
||||
|
||||
impl TestTokenKind {
|
||||
fn is_special(self) -> bool {
|
||||
matches!(self, Self::Special)
|
||||
}
|
||||
}
|
||||
|
||||
/// Decode behavior for token ids that are neither configured tokens nor byte ids.
|
||||
///
|
||||
/// The default is [`UnknownDecode::Error`] so tests notice missing tokenizer
|
||||
/// fixtures instead of silently accepting impossible ids. Individual tests can
|
||||
/// opt into empty or replacement output when they are explicitly modeling a
|
||||
/// lenient detokenization path.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum UnknownDecode {
|
||||
/// Return a tokenizer error on the first unknown id.
|
||||
Error,
|
||||
/// Drop unknown ids from decoded output.
|
||||
Empty,
|
||||
/// Emit U+FFFD for each unknown id.
|
||||
Replacement,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct TestToken {
|
||||
text: String,
|
||||
kind: TestTokenKind,
|
||||
}
|
||||
|
||||
/// Configurable tokenizer for Rust frontend tests.
|
||||
///
|
||||
/// `TestTokenizer` is intentionally small, but its methods obey the same basic
|
||||
/// contract as production tokenizers:
|
||||
///
|
||||
/// - ordinary text encodes as UTF-8 byte ids;
|
||||
/// - configured token ids start at 256, leaving `0..=255` for byte fallback;
|
||||
/// - configured token ids and token text are unique;
|
||||
/// - configured tokens are matched before ordinary bytes, using longest-prefix matching so
|
||||
/// multi-character markers such as `<think>` work naturally;
|
||||
/// - `token_to_id` and `id_to_token` are consistent for configured tokens;
|
||||
/// - `decode` is strict by default for ids outside the byte range and the configured token table;
|
||||
/// - `vocab_size` is an exclusive upper bound covering byte ids and configured token ids unless a
|
||||
/// test sets it explicitly.
|
||||
///
|
||||
/// Prefer this helper over ad-hoc fake tokenizers for tests that rely on
|
||||
/// tokenizer semantics. Keep dedicated tiny fakes for error injection or for
|
||||
/// tests that deliberately need a degenerate tokenizer.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TestTokenizer {
|
||||
token_to_id: BTreeMap<String, u32>,
|
||||
id_to_token: BTreeMap<u32, TestToken>,
|
||||
unknown_decode: UnknownDecode,
|
||||
vocab_size: Option<usize>,
|
||||
bos_token_id: Option<u32>,
|
||||
}
|
||||
|
||||
impl Default for TestTokenizer {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl TestTokenizer {
|
||||
/// Create a byte-level test tokenizer with strict unknown-id decode.
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
token_to_id: BTreeMap::new(),
|
||||
id_to_token: BTreeMap::new(),
|
||||
unknown_decode: UnknownDecode::Error,
|
||||
vocab_size: None,
|
||||
bos_token_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Add a configured token and return the updated tokenizer.
|
||||
///
|
||||
/// Configured tokens must use ids outside the byte range and may be marked
|
||||
/// special or regular.
|
||||
pub fn with_token(mut self, token: impl Into<String>, id: u32, kind: TestTokenKind) -> Self {
|
||||
self.insert_token(token, id, kind);
|
||||
self
|
||||
}
|
||||
|
||||
/// Add a special configured token and return the updated tokenizer.
|
||||
pub fn with_special_token(self, token: impl Into<String>, id: u32) -> Self {
|
||||
self.with_token(token, id, TestTokenKind::Special)
|
||||
}
|
||||
|
||||
/// Add a regular configured token and return the updated tokenizer.
|
||||
pub fn with_regular_token(self, token: impl Into<String>, id: u32) -> Self {
|
||||
self.with_token(token, id, TestTokenKind::Regular)
|
||||
}
|
||||
|
||||
/// Add a special BOS token inserted by `encode(..., true)`.
|
||||
///
|
||||
/// This also registers the token in the normal token/id maps so
|
||||
/// `token_to_id`, `id_to_token`, `decode`, and `is_special_id` stay
|
||||
/// consistent for the inserted id.
|
||||
pub fn with_bos_token(mut self, token: impl Into<String>, id: u32) -> Self {
|
||||
self.insert_token(token, id, TestTokenKind::Special);
|
||||
self.bos_token_id = Some(id);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set decode behavior for unknown non-byte ids.
|
||||
pub fn with_unknown_decode(mut self, behavior: UnknownDecode) -> Self {
|
||||
self.unknown_decode = behavior;
|
||||
self
|
||||
}
|
||||
|
||||
/// Set an explicit vocabulary size.
|
||||
///
|
||||
/// Use this when a test needs a model-like vocabulary bound that differs
|
||||
/// from the highest configured token id plus one.
|
||||
pub fn with_vocab_size(mut self, vocab_size: usize) -> Self {
|
||||
self.vocab_size = Some(vocab_size);
|
||||
self
|
||||
}
|
||||
|
||||
fn insert_token(&mut self, token: impl Into<String>, id: u32, kind: TestTokenKind) {
|
||||
let token = token.into();
|
||||
assert!(
|
||||
!token.is_empty(),
|
||||
"configured test token text must be non-empty"
|
||||
);
|
||||
assert!(
|
||||
id >= FIRST_CONFIGURED_TOKEN_ID,
|
||||
"configured test token id {id} overlaps byte fallback range 0..=255"
|
||||
);
|
||||
assert!(
|
||||
token.len() > 1,
|
||||
"configured test token text {token:?} overlaps byte fallback token text"
|
||||
);
|
||||
if self.token_to_id.insert(token.clone(), id).is_some() {
|
||||
panic!("configured test token text {token:?} was registered more than once");
|
||||
}
|
||||
if self.id_to_token.insert(id, TestToken { text: token, kind }).is_some() {
|
||||
panic!("configured test token id {id} was registered more than once");
|
||||
}
|
||||
}
|
||||
|
||||
fn byte_to_token(id: u32) -> Option<String> {
|
||||
u8::try_from(id).ok().map(|byte| String::from_utf8_lossy(&[byte]).into_owned())
|
||||
}
|
||||
|
||||
fn flush_bytes(bytes: &mut Vec<u8>, output: &mut String) {
|
||||
if !bytes.is_empty() {
|
||||
output.push_str(&String::from_utf8_lossy(bytes));
|
||||
bytes.clear();
|
||||
}
|
||||
}
|
||||
|
||||
fn configured_token_prefix(&self, text: &str) -> Option<(&str, u32)> {
|
||||
self.token_to_id
|
||||
.iter()
|
||||
.filter_map(|(token, &id)| text.starts_with(token).then_some((token.as_str(), id)))
|
||||
.max_by_key(|(token, _)| token.len())
|
||||
}
|
||||
|
||||
fn inferred_vocab_size(&self) -> usize {
|
||||
let max_configured =
|
||||
self.id_to_token.last_key_value().map(|(&id, _)| id as usize + 1).unwrap_or(0);
|
||||
256.max(max_configured)
|
||||
}
|
||||
}
|
||||
|
||||
impl Tokenizer for TestTokenizer {
|
||||
fn encode(&self, text: &str, add_special_tokens: bool) -> Result<Vec<u32>> {
|
||||
let mut ids = Vec::new();
|
||||
if add_special_tokens && let Some(bos_token_id) = self.bos_token_id {
|
||||
ids.push(bos_token_id);
|
||||
}
|
||||
|
||||
let mut rest = text;
|
||||
while !rest.is_empty() {
|
||||
if let Some((token, id)) = self.configured_token_prefix(rest) {
|
||||
ids.push(id);
|
||||
rest = &rest[token.len()..];
|
||||
continue;
|
||||
}
|
||||
|
||||
let ch = rest.chars().next().expect("rest is not empty");
|
||||
let mut buf = [0_u8; 4];
|
||||
ids.extend(ch.encode_utf8(&mut buf).bytes().map(u32::from));
|
||||
rest = &rest[ch.len_utf8()..];
|
||||
}
|
||||
|
||||
Ok(ids)
|
||||
}
|
||||
|
||||
fn decode(&self, token_ids: &[u32], skip_special_tokens: bool) -> Result<String> {
|
||||
let mut output = String::new();
|
||||
let mut pending_bytes = Vec::new();
|
||||
for &id in token_ids {
|
||||
if let Some(token) = self.id_to_token.get(&id) {
|
||||
Self::flush_bytes(&mut pending_bytes, &mut output);
|
||||
if !(skip_special_tokens && token.kind.is_special()) {
|
||||
output.push_str(&token.text);
|
||||
}
|
||||
} else if let Ok(byte) = u8::try_from(id) {
|
||||
pending_bytes.push(byte);
|
||||
} else {
|
||||
Self::flush_bytes(&mut pending_bytes, &mut output);
|
||||
match self.unknown_decode {
|
||||
UnknownDecode::Error => {
|
||||
return Err(TokenizerError(format!(
|
||||
"test tokenizer cannot decode unknown token id {id}"
|
||||
)));
|
||||
}
|
||||
UnknownDecode::Empty => {}
|
||||
UnknownDecode::Replacement => output.push('\u{FFFD}'),
|
||||
}
|
||||
}
|
||||
}
|
||||
Self::flush_bytes(&mut pending_bytes, &mut output);
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
fn token_to_id(&self, token: &str) -> Option<u32> {
|
||||
self.token_to_id.get(token).copied().or_else(|| {
|
||||
let bytes = token.as_bytes();
|
||||
(bytes.len() == 1).then(|| u32::from(bytes[0]))
|
||||
})
|
||||
}
|
||||
|
||||
fn id_to_token(&self, id: u32) -> Option<String> {
|
||||
self.id_to_token
|
||||
.get(&id)
|
||||
.map(|token| token.text.clone())
|
||||
.or_else(|| Self::byte_to_token(id))
|
||||
}
|
||||
|
||||
fn vocab_size(&self) -> usize {
|
||||
self.vocab_size.unwrap_or_else(|| self.inferred_vocab_size())
|
||||
}
|
||||
|
||||
fn is_special_id(&self, token_id: u32) -> bool {
|
||||
self.id_to_token.get(&token_id).is_some_and(|token| token.kind.is_special())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn byte_text_roundtrips_and_reports_byte_ids() {
|
||||
let tokenizer = TestTokenizer::new();
|
||||
|
||||
let ids = tokenizer.encode("hi", false).unwrap();
|
||||
assert_eq!(ids, vec![b'h' as u32, b'i' as u32]);
|
||||
assert_eq!(tokenizer.decode(&ids, false).unwrap(), "hi");
|
||||
assert_eq!(tokenizer.token_to_id("h"), Some(b'h' as u32));
|
||||
assert_eq!(tokenizer.id_to_token(b'h' as u32).as_deref(), Some("h"));
|
||||
assert_eq!(tokenizer.vocab_size(), 256);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn configured_tokens_use_longest_prefix_matching() {
|
||||
let tokenizer = TestTokenizer::new()
|
||||
.with_regular_token("<image>", 999)
|
||||
.with_regular_token("<image></image>", 1000);
|
||||
|
||||
assert_eq!(
|
||||
tokenizer.encode("a<image></image>b", false).unwrap(),
|
||||
vec![b'a' as u32, 1000, b'b' as u32,]
|
||||
);
|
||||
assert_eq!(
|
||||
tokenizer.decode(&[b'a' as u32, 1000, b'b' as u32], false).unwrap(),
|
||||
"a<image></image>b"
|
||||
);
|
||||
assert_eq!(tokenizer.token_to_id("<image>"), Some(999));
|
||||
assert_eq!(
|
||||
tokenizer.id_to_token(1000).as_deref(),
|
||||
Some("<image></image>")
|
||||
);
|
||||
assert_eq!(tokenizer.vocab_size(), 1001);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_ascii_text_roundtrips_through_buffered_byte_decode() {
|
||||
let tokenizer = TestTokenizer::new();
|
||||
let text = "你好, café, 🚀";
|
||||
|
||||
let ids = tokenizer.encode(text, false).unwrap();
|
||||
assert_eq!(
|
||||
ids,
|
||||
text.as_bytes().iter().copied().map(u32::from).collect::<Vec<_>>()
|
||||
);
|
||||
assert_eq!(tokenizer.decode(&ids, false).unwrap(), text);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn buffered_byte_decode_flushes_around_configured_tokens() {
|
||||
let tokenizer = TestTokenizer::new()
|
||||
.with_regular_token("<image>", 999)
|
||||
.with_special_token("<skip>", 1000);
|
||||
|
||||
assert_eq!(
|
||||
tokenizer.encode("你<image>好<skip>🚀", false).unwrap(),
|
||||
vec![228, 189, 160, 999, 229, 165, 189, 1000, 240, 159, 154, 128]
|
||||
);
|
||||
assert_eq!(
|
||||
tokenizer
|
||||
.decode(
|
||||
&[228, 189, 160, 999, 229, 165, 189, 1000, 240, 159, 154, 128],
|
||||
false
|
||||
)
|
||||
.unwrap(),
|
||||
"你<image>好<skip>🚀"
|
||||
);
|
||||
assert_eq!(
|
||||
tokenizer
|
||||
.decode(
|
||||
&[228, 189, 160, 999, 229, 165, 189, 1000, 240, 159, 154, 128],
|
||||
true
|
||||
)
|
||||
.unwrap(),
|
||||
"你<image>好🚀"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_utf8_bytes_decode_lossily_as_a_sequence() {
|
||||
let tokenizer = TestTokenizer::new();
|
||||
|
||||
assert_eq!(tokenizer.decode(&[0xE4, 0xBD], false).unwrap(), "\u{FFFD}");
|
||||
assert_eq!(
|
||||
tokenizer.decode(&[0xFF, b'a' as u32], false).unwrap(),
|
||||
"\u{FFFD}a"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn special_tokens_respect_skip_special_tokens() {
|
||||
let tokenizer = TestTokenizer::new()
|
||||
.with_bos_token("<bos>", 256)
|
||||
.with_special_token("<think>", 0xF001)
|
||||
.with_regular_token("</think>", 0xF002);
|
||||
|
||||
assert_eq!(
|
||||
tokenizer.encode("<think>x</think>", true).unwrap(),
|
||||
vec![256, 0xF001, b'x' as u32, 0xF002,]
|
||||
);
|
||||
assert_eq!(
|
||||
tokenizer.decode(&[256, 0xF001, b'x' as u32, 0xF002], false).unwrap(),
|
||||
"<bos><think>x</think>"
|
||||
);
|
||||
assert_eq!(
|
||||
tokenizer.decode(&[256, 0xF001, b'x' as u32, 0xF002], true).unwrap(),
|
||||
"x</think>"
|
||||
);
|
||||
assert!(tokenizer.is_special_id(0xF001));
|
||||
assert!(!tokenizer.is_special_id(0xF002));
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "configured test token id 255 overlaps byte fallback range 0..=255")]
|
||||
fn configured_token_id_must_stay_outside_byte_range() {
|
||||
let _ = TestTokenizer::new().with_regular_token("<token>", 255);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "configured test token text \"a\" overlaps byte fallback token text")]
|
||||
fn configured_token_text_must_not_shadow_byte_tokens() {
|
||||
let _ = TestTokenizer::new().with_regular_token("a", 256);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(
|
||||
expected = "configured test token text \"<token>\" was registered more than once"
|
||||
)]
|
||||
fn configured_token_text_must_be_unique() {
|
||||
let _ = TestTokenizer::new()
|
||||
.with_regular_token("<token>", 256)
|
||||
.with_regular_token("<token>", 257);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "configured test token id 256 was registered more than once")]
|
||||
fn configured_token_id_must_be_unique() {
|
||||
let _ = TestTokenizer::new()
|
||||
.with_regular_token("<token-a>", 256)
|
||||
.with_regular_token("<token-b>", 256);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_decode_is_strict_by_default_and_configurable() {
|
||||
let strict = TestTokenizer::new();
|
||||
assert!(strict.decode(&[300], false).is_err());
|
||||
assert_eq!(
|
||||
TestTokenizer::new()
|
||||
.with_unknown_decode(UnknownDecode::Empty)
|
||||
.decode(&[b'a' as u32, 300, b'b' as u32], false)
|
||||
.unwrap(),
|
||||
"ab"
|
||||
);
|
||||
assert_eq!(
|
||||
TestTokenizer::new()
|
||||
.with_unknown_decode(UnknownDecode::Replacement)
|
||||
.decode(&[300], false)
|
||||
.unwrap(),
|
||||
"\u{FFFD}"
|
||||
);
|
||||
assert_eq!(strict.id_to_token(300), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_vocab_size_overrides_inferred_bound() {
|
||||
let tokenizer = TestTokenizer::new()
|
||||
.with_regular_token("<high>", 10_000)
|
||||
.with_vocab_size(20_000);
|
||||
|
||||
assert_eq!(tokenizer.vocab_size(), 20_000);
|
||||
assert_eq!(tokenizer.id_to_token(10_000).as_deref(), Some("<high>"));
|
||||
}
|
||||
}
|
||||
@@ -777,6 +777,7 @@ class precompiled_wheel_utils:
|
||||
"vllm/vllm_flash_attn/_vllm_fa3_C.abi3.so",
|
||||
"vllm/cumem_allocator.abi3.so",
|
||||
"vllm/spinloop.abi3.so",
|
||||
"vllm/fs_io_C.abi3.so",
|
||||
# ROCm-specific libraries
|
||||
"vllm/_rocm_C.abi3.so",
|
||||
}
|
||||
@@ -1104,6 +1105,7 @@ if _is_cuda() or _is_hip():
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
ext_modules.append(CMakeExtension(name="vllm.spinloop"))
|
||||
ext_modules.append(CMakeExtension(name="vllm.fs_io_C"))
|
||||
|
||||
if _is_hip():
|
||||
ext_modules.append(CMakeExtension(name="vllm._rocm_C"))
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+7
-1
@@ -1587,7 +1587,13 @@ class AssetHandler(http.server.BaseHTTPRequestHandler):
|
||||
self.send_header("Content-Type", ctype)
|
||||
self.send_header("Content-Length", str(len(data)))
|
||||
self.end_headers()
|
||||
self.wfile.write(data)
|
||||
try:
|
||||
self.wfile.write(data)
|
||||
except (BrokenPipeError, ConnectionResetError) as e:
|
||||
logger.debug(
|
||||
"Client disconnected while serving test asset %s: %r", filename, e
|
||||
)
|
||||
self.close_connection = True
|
||||
|
||||
|
||||
def _find_free_port() -> int:
|
||||
|
||||
@@ -17,6 +17,7 @@ import ray
|
||||
from vllm import LLM
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.engine.arg_utils import EngineArgs
|
||||
from vllm.v1.executor import ray_executor_v2
|
||||
from vllm.v1.executor.ray_executor_v2 import RayExecutorV2
|
||||
|
||||
pytestmark = pytest.mark.usefixtures("enable_ray_v2_backend")
|
||||
@@ -95,6 +96,43 @@ def assert_executor(executor, tp_size, pp_size):
|
||||
assert handle.node_id is not None
|
||||
|
||||
|
||||
def test_select_tcpstore_port_seeds_disjoint_windows(monkeypatch):
|
||||
"""Co-located DP engines scan distinct, adjacent port windows, so two
|
||||
engines on a node cannot pick the same TCPStore port."""
|
||||
requested = []
|
||||
|
||||
def fake_get_open_port(start_port, max_attempts):
|
||||
requested.append((start_port, max_attempts))
|
||||
return start_port
|
||||
|
||||
monkeypatch.setattr(ray_executor_v2, "_get_open_port", fake_get_open_port)
|
||||
|
||||
ports = [
|
||||
RayExecutorV2._select_tcpstore_port(rank, master_port=29500)
|
||||
for rank in range(4)
|
||||
]
|
||||
|
||||
assert requested == [(29600, 32), (29632, 32), (29664, 32), (29696, 32)]
|
||||
assert len(set(ports)) == 4
|
||||
|
||||
|
||||
def test_select_tcpstore_port_non_dp_uses_random(monkeypatch):
|
||||
"""A non-DP engine has no local rank and uses a random port."""
|
||||
monkeypatch.setattr(ray_executor_v2, "get_open_port", lambda: 54321)
|
||||
assert RayExecutorV2._select_tcpstore_port(None, master_port=29500) == 54321
|
||||
|
||||
|
||||
def test_select_tcpstore_port_full_window_uses_random(monkeypatch):
|
||||
"""A fully occupied window falls back to a random port."""
|
||||
|
||||
def raise_full(start_port, max_attempts):
|
||||
raise RuntimeError("no open port")
|
||||
|
||||
monkeypatch.setattr(ray_executor_v2, "_get_open_port", raise_full)
|
||||
monkeypatch.setattr(ray_executor_v2, "get_open_port", lambda: 54321)
|
||||
assert RayExecutorV2._select_tcpstore_port(0, master_port=29500) == 54321
|
||||
|
||||
|
||||
@pytest.mark.parametrize("tp_size, pp_size", [(1, 1), (2, 1), (4, 1), (2, 2)])
|
||||
def test_ray_v2_executor(tp_size, pp_size):
|
||||
"""Validate RayExecutorV2 with various TP/PP configs."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
from collections.abc import Callable, Iterator
|
||||
from contextlib import contextmanager
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
# Test different image extensions (JPG/PNG) and formats (gray/RGB/RGBA)
|
||||
TEST_IMAGE_ASSETS = [
|
||||
@@ -8,3 +13,70 @@ TEST_IMAGE_ASSETS = [
|
||||
"1280px-Venn_diagram_rgb.svg.png", # "https://vllm-public-assets.s3.us-west-2.amazonaws.com/vision_model_images/1280px-Venn_diagram_rgb.svg.png",
|
||||
"RGBA_comp.png", # "https://vllm-public-assets.s3.us-west-2.amazonaws.com/vision_model_images/RGBA_comp.png",
|
||||
]
|
||||
|
||||
|
||||
def _shutdown_llm(llm: Any, gpu_memory_utilization: float) -> None:
|
||||
from vllm.distributed import cleanup_dist_env_and_memory
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
try:
|
||||
shutdown_timeout = 60.0 if current_platform.is_rocm() else None
|
||||
llm.llm_engine.engine_core.shutdown(timeout=shutdown_timeout)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
del llm
|
||||
|
||||
try:
|
||||
import torch
|
||||
|
||||
torch._dynamo.reset()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
if current_platform.is_rocm():
|
||||
from tests.utils import wait_for_rocm_memory_to_settle
|
||||
|
||||
wait_for_rocm_memory_to_settle(threshold_ratio=1.0 - gpu_memory_utilization)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def managed_llm(*args: Any, **kwargs: Any) -> Iterator[Any]:
|
||||
from vllm import LLM
|
||||
|
||||
llm = LLM(*args, **kwargs)
|
||||
gpu_memory_utilization = (
|
||||
llm.llm_engine.vllm_config.cache_config.gpu_memory_utilization
|
||||
)
|
||||
try:
|
||||
yield llm
|
||||
finally:
|
||||
_shutdown_llm(llm, gpu_memory_utilization)
|
||||
|
||||
|
||||
def _make_managed_llm_factory() -> Iterator[Callable[..., Any]]:
|
||||
from vllm import LLM
|
||||
|
||||
llms: list[tuple[Any, float]] = []
|
||||
|
||||
def make_llm(*args: Any, **kwargs: Any) -> Any:
|
||||
llm = LLM(*args, **kwargs)
|
||||
gpu_memory_utilization = (
|
||||
llm.llm_engine.vllm_config.cache_config.gpu_memory_utilization
|
||||
)
|
||||
llms.append((llm, gpu_memory_utilization))
|
||||
return llm
|
||||
|
||||
try:
|
||||
yield make_llm
|
||||
finally:
|
||||
while llms:
|
||||
llm, gpu_memory_utilization = llms.pop()
|
||||
_shutdown_llm(llm, gpu_memory_utilization)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def multimodal_llm_factory() -> Iterator[Callable[..., Any]]:
|
||||
yield from _make_managed_llm_factory()
|
||||
|
||||
@@ -1,19 +1,13 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import weakref
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.entrypoints.multimodal.conftest import TEST_IMAGE_ASSETS
|
||||
from vllm import LLM
|
||||
from vllm.distributed import cleanup_dist_env_and_memory
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def vision_llm():
|
||||
# pytest caches the fixture so we use weakref.proxy to
|
||||
# enable garbage collection
|
||||
llm = LLM(
|
||||
def vision_llm(multimodal_llm_factory):
|
||||
return multimodal_llm_factory(
|
||||
model="microsoft/Phi-3.5-vision-instruct",
|
||||
max_model_len=4096,
|
||||
max_num_seqs=5,
|
||||
@@ -23,12 +17,6 @@ def vision_llm():
|
||||
seed=0,
|
||||
)
|
||||
|
||||
yield weakref.proxy(llm)
|
||||
|
||||
del llm
|
||||
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"image_urls", [[TEST_IMAGE_ASSETS[0], TEST_IMAGE_ASSETS[1]]], indirect=True
|
||||
|
||||
@@ -69,6 +69,7 @@ def test_inject_into_mm_cache(
|
||||
image_urls,
|
||||
mm_processor_cache_type,
|
||||
caplog_vllm,
|
||||
multimodal_llm_factory,
|
||||
):
|
||||
"""Test that inject_into_mm_cache() injects pre-processed mm_kwargs into
|
||||
the processor cache and MM cache hit metrics are updated correctly.
|
||||
@@ -78,7 +79,7 @@ def test_inject_into_mm_cache(
|
||||
2. Extract cached kwargs, call inject_into_mm_cache with a new hash,
|
||||
then generate with a pre-rendered input -> verifies injection works
|
||||
"""
|
||||
llm = LLM(
|
||||
llm = multimodal_llm_factory(
|
||||
model="llava-hf/llava-1.5-7b-hf",
|
||||
max_model_len=4096,
|
||||
max_num_seqs=5,
|
||||
@@ -145,11 +146,12 @@ def test_inject_into_mm_cache(
|
||||
def test_inject_into_mm_cache_without_cache(
|
||||
num_gpus_available,
|
||||
image_urls,
|
||||
multimodal_llm_factory,
|
||||
):
|
||||
"""Test that inject_into_mm_cache works gracefully when processor cache
|
||||
is disabled (mm_processor_cache_gb=0). Should not crash.
|
||||
"""
|
||||
llm = LLM(
|
||||
llm = multimodal_llm_factory(
|
||||
model="llava-hf/llava-1.5-7b-hf",
|
||||
max_model_len=4096,
|
||||
max_num_seqs=5,
|
||||
|
||||
@@ -61,8 +61,9 @@ def test_mm_cache_stats(
|
||||
image_urls,
|
||||
mm_processor_cache_type,
|
||||
caplog_vllm,
|
||||
multimodal_llm_factory,
|
||||
):
|
||||
llm = LLM(
|
||||
llm = multimodal_llm_factory(
|
||||
model="llava-hf/llava-1.5-7b-hf",
|
||||
max_model_len=4096,
|
||||
max_num_seqs=5,
|
||||
|
||||
@@ -1,13 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import weakref
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.entrypoints.multimodal.conftest import managed_llm
|
||||
from vllm import LLM, SamplingParams
|
||||
from vllm.assets.image import ImageAsset
|
||||
from vllm.distributed import cleanup_dist_env_and_memory
|
||||
|
||||
MODEL = "llava-hf/llava-1.5-7b-hf"
|
||||
PROMPT = "USER: <image>\nDescribe this image briefly.\nASSISTANT:"
|
||||
@@ -17,20 +15,15 @@ TEXT_ONLY_PROMPT = "USER: What is 2 + 2?\nASSISTANT:"
|
||||
@pytest.fixture(scope="module")
|
||||
def llm():
|
||||
"""LLM with enable_mm_embeds=True and all modality limits zeroed out."""
|
||||
llm = LLM(
|
||||
with managed_llm(
|
||||
model=MODEL,
|
||||
max_model_len=2048,
|
||||
enforce_eager=True,
|
||||
gpu_memory_utilization=0.8,
|
||||
enable_mm_embeds=True,
|
||||
limit_mm_per_prompt={"image": 0},
|
||||
)
|
||||
|
||||
yield weakref.proxy(llm)
|
||||
|
||||
del llm
|
||||
|
||||
cleanup_dist_env_and_memory()
|
||||
) as llm:
|
||||
yield llm
|
||||
|
||||
|
||||
@pytest.mark.skip_global_cleanup
|
||||
|
||||
@@ -17,10 +17,9 @@ from vllm.entrypoints.openai.chat_completion.serving import OpenAIServingChat
|
||||
from vllm.entrypoints.openai.engine.protocol import GenerationError
|
||||
from vllm.entrypoints.openai.models.protocol import BaseModelPath
|
||||
from vllm.entrypoints.openai.models.serving import OpenAIServingModels
|
||||
from vllm.entrypoints.serve.render.serving import ServingRender
|
||||
from vllm.entrypoints.scale_out.render.serving import ServingRender
|
||||
from vllm.outputs import CompletionOutput, RequestOutput
|
||||
from vllm.renderers.hf import HfRenderer
|
||||
from vllm.renderers.online_derenderer import OnlineDerenderer
|
||||
from vllm.renderers.online_renderer import OnlineRenderer
|
||||
from vllm.tokenizers.registry import cached_tokenizer_from_config
|
||||
from vllm.v1.engine.async_llm import AsyncLLM
|
||||
@@ -206,15 +205,8 @@ def _build_serving_render(engine: AsyncLLM) -> ServingRender:
|
||||
chat_template=None,
|
||||
chat_template_content_format="auto",
|
||||
)
|
||||
online_derenderer = OnlineDerenderer(
|
||||
model_config=engine.model_config,
|
||||
renderer=engine.renderer,
|
||||
request_logger=None,
|
||||
chat_template=None,
|
||||
chat_template_content_format="auto",
|
||||
)
|
||||
|
||||
serving_render = ServingRender(models, online_renderer, online_derenderer)
|
||||
serving_render = ServingRender(models, online_renderer)
|
||||
|
||||
async def _fake_preprocess_chat(*args, **kwargs):
|
||||
# return conversation, engine_inputs
|
||||
|
||||
@@ -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
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user