"""Bounded pull-only HTTPS transport for the enterprise edge agent.""" from __future__ import annotations import hashlib import json import logging import re from collections.abc import Mapping from urllib.parse import quote, urlsplit from app.core.edge_gateway.contracts import ( EdgeContractError, EdgeEventContract, canonical_sha256, preflight_json, strict_json_bytes, ) from app.core.edge_gateway.policy import EdgeEgressPolicy, EdgePolicyError from app.edge_gateway.bootstrap import validate_server_crl LOGGER = logging.getLogger(__name__) _IDENTIFIER = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:-]{0,254}$") _SHA256 = re.compile(r"^[0-9a-f]{64}$") class EdgeTransportError(RuntimeError): """A bounded transport failure safe for logs and retry decisions.""" class EdgeAuthenticationStopped(EdgeTransportError): """Credential/certificate binding is no longer usable; stop the agent.""" class EdgeTransportConflict(EdgeTransportError): """The control plane rejected an idempotent replay with changed content.""" class EdgeTransport: MAX_REQUEST_BYTES = 262_144 MAX_RESPONSE_BYTES = 1_048_576 def __init__( self, *, base_url: str, gateway_id: str, environment: str, network_zone: str, generation: int, credential: str, certificate_sha256: str, allowed_control_hosts: set[str] | frozenset[str], allowed_proxy_hosts: set[str] | frozenset[str] | None = None, allowed_control_origins: set[str] | frozenset[str] | None = None, allowed_proxy_origins: set[str] | frozenset[str] | None = None, proxy_url: str | None = None, client_certificate_path: str, client_private_key_path: str, ca_bundle_path: str, client, server_crl_path: str | None = None, connect_timeout_seconds: int = 5, read_timeout_seconds: int = 20, ) -> None: self.policy = EdgeEgressPolicy( allowed_control_hosts=set(allowed_control_hosts), allowed_proxy_hosts=set(allowed_proxy_hosts or set()), allowed_control_origins=allowed_control_origins, allowed_proxy_origins=allowed_proxy_origins, ) self.policy.validate_destination(base_url, proxy_url=proxy_url) parsed = urlsplit(base_url) if parsed.query or parsed.fragment or parsed.path not in {"", "/"}: raise EdgePolicyError("control base URL must not contain a path or query") if not isinstance(gateway_id, str) or not _IDENTIFIER.fullmatch(gateway_id): raise ValueError("gateway_id is invalid") if not isinstance(certificate_sha256, str) or not _SHA256.fullmatch(certificate_sha256): raise ValueError("certificate_sha256 is invalid") if ( not isinstance(credential, str) or not credential.startswith("dopg_") or len(credential) > 160 ): raise ValueError("credential is invalid") if isinstance(generation, bool) or not isinstance(generation, int) or generation < 1: raise ValueError("generation is invalid") for label, timeout in { "connect_timeout_seconds": connect_timeout_seconds, "read_timeout_seconds": read_timeout_seconds, }.items(): if isinstance(timeout, bool) or not isinstance(timeout, int) or not 1 <= timeout <= 120: raise ValueError(f"{label} is invalid") if client is None or not callable(getattr(client, "request", None)): raise ValueError("an explicit bounded HTTP client is required") if hasattr(client, "trust_env") and client.trust_env is not False: raise ValueError("HTTP client must disable implicit environment proxies") self.base_url = base_url.rstrip("/") self.proxy_url = proxy_url self.gateway_id = gateway_id self.environment = environment self.network_zone = network_zone self.generation = generation self._credential = credential self.certificate_sha256 = certificate_sha256 self.client = client for label, path in { "client_certificate_path": client_certificate_path, "client_private_key_path": client_private_key_path, "ca_bundle_path": ca_bundle_path, }.items(): if not isinstance(path, str) or not path.startswith("/") or "\x00" in path: raise ValueError(f"{label} is invalid") self.client_certificate_path = client_certificate_path self._client_private_key_path = client_private_key_path self.ca_bundle_path = ca_bundle_path if server_crl_path is not None and ( not isinstance(server_crl_path, str) or not server_crl_path.startswith("/") or "\x00" in server_crl_path ): raise ValueError("server_crl_path is invalid") self.server_crl_path = server_crl_path if server_crl_path is not None: validate_server_crl(ca_bundle_path, server_crl_path) self.timeout = (connect_timeout_seconds, read_timeout_seconds) def _binding(self) -> dict[str, object]: return { "environment": self.environment, "network_zone": self.network_zone, "generation": self.generation, } def _request(self, path: str, payload: Mapping[str, object]) -> dict[str, object]: if self.server_crl_path is not None: try: validate_server_crl(self.ca_bundle_path, self.server_crl_path) except ValueError as exc: raise EdgeAuthenticationStopped("control-plane server CRL is not fresh") from exc if not path.startswith("/") or ".." in path or "//" in path: raise EdgeTransportError("control path is invalid") url = f"{self.base_url}{path}" self.policy.validate_destination(url, proxy_url=self.proxy_url) try: encoded = strict_json_bytes(payload) except EdgeContractError as exc: raise EdgeTransportError("control request is invalid") from exc if len(encoded) > self.MAX_REQUEST_BYTES: raise EdgeTransportError("control request exceeds the byte limit") headers = { "Accept": "application/json", "Content-Type": "application/json", "X-Edge-Credential": self._credential, "X-Edge-Certificate-SHA256": self.certificate_sha256, } try: response = self.client.request( "POST", url, data=encoded, headers=headers, timeout=self.timeout, allow_redirects=False, proxies=( {"https": self.proxy_url} if self.proxy_url is not None else {"http": None, "https": None} ), verify=self.ca_bundle_path, cert=(self.client_certificate_path, self._client_private_key_path), stream=True, ) except Exception as exc: LOGGER.warning( "edge control request failed category=%s", type(exc).__name__ ) raise EdgeTransportError("control request failed") from exc try: return self._read_response(response) finally: close = getattr(response, "close", None) if callable(close): close() def _read_response(self, response) -> dict[str, object]: status = getattr(response, "status_code", None) if not isinstance(status, int): raise EdgeTransportError("control response is invalid") if 300 <= status < 400: raise EdgeTransportError("control redirects are forbidden") headers_value = getattr(response, "headers", {}) content_length = headers_value.get("Content-Length") if isinstance(headers_value, Mapping) else None if content_length is not None: try: declared_length = int(content_length) except (TypeError, ValueError) as exc: raise EdgeTransportError("control response content length is invalid") from exc if declared_length < 0 or declared_length > self.MAX_RESPONSE_BYTES: raise EdgeTransportError("control response exceeds the byte limit") iterator = getattr(response, "iter_content", None) if not callable(iterator): raise EdgeTransportError("control response streaming is unavailable") chunks = bytearray() try: for chunk in iterator(chunk_size=65_536): if not isinstance(chunk, bytes): raise EdgeTransportError("control response chunk is invalid") chunks.extend(chunk) if len(chunks) > self.MAX_RESPONSE_BYTES: raise EdgeTransportError("control response exceeds the byte limit") except EdgeTransportError: raise except Exception as exc: raise EdgeTransportError("control response stream failed") from exc body = bytes(chunks) if len(body) > self.MAX_RESPONSE_BYTES: raise EdgeTransportError("control response exceeds the byte limit") if status in {401, 403}: raise EdgeAuthenticationStopped("edge credential or certificate binding was rejected") if status == 409: raise EdgeTransportConflict("control plane idempotency conflict") if not 200 <= status < 300: raise EdgeTransportError("control plane rejected the request") try: value = json.loads(body.decode("utf-8")) preflight_json( value, max_depth=32, max_nodes=5_000, max_string_bytes=131_072, max_total_bytes=self.MAX_RESPONSE_BYTES, ) except (UnicodeDecodeError, json.JSONDecodeError, EdgeContractError) as exc: raise EdgeTransportError("control response is invalid") from exc if not isinstance(value, dict): raise EdgeTransportError("control response is invalid") if set(value) == {"code", "message", "data"}: if value["code"] != 200: raise EdgeTransportError("control response envelope is invalid") value = value["data"] if not isinstance(value, dict): raise EdgeTransportError("control response data is invalid") return value def reconcile(self) -> dict[str, object]: cancelled: list[str] = [] releases: list[dict[str, object]] = [] release_digests: dict[str, str] = {} cancel_cursor: str | None = None release_cursor: str | None = None cancel_done = False release_done = False for _page in range(100): value = self._request( f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/reconcile", { **self._binding(), "limit": 50, "cancel_cursor": cancel_cursor, "release_cursor": release_cursor, }, ) required = { "cancelled_task_ids", "cancel_next_cursor", "release_offers", "release_next_cursor", "release_baseline", } if set(value) != required: raise EdgeTransportError("reconcile response schema is invalid") page_cancelled = value["cancelled_task_ids"] page_releases = value["release_offers"] if ( not isinstance(page_cancelled, list) or not isinstance(page_releases, list) or len(page_cancelled) > 50 or len(page_releases) > 50 ): raise EdgeTransportError("reconcile page exceeds safe limits") for task_id in page_cancelled: if not isinstance(task_id, str) or not _IDENTIFIER.fullmatch(task_id): raise EdgeTransportError("cancel reconciliation is invalid") if task_id not in cancelled: cancelled.append(task_id) baseline = value["release_baseline"] candidates = list(page_releases) if baseline is not None: candidates.append(baseline) for release in candidates: if not isinstance(release, Mapping): raise EdgeTransportError("release reconciliation is invalid") release_id = release.get("release_id") if not isinstance(release_id, str) or not _IDENTIFIER.fullmatch(release_id): raise EdgeTransportError("release reconciliation is invalid") digest = canonical_sha256(release) prior = release_digests.get(release_id) if prior is not None and prior != digest: raise EdgeTransportConflict("release page replay changed content") if prior is None: release_digests[release_id] = digest releases.append(dict(release)) next_cancel = value["cancel_next_cursor"] next_release = value["release_next_cursor"] if next_cancel is not None and ( not isinstance(next_cancel, str) or not _IDENTIFIER.fullmatch(next_cancel) ): raise EdgeTransportError("cancel cursor is invalid") if next_release is not None and ( not isinstance(next_release, str) or not _IDENTIFIER.fullmatch(next_release) ): raise EdgeTransportError("release cursor is invalid") cancel_done = next_cancel is None release_done = next_release is None cancel_cursor = next_cancel or ( page_cancelled[-1] if page_cancelled else cancel_cursor ) release_cursor = next_release or ( str(page_releases[-1].get("release_id")) if page_releases and isinstance(page_releases[-1], Mapping) else release_cursor ) if len(cancelled) > 5_000 or len(releases) > 5_000: raise EdgeTransportError("reconcile result exceeds safe limits") if cancel_done and release_done: return { "cancelled_task_ids": cancelled, "release_offers": releases, } raise EdgeTransportError("reconcile pagination did not converge") def cancel_requested(self, task_id: str) -> bool: if not isinstance(task_id, str) or not _IDENTIFIER.fullmatch(task_id): raise EdgeTransportError("cancel task identity is invalid") value = self.reconcile() cancelled = value.get("cancelled_task_ids") if not isinstance(cancelled, list) or len(cancelled) > 100: raise EdgeTransportError("cancel reconciliation is invalid") if any(not isinstance(item, str) or not _IDENTIFIER.fullmatch(item) for item in cancelled): raise EdgeTransportError("cancel reconciliation is invalid") return task_id in cancelled def pull_task(self): value = self._request( f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/tasks/pull", self._binding(), ) envelope = value.get("signed_task_envelope") task = value.get("task") if envelope is None and task is None: return None lease_token = value.get("lease_token") lease_expires_at = value.get("lease_expires_at") if ( not isinstance(envelope, Mapping) or not isinstance(task, Mapping) or envelope.get("task") != task or not isinstance(lease_token, str) or not isinstance(lease_expires_at, str) ): raise EdgeTransportError("pulled task response is invalid") return dict(envelope), lease_token, lease_expires_at def send_event(self, event: Mapping[str, object], lease_token: str) -> dict[str, object]: try: normalized_event = EdgeEventContract.from_mapping(event).to_mapping() except EdgeContractError as exc: raise EdgeTransportError("outbound event contract is invalid") from exc value = self._request( f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/events", { **self._binding(), "event": normalized_event, "lease_token": lease_token, }, ) expected_digest = hashlib.sha256(lease_token.encode("utf-8")).hexdigest() if ( value.get("status") != "accepted" or value.get("event_id") != normalized_event.get("event_id") or value.get("event_digest") != canonical_sha256(normalized_event) or value.get("remote_lease_digest") != expected_digest or not isinstance(value.get("received_at"), str) ): raise EdgeTransportError("event acknowledgement is invalid") return { "event_id": value["event_id"], "event_digest": value["event_digest"], "remote_lease_digest": value["remote_lease_digest"], "received_at": value["received_at"], "status": "accepted", } def task_outcome(self, task_id: str, outcome: str, lease_token: str, summary: Mapping[str, object]): value = self._request( f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/tasks/{quote(task_id, safe='')}/outcome", { **self._binding(), "outcome": outcome, "lease_token": lease_token, "safe_summary": dict(summary), }, ) if ( set(value) != {"task_id", "status", "replayed"} or value.get("task_id") != task_id or value.get("status") != outcome or not isinstance(value.get("replayed"), bool) ): raise EdgeTransportError("task outcome acknowledgement is invalid") return value def heartbeat(self, summary: Mapping[str, object]): return self._request( f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/heartbeat", {**self._binding(), "version": summary.get("version"), "safe_summary": dict(summary)}, ) def acknowledge_release(self, release_id: str, outcome: str, summary: Mapping[str, object]): return self._request( f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/releases/ack", { **self._binding(), "release_id": release_id, "outcome": outcome, "safe_summary": dict(summary), }, )