mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
guard spend-shaping params from strip to prevent overcharge
This commit is contained in:
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user