diff --git a/grpc_servicer/smg_grpc_servicer/sglang/request_manager.py b/grpc_servicer/smg_grpc_servicer/sglang/request_manager.py index 8814f7f702..d7a449bc66 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/request_manager.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/request_manager.py @@ -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 @@ -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() @@ -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 except Exception as e: logger.error(f"Handle loop error: {e}\n{get_exception_traceback()}") @@ -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") @@ -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, diff --git a/grpc_servicer/smg_grpc_servicer/sglang/server.py b/grpc_servicer/smg_grpc_servicer/sglang/server.py index 99230711bb..c9a3da6ccc 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/server.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/server.py @@ -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 @@ -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() stop_event.set() for sig in (signal.SIGTERM, signal.SIGINT): @@ -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: + 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 + ] + 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, + ) + await asyncio.sleep(5) + logger.info("Shutting down gRPC server") # Shutdown request manager first - this closes ZMQ sockets and stops background tasks diff --git a/grpc_servicer/smg_grpc_servicer/sglang/servicer.py b/grpc_servicer/smg_grpc_servicer/sglang/servicer.py index aa2ee916e4..647a2be155 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/servicer.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/servicer.py @@ -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", + ) + try: # Convert gRPC request to internal format tokenized_req = self._convert_generate_request(request) @@ -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) @@ -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")