Files
EvoScientist/tests/test_safe_transport.py
m4 c46ae17084 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.
2026-07-20 21:25:15 +08:00

297 lines
11 KiB
Python

"""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
)