From 20fbc7df2f164446311cab92f9dbf7ebb3bc7f14 Mon Sep 17 00:00:00 2001 From: Nikita Aksenov Date: Thu, 23 Jul 2026 17:24:29 +0300 Subject: [PATCH] fix: forecasts loading optimized --- src/routers/forecasts.py | 271 ++++++++++++++++++++++++-------- tests/benchmark_forecasts.py | 275 ++++++++++++++++++++++++++++++++ tests/test_forecasts.py | 287 ++++++++++++++++++++++++++++++++++ tests/test_snapshot_routes.py | 28 ++++ 4 files changed, 799 insertions(+), 62 deletions(-) create mode 100644 tests/benchmark_forecasts.py create mode 100644 tests/test_forecasts.py create mode 100644 tests/test_snapshot_routes.py diff --git a/src/routers/forecasts.py b/src/routers/forecasts.py index 6bf8ae3..4638da5 100644 --- a/src/routers/forecasts.py +++ b/src/routers/forecasts.py @@ -4,7 +4,8 @@ from typing import Annotated from fastapi import APIRouter, Depends, HTTPException, Query, status -from sqlalchemy import text, func +from sqlalchemy import func, text +from sqlalchemy.sql.elements import TextClause from sqlalchemy.orm import Session from ..database import get_db @@ -25,12 +26,10 @@ # Helpers # --------------------------------------------------------------------------- -def _predicted_free_count(f: Forecast, db: Session) -> int: - row = db.execute( - text("SELECT predicted_free_count FROM forecasts WHERE forecast_id = :id"), - {"id": f.forecast_id}, - ).one_or_none() - return row[0] if row else (f.capacity - f.predicted_occupied) +def _predicted_free_count(f: Forecast) -> int: + # The database column is generated from the same expression. Computing it + # from the already loaded values avoids one extra SELECT per forecast. + return f.capacity - f.predicted_occupied def _confidence_level(confidence: float) -> ConfidenceLevel | None: @@ -44,7 +43,7 @@ def _confidence_level(confidence: float) -> ConfidenceLevel | None: return ConfidenceLevel.very_low -def _serialize(f: Forecast, db: Session) -> ForecastPointResponse: +def _serialize(f: Forecast) -> ForecastPointResponse: return ForecastPointResponse( forecast_id=f.forecast_id, zone_id=f.zone_id, @@ -56,7 +55,7 @@ def _serialize(f: Forecast, db: Session) -> ForecastPointResponse: predicted_for=f.predicted_for, capacity=f.capacity, predicted_occupied=f.predicted_occupied, - predicted_free_count=_predicted_free_count(f, db), + predicted_free_count=_predicted_free_count(f), probability_free_space=f.probability_free_space, confidence=f.confidence, confidence_level=f.confidence_level.value if f.confidence_level else None, @@ -82,6 +81,181 @@ 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 _map_forecasts_statement( + *, + at: datetime, + 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, +) -> tuple[TextClause, dict[str, object]]: + """ + Build the map query around bounded index lookups. + + For every matching zone the two lateral branches read at most one candidate: + the nearest timestamp before ``at`` and the nearest timestamp at/after it. + The outer lateral query applies the same tie-breakers as the former global + ROW_NUMBER query. The composite index created by migration 000018 supports + these lookups by ``zone_id`` and ``predicted_for``. + """ + params: dict[str, object] = {"at": at} + forecast_filters = ["f.zone_id = z.parking_zone_id"] + zone_filters: list[str] = [] + + optional_forecast_filters = ( + ("camera_id", camera_id, "f.camera_id = :camera_id"), + ("partner_id", partner_id, "f.partner_id = :partner_id"), + ("model_type", model_type, "f.model_type = :model_type"), + ("generated_from", generated_from, "f.generated_at >= :generated_from"), + ("generated_to", generated_to, "f.generated_at <= :generated_to"), + ("from_", from_, "f.predicted_for >= :from_"), + ("to", to, "f.predicted_for <= :to"), + ) + + for name, value, clause in optional_forecast_filters: + if value is not None: + params[name] = value + forecast_filters.append(clause) + + if zone_id is not None: + params["zone_id"] = zone_id + zone_filters.append("z.parking_zone_id = :zone_id") + + if is_active is True: + zone_filters.append("z.is_active IS TRUE") + elif is_active is False: + zone_filters.append("z.is_active IS FALSE") + + if bbox is not None: + min_lon, min_lat, max_lon, max_lat = bbox + params.update( + { + "min_lon": min_lon, + "min_lat": min_lat, + "max_lon": max_lon, + "max_lat": max_lat, + } + ) + zone_filters.extend( + ( + "z.centroid_longitude >= :min_lon", + "z.centroid_longitude <= :max_lon", + "z.centroid_latitude >= :min_lat", + "z.centroid_latitude <= :max_lat", + ) + ) + + forecast_where = "\n AND ".join(forecast_filters) + zone_where = "\n AND ".join(zone_filters) if zone_filters else "TRUE" + + statement = text( + f""" + SELECT + z.parking_zone_id AS zone_id, + f.camera_id, + f.capacity, + f.predicted_occupied, + f.predicted_free_count, + f.probability_free_space, + f.confidence, + CAST(f.confidence_level AS TEXT) AS confidence_level, + f.predicted_for, + f.generated_at, + z.geometry, + z.pay, + CAST(z.zone_type AS TEXT) AS zone_type, + CAST(z.location_type AS TEXT) AS location_type, + z.is_accessible, + z.is_active + FROM parking_zones AS z + CROSS JOIN LATERAL ( + SELECT candidate.forecast_id + FROM ( + ( + SELECT + f.forecast_id, + f.predicted_for, + f.generated_at + FROM forecasts AS f + WHERE {forecast_where} + AND f.predicted_for >= :at + ORDER BY + f.predicted_for ASC, + f.generated_at DESC, + f.forecast_id DESC + LIMIT 1 + ) + UNION ALL + ( + SELECT + f.forecast_id, + f.predicted_for, + f.generated_at + FROM forecasts AS f + WHERE {forecast_where} + AND f.predicted_for < :at + ORDER BY + f.predicted_for DESC, + f.generated_at DESC, + f.forecast_id DESC + LIMIT 1 + ) + ) AS candidate + ORDER BY + ABS(EXTRACT(EPOCH FROM candidate.predicted_for - :at)) ASC, + candidate.generated_at DESC, + candidate.forecast_id DESC + LIMIT 1 + ) AS nearest + JOIN forecasts AS f ON f.forecast_id = nearest.forecast_id + WHERE {zone_where} + ORDER BY f.predicted_for ASC, z.parking_zone_id ASC + """ + ) + + return statement, params + + +def _list_map_forecasts( + *, + db: Session, + at: datetime, + 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: str | None, + is_active: bool | None, +) -> list[ForecastMapItem]: + bbox_bounds = _parse_bbox(bbox) if bbox is not None else None + statement, params = _map_forecasts_statement( + at=at, + 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, + ) + rows = db.execute(statement, params).mappings().all() + + return [ForecastMapItem(**row) for row in rows] + + # --------------------------------------------------------------------------- # GET /forecasts # --------------------------------------------------------------------------- @@ -106,12 +280,29 @@ def list_forecasts( at: datetime | None = None, latest_model_only: bool = False, bbox: str | None = None, + is_active: bool | None = None, view: str = "points", ): if view == "map" and at is None: raise HTTPException(status.HTTP_422_UNPROCESSABLE_ENTITY, detail={"error_description": "Parameter 'at' is required for view=map"}) + if view == "map": + return _list_map_forecasts( + db=db, + at=at, + 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, + is_active=is_active, + ) + query = db.query(Forecast) if zone_id is not None: @@ -130,6 +321,11 @@ def list_forecasts( query = query.filter(Forecast.predicted_for >= from_) if to: query = query.filter(Forecast.predicted_for <= to) + if is_active is not None: + query = query.join( + ParkingZone, + ParkingZone.parking_zone_id == Forecast.zone_id, + ).filter(ParkingZone.is_active == is_active) # Если указан at, независимо от view возвращаем по одному прогнозу на каждую зону. # Логика: @@ -183,27 +379,6 @@ def list_forecasts( & (Forecast.generated_at == latest_sq.c.max_gen), ) - # bbox имеет смысл только для карты, но теперь он не должен быть связан - # с выбором одного прогноза по at. - if view == "map" and bbox: - min_lon, min_lat, max_lon, max_lat = _parse_bbox(bbox) - - zone_ids_in_bbox = [] - zones = db.query(ParkingZone).all() - - for z in zones: - try: - coords = z.geometry["coordinates"][0] - z_lon = sum(c[0] for c in coords) / len(coords) - z_lat = sum(c[1] for c in coords) / len(coords) - - if min_lon <= z_lon <= max_lon and min_lat <= z_lat <= max_lat: - zone_ids_in_bbox.append(z.parking_zone_id) - except Exception: - pass - - query = query.filter(Forecast.zone_id.in_(zone_ids_in_bbox)) - forecasts = query.order_by(Forecast.predicted_for.asc()).all() if view == "series": @@ -211,7 +386,7 @@ def list_forecasts( ForecastSeriesPoint( predicted_for=f.predicted_for, predicted_occupied=f.predicted_occupied, - predicted_free_count=_predicted_free_count(f, db), + predicted_free_count=_predicted_free_count(f), capacity=f.capacity, probability_free_space=f.probability_free_space, confidence=f.confidence, @@ -222,36 +397,8 @@ def list_forecasts( for f in forecasts ] - if view == "map": - result = [] - for f in forecasts: - zone = db.query(ParkingZone).filter( - ParkingZone.parking_zone_id == f.zone_id - ).one_or_none() - if zone is None: - continue - result.append(ForecastMapItem( - zone_id=f.zone_id, - camera_id=f.camera_id, - capacity=f.capacity, - predicted_occupied=f.predicted_occupied, - predicted_free_count=_predicted_free_count(f, db), - probability_free_space=f.probability_free_space, - confidence=f.confidence, - confidence_level=f.confidence_level.value if f.confidence_level else None, - predicted_for=f.predicted_for, - generated_at=f.generated_at, - geometry=zone.geometry, - pay=zone.pay, - zone_type=zone.zone_type.value, - location_type=zone.location_type.value if zone.location_type else None, - is_accessible=zone.is_accessible, - is_active=zone.is_active, - )) - return result - # view=points (default) - return [_serialize(f, db) for f in forecasts] + return [_serialize(f) for f in forecasts] # --------------------------------------------------------------------------- @@ -318,7 +465,7 @@ def get_forecast( current_user: Annotated[User, require("forecasts.view")], db: Annotated[Session, Depends(get_db)], ): - return _serialize(_get_forecast_or_404(db, forecast_id), db) + return _serialize(_get_forecast_or_404(db, forecast_id)) # --------------------------------------------------------------------------- @@ -361,7 +508,7 @@ def update_forecast( db.commit() db.refresh(f) - return _serialize(f, db) + return _serialize(f) # --------------------------------------------------------------------------- diff --git a/tests/benchmark_forecasts.py b/tests/benchmark_forecasts.py new file mode 100644 index 0000000..bbd2c69 --- /dev/null +++ b/tests/benchmark_forecasts.py @@ -0,0 +1,275 @@ +"""PostgreSQL benchmark for the forecast map query. + +The benchmark creates transaction-local tables, fills them with synthetic +forecast history, compares the former global ranking with the optimized query, +and rolls everything back. It does not modify application data. +""" + +from __future__ import annotations + +import os +import sys +import time +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from sqlalchemy import URL, create_engine, text + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +database_url = URL.create( + "postgresql+psycopg2", + username=os.environ["POSTGRES_USER"], + password=os.environ["POSTGRES_PASSWORD"], + host=os.environ.get("FORECAST_BENCHMARK_HOST", "127.0.0.1"), + port=int(os.environ["POSTGRES_PORT"]), + database=os.environ["POSTGRES_DB"], +) +os.environ["DATABASE_URL"] = database_url.render_as_string(hide_password=False) + +from src.routers import forecasts # noqa: E402 + + +AT = datetime(2026, 7, 23, 16, 33, 5, tzinfo=timezone.utc) +BBOX = (34.3471, 61.78765, 34.36887, 61.79272) + + +def _plan_nodes(plan: dict[str, Any]): + yield plan + for child in plan.get("Plans", []): + yield from _plan_nodes(child) + + +def _explain_ms(connection, statement, params: dict[str, object]) -> tuple[float, set[str]]: + explain = text( + "EXPLAIN (ANALYZE, BUFFERS, FORMAT JSON) " + + statement.text + ) + document = connection.execute(explain, params).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) + + with engine.connect() as connection: + transaction = connection.begin() + try: + connection.execute( + text( + """ + CREATE TEMP TABLE parking_zones ( + parking_zone_id INTEGER PRIMARY KEY, + geometry JSONB NOT NULL, + centroid_latitude DOUBLE PRECISION, + centroid_longitude DOUBLE PRECISION, + pay INTEGER NOT NULL, + zone_type TEXT NOT NULL, + location_type TEXT, + is_accessible BOOLEAN, + is_active BOOLEAN NOT NULL + ) ON COMMIT DROP + """ + ) + ) + connection.execute( + text( + """ + CREATE TEMP TABLE forecasts ( + forecast_id BIGSERIAL PRIMARY KEY, + zone_id INTEGER NOT NULL, + camera_id INTEGER, + partner_id INTEGER, + model_type TEXT NOT NULL, + generated_at TIMESTAMPTZ NOT NULL, + predicted_for TIMESTAMPTZ NOT NULL, + capacity INTEGER NOT NULL, + predicted_occupied INTEGER NOT NULL, + predicted_free_count INTEGER GENERATED ALWAYS AS ( + capacity - predicted_occupied + ) STORED, + probability_free_space DOUBLE PRECISION NOT NULL, + confidence DOUBLE PRECISION NOT NULL, + confidence_level TEXT + ) ON COMMIT DROP + """ + ) + ) + connection.execute( + text( + """ + INSERT INTO parking_zones ( + parking_zone_id, + geometry, + centroid_latitude, + centroid_longitude, + pay, + zone_type, + location_type, + is_accessible, + is_active + ) + SELECT + zone_id, + jsonb_build_object( + 'type', 'Polygon', + 'coordinates', jsonb_build_array(jsonb_build_array( + jsonb_build_array(longitude, latitude) + )) + ), + latitude, + longitude, + 0, + 'parallel', + 'street', + TRUE, + zone_id <> 27 + FROM ( + SELECT + zone_id, + CASE + WHEN zone_id <= 27 THEN 61.788 + zone_id * 0.0001 + ELSE 62.5 + zone_id * 0.0001 + END AS latitude, + CASE + WHEN zone_id <= 27 THEN 34.348 + zone_id * 0.0005 + ELSE 36.0 + zone_id * 0.0005 + END AS longitude + FROM generate_series(1, 75) AS zone_id + ) AS zones + """ + ) + ) + + started = time.perf_counter() + connection.execute( + text( + """ + INSERT INTO forecasts ( + zone_id, + camera_id, + partner_id, + model_type, + generated_at, + predicted_for, + capacity, + predicted_occupied, + probability_free_space, + confidence, + confidence_level + ) + SELECT + ((item - 1) % 75) + 1, + 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', + 10, + item % 10, + 0.8, + 0.9, + 'high' + FROM generate_series(1, :row_count) AS item + """ + ), + {"row_count": row_count}, + ) + connection.execute( + text( + """ + CREATE INDEX idx_forecasts_routing_lookup + ON forecasts ( + zone_id, + predicted_for, + generated_at DESC, + forecast_id DESC + ) + """ + ) + ) + connection.execute(text("ANALYZE forecasts")) + connection.execute(text("ANALYZE parking_zones")) + seed_seconds = time.perf_counter() - started + + old_statement = text( + """ + SELECT f.forecast_id + FROM forecasts AS f + JOIN ( + SELECT + ranked_source.forecast_id, + ROW_NUMBER() OVER ( + PARTITION BY ranked_source.zone_id + ORDER BY + ABS(EXTRACT( + EPOCH FROM ranked_source.predicted_for - :at + )) ASC, + ranked_source.generated_at DESC, + ranked_source.forecast_id DESC + ) AS rn + FROM forecasts AS ranked_source + ) AS ranked ON ranked.forecast_id = f.forecast_id + WHERE ranked.rn = 1 + AND f.zone_id <= 27 + ORDER BY f.predicted_for ASC + """ + ) + new_statement, 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, + ) + + old_ms, _ = _explain_ms(connection, old_statement, {"at": AT}) + new_ms, indexes = _explain_ms(connection, new_statement, params) + rows = connection.execute(new_statement, params).mappings().all() + + if len(rows) != 26: + raise AssertionError(f"expected 26 active map zones, got {len(rows)}") + expected_time = datetime( + 2026, + 7, + 23, + 16, + 33, + tzinfo=timezone.utc, + ) + if any(row["predicted_for"] != expected_time for row in rows): + raise AssertionError("optimized query selected a non-nearest forecast") + if "idx_forecasts_routing_lookup" not in indexes: + 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") + + 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"speedup={old_ms / new_ms:.1f}x" + ) + finally: + transaction.rollback() + engine.dispose() + + +if __name__ == "__main__": + main() diff --git a/tests/test_forecasts.py b/tests/test_forecasts.py new file mode 100644 index 0000000..5f313da --- /dev/null +++ b/tests/test_forecasts.py @@ -0,0 +1,287 @@ +from __future__ import annotations + +import os +import unittest +from datetime import datetime, timezone +from typing import Any + +os.environ.setdefault("DATABASE_URL", "sqlite://") + +from src.db_models import Forecast # noqa: E402 +from src.routers import forecasts # noqa: E402 + + +AT = datetime(2026, 7, 23, 16, 33, 5, tzinfo=timezone.utc) + + +def _map_row() -> dict[str, Any]: + return { + "zone_id": 9, + "camera_id": 2, + "capacity": 8, + "predicted_occupied": 3, + "predicted_free_count": 5, + "probability_free_space": 0.95, + "confidence": 0.9, + "confidence_level": "high", + "predicted_for": datetime( + 2026, + 7, + 23, + 16, + 30, + tzinfo=timezone.utc, + ), + "generated_at": datetime( + 2026, + 7, + 23, + 15, + 0, + tzinfo=timezone.utc, + ), + "geometry": { + "type": "Polygon", + "coordinates": [[[34.35, 61.79], [34.36, 61.79]]], + }, + "pay": 0, + "zone_type": "parallel", + "location_type": "street", + "is_accessible": True, + "is_active": True, + } + + +def _statement( + **overrides: Any, +): + arguments = { + "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": None, + "is_active": None, + } + arguments.update(overrides) + return forecasts._map_forecasts_statement(**arguments) + + +class ForecastMapStatementTests(unittest.TestCase): + def test_uses_two_bounded_lateral_index_lookups(self): + statement, params = _statement( + bbox=(34.3471, 61.78765, 34.36887, 61.79272), + is_active=True, + ) + sql = " ".join(statement.text.lower().split()) + + self.assertIn("cross join lateral", sql) + self.assertNotIn("row_number", sql) + self.assertEqual(sql.count("f.zone_id = z.parking_zone_id"), 2) + self.assertIn("f.predicted_for >= :at", sql) + self.assertIn("f.predicted_for < :at", sql) + self.assertIn("z.is_active is true", sql) + self.assertIn("z.centroid_longitude >= :min_lon", sql) + self.assertIn("z.centroid_longitude <= :max_lon", sql) + self.assertIn("z.centroid_latitude >= :min_lat", sql) + self.assertIn("z.centroid_latitude <= :max_lat", sql) + self.assertEqual(params["at"], AT) + self.assertEqual(params["min_lon"], 34.3471) + self.assertEqual(params["max_lat"], 61.79272) + + def test_applies_forecast_filters_to_both_candidates(self): + generated_from = datetime(2026, 7, 22, tzinfo=timezone.utc) + generated_to = datetime(2026, 7, 23, tzinfo=timezone.utc) + from_ = datetime(2026, 7, 23, 15, tzinfo=timezone.utc) + to = datetime(2026, 7, 23, 18, tzinfo=timezone.utc) + + statement, params = _statement( + zone_id=9, + camera_id=3, + partner_id=4, + model_type="baseline", + generated_from=generated_from, + generated_to=generated_to, + from_=from_, + to=to, + is_active=False, + ) + sql = " ".join(statement.text.lower().split()) + + for clause in ( + "f.camera_id = :camera_id", + "f.partner_id = :partner_id", + "f.model_type = :model_type", + "f.generated_at >= :generated_from", + "f.generated_at <= :generated_to", + "f.predicted_for >= :from_", + "f.predicted_for <= :to", + ): + self.assertEqual(sql.count(clause), 2) + + self.assertIn("z.parking_zone_id = :zone_id", sql) + self.assertIn("z.is_active is false", sql) + self.assertEqual( + params, + { + "at": AT, + "zone_id": 9, + "camera_id": 3, + "partner_id": 4, + "model_type": "baseline", + "generated_from": generated_from, + "generated_to": generated_to, + "from_": from_, + "to": to, + }, + ) + + +class _FakeMappings: + def __init__(self, rows: list[dict[str, Any]]): + self._rows = rows + + def all(self) -> list[dict[str, Any]]: + return self._rows + + +class _FakeResult: + def __init__(self, rows: list[dict[str, Any]]): + self._rows = rows + + def mappings(self) -> _FakeMappings: + return _FakeMappings(self._rows) + + +class _FakeSession: + def __init__(self, rows: list[dict[str, Any]]): + self._rows = rows + self.calls: list[tuple[object, dict[str, object]]] = [] + + def execute( + self, + statement: object, + params: dict[str, object], + ) -> _FakeResult: + self.calls.append((statement, params)) + return _FakeResult(self._rows) + + +class ForecastMapExecutionTests(unittest.TestCase): + def test_builds_map_response_with_one_database_round_trip(self): + row = _map_row() + db = _FakeSession([row]) + + result = forecasts._list_map_forecasts( + db=db, # type: ignore[arg-type] + 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="34.3471,61.78765,34.36887,61.79272", + is_active=True, + ) + + self.assertEqual(len(db.calls), 1) + self.assertEqual(len(result), 1) + self.assertEqual(result[0].zone_id, 9) + self.assertEqual(result[0].predicted_free_count, 5) + self.assertEqual(result[0].predicted_for, row["predicted_for"]) + + def test_free_count_does_not_query_database(self): + item = Forecast(capacity=12, predicted_occupied=7) + + self.assertEqual(forecasts._predicted_free_count(item), 5) + + +class ForecastEndpointTests(unittest.TestCase): + def setUp(self): + self.db = _FakeSession([_map_row()]) + + def test_exact_map_request_uses_optimized_query(self): + request_at = datetime( + 2026, + 7, + 23, + 16, + 33, + 5, + 669000, + tzinfo=timezone.utc, + ) + result = forecasts.list_forecasts( + db=self.db, # type: ignore[arg-type] + zone_id=None, + camera_id=None, + partner_id=None, + model_type=None, + generated_from=None, + generated_to=None, + from_=None, + to=None, + at=request_at, + latest_model_only=False, + bbox="34.3471,61.78765,34.36887,61.79272", + is_active=True, + view="map", + ) + + self.assertEqual(result[0].zone_id, 9) + self.assertEqual(len(self.db.calls), 1) + + statement, params = self.db.calls[0] + sql = " ".join(statement.text.lower().split()) + self.assertIn("cross join lateral", sql) + self.assertIn("z.is_active is true", sql) + self.assertEqual(params["min_lon"], 34.3471) + self.assertEqual(params["max_lat"], 61.79272) + self.assertEqual(params["at"], request_at) + + def test_map_request_without_at_does_not_query_database(self): + with self.assertRaises(forecasts.HTTPException) as raised: + forecasts.list_forecasts( + db=self.db, # type: ignore[arg-type] + zone_id=None, + camera_id=None, + partner_id=None, + model_type=None, + generated_from=None, + generated_to=None, + from_=None, + to=None, + at=None, + latest_model_only=False, + bbox="34.3471,61.78765,34.36887,61.79272", + is_active=True, + view="map", + ) + + self.assertEqual(raised.exception.status_code, 422) + self.assertEqual(len(self.db.calls), 0) + + def test_fastapi_route_exposes_is_active_filter(self): + route = next( + route + for route in forecasts.router.routes + if route.path == "/forecasts" + ) + parameter_names = { + parameter.name + for parameter in route.dependant.query_params + } + + self.assertIn("is_active", parameter_names) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_snapshot_routes.py b/tests/test_snapshot_routes.py new file mode 100644 index 0000000..65164c3 --- /dev/null +++ b/tests/test_snapshot_routes.py @@ -0,0 +1,28 @@ +from __future__ import annotations + +import os +import unittest + +os.environ.setdefault("DATABASE_URL", "sqlite://") + +from src.routers import analytics, cameras # noqa: E402 + + +class SnapshotRouteContractTests(unittest.TestCase): + def test_snapshot_is_exposed_by_cameras_only(self): + analytics_paths = {route.path for route in analytics.router.routes} + camera_paths = {route.path for route in cameras.router.routes} + + self.assertNotIn( + "/admin/analytics/detections/{detection_run_id}/snapshot", + analytics_paths, + ) + self.assertNotIn( + "/admin/analytics/detections/{detection_run_id}/labels", + analytics_paths, + ) + self.assertIn("/cameras/{camera_id}/snapshot", camera_paths) + + +if __name__ == "__main__": + unittest.main()