This commit is contained in:
Luke Fox 2026-09-22 14:42:18 +05:00 • committed by GitHub
commit 2a1db171b3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 525 additions and 0 deletions

View 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()

View 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
)