from __future__ import annotations from datetime import datetime, timezone from typing import Any from uuid import UUID from geoalchemy2.functions import ST_Intersects, ST_MakeEnvelope from geoalchemy2.shape import to_shape from shapely.geometry import mapping from sqlalchemy.orm import Session from app.core.errors import AppError from app.models import Dataset, VectorFeature from app.schemas.temporal import ( TemporalComparisonRequest, TemporalComparisonResponse, TemporalDatasetRef, TemporalMetricComparison, TemporalObjectChanges, TemporalSeriesDataset, TemporalSeriesRead, ) from app.services.vector_feature_service import VectorFeatureService class TemporalAnalysisService: IDENTITY_COMPARISON_LIMIT = 5_000 @staticmethod def list_series(db: Session, project_id: UUID) -> list[TemporalSeriesRead]: rows = ( db.query(Dataset) .filter(Dataset.project_id == project_id) .filter(Dataset.temporal_series_key.isnot(None)) .filter(Dataset.observed_at.isnot(None)) .order_by(Dataset.temporal_series_key.asc(), Dataset.observed_at.asc()) .all() ) grouped: dict[str, list[Dataset]] = {} for row in rows: if row.temporal_series_key: grouped.setdefault(row.temporal_series_key, []).append(row) result: list[TemporalSeriesRead] = [] for key, datasets in grouped.items(): observed = [item.observed_at for item in datasets if item.observed_at is not None] if not observed: continue result.append( TemporalSeriesRead( temporal_series_key=key, source_name=datasets[-1].source_name, reference_layer_name=datasets[-1].reference_layer_name, dataset_count=len(datasets), first_observed_at=min(observed), last_observed_at=max(observed), datasets=[ TemporalSeriesDataset( id=item.id, name=item.name, observed_at=item.observed_at, source_version=item.source_version, feature_count=(item.metadata_json or {}).get("feature_count") if isinstance(item.metadata_json, dict) else None, ) for item in datasets if item.observed_at is not None ], ) ) return result @staticmethod def compare( db: Session, *, project_id: UUID, payload: TemporalComparisonRequest, ) -> TemporalComparisonResponse: if payload.earlier_dataset_id == payload.later_dataset_id: raise AppError( code="INVALID_TEMPORAL_COMPARISON", message="Choose two different dataset snapshots", status_code=400, ) earlier = TemporalAnalysisService._get_temporal_dataset(db, project_id, payload.earlier_dataset_id, "Earlier") later = TemporalAnalysisService._get_temporal_dataset(db, project_id, payload.later_dataset_id, "Later") if earlier.temporal_series_key != later.temporal_series_key: raise AppError( code="INCOMPATIBLE_TEMPORAL_SERIES", message="Dataset snapshots must belong to the same temporal series", details={ "earlier_series": earlier.temporal_series_key, "later_series": later.temporal_series_key, }, status_code=400, ) if earlier.observed_at >= later.observed_at: raise AppError( code="INVALID_TEMPORAL_ORDER", message="Earlier snapshot must have an observation date before the later snapshot", status_code=400, ) bbox = payload.bbox.model_dump() earlier_summary = VectorFeatureService.summarize_features_by_bbox(db, dataset=earlier, bbox=bbox) later_summary = VectorFeatureService.summarize_features_by_bbox(db, dataset=later, bbox=bbox) if ( earlier_summary["aggregation_method"] != later_summary["aggregation_method"] or earlier_summary["metric_unit"] != later_summary["metric_unit"] ): raise AppError( code="INCOMPATIBLE_TEMPORAL_AGGREGATION", message="Dataset snapshots use incompatible aggregation semantics", status_code=400, ) earlier_value = float(earlier_summary["metric_value"]) later_value = float(later_summary["metric_value"]) absolute_change = later_value - earlier_value percent_change = (absolute_change / earlier_value * 100.0) if earlier_value else None warnings = [ warning for warning in {earlier_summary.get("warning"), later_summary.get("warning")} if warning ] object_changes, geojson, identity_warnings = TemporalAnalysisService._compare_identity_features( db, earlier=earlier, later=later, bbox=bbox, preview_limit=payload.preview_limit, ) warnings.extend(identity_warnings) return TemporalComparisonResponse( temporal_series_key=earlier.temporal_series_key, earlier=TemporalDatasetRef( id=earlier.id, name=earlier.name, observed_at=earlier.observed_at, source_version=earlier.source_version, ), later=TemporalDatasetRef( id=later.id, name=later.name, observed_at=later.observed_at, source_version=later.source_version, ), selection_bbox=payload.bbox, metric=TemporalMetricComparison( label=str(later_summary["metric_label"]), unit=str(later_summary["metric_unit"]), aggregation_method=str(later_summary["aggregation_method"]), earlier_value=earlier_value, later_value=later_value, absolute_change=absolute_change, percent_change=percent_change, is_estimate=bool(earlier_summary["is_estimate"] or later_summary["is_estimate"]), ), object_changes=object_changes, geojson=geojson, warnings=warnings, generated_at=datetime.now(timezone.utc), ) @staticmethod def _get_temporal_dataset(db: Session, project_id: UUID, dataset_id: UUID, label: str) -> Dataset: dataset = db.get(Dataset, dataset_id) if not dataset or dataset.project_id != project_id: raise AppError(code="DATASET_NOT_FOUND", message=f"{label} dataset not found", status_code=404) if dataset.dataset_type not in {"vector", "geojson"}: raise AppError( code="DATASET_NOT_VECTOR", message="Temporal selection comparison currently requires vector datasets", status_code=400, ) if not dataset.temporal_series_key or not dataset.observed_at: raise AppError( code="TEMPORAL_METADATA_MISSING", message=f"{label} dataset has no explicit temporal series and observation date", status_code=400, ) return dataset @staticmethod def _compare_identity_features( db: Session, *, earlier: Dataset, later: Dataset, bbox: dict[str, Any], preview_limit: int, ) -> tuple[TemporalObjectChanges, dict[str, Any], list[str]]: earlier_config = earlier.source_metadata if isinstance(earlier.source_metadata, dict) else {} later_config = later.source_metadata if isinstance(later.source_metadata, dict) else {} if not earlier_config.get("identity_stable") or not later_config.get("identity_stable"): return ( TemporalObjectChanges(available=False), {"type": "FeatureCollection", "features": []}, ["Wijzigingen van individuele objecten kunnen voor deze bron niet betrouwbaar worden gevolgd."], ) normalized_bbox = VectorFeatureService._normalize_selection_bbox(bbox) envelope = ST_MakeEnvelope( normalized_bbox["min_x"], normalized_bbox["min_y"], normalized_bbox["max_x"], normalized_bbox["max_y"], 4326, ) def load(dataset_id: UUID) -> list[VectorFeature]: return ( db.query(VectorFeature) .filter(VectorFeature.dataset_id == dataset_id) .filter(ST_Intersects(VectorFeature.geometry, envelope)) .filter(VectorFeature.source_feature_id.isnot(None)) .order_by(VectorFeature.source_feature_id.asc()) .limit(TemporalAnalysisService.IDENTITY_COMPARISON_LIMIT + 1) .all() ) earlier_rows = load(earlier.id) later_rows = load(later.id) if ( len(earlier_rows) > TemporalAnalysisService.IDENTITY_COMPARISON_LIMIT or len(later_rows) > TemporalAnalysisService.IDENTITY_COMPARISON_LIMIT ): return ( TemporalObjectChanges(available=False), {"type": "FeatureCollection", "features": []}, ["Object-level preview was skipped because the selection exceeds the 5,000 feature safety limit."], ) earlier_by_id = {str(row.source_feature_id): row for row in earlier_rows if row.source_feature_id} later_by_id = {str(row.source_feature_id): row for row in later_rows if row.source_feature_id} earlier_ids = set(earlier_by_id) later_ids = set(later_by_id) added_ids = sorted(later_ids - earlier_ids) removed_ids = sorted(earlier_ids - later_ids) common_ids = sorted(earlier_ids & later_ids) comparison_property = str(later_config.get("comparison_property") or "").strip() or None modified_ids: list[str] = [] unchanged_ids: list[str] = [] for feature_id in common_ids: earlier_row = earlier_by_id[feature_id] later_row = later_by_id[feature_id] geometry_changed = not to_shape(earlier_row.geometry).equals(to_shape(later_row.geometry)) value_changed = False if comparison_property: value_changed = (earlier_row.properties_json or {}).get(comparison_property) != ( later_row.properties_json or {} ).get(comparison_property) (modified_ids if geometry_changed or value_changed else unchanged_ids).append(feature_id) features: list[dict[str, Any]] = [] for change_type, feature_ids, rows in ( ("added", added_ids, later_by_id), ("removed", removed_ids, earlier_by_id), ("modified", modified_ids, later_by_id), ): for feature_id in feature_ids: if len(features) >= preview_limit: break row = rows[feature_id] properties = dict(row.properties_json or {}) properties.update( { "change_type": change_type, "source_feature_id": feature_id, "earlier_dataset_id": str(earlier.id), "later_dataset_id": str(later.id), } ) if change_type == "modified" and comparison_property: before = (earlier_by_id[feature_id].properties_json or {}).get(comparison_property) after = (later_by_id[feature_id].properties_json or {}).get(comparison_property) properties.update({"value_before": before, "value_after": after}) if isinstance(before, (int, float)) and isinstance(after, (int, float)): properties["value_delta"] = after - before features.append( { "type": "Feature", "id": str(row.id), "geometry": mapping(to_shape(row.geometry)), "properties": properties, } ) warnings: list[str] = [] total_changes = len(added_ids) + len(removed_ids) + len(modified_ids) if total_changes > preview_limit: warnings.append( f"The map shows the first {preview_limit} of {total_changes} changed features; counts remain complete." ) return ( TemporalObjectChanges( available=True, added_count=len(added_ids), removed_count=len(removed_ids), modified_count=len(modified_ids), unchanged_count=len(unchanged_ids), ), {"type": "FeatureCollection", "features": features}, warnings, )