use litellm pre-existing functionality

This commit is contained in:
9qeklajc
2026-09-12 14:16:40 +02:00
parent 5920bb69da
commit 63a7227c2a
3 changed files with 363 additions and 204 deletions
+58 -134
View File
@@ -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(
+79
View File
@@ -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
+226 -70
View File
@@ -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",