diff --git a/bindings/python/README.md b/bindings/python/README.md index cb4f3da5ee..5506566052 100644 --- a/bindings/python/README.md +++ b/bindings/python/README.md @@ -2,6 +2,52 @@ This directory contains the Python bindings for SMG (Shepherd Model Gateway), built using [maturin](https://github.com/PyO3/maturin) and [PyO3](https://github.com/PyO3/pyo3). +## Quick Start + +### Installation + +```bash +pip install maturin +cd smg/bindings/python +maturin develop --features vendored-openssl +``` + +### Usage + +The `smg serve` command launches backend workers and the SMG router in a single command: + +```bash +# sglang with gRPC (default) +smg serve --backend sglang --model-path /path/to/model --port 8080 + +# sglang with HTTP +smg serve --backend sglang --model-path /path/to/model --port 8080 --connection-mode http + +# vLLM (gRPC only) +smg serve --backend vllm --model /path/to/model --port 8080 + +# TensorRT-LLM (gRPC only) +smg serve --backend trtllm --model /path/to/model --port 8080 + +# Multiple workers (data parallel) +smg serve --backend sglang --model-path /path/to/model --port 8080 --dp-size 4 +``` + +### Serve Options + +| Option | Default | Description | +|--------|---------|-------------| +| `--backend` | `sglang` | Backend to use: `sglang`, `vllm`, or `trtllm` | +| `--connection-mode` | `grpc` | Connection mode: `grpc` or `http`. vllm/trtllm only support grpc | +| `--host` | `127.0.0.1` | Host for the router | +| `--port` | `8080` | Port for the router | +| `--dp-size` | `1` | Data parallel size (number of worker replicas) | +| `--worker-host` | `127.0.0.1` | Host for worker processes | +| `--worker-base-port` | `31000` | Base port for workers | +| `--worker-startup-timeout` | `300` | Seconds to wait for workers to become healthy | + +Backend-specific options (e.g., `--tensor-parallel-size`, `--quantization`) are passed through to the backend. + ## Directory Structure ``` @@ -10,23 +56,15 @@ bindings/python/ │ ├── lib.rs # Rust/PyO3 bindings implementation │ └── smg/ # Python source code │ ├── __init__.py -│ ├── version.py +│ ├── cli.py # CLI entry point +│ ├── serve.py # smg serve implementation │ ├── launch_server.py │ ├── launch_router.py │ ├── router.py -│ ├── router_args.py -│ └── mini_lb.py +│ └── router_args.py ├── tests/ # Python unit tests -│ ├── conftest.py -│ ├── test_validation.py -│ ├── test_arg_parser.py -│ ├── test_router_config.py -│ └── test_startup_sequence.py -├── Cargo.toml # Rust package configuration for bindings +├── Cargo.toml # Rust package configuration ├── pyproject.toml # Python package configuration -├── setup.py # Setup configuration -├── MANIFEST.in # Package manifest -├── .coveragerc # Test coverage configuration └── README.md # This file ``` @@ -35,10 +73,7 @@ bindings/python/ ### Development Build ```bash -# Install maturin pip install maturin - -# Build and install in development mode cd smg/bindings/python maturin develop --features vendored-openssl ``` @@ -46,32 +81,14 @@ maturin develop --features vendored-openssl ### Production Build ```bash -# Build wheel cd smg/bindings/python maturin build --release --out dist --features vendored-openssl - -# Install the built wheel pip install dist/smg-*.whl ``` ## Testing ```bash -# Run Python unit tests (after maturin develop) cd smg/bindings/python pytest tests/ ``` - -## Configuration - -- **pyproject.toml**: Defines package metadata, dependencies, and build configuration -- **python-source**: Set to `"src"` indicating Python source uses the src layout -- **module-name**: `smg.smg_rs` - the Rust extension module name - -## Notes - -- The Rust bindings source code is located in `src/lib.rs` -- The bindings have their own `Cargo.toml` in this directory -- The main SMG library is located in `../../model_gateway/` and is used as a dependency -- The package includes both Python code and Rust extensions built with PyO3 -- PyO3 types are prefixed with `Py` in Rust but exposed to Python without the prefix using the `name` attribute diff --git a/bindings/python/pyproject.toml b/bindings/python/pyproject.toml index 2c1e22d220..7c39e69143 100644 --- a/bindings/python/pyproject.toml +++ b/bindings/python/pyproject.toml @@ -11,14 +11,13 @@ authors = [ {name = "Chang Su", email = "mckvtl@gmail.com"}, {name = "Keyang Ru", email = "rukeyang@gmail.com"} ] -requires-python = ">=3.8" +requires-python = ">=3.9" readme = "../../README.md" license = { text = "Apache-2.0" } classifiers = [ "Programming Language :: Python :: Implementation :: CPython", "Programming Language :: Rust", "Programming Language :: Python :: 3", - "Programming Language :: Python :: 3.8", "Programming Language :: Python :: 3.9", "Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.11", @@ -29,10 +28,8 @@ classifiers = [ dependencies = [ "setproctitle", - "aiohttp", - "orjson", - "uvicorn", - "fastapi", + "grpcio", + "grpcio-health-checking", ] [project.optional-dependencies] diff --git a/bindings/python/src/smg/serve.py b/bindings/python/src/smg/serve.py index c03bdeb104..8d0de0bdc5 100644 --- a/bindings/python/src/smg/serve.py +++ b/bindings/python/src/smg/serve.py @@ -38,15 +38,19 @@ def build_command( """Build the CLI command list to launch a worker.""" ... - @abstractmethod - def health_check(self, host: str, port: int, timeout: float) -> bool: + def health_check( + self, args: argparse.Namespace, host: str, port: int, timeout: float + ) -> bool: """Return True when the worker at host:port is healthy.""" - ... + if getattr(args, "connection_mode", "grpc") == "grpc": + return _grpc_health_check(host, port, timeout) + return _http_health_check(f"http://{host}:{port}/health", timeout) - @abstractmethod - def worker_url(self, host: str, port: int) -> str: + def worker_url(self, args: argparse.Namespace, host: str, port: int) -> str: """Return the URL used by the router to reach this worker.""" - ... + if getattr(args, "connection_mode", "grpc") == "grpc": + return f"grpc://{host}:{port}" + return f"http://{host}:{port}" def _get_tp_size(self, args: argparse.Namespace) -> int: """Return tensor-parallel size for GPU assignment. Default 1.""" @@ -97,19 +101,14 @@ def build_command( "--port", str(port), ] - if getattr(args, "grpc_mode", False): + if getattr(args, "connection_mode", "grpc") == "grpc": cmd.append("--grpc-mode") return cmd - def health_check(self, host: str, port: int, timeout: float) -> bool: - return _http_health_check(f"http://{host}:{port}/health", timeout) - - def worker_url(self, host: str, port: int) -> str: - return f"http://{host}:{port}" class VllmWorkerLauncher(WorkerLauncher): - """Launcher for vLLM inference workers (gRPC mode).""" + """Launcher for vLLM inference workers (gRPC mode only).""" def _get_tp_size(self, args: argparse.Namespace) -> int: return getattr(args, "tensor_parallel_size", 1) @@ -117,6 +116,8 @@ def _get_tp_size(self, args: argparse.Namespace) -> int: def build_command( self, args: argparse.Namespace, host: str, port: int ) -> List[str]: + if getattr(args, "connection_mode", "grpc") != "grpc": + raise ValueError("vLLM backend only supports grpc connection mode") return [ sys.executable, "-m", @@ -129,15 +130,10 @@ def build_command( str(port), ] - def health_check(self, host: str, port: int, timeout: float) -> bool: - return _grpc_health_check(host, port, timeout) - - def worker_url(self, host: str, port: int) -> str: - return f"grpc://{host}:{port}" class TrtllmWorkerLauncher(WorkerLauncher): - """Launcher for TensorRT-LLM inference workers (gRPC mode). + """Launcher for TensorRT-LLM inference workers (gRPC mode only). Uses ``python3 -m tensorrt_llm.commands.serve --grpc ...``. See https://github.com/NVIDIA/TensorRT-LLM/pull/11037 @@ -149,6 +145,8 @@ def _get_tp_size(self, args: argparse.Namespace) -> int: def build_command( self, args: argparse.Namespace, host: str, port: int ) -> List[str]: + if getattr(args, "connection_mode", "grpc") != "grpc": + raise ValueError("TensorRT-LLM backend only supports grpc connection mode") return [ sys.executable, "-m", @@ -161,11 +159,6 @@ def build_command( str(port), ] - def health_check(self, host: str, port: int, timeout: float) -> bool: - return _grpc_health_check(host, port, timeout) - - def worker_url(self, host: str, port: int) -> str: - return f"grpc://{host}:{port}" BACKEND_LAUNCHERS = { @@ -323,9 +316,30 @@ def add_serve_args(parser: argparse.ArgumentParser) -> None: help=f"Inference backend to use (default: {DEFAULT_BACKEND})", ) group.add_argument( + "--connection-mode", + default="grpc", + choices=["grpc", "http"], + help="Connection mode for workers (default: grpc). Note: vllm and trtllm only support grpc", + ) + # Router host/port - may be overridden by backend (e.g. sglang) + group.add_argument( + "--host", + default="127.0.0.1", + help="Host for the router (default: 127.0.0.1)", + ) + group.add_argument( + "--port", + type=int, + default=8080, + help="Port for the router (default: 8080)", + ) + # Data parallel size - may be overridden by backend + group.add_argument( + "--data-parallel-size", "--dp-size", type=int, default=1, + dest="data_parallel_size", help="Data parallel size (number of worker replicas)", ) group.add_argument( @@ -375,8 +389,10 @@ def parse_serve_args( backend = pre_args.backend # Pass 2: build full parser with backend-specific args + # Use conflict_handler='resolve' to let backend args override serve defaults parser = argparse.ArgumentParser( - description=f"Launch {backend} worker(s) + gateway router" + description=f"Launch {backend} worker(s) + gateway router", + conflict_handler="resolve", ) add_serve_args(parser) _import_backend_args(backend, parser) @@ -421,7 +437,7 @@ def run(self) -> None: # -- internal ----------------------------------------------------------- def _launch_workers(self) -> None: - ports = _find_available_ports(self.args.worker_base_port, self.args.dp_size) + ports = _find_available_ports(self.args.worker_base_port, self.args.data_parallel_size) host = self.args.worker_host for dp_rank, port in enumerate(ports): env = self.launcher.gpu_env(self.args, dp_rank) @@ -441,7 +457,7 @@ def _wait_healthy(self) -> None: raise RuntimeError( f"Worker on port {port} exited with code {proc.returncode}" ) - if self.launcher.health_check(host, port, timeout=5.0): + if self.launcher.health_check(self.args, host, port, timeout=5.0): logger.info("Worker on %s:%d is healthy", host, port) break time.sleep(2) @@ -453,7 +469,7 @@ def _wait_healthy(self) -> None: def _build_router_args(self) -> RouterArgs: worker_urls = [ - self.launcher.worker_url(self.args.worker_host, port) + self.launcher.worker_url(self.args, self.args.worker_host, port) for _, port in self.workers ] router_args = RouterArgs.from_cli_args(self.args, use_router_prefix=True) diff --git a/bindings/python/tests/test_serve.py b/bindings/python/tests/test_serve.py index 0b85690e9e..9ffa693aea 100644 --- a/bindings/python/tests/test_serve.py +++ b/bindings/python/tests/test_serve.py @@ -84,17 +84,29 @@ def test_backend_rejects_invalid_choice(self): with pytest.raises(SystemExit): parser.parse_args(["--backend", "nonexistent"]) - def test_adds_dp_size(self): + def test_adds_data_parallel_size(self): parser = argparse.ArgumentParser() add_serve_args(parser) args = parser.parse_args(["--dp-size", "4"]) - assert args.dp_size == 4 + assert args.data_parallel_size == 4 - def test_dp_size_default(self): + def test_data_parallel_size_default(self): parser = argparse.ArgumentParser() add_serve_args(parser) args = parser.parse_args([]) - assert args.dp_size == 1 + assert args.data_parallel_size == 1 + + def test_adds_connection_mode(self): + parser = argparse.ArgumentParser() + add_serve_args(parser) + args = parser.parse_args(["--connection-mode", "http"]) + assert args.connection_mode == "http" + + def test_connection_mode_default_is_grpc(self): + parser = argparse.ArgumentParser() + add_serve_args(parser) + args = parser.parse_args([]) + assert args.connection_mode == "grpc" def test_adds_worker_host(self): parser = argparse.ArgumentParser() @@ -132,6 +144,12 @@ def test_worker_startup_timeout_default(self): args = parser.parse_args([]) assert args.worker_startup_timeout == 300 + def test_host_default_is_localhost(self): + parser = argparse.ArgumentParser() + add_serve_args(parser) + args = parser.parse_args([]) + assert args.host == "127.0.0.1" + class TestImportBackendArgs: """Test _import_backend_args for each backend.""" @@ -185,10 +203,11 @@ def test_trtllm_basic(self): def test_trtllm_defaults(self): backend, args = parse_serve_args(["--backend", "trtllm"]) assert backend == "trtllm" - assert args.dp_size == 1 + assert args.data_parallel_size == 1 assert args.worker_host == "127.0.0.1" assert args.worker_base_port == 31000 assert args.worker_startup_timeout == 300 + assert args.connection_mode == "grpc" def test_trtllm_with_serve_args(self): backend, args = parse_serve_args([ @@ -199,7 +218,7 @@ def test_trtllm_with_serve_args(self): "--worker-startup-timeout", "600", ]) assert backend == "trtllm" - assert args.dp_size == 8 + assert args.data_parallel_size == 8 assert args.worker_host == "0.0.0.0" assert args.worker_base_port == 35000 assert args.worker_startup_timeout == 600 @@ -344,9 +363,10 @@ def test_gpu_env_none_copies_os_environ(self): class TestSglangWorkerLauncher: """Test SglangWorkerLauncher.build_command().""" - def test_build_command_basic(self): + def test_build_command_grpc_mode_default(self): + """Default connection_mode is grpc, so --grpc-mode should be present.""" launcher = SglangWorkerLauncher() - args = argparse.Namespace(model_path="/tmp/model", grpc_mode=False) + args = argparse.Namespace(model_path="/tmp/model", connection_mode="grpc") cmd = launcher.build_command(args, "127.0.0.1", 31000) assert "--model-path" in cmd assert "/tmp/model" in cmd @@ -354,22 +374,38 @@ def test_build_command_basic(self): assert "127.0.0.1" in cmd assert "--port" in cmd assert "31000" in cmd - assert "--grpc-mode" not in cmd + assert "--grpc-mode" in cmd - def test_build_command_grpc_mode(self): + def test_build_command_http_mode(self): + """When connection_mode is http, --grpc-mode should not be present.""" launcher = SglangWorkerLauncher() - args = argparse.Namespace(model_path="/tmp/model", grpc_mode=True) + args = argparse.Namespace(model_path="/tmp/model", connection_mode="http") cmd = launcher.build_command(args, "127.0.0.1", 31000) - assert "--grpc-mode" in cmd + assert "--grpc-mode" not in cmd - def test_worker_url(self): + def test_worker_url_grpc_mode(self): + launcher = SglangWorkerLauncher() + args = argparse.Namespace(connection_mode="grpc") + assert launcher.worker_url(args, "127.0.0.1", 31000) == "grpc://127.0.0.1:31000" + + def test_worker_url_http_mode(self): launcher = SglangWorkerLauncher() - assert launcher.worker_url("127.0.0.1", 31000) == "http://127.0.0.1:31000" + args = argparse.Namespace(connection_mode="http") + assert launcher.worker_url(args, "127.0.0.1", 31000) == "http://127.0.0.1:31000" - def test_health_check_delegates_to_http(self): + def test_health_check_grpc_mode(self): launcher = SglangWorkerLauncher() + args = argparse.Namespace(connection_mode="grpc") + with patch("smg.serve._grpc_health_check", return_value=True) as mock: + result = launcher.health_check(args, "127.0.0.1", 31000, 5.0) + assert result is True + mock.assert_called_once_with("127.0.0.1", 31000, 5.0) + + def test_health_check_http_mode(self): + launcher = SglangWorkerLauncher() + args = argparse.Namespace(connection_mode="http") with patch("smg.serve._http_health_check", return_value=True) as mock: - result = launcher.health_check("127.0.0.1", 31000, 5.0) + result = launcher.health_check(args, "127.0.0.1", 31000, 5.0) assert result is True mock.assert_called_once_with("http://127.0.0.1:31000/health", 5.0) @@ -379,7 +415,7 @@ class TestVllmWorkerLauncher: def test_build_command(self): launcher = VllmWorkerLauncher() - args = argparse.Namespace(model="/tmp/model") + args = argparse.Namespace(model="/tmp/model", connection_mode="grpc") cmd = launcher.build_command(args, "0.0.0.0", 32000) assert "vllm.entrypoints.grpc_server" in cmd assert "--model" in cmd @@ -389,14 +425,22 @@ def test_build_command(self): assert "--port" in cmd assert "32000" in cmd + def test_build_command_rejects_http_mode(self): + launcher = VllmWorkerLauncher() + args = argparse.Namespace(model="/tmp/model", connection_mode="http") + with pytest.raises(ValueError, match="vLLM backend only supports grpc"): + launcher.build_command(args, "0.0.0.0", 32000) + def test_worker_url(self): launcher = VllmWorkerLauncher() - assert launcher.worker_url("127.0.0.1", 32000) == "grpc://127.0.0.1:32000" + args = argparse.Namespace(connection_mode="grpc") + assert launcher.worker_url(args, "127.0.0.1", 32000) == "grpc://127.0.0.1:32000" def test_health_check_delegates_to_grpc(self): launcher = VllmWorkerLauncher() + args = argparse.Namespace(connection_mode="grpc") with patch("smg.serve._grpc_health_check", return_value=True) as mock: - result = launcher.health_check("127.0.0.1", 32000, 5.0) + result = launcher.health_check(args, "127.0.0.1", 32000, 5.0) assert result is True mock.assert_called_once_with("127.0.0.1", 32000, 5.0) @@ -406,7 +450,7 @@ class TestTrtllmWorkerLauncher: def test_build_command(self): launcher = TrtllmWorkerLauncher() - args = argparse.Namespace(model="/tmp/model") + args = argparse.Namespace(model="/tmp/model", connection_mode="grpc") cmd = launcher.build_command(args, "0.0.0.0", 50051) assert "tensorrt_llm.commands.serve" in cmd assert "/tmp/model" in cmd @@ -416,14 +460,22 @@ def test_build_command(self): assert "--port" in cmd assert "50051" in cmd + def test_build_command_rejects_http_mode(self): + launcher = TrtllmWorkerLauncher() + args = argparse.Namespace(model="/tmp/model", connection_mode="http") + with pytest.raises(ValueError, match="TensorRT-LLM backend only supports grpc"): + launcher.build_command(args, "0.0.0.0", 50051) + def test_worker_url(self): launcher = TrtllmWorkerLauncher() - assert launcher.worker_url("127.0.0.1", 50051) == "grpc://127.0.0.1:50051" + args = argparse.Namespace(connection_mode="grpc") + assert launcher.worker_url(args, "127.0.0.1", 50051) == "grpc://127.0.0.1:50051" def test_health_check_delegates_to_grpc(self): launcher = TrtllmWorkerLauncher() + args = argparse.Namespace(connection_mode="grpc") with patch("smg.serve._grpc_health_check", return_value=True) as mock: - result = launcher.health_check("127.0.0.1", 50051, 5.0) + result = launcher.health_check(args, "127.0.0.1", 50051, 5.0) assert result is True mock.assert_called_once_with("127.0.0.1", 50051, 5.0) @@ -528,7 +580,8 @@ def _make_args(**overrides): """Create a minimal argparse.Namespace for orchestrator tests.""" defaults = { "backend": "sglang", - "dp_size": 2, + "data_parallel_size": 2, + "connection_mode": "grpc", "worker_host": "127.0.0.1", "worker_base_port": 31000, "worker_startup_timeout": 10, @@ -545,8 +598,8 @@ def _make_args(**overrides): class TestServeOrchestrator: """Test ServeOrchestrator lifecycle methods.""" - def test_build_router_args_injects_worker_urls(self): - args = _make_args(dp_size=2) + def test_build_router_args_injects_worker_urls_grpc(self): + args = _make_args(data_parallel_size=2, connection_mode="grpc") orch = ServeOrchestrator("sglang", args) # Simulate workers already launched mock_proc1 = MagicMock() @@ -559,14 +612,30 @@ def test_build_router_args_injects_worker_urls(self): result = orch._build_router_args() mock_from_cli.assert_called_once_with(args, use_router_prefix=True) + assert mock_router_args.worker_urls == [ + "grpc://127.0.0.1:31000", + "grpc://127.0.0.1:31003", + ] + assert result is mock_router_args + + def test_build_router_args_http_mode(self): + args = _make_args(data_parallel_size=2, connection_mode="http") + orch = ServeOrchestrator("sglang", args) + mock_proc = MagicMock() + orch.workers = [(mock_proc, 31000), (mock_proc, 31003)] + + with patch("smg.serve.RouterArgs.from_cli_args") as mock_from_cli: + mock_router_args = MagicMock() + mock_from_cli.return_value = mock_router_args + orch._build_router_args() + assert mock_router_args.worker_urls == [ "http://127.0.0.1:31000", "http://127.0.0.1:31003", ] - assert result is mock_router_args def test_build_router_args_vllm_grpc_urls(self): - args = _make_args(backend="vllm", dp_size=2, model="/tmp/m") + args = _make_args(backend="vllm", data_parallel_size=2, model="/tmp/m", connection_mode="grpc") orch = ServeOrchestrator("vllm", args) mock_proc = MagicMock() orch.workers = [(mock_proc, 32000), (mock_proc, 32003)] @@ -638,7 +707,7 @@ def test_signal_handler_guard_prevents_reentry(self): orch._signal_handler(signal.SIGINT, None) def test_trtllm_orchestrator_launches_grpc_workers(self): - args = _make_args(backend="trtllm", dp_size=1, model="/tmp/m") + args = _make_args(backend="trtllm", data_parallel_size=1, model="/tmp/m", connection_mode="grpc") orch = ServeOrchestrator("trtllm", args) with patch("smg.serve._find_available_ports", return_value=[50051]): @@ -651,7 +720,7 @@ def test_trtllm_orchestrator_launches_grpc_workers(self): def test_launch_workers_passes_gpu_env(self): """_launch_workers passes CUDA_VISIBLE_DEVICES via gpu_env for each dp_rank.""" - args = _make_args(dp_size=2, tp_size=2) + args = _make_args(data_parallel_size=2, tp_size=2, connection_mode="grpc") orch = ServeOrchestrator("sglang", args) launched_envs = []