init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,153 @@
# Dynamic Bucket Load Balancer
A dynamic bucketing-based hybrid load balance proxy for [vLLM](https://github.com/vllm-project/vllm).
The proxy fronts multiple vLLM backend servers and distributes
OpenAI-compatible requests across them. It can run in two modes:
- **Plain load balancing** (default): for each request, estimate a load score
and forward it to the least-loaded backend instance.
- **Dynamic bucket load balancing** (`--enable-dynamic-bucket`): split the
backend pool into a **short-request group** and a **long-request group**, route
requests to a group by their length, and dynamically rebalance across groups
based on the load gap and length affinity.
## Files
- `dynamic_bucket_load_balancer.py` — the core algorithm (pure standard library).
Buckets requests by length, then dynamically adjusts bucket assignment using
bucket load and length affinity.
- `hybrid_proxy_server.py` — the FastAPI proxy server that uses the algorithm to
route requests to the backend servers.
## How It Works
1. **Static bucketing by length.** Each request is first mapped to its *standard
bucket* by request length. With dynamic bucketing enabled the proxy uses two
buckets: short `[0, --server-group-threshold)` and long
`[--server-group-threshold, --max-request-tokens)`.
2. **Server groups.** The ordered backend list is split into the same number of
groups as buckets, **in order**: the first instances form the short group, the
last instances form the long group. With 4 backends and 2 buckets, backends 0
and 1 serve the short bucket, backends 2 and 3 serve the long bucket.
- > **Tip:** configure the first two instances for short sequences and the
> last two for long sequences (e.g. smaller `max-model-len` / KV cache for
> the short group, larger for the long group) to get the best throughput.
3. **Dynamic rebalancing.** For a new request, the balancer looks at neighbor
buckets with a lighter load and computes a redirect probability
`(load-gap probability) × (length-affinity factor)`. If it exceeds the
threshold (`0.12`), the request is redirected to the neighbor bucket. This
means a large load gap is suppressed when the request length is far from the
neighbor bucket, while a modest gap can still trigger a redirect when the
length is close to the boundary.
4. **Within a group**, the least-loaded server (smallest active token count) is
picked via a min-heap, the load is accumulated for the duration of the
request, and released when streaming completes.
## Prerequisites
- Python 3.10+
- Install dependencies:
```bash
pip install "fastapi<0.124.0" httpx uvicorn
```
## Step 1: Start Your Backend Servers
Start at least two vLLM servers, each as a separate process on its own port. The
proxy also works with a single backend, but load balancing is only meaningful
with two or more.
```bash
vllm serve --host 0.0.0.0 --port 8100 ... # vLLM Server 0
vllm serve --host 0.0.0.0 --port 8101 ... # vLLM Server 1
```
## Step 2: Start the Proxy Server
From `examples/dynamic_bucket_load_balancer/`, point the proxy at each backend
with `--server-hosts` / `--server-ports`:
```bash
python hybrid_proxy_server.py \
--host 0.0.0.0 --port 8000 \
--server-hosts 127.0.0.1 127.0.0.1 \
--server-ports 8100 8101
```
This starts the proxy on port 8000 and load balances across the two backends.
### Enable Dynamic Bucket Load Balancing
Add `--enable-dynamic-bucket` to split the pool into short/long groups. The
server count must be `>= 2` so each bucket has at least one instance. With 4
servers the first two form the short group and the last two the long group:
```bash
python hybrid_proxy_server.py \
--host 0.0.0.0 --port 8000 \
--server-hosts 127.0.0.1 127.0.0.1 127.0.0.1 127.0.0.1 \
--server-ports 8100 8101 8102 8103 \
--enable-dynamic-bucket \
--server-group-threshold 32768
```
## Step 3: Send a Request to the Proxy
Send OpenAI-compatible requests to the proxy. For example:
```bash
curl -X POST http://localhost:8000/v1/completions \
-H "Content-Type: application/json" \
-d '{
"model": "your-model",
"prompt": "The quick brown fox jumps over the lazy dog",
"max_tokens": 16
}'
```
Or for chat completions:
```bash
curl -X POST http://localhost:8000/v1/chat/completions \
-H "Content-Type: application/json" \
-d '{
"model": "your-model",
"messages": [{"role": "user", "content": "Hello!"}],
"max_tokens": 16
}'
```
## Step 4: Health Check
Check that the proxy is running and how many backends it fronts:
```bash
curl http://localhost:8000/healthcheck
```
Returns a JSON object, e.g.:
```json
{"status": "ok", "server_instances": 2}
```
## Configuration
| Argument | Default | Description |
| --- | --- | --- |
| `--host` | `localhost` | Proxy listen host. |
| `--port` | `8000` | Proxy listen port. |
| `--server-hosts` | `localhost` | Hosts of the backend vLLM servers (one per server, in order). |
| `--server-ports` | `8001` | Ports of the backend vLLM servers (one per server, in order). |
| `--enable-dynamic-bucket` | `False` | Enable dynamic bucket load balancing. |
| `--server-group-threshold` | `32768` | Length boundary between the short and long buckets. |
| `--max-request-tokens` | `131072` | Upper bound of the long bucket (max request length). |
| `--max-retries` | `3` | Max retries for a backend HTTP request. |
| `--retry-delay` | `0.001` | Base delay (seconds) for exponential backoff retries. |
The number of `--server-hosts` must equal the number of `--server-ports`.

View File

@@ -0,0 +1,243 @@
import math
from collections import namedtuple
from typing import AnyStr
ServerInfo = namedtuple("ServerInfo", ["instance_type", "instance_idx"])
class Task:
"""Request task (carries request length info)."""
def __init__(self, task_id, task_length, task_load):
self.id = task_id
self.length = task_length
self.bucket_idx = -1
self.load = task_load
self.server_info: ServerInfo = ServerInfo("Unknown", -1)
def __repr__(self):
return (
f"Task(id={self.id}, length={self.length}, load={self.load}, "
f"instance_type={self.server_info.instance_type}, instance_idx={self.server_info.instance_idx})"
)
class Bucket:
"""A bucket."""
def __init__(self, bucket_ranges: tuple[int, int]):
self.min_length = bucket_ranges[0]
self.max_length = bucket_ranges[1]
self.task_count = 0
self.total_load = 0.0
class DynamicBucketLoadBalancer:
"""
Statically buckets requests by length first, then dynamically adjusts the
assignment of new requests based on bucket load and length affinity to
achieve load balancing.
"""
def __init__(
self, buckets: list[tuple[int, int]], sensitivity=1.0, affinity_strength=0.1, log_func=print, all_neighbor=False
):
"""
Initialize the load balancer.
:param buckets: length range of each bucket
:param sensitivity: sensitivity to the load gap (higher = more sensitive)
:param affinity_strength: strength of length affinity (higher = a request
stays more in its standard bucket)
:param log_func: logging function
:param all_neighbor: if False, only the left/right buckets are neighbors;
if True, all buckets are neighbors
"""
self.num_buckets = len(buckets)
self.sensitivity = sensitivity
self.affinity_strength = affinity_strength
self.log_func = log_func
self.all_neighbor = all_neighbor
self.buckets = {idx: Bucket(bucket_ranges) for idx, bucket_ranges in enumerate(buckets)}
bucket_boundaries = ", ".join(
f"bucket {idx}: [{bucket.min_length}, {bucket.max_length})" for idx, bucket in self.buckets.items()
)
self._log_info(f"Initialized {self.num_buckets} buckets: {bucket_boundaries}")
# Redirect only when the redirect probability exceeds this threshold
self.base_probability_threshold = 0.12
self._log_info(f"Load Balance base_probability_threshold: {self.base_probability_threshold:.2f} ")
# Track request tasks
self.tasks: dict[AnyStr, Task] = {} # type: ignore
# Statistics
self.redirected_tasks = 0
self.total_tasks = 0
def _log_info(self, msg, *args, **kwargs):
if self.log_func:
self.log_func(msg, *args, **kwargs)
def _get_standard_bucket_index(self, task_length):
"""Return the standard bucket index for the given request length."""
for bucket_idx, bucket in self.buckets.items():
if bucket.min_length <= task_length < bucket.max_length:
return bucket_idx
# Fall back to the last bucket if the length is outside every range
return self.num_buckets - 1
def _get_neighbor_indices(self, bucket_idx):
"""Return the left and right neighbor indices of the given bucket."""
if self.all_neighbor:
return list(range(self.num_buckets))
neighbors = []
if bucket_idx > 0:
neighbors.append(bucket_idx - 1)
if bucket_idx < self.num_buckets - 1:
neighbors.append(bucket_idx + 1)
return neighbors
def _calculate_length_affinity(self, task_length, neighbor_bucket_idx):
"""
Compute the affinity factor between the task length and the neighbor
bucket (0.0 to 1.0). 1.0 means right next to the neighbor bucket, 0.0
means far away from it.
"""
neighbor_bucket = self.buckets[neighbor_bucket_idx]
neighbor_bucket_min = neighbor_bucket.min_length
neighbor_bucket_max = neighbor_bucket.max_length
if neighbor_bucket_min < task_length < neighbor_bucket_max:
raise RuntimeError("task_length must be outside the neighbor bucket range")
neighbor_bucket_center = (neighbor_bucket_min + neighbor_bucket_max) / 2.0
neighbor_bucket_half_width = (neighbor_bucket_max - neighbor_bucket_min) / 2.0
distance_to_center = abs(task_length - neighbor_bucket_center)
# The closer to the neighbor bucket boundary, the closer the affinity to 1
if neighbor_bucket_half_width > 0:
# Relative distance from the neighbor bucket half-width
normalized_distance = (distance_to_center - neighbor_bucket_half_width) / neighbor_bucket_half_width
# Exponential decay, e.g. normalized_distance=0.1, affinity_strength=1.0 -> 0.9
neighbor_affinity = math.exp(-self.affinity_strength * normalized_distance)
else:
neighbor_affinity = 1.0 # Safeguard; unreachable in practice
# Clamp to [0, 1]
return max(0.0, min(neighbor_affinity, 1.0))
def _calculate_redirect_probability(self, task_length, standard_bucket_idx, neighbor_bucket_idx):
"""
Compute the probability of redirecting to the neighbor bucket based on
the load gap and length affinity.
"""
standard_load = self.buckets[standard_bucket_idx].total_load
neighbor_load = self.buckets[neighbor_bucket_idx].total_load
# --- 1. Base probability from the load gap ---
if standard_load <= 0:
load_probability = 0.0 # No load in the standard bucket -> no redirect
else:
load_ratio = neighbor_load / max(standard_load, 1e-9) # Guard against division by zero
# The smaller the neighbor load relative to the standard load, the
# higher the redirect probability
# e.g. neighbor/standard = 5/6, sensitivity=1.0 -> 1/6 ≈ 0.1667
load_probability = 1 - load_ratio**self.sensitivity
load_probability = max(0.0, min(load_probability, 1))
# --- 2. Length affinity factor ---
affinity_factor = self._calculate_length_affinity(task_length, neighbor_bucket_idx)
# --- 3. Final probability: load gap * length affinity ---
# A large load gap is suppressed when the request length is far from the
# neighbor bucket; a modest gap can still redirect when it is close.
final_probability = load_probability * affinity_factor
return final_probability
def dispatch_single_task(self, task_id: AnyStr, task_length: int, task_load):
return self.dispatch_task(Task(task_id, task_length, task_load))
def dispatch_task(self, cur_task):
"""
Assign a bucket to a new request, considering dynamic load balancing and
length affinity.
"""
self.total_tasks += 1
standard_bucket_idx = self._get_standard_bucket_index(cur_task.length)
neighbor_indices = self._get_neighbor_indices(standard_bucket_idx)
best_neighbor_idx = None
best_redirect_prob = 0.0
# Pick the neighbor with the highest redirect probability
for neighbor_idx in neighbor_indices:
# Only consider neighbors with lower load
if self.buckets[neighbor_idx].total_load < self.buckets[standard_bucket_idx].total_load:
prob = self._calculate_redirect_probability(cur_task.length, standard_bucket_idx, neighbor_idx)
if prob > best_redirect_prob:
best_redirect_prob = prob
best_neighbor_idx = neighbor_idx
# Decide the final bucket
final_bucket_idx = standard_bucket_idx
if best_neighbor_idx is not None and best_redirect_prob > 0:
if self.base_probability_threshold < best_redirect_prob:
final_bucket_idx = best_neighbor_idx
self.redirected_tasks += 1
self._log_info(
f"{cur_task} redirected from bucket {standard_bucket_idx} to {final_bucket_idx}"
f"(prob={best_redirect_prob:.4f})"
)
# Bookkeeping on the chosen bucket
self.buckets[final_bucket_idx].task_count += 1
self.buckets[final_bucket_idx].total_load += cur_task.load
cur_task.bucket_idx = final_bucket_idx
if cur_task.id in self.tasks:
raise RuntimeError(f"Task {cur_task.id} is existed!")
else:
self.tasks[cur_task.id] = cur_task
return final_bucket_idx, cur_task
def release_task(self, task_id):
"""Release the load of a request."""
if task_id in self.tasks:
found_task = self.tasks.pop(task_id)
if 0 <= found_task.bucket_idx < self.num_buckets:
self.buckets[found_task.bucket_idx].task_count -= 1
self.buckets[found_task.bucket_idx].total_load -= found_task.load
return True
else:
raise RuntimeError(f"Bucket {found_task.bucket_idx} not found")
else:
raise RuntimeError(f"Task {task_id} not found")
def release_all_tasks(self):
for bucket in self.buckets.values():
bucket.task_count = 0
bucket.total_load = 0
self.tasks.clear()
class NoStandardBucketLoadBalancer(DynamicBucketLoadBalancer):
"""Dispatch requests by load only, with no standard bucket."""
def __init__(self, num_buckets: int, max_length: int, log_func=print):
bucket_range = math.ceil(max_length / num_buckets)
start_length = 0
buckets = []
for _ in range(num_buckets):
end_length = start_length + bucket_range
if end_length > max_length:
end_length = max_length
buckets.append((start_length, end_length))
start_length += bucket_range
super().__init__(buckets=buckets, log_func=log_func, sensitivity=100, affinity_strength=0, all_neighbor=True)

View File

@@ -0,0 +1,455 @@
# Adapted from https://github.com/vllm-project/vllm/tests/v1/kv_connector/nixl_integration/toy_proxy_server.py
# SPDX-License-Identifier: Apache-2.0
#
# Dynamic bucketing-based hybrid load balance proxy server.
# See README.md in this directory for the tutorial and usage.
import argparse
import asyncio
import functools
import heapq
import os
import sys
import uuid
from contextlib import asynccontextmanager
from dataclasses import dataclass
from typing import Any
import httpx
from dynamic_bucket_load_balancer import DynamicBucketLoadBalancer, ServerInfo, Task
from fastapi import FastAPI, Request
from fastapi.responses import StreamingResponse
try:
from vllm.logger import init_logger
logger = init_logger(__name__)
except ImportError:
import logging
logger = logging.getLogger(__name__)
# Use uvloop for a faster event loop if available
try:
import uvloop
asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
except ImportError:
pass
class ServerState:
def __init__(self, host, port):
self.host = host
self.port = port
self.url = f"http://{host}:{port}/v1"
self.client = httpx.AsyncClient(
timeout=None,
base_url=self.url,
limits=httpx.Limits(max_connections=100000, max_keepalive_connections=100000),
)
self.active_tokens = 0
self.aborted_requests = set()
def __eq__(self, other):
self_host = self.host.replace("localhost", "0.0.0.0").replace("127.0.0.1", "0.0.0.0")
other_host = other.host.replace("localhost", "0.0.0.0").replace("127.0.0.1", "0.0.0.0")
return self_host == other_host and str(self.port) == str(other.port)
def __hash__(self):
self_host = self.host.replace("localhost", "0.0.0.0").replace("127.0.0.1", "0.0.0.0")
return hash((self_host, str(self.port)))
def __repr__(self):
return f"{self.host}:{self.port}"
@dataclass(order=True)
class ServerHeapItem:
priority: float
server_idx: int
server: ServerState
class ProxyState:
def __init__(self, server_instances):
self.infer_servers: list[ServerState] = [ServerState(h, p) for h, p in server_instances]
self.req_id_lock = asyncio.Lock()
# Dynamic bucket load balancer
self.bucket_load_balancer = None
if global_args.enable_dynamic_bucket:
self.num_buckets = 2 # Two buckets (short/long) when dynamic bucketing is enabled
self.server_group_threshold = global_args.server_group_threshold
buckets = [(0, self.server_group_threshold), (self.server_group_threshold, global_args.max_request_tokens)]
self.bucket_load_balancer = DynamicBucketLoadBalancer(buckets=buckets)
else:
self.num_buckets = 1 # No bucketing by default
# Priority queue per group; smaller score = higher priority (lower load)
server_heap_items = [ServerHeapItem(0.0, i, server) for i, server in enumerate(self.infer_servers)]
self.server_heaps: list[list[ServerHeapItem]] = self._group_servers(server_heap_items, self.num_buckets)
self.server_idx_to_group_idx = {}
# Heapify each group
for idx, cur_heap in enumerate(self.server_heaps):
for server_item in cur_heap:
self.server_idx_to_group_idx[server_item.server_idx] = idx
heapq.heapify(cur_heap)
logger.info(
"Dynamic bucket enabled: %s, number of groups: %s",
global_args.enable_dynamic_bucket,
len(self.server_heaps),
)
for group_idx, cur_heap in enumerate(self.server_heaps):
logger.info("Group %s: %s", group_idx, cur_heap)
@staticmethod
def _group_servers(servers: list[ServerHeapItem], num_groups: int):
"""
Split servers into num_groups groups.
Args:
servers (list): the server list to group.
num_groups (int): the number of groups.
Returns:
list[list]: the grouped server list.
Raises:
ValueError: when num_groups <= 0.
"""
if num_groups <= 0:
raise ValueError("Num of group is illegal")
if len(servers) < num_groups:
raise ValueError("Number of servers must greater than or equal to number of groups")
n = len(servers)
if n == 0:
return [[] for _ in range(num_groups)]
elif n == 1:
return [servers]
base_size = n // num_groups
remainder = n % num_groups
groups = []
start_index = 0
for i in range(num_groups):
group_size = base_size + 1 if i < remainder else base_size
end_index = start_index + group_size
groups.append(servers[start_index:end_index])
start_index = end_index
return groups
def _update_server_priority(self, server_idx: int):
"""Update the priority of a server in the heap."""
server = self.infer_servers[server_idx]
priority = server.active_tokens
# Remove the old entry, then add the new one
group_idx = self.server_idx_to_group_idx[server_idx]
self.server_heaps[group_idx] = [
server_heap_item
for server_heap_item in self.server_heaps[group_idx]
if server_heap_item.server_idx != server_idx
]
self.server_heaps[group_idx].append(ServerHeapItem(priority, server_idx, server))
heapq.heapify(self.server_heaps[group_idx])
async def next_req_id(self):
async with self.req_id_lock:
return str(uuid.uuid4())
def select_server(self, token_count, group_idx: int):
if not self.infer_servers:
raise RuntimeError("No inference servers available")
server_heap_item: ServerHeapItem = heapq.heappop(self.server_heaps[group_idx])
chosen_server_idx = server_heap_item.server_idx
# Update the chosen server (accumulate load)
self.infer_servers[chosen_server_idx].active_tokens += token_count
# Update priority and re-add to the heap
self._update_server_priority(chosen_server_idx)
return chosen_server_idx
def release_server(self, idx: int, token_count, req_id):
self.infer_servers[idx].active_tokens -= token_count
if global_args.enable_dynamic_bucket and req_id is not None and self.bucket_load_balancer is not None:
self.bucket_load_balancer.release_task(req_id)
# Update the priority queue after release
self._update_server_priority(idx)
def calculate_request_score(self, request_length: int, max_tokens: int = 16, ignore_eos: bool = False) -> float:
if ignore_eos:
return request_length + max_tokens
else:
# Note that 0.5 is an empirical value here because we don't know
# the actual number of tokens generated before EOS.
return request_length + 0.5 * max_tokens
def calculate_request_tokens(self, request_length: int) -> float:
return request_length / 4.0
def select_server_group(self, req_id: str, request_tokens, priority_score) -> tuple[int, Task | None]:
"""Pick the best group given the request length and the current load of each group."""
if global_args.enable_dynamic_bucket and self.bucket_load_balancer is not None:
group_idx, task = self.bucket_load_balancer.dispatch_single_task(req_id, request_tokens, priority_score)
return group_idx, task
else:
return 0, None
proxy_state = None
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--port", type=int, default=8000)
parser.add_argument("--host", type=str, default="localhost")
parser.add_argument("--server-hosts", type=str, nargs="+", default=["localhost"])
parser.add_argument("--server-ports", type=int, nargs="+", default=[8001])
parser.add_argument("--max-retries", type=int, default=3, help="Maximum number of retries for HTTP requests")
parser.add_argument(
"--retry-delay", type=float, default=0.001, help="Base delay (seconds) for exponential backoff retries"
)
parser.add_argument("--server-group-threshold", type=int, default=32 * 1024, help="Threshold of server groups")
parser.add_argument("--max-request-tokens", type=int, default=128 * 1024, help="Max tokens of request")
parser.add_argument(
"--enable-dynamic-bucket", action="store_true", default=False, help="Enable dynamic bucket load Balancer"
)
args = parser.parse_args()
if len(args.server_hosts) != len(args.server_ports):
raise ValueError("Number of dp hosts must match number of dp ports")
args.server_instances = list(zip(args.server_hosts, args.server_ports))
return args
@asynccontextmanager
async def lifespan(app: FastAPI):
global proxy_state
proxy_state = ProxyState(global_args.server_instances)
logger.debug("Initialized %s dp server clients.", len(proxy_state.infer_servers))
yield
for p in proxy_state.infer_servers:
await p.client.aclose()
async def listen_for_disconnect(request: Request) -> None:
"""Return when a disconnect message is received."""
while True:
message = await request.receive()
if message["type"] == "http.disconnect":
break
def with_cancellation(handler_func):
@functools.wraps(handler_func)
async def wrapper(*args, **kwargs):
request = kwargs["request"]
handler_task = asyncio.create_task(handler_func(*args, **kwargs))
cancellation_task = asyncio.create_task(listen_for_disconnect(request))
done, pending = await asyncio.wait([handler_task, cancellation_task], return_when=asyncio.FIRST_COMPLETED)
for task in pending:
task.cancel()
if handler_task in done:
return handler_task.result()
return None
return wrapper
app = FastAPI(lifespan=lifespan)
async def stream_service_response_with_retry(
client: httpx.AsyncClient,
endpoint: str,
req_data: dict,
request_id: str,
max_retries: int = 3,
base_delay: float = 0.2,
):
headers = {"Authorization": f"Bearer {os.environ.get('OPENAI_API_KEY')}", "X-Request-Id": request_id}
for attempt in range(1, max_retries + 1):
# Reset per retry to avoid leaking a stale True from a previous iteration
first_chunk_sent = False
try:
async with client.stream("POST", endpoint, json=req_data, headers=headers) as response:
response.raise_for_status()
async for chunk in response.aiter_bytes():
first_chunk_sent = True
yield chunk
return # Success; exit after streaming completes
except (httpx.RequestError, httpx.HTTPStatusError) as e:
# After the first chunk is forwarded, retry is forbidden (would duplicate/corrupt the stream).
if first_chunk_sent:
logger.error("Streaming to client interrupted after response started: %s", str(e))
return
if attempt < max_retries:
logger.warning("Attempt %s failed for streaming %s: %s", attempt, endpoint, str(e))
await asyncio.sleep(base_delay * (2 ** (attempt - 1)))
else:
logger.error("All %s attempts failed for streaming %s.", max_retries, endpoint)
raise e
except Exception as e:
# Same guard as above for non-HTTP exceptions
if first_chunk_sent:
logger.error("Streaming to client interrupted after response started: %s", str(e))
return
if attempt < max_retries:
logger.warning("Attempt %s failed for streaming %s: %s", attempt, endpoint, str(e))
await asyncio.sleep(base_delay * (2 ** (attempt - 1)))
else:
logger.error("All %s attempts failed for streaming %s.", max_retries, endpoint)
raise e
async def _select_instance(api: str, req_data: Any, request_length: int):
# refer to vLLM sampling_params: max_token default value
max_tokens = req_data.get("max_tokens", 16)
ignore_eos = req_data.get("ignore_eos", False)
priority_score = 0.0
if global_args.enable_dynamic_bucket:
priority_score = proxy_state.calculate_request_tokens(request_length)
else:
priority_score = proxy_state.calculate_request_score(
request_length, max_tokens=max_tokens, ignore_eos=ignore_eos
)
logger.debug(
"Request length: %s, max tokens: %s, ignore_eos: %s, Priority score: %s",
request_length,
max_tokens,
ignore_eos,
priority_score,
)
request_id = await proxy_state.next_req_id()
# Select server based on priority score
request_tokens = proxy_state.calculate_request_tokens(request_length)
group_idx, task = proxy_state.select_server_group(request_id, request_tokens, priority_score)
try:
server_idx = proxy_state.select_server(priority_score, group_idx)
except Exception:
if global_args.enable_dynamic_bucket and task is not None and proxy_state.bucket_load_balancer is not None:
proxy_state.bucket_load_balancer.release_task(task.id)
raise
if global_args.enable_dynamic_bucket and task is not None:
task.server_info = ServerInfo("DP", server_idx)
chosen_server = proxy_state.infer_servers[server_idx]
logger.debug(
"[group_idx=%s, server_idx=%s] Choose server %s to process request %s",
group_idx,
server_idx,
chosen_server.url,
request_id,
)
return InstanceInfo(
request_id=request_id, server_idx=server_idx, priority_score=priority_score, server_state=chosen_server
)
@dataclass
class InstanceInfo:
request_id: str
server_idx: int
priority_score: float
server_state: ServerState
async def _handle_completions(api: str, request: Request):
# streaming_started ensures release_server runs exactly once: in
# generate_stream's finally on the normal path, or below if it never started.
instance_info = None
streaming_started = False
try:
req_data = await request.json()
req_body = await request.body()
request_length = len(req_body)
instance_info = await _select_instance(api, req_data, request_length)
async def generate_stream():
nonlocal instance_info
try:
async for chunk in stream_service_response_with_retry(
instance_info.server_state.client, # type: ignore
api,
req_data,
request_id=instance_info.request_id, # type: ignore
max_retries=global_args.max_retries,
base_delay=global_args.retry_delay,
):
yield chunk
except Exception as e:
logger.error(
"Error during streaming from server %s: %s, the aborted request is: %s.",
instance_info.server_state.url, # type: ignore
str(e),
instance_info.request_id, # type: ignore
)
finally:
# Release load after streaming completes
proxy_state.release_server( # type: ignore
instance_info.server_idx, # type: ignore
instance_info.priority_score, # type: ignore
instance_info.request_id, # type: ignore
)
streaming_started = True
return StreamingResponse(generate_stream(), media_type="application/json")
except Exception as e:
import traceback
exc_info = sys.exc_info()
print(f"Error occurred in external dp proxy server - {api} endpoint")
print(e)
print("".join(traceback.format_exception(*exc_info)))
raise
finally:
# If streaming never started (client disconnect or selection error),
# release here to avoid leaking active_tokens / the bucket task; the
# normal path already released in generate_stream.
if instance_info is not None and not streaming_started:
proxy_state.release_server(instance_info.server_idx, instance_info.priority_score, instance_info.request_id)
@app.post("/v1/completions")
@with_cancellation
async def handle_completions(request: Request):
return await _handle_completions("/completions", request)
@app.post("/v1/chat/completions")
@with_cancellation
async def handle_chat_completions(request: Request):
return await _handle_completions("/chat/completions", request)
@app.get("/healthcheck")
async def healthcheck():
return {
"status": "ok",
"server_instances": len(proxy_state.infer_servers),
}
if __name__ == "__main__":
global global_args
global_args = parse_args()
import uvicorn
uvicorn.run(app, host=global_args.host, port=global_args.port)