mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(chatgpt): address review - hash unsafe cache keys, move tests, drop comments
- hash prompt_cache_keys containing header-unsafe chars instead of mapping them through _safe_header_value, whose replacement char could collapse distinct keys onto one cache shard and serialize unrelated requests (greptile review) - strip the explanatory comment blocks per AGENTS.md comment policy - move the tests into the existing mapped responses transformation test file per AGENTS.md test-placement rule
This commit is contained in:
parent
0ceb56820c
commit
ead8edc422
3 changed files with 54 additions and 100 deletions
|
|
@ -2,6 +2,7 @@
|
|||
Constants and helpers for ChatGPT subscription OAuth.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
import platform
|
||||
from typing import Any, Final
|
||||
|
|
@ -272,14 +273,6 @@ def _normalize_litellm_params(litellm_params: Any | None) -> dict:
|
|||
def get_chatgpt_session_id(litellm_params: object) -> str | None:
|
||||
params: Final = _normalize_litellm_params(litellm_params)
|
||||
metadata: Final = params.get("metadata")
|
||||
# A session id the proxy generated for a request that had none
|
||||
# (general_settings.missing_session_id: "generate") is per-request; using
|
||||
# it as identity pins every request to a different ChatGPT cache shard
|
||||
# and the prompt cache never hits (same guard as fireworks'
|
||||
# get_fireworks_session_id). The marker can sit in "metadata" or
|
||||
# "litellm_metadata" -- the LITELLM_METADATA_ROUTES (responses included)
|
||||
# carry internal metadata under the latter. Generated ids are skipped,
|
||||
# not returned, so a caller-supplied stable prompt_cache_key still wins.
|
||||
generated: Final = any(
|
||||
isinstance(params.get(name), dict) and params[name].get(SESSION_ID_GENERATED_METADATA_KEY)
|
||||
for name in ("metadata", "litellm_metadata")
|
||||
|
|
@ -293,16 +286,12 @@ def get_chatgpt_session_id(litellm_params: object) -> str | None:
|
|||
value = metadata.get("session_id")
|
||||
if value:
|
||||
return str(value)
|
||||
# ChatGPT derives prompt-cache affinity from the Responses session-id
|
||||
# header (the Codex CLI sends a stable per-conversation key derived from
|
||||
# prompt_cache_key). Callers of the responses API already send a stable
|
||||
# prompt_cache_key; retaining it as the session id lets repeat turns of a
|
||||
# conversation reuse the provider's prompt cache instead of landing on a
|
||||
# fresh shard every request. Explicit operator configuration above still
|
||||
# wins; litellm-internal per-request ids below still lose to it.
|
||||
prompt_cache_key: Final = params.get("prompt_cache_key")
|
||||
if prompt_cache_key:
|
||||
return _safe_header_value(str(prompt_cache_key)) or 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()
|
||||
if generated:
|
||||
return None
|
||||
for key in ("litellm_trace_id", "litellm_call_id"):
|
||||
|
|
|
|||
|
|
@ -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 unittest.mock import MagicMock, patch
|
||||
|
|
@ -12,6 +13,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
|
||||
|
|
@ -55,7 +60,6 @@ class TestChatGPTResponsesAPITransformation:
|
|||
assert isinstance(config, ChatGPTResponsesAPIConfig)
|
||||
assert config.custom_llm_provider == LlmProviders.CHATGPT
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[
|
||||
|
|
@ -92,14 +96,10 @@ class TestChatGPTResponsesAPITransformation:
|
|||
url = config.get_complete_url(api_base=None, litellm_params={})
|
||||
assert url == "https://chatgpt.example.com/responses"
|
||||
|
||||
custom_url = config.get_complete_url(
|
||||
api_base="https://custom.chatgpt.com", litellm_params={}
|
||||
)
|
||||
custom_url = config.get_complete_url(api_base="https://custom.chatgpt.com", litellm_params={})
|
||||
assert custom_url == "https://custom.chatgpt.com/responses"
|
||||
|
||||
url_with_slash = config.get_complete_url(
|
||||
api_base="https://chatgpt.example.com/", litellm_params={}
|
||||
)
|
||||
url_with_slash = config.get_complete_url(api_base="https://chatgpt.example.com/", litellm_params={})
|
||||
assert url_with_slash == "https://chatgpt.example.com/responses"
|
||||
|
||||
@patch("litellm.llms.chatgpt.responses.transformation.Authenticator")
|
||||
|
|
@ -162,9 +162,7 @@ class TestChatGPTResponsesAPITransformation:
|
|||
"user": "user_123",
|
||||
"temperature": 0.2,
|
||||
"top_p": 0.9,
|
||||
"context_management": [
|
||||
{"type": "compaction", "compact_threshold": 200000}
|
||||
],
|
||||
"context_management": [{"type": "compaction", "compact_threshold": 200000}],
|
||||
"metadata": {"foo": "bar"},
|
||||
"max_output_tokens": 123,
|
||||
"stream_options": {"include_usage": True},
|
||||
|
|
@ -203,9 +201,7 @@ class TestChatGPTResponsesAPITransformation:
|
|||
("chatgpt/gpt-5.3-codex", "gpt-5.3-codex"),
|
||||
],
|
||||
)
|
||||
def test_chatgpt_non_stream_sse_response_parsing(
|
||||
self, model_name: str, response_model: str
|
||||
):
|
||||
def test_chatgpt_non_stream_sse_response_parsing(self, model_name: str, response_model: str):
|
||||
config = ChatGPTResponsesAPIConfig()
|
||||
response_payload = {
|
||||
"id": "resp_test",
|
||||
|
|
@ -228,9 +224,7 @@ class TestChatGPTResponsesAPITransformation:
|
|||
"",
|
||||
]
|
||||
)
|
||||
raw_response = httpx.Response(
|
||||
200, headers={"content-type": "text/event-stream"}, text=sse_body
|
||||
)
|
||||
raw_response = httpx.Response(200, headers={"content-type": "text/event-stream"}, text=sse_body)
|
||||
logging_obj = MagicMock()
|
||||
|
||||
parsed = config.transform_response_api_response(
|
||||
|
|
@ -248,9 +242,7 @@ class TestChatGPTResponsesAPITransformation:
|
|||
("chatgpt/gpt-5.3-codex", "gpt-5.3-codex"),
|
||||
],
|
||||
)
|
||||
def test_chatgpt_non_stream_sse_response_recovers_output_items(
|
||||
self, model_name: str, response_model: str
|
||||
):
|
||||
def test_chatgpt_non_stream_sse_response_recovers_output_items(self, model_name: str, response_model: str):
|
||||
config = ChatGPTResponsesAPIConfig()
|
||||
response_payload = {
|
||||
"id": "resp_test",
|
||||
|
|
@ -273,9 +265,7 @@ class TestChatGPTResponsesAPITransformation:
|
|||
"",
|
||||
]
|
||||
)
|
||||
raw_response = httpx.Response(
|
||||
200, headers={"content-type": "text/event-stream"}, text=sse_body
|
||||
)
|
||||
raw_response = httpx.Response(200, headers={"content-type": "text/event-stream"}, text=sse_body)
|
||||
logging_obj = MagicMock()
|
||||
|
||||
parsed = config.transform_response_api_response(
|
||||
|
|
@ -315,9 +305,7 @@ class TestChatGPTResponsesAPITransformation:
|
|||
"",
|
||||
]
|
||||
)
|
||||
raw_response = httpx.Response(
|
||||
200, headers={"content-type": "text/event-stream"}, text=sse_body
|
||||
)
|
||||
raw_response = httpx.Response(200, headers={"content-type": "text/event-stream"}, text=sse_body)
|
||||
logging_obj = MagicMock()
|
||||
|
||||
parsed = config.transform_response_api_response(
|
||||
|
|
@ -350,9 +338,7 @@ class TestChatGPTResponsesAPITransformation:
|
|||
"",
|
||||
]
|
||||
)
|
||||
raw_response = httpx.Response(
|
||||
502, headers={"content-type": "text/event-stream"}, text=sse_body
|
||||
)
|
||||
raw_response = httpx.Response(502, headers={"content-type": "text/event-stream"}, text=sse_body)
|
||||
logging_obj = MagicMock()
|
||||
|
||||
with pytest.raises(OpenAIError) as exc_info:
|
||||
|
|
@ -364,3 +350,38 @@ class TestChatGPTResponsesAPITransformation:
|
|||
|
||||
assert "ChatGPT upstream failed" in str(exc_info.value)
|
||||
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_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_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_uuid4_fallback_without_any_key(self):
|
||||
assert ensure_chatgpt_session_id({}) != ensure_chatgpt_session_id({})
|
||||
|
|
|
|||
|
|
@ -1,56 +0,0 @@
|
|||
"""Session-id resolution for the ChatGPT provider.
|
||||
|
||||
ChatGPT partitions the Codex prompt cache by the Responses session-id
|
||||
header. These tests pin the resolution order:
|
||||
explicit operator config > prompt_cache_key > litellm-internal ids,
|
||||
with proxy-generated session ids skipped entirely.
|
||||
"""
|
||||
|
||||
from litellm.llms.chatgpt.common_utils import (
|
||||
ensure_chatgpt_session_id,
|
||||
get_chatgpt_session_id,
|
||||
)
|
||||
|
||||
|
||||
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():
|
||||
# litellm_trace_id / litellm_call_id are per-request; they must not
|
||||
# shadow a stable per-conversation key.
|
||||
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_is_sanitized_for_header_use():
|
||||
# Caller-controlled value headed for an HTTP header: control chars are
|
||||
# replaced (h11 field-value validation would otherwise 500 the request),
|
||||
# keeping the key stable.
|
||||
assert get_chatgpt_session_id({"prompt_cache_key": "a\nb"}) == "a_b"
|
||||
|
||||
|
||||
def test_generated_session_ids_are_skipped():
|
||||
# missing_session_id: "generate" stamps a per-request id in every channel;
|
||||
# it must not be used as identity, and the cache key must still win.
|
||||
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
|
||||
a, b = ensure_chatgpt_session_id(generated), ensure_chatgpt_session_id(generated)
|
||||
assert a != "gen" and b != "gen" and a != b
|
||||
|
||||
|
||||
def test_uuid4_fallback_without_any_key():
|
||||
a, b = ensure_chatgpt_session_id({}), ensure_chatgpt_session_id({})
|
||||
assert a and b and a != b
|
||||
Loading…
Add table
Reference in a new issue