mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
docs(proxy): harden duplicate burst guard fingerprint
This commit is contained in:
parent
a9d8b5f664
commit
30a6f1ae22
2 changed files with 209 additions and 45 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue