mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(proxy): make the invalid-model 403 path cheap under a burst of rejections (#39892)
* fix(proxy): make the invalid-model 403 path cheap under a burst of rejections Keep the wildcard pattern registry in specificity order at registration time so route() no longer re-sorts every pattern per lookup, and reuse the standardized failure payload across the async and threaded sync failure handlers regardless of what a callback did to log_event_type. Rejections are still logged and observable. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(router): wrap the filtered pattern tuple the way ruff format wants Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(router,logging): assert registry order and callback awaits instead of patching a class Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(router): inject the pattern sorter so the lookup test observes that route() never sorts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
3c0900b7c5
commit
a670a4621e
4 changed files with 70 additions and 10 deletions
|
|
@ -3234,8 +3234,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.model_call_details = {}
|
||||
|
||||
if (
|
||||
self.model_call_details.get("log_event_type") == "failed_api_call"
|
||||
and self.model_call_details.get("exception") is exception
|
||||
self.model_call_details.get("exception") is exception
|
||||
and self.model_call_details.get("standard_logging_object") is not None
|
||||
):
|
||||
return start_time, self.model_call_details["end_time"]
|
||||
|
|
|
|||
|
|
@ -56,8 +56,9 @@ class PatternMatchRouter:
|
|||
This class will store a mapping for regex pattern: List[Deployments]
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, pattern_utils: type[PatternUtils] = PatternUtils):
|
||||
self.patterns: dict[str, list] = {}
|
||||
self._pattern_utils: Final = pattern_utils
|
||||
|
||||
def add_pattern(self, pattern: str, llm_deployment: dict):
|
||||
"""
|
||||
|
|
@ -69,9 +70,10 @@ class PatternMatchRouter:
|
|||
"""
|
||||
# Convert the pattern to a regex
|
||||
regex: Final = self._pattern_to_regex(pattern)
|
||||
if regex not in self.patterns:
|
||||
self.patterns[regex] = []
|
||||
self.patterns[regex].append(llm_deployment)
|
||||
if regex in self.patterns:
|
||||
self.patterns[regex].append(llm_deployment)
|
||||
return
|
||||
self.patterns = dict(self._pattern_utils.sorted_patterns({**self.patterns, regex: [llm_deployment]}))
|
||||
|
||||
def remove_deployment(self, model_id: str) -> None:
|
||||
"""
|
||||
|
|
@ -138,11 +140,12 @@ class PatternMatchRouter:
|
|||
if request is None:
|
||||
return None
|
||||
|
||||
sorted_patterns: Final = PatternUtils.sorted_patterns(self.patterns)
|
||||
regex_filtered_model_names: Final = (
|
||||
[self._pattern_to_regex(m) for m in filtered_model_names] if filtered_model_names is not None else []
|
||||
tuple(self._pattern_to_regex(m) for m in filtered_model_names)
|
||||
if filtered_model_names is not None
|
||||
else ()
|
||||
)
|
||||
for pattern, llm_deployments in sorted_patterns:
|
||||
for pattern, llm_deployments in self.patterns.items():
|
||||
if filtered_model_names is not None and pattern not in regex_filtered_model_names:
|
||||
continue
|
||||
pattern_match = re.match(pattern, request)
|
||||
|
|
|
|||
|
|
@ -6077,6 +6077,34 @@ def test_failure_handler_helper_fn_builds_payload_once_per_exception():
|
|||
assert obj.model_call_details["standard_logging_object"] is not first_payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_failure_handler_reuses_payload_after_callable_async_callback():
|
||||
"""Regression for LIT-6886: the proxy runs async_failure_handler, then the threaded
|
||||
failure_handler, for every rejected request. A plain-function async callback (the
|
||||
Router registers one) is dispatched through CustomLogger.async_log_event, which
|
||||
restamps log_event_type on the shared model_call_details; the sync handler then
|
||||
rebuilt the standardized payload, doubling the redaction and payload cost of a 403."""
|
||||
router_style_callback = AsyncMock()
|
||||
obj = LitellmLogging(
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hey"}],
|
||||
stream=False,
|
||||
call_type="acompletion",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="lit-6886-1",
|
||||
function_id="f",
|
||||
dynamic_async_failure_callbacks=[router_style_callback],
|
||||
)
|
||||
exc = _raise_and_catch(_ClientError(status_code=403, message="key not allowed to access model"))
|
||||
await obj.async_failure_handler(exception=exc, traceback_exception="")
|
||||
first_payload = obj.model_call_details["standard_logging_object"]
|
||||
assert first_payload is not None
|
||||
assert router_style_callback.await_count == 1
|
||||
|
||||
obj.failure_handler(exc, "")
|
||||
assert obj.model_call_details["standard_logging_object"] is first_payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_hook_injection_marker_recorded_for_every_surface(logging_obj):
|
||||
"""The savings gate reads litellm_gateway_injected_cache from the request's
|
||||
|
|
|
|||
|
|
@ -2,8 +2,10 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import Mock
|
||||
|
||||
from litellm.router_utils import pattern_match_deployments
|
||||
from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
|
||||
from litellm.router_utils.pattern_match_deployments import PatternMatchRouter, PatternUtils
|
||||
|
||||
|
||||
def _wildcard_deployment(model_name: str) -> dict:
|
||||
|
|
@ -76,3 +78,31 @@ def test_get_pattern_still_resolves_unqualified_names(monkeypatch):
|
|||
router = PatternMatchRouter()
|
||||
router.add_pattern("openai/*", _wildcard_deployment("openai/*"))
|
||||
assert _matched_models(router.get_pattern("gpt-4o")) == ["openai/gpt-4o"]
|
||||
|
||||
|
||||
class _CountingPatternUtils(PatternUtils):
|
||||
sorted_patterns = staticmethod(Mock(wraps=PatternUtils.sorted_patterns))
|
||||
|
||||
|
||||
def test_route_never_sorts_and_the_most_specific_pattern_still_wins_after_registry_changes():
|
||||
"""Regression for LIT-6886: the auth layer walks the wildcard registry for every request, so an
|
||||
unmatched model name (an invalid-model 403) re-sorted every pattern by specificity per request and
|
||||
a burst of rejections saturated the worker CPU. Lookups must not sort; adding a pattern or removing
|
||||
a deployment must still leave the most specific pattern winning."""
|
||||
router = PatternMatchRouter(pattern_utils=_CountingPatternUtils)
|
||||
router.add_pattern("openai/*", _wildcard_deployment("openai/*"))
|
||||
router.add_pattern("anthropic/*", _wildcard_deployment("anthropic/*"))
|
||||
router.add_pattern("openai/gpt-*", {"model_name": "openai/gpt-*", "litellm_params": {"model": "azure/gpt-*"}})
|
||||
sorts_after_setup = _CountingPatternUtils.sorted_patterns.call_count
|
||||
|
||||
for _ in range(3):
|
||||
assert router.route("does-not-exist") is None
|
||||
assert _matched_models(router.route("openai/gpt-4o")) == ["azure/gpt-4o"]
|
||||
assert _matched_models(router.route("openai/o3")) == ["openai/o3"]
|
||||
assert _CountingPatternUtils.sorted_patterns.call_count == sorts_after_setup
|
||||
|
||||
router.add_pattern("openai/*", {**_wildcard_deployment("openai/*"), "model_info": {"id": "id-1"}})
|
||||
assert len(_matched_models(router.route("openai/o3"))) == 2
|
||||
router.remove_deployment("id-1")
|
||||
assert _matched_models(router.route("openai/gpt-4o")) == ["azure/gpt-4o"]
|
||||
assert _matched_models(router.route("openai/o3")) == ["openai/o3"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue