guard spend-shaping params from strip to prevent overcharge

This commit is contained in:
9qeklajc
2026-08-23 14:27:31 +02:00
parent 8520c5e458
commit 30a4a39393
2 changed files with 63 additions and 4 deletions
+30 -3
View File
@@ -84,13 +84,33 @@ def extract_error_message(response: Response) -> str:
return "" return ""
def strip_unsupported_param( # Spend-shaping fields bound how much work — and therefore cost — the upstream
body: dict, error_message: str # may perform. The reservation was priced with these caps in place; dropping one
) -> tuple[dict, str] | None: # and retrying would let the request run uncapped (or fan out) and bill above the
# caller's authorization. When the upstream names one of these, decline the strip
# and let the error propagate. Matched case-insensitively.
_SPEND_SHAPING_PARAMS = frozenset(
{
"max_tokens",
"max_completion_tokens",
"max_output_tokens",
"max_tokens_to_sample",
"n",
"best_of",
}
)
def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None:
"""Drop a top-level param the upstream named as unsupported/deprecated. """Drop a top-level param the upstream named as unsupported/deprecated.
Returns ``(new_body, param)`` (a new dict, original untouched) when the Returns ``(new_body, param)`` (a new dict, original untouched) when the
error names a top-level param present in the body, otherwise ``None``. error names a top-level param present in the body, otherwise ``None``.
Spend-shaping fields (output caps, fan-out counts) are never stripped:
removing one after the reservation was priced would uncap the retry and
overcharge. When the upstream names such a field, decline so the original
error propagates rather than silently resizing the request's cost.
""" """
match = _UNSUPPORTED_PARAM_RE.search(error_message) match = _UNSUPPORTED_PARAM_RE.search(error_message)
if not match: if not match:
@@ -98,6 +118,13 @@ def strip_unsupported_param(
param = match.group("param") param = match.group("param")
if param not in body: if param not in body:
return None return None
if param.lower() in _SPEND_SHAPING_PARAMS:
logger.warning(
"Upstream rejected spend-shaping param '%s'; refusing to strip it "
"(retrying uncapped would overcharge) — surfacing the error",
param,
)
return None
new_body = {k: v for k, v in body.items() if k != param} new_body = {k: v for k, v in body.items() if k != param}
return new_body, param return new_body, param
+33 -1
View File
@@ -63,7 +63,9 @@ class TestCorrectRequest:
assert correct_request(_body(temperature=1), "", set()) is None assert correct_request(_body(temperature=1), "", set()) is None
def test_returns_none_on_non_object_body(self) -> None: def test_returns_none_on_non_object_body(self) -> None:
assert correct_request(b"[1, 2, 3]", "`temperature` is deprecated", set()) is None assert (
correct_request(b"[1, 2, 3]", "`temperature` is deprecated", set()) is None
)
def test_deprecated_model_name_is_not_stripped_as_param(self) -> None: def test_deprecated_model_name_is_not_stripped_as_param(self) -> None:
"""A 'model is deprecated' error must not strip an unrelated body field. """A 'model is deprecated' error must not strip an unrelated body field.
@@ -106,6 +108,36 @@ class TestStripUnsupportedParam:
def test_declines_when_no_match(self) -> None: def test_declines_when_no_match(self) -> None:
assert strip_unsupported_param({"temperature": 1}, "nope") is None assert strip_unsupported_param({"temperature": 1}, "nope") is None
def test_never_strips_spend_shaping_params(self) -> None:
# Stripping an output cap after the reservation was priced would let
# the retry run uncapped and overcharge — the corrector must decline so
# the upstream error propagates instead.
for param in (
"max_tokens",
"max_completion_tokens",
"max_output_tokens",
"max_tokens_to_sample",
"n",
"best_of",
):
body = {"model": "m", param: 4, "messages": []}
assert (
strip_unsupported_param(body, f"`{param}` is not supported") is None
), param
# And through the full pipeline entry point.
assert (
correct_request(
json.dumps(body).encode(),
f"`{param}` is not supported",
set(),
)
is None
), param
def test_spend_shaping_guard_is_case_insensitive(self) -> None:
body = {"model": "m", "Max_Tokens": 4}
assert strip_unsupported_param(body, "`Max_Tokens` is deprecated") is None
class TestExtractErrorMessage: class TestExtractErrorMessage:
def test_extracts_nested_error_message(self) -> None: def test_extracts_nested_error_message(self) -> None: