From ca0a4a2153d49f5155a7814333baa8792aa9ca29 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Mon, 5 Oct 2026 20:23:08 +0000 Subject: [PATCH] fix: validate every redirect hop before fetching InputFile URLs Follow InputFile URL redirects manually and resolve/validate each hop before connecting so a public host cannot bounce the hosted server onto loopback, private, or link-local addresses. Pin TCP to the validated IP while keeping Host and TLS SNI on the original hostname. Return a generic fetch error to callers so redirect failures are not an SSRF oracle. Co-authored-by: Jake Barnby --- src/mcp_server_appwrite/server.py | 220 +++++++++++++----- tests/unit/test_server.py | 362 +++++++++++++++++++++++++++++- 2 files changed, 523 insertions(+), 59 deletions(-) diff --git a/src/mcp_server_appwrite/server.py b/src/mcp_server_appwrite/server.py index 5b57d98..a3ba90e 100644 --- a/src/mcp_server_appwrite/server.py +++ b/src/mcp_server_appwrite/server.py @@ -22,7 +22,7 @@ from pathlib import Path from types import UnionType from typing import Any, Union, cast, get_args, get_origin -from urllib.parse import unquote, urlsplit, urlunsplit +from urllib.parse import unquote, urljoin, urlsplit, urlunsplit import httpx import mcp.server.stdio @@ -504,12 +504,51 @@ def _configure_uploads(transport: str) -> None: _UPLOAD_TRANSPORT = transport -def _validate_fetch_url(url: str) -> None: - """Reject non-http(s) schemes and hosts that resolve to non-public addresses. +_REDIRECT_STATUS_CODES = frozenset({301, 302, 303, 307, 308}) + + +class _FetchFailed(Exception): + """Internal fetch failure. Logged server-side; callers see a generic ValueError.""" + + +def _ip_is_blocked(ip: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool: + if isinstance(ip, ipaddress.IPv6Address) and ip.ipv4_mapped is not None: + ip = ip.ipv4_mapped + return ( + ip.is_private + or ip.is_loopback + or ip.is_link_local + or ip.is_reserved + or ip.is_multicast + or ip.is_unspecified + ) + + +@dataclass(frozen=True) +class _PinnedFetchTarget: + """A URL rewritten to a validated IP, with Host/SNI kept on the original host.""" + + request_url: str + headers: dict[str, str] + extensions: dict[str, Any] - This is the SSRF guard for server-side URL fetches: it stops the model from making - the hosted server reach internal services, loopback, or the cloud metadata endpoint - (169.254.169.254). Note the resolve-then-reconnect DNS-rebinding gap is accepted. + +def _ascii_hostname(host: str) -> str: + try: + return host.encode("idna").decode("ascii") + except UnicodeError: + return host + + +def _pin_fetch_target(url: str) -> _PinnedFetchTarget: + """Resolve once, refuse non-public IPs, and pin TCP to a validated address. + + This is the SSRF guard for server-side URL fetches: it stops the model from + making the hosted server reach internal services, loopback, or the cloud + metadata endpoint (169.254.169.254). Connecting to the IP we just validated + (while sending the original hostname as Host and TLS SNI) closes the + resolve-then-reconnect DNS-rebinding gap without weakening certificate + verification. """ parts = urlsplit(url) if parts.scheme not in ("http", "https"): @@ -520,27 +559,70 @@ def _validate_fetch_url(url: str) -> None: if not host: raise ValueError("URL is missing a host.") - port = parts.port or (443 if parts.scheme == "https" else 80) try: - infos = socket.getaddrinfo(host, port) + port = parts.port or (443 if parts.scheme == "https" else 80) + except ValueError as exc: + raise ValueError("URL has an invalid port.") from exc + + try: + infos = socket.getaddrinfo(host, port, type=socket.SOCK_STREAM) except socket.gaierror as exc: raise ValueError(f"Could not resolve host '{host}'.") from exc + if not infos: + raise ValueError(f"Could not resolve host '{host}'.") - for info in infos: - ip = ipaddress.ip_address(info[4][0]) - if ( - ip.is_private - or ip.is_loopback - or ip.is_link_local - or ip.is_reserved - or ip.is_multicast - or ip.is_unspecified - ): + addrs = [str(info[4][0]) for info in infos] + for addr in addrs: + if _ip_is_blocked(ipaddress.ip_address(addr)): raise ValueError( "Refusing to fetch a URL that resolves to a private, loopback, or " "link-local address." ) + try: + ipaddress.ip_address(host) + except ValueError: + pass + else: + request_url = urlunsplit( + (parts.scheme, parts.netloc, parts.path, parts.query, "") + ) + return _PinnedFetchTarget(request_url, {}, {}) + + ipv4_addrs = [addr for addr in addrs if ":" not in addr] + pinned_ip = ipv4_addrs[0] if ipv4_addrs else addrs[0] + host_for_url = f"[{pinned_ip}]" if ":" in pinned_ip else pinned_ip + default_port = 443 if parts.scheme == "https" else 80 + netloc = host_for_url if port == default_port else f"{host_for_url}:{port}" + request_url = urlunsplit((parts.scheme, netloc, parts.path, parts.query, "")) + + ascii_host = _ascii_hostname(host) + host_header = f"{ascii_host}:{parts.port}" if parts.port is not None else ascii_host + headers = {"Host": host_header} + extensions: dict[str, Any] = {} + if parts.scheme == "https": + extensions["sni_hostname"] = ascii_host + return _PinnedFetchTarget(request_url, headers, extensions) + + +def _validate_fetch_url(url: str) -> None: + """Reject non-http(s) schemes and hosts that resolve to non-public addresses.""" + _pin_fetch_target(url) + + +def _response_header(resp: httpx.Response, name: str) -> str: + value = resp.headers.get(name) + if value is not None: + return value + for key, header_value in resp.headers.items(): + if str(key).lower() == name.lower(): + return str(header_value) + return "" + + +def _generic_fetch_error(param_name: str) -> ValueError: + return ValueError(f"Failed to fetch file from URL for '{param_name}'.") + def _derive_filename(resp: httpx.Response, url: str) -> str: """Best-effort filename from Content-Disposition, then the URL path, then a fallback.""" @@ -564,50 +646,76 @@ def _derive_filename(resp: httpx.Response, url: str) -> str: def _fetch_input_file(url: str, param_name: str) -> InputFile: - """Download a public URL (SSRF-guarded, size-capped) into an in-memory InputFile.""" - _validate_fetch_url(url) + """Download a public URL (SSRF-guarded, size-capped) into an in-memory InputFile. + + Redirects are followed manually. Every hop — the caller-supplied URL and each + ``Location`` — is resolved and validated *before* a connection to that hop is + opened, so a public host cannot bounce the server onto a private address. + """ + # Caller-supplied URL keeps a specific validation error (scheme / DNS / private). + target = _pin_fetch_target(url) + logical_url = url + try: with httpx.Client( timeout=FETCH_TIMEOUT_SECONDS, - follow_redirects=True, - max_redirects=FETCH_MAX_REDIRECTS, + follow_redirects=False, limits=httpx.Limits(max_connections=1), ) as client: - with client.stream("GET", url) as resp: - resp.raise_for_status() - # The final URL after redirects must also be public. - _validate_fetch_url(str(resp.url)) - - declared = resp.headers.get("content-length") - if declared is not None and declared.isdigit(): - if int(declared) > MAX_FETCH_BYTES: - raise ValueError( - f"File at URL for '{param_name}' is too large " - f"({declared} bytes); max is {MAX_FETCH_BYTES} bytes." - ) - - chunks: list[bytes] = [] - total = 0 - for chunk in resp.iter_bytes(): - total += len(chunk) - if total > MAX_FETCH_BYTES: - raise ValueError( - f"File at URL for '{param_name}' exceeds the max of " - f"{MAX_FETCH_BYTES} bytes." - ) - chunks.append(chunk) - - data = b"".join(chunks) - mime_type = ( - (resp.headers.get("content-type") or "").split(";")[0].strip() - ) - filename = _derive_filename(resp, url) - except httpx.HTTPError as exc: - raise ValueError( - f"Failed to fetch file from URL for '{param_name}': {exc}" - ) from exc + for redirects in range(FETCH_MAX_REDIRECTS + 1): + with client.stream( + "GET", + target.request_url, + headers=target.headers or None, + extensions=target.extensions or None, + ) as resp: + location = _response_header(resp, "location").strip() + if resp.status_code in _REDIRECT_STATUS_CODES and location: + if redirects >= FETCH_MAX_REDIRECTS: + raise _FetchFailed("too many redirects") + next_url = urljoin(logical_url, location) + try: + target = _pin_fetch_target(next_url) + except ValueError as exc: + raise _FetchFailed(str(exc)) from exc + logical_url = next_url + continue + + if resp.status_code >= 400: + raise _FetchFailed(f"upstream status {resp.status_code}") + + declared = resp.headers.get("content-length") + if declared is not None and declared.isdigit(): + if int(declared) > MAX_FETCH_BYTES: + raise ValueError( + f"File at URL for '{param_name}' is too large " + f"({declared} bytes); max is {MAX_FETCH_BYTES} bytes." + ) + + chunks: list[bytes] = [] + total = 0 + for chunk in resp.iter_bytes(): + total += len(chunk) + if total > MAX_FETCH_BYTES: + raise ValueError( + f"File at URL for '{param_name}' exceeds the max of " + f"{MAX_FETCH_BYTES} bytes." + ) + chunks.append(chunk) + + data = b"".join(chunks) + mime_type = ( + (resp.headers.get("content-type") or "").split(";")[0].strip() + ) + filename = _derive_filename(resp, url) + return InputFile.from_bytes(data, filename, mime_type or None) - return InputFile.from_bytes(data, filename, mime_type or None) + raise _FetchFailed("too many redirects") + except ValueError: + raise + except Exception as exc: + _log_startup(f"InputFile URL fetch failed for '{param_name}': {exc}") + raise _generic_fetch_error(param_name) from exc def _coerce_inline_content(value: Mapping, param_name: str) -> InputFile: diff --git a/tests/unit/test_server.py b/tests/unit/test_server.py index dceb504..773c416 100644 --- a/tests/unit/test_server.py +++ b/tests/unit/test_server.py @@ -1,10 +1,14 @@ import asyncio import base64 +import http.server import io +import ipaddress import json import os +import socket import sys import tempfile +import threading import time import unittest from pathlib import Path @@ -85,9 +89,11 @@ def __exit__(self, *args): class _FakeClient: - def __init__(self, response): + def __init__(self, response, responses=None): + self._responses = list(responses) if responses is not None else None self._response = response self.stream_kwargs = None + self.requested_urls = [] def __enter__(self): return self @@ -96,10 +102,89 @@ def __exit__(self, *args): return False def stream(self, method, url, **kwargs): + self.requested_urls.append(url) self.stream_kwargs = {"method": method, "url": url, **kwargs} + if self._responses is not None: + if not self._responses: + raise AssertionError(f"unexpected extra fetch: {url}") + return _FakeStream(self._responses.pop(0)) return _FakeStream(self._response) +class _SilentHTTPRequestHandler(http.server.BaseHTTPRequestHandler): + def log_message(self, format, *args): + return + + +def _start_local_http(do_get): + hits = [] + + class Handler(_SilentHTTPRequestHandler): + protocol_version = "HTTP/1.0" + + def do_GET(self): + hits.append({"path": self.path, "host": self.headers.get("Host")}) + do_get(self) + + server = http.server.ThreadingHTTPServer(("127.0.0.1", 0), Handler) + ready = threading.Event() + + def serve(): + ready.set() + server.serve_forever() + + thread = threading.Thread(target=serve, daemon=True) + thread.start() + ready.wait(timeout=1) + return server, hits + + +def _write_http(handler, status, headers=None, body=b""): + handler.send_response(status) + for key, value in (headers or {}).items(): + handler.send_header(key, value) + handler.send_header("Content-Length", str(len(body))) + handler.end_headers() + if body: + handler.wfile.write(body) + + +def _public_dns(host_map, default_ip="93.184.216.34"): + real = socket.getaddrinfo + + def wrapped(host, port, *args, **kwargs): + if host in host_map: + ip = host_map[host] + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", (ip, port or 0))] + if default_ip is not None and not _is_ip_literal(host): + return [ + (socket.AF_INET, socket.SOCK_STREAM, 6, "", (default_ip, port or 0)) + ] + return real(host, port, *args, **kwargs) + + return wrapped + + +def _is_ip_literal(host): + try: + ipaddress.ip_address(host) + return True + except ValueError: + return False + + +def _remap_create_connection(ip_to_addr): + real = socket.create_connection + + def wrapped(address, *args, **kwargs): + host = address[0] + if host in ip_to_addr: + address = ip_to_addr[host] + return real(address, *args, **kwargs) + + return wrapped + + class ServerHelperTests(unittest.TestCase): def test_parse_args_defaults_to_stdio(self): with patch.dict(os.environ, {}, clear=True): @@ -1484,17 +1569,37 @@ def setUp(self): def tearDown(self): _configure_uploads("stdio") - def _patch_fetch(self, response, addrinfo=_PUBLIC_ADDRINFO): + def _patch_fetch(self, response, addrinfo=_PUBLIC_ADDRINFO, responses=None): return ( patch( "mcp_server_appwrite.server.socket.getaddrinfo", return_value=addrinfo ), patch( "mcp_server_appwrite.server.httpx.Client", - return_value=_FakeClient(response), + return_value=_FakeClient(response, responses=responses), ), ) + def _assert_generic_fetch_error(self, exc, param="file"): + message = str(exc) + self.assertEqual(message, f"Failed to fetch file from URL for '{param}'.") + lowered = message.lower() + for leaked in ( + "127.0.0.1", + "169.254", + "10.0.0.", + "private", + "loopback", + "link-local", + "connection", + "refused", + "timeout", + "404", + "500", + "status", + ): + self.assertNotIn(leaked, lowered) + def test_url_object_uses_content_disposition_filename(self): response = _FakeResponse( data=b"\x89PNG\r\n", @@ -1588,6 +1693,257 @@ def test_http_instructions_mention_url_upload(self): self.assertIn("upload", http.lower()) self.assertNotIn("upload", stdio.lower()) + def test_url_fetch_pins_validated_ip_and_preserves_host_sni(self): + response = _FakeResponse(data=b"abc", headers={"content-type": "image/png"}) + addr, client = self._patch_fetch(response) + with addr, client as client_mock: + _coerce_argument("file", "https://example.com/dir/a.png", InputFile) + + fake = client_mock.return_value + self.assertEqual(fake.stream_kwargs["url"], "https://93.184.216.34/dir/a.png") + self.assertEqual(fake.stream_kwargs["headers"]["Host"], "example.com") + self.assertEqual( + fake.stream_kwargs["extensions"]["sni_hostname"], "example.com" + ) + + def test_public_to_public_redirect_downloads_final_body(self): + redirect = _FakeResponse( + status_code=302, + headers={"location": "https://cdn.example/file.png"}, + ) + final = _FakeResponse( + data=b"png-bytes", + headers={"content-type": "image/png"}, + url="https://cdn.example/file.png", + ) + addr, client = self._patch_fetch(redirect, responses=[redirect, final]) + with addr, client as client_mock: + coerced = _coerce_argument("file", "https://example.com/x", InputFile) + + self.assertEqual(coerced.data, b"png-bytes") + self.assertEqual(coerced.filename, "x") + self.assertEqual( + client_mock.return_value.requested_urls, + ["https://93.184.216.34/x", "https://93.184.216.34/file.png"], + ) + + def test_relative_redirect_joins_against_logical_url(self): + redirect = _FakeResponse( + status_code=302, + headers={"location": "/dir/final.png"}, + ) + final = _FakeResponse( + data=b"ok", + headers={"content-type": "image/png"}, + ) + addr, client = self._patch_fetch(redirect, responses=[redirect, final]) + with addr, client as client_mock: + coerced = _coerce_argument("file", "https://example.com/start", InputFile) + + self.assertEqual(coerced.data, b"ok") + self.assertEqual( + client_mock.return_value.requested_urls, + ["https://93.184.216.34/start", "https://93.184.216.34/dir/final.png"], + ) + + def test_redirect_cap_is_enforced(self): + redirect = _FakeResponse( + status_code=302, + headers={"location": "https://example.com/next"}, + ) + addr, client = self._patch_fetch(redirect) + with ( + addr, + client as client_mock, + patch.object(server_module, "FETCH_MAX_REDIRECTS", 1), + ): + with self.assertRaises(ValueError) as ctx: + _coerce_argument("file", "https://example.com/x", InputFile) + + self._assert_generic_fetch_error(ctx.exception) + self.assertEqual(len(client_mock.return_value.requested_urls), 2) + + def test_redirect_to_private_is_generic_and_does_not_follow(self): + redirect = _FakeResponse( + status_code=302, + headers={"location": "http://127.0.0.1/secret"}, + ) + addr = patch( + "mcp_server_appwrite.server.socket.getaddrinfo", + side_effect=_public_dns({}, default_ip="93.184.216.34"), + ) + client = patch( + "mcp_server_appwrite.server.httpx.Client", + return_value=_FakeClient(redirect), + ) + with addr, client as client_mock: + with self.assertRaises(ValueError) as ctx: + _coerce_argument("file", "https://example.com/x", InputFile) + + self._assert_generic_fetch_error(ctx.exception) + self.assertEqual( + client_mock.return_value.requested_urls, ["https://93.184.216.34/x"] + ) + + def test_http_error_and_connection_error_use_the_same_generic_message(self): + not_found = _FakeResponse(status_code=404, url="https://example.com/missing") + addr, client = self._patch_fetch(not_found) + with addr, client: + with self.assertRaises(ValueError) as http_ctx: + _coerce_argument("file", "https://example.com/missing", InputFile) + + class _BoomClient(_FakeClient): + def stream(self, method, url, **kwargs): + raise httpx.ConnectError("Connection refused to 10.0.0.1:8080") + + addr = patch( + "mcp_server_appwrite.server.socket.getaddrinfo", + return_value=_PUBLIC_ADDRINFO, + ) + client = patch( + "mcp_server_appwrite.server.httpx.Client", + return_value=_BoomClient(_FakeResponse()), + ) + with addr, client: + with self.assertRaises(ValueError) as conn_ctx: + _coerce_argument("file", "https://example.com/x", InputFile) + + self._assert_generic_fetch_error(http_ctx.exception) + self._assert_generic_fetch_error(conn_ctx.exception) + self.assertEqual(str(http_ctx.exception), str(conn_ctx.exception)) + + def test_url_fetch_rejects_ipv4_mapped_loopback(self): + response = _FakeResponse(data=b"secret") + addr, client = self._patch_fetch( + response, + addrinfo=[(None, None, None, None, ("::ffff:127.0.0.1", 80))], + ) + with addr, client as client_mock: + with self.assertRaises(ValueError) as ctx: + _coerce_argument("file", {"url": "https://evil.example/x"}, InputFile) + self.assertIn("private", str(ctx.exception).lower()) + client_mock.assert_not_called() + + def test_redirect_to_internal_does_not_hit_listener(self): + cases = ( + ("127.0.0.1", "http://127.0.0.1:{port}/secret"), + ("10.0.0.1", "http://10.0.0.1:{port}/secret"), + ("169.254.169.254", "http://169.254.169.254:{port}/latest"), + ) + for internal_host, location_template in cases: + with self.subTest(internal_host=internal_host): + self._assert_redirect_skips_internal_listener( + internal_host, location_template + ) + + def _assert_redirect_skips_internal_listener( + self, internal_host, location_template + ): + def public_get(handler): + location = location_template.format(port=internal_server.server_address[1]) + _write_http(handler, 302, {"Location": location}) + + def internal_get(handler): + _write_http(handler, 200, {"Content-Type": "text/plain"}, b"secret") + + public_server, public_hits = _start_local_http(public_get) + internal_server, internal_hits = _start_local_http(internal_get) + public_ip = "93.184.216.34" + try: + public_port = public_server.server_address[1] + dns = patch( + "mcp_server_appwrite.server.socket.getaddrinfo", + side_effect=_public_dns({"files.example": public_ip}, default_ip=None), + ) + connect = patch( + "socket.create_connection", + side_effect=_remap_create_connection( + { + public_ip: ("127.0.0.1", public_port), + "10.0.0.1": ( + "127.0.0.1", + internal_server.server_address[1], + ), + "169.254.169.254": ( + "127.0.0.1", + internal_server.server_address[1], + ), + } + ), + ) + with dns, connect: + with self.assertRaises(ValueError) as ctx: + _coerce_argument( + "file", + f"http://files.example:{public_port}/file", + InputFile, + ) + self._assert_generic_fetch_error(ctx.exception) + self.assertEqual(len(public_hits), 1) + self.assertEqual(internal_hits, []) + finally: + public_server.shutdown() + internal_server.shutdown() + public_server.server_close() + internal_server.server_close() + + def test_public_to_public_redirect_works_against_local_servers(self): + def first_get(handler): + location = f"http://cdn.example:{second_server.server_address[1]}/file.png" + _write_http(handler, 302, {"Location": location}) + + def second_get(handler): + _write_http( + handler, + 200, + { + "Content-Type": "image/png", + "Content-Disposition": 'filename="pic.png"', + }, + b"png-bytes", + ) + + first_server, first_hits = _start_local_http(first_get) + second_server, second_hits = _start_local_http(second_get) + first_ip = "93.184.216.34" + second_ip = "93.184.216.35" + try: + first_port = first_server.server_address[1] + second_port = second_server.server_address[1] + dns = patch( + "mcp_server_appwrite.server.socket.getaddrinfo", + side_effect=_public_dns( + {"files.example": first_ip, "cdn.example": second_ip}, + default_ip=None, + ), + ) + connect = patch( + "socket.create_connection", + side_effect=_remap_create_connection( + { + first_ip: ("127.0.0.1", first_port), + second_ip: ("127.0.0.1", second_port), + } + ), + ) + with dns, connect: + coerced = _coerce_argument( + "file", + f"http://files.example:{first_port}/start", + InputFile, + ) + self.assertEqual(coerced.data, b"png-bytes") + self.assertEqual(coerced.filename, "pic.png") + self.assertEqual(len(first_hits), 1) + self.assertEqual(first_hits[0]["host"], f"files.example:{first_port}") + self.assertEqual(len(second_hits), 1) + self.assertEqual(second_hits[0]["host"], f"cdn.example:{second_port}") + finally: + first_server.shutdown() + second_server.shutdown() + first_server.server_close() + second_server.server_close() + class RegionRoutingTests(unittest.TestCase): BASE = "https://cloud.appwrite.io/v1"