fix(proxy): estimate failed-request input tokens on /v1/messages and count system prompts

The Anthropic messages endpoint's exception handler passed the raw
request body dict to the failure hook, but request setup had already
replaced the processor's dict with one carrying the logging object, so
failure rows for /v1/messages never lifted recovered or estimated usage.
Pass the processor's dict instead.

The input-side estimate only counted the messages list, missing the
Anthropic top-level system prompt (string or text-block list) and the
Responses API instructions field, which live in optional_params. Count
them too.
This commit is contained in:
mateo-berri 2026-08-18 14:21:47 -07:00
parent d2fbaff2c9
commit 803113c63a
4 changed files with 140 additions and 9 deletions

View file

@ -179,7 +179,7 @@ async def anthropic_response(
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
request_data=data,
request_data=base_llm_response_processor.data,
)
body: Final = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=e.status_code,
@ -189,7 +189,7 @@ async def anthropic_response(
return JSONResponse(status_code=e.status_code, content=body)
except Exception as e:
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=base_llm_response_processor.data
)
verbose_proxy_logger.exception("litellm.proxy.proxy_server.anthropic_response(): Exception occured - %s", e)

View file

@ -403,24 +403,45 @@ def _exception_changes_request_flow(exc: BaseException) -> bool:
return isinstance(exc, (SensitiveDataRouteException, ModifyResponseException))
def _count_request_input_tokens(model: str, request_input: object) -> int:
def _prompt_block_text(block: object) -> str:
if isinstance(block, str):
return block
if not isinstance(block, dict):
return ""
block_text: Final = block.get("text")
return block_text if isinstance(block_text, str) else ""
def _system_prompt_text(system_input: object) -> str:
if isinstance(system_input, str):
return system_input
if not isinstance(system_input, list):
return ""
return "".join(_prompt_block_text(block) for block in system_input)
def _count_request_input_tokens(model: str, request_input: object, system_input: object) -> int:
system_text: Final = _system_prompt_text(system_input)
system_tokens: Final = litellm.token_counter(model=model, text=system_text) if system_text else 0
if isinstance(request_input, str):
return litellm.token_counter(model=model, text=request_input)
return system_tokens + litellm.token_counter(model=model, text=request_input)
if not isinstance(request_input, list) or not request_input:
return 0
return system_tokens
text_entries: Final = tuple(entry for entry in request_input if isinstance(entry, str))
if len(text_entries) == len(request_input):
return litellm.token_counter(model=model, text="".join(text_entries))
return litellm.token_counter(model=model, messages=request_input)
return system_tokens + litellm.token_counter(model=model, text="".join(text_entries))
return system_tokens + litellm.token_counter(model=model, messages=request_input)
def _estimate_dispatched_failure_usage(model: str, request_input: object) -> Usage | None:
def _estimate_dispatched_failure_usage(model: str, request_input: object, system_input: object) -> Usage | None:
"""A request that failed after dispatch consumed provider-billed input
tokens, but no provider usage ever came back. Estimate the input side with
the same tokenizer fallback interrupted streams use, so the spend log's
failure row records what was sent instead of zero."""
try:
input_tokens: Final = _count_request_input_tokens(model=model, request_input=request_input)
input_tokens: Final = _count_request_input_tokens(
model=model, request_input=request_input, system_input=system_input
)
except Exception:
return None
if input_tokens <= 0:
@ -440,9 +461,16 @@ def _failure_usage_to_lift(model_call_details: Mapping[str, object], dispatched:
return recovered_usage, model_call_details.get("response_cost")
if not dispatched or model_call_details.get(LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL):
return None
optional_params: Final = model_call_details.get("optional_params")
system_input: Final = (
(optional_params.get("system") or optional_params.get("instructions"))
if isinstance(optional_params, dict)
else None
)
estimated_usage: Final = _estimate_dispatched_failure_usage(
model=str(model_call_details.get("model") or ""),
request_input=model_call_details.get("messages"),
system_input=system_input,
)
if estimated_usage is None:
return None

View file

@ -164,6 +164,41 @@ class TestProxyExceptionPassthrough:
mock_logging.post_call_failure_hook.assert_awaited_once()
class TestFailureHookRequestData:
@pytest.mark.asyncio
async def test_failure_hook_gets_post_setup_data_with_logging_obj(self):
"""Request setup replaces the processor's data dict (adding the logging
object the failure hook needs to lift token usage from); the exception
handler must pass that replaced dict, not the raw request body dict."""
import litellm.proxy.anthropic_endpoints.endpoints as ep
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
captured = {}
async def fake_process(self, **kwargs):
self.data = {**self.data, "litellm_logging_obj": "logging-obj-sentinel"}
captured["processor_data"] = self.data
raise RuntimeError("provider timeout")
with (
patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})),
patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process),
patch.object(proxy_server, "proxy_logging_obj") as mock_logging,
):
mock_logging.post_call_failure_hook = AsyncMock()
with pytest.raises(ProxyException):
await ep.anthropic_response(
fastapi_response=MagicMock(),
request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(),
)
hook_request_data = mock_logging.post_call_failure_hook.await_args.kwargs["request_data"]
assert hook_request_data is captured["processor_data"]
assert hook_request_data["litellm_logging_obj"] == "logging-obj-sentinel"
class TestEventLoggingBatchEndpoint:
"""Test the stubbed event logging batch endpoint"""

View file

@ -616,6 +616,74 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens:
assert estimated.prompt_tokens > 0
assert estimated.completion_tokens == 0
def _dispatched_request_data(self, messages, optional_params):
from datetime import datetime
return {
"litellm_logging_obj": self._logging_obj(
{
"first_api_call_start_time": datetime.now(),
"model": "gpt-3.5-turbo",
"messages": messages,
"optional_params": optional_params,
}
),
"metadata": {},
}
@pytest.mark.asyncio
async def test_anthropic_system_prompt_counted_in_estimate(self):
import litellm as litellm_module
from litellm.types.utils import Usage
system_prompt = "You are a verbose historian who narrates every fact in exhaustive detail."
messages = [{"role": "user", "content": "write a short essay"}]
request_data = self._dispatched_request_data(messages, {"system": system_prompt, "max_tokens": 100})
await self._run(request_data)
estimated = request_data["combined_usage_object"]
assert isinstance(estimated, Usage)
expected = litellm_module.token_counter(model="gpt-3.5-turbo", messages=messages) + litellm_module.token_counter(
model="gpt-3.5-turbo", text=system_prompt
)
assert estimated.prompt_tokens == expected
@pytest.mark.asyncio
async def test_anthropic_system_text_blocks_counted_in_estimate(self):
import litellm as litellm_module
from litellm.types.utils import Usage
system_blocks = [
{"type": "text", "text": "part one of the system prompt. "},
{"type": "text", "text": "part two of the system prompt."},
]
messages = [{"role": "user", "content": "write a short essay"}]
request_data = self._dispatched_request_data(messages, {"system": system_blocks})
await self._run(request_data)
estimated = request_data["combined_usage_object"]
assert isinstance(estimated, Usage)
expected = litellm_module.token_counter(model="gpt-3.5-turbo", messages=messages) + litellm_module.token_counter(
model="gpt-3.5-turbo", text="part one of the system prompt. part two of the system prompt."
)
assert estimated.prompt_tokens == expected
@pytest.mark.asyncio
async def test_responses_instructions_counted_in_estimate(self):
import litellm as litellm_module
from litellm.types.utils import Usage
instructions = "Answer every question as a meticulous archivist."
request_data = self._dispatched_request_data("summarize the archive", {"instructions": instructions})
await self._run(request_data)
estimated = request_data["combined_usage_object"]
assert isinstance(estimated, Usage)
expected = litellm_module.token_counter(
model="gpt-3.5-turbo", text="summarize the archive"
) + litellm_module.token_counter(model="gpt-3.5-turbo", text=instructions)
assert estimated.prompt_tokens == expected
from typing import cast