diff --git a/src/tests/test_routing_metrics.py b/src/tests/test_routing_metrics.py new file mode 100644 index 000000000..cdfd9fcaf --- /dev/null +++ b/src/tests/test_routing_metrics.py @@ -0,0 +1,217 @@ +import importlib.util +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from vllm_router.routers.routing_logic import ( + DisaggregatedPrefillOrchestratedRouter, + KvawareRouter, + RoundRobinRouter, + SessionRouter, + cleanup_routing_logic, +) +from vllm_router.services.metrics_service import routing_decisions_total + +_LMCACHE_AVAILABLE = importlib.util.find_spec("lmcache") is not None +requires_lmcache = pytest.mark.skipif( + not _LMCACHE_AVAILABLE, reason="lmcache not installed" +) + + +class EndpointInfo: + def __init__( + self, + url: str, + model_label: str = "", + model_names: List[str] | None = None, + ): + self.url = url + self.model_label = model_label + self.model_names = model_names or ["test-model"] + + +class RequestStats: + def __init__(self, qps: float): + self.qps = qps + + +class Request: + def __init__(self, headers: Dict[str, str], body: Dict[str, Any] = None): + self.headers = headers + self.body = body + + +@pytest.fixture(autouse=True) +def _cleanup_router_singletons(): + cleanup_routing_logic() + yield + cleanup_routing_logic() + + +def _counter_value(server: str, model: str, algorithm: str, outcome: str) -> float: + """Return the current value of the routing decisions counter for the given labels.""" + for metric in routing_decisions_total.collect(): + for sample in metric.samples: + if not sample.name.endswith("_total"): + continue + labels = sample.labels + if ( + labels.get("server") == server + and labels.get("model") == model + and labels.get("algorithm") == algorithm + and labels.get("outcome") == outcome + ): + return sample.value + return 0.0 + + +def test_roundrobin_records_success_outcome_with_model(): + router = RoundRobinRouter() + endpoints = [EndpointInfo(url="http://engine1.com", model_names=["llama-3"])] + request = Request(headers={}) + + before = _counter_value("http://engine1.com", "llama-3", "roundrobin", "success") + router.route_request(endpoints, {}, {}, request) + after = _counter_value("http://engine1.com", "llama-3", "roundrobin", "success") + + assert after - before == 1 + + +@pytest.mark.asyncio +async def test_session_router_records_fallback_when_no_session_id(): + router = SessionRouter(session_key="session_id") + endpoints = [ + EndpointInfo(url="http://engine1.com"), + EndpointInfo(url="http://engine2.com"), + ] + request_stats = { + "http://engine1.com": RequestStats(qps=10), + "http://engine2.com": RequestStats(qps=5), + } + request = Request(headers={}) + request_json = {"model": "llama-3"} + + before = _counter_value("http://engine2.com", "llama-3", "session", "fallback") + url = await router.route_request( + endpoints, None, request_stats, request, request_json + ) + after = _counter_value("http://engine2.com", "llama-3", "session", "fallback") + + assert url == "http://engine2.com" + assert after - before == 1 + + +@pytest.mark.asyncio +async def test_session_router_records_success_with_session_id(): + router = SessionRouter(session_key="session_id") + endpoints = [ + EndpointInfo(url="http://engine1.com"), + EndpointInfo(url="http://engine2.com"), + ] + request_stats = { + "http://engine1.com": RequestStats(qps=10), + "http://engine2.com": RequestStats(qps=5), + } + request = Request(headers={"session_id": "abc123"}) + request_json = {"model": "llama-3"} + + before = { + ep.url: _counter_value(ep.url, "llama-3", "session", "success") + for ep in endpoints + } + url = await router.route_request( + endpoints, None, request_stats, request, request_json + ) + after = _counter_value(url, "llama-3", "session", "success") + + assert after - before[url] == 1 + + +@requires_lmcache +@pytest.mark.asyncio +async def test_kvaware_router_records_fallback_on_kv_miss(): + # __new__ + manual field setup bypasses the lmcache controller in __init__. + router = KvawareRouter.__new__(KvawareRouter) + router._initialized = True + router.session_key = "session_id" + router.threshold = 2000 + router.tokenizer = MagicMock() + router.tokenizer.encode = MagicMock(return_value=[1, 2, 3]) + router.instance_id_to_ip = {} + from uhashring import HashRing + + router.hash_ring = HashRing() + empty_layout = MagicMock() + empty_layout.layout_info = {} + router.query_manager = AsyncMock(return_value=empty_layout) + + endpoints = [ + EndpointInfo(url="http://engine1.com"), + EndpointInfo(url="http://engine2.com"), + ] + request_stats = { + "http://engine1.com": RequestStats(qps=10), + "http://engine2.com": RequestStats(qps=5), + } + request = Request(headers={"session_id": "abc123"}) + request_json = {"model": "llama-3", "prompt": "hi"} + + before = { + ep.url: _counter_value(ep.url, "llama-3", "kvaware", "fallback") + for ep in endpoints + } + url = await router.route_request( + endpoints, None, request_stats, request, request_json + ) + after = _counter_value(url, "llama-3", "kvaware", "fallback") + + assert after - before[url] == 1 + + +def test_disaggregated_prefill_orchestrated_records_two_decisions_per_request(): + router = DisaggregatedPrefillOrchestratedRouter( + prefill_model_labels=["prefill"], + decode_model_labels=["decode"], + ) + prefill_endpoints = [EndpointInfo(url="http://prefill1.com", model_label="prefill")] + decode_endpoints = [EndpointInfo(url="http://decode1.com", model_label="decode")] + + before_prefill = _counter_value( + "http://prefill1.com", + "llama-3", + "disaggregated_prefill_orchestrated", + "success", + ) + before_decode = _counter_value( + "http://decode1.com", + "llama-3", + "disaggregated_prefill_orchestrated", + "success", + ) + + router.select_prefill_endpoint(prefill_endpoints, "llama-3") + router.select_decode_endpoint(decode_endpoints, "llama-3") + + after_prefill = _counter_value( + "http://prefill1.com", + "llama-3", + "disaggregated_prefill_orchestrated", + "success", + ) + after_decode = _counter_value( + "http://decode1.com", + "llama-3", + "disaggregated_prefill_orchestrated", + "success", + ) + + assert after_prefill - before_prefill == 1 + assert after_decode - before_decode == 1 + + +def test_record_decision_handles_missing_model_label(): + router = RoundRobinRouter() + router._record_decision(server_url="", model="") + value = _counter_value("unknown", "unknown", "roundrobin", "success") + assert value >= 1 diff --git a/src/vllm_router/routers/routing_logic.py b/src/vllm_router/routers/routing_logic.py index 1f56bffb5..3c9dbe9f4 100644 --- a/src/vllm_router/routers/routing_logic.py +++ b/src/vllm_router/routers/routing_logic.py @@ -42,6 +42,7 @@ from vllm_router.log import init_logger from vllm_router.service_discovery import EndpointInfo +from vllm_router.services.metrics_service import routing_decisions_total from vllm_router.stats.engine_stats import EngineStats from vllm_router.stats.request_stats import RequestStats from vllm_router.utils import SingletonABCMeta @@ -59,6 +60,18 @@ class RoutingLogic(str, enum.Enum): class RoutingInterface(metaclass=SingletonABCMeta): + ALGORITHM_NAME: str = "unknown" + + def _record_decision( + self, server_url: str, model: str, outcome: str = "success" + ) -> None: + routing_decisions_total.labels( + server=server_url or "unknown", + model=model or "unknown", + algorithm=self.ALGORITHM_NAME, + outcome=outcome, + ).inc() + def _qps_routing( self, endpoints: List[EndpointInfo], request_stats: Dict[str, RequestStats] ) -> str: @@ -140,6 +153,8 @@ class RoundRobinRouter(RoutingInterface): # TODO (ApostaC): when available engines in the endpoints changes, the # algorithm may not be "perfectly" round-robin. + ALGORITHM_NAME = "roundrobin" + # Upper bound on cached endpoint-set entries to prevent unbounded memory # growth when endpoints change dynamically (add / remove / update). _MAX_CACHE_SIZE = 1024 @@ -192,7 +207,12 @@ def route_request( ): self._next_index.clear() self._next_index[endpoint_urls] = idx + 1 - return endpoint_urls[idx % len(endpoint_urls)] + url = endpoint_urls[idx % len(endpoint_urls)] + selected = next((e for e in endpoints if e.url == url), endpoints[0]) + model_names = getattr(selected, "model_names", None) or [] + model = model_names[0] if model_names else "unknown" + self._record_decision(url, model) + return url class SessionRouter(RoutingInterface): @@ -201,6 +221,8 @@ class SessionRouter(RoutingInterface): in the request headers """ + ALGORITHM_NAME = "session" + def __init__(self, session_key: str = None): if hasattr(self, "_initialized"): return @@ -239,12 +261,15 @@ async def route_request( # Update the hash ring with the current list of endpoints self._update_hash_ring(endpoints) + model = request_json.get("model", "unknown") if request_json else "unknown" if session_id is None: # Route based on QPS if no session ID is present url = self._qps_routing(endpoints, request_stats) + self._record_decision(url, model, outcome="fallback") else: # Use the hash ring to get the endpoint for the session ID url = self.hash_ring.get_node(session_id) + self._record_decision(url, model) return url @@ -255,6 +280,8 @@ class KvawareRouter(RoutingInterface): of the longest prefix match is found. """ + ALGORITHM_NAME = "kvaware" + def __init__( self, lmcache_controller_port: int, @@ -386,6 +413,7 @@ async def route_request( ] # Get the first key matched_tokens = instance_id.layout_info[matched_instance_id][1] + model = request_json.get("model", "unknown") if request_json else "unknown" if ( instance_id is None or len(instance_id.layout_info) == 0 @@ -396,11 +424,10 @@ async def route_request( # Update the hash ring with the current list of endpoints self._update_hash_ring(endpoints) if session_id is None: - # Route based on QPS if no session ID is present url = self._qps_routing(endpoints, request_stats) else: - # Use the hash ring to get the endpoint for the session ID url = self.hash_ring.get_node(session_id) + self._record_decision(url, model, outcome="fallback") return url else: queried_instance_ids = [info for info in instance_id.layout_info] @@ -425,7 +452,9 @@ async def route_request( logger.info( f"Routing request to {queried_instance_ids[0]} found by kvaware router" ) - return self.instance_id_to_ip[queried_instance_ids[0]] + url = self.instance_id_to_ip[queried_instance_ids[0]] + self._record_decision(url, model) + return url class PrefixAwareRouter(RoutingInterface): @@ -436,6 +465,8 @@ class PrefixAwareRouter(RoutingInterface): In this class, we assume that there is no eviction of prefix cache. """ + ALGORITHM_NAME = "prefixaware" + def __init__(self: int): if hasattr(self, "_initialized"): return @@ -504,6 +535,8 @@ async def route_request( await self.hashtrie.insert(prompt, selected_endpoint) + model = request_json.get("model", "unknown") if request_json else "unknown" + self._record_decision(selected_endpoint, model) return selected_endpoint @@ -513,6 +546,8 @@ class DisaggregatedPrefillRouter(RoutingInterface): First request goes to prefill endpoint, then second request goes to decode endpoint. """ + ALGORITHM_NAME = "disaggregated_prefill" + def __init__(self, prefill_model_labels: List[str], decode_model_labels: List[str]): self.prefill_model_labels = prefill_model_labels self.decode_model_labels = decode_model_labels @@ -565,6 +600,8 @@ class DisaggregatedPrefillOrchestratedRouter(RoutingInterface): Load balancing: Uses round-robin across available prefill and decode pods. """ + ALGORITHM_NAME = "disaggregated_prefill_orchestrated" + def __init__(self, prefill_model_labels: List[str], decode_model_labels: List[str]): if hasattr(self, "_initialized"): return @@ -617,7 +654,9 @@ def _find_endpoints(self, endpoints: List[EndpointInfo]): return prefiller_endpoints, decoder_endpoints def select_prefill_endpoint( - self, prefiller_endpoints: List[EndpointInfo] + self, + prefiller_endpoints: List[EndpointInfo], + model: str = "unknown", ) -> EndpointInfo: """Select prefill endpoint using round-robin load balancing.""" if not prefiller_endpoints: @@ -626,10 +665,13 @@ def select_prefill_endpoint( sorted_endpoints = sorted(prefiller_endpoints, key=lambda e: e.url) selected = sorted_endpoints[self.prefill_idx % len(sorted_endpoints)] self.prefill_idx += 1 + self._record_decision(selected.url, model) return selected def select_decode_endpoint( - self, decoder_endpoints: List[EndpointInfo] + self, + decoder_endpoints: List[EndpointInfo], + model: str = "unknown", ) -> EndpointInfo: """Select decode endpoint using round-robin load balancing.""" if not decoder_endpoints: @@ -638,6 +680,7 @@ def select_decode_endpoint( sorted_endpoints = sorted(decoder_endpoints, key=lambda e: e.url) selected = sorted_endpoints[self.decode_idx % len(sorted_endpoints)] self.decode_idx += 1 + self._record_decision(selected.url, model) return selected async def route_request( diff --git a/src/vllm_router/services/metrics_service/__init__.py b/src/vllm_router/services/metrics_service/__init__.py index c5645bb8f..f5e49d1b4 100644 --- a/src/vllm_router/services/metrics_service/__init__.py +++ b/src/vllm_router/services/metrics_service/__init__.py @@ -69,3 +69,10 @@ ["server", "model", "status"], buckets=(0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0, 60.0, 120.0), ) + +# --- Routing-level metrics --- +routing_decisions_total = Counter( + "vllm:routing_decisions_total", + "Total routing decisions made by the router.", + ["server", "model", "algorithm", "outcome"], +) diff --git a/src/vllm_router/services/request_service/request.py b/src/vllm_router/services/request_service/request.py index 90bb0ec20..ec9efb1f2 100644 --- a/src/vllm_router/services/request_service/request.py +++ b/src/vllm_router/services/request_service/request.py @@ -748,8 +748,11 @@ async def route_orchestrated_disaggregated_request( ) # Use round-robin load balancing to select prefill and decode endpoints - prefill_endpoint = router.select_prefill_endpoint(prefiller_endpoints) - decode_endpoint = router.select_decode_endpoint(decoder_endpoints) + requested_model = request_json.get("model", "unknown") + prefill_endpoint = router.select_prefill_endpoint( + prefiller_endpoints, requested_model + ) + decode_endpoint = router.select_decode_endpoint(decoder_endpoints, requested_model) prefill_url = prefill_endpoint.url decode_url = decode_endpoint.url @@ -917,6 +920,7 @@ async def route_disaggregated_prefill_request( # Same as vllm, Get request_id from X-Request-Id header if available request_id = request.headers.get("X-Request-Id") or str(uuid.uuid4()) request_json = await request.json() + requested_model = request_json.get("model", "unknown") # Save original request for decode phase orig_request_json = request_json.copy() @@ -928,6 +932,8 @@ async def route_disaggregated_prefill_request( request_json.pop("max_completion_tokens", None) st = time.time() + # str() — _base_url is a yarl.URL, not a plain string. + prefill_url = str(request.app.state.prefill_client._base_url) try: await send_request_to_prefiller( request.app.state.prefill_client, endpoint, request_json, request_id @@ -935,8 +941,9 @@ async def route_disaggregated_prefill_request( et = time.time() logger.info(f"{request_id} prefill time (TTFT): {et - st:.4f}") logger.info( - f"Routing request {request_id} with session id None to {request.app.state.prefill_client._base_url} at {et}, process time = {et - in_router_time:.4f}" + f"Routing request {request_id} with session id None to {prefill_url} at {et}, process time = {et - in_router_time:.4f}" ) + request.app.state.router._record_decision(prefill_url, requested_model) # Use original request for decode phase request_json = orig_request_json except aiohttp.ClientResponseError as e: @@ -1000,9 +1007,11 @@ async def generate_stream(): yield json.dumps(error_response).encode("utf-8") curr_time = time.time() + decode_url = str(request.app.state.decode_client._base_url) logger.info( - f"Routing request {request_id} with session id None to {request.app.state.decode_client._base_url} at {curr_time}, process time = {curr_time - et:.4f}" + f"Routing request {request_id} with session id None to {decode_url} at {curr_time}, process time = {curr_time - et:.4f}" ) + request.app.state.router._record_decision(decode_url, requested_model) return StreamingResponse( generate_stream(),