observability: publish the generated forward-pass-metrics endpoint to the bags (#33337)

When --forward-pass-metrics-ipc-name is left unset the reporter generates an
endpoint and has to hand it to external consumers (the documented contract is
that they read it back from the server config). That readback is the scheduler's
get_internal_state, which already reports get_context().resolved_server_args_dict(),
so the write moves to get_context().override and the read alongside it to
get_observability() — the endpoint still shows up in /server_info's
internal_states, and the ServerArgs instance stops being a message bus.

The test's server_args stand-in (a SimpleNamespace with a hand-rolled override)
becomes a real published config, so the reporter exercises the same accessors as
production.

Writer ratchet 19 -> 18.
This commit is contained in:
Cheng Wan
2026-08-02 21:24:11 -07:00
committed by GitHub
parent 0b3e8bedd1
commit aa3bbbc6e8
3 changed files with 30 additions and 22 deletions
@@ -26,7 +26,7 @@ from sglang.srt.observability.metrics_collector import (
SchedulerStats, SchedulerStats,
compute_routing_key_stats, compute_routing_key_stats,
) )
from sglang.srt.runtime_context import get_spec from sglang.srt.runtime_context import get_context, get_observability, get_spec
from sglang.srt.utils.device_timer import DeviceTimer from sglang.srt.utils.device_timer import DeviceTimer
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
@@ -215,11 +215,11 @@ class SchedulerMetricsReporter:
self.scheduler._fpm_worker_id = ( self.scheduler._fpm_worker_id = (
self.scheduler.server_args.forward_pass_metrics_worker_id self.scheduler.server_args.forward_pass_metrics_worker_id
) )
base_endpoint = self.scheduler.server_args.forward_pass_metrics_ipc_name base_endpoint = get_observability().forward_pass_metrics_ipc_name
if base_endpoint is None: if base_endpoint is None:
ipc_path = tempfile.NamedTemporaryFile(delete=False).name ipc_path = tempfile.NamedTemporaryFile(delete=False).name
base_endpoint = f"ipc://{ipc_path}" base_endpoint = f"ipc://{ipc_path}"
self.scheduler.server_args.override( get_context().override(
"metrics_reporter.ipc_endpoint", "metrics_reporter.ipc_endpoint",
forward_pass_metrics_ipc_name=base_endpoint, forward_pass_metrics_ipc_name=base_endpoint,
) )
@@ -1,3 +1,4 @@
from sglang.srt.runtime_context import get_context, get_observability
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu") register_cpu_ci(est_time=5, suite="base-a-test-cpu")
@@ -70,22 +71,19 @@ class _DummyPublisherThread:
pass pass
def _fake_server_args(**fields): def _publish_server_args(test, **fields):
"""server_args stand-in: carries fields and the override() entry point.""" """Publish a config for the reporter under test and return the instance."""
fields.setdefault("decode_log_interval", 40) fields.setdefault("decode_log_interval", 40)
ns = types.SimpleNamespace(**fields) override = get_context().override_server_args(**fields)
server_args = override.install()
def _override(source, **updates): test.addCleanup(override.restore)
for key, value in updates.items(): return server_args
setattr(ns, key, value)
ns.override = _override
return ns
def _make_reporter(scheduler) -> SchedulerMetricsReporter: def _make_reporter(test, scheduler) -> SchedulerMetricsReporter:
if not hasattr(scheduler, "server_args"): if not hasattr(scheduler, "server_args"):
scheduler.server_args = _fake_server_args( scheduler.server_args = _publish_server_args(
test,
enable_metrics=False, enable_metrics=False,
enable_metrics_for_all_schedulers=False, enable_metrics_for_all_schedulers=False,
kv_events_config=None, kv_events_config=None,
@@ -133,7 +131,7 @@ class TestForwardPassMetrics(unittest.TestCase):
self.scheduler._fpm_gpu_time_acc = 0.0 self.scheduler._fpm_gpu_time_acc = 0.0
self.scheduler.waiting_queue = [] self.scheduler.waiting_queue = []
self.scheduler.disaggregation_mode = DisaggregationMode.NULL self.scheduler.disaggregation_mode = DisaggregationMode.NULL
self.reporter = _make_reporter(self.scheduler) self.reporter = _make_reporter(self, self.scheduler)
self.scheduler.enable_fpm = True self.scheduler.enable_fpm = True
def _make_batch(self, **overrides): def _make_batch(self, **overrides):
@@ -274,7 +272,8 @@ class TestForwardPassMetrics(unittest.TestCase):
def test_init_metrics_uses_server_worker_id(self): def test_init_metrics_uses_server_worker_id(self):
scheduler = types.SimpleNamespace() scheduler = types.SimpleNamespace()
scheduler.server_args = _fake_server_args( scheduler.server_args = _publish_server_args(
self,
enable_metrics=False, enable_metrics=False,
enable_metrics_for_all_schedulers=False, enable_metrics_for_all_schedulers=False,
extra_metric_labels=None, extra_metric_labels=None,
@@ -290,7 +289,7 @@ class TestForwardPassMetrics(unittest.TestCase):
"sglang.srt.observability.forward_pass_metrics._FpmPublisherThread", "sglang.srt.observability.forward_pass_metrics._FpmPublisherThread",
_DummyPublisherThread, _DummyPublisherThread,
): ):
reporter = _make_reporter(scheduler) reporter = _make_reporter(self, scheduler)
self.assertTrue(scheduler.enable_fpm) self.assertTrue(scheduler.enable_fpm)
self.assertEqual(scheduler._fpm_worker_id, "endpoint-42") self.assertEqual(scheduler._fpm_worker_id, "endpoint-42")
@@ -298,11 +297,20 @@ class TestForwardPassMetrics(unittest.TestCase):
self.assertEqual(scheduler._fpm_publisher.worker_id, "endpoint-42") self.assertEqual(scheduler._fpm_publisher.worker_id, "endpoint-42")
self.assertEqual(scheduler._fpm_publisher.dp_rank, 2) self.assertEqual(scheduler._fpm_publisher.dp_rank, 2)
self.assertTrue(scheduler._fpm_publisher.endpoint.startswith("ipc://")) self.assertTrue(scheduler._fpm_publisher.endpoint.startswith("ipc://"))
self.assertIsNotNone(scheduler.server_args.forward_pass_metrics_ipc_name) # The bag is what makes the write a bag write: an instance mutation
# would still show up in the resolved dict through its ServerArgs base.
endpoint = get_observability().forward_pass_metrics_ipc_name
self.assertTrue(endpoint.startswith("ipc://"))
self.assertEqual(
get_context().resolved_server_args_dict()["forward_pass_metrics_ipc_name"],
endpoint,
)
self.assertIsNone(scheduler.server_args.forward_pass_metrics_ipc_name)
def test_init_fpm_disabled_on_non_last_pp_rank(self): def test_init_fpm_disabled_on_non_last_pp_rank(self):
scheduler = types.SimpleNamespace() scheduler = types.SimpleNamespace()
scheduler.server_args = _fake_server_args( scheduler.server_args = _publish_server_args(
self,
enable_metrics=False, enable_metrics=False,
enable_metrics_for_all_schedulers=False, enable_metrics_for_all_schedulers=False,
extra_metric_labels=None, extra_metric_labels=None,
@@ -318,7 +326,7 @@ class TestForwardPassMetrics(unittest.TestCase):
"sglang.srt.observability.forward_pass_metrics._FpmPublisherThread", "sglang.srt.observability.forward_pass_metrics._FpmPublisherThread",
_DummyPublisherThread, _DummyPublisherThread,
): ):
reporter = _make_reporter(scheduler) reporter = _make_reporter(self, scheduler)
self.assertFalse(scheduler.enable_fpm) self.assertFalse(scheduler.enable_fpm)
@@ -49,7 +49,7 @@ _EXCLUDED = (
"multimodal_gen", "multimodal_gen",
) )
_BASELINE = 19 _BASELINE = 18
class TestServerArgsWriterRatchet(CustomTestCase): class TestServerArgsWriterRatchet(CustomTestCase):