From 34ffc8f7b04e57c804b5fdd3bad5aee5647b4c1f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 4 Sep 2026 02:08:30 +0200 Subject: [PATCH] Validate destinations of pre-auth image URL fetches Cost estimation fetches caller-supplied image URLs before any credential is checked, so the node could be aimed at loopback, link-local, and private addresses. Restrict the fetch to globally reachable addresses (RFC 6890), resolve DNS once and validate every returned address, and rewrite plain HTTP to the validated address with the original Host header so a name cannot rebind between the check and the connection. HTTPS keeps its hostname because certificate validation already binds it. Cap the downloaded prefix and the number of URLs one request may fetch. --- routstr/payment/helpers.py | 100 +++++++++++++- tests/unit/test_image_url_fetch_guard.py | 167 +++++++++++++++++++++++ 2 files changed, 261 insertions(+), 6 deletions(-) create mode 100644 tests/unit/test_image_url_fetch_guard.py diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 1acd3d21..0c012be8 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -1,8 +1,12 @@ +import asyncio import base64 +import ipaddress import json import math +import socket from io import BytesIO from typing import Any +from urllib.parse import urlsplit, urlunsplit import httpx from fastapi import HTTPException, Response @@ -323,6 +327,14 @@ def estimate_prompt_tokens(body: dict) -> int: return _sum_string_chars(body) // 3 + _count_prompt_token_ids(body.get("prompt")) +IMAGE_FETCH_TIMEOUT_SECONDS = 10.0 +# Dimensions live in the header, so a prefix suffices and an endless body cannot +# pin memory. +IMAGE_FETCH_MAX_BYTES = 512 * 1024 +# Fetches are sequential, so an unbounded URL list is a request-time amplifier. +IMAGE_FETCH_MAX_PER_REQUEST = 8 + + def _get_image_dimensions(image_data: bytes) -> tuple[int, int]: """Extract image dimensions from image bytes.""" try: @@ -336,13 +348,81 @@ def _get_image_dimensions(image_data: bytes) -> tuple[int, int]: return (512, 512) -async def _fetch_image_from_url(url: str) -> bytes | None: - """Fetch image from URL.""" +def _is_blocked_address(address: str) -> bool: + """Allow only globally reachable addresses (RFC 6890).""" try: - async with httpx.AsyncClient(timeout=10.0) as client: - response = await client.get(url) - response.raise_for_status() - return response.content + ip = ipaddress.ip_address(address) + except ValueError: + return True + if isinstance(ip, ipaddress.IPv6Address): + # An embedded v4 address would otherwise smuggle a rejected target past + # the v6 checks. + for embedded in (ip.ipv4_mapped, ip.sixtofour): + if embedded is not None: + return _is_blocked_address(str(embedded)) + return not ip.is_global or ip.is_multicast + + +async def _validated_fetch_target(url: str) -> tuple[str, str]: + """Return the URL to request and its ``Host`` header. + + Cost estimation runs on the unauthenticated request body, so a caller can + otherwise aim the node at internal hosts. HTTP is rewritten to the resolved + address so the name cannot rebind between check and connect; HTTPS keeps its + hostname because certificate validation already binds the connection. + """ + parts = urlsplit(url) + if parts.scheme not in ("http", "https"): + raise ValueError(f"unsupported scheme: {parts.scheme or 'none'}") + host = parts.hostname + if not host: + raise ValueError("missing host") + + default_port = 443 if parts.scheme == "https" else 80 + port = parts.port or default_port + host_header = f"[{host}]" if ":" in host else host + if parts.port is not None: + host_header = f"{host_header}:{parts.port}" + + infos = await asyncio.get_running_loop().getaddrinfo( + host, port, proto=socket.IPPROTO_TCP + ) + if not infos: + raise ValueError("host did not resolve") + for info in infos: + if _is_blocked_address(str(info[4][0])): + raise ValueError("host resolves to a blocked address") + + if parts.scheme == "https": + return url, host_header + + family, _, _, _, sockaddr = infos[0] + address = str(sockaddr[0]) + pinned = f"[{address}]" if family == socket.AF_INET6 else address + if parts.port is not None: + pinned = f"{pinned}:{parts.port}" + return urlunsplit((parts.scheme, pinned, parts.path, parts.query, "")), host_header + + +async def _fetch_image_from_url(url: str) -> bytes | None: + """Fetch the leading bytes of an image, enough to read its dimensions.""" + try: + target, host_header = await _validated_fetch_target(url) + async with httpx.AsyncClient( + timeout=IMAGE_FETCH_TIMEOUT_SECONDS, follow_redirects=False + ) as client: + async with client.stream( + "GET", target, headers={"Host": host_header} + ) as response: + response.raise_for_status() + chunks: list[bytes] = [] + downloaded = 0 + async for chunk in response.aiter_bytes(): + chunks.append(chunk) + downloaded += len(chunk) + if downloaded >= IMAGE_FETCH_MAX_BYTES: + break + return b"".join(chunks)[:IMAGE_FETCH_MAX_BYTES] except Exception as e: logger.warning( "Failed to fetch image from URL", @@ -391,6 +471,7 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: Supports both base64 encoded images and image URLs. """ total_image_tokens = 0 + fetches = 0 for message in messages: if not isinstance(message, dict): @@ -452,7 +533,14 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: extra={"error": str(e)}, ) total_image_tokens += 85 + elif fetches >= IMAGE_FETCH_MAX_PER_REQUEST: + logger.warning( + "Skipping image URL fetch above per-request limit", + extra={"url": url[:100], "limit": IMAGE_FETCH_MAX_PER_REQUEST}, + ) + total_image_tokens += 85 else: + fetches += 1 image_bytes_or_none = await _fetch_image_from_url(url) if image_bytes_or_none: width, height = _get_image_dimensions(image_bytes_or_none) diff --git a/tests/unit/test_image_url_fetch_guard.py b/tests/unit/test_image_url_fetch_guard.py new file mode 100644 index 00000000..4cca987e --- /dev/null +++ b/tests/unit/test_image_url_fetch_guard.py @@ -0,0 +1,167 @@ +"""Guards for the pre-auth image URL fetch used by cost estimation.""" + +import socket +import threading +from http.server import BaseHTTPRequestHandler, HTTPServer +from typing import Any, Callable, Iterator + +import pytest + +from routstr.payment import helpers +from routstr.payment.helpers import ( + IMAGE_FETCH_MAX_BYTES, + IMAGE_FETCH_MAX_PER_REQUEST, + _fetch_image_from_url, + _is_blocked_address, + _validated_fetch_target, + estimate_image_tokens_in_messages, +) + +REQUESTED_PATHS: list[str] = [] + + +class _Loop: + def __init__(self, getaddrinfo: Callable[..., Any]) -> None: + self.getaddrinfo = getaddrinfo + + +class _Sink(BaseHTTPRequestHandler): + def do_GET(self) -> None: # noqa: N802 + REQUESTED_PATHS.append(self.path) + body = b"x" * (IMAGE_FETCH_MAX_BYTES * 4) + self.send_response(200) + self.send_header("Content-Type", "image/png") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + try: + self.wfile.write(body) + except BrokenPipeError: + # The client stops reading once the byte cap is reached. + pass + + def log_message(self, *args: object) -> None: + pass + + +@pytest.fixture +def sink() -> Iterator[str]: + REQUESTED_PATHS.clear() + server = HTTPServer(("127.0.0.1", 0), _Sink) + threading.Thread(target=server.serve_forever, daemon=True).start() + try: + yield f"http://127.0.0.1:{server.server_address[1]}" + finally: + server.shutdown() + server.server_close() + + +@pytest.mark.asyncio +async def test_loopback_url_is_not_fetched(sink: str) -> None: + assert await _fetch_image_from_url(f"{sink}/internal") is None + assert REQUESTED_PATHS == [] + + +@pytest.mark.asyncio +async def test_link_local_metadata_url_is_not_fetched() -> None: + assert await _fetch_image_from_url("http://169.254.169.254/latest/meta-data") is None + + +@pytest.mark.asyncio +async def test_non_http_scheme_is_rejected() -> None: + assert await _fetch_image_from_url("file:///etc/passwd") is None + + +@pytest.mark.parametrize( + "address", + [ + "127.0.0.1", + "10.0.0.1", + "169.254.169.254", + "100.64.0.1", # CGNAT: reachable inside many hosting networks + "192.0.0.1", + "224.0.0.1", + "::1", + "::ffff:127.0.0.1", + "2002:7f00:1::", # 6to4 wrapping 127.0.0.1 + ], +) +def test_non_global_addresses_are_blocked(address: str) -> None: + assert _is_blocked_address(address) is True + + +@pytest.mark.parametrize("address", ["8.8.8.8", "2001:4860:4860::8888"]) +def test_global_addresses_are_allowed(address: str) -> None: + assert _is_blocked_address(address) is False + + +@pytest.mark.asyncio +async def test_http_target_is_pinned_to_validated_address( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def fake_getaddrinfo(*args: object, **kwargs: object) -> list[tuple]: + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 80))] + + monkeypatch.setattr( + helpers.asyncio, "get_running_loop", lambda: _Loop(fake_getaddrinfo) + ) + + target, host_header = await _validated_fetch_target("http://example.com/cat.png") + + assert target == "http://93.184.216.34/cat.png" + assert host_header == "example.com" + + +@pytest.mark.asyncio +async def test_https_target_keeps_hostname_for_tls( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def fake_getaddrinfo(*args: object, **kwargs: object) -> list[tuple]: + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443))] + + monkeypatch.setattr( + helpers.asyncio, "get_running_loop", lambda: _Loop(fake_getaddrinfo) + ) + + target, host_header = await _validated_fetch_target("https://example.com/cat.png") + + assert target == "https://example.com/cat.png" + assert host_header == "example.com" + + +@pytest.fixture +def reachable_sink(sink: str, monkeypatch: pytest.MonkeyPatch) -> str: + """Let the local sink stand in for a public host, so cap tests keep the + address validation intact instead of disabling it.""" + + async def passthrough(url: str) -> tuple[str, str]: + return url, "images.example.com" + + monkeypatch.setattr(helpers, "_validated_fetch_target", passthrough) + return sink + + +@pytest.mark.asyncio +async def test_downloaded_bytes_are_capped(reachable_sink: str) -> None: + body = await _fetch_image_from_url(f"{reachable_sink}/allowed") + + assert body is not None + assert len(body) <= IMAGE_FETCH_MAX_BYTES + assert REQUESTED_PATHS == ["/allowed"] + + +@pytest.mark.asyncio +async def test_url_fetches_are_capped_per_request(reachable_sink: str) -> None: + urls = IMAGE_FETCH_MAX_PER_REQUEST + 3 + messages = [ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": f"{reachable_sink}/{index}"}} + for index in range(urls) + ], + } + ] + + await estimate_image_tokens_in_messages(messages) + + assert len(REQUESTED_PATHS) == IMAGE_FETCH_MAX_PER_REQUEST