Commit 28767629 authored by Sergio Gimenez's avatar Sergio Gimenez
Browse files

Merge branch 'feat/persistence-hardening' into 'develop'

Small fixes on the persistent model.

See merge request !19
parents b75c8efb 1879b265
Loading
Loading
Loading
Loading
Loading
+73 −4
Original line number Diff line number Diff line
from sqlalchemy import (
    CHAR,
    Boolean,
    CheckConstraint,
    Column,
    DateTime,
    ForeignKey,
@@ -19,6 +20,16 @@ from sqlalchemy.dialects.postgresql import UUID as PGUUID

metadata = MetaData()

DIRECTIONS = ("inbound", "outbound")


def _one_of(column: str, *allowed: str, nullable: bool = False) -> str:
    """CHECK body keeping a state column inside its documented values (RD §K.1)."""
    membership = f"{column} IN ({', '.join(repr(value) for value in allowed)})"

    return f"{column} IS NULL OR {membership}" if nullable else membership


# Subset of fm_db.partner_ops (RD §K.1): columns grow as features need them.
partner_ops = Table(
    "partner_ops",
@@ -32,9 +43,19 @@ partner_ops = Table(
    Column("token_endpoint", Text),
    Column("status", String(20), nullable=False, server_default="pending"),
    Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
    Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
    Column(
        "updated_at",
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
        onupdate=func.now(),
    ),
    Index("idx_partner_ops_oauth2_client_id", "oauth2_client_id"),
    Index("idx_partner_ops_status", "status"),
    CheckConstraint(
        _one_of("status", "pending", "active", "suspended", "decommissioned"),
        name="ck_partner_ops_status",
    ),
)

federation_agreements = Table(
@@ -50,9 +71,19 @@ federation_agreements = Table(
    Column("valid_until", DateTime(timezone=True)),
    Column("status", String(20), nullable=False, server_default="draft"),
    Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
    Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
    Column(
        "updated_at",
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
        onupdate=func.now(),
    ),
    Index("idx_federation_agreements_partner", "partner_op_id"),
    Index("idx_federation_agreements_status_validity", "status", "valid_from", "valid_until"),
    CheckConstraint(
        _one_of("status", "draft", "active", "suspended", "expired"),
        name="ck_federation_agreements_status",
    ),
)

federation_contexts = Table(
@@ -66,13 +97,32 @@ federation_contexts = Table(
    Column("status_callback_url", Text),
    Column("status", String(20), nullable=False),
    Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
    Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
    Column(
        "updated_at",
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
        onupdate=func.now(),
    ),
    UniqueConstraint(
        "partner_op_id",
        "direction",
        "federation_context_id",
        name="uq_federation_contexts_partner_direction_id",
    ),
    CheckConstraint(
        _one_of(
            "status",
            "available",
            "locked",
            "not_available",
            "temporary_failure",
            "failed",
            "terminated",
        ),
        name="ck_federation_contexts_status",
    ),
    CheckConstraint(_one_of("direction", *DIRECTIONS), name="ck_federation_contexts_direction"),
)

routing_rules = Table(
@@ -85,11 +135,21 @@ routing_rules = Table(
    Column("priority", Integer, nullable=False, server_default="100"),
    Column("is_active", Boolean, nullable=False, server_default=text("true")),
    Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
    Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
    Column(
        "updated_at",
        DateTime(timezone=True),
        nullable=False,
        server_default=func.now(),
        onupdate=func.now(),
    ),
    UniqueConstraint("identifier_type", "value_range", "priority", name="uq_routing_rules_rule"),
    Index("idx_routing_rules_type_value", "identifier_type", "value_range"),
    Index("idx_routing_rules_partner", "partner_op_id"),
    Index("idx_routing_rules_active", "is_active", postgresql_where=text("is_active = true")),
    CheckConstraint(
        _one_of("identifier_type", "msisdn_prefix", "ip_cidr"),
        name="ck_routing_rules_identifier_type",
    ),
)

federation_transactions = Table(
@@ -138,4 +198,13 @@ federation_transactions = Table(
        unique=True,
        postgresql_where=text("idempotency_key IS NOT NULL"),
    ),
    CheckConstraint(
        _one_of("status", "pending", "in_progress", "completed", "partially_completed", "failed"),
        name="ck_federation_transactions_status",
    ),
    CheckConstraint(_one_of("direction", *DIRECTIONS), name="ck_federation_transactions_direction"),
    CheckConstraint(
        _one_of("callback_status", "pending", "delivered", "failed", nullable=True),
        name="ck_federation_transactions_callback_status",
    ),
)
+108 −0
Original line number Diff line number Diff line
"""Constrain state columns to their documented values

The persistence model enumerates the legal values for every status, direction and
identifier-type column, but the tables accepted any string, so a typo reached the
database intact. These are CHECK constraints rather than PostgreSQL enum types
because the spec declares the columns as VARCHAR, and a CHECK can be amended in a
later revision without an ALTER TYPE.

Revision ID: 0002_state_checks
Revises: 0001_baseline
Create Date: 2026-09-22
"""

from collections.abc import Sequence

from alembic import op

revision: str = "0002_state_checks"
down_revision: str | None = "0001_baseline"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None

DIRECTIONS = ("inbound", "outbound")

# (table, constraint name, column, allowed values, nullable)
CHECKS: tuple[tuple[str, str, str, tuple[str, ...], bool], ...] = (
    (
        "partner_ops",
        "ck_partner_ops_status",
        "status",
        ("pending", "active", "suspended", "decommissioned"),
        False,
    ),
    (
        "federation_contexts",
        "ck_federation_contexts_status",
        "status",
        (
            "available",
            "locked",
            "not_available",
            "temporary_failure",
            "failed",
            "terminated",
        ),
        False,
    ),
    (
        "federation_contexts",
        "ck_federation_contexts_direction",
        "direction",
        DIRECTIONS,
        False,
    ),
    (
        "federation_agreements",
        "ck_federation_agreements_status",
        "status",
        ("draft", "active", "suspended", "expired"),
        False,
    ),
    (
        "routing_rules",
        "ck_routing_rules_identifier_type",
        "identifier_type",
        ("msisdn_prefix", "ip_cidr"),
        False,
    ),
    (
        "federation_transactions",
        "ck_federation_transactions_status",
        "status",
        ("pending", "in_progress", "completed", "partially_completed", "failed"),
        False,
    ),
    (
        "federation_transactions",
        "ck_federation_transactions_direction",
        "direction",
        DIRECTIONS,
        False,
    ),
    (
        "federation_transactions",
        "ck_federation_transactions_callback_status",
        "callback_status",
        ("pending", "delivered", "failed"),
        True,
    ),
)


def _condition(column: str, allowed: tuple[str, ...], nullable: bool) -> str:
    values = ", ".join(f"'{value}'" for value in allowed)
    membership = f"{column} IN ({values})"

    # A nullable column carries no state until it is set; NULL must stay legal.
    return f"{column} IS NULL OR {membership}" if nullable else membership


def upgrade() -> None:
    for table, name, column, allowed, nullable in CHECKS:
        op.create_check_constraint(name, table, _condition(column, allowed, nullable))


def downgrade() -> None:
    for table, name, _column, _allowed, _nullable in reversed(CHECKS):
        op.drop_constraint(name, table, type_="check")
+68 −0
Original line number Diff line number Diff line
import os
from uuid import uuid4

import pytest
from sqlalchemy import delete, insert, select, update

from federation_manager.adapters.database.core import (
    build_engine,
    build_session_maker,
    run_migrations,
)
from federation_manager.adapters.database.tables import partner_ops

pytestmark = pytest.mark.integration

URL = os.getenv("FM_POSTGRES_URL", "postgresql+asyncpg://fm:fm@localhost:5433/fm_db")


async def test_updated_at_moves_when_a_row_changes() -> None:
    """updated_at used to equal created_at forever: the column had no onupdate."""
    engine = build_engine(URL)
    await run_migrations(engine)
    session_maker = build_session_maker(engine)

    partner_id = uuid4()
    client_id = f"partner-{uuid4().hex[:8]}"

    async with session_maker() as session:
        await session.execute(
            insert(partner_ops).values(
                id=partner_id,
                mcc_mnc=uuid4().hex[:10],
                oauth2_client_id=client_id,
                base_url="https://partner.example",
                status="pending",
            )
        )
        await session.commit()

        created, first_updated = (
            await session.execute(
                select(partner_ops.c.created_at, partner_ops.c.updated_at).where(
                    partner_ops.c.id == partner_id
                )
            )
        ).one()
        assert first_updated == created

        await session.execute(
            update(partner_ops).where(partner_ops.c.id == partner_id).values(status="active")
        )
        await session.commit()

        second_created, second_updated = (
            await session.execute(
                select(partner_ops.c.created_at, partner_ops.c.updated_at).where(
                    partner_ops.c.id == partner_id
                )
            )
        ).one()

        assert second_created == created, "created_at must not move"
        assert second_updated > first_updated, "updated_at must advance on every change"

        await session.execute(delete(partner_ops).where(partner_ops.c.id == partner_id))
        await session.commit()

    await engine.dispose()