237 lines
8.5 KiB
Python
237 lines
8.5 KiB
Python
"""Production-shaped upgrade/cutover regressions for the RC audit migration."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import importlib.util
|
|
import json
|
|
import uuid
|
|
from datetime import UTC, datetime, timedelta
|
|
from pathlib import Path
|
|
from types import ModuleType
|
|
from unittest.mock import Mock
|
|
|
|
import pytest
|
|
from sqlalchemy import (
|
|
JSON,
|
|
BigInteger,
|
|
Column,
|
|
DateTime,
|
|
MetaData,
|
|
String,
|
|
Table,
|
|
Uuid,
|
|
create_engine,
|
|
inspect,
|
|
select,
|
|
)
|
|
from sqlalchemy.engine import Engine
|
|
from sqlalchemy.orm import Session
|
|
|
|
from alembic.migration import MigrationContext
|
|
from alembic.operations import Operations
|
|
from modelforge_api.domain.audit import AUDIT_EMPTY_LEGACY_PREFIX_SEAL
|
|
from modelforge_api.persistence.models import (
|
|
AuditChainHead,
|
|
AuditEvent,
|
|
)
|
|
from modelforge_api.services.audit import (
|
|
AUDIT_HASH_FORMAT_V1,
|
|
AuditWriter,
|
|
audit_chain_violations,
|
|
canonical_audit_payload_and_hash,
|
|
)
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
MIGRATION = (
|
|
ROOT
|
|
/ "alembic"
|
|
/ "versions"
|
|
/ "20260830_0024_audit_chain_checkpoint.py"
|
|
)
|
|
|
|
|
|
def _migration() -> ModuleType:
|
|
spec = importlib.util.spec_from_file_location("audit_migration_0024", MIGRATION)
|
|
assert spec is not None and spec.loader is not None
|
|
module = importlib.util.module_from_spec(spec)
|
|
spec.loader.exec_module(module)
|
|
return module
|
|
|
|
|
|
def _legacy_engine() -> tuple[Engine, Table]:
|
|
engine = create_engine("sqlite+pysqlite:///:memory:")
|
|
metadata = MetaData()
|
|
events = Table(
|
|
"audit_events",
|
|
metadata,
|
|
Column("id", Uuid(), primary_key=True),
|
|
Column("sequence", BigInteger(), nullable=False, unique=True),
|
|
Column("occurred_at", DateTime(timezone=True), nullable=False),
|
|
Column("correlation_id", String(64), nullable=False),
|
|
Column("actor_type", String(32), nullable=False),
|
|
Column("actor_id", String(255), nullable=False),
|
|
Column("action", String(128), nullable=False),
|
|
Column("resource_type", String(64), nullable=False),
|
|
Column("resource_id", String(64), nullable=True),
|
|
Column("outcome", String(32), nullable=False),
|
|
Column("details", JSON(), nullable=False),
|
|
Column("previous_event_hash", String(64), nullable=True),
|
|
Column("event_hash", String(64), nullable=False, unique=True),
|
|
)
|
|
metadata.create_all(engine)
|
|
return engine, events
|
|
|
|
|
|
def _insert_legacy_chain(engine: Engine, events: Table, count: int = 2) -> None:
|
|
previous_hash: str | None = None
|
|
occurred_at = datetime(2026, 8, 28, 21, 30, tzinfo=UTC)
|
|
rows: list[dict[str, object]] = []
|
|
for sequence in range(1, count + 1):
|
|
payload, event_hash = canonical_audit_payload_and_hash(
|
|
correlation_id=f"production-correlation-{sequence}",
|
|
actor_type="operator",
|
|
actor_id="production-operator",
|
|
action=f"PRODUCTION_ACTION_{sequence}",
|
|
resource_type="model_revision",
|
|
resource_id=str(uuid.uuid4()),
|
|
outcome="success",
|
|
details={"nested": {"approved": True}, "sequence": sequence},
|
|
previous_event_hash=previous_hash,
|
|
hash_format=AUDIT_HASH_FORMAT_V1,
|
|
)
|
|
rows.append(
|
|
{
|
|
"id": uuid.uuid4(),
|
|
"sequence": sequence,
|
|
"occurred_at": occurred_at + timedelta(microseconds=sequence),
|
|
"event_hash": event_hash,
|
|
**payload,
|
|
}
|
|
)
|
|
previous_hash = event_hash
|
|
with engine.begin() as connection:
|
|
connection.execute(events.insert(), rows)
|
|
|
|
|
|
def _run(module: ModuleType, engine: Engine, operation: str) -> None:
|
|
with engine.begin() as connection:
|
|
context = MigrationContext.configure(connection)
|
|
module.op = Operations(context)
|
|
getattr(module, operation)()
|
|
|
|
|
|
def test_upgrade_validates_and_seals_production_shaped_legacy_history() -> None:
|
|
engine, legacy_table = _legacy_engine()
|
|
_insert_legacy_chain(engine, legacy_table)
|
|
module = _migration()
|
|
|
|
_run(module, engine, "upgrade")
|
|
|
|
assert "hash_format" in {column["name"] for column in inspect(engine).get_columns("audit_events")}
|
|
with Session(engine) as session:
|
|
legacy = list(session.scalars(select(AuditEvent).order_by(AuditEvent.sequence)))
|
|
checkpoint = session.get(AuditChainHead, 1)
|
|
assert checkpoint is not None
|
|
assert [event.hash_format for event in legacy] == ["v1", "v1"]
|
|
assert checkpoint.event_count == 2
|
|
assert checkpoint.v2_start_sequence == 3
|
|
assert checkpoint.legacy_prefix_count == 2
|
|
assert checkpoint.legacy_prefix_seal != AUDIT_EMPTY_LEGACY_PREFIX_SEAL
|
|
assert audit_chain_violations(legacy, checkpoint) == []
|
|
|
|
AuditWriter(session, "operator", "post-cutover").write(
|
|
"POST_CUTOVER", "model_revision", str(uuid.uuid4()), {}
|
|
)
|
|
session.commit()
|
|
mixed = list(session.scalars(select(AuditEvent).order_by(AuditEvent.sequence)))
|
|
assert [event.hash_format for event in mixed] == ["v1", "v1", "v2"]
|
|
assert audit_chain_violations(mixed, session.get(AuditChainHead, 1)) == []
|
|
|
|
with pytest.raises(RuntimeError, match="cannot downgrade.*v2 events"):
|
|
_run(module, engine, "downgrade")
|
|
|
|
|
|
def test_upgrade_rejects_a_legacy_random_hash_before_schema_mutation() -> None:
|
|
engine, legacy_table = _legacy_engine()
|
|
_insert_legacy_chain(engine, legacy_table, count=1)
|
|
with engine.begin() as connection:
|
|
connection.execute(
|
|
legacy_table.update().values(
|
|
action="RECOVERY_RECONCILIATION_COMPLETED",
|
|
event_hash="a" * 64,
|
|
)
|
|
)
|
|
module = _migration()
|
|
|
|
with pytest.raises(RuntimeError, match="content hash is invalid"):
|
|
_run(module, engine, "upgrade")
|
|
|
|
assert "hash_format" not in {
|
|
column["name"] for column in inspect(engine).get_columns("audit_events")
|
|
}
|
|
assert "audit_chain_heads" not in inspect(engine).get_table_names()
|
|
|
|
|
|
@pytest.mark.parametrize("invalid_details", [["array"], "scalar", 42, True, None])
|
|
def test_upgrade_rejects_valid_hash_whose_legacy_details_are_not_an_object(
|
|
invalid_details: object,
|
|
) -> None:
|
|
engine, legacy_table = _legacy_engine()
|
|
_insert_legacy_chain(engine, legacy_table, count=1)
|
|
with engine.begin() as connection:
|
|
row = connection.execute(select(legacy_table)).mappings().one()
|
|
payload = {
|
|
"correlation_id": row["correlation_id"],
|
|
"actor_type": row["actor_type"],
|
|
"actor_id": row["actor_id"],
|
|
"action": row["action"],
|
|
"resource_type": row["resource_type"],
|
|
"resource_id": row["resource_id"],
|
|
"outcome": row["outcome"],
|
|
"details": invalid_details,
|
|
"previous_event_hash": row["previous_event_hash"],
|
|
}
|
|
valid_hash = hashlib.sha256(
|
|
json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8")
|
|
).hexdigest()
|
|
connection.execute(
|
|
legacy_table.update().values(details=invalid_details, event_hash=valid_hash)
|
|
)
|
|
|
|
with pytest.raises(RuntimeError, match="details must be a JSON object"):
|
|
_run(_migration(), engine, "upgrade")
|
|
|
|
assert "hash_format" not in {
|
|
column["name"] for column in inspect(engine).get_columns("audit_events")
|
|
}
|
|
assert "audit_chain_heads" not in inspect(engine).get_table_names()
|
|
|
|
|
|
def test_postgresql_migration_locks_legacy_and_current_writers_before_validation() -> None:
|
|
module = _migration()
|
|
connection = Mock()
|
|
connection.dialect.name = "postgresql"
|
|
|
|
module._lock_legacy_audit_chain(connection)
|
|
|
|
advisory_statement, advisory_parameters = connection.execute.call_args_list[0].args
|
|
table_lock_statement = connection.execute.call_args_list[1].args[0]
|
|
assert str(advisory_statement) == "select pg_advisory_xact_lock(:lock_key)"
|
|
assert advisory_parameters == {"lock_key": module._AUDIT_CHAIN_LOCK_KEY}
|
|
assert str(table_lock_statement) == "lock table audit_events in access exclusive mode"
|
|
|
|
|
|
def test_downgrade_is_supported_only_before_the_first_v2_event() -> None:
|
|
engine, legacy_table = _legacy_engine()
|
|
_insert_legacy_chain(engine, legacy_table, count=1)
|
|
module = _migration()
|
|
_run(module, engine, "upgrade")
|
|
|
|
_run(module, engine, "downgrade")
|
|
|
|
assert "hash_format" not in {
|
|
column["name"] for column in inspect(engine).get_columns("audit_events")
|
|
}
|
|
assert "audit_chain_heads" not in inspect(engine).get_table_names()
|