from __future__ import annotations import json from pathlib import Path from typing import Any, Iterable from uuid import UUID from geoalchemy2.functions import ST_Intersects, ST_MakeEnvelope from geoalchemy2.shape import from_shape from geoalchemy2.shape import to_shape from shapely.geometry import mapping from shapely.geometry import shape from shapely.ops import transform as transform_geometry from shapely.validation import make_valid from sqlalchemy import Float, cast, func from app.core.errors import AppError from app.models import Dataset, VectorFeature FULL_AREA_CLIPPED_OPERATOR_TOOLS = { "provision_mol_population_history.py", "provision_official_landuse_timeseries.py", "provision_regional_grb_buildings.py", "provision_regional_grb_context.py", } class VectorFeatureService: @staticmethod def can_use_full_area_fast_path(dataset: Dataset, selection_area_id: UUID | None) -> bool: if selection_area_id is None or dataset.area_id != selection_area_id: return False source_metadata = dataset.source_metadata if isinstance(dataset.source_metadata, dict) else {} if source_metadata.get("geometry_clipped_to_area") is True: return True provenance = dataset.provenance_metadata if isinstance(dataset.provenance_metadata, dict) else {} return provenance.get("operator_tool") in FULL_AREA_CLIPPED_OPERATOR_TOOLS @staticmethod def _feature_row(dataset_id: UUID, feature: dict[str, Any], index: int, feature_class: str | None) -> VectorFeature | None: geometry_payload = feature.get("geometry") if geometry_payload is None: return None try: geometry = shape(geometry_payload) except Exception as exc: raise AppError(code="INVALID_GEOJSON", message=f"Invalid feature geometry at index {index}", status_code=400) from exc if geometry.is_empty: return None if not geometry.is_valid: geometry = make_valid(geometry) if geometry.is_empty or not geometry.is_valid: raise AppError(code="INVALID_GEOMETRY", message=f"Invalid feature geometry at index {index}", status_code=400) if geometry.has_z: geometry = transform_geometry(lambda x, y, z=None: (x, y), geometry) properties = feature.get("properties") if isinstance(feature.get("properties"), dict) else {} source_feature_id = feature.get("id") if source_feature_id is None: source_feature_id = properties.get("id") or properties.get("source_feature_id") return VectorFeature( dataset_id=dataset_id, feature_class=feature_class, source_feature_id=str(source_feature_id) if source_feature_id is not None else None, properties_json=properties, geometry=from_shape(geometry, srid=4326), ) @staticmethod def _normalize_selection_bbox(bbox: dict[str, Any]) -> dict[str, float | str]: try: min_x = float(bbox["min_x"]) min_y = float(bbox["min_y"]) max_x = float(bbox["max_x"]) max_y = float(bbox["max_y"]) except (KeyError, TypeError, ValueError) as exc: raise AppError( code="INVALID_SELECTION_BBOX", message="Selection bbox must include numeric min_x, min_y, max_x and max_y values", status_code=400, ) from exc crs = str(bbox.get("crs") or "EPSG:4326").upper() if crs != "EPSG:4326": raise AppError( code="UNSUPPORTED_SELECTION_CRS", message="Map selection currently supports EPSG:4326 bbox coordinates only", details={"crs": crs}, status_code=400, ) if min_x >= max_x or min_y >= max_y: raise AppError( code="INVALID_SELECTION_BBOX", message="Selection bbox must have min_x < max_x and min_y < max_y", status_code=400, ) if min_x < -180 or max_x > 180 or min_y < -90 or max_y > 90: raise AppError( code="INVALID_SELECTION_BBOX", message="Selection bbox is outside EPSG:4326 longitude/latitude bounds", status_code=400, ) return {"min_x": min_x, "min_y": min_y, "max_x": max_x, "max_y": max_y, "crs": "EPSG:4326"} @staticmethod def _row_to_geojson_feature(row: VectorFeature) -> dict[str, Any]: geometry_value = row.geometry try: geometry = geometry_value if hasattr(geometry_value, "__geo_interface__") else to_shape(geometry_value) except Exception as exc: raise AppError( code="INVALID_VECTOR_FEATURE_GEOMETRY", message="Persisted vector feature geometry could not be converted to GeoJSON", details={"vector_feature_id": str(row.id)}, status_code=500, ) from exc properties = dict(row.properties_json or {}) properties.update( { "vector_feature_id": str(row.id), "dataset_id": str(row.dataset_id), "source_feature_id": row.source_feature_id, "feature_class": row.feature_class, } ) return { "type": "Feature", "id": str(row.id), "geometry": mapping(geometry), "properties": properties, } @staticmethod def select_features_by_bbox( db, dataset_id: UUID, bbox: dict[str, Any], limit: int = 100, dataset: Dataset | None = None, selection_geometry: Any | None = None, selection_area_id: UUID | None = None, full_dataset_area: bool = False, ) -> dict[str, Any]: normalized_bbox = VectorFeatureService._normalize_selection_bbox(bbox) safe_limit = max(1, min(int(limit), 1000)) selection_shape = selection_geometry if selection_shape is None: selection_shape = ST_MakeEnvelope( normalized_bbox["min_x"], normalized_bbox["min_y"], normalized_bbox["max_x"], normalized_bbox["max_y"], 4326, ) query = db.query(VectorFeature).filter(VectorFeature.dataset_id == dataset_id) if not full_dataset_area: query = query.filter(ST_Intersects(VectorFeature.geometry, selection_shape)) if hasattr(query, "count"): total_feature_count = int(query.count()) else: # Lightweight unit-test sessions do not always implement Query.count(). total_feature_count = len(query.all()) rows = ( query.order_by(VectorFeature.created_at.asc()) .limit(safe_limit + 1) .all() ) truncated = total_feature_count > safe_limit selected_rows = rows[:safe_limit] features = [VectorFeatureService._row_to_geojson_feature(row) for row in selected_rows] summary = None if dataset and isinstance(dataset.source_metadata, dict) and dataset.source_metadata.get("selection_aggregation"): summary = VectorFeatureService.summarize_features_by_bbox( db, dataset=dataset, bbox=normalized_bbox, total_feature_count=total_feature_count, selection_geometry=selection_geometry, full_dataset_area=full_dataset_area, ) result = { "selection_bbox": normalized_bbox, "feature_count": len(features), "total_feature_count": total_feature_count, "limit": safe_limit, "truncated": truncated, "geojson": { "type": "FeatureCollection", "features": features, }, "summary": summary, } if selection_area_id is not None: result["selection_area_id"] = str(selection_area_id) return result @staticmethod def summarize_features_by_bbox( db, *, dataset: Dataset, bbox: dict[str, Any], total_feature_count: int | None = None, selection_geometry: Any | None = None, full_dataset_area: bool = False, ) -> dict[str, Any]: normalized_bbox = VectorFeatureService._normalize_selection_bbox(bbox) selection_shape = selection_geometry if selection_shape is None: selection_shape = ST_MakeEnvelope( normalized_bbox["min_x"], normalized_bbox["min_y"], normalized_bbox["max_x"], normalized_bbox["max_y"], 4326, ) selection_filter = (VectorFeature.dataset_id == dataset.id,) if not full_dataset_area: selection_filter += (ST_Intersects(VectorFeature.geometry, selection_shape),) feature_count = total_feature_count if feature_count is None: feature_count = int(db.query(func.count(VectorFeature.id)).filter(*selection_filter).scalar() or 0) source_metadata = dataset.source_metadata if isinstance(dataset.source_metadata, dict) else {} config = source_metadata.get("selection_aggregation") if not isinstance(config, dict): config = {} method = str(config.get("method") or "feature_count") label = str(config.get("label") or "Objecten") unit = str(config.get("unit") or "objecten") warning = str(config["warning"]) if config.get("warning") else None is_estimate = bool(config.get("is_estimate", False)) metric_value = float(feature_count) if method == "intersection_area": measured_geometry = ( VectorFeature.geometry if full_dataset_area else func.ST_Intersection(VectorFeature.geometry, selection_shape) ) area_expression = func.ST_Area(func.ST_Transform(measured_geometry, 31370)) area_m2 = db.query(func.coalesce(func.sum(area_expression), 0.0)).filter(*selection_filter).scalar() divisor = 10_000.0 if unit == "ha" else 1.0 metric_value = float(area_m2 or 0.0) / divisor elif method == "intersection_length": measured_geometry = ( VectorFeature.geometry if full_dataset_area else func.ST_Intersection(VectorFeature.geometry, selection_shape) ) length_expression = func.ST_Length(func.ST_Transform(measured_geometry, 31370)) length_m = db.query(func.coalesce(func.sum(length_expression), 0.0)).filter(*selection_filter).scalar() divisor = 1_000.0 if unit == "km" else 1.0 metric_value = float(length_m or 0.0) / divisor elif method in {"sum", "area_weighted_sum"}: property_name = str(config.get("property") or "").strip() if not property_name: raise AppError( code="INVALID_SELECTION_AGGREGATION", message="Dataset selection aggregation requires a numeric property", details={"dataset_id": str(dataset.id), "method": method}, status_code=500, ) numeric_value = cast(VectorFeature.properties_json.op("->>")(property_name), Float) value_expression = numeric_value if method == "area_weighted_sum" and not full_dataset_area: source_area = func.ST_Area(func.ST_Transform(VectorFeature.geometry, 31370)) intersection_area = func.ST_Area( func.ST_Transform(func.ST_Intersection(VectorFeature.geometry, selection_shape), 31370) ) coverage_ratio = intersection_area / func.nullif(source_area, 0.0) value_expression = numeric_value * coverage_ratio aggregate_value = ( db.query(func.coalesce(func.sum(value_expression), 0.0)) .filter(*selection_filter) .filter(VectorFeature.properties_json.op("->>")(property_name).isnot(None)) .scalar() ) metric_value = float(aggregate_value or 0.0) if method == "area_weighted_sum" and not full_dataset_area: partial_feature_count = ( db.query(func.count(VectorFeature.id)) .filter(*selection_filter) .filter(coverage_ratio < 0.999999) .scalar() ) is_estimate = bool(partial_feature_count) if not is_estimate and config.get("warning_only_when_estimate", True): warning = None elif method == "area_weighted_sum": is_estimate = False if config.get("warning_only_when_estimate", True): warning = None elif method != "feature_count": raise AppError( code="INVALID_SELECTION_AGGREGATION", message="Unsupported dataset selection aggregation", details={"dataset_id": str(dataset.id), "method": method}, status_code=500, ) return { "metric_label": label, "metric_value": metric_value, "metric_unit": unit, "aggregation_method": method, "feature_count": feature_count, "is_estimate": is_estimate, "warning": warning, } @staticmethod def persist_geojson_features( db, dataset_id: UUID, payload: dict[str, Any], feature_class: str | None = None, *, commit: bool = True, ) -> list[VectorFeature]: features = payload.get("features") if payload.get("type") != "FeatureCollection" or not isinstance(features, list): raise AppError(code="INVALID_GEOJSON", message="GeoJSON payload must be a FeatureCollection", status_code=400) persisted: list[VectorFeature] = [] for index, feature in enumerate(features): if not isinstance(feature, dict): raise AppError(code="INVALID_GEOJSON", message=f"Feature {index} must be an object", status_code=400) row = VectorFeatureService._feature_row(dataset_id, feature, index, feature_class) if row is None: continue db.add(row) persisted.append(row) if commit: db.flush() db.commit() return persisted @staticmethod def persist_geojson_partitions( db, dataset_id: UUID, partition_paths: Iterable[str | Path], feature_class: str | None = None, *, batch_size: int = 1000, ) -> int: if batch_size <= 0: raise ValueError("batch_size must be positive") persisted_count = 0 source_feature_ids: set[str] = set() for partition_path in partition_paths: path = Path(partition_path) try: payload = json.loads(path.read_text(encoding="utf-8")) except (OSError, json.JSONDecodeError) as exc: raise AppError( code="INVALID_GEOJSON_PARTITION", message=f"Could not read GeoJSON partition {path.name}", status_code=400, ) from exc features = payload.get("features") if payload.get("type") != "FeatureCollection" or not isinstance(features, list): raise AppError( code="INVALID_GEOJSON_PARTITION", message=f"GeoJSON partition {path.name} must be a FeatureCollection", status_code=400, ) batch: list[VectorFeature] = [] for index, feature in enumerate(features): if not isinstance(feature, dict): raise AppError( code="INVALID_GEOJSON_PARTITION", message=f"Feature {index} in {path.name} must be an object", status_code=400, ) row = VectorFeatureService._feature_row(dataset_id, feature, index, feature_class) if row is None: continue if row.source_feature_id: if row.source_feature_id in source_feature_ids: raise AppError( code="DUPLICATE_SOURCE_FEATURE", message=f"Duplicate source feature {row.source_feature_id} across regional partitions", status_code=400, ) source_feature_ids.add(row.source_feature_id) db.add(row) batch.append(row) persisted_count += 1 if len(batch) >= batch_size: db.flush() for persisted in batch: db.expunge(persisted) batch.clear() if batch: db.flush() for persisted in batch: db.expunge(persisted) return persisted_count