from __future__ import annotations from hashlib import sha256 import json from math import isfinite from pathlib import Path from typing import Any from geoalchemy2.shape import to_shape from pyproj import CRS, Transformer from shapely.geometry import box from shapely.ops import transform as shapely_transform from shapely.ops import unary_union from app.core.errors import AppError from app.models import Area, Dataset, DatasetVersion from app.services.storage_service import StorageService class TileManifestService: """Versioned provenance and integrity contract for inference tile sets.""" CONTRACT_KEY = "geointel.raster.tile-manifest" CONTRACT_VERSION = "2.0.0" _BINDING_FIELDS = ( "source_dataset_id", "source_dataset_checksum_sha256", "source_dataset_size_bytes", "source_registry_id", "source_snapshot_id", "source_snapshot_checksum_sha256", "data_contract_key", "data_contract_version", "source_version", "dataset_version_id", "dataset_version", "dataset_version_checksum_sha256", "source_area_id", "source_area_geometry_sha256", ) _REQUIRED_INFERENCE_BINDING_FIELDS = ( "source_dataset_checksum_sha256", "source_registry_id", "source_snapshot_id", "source_snapshot_checksum_sha256", "data_contract_key", "data_contract_version", "dataset_version_id", "dataset_version", "dataset_version_checksum_sha256", ) _CHECKSUM_FIELDS = ( "source_dataset_checksum_sha256", "source_snapshot_checksum_sha256", "dataset_version_checksum_sha256", ) @staticmethod def file_sha256(path: str | Path) -> str: digest = sha256() with Path(path).open("rb") as stream: for chunk in iter(lambda: stream.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() @staticmethod def _latest_dataset_version(dataset: Dataset) -> DatasetVersion | None: versions = list(dataset.versions or []) if not versions: return None return max(versions, key=lambda item: (int(item.version or 0), str(item.id or ""))) @staticmethod def _source_snapshot_checksum(dataset: Dataset) -> str | None: snapshot = dataset.source_snapshot checksum = getattr(snapshot, "checksum_sha256", None) if snapshot is not None else None return str(checksum).lower() if checksum else None @staticmethod def _area_geometry_binding(db, dataset: Dataset) -> tuple[str | None, str | None]: if dataset.area_id is None: return None, None area = db.get(Area, dataset.area_id) if area is None or area.geometry is None: return str(dataset.area_id), None geometry = to_shape(area.geometry) return str(dataset.area_id), sha256(geometry.wkb).hexdigest() @classmethod def dataset_binding(cls, db, dataset: Dataset) -> dict[str, Any]: version = cls._latest_dataset_version(dataset) area_id, area_geometry_sha256 = cls._area_geometry_binding(db, dataset) return { "manifest_contract_key": cls.CONTRACT_KEY, "manifest_contract_version": cls.CONTRACT_VERSION, "source_dataset_id": str(dataset.id), "source_raster_id": str(dataset.id), "source_dataset_checksum_sha256": ( str(dataset.checksum_sha256).lower() if dataset.checksum_sha256 else None ), "source_dataset_size_bytes": dataset.size_bytes, "source_registry_id": str(dataset.source_registry_id) if dataset.source_registry_id else None, "source_snapshot_id": str(dataset.source_snapshot_id) if dataset.source_snapshot_id else None, "source_snapshot_checksum_sha256": cls._source_snapshot_checksum(dataset), "data_contract_key": dataset.data_contract_key, "data_contract_version": dataset.data_contract_version, "source_version": dataset.source_version, "dataset_version_id": str(version.id) if version is not None and version.id else None, "dataset_version": int(version.version) if version is not None and version.version is not None else None, "dataset_version_checksum_sha256": ( str(version.checksum_sha256).lower() if version is not None and version.checksum_sha256 else None ), "source_area_id": area_id, "source_area_geometry_sha256": area_geometry_sha256, } @staticmethod def tile_integrity(path: str | Path) -> dict[str, Any]: resolved = Path(path) return { "size_bytes": resolved.stat().st_size, "sha256": TileManifestService.file_sha256(resolved), } @staticmethod def _error( error_prefix: str, suffix: str, message: str, *, details: dict[str, Any] | None = None, ) -> AppError: return AppError( code=f"{error_prefix}_TILE_MANIFEST_{suffix}", message=message, details=details, status_code=422, ) @staticmethod def _bounds_values(value: Any) -> tuple[float, float, float, float] | None: if isinstance(value, dict): aliases = ( ("min_x", "min_y", "max_x", "max_y"), ("minx", "miny", "maxx", "maxy"), ("left", "bottom", "right", "top"), ) selected = next( ([value.get(key) for key in keys] for keys in aliases if all(key in value for key in keys)), None, ) elif isinstance(value, (list, tuple)) and len(value) == 4: selected = list(value) else: return None try: bounds = tuple(float(item) for item in selected) if selected is not None else None except (TypeError, ValueError): return None if bounds is None or not all(isfinite(item) for item in bounds): return None if bounds[0] >= bounds[2] or bounds[1] >= bounds[3]: return None return bounds @staticmethod def _to_epsg4326(bounds: tuple[float, float, float, float], raw_crs: Any): source_crs = CRS.from_user_input(raw_crs) geometry = box(*bounds) if not source_crs.equals(CRS.from_epsg(4326)): transformer = Transformer.from_crs(source_crs, "EPSG:4326", always_xy=True) geometry = shapely_transform(transformer.transform, geometry) if geometry.is_empty or not geometry.is_valid: raise ValueError("Bounds do not form a valid transformed geometry") if not all(isfinite(float(value)) for value in geometry.bounds): raise ValueError("Bounds transform to non-finite coordinates") return geometry @classmethod def _manifest_coverage(cls, manifest: dict[str, Any], *, error_prefix: str): tiles = manifest.get("tiles") default_crs = manifest.get("crs") or manifest.get("source_crs") or manifest.get("dataset_crs") parts = [] for index, tile in enumerate(tiles if isinstance(tiles, list) else []): if not isinstance(tile, dict): raise cls._error( error_prefix, "SCOPE_MISMATCH", "Tile manifest entries must be objects with explicit spatial metadata.", details={"tile_index": index}, ) bounds = cls._bounds_values(tile.get("bounds")) raw_crs = tile.get("crs") or default_crs if bounds is None or not raw_crs: raise cls._error( error_prefix, "SCOPE_MISMATCH", "Every inference tile requires finite bounds and an explicit CRS.", details={"tile_index": index}, ) try: parts.append(cls._to_epsg4326(bounds, raw_crs)) except Exception as exc: raise cls._error( error_prefix, "SCOPE_MISMATCH", "Inference tile bounds or CRS could not be normalized to EPSG:4326.", details={"tile_index": index, "reason": str(exc)}, ) from exc coverage = unary_union(parts) if coverage.is_empty or not coverage.is_valid: raise cls._error( error_prefix, "SCOPE_MISMATCH", "Inference tile union is empty or invalid.", ) min_x, min_y, max_x, max_y = coverage.bounds if min_x < -180 or min_y < -90 or max_x > 180 or max_y > 90: raise cls._error( error_prefix, "SCOPE_MISMATCH", "Inference tile union falls outside EPSG:4326 bounds.", details={"bounds": list(coverage.bounds)}, ) return coverage @classmethod def _validate_binding(cls, db, dataset: Dataset, manifest: dict[str, Any], *, error_prefix: str) -> dict[str, Any]: if ( manifest.get("manifest_contract_key") != cls.CONTRACT_KEY or manifest.get("manifest_contract_version") != cls.CONTRACT_VERSION ): raise cls._error( error_prefix, "PROVENANCE_MISMATCH", "Inference requires a versioned GeoIntel tile-manifest contract.", details={ "required_contract": f"{cls.CONTRACT_KEY}@{cls.CONTRACT_VERSION}", "manifest_contract": ( f"{manifest.get('manifest_contract_key')}@{manifest.get('manifest_contract_version')}" ), }, ) expected = cls.dataset_binding(db, dataset) missing = [ field for field in cls._REQUIRED_INFERENCE_BINDING_FIELDS if expected.get(field) in {None, ""} ] invalid_checksums = [ field for field in cls._CHECKSUM_FIELDS if len(str(expected.get(field) or "")) != 64 or any(character not in "0123456789abcdef" for character in str(expected.get(field) or "").lower()) ] if missing or invalid_checksums: raise cls._error( error_prefix, "PROVENANCE_MISMATCH", "The requested Dataset lacks complete immutable provenance for inference tiling.", details={ "missing_fields": missing, "invalid_checksum_fields": invalid_checksums, }, ) manifest_dataset_id = manifest.get("source_dataset_id") or manifest.get("source_raster_id") if str(manifest_dataset_id or "") != expected["source_dataset_id"]: raise cls._error( error_prefix, "DATASET_MISMATCH", "Tile manifest belongs to a different raster Dataset.", details={ "requested_dataset_id": expected["source_dataset_id"], "manifest_dataset_id": manifest_dataset_id, }, ) mismatches = {} for field in cls._BINDING_FIELDS: expected_value = expected.get(field) if expected_value is None or field == "source_dataset_id": continue observed_value = manifest.get(field) if str(observed_value) != str(expected_value): mismatches[field] = {"expected": expected_value, "observed": observed_value} if mismatches: raise cls._error( error_prefix, "PROVENANCE_MISMATCH", "Tile manifest provenance no longer matches the requested Dataset snapshot.", details={"mismatches": mismatches}, ) return expected @classmethod def _validate_tile_files( cls, manifest: dict[str, Any], manifest_path: Path, *, settings, error_prefix: str, ) -> list[str]: resolved_paths: list[str] = [] seen_paths: set[Path] = set() for index, tile in enumerate(manifest["tiles"]): raw_path = tile.get("path") if isinstance(tile, dict) else None if not isinstance(raw_path, str) or not raw_path.strip(): raise cls._error( error_prefix, "TILE_INTEGRITY_MISMATCH", "Every inference tile requires a path and immutable integrity evidence.", details={"tile_index": index}, ) candidate = Path(raw_path).expanduser() if not candidate.is_absolute(): candidate = manifest_path.parent / candidate candidate = StorageService.assert_within_storage_root( candidate, label="raster tile", settings=settings, ) if not candidate.is_file(): raise cls._error( error_prefix, "TILE_INTEGRITY_MISMATCH", "An inference tile referenced by the manifest does not exist.", details={"tile_index": index, "tile_path": str(candidate)}, ) if candidate in seen_paths: raise cls._error( error_prefix, "TILE_INTEGRITY_MISMATCH", "A tile path occurs more than once in the inference manifest.", details={"tile_index": index, "tile_path": str(candidate)}, ) seen_paths.add(candidate) observed_size = candidate.stat().st_size expected_size = tile.get("size_bytes") expected_checksum = str(tile.get("sha256") or "").strip().lower() if expected_size != observed_size or len(expected_checksum) != 64: raise cls._error( error_prefix, "TILE_INTEGRITY_MISMATCH", "Tile size/checksum evidence is missing or no longer matches the staged file.", details={ "tile_index": index, "expected_size_bytes": expected_size, "observed_size_bytes": observed_size, }, ) observed_checksum = cls.file_sha256(candidate) if observed_checksum != expected_checksum: raise cls._error( error_prefix, "TILE_INTEGRITY_MISMATCH", "Tile checksum no longer matches the immutable manifest evidence.", details={ "tile_index": index, "expected_sha256": expected_checksum, "observed_sha256": observed_checksum, }, ) resolved_paths.append(str(candidate)) declared_count = manifest.get("count") if declared_count != len(resolved_paths): raise cls._error( error_prefix, "TILE_INTEGRITY_MISMATCH", "Tile manifest count does not match its tile records.", details={"declared_count": declared_count, "tile_count": len(resolved_paths)}, ) return resolved_paths @classmethod def _validate_scope( cls, db, dataset: Dataset, manifest: dict[str, Any], coverage, *, error_prefix: str, ) -> None: manifest_bounds = cls._bounds_values(manifest.get("bounds")) manifest_crs = manifest.get("crs") or manifest.get("source_crs") or manifest.get("dataset_crs") dataset_bounds = cls._bounds_values(dataset.bounds_json) if dataset_bounds is None and isinstance(dataset.metadata_json, dict): dataset_bounds = cls._bounds_values( dataset.metadata_json.get("bounds_json") or dataset.metadata_json.get("bounds") ) if manifest_bounds is None or not manifest_crs or dataset_bounds is None or not dataset.crs: raise cls._error( error_prefix, "SCOPE_MISMATCH", "Dataset and tile manifest require explicit CRS and finite bounds for inference.", ) try: manifest_extent = cls._to_epsg4326(manifest_bounds, manifest_crs) dataset_extent = cls._to_epsg4326(dataset_bounds, dataset.crs) except Exception as exc: raise cls._error( error_prefix, "SCOPE_MISMATCH", "Dataset or manifest bounds could not be normalized to EPSG:4326.", details={"reason": str(exc)}, ) from exc tolerance = max(dataset_extent.bounds[2] - dataset_extent.bounds[0], dataset_extent.bounds[3] - dataset_extent.bounds[1]) * 1e-7 + 1e-10 if not manifest_extent.buffer(tolerance).covers(coverage): raise cls._error( error_prefix, "SCOPE_MISMATCH", "Tile union exceeds the extent declared by its manifest.", details={"tile_union_bounds": list(coverage.bounds), "manifest_bounds": list(manifest_extent.bounds)}, ) if not dataset_extent.buffer(tolerance).covers(coverage): raise cls._error( error_prefix, "SCOPE_MISMATCH", "Tile union exceeds the persisted Dataset extent.", details={"tile_union_bounds": list(coverage.bounds), "dataset_bounds": list(dataset_extent.bounds)}, ) if dataset.area_id is not None: area = db.get(Area, dataset.area_id) if area is None or area.geometry is None: raise cls._error( error_prefix, "SCOPE_MISMATCH", "Dataset references an Area that is unavailable for inference-scope validation.", details={"area_id": str(dataset.area_id)}, ) area_geometry = to_shape(area.geometry) if area_geometry.is_empty or not area_geometry.is_valid or not coverage.intersects(area_geometry): raise cls._error( error_prefix, "SCOPE_MISMATCH", "Tile union does not overlap the persisted Dataset Area.", details={"area_id": str(dataset.area_id), "tile_union_bounds": list(coverage.bounds)}, ) @classmethod def validate_for_inference( cls, db, dataset: Dataset, manifest: dict[str, Any], *, manifest_path: str | Path, settings, error_prefix: str, ) -> dict[str, Any]: resolved_manifest_path = StorageService.assert_within_storage_root( manifest_path, label="tile manifest", settings=settings, ) expected = cls._validate_binding(db, dataset, manifest, error_prefix=error_prefix) resolved_paths = cls._validate_tile_files( manifest, resolved_manifest_path, settings=settings, error_prefix=error_prefix, ) coverage = cls._manifest_coverage(manifest, error_prefix=error_prefix) cls._validate_scope(db, dataset, manifest, coverage, error_prefix=error_prefix) return { "manifest_contract_key": cls.CONTRACT_KEY, "manifest_contract_version": cls.CONTRACT_VERSION, "manifest_path": str(resolved_manifest_path), "manifest_sha256": cls.file_sha256(resolved_manifest_path), "source_dataset_id": expected["source_dataset_id"], "source_dataset_checksum_sha256": expected.get("source_dataset_checksum_sha256"), "source_snapshot_id": expected.get("source_snapshot_id"), "dataset_version_id": expected.get("dataset_version_id"), "source_area_id": expected.get("source_area_id"), "tile_count": len(resolved_paths), "tile_union_bounds_epsg4326": [float(value) for value in coverage.bounds], } def canonical_manifest_json(payload: dict[str, Any]) -> str: """Stable serializer shared by the writer and manifest-hash tests.""" return json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True)