fix(proxy): strip client standard_logging_object and zero-fill unknown recovered cost on the failure path

Auth and pass-through failures reach post_call_failure_hook with the raw request body unstripped, so a client-supplied standard_logging_object could feed the new attribution fallback when the logging object carries none. Pop the key before the lift so only the logging object may supply it. Also coalesce a None recovered cost to 0.0 so the lift always overwrites any client-supplied response_cost, matching the merge base's clobber semantics.
This commit is contained in:
mateo-berri 2026-08-19 16:19:26 -07:00
parent b477d0967a
commit dce207add4
2 changed files with 38 additions and 1 deletions

View file

@ -539,7 +539,7 @@ def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str,
_entries: Final = (
("first_api_call_start_time", _first_handoff),
("combined_usage_object", None if _usage_to_lift is None else _usage_to_lift[0]),
("response_cost", None if _usage_to_lift is None else _usage_to_lift[1]),
("response_cost", None if _usage_to_lift is None else (_usage_to_lift[1] or 0.0)),
("standard_logging_object", _model_call_details.get("standard_logging_object")),
)
return MappingProxyType({key: value for key, value in _entries if value is not None})
@ -2322,6 +2322,9 @@ class ProxyLogging:
original_exception=original_exception,
)
# Auth and pass-through failures reach this hook with the raw request
# body unstripped, so only the logging object may supply this key.
request_data.pop("standard_logging_object", None)
request_data.update(_failure_fields_to_lift(request_data))
# Remove before callbacks iterate — not serialisable

View file

@ -468,6 +468,23 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend:
assert request_data["response_cost"] == 3.5e-05
assert "litellm_logging_obj" not in request_data
@pytest.mark.asyncio
async def test_recovered_usage_without_cost_clobbers_client_cost_with_zero(self):
from litellm.types.utils import Usage
recovered_usage = Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31)
logging_obj = MagicMock()
logging_obj.model_call_details = {"combined_usage_object": recovered_usage}
request_data = {
"litellm_logging_obj": logging_obj,
"response_cost": 999.0,
"metadata": {},
}
await self._run(request_data)
assert request_data["combined_usage_object"] is recovered_usage
assert request_data["response_cost"] == 0.0
@pytest.mark.asyncio
async def test_no_recovered_usage_is_noop(self):
logging_obj = MagicMock()
@ -522,6 +539,23 @@ class TestPostCallFailureHookLiftsStandardLoggingObject:
await self._run(request_data)
assert request_data["standard_logging_object"] is authoritative
@pytest.mark.asyncio
async def test_client_supplied_key_is_stripped_when_logging_obj_supplies_none(self):
spoofed = {"model_id": "client-injected"}
request_data = {"standard_logging_object": spoofed, "metadata": {}}
await self._run(request_data)
assert "standard_logging_object" not in request_data
logging_obj = MagicMock()
logging_obj.model_call_details = {}
request_data_with_obj = {
"litellm_logging_obj": logging_obj,
"standard_logging_object": spoofed,
"metadata": {},
}
await self._run(request_data_with_obj)
assert "standard_logging_object" not in request_data_with_obj
@pytest.mark.asyncio
async def test_no_standard_logging_object_is_noop(self):
logging_obj = MagicMock()