Files
ModelForge/backend/tests/test_audit_integrity_rc.py
T

910 lines
31 KiB
Python

"""RC audit-chain integrity and compatibility regressions."""
from __future__ import annotations
import inspect as pyinspect
import uuid
from collections.abc import Iterator
from contextlib import contextmanager
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import Mock
import pytest
from sqlalchemy import create_engine, delete, event, func, insert, select, text, update
from sqlalchemy.dialects import postgresql
from sqlalchemy.engine import Connection, Engine
from sqlalchemy.orm import Session
import modelforge_api.domain.audit as audit_domain
import modelforge_api.persistence.models as persistence_models
from modelforge_api.db import build_engine
from modelforge_api.domain.audit import AUDIT_EMPTY_LEGACY_PREFIX_SEAL
from modelforge_api.persistence.models import (
AuditChainHead,
AuditEvent,
Base,
_protected_audit_dml_targets,
_textual_audit_dml_targets,
)
from modelforge_api.persistence.repositories import AuditRepository
from modelforge_api.services.audit import (
AUDIT_CHAIN_POSTGRES_LOCK_KEY,
AUDIT_CURRENT_HASH_FORMAT,
AUDIT_HASH_FORMAT_V1,
AuditContext,
AuditWriter,
_acquire_audit_write_lock,
audit_chain_violations,
canonical_audit_payload_and_hash,
legacy_audit_prefix_seal,
)
from modelforge_api.services.invariants import InvariantStatus, check_invariants
@pytest.fixture
def session() -> Session:
engine = create_engine("sqlite+pysqlite:///:memory:")
Base.metadata.create_all(engine)
with Session(engine) as value:
yield value
def _seed_chain(session: Session, count: int = 3) -> list[AuditEvent]:
writer = AuditWriter(session, "operator", "alice")
events = [
writer.write(
action=f"ACTION_{index}",
resource_type="model",
resource_id=f"model-{index}",
details={"index": index, "evidence": {"approved": True}},
)
for index in range(1, count + 1)
]
session.commit()
return events
def _audit_invariant(session: Session) -> Any:
return next(
result
for result in check_invariants(session).results
if result.key == "audit_chain_intact"
)
def _checkpoint(session: Session) -> AuditChainHead:
checkpoint = session.get(AuditChainHead, 1)
assert checkpoint is not None
return checkpoint
@contextmanager
def _privileged_database_connection(session: Session) -> Iterator[Connection]:
"""Use a direct Engine connection to model a database administrator/compromise."""
session.commit()
engine = session.get_bind()
assert isinstance(engine, Engine)
with engine.begin() as connection:
yield connection
session.expire_all()
def _privileged_tamper(session: Session, statement: Any) -> None:
with _privileged_database_connection(session) as connection:
connection.execute(statement)
def _application_engine(tmp_path: Any) -> Engine:
database = tmp_path / "application-audit-boundary.sqlite3"
url = f"sqlite+pysqlite:///{database.as_posix()}"
setup = create_engine(url)
Base.metadata.create_all(setup)
setup.dispose()
return build_engine(url)
def test_fresh_metadata_database_has_a_valid_empty_checkpoint(session: Session) -> None:
checkpoint = _checkpoint(session)
assert checkpoint.event_count == 0
assert checkpoint.last_sequence == 0
assert checkpoint.last_event_hash is None
assert checkpoint.legacy_prefix_seal == AUDIT_EMPTY_LEGACY_PREFIX_SEAL
assert audit_chain_violations([], checkpoint) == []
AuditWriter(session, "operator", "first-legitimate-writer").write(
"FIRST_EVENT", "model", "model-1", {}
)
session.commit()
events = list(session.scalars(select(AuditEvent)))
assert audit_chain_violations(events, _checkpoint(session)) == []
def test_a_canonical_chain_is_contiguous_linked_and_content_verified(session: Session) -> None:
_seed_chain(session)
events = list(session.scalars(select(AuditEvent).order_by(AuditEvent.sequence)))
assert [event.sequence for event in events] == [1, 2, 3]
assert events[0].previous_event_hash is None
assert events[1].previous_event_hash == events[0].event_hash
assert events[2].previous_event_hash == events[1].event_hash
assert audit_chain_violations(events, _checkpoint(session)) == []
assert _audit_invariant(session).status is InvariantStatus.HOLDS
def test_writer_stores_the_exact_detached_payload_that_it_hashes(session: Session) -> None:
details = {"evidence": {"approved": True}}
event = AuditWriter(session, "operator", "alice").write(
"APPROVE", "revision", "revision-1", details
)
details["evidence"]["approved"] = False
session.commit()
assert event.details == {"evidence": {"approved": True}}
assert audit_chain_violations([event], _checkpoint(session)) == []
def test_audit_context_preserves_request_correlation_and_principal(session: Session) -> None:
context = AuditContext(
actor_type="operator",
actor_id="principal-42",
correlation_id="request-correlation-42",
)
writer = AuditWriter(session, context=context)
writer.write("START", "operation", "op-1", {})
writer.write("COMPLETE", "operation", "op-1", {})
session.commit()
events = list(session.scalars(select(AuditEvent).order_by(AuditEvent.sequence)))
assert {event.actor_id for event in events} == {"principal-42"}
assert {event.correlation_id for event in events} == {"request-correlation-42"}
assert audit_chain_violations(events, _checkpoint(session)) == []
def test_legacy_audit_repository_uses_the_canonical_writer(session: Session) -> None:
event = AuditRepository(session).append(
correlation_id="repository-correlation",
actor_type="system",
actor_id="repository-test",
action="SYNC",
resource_type="project",
resource_id="project-1",
outcome="success",
details={"count": 2},
)
session.commit()
assert event.sequence == 1
assert audit_chain_violations([event], _checkpoint(session)) == []
@pytest.mark.parametrize(
("column", "tampered_value"),
[
("correlation_id", "tampered-correlation"),
("actor_type", "tampered-actor-type"),
("actor_id", "mallory"),
("action", "TAMPERED_ACTION"),
("resource_type", "tampered-resource"),
("resource_id", "tampered-resource-id"),
("outcome", "failure"),
("details", {"tampered": True}),
("previous_event_hash", "b" * 64),
("event_hash", "c" * 64),
],
)
def test_every_hashed_field_and_both_hash_columns_are_tamper_evident(
session: Session,
column: str,
tampered_value: object,
) -> None:
_seed_chain(session)
_privileged_tamper(
session,
update(AuditEvent)
.where(AuditEvent.sequence == 2)
.values({column: tampered_value}),
)
session.commit()
result = _audit_invariant(session)
assert result.status is InvariantStatus.VIOLATED
assert result.violations
def test_a_non_null_first_link_is_detected(session: Session) -> None:
_seed_chain(session)
_privileged_tamper(
session,
update(AuditEvent)
.where(AuditEvent.sequence == 1)
.values(previous_event_hash="d" * 64),
)
session.commit()
result = _audit_invariant(session)
assert result.status is InvariantStatus.VIOLATED
assert any("first event" in violation for violation in result.violations)
def test_a_sequence_gap_is_detected_even_when_remaining_numbers_are_unique(
session: Session,
) -> None:
_seed_chain(session)
_privileged_tamper(
session, delete(AuditEvent).where(AuditEvent.sequence == 2)
)
session.commit()
result = _audit_invariant(session)
assert result.status is InvariantStatus.VIOLATED
assert any("expected 2" in violation for violation in result.violations)
@pytest.mark.parametrize("field", ["id", "occurred_at"])
def test_v2_event_id_and_timestamp_are_tamper_evident(
session: Session, field: str
) -> None:
_seed_chain(session)
tampered = (
uuid.uuid4()
if field == "id"
else datetime.now(UTC) + timedelta(days=1)
)
_privileged_tamper(
session,
update(AuditEvent)
.where(AuditEvent.sequence == 2)
.values({field: tampered}),
)
session.commit()
assert _audit_invariant(session).status is InvariantStatus.VIOLATED
@pytest.mark.parametrize("retained", [2, 0])
def test_checkpoint_detects_tail_and_complete_chain_deletion(
session: Session, retained: int
) -> None:
_seed_chain(session)
_privileged_tamper(
session, delete(AuditEvent).where(AuditEvent.sequence > retained)
)
session.commit()
result = _audit_invariant(session)
assert result.status is InvariantStatus.VIOLATED
assert any("checkpoint" in violation for violation in result.violations)
def test_privileged_middle_tamper_is_caught_by_the_explicit_strict_gate(
session: Session,
) -> None:
_seed_chain(session)
_privileged_tamper(
session,
update(AuditEvent)
.where(AuditEvent.sequence == 2)
.values(details={"attacker": "changed-without-moving-head"}),
)
session.commit()
# The O(1) append gate intentionally does not rescan a sealed middle prefix. A privileged
# database edit remains visible to the explicit full-chain invariant/recovery verification.
AuditWriter(session, "operator", "legitimate-writer").write(
"LEGITIMATE_APPEND", "model", "4", {}
)
session.commit()
assert len(list(session.scalars(select(AuditEvent)))) == 4
assert _audit_invariant(session).status is InvariantStatus.VIOLATED
def test_privileged_tail_tamper_is_refused_by_the_constant_time_append_gate(
session: Session,
) -> None:
_seed_chain(session)
_privileged_tamper(
session,
update(AuditEvent)
.where(AuditEvent.sequence == 3)
.values(details={"attacker": "changed-tail"}),
)
session.commit()
with pytest.raises(RuntimeError, match="checkpoint/tail failed"):
AuditWriter(session, "operator", "legitimate-writer").write(
"REFUSED_AFTER_TAIL_TAMPER", "model", "4", {}
)
assert len(list(session.scalars(select(AuditEvent)))) == 3
def test_mixed_migrated_v1_prefix_and_v2_suffix_remain_verifiable(session: Session) -> None:
occurred_at = datetime(2026, 8, 29, 12, 0, tzinfo=UTC)
previous_hash: str | None = None
legacy: list[AuditEvent] = []
for sequence in (1, 2):
event_id = uuid.uuid4()
payload, event_hash = canonical_audit_payload_and_hash(
correlation_id=f"legacy-{sequence}",
actor_type="operator",
actor_id="legacy",
action=f"LEGACY_{sequence}",
resource_type="model",
resource_id=str(sequence),
outcome="success",
details={"sequence": sequence},
previous_event_hash=previous_hash,
hash_format=AUDIT_HASH_FORMAT_V1,
)
row = AuditEvent(
id=event_id,
sequence=sequence,
occurred_at=occurred_at + timedelta(seconds=sequence),
hash_format=AUDIT_HASH_FORMAT_V1,
event_hash=event_hash,
**payload,
)
legacy.append(row)
previous_hash = event_hash
with _privileged_database_connection(session) as connection:
connection.execute(
insert(AuditEvent),
[
{
"id": row.id,
"sequence": row.sequence,
"occurred_at": row.occurred_at,
"hash_format": row.hash_format,
"event_hash": row.event_hash,
"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": row.details,
"previous_event_hash": row.previous_event_hash,
}
for row in legacy
],
)
connection.execute(
update(AuditChainHead)
.where(AuditChainHead.singleton_id == 1)
.values(
event_count=2,
last_sequence=2,
last_event_hash=previous_hash,
hash_format=AUDIT_CURRENT_HASH_FORMAT,
v2_start_sequence=3,
legacy_prefix_count=2,
legacy_prefix_seal=legacy_audit_prefix_seal(legacy),
updated_at=datetime.now(UTC),
)
)
session.commit()
AuditWriter(session, "operator", "new-writer").write("V2", "model", "3", {})
session.commit()
events = list(session.scalars(select(AuditEvent).order_by(AuditEvent.sequence)))
assert [row.hash_format for row in events] == ["v1", "v1", "v2"]
assert audit_chain_violations(events, _checkpoint(session)) == []
_privileged_tamper(
session,
update(AuditEvent)
.where(AuditEvent.sequence == 1)
.values(occurred_at=datetime.now(UTC)),
)
session.commit()
assert any(
"prefix seal" in violation
for violation in audit_chain_violations(events, _checkpoint(session))
)
AuditWriter(session, "operator", "new-writer").write(
"APPEND_AFTER_PRIVILEGED_PREFIX_TAMPER", "model", "4", {}
)
session.commit()
assert _audit_invariant(session).status is InvariantStatus.VIOLATED
def test_ordinary_orm_code_cannot_update_the_checkpoint_separately(session: Session) -> None:
_seed_chain(session, count=1)
checkpoint = _checkpoint(session)
checkpoint.event_count = 0
with pytest.raises(ValueError, match="canonical audit writer"):
session.commit()
session.rollback()
with pytest.raises(ValueError, match="canonical audit writer"):
session.execute(
update(AuditChainHead)
.where(AuditChainHead.singleton_id == 1)
.values(event_count=0)
)
session.rollback()
def test_exact_bulk_delete_and_head_rewrite_poc_is_refused_but_writer_succeeds(
session: Session,
) -> None:
_seed_chain(session, count=3)
before = (
_checkpoint(session).event_count,
_checkpoint(session).last_sequence,
_checkpoint(session).last_event_hash,
)
with pytest.raises(ValueError, match="canonical audit writer"):
session.execute(delete(AuditEvent))
session.execute(
update(AuditChainHead)
.where(AuditChainHead.singleton_id == 1)
.values(event_count=0, last_sequence=0, last_event_hash=None)
)
session.rollback()
assert len(list(session.scalars(select(AuditEvent)))) == 3
head = _checkpoint(session)
assert (head.event_count, head.last_sequence, head.last_event_hash) == before
appended = AuditWriter(session, "operator", "legitimate-writer").write(
"LEGITIMATE_APPEND", "model", "model-4", {}
)
session.commit()
assert appended.sequence == 4
assert audit_chain_violations(
list(session.scalars(select(AuditEvent).order_by(AuditEvent.sequence))),
_checkpoint(session),
) == []
@pytest.mark.parametrize(
"statement",
[
delete(AuditChainHead),
insert(AuditChainHead).values(
singleton_id=1,
event_count=0,
last_sequence=0,
last_event_hash=None,
hash_format="v2",
v2_start_sequence=1,
legacy_prefix_count=0,
legacy_prefix_seal=AUDIT_EMPTY_LEGACY_PREFIX_SEAL,
),
update(AuditEvent).values(action="BULK_TAMPER"),
insert(AuditEvent).values(
id=uuid.uuid4(),
sequence=99,
occurred_at=datetime.now(UTC),
correlation_id="bulk",
actor_type="attacker",
actor_id="attacker",
action="BULK_INSERT",
resource_type="audit",
resource_id=None,
outcome="success",
details={},
previous_event_hash=None,
event_hash="a" * 64,
hash_format="v2",
),
text("delete from audit_events"),
],
)
def test_ordinary_session_dml_has_no_audit_bypass(session: Session, statement: Any) -> None:
with pytest.raises(ValueError, match="canonical audit writer"):
session.execute(statement)
session.rollback()
def test_ordinary_orm_and_legacy_bulk_inserts_cannot_create_audit_events(
session: Session,
) -> None:
payload = {
"id": uuid.uuid4(),
"sequence": 1,
"occurred_at": datetime.now(UTC),
"correlation_id": "ordinary",
"actor_type": "attacker",
"actor_id": "attacker",
"action": "INSERT",
"resource_type": "audit",
"resource_id": None,
"outcome": "success",
"details": {},
"previous_event_hash": None,
"event_hash": "a" * 64,
"hash_format": "v2",
}
session.add(AuditEvent(**payload))
with pytest.raises(ValueError, match="canonical audit writer"):
session.flush()
session.rollback()
with pytest.raises(ValueError, match="application session connection"):
session.bulk_insert_mappings(AuditEvent, [payload])
session.rollback()
def test_aliased_and_annotated_core_dml_resolve_the_protected_base_table(
session: Session,
) -> None:
event_alias = AuditEvent.__table__.alias("erased_events")
head_alias = AuditChainHead.__table__.alias("forged_head")
statements = [
delete(event_alias),
update(head_alias).values(
event_count=0,
last_sequence=0,
last_event_hash=None,
),
delete(AuditEvent.__table__._annotate({"reviewer": "alias-control"})),
update(
AuditChainHead.__table__._annotate({"reviewer": "alias-control"})
).values(event_count=0, last_sequence=0, last_event_hash=None),
]
for statement in statements:
assert _protected_audit_dml_targets(statement)
compiled = str(statement.compile(dialect=postgresql.dialect()))
assert _textual_audit_dml_targets(compiled)
with pytest.raises(ValueError, match="canonical audit writer"):
session.execute(statement)
session.rollback()
def test_exact_aliased_delete_and_head_rewrite_poc_is_refused(
session: Session,
) -> None:
_seed_chain(session, count=3)
event_alias = AuditEvent.__table__.alias("erased_events")
head_alias = AuditChainHead.__table__.alias("forged_head")
with pytest.raises(ValueError, match="canonical audit writer"):
session.execute(delete(event_alias))
session.rollback()
with pytest.raises(ValueError, match="canonical audit writer"):
session.execute(
update(head_alias).values(
event_count=0,
last_sequence=0,
last_event_hash=None,
)
)
session.rollback()
appended = AuditWriter(session, "operator", "legitimate-writer").write(
"APPEND_AFTER_ALIASED_POC", "model", "model-4", {}
)
session.commit()
assert appended.sequence == 4
assert audit_chain_violations(
list(session.scalars(select(AuditEvent).order_by(AuditEvent.sequence))),
_checkpoint(session),
) == []
@pytest.mark.parametrize(
"statement",
[
"/* leading /* nested */ comment */ DELETE\nFROM [main].[audit_events]",
"-- leading decoy\nUPDATE `main`.`audit_chain_heads` "
"SET event_count=0, last_sequence=0, last_event_hash=NULL",
'INSERT /* gap */ INTO "public"."audit_events" (id) VALUES (NULL)',
"WITH doomed AS (SELECT 1) DELETE FROM audit_events",
"TRUNCATE TABLE harmless_table, audit_events",
"MERGE INTO public.audit_chain_heads AS head USING incoming ON false "
"WHEN MATCHED THEN DELETE",
"COPY audit_events FROM STDIN",
"DROP TABLE harmless_table, audit_chain_heads",
"DO $$ BEGIN DELETE FROM audit_events; END $$",
"CALL rewrite_audit_chain()",
],
)
def test_session_connection_exec_driver_sql_blocks_obfuscated_audit_dml(
session: Session, statement: str
) -> None:
connection = session.connection()
with pytest.raises(ValueError, match="raw audit DML"):
connection.exec_driver_sql(statement)
session.rollback()
@pytest.mark.parametrize(
"statement",
[
"SELECT 'delete from audit_events' AS harmless",
"/* DELETE FROM audit_events */ SELECT 1",
"-- UPDATE audit_chain_heads SET event_count=0\nSELECT 1",
],
)
def test_raw_sql_lexer_ignores_non_executable_comments_and_strings(
session: Session, statement: str
) -> None:
assert _textual_audit_dml_targets(statement) == frozenset()
assert session.connection().exec_driver_sql(statement).scalar_one() in {
"delete from audit_events",
1,
}
session.rollback()
def test_raw_sql_lexer_ignores_postgresql_dollar_quoted_decoys() -> None:
assert (
_textual_audit_dml_targets(
"SELECT $audit$DELETE FROM audit_events$audit$, "
"$$UPDATE audit_chain_heads SET event_count=0$$"
)
== frozenset()
)
def test_raw_sql_lexer_classifies_postgresql_unicode_quoted_identifiers() -> None:
assert _textual_audit_dml_targets('DELETE FROM U&"audit_events"') == frozenset(
{"audit_events"}
)
def test_direct_engine_raw_tamper_is_outside_hook_boundary_but_strictly_detected(
session: Session,
) -> None:
_seed_chain(session, count=2)
with _privileged_database_connection(session) as connection:
connection.exec_driver_sql(
"UPDATE audit_events SET action='PRIVILEGED_RAW_TAMPER' WHERE sequence=1"
)
events = list(session.scalars(select(AuditEvent).order_by(AuditEvent.sequence)))
assert any(
"content hash" in violation
for violation in audit_chain_violations(events, _checkpoint(session))
)
def test_every_connection_from_the_application_engine_blocks_audit_dml(
tmp_path: Any,
) -> None:
engine = _application_engine(tmp_path)
try:
with Session(engine) as session:
AuditWriter(session, "operator", "legitimate").write(
"LEGITIMATE", "model", "model-1", {}
)
session.commit()
bind = session.bind
assert isinstance(bind, Engine)
with pytest.raises(ValueError, match="audit"), bind.begin() as connection:
connection.exec_driver_sql(
" /* ordinary connection */ DELETE FROM audit_events"
)
with pytest.raises(ValueError, match="audit"), engine.begin() as connection:
connection.exec_driver_sql(
'UPDATE "audit_chain_heads" SET event_count=0, last_sequence=0, '
"last_event_hash=NULL"
)
with Session(engine) as session:
assert session.scalar(select(func.count()).select_from(AuditEvent)) == 1
assert _checkpoint(session).event_count == 1
finally:
engine.dispose()
def test_late_statement_rewrite_is_rolled_back_before_commit(tmp_path: Any) -> None:
engine = _application_engine(tmp_path)
try:
with Session(engine) as session:
AuditWriter(session, "operator", "legitimate").write(
"LEGITIMATE", "model", "model-1", {}
)
session.commit()
def late_rewrite(
_connection: Any,
_cursor: Any,
statement: str,
parameters: Any,
_context: Any,
_executemany: bool,
) -> tuple[str, Any]:
if statement.strip().upper() == "SELECT 1":
return "DELETE FROM audit_events", parameters
return statement, parameters
event.listen(engine, "before_cursor_execute", late_rewrite, retval=True)
try:
with pytest.raises(ValueError, match="rolled back"), engine.begin() as connection:
connection.exec_driver_sql("SELECT 1")
finally:
event.remove(engine, "before_cursor_execute", late_rewrite)
with Session(engine) as session:
events = list(session.scalars(select(AuditEvent)))
assert len(events) == 1
assert audit_chain_violations(events, _checkpoint(session)) == []
finally:
engine.dispose()
@pytest.mark.parametrize(
"statement",
[
r'ALTER TABLE IF EXISTS U&"audit\005fevents" DISABLE TRIGGER ALL',
r'''ALTER TABLE IF EXISTS U&"audit!005fchain!005fheads" UESCAPE '!' DISABLE TRIGGER ALL''',
"CREATE OR REPLACE FUNCTION reset_chain() RETURNS void LANGUAGE SQL AS $$ "
"DELETE FROM audit_events $$",
"SELECT public.reset_chain()",
"CALL reset_chain()",
"DO $$ BEGIN EXECUTE 'DELETE FROM audit_events'; END $$",
"DELETE events, heads FROM audit_events AS events JOIN audit_chain_heads AS heads ON 1=1",
"DELETE FROM audit_events AS events USING audit_chain_heads AS heads",
],
)
def test_runtime_sql_classifier_defaults_procedural_and_obfuscated_mutation_to_deny(
statement: str,
) -> None:
assert _textual_audit_dml_targets(statement)
@pytest.mark.parametrize(
"statement",
[
"SELECT count(*) FROM audit_events",
"SELECT * FROM audit_chain_heads",
"SELECT * FROM modelforge_audit.append_event_v2(NULL, NULL, NULL, NULL, NULL, "
"NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL)",
],
)
def test_runtime_sql_classifier_keeps_read_and_canonical_append_controls(
statement: str,
) -> None:
assert _textual_audit_dml_targets(statement) == frozenset()
def test_old_visible_capabilities_cannot_be_imported_guessed_reused_or_crossed(
session: Session,
) -> None:
retired_names = {
"_AUDIT_INTERNAL_EXECUTION_TOKEN",
"_AUDIT_INTERNAL_EXECUTION_OPTION",
"_AUDIT_INTERNAL_CONNECTION_INFO_KEY",
"_allow_privileged_audit_connection",
"_allow_canonical_audit_event_insert",
"_mark_canonical_audit_head_update",
}
assert retired_names.isdisjoint(vars(audit_domain))
assert retired_names.isdisjoint(vars(persistence_models))
assert "_install_audit_dml_boundary" not in vars(persistence_models)
canonical_operation = persistence_models._append_canonical_audit_event
assert canonical_operation.__closure__ is None
assert pyinspect.getclosurevars(canonical_operation).nonlocals == {}
with pytest.raises(TypeError, match="unexpected keyword argument"):
canonical_operation(session, statement=delete(AuditEvent)) # type: ignore[call-arg]
connection = session.connection()
stolen_or_guessed = object()
for info in (session.info, connection.info):
info["_modelforge_audit_internal_execution"] = stolen_or_guessed
info["_modelforge_audit_privileged_connection"] = stolen_or_guessed
statement = delete(AuditEvent).execution_options(
_modelforge_audit_internal_execution=stolen_or_guessed
)
with pytest.raises(ValueError, match="canonical audit writer"):
session.execute(statement)
session.rollback()
engine = session.get_bind()
assert isinstance(engine, Engine)
with Session(engine) as other:
other.info["_modelforge_audit_internal_execution"] = stolen_or_guessed
other_connection = other.connection()
other_connection.info["_modelforge_audit_privileged_connection"] = stolen_or_guessed
with pytest.raises(ValueError, match="canonical audit writer"):
other.execute(
update(AuditChainHead)
.values(event_count=0, last_sequence=0, last_event_hash=None)
.execution_options(
_modelforge_audit_internal_execution=stolen_or_guessed
)
)
other.rollback()
def test_canonical_append_leaves_no_visible_or_reusable_connection_capability(
session: Session,
) -> None:
AuditWriter(session, "operator", "legitimate-writer").write(
"CAPABILITY_CLEANUP_CONTROL", "model", "model-1", {}
)
session.commit()
connection = session.connection()
forbidden_fragments = ("internal", "permit", "privileged", "token")
assert not any(
any(fragment in str(key).lower() for fragment in forbidden_fragments)
for key in session.info
)
assert not any(
any(fragment in str(key).lower() for fragment in forbidden_fragments)
for key in connection.info
)
with pytest.raises(ValueError, match="application session connection"):
connection.execute(delete(AuditEvent))
session.rollback()
def test_privileged_checkpoint_rewrite_is_refused_by_the_append_gate(
session: Session,
) -> None:
events = _seed_chain(session, count=3)
_privileged_tamper(
session,
update(AuditChainHead)
.where(AuditChainHead.singleton_id == 1)
.values(
event_count=2,
last_sequence=2,
last_event_hash=events[1].event_hash,
),
)
session.commit()
with pytest.raises(RuntimeError, match="checkpoint/tail failed"):
AuditWriter(session, "operator", "legitimate-writer").write(
"REFUSED_AFTER_HEAD_REWRITE", "model", "model-4", {}
)
def test_failed_event_insert_rolls_back_event_and_checkpoint(session: Session) -> None:
_seed_chain(session, count=1)
before = _checkpoint(session)
expected = (before.event_count, before.last_sequence, before.last_event_hash)
engine = session.get_bind()
assert isinstance(engine, Engine)
def fail_insert(
_connection: Connection,
_cursor: Any,
statement: str,
parameters: Any,
_context: Any,
_executemany: bool,
) -> None:
if statement.lstrip().lower().startswith("insert into audit_events") and (
"FAIL_INSERT" in repr(parameters)
):
raise RuntimeError("injected audit insert failure")
event.listen(engine, "before_cursor_execute", fail_insert)
try:
with pytest.raises(RuntimeError, match="injected"):
AuditWriter(session, "operator", "alice").write(
"FAIL_INSERT", "model", "2", {}
)
finally:
event.remove(engine, "before_cursor_execute", fail_insert)
rows = list(session.scalars(select(AuditEvent)))
after = _checkpoint(session)
assert len(rows) == 1
assert (after.event_count, after.last_sequence, after.last_event_hash) == expected
def test_postgresql_writers_take_the_shared_transaction_scoped_advisory_lock() -> None:
session = Mock()
session.get_bind.return_value = SimpleNamespace(
dialect=SimpleNamespace(name="postgresql")
)
_acquire_audit_write_lock(cast(Session, session))
statement, parameters = session.execute.call_args.args
assert str(statement) == "SELECT pg_advisory_xact_lock(:lock_key)"
assert parameters == {"lock_key": AUDIT_CHAIN_POSTGRES_LOCK_KEY}