|
|
|
|
@@ -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:
|
|
|
|
|
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 httpx.HTTPError:
|
|
|
|
|
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."})
|
|
|
|
|
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)
|
|
|
|
|
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."})
|
|
|
|
|
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"]
|
|
|
|
|
current = next_url
|
|
|
|
|
continue
|
|
|
|
|
return FetchResult(str(response.url), status, headers, body, chain)
|
|
|
|
|
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),
|
|
|
|
|
}
|
|
|
|
|
|