""" Router for CBPOA analysis endpoints Time series trends, district aggregation, and weather-health correlations """ from fastapi import APIRouter, HTTPException, Query from datetime import datetime, timedelta from typing import List, Literal from functools import lru_cache import random import pandas as pd from pydantic import BaseModel, Field from config import DATA_DIR, RISK_HIGH, PROJECT_ROOT, WUHAN_BOUNDS, LAT_STEP, LON_STEP from data.case_loader import load_cases_by_district_daily from utils.date_helpers import get_latest_date from utils.geojson import parse_geojson_file, load_districts from utils.geo import point_in_polygon from utils.risk import calculate_trend router = APIRouter(prefix="/api/analysis", tags=["analysis"]) class TrendResponse(BaseModel): """Response for trend data""" dates: List[str] = Field(..., description="Date labels") values: List[float] = Field(..., description="Risk values") trend: Literal["up", "down", "stable"] = Field(..., description="Trend direction") class DistrictRisk(BaseModel): """District-level risk aggregation""" district: str = Field(..., description="District name") avg_risk: float = Field(..., description="Average risk value") avg_aqi: float = Field(..., description="Average AQI from weather stations in this district") population: int = Field(..., description="Population (sum of 100m LandScan cells in district)") high_risk_count: int = Field(..., description="Count of high risk grids") total_grids: int = Field(..., description="Total grids in district") total_cases: int = Field(..., description="Total recorded cases (real, from cases_by_district_daily)") class DistrictsResponse(BaseModel): """Response for districts aggregation""" districts: List[DistrictRisk] = Field(..., description="District risk data") timestamp: str = Field(..., description="Response timestamp") class CorrelationFactor(BaseModel): """Correlation factor data""" factor: str = Field(..., description="Factor name") correlation: float = Field(..., description="Correlation coefficient (-1 to 1)") significance: Literal["high", "medium", "low"] = Field(..., description="Statistical significance") description: str = Field(..., description="Factor description") class CorrelationsResponse(BaseModel): """Response for correlations""" correlations: List[CorrelationFactor] = Field(..., description="Correlation factors") timestamp: str = Field(..., description="Response timestamp") @router.get("/trend", response_model=TrendResponse) async def get_trend(days: int = Query(default=7, ge=1, le=30)): """ Get time series trend data from ACTUAL historical observations Args: days: Number of days for trend (1-30) Returns: Trend data with dates, values, and trend direction """ latest_date = get_latest_date() try: base_date = datetime.strptime(latest_date, "%Y%m%d") except ValueError: raise HTTPException(status_code=500, detail="Invalid date format in data files") dates = [] values = [] for i in range(days): date = base_date - timedelta(days=days - 1 - i) date_str = date.strftime("%Y%m%d") filepath = DATA_DIR / f"risk_{date_str}.geojson" if filepath.exists(): grids = parse_geojson_file(filepath) if grids: avg_risk = sum(g["risk_value"] for g in grids) / len(grids) values.append(round(avg_risk, 4)) else: values.append(0) else: values.append(0) dates.append(date.strftime("%Y-%m-%d")) # Preserve the full requested date range: a "7天" request must return 7 # contiguous points. Days with no geojson (or empty grids) stay 0 rather # than being dropped, which previously produced fewer, non-contiguous points. trend_direction = calculate_trend(values) return TrendResponse( dates=dates, values=values, trend=trend_direction, ) @lru_cache(maxsize=1) def _grid_district_lookup() -> dict: """Map precomputed r{row}_c{col} grid id -> district name (loaded once).""" path = PROJECT_ROOT / "processed" / "grid_district_mapping.parquet" if not path.exists(): return {} df = pd.read_parquet(path) # Some grids have a null district_name; drop them so the lookup only ever # returns valid strings (missing keys fall back to "其他"). df = df.dropna(subset=["district_name"]) return dict(zip(df["grid_id"].astype(str), df["district_name"].astype(str))) @lru_cache(maxsize=1) def _district_population() -> dict: """Real population per district. Sums the LandScan-derived population_density of every 100m cell (grid_100m_with_dem_pop.parquet) grouped by district via the grid->district mapping. Returns {district_name: total_population}. """ pop_path = PROJECT_ROOT / "processed" / "grid_100m_with_dem_pop.parquet" map_path = PROJECT_ROOT / "processed" / "grid_district_mapping.parquet" if not pop_path.exists() or not map_path.exists(): return {} pop = pd.read_parquet(pop_path, columns=["grid_id", "population_density"]) mapping = pd.read_parquet(map_path).dropna(subset=["district_name"]) joined = pop.merge(mapping, on="grid_id", how="inner") by_d = joined.groupby("district_name")["population_density"].sum() return {str(k): int(round(v)) for k, v in by_d.items()} @lru_cache(maxsize=1) def _district_avg_aqi() -> dict: """Real average AQI per district from weather station daily data. Each station (with lat/lon) is assigned to a district using the same grid->district mapping (100m grid spacing of 1/1110 deg, the convention the mapping was built with), then AQI is averaged per district across all daily observations. Returns {district_name: avg_aqi}. Districts with no station fall back to the city-wide mean in the caller. """ map_path = PROJECT_ROOT / "processed" / "grid_district_mapping.parquet" station_path = PROJECT_ROOT / "processed" / "weather" / "station_daily_2022.parquet" if not map_path.exists() or not station_path.exists(): return {} mapping = pd.read_parquet(map_path).dropna(subset=["district_name"]) lookup = dict(zip(mapping["grid_id"].astype(str), mapping["district_name"].astype(str))) station = pd.read_parquet(station_path, columns=["station_id", "lat", "lon", "AQI"]) step = 1.0 / 1110.0 # mapping grid spacing in degrees min_lat = WUHAN_BOUNDS["min_lat"] min_lon = WUHAN_BOUNDS["min_lon"] coords = station[["station_id", "lat", "lon"]].drop_duplicates() station_to_district = {} for _, r in coords.iterrows(): row = int((r["lat"] - min_lat) / step) col = int((r["lon"] - min_lon) / step) station_to_district[r["station_id"]] = lookup.get(f"r{row}_c{col}", "其他") station = station.copy() station["district"] = station["station_id"].map(station_to_district) in_district = station[station["district"] != "其他"] by_d = in_district.groupby("district")["AQI"].mean() return {str(k): round(float(v), 1) for k, v in by_d.items()} @lru_cache(maxsize=1) def _district_total_cases() -> dict: """Real total recorded cases per district from cases_by_district_daily. District labels are normalized to the canonical 13 区-suffixed names at the data-access boundary (data.case_loader), so this is a plain per-district sum. Returns {district: cases}. """ try: df = load_cases_by_district_daily() except FileNotFoundError: return {} by_district = df.groupby("district")["total_cases"].sum() return {str(d): int(v) for d, v in by_district.items()} @lru_cache(maxsize=8) def _aggregate_districts(date: str) -> list: """Aggregate per-district risk for a date. Assigns each 100m risk grid to a district via the precomputed grid->district mapping (O(1) dict lookup per grid) instead of per-grid point-in-polygon (which is ~100x slower over 140k grids). Cached by date. """ grids = parse_geojson_file(DATA_DIR / f"risk_{date}.geojson") lookup = _grid_district_lookup() agg: dict = {} for g in grids: row = int((g["latitude"] - WUHAN_BOUNDS["min_lat"]) / LAT_STEP) col = int((g["longitude"] - WUHAN_BOUNDS["min_lon"]) / LON_STEP) name = lookup.get(f"r{row}_c{col}", "其他") a = agg.setdefault(name, {"sum": 0.0, "count": 0, "high": 0}) risk = g["risk_value"] a["sum"] += risk a["count"] += 1 if risk >= RISK_HIGH: a["high"] += 1 pop_by_district = _district_population() aqi_by_district = _district_avg_aqi() cases_by_district = _district_total_cases() # City-wide mean AQI as fallback for districts without a weather station. city_avg_aqi = round(sum(aqi_by_district.values()) / len(aqi_by_district), 1) if aqi_by_district else 0.0 result = [] for name, a in agg.items(): if a["count"] == 0: continue avg = a["sum"] / a["count"] result.append({ "district": name, "avg_risk": round(avg, 4), "avg_aqi": aqi_by_district.get(name, city_avg_aqi), "population": pop_by_district.get(name, 0), "high_risk_count": a["high"], "total_grids": a["count"], "total_cases": cases_by_district.get(name, 0), }) # '其他' (unassigned) last, otherwise by descending risk result.sort(key=lambda d: (d["district"] == "其他", -d["avg_risk"])) return result @router.get("/districts", response_model=DistrictsResponse) async def get_districts(): """Get district-level risk aggregation (cached per date).""" latest_date = get_latest_date() filepath = DATA_DIR / f"risk_{latest_date}.geojson" if not filepath.exists(): raise HTTPException(status_code=404, detail=f"No data found for date {latest_date}") districts = [DistrictRisk(**d) for d in _aggregate_districts(latest_date)] return DistrictsResponse( districts=districts, timestamp=datetime.now().isoformat(), ) @router.get("/correlations", response_model=CorrelationsResponse) async def get_correlations(): """ Get weather-health correlation analysis Returns: Correlation factors with coefficients and significance """ latest_date = get_latest_date() filepath = DATA_DIR / f"risk_{latest_date}.geojson" if not filepath.exists(): raise HTTPException(status_code=404, detail=f"No data found for date {latest_date}") grids = parse_geojson_file(filepath) if not grids: raise HTTPException(status_code=404, detail="No grid data found") # Calculate mock correlations based on risk patterns # In production, this would use actual weather and health data avg_risk = sum(g["risk_value"] for g in grids) / len(grids) risk_variance = sum((g["risk_value"] - avg_risk) ** 2 for g in grids) / len(grids) # Generate realistic correlation coefficients correlations = [ CorrelationFactor( factor="temperature", correlation=round(-0.45 - 0.1 * (avg_risk - 0.5), 3), significance="high" if risk_variance > 0.05 else "medium", description="Temperature vs risk: Lower temps correlate with higher risk" ), CorrelationFactor( factor="humidity", correlation=round(0.32 + 0.15 * (avg_risk - 0.5), 3), significance="medium", description="Humidity vs risk: Higher humidity slightly increases risk" ), CorrelationFactor( factor="PM2.5", correlation=round(0.58 + 0.1 * (avg_risk - 0.5), 3), significance="high", description="PM2.5 vs risk: Strong positive correlation" ), CorrelationFactor( factor="PM10", correlation=round(0.51 + 0.08 * (avg_risk - 0.5), 3), significance="high", description="PM10 vs risk: Moderate positive correlation" ), CorrelationFactor( factor="wind_speed", correlation=round(-0.28 - 0.05 * (avg_risk - 0.5), 3), significance="low", description="Wind speed vs risk: Higher wind disperses pollutants" ), CorrelationFactor( factor="population_density", correlation=round(0.42 + 0.12 * (avg_risk - 0.5), 3), significance="high", description="Population density vs risk: Dense areas show higher transmission" ), ] return CorrelationsResponse( correlations=correlations, timestamp=datetime.now().isoformat() )