"""Tests for M6-T04: EnergyContract CRUD + versioning + pricing-profile API. Coverage matrix --------------- GET /api/energy/profiles - unauthenticated → 401 - authenticated → 200, both 'manual' and 'tibber' profiles present GET /api/energy/contracts - unauthenticated → 401 - authenticated empty → 200, items=[] - after creating two contracts → items with active flag correct POST /api/energy/contracts - unauthenticated → 401 - missing CSRF → 403 - invalid values (missing required field) → 422, DB unchanged - unknown kind → 422 - valid manual → 201, ContractDetailResponse with versions - valid tibber → 201 GET /api/energy/contracts/{id} - unauthenticated → 401 - not found → 404 - success → versions ordered by effective_from PATCH /api/energy/contracts/{id} - unauthenticated → 401 - missing CSRF → 403 - not found → 404 - rename → name updated - activate → mutual exclusion (A active → activate B → A.active=False) - deactivate → active=False, others unchanged POST /api/energy/contracts/{id}/versions - unauthenticated → 401 - missing CSRF → 403 - not found → 404 - invalid values → 422, no rows written, old version unchanged - effective_from ≤ previous → 422 - success → old version closed (effective_to set), new open version added, old version values unchanged (only-append semantics) Auth/CSRF matrix tested for all write endpoints. """ from __future__ import annotations from datetime import UTC, datetime, timedelta from typing import Any import pytest from fastapi.testclient import TestClient from sqlalchemy import create_engine, select from sqlalchemy.orm import Session from app.models.energy import EnergyContract, EnergyContractVersion # --------------------------------------------------------------------------- # Shared helpers # --------------------------------------------------------------------------- _CSRF = "test-csrf-token" def _login(client: TestClient) -> None: resp = client.post( "/api/auth/login", json={"username": "admin", "password": "test-password"}, ) assert resp.status_code == 200, f"Login failed: {resp.status_code} {resp.text}" # Minimal valid values for each kind. _MANUAL_VALUES: dict[str, Any] = { "energy": { "buy": {"normal": 0.40, "dal": 0.30}, "sell": {"normal": 0.10, "dal": 0.10}, "energy_tax": 0.1108, "ode": 0.0, }, "standing": { "network_fee": 25.0, "management_fee": 9.87, }, "credits": { "heffingskorting": 600.0, }, } _TIBBER_VALUES: dict[str, Any] = { "energy": { "energy_tax": 0.1108, "sell_adjust": 0.0, }, "standing": { "management_fee": 5.99, "network_fee": 25.0, }, "credits": { "heffingskorting": 600.0, }, } def _manual_payload(**overrides) -> dict[str, Any]: base: dict[str, Any] = { "name": "Test Manual Contract", "kind": "manual", "currency": "EUR", "values": _MANUAL_VALUES, } base.update(overrides) return base def _tibber_payload(**overrides) -> dict[str, Any]: base: dict[str, Any] = { "name": "Test Tibber Contract", "kind": "tibber", "currency": "EUR", "values": _TIBBER_VALUES, } base.update(overrides) return base # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- @pytest.fixture() def contracts_client(auth_database): """TestClient + SQLAlchemy engine for EnergyContract API tests.""" from app.main import create_app app_url = auth_database["app_url"] engine = create_engine(app_url, connect_args={"check_same_thread": False}) fastapi_app = create_app() with TestClient(fastapi_app) as test_client: yield test_client, engine engine.dispose() # --------------------------------------------------------------------------- # GET /api/energy/profiles # --------------------------------------------------------------------------- def test_profiles_unauthenticated_returns_401(contracts_client): client, _ = contracts_client resp = client.get("/api/energy/profiles") assert resp.status_code == 401 def test_profiles_returns_both_kinds(contracts_client): client, _ = contracts_client _login(client) resp = client.get("/api/energy/profiles") assert resp.status_code == 200 body = resp.json() assert "profiles" in body kinds = {p["kind"] for p in body["profiles"]} assert "manual" in kinds assert "tibber" in kinds def test_profiles_contain_structure(contracts_client): """Each profile entry must have the nested structure the front-end needs.""" client, _ = contracts_client _login(client) resp = client.get("/api/energy/profiles") body = resp.json() for profile in body["profiles"]: assert "kind" in profile assert "label" in profile assert "energy" in profile assert "standing" in profile assert "credits" in profile # --------------------------------------------------------------------------- # GET /api/energy/contracts # --------------------------------------------------------------------------- def test_list_contracts_unauthenticated_returns_401(contracts_client): client, _ = contracts_client resp = client.get("/api/energy/contracts") assert resp.status_code == 401 def test_list_contracts_empty(contracts_client): client, _ = contracts_client _login(client) resp = client.get("/api/energy/contracts") assert resp.status_code == 200 body = resp.json() assert body["items"] == [] assert body["total"] == 0 def test_list_contracts_shows_active_flag(contracts_client): """After creating two contracts and activating one, active flag is correct.""" client, _ = contracts_client _login(client) # Create contract A resp_a = client.post( "/api/energy/contracts", json=_manual_payload(name="Contract A"), headers={"X-CSRF-Token": _CSRF}, ) assert resp_a.status_code == 201 id_a = resp_a.json()["id"] # Create contract B resp_b = client.post( "/api/energy/contracts", json=_manual_payload(name="Contract B"), headers={"X-CSRF-Token": _CSRF}, ) assert resp_b.status_code == 201 id_b = resp_b.json()["id"] # Activate A client.patch( f"/api/energy/contracts/{id_a}", json={"active": True}, headers={"X-CSRF-Token": _CSRF}, ) resp = client.get("/api/energy/contracts") items = {c["id"]: c for c in resp.json()["items"]} assert items[id_a]["active"] is True assert items[id_b]["active"] is False # --------------------------------------------------------------------------- # POST /api/energy/contracts # --------------------------------------------------------------------------- def test_create_contract_unauthenticated_returns_401(contracts_client): client, _ = contracts_client resp = client.post( "/api/energy/contracts", json=_manual_payload(), headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 401 def test_create_contract_missing_csrf_returns_403(contracts_client): client, _ = contracts_client _login(client) resp = client.post("/api/energy/contracts", json=_manual_payload()) assert resp.status_code == 403 def test_create_contract_invalid_values_returns_422_no_db_write(contracts_client): """Non-conforming values must return 422 without writing any rows.""" client, engine = contracts_client _login(client) # Missing required field energy.buy.normal bad_values = { "energy": { "buy": {"dal": 0.30}, # missing "normal" "sell": {"normal": 0.10, "dal": 0.10}, "energy_tax": 0.1108, "ode": 0.0, }, "standing": {"network_fee": 25.0, "management_fee": 9.87}, "credits": {"heffingskorting": 600.0}, } resp = client.post( "/api/energy/contracts", json=_manual_payload(values=bad_values), headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 422 # Confirm no rows were written. with Session(engine) as s: count = s.execute(select(EnergyContract)).scalars().all() assert len(count) == 0 def test_create_contract_unknown_kind_returns_422(contracts_client): client, _ = contracts_client _login(client) resp = client.post( "/api/energy/contracts", json={ "name": "Bad Kind", "kind": "nonexistent_kind_xyz", "values": {}, }, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 422 def test_create_manual_contract_success(contracts_client): client, engine = contracts_client _login(client) resp = client.post( "/api/energy/contracts", json=_manual_payload(), headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 201 body = resp.json() assert body["name"] == "Test Manual Contract" assert body["kind"] == "manual" assert body["active"] is False # new contracts are inactive by default assert "versions" in body assert len(body["versions"]) == 1 v = body["versions"][0] assert v["effective_to"] is None # open-ended first version # DB check with Session(engine) as s: contracts = s.execute(select(EnergyContract)).scalars().all() assert len(contracts) == 1 assert contracts[0].name == "Test Manual Contract" versions = s.execute(select(EnergyContractVersion)).scalars().all() assert len(versions) == 1 def test_create_tibber_contract_success(contracts_client): client, _ = contracts_client _login(client) resp = client.post( "/api/energy/contracts", json=_tibber_payload(), headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 201 body = resp.json() assert body["kind"] == "tibber" assert len(body["versions"]) == 1 def test_create_contract_defaults_effective_from(contracts_client): """When effective_from is omitted, the version is created with a recent timestamp. SQLite stores datetimes as naive UTC strings; we compare both sides as naive UTC. """ client, _ = contracts_client _login(client) # Use naive UTC (no tzinfo) to match what SQLite gives back via Pydantic. # Remove tzinfo from datetime.now(UTC) to get a comparable naive datetime. before = datetime.now(UTC).replace(tzinfo=None) resp = client.post( "/api/energy/contracts", json=_manual_payload(), headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 201 v = resp.json()["versions"][0] # Parse as naive (SQLite round-trip strips tz) raw_eff = datetime.fromisoformat(v["effective_from"]) # Strip tzinfo from raw_eff if present (shouldn't be, but guard anyway) eff_naive = raw_eff.replace(tzinfo=None) assert eff_naive >= before # --------------------------------------------------------------------------- # GET /api/energy/contracts/{id} # --------------------------------------------------------------------------- def test_get_contract_unauthenticated_returns_401(contracts_client): client, _ = contracts_client resp = client.get("/api/energy/contracts/1") assert resp.status_code == 401 def test_get_contract_not_found_returns_404(contracts_client): client, _ = contracts_client _login(client) resp = client.get("/api/energy/contracts/99999") assert resp.status_code == 404 def test_get_contract_returns_versions_ordered(contracts_client): """GET {id} must return versions ordered by effective_from.""" client, _ = contracts_client _login(client) t0 = datetime(2026, 1, 1, 0, 0, 0, tzinfo=UTC) resp = client.post( "/api/energy/contracts", json=_manual_payload(effective_from=t0.isoformat()), headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 201 contract_id = resp.json()["id"] t1 = datetime(2026, 6, 1, 0, 0, 0, tzinfo=UTC) new_values = dict(_MANUAL_VALUES) new_values = { **_MANUAL_VALUES, "energy": {**_MANUAL_VALUES["energy"], "buy": {"normal": 0.50, "dal": 0.40}}, } client.post( f"/api/energy/contracts/{contract_id}/versions", json={"effective_from": t1.isoformat(), "values": new_values}, headers={"X-CSRF-Token": _CSRF}, ) resp = client.get(f"/api/energy/contracts/{contract_id}") assert resp.status_code == 200 versions = resp.json()["versions"] assert len(versions) == 2 # Ordered by effective_from ascending eff0 = datetime.fromisoformat(versions[0]["effective_from"]) eff1 = datetime.fromisoformat(versions[1]["effective_from"]) assert eff0 < eff1 # --------------------------------------------------------------------------- # PATCH /api/energy/contracts/{id} # --------------------------------------------------------------------------- def test_patch_contract_unauthenticated_returns_401(contracts_client): client, _ = contracts_client resp = client.patch( "/api/energy/contracts/1", json={"name": "Renamed"}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 401 def test_patch_contract_missing_csrf_returns_403(contracts_client): client, _ = contracts_client _login(client) # Create a contract first resp = client.post( "/api/energy/contracts", json=_manual_payload(), headers={"X-CSRF-Token": _CSRF}, ) contract_id = resp.json()["id"] # PATCH without CSRF resp = client.patch(f"/api/energy/contracts/{contract_id}", json={"name": "Renamed"}) assert resp.status_code == 403 def test_patch_contract_not_found_returns_404(contracts_client): client, _ = contracts_client _login(client) resp = client.patch( "/api/energy/contracts/99999", json={"name": "Does Not Exist"}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 404 def test_patch_contract_rename(contracts_client): client, _ = contracts_client _login(client) resp = client.post( "/api/energy/contracts", json=_manual_payload(name="Original Name"), headers={"X-CSRF-Token": _CSRF}, ) contract_id = resp.json()["id"] resp = client.patch( f"/api/energy/contracts/{contract_id}", json={"name": "Renamed Contract"}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 200 assert resp.json()["name"] == "Renamed Contract" def test_activate_contract_mutual_exclusion(contracts_client): """Activating B deactivates A; at most one contract is active.""" client, engine = contracts_client _login(client) # Create A and B resp_a = client.post( "/api/energy/contracts", json=_manual_payload(name="A"), headers={"X-CSRF-Token": _CSRF}, ) id_a = resp_a.json()["id"] resp_b = client.post( "/api/energy/contracts", json=_manual_payload(name="B"), headers={"X-CSRF-Token": _CSRF}, ) id_b = resp_b.json()["id"] # Activate A resp = client.patch( f"/api/energy/contracts/{id_a}", json={"active": True}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 200 assert resp.json()["active"] is True # Now activate B — A must become inactive resp = client.patch( f"/api/energy/contracts/{id_b}", json={"active": True}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 200 assert resp.json()["active"] is True # Verify A is now inactive via list list_resp = client.get("/api/energy/contracts") items = {c["id"]: c for c in list_resp.json()["items"]} assert items[id_a]["active"] is False assert items[id_b]["active"] is True # DB check: only one active with Session(engine) as s: active_contracts = ( s.execute(select(EnergyContract).where(EnergyContract.active.is_(True))) .scalars() .all() ) assert len(active_contracts) == 1 assert active_contracts[0].id == id_b def test_deactivate_contract(contracts_client): """PATCH active=false deactivates the contract without touching others.""" client, _ = contracts_client _login(client) resp = client.post( "/api/energy/contracts", json=_manual_payload(), headers={"X-CSRF-Token": _CSRF}, ) contract_id = resp.json()["id"] # Activate it first client.patch( f"/api/energy/contracts/{contract_id}", json={"active": True}, headers={"X-CSRF-Token": _CSRF}, ) # Now deactivate resp = client.patch( f"/api/energy/contracts/{contract_id}", json={"active": False}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 200 assert resp.json()["active"] is False # --------------------------------------------------------------------------- # POST /api/energy/contracts/{id}/versions # --------------------------------------------------------------------------- def test_add_version_unauthenticated_returns_401(contracts_client): client, _ = contracts_client resp = client.post( "/api/energy/contracts/1/versions", json={"effective_from": datetime.now(UTC).isoformat(), "values": _MANUAL_VALUES}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 401 def test_add_version_missing_csrf_returns_403(contracts_client): client, _ = contracts_client _login(client) resp_c = client.post( "/api/energy/contracts", json=_manual_payload(), headers={"X-CSRF-Token": _CSRF}, ) cid = resp_c.json()["id"] # POST version without CSRF t1 = (datetime.now(UTC) + timedelta(days=1)).isoformat() resp = client.post( f"/api/energy/contracts/{cid}/versions", json={"effective_from": t1, "values": _MANUAL_VALUES}, ) assert resp.status_code == 403 def test_add_version_not_found_returns_404(contracts_client): client, _ = contracts_client _login(client) t1 = (datetime.now(UTC) + timedelta(days=1)).isoformat() resp = client.post( "/api/energy/contracts/99999/versions", json={"effective_from": t1, "values": _MANUAL_VALUES}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 404 def test_add_version_invalid_values_returns_422_no_db_write(contracts_client): """Non-conforming values for a new version must return 422; DB unchanged.""" client, engine = contracts_client _login(client) t0 = datetime(2026, 1, 1, tzinfo=UTC) resp_c = client.post( "/api/energy/contracts", json=_manual_payload(effective_from=t0.isoformat()), headers={"X-CSRF-Token": _CSRF}, ) cid = resp_c.json()["id"] bad_values = {"energy": {}} # missing almost everything t1 = datetime(2026, 6, 1, tzinfo=UTC) resp = client.post( f"/api/energy/contracts/{cid}/versions", json={"effective_from": t1.isoformat(), "values": bad_values}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 422 # Only the original version should exist in DB. with Session(engine) as s: versions = ( s.execute( select(EnergyContractVersion).where(EnergyContractVersion.contract_id == cid) ) .scalars() .all() ) assert len(versions) == 1 def test_add_version_effective_from_not_after_previous_returns_422(contracts_client): """effective_from ≤ previous open version's effective_from must return 422.""" client, _ = contracts_client _login(client) t0 = datetime(2026, 6, 1, tzinfo=UTC) resp_c = client.post( "/api/energy/contracts", json=_manual_payload(effective_from=t0.isoformat()), headers={"X-CSRF-Token": _CSRF}, ) cid = resp_c.json()["id"] # Attempt to add a version with t1 == t0 (not strictly after) resp = client.post( f"/api/energy/contracts/{cid}/versions", json={"effective_from": t0.isoformat(), "values": _MANUAL_VALUES}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 422 # Attempt to add a version with t1 < t0 t_before = datetime(2026, 1, 1, tzinfo=UTC) resp = client.post( f"/api/energy/contracts/{cid}/versions", json={"effective_from": t_before.isoformat(), "values": _MANUAL_VALUES}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 422 def test_add_version_closes_previous_and_appends(contracts_client): """Adding a new version must close the previous one and create a new open version. Verification: - Old version's effective_to == new version's effective_from. - Old version's values are unchanged (append-only; never overwritten). - New version has effective_to=None (open-ended). """ client, engine = contracts_client _login(client) t0 = datetime(2026, 1, 1, 0, 0, 0, tzinfo=UTC) old_values = _MANUAL_VALUES resp_c = client.post( "/api/energy/contracts", json=_manual_payload(effective_from=t0.isoformat(), values=old_values), headers={"X-CSRF-Token": _CSRF}, ) assert resp_c.status_code == 201 cid = resp_c.json()["id"] v1_id = resp_c.json()["versions"][0]["id"] t1 = datetime(2026, 6, 1, 0, 0, 0, tzinfo=UTC) new_values = { **old_values, "energy": {**old_values["energy"], "buy": {"normal": 0.50, "dal": 0.40}}, } resp = client.post( f"/api/energy/contracts/{cid}/versions", json={"effective_from": t1.isoformat(), "values": new_values}, headers={"X-CSRF-Token": _CSRF}, ) assert resp.status_code == 201 versions = resp.json()["versions"] assert len(versions) == 2 # Ordered ascending — first is the old version ver_old = versions[0] ver_new = versions[1] # Old version must be closed at t1. # SQLite round-trips datetimes as naive UTC strings; compare without tzinfo. assert ver_old["id"] == v1_id assert ver_old["effective_to"] is not None eff_to = datetime.fromisoformat(ver_old["effective_to"]).replace(tzinfo=None) assert eff_to == t1.replace(tzinfo=None) # Old version values must be unchanged assert ver_old["values"]["energy"]["buy"]["normal"] == old_values["energy"]["buy"]["normal"] # New version must be open-ended assert ver_new["effective_to"] is None assert ver_new["values"]["energy"]["buy"]["normal"] == 0.50 # Double-check via DB with Session(engine) as s: v1 = s.get(EnergyContractVersion, v1_id) assert v1 is not None assert v1.effective_to is not None assert v1.values["energy"]["buy"]["normal"] == old_values["energy"]["buy"]["normal"] all_versions = ( s.execute( select(EnergyContractVersion).where(EnergyContractVersion.contract_id == cid) ) .scalars() .all() ) assert len(all_versions) == 2