mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
5cec43cbb6
commit
92be052bf3
1 changed files with 76 additions and 67 deletions
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue