751 lines
23 KiB
Python
751 lines
23 KiB
Python
"""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
|