mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(guardrails): bill all chunks on mid-chunking block, strip client guardrail cost metadata, add cost map schema keys
A blocked chunk now logs the summed usage and cost of every ApplyGuardrail call AWS billed for the logical request, not just the blocking chunk. Client-supplied metadata.standard_logging_guardrail_information is stripped at the proxy boundary so callers cannot forge (even negative) guardrail cost into spend, and guardrail_information_cost ignores negative or non-finite entry costs as defense in depth. The cost map schema test now allows guardrail_cost_per_unit and the guardrail mode.
This commit is contained in:
parent
be594f5984
commit
354b0c3a45
8 changed files with 187 additions and 19 deletions
|
|
@ -1,3 +1,4 @@
|
|||
import math
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -46,6 +47,13 @@ def bedrock_guardrail_cost(usage_units: Mapping[str, int], aws_region_name: str
|
|||
return sum(units * pricing.guardrail_cost_per_unit.get(counter, 0.0) for counter, units in usage_units.items())
|
||||
|
||||
|
||||
def _billable_entry_cost(entry: GuardrailCostEntry) -> float:
|
||||
cost: Final = entry.guardrail_cost
|
||||
if cost is None or not math.isfinite(cost) or cost <= 0.0:
|
||||
return 0.0
|
||||
return cost
|
||||
|
||||
|
||||
def guardrail_information_cost(guardrail_information: object) -> float:
|
||||
try:
|
||||
parsed: Final = _GUARDRAIL_INFORMATION_ADAPTER.validate_python(guardrail_information)
|
||||
|
|
@ -54,8 +62,8 @@ def guardrail_information_cost(guardrail_information: object) -> float:
|
|||
if parsed is None:
|
||||
return 0.0
|
||||
if isinstance(parsed, GuardrailCostEntry):
|
||||
return parsed.guardrail_cost or 0.0
|
||||
return sum(entry.guardrail_cost or 0.0 for entry in parsed)
|
||||
return _billable_entry_cost(parsed)
|
||||
return sum(_billable_entry_cost(entry) for entry in parsed)
|
||||
|
||||
|
||||
def cost_breakdown_with_guardrail(cost_breakdown: CostBreakdown | None, guardrail_cost: float) -> CostBreakdown | None:
|
||||
|
|
|
|||
|
|
@ -873,6 +873,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
credentials, aws_region_name = self._load_credentials()
|
||||
allow_chunking: Final = not self._content_uses_contextual_grounding(content)
|
||||
|
||||
completed_chunk_usages: Final[list[BedrockGuardrailUsage]] = [] # mutable-ok: billed-chunk usage accumulator
|
||||
try:
|
||||
responses: Final = await self._apply_guardrail_content_with_chunking(
|
||||
content=content,
|
||||
|
|
@ -884,6 +885,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
event_type=event_type,
|
||||
start_time=start_time,
|
||||
allow_chunking=allow_chunking,
|
||||
completed_chunk_usages=completed_chunk_usages,
|
||||
)
|
||||
except HTTPException as exc:
|
||||
if not isinstance(exc.detail, dict):
|
||||
|
|
@ -915,6 +917,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
event_type: GuardrailEventHooks,
|
||||
start_time: "datetime",
|
||||
allow_chunking: bool,
|
||||
completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: billed-chunk usage accumulator
|
||||
) -> tuple[BedrockContentChunkResult, ...]:
|
||||
"""Post `content` to ApplyGuardrail, chunking only if AWS rejects it as too large.
|
||||
|
||||
|
|
@ -961,6 +964,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
request_data=request_data,
|
||||
event_type=event_type,
|
||||
start_time=start_time,
|
||||
completed_chunk_usages=completed_chunk_usages,
|
||||
)
|
||||
return (
|
||||
BedrockContentChunkResult(
|
||||
|
|
@ -991,6 +995,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
event_type=event_type,
|
||||
start_time=start_time,
|
||||
allow_chunking=allow_chunking,
|
||||
completed_chunk_usages=completed_chunk_usages,
|
||||
)
|
||||
for batch in batches
|
||||
]
|
||||
|
|
@ -1017,6 +1022,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
event_type=event_type,
|
||||
start_time=start_time,
|
||||
allow_chunking=allow_chunking,
|
||||
completed_chunk_usages=completed_chunk_usages,
|
||||
)
|
||||
second_results: Final = await self._apply_guardrail_content_with_chunking(
|
||||
content=second_half,
|
||||
|
|
@ -1028,6 +1034,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
event_type=event_type,
|
||||
start_time=start_time,
|
||||
allow_chunking=allow_chunking,
|
||||
completed_chunk_usages=completed_chunk_usages,
|
||||
)
|
||||
combined_results: Final = tuple(first_results) + tuple(second_results)
|
||||
if is_single_item_text_split:
|
||||
|
|
@ -1047,6 +1054,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
|
||||
event_type: GuardrailEventHooks,
|
||||
start_time: "datetime",
|
||||
completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: passed through to the single-call layer
|
||||
) -> BedrockGuardrailResponse:
|
||||
"""Post one ApplyGuardrail call for `content`, retrying with exponential
|
||||
backoff on AWS ThrottlingException (HTTP 429).
|
||||
|
|
@ -1074,6 +1082,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
request_data=request_data,
|
||||
event_type=event_type,
|
||||
start_time=start_time,
|
||||
completed_chunk_usages=completed_chunk_usages,
|
||||
)
|
||||
except HTTPException as exc:
|
||||
if (
|
||||
|
|
@ -1095,6 +1104,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
|
||||
event_type: GuardrailEventHooks,
|
||||
start_time: "datetime",
|
||||
completed_chunk_usages: list[BedrockGuardrailUsage], # mutable-ok: billed-chunk usage accumulator
|
||||
) -> BedrockGuardrailResponse:
|
||||
"""Make exactly one signed ApplyGuardrail HTTP call for `content` and
|
||||
parse the result. Raises HTTPException on a guardrail block or any
|
||||
|
|
@ -1110,7 +1120,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
|
||||
A block is logged here rather than by the caller: it ends the whole chunking
|
||||
flow immediately, with no further chunks attempted, so there is no later
|
||||
merged response for the caller to log instead.
|
||||
merged response for the caller to log instead. The logged usage still spans
|
||||
the whole logical request: chunks that passed before the block appended what
|
||||
AWS billed them to ``completed_chunk_usages``, and the attempt log sums those
|
||||
with the blocking call's own usage.
|
||||
"""
|
||||
bedrock_request_data: Final = { # mutable-ok: outbound JSON request body
|
||||
**base_request_data,
|
||||
|
|
@ -1154,10 +1167,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
event_type=event_type,
|
||||
start_time=start_time,
|
||||
aws_region_name=aws_region_name,
|
||||
completed_chunk_usages=completed_chunk_usages,
|
||||
)
|
||||
raise self._get_http_exception_for_blocked_guardrail(
|
||||
bedrock_guardrail_response, request_data=request_data
|
||||
)
|
||||
response_usage: Final = bedrock_guardrail_response.get("usage")
|
||||
if isinstance(response_usage, dict):
|
||||
completed_chunk_usages.append(
|
||||
response_usage
|
||||
) # rebind-ok: accumulator threaded from make_bedrock_api_request, recording this billed call
|
||||
return bedrock_guardrail_response
|
||||
|
||||
status_code, detail_message = self._parse_bedrock_guardrail_error_response(httpx_response)
|
||||
|
|
@ -1176,16 +1195,30 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
event_type: GuardrailEventHooks,
|
||||
start_time: "datetime",
|
||||
aws_region_name: str | None,
|
||||
completed_chunk_usages: Sequence[BedrockGuardrailUsage],
|
||||
) -> None:
|
||||
"""Log a single ApplyGuardrail HTTP attempt as-is (its own status,
|
||||
derived from its own response). Used only for the blocked-content
|
||||
case, which ends the whole chunking flow immediately."""
|
||||
"""Log the blocking ApplyGuardrail attempt, which ends the whole chunking
|
||||
flow immediately. Its status derives from its own response, but its usage
|
||||
(and so its cost) spans every billed call of the logical request: the
|
||||
chunks that passed before the block plus the blocking call itself."""
|
||||
blocking_usage: Final = json_response.get("usage")
|
||||
billed_usages: Final[tuple[BedrockGuardrailUsage, ...]] = tuple(completed_chunk_usages) + (
|
||||
(blocking_usage,) if isinstance(blocking_usage, dict) else ()
|
||||
)
|
||||
logged_json_response: Final = (
|
||||
{ # mutable-ok: raw AWS JSON payload carrying the total billed usage
|
||||
**json_response,
|
||||
"usage": self._sum_usage_counters(billed_usages),
|
||||
}
|
||||
if completed_chunk_usages
|
||||
else json_response
|
||||
)
|
||||
tracing_detail: Final = self._build_tracing_detail(
|
||||
BedrockGuardrailResponse(**json_response), aws_region_name=aws_region_name
|
||||
BedrockGuardrailResponse(**logged_json_response), aws_region_name=aws_region_name
|
||||
)
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_provider=self.guardrail_provider,
|
||||
guardrail_json_response=json_response,
|
||||
guardrail_json_response=logged_json_response,
|
||||
request_data=request_data or {}, # mutable-ok: logging helper requires a dict
|
||||
guardrail_status=self._get_bedrock_guardrail_response_status(response=httpx_response),
|
||||
start_time=start_time.timestamp(),
|
||||
|
|
@ -1511,15 +1544,20 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
Keys are taken from the responses rather than from a fixed list, so a counter
|
||||
this code does not know about (AWS has added several) is still summed and
|
||||
reported instead of being silently dropped to zero."""
|
||||
chunk_usages: Final = tuple(
|
||||
chunk_result.response.get("usage") or {} # mutable-ok: read-only empty fallback
|
||||
for chunk_result in chunk_results
|
||||
return BedrockGuardrail._sum_usage_counters(
|
||||
tuple(
|
||||
chunk_result.response.get("usage") or {} # mutable-ok: read-only empty fallback
|
||||
for chunk_result in chunk_results
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _sum_usage_counters(usages: Sequence[BedrockGuardrailUsage]) -> BedrockGuardrailUsage:
|
||||
return cast( # cast-ok: TypedDict assembled from a comprehension
|
||||
BedrockGuardrailUsage,
|
||||
{ # mutable-ok: builds the TypedDict payload
|
||||
key: sum(usage.get(key) or 0 for usage in chunk_usages)
|
||||
for key in dict.fromkeys(key for usage in chunk_usages for key in usage)
|
||||
key: sum(usage.get(key) or 0 for usage in usages)
|
||||
for key in dict.fromkeys(key for usage in usages for key in usage)
|
||||
},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -296,7 +296,10 @@ _ALLOW_CLIENT_MESSAGE_REDACTION_OPT_OUT_METADATA_KEY: Final = "allow_client_mess
|
|||
_CLIENT_PRICING_CONTROL_FIELDS: Final = frozenset(CustomPricingLiteLLMParams.model_fields.keys())
|
||||
# ``model_info`` carries the same pricing fields when read by
|
||||
# ``use_custom_pricing_for_model``; strip from metadata for the same reason.
|
||||
_CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info"})
|
||||
# ``standard_logging_guardrail_information`` is proxy-written telemetry summed
|
||||
# into response_cost and spend; a client seeding it forges (even negative)
|
||||
# guardrail cost.
|
||||
_CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logging_guardrail_information"})
|
||||
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override"
|
||||
|
||||
# Request fields whose value, when URL-valued, becomes the outbound destination
|
||||
|
|
|
|||
|
|
@ -89,6 +89,17 @@ def test_guardrail_information_cost_single_entry_and_garbage():
|
|||
assert guardrail_information_cost([{"guardrail_cost": "bad"}]) == 0.0
|
||||
|
||||
|
||||
def test_guardrail_information_cost_ignores_negative_and_non_finite():
|
||||
entries = [
|
||||
{"guardrail_name": "forged-negative", "guardrail_cost": -0.005},
|
||||
{"guardrail_name": "forged-nan", "guardrail_cost": float("nan")},
|
||||
{"guardrail_name": "forged-inf", "guardrail_cost": float("inf")},
|
||||
{"guardrail_name": "real", "guardrail_cost": 0.0003},
|
||||
]
|
||||
assert guardrail_information_cost(entries) == pytest.approx(0.0003)
|
||||
assert guardrail_information_cost({"guardrail_cost": -1.0}) == 0.0
|
||||
|
||||
|
||||
def test_cost_breakdown_with_guardrail_merges_and_creates():
|
||||
assert cost_breakdown_with_guardrail(None, 0.0) is None
|
||||
untouched = {"input_cost": 0.1, "total_cost": 0.4}
|
||||
|
|
|
|||
|
|
@ -4869,8 +4869,7 @@ def _guardrail_kwargs(response_cost):
|
|||
|
||||
|
||||
def test_payload_response_cost_includes_guardrail_cost(logging_obj):
|
||||
"""LIT-5651: guardrail invocations billed by the provider must count in
|
||||
response_cost so spend and budget enforcement see them like token cost."""
|
||||
"""LIT-5651: provider-billed guardrail cost must count in response_cost."""
|
||||
payload = _build_success_payload(logging_obj, _guardrail_kwargs(response_cost=0.0000429))
|
||||
|
||||
assert payload is not None
|
||||
|
|
|
|||
|
|
@ -5080,9 +5080,7 @@ async def test_apply_guardrail_failure_logs_a_dict_not_a_bare_string():
|
|||
|
||||
|
||||
def test_build_tracing_detail_surfaces_usage_counters_and_cost(monkeypatch):
|
||||
"""LIT-5650/LIT-5651: the billable usage block Bedrock returns per ApplyGuardrail
|
||||
call must land on the tracing detail as guardrail_usage, priced into
|
||||
guardrail_cost, so spend logs and budgets see what AWS bills."""
|
||||
"""LIT-5650/LIT-5651: AWS-billed usage must land as guardrail_usage priced into guardrail_cost."""
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"model_cost",
|
||||
|
|
@ -5119,3 +5117,84 @@ def test_build_tracing_detail_omits_guardrail_usage_when_bedrock_reports_none():
|
|||
):
|
||||
assert "guardrail_usage" not in detail
|
||||
assert "guardrail_cost" not in detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocked_chunk_logs_usage_and_cost_of_prior_passed_chunks(monkeypatch):
|
||||
"""LIT-5651 regression: a block on a later chunk must still bill the chunks AWS already processed."""
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"model_cost",
|
||||
{
|
||||
"bedrock/guardrails": {
|
||||
"guardrail_cost_per_unit": {
|
||||
"contentPolicyUnits": 0.00015,
|
||||
"wordPolicyUnits": 0.0,
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrailIdentifier="test-guardrail",
|
||||
guardrailVersion="DRAFT",
|
||||
chunk_budget_chars=40,
|
||||
)
|
||||
|
||||
too_large_response = MagicMock()
|
||||
too_large_response.status_code = 429
|
||||
too_large_response.json.return_value = {
|
||||
"message": "Input text size (60 text units) exceeds the maximum allowed (1 text units) for the content filter policy"
|
||||
}
|
||||
|
||||
passed_chunk_response = MagicMock()
|
||||
passed_chunk_response.status_code = 200
|
||||
passed_chunk_response.json.return_value = {
|
||||
"action": "NONE",
|
||||
"outputs": [],
|
||||
"assessments": [],
|
||||
"usage": {"contentPolicyUnits": 2, "wordPolicyUnits": 1},
|
||||
}
|
||||
|
||||
blocked_chunk_response = MagicMock()
|
||||
blocked_chunk_response.status_code = 200
|
||||
blocked_chunk_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"assessments": [{"contentPolicy": {"filters": [{"type": "HATE", "confidence": "HIGH", "action": "BLOCKED"}]}}],
|
||||
"outputs": [{"text": "Content blocked"}],
|
||||
"usage": {"contentPolicyUnits": 3},
|
||||
}
|
||||
|
||||
mock_credentials = MagicMock()
|
||||
mock_credentials.access_key = "test-access-key"
|
||||
mock_credentials.secret_key = "test-secret-key"
|
||||
mock_credentials.token = None
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [
|
||||
{"role": "user", "content": "a" * 30},
|
||||
{"role": "user", "content": "b" * 30},
|
||||
],
|
||||
}
|
||||
|
||||
with (
|
||||
patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post,
|
||||
patch.object(guardrail, "_load_credentials", return_value=(mock_credentials, "us-east-1")),
|
||||
patch.object(guardrail, "_prepare_request", return_value=MagicMock()),
|
||||
):
|
||||
mock_post.side_effect = [too_large_response, passed_chunk_response, blocked_chunk_response]
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
await guardrail.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=request_data["messages"],
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
assert mock_post.call_count == 3
|
||||
logged_entries = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert len(logged_entries) == 1
|
||||
logged = logged_entries[0]
|
||||
assert logged["guardrail_usage"] == {"contentPolicyUnits": 5, "wordPolicyUnits": 1}
|
||||
assert logged["guardrail_cost"] == pytest.approx(0.00075)
|
||||
assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 5, "wordPolicyUnits": 1}
|
||||
|
|
|
|||
|
|
@ -102,6 +102,30 @@ class TestStripClientPricingOverrides:
|
|||
assert data["metadata"] == {"user_session": "keep-me"}
|
||||
assert data["litellm_metadata"] == {}
|
||||
|
||||
def test_metadata_guardrail_information_dropped(self):
|
||||
# Client-seeded guardrail entries would otherwise be summed into
|
||||
# response_cost and spend, letting a caller forge (even negative)
|
||||
# guardrail cost against their own budget.
|
||||
data = {
|
||||
"model": "gpt-4",
|
||||
"metadata": {
|
||||
"user_session": "keep-me",
|
||||
"standard_logging_guardrail_information": [
|
||||
{
|
||||
"guardrail_name": "forged",
|
||||
"guardrail_status": "success",
|
||||
"guardrail_cost": -0.005,
|
||||
}
|
||||
],
|
||||
},
|
||||
"litellm_metadata": {
|
||||
"standard_logging_guardrail_information": [{"guardrail_cost": 5.0}],
|
||||
},
|
||||
}
|
||||
_strip_client_pricing_overrides(data)
|
||||
assert data["metadata"] == {"user_session": "keep-me"}
|
||||
assert data["litellm_metadata"] == {}
|
||||
|
||||
def test_non_pricing_fields_untouched(self):
|
||||
data = {
|
||||
"model": "gpt-4",
|
||||
|
|
@ -129,6 +153,7 @@ class TestStripClientPricingOverrides:
|
|||
|
||||
def test_metadata_field_set_contains_model_info(self):
|
||||
assert "model_info" in _CLIENT_PRICING_METADATA_FIELDS
|
||||
assert "standard_logging_guardrail_information" in _CLIENT_PRICING_METADATA_FIELDS
|
||||
|
||||
def test_strip_emits_debug_log_listing_dropped_fields(self, caplog):
|
||||
# Operators need a paper trail so they can diagnose why a previously
|
||||
|
|
|
|||
|
|
@ -869,6 +869,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"container",
|
||||
"image_edit",
|
||||
"embedding",
|
||||
"guardrail",
|
||||
"image_generation",
|
||||
"video_generation",
|
||||
"moderation",
|
||||
|
|
@ -976,6 +977,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"type": "string",
|
||||
},
|
||||
},
|
||||
"guardrail_cost_per_unit": {
|
||||
"type": "object",
|
||||
"additionalProperties": {"type": "number"},
|
||||
},
|
||||
"search_context_cost_per_query": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue