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()