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:
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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) ---
|
||||
|
||||
|
||||
|
||||
@@ -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 == []
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
)
|
||||
Reference in New Issue
Block a user