Spaces:
Sleeping
Sleeping
File size: 8,255 Bytes
93e2220 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 | """Database abstraction layer for TransitPulse.
Supports local DuckDB (direct Parquet querying) and BigQuery in cloud mode.
"""
from __future__ import annotations
import os
from pathlib import Path
from typing import Any
import duckdb
import pandas as pd
from config import CFG
class TransitPulseDB:
"""Manages queries to either local DuckDB or BigQuery."""
def __init__(self, mode: str = "local"):
self.mode = mode
self.duck_conn = None
if self.mode == "local":
# Initialize DuckDB local in-memory or file database
self.duck_conn = duckdb.connect(database=":memory:")
self.init_local_tables()
else:
# Cloud mode: initialize BigQuery client
from google.cloud import bigquery
self.bq_client = bigquery.Client(project=CFG.gcp_project)
print(f"BigQuery serving layer initialized in project: {CFG.gcp_project}")
def init_local_tables(self) -> None:
"""Create views in DuckDB over the output parquet files."""
scores_path = CFG.output_dir / "route_scores.parquet"
segment_path = CFG.output_dir / "daily_segment_metrics.parquet"
anomaly_path = CFG.output_dir / "anomaly_events.parquet"
if not scores_path.exists():
# If parquet doesn't exist, create empty tables or register mock schemas
print("WARNING: Parquet outputs not found. DuckDB views might fail until pipeline runs.")
# We can create dummy files so startup doesn't crash
return
# Register views pointing to parquet files directly
# This keeps the DuckDB database dynamically in sync with parquet outputs
self.duck_conn.execute(f"CREATE OR REPLACE VIEW route_scores AS SELECT * FROM read_parquet('{scores_path}')")
self.duck_conn.execute(f"CREATE OR REPLACE VIEW daily_segment_metrics AS SELECT * FROM read_parquet('{segment_path}')")
self.duck_conn.execute(f"CREATE OR REPLACE VIEW anomaly_events AS SELECT * FROM read_parquet('{anomaly_path}')")
print("DuckDB views created successfully.")
def run_query(self, query: str, params: dict[str, Any] | None = None) -> pd.DataFrame:
"""Helper to run SQL query on DuckDB or BigQuery."""
if self.mode == "local":
if self.duck_conn is None:
# Late-bind if connection was missed
self.duck_conn = duckdb.connect(database=":memory:")
self.init_local_tables()
# DuckDB supports named parameters with $name or standard python interpolation
# Let's execute and return dataframe
if params:
df = self.duck_conn.execute(query, params).df()
else:
df = self.duck_conn.execute(query).df()
else:
# Cloud mode (BigQuery)
# Standard BigQuery query
job_config = None
if params:
from google.cloud import bigquery
query_params = [
bigquery.ScalarQueryParameter(name, "STRING" if isinstance(val, str) else "FLOAT", val)
for name, val in params.items()
]
job_config = bigquery.QueryJobConfig(query_parameters=query_params)
# Map simple table names to dataset qualified names
bq_query = query.replace("route_scores", f"`{CFG.gcp_project}.{CFG.bq_dataset}.route_scores`")
bq_query = bq_query.replace("daily_segment_metrics", f"`{CFG.gcp_project}.{CFG.bq_dataset}.daily_segment_metrics`")
bq_query = bq_query.replace("anomaly_events", f"`{CFG.gcp_project}.{CFG.bq_dataset}.anomaly_events`")
query_job = self.bq_client.query(bq_query, job_config=job_config)
df = query_job.to_dataframe()
import numpy as np
return df.replace({np.nan: None})
# ββ Database Endpoint Operations βββββββββββββββββββββββββββββββββββ
def get_route_scores(self, sort_by: str = "reliability_score", ascending: bool = True) -> list[dict[str, Any]]:
"""Returns sorted route reliability scores."""
direction = "ASC" if ascending else "DESC"
# Validate column name to prevent SQL injection
valid_cols = ["route_id", "reliability_score", "mean_headway", "bunching_count", "gap_count", "mean_dwell_sec", "wow_trend", "date", "route_base_boardings"]
if sort_by not in valid_cols:
sort_by = "reliability_score"
# Get latest date available in the route_scores table
try:
latest_date_df = self.run_query("SELECT MAX(date) as max_date FROM route_scores")
latest_date = latest_date_df.iloc[0]["max_date"]
except Exception:
latest_date = None
if latest_date:
query = f"SELECT * FROM route_scores WHERE date = $latest_date ORDER BY {sort_by} {direction}"
params = {"latest_date": str(latest_date)}
else:
query = f"SELECT * FROM route_scores ORDER BY {sort_by} {direction}"
params = {}
df = self.run_query(query, params)
return df.to_dict(orient="records")
def get_route_timeline(self, route_id: str) -> list[dict[str, Any]]:
"""Returns historical scores timeline for a specific route."""
query = "SELECT date, reliability_score, mean_headway, mean_dwell_sec, bunching_count, gap_count FROM route_scores WHERE route_id = $route_id ORDER BY date ASC"
df = self.run_query(query, {"route_id": route_id})
return df.to_dict(orient="records")
def get_worst_segments(self, limit: int = 10) -> list[dict[str, Any]]:
"""Returns the bottom N route segments sorted by reliability score."""
try:
latest_date_df = self.run_query("SELECT MAX(date) as max_date FROM daily_segment_metrics")
latest_date = latest_date_df.iloc[0]["max_date"]
except Exception:
latest_date = None
if latest_date:
query = f"SELECT route_id, stop_id, reliability_score, total_trips, bunching_rate, gap_rate, mean_dwell_sec, route_base_boardings FROM daily_segment_metrics WHERE date = $latest_date ORDER BY reliability_score ASC LIMIT {limit}"
params = {"latest_date": str(latest_date)}
else:
query = f"SELECT route_id, stop_id, reliability_score, total_trips, bunching_rate, gap_rate, mean_dwell_sec, route_base_boardings FROM daily_segment_metrics ORDER BY reliability_score ASC LIMIT {limit}"
params = {}
df = self.run_query(query, params)
return df.to_dict(orient="records")
def get_anomalies(self, route_id: str | None = None, anomaly_type: str | None = None, limit: int = 100) -> list[dict[str, Any]]:
"""Returns recent anomaly events, optionally filtered."""
where_clauses = []
params = {}
if route_id:
where_clauses.append("route_id = $route_id")
params["route_id"] = route_id
if anomaly_type:
where_clauses.append("anomaly_type = $anomaly_type")
params["anomaly_type"] = anomaly_type
where_str = ""
if where_clauses:
where_str = "WHERE " + " AND ".join(where_clauses)
query = f"SELECT timestamp_str, vehicle_id, route_id, stop_id, stop_sequence, headway_sec, scheduled_headway_sec, anomaly_type FROM anomaly_events {where_str} ORDER BY timestamp_str DESC LIMIT {limit}"
df = self.run_query(query, params)
return df.to_dict(orient="records")
def compare_periods(self, route_id: str, date1: str, date2: str) -> dict[str, Any]:
"""Compares route performance metrics between two dates."""
query = "SELECT reliability_score, mean_headway, bunching_count, gap_count FROM route_scores WHERE route_id = $route_id AND date IN ($date1, $date2)"
df = self.run_query(query, {"route_id": route_id, "date1": date1, "date2": date2})
return df.to_dict(orient="records")
# Database Singleton
DB = TransitPulseDB(mode=CFG.mode)
|