diff --git a/src/tests/test_audio_speech_routing.py b/src/tests/test_audio_speech_routing.py index b4ca268da..2c2d2aa20 100644 --- a/src/tests/test_audio_speech_routing.py +++ b/src/tests/test_audio_speech_routing.py @@ -14,6 +14,7 @@ from vllm_router.routers.routing_logic import ( RoundRobinRouter, ) +from vllm_router.services.request_service.retry import RetryConfig from vllm_router.utils import SingletonABCMeta @@ -48,7 +49,6 @@ def cleanup_singletons(): def setup(): """Yield a (request, router) pair with all app-state dependencies patched.""" router = RoundRobinRouter() - router.max_instance_failover_reroute_attempts = 0 sd = MagicMock() sd.get_endpoint_info.return_value = ENDPOINTS @@ -63,6 +63,7 @@ def setup(): state.semantic_cache_available = False state.callbacks = None state.external_provider_registry = None + state.retry_config = RetryConfig(max_attempts=1) req = MagicMock() req.headers = {"content-type": "application/json"} diff --git a/src/tests/test_instance_failover.py b/src/tests/test_instance_failover.py deleted file mode 100644 index bf15f6107..000000000 --- a/src/tests/test_instance_failover.py +++ /dev/null @@ -1,223 +0,0 @@ -import json -from unittest.mock import MagicMock, patch - -import pytest - -from vllm_router.routers.routing_logic import ( - RoundRobinRouter, - RoutingLogic, - initialize_routing_logic, -) -from vllm_router.utils import SingletonABCMeta - - -class EndpointInfo: - def __init__(self, url, model_names=None, sleep=False, Id=None): - self.url = url - self.model_names = model_names or ["test-model"] - self.sleep = sleep - self.Id = Id - - -@pytest.fixture(autouse=True) -def cleanup_singletons(): - yield - for cls in list(SingletonABCMeta._instances.keys()): - del SingletonABCMeta._instances[cls] - - -ENDPOINTS = [EndpointInfo(url="http://engine1"), EndpointInfo(url="http://engine2")] - -MOCK_HEADERS = MagicMock() -MOCK_HEADERS.items.return_value = [("content-type", "text/event-stream")] - - -@pytest.fixture -def setup(): - """Yield (request, background_tasks) with all dependencies patched.""" - router = RoundRobinRouter() - router.max_instance_failover_reroute_attempts = 1 - - sd = MagicMock() - sd.get_endpoint_info.return_value = ENDPOINTS - sd.aliases = None - sd.has_ever_seen_model.return_value = True - - state = MagicMock() - state.router = router - state.engine_stats_scraper.get_engine_stats.return_value = {} - state.request_stats_monitor.get_request_stats.return_value = {} - state.otel_enabled = False - state.semantic_cache_available = False - state.callbacks = None - state.external_provider_registry = None - - req = MagicMock() - req.headers = {"content-type": "application/json"} - req.query_params = {} - req.method = "POST" - req.url = "http://router/v1/chat/completions" - req.app.state = state - - async def body(): - return json.dumps({"model": "test-model", "stream": False}).encode() - - req.body = body - - patches = [ - patch( - "vllm_router.services.request_service.request.get_service_discovery", - return_value=sd, - ), - patch( - "vllm_router.services.request_service.request.is_request_rewriter_initialized", - return_value=False, - ), - ] - for p in patches: - p.start() - yield req, router - for p in patches: - p.stop() - - -def test_initialize_sets_failover_attempts(): - assert ( - initialize_routing_logic( - RoutingLogic.ROUND_ROBIN - ).max_instance_failover_reroute_attempts - == 0 - ) - assert ( - initialize_routing_logic( - RoutingLogic.ROUND_ROBIN, max_instance_failover_reroute_attempts=3 - ).max_instance_failover_reroute_attempts - == 3 - ) - - -@pytest.mark.asyncio -async def test_no_retry_on_success(setup): - req, router = setup - - async def ok(*a, **kw): - yield MOCK_HEADERS, 200 - yield b"done" - - with patch( - "vllm_router.services.request_service.request.process_request", side_effect=ok - ) as mock: - from vllm_router.services.request_service.request import route_general_request - - resp = await route_general_request(req, "/v1/chat/completions", MagicMock()) - - assert resp.status_code == 200 - assert mock.call_count == 1 - - -@pytest.mark.asyncio -async def test_retries_on_failure_with_different_url(setup): - req, router = setup - urls_called = [] - - async def fail_then_ok(r, body, server_url, *a, **kw): - urls_called.append(server_url) - if len(urls_called) == 1: - raise ConnectionError("down") - yield MOCK_HEADERS, 200 - yield b"done" - - with patch( - "vllm_router.services.request_service.request.process_request", - side_effect=fail_then_ok, - ): - from vllm_router.services.request_service.request import route_general_request - - resp = await route_general_request(req, "/v1/chat/completions", MagicMock()) - - assert resp.status_code == 200 - assert len(urls_called) == 2 - assert urls_called[0] != urls_called[1] - - -@pytest.mark.asyncio -async def test_raises_after_all_attempts_exhausted(setup): - req, router = setup - - async def fail(*a, **kw): - raise ConnectionError("down") - yield - - with patch( - "vllm_router.services.request_service.request.process_request", side_effect=fail - ): - from vllm_router.services.request_service.request import route_general_request - - with pytest.raises(ConnectionError): - await route_general_request(req, "/v1/chat/completions", MagicMock()) - - -@pytest.mark.asyncio -async def test_no_retry_when_disabled(setup): - req, router = setup - router.max_instance_failover_reroute_attempts = 0 - call_count = 0 - - async def fail(*a, **kw): - nonlocal call_count - call_count += 1 - raise ConnectionError("down") - yield - - with patch( - "vllm_router.services.request_service.request.process_request", side_effect=fail - ): - from vllm_router.services.request_service.request import route_general_request - - with pytest.raises(ConnectionError): - await route_general_request(req, "/v1/chat/completions", MagicMock()) - - assert call_count == 1 - - -@pytest.mark.asyncio -async def test_breaks_when_no_remaining_endpoints(setup): - req, router = setup - router.max_instance_failover_reroute_attempts = 5 # more retries than endpoints - - async def fail(*a, **kw): - raise ConnectionError("down") - yield - - with patch( - "vllm_router.services.request_service.request.process_request", side_effect=fail - ): - from vllm_router.services.request_service.request import route_general_request - - with pytest.raises(ConnectionError): - await route_general_request(req, "/v1/chat/completions", MagicMock()) - - -@pytest.mark.asyncio -async def test_http_exception_not_retried(setup): - from fastapi import HTTPException - - req, router = setup - router.max_instance_failover_reroute_attempts = 3 - call_count = 0 - - async def fail(*a, **kw): - nonlocal call_count - call_count += 1 - raise HTTPException(status_code=400, detail="bad request") - yield - - with patch( - "vllm_router.services.request_service.request.process_request", side_effect=fail - ): - from vllm_router.services.request_service.request import route_general_request - - with pytest.raises(HTTPException): - await route_general_request(req, "/v1/chat/completions", MagicMock()) - - assert call_count == 1 diff --git a/src/tests/test_multipart_proxy.py b/src/tests/test_multipart_proxy.py index e597374c7..76a5e12d6 100644 --- a/src/tests/test_multipart_proxy.py +++ b/src/tests/test_multipart_proxy.py @@ -78,7 +78,6 @@ async def router_client(backend_url, model=AUDIO_MODEL): router = initialize_routing_logic( RoutingLogic.ROUND_ROBIN, - max_instance_failover_reroute_attempts=0, ) stack.callback(cleanup_routing_logic) diff --git a/src/tests/test_request_retry.py b/src/tests/test_request_retry.py new file mode 100644 index 000000000..83ef0e393 --- /dev/null +++ b/src/tests/test_request_retry.py @@ -0,0 +1,401 @@ +"""End-to-end behaviour of the retry policy inside ``route_general_request``. + +``process_request`` reports a backend error status by *yielding* it, not by +raising, so every stand-in here does the same. A stand-in that raises +``HTTPException`` for a backend status would not exercise the retry path at +all and would pass regardless of whether retrying works. +""" + +import json +from unittest.mock import MagicMock, patch + +import pytest + +from vllm_router.routers.routing_logic import RoundRobinRouter +from vllm_router.services.request_service import request as request_module +from vllm_router.services.request_service import retry as retry_module +from vllm_router.services.request_service.retry import RetryConfig +from vllm_router.utils import SingletonABCMeta + + +class EndpointInfo: + def __init__(self, url, model_names=None, sleep=False, Id=None): + self.url = url + self.model_names = model_names or ["test-model"] + self.sleep = sleep + self.Id = Id + + +ENDPOINTS = [EndpointInfo(url="http://engine1"), EndpointInfo(url="http://engine2")] + +MOCK_HEADERS = MagicMock() +MOCK_HEADERS.items.return_value = [("content-type", "text/event-stream")] + +# Keep the backoff out of the test clock; the maths is covered in +# test_retry_config.py. +FAST = {"initial_backoff_ms": 1} + + +@pytest.fixture(autouse=True) +def cleanup_singletons(): + yield + for cls in list(SingletonABCMeta._instances.keys()): + del SingletonABCMeta._instances[cls] + + +def _make_request(endpoints=ENDPOINTS, retry_config=None): + """Build a request whose app state is wired for routing.""" + router = RoundRobinRouter() + + sd = MagicMock() + sd.get_endpoint_info.return_value = list(endpoints) + sd.aliases = None + sd.has_ever_seen_model.return_value = True + + state = MagicMock() + state.router = router + state.engine_stats_scraper.get_engine_stats.return_value = {} + state.request_stats_monitor.get_request_stats.return_value = {} + state.otel_enabled = False + state.semantic_cache_available = False + state.callbacks = None + state.external_provider_registry = None + state.retry_config = retry_config or RetryConfig() + + req = MagicMock() + req.headers = {"content-type": "application/json"} + req.query_params = {} + req.method = "POST" + req.url = "http://router/v1/chat/completions" + req.app.state = state + + async def body(): + return json.dumps({"model": "test-model", "stream": False}).encode() + + req.body = body + return req, router, sd + + +@pytest.fixture +def setup(): + """A request against two engines, with retries disabled by default.""" + req, router, sd = _make_request() + patches = [ + patch.object(request_module, "get_service_discovery", return_value=sd), + patch.object( + request_module, "is_request_rewriter_initialized", return_value=False + ), + ] + for p in patches: + p.start() + yield req, router + for p in patches: + p.stop() + + +def _responds(status, calls, ok_after=None): + """A ``process_request`` stand-in that yields ``status`` then a body. + + With ``ok_after`` the first ``ok_after`` calls yield ``status`` and the + rest yield 200. + """ + + async def _impl(r, body, server_url, *a, **kw): + calls.append(server_url) + code = status if ok_after is None or len(calls) <= ok_after else 200 + yield MOCK_HEADERS, code + yield b"body" + + return _impl + + +def _fails(exc, calls, ok_after=None): + """A stand-in that raises ``exc`` as a transport failure.""" + + async def _impl(r, body, server_url, *a, **kw): + calls.append(server_url) + if ok_after is None or len(calls) <= ok_after: + raise exc + yield MOCK_HEADERS, 200 + yield b"body" + + return _impl + + +async def _route(req): + return await request_module.route_general_request( + req, "/v1/chat/completions", MagicMock() + ) + + +def _patch_process(impl): + return patch.object(request_module, "process_request", side_effect=impl) + + +def _patch_discovery(sd): + return ( + patch.object(request_module, "get_service_discovery", return_value=sd), + patch.object( + request_module, "is_request_rewriter_initialized", return_value=False + ), + ) + + +def _patch_sleep(recorder): + async def fake_sleep(delay): + recorder.append(delay) + + return patch.object(retry_module.asyncio, "sleep", fake_sleep) + + +# Default configuration: exactly one attempt. + + +@pytest.mark.asyncio +async def test_successful_request_is_attempted_once(setup): + req, _ = setup + calls = [] + + with _patch_process(_responds(200, calls)): + resp = await _route(req) + + assert resp.status_code == 200 + assert len(calls) == 1 + + +@pytest.mark.asyncio +async def test_retryable_status_is_not_retried_by_default(setup): + """Retries are opt-in, so the default must preserve fail-fast behaviour.""" + req, _ = setup + calls = [] + + with _patch_process(_responds(503, calls)): + resp = await _route(req) + + assert resp.status_code == 503 + assert len(calls) == 1 + + +@pytest.mark.asyncio +async def test_transport_failure_is_not_retried_by_default(setup): + req, _ = setup + calls = [] + + with _patch_process(_fails(ConnectionError("down"), calls)): + with pytest.raises(ConnectionError): + await _route(req) + + assert len(calls) == 1 + + +# Transport failures: exclude the engine, reroute immediately. + + +@pytest.mark.asyncio +async def test_transport_failure_reroutes_to_another_engine(setup): + req, _ = setup + req.app.state.retry_config = RetryConfig(max_attempts=3, **FAST) + calls = [] + + with _patch_process(_fails(ConnectionError("down"), calls, ok_after=1)): + resp = await _route(req) + + assert resp.status_code == 200 + assert calls[0] != calls[1], "a failed engine must not be retried immediately" + + +@pytest.mark.asyncio +async def test_transport_failure_reroute_is_not_delayed(setup): + """A dead engine must fail over at once, not wait out a retry backoff.""" + req, _ = setup + req.app.state.retry_config = RetryConfig( + max_attempts=5, initial_backoff_ms=10_000, backoff_multiplier=2.0 + ) + calls, slept = [], [] + + with ( + _patch_process(_fails(ConnectionError("down"), calls, ok_after=1)), + _patch_sleep(slept), + ): + resp = await _route(req) + + assert resp.status_code == 200 + assert slept == [], "another engine was available, so no backoff is due" + + +@pytest.mark.asyncio +async def test_error_is_raised_once_every_engine_failed(setup): + req, _ = setup + req.app.state.retry_config = RetryConfig(max_attempts=2, **FAST) + calls = [] + + with _patch_process(_fails(ConnectionError("down"), calls)): + with pytest.raises(ConnectionError): + await _route(req) + + assert len(calls) == 2 + + +@pytest.mark.asyncio +async def test_budget_is_not_spent_on_already_excluded_engines(setup): + """More attempts than engines must not loop pointlessly.""" + req, _ = setup + req.app.state.retry_config = RetryConfig(max_attempts=10, **FAST) + calls = [] + + with _patch_process(_fails(ConnectionError("down"), calls)): + with pytest.raises(ConnectionError): + await _route(req) + + # Two engines, ten attempts: the pool is retried, not spun on. + assert len(calls) == 10 + assert set(calls) == {"http://engine1", "http://engine2"} + + +@pytest.mark.asyncio +async def test_single_engine_transport_failure_is_retried(): + """The only engine is retried after a backoff rather than abandoned.""" + req, _, sd = _make_request( + endpoints=[EndpointInfo(url="http://engine1")], + retry_config=RetryConfig(max_attempts=3, **FAST), + ) + calls = [] + + discovery, rewriter = _patch_discovery(sd) + with ( + discovery, + rewriter, + _patch_process(_fails(ConnectionError("down"), calls, ok_after=1)), + ): + resp = await _route(req) + + assert resp.status_code == 200 + assert calls == ["http://engine1", "http://engine1"] + + +# Transient statuses: keep the engine, back off, re-issue. + + +@pytest.mark.asyncio +async def test_retryable_status_is_retried(setup): + req, _ = setup + req.app.state.retry_config = RetryConfig(max_attempts=4, **FAST) + calls = [] + + with _patch_process(_responds(503, calls, ok_after=1)): + resp = await _route(req) + + assert resp.status_code == 200 + assert len(calls) == 2 + + +@pytest.mark.asyncio +async def test_retryable_status_keeps_the_engine_eligible(): + """A busy engine is not broken, so a single-engine pool still retries.""" + req, _, sd = _make_request( + endpoints=[EndpointInfo(url="http://engine1")], + retry_config=RetryConfig(max_attempts=3, **FAST), + ) + calls = [] + + discovery, rewriter = _patch_discovery(sd) + with ( + discovery, + rewriter, + _patch_process(_responds(503, calls, ok_after=1)), + ): + resp = await _route(req) + + assert resp.status_code == 200 + assert calls == ["http://engine1", "http://engine1"] + + +@pytest.mark.asyncio +async def test_backend_response_passes_through_when_budget_is_spent(setup): + """The client gets the backend's own status, not a synthesised error.""" + req, _ = setup + req.app.state.retry_config = RetryConfig(max_attempts=3, **FAST) + calls = [] + + with _patch_process(_responds(503, calls)): + resp = await _route(req) + + assert resp.status_code == 503 + assert len(calls) == 3 + + +@pytest.mark.asyncio +async def test_non_retryable_status_is_returned_immediately(setup): + req, _ = setup + req.app.state.retry_config = RetryConfig(max_attempts=5, **FAST) + calls = [] + + with _patch_process(_responds(400, calls)): + resp = await _route(req) + + assert resp.status_code == 400 + assert len(calls) == 1 + + +@pytest.mark.asyncio +async def test_discarded_attempt_closes_the_upstream_generator(setup): + """A discarded response must be closed so the connection is released.""" + req, _ = setup + req.app.state.retry_config = RetryConfig(max_attempts=3, **FAST) + calls, closed = [], [] + + async def impl(r, body, server_url, *a, **kw): + calls.append(server_url) + try: + yield MOCK_HEADERS, 503 if len(calls) == 1 else 200 + yield b"body" + except GeneratorExit: + closed.append(server_url) + raise + + with _patch_process(impl): + resp = await _route(req) + + assert resp.status_code == 200 + assert len(closed) == 1 + + +@pytest.mark.asyncio +async def test_backoff_grows_across_successive_retries(setup): + req, _ = setup + req.app.state.retry_config = RetryConfig( + max_attempts=4, + initial_backoff_ms=100, + backoff_multiplier=2.0, + jitter_factor=0.0, + ) + calls, slept = [], [] + + with _patch_process(_responds(503, calls)), _patch_sleep(slept): + await _route(req) + + assert slept == [0.1, 0.2, 0.4] + + +# Client errors are never retried. + + +@pytest.mark.asyncio +async def test_http_exception_is_not_retried(setup): + """A malformed request is the caller's fault; retrying cannot help.""" + from fastapi import HTTPException + + req, _ = setup + req.app.state.retry_config = RetryConfig(max_attempts=5, **FAST) + calls = [] + + async def impl(*a, **kw): + calls.append(1) + raise HTTPException(status_code=400, detail="bad request") + yield + + with _patch_process(impl): + with pytest.raises(HTTPException): + await _route(req) + + assert len(calls) == 1 diff --git a/src/tests/test_retry_config.py b/src/tests/test_retry_config.py new file mode 100644 index 000000000..626289e31 --- /dev/null +++ b/src/tests/test_retry_config.py @@ -0,0 +1,167 @@ +"""Unit tests for the retry configuration itself: status classification, +backoff maths, validation, and construction from CLI arguments.""" + +from argparse import Namespace + +import pytest + +from vllm_router.services.request_service.retry import ( + RETRIES_DISABLED, + RETRYABLE_STATUS_CODES, + RetryConfig, + is_retryable_status, +) + + +class TestIsRetryableStatus: + @pytest.mark.parametrize("status", sorted(RETRYABLE_STATUS_CODES)) + def test_retryable(self, status): + assert is_retryable_status(status) + + @pytest.mark.parametrize( + "status", [200, 201, 204, 400, 401, 403, 404, 405, 409, 422, 499, 501] + ) + def test_not_retryable(self, status): + assert not is_retryable_status(status) + + def test_documented_set(self): + """The set is part of the CLI contract, so pin it explicitly.""" + assert RETRYABLE_STATUS_CODES == {408, 429, 500, 502, 503, 504} + + +class TestDefaults: + def test_retrying_is_opt_in(self): + """A default config must reproduce the historical single attempt.""" + config = RetryConfig() + assert config.max_attempts == 1 + assert not config.enabled + + def test_shared_disabled_instance(self): + assert RETRIES_DISABLED == RetryConfig() + assert not RETRIES_DISABLED.enabled + + def test_enabled_once_more_than_one_attempt(self): + assert RetryConfig(max_attempts=2).enabled + + def test_is_immutable(self): + config = RetryConfig() + with pytest.raises(Exception): + config.max_attempts = 5 + + +class TestBackoff: + def test_grows_by_the_multiplier(self): + config = RetryConfig( + max_attempts=5, + initial_backoff_ms=50, + backoff_multiplier=1.5, + jitter_factor=0.0, + ) + assert config.backoff_seconds(0) == pytest.approx(0.05) + assert config.backoff_seconds(1) == pytest.approx(0.075) + assert config.backoff_seconds(2) == pytest.approx(0.1125) + assert config.backoff_seconds(3) == pytest.approx(0.16875) + + def test_honours_a_custom_multiplier(self): + config = RetryConfig( + max_attempts=5, + initial_backoff_ms=100, + backoff_multiplier=2.0, + jitter_factor=0.0, + ) + assert [config.backoff_seconds(i) for i in range(4)] == pytest.approx( + [0.1, 0.2, 0.4, 0.8] + ) + + def test_is_capped(self): + config = RetryConfig(max_attempts=5, max_backoff_ms=100, jitter_factor=0.0) + assert config.backoff_seconds(10) == pytest.approx(0.1) + assert config.backoff_seconds(100) == pytest.approx(0.1) + + def test_jitter_spreads_delays_within_bounds(self): + """Jitter is what stops a fleet retrying in lockstep.""" + config = RetryConfig(max_attempts=5, jitter_factor=0.5) + delays = [config.backoff_seconds(0) for _ in range(200)] + + assert min(delays) < max(delays), "jitter must actually vary the delay" + assert all(0.025 <= d <= 0.075 for d in delays), "base 50ms +/- 50%" + + def test_jitter_can_be_switched_off(self): + config = RetryConfig(max_attempts=5, jitter_factor=0.0) + assert len({config.backoff_seconds(0) for _ in range(20)}) == 1 + + def test_delay_is_never_negative(self): + config = RetryConfig(max_attempts=5, jitter_factor=1.0) + assert all(config.backoff_seconds(0) >= 0.0 for _ in range(200)) + + +class TestValidation: + @pytest.mark.parametrize( + "kwargs, message", + [ + ({"max_attempts": 0}, "max attempts"), + ({"initial_backoff_ms": 0}, "initial backoff"), + ({"initial_backoff_ms": -1}, "initial backoff"), + ({"initial_backoff_ms": 500, "max_backoff_ms": 100}, "max backoff"), + ({"backoff_multiplier": 0.9}, "multiplier"), + ({"jitter_factor": -0.1}, "jitter"), + ({"jitter_factor": 1.1}, "jitter"), + ], + ) + def test_rejects_invalid(self, kwargs, message): + with pytest.raises(ValueError, match=message): + RetryConfig(**kwargs) + + def test_accepts_boundaries(self): + RetryConfig(max_attempts=1, jitter_factor=0.0, backoff_multiplier=1.0) + RetryConfig(max_attempts=2, jitter_factor=1.0) + RetryConfig(initial_backoff_ms=50, max_backoff_ms=50) + + +class TestFromArgs: + @staticmethod + def _args(**overrides): + args = Namespace( + enable_retries=True, + max_retries=5, + initial_backoff_ms=50, + max_backoff_ms=30000, + backoff_multiplier=1.5, + jitter_factor=0.2, + ) + for key, value in overrides.items(): + setattr(args, key, value) + return args + + def test_disabled_ignores_every_other_flag(self): + """A config template carrying retry values must not change behaviour.""" + config = RetryConfig.from_args( + self._args(enable_retries=False, max_retries=99, backoff_multiplier=0.1) + ) + assert config == RetryConfig() + assert not config.enabled + + def test_absent_flag_is_treated_as_disabled(self): + assert RetryConfig.from_args(Namespace()) == RetryConfig() + + def test_enabled_maps_every_flag(self): + config = RetryConfig.from_args( + self._args( + max_retries=10, + initial_backoff_ms=100, + max_backoff_ms=60000, + backoff_multiplier=2.0, + jitter_factor=0.1, + ) + ) + assert config == RetryConfig( + max_attempts=10, + initial_backoff_ms=100, + max_backoff_ms=60000, + backoff_multiplier=2.0, + jitter_factor=0.1, + ) + + def test_invalid_values_surface_as_value_error(self): + with pytest.raises(ValueError, match="jitter"): + RetryConfig.from_args(self._args(jitter_factor=5.0)) diff --git a/src/tests/test_transcription_streaming.py b/src/tests/test_transcription_streaming.py index ad45bb252..7fb951e84 100644 --- a/src/tests/test_transcription_streaming.py +++ b/src/tests/test_transcription_streaming.py @@ -51,7 +51,6 @@ def setup_mocks(): def _make_mock_request(): router = RoundRobinRouter() - router.max_instance_failover_reroute_attempts = 0 state = MagicMock() state.router = router diff --git a/src/vllm_router/README.md b/src/vllm_router/README.md index 125744965..edae441f5 100644 --- a/src/vllm_router/README.md +++ b/src/vllm_router/README.md @@ -61,6 +61,75 @@ The router can be configured using command-line arguments. Below are the availab - `--sentry-traces-sample-rate`: The sample rate for Sentry traces (0.0 to 1.0). Default is 0.1 (10%). - `--sentry-profile-session-sample-rate`: The sample rate for Sentry profiling sessions (0.0 to 1.0). Default is 1.0 (100%). +### Retry Configuration + +A forwarded request can fail in two ways, and the router responds to each differently: + +| Failure | Example | Response | +| --- | --- | --- | +| Transport failure | connection refused, timeout | The engine is excluded and the request is rerouted to another engine **immediately** | +| Retryable status | 408, 429, 500, 502, 503, 504 | The engine stays eligible and the request is re-issued **after a backoff** | + +Both draw on one budget, `--max-retries`. **Retrying is disabled by default**, so a +request is attempted exactly once unless `--enable-retries` is passed. + +- `--enable-retries`: Turn on retrying. Disabled by default. +- `--max-retries`: Maximum total attempts per request, counting the initial one, so 5 means one attempt plus up to four retries. Must be at least 2. Default is 5. +- `--initial-backoff-ms`: Backoff before the first retry. Default is 50. +- `--max-backoff-ms`: Upper bound on the backoff. Default is 30000. +- `--backoff-multiplier`: Growth factor between retries. Default is 1.5. +- `--jitter-factor`: Randomisation applied to each delay (0.0-1.0). Default is 0.2. + +#### Exponential backoff with jitter + +```text +delay = min(initial_backoff_ms x multiplier ^ retry, max_backoff_ms) +delay' = delay x (1 + U[-jitter_factor, +jitter_factor]) +``` + +Jitter is what prevents a [thundering herd](https://medium.com/@avnein4988/mitigating-the-thundering-herd-problem-exponential-backoff-with-jitter-b507cdf90d62): +without it, every router that backed off from the same incident retries in lockstep and +re-creates the overload it was backing off from. With the defaults, retries land at +roughly 50ms, 75ms, 112ms and 169ms, each spread across a +/-20% window. + +#### Behaviour worth knowing + +- **A retryable status does not blacklist the engine.** A busy engine is not a broken + one, so it remains a candidate after the backoff. This is what makes retrying useful + on a single-engine deployment. +- **Failover between healthy engines is never delayed.** The backoff applies only when + every engine has been excluded and the pool is given another chance. +- **The backend response is passed through once the budget is spent** - the client sees + the engine's own status and body, not a synthesised router error. +- **Non-retryable statuses return on the first attempt.** A 400 is the caller's problem + and retrying cannot fix it. +- **Retries only apply before streaming begins.** Once the first chunk has been + forwarded the response headers are already with the client, so a mid-stream error is + passed through untouched. Buffering whole responses to avoid this would defeat the + point of streaming; clients that need more can retry themselves. + +> **Replaces `--max-instance-failover-reroute-attempts`.** That flag rerouted a failed +> request to another engine and is now subsumed by `--max-retries`, which covers the +> same case and adds backoff. Its default was `0` (a single attempt), which is also the +> default here, so deployments that never set it are unaffected. Deployments that did +> set it should pass `--enable-retries --max-retries `. + +**Example with retry configuration:** + +```bash +vllm-router --port 8000 \ + --service-discovery static \ + --static-backends "http://localhost:9001,http://localhost:9002" \ + --static-models "facebook/opt-125m,facebook/opt-125m" \ + --routing-logic roundrobin \ + --enable-retries \ + --max-retries 5 \ + --initial-backoff-ms 100 \ + --max-backoff-ms 60000 \ + --backoff-multiplier 2.0 \ + --jitter-factor 0.1 +``` + ## Build docker image ```bash diff --git a/src/vllm_router/app.py b/src/vllm_router/app.py index 20673aa1c..acd16a7f6 100644 --- a/src/vllm_router/app.py +++ b/src/vllm_router/app.py @@ -45,6 +45,7 @@ from vllm_router.services.batch_service import initialize_batch_processor from vllm_router.services.callbacks_service.callbacks import configure_custom_callbacks from vllm_router.services.files_service import initialize_storage +from vllm_router.services.request_service.retry import RetryConfig from vllm_router.services.request_service.rewriter import ( get_request_rewriter, ) @@ -270,6 +271,8 @@ def initialize_all(app: FastAPI, args): if args.callbacks: configure_custom_callbacks(args.callbacks, app) + app.state.retry_config = RetryConfig.from_args(args) + initialize_routing_logic( args.routing_logic, session_key=args.session_key, @@ -285,7 +288,6 @@ def initialize_all(app: FastAPI, args): priority_field=args.priority_field, priority_default=args.priority_default, priority_threshold=args.priority_threshold, - max_instance_failover_reroute_attempts=args.max_instance_failover_reroute_attempts, lmcache_health_check_interval=args.lmcache_health_check_interval, lmcache_worker_timeout=args.lmcache_worker_timeout, ) diff --git a/src/vllm_router/parsers/parser.py b/src/vllm_router/parsers/parser.py index 73cee38d0..03c97798a 100644 --- a/src/vllm_router/parsers/parser.py +++ b/src/vllm_router/parsers/parser.py @@ -20,6 +20,7 @@ from vllm_router.parsers.yaml_utils import ( read_and_process_yaml_config_file, ) +from vllm_router.services.request_service.retry import RetryConfig from vllm_router.version import __version__ try: @@ -120,6 +121,14 @@ def validate_args(args): raise ValueError( "Sentry profile session sample rate must be between 0.0 and 1.0." ) + if args.enable_retries and args.max_retries < 2: + raise ValueError( + "--max-retries counts the initial attempt, so it must be at least 2 " + "when --enable-retries is set; 1 would disable retrying." + ) + # Surface the remaining RetryConfig invariants at startup, not on the + # first request. + RetryConfig.from_args(args) def parse_args(): @@ -508,11 +517,49 @@ def parse_args(): "Only used when --routing-logic=priority.", ) - parser.add_argument( - "--max-instance-failover-reroute-attempts", + retry_group = parser.add_argument_group( + "Retry Configuration", + "Configure retry behavior with exponential backoff (disabled by default)", + ) + retry_group.add_argument( + "--enable-retries", + action="store_true", + help="Retry requests that fail with a transient error, using exponential " + "backoff with jitter. Covers transport failures (rerouted to another engine) " + "and retryable statuses (408, 429, 500, 502, 503, 504). Disabled by default, " + "so a request is attempted exactly once.", + ) + retry_group.add_argument( + "--max-retries", type=int, - default=0, - help="Number of reroute attempts per failed request", + default=5, + help="Maximum total attempts per request, counting the initial one, so 5 " + "means one attempt plus up to four retries (default: 5). Only used with " + "--enable-retries.", + ) + retry_group.add_argument( + "--initial-backoff-ms", + type=int, + default=50, + help="Initial backoff duration in milliseconds (default: 50)", + ) + retry_group.add_argument( + "--max-backoff-ms", + type=int, + default=30000, + help="Maximum backoff duration in milliseconds (default: 30000)", + ) + retry_group.add_argument( + "--backoff-multiplier", + type=float, + default=1.5, + help="Exponential backoff multiplier (default: 1.5)", + ) + retry_group.add_argument( + "--jitter-factor", + type=float, + default=0.2, + help="Random jitter factor (0.0-1.0) to prevent thundering herd (default: 0.2)", ) parser.add_argument( diff --git a/src/vllm_router/routers/routing_logic.py b/src/vllm_router/routers/routing_logic.py index 81c9aa76b..b5b0c8e69 100644 --- a/src/vllm_router/routers/routing_logic.py +++ b/src/vllm_router/routers/routing_logic.py @@ -1180,9 +1180,6 @@ def initialize_routing_logic( else: raise ValueError(f"Invalid routing logic {routing_logic}") - router.max_instance_failover_reroute_attempts = kwargs.get( - "max_instance_failover_reroute_attempts", 0 - ) return router diff --git a/src/vllm_router/services/request_service/request.py b/src/vllm_router/services/request_service/request.py index 44b8cb498..be795005d 100644 --- a/src/vllm_router/services/request_service/request.py +++ b/src/vllm_router/services/request_service/request.py @@ -36,6 +36,11 @@ SessionRouter, ) from vllm_router.service_discovery import get_service_discovery +from vllm_router.services.request_service.retry import ( + RETRIES_DISABLED, + RetryConfig, + RetryState, +) from vllm_router.services.request_service.rewriter import ( get_request_rewriter, is_request_rewriter_initialized, @@ -318,6 +323,8 @@ async def process_request( timeout=aiohttp.ClientTimeout(total=None), ) as backend_response: http_status_code = backend_response.status + if http_status_code >= 400: + request_status = "error" # Set response status on span if tracing if span is not None: span.set_attribute("http.status_code", backend_response.status) @@ -342,9 +349,6 @@ async def process_request( backend_url, request_id, end_time ) - if http_status_code is not None and http_status_code >= 400: - request_status = "error" - # Track token usage for non-streaming requests if not is_streaming and full_response: try: @@ -387,6 +391,28 @@ async def process_request( end_span(span) if tracing_active else None +async def _select_backend( + request: Request, + candidates: list, + engine_stats, + request_stats, + request_json: dict, + request_endpoint, +) -> str: + """Return the URL of the engine to forward to.""" + if request_endpoint: + return candidates[0].url + + router = request.app.state.router + if isinstance( + router, (KvawareRouter, PrefixAwareRouter, SessionRouter, PriorityRouter) + ): + return await router.route_request( + candidates, engine_stats, request_stats, request, request_json + ) + return router.route_request(candidates, engine_stats, request_stats, request) + + async def route_general_request( request: Request, endpoint: str, background_tasks: BackgroundTasks ): @@ -572,24 +598,14 @@ async def route_general_request( ) logger.debug(f"Routing request {request_id} for model: {requested_model}") + server_url = await _select_backend( + request, endpoints, engine_stats, request_stats, request_json, request_endpoint + ) if request_endpoint: - server_url = endpoints[0].url logger.debug( f"Routing request {request_id} to engine with Id: {endpoints[0].Id}" ) - elif isinstance( - request.app.state.router, - (KvawareRouter, PrefixAwareRouter, SessionRouter, PriorityRouter), - ): - server_url = await request.app.state.router.route_request( - endpoints, engine_stats, request_stats, request, request_json - ) - else: - server_url = request.app.state.router.route_request( - endpoints, engine_stats, request_stats, request - ) - if isinstance(request.app.state.router, PriorityRouter): # PriorityRouter injects the resolved priority into request_json so # vLLM's own priority scheduler can preempt within the engine. @@ -621,31 +637,24 @@ async def route_general_request( "vllm.routing_logic", type(request.app.state.router).__name__ ) - error_urls = set() - last_error = None - max_attempts = request.app.state.router.max_instance_failover_reroute_attempts + 1 - - for attempt in range(max_attempts): - if attempt > 0: - remaining = [ep for ep in endpoints if ep.url not in error_urls] - if not remaining: - break - if request_endpoint: - server_url = remaining[0].url - elif isinstance( - request.app.state.router, - (KvawareRouter, PrefixAwareRouter, SessionRouter, PriorityRouter), - ): - server_url = await request.app.state.router.route_request( - remaining, engine_stats, request_stats, request, request_json - ) - else: - server_url = request.app.state.router.route_request( - remaining, engine_stats, request_stats, request - ) + retry_config = getattr(request.app.state, "retry_config", None) + if not isinstance(retry_config, RetryConfig): + retry_config = RETRIES_DISABLED + state = RetryState(retry_config, request_id) + + async for candidates in state.attempts(endpoints): + if state.attempt_number > 1: + server_url = await _select_backend( + request, + candidates, + engine_stats, + request_stats, + request_json, + request_endpoint, + ) logger.info( f"Routing request {request_id} to {server_url} " - f"(attempt {attempt + 1}/{max_attempts})" + f"(attempt {state.attempt_number}/{state.max_attempts})" ) if span is not None: span.set_attribute("vllm.backend_url", server_url) @@ -661,7 +670,14 @@ async def route_general_request( background_tasks, parent_span_context=span_context, ) + # process_request yields the backend status rather than raising it. headers, status = await anext(stream_generator) + if state.should_retry_status(status): + # Close to free the upstream connection. The engine is not + # excluded: it is busy, not broken. + await stream_generator.aclose() + state.record_transient_status(server_url, status) + continue media_type = headers.get("content-type", "text/event-stream") headers_dict = { key: value @@ -670,21 +686,19 @@ async def route_general_request( and key.lower() != "content-type" } headers_dict["X-Request-Id"] = request_id - last_error = None + state.record_response() break except HTTPException: + # Only raised for a malformed request, never for a backend status. raise - except Exception as e: - error_urls.add(server_url) - last_error = e - logger.warning( - f"Request {request_id} failed on {server_url} " - f"(attempt {attempt + 1}/{max_attempts}): {e}" - ) - - if last_error: - end_span(span, error=last_error, status_code=500) if tracing_active else None - raise last_error + except Exception as error: + state.record_transport_failure(server_url, error) + + if state.last_error: + status_code = getattr(state.last_error, "status_code", 500) + if tracing_active: + end_span(span, error=state.last_error, status_code=status_code) + raise state.last_error # Wrap the generator to end parent span when streaming completes async def traced_stream(): diff --git a/src/vllm_router/services/request_service/retry.py b/src/vllm_router/services/request_service/retry.py new file mode 100644 index 000000000..5bbc5fd00 --- /dev/null +++ b/src/vllm_router/services/request_service/retry.py @@ -0,0 +1,223 @@ +# Copyright 2024-2025 The vLLM Production Stack Authors. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Retry policy for requests forwarded to a backend engine. + +A forwarded request fails in one of two ways, which call for opposite responses: + +* **Transport failure** (connection refused, timeout). The engine could not serve + the request at all, so it is excluded and the request is rerouted to another + engine immediately -- waiting would only add latency. +* **Retryable status** (408, 429, 500, 502, 503, 504). The engine is reachable + but cannot take the request right now, so it stays eligible and the request is + re-issued after a backoff. + +Backoff is exponential with jitter:: + + delay = min(initial_backoff_ms * multiplier ** retry, max_backoff_ms) + delay' = delay * (1 + U[-jitter_factor, +jitter_factor]) + +Without the jitter, routers that backed off from the same incident retry in +lockstep and re-create the overload they backed off from. + +Retrying is opt-in; the default config allows a single attempt. +""" + +from __future__ import annotations + +import asyncio +import random +from dataclasses import dataclass +from typing import AsyncIterator, List, Optional, Sequence + +from fastapi import HTTPException + +from vllm_router.log import init_logger + +logger = init_logger(__name__) + +RETRYABLE_STATUS_CODES = frozenset({408, 429, 500, 502, 503, 504}) + + +def is_retryable_status(status_code: int) -> bool: + return status_code in RETRYABLE_STATUS_CODES + + +@dataclass(frozen=True) +class RetryConfig: + """Validated retry parameters. + + ``max_attempts`` counts the initial request, so ``1`` disables retrying and + ``5`` allows the initial attempt plus four retries. + """ + + max_attempts: int = 1 + initial_backoff_ms: int = 50 + max_backoff_ms: int = 30_000 + backoff_multiplier: float = 1.5 + jitter_factor: float = 0.2 + + def __post_init__(self) -> None: + if self.max_attempts < 1: + raise ValueError("Retry max attempts must be at least 1.") + if self.initial_backoff_ms <= 0: + raise ValueError("Retry initial backoff must be greater than 0.") + if self.max_backoff_ms < self.initial_backoff_ms: + raise ValueError( + "Retry max backoff must be greater than or equal to initial backoff." + ) + if self.backoff_multiplier < 1.0: + raise ValueError("Retry backoff multiplier must be at least 1.0.") + if not 0.0 <= self.jitter_factor <= 1.0: + raise ValueError("Retry jitter factor must be between 0.0 and 1.0.") + + @property + def enabled(self) -> bool: + return self.max_attempts > 1 + + def backoff_seconds(self, retry_index: int) -> float: + """Jittered delay before retry ``retry_index`` (0-based).""" + delay_ms = min( + self.initial_backoff_ms * self.backoff_multiplier**retry_index, + self.max_backoff_ms, + ) + if self.jitter_factor: + spread = random.uniform(-self.jitter_factor, self.jitter_factor) + delay_ms = max(0.0, delay_ms * (1 + spread)) + return delay_ms / 1000.0 + + @classmethod + def from_args(cls, args) -> "RetryConfig": + """Build a config from parsed CLI arguments. + + Without ``--enable-retries`` the other flags are ignored, so a shared + config template carrying them cannot change behaviour or block startup. + """ + if not getattr(args, "enable_retries", False): + return cls() + return cls( + max_attempts=args.max_retries, + initial_backoff_ms=args.initial_backoff_ms, + max_backoff_ms=args.max_backoff_ms, + backoff_multiplier=args.backoff_multiplier, + jitter_factor=args.jitter_factor, + ) + + +RETRIES_DISABLED = RetryConfig() + + +class TransientBackendStatus(HTTPException): + """A retryable backend status, as a raisable error. + + Subclasses ``HTTPException`` so that on the paths where the loop ends without + a response, the client sees the backend's status rather than a router 500. + """ + + def __init__(self, server_url: str, status_code: int) -> None: + super().__init__( + status_code=status_code, + detail=f"Backend {server_url} returned {status_code}", + ) + self.server_url = server_url + + +class RetryState: + """Per-request bookkeeping across attempts. + + Callers drive it with :meth:`attempts` and report each outcome. + """ + + def __init__(self, config: RetryConfig, request_id: str) -> None: + self._config = config + self._request_id = request_id + self._excluded: set = set() + self._attempt = 0 + self._retries_used = 0 + self._backoff_pending = False + self.last_error: Optional[BaseException] = None + + @property + def attempt_number(self) -> int: + """1-based, for logging.""" + return self._attempt + 1 + + @property + def max_attempts(self) -> int: + return self._config.max_attempts + + async def attempts(self, endpoints: Sequence) -> AsyncIterator[List]: + """Yield the endpoints eligible for each attempt, sleeping for the + backoff where one is due, until the budget is spent.""" + while self._attempt < self._config.max_attempts: + if self._attempt == 0: + candidates = list(endpoints) + else: + candidates = self._eligible(endpoints) + if not candidates: + if not self._config.enabled: + return + # Every engine is excluded, but the failures may be + # transient, so retry the pool rather than give up. + self._excluded.clear() + candidates = list(endpoints) + self._backoff_pending = True + if self._backoff_pending: + await self._backoff() + yield candidates + self._attempt += 1 + + def should_retry_status(self, status_code: int) -> bool: + return ( + self._config.enabled + and is_retryable_status(status_code) + and self._attempt + 1 < self._config.max_attempts + ) + + def record_transient_status(self, server_url: str, status_code: int) -> None: + """Note a retryable status; the engine stays eligible.""" + self.last_error = TransientBackendStatus(server_url, status_code) + self._backoff_pending = True + logger.warning( + f"Request {self._request_id} got retryable status {status_code} from " + f"{server_url}, retrying (attempt {self.attempt_number}/" + f"{self.max_attempts})" + ) + + def record_transport_failure(self, server_url: str, error: BaseException) -> None: + """Note an unusable engine; it is excluded from further attempts.""" + self._excluded.add(server_url) + self.last_error = error + logger.warning( + f"Request {self._request_id} failed on {server_url} " + f"(attempt {self.attempt_number}/{self.max_attempts}): {error}" + ) + + def record_response(self) -> None: + self.last_error = None + + def _eligible(self, endpoints: Sequence) -> List: + return [ep for ep in endpoints if ep.url not in self._excluded] + + async def _backoff(self) -> None: + delay = self._config.backoff_seconds(self._retries_used) + self._retries_used += 1 + self._backoff_pending = False + if delay <= 0: + return + logger.info( + f"Request {self._request_id} waiting {delay:.3f}s before retry " + f"{self._retries_used}/{self.max_attempts - 1}" + ) + await asyncio.sleep(delay)