diff --git a/app/main.py b/app/main.py index a00a5c0..9a1b0de 100644 --- a/app/main.py +++ b/app/main.py @@ -33,6 +33,7 @@ from app.services.public_ip import check_public_ipv4_and_notify from app.services.modbus_poll import poll_all_enabled_devices, BASE_POLL_TICK_SECONDS from app.services.ha_discovery import publish_discovery, publish_states from app.services.tibber_prices import refresh_prices +from app.services.energy_cost import compute_closed_periods from scripts.app_db_adopt import AppDatabaseAdoptionError, validate_app_runtime_db logger = logging.getLogger(__name__) @@ -92,6 +93,28 @@ def _run_scheduled_tibber_refresh() -> None: session.close() +def _run_scheduled_energy_cost() -> None: + """Scheduled job: compute billing records for all uncalculated closed 15-minute periods. + + Runs every minute so that a new period is picked up within 1 minute of + closing. The job is a no-op when: + - No active energy contract with a version covering the period exists. + - DSMR data has not yet arrived for the period boundaries. + - The period's billing record already exists and is not degraded. + + Any unexpected exceptions are caught and logged so that a single failure + does not crash the scheduler or affect the other background jobs. + """ + session_local = get_session_local() + session: Session = session_local() + try: + compute_closed_periods(session) + except Exception: + logger.exception("_run_scheduled_energy_cost: unexpected error") + finally: + session.close() + + def _run_scheduled_ha_state_publish() -> None: """Periodic job: publish discovery configs + state + availability for all enabled exposed entities. @@ -188,6 +211,17 @@ async def lifespan(_: FastAPI): max_instances=1, coalesce=True, ) + # Energy cost billing: compute uncalculated closed 15-minute periods every minute. + # The job is a no-op when no active contract or DSMR data is present, so it is + # safe to register unconditionally. + scheduler.add_job( + _run_scheduled_energy_cost, + trigger=IntervalTrigger(minutes=1), + id="energy-cost", + replace_existing=True, + max_instances=1, + coalesce=True, + ) scheduler.start() # MQTT: connect using DB-merged runtime settings so broker configured via UI diff --git a/app/services/energy_cost.py b/app/services/energy_cost.py new file mode 100644 index 0000000..8e4a2ce --- /dev/null +++ b/app/services/energy_cost.py @@ -0,0 +1,569 @@ +"""Billing engine for DSMR 15-minute energy metering periods. + +This module implements the two-layer billing model described in §3.4 of the +M6 design document: + +**Layer 1 — per-period metering cost (immutable, price-snapshot)** + ``compute_period(session, t0)`` computes the import cost, export revenue, and + net cost for the 15-minute period ``[t0, t0+15min)``. The result is written + to ``energy_cost_period`` with a full pricing snapshot so each row is + self-contained and auditable. Existing *successful* rows are never overwritten + by the normal tick path; only an explicit ``recompute_range`` call passes + ``overwrite=True``. + +**Layer 2 — summary (computed at read time, not stored)** + ``summarize(session, start, end)`` aggregates all non-degraded + ``energy_cost_period`` rows in ``[start, end)``, then adds the daily + standing charges (network_fee + management_fee, apportioned at EUR/month + ÷ 30 per day) and subtracts the energy-tax credit (heffingskorting, + apportioned at EUR/year ÷ 365 per day). + +Design notes +------------ +- **Decimal arithmetic throughout**: all monetary computations use + ``decimal.Decimal`` to avoid float binary rounding errors. Only when + writing to ``EnergyCostPeriod`` columns (Float) are values converted to + float. ``summarize`` converts back to Decimal for summation. +- **UTC quarter-hour grid**: period boundaries are aligned to UTC 00/15/30/45 + minutes (``floor_to_quarter``). NL local time (CET/CEST) is always a whole + number of hours from UTC, so the quarter-hour grid is the same in both + timezone representations. +- **Register keys**: DSMR payload uses JSON strings like ``"20915.154"`` + for cumulative kWh registers. ``register_at`` converts them to Decimal. +- **Degraded vs skip semantics**: + - *Missing readings* (``register_at`` returns None for start or end + boundary): write a ``degraded=True`` row with costs at 0 so the period is + tracked and can be retried by ``compute_closed_periods``. + - *Missing Tibber price* (``TibberPriceNotFoundError``): skip entirely (do + not write a row); the period will be retried once prices arrive. + - *Missing active contract version*: skip (no contract to compute against). +- **Lookback window in ``compute_closed_periods``**: to avoid scanning all + historical DSMR data on every tick, the function looks back at most 7 days + from the current time. This covers typical short outages (no data / no + contract) while staying bounded. Periods older than 7 days must be + recovered via an explicit ``recompute_range`` call. +""" + +from __future__ import annotations + +import logging +from datetime import UTC, datetime, timedelta +from decimal import Decimal +from typing import Any + +from sqlalchemy import select +from sqlalchemy.orm import Session + +from app.integrations.pricing.strategies import ( + PeriodDeltas, + TibberPriceNotFoundError, + get_strategy, +) +from app.models.energy import DsmrReading, EnergyCostPeriod +from app.services.contracts import active_contract_version_at + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + +_PERIOD_MINUTES = 15 +_LOOKBACK_DAYS = 7 # maximum lookback window for compute_closed_periods + +# DSMR payload register keys (cumulative kWh, JSON string values). +_KEY_D1 = "electricity_delivered_1" # delivered low-tariff (dal / _1) +_KEY_D2 = "electricity_delivered_2" # delivered high-tariff (normal / _2) +_KEY_R1 = "electricity_returned_1" # returned low-tariff +_KEY_R2 = "electricity_returned_2" # returned high-tariff + + +# --------------------------------------------------------------------------- +# Internal helpers +# --------------------------------------------------------------------------- + + +def floor_to_quarter(dt: datetime) -> datetime: + """Return *dt* floored to the nearest UTC quarter-hour boundary. + + The result always has seconds=0 and microseconds=0, and minutes in + {0, 15, 30, 45}. Timezone info is preserved if present. + """ + floored_minute = (dt.minute // _PERIOD_MINUTES) * _PERIOD_MINUTES + return dt.replace(minute=floored_minute, second=0, microsecond=0) + + +def _to_decimal(value: Any) -> Decimal: + """Convert *value* to Decimal via str() to avoid float binary rounding.""" + return Decimal(str(value)) + + +def _as_utc(dt: datetime) -> datetime: + """Attach UTC tzinfo to a naive datetime (SQLite read-back workaround).""" + if dt.tzinfo is None: + return dt.replace(tzinfo=UTC) + return dt + + +def _existing_period(session: Session, t0: datetime) -> EnergyCostPeriod | None: + """Return the EnergyCostPeriod row for period_start=t0, or None.""" + return session.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == t0) + ).scalar_one_or_none() + + +# --------------------------------------------------------------------------- +# register_at — boundary reading lookup +# --------------------------------------------------------------------------- + + +def register_at(session: Session, boundary: datetime) -> dict[str, Decimal] | None: + """Return the four cumulative kWh register values at *boundary*. + + Queries the most recent ``DsmrReading`` with ``recorded_at ≤ boundary`` + and extracts the four energy registers from ``payload``: + + d1 — electricity_delivered_1 (delivered low-tariff / dal) + d2 — electricity_delivered_2 (delivered high-tariff / normal) + r1 — electricity_returned_1 (returned low-tariff) + r2 — electricity_returned_2 (returned high-tariff) + + Returns + ------- + dict[str, Decimal] with keys ``d1``, ``d2``, ``r1``, ``r2``, or ``None`` + when: + - No ``DsmrReading`` row exists with ``recorded_at ≤ boundary``. + - Any of the four register keys is absent from the payload. + - Any of the four register values is ``None`` (null in JSON). + """ + row: DsmrReading | None = ( + session.execute( + select(DsmrReading) + .where(DsmrReading.recorded_at <= boundary) + .order_by(DsmrReading.recorded_at.desc()) + .limit(1) + ).scalar_one_or_none() + ) + + if row is None: + return None + + payload = row.payload or {} + try: + d1_raw = payload[_KEY_D1] + d2_raw = payload[_KEY_D2] + r1_raw = payload[_KEY_R1] + r2_raw = payload[_KEY_R2] + except KeyError: + return None + + if any(v is None for v in (d1_raw, d2_raw, r1_raw, r2_raw)): + return None + + return { + "d1": _to_decimal(d1_raw), + "d2": _to_decimal(d2_raw), + "r1": _to_decimal(r1_raw), + "r2": _to_decimal(r2_raw), + } + + +# --------------------------------------------------------------------------- +# compute_period — single 15-minute period +# --------------------------------------------------------------------------- + + +def compute_period(session: Session, t0: datetime, *, overwrite: bool = False) -> bool: + """Compute and upsert the billing record for the period ``[t0, t0+15min)``. + + Parameters + ---------- + session: + Active SQLAlchemy session. Caller is responsible for committing. + t0: + UTC start of the 15-minute period. **Must** lie on a quarter-hour + grid boundary (minutes ∈ {0, 15, 30, 45}, seconds=0, microseconds=0). + overwrite: + If ``True``, overwrite an existing *successful* row (i.e. re-compute + even when a non-degraded record already exists). The normal tick path + always passes ``False``; only ``recompute_range`` passes ``True``. + + Returns + ------- + bool + ``True`` if a record was written (inserted or updated), ``False`` if + the period was skipped (missing contract or missing Tibber price). + + Side-effects + ------------ + - Inserts or updates an ``EnergyCostPeriod`` row keyed on ``period_start=t0``. + - If readings are missing at either boundary: inserts/updates a degraded + row (costs=0, degraded=True). + - If the active contract version is missing: **skips** (returns False, no write). + - If the Tibber price is missing (TibberPriceNotFoundError): **skips** + (returns False, no write). + """ + t1 = t0 + timedelta(minutes=_PERIOD_MINUTES) + now = datetime.now(UTC) + + # Immutability guard: skip if a successful record already exists and we + # are not in overwrite mode. + existing = _existing_period(session, t0) + if existing is not None and not existing.degraded and not overwrite: + return False + + # --- Active contract version at t0 (checked before readings) --- + # If there is no active contract covering t0, skip the period entirely. + # We do not write a degraded row — there is no meaningful state to recover + # without a contract (we would not know which strategy to apply once data + # arrives). The period can be recovered via an explicit recompute_range once + # a contract is configured and activated. + version = active_contract_version_at(session, t0) + if version is None: + logger.debug("compute_period(%s): no active contract version — skipping.", t0.isoformat()) + return False + + # --- Boundary readings --- + start_regs = register_at(session, t0) + end_regs = register_at(session, t1) + + if start_regs is None or end_regs is None: + # Missing readings → write/update a degraded placeholder so the period + # is visible and can be retried by compute_closed_periods once data arrives. + _upsert_degraded(session, t0, now, existing) + return True # a record was written (degraded) + + # --- Compute deltas (end − start) --- + deltas = PeriodDeltas( + d1=end_regs["d1"] - start_regs["d1"], + d2=end_regs["d2"] - start_regs["d2"], + r1=end_regs["r1"] - start_regs["r1"], + r2=end_regs["r2"] - start_regs["r2"], + ) + + # --- Price strategy --- + strategy = get_strategy(version.contract.kind) + try: + result = strategy(deltas, t0, version.values, session) + except TibberPriceNotFoundError: + # Missing Tibber price → skip the period; it will be retried once the + # price arrives (e.g. after the next Tibber refresh job runs). + logger.debug( + "compute_period(%s): no Tibber price found — skipping.", t0.isoformat() + ) + return False + + # --- Upsert the billing record --- + import_cost: Decimal = result["import_cost"] + export_revenue: Decimal = result["export_revenue"] + net_cost: Decimal = result["net_cost"] + pricing: dict = result["pricing"] + + if existing is not None: + # Update in-place (overwrite=True or previous record was degraded). + existing.d1_kwh = float(deltas.d1) + existing.d2_kwh = float(deltas.d2) + existing.r1_kwh = float(deltas.r1) + existing.r2_kwh = float(deltas.r2) + existing.import_cost = float(import_cost) + existing.export_revenue = float(export_revenue) + existing.net_cost = float(net_cost) + existing.currency = version.contract.currency + existing.pricing = pricing + existing.contract_version_id = version.id + existing.degraded = False + existing.computed_at = now + else: + period = EnergyCostPeriod( + period_start=t0, + d1_kwh=float(deltas.d1), + d2_kwh=float(deltas.d2), + r1_kwh=float(deltas.r1), + r2_kwh=float(deltas.r2), + import_cost=float(import_cost), + export_revenue=float(export_revenue), + net_cost=float(net_cost), + currency=version.contract.currency, + pricing=pricing, + contract_version_id=version.id, + degraded=False, + computed_at=now, + ) + session.add(period) + + return True + + +def _upsert_degraded( + session: Session, + t0: datetime, + now: datetime, + existing: EnergyCostPeriod | None, +) -> None: + """Insert or update a degraded placeholder for period *t0*. + + When *existing* is not None (row was previously written — either degraded + or successful), the row is explicitly reset to the standard degraded state. + This is required for the ``recompute_range`` (overwrite=True) path: if the + row was previously a *successful* computation and the boundary readings have + since disappeared, the stale non-zero costs must be cleared so the row + accurately reflects the current "missing readings" state rather than + masquerading as a valid result. + """ + if existing is not None: + # Explicitly reset to degraded state — identical field values to the + # new-row path below. This covers the recompute-over-successful-row + # case where old non-zero costs must not survive the downgrade. + existing.d1_kwh = 0.0 + existing.d2_kwh = 0.0 + existing.r1_kwh = 0.0 + existing.r2_kwh = 0.0 + existing.import_cost = 0.0 + existing.export_revenue = 0.0 + existing.net_cost = 0.0 + existing.pricing = {} + existing.contract_version_id = None + existing.degraded = True + existing.computed_at = now + else: + period = EnergyCostPeriod( + period_start=t0, + d1_kwh=0.0, + d2_kwh=0.0, + r1_kwh=0.0, + r2_kwh=0.0, + import_cost=0.0, + export_revenue=0.0, + net_cost=0.0, + currency="EUR", # placeholder; real currency known after contract lookup + pricing={}, + contract_version_id=None, + degraded=True, + computed_at=now, + ) + session.add(period) + + +# --------------------------------------------------------------------------- +# compute_closed_periods — periodic tick +# --------------------------------------------------------------------------- + + +def compute_closed_periods(session: Session) -> int: + """Find and compute all uncalculated closed 15-minute periods. + + A period ``[t0, t1)`` is *closed* when ``t1 ≤ now``. This function: + + 1. Determines the lookback window: from ``now − LOOKBACK_DAYS`` to ``now``, + floored to the nearest quarter-hour. This avoids an unbounded full + historical scan on every tick while still covering the typical recovery + window (short outages, missing contract, etc.). Periods older than + ``LOOKBACK_DAYS`` must be recovered via an explicit ``recompute_range``. + 2. Iterates over all quarter-hour boundaries in that window where + ``t1 ≤ now`` (i.e. the period has already closed). + 3. For each boundary, calls ``compute_period(overwrite=False)``, which: + - Skips periods that already have a *successful* (non-degraded) record. + - Retries periods that have a *degraded* record. + - Skips periods for which no active contract version exists or the + Tibber price is unavailable (without writing a degraded row). + + Returns + ------- + int + Number of periods for which a record was written (inserted or updated). + Does not count skipped periods. + """ + now = datetime.now(UTC) + # Current period boundary (the one whose t1 has not yet passed). + current_t0 = floor_to_quarter(now) + # Earliest boundary to consider. + earliest_t0 = floor_to_quarter(now - timedelta(days=_LOOKBACK_DAYS)) + + written = 0 + t0 = earliest_t0 + while t0 < current_t0: + t1 = t0 + timedelta(minutes=_PERIOD_MINUTES) + if t1 <= now: + try: + did_write = compute_period(session, t0, overwrite=False) + if did_write: + written += 1 + except Exception: + logger.exception( + "compute_closed_periods: unexpected error for t0=%s — continuing.", + t0.isoformat(), + ) + t0 += timedelta(minutes=_PERIOD_MINUTES) + + if written: + session.commit() + logger.info("compute_closed_periods: wrote %d period(s).", written) + return written + + +# --------------------------------------------------------------------------- +# recompute_range — explicit full recompute +# --------------------------------------------------------------------------- + + +def recompute_range(session: Session, start: datetime, end: datetime) -> int: + """Recompute (overwrite) all 15-minute periods in ``[start, end)``. + + This is the *explicit opt-in* path for recovering from: + - Periods where readings or prices arrived late. + - Price corrections (new contract version retroactively applied). + - Any other reason to override the immutability guard. + + The function iterates over every UTC quarter-hour boundary in + ``[floor(start), end)`` and calls ``compute_period(overwrite=True)``. + Existing rows (including successful ones) are overwritten. + + Parameters + ---------- + session: + Active SQLAlchemy session. The function commits after all periods + have been processed. + start: + Inclusive start datetime (floored to the nearest quarter-hour internally). + end: + Exclusive end datetime. + + Returns + ------- + int + Number of periods for which a record was written (inserted or updated). + Periods skipped due to missing contract or missing Tibber price are + *not* counted. + """ + t0 = floor_to_quarter(_as_utc(start)) + end_utc = _as_utc(end) + + written = 0 + while t0 < end_utc: + try: + did_write = compute_period(session, t0, overwrite=True) + if did_write: + written += 1 + except Exception: + logger.exception( + "recompute_range: unexpected error for t0=%s — continuing.", + t0.isoformat(), + ) + t0 += timedelta(minutes=_PERIOD_MINUTES) + + session.commit() + logger.info( + "recompute_range(%s, %s): wrote %d period(s).", + start.isoformat(), + end.isoformat(), + written, + ) + return written + + +# --------------------------------------------------------------------------- +# summarize — layer-2 aggregation (read-time, not stored) +# --------------------------------------------------------------------------- + + +def summarize(session: Session, start: datetime, end: datetime) -> dict[str, Any]: + """Aggregate billing for the interval ``[start, end)``. + + Computes the total payable as: + + total_payable = Σ(net_cost) -- metered electricity + + fixed_costs -- (network_fee + management_fee) EUR/month ÷ 30 × days + - credits -- heffingskorting EUR/year ÷ 365 × days + + The fixed costs and credits are derived from the *currently active contract + version* at *end* (i.e. the version in effect at the end of the requested + interval). When no active contract version exists, fixed costs and credits + are both 0; only the metered sum is returned. + + All arithmetic uses Decimal; the returned dict contains Python floats for + JSON-serialisation convenience. + + Parameters + ---------- + session: + Active read-only SQLAlchemy session. + start: + Inclusive start of the summary interval. + end: + Exclusive end of the summary interval. + + Returns + ------- + dict with keys: + + currency str ISO 4217 currency (from contract, or "EUR" fallback) + metered_import float Σ import_cost from non-degraded periods + metered_export float Σ export_revenue from non-degraded periods + metered_net float Σ net_cost from non-degraded periods + fixed_costs float standing charges apportioned over the interval + credits float energy-tax credit apportioned over the interval + total_payable float metered_net + fixed_costs − credits + period_count int number of non-degraded periods in range + degraded_count int number of degraded periods in range + days float interval length in days (total_seconds / 86400) + """ + start_utc = _as_utc(start) + end_utc = _as_utc(end) + + # --- Fetch all EnergyCostPeriod rows in [start, end) --- + rows = session.execute( + select(EnergyCostPeriod).where( + EnergyCostPeriod.period_start >= start_utc, + EnergyCostPeriod.period_start < end_utc, + ) + ).scalars().all() + + good_rows = [r for r in rows if not r.degraded] + degraded_rows = [r for r in rows if r.degraded] + + # Σ monetary amounts (Decimal arithmetic). + sum_import = sum((_to_decimal(r.import_cost) for r in good_rows), Decimal("0")) + sum_export = sum((_to_decimal(r.export_revenue) for r in good_rows), Decimal("0")) + sum_net = sum((_to_decimal(r.net_cost) for r in good_rows), Decimal("0")) + + # --- Interval length in days --- + total_seconds = (end_utc - start_utc).total_seconds() + days = _to_decimal(str(total_seconds)) / _to_decimal("86400") + + # --- Active contract version at *end* for standing charges --- + version = active_contract_version_at(session, end_utc) + + fixed_dec = Decimal("0") + credits_dec = Decimal("0") + currency = "EUR" + + if version is not None: + currency = version.contract.currency + vals: dict = version.values or {} + standing: dict = vals.get("standing", {}) + creds: dict = vals.get("credits", {}) + + network_fee = _to_decimal(standing.get("network_fee", 0)) + management_fee = _to_decimal(standing.get("management_fee", 0)) + heffingskorting = _to_decimal(creds.get("heffingskorting", 0)) + + # Standing charges: EUR/month → EUR/day (÷ 30) × days. + fixed_dec = (network_fee + management_fee) / Decimal("30") * days + + # Energy-tax credit: EUR/year → EUR/day (÷ 365) × days. + credits_dec = heffingskorting / Decimal("365") * days + + total_payable = sum_net + fixed_dec - credits_dec + + return { + "currency": currency, + "metered_import": float(sum_import), + "metered_export": float(sum_export), + "metered_net": float(sum_net), + "fixed_costs": float(fixed_dec), + "credits": float(credits_dec), + "total_payable": float(total_payable), + "period_count": len(good_rows), + "degraded_count": len(degraded_rows), + "days": float(days), + } diff --git a/docs/design/m6-tibber-dynamic-energy.md b/docs/design/m6-tibber-dynamic-energy.md index c5d9d0f..577a4aa 100644 --- a/docs/design/m6-tibber-dynamic-energy.md +++ b/docs/design/m6-tibber-dynamic-energy.md @@ -378,7 +378,7 @@ Phase D(API + 前端) - **Reviewer checklist**: on_message 网络线程内只做短事务、吞异常不崩连接;整帧存(含 gas/各相);null 容错;`source_id` 幂等;10s 降采样正确。 ### M6-T07 — 计费引擎 + 周期 job + 汇总 + 重算 -- **Status**: `todo` · **Depends**: M6-T01, M6-T03, M6-T04, M6-T05, M6-T06 +- **Status**: `done` · **Depends**: M6-T01, M6-T03, M6-T04, M6-T05, M6-T06 - **Context**: 每 15min 用寄存器差 × active 合同 strategy 算计量电费(快照、不可变);汇总加固定费 − 抵扣;周期 job + 重算。 - **Files**: `create app/services/energy_cost.py`;`modify app/main.py`(周期 job);`create tests/test_energy_cost.py` - **Steps**: diff --git a/tests/test_energy_cost.py b/tests/test_energy_cost.py new file mode 100644 index 0000000..a078c3c --- /dev/null +++ b/tests/test_energy_cost.py @@ -0,0 +1,1197 @@ +"""Tests for M6-T07: billing engine (app/services/energy_cost.py). + +Acceptance criteria covered +---------------------------- +1. ``compute_period`` with a ``manual`` dual-tariff contract: correct + import_cost / export_revenue / net_cost; pricing snapshot and + contract_version_id stored. +2. ``compute_period`` idempotency: calling twice with overwrite=False leaves + the row unchanged; overwrite=True re-computes and updates it. +3. ``compute_period`` with a ``tibber`` contract + a matching TibberPrice: + correct calculation using ``total`` as buy price. +4. Missing Tibber price → period is **skipped** (no row written). +5. Missing readings (no DsmrReading covering a boundary) → degraded row. +6. Cross-version selection: two contract versions with different effective + dates; each t0 picks the correct version (asserts contract_version_id). +7. ``summarize``: Σnet + standing charges (month/30 × days) − heffingskorting + (year/365 × days) — hand-calculated against known values. +8. ``compute_closed_periods``: only processes closed and uncalculated periods; + does not touch already-computed (non-degraded) rows. +9. ``recompute_range``: explicitly overwrites all periods in [start, end) + including already-computed rows. +10. ``floor_to_quarter`` helper returns correct UTC quarter-hour boundaries. +""" + +from __future__ import annotations + +from datetime import UTC, datetime, timedelta +from decimal import Decimal +from pathlib import Path + +import pytest +from alembic import command +from alembic.config import Config +from sqlalchemy import create_engine, select +from sqlalchemy.orm import Session + +from app.integrations.pricing.strategies import PeriodDeltas # noqa: F401 +from app.models.energy import ( + DsmrReading, + EnergyContract, + EnergyContractVersion, + EnergyCostPeriod, + TibberPrice, +) +from app.services.energy_cost import ( + compute_closed_periods, + compute_period, + floor_to_quarter, + register_at, + recompute_range, + summarize, +) + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +def _make_app_alembic_config(database_url: str) -> Config: + cfg = Config("alembic_app.ini") + cfg.set_main_option("sqlalchemy.url", database_url) + return cfg + + +@pytest.fixture() +def energy_db(tmp_path: Path): + """Temporary SQLite DB upgraded to head; yields an open Session.""" + db_path = tmp_path / "energy_cost_test.db" + db_url = f"sqlite:///{db_path}" + alembic_cfg = _make_app_alembic_config(db_url) + command.upgrade(alembic_cfg, "head") + engine = create_engine(db_url, connect_args={"check_same_thread": False}) + session = Session(engine) + yield session + session.close() + engine.dispose() + + +# --------------------------------------------------------------------------- +# Data-builder helpers +# --------------------------------------------------------------------------- + +_UTC = UTC + + +def _ts(hour: int, minute: int = 0, second: int = 0) -> datetime: + """Return a UTC datetime on 2026-06-23 at the given time.""" + return datetime(2026, 6, 23, hour, minute, second, tzinfo=_UTC) + + +# Manual contract values used across most tests. +_MANUAL_VALUES = { + "energy": { + "buy": {"normal": 0.133, "dal": 0.127}, + "sell": {"normal": 0.05, "dal": 0.05}, + "energy_tax": 0.11, + "ode": 0.0, + }, + "standing": { + "network_fee": 9.87, + "management_fee": 9.87, + }, + "credits": { + "heffingskorting": 600.0, + }, +} + + +def _make_contract( + session: Session, + *, + kind: str = "manual", + active: bool = True, + currency: str = "EUR", +) -> EnergyContract: + """Insert and flush an EnergyContract; return the ORM object.""" + now = datetime.now(_UTC) + c = EnergyContract( + name=f"Test Contract ({kind})", + kind=kind, + active=active, + currency=currency, + created_at=now, + updated_at=now, + ) + session.add(c) + session.flush() + return c + + +def _make_version( + session: Session, + contract: EnergyContract, + values: dict, + *, + effective_from: datetime, + effective_to: datetime | None = None, +) -> EnergyContractVersion: + """Insert and flush an EnergyContractVersion; return the ORM object.""" + v = EnergyContractVersion( + contract_id=contract.id, + effective_from=effective_from, + effective_to=effective_to, + values=values, + created_at=datetime.now(_UTC), + ) + session.add(v) + session.flush() + return v + + +def _make_reading( + session: Session, + *, + recorded_at: datetime, + d1: str = "20000.000", + d2: str = "10000.000", + r1: str = "5000.000", + r2: str = "3000.000", + source_id: int | None = None, +) -> DsmrReading: + """Insert and flush a DsmrReading with the given register values.""" + r = DsmrReading( + recorded_at=recorded_at, + source_id=source_id, + payload={ + "electricity_delivered_1": d1, + "electricity_delivered_2": d2, + "electricity_returned_1": r1, + "electricity_returned_2": r2, + "current_electricity_usage": "1.234", + }, + ) + session.add(r) + session.flush() + return r + + +def _make_tibber_price( + session: Session, + *, + starts_at: datetime, + total: float, + energy: float = 0.18, + currency: str = "EUR", +) -> TibberPrice: + """Insert and flush a TibberPrice row; return the ORM object.""" + tax = round(total - energy, 6) + row = TibberPrice( + starts_at=starts_at, + resolution="QUARTER_HOURLY", + energy=energy, + tax=tax, + total=total, + level="NORMAL", + currency=currency, + fetched_at=datetime.now(_UTC), + ) + session.add(row) + session.flush() + return row + + +# --------------------------------------------------------------------------- +# 10. floor_to_quarter helper +# --------------------------------------------------------------------------- + + +class TestFloorToQuarter: + def test_already_on_boundary(self) -> None: + dt = _ts(10, 0) + assert floor_to_quarter(dt) == dt + + def test_15_boundary(self) -> None: + dt = _ts(10, 15) + assert floor_to_quarter(dt) == dt + + def test_floor_mid_period(self) -> None: + dt = _ts(10, 7, 30) + assert floor_to_quarter(dt) == _ts(10, 0) + + def test_floor_minute_16(self) -> None: + dt = _ts(10, 16) + assert floor_to_quarter(dt) == _ts(10, 15) + + def test_floor_minute_44(self) -> None: + dt = _ts(10, 44, 59) + assert floor_to_quarter(dt) == _ts(10, 30) + + def test_floor_minute_45(self) -> None: + dt = _ts(10, 45) + assert floor_to_quarter(dt) == _ts(10, 45) + + def test_seconds_zeroed(self) -> None: + dt = _ts(10, 3, 45) + result = floor_to_quarter(dt) + assert result.second == 0 + assert result.microsecond == 0 + + def test_timezone_preserved(self) -> None: + dt = _ts(10, 7) + result = floor_to_quarter(dt) + assert result.tzinfo is _UTC + + +# --------------------------------------------------------------------------- +# register_at tests +# --------------------------------------------------------------------------- + + +class TestRegisterAt: + def test_returns_none_when_no_readings(self, energy_db: Session) -> None: + result = register_at(energy_db, _ts(10, 0)) + assert result is None + + def test_returns_most_recent_at_or_before_boundary(self, energy_db: Session) -> None: + # Insert two readings: one before boundary, one after. + _make_reading(energy_db, recorded_at=_ts(9, 55), d1="100.0", d2="200.0", r1="10.0", r2="20.0", source_id=1) + _make_reading(energy_db, recorded_at=_ts(10, 5), d1="999.0", d2="999.0", r1="999.0", r2="999.0", source_id=2) + energy_db.commit() + + result = register_at(energy_db, _ts(10, 0)) + assert result is not None + assert result["d1"] == Decimal("100.0") + assert result["d2"] == Decimal("200.0") + + def test_exact_boundary_included(self, energy_db: Session) -> None: + _make_reading(energy_db, recorded_at=_ts(10, 0), d1="500.0", d2="600.0", r1="50.0", r2="60.0", source_id=1) + energy_db.commit() + + result = register_at(energy_db, _ts(10, 0)) + assert result is not None + assert result["d1"] == Decimal("500.0") + + def test_missing_register_key_returns_none(self, energy_db: Session) -> None: + r = DsmrReading( + recorded_at=_ts(10, 0), + source_id=99, + payload={"electricity_delivered_1": "100.0"}, # missing d2, r1, r2 + ) + energy_db.add(r) + energy_db.commit() + + result = register_at(energy_db, _ts(10, 0)) + assert result is None + + def test_null_register_value_returns_none(self, energy_db: Session) -> None: + r = DsmrReading( + recorded_at=_ts(10, 0), + source_id=88, + payload={ + "electricity_delivered_1": None, + "electricity_delivered_2": "200.0", + "electricity_returned_1": "10.0", + "electricity_returned_2": "20.0", + }, + ) + energy_db.add(r) + energy_db.commit() + + result = register_at(energy_db, _ts(10, 0)) + assert result is None + + def test_values_are_decimal(self, energy_db: Session) -> None: + _make_reading(energy_db, recorded_at=_ts(10, 0), d1="20915.154", d2="18372.099", + r1="1234.567", r2="890.123", source_id=1) + energy_db.commit() + + result = register_at(energy_db, _ts(10, 0)) + assert result is not None + assert isinstance(result["d1"], Decimal) + assert result["d1"] == Decimal("20915.154") + assert result["r2"] == Decimal("890.123") + + +# --------------------------------------------------------------------------- +# 1-2. compute_period — manual dual-tariff +# --------------------------------------------------------------------------- + +# Hand-calculation reference for the tests below: +# +# Start registers (t0 = 10:00): d1=20000.000, d2=10000.000, r1=5000.000, r2=3000.000 +# End registers (t1 = 10:15): d1=20000.500, d2=10001.200, r1=5000.000, r2=3000.100 +# +# Deltas: Δd1=0.500, Δd2=1.200, Δr1=0.000, Δr2=0.100 +# +# buy_dal = 0.127 + 0.11 + 0.0 = 0.237 +# buy_normal = 0.133 + 0.11 + 0.0 = 0.243 +# sell_dal = 0.05 +# sell_normal = 0.05 +# +# import_cost = 0.500×0.237 + 1.200×0.243 = 0.1185 + 0.2916 = 0.4101 +# export_revenue = 0.000×0.05 + 0.100×0.05 = 0.000 + 0.005 = 0.005 +# net_cost = 0.4101 − 0.005 = 0.4051 + +_T0 = _ts(10, 0) +_T1 = _ts(10, 15) + +_START_D1 = "20000.000" +_START_D2 = "10000.000" +_START_R1 = "5000.000" +_START_R2 = "3000.000" + +_END_D1 = "20000.500" +_END_D2 = "10001.200" +_END_R1 = "5000.000" +_END_R2 = "3000.100" + + +def _setup_manual_scenario(session: Session) -> EnergyContractVersion: + """Create active manual contract + two boundary readings; return the version.""" + contract = _make_contract(session, kind="manual", active=True) + version = _make_version( + session, contract, _MANUAL_VALUES, + effective_from=_ts(0, 0), # covers t0=10:00 + ) + # Start reading (at t0) + _make_reading(session, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + # End reading (at t1) + _make_reading(session, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + session.commit() + return version + + +class TestComputePeriodManual: + def test_correct_import_cost(self, energy_db: Session) -> None: + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # import_cost = 0.500×0.237 + 1.200×0.243 = 0.4101 + assert abs(row.import_cost - 0.4101) < 1e-9 + + def test_correct_export_revenue(self, energy_db: Session) -> None: + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # export_revenue = 0.000×0.05 + 0.100×0.05 = 0.005 + assert abs(row.export_revenue - 0.005) < 1e-9 + + def test_correct_net_cost(self, energy_db: Session) -> None: + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # net_cost = 0.4101 − 0.005 = 0.4051 + assert abs(row.net_cost - 0.4051) < 1e-9 + + def test_pricing_snapshot_stored(self, energy_db: Session) -> None: + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.pricing["kind"] == "manual" + assert "buy_dal" in row.pricing + assert "buy_normal" in row.pricing + + def test_contract_version_id_stored(self, energy_db: Session) -> None: + version = _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.contract_version_id == version.id + + def test_not_degraded(self, energy_db: Session) -> None: + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.degraded is False + + def test_kwh_deltas_stored(self, energy_db: Session) -> None: + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert abs(row.d1_kwh - 0.5) < 1e-9 + assert abs(row.d2_kwh - 1.2) < 1e-9 + assert abs(row.r1_kwh - 0.0) < 1e-9 + assert abs(row.r2_kwh - 0.1) < 1e-9 + + def test_currency_stored(self, energy_db: Session) -> None: + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.currency == "EUR" + + +# --------------------------------------------------------------------------- +# 2. Idempotency (overwrite=False / overwrite=True) +# --------------------------------------------------------------------------- + + +class TestComputePeriodIdempotency: + def test_no_overwrite_does_not_change_row(self, energy_db: Session) -> None: + """Calling compute_period twice with overwrite=False must leave the row unchanged.""" + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + row_before = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + first_computed_at = row_before.computed_at + + # Modify an end reading — but overwrite=False should ignore it. + _make_reading(energy_db, recorded_at=_T1 + timedelta(seconds=1), + d1="99999.0", d2="99999.0", r1="99999.0", r2="99999.0", + source_id=99) + energy_db.commit() + + result = compute_period(energy_db, _T0, overwrite=False) + assert result is False # did not overwrite + + row_after = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # computed_at must be unchanged (row not touched). + assert row_after.computed_at == first_computed_at + + def test_overwrite_true_updates_row(self, energy_db: Session) -> None: + """compute_period with overwrite=True must write to an existing successful row. + + We verify overwrite=True by checking that the function returns True + (a write occurred) even though a non-degraded row already exists. + The row's computed_at timestamp will change because we call compute_period + again — that is the observable side-effect of overwriting. + """ + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + + row_before = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row_before.degraded is False, "Pre-condition: row must be non-degraded" + + # overwrite=True: must return True even though row already exists + result = compute_period(energy_db, _T0, overwrite=True) + assert result is True + + def test_only_one_row_per_period(self, energy_db: Session) -> None: + """There must never be two EnergyCostPeriod rows for the same period_start.""" + _setup_manual_scenario(energy_db) + compute_period(energy_db, _T0) + compute_period(energy_db, _T0, overwrite=True) + + rows = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalars().all() + assert len(rows) == 1 + + +# --------------------------------------------------------------------------- +# 3. compute_period — tibber contract + tibber_price +# --------------------------------------------------------------------------- + +# Hand-calculation: +# Δd1=0.500, Δd2=1.200 → total_delivered = 1.700 kWh +# Δr1=0.000, Δr2=0.100 → total_returned = 0.100 kWh +# tibber total = 0.25 EUR/kWh +# buy = 0.25; sell = 0.25 − 0.10 (energy_tax) − 0.0 (sell_adjust) = 0.15 +# import_cost = 1.700 × 0.25 = 0.425 +# export_revenue = 0.100 × 0.15 = 0.015 +# net_cost = 0.425 − 0.015 = 0.410 + +_TIBBER_VALUES = { + "energy": { + "energy_tax": 0.10, + "sell_adjust": 0.0, + }, + "standing": { + "management_fee": 5.99, + "network_fee": 9.87, + }, + "credits": { + "heffingskorting": 600.0, + }, +} + + +class TestComputePeriodTibber: + def _setup(self, session: Session, total: float = 0.25) -> tuple[EnergyContractVersion, TibberPrice]: + contract = _make_contract(session, kind="tibber", active=True) + version = _make_version(session, contract, _TIBBER_VALUES, effective_from=_ts(0, 0)) + _make_reading(session, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(session, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + price = _make_tibber_price(session, starts_at=_ts(9, 45), total=total) + session.commit() + return version, price + + def test_correct_import_cost(self, energy_db: Session) -> None: + self._setup(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # import_cost = 1.700 × 0.25 = 0.425 + assert abs(row.import_cost - 0.425) < 1e-9 + + def test_correct_export_revenue(self, energy_db: Session) -> None: + self._setup(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # sell = 0.25 - 0.10 = 0.15; export_revenue = 0.100 × 0.15 = 0.015 + assert abs(row.export_revenue - 0.015) < 1e-9 + + def test_correct_net_cost(self, energy_db: Session) -> None: + self._setup(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # net_cost = 0.425 - 0.015 = 0.410 + assert abs(row.net_cost - 0.410) < 1e-9 + + def test_contract_version_id_stored(self, energy_db: Session) -> None: + version, _ = self._setup(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.contract_version_id == version.id + + def test_pricing_snapshot_has_tibber_fields(self, energy_db: Session) -> None: + self._setup(energy_db) + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.pricing["kind"] == "tibber" + assert "buy" in row.pricing + assert "sell" in row.pricing + assert "tibber_price_starts_at" in row.pricing + + +# --------------------------------------------------------------------------- +# 4. Missing Tibber price → skip (no row written) +# --------------------------------------------------------------------------- + + +class TestComputePeriodMissingTibberPrice: + def test_no_row_written_when_price_missing(self, energy_db: Session) -> None: + """When Tibber price is absent for the period, no EnergyCostPeriod is written.""" + contract = _make_contract(energy_db, kind="tibber", active=True) + _make_version(energy_db, contract, _TIBBER_VALUES, effective_from=_ts(0, 0)) + _make_reading(energy_db, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(energy_db, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + # No TibberPrice inserted. + energy_db.commit() + + result = compute_period(energy_db, _T0) + assert result is False + + rows = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalars().all() + assert len(rows) == 0, "No row should be written when Tibber price is missing" + + def test_tibber_price_after_t0_counts_as_missing(self, energy_db: Session) -> None: + """A TibberPrice with starts_at > t0 must NOT be used; period is skipped.""" + contract = _make_contract(energy_db, kind="tibber", active=True) + _make_version(energy_db, contract, _TIBBER_VALUES, effective_from=_ts(0, 0)) + _make_reading(energy_db, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(energy_db, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + # Price starts AFTER t0 — must not be used. + _make_tibber_price(energy_db, starts_at=_ts(10, 15), total=0.25) + energy_db.commit() + + result = compute_period(energy_db, _T0) + assert result is False + + rows = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalars().all() + assert len(rows) == 0 + + +# --------------------------------------------------------------------------- +# 5. Missing readings → degraded row +# --------------------------------------------------------------------------- + + +class TestComputePeriodMissingReadings: + def test_degraded_when_start_reading_missing(self, energy_db: Session) -> None: + """No DsmrReading at or before t0 → degraded row written.""" + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + # Only an end reading; no start reading. + _make_reading(energy_db, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + energy_db.commit() + + result = compute_period(energy_db, _T0) + assert result is True + + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.degraded is True + assert row.import_cost == 0.0 + assert row.export_revenue == 0.0 + assert row.net_cost == 0.0 + + def test_degraded_when_end_reading_missing(self, energy_db: Session) -> None: + """No DsmrReading at or before t1 → degraded row written. + + We place a reading BEFORE t0 (so t0 boundary has data) but the first + reading AT OR AFTER t1 is only after t1+5min, leaving the t1 boundary + without a reading ≤ t1. This forces ``register_at(t1)`` to return the + same row as ``register_at(t0)`` — both map to the same pre-t0 reading — + which means deltas = 0 but NOT degraded. + + Actually the degraded condition for a *missing end reading* occurs when + there is NO DsmrReading in the DB at all with ``recorded_at ≤ t1``. + To achieve that, we only add a reading after t1. + """ + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + # Only a reading AFTER t1 — no reading at or before t1. + _make_reading(energy_db, recorded_at=_ts(10, 20), d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + energy_db.commit() + + result = compute_period(energy_db, _T0) + assert result is True + + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.degraded is True + + def test_degraded_has_no_contract_version_id(self, energy_db: Session) -> None: + """Degraded rows written due to missing readings have contract_version_id=None.""" + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + # No readings at all. + energy_db.commit() + + compute_period(energy_db, _T0) + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row.contract_version_id is None + + def test_degraded_retried_on_second_compute(self, energy_db: Session) -> None: + """A degraded row is retried (overwritten) when readings become available. + + First pass: no readings at all → degraded row (both start and end missing). + Second pass: add both boundary readings → row transitions to non-degraded. + """ + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + # First compute: no readings at all → degraded. + energy_db.commit() + + compute_period(energy_db, _T0) + row_degraded = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row_degraded.degraded is True + + # Now add both boundary readings and retry with overwrite=False. + # Degraded rows are retried by compute_period (overwrite=False still re-tries degraded). + _make_reading(energy_db, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(energy_db, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + energy_db.commit() + + result = compute_period(energy_db, _T0, overwrite=False) + assert result is True + + row_fixed = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row_fixed.degraded is False + assert row_fixed.import_cost > 0 + + +# --------------------------------------------------------------------------- +# 6. Cross-version selection +# --------------------------------------------------------------------------- + + +class TestCrossVersionSelection: + """Two contract versions with non-overlapping effective dates. + + Version 1: effective [00:00, 08:00) → lower buy prices + Version 2: effective [08:00, ∞) → higher buy prices + + Periods before 08:00 must use version 1; periods at or after 08:00 use v2. + """ + + # Different rate sets so we can distinguish which version was used. + _VALUES_V1 = { + "energy": { + "buy": {"normal": 0.10, "dal": 0.10}, + "sell": {"normal": 0.05, "dal": 0.05}, + "energy_tax": 0.0, + "ode": 0.0, + }, + "standing": {"network_fee": 5.0, "management_fee": 5.0}, + "credits": {"heffingskorting": 300.0}, + } + _VALUES_V2 = { + "energy": { + "buy": {"normal": 0.50, "dal": 0.50}, + "sell": {"normal": 0.10, "dal": 0.10}, + "energy_tax": 0.0, + "ode": 0.0, + }, + "standing": {"network_fee": 5.0, "management_fee": 5.0}, + "credits": {"heffingskorting": 300.0}, + } + + def _setup(self, session: Session) -> tuple[EnergyContractVersion, EnergyContractVersion]: + contract = _make_contract(session, kind="manual", active=True) + v1 = _make_version( + session, contract, self._VALUES_V1, + effective_from=_ts(0, 0), + effective_to=_ts(8, 0), + ) + v2 = _make_version( + session, contract, self._VALUES_V2, + effective_from=_ts(8, 0), + effective_to=None, + ) + session.commit() + return v1, v2 + + def test_period_before_version_boundary_uses_v1(self, energy_db: Session) -> None: + v1, v2 = self._setup(energy_db) + t0_early = _ts(7, 0) + t1_early = _ts(7, 15) + _make_reading(energy_db, recorded_at=t0_early, d1="1000.0", d2="1000.0", + r1="0.0", r2="0.0", source_id=10) + _make_reading(energy_db, recorded_at=t1_early, d1="1001.0", d2="1000.0", + r1="0.0", r2="0.0", source_id=11) + energy_db.commit() + + compute_period(energy_db, t0_early) + + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == t0_early) + ).scalar_one() + assert row.contract_version_id == v1.id, ( + f"Expected version v1 (id={v1.id}), got {row.contract_version_id}" + ) + # import_cost = 1.0 kWh × buy_normal(v1)=0.10 = 0.10 + assert abs(row.import_cost - 0.10) < 1e-9 + + def test_period_at_version_boundary_uses_v2(self, energy_db: Session) -> None: + v1, v2 = self._setup(energy_db) + t0_late = _ts(8, 0) + t1_late = _ts(8, 15) + _make_reading(energy_db, recorded_at=t0_late, d1="2000.0", d2="2000.0", + r1="0.0", r2="0.0", source_id=20) + _make_reading(energy_db, recorded_at=t1_late, d1="2001.0", d2="2000.0", + r1="0.0", r2="0.0", source_id=21) + energy_db.commit() + + compute_period(energy_db, t0_late) + + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == t0_late) + ).scalar_one() + assert row.contract_version_id == v2.id, ( + f"Expected version v2 (id={v2.id}), got {row.contract_version_id}" + ) + # import_cost = 1.0 kWh × buy_normal(v2)=0.50 = 0.50 + assert abs(row.import_cost - 0.50) < 1e-9 + + +# --------------------------------------------------------------------------- +# 7. summarize — hand-verified +# --------------------------------------------------------------------------- + +# Hand-calculation for TestSummarize: +# +# Interval: 2026-06-23 00:00 UTC to 2026-06-24 00:00 UTC (exactly 1 day) +# days = 1.0 +# +# Active contract values (from _MANUAL_VALUES): +# standing: network_fee=9.87, management_fee=9.87 +# credits: heffingskorting=600.0 +# +# Fixed costs per day = (9.87 + 9.87) / 30 × 1.0 = 19.74 / 30 = 0.658 +# Credits per day = 600.0 / 365 × 1.0 = 600.0 / 365 ≈ 1.6438... +# +# Metered net (from two periods, each net_cost=0.4051): +# Σnet = 2 × 0.4051 = 0.8102 +# (slightly approximate: actual computed floats from compute_period) +# +# total_payable = 0.8102 + 0.658 − 1.6438... ≈ −0.1756... + + +class TestSummarize: + def _setup_two_periods(self, session: Session) -> None: + """Insert contract + readings for two consecutive 15-min periods and compute them.""" + contract = _make_contract(session, kind="manual", active=True) + _make_version(session, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + + # Period 1: [10:00, 10:15) + _make_reading(session, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(session, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + + # Period 2: [10:15, 10:30) — same delta as period 1 + t2 = _ts(10, 30) + _make_reading(session, recorded_at=t2, + d1=str(float(_END_D1) + 0.5), + d2=str(float(_END_D2) + 1.2), + r1=str(float(_END_R1)), + r2=str(float(_END_R2) + 0.1), + source_id=3) + + session.commit() + compute_period(session, _T0) + compute_period(session, _T1) + session.commit() + + def test_metered_net_sum(self, energy_db: Session) -> None: + self._setup_two_periods(energy_db) + # Summarise over the two computed periods. + result = summarize(energy_db, _ts(10, 0), _ts(10, 30)) + # Σnet ≈ 2 × 0.4051 = 0.8102 + assert abs(result["metered_net"] - 0.8102) < 1e-6 + + def test_period_count(self, energy_db: Session) -> None: + self._setup_two_periods(energy_db) + result = summarize(energy_db, _ts(10, 0), _ts(10, 30)) + assert result["period_count"] == 2 + assert result["degraded_count"] == 0 + + def test_fixed_costs_formula(self, energy_db: Session) -> None: + """Fixed costs = (network_fee + management_fee) / 30 × days.""" + self._setup_two_periods(energy_db) + # Interval is 30 minutes = 0.5/48 day = 0.020833... days + days = Decimal("1800") / Decimal("86400") + expected_fixed = (Decimal("9.87") + Decimal("9.87")) / Decimal("30") * days + result = summarize(energy_db, _ts(10, 0), _ts(10, 30)) + assert abs(Decimal(str(result["fixed_costs"])) - expected_fixed) < Decimal("1e-9") + + def test_credits_formula(self, energy_db: Session) -> None: + """Credits = heffingskorting / 365 × days.""" + self._setup_two_periods(energy_db) + days = Decimal("1800") / Decimal("86400") + expected_credits = Decimal("600.0") / Decimal("365") * days + result = summarize(energy_db, _ts(10, 0), _ts(10, 30)) + assert abs(Decimal(str(result["credits"])) - expected_credits) < Decimal("1e-9") + + def test_total_payable_formula(self, energy_db: Session) -> None: + """total_payable = metered_net + fixed_costs − credits.""" + self._setup_two_periods(energy_db) + result = summarize(energy_db, _ts(10, 0), _ts(10, 30)) + expected = result["metered_net"] + result["fixed_costs"] - result["credits"] + assert abs(result["total_payable"] - expected) < 1e-9 + + def test_no_active_contract_returns_zero_standing(self, energy_db: Session) -> None: + """When no active contract exists, fixed_costs and credits are both 0.""" + # No contract at all — just compute a period manually and summarise. + result = summarize(energy_db, _ts(10, 0), _ts(10, 30)) + assert result["fixed_costs"] == 0.0 + assert result["credits"] == 0.0 + assert result["metered_net"] == 0.0 + assert result["total_payable"] == 0.0 + + def test_currency_from_contract(self, energy_db: Session) -> None: + self._setup_two_periods(energy_db) + result = summarize(energy_db, _ts(10, 0), _ts(10, 30)) + assert result["currency"] == "EUR" + + def test_degraded_excluded_from_metered_sum(self, energy_db: Session) -> None: + """Degraded rows must not contribute to the metered sums. + + We produce a degraded row by calling compute_period when no readings + exist at all for the period boundaries. + """ + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + # No readings at all → degraded period. + energy_db.commit() + compute_period(energy_db, _T0) + energy_db.commit() + + result = summarize(energy_db, _T0, _T1) + assert result["metered_net"] == 0.0 + assert result["degraded_count"] == 1 + assert result["period_count"] == 0 + + def test_one_day_summarize_hand_calc(self, energy_db: Session) -> None: + """Full 1-day hand-calculation: fixed_costs/30 and credits/365 with 1-day interval.""" + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + energy_db.commit() + + # Summarize over exactly 1 day (no periods in DB — only standing/credits). + start = datetime(2026, 6, 23, 0, 0, 0, tzinfo=_UTC) + end = datetime(2026, 6, 24, 0, 0, 0, tzinfo=_UTC) + result = summarize(energy_db, start, end) + + # fixed_costs = (9.87 + 9.87) / 30 × 1.0 = 0.658 + assert abs(result["fixed_costs"] - (9.87 + 9.87) / 30) < 1e-9 + # credits = 600 / 365 + assert abs(result["credits"] - 600.0 / 365) < 1e-9 + # total = 0 + 0.658 − (600/365) + expected_total = (9.87 + 9.87) / 30 - 600.0 / 365 + assert abs(result["total_payable"] - expected_total) < 1e-9 + + +# --------------------------------------------------------------------------- +# 8. compute_closed_periods +# --------------------------------------------------------------------------- + + +class TestComputeClosedPeriods: + def test_skips_future_periods(self, energy_db: Session) -> None: + """Periods whose t1 > now must not be computed.""" + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=datetime.now(_UTC)) + energy_db.commit() + + # The period [now, now+15min) is still open — should not be computed. + written = compute_closed_periods(energy_db) + assert written == 0 # nothing written (no closed periods with data) + + def test_does_not_overwrite_successful_period(self, energy_db: Session) -> None: + """A non-degraded row must not be overwritten by the regular tick. + + Uses a period from 2 days ago so it is definitely closed and within + the 7-day lookback window. + """ + past_t0 = floor_to_quarter(datetime.now(_UTC) - timedelta(days=2)) + past_t1 = past_t0 + timedelta(minutes=15) + + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=past_t0 - timedelta(hours=1)) + _make_reading(energy_db, recorded_at=past_t0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(energy_db, recorded_at=past_t1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + energy_db.commit() + + # First compute: direct call. + compute_period(energy_db, past_t0) + energy_db.commit() + + row_before = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == past_t0) + ).scalar_one() + original_computed_at = row_before.computed_at + + # compute_closed_periods runs — the row should NOT be touched. + compute_closed_periods(energy_db) + + row_after = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == past_t0) + ).scalar_one() + assert row_after.computed_at == original_computed_at + + def test_retries_degraded_rows(self, energy_db: Session) -> None: + """A degraded row is retried when readings become available on the next tick. + + We use a period from 2 days ago so it is definitely closed and within the + 7-day lookback window, regardless of what time the test runs today. + """ + # Use a period that is definitely closed (2 days ago, early morning UTC). + past_t0 = floor_to_quarter(datetime.now(_UTC) - timedelta(days=2)) + past_t1 = past_t0 + timedelta(minutes=15) + + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=past_t0 - timedelta(hours=1)) + # First compute: no readings → degraded. + energy_db.commit() + compute_period(energy_db, past_t0) + energy_db.commit() + + row_before = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == past_t0) + ).scalar_one() + assert row_before.degraded is True + + # Now add both boundary readings. + _make_reading(energy_db, recorded_at=past_t0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(energy_db, recorded_at=past_t1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + energy_db.commit() + + # compute_closed_periods should retry the degraded row. + compute_closed_periods(energy_db) + + row_after = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == past_t0) + ).scalar_one() + assert row_after.degraded is False + assert row_after.import_cost > 0 + + +# --------------------------------------------------------------------------- +# 9. recompute_range +# --------------------------------------------------------------------------- + + +class TestRecomputeRange: + 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) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + _make_reading(energy_db, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(energy_db, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + energy_db.commit() + + # Initial compute. + compute_period(energy_db, _T0) + energy_db.commit() + + # Simulate "new" end reading with higher values. + _make_reading(energy_db, recorded_at=_T1 - timedelta(seconds=5), + d1="20020.0", d2="10020.0", r1="5000.0", r2="3000.0", + source_id=77) + energy_db.commit() + + count = recompute_range(energy_db, _T0, _T1) + assert count == 1 + + row = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # The row must have been updated (import_cost changes because deltas changed). + assert row.degraded is False + + def test_returns_count_of_written_periods(self, energy_db: Session) -> None: + """recompute_range returns the number of periods actually written.""" + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + + # Add readings for two consecutive periods: [10:00,10:15), [10:15,10:30). + _make_reading(energy_db, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(energy_db, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + t2 = _ts(10, 30) + _make_reading(energy_db, recorded_at=t2, + d1=str(float(_END_D1) + 0.5), + d2=str(float(_END_D2) + 1.2), + r1=_END_R1, r2=str(float(_END_R2) + 0.1), + source_id=3) + energy_db.commit() + + count = recompute_range(energy_db, _T0, t2) + assert count == 2 + + def test_idempotent_on_multiple_calls(self, energy_db: Session) -> None: + """Calling recompute_range twice must not create duplicate rows.""" + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + _make_reading(energy_db, recorded_at=_T0, d1=_START_D1, d2=_START_D2, + r1=_START_R1, r2=_START_R2, source_id=1) + _make_reading(energy_db, recorded_at=_T1, d1=_END_D1, d2=_END_D2, + r1=_END_R1, r2=_END_R2, source_id=2) + energy_db.commit() + + recompute_range(energy_db, _T0, _T1) + recompute_range(energy_db, _T0, _T1) + + rows = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalars().all() + assert len(rows) == 1 + + def test_recompute_downgrades_successful_row_when_readings_disappear( + self, energy_db: Session + ) -> None: + """recompute_range must downgrade a previously successful row to degraded + when boundary readings no longer exist. + + Scenario (reproduces REWORK 1 from reviewer probe3.py): + 1. Compute period [10:00, 10:15) successfully — row is non-degraded, costs > 0. + 2. Delete both boundary readings (simulate data loss / correction). + 3. Call recompute_range over that window (overwrite=True path). + 4. Row must now be degraded=True, all cost/kWh fields = 0, + contract_version_id = None, pricing = {}. + + This verifies the "缺读数→degraded" contract holds even for the explicit + recompute path, i.e. stale successful values are never silently preserved. + """ + # Step 1: successful compute. + contract = _make_contract(energy_db, kind="manual", active=True) + _make_version(energy_db, contract, _MANUAL_VALUES, effective_from=_ts(0, 0)) + start_reading = _make_reading( + energy_db, recorded_at=_T0, + d1=_START_D1, d2=_START_D2, r1=_START_R1, r2=_START_R2, source_id=1, + ) + end_reading = _make_reading( + energy_db, recorded_at=_T1, + d1=_END_D1, d2=_END_D2, r1=_END_R1, r2=_END_R2, source_id=2, + ) + energy_db.commit() + + compute_period(energy_db, _T0) + energy_db.commit() + + row_initial = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + # Pre-condition: row is successful with non-zero import_cost. + assert row_initial.degraded is False + assert row_initial.import_cost > 0, "Pre-condition: import_cost must be non-zero" + assert row_initial.contract_version_id is not None + + # Step 2: delete the boundary readings to simulate data loss. + energy_db.delete(start_reading) + energy_db.delete(end_reading) + energy_db.commit() + + # Step 3: explicit recompute (overwrite=True path). + count = recompute_range(energy_db, _T0, _T1) + assert count == 1, "recompute_range must report one period written (degraded)" + + # Step 4: row must now reflect degraded state — no stale costs. + energy_db.expire_all() + row_after = energy_db.execute( + select(EnergyCostPeriod).where(EnergyCostPeriod.period_start == _T0) + ).scalar_one() + assert row_after.degraded is True + assert row_after.import_cost == 0.0 + assert row_after.export_revenue == 0.0 + assert row_after.net_cost == 0.0 + assert row_after.d1_kwh == 0.0 + assert row_after.d2_kwh == 0.0 + assert row_after.r1_kwh == 0.0 + assert row_after.r2_kwh == 0.0 + assert row_after.contract_version_id is None + assert row_after.pricing == {}