main.py

#!/usr/bin/env python3
"""
Causal ROI Dashboard API
FastAPI backend for serving metrics to React UI
"""

import logging
import os
from datetime import datetime, timedelta
from typing import List, Optional

from bigquery_client import BigQueryClient
from fastapi import BackgroundTasks, FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware

from models import (DiagnosticsResponse, EdgeHistory, JobResponse, PathDetail,
                    PathLeaderboard, SystemSummary)

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

app = FastAPI(
    title="Causal ROI Dashboard API",
    description="Self-learning revenue attribution system",
    version="1.0.0",
)

# CORS for React frontend
app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],  # Restrict in production
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)

# Initialize BigQuery client
bq = BigQueryClient(
    project_id=os.getenv("GCP_PROJECT_ID"),
    dataset_id=os.getenv("BQ_DATASET", "causal_roi"),
)


@app.get("/")
async def root():
    return {
        "service": "Causal ROI Dashboard API",
        "status": "running",
        "endpoints": [
            "/api/summary",
            "/api/paths/top",
            "/api/path/{path_id}",
            "/api/edge/{edge_id}/history",
            "/api/diagnostics",
        ],
    }


@app.get("/api/summary", response_model=SystemSummary)
async def get_system_summary():
    """
    Top bar metrics: Net ΔROI, Mean Velocity, System PCR
    """
    try:
        result = bq.fetch_one(
            """
            SELECT
                snapshot_date,
                net_delta_roi,
                mean_velocity_ms,
                system_pcr,
                delta_roi_change,
                velocity_change,
                pcr_change
            FROM `{dataset}.v_system_summary`
            ORDER BY snapshot_date DESC
            LIMIT 1
        """
        )

        if not result:
            raise HTTPException(status_code=404, detail="No summary data available")

        return SystemSummary(**result)

    except Exception as e:
        logger.error(f"Error fetching summary: {e}")
        raise HTTPException(status_code=500, detail=str(e))


@app.get("/api/paths/top", response_model=List[PathLeaderboard])
async def get_top_paths(window_days: int = 7, limit: int = 20, min_actions: int = 10):
    """
    Path leaderboard: top paths by ΔROI × P(causal)
    """
    try:
        results = bq.fetch_all(
            f"""
            SELECT
                path_id,
                avg_roi_delta AS delta_roi_7d,
                avg_velocity_ms AS velocity_ms,
                avg_pcr AS pcr,
                total_actions,
                total_outcomes,
                stability_score,
                rank
            FROM `{{dataset}}.v_top_paths_7d`
            WHERE total_actions >= {min_actions}
            ORDER BY rank
            LIMIT {limit}
        """
        )

        return [PathLeaderboard(**r) for r in results]

    except Exception as e:
        logger.error(f"Error fetching top paths: {e}")
        raise HTTPException(status_code=500, detail=str(e))


@app.get("/api/path/{path_id}", response_model=PathDetail)
async def get_path_detail(path_id: str):
    """
    Path detail: graph structure + edges + metrics timeseries
    """
    try:
        # Get path metrics
        path_metrics = bq.fetch_one(
            f"""
            SELECT
                path_id,
                avg_velocity_ms AS velocity_ms,
                avg_p_causal AS p_causal,
                avg_pcr AS pcr,
                avg_roi_delta AS roi_delta
            FROM `{{dataset}}.v_top_paths_7d`
            WHERE path_id = '{path_id}'
        """
        )

        if not path_metrics:
            raise HTTPException(status_code=404, detail=f"Path {path_id} not found")

        # Get edges for this path
        edges = bq.fetch_all(
            f"""
            SELECT
                edge_id,
                src_node_id,
                dst_node_id,
                src_label,
                dst_label,
                edge_type,
                latest_delta_roi,
                latest_latency_ms,
                decay_weight,
                avg_delta_roi_7d,
                total_spend_7d
            FROM `{{dataset}}.v_edge_details`
            WHERE edge_id IN (
                SELECT DISTINCT edge_id
                FROM `{{dataset}}.edge_metrics`
                WHERE path_id = '{path_id}'
            )
        """
        )

        # Get timeseries
        timeseries = bq.fetch_all(
            f"""
            SELECT
                date,
                weighted_roi,
                daily_spend,
                avg_latency,
                roi_7d_ma
            FROM `{{dataset}}.v_path_detail`
            WHERE path_id = '{path_id}'
            ORDER BY date DESC
            LIMIT 30
        """
        )

        return PathDetail(
            path_id=path_id, metrics=path_metrics, edges=edges, timeseries=timeseries
        )

    except HTTPException:
        raise
    except Exception as e:
        logger.error(f"Error fetching path detail: {e}")
        raise HTTPException(status_code=500, detail=str(e))


@app.get("/api/edge/{edge_id}/history", response_model=EdgeHistory)
async def get_edge_history(edge_id: str, days: int = 7):
    """
    Edge sparkline data: ΔROI over time
    """
    try:
        history = bq.fetch_all(
            f"""
            SELECT
                date,
                delta_roi,
                n_obs,
                spend,
                roi_3d_ma,
                roi_7d_stddev
            FROM `{{dataset}}.v_edge_history_7d`
            WHERE edge_id = '{edge_id}'
            ORDER BY date DESC
            LIMIT {days}
        """
        )

        if not history:
            raise HTTPException(status_code=404, detail=f"Edge {edge_id} not found")

        return EdgeHistory(edge_id=edge_id, history=history)

    except HTTPException:
        raise
    except Exception as e:
        logger.error(f"Error fetching edge history: {e}")
        raise HTTPException(status_code=500, detail=str(e))


@app.get("/api/diagnostics", response_model=DiagnosticsResponse)
async def get_diagnostics():
    """
    System health: data freshness, coverage, control quality
    """
    try:
        diagnostics = bq.fetch_all(
            """
            SELECT
                date,
                forecasts_logged,
                actions_logged,
                outcomes_logged,
                last_data_insert,
                attributed_paths,
                avg_attribution_weight,
                n_controls,
                control_quality_score
            FROM `{dataset}.v_diagnostics_daily`
            ORDER BY date DESC
            LIMIT 7
        """
        )

        # Compute data freshness
        latest = diagnostics[0] if diagnostics else None
        data_age_hours = None
        if latest and latest.get("last_data_insert"):
            age = datetime.now() - latest["last_data_insert"]
            data_age_hours = age.total_seconds() / 3600

        return DiagnosticsResponse(
            daily_stats=diagnostics,
            data_age_hours=data_age_hours,
            status="healthy" if data_age_hours and data_age_hours < 24 else "stale",
        )

    except Exception as e:
        logger.error(f"Error fetching diagnostics: {e}")
        raise HTTPException(status_code=500, detail=str(e))


@app.post("/jobs/aggregate", response_model=JobResponse)
async def trigger_aggregation(
    background_tasks: BackgroundTasks, run_date: Optional[str] = None
):
    """
    Manually trigger daily aggregation job.
    Called by Cloud Scheduler or for ad-hoc runs.
    """
    try:
        from jobs.daily_aggregator import run_daily_aggregation

        if run_date is None:
            run_date = (datetime.now() - timedelta(days=1)).strftime("%Y-%m-%d")

        # Run in background
        background_tasks.add_task(run_daily_aggregation, run_date)

        return JobResponse(
            job_id=f"aggregate_{run_date}",
            status="started",
            run_date=run_date,
            message=f"Aggregation job started for {run_date}",
        )

    except Exception as e:
        logger.error(f"Error triggering aggregation: {e}")
        raise HTTPException(status_code=500, detail=str(e))


@app.get("/health")
async def health_check():
    """Health check for Cloud Run"""
    try:
        # Verify BigQuery connection
        bq.fetch_one("SELECT 1 as alive")
        return {"status": "healthy", "timestamp": datetime.now().isoformat()}
    except Exception as e:
        logger.error(f"Health check failed: {e}")
        raise HTTPException(status_code=503, detail="Unhealthy")


if __name__ == "__main__":
    import uvicorn

    uvicorn.run(app, host="0.0.0.0", port=8080)
← All docsView source on GitHub →