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:
mateo-berri 2026-08-18 14:52:23 -07:00
parent be594f5984
commit 354b0c3a45
8 changed files with 187 additions and 19 deletions

View file

@ -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:

View file

@ -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)
},
)

View file

@ -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

View file

@ -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}

View file

@ -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

View file

@ -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}

View file

@ -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

View file

@ -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": {