| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427 |
- """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),
- },
- )
|