mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge 7b23f60be3 into 2dccc0dc79
This commit is contained in:
commit
0b5ba872cc
2 changed files with 128 additions and 12 deletions
|
|
@ -510,8 +510,7 @@ class BaseResponsesAPIStreamingIterator:
|
|||
if self._persist_completed_response_before_logging:
|
||||
self._persist_completed_response_to_cache(is_async=is_async)
|
||||
|
||||
logging_response: Final[object] = _logging_copy(self.completed_response)
|
||||
self._restore_provider_response_headers(logging_response)
|
||||
logging_response: Final = self._response_for_success_logging()
|
||||
|
||||
end_time: Final = datetime.now()
|
||||
if is_async:
|
||||
|
|
@ -546,6 +545,16 @@ class BaseResponsesAPIStreamingIterator:
|
|||
)
|
||||
self._run_post_success_hooks(end_time=end_time)
|
||||
|
||||
def _response_for_success_logging(self) -> object:
|
||||
"""The handlers only unwrap the completion event in their assembled-stream branch, which a
|
||||
non-streaming caller draining this iterator never reaches, so unwrap the logging copy here."""
|
||||
copied: Final = _logging_copy(self.completed_response)
|
||||
self._restore_provider_response_headers(copied)
|
||||
unwrapped: Final = getattr(copied, "response", None)
|
||||
if getattr(self.logging_obj, "stream", None) is not True and isinstance(unwrapped, ResponsesAPIResponse):
|
||||
return unwrapped
|
||||
return copied
|
||||
|
||||
def _restore_provider_response_headers(self, logging_response: object) -> None:
|
||||
"""Re-apply the provider's response headers to the copy handed to logging callbacks.
|
||||
|
||||
|
|
|
|||
|
|
@ -251,6 +251,104 @@ def test_sync_transport_error_before_completed_event_raises():
|
|||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("stream", [False, True])
|
||||
async def test_completed_event_is_unwrapped_only_for_non_streaming_callers(stream: bool):
|
||||
"""A non-streaming caller can still drain this iterator (chat -> responses bridge
|
||||
against a provider that always answers with SSE). The success handlers only unwrap
|
||||
the completion event when logging in stream mode, so a wrapped event there yields no
|
||||
standard_logging_object and the SpendLogs row is dropped (#36426)."""
|
||||
logging_obj = _logging_obj_stub()
|
||||
logging_obj.stream = stream
|
||||
|
||||
iterator = _make_header_iterator(headers={}, config=_headers_config(), logging_obj=logging_obj)
|
||||
async for _ in iterator:
|
||||
pass
|
||||
|
||||
logged = logging_obj.dispatch_success_handlers.call_args.args[0]
|
||||
if stream:
|
||||
assert logged.type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
else:
|
||||
assert isinstance(logged, ResponsesAPIResponse)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("stream", [False, True])
|
||||
async def test_completed_response_builds_real_spend_log_payload(stream: bool):
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
|
||||
from litellm.types.llms.openai import ResponseAPIUsage
|
||||
|
||||
model: Final = "gpt-4o-mini"
|
||||
request_id: Final = "resp_spend_regression"
|
||||
prompt_tokens: Final = 11
|
||||
completion_tokens: Final = 7
|
||||
total_tokens: Final = prompt_tokens + completion_tokens
|
||||
input_cost_per_token: Final = 0.00000015
|
||||
output_cost_per_token: Final = 0.0000006
|
||||
expected_cost: Final = prompt_tokens * input_cost_per_token + completion_tokens * output_cost_per_token
|
||||
expected_status: Final = "success"
|
||||
start_time: Final = datetime.now()
|
||||
logging_obj: Final = LiteLLMLoggingObj(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
stream=stream,
|
||||
call_type="aresponses",
|
||||
start_time=start_time,
|
||||
litellm_call_id=request_id,
|
||||
function_id=request_id,
|
||||
)
|
||||
logging_obj.update_environment_variables(
|
||||
litellm_params={
|
||||
"aresponses": True,
|
||||
"input_cost_per_token": input_cost_per_token,
|
||||
"output_cost_per_token": output_cost_per_token,
|
||||
},
|
||||
optional_params={},
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
logging_obj._on_deferred_stream_complete = lambda: None
|
||||
response: Final = ResponsesAPIResponse(
|
||||
id=request_id,
|
||||
created_at=1,
|
||||
model=model,
|
||||
object="response",
|
||||
status="completed",
|
||||
output=[],
|
||||
usage=ResponseAPIUsage(
|
||||
input_tokens=prompt_tokens,
|
||||
output_tokens=completion_tokens,
|
||||
total_tokens=total_tokens,
|
||||
),
|
||||
)
|
||||
iterator: Final = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model=model,
|
||||
responses_api_provider_config=Mock(spec=BaseResponsesAPIConfig),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
iterator.completed_response = ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=response,
|
||||
)
|
||||
iterator._log_completed_response(is_async=True)
|
||||
await logging_obj._deferred_stream_complete_args[0]
|
||||
|
||||
standard: Final = logging_obj.model_call_details["standard_logging_object"]
|
||||
assert standard["prompt_tokens"] == prompt_tokens
|
||||
assert standard["completion_tokens"] == completion_tokens
|
||||
assert standard["total_tokens"] == total_tokens
|
||||
assert standard["response_cost"] == pytest.approx(expected_cost)
|
||||
assert standard["status"] == expected_status
|
||||
payload: Final = get_logging_payload(
|
||||
logging_obj.model_call_details, response, start_time, datetime.now()
|
||||
)
|
||||
assert payload["prompt_tokens"] == prompt_tokens
|
||||
assert payload["completion_tokens"] == completion_tokens
|
||||
assert payload["total_tokens"] == total_tokens
|
||||
assert payload["spend"] == pytest.approx(expected_cost)
|
||||
assert payload["status"] == expected_status
|
||||
|
||||
|
||||
def test_stream_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatch):
|
||||
"""
|
||||
Regression test for LIT-6184 on the /v1/responses streaming surface: the
|
||||
|
|
@ -592,10 +690,12 @@ def _make_header_iterator(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_logging_response_carries_provider_response_headers():
|
||||
@pytest.mark.parametrize("stream", [False, True])
|
||||
async def test_streaming_logging_response_carries_provider_response_headers(stream: bool):
|
||||
"""LIT-6055: the provider headers the iterator captured must reach the logged response, so
|
||||
custom loggers can read Azure's apim-request-id from the callback payload."""
|
||||
logging_obj = _logging_obj_stub()
|
||||
logging_obj.stream = stream
|
||||
logged: list[object] = []
|
||||
logging_obj.dispatch_success_handlers = _capture_dispatch(logged)
|
||||
|
||||
|
|
@ -610,7 +710,8 @@ async def test_streaming_logging_response_carries_provider_response_headers():
|
|||
pass
|
||||
|
||||
assert len(logged) == 1
|
||||
hidden_params = logged[0].response._hidden_params
|
||||
logged_response = logged[0].response if stream else logged[0]
|
||||
hidden_params = logged_response._hidden_params
|
||||
assert hidden_params["additional_headers"]["llm_provider-apim-request-id"] == "azure-correlation-1"
|
||||
assert hidden_params["additional_headers"]["llm_provider-x-ms-region"] == "East US 2"
|
||||
assert hidden_params["headers"]["apim-request-id"] == "azure-correlation-1"
|
||||
|
|
@ -621,10 +722,12 @@ async def test_streaming_logging_response_carries_provider_response_headers():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_logging_copy_preserves_transform_hidden_params():
|
||||
@pytest.mark.parametrize("stream", [False, True])
|
||||
async def test_streaming_logging_copy_preserves_transform_hidden_params(stream: bool):
|
||||
"""LIT-6055: model_validate(model_dump()) drops pydantic private attributes, so headers a
|
||||
provider transform already set on the response (fake_stream) must be re-applied."""
|
||||
logging_obj = _logging_obj_stub()
|
||||
logging_obj.stream = stream
|
||||
logged: list[object] = []
|
||||
logging_obj.dispatch_success_handlers = _capture_dispatch(logged)
|
||||
|
||||
|
|
@ -645,7 +748,8 @@ async def test_streaming_logging_copy_preserves_transform_hidden_params():
|
|||
pass
|
||||
|
||||
assert len(logged) == 1
|
||||
hidden_params = logged[0].response._hidden_params
|
||||
logged_response = logged[0].response if stream else logged[0]
|
||||
hidden_params = logged_response._hidden_params
|
||||
assert hidden_params["additional_headers"]["llm_provider-apim-request-id"] == "from-transform"
|
||||
assert hidden_params["headers"]["apim-request-id"] == "from-transform"
|
||||
assert iterator.completed_response is not logged[0]
|
||||
|
|
@ -654,16 +758,18 @@ async def test_streaming_logging_copy_preserves_transform_hidden_params():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_logging_copy_fallback_leaves_caller_event_untouched():
|
||||
"""LIT-6055: when the logging copy falls back to the original event, the header restore must
|
||||
not stamp logging-only state onto the object the caller is iterating."""
|
||||
@pytest.mark.parametrize("stream", [False, True])
|
||||
async def test_streaming_logging_copy_fallback_leaves_caller_event_untouched(stream: bool):
|
||||
"""Serialization failures must keep logging headers separate from the caller's response."""
|
||||
expected_request_id: Final = "azure-correlation-1"
|
||||
logging_obj = _logging_obj_stub()
|
||||
logging_obj.stream = stream
|
||||
logged: list[object] = []
|
||||
logging_obj.dispatch_success_handlers = _capture_dispatch(logged)
|
||||
logging_obj._on_deferred_stream_complete = None
|
||||
|
||||
iterator = _make_header_iterator(
|
||||
headers={"apim-request-id": "azure-correlation-1"},
|
||||
headers={"apim-request-id": expected_request_id},
|
||||
config=_headers_config(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
|
@ -678,8 +784,9 @@ async def test_streaming_logging_copy_fallback_leaves_caller_event_untouched():
|
|||
|
||||
assert len(logged) == 1
|
||||
assert logged[0] is not iterator.completed_response
|
||||
assert logged[0].response is not iterator.completed_response.response
|
||||
assert logged[0].response._hidden_params["headers"]["apim-request-id"] == "azure-correlation-1"
|
||||
logged_response = logged[0].response if stream else logged[0]
|
||||
assert logged_response is not iterator.completed_response.response
|
||||
assert logged_response._hidden_params["headers"]["apim-request-id"] == expected_request_id
|
||||
assert iterator.completed_response.response._hidden_params == {}
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue