From 730efec7ecc38a065ed10cb8cc9f91eed33b0244 Mon Sep 17 00:00:00 2001 From: Nikita Aksenov Date: Thu, 23 Jul 2026 18:35:34 +0300 Subject: [PATCH] fix: forecasts endpoint with latest_model_only=true --- src/routers/forecasts.py | 226 ++++++++++++++++++++++++++++------- tests/benchmark_forecasts.py | 134 +++++++++++++++++++-- tests/test_forecasts.py | 77 ++++++++++++ 3 files changed, 389 insertions(+), 48 deletions(-) diff --git a/src/routers/forecasts.py b/src/routers/forecasts.py index 4638da5..704ec77 100644 --- a/src/routers/forecasts.py +++ b/src/routers/forecasts.py @@ -4,9 +4,9 @@ from typing import Annotated from fastapi import APIRouter, Depends, HTTPException, Query, status -from sqlalchemy import func, text +from sqlalchemy import func, select, text, true from sqlalchemy.sql.elements import TextClause -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, aliased from ..database import get_db from ..db_models import ConfidenceLevel, Forecast, ParkingZone, User @@ -81,6 +81,121 @@ def _parse_bbox(bbox: str) -> tuple[float, float, float, float]: detail={"error_description": "bbox must be min_lon,min_lat,max_lon,max_lat"}) +def _apply_forecast_filters( + query, + forecast, + *, + camera_id: int | None, + partner_id: int | None, + model_type: str | None, + generated_from: datetime | None, + generated_to: datetime | None, + from_: datetime | None, + to: datetime | None, +): + if camera_id is not None: + query = query.filter(forecast.camera_id == camera_id) + if partner_id is not None: + query = query.filter(forecast.partner_id == partner_id) + if model_type is not None: + query = query.filter(forecast.model_type == model_type) + if generated_from is not None: + query = query.filter(forecast.generated_at >= generated_from) + if generated_to is not None: + query = query.filter(forecast.generated_at <= generated_to) + if from_ is not None: + query = query.filter(forecast.predicted_for >= from_) + if to is not None: + query = query.filter(forecast.predicted_for <= to) + return query + + +def _latest_generation_query( + *, + db: Session, + zone_id: int | None, + camera_id: int | None, + partner_id: int | None, + model_type: str | None, + generated_from: datetime | None, + generated_to: datetime | None, + from_: datetime | None, + to: datetime | None, + bbox: tuple[float, float, float, float] | None, + is_active: bool | None, +): + """ + Return only points from the latest matching generation in each zone. + + Starting from the much smaller parking_zones table lets PostgreSQL perform + one bounded index lookup per zone instead of grouping or sorting the entire + forecasts table. The uq_forecast_point index starts with + (zone_id, generated_at), so ORDER BY generated_at DESC LIMIT 1 is an + index-backed lookup. + """ + zones = select(ParkingZone.parking_zone_id.label("zone_id")) + + if zone_id is not None: + zones = zones.filter(ParkingZone.parking_zone_id == zone_id) + if is_active is not None: + zones = zones.filter(ParkingZone.is_active == is_active) + if bbox is not None: + min_lon, min_lat, max_lon, max_lat = bbox + zones = zones.filter( + ParkingZone.centroid_longitude >= min_lon, + ParkingZone.centroid_longitude <= max_lon, + ParkingZone.centroid_latitude >= min_lat, + ParkingZone.centroid_latitude <= max_lat, + ) + + matching_zones = zones.subquery("matching_zones") + generation_candidate = aliased(Forecast, name="generation_candidate") + latest_generation = select( + generation_candidate.generated_at.label("generated_at") + ).filter( + generation_candidate.zone_id == matching_zones.c.zone_id + ) + latest_generation = _apply_forecast_filters( + latest_generation, + generation_candidate, + camera_id=camera_id, + partner_id=partner_id, + model_type=model_type, + generated_from=generated_from, + generated_to=generated_to, + from_=from_, + to=to, + ) + latest_generation = ( + latest_generation + .order_by(generation_candidate.generated_at.desc()) + .limit(1) + .lateral("latest_generation") + ) + + query = ( + db.query(Forecast) + .select_from(matching_zones) + .join(latest_generation, true()) + .join( + Forecast, + (Forecast.zone_id == matching_zones.c.zone_id) + & (Forecast.generated_at == latest_generation.c.generated_at), + ) + ) + return _apply_forecast_filters( + query, + Forecast, + camera_id=camera_id, + partner_id=partner_id, + model_type=model_type, + generated_from=generated_from, + generated_to=generated_to, + from_=from_, + to=to, + ) + + def _map_forecasts_statement( *, at: datetime, @@ -94,6 +209,7 @@ def _map_forecasts_statement( to: datetime | None, bbox: tuple[float, float, float, float] | None, is_active: bool | None, + latest_model_only: bool, ) -> tuple[TextClause, dict[str, object]]: """ Build the map query around bounded index lookups. @@ -151,6 +267,26 @@ def _map_forecasts_statement( ) ) + latest_generation_join = "" + if latest_model_only: + latest_filters = [ + clause.replace("f.", "generation_candidate.") + for clause in forecast_filters + ] + latest_where = "\n AND ".join(latest_filters) + latest_generation_join = f""" + CROSS JOIN LATERAL ( + SELECT generation_candidate.generated_at + FROM forecasts AS generation_candidate + WHERE {latest_where} + ORDER BY generation_candidate.generated_at DESC + LIMIT 1 + ) AS latest_generation + """ + forecast_filters.append( + "f.generated_at = latest_generation.generated_at" + ) + forecast_where = "\n AND ".join(forecast_filters) zone_where = "\n AND ".join(zone_filters) if zone_filters else "TRUE" @@ -174,6 +310,7 @@ def _map_forecasts_statement( z.is_accessible, z.is_active FROM parking_zones AS z + {latest_generation_join} CROSS JOIN LATERAL ( SELECT candidate.forecast_id FROM ( @@ -236,6 +373,7 @@ def _list_map_forecasts( to: datetime | None, bbox: str | None, is_active: bool | None, + latest_model_only: bool, ) -> list[ForecastMapItem]: bbox_bounds = _parse_bbox(bbox) if bbox is not None else None statement, params = _map_forecasts_statement( @@ -250,6 +388,7 @@ def _list_map_forecasts( to=to, bbox=bbox_bounds, is_active=is_active, + latest_model_only=latest_model_only, ) rows = db.execute(statement, params).mappings().all() @@ -301,31 +440,56 @@ def list_forecasts( to=to, bbox=bbox, is_active=is_active, + latest_model_only=latest_model_only, ) - query = db.query(Forecast) + bbox_bounds = _parse_bbox(bbox) if bbox is not None else None - if zone_id is not None: - query = query.filter(Forecast.zone_id == zone_id) - if camera_id is not None: - query = query.filter(Forecast.camera_id == camera_id) - if partner_id is not None: - query = query.filter(Forecast.partner_id == partner_id) - if model_type is not None: - query = query.filter(Forecast.model_type == model_type) - if generated_from: - query = query.filter(Forecast.generated_at >= generated_from) - if generated_to: - query = query.filter(Forecast.generated_at <= generated_to) - if from_: - query = query.filter(Forecast.predicted_for >= from_) - if to: - query = query.filter(Forecast.predicted_for <= to) - if is_active is not None: + if latest_model_only: + query = _latest_generation_query( + db=db, + zone_id=zone_id, + camera_id=camera_id, + partner_id=partner_id, + model_type=model_type, + generated_from=generated_from, + generated_to=generated_to, + from_=from_, + to=to, + bbox=bbox_bounds, + is_active=is_active, + ) + else: + query = db.query(Forecast) + if zone_id is not None: + query = query.filter(Forecast.zone_id == zone_id) + query = _apply_forecast_filters( + query, + Forecast, + camera_id=camera_id, + partner_id=partner_id, + model_type=model_type, + generated_from=generated_from, + generated_to=generated_to, + from_=from_, + to=to, + ) + + if not latest_model_only and (is_active is not None or bbox_bounds is not None): query = query.join( ParkingZone, ParkingZone.parking_zone_id == Forecast.zone_id, - ).filter(ParkingZone.is_active == is_active) + ) + if is_active is not None: + query = query.filter(ParkingZone.is_active == is_active) + if bbox_bounds is not None: + min_lon, min_lat, max_lon, max_lat = bbox_bounds + query = query.filter( + ParkingZone.centroid_longitude >= min_lon, + ParkingZone.centroid_longitude <= max_lon, + ParkingZone.centroid_latitude >= min_lat, + ParkingZone.centroid_latitude <= max_lat, + ) # Если указан at, независимо от view возвращаем по одному прогнозу на каждую зону. # Логика: @@ -359,26 +523,6 @@ def list_forecasts( .filter(ranked_sq.c.rn == 1) ) - # Если at не указан, но requested latest_model_only, - # оставляем старую логику: последняя генерация по каждой зоне + predicted_for. - elif latest_model_only: - latest_sq = ( - query.with_entities( - Forecast.zone_id.label("zone_id"), - Forecast.predicted_for.label("predicted_for"), - func.max(Forecast.generated_at).label("max_gen"), - ) - .group_by(Forecast.zone_id, Forecast.predicted_for) - .subquery() - ) - - query = query.join( - latest_sq, - (Forecast.zone_id == latest_sq.c.zone_id) - & (Forecast.predicted_for == latest_sq.c.predicted_for) - & (Forecast.generated_at == latest_sq.c.max_gen), - ) - forecasts = query.order_by(Forecast.predicted_for.asc()).all() if view == "series": diff --git a/tests/benchmark_forecasts.py b/tests/benchmark_forecasts.py index bbd2c69..493a706 100644 --- a/tests/benchmark_forecasts.py +++ b/tests/benchmark_forecasts.py @@ -15,6 +15,7 @@ from typing import Any from sqlalchemy import URL, create_engine, text +from sqlalchemy.orm import Session sys.path.insert(0, str(Path(__file__).resolve().parents[1])) @@ -56,6 +57,23 @@ def _explain_ms(connection, statement, params: dict[str, object]) -> tuple[float return float(root["Execution Time"]), indexes +def _explain_select_ms(connection, statement) -> tuple[float, set[str]]: + compiled = statement.compile( + dialect=connection.dialect, + compile_kwargs={"literal_binds": True}, + ) + document = connection.execute( + text("EXPLAIN (ANALYZE, BUFFERS, FORMAT JSON) " + str(compiled)) + ).scalar_one() + root = document[0] + indexes = { + node["Index Name"] + for node in _plan_nodes(root["Plan"]) + if "Index Name" in node + } + return float(root["Execution Time"]), indexes + + def main() -> None: row_count = int(os.environ.get("FORECAST_BENCHMARK_ROWS", "1000000")) engine = create_engine(database_url) @@ -89,6 +107,7 @@ def main() -> None: camera_id INTEGER, partner_id INTEGER, model_type TEXT NOT NULL, + model_version TEXT, generated_at TIMESTAMPTZ NOT NULL, predicted_for TIMESTAMPTZ NOT NULL, capacity INTEGER NOT NULL, @@ -98,7 +117,9 @@ def main() -> None: ) STORED, probability_free_space DOUBLE PRECISION NOT NULL, confidence DOUBLE PRECISION NOT NULL, - confidence_level TEXT + confidence_level TEXT, + metadata JSONB, + created_by_user_id INTEGER ) ON COMMIT DROP """ ) @@ -167,23 +188,38 @@ def main() -> None: confidence_level ) SELECT - ((item - 1) % 75) + 1, + zone_id, 1, 1, 'baseline', - TIMESTAMPTZ '2026-07-01 00:00:00+00', - TIMESTAMPTZ '2026-07-15 00:00:00+00' - + ((item - 1) / 75) * INTERVAL '1 minute', + TIMESTAMPTZ '2026-07-01 00:00:00+00' + + (sequence_number / 288) * INTERVAL '6 hours', + TIMESTAMPTZ '2026-07-23 00:00:00+00' + + (sequence_number % 288) * INTERVAL '5 minutes', 10, item % 10, 0.8, 0.9, 'high' - FROM generate_series(1, :row_count) AS item + FROM ( + SELECT + item, + ((item - 1) % 75) + 1 AS zone_id, + (item - 1) / 75 AS sequence_number + FROM generate_series(1, :row_count) AS item + ) AS generated_points """ ), {"row_count": row_count}, ) + connection.execute( + text( + """ + CREATE UNIQUE INDEX uq_forecast_point + ON forecasts (zone_id, generated_at, predicted_for) + """ + ) + ) connection.execute( text( """ @@ -236,20 +272,67 @@ def main() -> None: to=None, bbox=BBOX, is_active=True, + latest_model_only=False, + ) + latest_statement, latest_params = forecasts._map_forecasts_statement( + at=AT, + zone_id=None, + camera_id=None, + partner_id=None, + model_type=None, + generated_from=None, + generated_to=None, + from_=None, + to=None, + bbox=BBOX, + is_active=True, + latest_model_only=True, ) old_ms, _ = _explain_ms(connection, old_statement, {"at": AT}) new_ms, indexes = _explain_ms(connection, new_statement, params) + latest_ms, latest_indexes = _explain_ms( + connection, + latest_statement, + latest_params, + ) + latest_series_query = forecasts._latest_generation_query( + db=Session(bind=connection), + zone_id=None, + camera_id=None, + partner_id=None, + model_type=None, + generated_from=None, + generated_to=None, + from_=None, + to=None, + bbox=BBOX, + is_active=True, + ) + latest_series_ms, latest_series_indexes = _explain_select_ms( + connection, + latest_series_query.statement, + ) rows = connection.execute(new_statement, params).mappings().all() + latest_rows = ( + connection.execute(latest_statement, latest_params) + .mappings() + .all() + ) + latest_series_rows = latest_series_query.all() if len(rows) != 26: raise AssertionError(f"expected 26 active map zones, got {len(rows)}") + if len(latest_rows) != 26: + raise AssertionError( + f"expected 26 latest-model map zones, got {len(latest_rows)}" + ) expected_time = datetime( 2026, 7, 23, 16, - 33, + 35, tzinfo=timezone.utc, ) if any(row["predicted_for"] != expected_time for row in rows): @@ -258,12 +341,49 @@ def main() -> None: raise AssertionError(f"forecast lookup index was not used: {sorted(indexes)}") if new_ms >= 2_000: raise AssertionError(f"optimized query exceeded 2 seconds: {new_ms:.3f} ms") + if "uq_forecast_point" not in latest_indexes: + raise AssertionError( + "latest-generation index was not used: " + f"{sorted(latest_indexes)}" + ) + if latest_ms >= 2_000: + raise AssertionError( + "latest-model query exceeded 2 seconds: " + f"{latest_ms:.3f} ms" + ) + generations_by_zone: dict[int, set[datetime]] = {} + for forecast in latest_series_rows: + generations_by_zone.setdefault(forecast.zone_id, set()).add( + forecast.generated_at + ) + if len(generations_by_zone) != 26: + raise AssertionError( + "expected latest series for 26 active zones, got " + f"{len(generations_by_zone)}" + ) + if any(len(generations) != 1 for generations in generations_by_zone.values()): + raise AssertionError("latest series contains multiple generations") + if len(latest_series_rows) > 26 * 288: + raise AssertionError("latest series returned more than one horizon per zone") + if "uq_forecast_point" not in latest_series_indexes: + raise AssertionError( + "latest-series index was not used: " + f"{sorted(latest_series_indexes)}" + ) + if latest_series_ms >= 2_000: + raise AssertionError( + "latest-series query exceeded 2 seconds: " + f"{latest_series_ms:.3f} ms" + ) print( "forecast_map " f"rows={row_count} result_zones={len(rows)} " f"seed_seconds={seed_seconds:.3f} " f"old_ms={old_ms:.3f} new_ms={new_ms:.3f} " + f"latest_ms={latest_ms:.3f} " + f"latest_series_rows={len(latest_series_rows)} " + f"latest_series_ms={latest_series_ms:.3f} " f"speedup={old_ms / new_ms:.1f}x" ) finally: diff --git a/tests/test_forecasts.py b/tests/test_forecasts.py index 5f313da..d7d23b7 100644 --- a/tests/test_forecasts.py +++ b/tests/test_forecasts.py @@ -7,6 +7,10 @@ os.environ.setdefault("DATABASE_URL", "sqlite://") +from sqlalchemy import create_mock_engine # noqa: E402 +from sqlalchemy.dialects import postgresql # noqa: E402 +from sqlalchemy.orm import Session # noqa: E402 + from src.db_models import Forecast # noqa: E402 from src.routers import forecasts # noqa: E402 @@ -67,6 +71,7 @@ def _statement( "to": None, "bbox": None, "is_active": None, + "latest_model_only": False, } arguments.update(overrides) return forecasts._map_forecasts_statement(**arguments) @@ -77,6 +82,7 @@ def test_uses_two_bounded_lateral_index_lookups(self): statement, params = _statement( bbox=(34.3471, 61.78765, 34.36887, 61.79272), is_active=True, + latest_model_only=False, ) sql = " ".join(statement.text.lower().split()) @@ -141,6 +147,76 @@ def test_applies_forecast_filters_to_both_candidates(self): }, ) + def test_latest_model_limits_candidates_to_latest_zone_generation(self): + statement, _ = _statement( + latest_model_only=True, + model_type="baseline", + ) + sql = " ".join(statement.text.lower().split()) + + self.assertIn( + "select generation_candidate.generated_at", + sql, + ) + self.assertIn( + "order by generation_candidate.generated_at desc limit 1", + sql, + ) + self.assertEqual( + sql.count("f.generated_at = latest_generation.generated_at"), + 2, + ) + self.assertNotIn("group by", sql) + + +class ForecastLatestGenerationQueryTests(unittest.TestCase): + def setUp(self): + engine = create_mock_engine( + "postgresql+psycopg2://", + lambda *args, **kwargs: None, + ) + self.db = Session(bind=engine) + + def tearDown(self): + self.db.close() + + def test_uses_one_bounded_lateral_lookup_per_matching_zone(self): + query = forecasts._latest_generation_query( + db=self.db, + zone_id=9, + camera_id=3, + partner_id=4, + model_type="baseline", + generated_from=None, + generated_to=None, + from_=None, + to=None, + bbox=(34.3471, 61.78765, 34.36887, 61.79272), + is_active=True, + ) + sql = " ".join( + str( + query.statement.compile( + dialect=postgresql.dialect(), + compile_kwargs={"literal_binds": True}, + ) + ).lower().split() + ) + + self.assertIn("join lateral", sql) + self.assertIn( + "order by generation_candidate.generated_at desc limit 1", + sql, + ) + self.assertIn( + "forecasts.generated_at = latest_generation.generated_at", + sql, + ) + self.assertIn("parking_zones.parking_zone_id = 9", sql) + self.assertIn("parking_zones.centroid_longitude >=", sql) + self.assertNotIn("group by", sql) + self.assertNotIn("row_number", sql) + class _FakeMappings: def __init__(self, rows: list[dict[str, Any]]): @@ -190,6 +266,7 @@ def test_builds_map_response_with_one_database_round_trip(self): to=None, bbox="34.3471,61.78765,34.36887,61.79272", is_active=True, + latest_model_only=False, ) self.assertEqual(len(db.calls), 1)