diff --git a/agent.py b/agent.py index 6a9b865..86f6523 100644 --- a/agent.py +++ b/agent.py @@ -1,189 +1,603 @@ -"""launch-check-studio-v1 agent. - -Starter stack: - - DeepAgents for tool-calling orchestration - - Caller-provided LLM credentials via ctx.llm - - A tiny model-call middleware hook you can replace with tracing, - routing, rate limits, or policy checks -""" +"""LaunchCheck Studio: SSRF-safe launch readiness audits with user-scoped persistence.""" from __future__ import annotations +import asyncio +import hashlib +import ipaddress import json -from pathlib import Path +import os +import re +import socket +from dataclasses import dataclass +from datetime import datetime, timezone +from html.parser import HTMLParser from typing import Any +from urllib.parse import urljoin, urlparse, urlunparse -from pydantic import BaseModel +from pydantic import BaseModel, Field import a2a_pack as a2a from a2a_pack import ( A2AAgent, - LLMProvisioning, - {{ auth_type }}, + AgentDatabase, + AgentDatabaseEnv, + AgentDatabaseMigrations, + AgentPlatformResources, + EgressPolicy, + PlatformUserAuth, Pricing, + Resources, RunContext, - WorkspaceAccess, - WorkspaceMode, + State, ) -from a2a_pack.context import LLMCreds + +MAX_REDIRECTS = 4 +MAX_RESPONSE_BYTES = 512_000 +MAX_ROBOTS_BYTES = 64_000 +DNS_TIMEOUT_SECONDS = 2.0 +CONNECT_TIMEOUT_SECONDS = 3.0 +READ_TIMEOUT_SECONDS = 5.0 +TOTAL_TIMEOUT_SECONDS = 12.0 +USER_AGENT = "LaunchCheckStudio/0.1 (+https://a2acloud.io)" +DB_CONNECT_OPTIONS = "-c statement_timeout=5000 -c lock_timeout=2000 -c idle_in_transaction_session_timeout=5000" class LaunchCheckStudioV1Config(BaseModel): - pass + max_redirects: int = Field(default=MAX_REDIRECTS, ge=0, le=5) + max_response_bytes: int = Field(default=MAX_RESPONSE_BYTES, ge=32_000, le=1_000_000) -SYSTEM_PROMPT = """\ -You are a compact tool-calling agent. - -Use the text_stats tool when the user asks about text, counts, summaries, -or anything where exact length/word numbers would help. Mention tool results -briefly instead of dumping raw JSON. -""" - -RUNTIME_SKILLS_DIR = "launch-check-studio-v1/.deepagents/skills/" -DEEPAGENTS_RECURSION_LIMIT = 500 +@dataclass(frozen=True) +class FetchResult: + final_url: str + status_code: int | None + headers: dict[str, str] + body: bytes + redirect_chain: list[dict[str, Any]] + error: dict[str, Any] | None = None -class LaunchCheckStudioV1(A2AAgent[LaunchCheckStudioV1Config, {{ auth_type }}]): +class _HeadParser(HTMLParser): + def __init__(self) -> None: + super().__init__() + self.in_title = False + self.title_parts: list[str] = [] + self.description: str | None = None + self.canonical: str | None = None + self.viewport: str | None = None + + def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None: + tag = tag.lower() + attr = {k.lower(): (v or "") for k, v in attrs} + if tag == "title": + self.in_title = True + elif tag == "meta": + name = attr.get("name", "").lower() + prop = attr.get("property", "").lower() + content = attr.get("content", "").strip() + if name == "description" and content and self.description is None: + self.description = content[:500] + if name == "viewport" and content and self.viewport is None: + self.viewport = content[:300] + if prop == "og:description" and content and self.description is None: + self.description = content[:500] + elif tag == "link" and attr.get("rel", "").lower() == "canonical": + href = attr.get("href", "").strip() + if href and self.canonical is None: + self.canonical = href[:1000] + + def handle_endtag(self, tag: str) -> None: + if tag.lower() == "title": + self.in_title = False + + def handle_data(self, data: str) -> None: + if self.in_title: + self.title_parts.append(data) + + @property + def title(self) -> str | None: + value = " ".join(" ".join(self.title_parts).split()).strip() + return value[:300] if value else None + + +class LaunchCheckStudioV1(A2AAgent[LaunchCheckStudioV1Config, PlatformUserAuth]): name = "launch-check-studio-v1" - description = "LaunchCheck is a one-page full-stack startup for SSRF-safe launch-readiness audits of public URLs with user-scoped persistence." + description = ( + "LaunchCheck audits public HTTP(S) launch pages with bounded SSRF-safe " + "network checks and user-scoped Postgres persistence." + ) version = "0.1.0" config_model = LaunchCheckStudioV1Config - auth_model = {{ auth_type }} + auth_model = PlatformUserAuth - # Hosted generated agents read the caller's saved LLM credential through - # ctx.llm. The platform may proxy that credential through LiteLLM, but agent - # code never reads provider keys, LiteLLM master keys, or OPENAI_API_KEY - # directly. - llm_provisioning = LLMProvisioning.PLATFORM + state = State.DURABLE + resources = Resources(cpu="500m", memory="512Mi", max_runtime_seconds=120) + egress = EgressPolicy(deny_internet_by_default=False) + tools_used = ("httpx", "psycopg", "postgres") pricing = Pricing( price_per_call_usd=0.0, - caller_pays_llm=True, - notes="Starter agent uses the caller's saved LLM credential via ctx.llm.", + caller_pays_llm=False, + notes="Deterministic bounded audit; no LLM credential required.", ) - workspace_access = WorkspaceAccess.dynamic( - max_files=64, - allowed_modes=(WorkspaceMode.READ_ONLY, WorkspaceMode.READ_WRITE_OVERLAY), - require_reason=False, - ) - tools_used = ("deepagents", "langchain") - - @a2a.tool(description="Ask the starter DeepAgent to answer with tool calls when useful") - async def ask(self, ctx: RunContext[{{ auth_type }}], prompt: str) -> str: - creds = ctx.llm - await ctx.emit_progress(f"llm: {creds.model} via {creds.source}") - if not creds.api_key: - return ( - "LLM key required. Add an LLM credential in Settings > LLM " - "credentials before running this agent; for local --invoke " - "runs set AGENT_LLM_KEY." - ) - graph = self._build_deep_agent(ctx=ctx, creds=creds) - state = await graph.ainvoke( - {"messages": [{"role": "user", "content": prompt}]}, - config={"recursion_limit": DEEPAGENTS_RECURSION_LIMIT}, + platform_resources = AgentPlatformResources( + databases=( + AgentDatabase( + name="launch-check-studio-v1-data", + scope="user", + access_mode="read_write", + env=AgentDatabaseEnv(url="DATABASE_URL"), + migrations=AgentDatabaseMigrations(path="db/migrations"), + ), ) - await ctx.emit_progress("deepagent finished") - return _last_message_text(state) + ) - def _build_deep_agent( + @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, + cost_class="network-read", + ) + async def audit_url( self, - *, - ctx: RunContext[{{ auth_type }}], - creds: LLMCreds, - ) -> Any: - # Lazy imports keep `a2a card` usable before local dependencies are - # installed. `a2a deploy` installs requirements.txt during the build. - from a2a_pack.deepagents import create_a2a_deep_agent - from langchain.agents.middleware import wrap_model_call - from langchain_core.tools import tool + ctx: RunContext[PlatformUserAuth], + audit_id: str, + url: str, + ) -> dict[str, Any]: + tenant = _tenant_key(ctx) + audit_id = _clean_audit_id(audit_id) + started_at = _now_iso() + await ctx.emit_progress("Validating public target and resolving DNS") - @tool - def text_stats(text: str) -> str: - """Return exact word, character, and line counts for text.""" - words = [part for part in text.split() if part.strip()] - return json.dumps( - { - "characters": len(text), - "words": len(words), - "lines": len(text.splitlines()) or 1, - } + validation = await _validate_public_url(url) + if not validation["ok"]: + failure = _failure_result(audit_id, url, validation["code"], validation["message"]) + await _persist_failure_if_possible(tenant, audit_id, url, failure, started_at) + return failure + + try: + fetch = await _fetch_bounded(validation["url"], self.config.max_redirects, self.config.max_response_bytes) + if fetch.error: + result = _failure_result(audit_id, validation["url"], fetch.error["code"], fetch.error["message"]) + result["target_url"] = validation["url"] + await _persist_failure_if_possible(tenant, audit_id, validation["url"], result, started_at) + return result + + await ctx.emit_progress("Analyzing metadata, robots, and security headers") + robots = await _check_robots(fetch.final_url) + report = _build_report(fetch, robots) + receipt = _build_app_receipt(tenant, audit_id, validation["url"], report, started_at) + payload = { + "audit_id": audit_id, + "target_url": validation["url"], + "ok": True, + "report": report, + "receipt": receipt, + "created_at": started_at, + "updated_at": _now_iso(), + } + await _save_audit_and_receipt(tenant, audit_id, payload, receipt) + return payload + except Exception: # noqa: BLE001 - structured public failure, no raw internals + result = _failure_result( + audit_id, + validation.get("url", url), + "audit_failed", + "The audit could not be completed within the bounded execution policy.", ) + await _persist_failure_if_possible(tenant, audit_id, validation.get("url", url), result, started_at) + return result - @wrap_model_call - async def log_model_call(request: Any, handler: Any) -> Any: - messages = request.state.get("messages", []) - print( - "[middleware] model_call " - f"model={creds.model} source={creds.source} messages={len(messages)}" - ) - return await handler(request) + @a2a.tool( + description="Reopen a previously persisted launch-readiness audit for the signed-in user.", + timeout_seconds=20, + idempotent=True, + cost_class="db-read", + ) + async def get_audit( + self, + ctx: RunContext[PlatformUserAuth], + audit_id: str, + ) -> dict[str, Any]: + tenant = _tenant_key(ctx) + audit_id = _clean_audit_id(audit_id) + row = await _load_audit(tenant, audit_id) + if row is None: + return { + "ok": False, + "code": "audit_not_found", + "message": "No audit with that id exists for this signed-in user.", + "audit_id": audit_id, + } + return row - backend = ctx.workspace_backend() - skill_sources = _seed_runtime_skills(backend, ctx) - # create_a2a_deep_agent resolves provider:model strings with - # langchain.init_chat_model from ctx.llm, preserving LiteLLM routing, - # provider-specific extra body, and runtime model overrides. - return create_a2a_deep_agent( - ctx, - creds=creds, - backend=backend, - skills=skill_sources or None, - tools=[text_stats], - middleware=[log_model_call], - system_prompt=SYSTEM_PROMPT, + +def _tenant_key(ctx: RunContext[PlatformUserAuth]) -> str: + stable_id = getattr(ctx.auth, "user_id", None) + if stable_id is not None: + return f"user:{stable_id}" + sub = str(getattr(ctx.auth, "sub", "") or "").strip() + if sub: + return f"user:{sub}" + raise PermissionError("stable platform identity required") + + +def _clean_audit_id(value: str) -> str: + cleaned = str(value or "").strip() + if not re.fullmatch(r"[A-Za-z0-9][A-Za-z0-9_.:-]{0,127}", cleaned): + raise ValueError("audit_id must be 1-128 safe characters") + return cleaned + + +def _now_iso() -> str: + return datetime.now(timezone.utc).replace(microsecond=0).isoformat() + + +def _failure_result(audit_id: str, url: str, code: str, message: str) -> dict[str, Any]: + return { + "ok": False, + "code": code, + "message": message, + "audit_id": audit_id, + "target_url": _safe_url_echo(url), + } + + +def _safe_url_echo(url: str) -> str: + try: + parsed = urlparse(url) + netloc = parsed.hostname or "" + if parsed.port: + netloc = f"{netloc}:{parsed.port}" + return urlunparse((parsed.scheme, netloc, parsed.path or "", "", parsed.query[:200], ""))[:1000] + except Exception: # noqa: BLE001 + return "" + + +async def _validate_public_url(raw_url: str) -> dict[str, Any]: + value = str(raw_url or "").strip() + if len(value) > 2048: + return {"ok": False, "code": "invalid_url", "message": "URL is too long."} + parsed = urlparse(value) + if parsed.scheme.lower() not in {"http", "https"}: + return {"ok": False, "code": "invalid_scheme", "message": "Only http and https URLs are supported."} + if parsed.username or parsed.password or "@" in parsed.netloc: + return {"ok": False, "code": "credentials_not_allowed", "message": "URLs with embedded credentials are not allowed."} + if not parsed.hostname: + return {"ok": False, "code": "invalid_url", "message": "URL must include a hostname."} + 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): + 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: + 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}" + normalized = urlunparse((parsed.scheme.lower(), netloc, parsed.path or "/", "", parsed.query, "")) + return {"ok": True, "url": normalized} + + +async def _host_is_disallowed(host: str, port: int) -> bool: + try: + ip = ipaddress.ip_address(host.strip("[]")) + return _ip_disallowed(ip) + except ValueError: + pass + lowered = host.lower() + if lowered in {"localhost", "metadata.google.internal"} or lowered.endswith(".localhost"): + return True + 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 + 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) -def _runtime_skills_root(ctx: RunContext[Any]) -> str: - workspace = getattr(ctx, "_workspace", None) - prefixes = tuple(getattr(workspace, "write_prefixes", ()) or ()) - if not prefixes: - outputs_prefix = getattr(workspace, "outputs_prefix", None) - prefixes = (outputs_prefix or "outputs/",) - prefix = str(prefixes[0]).strip("/") - return f"/{prefix}/{RUNTIME_SKILLS_DIR}" if prefix else f"/{RUNTIME_SKILLS_DIR}" +def _ip_disallowed(ip: ipaddress._BaseAddress) -> bool: + metadata = ipaddress.ip_address("169.254.169.254") + return bool( + ip.is_private + or ip.is_loopback + or ip.is_link_local + or ip.is_multicast + or ip.is_reserved + or ip.is_unspecified + or ip == metadata + ) -def _seed_runtime_skills(backend: Any, ctx: RunContext[Any]) -> list[str]: - """Copy packaged DeepAgents skills into the invocation workspace. +async def _fetch_bounded(url: str, max_redirects: int, max_bytes: int) -> FetchResult: + import httpx - DeepAgents loads skills from its backend, while source-controlled - ``skills/`` folders live in the image. This bridge lets generated agents - ship reusable SKILL.md bundles without giving up durable A2A workspace - files. - """ - root = Path(__file__).parent / "skills" - if not root.exists(): - return [] - runtime_skills_root = _runtime_skills_root(ctx) - uploads: list[tuple[str, bytes]] = [] - for path in root.rglob("*"): - if path.is_file(): - rel = path.relative_to(root).as_posix() - uploads.append((runtime_skills_root + rel, path.read_bytes())) - if uploads: - backend.upload_files(uploads) - return [runtime_skills_root] - return [] + 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) + return FetchResult(current, None, {}, b"", chain, {"code": "too_many_redirects", "message": "Redirect count exceeded the audit limit."}) -def _last_message_text(state: dict[str, Any]) -> str: - messages = state.get("messages") or [] - if not messages: - return json.dumps(state, default=str) +async def _read_limited(response: Any, max_bytes: int) -> bytes: + chunks: list[bytes] = [] + total = 0 + async for chunk in response.aiter_bytes(): + total += len(chunk) + if total > max_bytes: + remaining = max_bytes - (total - len(chunk)) + if remaining > 0: + chunks.append(chunk[:remaining]) + break + chunks.append(chunk) + return b"".join(chunks) - content = getattr(messages[-1], "content", None) - if isinstance(content, str): - return content - if isinstance(content, list): - parts: list[str] = [] - for item in content: - if isinstance(item, dict): - text = item.get("text") or item.get("content") - if text: - parts.append(str(text)) - elif item: - parts.append(str(item)) - return "\n".join(parts) if parts else json.dumps(content, default=str) - return str(content or messages[-1]) + +async def _check_robots(final_url: str) -> dict[str, Any]: + parsed = urlparse(final_url) + robots_url = urlunparse((parsed.scheme, parsed.netloc, "/robots.txt", "", "", "")) + 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) + return { + "present": response.status_code == 200, + "status_code": response.status_code, + "checked_url": validation["url"], + "bytes_read": len(body), + } + except Exception: # noqa: BLE001 + return {"present": False, "status_code": None, "checked_url": validation["url"]} + + +def _build_report(fetch: FetchResult, robots: dict[str, Any]) -> dict[str, Any]: + text = fetch.body.decode("utf-8", errors="replace") + parser = _HeadParser() + parser.feed(text[:200_000]) + headers = fetch.headers + security = { + "strict_transport_security": bool(headers.get("strict-transport-security")), + "content_security_policy": bool(headers.get("content-security-policy")), + "x_content_type_options": headers.get("x-content-type-options", "").lower() == "nosniff", + "referrer_policy": bool(headers.get("referrer-policy")), + "permissions_policy": bool(headers.get("permissions-policy")), + "x_frame_options": bool(headers.get("x-frame-options")), + } + evidence = { + "status_code": fetch.status_code, + "final_url": fetch.final_url, + "redirect_chain": fetch.redirect_chain, + "https": urlparse(fetch.final_url).scheme == "https", + "title": parser.title, + "description": parser.description, + "canonical": urljoin(fetch.final_url, parser.canonical) if parser.canonical else None, + "viewport": parser.viewport, + "robots": robots, + "security_headers": security, + "bytes_read": len(fetch.body), + } + checks = { + "reachable": fetch.status_code is not None and 200 <= int(fetch.status_code) < 400, + "https": evidence["https"], + "has_title": bool(parser.title), + "has_description": bool(parser.description), + "has_canonical": bool(parser.canonical), + "has_viewport": bool(parser.viewport), + "robots_present": bool(robots.get("present")), + "security_headers": security, + } + fixes = _recommended_fixes(checks) + return { + "reachable": checks["reachable"], + "final_url": fetch.final_url, + "checks": checks, + "evidence": evidence, + "fixes": fixes, + "summary": _summary(checks, fixes), + } + + +def _recommended_fixes(checks: dict[str, Any]) -> list[dict[str, str]]: + fixes: list[dict[str, str]] = [] + if not checks["https"]: + fixes.append({"category": "security", "priority": "high", "fix": "Serve the launch page over HTTPS and redirect HTTP to HTTPS."}) + if not checks["has_title"]: + fixes.append({"category": "metadata", "priority": "high", "fix": "Add a concise, unique