diff --git a/.gitignore b/.gitignore index 3b709b7c..226b899f 100644 --- a/.gitignore +++ b/.gitignore @@ -34,6 +34,21 @@ __pycache__/ *.py[cod] *$py.class +# Generated gRPC stubs (compiled from proto/vllm_engine.proto by setup.py) +py_src/vllm_router/proto/vllm_engine_pb2.py +py_src/vllm_router/proto/vllm_engine_pb2_grpc.py +py_src/vllm_router/proto/engine_client_pb2.py +py_src/vllm_router/proto/engine_client_pb2_grpc.py + +# Local bench / scrape / driver dumps. Keep them outside this repo +# (e.g. ../logs or ../bench_artifacts), not on the feature branch. +logs/ +logs_*/ +*.log +*.err +*.prom +*.prom.err + # C extensions *.so diff --git a/MANIFEST.in b/MANIFEST.in index e1d6e7a9..9cbe22e8 100644 --- a/MANIFEST.in +++ b/MANIFEST.in @@ -1,3 +1,4 @@ # Must include: include Cargo.toml # Rust project configuration recursive-include src *.rs # Rust source files +include proto/vllm_engine.proto diff --git a/examples/simulate_consistent_hash.rs b/examples/simulate_consistent_hash.rs index 5956b90c..712f5009 100644 --- a/examples/simulate_consistent_hash.rs +++ b/examples/simulate_consistent_hash.rs @@ -309,7 +309,7 @@ fn analyze_ring_distribution( for url in worker_urls { vnodes_per_url.insert(url.clone(), 0); } - for (_, url) in ring.iter() { + for url in ring.values() { *vnodes_per_url.get_mut(url).unwrap() += 1; } diff --git a/proto/vllm_engine.proto b/proto/vllm_engine.proto new file mode 100644 index 00000000..aeb65ce7 --- /dev/null +++ b/proto/vllm_engine.proto @@ -0,0 +1,109 @@ +syntax = "proto3"; + +package vllm.engine.v1; + +// vLLM-specific Engine gRPC service implemented by VllmEngineServicer. +service VllmEngine { + // Bi-directional / Server streaming for token generation + rpc GenerateStream (GenerateRequest) returns (stream GenerateStreamResponse); + + // Unary generation (returns complete output after finishing) + rpc Generate (GenerateRequest) returns (GenerateResponse); + + // Model and backend capability discovery + rpc GetModelInfo (ModelInfoRequest) returns (ModelInfoResponse); + + // Health and readiness probing + rpc HealthCheck (HealthCheckRequest) returns (HealthCheckResponse); + + // Prefix KV cache admin. Start/Stop profile RPCs are omitted until + // the Servicer can actually enable vLLM's profiler (--profiler-config). + rpc ResetPrefixCache (EmptyRequest) returns (AdminResponse); +} + +// Request message for generation +message GenerateRequest { + string request_id = 1; + repeated uint32 prompt_token_ids = 2; + uint32 dp_rank = 3; + SamplingParams sampling_params = 4; + repeated MultimodalItem multimodal_data = 5; + ExecutionMode execution_mode = 6; + string prompt_text = 7; // Optional fallback text prompt +} + +message SamplingParams { + float temperature = 1; + float top_p = 2; + int32 max_tokens = 3; + repeated string stop_sequences = 4; + repeated uint32 stop_token_ids = 5; + float frequency_penalty = 6; + float presence_penalty = 7; + bool ignore_eos = 8; + optional uint64 seed = 9; + int32 top_k = 10; +} + +enum ExecutionMode { + NORMAL = 0; + PREFILL_ONLY = 1; + DECODE_ONLY = 2; +} + +message MultimodalItem { + string modality_type = 1; // "image", "audio", "video" + bytes raw_data = 2; // Raw byte buffer (avoids Base64 overhead) + repeated int64 shape = 3; // Tensor dimensions if preprocessed +} + +message GenerateStreamResponse { + string request_id = 1; + optional uint32 token_id = 2; + string text_delta = 3; + bool is_finished = 4; + string finish_reason = 5; + WorkerMetrics metrics = 6; +} + +message GenerateResponse { + string request_id = 1; + repeated uint32 output_token_ids = 2; + string output_text = 3; + string finish_reason = 4; + WorkerMetrics metrics = 5; +} + +message WorkerMetrics { + uint32 running_requests = 1; + uint32 waiting_requests = 2; + float kv_cache_usage_percent = 3; +} + +message ModelInfoRequest {} + +message ModelInfoResponse { + string model_name = 1; + uint32 max_model_len = 2; + uint32 dp_size = 3; + uint32 block_size = 4; + repeated string stop_tokens = 5; +} + +message HealthCheckRequest {} + +message HealthCheckResponse { + enum ServingStatus { + UNKNOWN = 0; + SERVING = 1; + NOT_SERVING = 2; + } + ServingStatus status = 1; +} + +message EmptyRequest {} + +message AdminResponse { + bool success = 1; + string message = 2; +} diff --git a/py_src/vllm_router/__init__.py b/py_src/vllm_router/__init__.py index 42762cb7..f06b4393 100644 --- a/py_src/vllm_router/__init__.py +++ b/py_src/vllm_router/__init__.py @@ -2,8 +2,9 @@ try: from vllm_router.router import Router - - __all__ = ["__version__", "Router"] except ImportError: - # Router is not available if Rust extension is not built - __all__ = ["__version__"] + Router = None + +__all__ = ["__version__"] +if Router is not None: + __all__.append("Router") diff --git a/py_src/vllm_router/proto/__init__.py b/py_src/vllm_router/proto/__init__.py new file mode 100644 index 00000000..5b88a092 --- /dev/null +++ b/py_src/vllm_router/proto/__init__.py @@ -0,0 +1,10 @@ +try: + from . import vllm_engine_pb2, vllm_engine_pb2_grpc +except ImportError as exc: + raise ImportError( + "gRPC stubs are missing. They are generated from " + "proto/vllm_engine.proto during `pip install -e .` / wheel build. " + "Re-install the package so setup.py can run grpcio-tools." + ) from exc + +__all__ = ["vllm_engine_pb2", "vllm_engine_pb2_grpc"] diff --git a/py_src/vllm_router/vllm_servicer.py b/py_src/vllm_router/vllm_servicer.py new file mode 100644 index 00000000..0b1a0c86 --- /dev/null +++ b/py_src/vllm_router/vllm_servicer.py @@ -0,0 +1,630 @@ +""" +vLLM gRPC Servicer for vllm-router. + +Exposes AsyncLLMEngine via gRPC/HTTP2 using the VllmEngine contract. +Enables high-performance binary transport, prompt_token_ids ingestion, +and native cluster admin controls. +""" + +import argparse +import asyncio +import logging +import signal +import uuid +from dataclasses import dataclass +from typing import Any, AsyncGenerator + +import grpc + +from vllm_router.proto import vllm_engine_pb2, vllm_engine_pb2_grpc + +try: + from vllm import SamplingParams, TokensPrompt + from vllm.engine.arg_utils import AsyncEngineArgs + from vllm.engine.async_llm_engine import AsyncLLMEngine + + HAS_VLLM = True +except ImportError: + HAS_VLLM = False + + +try: + from grpc_health.v1 import health, health_pb2, health_pb2_grpc + + HAS_GRPC_HEALTH = True +except ImportError: + HAS_GRPC_HEALTH = False + +try: + import setproctitle +except ImportError: + setproctitle = None + +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s [%(levelname)s] [%(name)s] %(message)s", +) +logger = logging.getLogger("vllm_servicer") + + +class VllmEngineServicer(vllm_engine_pb2_grpc.VllmEngineServicer): + """ + gRPC Servicer implementing vllm.engine.v1.VllmEngine. + Dispatches generation, metadata discovery, health checks, and admin controls. + """ + + def __init__( + self, + engine, + model_name: str, + max_model_len: int = 4096, + dp_size: int = 1, + block_size: int = 16, + ): + self.engine = engine + self.model_name = model_name + self.max_model_len = max_model_len + self.dp_size = dp_size + self.block_size = block_size + self._running_requests: int = 0 + self._waiting_requests: int = 0 + + def _get_metrics(self) -> vllm_engine_pb2.WorkerMetrics: + """Collect current queue and VRAM metrics from engine or internal counters.""" + kv_usage = 0.0 + running = self._running_requests + waiting = self._waiting_requests + + try: + # Query stats from engine if available + if hasattr(self.engine, "get_stats"): + stats = self.engine.get_stats() + kv_usage = getattr(stats, "gpu_cache_usage", 0.0) + waiting = getattr(stats, "num_waiting_sys", waiting) + running = getattr(stats, "num_running_sys", running) + elif ( + hasattr(self.engine, "stat_logger") + and self.engine.stat_logger is not None + ): + stats = getattr(self.engine.stat_logger, "stats", None) + if stats: + kv_usage = getattr(stats, "gpu_cache_usage", 0.0) + waiting = getattr(stats, "num_waiting_sys", waiting) + running = getattr(stats, "num_running_sys", running) + except Exception: + pass + + return vllm_engine_pb2.WorkerMetrics( + running_requests=running, + waiting_requests=waiting, + kv_cache_usage_percent=float(kv_usage), + ) + + def _build_sampling_params(self, pb_params: vllm_engine_pb2.SamplingParams): + """Convert Protobuf SamplingParams to vLLM SamplingParams.""" + if getattr(self.engine, "is_mock", False): + return None + if not HAS_VLLM: + logger.critical( + "vLLM is not installed. Real engine requires vLLM SamplingParams." + ) + raise RuntimeError( + "vLLM is not installed in the current Python environment." + ) + + kwargs = { + "temperature": pb_params.temperature if pb_params.temperature > 0 else 0.0, + "top_p": pb_params.top_p if pb_params.top_p > 0 else 1.0, + "max_tokens": pb_params.max_tokens if pb_params.max_tokens > 0 else 16, + "frequency_penalty": pb_params.frequency_penalty, + "presence_penalty": pb_params.presence_penalty, + "ignore_eos": pb_params.ignore_eos, + } + if pb_params.stop_sequences: + kwargs["stop"] = list(pb_params.stop_sequences) + if pb_params.stop_token_ids: + kwargs["stop_token_ids"] = list(pb_params.stop_token_ids) + if pb_params.top_k > 0: + kwargs["top_k"] = pb_params.top_k + if pb_params.HasField("seed"): + kwargs["seed"] = pb_params.seed + + return SamplingParams(**kwargs) + + async def GenerateStream( + self, + request: vllm_engine_pb2.GenerateRequest, + context: grpc.aio.ServicerContext, + ) -> AsyncGenerator[vllm_engine_pb2.GenerateStreamResponse, None]: + """ + Stream generated tokens for incoming GenerateRequest over gRPC HTTP/2. + Supports pre-tokenized prompt_token_ids, cancellation propagation, + and multi-token delta emission for speculative/multi-step outputs. + """ + # Validate data-parallel rank + if request.dp_rank >= self.dp_size: + await context.abort( + grpc.StatusCode.INVALID_ARGUMENT, + f"Requested dp_rank {request.dp_rank} is out of bounds for worker with dp_size {self.dp_size}.", + ) + return + + # Validate execution mode (P/D disaggregation scheduled for later PRs) + if request.execution_mode != vllm_engine_pb2.ExecutionMode.NORMAL: + await context.abort( + grpc.StatusCode.UNIMPLEMENTED, + f"ExecutionMode {vllm_engine_pb2.ExecutionMode.Name(request.execution_mode)} " + "is not supported in PR 1 (scheduled for disaggregated P/D in future PRs).", + ) + return + + # Validate multimodal input (scheduled for later PRs) + if request.multimodal_data: + await context.abort( + grpc.StatusCode.UNIMPLEMENTED, + "Multimodal generation is not supported in PR 1.", + ) + return + + request_id = request.request_id or f"req-{uuid.uuid4().hex[:12]}" + prompt_token_ids = list(request.prompt_token_ids) + + if not prompt_token_ids and request.prompt_text: + prompt = request.prompt_text + elif prompt_token_ids: + if getattr(self.engine, "is_mock", False): + prompt = {"prompt_token_ids": prompt_token_ids} + elif HAS_VLLM: + prompt = TokensPrompt(prompt_token_ids=prompt_token_ids) + else: + await context.abort( + grpc.StatusCode.INTERNAL, + "vLLM is not installed on this worker.", + ) + return + else: + await context.abort( + grpc.StatusCode.INVALID_ARGUMENT, + "Either prompt_token_ids or prompt_text must be provided.", + ) + return + + sampling_params = self._build_sampling_params(request.sampling_params) + + self._waiting_requests += 1 + is_first_chunk = True + prev_text = "" + prev_token_count = 0 + + try: + generator = self.engine.generate( + prompt=prompt, + sampling_params=sampling_params, + request_id=request_id, + ) + + async for request_output in generator: + if is_first_chunk: + self._waiting_requests = max(0, self._waiting_requests - 1) + self._running_requests += 1 + is_first_chunk = False + + # Check client cancellation + if context.cancelled(): + logger.info(f"Client cancelled stream for request {request_id}") + if hasattr(self.engine, "abort"): + await self.engine.abort(request_id) + break + + outputs = request_output.outputs + if not outputs: + continue + + output = outputs[0] + current_text = output.text + current_token_ids = output.token_ids + + # Compute text and token deltas + text_delta = current_text[len(prev_text) :] + prev_text = current_text + + new_token_ids = current_token_ids[prev_token_count:] + prev_token_count = len(current_token_ids) + + is_finished = request_output.finished + finish_reason = output.finish_reason or "" + + if not new_token_ids: + response = vllm_engine_pb2.GenerateStreamResponse( + request_id=request_id, + text_delta=text_delta, + is_finished=is_finished, + finish_reason=finish_reason, + metrics=self._get_metrics(), + ) + yield response + elif len(new_token_ids) == 1: + response = vllm_engine_pb2.GenerateStreamResponse( + request_id=request_id, + token_id=new_token_ids[0], + text_delta=text_delta, + is_finished=is_finished, + finish_reason=finish_reason, + metrics=self._get_metrics(), + ) + yield response + else: + # Multi-step or speculative decoding produced multiple tokens in one step. + # Emit each token so no tokens are dropped from the stream. + for idx, tok_id in enumerate(new_token_ids): + is_last = idx == len(new_token_ids) - 1 + response = vllm_engine_pb2.GenerateStreamResponse( + request_id=request_id, + token_id=tok_id, + text_delta=text_delta if is_last else "", + is_finished=is_finished if is_last else False, + finish_reason=finish_reason if is_last else "", + metrics=self._get_metrics(), + ) + yield response + + except asyncio.CancelledError: + logger.info(f"Stream cancelled by runtime for request {request_id}") + if hasattr(self.engine, "abort"): + await self.engine.abort(request_id) + raise + except Exception as e: + logger.exception( + f"Error during GenerateStream for request {request_id}: {e}" + ) + await context.abort(grpc.StatusCode.INTERNAL, str(e)) + finally: + if is_first_chunk: + self._waiting_requests = max(0, self._waiting_requests - 1) + else: + self._running_requests = max(0, self._running_requests - 1) + + async def Generate( + self, + request: vllm_engine_pb2.GenerateRequest, + context: grpc.aio.ServicerContext, + ) -> vllm_engine_pb2.GenerateResponse: + """Unary generation returning full output in a single response.""" + accumulated_text = [] + accumulated_token_ids = [] + finish_reason = "" + if not request.request_id: + request.request_id = f"req-{uuid.uuid4().hex[:12]}" + request_id = request.request_id + + async for chunk in self.GenerateStream(request, context): + if chunk.text_delta: + accumulated_text.append(chunk.text_delta) + if chunk.HasField("token_id"): + accumulated_token_ids.append(chunk.token_id) + if chunk.is_finished: + finish_reason = chunk.finish_reason + + return vllm_engine_pb2.GenerateResponse( + request_id=request_id, + output_token_ids=accumulated_token_ids, + output_text="".join(accumulated_text), + finish_reason=finish_reason, + metrics=self._get_metrics(), + ) + + async def GetModelInfo( + self, + request: vllm_engine_pb2.ModelInfoRequest, + context: grpc.aio.ServicerContext, + ) -> vllm_engine_pb2.ModelInfoResponse: + """Returns model metadata and capability info for router discovery.""" + return vllm_engine_pb2.ModelInfoResponse( + model_name=self.model_name, + max_model_len=self.max_model_len, + dp_size=self.dp_size, + block_size=self.block_size, + stop_tokens=[], + ) + + async def HealthCheck( + self, + request: vllm_engine_pb2.HealthCheckRequest, + context: grpc.aio.ServicerContext, + ) -> vllm_engine_pb2.HealthCheckResponse: + """Probes engine health and serving status.""" + status = vllm_engine_pb2.HealthCheckResponse.ServingStatus.SERVING + try: + if hasattr(self.engine, "check_health"): + await self.engine.check_health() + except Exception as e: + logger.warning(f"Engine health check failed: {e}") + status = vllm_engine_pb2.HealthCheckResponse.ServingStatus.NOT_SERVING + + return vllm_engine_pb2.HealthCheckResponse(status=status) + + async def ResetPrefixCache( + self, + request: vllm_engine_pb2.EmptyRequest, + context: grpc.aio.ServicerContext, + ) -> vllm_engine_pb2.AdminResponse: + """Clears prefix KV cache blocks in VRAM.""" + try: + if hasattr(self.engine, "reset_prefix_cache"): + await self.engine.reset_prefix_cache() + return vllm_engine_pb2.AdminResponse( + success=True, message="Prefix cache reset successfully" + ) + return vllm_engine_pb2.AdminResponse( + success=False, message="Engine does not support reset_prefix_cache" + ) + except Exception as e: + logger.error(f"Failed to reset prefix cache: {e}") + return vllm_engine_pb2.AdminResponse(success=False, message=str(e)) + + +def parse_args(): + parser = argparse.ArgumentParser(description="vLLM gRPC Engine Servicer") + parser.add_argument( + "--model", type=str, required=True, help="Model name or local filesystem path" + ) + parser.add_argument( + "--host", type=str, default="0.0.0.0", help="Host interface to bind gRPC server" + ) + parser.add_argument( + "--port", + type=int, + default=50051, + help="Port to listen for incoming gRPC HTTP/2 connections", + ) + parser.add_argument( + "--dp-size", type=int, default=1, help="Data parallel size on this node" + ) + parser.add_argument( + "--tensor-parallel-size", + type=int, + default=1, + help="Tensor parallel size per instance", + ) + parser.add_argument( + "--gpu-memory-utilization", + type=float, + default=0.90, + help="GPU memory fraction for vLLM", + ) + parser.add_argument( + "--max-model-len", type=int, default=None, help="Maximum context length" + ) + parser.add_argument( + "--trust-remote-code", + action="store_true", + help="Trust remote code from HuggingFace", + ) + parser.add_argument( + "--enforce-eager", action="store_true", help="Enforce eager execution mode" + ) + parser.add_argument( + "--enable-prefix-caching", + action=argparse.BooleanOptionalAction, + default=True, + help="Enable/disable automatic prefix caching (default: enabled)", + ) + parser.add_argument( + "--block-size", type=int, default=16, help="Token block size for PagedAttention" + ) + parser.add_argument( + "--mock-engine", + action="store_true", + help="Run with mock engine for offline unit testing without GPU", + ) + return parser.parse_args() + + +@dataclass +class _MockOutput: + index: int + text: str + token_ids: list[int] + cumulative_logprob: float = 0.0 + logprobs: Any = None + finish_reason: str | None = None + + +@dataclass +class _MockRequestOutput: + request_id: str + prompt: Any + prompt_token_ids: list[int] + prompt_logprobs: Any + outputs: list[_MockOutput] + finished: bool + + +class MockAsyncEngine: + """Mock AsyncLLMEngine for GPU-free local unit testing and CI validation.""" + + is_mock: bool = True + + def __init__(self, model_name: str): + self.model_name = model_name + + def get_stats(self): + """Mock engine statistics for metrics testing.""" + + class MockStats: + gpu_cache_usage = 0.15 + num_waiting_sys = 0 + num_running_sys = 1 + + return MockStats() + + async def generate(self, prompt, sampling_params, request_id: str): + words = [ + "Hello", + " world", + "!", + " This", + " is", + " a", + " gRPC", + " streaming", + " test", + ".", + ] + token_ids: list[int] = [] + current_text = "" + for i, word in enumerate(words): + await asyncio.sleep(0.01) + token_ids.append(1000 + i) + current_text += word + output = _MockOutput( + index=0, + text=current_text, + token_ids=list(token_ids), + cumulative_logprob=-0.1 * (i + 1), + logprobs=None, + finish_reason="stop" if i == len(words) - 1 else None, + ) + yield _MockRequestOutput( + request_id=request_id, + prompt=None, + prompt_token_ids=[1, 2, 3], + prompt_logprobs=None, + outputs=[output], + finished=(i == len(words) - 1), + ) + + async def check_health(self): + return True + + async def reset_prefix_cache(self): + logger.info("[MockEngine] reset_prefix_cache invoked") + + async def abort(self, request_id: str): + logger.info(f"[MockEngine] abort invoked for request {request_id}") + + +async def serve(args): + if setproctitle: + setproctitle.setproctitle(f"vllm::servicer:{args.port}") + + max_model_len = args.max_model_len or 4096 + block_size = args.block_size + + if args.mock_engine: + logger.info("Initializing MockAsyncEngine for offline testing...") + engine = MockAsyncEngine(model_name=args.model) + else: + if not HAS_VLLM: + logger.critical( + "vLLM is not installed in the current Python environment. " + "Cannot start the vLLM Servicer with a real engine. " + "Please install vllm (`pip install vllm`) or pass --mock-engine for testing." + ) + raise SystemExit("Error: vLLM is not installed.") + + logger.info(f"Initializing AsyncLLMEngine for model: {args.model}") + + engine_args_kwargs = { + "model": args.model, + "tensor_parallel_size": args.tensor_parallel_size, + "gpu_memory_utilization": args.gpu_memory_utilization, + "trust_remote_code": args.trust_remote_code, + "enforce_eager": args.enforce_eager, + "enable_prefix_caching": args.enable_prefix_caching, + "block_size": block_size, + } + if args.dp_size > 1: + import inspect + + sig = inspect.signature(AsyncEngineArgs) + if "data_parallel_size" in sig.parameters: + engine_args_kwargs["data_parallel_size"] = args.dp_size + else: + logger.warning( + f"--dp-size={args.dp_size} specified, but AsyncEngineArgs does not accept " + "'data_parallel_size'. For multi-process DP, launch separate Servicer " + "processes per GPU rank or use tensor parallelism." + ) + + if args.max_model_len is not None: + engine_args_kwargs["max_model_len"] = args.max_model_len + + engine_args = AsyncEngineArgs(**engine_args_kwargs) + engine = AsyncLLMEngine.from_engine_args(engine_args) + + # Retrieve resolved max_model_len and block_size from engine config if available + try: + if hasattr(engine, "engine") and hasattr(engine.engine, "model_config"): + max_model_len = engine.engine.model_config.max_model_len + if hasattr(engine, "engine") and hasattr(engine.engine, "cache_config"): + block_size = engine.engine.cache_config.block_size + except Exception as e: + logger.debug(f"Could not inspect internal engine config: {e}") + + server = grpc.aio.server( + options=[ + ("grpc.max_send_message_length", 128 * 1024 * 1024), + ("grpc.max_receive_message_length", 128 * 1024 * 1024), + ("grpc.http2.max_pings_without_data", 0), + ("grpc.keepalive_time_ms", 10000), + ("grpc.keepalive_timeout_ms", 5000), + ("grpc.keepalive_permit_without_calls", True), + ] + ) + + servicer = VllmEngineServicer( + engine=engine, + model_name=args.model, + max_model_len=max_model_len, + dp_size=args.dp_size, + block_size=block_size, + ) + vllm_engine_pb2_grpc.add_VllmEngineServicer_to_server(servicer, server) + + if HAS_GRPC_HEALTH: + health_servicer = health.HealthServicer() + health_pb2_grpc.add_HealthServicer_to_server(health_servicer, server) + health_servicer.set("", health_pb2.HealthCheckResponse.SERVING) + health_servicer.set( + "vllm.engine.v1.VllmEngine", health_pb2.HealthCheckResponse.SERVING + ) + + listen_addr = f"{args.host}:{args.port}" + server.add_insecure_port(listen_addr) + + logger.info( + f"Starting vLLM Engine Servicer on {listen_addr} (model={args.model}, dp_size={args.dp_size})" + ) + await server.start() + + stop_event = asyncio.Event() + + def signal_handler(): + logger.info("Shutdown signal received. Stopping gRPC server...") + stop_event.set() + + loop = asyncio.get_running_loop() + for sig in (signal.SIGINT, signal.SIGTERM): + try: + loop.add_signal_handler(sig, signal_handler) + except NotImplementedError: + # Signal handling on non-Unix platforms + pass + + await stop_event.wait() + logger.info("Draining gRPC server streams...") + await server.stop(grace=5.0) + logger.info("vLLM Servicer terminated cleanly.") + + +def main(): + args = parse_args() + try: + asyncio.run(serve(args)) + except (KeyboardInterrupt, SystemExit): + pass + + +if __name__ == "__main__": + main() diff --git a/py_test/e2e/test_differential_grpc_vs_http.py b/py_test/e2e/test_differential_grpc_vs_http.py new file mode 100644 index 00000000..ac8e378a --- /dev/null +++ b/py_test/e2e/test_differential_grpc_vs_http.py @@ -0,0 +1,384 @@ +""" +Differential harness: stock HTTP API server vs gRPC Servicer. + +Skipped in CPU CI (`pytest --ignore=py_test/e2e`) and skipped whenever +vLLM is not installed. Run manually against two GPUs: + + pytest py_test/e2e/test_differential_grpc_vs_http.py -v + # or: + python py_test/e2e/test_differential_grpc_vs_http.py +""" + +import asyncio +import json +import os +import subprocess +import sys +import time +from pathlib import Path +from typing import Dict, List, Tuple + +import pytest + +from vllm_router.vllm_servicer import HAS_VLLM + +pytestmark = [ + pytest.mark.e2e, + pytest.mark.skipif(not HAS_VLLM, reason="vLLM is required"), +] + +ROUTER_DIR = Path(__file__).resolve().parents[2] +MODEL_PATH = os.environ.get("MODEL_PATH", "Qwen/Qwen3.5-4B") +HTTP_PORT = 18200 +GRPC_PORT = 50055 +HOST = "127.0.0.1" + + +async def wait_for_http_health( + url: str, proc: subprocess.Popen, timeout_s: int = 360 +) -> bool: + import requests + + print(f"Waiting for HTTP server at {url} to become healthy...") + t0 = time.time() + while time.time() - t0 < timeout_s: + if proc.poll() is not None: + raise RuntimeError( + f"HTTP server process died unexpectedly with code {proc.poll()}" + ) + try: + r = requests.get(f"{url}/health", timeout=1.0) + if r.status_code == 200: + print(f"HTTP server healthy after {time.time() - t0:.1f}s!") + return True + except Exception: + pass + await asyncio.sleep(2.0) + return False + + +async def wait_for_grpc_health(client, proc: subprocess.Popen, timeout_s: int = 360): + from vllm_router.proto import vllm_engine_pb2 + + print("Waiting for gRPC server to become healthy...") + t0 = time.time() + while time.time() - t0 < timeout_s: + if proc.poll() is not None: + raise RuntimeError( + f"gRPC server process died unexpectedly with code {proc.poll()}" + ) + try: + res = await asyncio.wait_for( + client.HealthCheck(vllm_engine_pb2.HealthCheckRequest()), timeout=2.0 + ) + if res.status == vllm_engine_pb2.HealthCheckResponse.ServingStatus.SERVING: + print(f"gRPC server healthy after {time.time() - t0:.1f}s!") + return True + except Exception: + pass + await asyncio.sleep(2.0) + return False + + +def query_http_stream( + prompt: str, max_tokens: int = 30 +) -> Tuple[str, str, List[str], Dict]: + import requests + + url = f"http://{HOST}:{HTTP_PORT}/v1/completions" + payload = { + "model": MODEL_PATH, + "prompt": prompt, + "temperature": 0.0, + "seed": 42, + "max_tokens": max_tokens, + "stream": True, + } + + t0 = time.perf_counter() + r = requests.post(url, json=payload, stream=True, timeout=30.0) + r.raise_for_status() + + chunks = [] + finish_reason = "" + sample_metadata = {} + + for line in r.iter_lines(): + if not line: + continue + line_str = line.decode("utf-8") + if line_str.startswith("data: "): + data_content = line_str[6:].strip() + if data_content == "[DONE]": + break + chunk_obj = json.loads(data_content) + if not sample_metadata: + sample_metadata = { + "id": chunk_obj.get("id"), + "object": chunk_obj.get("object"), + "model": chunk_obj.get("model"), + } + choice = chunk_obj["choices"][0] + text = choice.get("text", "") + if text: + chunks.append(text) + if choice.get("finish_reason"): + finish_reason = choice.get("finish_reason") + + latency = time.perf_counter() - t0 + full_text = "".join(chunks) + return full_text, finish_reason, chunks, {"latency": latency, **sample_metadata} + + +async def query_grpc_stream( + client, + prompt_token_ids: List[int], + max_tokens: int = 30, +) -> Tuple[str, str, List[str], List[int], Dict]: + from vllm_router.proto import vllm_engine_pb2 + + req = vllm_engine_pb2.GenerateRequest( + request_id=f"diff-test-{int(time.time() * 1000)}", + prompt_token_ids=prompt_token_ids, + sampling_params=vllm_engine_pb2.SamplingParams( + temperature=0.0, + seed=42, + max_tokens=max_tokens, + ), + ) + + t0 = time.perf_counter() + stream = client.GenerateStream(req) + + text_chunks = [] + token_ids = [] + finish_reason = "" + sample_metrics = {} + + async for chunk in stream: + if chunk.text_delta: + text_chunks.append(chunk.text_delta) + if chunk.HasField("token_id"): + token_ids.append(chunk.token_id) + if chunk.is_finished: + finish_reason = chunk.finish_reason + if chunk.metrics and not sample_metrics: + sample_metrics = { + "running_requests": chunk.metrics.running_requests, + "waiting_requests": chunk.metrics.waiting_requests, + "kv_cache_usage_percent": chunk.metrics.kv_cache_usage_percent, + } + + latency = time.perf_counter() - t0 + full_text = "".join(text_chunks) + return ( + full_text, + finish_reason, + text_chunks, + token_ids, + {"latency": latency, **sample_metrics}, + ) + + +async def run_differential_harness(): + import grpc + from transformers import AutoTokenizer + + from vllm_router.proto import vllm_engine_pb2_grpc + + print("=" * 80) + print("DIFFERENTIAL TESTING: HTTP API SERVER vs. gRPC WORKER DAEMON") + print(f"Model: {MODEL_PATH}") + print("=" * 80) + + py_bin = sys.executable + + env_http = os.environ.copy() + env_http["CUDA_VISIBLE_DEVICES"] = "1" + http_log = open("/tmp/diff_test_http.log", "w") + http_cmd = [ + py_bin, + "-m", + "vllm.entrypoints.openai.api_server", + "--model", + MODEL_PATH, + "--host", + HOST, + "--port", + str(HTTP_PORT), + "--gpu-memory-utilization", + "0.70", + "--max-model-len", + "4096", + "--trust-remote-code", + "--enforce-eager", + ] + print(f"Launching HTTP API Server on GPU 1 (Port {HTTP_PORT})...") + proc_http = subprocess.Popen( + http_cmd, env=env_http, stdout=http_log, stderr=subprocess.STDOUT + ) + + env_grpc = os.environ.copy() + env_grpc["CUDA_VISIBLE_DEVICES"] = "0" + env_grpc["PYTHONPATH"] = ( + f"{ROUTER_DIR / 'py_src'}{os.pathsep}{env_grpc.get('PYTHONPATH', '')}" + ) + grpc_log = open("/tmp/diff_test_grpc.log", "w") + grpc_cmd = [ + py_bin, + "-m", + "vllm_router.vllm_servicer", + "--model", + MODEL_PATH, + "--host", + HOST, + "--port", + str(GRPC_PORT), + "--gpu-memory-utilization", + "0.70", + "--max-model-len", + "4096", + "--trust-remote-code", + "--enforce-eager", + ] + print(f"Launching gRPC Servicer on GPU 0 (Port {GRPC_PORT})...") + proc_grpc = subprocess.Popen( + grpc_cmd, env=env_grpc, stdout=grpc_log, stderr=subprocess.STDOUT + ) + + grpc_channel = grpc.aio.insecure_channel(f"{HOST}:{GRPC_PORT}") + grpc_client = vllm_engine_pb2_grpc.VllmEngineStub(grpc_channel) + + try: + print( + "\nWaiting for both servers to finish model initialization and kernel warmup..." + ) + http_ok, grpc_ok = await asyncio.gather( + wait_for_http_health(f"http://{HOST}:{HTTP_PORT}", proc_http), + wait_for_grpc_health(grpc_client, proc_grpc), + ) + + assert http_ok and grpc_ok, "Both servers must be healthy!" + print("\n>>> BOTH SERVERS ARE READY! INITIATING DIFFERENTIAL TEST SUITE <<<\n") + + tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH, trust_remote_code=True) + + test_cases = [ + ("Case 1 (Factual Knowledge)", "Q: What is the capital of France?\nA:", 30), + ( + "Case 2 (Arithmetic Calculation)", + "Calculate step by step: 25 * 14 = ", + 35, + ), + ( + "Case 3 (Python Code Generation)", + 'def is_prime(n: int) -> bool:\n """Return True if n is prime."""\n', + 40, + ), + ] + + all_passed = True + + for title, prompt, max_tokens in test_cases: + print("-" * 80) + print(f"RUNNING: {title}") + print(f"Prompt: {repr(prompt)}") + prompt_token_ids = tokenizer.encode(prompt) + print( + f"Pre-tokenized Token IDs ({len(prompt_token_ids)} tokens): {prompt_token_ids[:10]}..." + ) + + http_text, http_finish, http_chunks, http_meta = query_http_stream( + prompt, max_tokens=max_tokens + ) + + grpc_text, grpc_finish, grpc_chunks, grpc_tokens, grpc_meta = ( + await query_grpc_stream( + grpc_client, prompt_token_ids, max_tokens=max_tokens + ) + ) + + print("\n[HTTP Server Result (GPU 1)]") + print(f" Generated Text: {repr(http_text)}") + print(f" Finish Reason : {http_finish}") + print(f" Chunks Count : {len(http_chunks)}") + print(f" Latency : {http_meta['latency']:.3f}s") + print( + f" OpenAI Meta : id={http_meta.get('id')}, model={http_meta.get('model')}" + ) + + print("\n[gRPC Worker Result (GPU 0)]") + print(f" Generated Text: {repr(grpc_text)}") + print(f" Finish Reason : {grpc_finish}") + print(f" Chunks Count : {len(grpc_chunks)}") + print(f" Tokens Decoded: {len(grpc_tokens)}") + print(f" Latency : {grpc_meta['latency']:.3f}s") + print( + f" Protobuf Meta : running_reqs={grpc_meta.get('running_requests')}, kv_cache={grpc_meta.get('kv_cache_usage_percent')}" + ) + + http_tokens = tokenizer.encode(http_text, add_special_tokens=False) + + text_match = http_text == grpc_text + finish_match = http_finish == grpc_finish + token_match = http_tokens == grpc_tokens + + print("\n--- VERIFICATION VERDICT ---") + print( + f" Text Match (Byte-for-Byte) : {'PASS (IDENTICAL)' if text_match else 'FAIL'}" + ) + print( + f" Token Sequence Match : {'PASS (IDENTICAL)' if token_match else 'FAIL'}" + ) + print( + f" Finish Reason Match : {'PASS (IDENTICAL)' if finish_match else 'FAIL'}" + ) + + if not (text_match and finish_match and token_match): + all_passed = False + print("MISMATCH DETECTED!") + print(f" HTTP Text : {http_text}") + print(f" gRPC Text : {grpc_text}") + print(f" HTTP Tokens : {http_tokens}") + print(f" gRPC Tokens : {grpc_tokens}") + + assert text_match, f"Text mismatch in {title}!" + assert finish_match, f"Finish reason mismatch in {title}!" + assert token_match, f"Token sequence mismatch in {title}!" + + print("\n" + "=" * 80) + if all_passed: + print( + ">>> ALL DIFFERENTIAL PARITY TESTS PASSED: 100% IDENTICAL OUTPUT! <<<" + ) + else: + print(">>> SOME TESTS FAILED! <<<") + print("=" * 80) + + finally: + print("\nTearing down servers...") + await grpc_channel.close() + proc_http.terminate() + proc_grpc.terminate() + try: + proc_http.wait(timeout=10) + except subprocess.TimeoutExpired: + proc_http.kill() + try: + proc_grpc.wait(timeout=10) + except subprocess.TimeoutExpired: + proc_grpc.kill() + http_log.close() + grpc_log.close() + print("Servers terminated cleanly.") + + +@pytest.mark.asyncio +async def test_differential_grpc_vs_http(): + await run_differential_harness() + + +if __name__ == "__main__": + if not HAS_VLLM: + raise SystemExit("vLLM is required for this GPU harness.") + asyncio.run(run_differential_harness()) diff --git a/py_test/e2e/test_vllm_servicer_real_gpu.py b/py_test/e2e/test_vllm_servicer_real_gpu.py new file mode 100644 index 00000000..17f3ddae --- /dev/null +++ b/py_test/e2e/test_vllm_servicer_real_gpu.py @@ -0,0 +1,170 @@ +""" +Live GPU integration test for the vLLM gRPC Servicer. + +Skipped in CPU CI (`pytest --ignore=py_test/e2e`) and skipped whenever +vLLM is not installed. Run manually against a real GPU: + + pytest py_test/e2e/test_vllm_servicer_real_gpu.py -v + # or: + python py_test/e2e/test_vllm_servicer_real_gpu.py +""" + +import asyncio +import os +import subprocess +import sys +import time +from pathlib import Path + +import pytest + +from vllm_router.vllm_servicer import HAS_VLLM + +pytestmark = [ + pytest.mark.e2e, + pytest.mark.skipif(not HAS_VLLM, reason="vLLM is required"), +] + +ROUTER_DIR = Path(__file__).resolve().parents[2] +MODEL_PATH = os.environ.get("MODEL_PATH", "Qwen/Qwen3.5-4B") +PORT = 50055 +HOST = "127.0.0.1" + + +async def run_real_gpu_harness(): + import grpc + from transformers import AutoTokenizer + + from vllm_router.proto import vllm_engine_pb2, vllm_engine_pb2_grpc + + print(f"=== Starting Servicer with real model {MODEL_PATH} on GPU 0 ===") + env = os.environ.copy() + env["CUDA_VISIBLE_DEVICES"] = "0" + env["PYTHONPATH"] = ( + f"{ROUTER_DIR / 'py_src'}{os.pathsep}{env.get('PYTHONPATH', '')}" + ) + + cmd = [ + sys.executable, + "-m", + "vllm_router.vllm_servicer", + "--model", + MODEL_PATH, + "--host", + HOST, + "--port", + str(PORT), + "--gpu-memory-utilization", + "0.75", + "--max-model-len", + "4096", + "--trust-remote-code", + "--enforce-eager", + ] + + log_file = open("/tmp/vllm_servicer_gpu_test.log", "w") + proc = subprocess.Popen(cmd, env=env, stdout=log_file, stderr=subprocess.STDOUT) + + target = f"{HOST}:{PORT}" + channel = grpc.aio.insecure_channel(target) + client = vllm_engine_pb2_grpc.VllmEngineStub(channel) + + try: + print("Waiting for Servicer to initialize and load model into VRAM...") + healthy = False + for attempt in range(300): + try: + res = await asyncio.wait_for( + client.HealthCheck(vllm_engine_pb2.HealthCheckRequest()), + timeout=2.0, + ) + if ( + res.status + == vllm_engine_pb2.HealthCheckResponse.ServingStatus.SERVING + ): + healthy = True + print(f"Servicer healthy after {attempt * 2}s!") + break + except Exception: + pass + if proc.poll() is not None: + raise RuntimeError( + "Servicer process died unexpectedly! Check /tmp/vllm_servicer_gpu_test.log" + ) + await asyncio.sleep(2.0) + + if not healthy: + raise TimeoutError("Servicer failed to become healthy within timeout.") + + print("\n--- Testing RPC: GetModelInfo ---") + info = await client.GetModelInfo(vllm_engine_pb2.ModelInfoRequest()) + print(f"Model Name: {info.model_name}") + print(f"Max Model Len: {info.max_model_len}") + print(f"DP Size: {info.dp_size}") + print(f"Block Size: {info.block_size}") + assert info.model_name == MODEL_PATH + + print( + "\n--- Testing RPC: GenerateStream with pre-tokenized prompt_token_ids ---" + ) + tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH, trust_remote_code=True) + prompt_text = "Q: What is the capital of France?\nA:" + token_ids = tokenizer.encode(prompt_text) + print(f"Prompt text: '{prompt_text}'") + print(f"Tokenized IDs ({len(token_ids)} tokens): {token_ids}") + + req = vllm_engine_pb2.GenerateRequest( + request_id="real-gpu-test-01", + prompt_token_ids=token_ids, + sampling_params=vllm_engine_pb2.SamplingParams( + temperature=0.0, + max_tokens=30, + ), + ) + + stream = client.GenerateStream(req) + output_tokens = [] + output_text_chunks = [] + t0 = time.perf_counter() + + async for chunk in stream: + if chunk.HasField("token_id"): + output_tokens.append(chunk.token_id) + output_text_chunks.append(chunk.text_delta) + print(chunk.text_delta, end="", flush=True) + + ttft = time.perf_counter() - t0 + full_generated_text = "".join(output_text_chunks) + print(f"\n[Generated {len(output_tokens)} tokens in {ttft:.3f}s]") + print(f"Full text: '{full_generated_text.strip()}'") + assert "Paris" in full_generated_text + + print("\n--- Testing RPC: ResetPrefixCache ---") + cache_reset = await client.ResetPrefixCache(vllm_engine_pb2.EmptyRequest()) + print( + f"ResetPrefixCache: success={cache_reset.success}, message='{cache_reset.message}'" + ) + assert cache_reset.success is True + + print("\n=== ALL REAL GPU gRPC TESTS PASSED PERFECTLY! ===") + + finally: + await channel.close() + print("Stopping Servicer process...") + proc.terminate() + try: + proc.wait(timeout=10) + except subprocess.TimeoutExpired: + proc.kill() + log_file.close() + + +@pytest.mark.asyncio +async def test_vllm_servicer_real_gpu(): + await run_real_gpu_harness() + + +if __name__ == "__main__": + if not HAS_VLLM: + raise SystemExit("vLLM is required for this GPU harness.") + asyncio.run(run_real_gpu_harness()) diff --git a/py_test/test_vllm_servicer.py b/py_test/test_vllm_servicer.py new file mode 100644 index 00000000..dd27398b --- /dev/null +++ b/py_test/test_vllm_servicer.py @@ -0,0 +1,467 @@ +""" +Unit and integration tests for the vLLM gRPC Servicer. + +Tests RPC methods: +- GenerateStream (streaming token generation with prompt_token_ids) +- Generate (unary generation) +- GetModelInfo (metadata discovery) +- HealthCheck (serving status) +- ResetPrefixCache (prefix KV cache admin) +""" + +import asyncio +import sys +from dataclasses import dataclass, field +from typing import Any + +import grpc +import pytest +import pytest_asyncio + +from vllm_router.proto import vllm_engine_pb2, vllm_engine_pb2_grpc +from vllm_router.vllm_servicer import ( + HAS_VLLM, + VllmEngineServicer, + MockAsyncEngine, + parse_args, +) + +try: + from grpc_health.v1 import health, health_pb2, health_pb2_grpc + + HAS_GRPC_HEALTH = True +except ImportError: + HAS_GRPC_HEALTH = False + + +@pytest_asyncio.fixture +async def grpc_test_server(): + """Starts an in-process gRPC test server with MockAsyncEngine on an ephemeral port.""" + server = grpc.aio.server() + engine = MockAsyncEngine(model_name="mock-model/test-4b") + servicer = VllmEngineServicer( + engine=engine, + model_name="mock-model/test-4b", + max_model_len=8192, + dp_size=2, + block_size=16, + ) + vllm_engine_pb2_grpc.add_VllmEngineServicer_to_server(servicer, server) + + if HAS_GRPC_HEALTH: + health_servicer = health.HealthServicer() + health_pb2_grpc.add_HealthServicer_to_server(health_servicer, server) + health_servicer.set("", health_pb2.HealthCheckResponse.SERVING) + health_servicer.set( + "vllm.engine.v1.VllmEngine", health_pb2.HealthCheckResponse.SERVING + ) + + port = server.add_insecure_port("127.0.0.1:0") + await server.start() + + channel = grpc.aio.insecure_channel(f"127.0.0.1:{port}") + client = vllm_engine_pb2_grpc.VllmEngineStub(channel) + + yield { + "client": client, + "channel": channel, + "server": server, + "engine": engine, + "servicer": servicer, + "port": port, + } + + await channel.close() + await server.stop(grace=0.5) + + +@pytest.mark.asyncio +async def test_health_check(grpc_test_server): + """Verifies that HealthCheck returns SERVING status.""" + client = grpc_test_server["client"] + req = vllm_engine_pb2.HealthCheckRequest() + res = await client.HealthCheck(req) + assert res.status == vllm_engine_pb2.HealthCheckResponse.ServingStatus.SERVING + + +@pytest.mark.asyncio +async def test_get_model_info(grpc_test_server): + """Verifies model metadata discovery.""" + client = grpc_test_server["client"] + req = vllm_engine_pb2.ModelInfoRequest() + res = await client.GetModelInfo(req) + assert res.model_name == "mock-model/test-4b" + assert res.max_model_len == 8192 + assert res.dp_size == 2 + assert res.block_size == 16 + + +@pytest.mark.asyncio +async def test_generate_stream_prompt_token_ids(grpc_test_server): + """Verifies streaming token generation with pre-tokenized prompt_token_ids.""" + client = grpc_test_server["client"] + req = vllm_engine_pb2.GenerateRequest( + request_id="req-stream-test-01", + prompt_token_ids=[101, 2054, 2003], + dp_rank=0, + sampling_params=vllm_engine_pb2.SamplingParams( + temperature=0.7, + max_tokens=20, + ), + ) + + chunks = [] + async for chunk in client.GenerateStream(req): + chunks.append(chunk) + + assert len(chunks) > 0 + full_text = "".join(c.text_delta for c in chunks) + assert full_text == "Hello world! This is a gRPC streaming test." + assert chunks[-1].is_finished is True + assert chunks[-1].finish_reason == "stop" + assert chunks[-1].metrics.running_requests >= 0 + + +@pytest.mark.asyncio +async def test_generate_stream_text_fallback(grpc_test_server): + """Verifies streaming generation with prompt_text fallback.""" + client = grpc_test_server["client"] + req = vllm_engine_pb2.GenerateRequest( + request_id="req-stream-text-02", + prompt_text="Hello from text prompt", + dp_rank=1, + sampling_params=vllm_engine_pb2.SamplingParams( + temperature=0.0, + max_tokens=10, + ), + ) + + chunks = [] + async for chunk in client.GenerateStream(req): + chunks.append(chunk) + + assert len(chunks) > 0 + full_text = "".join(c.text_delta for c in chunks) + assert "Hello world!" in full_text + + +@pytest.mark.asyncio +async def test_generate_unary(grpc_test_server): + """Verifies unary Generate RPC accumulating the full output.""" + client = grpc_test_server["client"] + req = vllm_engine_pb2.GenerateRequest( + request_id="req-unary-03", + prompt_token_ids=[1, 2, 3], + sampling_params=vllm_engine_pb2.SamplingParams( + temperature=0.0, + max_tokens=15, + ), + ) + + res = await client.Generate(req) + assert res.request_id == "req-unary-03" + assert res.output_text == "Hello world! This is a gRPC streaming test." + assert len(res.output_token_ids) == 10 + assert res.finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_reset_prefix_cache(grpc_test_server): + """Verifies ResetPrefixCache forwards to the engine and reports success.""" + client = grpc_test_server["client"] + empty = vllm_engine_pb2.EmptyRequest() + + reset_res = await client.ResetPrefixCache(empty) + assert reset_res.success is True + assert "reset" in reset_res.message.lower() + + +@pytest.mark.asyncio +async def test_stream_cancellation(grpc_test_server): + """Verifies that client stream cancellation is handled cleanly.""" + client = grpc_test_server["client"] + req = vllm_engine_pb2.GenerateRequest( + request_id="req-cancel-04", + prompt_token_ids=[100, 200], + sampling_params=vllm_engine_pb2.SamplingParams(max_tokens=50), + ) + + call = client.GenerateStream(req) + chunks_received = 0 + try: + async for chunk in call: + chunks_received += 1 + if chunks_received == 2: + call.cancel() + break + except (grpc.aio.AioRpcError, asyncio.CancelledError): + pass + + assert chunks_received == 2 + + +@pytest.mark.asyncio +async def test_unsupported_execution_mode_rejected(grpc_test_server): + """Verifies that PREFILL_ONLY or DECODE_ONLY modes are rejected with UNIMPLEMENTED.""" + client = grpc_test_server["client"] + req = vllm_engine_pb2.GenerateRequest( + request_id="req-exec-mode-05", + prompt_token_ids=[1, 2, 3], + execution_mode=vllm_engine_pb2.ExecutionMode.PREFILL_ONLY, + ) + + with pytest.raises(grpc.aio.AioRpcError) as exc_info: + async for _ in client.GenerateStream(req): + pass + + assert exc_info.value.code() == grpc.StatusCode.UNIMPLEMENTED + assert "not supported in PR 1" in exc_info.value.details() + + +@pytest.mark.asyncio +async def test_unsupported_multimodal_rejected(grpc_test_server): + """Verifies that multimodal input is rejected with UNIMPLEMENTED in PR 1.""" + client = grpc_test_server["client"] + req = vllm_engine_pb2.GenerateRequest( + request_id="req-mm-06", + prompt_token_ids=[1, 2, 3], + multimodal_data=[ + vllm_engine_pb2.MultimodalItem( + modality_type="image", + raw_data=b"fake-image-bytes", + ) + ], + ) + + with pytest.raises(grpc.aio.AioRpcError) as exc_info: + async for _ in client.GenerateStream(req): + pass + + assert exc_info.value.code() == grpc.StatusCode.UNIMPLEMENTED + assert "Multimodal generation is not supported in PR 1" in exc_info.value.details() + + +@pytest.mark.asyncio +async def test_invalid_dp_rank_rejected(grpc_test_server): + """Verifies that out-of-bounds dp_rank (>= dp_size) is rejected with INVALID_ARGUMENT.""" + client = grpc_test_server["client"] + # grpc_test_server has dp_size=2, so dp_rank=5 is invalid + req = vllm_engine_pb2.GenerateRequest( + request_id="req-dp-07", + prompt_token_ids=[1, 2, 3], + dp_rank=5, + ) + + with pytest.raises(grpc.aio.AioRpcError) as exc_info: + async for _ in client.GenerateStream(req): + pass + + assert exc_info.value.code() == grpc.StatusCode.INVALID_ARGUMENT + assert "out of bounds" in exc_info.value.details() + + +@dataclass +class _TestChunkOutput: + index: int = 0 + text: str = "" + token_ids: list[int] = field(default_factory=list) + cumulative_logprob: float = 0.0 + logprobs: Any = None + finish_reason: str | None = None + + +@dataclass +class _TestStepOutput: + request_id: str + prompt: Any = None + prompt_token_ids: list[int] = field(default_factory=list) + prompt_logprobs: Any = None + outputs: list[_TestChunkOutput] = field(default_factory=list) + finished: bool = False + + +@pytest.mark.asyncio +async def test_multi_token_step_emission(grpc_test_server): + """Verifies that multi-step / speculative decoding yielding multiple tokens in one step does not drop tokens.""" + + class MultiStepEngine: + is_mock: bool = True + + def __init__(self): + pass + + async def generate(self, prompt, sampling_params, request_id: str): + # Step 1: emits 3 tokens at once: [501, 502, 503] + yield _TestStepOutput( + request_id=request_id, + prompt=None, + prompt_token_ids=[1, 2], + prompt_logprobs=None, + outputs=[ + _TestChunkOutput( + index=0, + text="Batch one", + token_ids=[501, 502, 503], + finish_reason=None, + ) + ], + finished=False, + ) + # Step 2: emits 2 more tokens: [504, 505] + yield _TestStepOutput( + request_id=request_id, + prompt=None, + prompt_token_ids=[1, 2], + prompt_logprobs=None, + outputs=[ + _TestChunkOutput( + index=0, + text="Batch one and two", + token_ids=[501, 502, 503, 504, 505], + finish_reason="stop", + ) + ], + finished=True, + ) + + server = grpc.aio.server() + engine = MultiStepEngine() + servicer = VllmEngineServicer( + engine=engine, + model_name="mock-multistep", + max_model_len=4096, + dp_size=1, + ) + vllm_engine_pb2_grpc.add_VllmEngineServicer_to_server(servicer, server) + port = server.add_insecure_port("127.0.0.1:0") + await server.start() + + channel = grpc.aio.insecure_channel(f"127.0.0.1:{port}") + client = vllm_engine_pb2_grpc.VllmEngineStub(channel) + + try: + req = vllm_engine_pb2.GenerateRequest( + request_id="req-multistep-08", + prompt_token_ids=[1, 2], + ) + + emitted_tokens = [] + async for chunk in client.GenerateStream(req): + if chunk.HasField("token_id"): + emitted_tokens.append(chunk.token_id) + + # All 5 tokens must be emitted in order, none dropped! + assert emitted_tokens == [501, 502, 503, 504, 505] + + # Verify unary Generate also receives all 5 tokens + res = await client.Generate(req) + assert list(res.output_token_ids) == [501, 502, 503, 504, 505] + assert res.output_text == "Batch one and two" + finally: + await channel.close() + await server.stop(grace=0.5) + + +@pytest.mark.asyncio +async def test_stream_token_zero_preserved(): + """Verifies that token_id=0 is preserved and not dropped by proto presence check or sentinel confusion.""" + + class ZeroTokenEngine: + is_mock: bool = True + + async def generate(self, prompt, sampling_params, request_id: str): + # Token sequence includes 0: [0, 42, 0, 99] + token_sequence = [0, 42, 0, 99] + for i, tok in enumerate(token_sequence): + is_last = i == len(token_sequence) - 1 + yield _TestStepOutput( + request_id=request_id, + prompt=None, + prompt_token_ids=[1], + prompt_logprobs=None, + outputs=[ + _TestChunkOutput( + index=0, + text=f"tok_{tok} ", + token_ids=token_sequence[: i + 1], + finish_reason="stop" if is_last else None, + ) + ], + finished=is_last, + ) + + server = grpc.aio.server() + engine = ZeroTokenEngine() + servicer = VllmEngineServicer( + engine=engine, + model_name="mock-zero-token", + max_model_len=4096, + dp_size=1, + ) + vllm_engine_pb2_grpc.add_VllmEngineServicer_to_server(servicer, server) + port = server.add_insecure_port("127.0.0.1:0") + await server.start() + + channel = grpc.aio.insecure_channel(f"127.0.0.1:{port}") + client = vllm_engine_pb2_grpc.VllmEngineStub(channel) + + try: + req = vllm_engine_pb2.GenerateRequest( + request_id="req-zero-token-01", + prompt_token_ids=[1], + ) + + emitted_tokens = [] + async for chunk in client.GenerateStream(req): + if chunk.HasField("token_id"): + emitted_tokens.append(chunk.token_id) + + # Token 0 MUST be retained and correctly distinguished from 'no token' + assert emitted_tokens == [0, 42, 0, 99] + + res = await client.Generate(req) + assert list(res.output_token_ids) == [0, 42, 0, 99] + finally: + await channel.close() + await server.stop(grace=0.5) + + +def test_prefix_caching_default_enabled(monkeypatch): + monkeypatch.setattr( + sys, "argv", ["vllm-servicer", "--model", "dummy", "--mock-engine"] + ) + args = parse_args() + assert args.enable_prefix_caching is True + + +def test_prefix_caching_can_be_disabled(monkeypatch): + monkeypatch.setattr( + sys, + "argv", + [ + "vllm-servicer", + "--model", + "dummy", + "--mock-engine", + "--no-enable-prefix-caching", + ], + ) + args = parse_args() + assert args.enable_prefix_caching is False + + +@pytest.mark.skipif(HAS_VLLM, reason="only checks the no-vLLM fail-fast path") +def test_serve_without_vllm_exits_unless_mock(): + """Real-engine startup must abort loudly when vLLM is not installed.""" + from types import SimpleNamespace + + from vllm_router.vllm_servicer import serve + + args = SimpleNamespace( + mock_engine=False, + port=50051, + max_model_len=None, + block_size=16, + ) + with pytest.raises(SystemExit, match="vLLM is not installed"): + asyncio.run(serve(args)) diff --git a/pyproject.toml b/pyproject.toml index 1c0a164d..614ffc9d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,5 +1,10 @@ [build-system] -requires = ["setuptools>=45", "wheel", "setuptools-rust>=1.5.2"] +requires = [ + "setuptools>=45", + "wheel", + "setuptools-rust>=1.5.2", + "grpcio-tools>=1.60.0", +] build-backend = "setuptools.build_meta" [project] @@ -23,6 +28,8 @@ dependencies = [ "uvicorn", "fastapi", "requests>=2.25.0", + "grpcio>=1.60.0", + "protobuf>=5.0.0", ] [project.optional-dependencies] @@ -30,10 +37,12 @@ dev = [ "pytest>=7.0.0", "pytest-asyncio>=0.21.0", "pytest-cov>=4.0.0", + "grpcio-tools>=1.60.0", ] [project.scripts] vllm-router = "vllm_router.launch_router:main" +vllm-servicer = "vllm_router.vllm_servicer:main" # https://github.com/PyO3/setuptools-rust?tab=readme-ov-file diff --git a/setup.py b/setup.py index 377d18b9..92e4e9f5 100644 --- a/setup.py +++ b/setup.py @@ -1,7 +1,99 @@ import os +import re +from pathlib import Path from setuptools import setup +ROOT = Path(__file__).resolve().parent +PROTO_DIR = ROOT / "proto" +PROTO_OUT_DIR = ROOT / "py_src" / "vllm_router" / "proto" +PROTO_NAME = "vllm_engine" +# Keep generated stubs importable against the runtime floors in pyproject.toml. +RUNTIME_GRPCIO_FLOOR = "1.60.0" + + +def generate_grpc_stubs() -> None: + """Compile proto/vllm_engine.proto into package-local Python stubs. + + Invoked during wheel builds and `pip install -e` so generated + ``*_pb2.py`` / ``*_pb2_grpc.py`` files are not checked into git. + """ + proto_file = PROTO_DIR / f"{PROTO_NAME}.proto" + if not proto_file.is_file(): + raise FileNotFoundError(f"Missing protobuf contract: {proto_file}") + + try: + from grpc_tools import protoc + except ImportError as exc: + raise RuntimeError( + "grpcio-tools is required to generate gRPC stubs. " + "It is declared in [build-system] requires; re-run " + "`pip install -e .` from a PEP 517 isolated build." + ) from exc + + PROTO_OUT_DIR.mkdir(parents=True, exist_ok=True) + rc = protoc.main( + [ + "grpc_tools.protoc", + f"-I{PROTO_DIR}", + f"--python_out={PROTO_OUT_DIR}", + f"--grpc_python_out={PROTO_OUT_DIR}", + str(proto_file), + ] + ) + if rc != 0: + raise RuntimeError(f"protoc failed with exit code {rc}") + + _patch_generated_stubs() + + +def _patch_generated_stubs() -> None: + """Make generated stubs package-relative and runtime-floor compatible.""" + grpc_path = PROTO_OUT_DIR / f"{PROTO_NAME}_pb2_grpc.py" + grpc_text = grpc_path.read_text() + grpc_text = grpc_text.replace("import warnings\n", "") + grpc_text = grpc_text.replace( + f"import {PROTO_NAME}_pb2 as {PROTO_NAME.replace('_', '__')}_pb2", + f"from . import {PROTO_NAME}_pb2 as {PROTO_NAME.replace('_', '__')}_pb2", + ) + # protoc emits `import vllm_engine_pb2 as vllm__engine__pb2` + grpc_text = grpc_text.replace( + "import vllm_engine_pb2 as vllm__engine__pb2", + "from . import vllm_engine_pb2 as vllm__engine__pb2", + ) + grpc_text = re.sub( + r'GRPC_GENERATED_VERSION = ["\'][^"\']+["\']', + f'GRPC_GENERATED_VERSION = "{RUNTIME_GRPCIO_FLOOR}"', + grpc_text, + count=1, + ) + grpc_text = grpc_text.replace( + "_version_not_supported = True", + "_version_not_supported = False", + ) + grpc_path.write_text(grpc_text) + + pb2_path = PROTO_OUT_DIR / f"{PROTO_NAME}_pb2.py" + pb2_text = pb2_path.read_text() + pb2_text = re.sub( + r"_runtime_version\.ValidateProtobufRuntimeVersion\([\s\S]*?\)", + ( + "try:\n" + " _runtime_version.ValidateProtobufRuntimeVersion(\n" + " _runtime_version.Domain.PUBLIC, 5, 0, 0, " + f'"", "{PROTO_NAME}.proto"\n' + " )\n" + "except Exception:\n" + " pass" + ), + pb2_text, + count=1, + ) + pb2_path.write_text(pb2_text) + + +generate_grpc_stubs() + no_rust = os.environ.get("VLLM_ROUTER_BUILD_NO_RUST") == "1" rust_extensions = []