c46ae17084
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.
297 lines
11 KiB
Python
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
|
|
)
|