From 63a7227c2abd431a999b74b97968c262bc454042 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 12 Sep 2026 14:16:40 +0200 Subject: [PATCH] use litellm pre-existing functionality --- routstr/payment/helpers.py | 192 ++++++------------- routstr/payment/responses_input.py | 79 ++++++++ tests/unit/test_payment_helpers.py | 296 ++++++++++++++++++++++------- 3 files changed, 363 insertions(+), 204 deletions(-) create mode 100644 routstr/payment/responses_input.py diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 05304684..d088151f 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -24,6 +24,12 @@ from ..wallet import ( deserialize_token_from_string, is_trusted_source_mint, ) +from .responses_input import ( + FILE_ID_URL_PREFIX, + count_input_images, + input_image_part_to_image_url, + responses_input_to_messages, +) logger = get_logger(__name__) @@ -237,8 +243,12 @@ async def calculate_discounted_max_cost( if isinstance(messages, list): image_tokens += await estimate_image_tokens_in_messages(messages) input_data = body.get("input") - if input_data is not None: - image_tokens += await estimate_image_tokens_from_input(input_data) + if isinstance(input_data, list): + converted = responses_input_to_messages(input_data) + if converted is None: + image_tokens += count_input_images(input_data) * _MAX_ORIGINAL_IMAGE_TOKENS + else: + image_tokens += await estimate_image_tokens_in_messages(converted) if image_tokens > 0: logger.debug( "Found images in request", @@ -553,17 +563,9 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: continue content_type = content_item.get("type") - if content_type not in ("image_url", "input_image"): - continue - - # Responses-style ``input_image`` parts carry their detail and - # file_id as siblings of the image reference; route them through - # the input_image estimator so original detail / file_id are - # honored on the chat path too. if content_type == "input_image": - total_image_tokens += await _estimate_input_image_tokens( - content_item - ) + content_item = input_image_part_to_image_url(content_item) + elif content_type != "image_url": continue image_url_data = content_item.get("image_url") @@ -575,7 +577,7 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: detail = "auto" elif isinstance(image_url_data, dict): url = image_url_data.get("url", "") - detail = image_url_data.get("detail", "auto") + detail = image_url_data.get("detail") or "auto" else: continue @@ -583,143 +585,65 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: continue if url.startswith("data:image/"): - try: - header, base64_data = url.split(",", 1) - image_bytes = base64.b64decode(base64_data) - width, height = _get_image_dimensions(image_bytes) - tokens = _calculate_image_tokens(width, height, detail) - total_image_tokens += tokens - logger.debug( - "Calculated tokens for base64 image", - extra={ - "width": width, - "height": height, - "detail": detail, - "tokens": tokens, - }, - ) - except Exception as e: - logger.warning( - "Failed to process base64 image", - extra={"error": str(e)}, - ) - total_image_tokens += 85 + total_image_tokens += _data_url_image_tokens(url, detail) + elif url.startswith(FILE_ID_URL_PREFIX): + total_image_tokens += _worst_case_image_tokens(detail) elif fetches >= IMAGE_FETCH_MAX_PER_REQUEST: logger.warning( "Skipping image URL fetch above per-request limit", extra={"url": url[:100], "limit": IMAGE_FETCH_MAX_PER_REQUEST}, ) - total_image_tokens += 85 + total_image_tokens += _worst_case_image_tokens(detail) else: fetches += 1 image_bytes_or_none = await _fetch_image_from_url(url) - if image_bytes_or_none: - width, height = _get_image_dimensions(image_bytes_or_none) - tokens = _calculate_image_tokens(width, height, detail) - total_image_tokens += tokens - logger.debug( - "Calculated tokens for URL image", - extra={ - "url": url[:100], - "width": width, - "height": height, - "detail": detail, - "tokens": tokens, - }, - ) - else: - total_image_tokens += 85 + total_image_tokens += _image_bytes_tokens( + image_bytes_or_none, detail, source=url[:100] + ) return total_image_tokens -async def _estimate_input_image_tokens(item: dict) -> int: - """Estimate tokens for a Responses API ``input_image`` item. - - Honors the item-level ``detail``. Data-URL images are measured from their - decoded bytes; remote URLs are fetched and measured like the chat path. - Only ``file_id`` references (whose dimensions cannot be fetched here) and - unfetchable/broken images fall back to conservative estimates: the - max-size tile math for high/auto and the 30,000-patch worst case (36,000 - tokens) for original, so we never under-reserve. - """ - detail = item.get("detail") or "auto" - image_url = item.get("image_url") - if isinstance(image_url, dict): - image_url = image_url.get("url", "") - - def _worst_case() -> int: - # Dimensions unknown: reserve the worst case for the declared detail so - # a broken/unreadable image still covers what the upstream could bill. - if detail == "original": - return _MAX_ORIGINAL_IMAGE_TOKENS - return _calculate_image_tokens(2048, 2048, detail) - - if isinstance(image_url, str) and image_url: - if image_url.startswith("data:image/"): - image_bytes = None - try: - _, base64_data = image_url.split(",", 1) - image_bytes = base64.b64decode(base64_data, validate=True) - except Exception as e: - logger.warning( - "Failed to decode base64 image", extra={"error": str(e)} - ) - if image_bytes is not None: - try: - img = Image.open(BytesIO(image_bytes)) - return _calculate_image_tokens(img.size[0], img.size[1], detail) - except Exception as e: - logger.warning( - "Failed to read image dimensions", extra={"error": str(e)} - ) - # Undecodable / unreadable data URL: reserve the worst case. - return _worst_case() - # Remote URL: fetch and measure like the chat path so a small image - # does not reserve the original-detail worst case. - image_bytes = await _fetch_image_from_url(image_url) - if image_bytes: - try: - img = Image.open(BytesIO(image_bytes)) - return _calculate_image_tokens(img.size[0], img.size[1], detail) - except Exception as e: - logger.warning( - "Failed to read image dimensions", extra={"error": str(e)} - ) - # Unfetchable or unreadable: fall through to the conservative estimate. - - if item.get("file_id") or image_url: - return _worst_case() - return 0 +def _worst_case_image_tokens(detail: str) -> int: + """Dimensions unknown: reserve the most ``detail`` can bill.""" + if detail == "original": + return _MAX_ORIGINAL_IMAGE_TOKENS + return _calculate_image_tokens(2048, 2048, detail) -async def estimate_image_tokens_from_input(input_data: Any) -> int: - """Estimate total tokens for images embedded in a Responses API ``input``. +def _data_url_image_tokens(url: str, detail: str) -> int: + try: + _, base64_data = url.split(",", 1) + image_bytes = base64.b64decode(base64_data, validate=True) + except Exception as e: + logger.warning("Failed to decode base64 image", extra={"error": str(e)}) + return _worst_case_image_tokens(detail) + return _image_bytes_tokens(image_bytes, detail, source="data-url") - Recognizes ``input_image`` items at the top level of the input list and - inside ``message`` content parts. - """ - if not isinstance(input_data, list): - return 0 - total_image_tokens = 0 - for item in input_data: - if not isinstance(item, dict): - continue - - if item.get("type") == "input_image": - total_image_tokens += await _estimate_input_image_tokens(item) - continue - - content = item.get("content") - if not isinstance(content, list): - continue - - for part in content: - if isinstance(part, dict) and part.get("type") == "input_image": - total_image_tokens += await _estimate_input_image_tokens(part) - - return total_image_tokens +def _image_bytes_tokens(image_bytes: bytes | None, detail: str, source: str) -> int: + if not image_bytes: + return _worst_case_image_tokens(detail) + try: + width, height = Image.open(BytesIO(image_bytes)).size + except Exception as e: + logger.warning( + "Failed to read image dimensions", + extra={"error": str(e), "source": source}, + ) + return _worst_case_image_tokens(detail) + tokens = _calculate_image_tokens(width, height, detail) + logger.debug( + "Calculated image tokens", + extra={ + "source": source, + "width": width, + "height": height, + "detail": detail, + "tokens": tokens, + }, + ) + return tokens def create_error_response( diff --git a/routstr/payment/responses_input.py b/routstr/payment/responses_input.py new file mode 100644 index 00000000..6a90a1e0 --- /dev/null +++ b/routstr/payment/responses_input.py @@ -0,0 +1,79 @@ +"""Convert a Responses API ``input`` into chat ``messages`` via litellm. + +litellm drops ``file_id`` (emits ``url: ""``) and nests a dict-form ``image_url`` +as-is, so ``input_image`` parts are flattened to ``{image_url: str, detail}`` first. +``file_id`` becomes a sentinel URL the image walker treats as unfetchable. +""" + +from typing import Any + +from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, +) + +from ..core import get_logger + +logger = get_logger(__name__) + +FILE_ID_URL_PREFIX = "file-id:" + + +def _flatten_input_image(part: dict[str, Any]) -> tuple[str, str]: + raw = part.get("image_url") + url = raw.get("url", "") if isinstance(raw, dict) else raw + detail = part.get("detail") or ( + raw.get("detail") if isinstance(raw, dict) else None + ) + if not url and part.get("file_id"): + url = f"{FILE_ID_URL_PREFIX}{part['file_id']}" + return (url if isinstance(url, str) else ""), (detail or "auto") + + +def _normalize_item(item: Any) -> Any: + if not isinstance(item, dict): + return item + if item.get("type") == "input_image": + url, detail = _flatten_input_image(item) + return {**item, "image_url": url, "detail": detail} + content = item.get("content") + if isinstance(content, list): + return {**item, "content": [_normalize_item(part) for part in content]} + return item + + +def input_image_part_to_image_url(part: dict[str, Any]) -> dict[str, Any]: + """Reshape an ``input_image`` part found inside chat ``messages``.""" + url, detail = _flatten_input_image(part) + return {"type": "image_url", "image_url": {"url": url, "detail": detail}} + + +def count_input_images(input_data: Any) -> int: + if isinstance(input_data, dict): + own = 1 if input_data.get("type") == "input_image" else 0 + return own + count_input_images(input_data.get("content")) + if isinstance(input_data, list): + return sum(count_input_images(item) for item in input_data) + return 0 + + +def responses_input_to_messages(input_data: Any) -> list[dict[str, Any]] | None: + """Returns ``None`` when the transform fails so the caller can worst-case.""" + if isinstance(input_data, str): + return [{"role": "user", "content": input_data}] + if not isinstance(input_data, list): + return [] + try: + normalized = [_normalize_item(item) for item in input_data] + converted = ( + LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=normalized, # type: ignore[arg-type] + responses_api_request={}, + ) + ) + return [dict(message) for message in converted] + except Exception as e: + logger.warning( + "Responses input transform failed; using conservative image fallback", + extra={"error": str(e)}, + ) + return None diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index 84ab8959..343d9cdc 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -291,12 +291,6 @@ async def test_discount_cannot_be_dodged_by_hiding_prompt_in_tools() -> None: async def test_discounted_max_cost_counts_responses_input_images() -> None: - """A Responses ``input_image`` must add image tokens to the reservation. - - Regression for the review finding that ``estimate_image_tokens_from_input`` - was defined but never called: a Responses body carries ``input``, not - ``messages``, so its images were previously reserved at zero tokens. - """ import base64 from io import BytesIO @@ -342,7 +336,9 @@ async def test_discounted_max_cost_counts_responses_input_images() -> None: patch.object(settings, "tolerance_percentage", 0), patch.object(settings, "min_request_msat", 1000), ): - cost_no_image = await calculate_discounted_max_cost(100_000, no_image, model_obj) + cost_no_image = await calculate_discounted_max_cost( + 100_000, no_image, model_obj + ) cost_with_image = await calculate_discounted_max_cost( 100_000, with_image, model_obj ) @@ -352,61 +348,165 @@ async def test_discounted_max_cost_counts_responses_input_images() -> None: assert cost_with_image > cost_no_image -async def test_estimate_input_image_tokens_remote_original_fetches() -> None: - """A remote ``original`` image is fetched and measured, not worst-cased. +def _responses_image(url: str, detail: str | None = "original") -> list[dict[str, Any]]: + return [ + { + "role": "user", + "content": [{"type": "input_image", "image_url": url, "detail": detail}], + } + ] - Regression for the review finding that any non-data URL with - ``detail: \"original\"`` reserved the 36,000-token worst case without - trying to fetch — a 512x512 image reserved ~117x its real cost. - """ + +async def _responses_image_tokens(input_data: list[dict[str, Any]]) -> int: + from routstr.payment.helpers import estimate_image_tokens_in_messages + from routstr.payment.responses_input import responses_input_to_messages + + messages = responses_input_to_messages(input_data) + assert messages is not None + return await estimate_image_tokens_in_messages(messages) + + +async def test_remote_original_image_is_fetched_not_worst_cased() -> None: from io import BytesIO - from unittest.mock import patch as mock_patch from PIL import Image - from routstr.payment.helpers import _estimate_input_image_tokens - image = Image.new("RGB", (512, 512), "red") buffer = BytesIO() image.save(buffer, format="JPEG") image_bytes = buffer.getvalue() - with mock_patch( + with patch( "routstr.payment.helpers._fetch_image_from_url", new=AsyncMock(return_value=image_bytes), ): - # 512x512 original -> 16x16 = 256 patches -> ceil(256 * 1.2) = 308 tokens, - # far below the 36,000 worst case a blind fallback would reserve. - assert await _estimate_input_image_tokens( - {"type": "input_image", "image_url": "https://x.test/i.jpg", "detail": "original"} - ) == 308 + # 256 patches * 1.2 + assert ( + await _responses_image_tokens(_responses_image("https://x.test/i.jpg")) + == 308 + ) - # When the fetch fails, fall back to the original-detail worst case. - with mock_patch( + with patch( "routstr.payment.helpers._fetch_image_from_url", new=AsyncMock(return_value=None), ): - assert await _estimate_input_image_tokens( - {"type": "input_image", "image_url": "https://x.test/i.jpg", "detail": "original"} - ) == 36_000 + assert ( + await _responses_image_tokens(_responses_image("https://x.test/i.jpg")) + == 36_000 + ) -async def test_estimate_input_image_tokens_broken_original_data_url() -> None: - """A broken data URL with ``detail: \"original\"`` reserves the worst case. +async def test_broken_data_url_reserves_declared_detail_worst_case() -> None: + from routstr.payment.helpers import estimate_image_tokens_in_messages - Regression for the review finding that the ``except`` branch returned 85 - (low-detail) regardless of the declared detail. - """ - from routstr.payment.helpers import _estimate_input_image_tokens + broken = "data:image/jpeg;base64,!!!" + assert await _responses_image_tokens(_responses_image(broken)) == 36_000 + assert await _responses_image_tokens(_responses_image(broken, "high")) == 85 + ( + 170 * 4 + ) - # "!!!" is not valid base64, so decoding raises before any dimension read. - assert await _estimate_input_image_tokens( - {"type": "input_image", "image_url": "data:image/jpeg;base64,!!!", "detail": "original"} - ) == 36_000 - # Non-original details fall back to the max-size tile math, not the 85 floor. - assert await _estimate_input_image_tokens( - {"type": "input_image", "image_url": "data:image/jpeg;base64,!!!", "detail": "high"} - ) == 85 + (170 * 4) + chat = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": broken, "detail": "original"}, + } + ], + } + ] + assert await estimate_image_tokens_in_messages(chat) == 36_000 + + +async def test_chat_original_image_fetch_failure_reserves_worst_case() -> None: + from routstr.payment.helpers import estimate_image_tokens_in_messages + + chat = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": "https://x.test/i.jpg", "detail": "original"}, + } + ], + } + ] + with patch( + "routstr.payment.helpers._fetch_image_from_url", + new=AsyncMock(return_value=None), + ): + assert await estimate_image_tokens_in_messages(chat) == 36_000 + + +async def test_responses_images_share_per_request_fetch_cap() -> None: + from routstr.payment.helpers import IMAGE_FETCH_MAX_PER_REQUEST + + input_data = [ + { + "role": "user", + "content": [ + { + "type": "input_image", + "image_url": f"https://x.test/{i}.jpg", + "detail": "original", + } + for i in range(IMAGE_FETCH_MAX_PER_REQUEST + 1) + ], + } + ] + fetch = AsyncMock(return_value=None) + with patch("routstr.payment.helpers._fetch_image_from_url", new=fetch): + tokens = await _responses_image_tokens(input_data) + + assert fetch.await_count == IMAGE_FETCH_MAX_PER_REQUEST + assert tokens == 36_000 * (IMAGE_FETCH_MAX_PER_REQUEST + 1) + + +async def test_responses_transform_failure_falls_back_to_worst_case() -> None: + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.001 + pricing.max_prompt_cost = 100.0 + pricing.max_completion_cost = 0.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + body = { + "model": "test-model", + "input": [ + { + "role": "user", + "content": [ + {"type": "input_image", "image_url": "https://x.test/a.jpg"}, + {"type": "input_image", "image_url": "https://x.test/b.jpg"}, + ], + } + ], + } + fetch = AsyncMock(return_value=None) + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + patch("routstr.payment.helpers._fetch_image_from_url", new=fetch), + patch( + "routstr.payment.responses_input.LiteLLMCompletionResponsesConfig." + "transform_responses_api_input_to_messages", + side_effect=RuntimeError("boom"), + ), + ): + cost = await calculate_discounted_max_cost(100_000, body, model_obj) + + fetch.assert_not_awaited() + # 2 * 36,000 tokens * 0.001 sats = 72 sats reserved + assert 72_000 <= cost < 100_000 async def test_discounted_max_cost_body_max_output_tokens_fallback() -> None: @@ -535,22 +635,28 @@ async def test_discounted_max_cost_invalid_completion_cap_ignored() -> None: assert cost == 80_000 -async def test_estimate_image_tokens_from_input_detail_and_file_id() -> None: +def _responses_file_image(detail: str | None) -> list[dict[str, Any]]: + part: dict[str, Any] = {"type": "input_image", "file_id": "file-1"} + if detail is not None: + part["detail"] = detail + return [{"role": "user", "content": [part]}] + + +async def test_responses_input_detail_and_file_id() -> None: import base64 from io import BytesIO from PIL import Image - from routstr.payment.helpers import estimate_image_tokens_from_input - # file_id: dimensions can't be fetched, so use a conservative max-size # estimate (4 tiles for auto/high) and honor the detail sibling for low. - assert await estimate_image_tokens_from_input( - [{"type": "input_image", "file_id": "file-1"}] - ) == 85 + (170 * 4) - assert await estimate_image_tokens_from_input( - [{"type": "input_image", "file_id": "file-1", "detail": "low"}] - ) == 85 + fetch = AsyncMock(return_value=None) + with patch("routstr.payment.helpers._fetch_image_from_url", new=fetch): + assert await _responses_image_tokens(_responses_file_image(None)) == 85 + ( + 170 * 4 + ) + assert await _responses_image_tokens(_responses_file_image("low")) == 85 + fetch.assert_not_awaited() # image_url honors the sibling detail instead of always defaulting to auto. image = Image.new("RGB", (512, 512), "red") @@ -558,12 +664,10 @@ async def test_estimate_image_tokens_from_input_detail_and_file_id() -> None: image.save(buffer, format="JPEG") data_url = "data:image/jpeg;base64," + base64.b64encode(buffer.getvalue()).decode() - assert await estimate_image_tokens_from_input( - [{"type": "input_image", "image_url": data_url, "detail": "low"}] - ) == 85 - assert await estimate_image_tokens_from_input( - [{"type": "input_image", "image_url": data_url, "detail": "high"}] - ) == 85 + 170 # 512x512 = 1 tile + assert await _responses_image_tokens(_responses_image(data_url, "low")) == 85 + assert ( + await _responses_image_tokens(_responses_image(data_url, "high")) == 85 + 170 + ) # 512x512 = 1 tile def test_calculate_image_tokens_original_detail() -> None: @@ -579,14 +683,12 @@ def test_calculate_image_tokens_original_detail() -> None: assert _calculate_image_tokens(10_000, 10_000, "original") == 36_000 -async def test_estimate_image_tokens_from_input_original_detail() -> None: +async def test_responses_input_original_detail() -> None: import base64 from io import BytesIO from PIL import Image - from routstr.payment.helpers import estimate_image_tokens_from_input - image = Image.new("RGB", (2048, 2048), "red") buffer = BytesIO() image.save(buffer, format="JPEG") @@ -594,19 +696,75 @@ async def test_estimate_image_tokens_from_input_original_detail() -> None: # image_url: billed at the decoded original resolution (4,096 patches), # not the 765-token tile cap. - assert await estimate_image_tokens_from_input( - [{"type": "input_image", "image_url": data_url, "detail": "original"}] - ) == 4_916 + assert await _responses_image_tokens(_responses_image(data_url)) == 4_916 # file_id: dimensions unknown, so use the 30,000-patch worst case. - assert await estimate_image_tokens_from_input( - [{"type": "input_image", "file_id": "file-1", "detail": "original"}] - ) == 36_000 + assert await _responses_image_tokens(_responses_file_image("original")) == 36_000 # Explicit null detail behaves like the auto default (tiled math). - assert await estimate_image_tokens_from_input( - [{"type": "input_image", "file_id": "file-1", "detail": None}] - ) == 85 + (170 * 4) + assert await _responses_image_tokens(_responses_image(data_url, None)) == 85 + ( + 170 * 4 + ) + + +def test_responses_input_to_messages_shapes() -> None: + from routstr.payment.responses_input import ( + FILE_ID_URL_PREFIX, + count_input_images, + responses_input_to_messages, + ) + + input_data = [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "hi"}, + {"type": "input_image", "file_id": "file-1", "detail": "original"}, + ], + }, + {"type": "function_call_output", "call_id": "c1", "output": "out"}, + ] + messages = responses_input_to_messages(input_data) + assert messages is not None + assert messages[0]["role"] == "user" + parts = messages[0]["content"] + assert parts[0] == {"type": "text", "text": "hi"} + assert parts[1]["type"] == "image_url" + assert parts[1]["image_url"] == { + "url": f"{FILE_ID_URL_PREFIX}file-1", + "detail": "original", + } + assert messages[1]["role"] == "tool" + + # dict-form image_url: litellm nests it verbatim, so it is flattened first. + nested = responses_input_to_messages( + [ + { + "type": "message", + "role": "user", + "content": [ + { + "type": "input_image", + "image_url": { + "url": "https://x.test/a.jpg", + "detail": "original", + }, + } + ], + } + ] + ) + assert nested is not None + assert nested[0]["content"][0]["image_url"] == { + "url": "https://x.test/a.jpg", + "detail": "original", + } + + assert responses_input_to_messages("plain") == [ + {"role": "user", "content": "plain"} + ] + assert responses_input_to_messages(None) == [] + assert count_input_images(input_data) == 1 async def test_estimate_image_tokens_in_messages_original_detail() -> None: @@ -637,8 +795,6 @@ async def test_estimate_image_tokens_in_messages_original_detail() -> None: # 640x640 -> 20x20 = 400 patches -> ceil(400 * 1.2) = 480 tokens. assert await estimate_image_tokens_in_messages(messages) == 480 - # input_image parts inside messages honor their sibling detail and - # file_id through the Responses estimator as well. messages = [ { "role": "user",