From 50ab9a3211254c1b8031fa0014cdcbb79718e6c4 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 8 Oct 2026 21:20:46 -0700 Subject: [PATCH] fix(token_counter): price a base64 PDF document per page instead of as one image (#45301) * fix(token_counter): price a base64 PDF document per page instead of as one image A `document` or `file` block carrying inline PDF bytes was priced like a single image (85 tokens), so when a provider's count-tokens endpoint rejected the model (Bedrock Opus) the local fallback answered 116 for a 12-page PDF the provider then billed at 35941 input tokens. The counter now reads the PDF with pypdf and prices each page as its extracted text plus the image Anthropic renders it to (1568 px long edge, 1.15 MP, 750 pixels per token), falling back to the old image pricing when pypdf is missing or the bytes are not a readable PDF. * refactor(token_counter): count PDF pages as they are read Sum each page's text and rendered-image tokens straight from the pypdf reader instead of materializing a page list first, keep the fallback to image pricing atomic when a page cannot be read, and annotate the new tests' locals as Final * test(integration): audit cells for page-priced PDF documents in count_tokens, pre-call checks and spend * test(integration): release the held peer when the concurrent budget wait times out --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/constants.py | 4 + litellm/litellm_core_utils/token_counter.py | 60 +++- tests/integration/_support/pdf_document.py | 149 +++++++++ .../test_pdf_document_count_tokens_chaos.py | 226 ++++++++++++++ .../test_pdf_document_count_tokens_wire.py | 273 ++++++++++++++++ ...df_document_pre_call_checks_owned_proxy.py | 292 ++++++++++++++++++ .../litellm_core_utils/test_token_counter.py | 104 ++++++- 7 files changed, 1098 insertions(+), 10 deletions(-) create mode 100644 tests/integration/_support/pdf_document.py create mode 100644 tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_chaos.py create mode 100644 tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_wire.py create mode 100644 tests/integration/routing/test_pdf_document_pre_call_checks_owned_proxy.py diff --git a/litellm/constants.py b/litellm/constants.py index 74273b9ecb9..699872f32cc 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -680,6 +680,10 @@ EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE: Final = float( ANTHROPIC_TOKEN_COUNTING_BETA_VERSION = os.getenv("ANTHROPIC_TOKEN_COUNTING_BETA_VERSION", "token-counting-2024-11-01") ANTHROPIC_SKILLS_API_BETA_VERSION: Final = "skills-2025-10-02" ANTHROPIC_BATCHES_ROUTE: Final = "/v1/messages/batches" +ANTHROPIC_IMAGE_MAX_LONG_EDGE_PX: Final = 1568 +ANTHROPIC_IMAGE_MAX_PIXELS: Final = 1_150_000 +ANTHROPIC_IMAGE_PIXELS_PER_TOKEN: Final = 750 +PDF_DATA_URL_PREFIX: Final = "data:application/pdf;base64," VERTEX_BATCH_PREDICTION_JOBS_ROUTE: Final = "batchPredictionJobs" ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES: Final = { "low": 1, diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 151830f1589..7ac3e9c80c5 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -2,6 +2,7 @@ ## Helper utilities for token counting import base64 import io +import math import struct from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from itertools import accumulate @@ -17,6 +18,9 @@ import litellm from litellm import verbose_logger from litellm._lazy_imports import get_default_encoding from litellm.constants import ( + ANTHROPIC_IMAGE_MAX_LONG_EDGE_PX, + ANTHROPIC_IMAGE_MAX_PIXELS, + ANTHROPIC_IMAGE_PIXELS_PER_TOKEN, DEFAULT_IMAGE_HEIGHT, DEFAULT_IMAGE_TOKEN_COUNT, DEFAULT_IMAGE_WIDTH, @@ -25,6 +29,7 @@ from litellm.constants import ( MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES, MAX_TILE_HEIGHT, MAX_TILE_WIDTH, + PDF_DATA_URL_PREFIX, TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS, TOKEN_COUNTER_MAX_CONCURRENT_COUNTS, TOKEN_COUNTER_MAX_EXACT_CHARS, @@ -805,6 +810,49 @@ def _anthropic_image_source_data( return "" +def _anthropic_rendered_page_image_tokens(width: float, height: float) -> int: + """A PDF page is rasterized within Anthropic's image limits, then billed by its pixel area.""" + if width <= 0 or height <= 0: + return 0 + area_at_max_edge: Final = ANTHROPIC_IMAGE_MAX_LONG_EDGE_PX**2 * min(width, height) / max(width, height) + return math.ceil(min(float(ANTHROPIC_IMAGE_MAX_PIXELS), area_at_max_edge) / ANTHROPIC_IMAGE_PIXELS_PER_TOKEN) + + +def _count_inline_pdf_tokens(data_url: str, count_function: TokenCounterFunction) -> int | None: + if not data_url.startswith(PDF_DATA_URL_PREFIX): + return None + try: + from pypdf import PdfReader + except ImportError: + verbose_logger.debug("pypdf is not installed, so the PDF document is priced like one image") + return None + try: + reader: Final = PdfReader(io.BytesIO(base64.b64decode(data_url[len(PDF_DATA_URL_PREFIX) :]))) + return sum( + count_function(page.extract_text() or "") + + _anthropic_rendered_page_image_tokens(float(page.mediabox.width), float(page.mediabox.height)) + for page in reader.pages + ) + except Exception as e: + verbose_logger.debug("Could not read the PDF document's pages (%s), so it is priced like one image", e) + return None + + +def _count_opaque_document_tokens( + data_url: str, + count_function: TokenCounterFunction, + use_default_image_token_count: bool, +) -> int: + pdf_tokens: Final = _count_inline_pdf_tokens(data_url, count_function) + if pdf_tokens is not None: + return pdf_tokens + return calculate_img_tokens( + data=data_url, + mode="auto", + use_default_image_token_count=use_default_image_token_count, + ) + + def _count_document_tokens( document: ChatCompletionDocumentObject | AnthropicMessagesDocumentParam, count_function: TokenCounterFunction, @@ -824,10 +872,8 @@ def _count_document_tokens( return metadata_tokens + _count_content_list( count_function, content, use_default_image_token_count, default_token_count ) - return metadata_tokens + calculate_img_tokens( - data=_anthropic_image_source_data(source), - mode="auto", - use_default_image_token_count=use_default_image_token_count, + return metadata_tokens + _count_opaque_document_tokens( + _anthropic_image_source_data(source), count_function, use_default_image_token_count ) @@ -844,11 +890,7 @@ def _count_file_tokens( name_tokens: Final = count_function(filename) if isinstance(filename, str) and filename else 0 if not isinstance(file_data, str) or not file_data: return name_tokens - return name_tokens + calculate_img_tokens( - data=file_data, - mode="auto", - use_default_image_token_count=use_default_image_token_count, - ) + return name_tokens + _count_opaque_document_tokens(file_data, count_function, use_default_image_token_count) def _count_anthropic_content( diff --git a/tests/integration/_support/pdf_document.py b/tests/integration/_support/pdf_document.py new file mode 100644 index 00000000000..5c6ad82b85f --- /dev/null +++ b/tests/integration/_support/pdf_document.py @@ -0,0 +1,149 @@ +import base64 +import json +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from itertools import accumulate +from types import MappingProxyType +from typing import Final + +from integration._support.wire import Reply +from pydantic import JsonValue + +PDF_MEDIA_TYPE: Final = "application/pdf" +COUNT_TOKENS_TARGET: Final = "/v1/messages/count_tokens" +COUNT_REFUSED: Final = Reply( + status=400, + body=json.dumps( + {"type": "error", "error": {"type": "invalid_request_error", "message": "count_tokens is not supported"}} + ).encode(), +) + + +@dataclass(frozen=True, slots=True) +class Page: + width: int = 612 + height: int = 792 + text: str | None = None + + +LETTER: Final = Page() +NARROW: Final = Page(width=100, height=1000) +_RENDERED_TOKENS: Final = MappingProxyType({(612, 792): 1534, (100, 1000): 328}) + + +def rendered_tokens(pages: Sequence[Page]) -> int: + return sum(_RENDERED_TOKENS[(page.width, page.height)] for page in pages) + + +def _escaped(text: str) -> str: + return text.replace("\\", "\\\\").replace("(", "\\(").replace(")", "\\)") + + +def _content_stream(page: Page) -> bytes: + stream: Final = f"BT /F1 12 Tf 72 {page.height - 72} Td ({_escaped(page.text or '')}) Tj ET".encode() + return b"<< /Length " + str(len(stream)).encode() + b" >>\nstream\n" + stream + b"\nendstream" + + +def _page_object(page: Page, contents_id: int | None) -> bytes: + contents: Final = f" /Contents {contents_id} 0 R" if contents_id is not None else "" + return ( + f"<< /Type /Page /Parent 2 0 R /MediaBox [0 0 {page.width} {page.height}]" + f" /Resources << /Font << /F1 1 0 R >> >>{contents} >>" + ).encode() + + +def _page_bodies(pages: Sequence[Page], starts: Sequence[int]) -> Iterator[bytes]: + for page, start in zip(pages, starts): + if page.text is None: + yield _page_object(page, None) + else: + yield _content_stream(page) + yield _page_object(page, start) + + +def pdf_bytes(pages: Sequence[Page]) -> bytes: + starts: Final = tuple(accumulate((1 if page.text is None else 2 for page in pages), initial=3)) + page_ids: Final = tuple(start + (0 if page.text is None else 1) for start, page in zip(starts, pages)) + kids: Final = " ".join(f"{identity} 0 R" for identity in page_ids) + bodies: Final = ( + b"<< /Type /Font /Subtype /Type1 /BaseFont /Helvetica >>", + f"<< /Type /Pages /Kids [{kids}] /Count {len(pages)} >>".encode(), + *_page_bodies(pages, starts), + b"<< /Type /Catalog /Pages 2 0 R >>", + ) + header: Final = b"%PDF-1.4\n" + objects: Final = tuple( + f"{number} 0 obj\n".encode() + body + b"\nendobj\n" for number, body in enumerate(bodies, start=1) + ) + offsets: Final = tuple(accumulate((len(chunk) for chunk in objects), initial=len(header))) + entries: Final = b"".join(f"{offset:010d} 00000 n \n".encode() for offset in offsets[:-1]) + trailer: Final = ( + f"xref\n0 {len(bodies) + 1}\n".encode() + + b"0000000000 65535 f \n" + + entries + + f"trailer\n<< /Size {len(bodies) + 1} /Root {starts[-1]} 0 R >>\nstartxref\n{offsets[-1]}\n%%EOF\n".encode() + ) + return header + b"".join(objects) + trailer + + +def encoded(raw: bytes) -> str: + return base64.b64encode(raw).decode() + + +def data_url(media_type: str, raw: bytes) -> str: + return f"data:{media_type};base64,{encoded(raw)}" + + +def pdf_data_url(pages: Sequence[Page]) -> str: + return data_url(PDF_MEDIA_TYPE, pdf_bytes(pages)) + + +def base64_source(data: JsonValue, media_type: JsonValue = PDF_MEDIA_TYPE) -> dict[str, JsonValue]: + return {"type": "base64", "media_type": media_type, "data": data} + + +def document(source: Mapping[str, JsonValue], **fields: JsonValue) -> dict[str, JsonValue]: + return {"type": "document", "source": dict(source), **fields} + + +def pdf_document(pages: Sequence[Page], **fields: JsonValue) -> dict[str, JsonValue]: + return document(base64_source(encoded(pdf_bytes(pages))), **fields) + + +def text_document(text: str) -> dict[str, JsonValue]: + return document({"type": "text", "media_type": "text/plain", "data": text}) + + +def chat_file(file_data: JsonValue, filename: JsonValue = "document.pdf") -> dict[str, JsonValue]: + return {"type": "file", "file": {"filename": filename, "file_data": file_data}} + + +def responses_input_file(file_data: JsonValue, filename: JsonValue = "document.pdf") -> dict[str, JsonValue]: + return {"type": "input_file", "filename": filename, "file_data": file_data} + + +def messages_body(model: str, blocks: Sequence[JsonValue], text: str, **fields: JsonValue) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 64, + "messages": [{"role": "user", "content": [*blocks, {"type": "text", "text": text}]}], + **fields, + } + + +def chat_body(model: str, parts: Sequence[JsonValue], text: str, **fields: JsonValue) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 64, + "messages": [{"role": "user", "content": [*parts, {"type": "text", "text": text}]}], + **fields, + } + + +def responses_body(model: str, items: Sequence[JsonValue], text: str, **fields: JsonValue) -> dict[str, JsonValue]: + return { + "model": model, + "max_output_tokens": 64, + "input": [{"role": "user", "content": [*items, {"type": "input_text", "text": text}]}], + **fields, + } diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_chaos.py b/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_chaos.py new file mode 100644 index 00000000000..f196018f117 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_chaos.py @@ -0,0 +1,226 @@ +import asyncio +import os +import re +import signal +import threading +import uuid +from collections.abc import Callable +from dataclasses import dataclass +from pathlib import Path +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.pdf_document import ( + COUNT_REFUSED, + COUNT_TOKENS_TARGET, + LETTER, + messages_body, + pdf_document, + rendered_tokens, +) +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from integration.providers._cache_control_marks_support import anthropic_peer, owned_config +from pydantic import JsonValue, TypeAdapter + +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_MODEL: Final = "burst-claude" +_PROVIDER_KEY: Final = "synthetic-provider-key" +_ASK: Final = "Summarize the attached report in one sentence." +_COUNT_BURST: Final = 24 +_MESSAGE_BURST: Final = 12 +_PAGE_COUNTS: Final = (1, 3, 12) +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +@dataclass(frozen=True, slots=True) +class _Counted: + pages: int + status: int + text: str + input_tokens: int | None + + +@dataclass(frozen=True, slots=True) +class _Sent: + marker: str + status: int + text: str + call_id: str + + +def _peer(request: Request) -> Reply: + if request.target == COUNT_TOKENS_TARGET: + return COUNT_REFUSED + return anthropic_peer(request) + + +def _held(release: threading.Event) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert release.wait(timeout=120), "Held peer was never released" + return _peer(request) + + return respond + + +def _deployment(api_base: str) -> dict[str, JsonValue]: + return { + "model_name": _MODEL, + "litellm_params": {"model": "anthropic/claude-opus-5-5", "api_base": api_base, "api_key": _PROVIDER_KEY}, + } + + +def _count_body(pages: int) -> dict[str, JsonValue]: + blocks: Final[list[JsonValue]] = [pdf_document((LETTER,) * pages)] if pages else [] + return {"model": _MODEL, "messages": messages_body(_MODEL, blocks, _ASK)["messages"]} + + +def _input_tokens(response: httpx.Response) -> int | None: + if response.status_code != 200: + return None + counted: Final = _JSON_OBJECT.validate_json(response.content).get("input_tokens") + return counted if isinstance(counted, int) else None + + +async def _fire_counts(url: str, key: str, *, tolerate_transport_errors: bool = False) -> tuple[_Counted, ...]: + async def one(client: httpx.AsyncClient, index: int) -> _Counted: + pages: Final = _PAGE_COUNTS[index % len(_PAGE_COUNTS)] + response: Final = await client.post( + "/v1/messages/count_tokens", json=_count_body(pages), headers={"Authorization": f"Bearer {key}"} + ) + return _Counted(pages, response.status_code, response.text, _input_tokens(response)) + + async with httpx.AsyncClient(base_url=url, timeout=90, trust_env=False) as client: + results: Final = await asyncio.gather( + *(one(client, index) for index in range(_COUNT_BURST)), return_exceptions=tolerate_transport_errors + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Counted)) + + +async def _fire_messages(url: str, key: str) -> tuple[_Sent, ...]: + async def one(client: httpx.AsyncClient) -> _Sent: + marker: Final = uuid.uuid4().hex + body: Final = messages_body(_MODEL, [pdf_document((LETTER,))], f"{_ASK} marker-{marker}") + response: Final = await client.post("/v1/messages", json=body, headers={"Authorization": f"Bearer {key}"}) + return _Sent(marker, response.status_code, response.text, response.headers.get("x-litellm-call-id", "")) + + async with httpx.AsyncClient(base_url=url, timeout=90, trust_env=False) as client: + return tuple(await asyncio.gather(*(one(client) for _ in range(_MESSAGE_BURST)))) + + +def _baseline(gateway: Gateway) -> int: + response: Final = gateway.request("POST", "/v1/messages/count_tokens", _count_body(0)) + assert response.status_code == 200, response.text + counted: Final = _input_tokens(response) + assert counted is not None, response.text + return counted + + +def _assert_exact(counts: tuple[_Counted, ...], baseline: int) -> None: + for item in counts: + assert item.status == 200, (item.pages, item.status, item.text) + assert item.input_tokens is not None and item.input_tokens - baseline == rendered_tokens( + (LETTER,) * item.pages + ), ( + item.pages, + item.input_tokens, + baseline, + ) + + +def _single_spend_row(item: _Sent) -> None: + assert item.status == 200, (item.status, item.text) + assert item.call_id, item.text + rows: Final = eventually( + lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (item.call_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + assert len(rows) == 1, item.call_id + + +def _held_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@pytest.mark.timeout(2 * graceful_stop_seconds() + 180) +async def test_pdf_count_burst_across_two_workers_prices_every_page_and_logs_each_message_once( + gateway: Gateway, tmp_path: Path +) -> None: + with wire_server(_peer) as wire: + config: Final = owned_config(tmp_path, [_deployment(wire.url)], litellm_settings={"cache": False}) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + eventually( + lambda: _STARTED_WORKER.findall(owned.log.read_text()), + lambda pids: len(pids) == 2, + seconds=graceful_stop_seconds(), + ) + owned_url: Final = str(owned.gateway.client.base_url) + baseline: Final = _baseline(owned.gateway) + wire.drain() + counts, messages = await asyncio.gather( + _fire_counts(owned_url, owned.gateway.key), _fire_messages(owned_url, owned.gateway.key) + ) + received: Final = wire.drain() + targets: Final = [request.target for request in received] + assert targets.count(COUNT_TOKENS_TARGET) == _COUNT_BURST, targets + assert targets.count("/v1/messages") == _MESSAGE_BURST, targets + _assert_exact(counts, baseline) + assert {item.marker for item in messages} == { + _JSON_OBJECT.validate_json(item.text)["id"][len("msg_") :] for item in messages if item.status == 200 + } + for item in messages: + _single_spend_row(item) + + +@pytest.mark.timeout(2 * graceful_stop_seconds() + 180) +async def test_worker_sigkill_mid_pdf_count_burst_leaves_the_sibling_pricing_pages( + gateway: Gateway, tmp_path: Path +) -> None: + release: Final = threading.Event() + with wire_server(_held(release)) as wire: + config: Final = owned_config(tmp_path, [_deployment(wire.url)], litellm_settings={"cache": False}) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + workers: Final[tuple[int, ...]] = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=graceful_stop_seconds(), + ) + owned_url: Final = str(owned.gateway.client.base_url) + burst: Final = asyncio.create_task( + _fire_counts(owned_url, owned.gateway.key, tolerate_transport_errors=True) + ) + try: + await asyncio.to_thread( + eventually, lambda: wire.received.qsize(), lambda size: size >= _COUNT_BURST, 90 + ) + held_by: Final = MappingProxyType({pid: _held_upstream_connections(pid, wire.url) for pid in workers}) + victim: Final = max(workers, key=held_by.__getitem__) + os.kill(victim, signal.SIGKILL) + finally: + release.set() + served: Final = await burst + wire.drain() + baseline: Final = _baseline(owned.gateway) + after: Final = await _fire_counts(owned_url, owned.gateway.key) + after_received: Final = wire.drain() + assert sum(held_by.values()) == _COUNT_BURST, held_by + assert held_by[victim] > 0, held_by + assert len(served) == _COUNT_BURST - held_by[victim], (len(served), held_by) + assert len(after) == _COUNT_BURST, len(after) + assert [request.target for request in after_received] == [COUNT_TOKENS_TARGET] * (_COUNT_BURST + 1) + _assert_exact(served, baseline) + _assert_exact(after, baseline) diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_wire.py new file mode 100644 index 00000000000..41d5f3d8685 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_pdf_document_count_tokens_wire.py @@ -0,0 +1,273 @@ +import io +import json +import uuid +from collections.abc import Sequence +from types import MappingProxyType +from typing import Final + +import pytest +from integration._support.client import Gateway, eventually, object_value +from integration._support.database import read_rows +from integration._support.pdf_document import ( + COUNT_REFUSED, + COUNT_TOKENS_TARGET, + LETTER, + NARROW, + Page, + base64_source, + chat_file, + document, + encoded, + messages_body, + pdf_bytes, + pdf_data_url, + pdf_document, + rendered_tokens, + responses_input_file, + text_document, +) +from integration._support.provider import SharedProvider +from integration._support.wire import Reply +from pydantic import JsonValue, TypeAdapter +from pypdf import PdfReader, PdfWriter + +_MODEL: Final = "anthropic/claude-opus-5-5" +_ASK: Final = "Summarize the attached report in one sentence." +_TEXT: Final = "The quarterly report covers revenue, margins and headcount." +_TITLE: Final = "Quarterly report" +_CONTEXT: Final = "Board pack, page one" +_PEER_INPUT_TOKENS: Final = 12 +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_PNG_HEADER: Final = b"\x89PNG\r\n\x1a\n\x00\x00\x00\x0dIHDR" + (1568).to_bytes(4, "big") * 2 + b"\x08\x06\x00\x00\x00" + + +def _locked(raw: bytes, user_password: str) -> bytes: + writer: Final = PdfWriter(clone_from=PdfReader(io.BytesIO(raw))) + writer.encrypt(user_password=user_password, owner_password="owner", algorithm="RC4-128") + out: Final = io.BytesIO() + writer.write(out) + return out.getvalue() + + +_UNREADABLE: Final[MappingProxyType[str, JsonValue]] = MappingProxyType( + { + "garbage-5kb": encoded(bytes(range(256)) * 20), + "cut-mid-stream": encoded(pdf_bytes((LETTER, LETTER))[:200]), + "png-bytes": encoded(_PNG_HEADER), + "user-password": encoded(_locked(pdf_bytes((LETTER, LETTER)), "reader")), + "empty": "", + "int": 1234, + } +) + + +def _int(value: JsonValue) -> int: + assert isinstance(value, int), value + return value + + +def _turn(blocks: Sequence[JsonValue]) -> list[JsonValue]: + return messages_body(_MODEL, blocks, _ASK)["messages"] + + +def _count(gateway: Gateway, provider: SharedProvider, blocks: Sequence[JsonValue]) -> int: + provider.expect(COUNT_REFUSED) + response: Final = gateway.request("POST", "/v1/messages/count_tokens", {"model": _MODEL, "messages": _turn(blocks)}) + assert response.status_code == 200, response.text + assert [(request.method, request.target) for request in provider.received()] == [("POST", COUNT_TOKENS_TARGET)] + return _int(_JSON_OBJECT.validate_json(response.content)["input_tokens"]) + + +def _local(gateway: Gateway, parts: Sequence[JsonValue]) -> int: + response: Final = gateway.request( + "POST", "/utils/token_counter", {"model": _MODEL, "messages": _turn(parts)}, params={"call_endpoint": "false"} + ) + assert response.status_code == 200, response.text + return _int(_JSON_OBJECT.validate_json(response.content)["total_tokens"]) + + +def _responses_count(gateway: Gateway, provider: SharedProvider, items: Sequence[JsonValue]) -> int: + provider.expect(COUNT_REFUSED) + response: Final = gateway.request( + "POST", + "/v1/responses/input_tokens", + {"model": _MODEL, "input": [{"role": "user", "content": [*items, {"type": "input_text", "text": _ASK}]}]}, + ) + assert response.status_code == 200, response.text + assert [(request.method, request.target) for request in provider.received()] == [("POST", COUNT_TOKENS_TARGET)] + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["object"] == "response.input_tokens", response.text + return _int(payload["input_tokens"]) + + +def _cost(gateway: Gateway, parts: Sequence[JsonValue]) -> float: + payload: Final = gateway.post("/spend/calculate", {"model": _MODEL, "messages": _turn(parts)}) + cost: Final = payload["cost"] + assert isinstance(cost, (int, float)), payload + return float(cost) + + +def _input_price(gateway: Gateway) -> float: + listed: Final = gateway.get("/model/info")["data"] + assert isinstance(listed, list), listed + rows: Final = [object_value(row) for row in listed if isinstance(row, dict) and row.get("model_name") == _MODEL] + assert len(rows) == 1, rows + price: Final = object_value(rows[0]["model_info"])["input_cost_per_token"] + assert isinstance(price, float) and price > 0, price + return price + + +def _peer_message(identity: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": identity, + "type": "message", + "role": "assistant", + "model": "claude-opus-5-5", + "content": [{"type": "text", "text": "One sentence."}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": _PEER_INPUT_TOKENS, "output_tokens": 3}, + } + ).encode() + ) + + +def _prompt_tokens(call_id: str) -> int: + rows: Final = eventually( + lambda: read_rows('SELECT prompt_tokens FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + return _int(rows[0]["prompt_tokens"]) + + +@pytest.mark.parametrize("pages", [1, 3, 12]) +def test_count_tokens_fallback_prices_each_blank_page_as_a_rendered_image( + gateway: Gateway, provider: SharedProvider, pages: int +) -> None: + letter_pages: Final = (LETTER,) * pages + with_document: Final = _count(gateway, provider, [pdf_document(letter_pages)]) + without: Final = _count(gateway, provider, []) + assert with_document - without == rendered_tokens(letter_pages), (with_document, without) + + +def test_count_tokens_fallback_adds_a_page_text_to_its_rendered_image( + gateway: Gateway, provider: SharedProvider +) -> None: + text_page: Final = _count(gateway, provider, [pdf_document((Page(text=_TEXT),))]) + blank_page: Final = _count(gateway, provider, [pdf_document((LETTER,))]) + text_source: Final = _count(gateway, provider, [text_document(_TEXT)]) + without: Final = _count(gateway, provider, []) + assert text_source - without > 0, (text_source, without) + assert text_page - blank_page == text_source - without, (text_page, blank_page, text_source, without) + + +def test_count_tokens_fallback_prices_a_page_by_its_rendered_area(gateway: Gateway, provider: SharedProvider) -> None: + mixed: Final = (LETTER, NARROW, LETTER) + without: Final = _count(gateway, provider, []) + narrow: Final = _count(gateway, provider, [pdf_document((NARROW,))]) + assert narrow - without == rendered_tokens((NARROW,)), (narrow, without) + assert _count(gateway, provider, [pdf_document(mixed)]) - without == rendered_tokens(mixed) + + +def test_count_tokens_fallback_prices_every_document_in_the_turn(gateway: Gateway, provider: SharedProvider) -> None: + both: Final = _count(gateway, provider, [pdf_document((LETTER,) * 3), pdf_document((LETTER,))]) + without: Final = _count(gateway, provider, []) + assert both - without == rendered_tokens((LETTER,) * 4), (both, without) + + +def test_count_tokens_fallback_prices_a_pdf_without_pages_as_nothing( + gateway: Gateway, provider: SharedProvider +) -> None: + assert _count(gateway, provider, [pdf_document(())]) == _count(gateway, provider, []) + + +def test_count_tokens_fallback_adds_title_and_context_once_per_document( + gateway: Gateway, provider: SharedProvider +) -> None: + without: Final = _count(gateway, provider, []) + plain: Final = _count(gateway, provider, [pdf_document((LETTER,))]) + annotated: Final = _count(gateway, provider, [pdf_document((LETTER,), title=_TITLE, context=_CONTEXT)]) + title: Final = _count(gateway, provider, [text_document(_TITLE)]) - without + context: Final = _count(gateway, provider, [text_document(_CONTEXT)]) - without + assert title > 0 and context > 0, (title, context) + assert annotated - plain == title + context, (annotated, plain, title, context) + + +@pytest.mark.parametrize("label", list(_UNREADABLE)) +def test_count_tokens_fallback_prices_unreadable_pdf_bytes_like_one_image( + gateway: Gateway, provider: SharedProvider, label: str +) -> None: + data: Final = _UNREADABLE[label] + as_pdf: Final = _count(gateway, provider, [document(base64_source(data))]) + as_png: Final = _count(gateway, provider, [document(base64_source(data, media_type="image/png"))]) + without: Final = _count(gateway, provider, []) + assert as_pdf == as_png, (as_pdf, as_png) + assert 0 <= as_pdf - without < rendered_tokens((LETTER,)), (as_pdf, without) + + +def test_count_tokens_fallback_reads_a_list_wrapped_base64_string_like_the_bare_string( + gateway: Gateway, provider: SharedProvider +) -> None: + raw: Final = encoded(pdf_bytes((LETTER,))) + wrapped: Final = _count(gateway, provider, [document(base64_source([raw]))]) + assert wrapped == _count(gateway, provider, [document(base64_source(raw))]), wrapped + + +def test_count_tokens_fallback_reads_an_owner_locked_pdf(gateway: Gateway, provider: SharedProvider) -> None: + pages: Final = (LETTER, LETTER) + locked: Final = document(base64_source(encoded(_locked(pdf_bytes(pages), "")))) + assert _count(gateway, provider, [locked]) - _count(gateway, provider, []) == rendered_tokens(pages) + + +def test_count_tokens_fallback_parses_only_the_pdf_media_type(gateway: Gateway, provider: SharedProvider) -> None: + raw: Final = encoded(pdf_bytes((LETTER,))) + without: Final = _count(gateway, provider, []) + labelled: Final = _count(gateway, provider, [document(base64_source(raw))]) + assert labelled - without == rendered_tokens((LETTER,)), (labelled, without) + mislabelled: Final = _count(gateway, provider, [document(base64_source(raw, media_type="application/x-pdf"))]) + assert mislabelled == _count(gateway, provider, [document(base64_source(raw, media_type="image/png"))]) + + +def test_utils_token_counter_prices_a_chat_file_by_its_pages(gateway: Gateway) -> None: + one: Final = _local(gateway, [chat_file(pdf_data_url((LETTER,)))]) + three: Final = _local(gateway, [chat_file(pdf_data_url((LETTER,) * 3))]) + assert three - one == rendered_tokens((LETTER,) * 2), (three, one) + text_page: Final = _local(gateway, [chat_file(pdf_data_url((Page(text=_TEXT),)))]) + text_part: Final = _local(gateway, [chat_file(pdf_data_url((LETTER,))), {"type": "text", "text": _TEXT}]) + assert text_part - one > 0, (text_part, one) + assert text_page - one == text_part - one, (text_page, text_part, one) + + +def test_responses_input_tokens_fallback_prices_an_input_file_by_its_pages( + gateway: Gateway, provider: SharedProvider +) -> None: + one: Final = _responses_count(gateway, provider, [responses_input_file(pdf_data_url((LETTER,)))]) + three: Final = _responses_count(gateway, provider, [responses_input_file(pdf_data_url((LETTER,) * 3))]) + assert three - one == rendered_tokens((LETTER,) * 2), (three, one) + + +def test_spend_calculate_prices_a_chat_file_by_its_pages(gateway: Gateway) -> None: + twelve: Final = _cost(gateway, [chat_file(pdf_data_url((LETTER,) * 12))]) + one: Final = _cost(gateway, [chat_file(pdf_data_url((LETTER,)))]) + assert twelve - one == pytest.approx(rendered_tokens((LETTER,) * 11) * _input_price(gateway)), (twelve, one) + + +def test_a_cached_pdf_message_keeps_the_peer_usage_on_both_spend_rows( + gateway: Gateway, provider: SharedProvider +) -> None: + marker: Final = uuid.uuid4().hex + body: Final = messages_body(_MODEL, [pdf_document((LETTER,))], f"{_ASK} marker-{marker}") + provider.expect(_peer_message(f"msg_{marker}")) + first: Final = gateway.request("POST", "/v1/messages", body) + second: Final = gateway.request("POST", "/v1/messages", body) + assert (first.status_code, second.status_code) == (200, 200), (first.text, second.text) + assert [(request.method, request.target) for request in provider.received()] == [("POST", "/v1/messages")] + assert _JSON_OBJECT.validate_json(second.content)["id"] == f"msg_{marker}", second.text + assert first.headers["x-litellm-call-id"] != second.headers["x-litellm-call-id"] + assert [_prompt_tokens(response.headers["x-litellm-call-id"]) for response in (first, second)] == [ + _PEER_INPUT_TOKENS, + _PEER_INPUT_TOKENS, + ] diff --git a/tests/integration/routing/test_pdf_document_pre_call_checks_owned_proxy.py b/tests/integration/routing/test_pdf_document_pre_call_checks_owned_proxy.py new file mode 100644 index 00000000000..9fad73f8080 --- /dev/null +++ b/tests/integration/routing/test_pdf_document_pre_call_checks_owned_proxy.py @@ -0,0 +1,292 @@ +import os +import socket +import threading +import uuid +from collections.abc import Callable, Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows, scratch_database +from integration._support.pdf_document import ( + COUNT_REFUSED, + COUNT_TOKENS_TARGET, + LETTER, + chat_body, + chat_file, + messages_body, + pdf_data_url, + pdf_document, + responses_body, + responses_input_file, +) +from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.providers._cache_control_marks_support import anthropic_peer, marker_of, owned_config +from pydantic import JsonValue, TypeAdapter +from redis import Redis + +_WINDOW: Final = "window-claude" +_ITPM: Final = "itpm-claude" +_BUDGET: Final = "budget-claude" +_AFFINITY: Final = "affinity-claude" +_LIMIT: Final = 5000 +_PAGE_COUNT: Final = 12 +_PAGES: Final = (LETTER,) * _PAGE_COUNT +_ASK: Final = "Summarize the attached report in one sentence." +_FOLLOW_UPS: Final = 9 +_TRANSCRIPT: Final = " ".join(f"Line {number}: revenue, margins and headcount moved." for number in range(160)) +_PIN_KEYS: Final = "*:prompt_caching" +_PROVIDER_KEY: Final = "synthetic-provider-key" +_SURFACES: Final = ("messages", "messages-stream", "chat-file", "responses-input-file") +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +@dataclass(frozen=True, slots=True) +class _Rig: + owned: OwnedProxy + wire: Wire + held_port: int + + +def _marker() -> str: + return uuid.uuid4().hex + + +def _free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return int(reserve.getsockname()[1]) + + +def _deployment( + name: str, + backend: str, + api_base: str, + *, + model_info: Mapping[str, JsonValue] | None = None, + **params: JsonValue, +) -> dict[str, JsonValue]: + return { + "model_name": name, + "litellm_params": {"model": f"anthropic/{backend}", "api_base": api_base, "api_key": _PROVIDER_KEY, **params}, + "model_info": dict(model_info or {}), + } + + +def _peer(request: Request) -> Reply: + if request.target == COUNT_TOKENS_TARGET: + return COUNT_REFUSED + return anthropic_peer(request) + + +def _held(release: threading.Event) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert release.wait(timeout=120), "Held peer was never released" + return _peer(request) + + return respond + + +def _request(surface: str, model: str, pages: int, marker: str) -> tuple[str, dict[str, JsonValue]]: + text: Final = f"{_ASK} marker-{marker}" + letter_pages: Final = (LETTER,) * pages + if surface == "messages": + return "/v1/messages", messages_body(model, [pdf_document(letter_pages)], text) + if surface == "messages-stream": + return "/v1/messages", messages_body(model, [pdf_document(letter_pages)], text, stream=True) + if surface == "chat-file": + return "/v1/chat/completions", chat_body(model, [chat_file(pdf_data_url(letter_pages))], text) + return "/v1/responses", responses_body(model, [responses_input_file(pdf_data_url(letter_pages))], text) + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]: + directory: Final = tmp_path_factory.mktemp("pdf-pre-call") + held_port: Final = _free_port() + with gateway_from_environment() as gateway, wire_server(_peer) as wire: + config: Final = owned_config( + directory, + [ + _deployment(_WINDOW, "claude-haiku-5-5", wire.url, model_info={"max_input_tokens": _LIMIT}), + _deployment(_ITPM, "claude-sonnet-4-6", wire.url, itpm=_LIMIT), + _deployment(_BUDGET, "claude-opus-4-8", f"http://127.0.0.1:{held_port}"), + ], + litellm_settings={"cache": False}, + router_settings={ + "enable_pre_call_checks": True, + "optional_pre_call_checks": ["enforce_model_rate_limits"], + }, + ) + with owned_proxy_process(gateway, directory, {}, config=config, workers=2) as owned: + yield _Rig(owned, wire, held_port) + + +def _spend_row(call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + return rows[0] + + +@pytest.mark.parametrize("surface", _SURFACES) +def test_context_window_check_rejects_a_pdf_whose_pages_exceed_max_input_tokens(rig: _Rig, surface: str) -> None: + rig.wire.drain() + path, body = _request(surface, _WINDOW, _PAGE_COUNT, _marker()) + response: Final = rig.owned.gateway.request("POST", path, body) + assert response.status_code == 400, (response.status_code, response.text) + assert "Context Window exceeded" in response.text, response.text + assert [request.target for request in rig.wire.drain()] == [] + + +def test_context_window_check_passes_a_one_page_pdf_and_forwards_its_bytes(rig: _Rig) -> None: + rig.wire.drain() + path, body = _request("messages", _WINDOW, 1, _marker()) + response: Final = rig.owned.gateway.request("POST", path, body) + assert response.status_code == 200, response.text + received: Final = rig.wire.drain() + assert [request.target for request in received] == ["/v1/messages"] + forwarded: Final = _JSON_OBJECT.validate_json(received[0].body)["messages"] + assert isinstance(forwarded, list) and len(forwarded) == 1, forwarded + blocks: Final = object_value(forwarded[0])["content"] + assert isinstance(blocks, list) and pdf_document((LETTER,)) in blocks, blocks + + +def test_model_itpm_check_rejects_a_pdf_whose_pages_exceed_the_limit(rig: _Rig) -> None: + rig.wire.drain() + path, body = _request("messages", _ITPM, _PAGE_COUNT, _marker()) + response: Final = rig.owned.gateway.request("POST", path, body) + assert response.status_code == 429, (response.status_code, response.text) + assert '"error"' in response.text, response.text + assert [request.target for request in rig.wire.drain()] == [] + + +def test_key_budget_reservation_rejects_the_second_concurrent_pdf_request(rig: _Rig) -> None: + release: Final = threading.Event() + requests: Final = tuple(_request("messages", _BUDGET, _PAGE_COUNT, _marker()) for _ in range(2)) + with wire_server(_held(release), port=rig.held_port) as wire, rig.owned.gateway.scenario() as scenario: + key: Final = scenario.key(max_budget=0.01) + headers: Final = {"Authorization": f"Bearer {key}"} + with ( + httpx.Client(base_url=str(rig.owned.gateway.client.base_url), timeout=90, trust_env=False) as client, + ThreadPoolExecutor(max_workers=2) as pool, + ): + futures: Final = tuple( + pool.submit(client.post, path, json=body, headers=headers) for path, body in requests + ) + try: + eventually( + lambda: wire.received.qsize() + sum(1 for future in futures if future.done()), + lambda settled: settled >= 2, + seconds=60, + ) + finally: + release.set() + responses: Final = tuple(future.result(timeout=90) for future in futures) + received: Final = wire.drain() + assert sorted(response.status_code for response in responses) == [200, 422], [ + response.text for response in responses + ] + rejected: Final = next(response for response in responses if response.status_code != 200) + assert "udget" in rejected.text, rejected.text + assert [request.target for request in received] == ["/v1/messages"] + + +def test_a_pre_call_rejection_logs_one_failure_row_and_calls_no_upstream(rig: _Rig) -> None: + path, body = _request("messages", _WINDOW, _PAGE_COUNT, _marker()) + response: Final = rig.owned.gateway.request("POST", path, body) + assert response.status_code == 400, (response.status_code, response.text) + assert _spend_row(response.headers["x-litellm-call-id"])["status"] == "failure" + assert rig.wire.drain() == () + + +def _first_turn(session: str, marker: str) -> list[JsonValue]: + return [ + {"role": "system", "content": f"Answer from the attached report. session {session}"}, + *chat_body( + _AFFINITY, + [{"type": "text", "text": _TRANSCRIPT}, chat_file(pdf_data_url(_PAGES))], + f"{_ASK} marker-{marker}", + )["messages"], + ] + + +def _follow_up(first_turn: list[JsonValue], marker: str) -> list[JsonValue]: + return [ + *first_turn, + {"role": "assistant", "content": "One sentence."}, + {"role": "user", "content": f"And the margins? marker-{marker}"}, + ] + + +def _served(wire: Wire) -> frozenset[str]: + return frozenset(marker_of(request) for request in wire.drain()) + + +def _pin_count(cache: Redis) -> int: + return len(cache.keys(_PIN_KEYS)) + + +@pytest.mark.timeout(2 * graceful_stop_seconds() + 180) +def test_prompt_caching_affinity_pins_follow_ups_of_a_cached_turn_that_carries_a_pdf( + gateway: Gateway, tmp_path: Path +) -> None: + session: Final = _marker() + first_marker: Final = _marker() + follow_up_markers: Final = tuple(_marker() for _ in range(_FOLLOW_UPS)) + first_turn: Final = _first_turn(session, first_marker) + with scratch_database() as database_url, wire_server(_peer) as left, wire_server(_peer) as right: + config: Final = owned_config( + tmp_path, + [ + _deployment(_AFFINITY, "claude-opus-5-5", left.url, model_info={"id": f"pdf-left-{session}"}), + _deployment(_AFFINITY, "claude-opus-5-5", right.url, model_info={"id": f"pdf-right-{session}"}), + ], + litellm_settings={"cache": False}, + router_settings={ + "optional_pre_call_checks": ["prompt_caching"], + "redis_host": os.environ["REDIS_HOST"], + "redis_port": int(os.environ["REDIS_PORT"]), + }, + ) + with ( + Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache, + owned_proxy_process( + gateway, + tmp_path, + {"DATABASE_URL": database_url}, + config=config, + remove_environment=("DATABASE_URL_READ_REPLICA",), + workers=2, + ) as owned, + owned.gateway.scenario() as scenario, + ): + key: Final = scenario.key(metadata={"enable_prompt_caching": True}) + pins_before: Final = _pin_count(cache) + first: Final = owned.gateway.request( + "POST", "/v1/chat/completions", {"model": _AFFINITY, "messages": first_turn, "max_tokens": 64}, key=key + ) + assert first.status_code == 200, first.text + eventually(lambda: _pin_count(cache), lambda count: count > pins_before, seconds=70) + follow_ups: Final = tuple( + owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": _AFFINITY, "messages": _follow_up(first_turn, marker), "max_tokens": 64}, + key=key, + ) + for marker in follow_up_markers + ) + served: Final = {"left": _served(left), "right": _served(right)} + assert [response.status_code for response in follow_ups] == [200] * _FOLLOW_UPS, [ + response.text for response in follow_ups + ] + assert served["left"] | served["right"] == {first_marker, *follow_up_markers}, served + first_side: Final = next(side for side in served if first_marker in served[side]) + assert served[first_side] == {first_marker, *follow_up_markers}, served diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index 163897c2609..a646135efc3 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -11,6 +11,7 @@ import threading import time from collections.abc import Mapping from concurrent.futures import Future, wait +from itertools import accumulate, chain from pathlib import Path from typing import Final from unittest.mock import MagicMock @@ -1444,7 +1445,7 @@ def _count_user_content(content: list[dict]) -> int: ids=["base64", "url", "file"], ) def test_anthropic_document_block_with_opaque_source_is_priced_like_an_image(source: dict[str, str]): - """A `document` whose bytes can't be tokenized locally is priced like an `image`, not raised on.""" + """A `document` with no readable pages (a bare PDF header, a URL, a file id) is priced like an `image`, not raised on.""" prompt = {"type": "text", "text": "Summarize this file."} assert _count_user_content([prompt, {"type": "document", "source": source}]) == _count_user_content( @@ -1504,6 +1505,12 @@ def test_openai_file_block_prices_like_the_equivalent_anthropic_document(): assert _count_user_content([prompt, inline_file]) == _count_user_content([prompt, document]) assert _count_user_content([prompt, inline_file]) > _count_user_content([prompt]) + readable: Final = _pdf_base64(("Revenue grew eleven percent while churn fell to two percent.",)) + readable_file: Final = {"type": "file", "file": {"filename": "report.pdf", "file_data": "data:application/pdf;base64," + readable}} + readable_document: Final = {"type": "document", "title": "report.pdf", "source": _pdf_source(readable)} + assert _count_user_content([prompt, readable_file]) == _count_user_content([prompt, readable_document]) + assert _count_user_content([prompt, readable_file]) > _count_user_content([prompt, inline_file]) + def test_openai_file_block_without_inline_bytes_counts_what_it_carries(): """A `file` block naming an uploaded file has no bytes to price, so it adds only the filename's tokens.""" @@ -1518,6 +1525,101 @@ def test_openai_file_block_without_inline_bytes_counts_what_it_carries(): ) +def _pdf_base64(pages: tuple[str, ...], width: int = 612, height: int = 792) -> str: + def page_objects(index: int, text: str) -> tuple[bytes, bytes]: + escaped: Final = text.replace("\\", "\\\\").replace("(", "\\(").replace(")", "\\)") + stream: Final = f"BT /F1 12 Tf 72 720 Td ({escaped}) Tj ET".encode("latin-1") + content: Final = f"<< /Length {len(stream)} >>\nstream\n".encode() + stream + b"\nendstream" + page: Final = ( + f"<< /Type /Page /Parent 2 0 R /MediaBox [0 0 {width} {height}] " + f"/Resources << /Font << /F1 3 0 R >> >> /Contents {4 + 2 * index} 0 R >>" + ).encode() + return content, page + + kids: Final = " ".join(f"{5 + 2 * index} 0 R" for index in range(len(pages))) + bodies: Final = ( + b"<< /Type /Catalog /Pages 2 0 R >>", + f"<< /Type /Pages /Count {len(pages)} /Kids [ {kids} ] >>".encode(), + b"<< /Type /Font /Subtype /Type1 /BaseFont /Helvetica >>", + *chain.from_iterable(page_objects(index, text) for index, text in enumerate(pages)), + ) + header: Final = b"%PDF-1.4\n" + objects: Final = tuple( + f"{number} 0 obj\n".encode() + body + b"\nendobj\n" for number, body in enumerate(bodies, start=1) + ) + offsets: Final = accumulate((len(header), *(len(obj) for obj in objects[:-1]))) + xref: Final = f"xref\n0 {len(objects) + 1}\n0000000000 65535 f \n".encode() + b"".join( + f"{offset:010d} 00000 n \n".encode() for offset in offsets + ) + trailer: Final = ( + f"trailer\n<< /Size {len(objects) + 1} /Root 1 0 R >>\n" + f"startxref\n{len(header) + sum(len(obj) for obj in objects)}\n%%EOF\n" + ).encode() + return base64.b64encode(header + b"".join(objects) + xref + trailer).decode() + + +def _pdf_source(pdf_base64: str) -> dict[str, str]: + return {"type": "base64", "media_type": "application/pdf", "data": pdf_base64} + + +def test_base64_pdf_document_counts_every_page_text_and_rendering(): + """A base64 PDF is read page by page: each page costs its text plus the image Anthropic renders it to. + + Before the fix the whole document was priced as one 85-token image, so a count_tokens call that fell back + to the local counter answered 116 for a 12-page PDF the provider then billed at 35941 input tokens. + """ + prompt: Final = {"type": "text", "text": "Summarize this file."} + first: Final = "Revenue grew eleven percent while churn fell to two percent." + second: Final = "Headcount is flat and the office lease was renewed for three years." + base: Final = _count_user_content([prompt]) + blank_page: Final = _count_user_content([prompt, {"type": "document", "source": _pdf_source(_pdf_base64(("",)))}]) - base + + assert blank_page > 0 + assert _count_user_content([prompt, {"type": "document", "source": _pdf_source(_pdf_base64((first, second)))}]) == ( + _count_user_content([prompt, {"type": "text", "text": first}, {"type": "text", "text": second}]) + 2 * blank_page + ) + one_page: Final = _count_user_content([prompt, {"type": "document", "source": _pdf_source(_pdf_base64((first,)))}]) - base + twelve_pages: Final = ( + _count_user_content([prompt, {"type": "document", "source": _pdf_source(_pdf_base64((first,) * 12))}]) - base + ) + assert one_page > _count_user_content([prompt, {"type": "text", "text": first}]) - base + assert twelve_pages == 12 * one_page + + +@pytest.mark.parametrize("fields", [{}, {"title": "Q3 board packet", "context": "Shared by finance"}], ids=["bare", "described"]) +def test_pdf_page_rendering_cost_follows_anthropic_image_scaling(fields: dict[str, str]): + """Anthropic rasterizes each PDF page within its image limits (1568 px long edge, 1.15 MP) and bills + width * height / 750 tokens for it: https://platform.claude.com/docs/en/build-with-claude/pdf-support and + https://platform.claude.com/docs/en/build-with-claude/vision, read 2026-10-07, when Bedrock billed about + 1550 tokens per blank Letter page on both Sonnet 4.6 and Opus 4.8. + """ + prompt: Final = {"type": "text", "text": "Summarize this file."} + + def blank_page_cost(width: int, height: int) -> int: + with_page: Final = {"type": "document", "source": _pdf_source(_pdf_base64(("",), width, height)), **fields} + without_page: Final = {"type": "document", "source": _pdf_source(_pdf_base64((), width, height)), **fields} + return _count_user_content([prompt, with_page]) - _count_user_content([prompt, without_page]) + + letter: Final = blank_page_cost(612, 792) + poster: Final = blank_page_cost(2448, 3168) + strip: Final = blank_page_cost(1000, 100) + + assert letter == poster == 1534 + assert strip == 328 + + +def test_pdf_document_without_pypdf_is_priced_like_an_image(monkeypatch: pytest.MonkeyPatch): + prompt: Final = {"type": "text", "text": "Summarize this file."} + source: Final = _pdf_source(_pdf_base64(("Revenue grew eleven percent while churn fell to two percent.",))) + priced_by_page: Final = _count_user_content([prompt, {"type": "document", "source": source}]) + + monkeypatch.setitem(sys.modules, "pypdf", None) + + priced_as_image: Final = _count_user_content([prompt, {"type": "document", "source": source}]) + assert priced_as_image == _count_user_content([prompt, {"type": "image", "source": source}]) + assert priced_as_image < priced_by_page + + def _png_data_url(width: int, height: int) -> str: ihdr = b"\x89PNG\r\n\x1a\n" + (13).to_bytes(4, "big") + b"IHDR" + width.to_bytes(4, "big") + height.to_bytes(4, "big") return "data:image/png;base64," + base64.b64encode(ihdr + b"\x08\x06\x00\x00\x00").decode()