fix(router): guard latency cache RMW with per-model-group asyncio locks (#24720)

`async_log_success_event` and `async_log_failure_event` in
`LowestLatencyLoggingHandler` both did an unguarded read-modify-write
on the shared `{model_group}_map` cache key. Under concurrent
completions each coroutine reads the same stale snapshot, modifies it
locally, then writes back — the last writer wins and all other updates
are silently dropped.

Result: deployments with lost latency data fall back to `latency: [0]`
(treated as fastest), and are randomly selected — making
`latency-based-routing` behave identically to `simple-shuffle`
weighted by deployment count.

Fix: add a `defaultdict(asyncio.Lock)` keyed by `latency_key` to the
handler. Wrap the entire read-modify-write block in each async method
with `async with self._cache_locks[latency_key]`. Per-key locking
means concurrent updates for *different* model groups are unaffected.

Fixes #24720
This commit is contained in:
Ignazio De Santis 2026-03-31 03:55:18 +08:00
parent 5cec43cbb6
commit 92be052bf3

View file

@ -1,6 +1,8 @@
#### What this does ####
# picks based on response time (for streaming, this is time to first token)
import asyncio
import random
from collections import defaultdict
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
@ -34,6 +36,10 @@ class LowestLatencyLoggingHandler(CustomLogger):
def __init__(self, router_cache: DualCache, routing_args: dict = {}):
self.router_cache = router_cache
self.routing_args = RoutingArgs(**routing_args)
# Per-model-group locks prevent the lost-update race in async_log_success_event
# and async_log_failure_event where concurrent coroutines read a stale snapshot,
# each overwriting the other's updates. See https://github.com/BerriAI/litellm/issues/24720
self._cache_locks: defaultdict[str, asyncio.Lock] = defaultdict(asyncio.Lock)
def log_success_event( # noqa: PLR0915
self, kwargs, response_obj, start_time, end_time
@ -228,29 +234,30 @@ class LowestLatencyLoggingHandler(CustomLogger):
}
"""
latency_key = f"{model_group}_map"
request_count_dict = (
await self.router_cache.async_get_cache(key=latency_key) or {}
)
async with self._cache_locks[latency_key]:
request_count_dict = (
await self.router_cache.async_get_cache(key=latency_key) or {}
)
if id not in request_count_dict:
request_count_dict[id] = {}
if id not in request_count_dict:
request_count_dict[id] = {}
## Latency - give 1000s penalty for failing
if (
len(request_count_dict[id].get("latency", []))
< self.routing_args.max_latency_list_size
):
request_count_dict[id].setdefault("latency", []).append(1000.0)
else:
request_count_dict[id]["latency"] = request_count_dict[id][
"latency"
][: self.routing_args.max_latency_list_size - 1] + [1000.0]
## Latency - give 1000s penalty for failing
if (
len(request_count_dict[id].get("latency", []))
< self.routing_args.max_latency_list_size
):
request_count_dict[id].setdefault("latency", []).append(1000.0)
else:
request_count_dict[id]["latency"] = request_count_dict[id][
"latency"
][: self.routing_args.max_latency_list_size - 1] + [1000.0]
await self.router_cache.async_set_cache(
key=latency_key,
value=request_count_dict,
ttl=self.routing_args.ttl,
) # reset map within window
await self.router_cache.async_set_cache(
key=latency_key,
value=request_count_dict,
ttl=self.routing_args.ttl,
) # reset map within window
else:
# do nothing if it's not a timeout error
return
@ -347,66 +354,68 @@ class LowestLatencyLoggingHandler(CustomLogger):
ttft_seconds, completion_tokens
)
# ------------
# Update usage
# Update usage (guarded by a per-model-group lock to prevent
# lost-update races under concurrent completions)
# ------------
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
request_count_dict = (
await self.router_cache.async_get_cache(
key=latency_key,
parent_otel_span=parent_otel_span,
local_only=True,
async with self._cache_locks[latency_key]:
request_count_dict = (
await self.router_cache.async_get_cache(
key=latency_key,
parent_otel_span=parent_otel_span,
local_only=True,
)
or {}
)
or {}
)
if id not in request_count_dict:
request_count_dict[id] = {}
if id not in request_count_dict:
request_count_dict[id] = {}
## Latency
if (
len(request_count_dict[id].get("latency", []))
< self.routing_args.max_latency_list_size
):
request_count_dict[id].setdefault("latency", []).append(final_value)
else:
request_count_dict[id]["latency"] = request_count_dict[id][
"latency"
][: self.routing_args.max_latency_list_size - 1] + [final_value]
## Time to first token
if time_to_first_token is not None:
## Latency
if (
len(request_count_dict[id].get("time_to_first_token", []))
len(request_count_dict[id].get("latency", []))
< self.routing_args.max_latency_list_size
):
request_count_dict[id].setdefault(
"time_to_first_token", []
).append(time_to_first_token)
request_count_dict[id].setdefault("latency", []).append(final_value)
else:
request_count_dict[id][
"time_to_first_token"
] = request_count_dict[id]["time_to_first_token"][
: self.routing_args.max_latency_list_size - 1
] + [
time_to_first_token
]
request_count_dict[id]["latency"] = request_count_dict[id][
"latency"
][: self.routing_args.max_latency_list_size - 1] + [final_value]
if precise_minute not in request_count_dict[id]:
request_count_dict[id][precise_minute] = {}
## Time to first token
if time_to_first_token is not None:
if (
len(request_count_dict[id].get("time_to_first_token", []))
< self.routing_args.max_latency_list_size
):
request_count_dict[id].setdefault(
"time_to_first_token", []
).append(time_to_first_token)
else:
request_count_dict[id][
"time_to_first_token"
] = request_count_dict[id]["time_to_first_token"][
: self.routing_args.max_latency_list_size - 1
] + [
time_to_first_token
]
## TPM
request_count_dict[id][precise_minute]["tpm"] = (
request_count_dict[id][precise_minute].get("tpm", 0) + total_tokens
)
if precise_minute not in request_count_dict[id]:
request_count_dict[id][precise_minute] = {}
## RPM
request_count_dict[id][precise_minute]["rpm"] = (
request_count_dict[id][precise_minute].get("rpm", 0) + 1
)
## TPM
request_count_dict[id][precise_minute]["tpm"] = (
request_count_dict[id][precise_minute].get("tpm", 0) + total_tokens
)
await self.router_cache.async_set_cache(
key=latency_key, value=request_count_dict, ttl=self.routing_args.ttl
) # reset map within window
## RPM
request_count_dict[id][precise_minute]["rpm"] = (
request_count_dict[id][precise_minute].get("rpm", 0) + 1
)
await self.router_cache.async_set_cache(
key=latency_key, value=request_count_dict, ttl=self.routing_args.ttl
) # reset map within window
### TESTING ###
if self.test_flag: