From 30a4a393933fc798c3f26ff0e8d74ca2d7cf6c0e Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 23 Aug 2026 14:25:26 +0200 Subject: [PATCH] guard spend-shaping params from strip to prevent overcharge --- routstr/upstream/request_correction.py | 33 ++++++++++++++++++++++--- tests/unit/test_request_correction.py | 34 +++++++++++++++++++++++++- 2 files changed, 63 insertions(+), 4 deletions(-) diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py index c2ea5b1d..d959a015 100644 --- a/routstr/upstream/request_correction.py +++ b/routstr/upstream/request_correction.py @@ -84,13 +84,33 @@ def extract_error_message(response: Response) -> str: return "" -def strip_unsupported_param( - body: dict, error_message: str -) -> tuple[dict, str] | None: +# Spend-shaping fields bound how much work — and therefore cost — the upstream +# may perform. The reservation was priced with these caps in place; dropping one +# 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. Returns ``(new_body, param)`` (a new dict, original untouched) when the 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) if not match: @@ -98,6 +118,13 @@ def strip_unsupported_param( param = match.group("param") if param not in body: 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} return new_body, param diff --git a/tests/unit/test_request_correction.py b/tests/unit/test_request_correction.py index 903fe16e..3b1110a2 100644 --- a/tests/unit/test_request_correction.py +++ b/tests/unit/test_request_correction.py @@ -63,7 +63,9 @@ class TestCorrectRequest: assert correct_request(_body(temperature=1), "", set()) is 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: """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: 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: def test_extracts_nested_error_message(self) -> None: