from __future__ import annotations import unittest from unittest.mock import patch from sqlalchemy import create_engine, event from sqlalchemy.orm import sessionmaker from govoplan_admin.backend.db.models import GovernanceTemplate, GovernanceTemplateAssignment from govoplan_admin.backend.governance import set_template_assignments, synchronize_templates from govoplan_core.core.access import GovernanceProjectionOutcome, GovernanceProjectionResult from govoplan_core.db.base import Base from govoplan_core.tenancy.scope import Tenant, create_scope_tables, scope_registry class _Projection: def __init__(self) -> None: self.batches = [] def reconcile(self, session: object, batch): del session self.batches.append(batch) return GovernanceProjectionResult( operation_id=batch.operation_id, dry_run=batch.dry_run, outcomes=tuple( GovernanceProjectionOutcome( assignment_id=command.assignment_id, template_id=command.template.template_id, tenant_id=command.template.tenant_id, kind=command.template.kind, operation=command.operation, status=("removed" if command.operation == "remove" else "created"), resource_id=f"resource-{command.assignment_id}", provenance=dict(command.provenance), ) for command in batch.commands ), ) class _Registry: def __init__(self, projection: _Projection) -> None: self.projection = projection def has_capability(self, name: str) -> bool: return name == "access.governanceProjection.v1" def require_capability(self, name: str) -> object: if not self.has_capability(name): raise KeyError(name) return self.projection class GovernanceBulkSyncTests(unittest.TestCase): def setUp(self) -> None: self.engine = create_engine("sqlite:///:memory:") create_scope_tables(self.engine) Base.metadata.create_all(bind=self.engine) self.Session = sessionmaker(bind=self.engine) self.session = self.Session() self.projection = _Projection() self.registry_patch = patch( "govoplan_admin.backend.governance.get_registry", return_value=_Registry(self.projection), ) self.registry_patch.start() def tearDown(self) -> None: self.registry_patch.stop() self.session.close() Base.metadata.drop_all(bind=self.engine) scope_registry.metadata.drop_all(bind=self.engine) self.engine.dispose() def _template(self) -> GovernanceTemplate: item = GovernanceTemplate( id="template-1", kind="role", slug="reviewer", name="Reviewer", permissions=["access:role:read"], ) self.session.add(item) self.session.flush() return item def test_assignment_validation_and_projection_use_bounded_reads(self) -> None: item = self._template() tenants = [ Tenant(id=f"tenant-{index}", slug=f"tenant-{index}", name=f"Tenant {index}") for index in range(200) ] self.session.add_all(tenants) self.session.commit() self.session.refresh(item) select_count = 0 def record_select(_conn, _cursor, statement, _parameters, _context, _executemany): nonlocal select_count if statement.lstrip().upper().startswith("SELECT"): select_count += 1 event.listen(self.engine, "before_cursor_execute", record_select) try: result = set_template_assignments( self.session, item, [ {"tenant_id": f"tenant-{index}", "mode": "required" if index % 2 else "available"} for index in range(len(tenants)) ], ) finally: event.remove(self.engine, "before_cursor_execute", record_select) self.assertEqual(2, select_count) self.assertEqual(200, len(result.outcomes)) self.assertEqual(1, len(self.projection.batches)) self.assertEqual(200, len(self.projection.batches[0].commands)) self.assertEqual(200, self.session.query(GovernanceTemplateAssignment).count()) def test_bulk_synchronization_reports_stale_tenant_without_skipping_valid_assignment(self) -> None: item = self._template() self.session.add(Tenant(id="tenant-valid", slug="valid", name="Valid")) self.session.add_all( [ GovernanceTemplateAssignment( id="assignment-valid", template_id=item.id, tenant_id="tenant-valid", mode="required", ), GovernanceTemplateAssignment( id="assignment-stale", template_id=item.id, tenant_id="tenant-missing", mode="available", ), ] ) self.session.commit() result = synchronize_templates( self.session, template_ids=(item.id,), dry_run=False, ) self.assertEqual( ["blocked", "created"], sorted(item.status for item in result.outcomes), ) stale = next(item for item in result.outcomes if item.assignment_id == "assignment-stale") self.assertEqual(("unknown_tenant",), stale.blocker_codes) self.assertEqual( ["assignment-valid"], [command.assignment_id for command in self.projection.batches[-1].commands], ) if __name__ == "__main__": unittest.main()