Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions components/src/dynamo/vllm/llm_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,7 @@

from .handlers import (
VllmEnginePauseController,
_apply_nvext_cache_salt,
build_sampling_params,
get_dp_range_for_worker,
)
Expand Down Expand Up @@ -366,6 +367,7 @@ async def generate(

token_ids = request.get("token_ids", [])
prompt = TokensPrompt(prompt_token_ids=token_ids)
_apply_nvext_cache_salt(request, prompt)

# TODO: remove dict() once build_sampling_params accepts GenerateRequest
sampling_params = build_sampling_params(
Expand Down
45 changes: 45 additions & 0 deletions components/src/dynamo/vllm/tests/test_vllm_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -410,6 +410,51 @@ async def run_generate():
assert captured["enable_rl"] is True


def test_unified_generate_applies_nvext_cache_salt(monkeypatch):
from dynamo.common.constants import DisaggregationMode as CommonDisaggregationMode
from dynamo.vllm import llm_engine

captured = {}

def fake_build_sampling_params(
request, default_sampling_params, model_max_len=None, *, enable_rl=False
):
return SimpleNamespace(extra_args=None)

async def empty_generation():
if False:
yield None

def fake_generate(prompt, *args, **kwargs):
captured["prompt"] = prompt
return empty_generation()

engine = llm_engine.VllmLLMEngine(
SimpleNamespace(),
CommonDisaggregationMode.AGGREGATED,
served_model_name="test-model",
component="backend",
)
engine.engine_client = SimpleNamespace(generate=fake_generate)
engine._default_sampling_params = {}
engine._model_max_len = 4096

monkeypatch.setattr(llm_engine, "build_sampling_params", fake_build_sampling_params)

async def run_generate():
request = {
"token_ids": [1, 2, 3],
"extra_args": {"nvext": {"cache_salt": "tenant-a"}},
}
context = SimpleNamespace(id=lambda: "req", trace_headers=lambda: None)
async for _ in engine.generate(request, context):
pass

asyncio.run(run_generate())

assert captured["prompt"]["cache_salt"] == "tenant-a"


@pytest.mark.asyncio
async def test_unified_start_returns_normalized_served_model_name(monkeypatch):
"""Return the Dynamo-normalized served model name from EngineConfig."""
Expand Down
Loading