456 lines
17 KiB
Python
456 lines
17 KiB
Python
# 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)
|