mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
Merge 67798ccd91 into 252c71c0b2
This commit is contained in:
commit
2a1db171b3
2 changed files with 525 additions and 0 deletions
162
cookbook/litellm_proxy_server/duplicate_burst_guard.py
Normal file
162
cookbook/litellm_proxy_server/duplicate_burst_guard.py
Normal file
|
|
@ -0,0 +1,162 @@
|
|||
"""
|
||||
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.
|
||||
|
||||
The request fingerprint uses a whitelist of model-affecting fields and scopes
|
||||
duplicates by authenticated API-key owner plus request user/session. Adjust
|
||||
those fields if your deployment needs different duplicate semantics.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
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
|
||||
|
||||
FINGERPRINT_REQUEST_KEYS = (
|
||||
"model",
|
||||
"system",
|
||||
"messages",
|
||||
"prompt",
|
||||
"input",
|
||||
"query",
|
||||
"documents",
|
||||
"instructions",
|
||||
"previous_response_id",
|
||||
"temperature",
|
||||
"max_tokens",
|
||||
"max_completion_tokens",
|
||||
"top_p",
|
||||
"stop",
|
||||
"stream",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"functions",
|
||||
"function_call",
|
||||
"response_format",
|
||||
"n",
|
||||
"presence_penalty",
|
||||
"frequency_penalty",
|
||||
"logit_bias",
|
||||
"seed",
|
||||
"modalities",
|
||||
"audio",
|
||||
"reasoning_effort",
|
||||
"thinking",
|
||||
"extra_body",
|
||||
"vector_store_id",
|
||||
)
|
||||
SKIP_CALL_TYPES: frozenset[str] = frozenset({"aimage_edit"})
|
||||
|
||||
|
||||
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]:
|
||||
if call_type in SKIP_CALL_TYPES:
|
||||
return data
|
||||
|
||||
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:
|
||||
metadata_value = data.get("metadata")
|
||||
metadata: Mapping[str, object]
|
||||
if isinstance(metadata_value, Mapping):
|
||||
metadata = cast(Mapping[str, object], metadata_value)
|
||||
else:
|
||||
metadata = {}
|
||||
key_user_id: object = None
|
||||
if user_api_key_dict is not None:
|
||||
key_user_id = user_api_key_dict.user_id
|
||||
request_user_id: object = (
|
||||
data.get("user")
|
||||
or metadata.get("user_id")
|
||||
or metadata.get("session_id")
|
||||
or "anonymous_request"
|
||||
)
|
||||
request = {key: data[key] for key in FINGERPRINT_REQUEST_KEYS if key in data}
|
||||
raw = json.dumps(
|
||||
{
|
||||
"api_key_identity": key_user_id or "anonymous_key",
|
||||
"call_type": call_type,
|
||||
"request": request,
|
||||
"request_identity": request_user_id,
|
||||
},
|
||||
default=str,
|
||||
ensure_ascii=False,
|
||||
separators=(",", ":"),
|
||||
sort_keys=True,
|
||||
)
|
||||
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
duplicate_burst_guard = DuplicateBurstGuard()
|
||||
363
tests/test_litellm/test_duplicate_burst_guard_example.py
Normal file
363
tests/test_litellm/test_duplicate_burst_guard_example.py
Normal file
|
|
@ -0,0 +1,363 @@
|
|||
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")
|
||||
OTHER_USER_API_KEY = UserAPIKeyAuth(user_id="user-2")
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_avoids_delimiter_collisions() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a = _chat_request("c")
|
||||
request_a["messages"] = [
|
||||
{"role": "system", "content": "a|b"},
|
||||
{"role": "user", "content": "c"},
|
||||
]
|
||||
request_b = _chat_request("b|c")
|
||||
request_b["messages"] = [
|
||||
{"role": "system", "content": "a"},
|
||||
{"role": "user", "content": "b|c"},
|
||||
]
|
||||
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_multimodal_content() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a: dict[str, object] = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What is in this image?"},
|
||||
{"type": "image_url", "image_url": {"url": "https://a.test/1.png"}},
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
request_b: dict[str, object] = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What is in this image?"},
|
||||
{"type": "image_url", "image_url": {"url": "https://a.test/2.png"}},
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_prefers_request_user_identity() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a = _chat_request("Summarize invoice A")
|
||||
request_a["user"] = "end-user-a"
|
||||
request_b = _chat_request("Summarize invoice A")
|
||||
request_b["user"] = "end-user-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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_full_chat_context() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a = _chat_request("What should I do next?")
|
||||
request_a["messages"] = [
|
||||
{"role": "system", "content": "Use the shared operating policy."},
|
||||
{"role": "user", "content": "Invoice A is overdue."},
|
||||
{"role": "assistant", "content": "Ask for payment status."},
|
||||
{"role": "user", "content": "What should I do next?"},
|
||||
]
|
||||
request_b = _chat_request("What should I do next?")
|
||||
request_b["messages"] = [
|
||||
{"role": "system", "content": "Use the shared operating policy."},
|
||||
{"role": "user", "content": "Invoice B has a credit memo."},
|
||||
{"role": "assistant", "content": "Check whether it was applied."},
|
||||
{"role": "user", "content": "What should I do next?"},
|
||||
]
|
||||
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_output_options() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a = _chat_request("Summarize invoice A")
|
||||
request_a["response_format"] = {"type": "json_object"}
|
||||
request_b = _chat_request("Summarize invoice A")
|
||||
request_b["tools"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "summarize_invoice",
|
||||
"parameters": {"type": "object"},
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_non_chat_input() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a: dict[str, object] = {
|
||||
"model": "text-embedding-3-small",
|
||||
"input": "invoice A",
|
||||
}
|
||||
request_b: dict[str, object] = {
|
||||
"model": "text-embedding-3-small",
|
||||
"input": "invoice B",
|
||||
}
|
||||
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request_a, "aembedding")
|
||||
accepted = await guard.async_pre_call_hook(
|
||||
USER_API_KEY, CACHE, request_b, "aembedding"
|
||||
)
|
||||
|
||||
assert accepted == request_b
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_ignores_ephemeral_call_ids() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a = _chat_request("Summarize invoice A")
|
||||
request_a["litellm_call_id"] = "call-a"
|
||||
request_b = _chat_request("Summarize invoice A")
|
||||
request_b["litellm_call_id"] = "call-b"
|
||||
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request_a, CALL_TYPE)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request_b, CALL_TYPE)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_authenticated_key_identity() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request = _chat_request("Summarize invoice A")
|
||||
request["metadata"] = {"session_id": "shared-session"}
|
||||
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request, CALL_TYPE)
|
||||
accepted = await guard.async_pre_call_hook(
|
||||
OTHER_USER_API_KEY, CACHE, request, CALL_TYPE
|
||||
)
|
||||
|
||||
assert accepted == request
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_rerank_fields() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a: dict[str, object] = {
|
||||
"model": "rerank-english-v3.0",
|
||||
"query": "invoice A",
|
||||
"documents": ["invoice A is overdue"],
|
||||
}
|
||||
request_b: dict[str, object] = {
|
||||
"model": "rerank-english-v3.0",
|
||||
"query": "invoice B",
|
||||
"documents": ["invoice B has a credit memo"],
|
||||
}
|
||||
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request_a, "rerank")
|
||||
accepted = await guard.async_pre_call_hook(USER_API_KEY, CACHE, request_b, "rerank")
|
||||
|
||||
assert accepted == request_b
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_responses_context() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a: dict[str, object] = {
|
||||
"model": "gpt-4o-mini",
|
||||
"input": "continue",
|
||||
"previous_response_id": "resp-a",
|
||||
"instructions": "Use policy A",
|
||||
}
|
||||
request_b: dict[str, object] = {
|
||||
"model": "gpt-4o-mini",
|
||||
"input": "continue",
|
||||
"previous_response_id": "resp-b",
|
||||
"instructions": "Use policy B",
|
||||
}
|
||||
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request_a, "aresponses")
|
||||
accepted = await guard.async_pre_call_hook(
|
||||
USER_API_KEY, CACHE, request_b, "aresponses"
|
||||
)
|
||||
|
||||
assert accepted == request_b
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_top_level_system_prompt() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a: dict[str, object] = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": "Use policy A",
|
||||
"messages": [{"role": "user", "content": "Summarize invoice A"}],
|
||||
}
|
||||
request_b: dict[str, object] = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": "Use policy B",
|
||||
"messages": [{"role": "user", "content": "Summarize invoice A"}],
|
||||
}
|
||||
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_includes_search_fields() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request_a: dict[str, object] = {
|
||||
"model": "search-model",
|
||||
"query": "invoice A",
|
||||
"vector_store_id": "store-a",
|
||||
}
|
||||
request_b: dict[str, object] = {
|
||||
"model": "search-model",
|
||||
"query": "invoice B",
|
||||
"vector_store_id": "store-b",
|
||||
}
|
||||
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request_a, "asearch")
|
||||
accepted = await guard.async_pre_call_hook(
|
||||
USER_API_KEY, CACHE, request_b, "asearch"
|
||||
)
|
||||
|
||||
assert accepted == request_b
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_burst_guard_skips_image_edits() -> None:
|
||||
guard = DuplicateBurstGuard(max_calls=1, window_seconds=60)
|
||||
request: dict[str, object] = {
|
||||
"model": "gpt-image-1",
|
||||
"prompt": "Add the logo",
|
||||
"image": object(),
|
||||
"mask": object(),
|
||||
}
|
||||
|
||||
assert (
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request, "aimage_edit")
|
||||
== request
|
||||
)
|
||||
assert (
|
||||
await guard.async_pre_call_hook(USER_API_KEY, CACHE, request, "aimage_edit")
|
||||
== request
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue