|
| 1 | +"""Materialization helpers for external market signal artifact trees.""" |
| 2 | + |
| 3 | +from __future__ import annotations |
| 4 | + |
| 5 | +import hashlib |
| 6 | +import json |
| 7 | +import posixpath |
| 8 | +from pathlib import Path |
| 9 | +from typing import Any, Iterable, Mapping |
| 10 | + |
| 11 | +from .strategy_plugin_artifacts import download_gcs_object, parse_gcs_uri |
| 12 | + |
| 13 | + |
| 14 | +MARKET_SIGNAL_ARTIFACT_LINK_FIELDS = frozenset( |
| 15 | + { |
| 16 | + "bundle_path", |
| 17 | + "catalog_path", |
| 18 | + "consumer_contract_registry_manifest_path", |
| 19 | + "handoff_manifest_path", |
| 20 | + "quality_report_path", |
| 21 | + "registry_path", |
| 22 | + "signal_bundle_manifest_path", |
| 23 | + "source_family_catalog_manifest_path", |
| 24 | + } |
| 25 | +) |
| 26 | + |
| 27 | + |
| 28 | +def materialize_market_signal_artifact_tree( |
| 29 | + reference: str, |
| 30 | + *, |
| 31 | + cache_dir: Path, |
| 32 | + client_factory: Any = None, |
| 33 | + link_fields: Iterable[str] | None = None, |
| 34 | +) -> tuple[Path, dict[str, Any]]: |
| 35 | + """Return a local path for a market signal artifact and its linked JSON tree.""" |
| 36 | + |
| 37 | + raw_reference = _required_string(reference, field_name="reference") |
| 38 | + if not raw_reference.startswith("gs://"): |
| 39 | + local_path = Path(raw_reference).expanduser() |
| 40 | + return local_path, { |
| 41 | + "source_uri": None, |
| 42 | + "local_path": raw_reference, |
| 43 | + "cache_dir": None, |
| 44 | + "materialized_count": 0, |
| 45 | + "materialized_paths": (), |
| 46 | + } |
| 47 | + |
| 48 | + fields = frozenset(link_fields or MARKET_SIGNAL_ARTIFACT_LINK_FIELDS) |
| 49 | + cache_root = cache_root_for_market_signal_artifact_tree( |
| 50 | + raw_reference, |
| 51 | + cache_dir=cache_dir, |
| 52 | + ) |
| 53 | + visited: dict[str, Path] = {} |
| 54 | + _materialize_gcs_json_tree( |
| 55 | + raw_reference, |
| 56 | + cache_root=cache_root, |
| 57 | + client_factory=client_factory, |
| 58 | + link_fields=fields, |
| 59 | + visited=visited, |
| 60 | + ) |
| 61 | + _, object_name = parse_gcs_uri(raw_reference) |
| 62 | + local_path = local_path_for_gcs_object(cache_root, object_name) |
| 63 | + return local_path, { |
| 64 | + "source_uri": raw_reference, |
| 65 | + "local_path": str(local_path), |
| 66 | + "cache_dir": str(cache_root), |
| 67 | + "materialized_count": len(visited), |
| 68 | + "materialized_paths": tuple(str(path) for path in visited.values()), |
| 69 | + } |
| 70 | + |
| 71 | + |
| 72 | +def cache_root_for_market_signal_artifact_tree(reference: str, *, cache_dir: Path) -> Path: |
| 73 | + raw_reference = _required_string(reference, field_name="reference") |
| 74 | + digest = hashlib.sha256(raw_reference.encode("utf-8")).hexdigest()[:16] |
| 75 | + return cache_dir / digest |
| 76 | + |
| 77 | + |
| 78 | +def local_path_for_gcs_object(cache_root: Path, object_name: str) -> Path: |
| 79 | + raw_object_name = _required_string(object_name, field_name="object_name") |
| 80 | + if raw_object_name.startswith("/"): |
| 81 | + raise ValueError(f"GCS object name must be relative: {raw_object_name}") |
| 82 | + normalized = posixpath.normpath(raw_object_name) |
| 83 | + if normalized in {"", ".", ".."} or normalized.startswith("../"): |
| 84 | + raise ValueError(f"GCS object name escapes the cache root: {raw_object_name}") |
| 85 | + return cache_root.joinpath(*normalized.split("/")) |
| 86 | + |
| 87 | + |
| 88 | +def resolve_gcs_artifact_reference(base_uri: str, reference: str) -> str: |
| 89 | + raw_reference = _required_string(reference, field_name="reference") |
| 90 | + if raw_reference.startswith("gs://"): |
| 91 | + parse_gcs_uri(raw_reference) |
| 92 | + return raw_reference |
| 93 | + if "://" in raw_reference: |
| 94 | + raise ValueError(f"Unsupported market signal artifact reference: {raw_reference}") |
| 95 | + if raw_reference.startswith("/"): |
| 96 | + raise ValueError( |
| 97 | + "GCS market signal artifacts must use relative linked paths or gs:// URIs: " |
| 98 | + f"{raw_reference}" |
| 99 | + ) |
| 100 | + |
| 101 | + bucket_name, object_name = parse_gcs_uri(base_uri) |
| 102 | + base_dir = posixpath.dirname(object_name) |
| 103 | + resolved = posixpath.normpath(posixpath.join(base_dir, raw_reference)) |
| 104 | + if resolved in {"", ".", ".."} or resolved.startswith("../"): |
| 105 | + raise ValueError( |
| 106 | + "GCS market signal artifact reference escapes the bucket root: " |
| 107 | + f"{raw_reference}" |
| 108 | + ) |
| 109 | + return f"gs://{bucket_name}/{resolved}" |
| 110 | + |
| 111 | + |
| 112 | +def _materialize_gcs_json_tree( |
| 113 | + uri: str, |
| 114 | + *, |
| 115 | + cache_root: Path, |
| 116 | + client_factory: Any, |
| 117 | + link_fields: frozenset[str], |
| 118 | + visited: dict[str, Path], |
| 119 | +) -> None: |
| 120 | + if uri in visited: |
| 121 | + return |
| 122 | + |
| 123 | + _, object_name = parse_gcs_uri(uri) |
| 124 | + local_path = local_path_for_gcs_object(cache_root, object_name) |
| 125 | + download_gcs_object(uri, local_path, client_factory=client_factory) |
| 126 | + visited[uri] = local_path |
| 127 | + |
| 128 | + payload = _read_json_object(local_path) |
| 129 | + if payload is None: |
| 130 | + return |
| 131 | + for linked_uri in _iter_linked_gcs_artifact_uris( |
| 132 | + payload, |
| 133 | + base_uri=uri, |
| 134 | + link_fields=link_fields, |
| 135 | + ): |
| 136 | + _materialize_gcs_json_tree( |
| 137 | + linked_uri, |
| 138 | + cache_root=cache_root, |
| 139 | + client_factory=client_factory, |
| 140 | + link_fields=link_fields, |
| 141 | + visited=visited, |
| 142 | + ) |
| 143 | + |
| 144 | + |
| 145 | +def _read_json_object(path: Path) -> Mapping[str, Any] | list[Any] | None: |
| 146 | + if path.suffix.lower() != ".json": |
| 147 | + return None |
| 148 | + try: |
| 149 | + payload = json.loads(path.read_text(encoding="utf-8")) |
| 150 | + except json.JSONDecodeError as exc: |
| 151 | + raise ValueError(f"Invalid JSON market signal artifact: {path}") from exc |
| 152 | + if not isinstance(payload, (dict, list)): |
| 153 | + return None |
| 154 | + return payload |
| 155 | + |
| 156 | + |
| 157 | +def _iter_linked_gcs_artifact_uris( |
| 158 | + payload: Any, |
| 159 | + *, |
| 160 | + base_uri: str, |
| 161 | + link_fields: frozenset[str], |
| 162 | +) -> Iterable[str]: |
| 163 | + if isinstance(payload, Mapping): |
| 164 | + for key, value in payload.items(): |
| 165 | + if key in link_fields and isinstance(value, str) and value.strip(): |
| 166 | + yield resolve_gcs_artifact_reference(base_uri, value.strip()) |
| 167 | + yield from _iter_linked_gcs_artifact_uris( |
| 168 | + value, |
| 169 | + base_uri=base_uri, |
| 170 | + link_fields=link_fields, |
| 171 | + ) |
| 172 | + elif isinstance(payload, list): |
| 173 | + for item in payload: |
| 174 | + yield from _iter_linked_gcs_artifact_uris( |
| 175 | + item, |
| 176 | + base_uri=base_uri, |
| 177 | + link_fields=link_fields, |
| 178 | + ) |
| 179 | + |
| 180 | + |
| 181 | +def _required_string(value: Any, *, field_name: str) -> str: |
| 182 | + text = str(value or "").strip() |
| 183 | + if not text: |
| 184 | + raise ValueError(f"{field_name} must be a non-empty string") |
| 185 | + return text |
0 commit comments