Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
220 changes: 164 additions & 56 deletions src/mcp_server_appwrite/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"):
Expand All @@ -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."""
Expand All @@ -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:
Expand Down
Loading
Loading