import asyncio import logging import socket import time from collections.abc import Callable, Sequence from dataclasses import dataclass from typing import Protocol from urllib.parse import urljoin import aiohttp from aiohttp.abc import AbstractResolver from app.config import Settings from app.models import CurrentStatus, HealthState, utc_now from app.security import Resolver, TargetPolicyError, ValidatedTarget, redact_url, system_resolver, validate_target logger = logging.getLogger("monitor.checker") @dataclass(frozen=True) class HopResponse: status: int location: str | None class Fetcher(Protocol): async def fetch(self, target: ValidatedTarget, settings: Settings, remaining: float) -> HopResponse: ... class PinnedResolver(AbstractResolver): def __init__(self, hostname: str, addresses: Sequence[str]) -> None: self.hostname = hostname self.addresses = addresses async def resolve(self, host: str, port: int = 0, family: int = socket.AF_UNSPEC) -> list[dict[str, object]]: if host.rstrip(".").lower() != self.hostname: raise OSError("unpinned hostname requested") result: list[dict[str, object]] = [] for address in self.addresses: address_family = socket.AF_INET6 if ":" in address else socket.AF_INET if family in (socket.AF_UNSPEC, address_family): result.append({"hostname": host, "host": address, "port": port, "family": address_family, "proto": 0, "flags": 0}) return result async def close(self) -> None: return None class AioHttpFetcher: async def fetch(self, target: ValidatedTarget, settings: Settings, remaining: float) -> HopResponse: resolver = PinnedResolver(target.hostname, target.addresses) connector = aiohttp.TCPConnector(resolver=resolver, use_dns_cache=True, force_close=True) timeout = aiohttp.ClientTimeout( total=remaining, connect=min(settings.connect_timeout_seconds, remaining), sock_read=min(settings.read_timeout_seconds, remaining), ) headers = {"User-Agent": settings.user_agent, "Accept": "*/*"} async with aiohttp.ClientSession(connector=connector, timeout=timeout, trust_env=False) as session: async with session.get(target.url, headers=headers, allow_redirects=False) as response: await response.content.read(1) return HopResponse(response.status, response.headers.get("Location")) _ERROR_MESSAGES = { "blocked_target": "target rejected by network policy", "timeout": "endpoint check timed out", "dns_or_connect": "endpoint connection failed", "tls": "endpoint TLS validation failed", "protocol": "endpoint returned an invalid HTTP response", "redirect_limit": "endpoint exceeded the redirect limit", } class EndpointChecker: def __init__( self, settings: Settings, fetcher: Fetcher | None = None, resolver: Resolver = system_resolver, clock: Callable[[], float] = time.monotonic, ) -> None: self.settings = settings self.fetcher = fetcher or AioHttpFetcher() self.resolver = resolver self.clock = clock async def check(self, url: str, monitor_id: str = "unknown") -> CurrentStatus: started = self.clock() safe_url = redact_url(url) error_code: str | None = None http_status: int | None = None try: async with asyncio.timeout(self.settings.total_timeout_seconds): current = url for redirects in range(self.settings.max_redirects + 1): target = await validate_target(current, self.resolver) remaining = max(0.001, self.settings.total_timeout_seconds - (self.clock() - started)) response = await self.fetcher.fetch(target, self.settings, remaining) http_status = response.status if response.status not in {301, 302, 303, 307, 308} or not response.location: break if redirects == self.settings.max_redirects: error_code = "redirect_limit" break current = urljoin(current, response.location) state = HealthState.up if error_code is None and http_status is not None and 200 <= http_status < 400 else HealthState.down except TargetPolicyError: error_code, state = "blocked_target", HealthState.down except TimeoutError: error_code, state = "timeout", HealthState.down except aiohttp.ClientConnectorCertificateError: error_code, state = "tls", HealthState.down except (aiohttp.ClientConnectorError, socket.gaierror, OSError): error_code, state = "dns_or_connect", HealthState.down except aiohttp.ClientError: error_code, state = "protocol", HealthState.down latency = max(0.0, (self.clock() - started) * 1000) status = CurrentStatus( state=state, checked_at=utc_now(), latency_ms=round(latency, 3), http_status=http_status, error_code=error_code, error_message=_ERROR_MESSAGES.get(error_code) if error_code else None, ) logger.info( "endpoint check completed", extra={"event": "check.completed", "monitor_id": monitor_id, "url": safe_url, "outcome": state, "error_code": error_code, "latency_ms": status.latency_ms}, ) return status