docs(proxy): harden duplicate burst guard fingerprint

This commit is contained in:
foxjl85-dev 2026-06-29 15:41:29 -05:00
parent a9d8b5f664
commit 30a6f1ae22
2 changed files with 209 additions and 45 deletions

View file

@ -9,10 +9,15 @@ litellm_settings:
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 request user/session before API-key owner. 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
@ -25,6 +30,34 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import CallTypesLiteral
FINGERPRINT_REQUEST_KEYS = (
"model",
"messages",
"prompt",
"input",
"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",
)
class DuplicateBurstGuard(CustomLogger):
def __init__(self, max_calls: int = 2, window_seconds: float = 10.0) -> None:
@ -85,64 +118,35 @@ class DuplicateBurstGuard(CustomLogger):
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 = {}
key_user_id: object = None
if user_api_key_dict is not None:
key_user_id = user_api_key_dict.user_id
user_id: object = (
(user_api_key_dict.user_id if user_api_key_dict is not None else None)
or data.get("user")
data.get("user")
or metadata.get("user_id")
or metadata.get("session_id")
or key_user_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,
)
request = {key: data[key] for key in FINGERPRINT_REQUEST_KEYS if key in data}
raw = json.dumps(
{
"call_type": call_type,
"request": request,
"request_identity": user_id,
},
default=str,
ensure_ascii=False,
separators=(",", ":"),
sort_keys=True,
)
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()

View file

@ -78,3 +78,163 @@ async def test_duplicate_burst_guard_uses_completion_prompt() -> None:
)
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