fix(chatgpt): hash surrogate cache keys, type session params, flatten tests

This commit is contained in:
Onur Cakmak 2026-10-02 15:08:40 -04:00
parent 84ed22e102
commit 6615bd17f8
2 changed files with 50 additions and 44 deletions

View file

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

View file

@ -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({})