diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index a67e2034c..7ed41ba8c 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -263,6 +263,11 @@ class GenerateReqInput: # For DP routing — external router assigns a specific DP worker routed_dp_rank: Optional[int] = None + # Deprecated alias for `routed_dp_rank`, still accepted because + # sgl-model-gateway's dp-aware mode injects this spelling into every + # request it forwards (DPAwareWorker::prepare_request), and the OpenAI + # entrypoints and Engine.generate() accept it as well. + data_parallel_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 # Routing key for routing-key schedule policy @@ -365,6 +370,18 @@ class GenerateReqInput: ValueError: If inputs are not properly specified (e.g., none or all of text, input_ids, input_embeds are provided) """ + if self.data_parallel_rank is not None: + import warnings + + warnings.warn( + "'data_parallel_rank' is deprecated, use 'routed_dp_rank' instead.", + DeprecationWarning, + stacklevel=2, + ) + if self.routed_dp_rank is None: + self.routed_dp_rank = self.data_parallel_rank + self.data_parallel_rank = None + self._validate_inputs() self._determine_batch_size() if self.session_id is not None and self.session_params is not None: diff --git a/test/registered/unit/managers/test_io_struct.py b/test/registered/unit/managers/test_io_struct.py index 50d232822..3b46ed8e5 100644 --- a/test/registered/unit/managers/test_io_struct.py +++ b/test/registered/unit/managers/test_io_struct.py @@ -703,6 +703,30 @@ class TestGenerateReqInputNormalization(CustomTestCase): ) req.normalize_batch_and_arguments() + def test_data_parallel_rank_alias_maps_to_routed_dp_rank(self): + req = GenerateReqInput(text="Hello", sampling_params={}, data_parallel_rank=2) + req.normalize_batch_and_arguments() + self.assertEqual(req.routed_dp_rank, 2) + self.assertIsNone(req.data_parallel_rank) + + def test_data_parallel_rank_alias_does_not_override_routed_dp_rank(self): + req = GenerateReqInput( + text="Hello", sampling_params={}, data_parallel_rank=2, routed_dp_rank=1 + ) + req.normalize_batch_and_arguments() + self.assertEqual(req.routed_dp_rank, 1) + + def test_data_parallel_rank_alias_propagates_to_batch_items(self): + req = GenerateReqInput( + text=["Hello", "World"], + sampling_params=[{}, {}], + rid=["id1", "id2"], + data_parallel_rank=3, + ) + req.normalize_batch_and_arguments() + self.assertEqual(req[0].routed_dp_rank, 3) + self.assertEqual(req[1].routed_dp_rank, 3) + class TestEmbeddingReqInputGetItem(CustomTestCase): """Test EmbeddingReqInput.__getitem__."""