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.
This commit is contained in:
m4
2026-07-20 21:25:15 +08:00
parent c21fc0a272
commit c46ae17084
8 changed files with 1143 additions and 2 deletions
+24 -2
View File
@@ -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",
]
@@ -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
+2
View File
@@ -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,
@@ -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,
)
+26
View File
@@ -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) ---
+296
View File
@@ -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 == []
+1
View File
@@ -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,
+296
View File
@@ -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
)