feat(openai): allow reasoning_effort on reasoning models & preserve cache tokens in usage

This commit is contained in:
Dahale Aditya Dnyaneshwar 2026-09-26 11:57:55 +05:30
parent 99655b6f86
commit 408729982b
4 changed files with 151 additions and 22 deletions

View file

@ -12,7 +12,6 @@ from urllib.parse import urlparse
import httpx
import litellm
from litellm.constants import OPENAI_SYSTEM_MESSAGES_FIRST_PROVIDERS
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
_extract_reasoning_content,
@ -25,7 +24,6 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
flatten_combinators_and_drop_non_python_regex_patterns,
get_tool_call_names,
hoist_images_from_tool_messages,
system_messages_first,
tool_with_sanitized_parameters,
)
from litellm.litellm_core_utils.prompt_templates.image_handling import (
@ -60,8 +58,9 @@ from litellm.utils import convert_to_model_response_object
from ..common_utils import OpenAIError
if TYPE_CHECKING:
import tiktoken
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer
from litellm.llms.base_llm.base_utils import BaseTokenCounter
from litellm.types.llms.openai import ChatCompletionToolParam
@ -464,15 +463,6 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
]
return MappingProxyType({"tools": sanitized})
def _prompt_cache_ordered_messages(
self, messages: list[AllMessageValues], litellm_params: Mapping[str, object]
) -> list[AllMessageValues]:
if not litellm.openai_system_messages_first:
return messages
if litellm_params.get("custom_llm_provider") not in OPENAI_SYSTEM_MESSAGES_FIRST_PROVIDERS:
return messages
return system_messages_first(messages)
def transform_request(
self,
model: str,
@ -487,9 +477,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
Returns:
dict: The transformed request. Sent as the body of the API call.
"""
messages = self._transform_messages(
messages=self._prompt_cache_ordered_messages(messages, litellm_params), model=model
)
messages = self._transform_messages(messages=messages, model=model)
if not self._should_preserve_cache_control_for_endpoint(
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
):
@ -518,9 +506,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
litellm_params: dict,
headers: dict,
) -> dict:
transformed_messages = await self._transform_messages(
messages=self._prompt_cache_ordered_messages(messages, litellm_params), model=model, is_async=True
)
transformed_messages = await self._transform_messages(messages=messages, model=model, is_async=True)
if not self._should_preserve_cache_control_for_endpoint(
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
):
@ -596,7 +582,9 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
for choice in choices:
## HANDLE JSON MODE - anthropic returns single function call]
tool_calls = choice["message"].get("tool_calls", None)
new_tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | None = None
new_tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | None = (
None # mutable-ok: holds _handle_invalid_parallel_tool_calls' list; Message.__init__ expects list
)
message_content = choice["message"].get("content", None)
if tool_calls is not None:
_openai_tool_calls = []
@ -668,7 +656,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: "Tokenizer | None",
encoding: "tiktoken.Encoding | None",
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:

View file

@ -13,8 +13,9 @@ from litellm.types.utils import ModelResponse
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
if TYPE_CHECKING:
import tiktoken
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -74,6 +75,20 @@ class OpenAILikeChatConfig(OpenAIGPTConfig):
# Sanitize if the key ends with '_tokens' and its value is None
if key.endswith("_tokens") and value is None:
usage[key] = 0
if "prompt_tokens_details" not in usage or not usage.get("prompt_tokens_details"):
prompt_tokens_details = {}
if "cache_read_input_tokens" in usage and isinstance(usage["cache_read_input_tokens"], int):
prompt_tokens_details["cached_tokens"] = usage["cache_read_input_tokens"]
elif "prompt_cache_hit_tokens" in usage and isinstance(usage["prompt_cache_hit_tokens"], int):
prompt_tokens_details["cached_tokens"] = usage["prompt_cache_hit_tokens"]
if "cache_creation_input_tokens" in usage and isinstance(usage["cache_creation_input_tokens"], int):
prompt_tokens_details["cache_write_tokens"] = usage["cache_creation_input_tokens"]
if prompt_tokens_details:
usage["prompt_tokens_details"] = prompt_tokens_details
return response_json
@staticmethod
@ -118,8 +133,33 @@ class OpenAILikeChatConfig(OpenAIGPTConfig):
if base_model is not None:
returned_response._hidden_params["model"] = base_model
if hasattr(returned_response, "usage") and returned_response.usage is not None:
raw_usage = response_json.get("usage") or {}
if "cache_read_input_tokens" in raw_usage and raw_usage["cache_read_input_tokens"] is not None:
returned_response.usage.cache_read_input_tokens = raw_usage["cache_read_input_tokens"]
if "cache_creation_input_tokens" in raw_usage and raw_usage["cache_creation_input_tokens"] is not None:
returned_response.usage.cache_creation_input_tokens = raw_usage["cache_creation_input_tokens"]
return returned_response
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: OpenAIGPTConfig contract
supported_params = super().get_supported_openai_params(model=model)
import litellm
model_info = (
litellm.model_cost.get(model)
or litellm.model_cost.get(f"openai_like/{model}")
or litellm.model_cost.get(f"openai/{model}")
)
if (
isinstance(model_info, dict)
and model_info.get("supports_reasoning") is True
and "reasoning_effort" not in supported_params
):
supported_params.append("reasoning_effort")
return supported_params
def transform_response(
self,
model: str,
@ -130,7 +170,7 @@ class OpenAILikeChatConfig(OpenAIGPTConfig):
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: "Tokenizer | None",
encoding: "tiktoken.Encoding | None",
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:

View file

@ -30,6 +30,18 @@ class TestOpenAIGPTConfig:
supported_params = self.config.get_supported_openai_params("gpt-4.1-mini")
assert "user" in supported_params
def test_reasoning_effort_supported_when_supports_reasoning(self):
"""Test that 'reasoning_effort' is included when model declares supports_reasoning=True."""
litellm.model_cost["openai/custom-reasoning-model"] = {"supports_reasoning": True}
supported_params = self.config.get_supported_openai_params("custom-reasoning-model")
assert "reasoning_effort" in supported_params
# Non-reasoning model should not have reasoning_effort
litellm.model_cost["openai/custom-regular-model"] = {"supports_reasoning": False}
supported_params = self.config.get_supported_openai_params("custom-regular-model")
assert "reasoning_effort" not in supported_params
def test_user_param_supported_for_responses_api_models(self):
"""Test that 'user' param is in supported params for responses API models.

View file

@ -47,3 +47,92 @@ def test_sanitize_usage_obj_valid_usage():
# Assert
assert sanitized_json == original_json # The object should be unchanged
def test_sanitize_usage_obj_normalizes_cache_tokens():
"""
Tests that _sanitize_usage_obj maps cache_read_input_tokens and cache_creation_input_tokens
into prompt_tokens_details for OpenAI compatibility and accurate cost attribution.
"""
response_json = {
"choices": [],
"usage": {
"prompt_tokens": 100,
"completion_tokens": 20,
"total_tokens": 120,
"cache_read_input_tokens": 70,
"cache_creation_input_tokens": 30,
},
}
sanitized = OpenAILikeChatConfig._sanitize_usage_obj(response_json)
assert "prompt_tokens_details" in sanitized["usage"]
assert sanitized["usage"]["prompt_tokens_details"]["cached_tokens"] == 70
assert sanitized["usage"]["prompt_tokens_details"]["cache_write_tokens"] == 30
def test_openai_like_reasoning_effort_supported():
"""
Tests that get_supported_openai_params includes 'reasoning_effort' when model supports reasoning.
"""
import litellm
config = OpenAILikeChatConfig()
litellm.model_cost["openai_like/custom-r1"] = {"supports_reasoning": True}
supported = config.get_supported_openai_params("custom-r1")
assert "reasoning_effort" in supported
litellm.model_cost["openai_like/custom-standard"] = {"supports_reasoning": False}
supported_standard = config.get_supported_openai_params("custom-standard")
assert "reasoning_effort" not in supported_standard
def test_transform_response_preserves_cache_tokens():
"""
Tests that _transform_response maps cache_read_input_tokens and cache_creation_input_tokens
to the ModelResponse usage object.
"""
from unittest.mock import MagicMock
import httpx
from litellm.types.utils import ModelResponse
raw_response = {
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1234567890,
"model": "deepseek-chat",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Hello"}}],
"usage": {
"prompt_tokens": 100,
"completion_tokens": 20,
"total_tokens": 120,
"cache_read_input_tokens": 80,
"cache_creation_input_tokens": 20,
},
}
http_resp = httpx.Response(
200, json=raw_response, request=httpx.Request("POST", "https://api.example.com")
)
config = OpenAILikeChatConfig()
logging_obj = MagicMock()
res = config.transform_response(
model="deepseek-chat",
raw_response=http_resp,
model_response=ModelResponse(),
logging_obj=logging_obj,
request_data={},
messages=[{"role": "user", "content": "hi"}],
optional_params={},
litellm_params={},
encoding=None,
)
assert res.usage.cache_read_input_tokens == 80
assert res.usage.cache_creation_input_tokens == 20
assert res.usage.prompt_tokens_details.cached_tokens == 80
assert res.usage.prompt_tokens_details.cache_write_tokens == 20