diff --git a/app/api/routes/api/meters.py b/app/api/routes/api/meters.py index 146c5c8..80069f7 100644 --- a/app/api/routes/api/meters.py +++ b/app/api/routes/api/meters.py @@ -151,7 +151,7 @@ def _trigger_recompute(db: Session, start: datetime, label: str) -> int: # started_at is in the future — nothing to recompute. logger.info("%s: started_at (%s) is in the future, skipping recompute.", label, start) return 0 - n = recompute_range(db, start, end) + n = recompute_range(db, start, end, commit=False) logger.info( "%s: recomputed %d period(s) in window [%s, %s).", label, @@ -352,23 +352,27 @@ def patch_energy_meter( note=body.note, started_at=new_started_at_utc, ) + + # Retroactive recompute if started_at was changed. + if new_started_at_utc is not None and old_started_at is not None: + # Normalise old_started_at to UTC-aware for comparison. + if old_started_at.tzinfo is None: + old_started_at = old_started_at.replace(tzinfo=UTC) + # Window = [min(old, new), now) — covers all periods whose attribution + # may have changed due to the boundary shift in either direction. + window_start = min(old_started_at, new_started_at_utc) + _trigger_recompute(db, window_start, f"PATCH /api/energy/meters/{meter_id}") + + db.commit() except MeterIntervalError as exc: + db.rollback() raise HTTPException( status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc), ) - - # Retroactive recompute if started_at was changed. - if new_started_at_utc is not None and old_started_at is not None: - # Normalise old_started_at to UTC-aware for comparison. - if old_started_at.tzinfo is None: - old_started_at = old_started_at.replace(tzinfo=UTC) - # Window = [min(old, new), now) — covers all periods whose attribution - # may have changed due to the boundary shift in either direction. - window_start = min(old_started_at, new_started_at_utc) - _trigger_recompute(db, window_start, f"PATCH /api/energy/meters/{meter_id}") - - db.commit() + except Exception: + db.rollback() + raise db.refresh(meter) # Trigger HA discovery re-publish so label renames on the active meter diff --git a/app/services/energy_cost.py b/app/services/energy_cost.py index 6e02173..fd1cfdd 100644 --- a/app/services/energy_cost.py +++ b/app/services/energy_cost.py @@ -665,7 +665,9 @@ def compute_closed_periods(session: Session) -> int: # --------------------------------------------------------------------------- -def recompute_range(session: Session, start: datetime, end: datetime) -> int: +def recompute_range( + session: Session, start: datetime, end: datetime, *, commit: bool = True +) -> int: """Recompute (overwrite) all 15-minute periods in ``[start, end)``. This is the *explicit opt-in* path for recovering from: @@ -686,8 +688,11 @@ def recompute_range(session: Session, start: datetime, end: datetime) -> int: Parameters ---------- session: - Active SQLAlchemy session. The function commits after all periods - have been processed. + Active SQLAlchemy session. + commit: + When true (the default), commit after all periods have been processed. + Callers composing this recompute with other writes may pass false and + own the surrounding transaction themselves. start: Inclusive start datetime (floored to the nearest quarter-hour internally). end: @@ -724,7 +729,8 @@ def recompute_range(session: Session, start: datetime, end: datetime) -> int: ) t0 += timedelta(minutes=_PERIOD_MINUTES) - session.commit() + if commit: + session.commit() logger.info( "recompute_range(%s, %s): wrote %d period(s).", start.isoformat(), diff --git a/tests/test_api_meters.py b/tests/test_api_meters.py index d5fd1e3..0a2e7fd 100644 --- a/tests/test_api_meters.py +++ b/tests/test_api_meters.py @@ -36,10 +36,10 @@ from unittest.mock import patch import pytest from fastapi.testclient import TestClient -from sqlalchemy import create_engine, select +from sqlalchemy import create_engine, event, select from sqlalchemy.orm import Session -from app.models.energy import Meter +from app.models.energy import EnergyCostPeriod, Meter from app.models.meter_source import MeterSource, MeterSourceBinding, MeterSourceChannel # --------------------------------------------------------------------------- @@ -443,6 +443,54 @@ def test_declare_meter_recompute_failure_rolls_back_handoff(meters_client): assert session.execute(select(MeterSourceBinding)).scalar_one().ended_at is None +def test_declare_meter_final_commit_failure_rolls_back_handoff_and_recompute(meters_client): + """A real recompute remains uncommitted until the route's final commit succeeds.""" + client, engine = meters_client + _login(client) + boundary = datetime.now(UTC) - timedelta(minutes=45) + boundary = boundary.replace(minute=boundary.minute - boundary.minute % 15, second=0, microsecond=0) + old_start = boundary - timedelta(days=1) + + with Session(engine) as session: + old_meter = Meter( + label="Old meter", + commodity="electricity", + started_at=old_start, + reason="initial", + created_at=old_start, + ) + session.add(old_meter) + session.commit() + old_id = old_meter.id + channel_uuid = _add_bound_channel(engine, meter_id=old_id, started_at=old_start) + + def fail_final_commit(_session: Session) -> None: + raise RuntimeError("final commit failed") + + event.listen(Session, "before_commit", fail_final_commit) + try: + with pytest.raises(RuntimeError, match="final commit failed"): + client.post( + "/api/energy/meters", + json=_declare_payload( + label="New meter", + started_at=boundary.isoformat(), + reason="meter_swap", + source_channel_uuid=channel_uuid, + ), + headers={"X-CSRF-Token": _CSRF}, + ) + finally: + event.remove(Session, "before_commit", fail_final_commit) + + with Session(engine) as session: + assert session.execute(select(Meter).where(Meter.label == "New meter")).scalar_one_or_none() is None + assert session.get(Meter, old_id).ended_at is None + binding = session.execute(select(MeterSourceBinding)).scalar_one() + assert binding.ended_at is None + assert session.execute(select(EnergyCostPeriod)).scalars().all() == [] + + def test_declare_meter_overlap_returns_422(meters_client): """Declaring a meter with started_at before active meter's started_at → 422.""" client, _ = meters_client @@ -546,6 +594,7 @@ def test_declare_meter_retroactive_triggers_recompute(meters_client): assert resp.status_code == 201 # recompute_range should have been called with start == t_past assert mock_recompute.called + assert mock_recompute.call_args.kwargs["commit"] is False call_args = mock_recompute.call_args recompute_start = call_args[0][1] # positional arg index 1 (session is 0) # Normalise for comparison @@ -685,6 +734,7 @@ def test_patch_meter_started_at_retroactive_triggers_recompute(meters_client): assert resp.status_code == 200 # recompute should be triggered assert mock_recompute.called + assert mock_recompute.call_args.kwargs["commit"] is False call_args = mock_recompute.call_args recompute_start = call_args[0][1] if recompute_start.tzinfo is None: diff --git a/tests/test_energy_cost.py b/tests/test_energy_cost.py index 6fad5cb..a3ae6e1 100644 --- a/tests/test_energy_cost.py +++ b/tests/test_energy_cost.py @@ -2176,6 +2176,32 @@ class TestComputeClosedPeriods: class TestRecomputeRange: + def test_default_commit_is_visible_to_a_new_session(self, energy_db: Session) -> None: + """The public recompute API retains its standalone commit behaviour.""" + _setup_manual_scenario(energy_db) + + assert recompute_range(energy_db, _T0, _T1) == 1 + assert energy_db.bind is not None + with Session(energy_db.bind) as observer: + row = observer.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.degraded is False + + def test_commit_false_leaves_recompute_for_caller_to_rollback(self, energy_db: Session) -> None: + """Caller-owned recompute writes disappear entirely after rollback.""" + _setup_manual_scenario(energy_db) + + assert recompute_range(energy_db, _T0, _T1, commit=False) == 1 + energy_db.rollback() + + assert energy_db.bind is not None + with Session(energy_db.bind) as observer: + rows = observer.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalars().all() + assert rows == [] + def test_overwrites_existing_rows(self, energy_db: Session) -> None: """recompute_range must overwrite non-degraded rows (explicit opt-in).""" contract = _make_contract(energy_db, kind="manual", active=True)