diff --git a/src/tests/test_request_stats.py b/src/tests/test_request_stats.py new file mode 100644 index 000000000..12993afef --- /dev/null +++ b/src/tests/test_request_stats.py @@ -0,0 +1,24 @@ +import pytest + +from vllm_router.stats.request_stats import RequestStatsMonitor, SingletonMeta + + +@pytest.fixture(autouse=True) +def reset_request_stats_monitor(): + SingletonMeta._instances.pop(RequestStatsMonitor, None) + yield + SingletonMeta._instances.pop(RequestStatsMonitor, None) + + +def test_avg_decoding_length_tracks_decode_duration(): + monitor = RequestStatsMonitor(sliding_window_size=60) + engine_url = "http://engine" + + monitor.on_new_request(engine_url, "request-1", 100.0) + monitor.on_request_response(engine_url, "request-1", 101.0) + assert monitor.get_request_stats(101.0)[engine_url].avg_decoding_length == -1 + + monitor.on_request_complete(engine_url, "request-1", 105.0) + assert monitor.get_request_stats(105.0)[engine_url].avg_decoding_length == 4.0 + assert (engine_url, "request-1") not in monitor.first_token_time + assert monitor.get_request_stats(166.0)[engine_url].avg_decoding_length == -1 diff --git a/src/vllm_router/stats/request_stats.py b/src/vllm_router/stats/request_stats.py index f0409b912..825bc69cd 100644 --- a/src/vllm_router/stats/request_stats.py +++ b/src/vllm_router/stats/request_stats.py @@ -216,6 +216,16 @@ def on_request_complete(self, engine_url: str, request_id: str, timestamp: float ) self.finished_requests[engine_url] += 1 + first_token_time = self.first_token_time.pop((engine_url, request_id), None) + if first_token_time is not None: + if engine_url not in self.decoding_length_monitors: + self.decoding_length_monitors[engine_url] = MovingAverageMonitor( + self.sliding_window_size + ) + self.decoding_length_monitors[engine_url].update( + timestamp, timestamp - first_token_time + ) + if request_start_time := self.request_start_time.get((engine_url, request_id)): self.latency_monitors[engine_url].update( timestamp, time.time() - request_start_time @@ -272,6 +282,7 @@ def get_request_stats(self, current_time: float) -> Dict[str, RequestStats]: finished = self.finished_requests.get(engine_url, 0) if engine_url in self.decoding_length_monitors: + self.decoding_length_monitors[engine_url].update_no_value(current_time) avg_dec_len = self.decoding_length_monitors[engine_url].get_average() else: avg_dec_len = -1