diff --git a/alembic_app/versions/20260822_18_contract_scopes.py b/alembic_app/versions/20260822_18_contract_scopes.py new file mode 100644 index 0000000..efe1fa1 --- /dev/null +++ b/alembic_app/versions/20260822_18_contract_scopes.py @@ -0,0 +1,107 @@ +"""add a billing scope to energy contracts + +Revision ID: 20260822_18_contract_scopes +Revises: 20260822_17_warmtelink_readings +Create Date: 2026-08-22 00:00:00.000000 + +The upgrade preserves every existing contract, version and cost row. Existing +contracts predate scopes and therefore deterministically belong to electricity. +""" + +from typing import Sequence, Union + +import sqlalchemy as sa +from alembic import op + + +revision: str = "20260822_18_contract_scopes" +down_revision: Union[str, None] = "20260822_17_warmtelink_readings" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def _count(connection: sa.Connection, table: str) -> int: + return int(connection.execute(sa.text(f"SELECT COUNT(*) FROM {table}")).scalar_one()) + + +def _orphan_count(connection: sa.Connection) -> int: + version_orphans = connection.execute( + sa.text( + "SELECT COUNT(*) FROM energy_contract_version v " + "LEFT JOIN energy_contract c ON c.id = v.contract_id WHERE c.id IS NULL" + ) + ).scalar_one() + cost_orphans = connection.execute( + sa.text( + "SELECT COUNT(*) FROM energy_cost_period p " + "LEFT JOIN energy_contract_version v ON v.id = p.contract_version_id " + "WHERE p.contract_version_id IS NOT NULL AND v.id IS NULL" + ) + ).scalar_one() + return int(version_orphans) + int(cost_orphans) + + +def _audit_scope_upgrade(connection: sa.Connection, before: dict[str, int], orphan_before: int) -> None: + after = {table: _count(connection, table) for table in before} + if after != before: + raise RuntimeError("contract scope migration row-count audit failed") + if _orphan_count(connection) != orphan_before: + raise RuntimeError("contract scope migration FK audit failed") + invalid_scope_count = connection.execute( + sa.text("SELECT COUNT(*) FROM energy_contract WHERE scope IS NULL OR scope != 'electricity'") + ).scalar_one() + if invalid_scope_count: + raise RuntimeError("contract scope migration backfill audit failed") + + # Kept on Alembic's Config attributes rather than an environment switch so + # isolated migration tests can deterministically exercise the rollback + # boundary without changing production behavior. + failure_injector = op.get_context().config.attributes.get("m8_t12_post_ddl_audit_failure") + if callable(failure_injector): + failure_injector() + + +def _apply_scope_schema() -> None: + # SQLite batch mode reconstructs the table. The server default gives every + # historical row its deterministic value during reconstruction. + with op.batch_alter_table("energy_contract", schema=None) as batch_op: + batch_op.add_column( + sa.Column("scope", sa.String(length=32), nullable=False, server_default="electricity") + ) + batch_op.create_index("ix_energy_contract_scope", ["scope"]) + + +def upgrade() -> None: + connection = op.get_bind() + before = { + table: _count(connection, table) + for table in ("energy_contract", "energy_contract_version", "energy_cost_period") + } + orphan_before = _orphan_count(connection) + + if connection.dialect.name != "sqlite": + _apply_scope_schema() + _audit_scope_upgrade(connection, before, orphan_before) + return + + # Alembic marks SQLite batch DDL as non-transactional. SQLite itself can + # nevertheless atomically roll back CREATE/COPY/DROP/RENAME when an + # explicit transaction owns the complete batch operation. Keep the audit + # inside that boundary so a failed audit cannot strand a revision-17 DB + # with a revision-18 table shape. + connection.exec_driver_sql("BEGIN IMMEDIATE") + try: + _apply_scope_schema() + _audit_scope_upgrade(connection, before, orphan_before) + except BaseException: + connection.exec_driver_sql("ROLLBACK") + raise + else: + connection.exec_driver_sql("COMMIT") + + +def downgrade() -> None: + # Schema reversibility is only exercised against isolated temporary test DBs. + with op.batch_alter_table("energy_contract", schema=None) as batch_op: + batch_op.drop_index("ix_energy_contract_scope") + batch_op.drop_column("scope") diff --git a/app/api/routes/api/energy_contracts.py b/app/api/routes/api/energy_contracts.py index b204c02..7a4630d 100644 --- a/app/api/routes/api/energy_contracts.py +++ b/app/api/routes/api/energy_contracts.py @@ -46,7 +46,9 @@ from app.schemas.energy_contract import ( from app.services.auth import AuthenticatedSession from app.services import timezone as _tz_mod from app.services.contracts import ( + CONTRACT_KIND_SCOPES, ContractVersionError, + ContractScopeError, activate_contract, add_version, create_contract, @@ -98,6 +100,7 @@ def _contract_detail(db: Session, contract) -> ContractDetailResponse: id=contract.id, name=contract.name, kind=contract.kind, + scope=contract.scope, active=contract.active, currency=contract.currency, created_at=contract.created_at, @@ -163,16 +166,22 @@ def get_profiles( @router.get("/contracts", response_model=ContractListResponse) def list_energy_contracts( + scope: str = "electricity", db: Session = Depends(get_db), _auth: AuthenticatedSession = Depends(require_session), ) -> ContractListResponse: """List all energy contracts with their active status. - Returns a flat list (no embedded version history); use + Scope defaults to ``electricity`` for old clients. Returns a flat list (no embedded version history); use GET /api/energy/contracts/{id} to fetch the full version history for a specific contract. """ - contracts = list_contracts(db) + if scope not in set(CONTRACT_KIND_SCOPES.values()): + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, + detail=f"Unknown energy contract scope: {scope!r}", + ) + contracts = list_contracts(db, scope=scope) items = [ContractResponse.model_validate(c) for c in contracts] return ContractListResponse(items=items, total=len(items)) @@ -208,10 +217,11 @@ def create_energy_contract( name=body.name, kind=body.kind, currency=body.currency, + scope=body.scope, values=body.values, effective_from=effective_from, ) - except (ProfileNotFoundError, ProfileValidationError) as exc: + except (ProfileNotFoundError, ProfileValidationError, ContractScopeError) as exc: _raise_422_for_profile_error(exc) db.commit() @@ -256,11 +266,11 @@ def patch_energy_contract( """Partially update a contract: rename or change activation status. - ``name``: updates the human-readable label. - - ``active=true``: activates this contract (all others are deactivated). + - ``active=true``: activates this contract (same-scope contracts are deactivated). - ``active=false``: deactivates this contract (no effect on others). - At most one contract may be active at any time; the service layer enforces - mutual exclusion. + At most one contract may be active per scope; the service layer enforces + scope-local mutual exclusion. """ contract = _get_contract_or_404(db, contract_id) diff --git a/app/models/energy.py b/app/models/energy.py index 25ab153..97fcc52 100644 --- a/app/models/energy.py +++ b/app/models/energy.py @@ -163,8 +163,10 @@ class EnergyContract(Base): ``kind`` determines which price strategy is used (``manual`` for fixed dual-tariff rates entered by the user, ``tibber`` for dynamic API prices). - Only one contract may be ``active`` at a time; the service layer enforces - mutual exclusion. Specific pricing values live in ``EnergyContractVersion`` + A contract belongs to an energy ``scope`` (currently electricity; thermal + profiles are reserved for the next milestone). Only one contract may be + ``active`` per scope; the service layer enforces mutual exclusion. Specific + pricing values live in ``EnergyContractVersion`` so that price changes can be tracked without modifying historical records. """ @@ -180,6 +182,12 @@ class EnergyContract(Base): # migration simple and the strategy registry extensible. kind: Mapped[str] = mapped_column(String(32), nullable=False) + # Billing domain. The service registry derives this from ``kind`` so API + # callers cannot move a pricing strategy into an incompatible domain. + scope: Mapped[str] = mapped_column( + String(32), nullable=False, default="electricity", index=True + ) + # Whether this is the currently active contract (at most one should be True). active: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) diff --git a/app/schemas/energy_contract.py b/app/schemas/energy_contract.py index 8dd63d4..3e94ada 100644 --- a/app/schemas/energy_contract.py +++ b/app/schemas/energy_contract.py @@ -51,6 +51,7 @@ class ContractResponse(BaseModel): id: int name: str kind: str + scope: str active: bool currency: str created_at: datetime @@ -69,6 +70,7 @@ class ContractDetailResponse(BaseModel): id: int name: str kind: str + scope: str active: bool currency: str created_at: datetime @@ -101,6 +103,7 @@ class ContractCreate(BaseModel): name: str = Field(..., min_length=1, max_length=255) kind: str = Field(..., min_length=1, max_length=32) + scope: str | None = Field(default=None, min_length=1, max_length=32) currency: str = Field(default="EUR", min_length=1, max_length=8) values: dict[str, Any] effective_from: datetime | None = Field( diff --git a/app/services/contracts.py b/app/services/contracts.py index d8396d7..aca7d4b 100644 --- a/app/services/contracts.py +++ b/app/services/contracts.py @@ -11,8 +11,8 @@ Design decisions setting its ``effective_to`` to the new version's ``effective_from``; raises ``ContractVersionError`` if the new date is strictly earlier than the previous version's ``effective_from``. -- ``activate_contract``: mutual-exclusion; sets all other contracts' ``active`` - to False, then sets the given contract's ``active`` to True. +- ``activate_contract``: scope-local mutual exclusion; sets other contracts in + the target scope inactive, then sets the given contract active. - ``active_contract_version_at``: returns the single version of the currently active contract that covers *ts* (``effective_from ≤ ts < effective_to``, or open-ended when ``effective_to`` is None). @@ -32,7 +32,7 @@ import logging from datetime import UTC, datetime from typing import Any -from sqlalchemy import select +from sqlalchemy import select, update from sqlalchemy.orm import Session from app.integrations.pricing.profiles import validate_values @@ -41,6 +41,32 @@ from app.models.energy import EnergyContract, EnergyContractVersion logger = logging.getLogger(__name__) +# This is deliberately separate from the pricing-profile loader. T12 needs to +# reserve the thermal domain before T13 supplies its actual profile. +CONTRACT_KIND_SCOPES: dict[str, str] = { + "manual": "electricity", + "tibber": "electricity", + "district_heating": "thermal", +} + + +class ContractScopeError(ValueError): + """Raised when a contract kind is unknown or its supplied scope disagrees.""" + + +def contract_scope_for_kind(kind: str, requested_scope: str | None = None) -> str: + """Return the registry-owned scope for *kind*, rejecting client mismatches.""" + try: + scope = CONTRACT_KIND_SCOPES[kind] + except KeyError as exc: + raise ContractScopeError(f"Unknown energy contract kind: {kind!r}") from exc + if requested_scope is not None and requested_scope != scope: + raise ContractScopeError( + f"Contract kind {kind!r} belongs to scope {scope!r}, not {requested_scope!r}." + ) + return scope + + # --------------------------------------------------------------------------- # Internal helpers # --------------------------------------------------------------------------- @@ -86,10 +112,14 @@ def get_contract_or_none(session: Session, contract_id: int) -> EnergyContract | ).scalar_one_or_none() -def list_contracts(session: Session) -> list[EnergyContract]: - """Return all contracts ordered by id (ascending).""" +def list_contracts(session: Session, *, scope: str = "electricity") -> list[EnergyContract]: + """Return contracts in one scope, ordered by id (ascending).""" return list( - session.execute(select(EnergyContract).order_by(EnergyContract.id)).scalars().all() + session.execute( + select(EnergyContract) + .where(EnergyContract.scope == scope) + .order_by(EnergyContract.id) + ).scalars().all() ) @@ -117,6 +147,7 @@ def create_contract( name: str, kind: str, currency: str = "EUR", + scope: str | None = None, values: dict[str, Any], effective_from: datetime, ) -> EnergyContract: @@ -150,6 +181,7 @@ def create_contract( ProfileValidationError If *values* does not conform to the profile structure. """ + resolved_scope = contract_scope_for_kind(kind, scope) # Validate (and fill defaults) before any DB write. filled_values = validate_values(kind, values) @@ -157,6 +189,7 @@ def create_contract( contract = EnergyContract( name=name, kind=kind, + scope=resolved_scope, currency=currency, active=False, # New contracts are inactive; caller must explicitly activate. created_at=now, @@ -262,15 +295,18 @@ def add_version( def activate_contract(session: Session, contract: EnergyContract) -> None: """Activate a contract with mutual exclusion. - Sets every other contract's ``active`` flag to False, then sets the given - contract's ``active`` to True. This guarantees at most one active contract - at any time. + Sets every other contract in the same scope inactive, then sets the given + contract active. This guarantees at most one active contract per scope. Caller must commit after this returns. """ - # Deactivate all contracts (including the target; we re-activate below). - for other in session.execute(select(EnergyContract)).scalars().all(): - other.active = False + # This bulk update is a single write statement inside the caller's + # transaction. SQLite serializes writers, and another scope is never touched. + session.execute( + update(EnergyContract) + .where(EnergyContract.scope == contract.scope, EnergyContract.id != contract.id) + .values(active=False) + ) contract.active = True contract.updated_at = datetime.now(UTC) logger.info("Activated contract %r (id=%d)", contract.name, contract.id) @@ -286,7 +322,9 @@ def deactivate_contract(session: Session, contract: EnergyContract) -> None: logger.info("Deactivated contract %r (id=%d)", contract.name, contract.id) -def active_contract_versions(session: Session) -> list[EnergyContractVersion]: +def active_contract_versions( + session: Session, *, scope: str = "electricity" +) -> list[EnergyContractVersion]: """Return all versions of the currently active contract, ordered by effective_from ascending. Returns an empty list when there is no active contract. The list spans the @@ -295,7 +333,9 @@ def active_contract_versions(session: Session) -> list[EnergyContractVersion]: cost / credit accumulation (Principle C). """ active = session.execute( - select(EnergyContract).where(EnergyContract.active.is_(True)).limit(1) + select(EnergyContract) + .where(EnergyContract.active.is_(True), EnergyContract.scope == scope) + .limit(1) ).scalar_one_or_none() if active is None: @@ -313,7 +353,7 @@ def active_contract_versions(session: Session) -> list[EnergyContractVersion]: def active_contract_version_at( - session: Session, ts: datetime + session: Session, ts: datetime, *, scope: str = "electricity" ) -> EnergyContractVersion | None: """Return the active contract's version that covers *ts*. @@ -336,7 +376,9 @@ def active_contract_version_at( EnergyContractVersion | None """ active = session.execute( - select(EnergyContract).where(EnergyContract.active.is_(True)).limit(1) + select(EnergyContract) + .where(EnergyContract.active.is_(True), EnergyContract.scope == scope) + .limit(1) ).scalar_one_or_none() if active is None: diff --git a/docs/design/m8-warmtelink-energy.md b/docs/design/m8-warmtelink-energy.md index 3d475f1..487b5ee 100644 --- a/docs/design/m8-warmtelink-energy.md +++ b/docs/design/m8-warmtelink-energy.md @@ -727,7 +727,7 @@ T01~T06 先把现有 DSMR 安全迁到统一 source/binding;T07~T11 再接 ### M8-T12 — 合同 Scope 与按 Scope 激活 [structural] -- **Status**: `todo` +- **Status**: `done` - **Depends**: M8-T08 - **Context**: electricity 与 thermal 必须能各有一份 active 合同,同时保持旧客户端默认看 electricity。 @@ -740,6 +740,7 @@ T01~T06 先把现有 DSMR 安全迁到统一 source/binding;T07~T11 再接 - `modify scripts/app_db_adopt.py` - `modify tests/test_api_energy_contracts.py` - `modify tests/test_energy_models.py` +- `modify tests/test_warmtelink_models.py` - `modify openapi/openapi.json` - `modify openapi/openapi.yaml` - `modify frontend/src/api/schema.d.ts` diff --git a/frontend/src/api/schema.d.ts b/frontend/src/api/schema.d.ts index 85f9ae9..0990a24 100644 --- a/frontend/src/api/schema.d.ts +++ b/frontend/src/api/schema.d.ts @@ -472,7 +472,7 @@ export interface paths { * List Energy Contracts * @description List all energy contracts with their active status. * - * Returns a flat list (no embedded version history); use + * Scope defaults to ``electricity`` for old clients. Returns a flat list (no embedded version history); use * GET /api/energy/contracts/{id} to fetch the full version history for a * specific contract. */ @@ -519,11 +519,11 @@ export interface paths { * @description Partially update a contract: rename or change activation status. * * - ``name``: updates the human-readable label. - * - ``active=true``: activates this contract (all others are deactivated). + * - ``active=true``: activates this contract (same-scope contracts are deactivated). * - ``active=false``: deactivates this contract (no effect on others). * - * At most one contract may be active at any time; the service layer enforces - * mutual exclusion. + * At most one contract may be active per scope; the service layer enforces + * scope-local mutual exclusion. */ patch: operations["patch_energy_contract_api_energy_contracts__contract_id__patch"]; trace?: never; @@ -1549,6 +1549,8 @@ export interface components { name: string; /** Kind */ kind: string; + /** Scope */ + scope?: string | null; /** * Currency * @default EUR @@ -1578,6 +1580,8 @@ export interface components { name: string; /** Kind */ kind: string; + /** Scope */ + scope: string; /** Active */ active: boolean; /** Currency */ @@ -1632,6 +1636,8 @@ export interface components { name: string; /** Kind */ kind: string; + /** Scope */ + scope: string; /** Active */ active: boolean; /** Currency */ @@ -3494,7 +3500,9 @@ export interface operations { }; list_energy_contracts_api_energy_contracts_get: { parameters: { - query?: never; + query?: { + scope?: string; + }; header?: never; path?: never; cookie?: never; @@ -3510,6 +3518,15 @@ export interface operations { "application/json": components["schemas"]["ContractListResponse"]; }; }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; }; }; create_energy_contract_api_energy_contracts_post: { diff --git a/openapi/openapi.json b/openapi/openapi.json index fb47f40..2705bc6 100644 --- a/openapi/openapi.json +++ b/openapi/openapi.json @@ -1126,8 +1126,20 @@ "api-energy-contracts" ], "summary": "List Energy Contracts", - "description": "List all energy contracts with their active status.\n\nReturns a flat list (no embedded version history); use\nGET /api/energy/contracts/{id} to fetch the full version history for a\nspecific contract.", + "description": "List all energy contracts with their active status.\n\nScope defaults to ``electricity`` for old clients. Returns a flat list (no embedded version history); use\nGET /api/energy/contracts/{id} to fetch the full version history for a\nspecific contract.", "operationId": "list_energy_contracts_api_energy_contracts_get", + "parameters": [ + { + "name": "scope", + "in": "query", + "required": false, + "schema": { + "type": "string", + "default": "electricity", + "title": "Scope" + } + } + ], "responses": { "200": { "description": "Successful Response", @@ -1138,6 +1150,16 @@ } } } + }, + "422": { + "description": "Validation Error", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + } } } }, @@ -1247,7 +1269,7 @@ "api-energy-contracts" ], "summary": "Patch Energy Contract", - "description": "Partially update a contract: rename or change activation status.\n\n- ``name``: updates the human-readable label.\n- ``active=true``: activates this contract (all others are deactivated).\n- ``active=false``: deactivates this contract (no effect on others).\n\nAt most one contract may be active at any time; the service layer enforces\nmutual exclusion.", + "description": "Partially update a contract: rename or change activation status.\n\n- ``name``: updates the human-readable label.\n- ``active=true``: activates this contract (same-scope contracts are deactivated).\n- ``active=false``: deactivates this contract (no effect on others).\n\nAt most one contract may be active per scope; the service layer enforces\nscope-local mutual exclusion.", "operationId": "patch_energy_contract_api_energy_contracts__contract_id__patch", "parameters": [ { @@ -3705,6 +3727,19 @@ "minLength": 1, "title": "Kind" }, + "scope": { + "anyOf": [ + { + "type": "string", + "maxLength": 32, + "minLength": 1 + }, + { + "type": "null" + } + ], + "title": "Scope" + }, "currency": { "type": "string", "maxLength": 8, @@ -3754,6 +3789,10 @@ "type": "string", "title": "Kind" }, + "scope": { + "type": "string", + "title": "Scope" + }, "active": { "type": "boolean", "title": "Active" @@ -3785,6 +3824,7 @@ "id", "name", "kind", + "scope", "active", "currency", "created_at", @@ -3861,6 +3901,10 @@ "type": "string", "title": "Kind" }, + "scope": { + "type": "string", + "title": "Scope" + }, "active": { "type": "boolean", "title": "Active" @@ -3885,6 +3929,7 @@ "id", "name", "kind", + "scope", "active", "currency", "created_at", diff --git a/openapi/openapi.yaml b/openapi/openapi.yaml index e85d463..e5b0029 100644 --- a/openapi/openapi.yaml +++ b/openapi/openapi.yaml @@ -877,12 +877,21 @@ paths: description: 'List all energy contracts with their active status. - Returns a flat list (no embedded version history); use + Scope defaults to ``electricity`` for old clients. Returns a flat list (no + embedded version history); use GET /api/energy/contracts/{id} to fetch the full version history for a specific contract.' operationId: list_energy_contracts_api_energy_contracts_get + parameters: + - name: scope + in: query + required: false + schema: + type: string + default: electricity + title: Scope responses: '200': description: Successful Response @@ -890,6 +899,12 @@ paths: application/json: schema: $ref: '#/components/schemas/ContractListResponse' + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' post: tags: - api-energy-contracts @@ -974,14 +989,14 @@ paths: - ``name``: updates the human-readable label. - - ``active=true``: activates this contract (all others are deactivated). + - ``active=true``: activates this contract (same-scope contracts are deactivated). - ``active=false``: deactivates this contract (no effect on others). - At most one contract may be active at any time; the service layer enforces + At most one contract may be active per scope; the service layer enforces - mutual exclusion.' + scope-local mutual exclusion.' operationId: patch_energy_contract_api_energy_contracts__contract_id__patch parameters: - name: contract_id @@ -2768,6 +2783,13 @@ components: maxLength: 32 minLength: 1 title: Kind + scope: + anyOf: + - type: string + maxLength: 32 + minLength: 1 + - type: 'null' + title: Scope currency: type: string maxLength: 8 @@ -2813,6 +2835,9 @@ components: kind: type: string title: Kind + scope: + type: string + title: Scope active: type: boolean title: Active @@ -2837,6 +2862,7 @@ components: - id - name - kind + - scope - active - currency - created_at @@ -2901,6 +2927,9 @@ components: kind: type: string title: Kind + scope: + type: string + title: Scope active: type: boolean title: Active @@ -2920,6 +2949,7 @@ components: - id - name - kind + - scope - active - currency - created_at diff --git a/scripts/app_db_adopt.py b/scripts/app_db_adopt.py index beb9943..5a3d218 100644 --- a/scripts/app_db_adopt.py +++ b/scripts/app_db_adopt.py @@ -15,7 +15,7 @@ if str(PROJECT_ROOT) not in sys.path: from app.config import get_settings -APP_BASELINE_REVISION = "20260822_17_warmtelink_readings" +APP_BASELINE_REVISION = "20260822_18_contract_scopes" class AppDatabaseAdoptionError(RuntimeError): diff --git a/tests/test_api_energy_contracts.py b/tests/test_api_energy_contracts.py index a2c8116..e33af6d 100644 --- a/tests/test_api_energy_contracts.py +++ b/tests/test_api_energy_contracts.py @@ -55,6 +55,7 @@ from sqlalchemy import create_engine, select from sqlalchemy.orm import Session from app.models.energy import EnergyContract, EnergyContractVersion +from app.services.contracts import activate_contract # --------------------------------------------------------------------------- # Shared helpers @@ -567,6 +568,72 @@ def test_deactivate_contract(contracts_client): assert resp.json()["active"] is False +def test_scope_defaults_filtering_and_kind_mismatch(contracts_client): + """Old clients default to electricity; a supplied incompatible scope is rejected.""" + client, engine = contracts_client + _login(client) + created = client.post( + "/api/energy/contracts", + json=_manual_payload(), + headers={"X-CSRF-Token": _CSRF}, + ) + assert created.status_code == 201 + assert created.json()["scope"] == "electricity" + assert client.get("/api/energy/contracts").json()["total"] == 1 + assert client.get("/api/energy/contracts?scope=thermal").json()["items"] == [] + + mismatch = client.post( + "/api/energy/contracts", + json=_manual_payload(scope="thermal"), + headers={"X-CSRF-Token": _CSRF}, + ) + assert mismatch.status_code == 422 + with Session(engine) as session: + assert len(session.execute(select(EnergyContract)).scalars().all()) == 1 + + +def test_activation_is_scope_local_and_transaction_rollback_is_safe(contracts_client): + """A thermal activation neither deactivates electricity nor survives rollback.""" + client, engine = contracts_client + _login(client) + electricity = client.post( + "/api/energy/contracts", + json=_manual_payload(name="Electricity"), + headers={"X-CSRF-Token": _CSRF}, + ).json() + client.patch( + f"/api/energy/contracts/{electricity['id']}", + json={"active": True}, + headers={"X-CSRF-Token": _CSRF}, + ) + now = datetime.now(UTC) + with Session(engine) as session: + thermal = EnergyContract( + name="Future thermal", kind="district_heating", scope="thermal", active=False, + currency="EUR", created_at=now, updated_at=now, + ) + session.add(thermal) + session.commit() + thermal_id = thermal.id + + with Session(engine) as session: + thermal = session.get(EnergyContract, thermal_id) + assert thermal is not None + activate_contract(session, thermal) + session.rollback() # Simulate a later write failure in this transaction. + + with Session(engine) as session: + rows = {row.scope: row for row in session.execute(select(EnergyContract)).scalars()} + assert rows["electricity"].active is True + assert rows["thermal"].active is False + activate_contract(session, rows["thermal"]) + session.commit() + + with Session(engine) as session: + active = session.execute(select(EnergyContract).where(EnergyContract.active.is_(True))).scalars().all() + assert {row.scope for row in active} == {"electricity", "thermal"} + + # --------------------------------------------------------------------------- # POST /api/energy/contracts/{id}/versions # --------------------------------------------------------------------------- diff --git a/tests/test_energy_models.py b/tests/test_energy_models.py index 25a0a16..6db0890 100644 --- a/tests/test_energy_models.py +++ b/tests/test_energy_models.py @@ -157,11 +157,15 @@ def test_energy_contract_columns(energy_db): inspector = inspect(energy_db) columns = {col["name"]: col for col in inspector.get_columns("energy_contract")} - required_non_nullable = {"id", "name", "kind", "active", "currency", "created_at", "updated_at"} + required_non_nullable = { + "id", "name", "kind", "scope", "active", "currency", "created_at", "updated_at" + } for col_name in required_non_nullable: assert col_name in columns, f"Missing column: {col_name}" assert not columns[col_name]["nullable"], f"{col_name} should be NOT NULL" + assert any(index["name"] == "ix_energy_contract_scope" for index in inspector.get_indexes("energy_contract")) + def test_energy_contract_version_columns(energy_db): """energy_contract_version must have all required columns with correct nullability.""" @@ -537,6 +541,7 @@ def test_energy_contract_insert_and_retrieve(energy_db): assert fetched is not None assert fetched.name == "My Manual Contract" assert fetched.kind == "manual" + assert fetched.scope == "electricity" assert fetched.active is True assert fetched.currency == "EUR" @@ -1208,3 +1213,157 @@ def test_migration_downgrade_removes_meter_table(tmp_path: Path): "meter_id must be removed from energy_cost_period after downgrade" ) engine.dispose() + + +def test_contract_scope_migration_preserves_historical_contract_audit(tmp_path: Path): + """A revision-17 fixture upgrades/downgrades without altering contract audit rows.""" + db_url = f"sqlite:///{tmp_path / 'contract_scope_history.db'}" + cfg = _make_app_alembic_config(db_url) + command.upgrade(cfg, "20260822_17_warmtelink_readings") + engine = create_engine(db_url, connect_args={"check_same_thread": False}) + now = datetime.now(tz=timezone.utc).replace(tzinfo=None) + values = '{"energy":{"buy":{"normal":0.4}}}' + pricing = '{"historic":"unchanged"}' + with engine.begin() as connection: + contract_id = connection.execute( + text( + "INSERT INTO energy_contract (name, kind, active, currency, created_at, updated_at) " + "VALUES ('Historic', 'manual', 1, 'EUR', :now, :now)" + ), + {"now": now}, + ).lastrowid + version_id = connection.execute( + text( + "INSERT INTO energy_contract_version " + "(contract_id, effective_from, effective_to, \"values\", created_at) " + "VALUES (:contract_id, :now, NULL, :values, :now)" + ), + {"contract_id": contract_id, "now": now, "values": values}, + ).lastrowid + connection.execute( + text( + "INSERT INTO energy_cost_period " + "(period_start, d1_kwh, d2_kwh, r1_kwh, r2_kwh, import_cost, export_revenue, " + "net_cost, currency, pricing, contract_version_id, degraded, computed_at, meter_id, " + "source_binding_id) VALUES (:now, 1, 2, 0, 0, 3, 0, 3, 'EUR', :pricing, :version_id, " + "0, :now, NULL, NULL)" + ), + {"now": now, "pricing": pricing, "version_id": version_id}, + ) + + command.upgrade(cfg, "head") + command.upgrade(cfg, "head") + with engine.connect() as connection: + assert connection.execute(text("SELECT version_num FROM alembic_version")).scalar_one() == ( + "20260822_18_contract_scopes" + ) + assert connection.execute(text("SELECT scope FROM energy_contract")).scalar_one() == "electricity" + assert connection.execute(text("SELECT COUNT(*) FROM energy_contract")).scalar_one() == 1 + assert connection.execute(text("SELECT COUNT(*) FROM energy_contract_version")).scalar_one() == 1 + assert connection.execute(text("SELECT COUNT(*) FROM energy_cost_period")).scalar_one() == 1 + assert connection.execute(text("SELECT \"values\" FROM energy_contract_version")).scalar_one() == values + assert connection.execute(text("SELECT pricing FROM energy_cost_period")).scalar_one() == pricing + assert connection.execute( + text( + "SELECT COUNT(*) FROM energy_contract_version v LEFT JOIN energy_contract c " + "ON c.id = v.contract_id WHERE c.id IS NULL" + ) + ).scalar_one() == 0 + inspector = inspect(connection) + assert any(item["name"] == "ix_energy_contract_scope" for item in inspector.get_indexes("energy_contract")) + + command.downgrade(cfg, "20260822_17_warmtelink_readings") + with engine.connect() as connection: + assert "scope" not in {item["name"] for item in inspect(connection).get_columns("energy_contract")} + assert connection.execute(text("SELECT COUNT(*) FROM energy_contract_version")).scalar_one() == 1 + assert connection.execute(text("SELECT COUNT(*) FROM energy_cost_period")).scalar_one() == 1 + engine.dispose() + + +def test_contract_scope_migration_audit_failure_restores_revision_17(tmp_path: Path): + """A post-DDL audit failure leaves no SQLite batch-migration residue.""" + db_url = f"sqlite:///{tmp_path / 'contract_scope_audit_failure.db'}" + cfg = _make_app_alembic_config(db_url) + command.upgrade(cfg, "20260822_17_warmtelink_readings") + engine = _engine_with_fk(db_url) + now = datetime.now(tz=timezone.utc).replace(tzinfo=None) + values = '{"energy":{"buy":{"normal":0.4}}}' + pricing = '{"historic":"unchanged"}' + with engine.begin() as connection: + contract_id = connection.execute( + text( + "INSERT INTO energy_contract (name, kind, active, currency, created_at, updated_at) " + "VALUES ('Historic', 'manual', 1, 'EUR', :now, :now)" + ), + {"now": now}, + ).lastrowid + version_id = connection.execute( + text( + "INSERT INTO energy_contract_version " + "(contract_id, effective_from, effective_to, \"values\", created_at) " + "VALUES (:contract_id, :now, NULL, :values, :now)" + ), + {"contract_id": contract_id, "now": now, "values": values}, + ).lastrowid + connection.execute( + text( + "INSERT INTO energy_cost_period " + "(period_start, d1_kwh, d2_kwh, r1_kwh, r2_kwh, import_cost, export_revenue, " + "net_cost, currency, pricing, contract_version_id, degraded, computed_at, meter_id, " + "source_binding_id) VALUES (:now, 1, 2, 0, 0, 3, 0, 3, 'EUR', :pricing, :version_id, " + "0, :now, NULL, NULL)" + ), + {"now": now, "pricing": pricing, "version_id": version_id}, + ) + engine.dispose() + + def _raise_after_ddl() -> None: + raise RuntimeError("injected post-DDL audit failure") + + cfg.attributes["m8_t12_post_ddl_audit_failure"] = _raise_after_ddl + with pytest.raises(RuntimeError, match="injected post-DDL audit failure"): + command.upgrade(cfg, "head") + + engine = _engine_with_fk(db_url) + with engine.connect() as connection: + assert connection.execute(text("SELECT version_num FROM alembic_version")).scalar_one() == ( + "20260822_17_warmtelink_readings" + ) + assert "scope" not in {item["name"] for item in inspect(connection).get_columns("energy_contract")} + assert not any( + item["name"] == "ix_energy_contract_scope" + for item in inspect(connection).get_indexes("energy_contract") + ) + assert connection.execute(text("SELECT COUNT(*) FROM energy_contract")).scalar_one() == 1 + assert connection.execute(text("SELECT COUNT(*) FROM energy_contract_version")).scalar_one() == 1 + assert connection.execute(text("SELECT COUNT(*) FROM energy_cost_period")).scalar_one() == 1 + assert connection.execute(text("SELECT \"values\" FROM energy_contract_version")).scalar_one() == values + assert connection.execute(text("SELECT pricing FROM energy_cost_period")).scalar_one() == pricing + assert connection.execute( + text( + "SELECT COUNT(*) FROM energy_contract_version v LEFT JOIN energy_contract c " + "ON c.id = v.contract_id WHERE c.id IS NULL" + ) + ).scalar_one() == 0 + assert connection.execute( + text( + "SELECT COUNT(*) FROM energy_cost_period p LEFT JOIN energy_contract_version v " + "ON v.id = p.contract_version_id " + "WHERE p.contract_version_id IS NOT NULL AND v.id IS NULL" + ) + ).scalar_one() == 0 + assert connection.execute(text("PRAGMA foreign_key_check")).all() == [] + assert connection.execute( + text("SELECT name FROM sqlite_master WHERE name LIKE '_alembic_tmp_%'") + ).all() == [] + engine.dispose() + + del cfg.attributes["m8_t12_post_ddl_audit_failure"] + command.upgrade(cfg, "head") + engine = _engine_with_fk(db_url) + with engine.connect() as connection: + assert connection.execute(text("SELECT version_num FROM alembic_version")).scalar_one() == ( + "20260822_18_contract_scopes" + ) + assert connection.execute(text("SELECT scope FROM energy_contract")).scalar_one() == "electricity" + engine.dispose() diff --git a/tests/test_warmtelink_models.py b/tests/test_warmtelink_models.py index 5618533..d61472c 100644 --- a/tests/test_warmtelink_models.py +++ b/tests/test_warmtelink_models.py @@ -14,8 +14,6 @@ from sqlalchemy import create_engine, event, inspect, text from sqlalchemy.orm import Session from app.models.meter_source import MeterSourceChannel, WarmteLinkReading -from scripts.app_db_adopt import APP_BASELINE_REVISION - REVISION_16 = "20260822_16_dsmr_source_adoption" REVISION_17 = "20260822_17_warmtelink_readings" @@ -78,8 +76,8 @@ def _insert_reading(connection, channel_id: int, timestamp: datetime, value: Dec def test_empty_database_and_revision_16_upgrade_are_additive_and_idempotent(tmp_path: Path): empty_url = f"sqlite:///{tmp_path / 'warmtelink_empty.db'}" empty_config = _config(empty_url) - command.upgrade(empty_config, "head") - command.upgrade(empty_config, "head") + command.upgrade(empty_config, REVISION_17) + command.upgrade(empty_config, REVISION_17) empty_engine = _engine(empty_url) try: with empty_engine.connect() as connection: @@ -106,8 +104,8 @@ def test_empty_database_and_revision_16_upgrade_are_additive_and_idempotent(tmp_ finally: engine.dispose() - command.upgrade(config, "head") - command.upgrade(config, "head") + command.upgrade(config, REVISION_17) + command.upgrade(config, REVISION_17) engine = _engine(database_url) try: with engine.connect() as connection: @@ -125,7 +123,7 @@ def test_empty_database_and_revision_16_upgrade_are_additive_and_idempotent(tmp_ def test_warmtelink_reading_constraints_indexes_and_decimal_round_trip(tmp_path: Path): database_url = f"sqlite:///{tmp_path / 'warmtelink_constraints.db'}" config = _config(database_url) - command.upgrade(config, "head") + command.upgrade(config, REVISION_17) timestamp = datetime(2026, 8, 22, 10, 30, tzinfo=timezone.utc) engine = _engine(database_url) try: @@ -233,13 +231,12 @@ def test_warmtelink_reading_model_uses_restrictive_relationship_and_aware_column assert "delete-orphan" not in relationship.cascade assert WarmteLinkReading.__table__.c.recorded_at.type.timezone is True assert WarmteLinkReading.__table__.c.received_at.type.timezone is True - assert APP_BASELINE_REVISION == REVISION_17 def test_warmtelink_reading_downgrade_is_schema_only_on_temporary_database(tmp_path: Path): database_url = f"sqlite:///{tmp_path / 'warmtelink_downgrade.db'}" config = _config(database_url) - command.upgrade(config, "head") + command.upgrade(config, REVISION_17) command.downgrade(config, REVISION_16) engine = _engine(database_url)