mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
b477d0967a
commit
dce207add4
2 changed files with 38 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue