mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
d2fbaff2c9
commit
803113c63a
4 changed files with 140 additions and 9 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue