mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(responses): build the billed terminal response immutably and guard the cache dump
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
7121e64db4
commit
7cc07d437a
2 changed files with 112 additions and 28 deletions
|
|
@ -427,39 +427,45 @@ class BaseResponsesAPIStreamingIterator:
|
|||
openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE,
|
||||
openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED,
|
||||
):
|
||||
self.completed_response = openai_responses_api_chunk
|
||||
_response_obj: Final[object] = getattr(openai_responses_api_chunk, "response", None)
|
||||
_typed_response: Final[ResponsesAPIResponse | None] = (
|
||||
ResponsesAPIResponse.model_construct(**_response_obj) # pyright: ignore[reportUnknownArgumentType] # the model_constructed terminal event leaves response as an untyped dict
|
||||
if isinstance(_response_obj, dict)
|
||||
else _response_obj
|
||||
if isinstance(_response_obj, ResponsesAPIResponse)
|
||||
else None
|
||||
_estimate_wanted: Final[bool] = _chunk_type in (
|
||||
openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE,
|
||||
)
|
||||
if (
|
||||
_typed_response is not None
|
||||
and _chunk_type
|
||||
in (
|
||||
openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE,
|
||||
_billed_response: Final[ResponsesAPIResponse | None] = _billed_terminal_response(
|
||||
_response_obj,
|
||||
(
|
||||
lambda: (
|
||||
_estimate_usage_safely(
|
||||
self.model or "",
|
||||
self.request_data.get("input"),
|
||||
self.request_data,
|
||||
self._generated_content + self._generated_tool_arguments,
|
||||
)
|
||||
if _estimate_wanted
|
||||
else None
|
||||
)
|
||||
),
|
||||
)
|
||||
_terminal_chunk: Final = (
|
||||
openai_responses_api_chunk
|
||||
if _billed_response is None or _billed_response is _response_obj
|
||||
else (
|
||||
openai_responses_api_chunk.model_copy(update={"response": _billed_response})
|
||||
if issubclass(type(openai_responses_api_chunk), BaseModel) # pyright: ignore[reportUnnecessaryIsInstance] # test stubs use spec'd Mocks whose __class__ reports BaseModel but whose model_copy returns a Mock
|
||||
else _replace_response(openai_responses_api_chunk, _billed_response)
|
||||
)
|
||||
and _typed_response.usage is None
|
||||
):
|
||||
_typed_response.usage = _estimate_usage_safely(
|
||||
self.model or "",
|
||||
self.request_data.get("input"),
|
||||
self.request_data,
|
||||
self._generated_content + self._generated_tool_arguments,
|
||||
)
|
||||
if _typed_response is not None and _typed_response is not _response_obj:
|
||||
openai_responses_api_chunk.response = _typed_response # pyright: ignore[reportAttributeAccessIssue] # reached only on the dict path, which only response-carrying terminal events produce
|
||||
_stamp_responses_usage_cost(_typed_response, self.logging_obj)
|
||||
)
|
||||
self.completed_response = _terminal_chunk
|
||||
_stamp_responses_usage_cost(_billed_response, self.logging_obj)
|
||||
|
||||
if _chunk_type == openai_types.ResponsesAPIStreamEvents.RESPONSE_FAILED:
|
||||
self._handle_logging_failed_response()
|
||||
else:
|
||||
self._handle_logging_completed_response()
|
||||
|
||||
return _terminal_chunk
|
||||
|
||||
return openai_responses_api_chunk
|
||||
|
||||
return None
|
||||
|
|
@ -688,7 +694,9 @@ class BaseResponsesAPIStreamingIterator:
|
|||
if cache is None:
|
||||
return
|
||||
|
||||
cached_response: Final = response_obj.model_dump_json()
|
||||
cached_response: Final = _dump_json_safely(response_obj)
|
||||
if cached_response is None:
|
||||
return
|
||||
if is_async:
|
||||
from litellm.caching.caching_handler import create_cache_write_task
|
||||
|
||||
|
|
@ -1334,6 +1342,38 @@ def _add_text_like_part_events(
|
|||
)
|
||||
|
||||
|
||||
def _billed_terminal_response(
|
||||
response_obj: object, estimate: Callable[[], ResponseAPIUsage | None] | None
|
||||
) -> ResponsesAPIResponse | None:
|
||||
if isinstance(response_obj, ResponsesAPIResponse):
|
||||
return (
|
||||
response_obj
|
||||
if response_obj.usage is not None or estimate is None
|
||||
else response_obj.model_copy(update={"usage": estimate()})
|
||||
)
|
||||
if not isinstance(response_obj, dict):
|
||||
return None
|
||||
usage: Final[object] = response_obj.get("usage") # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # a model_constructed terminal event leaves response as an untyped dict
|
||||
return ResponsesAPIResponse.model_construct(
|
||||
**{**response_obj, "usage": usage if usage is not None or estimate is None else estimate()} # pyright: ignore[reportUnknownArgumentType, reportArgumentType] # same untyped dict spread
|
||||
)
|
||||
|
||||
|
||||
def _replace_response(
|
||||
event: ResponsesAPIStreamingResponse, response: ResponsesAPIResponse
|
||||
) -> ResponsesAPIStreamingResponse:
|
||||
setattr(event, "response", response)
|
||||
return event
|
||||
|
||||
|
||||
def _dump_json_safely(response: BaseModel) -> str | None:
|
||||
try:
|
||||
return response.model_dump_json()
|
||||
except Exception as exc:
|
||||
verbose_logger.debug("could not serialize completed response for cache: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
def _logging_copy(event: object) -> object:
|
||||
"""Hand logging callbacks a copy, so their usage rewrite (Responses shape to chat shape) never
|
||||
reaches the event the caller is iterating. The round trip through ``model_dump`` sidesteps the
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from unittest.mock import Mock, patch
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic_core import PydanticSerializationError
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -887,10 +888,11 @@ async def test_completed_event_with_a_dict_response_is_typed_and_billed():
|
|||
request_data={"input": "count these input tokens please"},
|
||||
)
|
||||
|
||||
async for _ in iterator:
|
||||
pass
|
||||
yielded: Final = [chunk async for chunk in iterator]
|
||||
|
||||
completed_response: Final = iterator.completed_response.response
|
||||
terminal_event: Final = iterator.completed_response
|
||||
assert yielded[-1] is terminal_event
|
||||
completed_response: Final = terminal_event.response
|
||||
assert isinstance(completed_response, ResponsesAPIResponse)
|
||||
usage: Final = completed_response.usage
|
||||
assert usage is not None
|
||||
|
|
@ -898,3 +900,45 @@ async def test_completed_event_with_a_dict_response_is_typed_and_billed():
|
|||
assert usage.output_tokens > 0
|
||||
assert usage.cost == pytest.approx(0.000704)
|
||||
logging_obj._response_cost_calculator.assert_any_call(result=completed_response)
|
||||
|
||||
|
||||
def test_billed_terminal_response_keeps_a_response_that_already_has_usage():
|
||||
from litellm.responses.streaming_iterator import _billed_terminal_response
|
||||
|
||||
response: Final = _responses_api_response_with_usage()
|
||||
|
||||
assert _billed_terminal_response(response, None) is response
|
||||
|
||||
|
||||
def test_billed_terminal_response_copies_when_estimating_and_leaves_the_original_untouched():
|
||||
from litellm.responses.streaming_iterator import _billed_terminal_response
|
||||
|
||||
response: Final = _responses_api_response_without_usage()
|
||||
estimated: Final = ResponseAPIUsage(input_tokens=3, output_tokens=4, total_tokens=7)
|
||||
|
||||
billed: Final = _billed_terminal_response(response, lambda: estimated)
|
||||
|
||||
assert billed is not response
|
||||
assert billed.usage is estimated
|
||||
assert response.usage is None
|
||||
|
||||
|
||||
def test_persist_completed_response_to_cache_survives_an_unserializable_response(monkeypatch):
|
||||
bad_response: Final = ResponsesAPIResponse.model_construct(id="r", output=[object()], usage=None)
|
||||
with pytest.raises(PydanticSerializationError):
|
||||
bad_response.model_dump_json()
|
||||
|
||||
logging_obj: Final = _logging_obj_stub()
|
||||
caching_handler: Final = Mock()
|
||||
caching_handler.request_kwargs = {"stream": True}
|
||||
logging_obj._llm_caching_handler = caching_handler
|
||||
iterator: Final = _make_iterator(sse_events=[], logging_obj=logging_obj)
|
||||
iterator.completed_response = ResponseCompletedEvent.model_construct(
|
||||
type="response.completed", response=bad_response
|
||||
)
|
||||
cache: Final = Mock()
|
||||
monkeypatch.setattr(litellm, "cache", cache)
|
||||
|
||||
iterator._persist_completed_response_to_cache(is_async=False)
|
||||
|
||||
cache.add_cache.assert_not_called()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue