Skip to content
Merged
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
83 changes: 50 additions & 33 deletions bindings/python/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 |

Comment thread
coderabbitai[bot] marked this conversation as resolved.
Backend-specific options (e.g., `--tensor-parallel-size`, `--quantization`) are passed through to the backend.

## Directory Structure

```
Expand All @@ -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
```

Expand All @@ -35,43 +73,22 @@ 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
```

### 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/
```
Comment thread
slin1237 marked this conversation as resolved.

## 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
9 changes: 3 additions & 6 deletions bindings/python/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -29,10 +28,8 @@ classifiers = [

dependencies = [
"setproctitle",
"aiohttp",
"orjson",
"uvicorn",
"fastapi",
"grpcio",
"grpcio-health-checking",
]
Comment thread
coderabbitai[bot] marked this conversation as resolved.

[project.optional-dependencies]
Expand Down
72 changes: 44 additions & 28 deletions bindings/python/src/smg/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -97,26 +101,23 @@ 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)

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",
Expand All @@ -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 <model> --grpc ...``.
See https://github.com/NVIDIA/TensorRT-LLM/pull/11037
Expand All @@ -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",
Expand All @@ -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 = {
Expand Down Expand Up @@ -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)",
)
Comment thread
slin1237 marked this conversation as resolved.
# 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(
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand All @@ -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)
Expand Down
Loading
Loading