From 5aab054ec8ce6b6100fbfb7aafe67d632a7df3aa Mon Sep 17 00:00:00 2001 From: Sage <80211083+sagearc@users.noreply.github.com> Date: Tue, 8 Sep 2026 08:22:32 +0300 Subject: [PATCH] [rust-server] fix p/d bootstrap across dp listeners (#36234) --- python/sglang/srt/disaggregation/base/conn.py | 2 + .../sglang/srt/disaggregation/common/conn.py | 22 +++++ python/sglang/srt/disaggregation/prefill.py | 5 ++ python/sglang/srt/rust_server/server.py | 4 +- .../test_register_to_bootstrap.py | 85 +++++++++++++++++++ 5 files changed, 117 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 358ff95c4..2dea4485d 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -74,6 +74,8 @@ class KVArgs: page_size: int # for system dp system_dp_rank: int + # Local Rust /route registry port; None on scheduler ranks without a listener. + rust_http_port: Optional[int] # for pp prefill pp_rank: int prefill_start_layer: int diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 39763024a..7f486522e 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -828,6 +828,28 @@ class CommonKVManager(BaseKVManager): "prefill_http_port": get_serving().port, } + if envs.SGLANG_RUST_SERVER.get() and self.attn_dp_size > 1: + topology_rows = get_world_group().all_gather_object(payload) + # Every scheduler contributes a topology row. Only the scheduler + # ranks that own a Rust listener populate their local registry. + if self.kv_args.rust_http_port is None: + return + registry_host = {"0.0.0.0": "127.0.0.1", "::": "::1"}.get( + self.bootstrap_host, self.bootstrap_host + ) + registry = NetworkAddress(registry_host, self.kv_args.rust_http_port) + url = f"{registry.to_url()}/route" + for topology_row in topology_rows: + registry_row = { + **topology_row, + "prefill_http_port": self.kv_args.rust_http_port, + } + self._register_topology_row(url, registry_row) + return + + self._register_topology_row(url, payload) + + def _register_topology_row(self, url: str, payload: Dict) -> None: max_retries, initial_delay, max_delay = 5, 1.0, 30.0 for attempt in range(max_retries): try: diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 20e616590..9ecdb2775 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -208,6 +208,11 @@ class PrefillBootstrapQueue: kv_args.engine_rank = self.tp_rank kv_args.pp_rank = self.pp_rank kv_args.system_dp_rank = self.scheduler.ps.dp_rank + kv_args.rust_http_port = ( + self.scheduler.rust_server.http_port + if self.scheduler.rust_server is not None + else None + ) kv_args.kv_cache_dtype_str = ( self.scheduler.tp_worker.model_runner.kv_cache_dtype_str ) diff --git a/python/sglang/srt/rust_server/server.py b/python/sglang/srt/rust_server/server.py index 64317d04d..8dd564f7b 100644 --- a/python/sglang/srt/rust_server/server.py +++ b/python/sglang/srt/rust_server/server.py @@ -55,10 +55,12 @@ class RustServer: def __init__( self, server: Server, + http_port: int, mm_spec: Optional[RustMmSpec] = None, max_per_poll: int = 256, ): self.server = server + self.http_port = http_port self.mm_spec = mm_spec self._max_per_poll = max_per_poll @@ -151,7 +153,7 @@ class RustServer: dp_note, ) - return cls(server, mm_spec=mm_spec) + return cls(server, http_port=listen_port, mm_spec=mm_spec) def wait_request(self, timeout_ms: int) -> None: """Block until a request is pushed into the in-process ring or the timeout diff --git a/test/registered/unit/disaggregation/test_register_to_bootstrap.py b/test/registered/unit/disaggregation/test_register_to_bootstrap.py index 76b55c27a..285d3f425 100644 --- a/test/registered/unit/disaggregation/test_register_to_bootstrap.py +++ b/test/registered/unit/disaggregation/test_register_to_bootstrap.py @@ -8,6 +8,7 @@ register_cpu_ci(est_time=11, suite="base-a-test-cpu") import unittest from unittest.mock import MagicMock, call, patch +from sglang.srt.environ import envs from sglang.srt.runtime_context import get_context from sglang.test.test_utils import CustomTestCase @@ -199,6 +200,86 @@ class TestRegisterToBootstrap(CustomTestCase): url_used = mock_put.call_args[0][0] self.assertIn("10.0.0.1", url_used) + @patch("sglang.srt.disaggregation.common.conn.requests.put") + @patch("sglang.srt.disaggregation.common.conn.get_world_group") + def test_rust_attention_dp_replicates_complete_topology_across_hosts( + self, mock_world_group, mock_put + ): + success_resp = MagicMock() + success_resp.status_code = 200 + mock_put.return_value = success_resp + + schedulers = ( + (0, 0, "10.0.0.1", 17000, 8765), + (0, 1, "10.0.0.1", 17001, None), + (1, 0, "10.0.0.2", 17002, 8766), + (1, 1, "10.0.0.2", 17003, None), + ) + + def gather_topology(payload): + return [ + { + **payload, + "attn_dp_rank": dp_rank, + "attn_tp_rank": tp_rank, + "rank_ip": host, + "rank_port": rank_port, + } + for dp_rank, tp_rank, host, rank_port, _ in schedulers + ] + + mock_world_group.return_value.all_gather_object.side_effect = gather_topology + + with envs.SGLANG_RUST_SERVER.override(True): + for dp_rank, tp_rank, local_ip, _, rust_http_port in schedulers: + manager = self._make_manager() + manager.attn_dp_size = 2 + manager.attn_dp_rank = dp_rank + manager.attn_tp_size = 2 + manager.attn_tp_rank = tp_rank + manager.local_ip = local_ip + manager.bootstrap_host = local_ip + manager.kv_args.rust_http_port = rust_http_port + manager.register_to_bootstrap() + + topology_by_registry = {} + for put_call in mock_put.call_args_list: + payload = put_call.kwargs["json"] + topology_by_registry.setdefault(put_call.args[0], {})[ + (payload["attn_dp_rank"], payload["attn_tp_rank"]) + ] = (payload["rank_ip"], payload["rank_port"]) + complete_topology = { + (dp, tp): (host, rank_port) for dp, tp, host, rank_port, _ in schedulers + } + self.assertEqual( + topology_by_registry, + { + "http://10.0.0.1:8765/route": complete_topology, + "http://10.0.0.2:8766/route": complete_topology, + }, + ) + self.assertEqual(mock_put.call_count, 8) + self.assertEqual( + { + (put_call.args[0], put_call.kwargs["json"]["prefill_http_port"]) + for put_call in mock_put.call_args_list + }, + { + ("http://10.0.0.1:8765/route", 8765), + ("http://10.0.0.2:8766/route", 8766), + }, + ) + self.assertEqual( + [ + ( + gather_call.args[0]["attn_dp_rank"], + gather_call.args[0]["attn_tp_rank"], + ) + for gather_call in mock_world_group.return_value.all_gather_object.call_args_list + ], + [(dp, tp) for dp, tp, _, _, _ in schedulers], + ) + @patch("sglang.srt.disaggregation.common.conn.time") @patch("sglang.srt.disaggregation.common.conn.requests.put") def test_wildcard_host_0000_uses_ipv4_loopback(self, mock_put, mock_time): @@ -258,6 +339,9 @@ class TestRegisterToBootstrap(CustomTestCase): mgr.register_to_bootstrap = CommonKVManager.register_to_bootstrap.__get__( mgr, CommonKVManager ) + mgr._register_topology_row = CommonKVManager._register_topology_row.__get__( + mgr, CommonKVManager + ) # Set attributes that register_to_bootstrap reads mgr.dist_init_addr = dist_init_addr @@ -278,6 +362,7 @@ class TestRegisterToBootstrap(CustomTestCase): mgr.kv_args = MagicMock() mgr.kv_args.page_size = 16 + mgr.kv_args.rust_http_port = None # Resolved per-runner value threaded through KVArgs (the payload field). mgr.kv_cache_dtype_str = "auto"