mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
docs(proxy): add duplicate burst guard cookbook example
This commit is contained in:
parent
b443037783
commit
a9d8b5f664
2 changed files with 228 additions and 0 deletions
148
cookbook/litellm_proxy_server/duplicate_burst_guard.py
Normal file
148
cookbook/litellm_proxy_server/duplicate_burst_guard.py
Normal file
|
|
@ -0,0 +1,148 @@
|
|||
"""
|
||||
Example LiteLLM Proxy callback that blocks short duplicate request bursts.
|
||||
|
||||
Register this handler in proxy_config.yaml:
|
||||
|
||||
litellm_settings:
|
||||
callbacks:
|
||||
- duplicate_burst_guard.duplicate_burst_guard
|
||||
|
||||
This is an in-memory, single-process example. Use a shared store if your proxy
|
||||
runs multiple workers or needs duplicate detection across instances.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import time
|
||||
from collections import defaultdict, deque
|
||||
from collections.abc import Mapping
|
||||
from typing import Optional, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
|
||||
|
||||
class DuplicateBurstGuard(CustomLogger):
|
||||
def __init__(self, max_calls: int = 2, window_seconds: float = 10.0) -> None:
|
||||
self.max_calls = max_calls
|
||||
self.window_seconds = window_seconds
|
||||
self._calls: defaultdict[str, deque[float]] = defaultdict(
|
||||
deque
|
||||
) # mutable-ok: sliding window state; use Redis ZSET for multi-worker
|
||||
self._last_gc = 0.0 # mutable-ok: local GC watermark for in-process state
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict[str, object],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> dict[str, object]:
|
||||
fingerprint = self._fingerprint(
|
||||
data=data, user_api_key_dict=user_api_key_dict, call_type=call_type
|
||||
)
|
||||
count = await self._record_and_count(fingerprint)
|
||||
if count > self.max_calls:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail={
|
||||
"error": "Duplicate request burst detected",
|
||||
"fingerprint": fingerprint,
|
||||
},
|
||||
)
|
||||
return data
|
||||
|
||||
async def _record_and_count(self, fingerprint: str) -> int:
|
||||
now = time.monotonic()
|
||||
async with self._lock:
|
||||
if now - self._last_gc > self.window_seconds:
|
||||
self._prune(now)
|
||||
self._last_gc = now
|
||||
|
||||
timestamps = self._calls[fingerprint]
|
||||
timestamps.append(now)
|
||||
self._prune_key(timestamps, now)
|
||||
return len(timestamps)
|
||||
|
||||
def _prune(self, now: float) -> None:
|
||||
for fingerprint, timestamps in list(self._calls.items()):
|
||||
self._prune_key(timestamps, now)
|
||||
if not timestamps:
|
||||
self._calls.pop(fingerprint, None)
|
||||
|
||||
def _prune_key(self, timestamps: deque[float], now: float) -> None:
|
||||
while timestamps and now - timestamps[0] > self.window_seconds:
|
||||
timestamps.popleft()
|
||||
|
||||
def _fingerprint(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
user_api_key_dict: Optional[UserAPIKeyAuth],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> str:
|
||||
messages_value = data.get("messages")
|
||||
system_prompt = ""
|
||||
last_user_prompt = ""
|
||||
prompt = ""
|
||||
|
||||
if isinstance(messages_value, list):
|
||||
messages = cast(list[object], messages_value)
|
||||
for message in messages:
|
||||
if not isinstance(message, Mapping):
|
||||
continue
|
||||
message_data = cast(Mapping[str, object], message)
|
||||
if message_data.get("role") == "system" and not system_prompt:
|
||||
system_prompt = self._content_text(message_data.get("content"))
|
||||
if message_data.get("role") == "user":
|
||||
last_user_prompt = self._content_text(message_data.get("content"))
|
||||
else:
|
||||
prompt = self._content_text(data.get("prompt"))
|
||||
|
||||
metadata_value = data.get("metadata")
|
||||
metadata: Mapping[str, object]
|
||||
if isinstance(metadata_value, Mapping):
|
||||
metadata = cast(Mapping[str, object], metadata_value)
|
||||
else:
|
||||
metadata = {}
|
||||
user_id: object = (
|
||||
(user_api_key_dict.user_id if user_api_key_dict is not None else None)
|
||||
or data.get("user")
|
||||
or metadata.get("user_id")
|
||||
or metadata.get("session_id")
|
||||
or "anonymous"
|
||||
)
|
||||
raw = "|".join(
|
||||
(
|
||||
str(user_id),
|
||||
str(call_type),
|
||||
str(data.get("model", "")),
|
||||
str(data.get("temperature", "")),
|
||||
str(data.get("max_tokens", "")),
|
||||
system_prompt,
|
||||
last_user_prompt,
|
||||
prompt,
|
||||
)
|
||||
)
|
||||
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
|
||||
|
||||
def _content_text(self, content: object) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
content_items = cast(list[object], content)
|
||||
return "\n".join(
|
||||
text
|
||||
for item in content_items
|
||||
if isinstance(item, Mapping)
|
||||
for text in [cast(Mapping[str, object], item).get("text")]
|
||||
if isinstance(text, str)
|
||||
)
|
||||
return ""
|
||||
|
||||
|
||||
duplicate_burst_guard = DuplicateBurstGuard()
|
||||
80
tests/test_litellm/test_duplicate_burst_guard_example.py
Normal file
80
tests/test_litellm/test_duplicate_burst_guard_example.py
Normal file
|
|
@ -0,0 +1,80 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from cookbook.litellm_proxy_server.duplicate_burst_guard import DuplicateBurstGuard
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
|
||||
USER_API_KEY = UserAPIKeyAuth(user_id="user-1")
|
||||
CACHE = DualCache()
|
||||
CALL_TYPE: CallTypesLiteral = "completion"
|
||||
|
||||
|
||||
def _chat_request(user_prompt: str) -> dict[str, object]:
|
||||
return {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [
|
||||
{"role": "system", "content": "Use the shared operating policy."},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
"temperature": 0,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_blocks_repeated_prompt() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request = _chat_request("Summarize invoice A")
|
||||
|
||||
assert (
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request, CALL_TYPE)
|
||||
== request
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request, CALL_TYPE)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
detail = exc_info.value.detail
|
||||
assert isinstance(detail, Mapping)
|
||||
detail_data = cast(Mapping[str, object], detail)
|
||||
assert detail_data.get("error") == "Duplicate request burst detected"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_uses_last_user_prompt() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_b = _chat_request("Summarize invoice B")
|
||||
|
||||
await guard.async_pre_call_hook(
|
||||
USER_API_KEY, CACHE, _chat_request("Summarize invoice A"), CALL_TYPE
|
||||
)
|
||||
accepted = await guard.async_pre_call_hook(
|
||||
USER_API_KEY, CACHE, request_b, CALL_TYPE
|
||||
)
|
||||
|
||||
assert accepted == request_b
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_uses_completion_prompt() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a: dict[str, object] = {
|
||||
"model": "gpt-4o-mini",
|
||||
"prompt": "Summarize invoice A",
|
||||
}
|
||||
request_b: dict[str, object] = {
|
||||
"model": "gpt-4o-mini",
|
||||
"prompt": "Summarize invoice B",
|
||||
}
|
||||
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request_a, CALL_TYPE)
|
||||
accepted = await guard.async_pre_call_hook(
|
||||
USER_API_KEY, CACHE, request_b, CALL_TYPE
|
||||
)
|
||||
|
||||
assert accepted == request_b
|
||||
Loading…
Add table
Reference in a new issue