From d79162f220669c6a4b6c1fa673abec818a01e60e Mon Sep 17 00:00:00 2001 From: Ricky-7-Yan <2314530442@qq.com> Date: Mon, 31 Aug 2026 02:31:23 +0800 Subject: [PATCH] fix: stop uvicorn threads during shutdown --- src/litserve/server.py | 25 +++++++++++++++++++++++-- tests/unit/test_lit_server.py | 13 ++++++++++++- 2 files changed, 35 insertions(+), 3 deletions(-) diff --git a/src/litserve/server.py b/src/litserve/server.py index 9a6c253d..797c7947 100644 --- a/src/litserve/server.py +++ b/src/litserve/server.py @@ -421,6 +421,24 @@ def run(self, worker_id: int, sockets: Union[list[socket.socket], None] = None) return super().run(sockets) +class _UvicornThread(threading.Thread): + def __init__( + self, + server: _Server, + worker_id: int, + sockets: Union[list[socket.socket], None], + name: str, + ) -> None: + super().__init__(target=server.run, args=(worker_id, sockets), name=name) + self.server = server + + def terminate(self) -> None: + self.server.should_exit = True + + def kill(self) -> None: + self.server.force_exit = True + + class LitServer: """Initialize a LitServer for high-performance AI model serving. @@ -1575,8 +1593,11 @@ def _start_server(self, port, num_uvicorn_servers, log_level, sockets, uvicorn_w target=server.run, args=(response_queue_id, sockets), name=f"LitServer-{response_queue_id}" ) elif uvicorn_worker_type == "thread": - w = threading.Thread( - target=server.run, args=(response_queue_id, sockets), name=f"LitServer-{response_queue_id}" + w = _UvicornThread( + server=server, + worker_id=response_queue_id, + sockets=sockets, + name=f"LitServer-{response_queue_id}", ) else: raise ValueError("Invalid value for api_server_worker_type. Must be 'process' or 'thread'") diff --git a/tests/unit/test_lit_server.py b/tests/unit/test_lit_server.py index 9a085bd7..4f115b9b 100644 --- a/tests/unit/test_lit_server.py +++ b/tests/unit/test_lit_server.py @@ -31,7 +31,7 @@ import litserve as ls from litserve import LitAPI from litserve.connector import _Connector -from litserve.server import LitServer +from litserve.server import LitServer, _UvicornThread from litserve.utils import WorkerSetupStatus, wrap_litserve_start @@ -247,6 +247,17 @@ def test_start_server(mock_server): assert server.lit_api.spec.response_queue_id is not None, "response_queue_id must be generated" +def test_uvicorn_thread_shutdown_signals(): + uvicorn_server = MagicMock() + worker = _UvicornThread(uvicorn_server, worker_id=0, sockets=None, name="LitServer-0") + + worker.terminate() + assert uvicorn_server.should_exit is True + + worker.kill() + assert uvicorn_server.force_exit is True + + @pytest.fixture def server_for_api_worker_test(simple_litapi): server = ls.LitServer(simple_litapi, devices=1)