transport.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427
  1. """Bounded pull-only HTTPS transport for the enterprise edge agent."""
  2. from __future__ import annotations
  3. import hashlib
  4. import json
  5. import logging
  6. import re
  7. from collections.abc import Mapping
  8. from urllib.parse import quote, urlsplit
  9. from app.core.edge_gateway.contracts import (
  10. EdgeContractError,
  11. EdgeEventContract,
  12. canonical_sha256,
  13. preflight_json,
  14. strict_json_bytes,
  15. )
  16. from app.core.edge_gateway.policy import EdgeEgressPolicy, EdgePolicyError
  17. from app.edge_gateway.bootstrap import validate_server_crl
  18. LOGGER = logging.getLogger(__name__)
  19. _IDENTIFIER = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:-]{0,254}$")
  20. _SHA256 = re.compile(r"^[0-9a-f]{64}$")
  21. class EdgeTransportError(RuntimeError):
  22. """A bounded transport failure safe for logs and retry decisions."""
  23. class EdgeAuthenticationStopped(EdgeTransportError):
  24. """Credential/certificate binding is no longer usable; stop the agent."""
  25. class EdgeTransportConflict(EdgeTransportError):
  26. """The control plane rejected an idempotent replay with changed content."""
  27. class EdgeTransport:
  28. MAX_REQUEST_BYTES = 262_144
  29. MAX_RESPONSE_BYTES = 1_048_576
  30. def __init__(
  31. self,
  32. *,
  33. base_url: str,
  34. gateway_id: str,
  35. environment: str,
  36. network_zone: str,
  37. generation: int,
  38. credential: str,
  39. certificate_sha256: str,
  40. allowed_control_hosts: set[str] | frozenset[str],
  41. allowed_proxy_hosts: set[str] | frozenset[str] | None = None,
  42. allowed_control_origins: set[str] | frozenset[str] | None = None,
  43. allowed_proxy_origins: set[str] | frozenset[str] | None = None,
  44. proxy_url: str | None = None,
  45. client_certificate_path: str,
  46. client_private_key_path: str,
  47. ca_bundle_path: str,
  48. client,
  49. server_crl_path: str | None = None,
  50. connect_timeout_seconds: int = 5,
  51. read_timeout_seconds: int = 20,
  52. ) -> None:
  53. self.policy = EdgeEgressPolicy(
  54. allowed_control_hosts=set(allowed_control_hosts),
  55. allowed_proxy_hosts=set(allowed_proxy_hosts or set()),
  56. allowed_control_origins=allowed_control_origins,
  57. allowed_proxy_origins=allowed_proxy_origins,
  58. )
  59. self.policy.validate_destination(base_url, proxy_url=proxy_url)
  60. parsed = urlsplit(base_url)
  61. if parsed.query or parsed.fragment or parsed.path not in {"", "/"}:
  62. raise EdgePolicyError("control base URL must not contain a path or query")
  63. if not isinstance(gateway_id, str) or not _IDENTIFIER.fullmatch(gateway_id):
  64. raise ValueError("gateway_id is invalid")
  65. if not isinstance(certificate_sha256, str) or not _SHA256.fullmatch(certificate_sha256):
  66. raise ValueError("certificate_sha256 is invalid")
  67. if (
  68. not isinstance(credential, str)
  69. or not credential.startswith("dopg_")
  70. or len(credential) > 160
  71. ):
  72. raise ValueError("credential is invalid")
  73. if isinstance(generation, bool) or not isinstance(generation, int) or generation < 1:
  74. raise ValueError("generation is invalid")
  75. for label, timeout in {
  76. "connect_timeout_seconds": connect_timeout_seconds,
  77. "read_timeout_seconds": read_timeout_seconds,
  78. }.items():
  79. if isinstance(timeout, bool) or not isinstance(timeout, int) or not 1 <= timeout <= 120:
  80. raise ValueError(f"{label} is invalid")
  81. if client is None or not callable(getattr(client, "request", None)):
  82. raise ValueError("an explicit bounded HTTP client is required")
  83. if hasattr(client, "trust_env") and client.trust_env is not False:
  84. raise ValueError("HTTP client must disable implicit environment proxies")
  85. self.base_url = base_url.rstrip("/")
  86. self.proxy_url = proxy_url
  87. self.gateway_id = gateway_id
  88. self.environment = environment
  89. self.network_zone = network_zone
  90. self.generation = generation
  91. self._credential = credential
  92. self.certificate_sha256 = certificate_sha256
  93. self.client = client
  94. for label, path in {
  95. "client_certificate_path": client_certificate_path,
  96. "client_private_key_path": client_private_key_path,
  97. "ca_bundle_path": ca_bundle_path,
  98. }.items():
  99. if not isinstance(path, str) or not path.startswith("/") or "\x00" in path:
  100. raise ValueError(f"{label} is invalid")
  101. self.client_certificate_path = client_certificate_path
  102. self._client_private_key_path = client_private_key_path
  103. self.ca_bundle_path = ca_bundle_path
  104. if server_crl_path is not None and (
  105. not isinstance(server_crl_path, str)
  106. or not server_crl_path.startswith("/")
  107. or "\x00" in server_crl_path
  108. ):
  109. raise ValueError("server_crl_path is invalid")
  110. self.server_crl_path = server_crl_path
  111. if server_crl_path is not None:
  112. validate_server_crl(ca_bundle_path, server_crl_path)
  113. self.timeout = (connect_timeout_seconds, read_timeout_seconds)
  114. def _binding(self) -> dict[str, object]:
  115. return {
  116. "environment": self.environment,
  117. "network_zone": self.network_zone,
  118. "generation": self.generation,
  119. }
  120. def _request(self, path: str, payload: Mapping[str, object]) -> dict[str, object]:
  121. if self.server_crl_path is not None:
  122. try:
  123. validate_server_crl(self.ca_bundle_path, self.server_crl_path)
  124. except ValueError as exc:
  125. raise EdgeAuthenticationStopped("control-plane server CRL is not fresh") from exc
  126. if not path.startswith("/") or ".." in path or "//" in path:
  127. raise EdgeTransportError("control path is invalid")
  128. url = f"{self.base_url}{path}"
  129. self.policy.validate_destination(url, proxy_url=self.proxy_url)
  130. try:
  131. encoded = strict_json_bytes(payload)
  132. except EdgeContractError as exc:
  133. raise EdgeTransportError("control request is invalid") from exc
  134. if len(encoded) > self.MAX_REQUEST_BYTES:
  135. raise EdgeTransportError("control request exceeds the byte limit")
  136. headers = {
  137. "Accept": "application/json",
  138. "Content-Type": "application/json",
  139. "X-Edge-Credential": self._credential,
  140. "X-Edge-Certificate-SHA256": self.certificate_sha256,
  141. }
  142. try:
  143. response = self.client.request(
  144. "POST",
  145. url,
  146. data=encoded,
  147. headers=headers,
  148. timeout=self.timeout,
  149. allow_redirects=False,
  150. proxies=(
  151. {"https": self.proxy_url}
  152. if self.proxy_url is not None
  153. else {"http": None, "https": None}
  154. ),
  155. verify=self.ca_bundle_path,
  156. cert=(self.client_certificate_path, self._client_private_key_path),
  157. stream=True,
  158. )
  159. except Exception as exc:
  160. LOGGER.warning(
  161. "edge control request failed category=%s", type(exc).__name__
  162. )
  163. raise EdgeTransportError("control request failed") from exc
  164. try:
  165. return self._read_response(response)
  166. finally:
  167. close = getattr(response, "close", None)
  168. if callable(close):
  169. close()
  170. def _read_response(self, response) -> dict[str, object]:
  171. status = getattr(response, "status_code", None)
  172. if not isinstance(status, int):
  173. raise EdgeTransportError("control response is invalid")
  174. if 300 <= status < 400:
  175. raise EdgeTransportError("control redirects are forbidden")
  176. headers_value = getattr(response, "headers", {})
  177. content_length = headers_value.get("Content-Length") if isinstance(headers_value, Mapping) else None
  178. if content_length is not None:
  179. try:
  180. declared_length = int(content_length)
  181. except (TypeError, ValueError) as exc:
  182. raise EdgeTransportError("control response content length is invalid") from exc
  183. if declared_length < 0 or declared_length > self.MAX_RESPONSE_BYTES:
  184. raise EdgeTransportError("control response exceeds the byte limit")
  185. iterator = getattr(response, "iter_content", None)
  186. if not callable(iterator):
  187. raise EdgeTransportError("control response streaming is unavailable")
  188. chunks = bytearray()
  189. try:
  190. for chunk in iterator(chunk_size=65_536):
  191. if not isinstance(chunk, bytes):
  192. raise EdgeTransportError("control response chunk is invalid")
  193. chunks.extend(chunk)
  194. if len(chunks) > self.MAX_RESPONSE_BYTES:
  195. raise EdgeTransportError("control response exceeds the byte limit")
  196. except EdgeTransportError:
  197. raise
  198. except Exception as exc:
  199. raise EdgeTransportError("control response stream failed") from exc
  200. body = bytes(chunks)
  201. if len(body) > self.MAX_RESPONSE_BYTES:
  202. raise EdgeTransportError("control response exceeds the byte limit")
  203. if status in {401, 403}:
  204. raise EdgeAuthenticationStopped("edge credential or certificate binding was rejected")
  205. if status == 409:
  206. raise EdgeTransportConflict("control plane idempotency conflict")
  207. if not 200 <= status < 300:
  208. raise EdgeTransportError("control plane rejected the request")
  209. try:
  210. value = json.loads(body.decode("utf-8"))
  211. preflight_json(
  212. value,
  213. max_depth=32,
  214. max_nodes=5_000,
  215. max_string_bytes=131_072,
  216. max_total_bytes=self.MAX_RESPONSE_BYTES,
  217. )
  218. except (UnicodeDecodeError, json.JSONDecodeError, EdgeContractError) as exc:
  219. raise EdgeTransportError("control response is invalid") from exc
  220. if not isinstance(value, dict):
  221. raise EdgeTransportError("control response is invalid")
  222. if set(value) == {"code", "message", "data"}:
  223. if value["code"] != 200:
  224. raise EdgeTransportError("control response envelope is invalid")
  225. value = value["data"]
  226. if not isinstance(value, dict):
  227. raise EdgeTransportError("control response data is invalid")
  228. return value
  229. def reconcile(self) -> dict[str, object]:
  230. cancelled: list[str] = []
  231. releases: list[dict[str, object]] = []
  232. release_digests: dict[str, str] = {}
  233. cancel_cursor: str | None = None
  234. release_cursor: str | None = None
  235. cancel_done = False
  236. release_done = False
  237. for _page in range(100):
  238. value = self._request(
  239. f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/reconcile",
  240. {
  241. **self._binding(),
  242. "limit": 50,
  243. "cancel_cursor": cancel_cursor,
  244. "release_cursor": release_cursor,
  245. },
  246. )
  247. required = {
  248. "cancelled_task_ids", "cancel_next_cursor", "release_offers",
  249. "release_next_cursor", "release_baseline",
  250. }
  251. if set(value) != required:
  252. raise EdgeTransportError("reconcile response schema is invalid")
  253. page_cancelled = value["cancelled_task_ids"]
  254. page_releases = value["release_offers"]
  255. if (
  256. not isinstance(page_cancelled, list)
  257. or not isinstance(page_releases, list)
  258. or len(page_cancelled) > 50
  259. or len(page_releases) > 50
  260. ):
  261. raise EdgeTransportError("reconcile page exceeds safe limits")
  262. for task_id in page_cancelled:
  263. if not isinstance(task_id, str) or not _IDENTIFIER.fullmatch(task_id):
  264. raise EdgeTransportError("cancel reconciliation is invalid")
  265. if task_id not in cancelled:
  266. cancelled.append(task_id)
  267. baseline = value["release_baseline"]
  268. candidates = list(page_releases)
  269. if baseline is not None:
  270. candidates.append(baseline)
  271. for release in candidates:
  272. if not isinstance(release, Mapping):
  273. raise EdgeTransportError("release reconciliation is invalid")
  274. release_id = release.get("release_id")
  275. if not isinstance(release_id, str) or not _IDENTIFIER.fullmatch(release_id):
  276. raise EdgeTransportError("release reconciliation is invalid")
  277. digest = canonical_sha256(release)
  278. prior = release_digests.get(release_id)
  279. if prior is not None and prior != digest:
  280. raise EdgeTransportConflict("release page replay changed content")
  281. if prior is None:
  282. release_digests[release_id] = digest
  283. releases.append(dict(release))
  284. next_cancel = value["cancel_next_cursor"]
  285. next_release = value["release_next_cursor"]
  286. if next_cancel is not None and (
  287. not isinstance(next_cancel, str) or not _IDENTIFIER.fullmatch(next_cancel)
  288. ):
  289. raise EdgeTransportError("cancel cursor is invalid")
  290. if next_release is not None and (
  291. not isinstance(next_release, str) or not _IDENTIFIER.fullmatch(next_release)
  292. ):
  293. raise EdgeTransportError("release cursor is invalid")
  294. cancel_done = next_cancel is None
  295. release_done = next_release is None
  296. cancel_cursor = next_cancel or (
  297. page_cancelled[-1] if page_cancelled else cancel_cursor
  298. )
  299. release_cursor = next_release or (
  300. str(page_releases[-1].get("release_id"))
  301. if page_releases and isinstance(page_releases[-1], Mapping)
  302. else release_cursor
  303. )
  304. if len(cancelled) > 5_000 or len(releases) > 5_000:
  305. raise EdgeTransportError("reconcile result exceeds safe limits")
  306. if cancel_done and release_done:
  307. return {
  308. "cancelled_task_ids": cancelled,
  309. "release_offers": releases,
  310. }
  311. raise EdgeTransportError("reconcile pagination did not converge")
  312. def cancel_requested(self, task_id: str) -> bool:
  313. if not isinstance(task_id, str) or not _IDENTIFIER.fullmatch(task_id):
  314. raise EdgeTransportError("cancel task identity is invalid")
  315. value = self.reconcile()
  316. cancelled = value.get("cancelled_task_ids")
  317. if not isinstance(cancelled, list) or len(cancelled) > 100:
  318. raise EdgeTransportError("cancel reconciliation is invalid")
  319. if any(not isinstance(item, str) or not _IDENTIFIER.fullmatch(item) for item in cancelled):
  320. raise EdgeTransportError("cancel reconciliation is invalid")
  321. return task_id in cancelled
  322. def pull_task(self):
  323. value = self._request(
  324. f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/tasks/pull",
  325. self._binding(),
  326. )
  327. envelope = value.get("signed_task_envelope")
  328. task = value.get("task")
  329. if envelope is None and task is None:
  330. return None
  331. lease_token = value.get("lease_token")
  332. lease_expires_at = value.get("lease_expires_at")
  333. if (
  334. not isinstance(envelope, Mapping)
  335. or not isinstance(task, Mapping)
  336. or envelope.get("task") != task
  337. or not isinstance(lease_token, str)
  338. or not isinstance(lease_expires_at, str)
  339. ):
  340. raise EdgeTransportError("pulled task response is invalid")
  341. return dict(envelope), lease_token, lease_expires_at
  342. def send_event(self, event: Mapping[str, object], lease_token: str) -> dict[str, object]:
  343. try:
  344. normalized_event = EdgeEventContract.from_mapping(event).to_mapping()
  345. except EdgeContractError as exc:
  346. raise EdgeTransportError("outbound event contract is invalid") from exc
  347. value = self._request(
  348. f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/events",
  349. {
  350. **self._binding(),
  351. "event": normalized_event,
  352. "lease_token": lease_token,
  353. },
  354. )
  355. expected_digest = hashlib.sha256(lease_token.encode("utf-8")).hexdigest()
  356. if (
  357. value.get("status") != "accepted"
  358. or value.get("event_id") != normalized_event.get("event_id")
  359. or value.get("event_digest") != canonical_sha256(normalized_event)
  360. or value.get("remote_lease_digest") != expected_digest
  361. or not isinstance(value.get("received_at"), str)
  362. ):
  363. raise EdgeTransportError("event acknowledgement is invalid")
  364. return {
  365. "event_id": value["event_id"],
  366. "event_digest": value["event_digest"],
  367. "remote_lease_digest": value["remote_lease_digest"],
  368. "received_at": value["received_at"],
  369. "status": "accepted",
  370. }
  371. def task_outcome(self, task_id: str, outcome: str, lease_token: str, summary: Mapping[str, object]):
  372. value = self._request(
  373. f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/tasks/{quote(task_id, safe='')}/outcome",
  374. {
  375. **self._binding(),
  376. "outcome": outcome,
  377. "lease_token": lease_token,
  378. "safe_summary": dict(summary),
  379. },
  380. )
  381. if (
  382. set(value) != {"task_id", "status", "replayed"}
  383. or value.get("task_id") != task_id
  384. or value.get("status") != outcome
  385. or not isinstance(value.get("replayed"), bool)
  386. ):
  387. raise EdgeTransportError("task outcome acknowledgement is invalid")
  388. return value
  389. def heartbeat(self, summary: Mapping[str, object]):
  390. return self._request(
  391. f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/heartbeat",
  392. {**self._binding(), "version": summary.get("version"), "safe_summary": dict(summary)},
  393. )
  394. def acknowledge_release(self, release_id: str, outcome: str, summary: Mapping[str, object]):
  395. return self._request(
  396. f"/api/datasource/edge/gateways/{quote(self.gateway_id, safe='')}/releases/ack",
  397. {
  398. **self._binding(),
  399. "release_id": release_id,
  400. "outcome": outcome,
  401. "safe_summary": dict(summary),
  402. },
  403. )