Skip to content
Open
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
25 changes: 23 additions & 2 deletions src/litserve/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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'")
Expand Down
13 changes: 12 additions & 1 deletion tests/unit/test_lit_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
Expand Down
Loading