diff --git a/litellm/llms/chatgpt/common_utils.py b/litellm/llms/chatgpt/common_utils.py index 3946ce53975..661de929552 100644 --- a/litellm/llms/chatgpt/common_utils.py +++ b/litellm/llms/chatgpt/common_utils.py @@ -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()) diff --git a/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py b/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py index ed64e4b73c0..be0678e3571 100644 --- a/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py +++ b/tests/unit/llms/chatgpt/responses/test_chatgpt_responses_transformation.py @@ -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({})