artifacts.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353
  1. """Digest-bound, bounded Parquet artifacts owned by the DataOps Runner."""
  2. from __future__ import annotations
  3. import hashlib
  4. import io
  5. import json
  6. import os
  7. import re
  8. import tempfile
  9. from collections.abc import Mapping
  10. from contextlib import suppress
  11. from datetime import UTC, datetime, timedelta
  12. from typing import Any
  13. import polars as pl
  14. from sqlalchemy import text
  15. from app.core.common.identifiers import (
  16. ensure_governance_uid,
  17. new_governance_uid,
  18. )
  19. PARQUET_CONTENT_TYPE = "application/x-parquet"
  20. _DIGEST = re.compile(r"^[0-9a-f]{64}$")
  21. def _now_utc(clock) -> datetime:
  22. value = clock()
  23. if not isinstance(value, datetime):
  24. raise ValueError("artifact clock must return a datetime")
  25. if value.tzinfo is None:
  26. value = value.replace(tzinfo=UTC)
  27. return value.astimezone(UTC)
  28. def _timestamp(value: datetime) -> str:
  29. return value.astimezone(UTC).isoformat().replace("+00:00", "Z")
  30. def _parse_timestamp(value: Any) -> datetime:
  31. if not isinstance(value, str) or not value.endswith("Z"):
  32. raise ValueError("artifact expiry metadata is invalid")
  33. try:
  34. parsed = datetime.fromisoformat(value[:-1] + "+00:00")
  35. except ValueError as exc:
  36. raise ValueError("artifact expiry metadata is invalid") from exc
  37. return parsed.astimezone(UTC)
  38. def _uid(value: Any, label: str) -> str:
  39. try:
  40. return ensure_governance_uid({"uid": str(value)})
  41. except ValueError as exc:
  42. raise ValueError(f"{label} must be a valid UUIDv7") from exc
  43. def _schema_hash(frame: pl.DataFrame | pl.LazyFrame) -> str:
  44. schema = (
  45. frame.collect_schema()
  46. if isinstance(frame, pl.LazyFrame)
  47. else frame.schema
  48. )
  49. canonical = [
  50. {"name": name, "dtype": str(dtype)}
  51. for name, dtype in schema.items()
  52. ]
  53. encoded = json.dumps(
  54. canonical,
  55. sort_keys=True,
  56. separators=(",", ":"),
  57. ensure_ascii=False,
  58. )
  59. return hashlib.sha256(encoded.encode("utf-8")).hexdigest()
  60. def _metadata(value: Any) -> dict[str, str]:
  61. if not isinstance(value, Mapping):
  62. raise ValueError("artifact content metadata is missing")
  63. normalized = {}
  64. for key, item in value.items():
  65. name = str(key).lower()
  66. if name.startswith("x-amz-meta-"):
  67. name = name[len("x-amz-meta-") :]
  68. if name in {
  69. "sha256",
  70. "row-count",
  71. "schema-sha256",
  72. "expires-at",
  73. "artifact-bytes",
  74. }:
  75. normalized[name] = str(item)
  76. required = {
  77. "sha256",
  78. "row-count",
  79. "schema-sha256",
  80. "expires-at",
  81. "artifact-bytes",
  82. }
  83. if set(normalized) != required:
  84. raise ValueError("artifact content metadata is incomplete")
  85. return normalized
  86. class ArtifactStore:
  87. """Read and write only server-owned, bounded Parquet artifacts."""
  88. def __init__(
  89. self,
  90. client,
  91. *,
  92. bucket: str,
  93. max_artifact_bytes: int,
  94. max_rows: int,
  95. memory_limit_bytes: int,
  96. max_ttl_seconds: int = 86400,
  97. clock=None,
  98. ):
  99. if not re.fullmatch(r"[a-z0-9][a-z0-9.-]{1,61}[a-z0-9]", bucket):
  100. raise ValueError("artifact bucket name is invalid")
  101. self.client = client
  102. self.bucket = bucket
  103. self.max_artifact_bytes = int(max_artifact_bytes)
  104. self.max_rows = int(max_rows)
  105. self.memory_limit_bytes = int(memory_limit_bytes)
  106. self.max_ttl_seconds = int(max_ttl_seconds)
  107. self.clock = clock or (lambda: datetime.now(UTC))
  108. if (
  109. self.max_artifact_bytes < 1024
  110. or self.max_rows < 1
  111. or self.memory_limit_bytes < self.max_artifact_bytes
  112. or self.max_ttl_seconds < 1
  113. ):
  114. raise ValueError("artifact resource limits are invalid")
  115. if not self.client.bucket_exists(self.bucket):
  116. raise ValueError("artifact bucket does not exist")
  117. def _parse_ref(self, ref: Any) -> str:
  118. prefix = f"minio://{self.bucket}/"
  119. if not isinstance(ref, str) or not ref.startswith(prefix):
  120. raise ValueError("artifact reference is not owned by this store")
  121. key = ref[len(prefix) :]
  122. match = re.fullmatch(
  123. r"rules/([0-9a-f-]{36})/([0-9a-f-]{36})\.parquet",
  124. key,
  125. )
  126. if match is None:
  127. raise ValueError("artifact reference is invalid")
  128. _uid(match.group(1), "artifact correlation id")
  129. _uid(match.group(2), "artifact id")
  130. return key
  131. def _validated_stat(
  132. self,
  133. key: str,
  134. *,
  135. expected_digest: str | None = None,
  136. ) -> tuple[Any, dict[str, str]]:
  137. stat = self.client.stat_object(self.bucket, key)
  138. size = int(getattr(stat, "size", -1))
  139. if size < 1 or size > self.max_artifact_bytes:
  140. raise ValueError("artifact size exceeds the configured limit")
  141. if str(getattr(stat, "content_type", "")).lower() != PARQUET_CONTENT_TYPE:
  142. raise ValueError("artifact content type is invalid")
  143. metadata = _metadata(getattr(stat, "metadata", None))
  144. digest = metadata["sha256"]
  145. if _DIGEST.fullmatch(digest) is None:
  146. raise ValueError("artifact digest metadata is invalid")
  147. if expected_digest is not None and digest != expected_digest:
  148. raise ValueError("artifact digest does not match")
  149. try:
  150. row_count = int(metadata["row-count"])
  151. metadata_size = int(metadata["artifact-bytes"])
  152. except (TypeError, ValueError) as exc:
  153. raise ValueError("artifact count metadata is invalid") from exc
  154. if row_count < 0 or row_count > self.max_rows:
  155. raise ValueError("artifact row count exceeds the configured limit")
  156. if metadata_size != size:
  157. raise ValueError("artifact size metadata does not match")
  158. if _DIGEST.fullmatch(metadata["schema-sha256"]) is None:
  159. raise ValueError("artifact schema metadata is invalid")
  160. if _parse_timestamp(metadata["expires-at"]) <= _now_utc(self.clock):
  161. raise ValueError("artifact has expired")
  162. return stat, metadata
  163. def write(
  164. self,
  165. frame: pl.LazyFrame | pl.DataFrame,
  166. correlation_id: str,
  167. ttl_seconds: int,
  168. ) -> dict[str, Any]:
  169. correlation = _uid(correlation_id, "correlation_id")
  170. if (
  171. isinstance(ttl_seconds, bool)
  172. or not isinstance(ttl_seconds, int)
  173. or ttl_seconds < 1
  174. or ttl_seconds > self.max_ttl_seconds
  175. ):
  176. raise ValueError("artifact TTL is outside the configured limit")
  177. if isinstance(frame, pl.DataFrame):
  178. lazy = frame.lazy()
  179. elif isinstance(frame, pl.LazyFrame):
  180. lazy = frame
  181. else:
  182. raise ValueError("artifact frame must be a Polars frame")
  183. collected = lazy.head(self.max_rows + 1).collect(engine="streaming")
  184. if collected.height > self.max_rows:
  185. raise ValueError("artifact row count exceeds the configured limit")
  186. if collected.estimated_size() > self.memory_limit_bytes:
  187. raise ValueError("artifact frame exceeds the configured memory limit")
  188. schema_digest = _schema_hash(collected)
  189. expires_at = _timestamp(
  190. _now_utc(self.clock) + timedelta(seconds=ttl_seconds)
  191. )
  192. artifact_id = new_governance_uid()
  193. key = f"rules/{correlation}/{artifact_id}.parquet"
  194. path = None
  195. try:
  196. with tempfile.NamedTemporaryFile(
  197. prefix="dataops-rule-artifact-",
  198. suffix=".parquet",
  199. delete=False,
  200. ) as handle:
  201. path = handle.name
  202. collected.write_parquet(path)
  203. size = os.path.getsize(path)
  204. if size < 1 or size > self.max_artifact_bytes:
  205. raise ValueError("artifact size exceeds the configured limit")
  206. digest = hashlib.sha256()
  207. with open(path, "rb") as handle:
  208. while chunk := handle.read(1024 * 1024):
  209. digest.update(chunk)
  210. digest_hex = digest.hexdigest()
  211. with open(path, "rb") as handle:
  212. self.client.put_object(
  213. self.bucket,
  214. key,
  215. handle,
  216. size,
  217. content_type=PARQUET_CONTENT_TYPE,
  218. metadata={
  219. "sha256": digest_hex,
  220. "row-count": str(collected.height),
  221. "schema-sha256": schema_digest,
  222. "expires-at": expires_at,
  223. "artifact-bytes": str(size),
  224. },
  225. )
  226. self._validated_stat(key, expected_digest=digest_hex)
  227. artifact_ref = f"minio://{self.bucket}/{key}"
  228. self.read(artifact_ref, digest_hex)
  229. finally:
  230. if path is not None:
  231. with suppress(FileNotFoundError):
  232. os.unlink(path)
  233. return {
  234. "artifact_ref": artifact_ref,
  235. "digest": digest_hex,
  236. "row_count": collected.height,
  237. "schema_hash": schema_digest,
  238. "expires_at": expires_at,
  239. }
  240. def describe(self, ref: str) -> dict[str, Any]:
  241. """Return validated object metadata without exposing MinIO credentials."""
  242. key = self._parse_ref(ref)
  243. _stat, metadata = self._validated_stat(key)
  244. return {
  245. "artifact_ref": ref,
  246. "digest": metadata["sha256"],
  247. "row_count": int(metadata["row-count"]),
  248. "schema_hash": metadata["schema-sha256"],
  249. "expires_at": metadata["expires-at"],
  250. }
  251. def read(self, ref: str, expected_digest: str) -> pl.LazyFrame:
  252. if _DIGEST.fullmatch(str(expected_digest or "")) is None:
  253. raise ValueError("expected artifact digest is invalid")
  254. key = self._parse_ref(ref)
  255. _stat, metadata = self._validated_stat(
  256. key, expected_digest=expected_digest
  257. )
  258. response = self.client.get_object(self.bucket, key)
  259. digest = hashlib.sha256()
  260. payload = io.BytesIO()
  261. size = 0
  262. try:
  263. while chunk := response.read(1024 * 1024):
  264. size += len(chunk)
  265. if size > self.max_artifact_bytes:
  266. raise ValueError(
  267. "artifact size exceeds the configured limit"
  268. )
  269. digest.update(chunk)
  270. payload.write(chunk)
  271. finally:
  272. response.close()
  273. release = getattr(response, "release_conn", None)
  274. if callable(release):
  275. release()
  276. if digest.hexdigest() != expected_digest:
  277. raise ValueError("artifact digest does not match content")
  278. payload.seek(0)
  279. try:
  280. frame = pl.read_parquet(payload)
  281. except Exception as exc:
  282. raise ValueError("artifact is not valid Parquet") from exc
  283. if frame.height != int(metadata["row-count"]):
  284. raise ValueError("artifact row count does not match metadata")
  285. if frame.height > self.max_rows:
  286. raise ValueError("artifact row count exceeds the configured limit")
  287. if frame.estimated_size() > self.memory_limit_bytes:
  288. raise ValueError("artifact frame exceeds the configured memory limit")
  289. if _schema_hash(frame) != metadata["schema-sha256"]:
  290. raise ValueError("artifact schema does not match metadata")
  291. return frame.lazy()
  292. class PostgresArtifactResolver:
  293. """Resolve a canonical artifact binding without accepting caller paths."""
  294. def __init__(self, engine, artifact_store: ArtifactStore):
  295. self.engine = engine
  296. self.artifact_store = artifact_store
  297. def resolve(self, *, binding_id: str, correlation_id: str) -> dict[str, Any]:
  298. binding = _uid(binding_id, "artifact binding id")
  299. correlation = _uid(correlation_id, "artifact correlation id")
  300. statement = text(
  301. """
  302. SELECT object_ref, binding_hash
  303. FROM public.dataflow_dataset_bindings
  304. WHERE id = CAST(:binding_id AS uuid)
  305. AND object_kind = 'parquet_artifact'
  306. AND access_mode IN ('read', 'read_write')
  307. """
  308. )
  309. with self.engine.connect() as connection:
  310. row = connection.execute(
  311. statement, {"binding_id": binding}
  312. ).mappings().one_or_none()
  313. if row is None:
  314. raise ValueError("canonical artifact binding was not found")
  315. artifact_ref = str(row["object_ref"])
  316. if f"/rules/{correlation}/" not in artifact_ref:
  317. raise ValueError(
  318. "artifact binding does not match the execution correlation"
  319. )
  320. return {
  321. **self.artifact_store.describe(artifact_ref),
  322. "binding_hash": str(row["binding_hash"]),
  323. }