From f69d6fc28a9fb5c02498e1b33824f0e5c2dac387 Mon Sep 17 00:00:00 2001 From: Jeremy Zhang Date: Sat, 12 Sep 2026 01:33:38 +0800 Subject: [PATCH] [OpenAI] Propagate PD routing metadata through /v1/responses (#35503) --- .../sglang/srt/entrypoints/openai/protocol.py | 17 +++++++++++ .../entrypoints/openai/serving_responses.py | 9 ++++++ .../openai/test_responses_protocol.py | 19 ++++++++++++ .../openai/test_serving_responses.py | 30 +++++++++++++++++-- 4 files changed, 73 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index 3b99baea6..5199cf7f8 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -1652,6 +1652,18 @@ class ResponsesRequest(BaseModel): default=None, description="Cache salt for request caching" ) + # For PD disaggregation + bootstrap_host: Optional[Union[List[str], str]] = None + bootstrap_port: Optional[Union[List[Optional[int]], int]] = None + bootstrap_room: Optional[Union[List[int], int]] = None + + # For DP routing — external router assigns a specific DP worker + routed_dp_rank: Optional[int] = None + # For PD disagg — hint telling decode which prefill DP worker has the KV cache + disagg_prefill_dp_rank: Optional[int] = None + # Deprecated: use routed_dp_rank instead + data_parallel_rank: Optional[int] = None + # SGLang sampling extras. ``None`` defers to ``--preferred-sampling-params``. frequency_penalty: float = 0.0 presence_penalty: float = 0.0 @@ -1669,6 +1681,11 @@ class ResponsesRequest(BaseModel): "repetition_penalty": 1.0, } + @model_validator(mode="before") + @classmethod + def _handle_deprecated_dp_rank(cls, values): + return _migrate_deprecated_dp_rank(values) + @model_validator(mode="before") @classmethod def normalize_reasoning_to_thinking(cls, values): diff --git a/python/sglang/srt/entrypoints/openai/serving_responses.py b/python/sglang/srt/entrypoints/openai/serving_responses.py index cdbc2eb38..052f4c3f4 100644 --- a/python/sglang/srt/entrypoints/openai/serving_responses.py +++ b/python/sglang/srt/entrypoints/openai/serving_responses.py @@ -435,6 +435,10 @@ class OpenAIServingResponses(OpenAIServingChat): else {} ) + effective_routed_dp_rank = self.extract_routed_dp_rank_from_header( + raw_request, request.routed_dp_rank + ) + adapted_request = GenerateReqInput( **prompt_kwargs, **logprob_kwargs, @@ -464,6 +468,11 @@ class OpenAIServingResponses(OpenAIServingChat): session_id=request.session_id, extra_key=request.extra_key, cache_salt=request.cache_salt, + bootstrap_host=request.bootstrap_host, + bootstrap_port=request.bootstrap_port, + bootstrap_room=request.bootstrap_room, + routed_dp_rank=effective_routed_dp_rank, + disagg_prefill_dp_rank=request.disagg_prefill_dp_rank, # background+stream streams on this connection, so don't detach. background=request.background and not request.stream, require_reasoning=require_reasoning, diff --git a/test/registered/unit/entrypoints/openai/test_responses_protocol.py b/test/registered/unit/entrypoints/openai/test_responses_protocol.py index 3c10962d6..f186f31bb 100644 --- a/test/registered/unit/entrypoints/openai/test_responses_protocol.py +++ b/test/registered/unit/entrypoints/openai/test_responses_protocol.py @@ -28,6 +28,25 @@ def _in_progress_response(request: ResponsesRequest) -> ResponsesResponse: class ResponsesRequestTestCase(CustomTestCase): + def test_pd_routing_fields(self): + with self.assertWarns(DeprecationWarning): + request = ResponsesRequest( + model="x", + input="hi", + bootstrap_host="10.0.0.1", + bootstrap_port=8998, + bootstrap_room=42, + data_parallel_rank=1, + disagg_prefill_dp_rank=0, + store=False, + ) + + self.assertEqual(request.bootstrap_host, "10.0.0.1") + self.assertEqual(request.bootstrap_port, 8998) + self.assertEqual(request.bootstrap_room, 42) + self.assertEqual(request.routed_dp_rank, 1) + self.assertEqual(request.disagg_prefill_dp_rank, 0) + def test_function_tool_accepted(self): request = ResponsesRequest( model="x", diff --git a/test/registered/unit/entrypoints/openai/test_serving_responses.py b/test/registered/unit/entrypoints/openai/test_serving_responses.py index bac16c62c..c3b27c9ff 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_responses.py +++ b/test/registered/unit/entrypoints/openai/test_serving_responses.py @@ -922,7 +922,7 @@ class EnginePassthroughTestCase(CustomTestCase): """Both flags cross hops with no type contract, and dropping either fails silently.""" - def _capture(self, serving, request): + def _capture(self, serving, request, raw_request=None): # Let the real _process_messages run: it is the hop that turns # skip_special_tokens off, so mocking it would make that assertion vacuous. # chat_template_name=None routes it through the tokenizer's template @@ -954,9 +954,35 @@ class EnginePassthroughTestCase(CustomTestCase): yield context serving._generate_with_builtin_tools = fake_generate - asyncio.run(serving.create_responses(request)) + asyncio.run(serving.create_responses(request, raw_request=raw_request)) return captured + def test_pd_routing_fields_forwarded_to_engine(self): + serving = make_serving() + raw_request = Mock(headers={"x-data-parallel-rank": "2"}, state=Mock()) + + captured = self._capture( + serving, + ResponsesRequest( + model="x", + input="hi", + bootstrap_host="10.0.0.1", + bootstrap_port=8998, + bootstrap_room=42, + routed_dp_rank=1, + disagg_prefill_dp_rank=0, + store=False, + ), + raw_request=raw_request, + ) + + adapted_request = captured["adapted_request"] + self.assertEqual(adapted_request.bootstrap_host, "10.0.0.1") + self.assertEqual(adapted_request.bootstrap_port, 8998) + self.assertEqual(adapted_request.bootstrap_room, 42) + self.assertEqual(adapted_request.routed_dp_rank, 2) + self.assertEqual(adapted_request.disagg_prefill_dp_rank, 0) + def test_require_reasoning_forwarded_when_reasoning_parser_configured(self): serving = make_serving() serving.reasoning_parser = "deepseek-r1"