diff --git a/components/src/dynamo/vllm/llm_engine.py b/components/src/dynamo/vllm/llm_engine.py index 69101366a2d9..a6f53835a435 100644 --- a/components/src/dynamo/vllm/llm_engine.py +++ b/components/src/dynamo/vllm/llm_engine.py @@ -73,6 +73,7 @@ from .handlers import ( VllmEnginePauseController, + _apply_nvext_cache_salt, build_sampling_params, get_dp_range_for_worker, ) @@ -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( diff --git a/components/src/dynamo/vllm/tests/test_vllm_unit.py b/components/src/dynamo/vllm/tests/test_vllm_unit.py index f1dd0021dd43..24eeb16b2d50 100644 --- a/components/src/dynamo/vllm/tests/test_vllm_unit.py +++ b/components/src/dynamo/vllm/tests/test_vllm_unit.py @@ -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."""