Skip to content
Closed
29 changes: 29 additions & 0 deletions bindings/python/src/smg/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,9 @@
import logging
import os
import random
import shutil
import signal
import tempfile
import socket
import subprocess
import sys
Expand Down Expand Up @@ -578,6 +580,7 @@ def __init__(self, backend: str, args: argparse.Namespace, backend_args: list[st
self.launcher: WorkerLauncher = BACKEND_LAUNCHERS[backend]()
self.workers: list[tuple[subprocess.Popen, int]] = []
self._shutting_down = False
self._prometheus_dir: str | None = None

# -- public API ---------------------------------------------------------

Expand All @@ -599,6 +602,15 @@ def run(self) -> None:
def _launch_workers(self) -> None:
ports = _find_available_ports(self.args.worker_base_port, self.args.data_parallel_size)
host = self.args.worker_host

if getattr(self.args, "connection_mode", "grpc") == "grpc":
self._prometheus_dir = tempfile.mkdtemp(prefix="smg_prometheus_")
os.environ["PROMETHEUS_MULTIPROC_DIR"] = self._prometheus_dir
logger.info(
"Set PROMETHEUS_MULTIPROC_DIR=%s for gRPC metrics collection",
self._prometheus_dir,
)

for dp_rank, port in enumerate(ports):
env = self.launcher.gpu_env(self.args, dp_rank)
proc = self.launcher.launch(self.args, self.backend_args, host, port, env)
Expand Down Expand Up @@ -669,6 +681,23 @@ def _cleanup_workers(self) -> None:
except (ProcessLookupError, OSError):
pass

self._cleanup_prometheus_dir()

def _cleanup_prometheus_dir(self) -> None:
"""Remove the temporary prometheus multiprocess directory and its .db files."""
if self._prometheus_dir is None:
return
try:
shutil.rmtree(self._prometheus_dir)
logger.info("Cleaned up PROMETHEUS_MULTIPROC_DIR=%s", self._prometheus_dir)
except OSError as e:
logger.warning(
"Failed to clean up PROMETHEUS_MULTIPROC_DIR=%s: %s",
self._prometheus_dir,
e,
)
self._prometheus_dir = None
Comment on lines +686 to +699

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Nit: After removing the directory, the PROMETHEUS_MULTIPROC_DIR environment variable still points to the now-deleted path. If any code (e.g., a straggling atexit handler or a library import) reads the env var after cleanup, it will reference a non-existent directory, which can cause confusing errors.

Consider also clearing the env var:

Suggested change
def _cleanup_prometheus_dir(self) -> None:
"""Remove the temporary prometheus multiprocess directory and its .db files."""
if self._prometheus_dir is None:
return
try:
shutil.rmtree(self._prometheus_dir)
logger.info("Cleaned up PROMETHEUS_MULTIPROC_DIR=%s", self._prometheus_dir)
except OSError as e:
logger.warning(
"Failed to clean up PROMETHEUS_MULTIPROC_DIR=%s: %s",
self._prometheus_dir,
e,
)
self._prometheus_dir = None
def _cleanup_prometheus_dir(self) -> None:
"""Remove the temporary prometheus multiprocess directory and its .db files."""
if self._prometheus_dir is None:
return
try:
shutil.rmtree(self._prometheus_dir)
logger.info("Cleaned up PROMETHEUS_MULTIPROC_DIR=%s", self._prometheus_dir)
except OSError as e:
logger.warning(
"Failed to clean up PROMETHEUS_MULTIPROC_DIR=%s: %s",
self._prometheus_dir,
e,
)
os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None)
self._prometheus_dir = None



# ---------------------------------------------------------------------------
# Entry point
Expand Down
66 changes: 65 additions & 1 deletion grpc_servicer/smg_grpc_servicer/sglang/request_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,6 +142,7 @@ class GrpcReqState:
# Metrics (same as TokenizerManager's ReqState)
time_stats: APIServerReqTimeStats
last_completion_tokens: int = 1
ttft_observed: bool = False

# Streaming state
stream_finished: bool = False
Expand Down Expand Up @@ -205,6 +206,33 @@ def __init__(

# Metrics
self.last_receive_tstamp = real_time()
self.metrics_collector = None
if server_args.enable_metrics:
try:
from sglang.srt.observability.metrics_collector import (
TokenizerMetricsCollector,
)

labels = {
"model_name": server_args.served_model_name,
}
self.metrics_collector = TokenizerMetricsCollector(
server_args=server_args,
labels=labels,
bucket_time_to_first_token=server_args.bucket_time_to_first_token,
bucket_e2e_request_latency=server_args.bucket_e2e_request_latency,
bucket_inter_token_latency=getattr(
server_args, "bucket_inter_token_latency", None
),
)
self._metrics_labels = labels
logger.info("TokenizerMetricsCollector initialized for gRPC request-level metrics")
except Exception:
logger.warning(
"Failed to initialize TokenizerMetricsCollector, "
"request-level metrics will be unavailable",
exc_info=True,
)
Comment on lines +209 to +235

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🧹 Nitpick | 🔵 Trivial

Metrics init: minor hardening — _metrics_labels default and narrower except.

Two small points:

  1. self._metrics_labels is only assigned inside the successful try branch. All current read sites are gated by self.metrics_collector is not None, so this is safe today, but any future reader that forgets the guard will hit AttributeError. Defining self._metrics_labels = {} before the try makes the invariant robust to refactors.
  2. except Exception: swallows everything including KeyboardInterrupt subclasses (not in Py3, but) and, more importantly, masks programming errors like TypeError from passing an unexpected kwarg (e.g., bucket_inter_token_latency not being accepted by older sglang versions). Consider logging the exception type at least, or catching ImportError + a narrower runtime error separately so API-mismatch bugs surface during development.
Proposed diff
         self.metrics_collector = None
+        self._metrics_labels: dict[str, str] = {}
         if server_args.enable_metrics:
             try:
                 from sglang.srt.observability.metrics_collector import (
                     TokenizerMetricsCollector,
                 )
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@grpc_servicer/smg_grpc_servicer/sglang/request_manager.py` around lines 209 -
235, Initialize self._metrics_labels = {} before the try so it's always defined,
then in the try continue to assign labels and create TokenizerMetricsCollector
(using server_args, labels, bucket_time_to_first_token,
bucket_e2e_request_latency, and optional bucket_inter_token_latency) and set
self.metrics_collector and self._metrics_labels on success; replace the broad
except Exception: with narrower handling—catch ImportError (log a warning with
exc_info) and catch TypeError (or other specific runtime/API-mismatch errors)
separately so you log the exception type and message via logger.warning,
ensuring programming errors aren't silently swallowed while preserving graceful
degradation when observability imports or APIs are missing.


# Crash dump for debugging
self.crash_dump_request_list = []
Expand Down Expand Up @@ -576,12 +604,34 @@ async def _handle_batch_output(self, batch_out: BatchTokenIDOutput):
logger.debug(f"Skipping output for aborted request {rid}")
continue

# Update metrics
# Update timing
if state.time_stats.first_token_time == 0.0:
state.time_stats.set_first_token_time()
else:
state.time_stats.set_last_time()

# Observe request-level Prometheus metrics
if self.metrics_collector is not None:
completion_tokens = (
batch_out.completion_tokens[i] if batch_out.completion_tokens else 0
)
labels = dict(self._metrics_labels)
if not state.ttft_observed:
state.ttft_observed = True
state.last_completion_tokens = completion_tokens
self.metrics_collector.observe_time_to_first_token(
labels, state.time_stats.get_first_token_latency()
)
else:
num_new_tokens = completion_tokens - state.last_completion_tokens
if num_new_tokens > 0:
self.metrics_collector.observe_inter_token_latency(
labels,
state.time_stats.get_interval(),
num_new_tokens,
)
state.last_completion_tokens = completion_tokens

# Extract output for this request
output_data = {
"request_id": rid,
Expand Down Expand Up @@ -671,6 +721,20 @@ def get_part(attr_name):
state.stream_finished = True
state.event.set()

if self.metrics_collector is not None:
prompt_tokens = output_data["meta_info"].get("prompt_tokens", 0)
compl_tokens = output_data["meta_info"].get("completion_tokens", 0)
cached_tokens = output_data["meta_info"].get("cached_tokens", 0)
self.metrics_collector.observe_one_finished_request(
dict(self._metrics_labels),
prompt_tokens,
compl_tokens,
cached_tokens,
state.time_stats.get_e2e_latency(),
False,
0,
)

# Remove from tracking after a delay
async def cleanup(request_id):
await asyncio.sleep(5.0)
Expand Down
Loading