Pin LaunchCheck requests to validated DNS

This commit is contained in:
2026-07-18 06:49:58 -03:00
parent 8aae6c3b8f
commit cc71329525
6 changed files with 190 additions and 67 deletions

View File

@@ -1,5 +1,5 @@
name: launch-check-studio-v1
version: 0.1.0
version: 0.1.1
entrypoint: agent:LaunchCheckStudioV1
expose:
public: false
@@ -29,4 +29,4 @@ runtime:
max_runtime_seconds: 120
egress:
deny_internet_by_default: false
tools_used: [httpx, psycopg, postgres]
tools_used: [aiohttp, psycopg, postgres]

216
agent.py
View File

@@ -11,7 +11,7 @@ import socket
from dataclasses import dataclass
from datetime import datetime, timezone
from html.parser import HTMLParser
from typing import Any
from typing import Annotated, Any
from urllib.parse import urljoin, urlparse, urlunparse
from pydantic import BaseModel, Field
@@ -61,6 +61,41 @@ class FetchResult:
error: dict[str, Any] | None = None
class _PinnedResolver:
"""Resolve one hostname only to addresses validated before the request."""
def __init__(self, hostname: str, addresses: tuple[str, ...]) -> None:
self._hostname = hostname.lower().rstrip(".")
self._addresses = addresses
async def resolve(
self,
host: str,
port: int = 0,
family: int = socket.AF_UNSPEC,
) -> list[dict[str, Any]]:
if host.lower().rstrip(".") != self._hostname:
raise OSError("resolver hostname mismatch")
return [
{
"hostname": host,
"host": address,
"port": port,
"family": (
socket.AF_INET6
if ipaddress.ip_address(address).version == 6
else socket.AF_INET
),
"proto": 0,
"flags": 0,
}
for address in self._addresses
]
async def close(self) -> None:
return None
class _HeadParser(HTMLParser):
def __init__(self) -> None:
super().__init__()
@@ -110,7 +145,7 @@ class LaunchCheckStudioV1(A2AAgent[LaunchCheckStudioV1Config, PlatformUserAuth])
"LaunchCheck audits public HTTP(S) launch pages with bounded SSRF-safe "
"network checks and user-scoped Postgres persistence."
)
version = "0.1.0"
version = "0.1.1"
config_model = LaunchCheckStudioV1Config
auth_model = PlatformUserAuth
@@ -119,7 +154,7 @@ class LaunchCheckStudioV1(A2AAgent[LaunchCheckStudioV1Config, PlatformUserAuth])
state_model = LaunchCheckStudioV1State
resources = Resources(cpu="500m", memory="512Mi", max_runtime_seconds=120)
egress = EgressPolicy(deny_internet_by_default=False)
tools_used = ("httpx", "psycopg", "postgres")
tools_used = ("aiohttp", "psycopg", "postgres")
pricing = Pricing(
price_per_call_usd=0.0,
caller_pays_llm=False,
@@ -140,14 +175,17 @@ class LaunchCheckStudioV1(A2AAgent[LaunchCheckStudioV1Config, PlatformUserAuth])
@a2a.tool(
description="Run a bounded SSRF-safe launch readiness audit for a public HTTP(S) URL and persist the report.",
timeout_seconds=60,
idempotent=True,
idempotent=False,
cost_class="network-read",
)
async def audit_url(
self,
ctx: RunContext[PlatformUserAuth],
audit_id: str,
url: str,
audit_id: Annotated[
str,
Field(min_length=1, max_length=128, pattern=r"^[A-Za-z0-9][A-Za-z0-9_.:-]{0,127}$"),
],
url: Annotated[str, Field(min_length=8, max_length=2048)],
) -> dict[str, Any]:
tenant = _tenant_key(ctx)
audit_id = _clean_audit_id(audit_id)
@@ -202,7 +240,10 @@ class LaunchCheckStudioV1(A2AAgent[LaunchCheckStudioV1Config, PlatformUserAuth])
async def get_audit(
self,
ctx: RunContext[PlatformUserAuth],
audit_id: str,
audit_id: Annotated[
str,
Field(min_length=1, max_length=128, pattern=r"^[A-Za-z0-9][A-Za-z0-9_.:-]{0,127}$"),
],
) -> dict[str, Any]:
tenant = _tenant_key(ctx)
audit_id = _clean_audit_id(audit_id)
@@ -273,42 +314,66 @@ async def _validate_public_url(raw_url: str) -> dict[str, Any]:
host = parsed.hostname.strip().rstrip(".")
if not host:
return {"ok": False, "code": "invalid_url", "message": "URL must include a hostname."}
if parsed.port is not None and not (1 <= parsed.port <= 65535):
try:
port = parsed.port
except ValueError:
return {"ok": False, "code": "invalid_url", "message": "URL port is invalid."}
blocked = await _host_is_disallowed(host, parsed.port or (443 if parsed.scheme == "https" else 80))
if blocked:
if port is not None and not (1 <= port <= 65535):
return {"ok": False, "code": "invalid_url", "message": "URL port is invalid."}
try:
host = str(ipaddress.ip_address(host.strip("[]")))
except ValueError:
try:
host = host.encode("idna").decode("ascii").lower()
except UnicodeError:
return {"ok": False, "code": "invalid_url", "message": "URL hostname is invalid."}
addresses = await _resolve_public_addresses(
host,
port or (443 if parsed.scheme == "https" else 80),
)
if addresses is None:
return {"ok": False, "code": "disallowed_private_target", "message": "Target resolves to a non-public network address."}
netloc = host
if parsed.port:
netloc = f"{host}:{parsed.port}"
if not addresses:
return {"ok": False, "code": "dns_resolution_failed", "message": "Target hostname could not be resolved safely."}
netloc = f"[{host}]" if ":" in host else host
if port:
netloc = f"{netloc}:{port}"
normalized = urlunparse((parsed.scheme.lower(), netloc, parsed.path or "", "", parsed.query, ""))
return {"ok": True, "url": normalized}
return {
"ok": True,
"url": normalized,
"hostname": host,
"port": port or (443 if parsed.scheme == "https" else 80),
"addresses": addresses,
}
async def _host_is_disallowed(host: str, port: int) -> bool:
async def _resolve_public_addresses(host: str, port: int) -> tuple[str, ...] | None:
try:
ip = ipaddress.ip_address(host.strip("[]"))
return _ip_disallowed(ip)
return None if _ip_disallowed(ip) else (str(ip),)
except ValueError:
pass
lowered = host.lower()
if lowered in {"localhost", "metadata.google.internal"} or lowered.endswith(".localhost"):
return True
return None
try:
infos = await asyncio.wait_for(
asyncio.get_running_loop().getaddrinfo(host, port, type=socket.SOCK_STREAM),
timeout=DNS_TIMEOUT_SECONDS,
)
except Exception: # DNS failure is not private, fetch will return a structured network error
return False
except Exception:
return ()
addresses: set[str] = set()
for info in infos:
sockaddr = info[4]
if sockaddr:
addresses.add(str(sockaddr[0]))
if not addresses:
return True
return any(_ip_disallowed(ipaddress.ip_address(address)) for address in addresses)
return ()
if any(_ip_disallowed(ipaddress.ip_address(address)) for address in addresses):
return None
return tuple(sorted(addresses))
def _ip_disallowed(ip: ipaddress._BaseAddress) -> bool:
@@ -325,44 +390,76 @@ def _ip_disallowed(ip: ipaddress._BaseAddress) -> bool:
async def _fetch_bounded(url: str, max_redirects: int, max_bytes: int) -> FetchResult:
import httpx
current = url
chain: list[dict[str, Any]] = []
timeout = httpx.Timeout(TOTAL_TIMEOUT_SECONDS, connect=CONNECT_TIMEOUT_SECONDS, read=READ_TIMEOUT_SECONDS)
async with httpx.AsyncClient(timeout=timeout, follow_redirects=False, trust_env=False, headers={"user-agent": USER_AGENT}) as client:
for hop in range(max_redirects + 1):
validation = await _validate_public_url(current)
if not validation["ok"]:
return FetchResult(current, None, {}, b"", chain, {"code": validation["code"], "message": validation["message"]})
try:
async with client.stream("GET", validation["url"]) as response:
body = await _read_limited(response, max_bytes)
except httpx.TimeoutException:
return FetchResult(current, None, {}, b"", chain, {"code": "network_timeout", "message": "The target did not respond within the bounded timeout."})
except httpx.HTTPError:
return FetchResult(current, None, {}, b"", chain, {"code": "network_error", "message": "The target could not be reached safely."})
headers = {k.lower(): v[:1000] for k, v in response.headers.items()}
status = response.status_code
if status in {301, 302, 303, 307, 308} and response.headers.get("location"):
location = response.headers["location"]
next_url = urljoin(str(response.url), location)
chain.append({"from": validation["url"], "to": _safe_url_echo(next_url), "status_code": status})
if hop >= max_redirects:
return FetchResult(validation["url"], status, headers, body, chain, {"code": "too_many_redirects", "message": "Redirect count exceeded the audit limit."})
next_validation = await _validate_public_url(next_url)
if not next_validation["ok"]:
return FetchResult(validation["url"], status, headers, body, chain, {"code": next_validation["code"], "message": next_validation["message"]})
current = next_validation["url"]
continue
return FetchResult(str(response.url), status, headers, body, chain)
for hop in range(max_redirects + 1):
validation = await _validate_public_url(current)
if not validation["ok"]:
return FetchResult(current, None, {}, b"", chain, {"code": validation["code"], "message": validation["message"]})
try:
status, headers, body, response_url = await _request_pinned(
validation,
max_bytes=max_bytes,
total_timeout=TOTAL_TIMEOUT_SECONDS,
connect_timeout=CONNECT_TIMEOUT_SECONDS,
read_timeout=READ_TIMEOUT_SECONDS,
)
except TimeoutError:
return FetchResult(current, None, {}, b"", chain, {"code": "network_timeout", "message": "The target did not respond within the bounded timeout."})
except Exception: # noqa: BLE001 - return a stable public network error
return FetchResult(current, None, {}, b"", chain, {"code": "network_error", "message": "The target could not be reached safely."})
if status in {301, 302, 303, 307, 308} and headers.get("location"):
next_url = urljoin(response_url, headers["location"])
chain.append({"from": validation["url"], "to": _safe_url_echo(next_url), "status_code": status})
if hop >= max_redirects:
return FetchResult(validation["url"], status, headers, body, chain, {"code": "too_many_redirects", "message": "Redirect count exceeded the audit limit."})
current = next_url
continue
return FetchResult(response_url, status, headers, body, chain)
return FetchResult(current, None, {}, b"", chain, {"code": "too_many_redirects", "message": "Redirect count exceeded the audit limit."})
async def _request_pinned(
validation: dict[str, Any],
*,
max_bytes: int,
total_timeout: float,
connect_timeout: float,
read_timeout: float,
) -> tuple[int, dict[str, str], bytes, str]:
import aiohttp
resolver = _PinnedResolver(
str(validation["hostname"]),
tuple(str(item) for item in validation["addresses"]),
)
connector = aiohttp.TCPConnector(resolver=resolver, use_dns_cache=False)
timeout = aiohttp.ClientTimeout(
total=total_timeout,
connect=connect_timeout,
sock_connect=connect_timeout,
sock_read=read_timeout,
)
try:
async with aiohttp.ClientSession(
connector=connector,
connector_owner=True,
timeout=timeout,
trust_env=False,
headers={"user-agent": USER_AGENT, "accept-encoding": "identity"},
) as client:
async with client.get(validation["url"], allow_redirects=False) as response:
body = await _read_limited(response, max_bytes)
headers = {key.lower(): value[:1000] for key, value in response.headers.items()}
return response.status, headers, body, str(response.url)
except asyncio.TimeoutError as exc:
raise TimeoutError("bounded request timed out") from exc
async def _read_limited(response: Any, max_bytes: int) -> bytes:
chunks: list[bytes] = []
total = 0
async for chunk in response.aiter_bytes():
async for chunk in response.content.iter_chunked(64 * 1024):
total += len(chunk)
if total > max_bytes:
remaining = max_bytes - (total - len(chunk))
@@ -379,16 +476,17 @@ async def _check_robots(final_url: str) -> dict[str, Any]:
validation = await _validate_public_url(robots_url)
if not validation["ok"]:
return {"present": False, "status_code": None, "checked_url": _safe_url_echo(robots_url)}
import httpx
timeout = httpx.Timeout(6.0, connect=2.0, read=3.0)
try:
async with httpx.AsyncClient(timeout=timeout, follow_redirects=False, trust_env=False, headers={"user-agent": USER_AGENT}) as client:
async with client.stream("GET", validation["url"]) as response:
body = await _read_limited(response, MAX_ROBOTS_BYTES)
status, _headers, body, _response_url = await _request_pinned(
validation,
max_bytes=MAX_ROBOTS_BYTES,
total_timeout=6.0,
connect_timeout=2.0,
read_timeout=3.0,
)
return {
"present": response.status_code == 200,
"status_code": response.status_code,
"present": status == 200,
"status_code": status,
"checked_url": validation["url"],
"bytes_read": len(body),
}

View File

@@ -1,7 +1,7 @@
{
"name": "launch-check-studio-v1-frontend",
"private": true,
"version": "0.1.0",
"version": "0.1.1",
"type": "module",
"scripts": {
"dev": "vite --host 0.0.0.0 --port 5173",

View File

@@ -96,7 +96,7 @@ export function App() {
HTTPS, page metadata, robots.txt, and key security headers, then saves the evidence for your account.
</p>
<div className="trust-bar" aria-label="Safety guarantees">
<span>SSRF-safe DNS + redirects</span>
<span>Pinned public DNS + safe redirects</span>
<span>User-scoped Postgres</span>
<span>No LLM or secrets in browser</span>
</div>

View File

@@ -1,4 +1,4 @@
# a2a-pack is auto-installed by the deploy build.
httpx>=0.27
aiohttp>=3.10,<4
psycopg[binary]>=3.2
pytest>=8.0

View File

@@ -57,8 +57,8 @@ async def test_acceptance_success_reload_and_receipt(monkeypatch):
os.environ.pop("DATABASE_URL", None)
stored = {}
async def fake_host_allowed(host, port):
return False
async def fake_public_addresses(host, port):
return ("93.184.216.34",)
async def fake_fetch(url, max_redirects, max_bytes):
return FetchResult(
@@ -92,7 +92,7 @@ async def test_acceptance_success_reload_and_receipt(monkeypatch):
async def fake_load(tenant, audit_id):
return stored.get((tenant, audit_id))
monkeypatch.setattr(launch_agent, "_host_is_disallowed", fake_host_allowed)
monkeypatch.setattr(launch_agent, "_resolve_public_addresses", fake_public_addresses)
monkeypatch.setattr(launch_agent, "_fetch_bounded", fake_fetch)
monkeypatch.setattr(launch_agent, "_check_robots", fake_robots)
monkeypatch.setattr(launch_agent, "_save_audit_and_receipt", fake_save)
@@ -146,3 +146,28 @@ async def test_url_policy_rejects_dangerous_inputs(monkeypatch):
monkeypatch.setattr(launch_agent.asyncio, "get_running_loop", lambda: FakeLoop())
assert (await launch_agent._validate_public_url("https://rebinding.test"))["code"] == "disallowed_private_target"
@pytest.mark.asyncio
async def test_validated_dns_addresses_are_pinned_for_connection(monkeypatch):
calls = {"count": 0}
async def fake_getaddrinfo(host, port, type):
calls["count"] += 1
return [(None, None, None, None, ("93.184.216.34", port))]
class FakeLoop:
def getaddrinfo(self, host, port, type):
return fake_getaddrinfo(host, port, type)
monkeypatch.setattr(launch_agent.asyncio, "get_running_loop", lambda: FakeLoop())
validation = await launch_agent._validate_public_url("https://rebinding.test/path")
assert validation["ok"] is True
assert validation["addresses"] == ("93.184.216.34",)
resolver = launch_agent._PinnedResolver(
validation["hostname"], validation["addresses"]
)
resolved = await resolver.resolve(validation["hostname"], validation["port"])
assert [item["host"] for item in resolved] == ["93.184.216.34"]
assert calls["count"] == 1