Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 30 additions & 17 deletions grpc_servicer/smg_grpc_servicer/sglang/request_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,6 +195,10 @@ def __init__(
# State Management (from TokenizerManager)
self.rid_to_state: dict[str, GrpcReqState] = {}
self.asyncio_tasks: set = set()
# Separate handle_loop ref so the drain loop can detect if scheduler
# output forwarding has died — without it, rid_to_state would never
# drain and the pod would wait for SIGKILL.
self.handle_loop_task: asyncio.Task | None = None
self.gracefully_exit = False
self.no_create_loop = False
self.event_loop = None
Expand Down Expand Up @@ -479,8 +483,13 @@ async def handle_loop(self):
"""
Main event loop - processes outputs from scheduler.
Mimics TokenizerManager's handle_loop.

Runs until the task is cancelled (by shutdown()) or an unrecoverable
ZMQ error. It must not gate on ``gracefully_exit`` — during drain
the health flag is set but scheduler outputs must still be forwarded
so in-flight streaming requests can finish their token sequences.
"""
while not self.gracefully_exit:
while True:
try:
# Receive from scheduler
recv_obj = await self.recv_from_scheduler.recv_pyobj()
Expand Down Expand Up @@ -508,16 +517,14 @@ async def handle_loop(self):
logger.warning(f"Unknown output type: {type(recv_obj)}")

except zmq.error.Again:
# Timeout, check if we should exit
if self.gracefully_exit:
break
# Timeout on non-blocking recv; keep polling.
continue
except zmq.error.ZMQError as e:
# Socket closed or other ZMQ error - exit cleanly if shutting down
# Socket closed or unrecoverable ZMQ error.
if self.gracefully_exit:
logger.debug(f"ZMQ recv interrupted during shutdown: {e}")
break
logger.error(f"ZMQ error in handle loop: {e}\n{get_exception_traceback()}")
else:
logger.error(f"ZMQ error in handle loop: {e}\n{get_exception_traceback()}")
break
Comment on lines 524 to 528

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge Avoid treating drain mode as terminal ZMQ shutdown

begin_drain() now sets gracefully_exit before real shutdown, but this branch still treats any ZMQError under that flag as shutdown and exits handle_loop. During a long drain window, a transient socket error would stop forwarding scheduler outputs to request queues, which can leave in-flight streams stuck and prevent the drain loop from ever completing cleanly. Drain and terminal-shutdown states should be separated for this error path.

Useful? React with 👍 / 👎.

except Exception as e:
logger.error(f"Handle loop error: {e}\n{get_exception_traceback()}")
Expand Down Expand Up @@ -817,6 +824,16 @@ def record_request_for_crash_dump(self, obj):
}
)

def begin_drain(self) -> None:
"""Mark the manager as draining.

Idempotent and non-destructive: flips gracefully_exit so the health
servicer reports NOT_SERVING and the server-side drain loop can
begin, but does not cancel tasks or enqueue shutdown errors on
in-flight requests. Safe to call from a synchronous signal handler.
"""
self.gracefully_exit = True

async def shutdown(self):
"""Gracefully shutdown the request manager."""
logger.info("Shutting down GrpcRequestManager")
Expand Down Expand Up @@ -903,25 +920,21 @@ def auto_create_handle_loop(self):

self.no_create_loop = True
loop = get_or_create_event_loop()
self.asyncio_tasks.add(loop.create_task(print_exception_wrapper(self.handle_loop)))
self.handle_loop_task = loop.create_task(print_exception_wrapper(self.handle_loop))
self.asyncio_tasks.add(self.handle_loop_task)

self.event_loop = loop

# We only add signal handler when the tokenizer manager is in the main thread
# due to the CPython limitation.
# due to the CPython limitation. The SIGTERM handler here is a startup-window
# fallback; once serve_grpc runs, it overrides SIGTERM with the drain-aware
# signal_handler in server.py. SIGQUIT stays owned here to forward scheduler
# crashes to a process-tree kill.
if threading.current_thread() is threading.main_thread():
signal_handler = GrpcSignalHandler(self)
loop.add_signal_handler(signal.SIGTERM, signal_handler.sigterm_handler)
# Update the signal handler for the process. It overrides the sigquit handler in the launch phase.
loop.add_signal_handler(signal.SIGQUIT, signal_handler.running_phase_sigquit_handler)

self.asyncio_tasks.add(loop.create_task(print_exception_wrapper(self.sigterm_watchdog)))

async def sigterm_watchdog(self):
"""Watchdog to handle SIGTERM gracefully, matching TokenizerManager pattern."""
while not self.gracefully_exit:
await asyncio.sleep(1.0)

def _req_stats_init(
self,
obj: TokenizedGenerateReqInput | TokenizedEmbeddingReqInput,
Expand Down
53 changes: 52 additions & 1 deletion grpc_servicer/smg_grpc_servicer/sglang/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode
from sglang.srt.managers.disagg_service import start_disagg_service
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import kill_process_tree
from sglang.srt.utils import get_bool_env_var, kill_process_tree
from sglang.utils import get_exception_traceback
from smg_grpc_proto import sglang_scheduler_pb2, sglang_scheduler_pb2_grpc

Expand Down Expand Up @@ -276,6 +276,10 @@ def _cert_config_fetcher():

def signal_handler():
logger.info("Received shutdown signal")
# Flip health to NOT_SERVING and mark the request manager as
# draining so K8s stops routing new traffic to this pod and
# in-flight requests are allowed to finish.
servicer.begin_drain()

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge Prevent warmup from restoring SERVING after drain begins

Calling begin_drain() here marks health as NOT_SERVING, but the warmup thread later calls health_servicer.set_serving() unconditionally in _wait_and_warmup_grpc. If SIGTERM arrives during warmup and drain lasts for in-flight requests, readiness can flip back to SERVING during drain, so Kubernetes may resume routing traffic to a pod that is intentionally rejecting inference RPCs with UNAVAILABLE. Drain mode should make health status sticky (or warmup should skip set_serving once gracefully_exit is true).

Useful? React with 👍 / 👎.

stop_event.set()

for sig in (signal.SIGTERM, signal.SIGINT):
Expand All @@ -284,6 +288,53 @@ def signal_handler():
try:
await stop_event.wait()
finally:
# Drain phase: wait for in-flight requests to finish naturally.
# Mirrors tokenizer_manager.sigterm_watchdog. No in-process
# timeout; K8s terminationGracePeriodSeconds SIGKILL is the backstop.
# rid_to_state is only mutated from the event-loop thread, so a
# list() snapshot is safe without additional guarding.
while True:

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge Stop accepting new RPCs before waiting for drain completion

The new drain loop runs before the server is stopped, so the process still accepts Generate/Embed RPCs while waiting for rid_to_state to reach zero. In environments with persistent gRPC clients (or direct pod traffic), new requests can continue to enter during drain and keep remain_num_req non-zero, causing termination to stall until external SIGKILL and defeating the graceful-shutdown goal. Move listener shutdown/rejection earlier (or explicitly reject inference RPCs once draining starts) so the drain set is bounded.

Useful? React with 👍 / 👎.

if get_bool_env_var("SGL_FORCE_SHUTDOWN"):
logger.warning("SGL_FORCE_SHUTDOWN set; skipping drain")
break

# If handle_loop has died (e.g. fatal ZMQError, scheduler
# crash), scheduler outputs will never be forwarded and
# in-flight rids can never be marked finished — abort the
# drain instead of blocking until SIGKILL.
handle_loop_task = servicer.request_manager.handle_loop_task
if handle_loop_task is not None and handle_loop_task.done():
logger.warning(
"handle_loop task has terminated; aborting drain and proceeding to shutdown"
)
break

# Finished requests linger in rid_to_state for 5 s via the
# cleanup() task in _handle_batch_output; exclude them so
# the drain loop exits promptly once real work is done.
remaining_rids = [
rid
for rid, state in servicer.request_manager.rid_to_state.items()
if not state.finished
]
Comment on lines +315 to +319

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge Include active RPC handlers in drain completion check

The drain loop decides shutdown readiness solely from request_manager.rid_to_state entries marked unfinished, but that map does not cover every still-active RPC handler. In particular, generate_request can be mid-flight with no active RID between its n>1 prefix phase and launching phase-2 generators, so this check can read zero and proceed to servicer.shutdown() while a client request is still being processed, causing that request to be aborted despite graceful-drain mode. Consider tracking handler-level in-flight RPCs (or a broader active-request counter) instead of relying only on RID state.

Useful? React with 👍 / 👎.

remain_num_req = len(remaining_rids)
if remain_num_req == 0:
logger.info("Drain complete; no in-flight requests")
break

# Truncate rid list in logs: under high load there can be
# thousands of in-flight requests and the full list would
# flood the log every poll interval.
log_rids = remaining_rids[:10]
if remain_num_req > 10:
log_rids = log_rids + ["..."]
logger.info(
"Gracefully exiting... Remaining number of requests %d. Remaining requests %s",
remain_num_req,
log_rids,
)
Comment on lines +331 to +335

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

medium

Logging the full list of remaining_rids can be extremely verbose in high-throughput scenarios where many requests are in-flight during a scale-down. This could lead to log flooding and performance degradation. It is better to truncate the list in the log message.

Suggested change
logger.info(
"Gracefully exiting... Remaining number of requests %d. Remaining requests %s",
remain_num_req,
remaining_rids,
)
logger.info(
"Gracefully exiting... Remaining number of requests %d. Remaining requests %s",
remain_num_req,
remaining_rids[:10] + (["..."] if remain_num_req > 10 else []),
)

await asyncio.sleep(5)
Comment on lines +296 to +336

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

high

There is a potential hang in this drain loop if the background handle_loop task in GrpcRequestManager terminates unexpectedly (for example, due to a ZMQError as seen in request_manager.py line 524). If handle_loop stops, scheduler outputs will no longer be processed, and in-flight requests will never be marked as finished. This would cause the drain loop to wait indefinitely until the pod is forcefully killed by the orchestrator (e.g., Kubernetes SIGKILL after terminationGracePeriodSeconds). Consider adding a mechanism to detect if the background processing task is still alive during the drain phase.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

medium

A 5-second sleep interval for the drain loop is relatively long. It can delay the final termination of the pod by several seconds even after all in-flight requests have finished. Reducing this to 1 second would make the shutdown process more responsive without significantly increasing CPU overhead.

Suggested change
await asyncio.sleep(5)
await asyncio.sleep(1)


logger.info("Shutting down gRPC server")

# Shutdown request manager first - this closes ZMQ sockets and stops background tasks
Expand Down
28 changes: 28 additions & 0 deletions grpc_servicer/smg_grpc_servicer/sglang/servicer.py
Original file line number Diff line number Diff line change
Expand Up @@ -212,6 +212,16 @@ async def Generate(
"""Handle generation requests with streaming responses."""
logger.info(f"Receive generation request: {request.request_id}")

# Reject new RPCs once drain has started. K8s removes the pod from
# Service Endpoints when health flips to NOT_SERVING, but persistent
# clients or direct pod traffic could otherwise keep feeding work
# into rid_to_state and stall the drain loop indefinitely.
if self.request_manager.gracefully_exit:
await context.abort(
grpc.StatusCode.UNAVAILABLE,
"Server is shutting down",
)
Comment on lines +219 to +223

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge Exempt internal warmup RPCs from drain-time rejection

The new gracefully_exit guard aborts every Generate/Embed call, including the server’s own warmup calls, so if SIGTERM arrives before warmup finishes the warmup thread can receive UNAVAILABLE and follow its existing fatal path in server.py (_execute_grpc_server_warmup catches the RPC error and calls kill_process_tree). In that startup/shutdown overlap, the process is force-killed instead of completing the intended drain, which can still drop in-flight direct-to-pod requests.

Useful? React with 👍 / 👎.


try:
# Convert gRPC request to internal format
tokenized_req = self._convert_generate_request(request)
Expand Down Expand Up @@ -271,6 +281,13 @@ async def Embed(
"""Handle embedding requests."""
logger.info(f"Receive embedding request: {request.request_id}")

# Reject new RPCs once drain has started (same rationale as Generate).
if self.request_manager.gracefully_exit:
await context.abort(
grpc.StatusCode.UNAVAILABLE,
"Server is shutting down",
)

try:
tokenized_req = self._convert_embed_request(request)

Expand Down Expand Up @@ -1098,6 +1115,17 @@ def _create_completion_response(
),
)

def begin_drain(self) -> None:
"""Mark the service as draining so health flips to NOT_SERVING.

Non-destructive: does not cancel in-flight requests. Safe to call
from a synchronous signal handler. Must be followed by shutdown()
(after the drain loop completes) for final cleanup.
"""
if self.health_servicer:
self.health_servicer.set_not_serving()
self.request_manager.begin_drain()

async def shutdown(self):
"""Shutdown the service."""
logger.info("Shutting down gRPC service")
Expand Down
Loading