mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(chatgpt): hash surrogate cache keys, type session params, flatten tests
This commit is contained in:
parent
84ed22e102
commit
6615bd17f8
2 changed files with 50 additions and 44 deletions
|
|
@ -5,13 +5,14 @@ Constants and helpers for ChatGPT subscription OAuth.
|
|||
import hashlib
|
||||
import os
|
||||
import platform
|
||||
from typing import Any, Final
|
||||
from typing import Final
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.constants import SESSION_ID_GENERATED_METADATA_KEY
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
# OAuth + API constants (derived from openai/codex)
|
||||
CHATGPT_AUTH_BASE: Final = "https://auth.openai.com"
|
||||
|
|
@ -252,25 +253,18 @@ def get_chatgpt_default_instructions() -> str:
|
|||
return os.getenv("CHATGPT_DEFAULT_INSTRUCTIONS") or CHATGPT_DEFAULT_INSTRUCTIONS
|
||||
|
||||
|
||||
def _normalize_litellm_params(litellm_params: Any | None) -> dict[str, object]:
|
||||
def _normalize_litellm_params(litellm_params: dict[str, object] | GenericLiteLLMParams | None) -> dict[str, object]:
|
||||
if litellm_params is None:
|
||||
return {}
|
||||
if isinstance(litellm_params, dict):
|
||||
return litellm_params
|
||||
if hasattr(litellm_params, "model_dump"):
|
||||
try:
|
||||
return litellm_params.model_dump()
|
||||
except Exception:
|
||||
return {}
|
||||
if hasattr(litellm_params, "dict"):
|
||||
try:
|
||||
return litellm_params.dict()
|
||||
except Exception:
|
||||
return {}
|
||||
return {}
|
||||
try:
|
||||
return litellm_params.model_dump()
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def get_chatgpt_session_id(litellm_params: object) -> str | None:
|
||||
def get_chatgpt_session_id(litellm_params: dict[str, object] | GenericLiteLLMParams | None) -> str | None:
|
||||
params: Final = _normalize_litellm_params(litellm_params)
|
||||
metadata: Final = params.get("metadata")
|
||||
generated: Final = any(
|
||||
|
|
@ -292,7 +286,7 @@ def get_chatgpt_session_id(litellm_params: object) -> str | None:
|
|||
key = str(prompt_cache_key)
|
||||
safe = _safe_header_value(key)
|
||||
# hashing avoids collisions from _safe_header_value's replacement char
|
||||
return safe if safe == key else hashlib.sha256(key.encode()).hexdigest()
|
||||
return safe if safe == key else hashlib.sha256(key.encode("utf-8", "surrogatepass")).hexdigest()
|
||||
if generated:
|
||||
return None
|
||||
for key in ("litellm_trace_id", "litellm_call_id"):
|
||||
|
|
@ -302,5 +296,5 @@ def get_chatgpt_session_id(litellm_params: object) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def ensure_chatgpt_session_id(litellm_params: object) -> str:
|
||||
def ensure_chatgpt_session_id(litellm_params: dict[str, object] | GenericLiteLLMParams | None) -> str:
|
||||
return get_chatgpt_session_id(litellm_params) or str(uuid4())
|
||||
|
|
|
|||
|
|
@ -371,36 +371,48 @@ class TestChatGPTResponsesAPITransformation:
|
|||
assert exc_info.value.status_code == 502
|
||||
|
||||
|
||||
class TestChatGPTSessionId:
|
||||
def test_explicit_session_ids_win(self):
|
||||
assert get_chatgpt_session_id({"session_id": "s", "prompt_cache_key": "k"}) == "s"
|
||||
assert get_chatgpt_session_id({"litellm_session_id": "ls", "prompt_cache_key": "k"}) == "ls"
|
||||
assert get_chatgpt_session_id({"metadata": {"session_id": "ms"}, "prompt_cache_key": "k"}) == "ms"
|
||||
def test_explicit_session_ids_win():
|
||||
assert get_chatgpt_session_id({"session_id": "s", "prompt_cache_key": "k"}) == "s"
|
||||
assert get_chatgpt_session_id({"litellm_session_id": "ls", "prompt_cache_key": "k"}) == "ls"
|
||||
assert get_chatgpt_session_id({"metadata": {"session_id": "ms"}, "prompt_cache_key": "k"}) == "ms"
|
||||
|
||||
def test_prompt_cache_key_becomes_session_id(self):
|
||||
assert get_chatgpt_session_id({"prompt_cache_key": "conv-abc"}) == "conv-abc"
|
||||
assert ensure_chatgpt_session_id({"prompt_cache_key": "conv-abc"}) == "conv-abc"
|
||||
|
||||
def test_prompt_cache_key_beats_internal_request_ids(self):
|
||||
assert get_chatgpt_session_id({"litellm_call_id": "c", "prompt_cache_key": "k"}) == "k"
|
||||
assert get_chatgpt_session_id({"litellm_trace_id": "t", "prompt_cache_key": "k"}) == "k"
|
||||
def test_prompt_cache_key_becomes_session_id():
|
||||
assert get_chatgpt_session_id({"prompt_cache_key": "conv-abc"}) == "conv-abc"
|
||||
assert ensure_chatgpt_session_id({"prompt_cache_key": "conv-abc"}) == "conv-abc"
|
||||
|
||||
def test_unsafe_cache_keys_are_hashed_not_collapsed(self):
|
||||
mangled_a = get_chatgpt_session_id({"prompt_cache_key": "a\nb"})
|
||||
mangled_b = get_chatgpt_session_id({"prompt_cache_key": "a\tb"})
|
||||
assert mangled_a == hashlib.sha256(b"a\nb").hexdigest()
|
||||
assert mangled_b == hashlib.sha256(b"a\tb").hexdigest()
|
||||
assert mangled_a != mangled_b
|
||||
|
||||
def test_generated_session_ids_are_skipped(self):
|
||||
for metadata_key in ("metadata", "litellm_metadata"):
|
||||
generated = {
|
||||
"litellm_session_id": "gen",
|
||||
metadata_key: {"session_id": "gen", "litellm_session_id_generated": True},
|
||||
}
|
||||
assert get_chatgpt_session_id({**generated, "prompt_cache_key": "k"}) == "k"
|
||||
assert get_chatgpt_session_id(generated) is None
|
||||
assert ensure_chatgpt_session_id(generated) != ensure_chatgpt_session_id(generated) != "gen"
|
||||
def test_prompt_cache_key_beats_internal_request_ids():
|
||||
assert get_chatgpt_session_id({"litellm_call_id": "c", "prompt_cache_key": "k"}) == "k"
|
||||
assert get_chatgpt_session_id({"litellm_trace_id": "t", "prompt_cache_key": "k"}) == "k"
|
||||
|
||||
def test_uuid4_fallback_without_any_key(self):
|
||||
assert ensure_chatgpt_session_id({}) != ensure_chatgpt_session_id({})
|
||||
|
||||
def test_unsafe_cache_keys_are_hashed_not_collapsed():
|
||||
mangled_a = get_chatgpt_session_id({"prompt_cache_key": "a\nb"})
|
||||
mangled_b = get_chatgpt_session_id({"prompt_cache_key": "a\tb"})
|
||||
assert mangled_a == hashlib.sha256(b"a\nb").hexdigest()
|
||||
assert mangled_b == hashlib.sha256(b"a\tb").hexdigest()
|
||||
assert mangled_a != mangled_b
|
||||
|
||||
|
||||
def test_lone_surrogate_cache_keys_hash_to_distinct_session_ids():
|
||||
high = get_chatgpt_session_id({"prompt_cache_key": json.loads('"conv-\\ud800"')})
|
||||
low = get_chatgpt_session_id({"prompt_cache_key": json.loads('"conv-\\udc00"')})
|
||||
assert high == hashlib.sha256(b"conv-\xed\xa0\x80").hexdigest()
|
||||
assert low == hashlib.sha256(b"conv-\xed\xb0\x80").hexdigest()
|
||||
assert high != low
|
||||
|
||||
|
||||
def test_generated_session_ids_are_skipped():
|
||||
for metadata_key in ("metadata", "litellm_metadata"):
|
||||
generated = {
|
||||
"litellm_session_id": "gen",
|
||||
metadata_key: {"session_id": "gen", "litellm_session_id_generated": True},
|
||||
}
|
||||
assert get_chatgpt_session_id({**generated, "prompt_cache_key": "k"}) == "k"
|
||||
assert get_chatgpt_session_id(generated) is None
|
||||
assert ensure_chatgpt_session_id(generated) != ensure_chatgpt_session_id(generated) != "gen"
|
||||
|
||||
|
||||
def test_uuid4_fallback_without_any_key():
|
||||
assert ensure_chatgpt_session_id({}) != ensure_chatgpt_session_id({})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue