M8-R03B: keep meter swap recompute in caller transaction
This commit is contained in:
@@ -151,7 +151,7 @@ def _trigger_recompute(db: Session, start: datetime, label: str) -> int:
|
|||||||
# started_at is in the future — nothing to recompute.
|
# started_at is in the future — nothing to recompute.
|
||||||
logger.info("%s: started_at (%s) is in the future, skipping recompute.", label, start)
|
logger.info("%s: started_at (%s) is in the future, skipping recompute.", label, start)
|
||||||
return 0
|
return 0
|
||||||
n = recompute_range(db, start, end)
|
n = recompute_range(db, start, end, commit=False)
|
||||||
logger.info(
|
logger.info(
|
||||||
"%s: recomputed %d period(s) in window [%s, %s).",
|
"%s: recomputed %d period(s) in window [%s, %s).",
|
||||||
label,
|
label,
|
||||||
@@ -352,11 +352,6 @@ def patch_energy_meter(
|
|||||||
note=body.note,
|
note=body.note,
|
||||||
started_at=new_started_at_utc,
|
started_at=new_started_at_utc,
|
||||||
)
|
)
|
||||||
except MeterIntervalError as exc:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
|
||||||
detail=str(exc),
|
|
||||||
)
|
|
||||||
|
|
||||||
# Retroactive recompute if started_at was changed.
|
# Retroactive recompute if started_at was changed.
|
||||||
if new_started_at_utc is not None and old_started_at is not None:
|
if new_started_at_utc is not None and old_started_at is not None:
|
||||||
@@ -369,6 +364,15 @@ def patch_energy_meter(
|
|||||||
_trigger_recompute(db, window_start, f"PATCH /api/energy/meters/{meter_id}")
|
_trigger_recompute(db, window_start, f"PATCH /api/energy/meters/{meter_id}")
|
||||||
|
|
||||||
db.commit()
|
db.commit()
|
||||||
|
except MeterIntervalError as exc:
|
||||||
|
db.rollback()
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||||
|
detail=str(exc),
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
db.rollback()
|
||||||
|
raise
|
||||||
db.refresh(meter)
|
db.refresh(meter)
|
||||||
|
|
||||||
# Trigger HA discovery re-publish so label renames on the active meter
|
# Trigger HA discovery re-publish so label renames on the active meter
|
||||||
|
|||||||
@@ -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)``.
|
"""Recompute (overwrite) all 15-minute periods in ``[start, end)``.
|
||||||
|
|
||||||
This is the *explicit opt-in* path for recovering from:
|
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
|
Parameters
|
||||||
----------
|
----------
|
||||||
session:
|
session:
|
||||||
Active SQLAlchemy session. The function commits after all periods
|
Active SQLAlchemy session.
|
||||||
have been processed.
|
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:
|
start:
|
||||||
Inclusive start datetime (floored to the nearest quarter-hour internally).
|
Inclusive start datetime (floored to the nearest quarter-hour internally).
|
||||||
end:
|
end:
|
||||||
@@ -724,6 +729,7 @@ def recompute_range(session: Session, start: datetime, end: datetime) -> int:
|
|||||||
)
|
)
|
||||||
t0 += timedelta(minutes=_PERIOD_MINUTES)
|
t0 += timedelta(minutes=_PERIOD_MINUTES)
|
||||||
|
|
||||||
|
if commit:
|
||||||
session.commit()
|
session.commit()
|
||||||
logger.info(
|
logger.info(
|
||||||
"recompute_range(%s, %s): wrote %d period(s).",
|
"recompute_range(%s, %s): wrote %d period(s).",
|
||||||
|
|||||||
@@ -36,10 +36,10 @@ from unittest.mock import patch
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
from sqlalchemy import create_engine, select
|
from sqlalchemy import create_engine, event, select
|
||||||
from sqlalchemy.orm import Session
|
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
|
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
|
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):
|
def test_declare_meter_overlap_returns_422(meters_client):
|
||||||
"""Declaring a meter with started_at before active meter's started_at → 422."""
|
"""Declaring a meter with started_at before active meter's started_at → 422."""
|
||||||
client, _ = meters_client
|
client, _ = meters_client
|
||||||
@@ -546,6 +594,7 @@ def test_declare_meter_retroactive_triggers_recompute(meters_client):
|
|||||||
assert resp.status_code == 201
|
assert resp.status_code == 201
|
||||||
# recompute_range should have been called with start == t_past
|
# recompute_range should have been called with start == t_past
|
||||||
assert mock_recompute.called
|
assert mock_recompute.called
|
||||||
|
assert mock_recompute.call_args.kwargs["commit"] is False
|
||||||
call_args = mock_recompute.call_args
|
call_args = mock_recompute.call_args
|
||||||
recompute_start = call_args[0][1] # positional arg index 1 (session is 0)
|
recompute_start = call_args[0][1] # positional arg index 1 (session is 0)
|
||||||
# Normalise for comparison
|
# Normalise for comparison
|
||||||
@@ -685,6 +734,7 @@ def test_patch_meter_started_at_retroactive_triggers_recompute(meters_client):
|
|||||||
assert resp.status_code == 200
|
assert resp.status_code == 200
|
||||||
# recompute should be triggered
|
# recompute should be triggered
|
||||||
assert mock_recompute.called
|
assert mock_recompute.called
|
||||||
|
assert mock_recompute.call_args.kwargs["commit"] is False
|
||||||
call_args = mock_recompute.call_args
|
call_args = mock_recompute.call_args
|
||||||
recompute_start = call_args[0][1]
|
recompute_start = call_args[0][1]
|
||||||
if recompute_start.tzinfo is None:
|
if recompute_start.tzinfo is None:
|
||||||
|
|||||||
@@ -2176,6 +2176,32 @@ class TestComputeClosedPeriods:
|
|||||||
|
|
||||||
|
|
||||||
class TestRecomputeRange:
|
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:
|
def test_overwrites_existing_rows(self, energy_db: Session) -> None:
|
||||||
"""recompute_range must overwrite non-degraded rows (explicit opt-in)."""
|
"""recompute_range must overwrite non-degraded rows (explicit opt-in)."""
|
||||||
contract = _make_contract(energy_db, kind="manual", active=True)
|
contract = _make_contract(energy_db, kind="manual", active=True)
|
||||||
|
|||||||
Reference in New Issue
Block a user