mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #35114 from BerriAI/litellm_fix_messages_stream_cost_cache_tokens
fix(cost): match streamed Messages usage cost to the recorded spend
This commit is contained in:
commit
a4dd1be53b
4 changed files with 364 additions and 48 deletions
|
|
@ -2216,7 +2216,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def calculate_usage(
|
||||
self,
|
||||
usage_object: dict,
|
||||
usage_object: Mapping[str, Any],
|
||||
reasoning_content: str | None,
|
||||
completion_response: dict | None = None,
|
||||
speed: str | None = None,
|
||||
|
|
|
|||
|
|
@ -158,7 +158,7 @@ ProxyRouteType: TypeAlias = Literal[
|
|||
"acancel_run",
|
||||
"adelete_run",
|
||||
]
|
||||
from litellm.types.utils import ServerToolUse
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
# Type alias for streaming chunk serializer (chunk after hooks + cost injection -> wire format)
|
||||
StreamChunkSerializer = Callable[[Any], str]
|
||||
|
|
@ -3321,7 +3321,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
str_so_far += str(chunk.get("content", ""))
|
||||
|
||||
model_name = request_data.get("model", "")
|
||||
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, model_name)
|
||||
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
chunk, model_name, request_data.get("litellm_logging_obj")
|
||||
)
|
||||
|
||||
# Set before the yield: an async generator suspends at the yield,
|
||||
# so a GeneratorExit on client disconnect is raised there and any
|
||||
|
|
@ -3418,20 +3420,27 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
@overload
|
||||
@staticmethod
|
||||
def _process_chunk_with_cost_injection(chunk: bytes, model_name: str) -> bytes: ...
|
||||
def _process_chunk_with_cost_injection(
|
||||
chunk: bytes, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
|
||||
) -> bytes: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def _process_chunk_with_cost_injection(chunk: object, model_name: str) -> object: ...
|
||||
def _process_chunk_with_cost_injection(
|
||||
chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
|
||||
) -> object: ...
|
||||
|
||||
@staticmethod
|
||||
def _process_chunk_with_cost_injection(chunk: object, model_name: str) -> object:
|
||||
def _process_chunk_with_cost_injection(
|
||||
chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
|
||||
) -> object:
|
||||
"""
|
||||
Process a streaming chunk and inject cost information if enabled.
|
||||
|
||||
Args:
|
||||
chunk: The streaming chunk (dict, str, bytes, or bytearray)
|
||||
model_name: Model name for cost calculation
|
||||
litellm_logging_obj: The call's logging object, used for pricing
|
||||
|
||||
Returns:
|
||||
The processed chunk with cost information injected if applicable
|
||||
|
|
@ -3441,21 +3450,27 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
try:
|
||||
if isinstance(chunk, dict):
|
||||
maybe_modified: Final = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(chunk, model_name)
|
||||
maybe_modified: Final = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(
|
||||
chunk, model_name, litellm_logging_obj
|
||||
)
|
||||
if maybe_modified is not None:
|
||||
return maybe_modified
|
||||
elif isinstance(chunk, (bytes, bytearray)):
|
||||
try:
|
||||
s: Final = chunk.decode("utf-8")
|
||||
if s.endswith(("\n\n", "\r\n\r\n")):
|
||||
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(s, model_name)
|
||||
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(
|
||||
s, model_name, litellm_logging_obj
|
||||
)
|
||||
if maybe_mod is not None:
|
||||
return maybe_mod.encode("utf-8")
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(chunk, str):
|
||||
# Try to parse SSE frame and inject cost into the data line
|
||||
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(chunk, model_name)
|
||||
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(
|
||||
chunk, model_name, litellm_logging_obj
|
||||
)
|
||||
if maybe_mod is not None:
|
||||
# Ensure trailing frame separator
|
||||
return maybe_mod if maybe_mod.endswith("\n\n") else (maybe_mod + "\n\n")
|
||||
|
|
@ -3466,13 +3481,16 @@ class ProxyBaseLLMRequestProcessing:
|
|||
return chunk
|
||||
|
||||
@staticmethod
|
||||
def _inject_cost_into_sse_frame_str(frame_str: str, model_name: str) -> str | None:
|
||||
def _inject_cost_into_sse_frame_str(
|
||||
frame_str: str, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
|
||||
) -> str | None:
|
||||
"""
|
||||
Inject cost information into an SSE frame string by modifying the JSON in the 'data:' line.
|
||||
|
||||
Args:
|
||||
frame_str: SSE frame string that may contain multiple lines
|
||||
model_name: Model name for cost calculation
|
||||
litellm_logging_obj: The call's logging object, forwarded for pricing
|
||||
|
||||
Returns:
|
||||
Modified SSE frame string with cost injected, or None if no modification needed
|
||||
|
|
@ -3486,7 +3504,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
json_part = stripped_ln.split("data:", 1)[1].strip()
|
||||
if json_part and json_part != "[DONE]":
|
||||
obj = json.loads(json_part)
|
||||
maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(obj, model_name)
|
||||
maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(
|
||||
obj, model_name, litellm_logging_obj
|
||||
)
|
||||
if maybe_modified is not None:
|
||||
lines[idx] = "data: " + safe_dumps(maybe_modified) + ("\r" if ln.endswith("\r") else "")
|
||||
return "\n".join(lines)
|
||||
|
|
@ -3494,34 +3514,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
except Exception:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _anthropic_stream_usage_kwargs(usage: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
prompt_tokens: Final = int(usage.get("input_tokens", 0) or 0)
|
||||
completion_tokens: Final = int(usage.get("output_tokens", 0) or 0)
|
||||
total_tokens: Final = int(
|
||||
usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens)
|
||||
)
|
||||
web_search_requests: Final = usage.get("web_search_requests")
|
||||
server_tool_use: Final = (
|
||||
ServerToolUse(web_search_requests=web_search_requests) if web_search_requests is not None else None
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (
|
||||
("prompt_tokens", prompt_tokens),
|
||||
("completion_tokens", completion_tokens),
|
||||
("total_tokens", total_tokens),
|
||||
("completion_tokens_details", usage.get("completion_tokens_details")),
|
||||
("prompt_tokens_details", usage.get("prompt_tokens_details")),
|
||||
("cache_creation_input_tokens", usage.get("cache_creation_input_tokens")),
|
||||
("cache_read_input_tokens", usage.get("cache_read_input_tokens")),
|
||||
("server_tool_use", server_tool_use),
|
||||
)
|
||||
if value is not None
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _openai_stream_usage_kwargs(usage: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
prompt_tokens: Final = int(usage.get("prompt_tokens", 0) or 0)
|
||||
|
|
@ -3544,11 +3536,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _stream_usage_kwargs_for_event(obj: Mapping[str, object], usage: Mapping[str, Any]) -> Mapping[str, Any] | None:
|
||||
def _stream_usage_for_event(obj: Mapping[str, object], usage: Mapping[str, Any]) -> Usage | None:
|
||||
# Anthropic reports input_tokens excluding cache tokens, so reuse the non-streaming
|
||||
# transformation to total the prompt and keep the 5m/1h cache creation split
|
||||
if obj.get("type") == "message_delta":
|
||||
return ProxyBaseLLMRequestProcessing._anthropic_stream_usage_kwargs(usage)
|
||||
return AnthropicConfig().calculate_usage(usage_object=usage, reasoning_content=None)
|
||||
if obj.get("object") == "chat.completion.chunk":
|
||||
return ProxyBaseLLMRequestProcessing._openai_stream_usage_kwargs(usage)
|
||||
return Usage(**ProxyBaseLLMRequestProcessing._openai_stream_usage_kwargs(usage))
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -3563,7 +3557,54 @@ class ProxyBaseLLMRequestProcessing:
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def _inject_cost_into_usage_dict(obj: dict, model_name: str) -> dict | None:
|
||||
def _logging_obj_cost_or_none(
|
||||
model_response: ModelResponse, litellm_logging_obj: LiteLLMLoggingObj
|
||||
) -> float | None:
|
||||
# Pricing a frame stamps cost_breakdown and, on failure, the cost-failure debug key onto
|
||||
# the live logging object. The pass-through handlers never recompute either one, so a
|
||||
# frame-derived breakdown would outlive the stream and land in the spend log. Snapshot
|
||||
# both and put them back, so pricing here stays a read as far as the request is concerned
|
||||
breakdown_before: Final = getattr(litellm_logging_obj, "cost_breakdown", None)
|
||||
call_details: Final = getattr(litellm_logging_obj, "model_call_details", None)
|
||||
debug_key: Final = "response_cost_failure_debug_information"
|
||||
debug_missing: Final = object()
|
||||
debug_before: Final = call_details.get(debug_key, debug_missing) if isinstance(call_details, dict) else None
|
||||
try:
|
||||
cost: Final = litellm_logging_obj._response_cost_calculator(result=model_response) # pyright: ignore[reportPrivateUsage] # reuse the call's own cost calc for pricing parity with the logging callback
|
||||
except Exception: # noqa: BLE001 # a pricing failure falls back to model-name pricing instead of breaking the stream
|
||||
return None
|
||||
finally:
|
||||
if hasattr(litellm_logging_obj, "cost_breakdown"):
|
||||
litellm_logging_obj.cost_breakdown = breakdown_before
|
||||
if isinstance(call_details, dict):
|
||||
if debug_before is debug_missing:
|
||||
call_details.pop(debug_key, None)
|
||||
else:
|
||||
call_details[debug_key] = debug_before
|
||||
return float(cost) if isinstance(cost, (int, float)) and not isinstance(cost, bool) else None
|
||||
|
||||
@staticmethod
|
||||
def _streamed_usage_cost(
|
||||
model_response: ModelResponse,
|
||||
model_name: str,
|
||||
service_tier: str | None,
|
||||
litellm_logging_obj: LiteLLMLoggingObj | None,
|
||||
) -> float | None:
|
||||
# Pricing via the logging object inherits the deployment's custom pricing, so the
|
||||
# streamed cost matches what the logging callback records instead of sticker price
|
||||
cost_from_logging_obj: Final = (
|
||||
ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, litellm_logging_obj)
|
||||
if litellm_logging_obj is not None
|
||||
else None
|
||||
)
|
||||
if cost_from_logging_obj is not None:
|
||||
return cost_from_logging_obj
|
||||
return ProxyBaseLLMRequestProcessing._completion_cost_or_none(model_response, model_name, service_tier)
|
||||
|
||||
@staticmethod
|
||||
def _inject_cost_into_usage_dict(
|
||||
obj: dict, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
|
||||
) -> dict | None:
|
||||
"""
|
||||
Inject cost information into the usage object of a streamed usage event
|
||||
(Anthropic ``message_delta`` or OpenAI ``chat.completion.chunk``).
|
||||
|
|
@ -3571,6 +3612,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
Args:
|
||||
obj: Dictionary containing the SSE event data
|
||||
model_name: Model name for cost calculation
|
||||
litellm_logging_obj: The call's logging object, used for pricing
|
||||
|
||||
Returns:
|
||||
Modified dictionary with cost injected, or None if no modification needed
|
||||
|
|
@ -3578,14 +3620,15 @@ class ProxyBaseLLMRequestProcessing:
|
|||
usage: Final = obj.get("usage")
|
||||
if not isinstance(usage, dict):
|
||||
return None
|
||||
usage_kwargs: Final = ProxyBaseLLMRequestProcessing._stream_usage_kwargs_for_event(obj, usage)
|
||||
if usage_kwargs is None:
|
||||
stream_usage: Final = ProxyBaseLLMRequestProcessing._stream_usage_for_event(obj, usage)
|
||||
if stream_usage is None:
|
||||
return None
|
||||
service_tier: Final = obj.get("service_tier")
|
||||
cost_val: Final = ProxyBaseLLMRequestProcessing._completion_cost_or_none(
|
||||
ModelResponse(usage=Usage(**usage_kwargs)),
|
||||
cost_val: Final = ProxyBaseLLMRequestProcessing._streamed_usage_cost(
|
||||
ModelResponse(usage=stream_usage),
|
||||
model_name,
|
||||
service_tier if isinstance(service_tier, str) else None,
|
||||
litellm_logging_obj,
|
||||
)
|
||||
if cost_val is None:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -106,7 +106,7 @@ class PassThroughStreamingHandler:
|
|||
) # rebind-ok: SSE frame reassembly buffer across transport chunks
|
||||
if complete_frames:
|
||||
yield ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
complete_frames, resolved_model_name
|
||||
complete_frames, resolved_model_name, litellm_logging_obj
|
||||
)
|
||||
if pending:
|
||||
yield pending
|
||||
|
|
|
|||
|
|
@ -6082,6 +6082,254 @@ class TestInjectCostIntoUsageDict:
|
|||
injected = json.loads(result.split("\n")[0].split("data:", 1)[1].strip())
|
||||
assert injected["usage"]["cost"] == pytest.approx(self._expected_cost("gpt-4o-mini", 11, 4))
|
||||
|
||||
def test_message_delta_cost_charges_the_non_cached_input_tokens(self):
|
||||
"""Anthropic reports ``input_tokens`` excluding cache tokens, so reading it as the whole
|
||||
prompt total drops the non-cached input from the bill on every cache hit."""
|
||||
model = "claude-haiku-4-5"
|
||||
pricing = litellm.model_cost[model]
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {
|
||||
"input_tokens": 14,
|
||||
"output_tokens": 8,
|
||||
"cache_read_input_tokens": 3202,
|
||||
"cache_creation_input_tokens": 0,
|
||||
},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model)
|
||||
|
||||
assert result is not None
|
||||
expected = (
|
||||
14 * pricing["input_cost_per_token"]
|
||||
+ 3202 * pricing["cache_read_input_token_cost"]
|
||||
+ 8 * pricing["output_cost_per_token"]
|
||||
)
|
||||
dropped_input = expected - 14 * pricing["input_cost_per_token"]
|
||||
assert result["usage"]["cost"] == pytest.approx(expected)
|
||||
assert result["usage"]["cost"] > dropped_input
|
||||
|
||||
def test_message_delta_prices_1h_cache_creation_above_the_5m_rate(self):
|
||||
"""The ``cache_creation`` 5m/1h split has to survive into ``prompt_tokens_details``,
|
||||
otherwise a 1h write is billed at the cheaper 5m rate."""
|
||||
model = "claude-haiku-4-5"
|
||||
pricing = litellm.model_cost[model]
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {
|
||||
"input_tokens": 14,
|
||||
"output_tokens": 8,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 2000,
|
||||
"cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 2000},
|
||||
},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model)
|
||||
|
||||
assert result is not None
|
||||
base = 14 * pricing["input_cost_per_token"] + 8 * pricing["output_cost_per_token"]
|
||||
expected_1h = base + 2000 * pricing["cache_creation_input_token_cost_above_1hr"]
|
||||
flat_5m = base + 2000 * pricing["cache_creation_input_token_cost"]
|
||||
assert expected_1h != pytest.approx(flat_5m)
|
||||
assert result["usage"]["cost"] == pytest.approx(expected_1h)
|
||||
|
||||
def test_message_delta_prices_through_the_logging_obj_so_custom_pricing_applies(self):
|
||||
"""Costing by model name alone yields sticker price, so a deployment with a negotiated
|
||||
discount streamed a ``usage.cost`` that disagreed with the callback's ``response_cost``."""
|
||||
|
||||
class _StubLoggingObj:
|
||||
def __init__(self, cost):
|
||||
self._cost = cost
|
||||
self.captured_result = None
|
||||
|
||||
def _response_cost_calculator(self, result):
|
||||
self.captured_result = result
|
||||
return self._cost
|
||||
|
||||
model = "claude-haiku-4-5"
|
||||
discounted_cost = 0.00099
|
||||
stub = _StubLoggingObj(discounted_cost)
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {
|
||||
"input_tokens": 14,
|
||||
"output_tokens": 8,
|
||||
"cache_read_input_tokens": 3202,
|
||||
"cache_creation_input_tokens": 500,
|
||||
"cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 400},
|
||||
},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model, stub)
|
||||
|
||||
assert result is not None
|
||||
assert result["usage"]["cost"] == discounted_cost
|
||||
assert result["usage"]["cost"] != pytest.approx(self._expected_cost(model, 14 + 500 + 3202, 8))
|
||||
usage = stub.captured_result.usage
|
||||
assert usage.prompt_tokens == 14 + 500 + 3202
|
||||
details = usage.prompt_tokens_details.cache_creation_token_details
|
||||
assert details.ephemeral_5m_input_tokens == 100
|
||||
assert details.ephemeral_1h_input_tokens == 400
|
||||
|
||||
def test_message_delta_falls_back_to_model_pricing_when_the_logging_obj_returns_no_cost(self):
|
||||
class _StubLoggingObj:
|
||||
def _response_cost_calculator(self, result):
|
||||
return None
|
||||
|
||||
model = "claude-haiku-4-5"
|
||||
pricing = litellm.model_cost[model]
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"input_tokens": 14, "output_tokens": 8, "cache_read_input_tokens": 3202},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model, _StubLoggingObj())
|
||||
|
||||
assert result is not None
|
||||
assert result["usage"]["cost"] == pytest.approx(
|
||||
14 * pricing["input_cost_per_token"]
|
||||
+ 3202 * pricing["cache_read_input_token_cost"]
|
||||
+ 8 * pricing["output_cost_per_token"]
|
||||
)
|
||||
|
||||
def test_message_delta_falls_back_to_model_pricing_when_the_logging_obj_raises(self):
|
||||
"""A pricing failure mid-stream must not break the frame, so the raise falls back to
|
||||
model-name pricing rather than propagating into the response body."""
|
||||
|
||||
class _StubLoggingObj:
|
||||
def _response_cost_calculator(self, result):
|
||||
raise ValueError("no pricing for this deployment")
|
||||
|
||||
model = "claude-haiku-4-5"
|
||||
pricing = litellm.model_cost[model]
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"input_tokens": 14, "output_tokens": 8, "cache_read_input_tokens": 3202},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, model, _StubLoggingObj())
|
||||
|
||||
assert result is not None
|
||||
assert result["usage"]["cost"] == pytest.approx(
|
||||
14 * pricing["input_cost_per_token"]
|
||||
+ 3202 * pricing["cache_read_input_token_cost"]
|
||||
+ 8 * pricing["output_cost_per_token"]
|
||||
)
|
||||
|
||||
def test_pricing_a_frame_leaves_the_real_logging_obj_unchanged(self):
|
||||
"""Pricing runs against the live logging object, and the pass-through handlers never
|
||||
recompute cost_breakdown, so a frame-derived breakdown would reach the spend log."""
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
Logging as LiteLLMLoggingObj,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model="claude-haiku-4-5",
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="lit4902-breakdown-test",
|
||||
function_id="lit4902-breakdown-test",
|
||||
)
|
||||
logging_obj.update_environment_variables(litellm_params={}, optional_params={})
|
||||
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
|
||||
assert logging_obj.cost_breakdown is None
|
||||
|
||||
model_response = ModelResponse(
|
||||
usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224)
|
||||
)
|
||||
cost = ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, logging_obj)
|
||||
|
||||
assert cost is not None and cost > 0
|
||||
assert logging_obj.cost_breakdown is None
|
||||
assert "response_cost_failure_debug_information" not in logging_obj.model_call_details
|
||||
|
||||
def test_pricing_a_frame_restores_a_breakdown_the_request_already_had(self):
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
Logging as LiteLLMLoggingObj,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model="claude-haiku-4-5",
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="lit4902-breakdown-restore",
|
||||
function_id="lit4902-breakdown-restore",
|
||||
)
|
||||
logging_obj.update_environment_variables(litellm_params={}, optional_params={})
|
||||
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
|
||||
logging_obj.set_cost_breakdown(
|
||||
input_cost=0.5, output_cost=0.25, total_cost=0.75, cost_for_built_in_tools_cost_usd_dollar=0.0
|
||||
)
|
||||
existing = logging_obj.cost_breakdown
|
||||
|
||||
model_response = ModelResponse(
|
||||
usage=Usage(prompt_tokens=3216, completion_tokens=8, total_tokens=3224)
|
||||
)
|
||||
ProxyBaseLLMRequestProcessing._logging_obj_cost_or_none(model_response, logging_obj)
|
||||
|
||||
assert logging_obj.cost_breakdown is existing
|
||||
assert logging_obj.cost_breakdown["total_cost"] == 0.75
|
||||
|
||||
def test_openai_chunk_prices_through_the_logging_obj_so_custom_pricing_applies(self):
|
||||
"""The chat.completion.chunk path rides the same pricer, so a discounted deployment
|
||||
streaming /v1/chat/completions gets its negotiated price instead of sticker."""
|
||||
|
||||
class _StubLoggingObj:
|
||||
def __init__(self, cost):
|
||||
self._cost = cost
|
||||
self.captured_result = None
|
||||
|
||||
def _response_cost_calculator(self, result):
|
||||
self.captured_result = result
|
||||
return self._cost
|
||||
|
||||
discounted_cost = 0.00031
|
||||
stub = _StubLoggingObj(discounted_cost)
|
||||
event = {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion.chunk",
|
||||
"choices": [],
|
||||
"usage": {"prompt_tokens": 1000, "completion_tokens": 100, "total_tokens": 1100},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "gpt-4o-mini", stub)
|
||||
|
||||
assert result is not None
|
||||
assert result["usage"]["cost"] == discounted_cost
|
||||
assert result["usage"]["cost"] != pytest.approx(self._expected_cost("gpt-4o-mini", 1000, 100))
|
||||
usage = stub.captured_result.usage
|
||||
assert usage.prompt_tokens == 1000
|
||||
assert usage.completion_tokens == 100
|
||||
|
||||
def test_openai_chunk_falls_back_to_model_pricing_when_the_logging_obj_returns_no_cost(self):
|
||||
class _StubLoggingObj:
|
||||
def _response_cost_calculator(self, result):
|
||||
return None
|
||||
|
||||
event = {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion.chunk",
|
||||
"choices": [],
|
||||
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "gpt-4o-mini", _StubLoggingObj())
|
||||
|
||||
assert result is not None
|
||||
assert result["usage"]["cost"] == pytest.approx(self._expected_cost("gpt-4o-mini", 11, 4))
|
||||
|
||||
|
||||
class TestProcessChunkWithCostInjection:
|
||||
def test_complete_usage_frame_chunk_is_injected(self, monkeypatch):
|
||||
|
|
@ -6116,6 +6364,31 @@ class TestProcessChunkWithCostInjection:
|
|||
|
||||
assert ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, "gpt-4o-mini") == chunk
|
||||
|
||||
def test_message_delta_frame_is_priced_with_the_logging_obj(self, monkeypatch):
|
||||
"""Pins that the logging object reaches the pricer through the byte-frame entry point,
|
||||
which is how the proxy actually calls this on a streamed Messages API request."""
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
|
||||
class _StubLoggingObj:
|
||||
def _response_cost_calculator(self, result):
|
||||
return 0.00042
|
||||
|
||||
chunk = (
|
||||
b"event: message_delta\n"
|
||||
b'data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},'
|
||||
b'"usage":{"input_tokens":14,"output_tokens":8,"cache_read_input_tokens":3202}}\n\n'
|
||||
)
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
chunk, "claude-haiku-4-5", _StubLoggingObj()
|
||||
)
|
||||
|
||||
assert result != chunk
|
||||
data_line = next(ln for ln in result.decode("utf-8").splitlines() if ln.startswith("data:"))
|
||||
payload = json.loads(data_line.split("data:", 1)[1].strip())
|
||||
assert payload["usage"]["cost"] == 0.00042
|
||||
assert payload["usage"]["cache_read_input_tokens"] == 3202
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SSE keepalive during the time-to-first-token (issue #34819)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue