diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 702ab527..4d4696f2 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -196,9 +196,13 @@ async def calculate_discounted_max_cost( adjusted = max_cost_for_model - if messages := body.get("messages"): - prompt_tokens = estimate_tokens(messages) + messages = body.get("messages") + # Estimated over the whole body: a discount driven by message text alone lets + # a caller hide prompt weight elsewhere, shrink the reservation, and be billed + # for work the reservation never covered. + prompt_tokens = estimate_prompt_tokens(body) + if isinstance(messages, list): image_tokens = await estimate_image_tokens_in_messages(messages) if image_tokens > 0: logger.debug( @@ -210,6 +214,7 @@ async def calculate_discounted_max_cost( ) prompt_tokens += image_tokens + if prompt_tokens > 0: estimated_prompt_delta_sats = ( max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt ) @@ -262,6 +267,38 @@ def estimate_tokens(messages: list) -> int: return total // 3 +def _sum_string_chars(node: Any) -> int: + """Recursively sum the length of every string in the tree, keys included. + + Nothing is excluded. Keys count because JSON-schema property names are + forwarded to the provider, and no exclusion rule can be trusted here: every + part of the body is caller-controlled, so any carve-out (by key name or by + value shape) is a place to hide prompt weight for free. Inline image data is + therefore counted as text too, which only makes the discount smaller. + """ + if isinstance(node, str): + return len(node) + if isinstance(node, dict): + return sum( + len(str(key)) + _sum_string_chars(value) for key, value in node.items() + ) + if isinstance(node, list): + return sum(_sum_string_chars(item) for item in node) + return 0 + + +def estimate_prompt_tokens(body: dict) -> int: + """Conservatively estimate prompt tokens for the whole provider-bound body. + + Unlike ``estimate_tokens`` (message text only), this walks every field, so + prompt weight hidden in tool schemas, tool-call arguments, ``system``, or + any field forwarded in future cannot escape the reservation estimate. It + over-estimates rather than under-estimates: the result only shrinks a + discount against a reservation that settlement later refunds. + """ + return _sum_string_chars(body) // 3 + + def _get_image_dimensions(image_data: bytes) -> tuple[int, int]: """Extract image dimensions from image bytes.""" try: diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index 6809d94c..e5dbf136 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -155,3 +155,110 @@ async def test_discounted_max_cost_floors_at_min_request_msat() -> None: cost = await calculate_discounted_max_cost(150_000, body, model_obj) assert cost == 1000 + + +def test_estimate_prompt_tokens_counts_every_string_in_the_body() -> None: + from routstr.payment.helpers import estimate_prompt_tokens, estimate_tokens + + hidden = "x" * 3_000 # ~1000 tokens of prompt hidden from the text estimator + body: dict[str, Any] = { + "messages": [{"role": "user", "content": "hi"}], + "tools": [ + { + "type": "function", + "function": { + "name": "f", + "description": hidden, + "parameters": {"type": "object", "properties": {hidden: {}}}, + }, + } + ], + } + + # The text-only estimator sees almost nothing; the conservative one sees it. + assert estimate_tokens(body["messages"]) < 10 + assert estimate_prompt_tokens(body) >= 1_000 + + # No carve-out is exempt: neither a caller-chosen key name nor a caller-chosen + # value prefix can buy a discount, so both still count in full. + assert estimate_prompt_tokens({"tools": [{"data": hidden}]}) >= 1_000 + assert estimate_prompt_tokens({"system": "data:" + hidden}) >= 1_000 + + +async def test_discount_cannot_be_dodged_by_hiding_prompt_in_tools() -> None: + """A large prompt moved from messages into tool schemas must reserve the + same cost — otherwise a caller undercharges by hiding weight from the + estimator.""" + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.5 + pricing.completion = 0.01 + pricing.max_prompt_cost = 100.0 + pricing.max_completion_cost = 100.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + big_text = "word " * 2_000 + base = {"model": "test-model", "max_tokens": 10} + in_messages = { + **base, + "messages": [{"role": "user", "content": big_text}], + } + hiding_places = { + "tools": { + **base, + "messages": [{"role": "user", "content": "hi"}], + "tools": [ + {"type": "function", "function": {"name": "f", "description": big_text}} + ], + }, + # Anthropic forwards a top-level system prompt; it is billed like any other. + "system": { + **base, + "messages": [{"role": "user", "content": "hi"}], + "system": big_text, + }, + # A key named like an image field must not win an image exclusion. + "image-named key": { + **base, + "messages": [{"role": "user", "content": "hi"}], + "tools": [{"function": {"parameters": {"data": big_text}}}], + }, + # Nor may a caller-chosen "data:" prefix, in any field the body allows. + "data-prefixed content": { + **base, + "messages": [{"role": "user", "content": "data:" + big_text}], + }, + "data-prefixed text block": { + **base, + "messages": [ + { + "role": "user", + "content": [{"type": "text", "text": "data:" + big_text}], + } + ], + }, + "data-prefixed system": { + **base, + "messages": [{"role": "user", "content": "hi"}], + "system": "data:" + big_text, + }, + } + + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost_messages = await calculate_discounted_max_cost( + 150_000, in_messages, model_obj + ) + for where, body in hiding_places.items(): + cost = await calculate_discounted_max_cost(150_000, body, model_obj) + # Same prompt weight → at least the same reservation, never the floor. + assert cost >= cost_messages, where + assert cost > 1000, where