This commit is contained in:
Onur C. Cakmak 2026-10-03 16:24:38 -04:00 • committed by GitHub
commit 34b8de278a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 85 additions and 23 deletions

View file

@ -2,14 +2,17 @@
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"
@ -250,35 +253,42 @@ 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:
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)
for key in ("litellm_session_id", "session_id"):
value = params.get(key)
if value:
return str(value)
metadata: Final = params.get("metadata")
if isinstance(metadata, dict):
value = metadata.get("session_id")
if value:
return str(value)
generated: Final = any(
True
for session_metadata in (metadata, params.get("litellm_metadata"))
if isinstance(session_metadata, dict) and session_metadata.get(SESSION_ID_GENERATED_METADATA_KEY)
)
if not generated:
for key in ("litellm_session_id", "session_id"):
value = params.get(key)
if value:
return str(value)
if isinstance(metadata, dict):
value = metadata.get("session_id")
if value:
return str(value)
prompt_cache_key: Final[object] = params.get("prompt_cache_key")
if prompt_cache_key:
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("utf-8", "surrogatepass")).hexdigest()
if generated:
return None
for key in ("litellm_trace_id", "litellm_call_id"):
value = params.get(key)
if value:
@ -286,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

@ -4,6 +4,7 @@ Tests for ChatGPT subscription Responses API transformation
Source: litellm/llms/chatgpt/responses/transformation.py
"""
import hashlib
import json
from collections.abc import Generator
from typing import Final
@ -13,6 +14,10 @@ import httpx
import pytest
import litellm
from litellm.llms.chatgpt.common_utils import (
ensure_chatgpt_session_id,
get_chatgpt_session_id,
)
from litellm.llms.chatgpt.responses.transformation import ChatGPTResponsesAPIConfig
from litellm.llms.openai.common_utils import OpenAIError
from litellm.main import responses_api_bridge_check
@ -386,3 +391,50 @@ class TestChatGPTResponsesAPITransformation:
assert "ChatGPT upstream failed" in str(exc_info.value)
assert exc_info.value.status_code == 502
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():
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():
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_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({})