From c46ae17084fdfd4772a1b8126ada672907520e37 Mon Sep 17 00:00:00 2001 From: m4 Date: Mon, 20 Jul 2026 21:25:15 +0800 Subject: [PATCH] feat(model-registry): add EndpointPolicy and SafeHttpTransport SSRF defenses EndpointPolicy validates provider base URLs (section 4.3): public https endpoints with hostname and optional port pass; loopback, private, link-local, multicast, unspecified, and cloud-metadata addresses are denied unless the normalized URL exactly matches a registered development_endpoints entry (no prefix or wildcard matching). URLs with user info, fragments, or non-http(s) schemes are rejected with the new stable 422 code ENDPOINT_NOT_ALLOWED. SafeHttpTransport is the single network egress for adapters: a custom httpcore NetworkBackend resolves DNS under control on every connect (retries included), filters denied ranges, and connects directly to the selected IP, while TLS SNI/certificate checks and the HTTP Host header keep the original hostname. Redirects and env proxies are disabled; every request origin re-passes URL-layer validation before any I/O. --- EvoScientist/model_registry/__init__.py | 26 +- .../model_registry/endpoint_policy.py | 206 ++++++++++++ EvoScientist/model_registry/errors.py | 2 + EvoScientist/model_registry/safe_transport.py | 292 +++++++++++++++++ EvoScientist/model_registry/schemas.py | 26 ++ tests/test_endpoint_policy.py | 296 ++++++++++++++++++ tests/test_model_registry_schemas.py | 1 + tests/test_safe_transport.py | 296 ++++++++++++++++++ 8 files changed, 1143 insertions(+), 2 deletions(-) create mode 100644 EvoScientist/model_registry/endpoint_policy.py create mode 100644 EvoScientist/model_registry/safe_transport.py create mode 100644 tests/test_endpoint_policy.py create mode 100644 tests/test_safe_transport.py diff --git a/EvoScientist/model_registry/__init__.py b/EvoScientist/model_registry/__init__.py index 7c9e583..13b79ae 100644 --- a/EvoScientist/model_registry/__init__.py +++ b/EvoScientist/model_registry/__init__.py @@ -1,14 +1,25 @@ -"""Unified model registry: RegistryV4 schema, error taxonomy, and SQLite store. +"""Unified model registry: schema, errors, store, and network egress policy. This subpackage is the single authority for the version 4 model registry (design doc sections 4.2, 4.3, 5.1, 8.2, 9.5). HTTP JSON, SQLite JSON, import -tooling, and test fixtures all reuse these Pydantic models. +tooling, and test fixtures all reuse these Pydantic models. EndpointPolicy +and SafeHttpTransport form the SSRF defense and the only network egress +for adapters. """ from __future__ import annotations +from .endpoint_policy import EndpointPolicy from .errors import ERROR_HTTP_STATUS, ErrorDetail, ErrorPayload, ModelRegistryError from .hashing import configuration_hash +from .safe_transport import ( + AsyncSafeHttpTransport, + AsyncSafeNetworkBackend, + SafeHttpTransport, + SafeNetworkBackend, + build_safe_async_http_client, + build_safe_http_client, +) from .schemas import ( AdapterParameterSpec, AuthConfig, @@ -17,6 +28,8 @@ from .schemas import ( Capabilities, CredentialStatus, CredentialWrite, + DevelopmentEndpoint, + EndpointPolicyPublic, ModelAvailability, ModelConfig, ModelRef, @@ -33,12 +46,17 @@ from .store import ModelRuntimeStore, SharedStorageError __all__ = [ "ERROR_HTTP_STATUS", "AdapterParameterSpec", + "AsyncSafeHttpTransport", + "AsyncSafeNetworkBackend", "AuthConfig", "AuthRef", "AuthSpec", "Capabilities", "CredentialStatus", "CredentialWrite", + "DevelopmentEndpoint", + "EndpointPolicy", + "EndpointPolicyPublic", "ErrorDetail", "ErrorPayload", "ModelAvailability", @@ -52,7 +70,11 @@ __all__ = [ "ProviderRuntimeConfig", "RegistryV4", "ResolvedModelConfig", + "SafeHttpTransport", + "SafeNetworkBackend", "SharedStorageError", "VerificationInfo", + "build_safe_async_http_client", + "build_safe_http_client", "configuration_hash", ] diff --git a/EvoScientist/model_registry/endpoint_policy.py b/EvoScientist/model_registry/endpoint_policy.py new file mode 100644 index 0000000..097a97d --- /dev/null +++ b/EvoScientist/model_registry/endpoint_policy.py @@ -0,0 +1,206 @@ +"""EndpointPolicy: the SSRF allow rules for provider base URLs (section 4.3). + +Public endpoints require https plus a hostname with an optional port. The +deny list covers loopback, private, link-local, multicast, unspecified, and +cloud-metadata (169.254.169.254) addresses. Local addresses over http or +https are only allowed when the normalized URL exactly matches a platform +``development_endpoints`` entry — an exact full-string match, never a +prefix or wildcard match. + +Normalization lowercases the scheme and host, drops the default port, and +strips trailing ``/`` characters. URLs carrying user info, a fragment, or a +non-http(s) scheme are always rejected. +""" + +from __future__ import annotations + +import ipaddress +import typing +from collections.abc import Iterable +from urllib.parse import SplitResult, urlsplit + +from .errors import ENDPOINT_NOT_ALLOWED, ModelRegistryError +from .schemas import DevelopmentEndpoint, EndpointPolicyPublic + +_DEFAULT_PORTS = {"http": 80, "https": 443} + +CLOUD_METADATA_IP = ipaddress.ip_address("169.254.169.254") + + +def denied_network_reason(ip: str) -> str | None: + """Return the deny-list category for ``ip``, or ``None`` if allowed.""" + try: + address = ipaddress.ip_address(ip) + except ValueError: + return None + if address == CLOUD_METADATA_IP: + return "cloud_metadata" + if address.is_loopback: + return "loopback" + if address.is_link_local: + return "link_local" + if address.is_multicast: + return "multicast" + if address.is_unspecified: + return "unspecified" + if address.is_private: + return "private" + return None + + +def _reject(message: str) -> ModelRegistryError: + return ModelRegistryError(ENDPOINT_NOT_ALLOWED, message) + + +def _host_text(host: str) -> str: + return f"[{host}]" if ":" in host else host + + +def _normalize_parts( + scheme: str, + host: str, + port: int | None, + path: str, + query: str, +) -> str: + netloc = _host_text(host) + if port is not None and port != _DEFAULT_PORTS[scheme]: + netloc = f"{netloc}:{port}" + normalized = f"{scheme}://{netloc}{path.rstrip('/')}" + if query: + normalized = f"{normalized}?{query}" + return normalized + + +class EndpointPolicy: + """Validates provider base URLs against the platform endpoint rules.""" + + def __init__( + self, development_endpoints: Iterable[DevelopmentEndpoint] = () + ) -> None: + self._entries = list(development_endpoints) + self._registered: dict[str, DevelopmentEndpoint] = {} + for entry in self._entries: + try: + normalized = self._parse_and_normalize(entry.url) + except ModelRegistryError as exc: + raise ValueError( + f"Invalid development endpoint {entry.id!r}: {exc.message}" + ) from exc + if normalized in self._registered: + raise ValueError( + f"Duplicate development endpoint registration: {entry.url!r}." + ) + self._registered[normalized] = entry + # Request origins (scheme://host:port) implied by the registered + # entries. Base URLs must match an entry exactly, but the traffic an + # allowed base URL produces necessarily spans the whole origin. + self._registered_origins = { + self._origin_of(normalized) for normalized in self._registered + } + + @staticmethod + def _origin_of(normalized: str) -> str: + parts = urlsplit(normalized) + return _normalize_parts(parts.scheme, parts.hostname or "", parts.port, "", "") + + @staticmethod + def _parse_and_normalize(url: str) -> str: + parts = _checked_split(url) + return _normalize_parts( + parts.scheme.lower(), + parts.hostname or "", + parts.port, + parts.path, + parts.query, + ) + + def validate_base_url(self, url: str) -> str: + """Return the normalized base URL, or raise ``ENDPOINT_NOT_ALLOWED``.""" + parts = _checked_split(url) + scheme = parts.scheme.lower() + host = parts.hostname or "" + normalized = _normalize_parts(scheme, host, parts.port, parts.path, parts.query) + return self._validate(normalized, scheme, host, self._registered) + + def validate_request_origin(self, origin: str) -> str: + """Validate a request origin (``scheme://host[:port]``) before any I/O. + + An origin is allowed when a registered development entry implies it or + when it satisfies the public https rules on its own. + """ + parts = _checked_split(origin) + scheme = parts.scheme.lower() + host = parts.hostname or "" + normalized = _normalize_parts(scheme, host, parts.port, "", "") + return self._validate(normalized, scheme, host, self._registered_origins) + + @staticmethod + def _validate( + normalized: str, + scheme: str, + host: str, + allowed: typing.Container[str], + ) -> str: + if normalized in allowed: + return normalized + if scheme != "https": + raise _reject( + "Base URL must use https unless it exactly matches a registered " + "development endpoint." + ) + reason = denied_network_reason(host) + if reason is not None: + raise _reject( + f"Base URL host falls in the denied {reason} range and is not a " + "registered development endpoint." + ) + return normalized + + def is_development_endpoint(self, url: str) -> bool: + """Return True when ``url`` normalizes to a registered entry.""" + try: + normalized = self._parse_and_normalize(url) + except ModelRegistryError: + return False + return normalized in self._registered + + def allows_denied_network(self, host: str, port: int) -> bool: + """Return True when ``host:port`` is a registered development target. + + Used by the transport at connect time: registered local endpoints may + resolve to deny-listed addresses; every other target is filtered. + The scheme is unknown at connect time, so both default-port spellings + are considered. + """ + host = host.lower() + for scheme, default_port in _DEFAULT_PORTS.items(): + normalized = _normalize_parts( + scheme, host, None if port == default_port else port, "", "" + ) + if normalized in self._registered_origins: + return True + return False + + def public_view(self) -> EndpointPolicyPublic: + """Return the section 9.1 browser-facing view of this policy.""" + return EndpointPolicyPublic(development_endpoints=list(self._entries)) + + +def _checked_split(url: str) -> SplitResult: + """Split ``url`` and reject structurally disallowed forms.""" + try: + parts = urlsplit(url) + # Accessing .port raises ValueError for out-of-range or text ports. + _ = parts.port + except ValueError as exc: + raise _reject("Base URL is malformed.") from exc + if parts.scheme.lower() not in _DEFAULT_PORTS: + raise _reject("Base URL must use the http or https scheme.") + if parts.username is not None or parts.password is not None: + raise _reject("Base URL must not contain user info.") + if parts.fragment: + raise _reject("Base URL must not contain a fragment.") + if not parts.hostname: + raise _reject("Base URL must include a hostname.") + return parts diff --git a/EvoScientist/model_registry/errors.py b/EvoScientist/model_registry/errors.py index 75bea91..fac815e 100644 --- a/EvoScientist/model_registry/errors.py +++ b/EvoScientist/model_registry/errors.py @@ -39,6 +39,7 @@ CREDENTIAL_REJECTED = "CREDENTIAL_REJECTED" RUN_CREDENTIAL_REVISION_UNAVAILABLE = "RUN_CREDENTIAL_REVISION_UNAVAILABLE" AUTH_MODE_UNSUPPORTED = "AUTH_MODE_UNSUPPORTED" ADAPTER_NOT_SUPPORTED = "ADAPTER_NOT_SUPPORTED" +ENDPOINT_NOT_ALLOWED = "ENDPOINT_NOT_ALLOWED" CAPABILITY_UNSUPPORTED_BY_ADAPTER = "CAPABILITY_UNSUPPORTED_BY_ADAPTER" MODEL_CAPABILITY_UNAVAILABLE = "MODEL_CAPABILITY_UNAVAILABLE" UNSUPPORTED_RUNTIME_PARAMETER = "UNSUPPORTED_RUNTIME_PARAMETER" @@ -64,6 +65,7 @@ ERROR_HTTP_STATUS: dict[str, int] = { RUN_CREDENTIAL_REVISION_UNAVAILABLE: 422, AUTH_MODE_UNSUPPORTED: 422, ADAPTER_NOT_SUPPORTED: 422, + ENDPOINT_NOT_ALLOWED: 422, CAPABILITY_UNSUPPORTED_BY_ADAPTER: 422, MODEL_CAPABILITY_UNAVAILABLE: 422, UNSUPPORTED_RUNTIME_PARAMETER: 422, diff --git a/EvoScientist/model_registry/safe_transport.py b/EvoScientist/model_registry/safe_transport.py new file mode 100644 index 0000000..13761e3 --- /dev/null +++ b/EvoScientist/model_registry/safe_transport.py @@ -0,0 +1,292 @@ +"""SafeHttpTransport: the single network egress for all adapters (section 4.3). + +Provider tests and real model calls must share this transport instead of an +SDK's default HTTP client. The policy is enforced in two layers, in order: + +1. URL layer — every request origin passes + ``EndpointPolicy.validate_request_origin`` before any I/O, so unregistered + local endpoints never open a socket. +2. IP layer — a custom ``httpcore.NetworkBackend`` performs controlled DNS + resolution inside ``connect_tcp`` on every connect (retries included), + filters the deny-listed ranges, and connects directly to the selected IP. + TLS then runs on the origin hostname, so SNI and certificate hostname + checks keep the original name, and HTTP keeps the original Host header. + +Redirects are disabled on the built clients; even if a caller enables them, +every redirected origin re-passes both layers. +""" + +from __future__ import annotations + +import functools +import socket +import ssl +import typing + +import anyio +import httpcore +import httpx +from httpcore._backends.anyio import AnyIOStream +from httpcore._backends.base import ( + SOCKET_OPTION, + AsyncNetworkBackend, + AsyncNetworkStream, + NetworkBackend, + NetworkStream, +) +from httpcore._backends.sync import SyncStream +from httpcore._exceptions import ( + ConnectError, + ConnectTimeout, + ExceptionMapping, + map_exceptions, +) +from httpx._config import DEFAULT_LIMITS + +from .endpoint_policy import EndpointPolicy, denied_network_reason + +GetAddrInfo = typing.Callable[..., list] + + +def _select_address( + policy: EndpointPolicy, + host: str, + port: int, + infos: list, +) -> tuple[int, tuple]: + """Pick the first policy-compliant ``getaddrinfo`` result.""" + allow_denied = policy.allows_denied_network(host, port) + for family, _type, _proto, _canonname, sockaddr in infos: + if allow_denied or denied_network_reason(sockaddr[0]) is None: + return family, sockaddr + raise ConnectError( + "DNS resolution for the endpoint yielded only denied network addresses." + ) + + +class SafeNetworkBackend(NetworkBackend): + """Sync ``httpcore.NetworkBackend`` with controlled DNS and IP filtering.""" + + def __init__( + self, + policy: EndpointPolicy, + *, + getaddrinfo: GetAddrInfo = socket.getaddrinfo, + ) -> None: + self._policy = policy + self._getaddrinfo = getaddrinfo + # Test hook: the (ip, port) selected by the most recent connect. + self.last_selected: tuple[str, int] | None = None + + def resolve(self, host: str, port: int) -> tuple[int, tuple]: + """Resolve ``host`` and return the selected ``(family, sockaddr)``.""" + infos = self._getaddrinfo(host, port, type=socket.SOCK_STREAM) + return _select_address(self._policy, host, port, infos) + + def connect_tcp( + self, + host: str, + port: int, + timeout: float | None = None, + local_address: str | None = None, + socket_options: typing.Iterable[SOCKET_OPTION] | None = None, + ) -> NetworkStream: + family, sockaddr = self.resolve(host, port) + self.last_selected = (sockaddr[0], sockaddr[1]) + exc_map: ExceptionMapping = { + socket.timeout: ConnectTimeout, + OSError: ConnectError, + } + with map_exceptions(exc_map): + sock = socket.socket(family, socket.SOCK_STREAM) + try: + if local_address is not None: + sock.bind((local_address, 0)) + sock.settimeout(timeout) + # Connect to the validated IP directly — no second DNS lookup. + sock.connect(sockaddr) + for option in socket_options or (): + sock.setsockopt(*option) + sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + except BaseException: + sock.close() + raise + return SyncStream(sock) + + def connect_unix_socket( + self, + path: str, + timeout: float | None = None, + socket_options: typing.Iterable[SOCKET_OPTION] | None = None, + ) -> NetworkStream: + raise ConnectError("Unix domain sockets are not an allowed endpoint.") + + +class AsyncSafeNetworkBackend(AsyncNetworkBackend): + """Async ``httpcore.AsyncNetworkBackend`` with the same controls.""" + + def __init__( + self, + policy: EndpointPolicy, + *, + getaddrinfo: GetAddrInfo = socket.getaddrinfo, + ) -> None: + self._policy = policy + self._getaddrinfo = getaddrinfo + self.last_selected: tuple[str, int] | None = None + + async def resolve(self, host: str, port: int) -> tuple[int, tuple]: + lookup = functools.partial( + self._getaddrinfo, host, port, type=socket.SOCK_STREAM + ) + infos = await anyio.to_thread.run_sync(lookup) + return _select_address(self._policy, host, port, infos) + + async def connect_tcp( + self, + host: str, + port: int, + timeout: float | None = None, + local_address: str | None = None, + socket_options: typing.Iterable[SOCKET_OPTION] | None = None, + ) -> AsyncNetworkStream: + _family, sockaddr = await self.resolve(host, port) + self.last_selected = (sockaddr[0], sockaddr[1]) + exc_map: ExceptionMapping = { + TimeoutError: ConnectTimeout, + OSError: ConnectError, + anyio.BrokenResourceError: ConnectError, + } + with map_exceptions(exc_map): + with anyio.fail_after(timeout): + # anyio skips DNS for IP literals, so the validated IP is the + # actual connection peer. + stream: anyio.abc.ByteStream = await anyio.connect_tcp( + remote_host=sockaddr[0], + remote_port=sockaddr[1], + local_host=local_address, + ) + for option in socket_options or (): + stream._raw_socket.setsockopt(*option) # type: ignore[attr-defined] + return AnyIOStream(stream) + + async def connect_unix_socket( + self, + path: str, + timeout: float | None = None, + socket_options: typing.Iterable[SOCKET_OPTION] | None = None, + ) -> AsyncNetworkStream: + raise ConnectError("Unix domain sockets are not an allowed endpoint.") + + +def _origin_string(url: httpx.URL) -> str: + host = url.host + if ":" in host: + host = f"[{host}]" + netloc = host if url.port is None else f"{host}:{url.port}" + return f"{url.scheme}://{netloc}" + + +class SafeHttpTransport(httpx.HTTPTransport): + """Sync httpx transport enforcing EndpointPolicy at the URL and IP layer.""" + + def __init__( + self, + policy: EndpointPolicy, + *, + retries: int = 0, + limits: httpx.Limits = DEFAULT_LIMITS, + ssl_context: ssl.SSLContext | None = None, + backend: SafeNetworkBackend | None = None, + getaddrinfo: GetAddrInfo = socket.getaddrinfo, + ) -> None: + ssl_context = ssl_context or ssl.create_default_context() + super().__init__( + verify=ssl_context, limits=limits, retries=retries, trust_env=False + ) + self._policy = policy + # httpx.HTTPTransport has no network_backend hook, so the pool is + # rebuilt with the safe backend. The request/response mapping stays + # entirely with httpx. + self._pool = httpcore.ConnectionPool( + ssl_context=ssl_context, + max_connections=limits.max_connections, + max_keepalive_connections=limits.max_keepalive_connections, + keepalive_expiry=limits.keepalive_expiry, + retries=retries, + network_backend=backend + if backend is not None + else SafeNetworkBackend(policy, getaddrinfo=getaddrinfo), + ) + + def handle_request(self, request: httpx.Request) -> httpx.Response: + self._policy.validate_request_origin(_origin_string(request.url)) + return super().handle_request(request) + + +class AsyncSafeHttpTransport(httpx.AsyncHTTPTransport): + """Async httpx transport enforcing EndpointPolicy at the URL and IP layer.""" + + def __init__( + self, + policy: EndpointPolicy, + *, + retries: int = 0, + limits: httpx.Limits = DEFAULT_LIMITS, + ssl_context: ssl.SSLContext | None = None, + backend: AsyncSafeNetworkBackend | None = None, + getaddrinfo: GetAddrInfo = socket.getaddrinfo, + ) -> None: + ssl_context = ssl_context or ssl.create_default_context() + super().__init__( + verify=ssl_context, limits=limits, retries=retries, trust_env=False + ) + self._policy = policy + self._pool = httpcore.AsyncConnectionPool( + ssl_context=ssl_context, + max_connections=limits.max_connections, + max_keepalive_connections=limits.max_keepalive_connections, + keepalive_expiry=limits.keepalive_expiry, + retries=retries, + network_backend=backend + if backend is not None + else AsyncSafeNetworkBackend(policy, getaddrinfo=getaddrinfo), + ) + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + self._policy.validate_request_origin(_origin_string(request.url)) + return await super().handle_async_request(request) + + +def build_safe_http_client( + policy: EndpointPolicy, + timeout: float | httpx.Timeout = 120.0, + *, + retries: int = 0, + getaddrinfo: GetAddrInfo = socket.getaddrinfo, +) -> httpx.Client: + """Build the sync client every adapter and provider test must share.""" + return httpx.Client( + transport=SafeHttpTransport(policy, retries=retries, getaddrinfo=getaddrinfo), + follow_redirects=False, + trust_env=False, + timeout=timeout, + ) + + +def build_safe_async_http_client( + policy: EndpointPolicy, + timeout: float | httpx.Timeout = 120.0, + *, + retries: int = 0, + getaddrinfo: GetAddrInfo = socket.getaddrinfo, +) -> httpx.AsyncClient: + """Build the async client every adapter and provider test must share.""" + return httpx.AsyncClient( + transport=AsyncSafeHttpTransport( + policy, retries=retries, getaddrinfo=getaddrinfo + ), + follow_redirects=False, + trust_env=False, + timeout=timeout, + ) diff --git a/EvoScientist/model_registry/schemas.py b/EvoScientist/model_registry/schemas.py index 28ba92f..5d44ceb 100644 --- a/EvoScientist/model_registry/schemas.py +++ b/EvoScientist/model_registry/schemas.py @@ -339,6 +339,32 @@ class ResolvedModelConfig(BaseModel): effective_capabilities: Capabilities +# --- Endpoint policy (sections 4.3, 9.1) --- + + +class DevelopmentEndpoint(BaseModel): + """One local endpoint pre-registered by the deployment administrator. + + Only base URLs that exactly match a registered entry (after EndpointPolicy + normalization) may point at loopback or private network addresses. + """ + + id: NonEmptyString + url: NonEmptyString + label: NonEmptyString + + +class EndpointPolicyPublic(BaseModel): + """The browser-facing view of the endpoint policy (section 9.1). + + Exposes only the selectable local endpoints; network rules and secrets + are never part of this shape. + """ + + public_https_allowed: Literal[True] = True + development_endpoints: list[DevelopmentEndpoint] = Field(default_factory=list) + + # --- API helper models (sections 9.1, 9.2) --- diff --git a/tests/test_endpoint_policy.py b/tests/test_endpoint_policy.py new file mode 100644 index 0000000..588c0f7 --- /dev/null +++ b/tests/test_endpoint_policy.py @@ -0,0 +1,296 @@ +"""Tests for the EndpointPolicy URL allow rules (design doc section 4.3). + +Public endpoints require https plus a hostname with an optional port. +Loopback, private, link-local, multicast, unspecified, and cloud-metadata +addresses are denied unless the exact normalized URL is registered as a +platform ``development_endpoints`` entry. Matching is an exact full-string +match after normalization — no prefix or wildcard matching. +""" + +from __future__ import annotations + +import pytest + +from EvoScientist.model_registry.endpoint_policy import EndpointPolicy +from EvoScientist.model_registry.errors import ENDPOINT_NOT_ALLOWED, ModelRegistryError +from EvoScientist.model_registry.schemas import ( + DevelopmentEndpoint, + EndpointPolicyPublic, +) + + +def _ollama_policy() -> EndpointPolicy: + return EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="ollama", url="http://localhost:11434", label="Local Ollama" + ) + ] + ) + + +class TestPublicEndpoints: + @pytest.mark.parametrize( + ("url", "normalized"), + [ + ("https://api.example.com", "https://api.example.com"), + ("HTTPS://API.Example.COM/", "https://api.example.com"), + ("https://api.example.com:443", "https://api.example.com"), + ("https://api.example.com:8443/v1/", "https://api.example.com:8443/v1"), + ("https://8.8.8.8", "https://8.8.8.8"), + ], + ) + def test_public_https_passes(self, url, normalized): + policy = EndpointPolicy() + assert policy.validate_base_url(url) == normalized + + def test_http_public_rejected(self): + policy = EndpointPolicy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url("http://api.example.com") + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + assert excinfo.value.http_status == 422 + + @pytest.mark.parametrize( + "url", + [ + # loopback + "https://127.0.0.1", + "https://127.0.0.1:8443", + "http://127.0.0.1:8080", + "https://[::1]", + # private + "https://10.0.0.5", + "https://172.16.0.1", + "https://192.168.1.1", + "https://[fd00::1]", + # link-local + "https://169.254.1.1", + "https://[fe80::1]", + # multicast + "https://224.0.0.1", + # unspecified + "https://0.0.0.0", + "https://[::]", + # cloud metadata + "https://169.254.169.254", + "http://169.254.169.254", + ], + ) + def test_denied_ip_literals_rejected(self, url): + policy = EndpointPolicy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url(url) + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + + @pytest.mark.parametrize( + "url", + [ + "https://user:pass@api.example.com", + "https://user@api.example.com", + "http://user@localhost:11434", + ], + ) + def test_userinfo_rejected(self, url): + policy = _ollama_policy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url(url) + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + + def test_fragment_rejected(self): + policy = EndpointPolicy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url("https://api.example.com/#section") + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + + @pytest.mark.parametrize( + "url", + [ + "ftp://example.com", + "file:///etc/passwd", + "wss://example.com", + "example.com", + ], + ) + def test_disallowed_scheme_rejected(self, url): + policy = EndpointPolicy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url(url) + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + + @pytest.mark.parametrize("url", ["https://", "https://:8443", ""]) + def test_missing_host_rejected(self, url): + policy = EndpointPolicy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url(url) + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + + def test_error_message_is_safe(self): + policy = EndpointPolicy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url("https://user:secret@127.0.0.1/x") + assert "secret" not in excinfo.value.message + + +class TestDevelopmentEndpoints: + def test_registered_local_endpoint_passes(self): + policy = _ollama_policy() + assert ( + policy.validate_base_url("http://localhost:11434") + == "http://localhost:11434" + ) + + def test_unregistered_local_endpoint_rejected(self): + policy = EndpointPolicy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url("http://localhost:11434") + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + + @pytest.mark.parametrize( + "url", + [ + # different port — exact match, not prefix or wildcard + "http://localhost:11435", + # extra path segment — no prefix matching + "http://localhost:11434/api", + # a different host that merely shares the prefix + "http://localhost:11434.evil.com", + ], + ) + def test_exact_match_only(self, url): + policy = _ollama_policy() + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_base_url(url) + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + + @pytest.mark.parametrize( + "url", + [ + "HTTP://LOCALHOST:11434/", + "http://Localhost:11434", + ], + ) + def test_normalization_equivalence_hits_same_entry(self, url): + policy = _ollama_policy() + assert policy.validate_base_url(url) == "http://localhost:11434" + + def test_default_port_normalization_hits_same_entry(self): + policy = EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="local-http", url="http://127.0.0.1", label="Local HTTP" + ) + ] + ) + assert policy.validate_base_url("http://127.0.0.1:80/") == "http://127.0.0.1" + + def test_registered_private_ip_passes(self): + policy = EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="lan", url="http://192.168.1.10:8080", label="LAN box" + ) + ] + ) + assert ( + policy.validate_base_url("http://192.168.1.10:8080/") + == "http://192.168.1.10:8080" + ) + + def test_registration_normalization_applies_to_entries(self): + policy = EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="ollama", url="HTTP://LOCALHOST:11434/", label="Local Ollama" + ) + ] + ) + assert ( + policy.validate_base_url("http://localhost:11434") + == "http://localhost:11434" + ) + + def test_entry_with_path_matches_exactly(self): + policy = EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="lan", url="http://localhost:11434/api", label="LAN API" + ) + ] + ) + assert ( + policy.validate_base_url("http://localhost:11434/api/") + == "http://localhost:11434/api" + ) + with pytest.raises(ModelRegistryError): + policy.validate_base_url("http://localhost:11434/other") + + def test_entry_with_path_allows_whole_origin_for_requests(self): + policy = EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="lan", url="http://localhost:11434/api", label="LAN API" + ) + ] + ) + assert ( + policy.validate_request_origin("http://localhost:11434") + == "http://localhost:11434" + ) + + def test_request_origin_uses_public_rules_without_registration(self): + policy = _ollama_policy() + assert ( + policy.validate_request_origin("https://api.example.com") + == "https://api.example.com" + ) + with pytest.raises(ModelRegistryError) as excinfo: + policy.validate_request_origin("http://localhost:11435") + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + + def test_duplicate_registration_rejected(self): + with pytest.raises(ValueError, match=r"[Dd]uplicate"): + EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="a", url="http://localhost:11434", label="A" + ), + DevelopmentEndpoint( + id="b", url="HTTP://LOCALHOST:11434/", label="B" + ), + ] + ) + + @pytest.mark.parametrize( + "url", + ["ftp://localhost:11434", "https://user@localhost:11434", "not-a-url"], + ) + def test_invalid_registration_rejected(self, url): + with pytest.raises(ValueError, match="development endpoint"): + EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint(id="bad", url=url, label="Bad") + ] + ) + + +class TestPublicView: + def test_public_view_shape(self): + policy = _ollama_policy() + view = policy.public_view() + assert isinstance(view, EndpointPolicyPublic) + assert view.model_dump() == { + "public_https_allowed": True, + "development_endpoints": [ + { + "id": "ollama", + "url": "http://localhost:11434", + "label": "Local Ollama", + } + ], + } + + def test_public_view_empty(self): + view = EndpointPolicy().public_view() + assert view.public_https_allowed is True + assert view.development_endpoints == [] diff --git a/tests/test_model_registry_schemas.py b/tests/test_model_registry_schemas.py index a86a5db..fe67eff 100644 --- a/tests/test_model_registry_schemas.py +++ b/tests/test_model_registry_schemas.py @@ -59,6 +59,7 @@ ALL_ERROR_CODES = { "RUN_CREDENTIAL_REVISION_UNAVAILABLE": 422, "AUTH_MODE_UNSUPPORTED": 422, "ADAPTER_NOT_SUPPORTED": 422, + "ENDPOINT_NOT_ALLOWED": 422, "CAPABILITY_UNSUPPORTED_BY_ADAPTER": 422, "MODEL_CAPABILITY_UNAVAILABLE": 422, "UNSUPPORTED_RUNTIME_PARAMETER": 422, diff --git a/tests/test_safe_transport.py b/tests/test_safe_transport.py new file mode 100644 index 0000000..0564160 --- /dev/null +++ b/tests/test_safe_transport.py @@ -0,0 +1,296 @@ +"""Tests for SafeHttpTransport: the single network egress (design doc 4.3). + +A controlled fake ``getaddrinfo`` plus a local threaded HTTP server cover +the DNS-rebinding defenses: resolution and network filtering run on every +connect, denied ranges never receive bytes, registered development +endpoints stay reachable with the original Host header, and redirects are +never followed. +""" + +from __future__ import annotations + +import json +import socket +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +import httpcore +import httpx +import pytest + +from EvoScientist.model_registry.endpoint_policy import EndpointPolicy +from EvoScientist.model_registry.errors import ENDPOINT_NOT_ALLOWED, ModelRegistryError +from EvoScientist.model_registry.safe_transport import ( + AsyncSafeHttpTransport, + AsyncSafeNetworkBackend, + SafeHttpTransport, + SafeNetworkBackend, + build_safe_async_http_client, + build_safe_http_client, +) +from EvoScientist.model_registry.schemas import DevelopmentEndpoint + + +class _Handler(BaseHTTPRequestHandler): + def do_GET(self): + self.server.requests.append( + {"path": self.path, "host": self.headers.get("Host")} + ) + if self.path == "/redirect": + self.send_response(302) + self.send_header("Location", "/target") + self.send_header("Content-Length", "0") + self.end_headers() + return + body = json.dumps({"ok": True}).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, *args): + pass + + +@pytest.fixture +def local_server(): + server = ThreadingHTTPServer(("127.0.0.1", 0), _Handler) + server.requests = [] + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + yield server + server.shutdown() + server.server_close() + thread.join() + + +def _make_fake_getaddrinfo(mapping, counter): + """Return a getaddrinfo stand-in mapping hostnames to fixed IP lists.""" + + def fake(host, port, *args, **kwargs): + counter["calls"] += 1 + return [ + (socket.AF_INET, socket.SOCK_STREAM, 6, "", (ip, port)) + for ip in mapping[host] + ] + + return fake + + +def _dev_policy(port: int) -> EndpointPolicy: + return EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="local", url=f"http://localhost:{port}", label="Local server" + ) + ] + ) + + +class TestResolve: + def test_resolve_filters_denied_ips(self): + counter = {"calls": 0} + backend = SafeNetworkBackend( + EndpointPolicy(), + getaddrinfo=_make_fake_getaddrinfo( + {"host.test": ["10.0.0.1", "8.8.8.8"]}, counter + ), + ) + _family, sockaddr = backend.resolve("host.test", 443) + assert sockaddr == ("8.8.8.8", 443) + + def test_resolve_rejects_when_only_denied_ips(self): + counter = {"calls": 0} + backend = SafeNetworkBackend( + EndpointPolicy(), + getaddrinfo=_make_fake_getaddrinfo( + {"host.test": ["169.254.169.254", "127.0.0.1"]}, counter + ), + ) + with pytest.raises(httpcore.ConnectError): + backend.resolve("host.test", 443) + + def test_resolve_allows_denied_ip_for_registered_development_endpoint(self): + policy = _dev_policy(port=11434) + counter = {"calls": 0} + backend = SafeNetworkBackend( + policy, + getaddrinfo=_make_fake_getaddrinfo({"localhost": ["127.0.0.1"]}, counter), + ) + _family, sockaddr = backend.resolve("localhost", 11434) + assert sockaddr == ("127.0.0.1", 11434) + + +class TestSyncTransport: + def test_denied_ip_fails_without_sending_request(self, local_server): + counter = {"calls": 0} + policy = EndpointPolicy() + transport = SafeHttpTransport( + policy, + getaddrinfo=_make_fake_getaddrinfo( + {"example.test": ["127.0.0.1"]}, counter + ), + ) + client = httpx.Client(transport=transport, follow_redirects=False) + with pytest.raises(httpx.ConnectError): + client.get("https://example.test/") + assert counter["calls"] == 1 + assert local_server.requests == [] + + def test_cloud_metadata_ip_blocked(self): + counter = {"calls": 0} + transport = SafeHttpTransport( + EndpointPolicy(), + getaddrinfo=_make_fake_getaddrinfo( + {"example.test": ["169.254.169.254"]}, counter + ), + ) + client = httpx.Client(transport=transport, follow_redirects=False) + with pytest.raises(httpx.ConnectError): + client.get("https://example.test/latest/meta-data") + + def test_every_retry_resolves_again(self): + counter = {"calls": 0} + transport = SafeHttpTransport( + EndpointPolicy(), + retries=1, + getaddrinfo=_make_fake_getaddrinfo({"retry.test": ["10.0.0.1"]}, counter), + ) + client = httpx.Client(transport=transport, follow_redirects=False) + with pytest.raises(httpx.ConnectError): + client.get("https://retry.test/") + # Initial attempt plus one retry; each connect re-runs DNS and checks. + assert counter["calls"] == 2 + + def test_development_endpoint_connects_with_original_host(self, local_server): + port = local_server.server_address[1] + policy = _dev_policy(port) + counter = {"calls": 0} + backend = SafeNetworkBackend( + policy, + getaddrinfo=_make_fake_getaddrinfo({"localhost": ["127.0.0.1"]}, counter), + ) + transport = SafeHttpTransport(policy, backend=backend) + client = httpx.Client(transport=transport, follow_redirects=False) + response = client.get(f"http://localhost:{port}/models") + assert response.status_code == 200 + # The connection peer is exactly the selected, policy-checked IP. + assert backend.last_selected == ("127.0.0.1", port) + # The HTTP Host header keeps the original hostname, not the IP. + assert local_server.requests == [ + {"path": "/models", "host": f"localhost:{port}"} + ] + + def test_unregistered_local_endpoint_rejected_before_connect(self, local_server): + port = local_server.server_address[1] + counter = {"calls": 0} + transport = SafeHttpTransport( + EndpointPolicy(), + getaddrinfo=_make_fake_getaddrinfo({"localhost": ["127.0.0.1"]}, counter), + ) + client = httpx.Client(transport=transport, follow_redirects=False) + with pytest.raises(ModelRegistryError) as excinfo: + client.get(f"http://localhost:{port}/") + assert excinfo.value.code == ENDPOINT_NOT_ALLOWED + assert counter["calls"] == 0 + assert local_server.requests == [] + + def test_development_entry_with_path_allows_origin_traffic(self, local_server): + port = local_server.server_address[1] + policy = EndpointPolicy( + development_endpoints=[ + DevelopmentEndpoint( + id="local", + url=f"http://localhost:{port}/api", + label="Local server", + ) + ] + ) + transport = SafeHttpTransport( + policy, + getaddrinfo=_make_fake_getaddrinfo( + {"localhost": ["127.0.0.1"]}, {"calls": 0} + ), + ) + client = httpx.Client(transport=transport, follow_redirects=False) + response = client.get(f"http://localhost:{port}/api/models") + assert response.status_code == 200 + + def test_redirect_not_followed(self, local_server): + port = local_server.server_address[1] + policy = _dev_policy(port) + client = build_safe_http_client( + policy, + getaddrinfo=_make_fake_getaddrinfo( + {"localhost": ["127.0.0.1"]}, {"calls": 0} + ), + ) + response = client.get(f"http://localhost:{port}/redirect") + assert response.status_code == 302 + assert local_server.requests == [ + {"path": "/redirect", "host": f"localhost:{port}"} + ] + client.close() + + def test_build_safe_http_client_defaults(self): + policy = EndpointPolicy() + client = build_safe_http_client(policy, timeout=5.0) + assert isinstance(client, httpx.Client) + assert isinstance(client._transport, SafeHttpTransport) + assert client.follow_redirects is False + assert client.trust_env is False + assert client.timeout == httpx.Timeout(5.0) + client.close() + + +class TestAsyncTransport: + async def test_development_endpoint_connects_with_original_host(self, local_server): + port = local_server.server_address[1] + policy = _dev_policy(port) + counter = {"calls": 0} + transport = AsyncSafeHttpTransport( + policy, + getaddrinfo=_make_fake_getaddrinfo({"localhost": ["127.0.0.1"]}, counter), + ) + async with httpx.AsyncClient( + transport=transport, follow_redirects=False + ) as client: + response = await client.get(f"http://localhost:{port}/models") + assert response.status_code == 200 + assert local_server.requests == [ + {"path": "/models", "host": f"localhost:{port}"} + ] + + async def test_denied_ip_fails_without_sending_request(self, local_server): + counter = {"calls": 0} + transport = AsyncSafeHttpTransport( + EndpointPolicy(), + getaddrinfo=_make_fake_getaddrinfo( + {"example.test": ["169.254.169.254"]}, counter + ), + ) + async with httpx.AsyncClient( + transport=transport, follow_redirects=False + ) as client: + with pytest.raises(httpx.ConnectError): + await client.get("https://example.test/") + assert local_server.requests == [] + + async def test_build_safe_async_http_client_defaults(self): + client = build_safe_async_http_client(EndpointPolicy(), timeout=5.0) + assert isinstance(client, httpx.AsyncClient) + assert isinstance(client._transport, AsyncSafeHttpTransport) + assert client.follow_redirects is False + assert client.trust_env is False + await client.aclose() + + +class TestBackendSelectionVisibleToTests: + def test_sync_backend_is_httpcore_backend(self): + assert isinstance(SafeNetworkBackend(EndpointPolicy()), httpcore.NetworkBackend) + + def test_async_backend_is_httpcore_backend(self): + assert isinstance( + AsyncSafeNetworkBackend(EndpointPolicy()), httpcore.AsyncNetworkBackend + )