mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Merge 9572b60b80 into 3746ba58d7
This commit is contained in:
commit
cdb87d72eb
3 changed files with 100 additions and 24 deletions
|
|
@ -227,9 +227,11 @@ class BaseResponsesAPIStreamingIterator:
|
|||
"api_base": _api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
}
|
||||
self._hidden_params["additional_headers"] = process_response_headers(
|
||||
self.response.headers or {}
|
||||
) # GUARANTEE OPENAI HEADERS IN RESPONSE
|
||||
_raw_headers: Final = dict(self.response.headers or {}) # mutable-ok: process_response_headers takes a dict
|
||||
_processed_headers: Final = process_response_headers(_raw_headers)
|
||||
self._hidden_params["additional_headers"] = _processed_headers # GUARANTEE OPENAI HEADERS IN RESPONSE
|
||||
self._raw_headers: Final[Mapping[str, str]] = MappingProxyType(_raw_headers)
|
||||
self._processed_headers: Final[Mapping[str, object]] = MappingProxyType(_processed_headers)
|
||||
|
||||
def _check_max_streaming_duration(self) -> None:
|
||||
"""Raise litellm.Timeout if the stream has exceeded LITELLM_MAX_STREAMING_DURATION_SECONDS."""
|
||||
|
|
@ -407,18 +409,7 @@ class BaseResponsesAPIStreamingIterator:
|
|||
if self._persist_completed_response_before_logging:
|
||||
self._persist_completed_response_to_cache(is_async=is_async)
|
||||
|
||||
# Create a copy for logging to avoid modifying the response object that will be returned to the user
|
||||
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
|
||||
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
|
||||
# Use model_dump + model_validate instead of deepcopy to avoid pickle errors with
|
||||
# Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192)
|
||||
logging_response = self.completed_response
|
||||
if self.completed_response is not None and hasattr(self.completed_response, "model_dump"):
|
||||
try:
|
||||
logging_response = type(self.completed_response).model_validate(self.completed_response.model_dump())
|
||||
except Exception:
|
||||
# Fallback to original if serialization fails
|
||||
pass
|
||||
logging_response: Final = self._build_logging_response()
|
||||
|
||||
end_time: Final = datetime.now()
|
||||
if is_async:
|
||||
|
|
@ -448,6 +439,30 @@ class BaseResponsesAPIStreamingIterator:
|
|||
)
|
||||
self._run_post_success_hooks(end_time=end_time)
|
||||
|
||||
def _build_logging_response(self) -> ResponsesAPIStreamingResponse | None:
|
||||
completed: Final = self.completed_response
|
||||
if completed is None:
|
||||
return None
|
||||
copied: Final = self._copy_for_logging(completed)
|
||||
inner_response: Final = getattr(copied, "response", None)
|
||||
if isinstance(inner_response, ResponsesAPIResponse):
|
||||
inner_response.store_provider_response_headers(
|
||||
processed_headers=self._processed_headers,
|
||||
raw_headers=self._raw_headers,
|
||||
)
|
||||
return copied
|
||||
|
||||
@staticmethod
|
||||
def _copy_for_logging(completed: ResponsesAPIStreamingResponse) -> ResponsesAPIStreamingResponse:
|
||||
"""
|
||||
model_dump + model_validate, instead of deepcopy, avoids pickle errors with Pydantic
|
||||
ValidatorIterator on tool_choice.allowed_tools (fixes #17192)
|
||||
"""
|
||||
try:
|
||||
return type(completed).model_validate(completed.model_dump())
|
||||
except Exception:
|
||||
return completed
|
||||
|
||||
def _handle_logging_completed_response(self):
|
||||
"""Base implementation - should be overridden by subclasses"""
|
||||
|
||||
|
|
|
|||
|
|
@ -1347,6 +1347,16 @@ class ResponsesAPIResponse(BaseLiteLLMOpenAIResponseObject):
|
|||
# Define private attributes using PrivateAttr
|
||||
_hidden_params: dict = PrivateAttr(default_factory=dict)
|
||||
|
||||
def store_provider_response_headers(
|
||||
self, *, processed_headers: Mapping[str, object], raw_headers: Mapping[str, str]
|
||||
) -> None:
|
||||
"""Keep provider headers on a response rebuilt from `model_dump()`, which drops private attrs."""
|
||||
self._hidden_params = { # mutable-ok: the private attr is typed as dict
|
||||
**self._hidden_params,
|
||||
"additional_headers": dict(processed_headers), # mutable-ok: logging callbacks expect a plain dict
|
||||
"headers": dict(raw_headers), # mutable-ok: logging callbacks expect a plain dict
|
||||
}
|
||||
|
||||
@field_validator("reasoning", mode="before")
|
||||
@classmethod
|
||||
def validate_reasoning_to_dict(cls, value: Any) -> dict[str, Any] | None:
|
||||
|
|
|
|||
|
|
@ -18,7 +18,11 @@ from litellm.litellm_core_utils import thread_pool_executor as thread_pool_execu
|
|||
from litellm.responses import streaming_iterator as responses_streaming_iterator_module
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
|
||||
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
|
||||
class RecordingCustomLogger(CustomLogger):
|
||||
|
|
@ -93,14 +97,8 @@ def _make_logging_obj() -> LitellmLogging:
|
|||
return logging_obj
|
||||
|
||||
|
||||
def _make_iterator(logging_obj: LitellmLogging) -> ResponsesAPIStreamingIterator:
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="gpt-5.4-nano",
|
||||
responses_api_provider_config=None,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
iterator.completed_response = ResponsesAPIResponse(
|
||||
def _make_completed_response() -> ResponsesAPIResponse:
|
||||
return ResponsesAPIResponse(
|
||||
id="resp_lit4210",
|
||||
created_at=1700000000.0,
|
||||
model="gpt-5.4-nano",
|
||||
|
|
@ -116,6 +114,16 @@ def _make_iterator(logging_obj: LitellmLogging) -> ResponsesAPIStreamingIterator
|
|||
temperature=1.0,
|
||||
top_p=1.0,
|
||||
)
|
||||
|
||||
|
||||
def _make_iterator(logging_obj: LitellmLogging) -> ResponsesAPIStreamingIterator:
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200),
|
||||
model="gpt-5.4-nano",
|
||||
responses_api_provider_config=None,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
iterator.completed_response = _make_completed_response()
|
||||
return iterator
|
||||
|
||||
|
||||
|
|
@ -156,3 +164,46 @@ async def test_sync_callbacks_run_only_after_async_handler_completes(recording_e
|
|||
submit_times = recording_executor.submit_times_for(logging_obj)
|
||||
assert len(submit_times) == 1
|
||||
assert submit_times[0] >= recorder.async_hook_finished
|
||||
|
||||
|
||||
class HeaderCapturingLogger(CustomLogger):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.hidden_params: dict | None = None
|
||||
self.standard_logging_payload: dict | None = None
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
self.hidden_params = getattr(response_obj, "_hidden_params", None)
|
||||
self.standard_logging_payload = kwargs.get("standard_logging_object")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_response_headers_reach_success_callbacks(monkeypatch, recording_executor):
|
||||
"""Regression test for #38090: provider headers (e.g. Azure apim-request-id) must survive
|
||||
the completed-response copy made for logging"""
|
||||
recorder = HeaderCapturingLogger()
|
||||
monkeypatch.setattr(litellm, "success_callback", [recorder])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [recorder])
|
||||
|
||||
logging_obj = _make_logging_obj()
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=httpx.Response(200, headers={"apim-request-id": "req-38090", "x-ms-region": "East US 2"}),
|
||||
model="gpt-5.4-nano",
|
||||
responses_api_provider_config=None,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
iterator.completed_response = ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=_make_completed_response(),
|
||||
)
|
||||
|
||||
iterator._log_completed_response(is_async=True)
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
assert recorder.hidden_params is not None
|
||||
assert recorder.hidden_params["additional_headers"]["llm_provider-apim-request-id"] == "req-38090"
|
||||
assert recorder.hidden_params["headers"]["apim-request-id"] == "req-38090"
|
||||
assert recorder.standard_logging_payload is not None
|
||||
logged_headers = recorder.standard_logging_payload["hidden_params"]["additional_headers"]
|
||||
assert logged_headers["llm_provider-apim-request-id"] == "req-38090"
|
||||
assert logged_headers["llm_provider-x-ms-region"] == "East US 2"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue