from __future__ import annotations from datetime import UTC, datetime from types import SimpleNamespace import unittest from sqlalchemy import create_engine from sqlalchemy.orm import Session from govoplan_core.auth import ApiPrincipal from govoplan_core.core.access import PrincipalRef from govoplan_core.core.search import ( SearchAuthorizationRequest, SearchBackfillRequest, SearchResourceReference, ) from govoplan_core.db.base import Base from govoplan_risk_compliance.backend.assurance import ( AssuranceEdgeInput, AssuranceNodeInput, RiskAssuranceAccessError, RiskAssuranceConflictError, RiskAssuranceError, RiskAssuranceNotFoundError, assurance_graph, assurance_summary, create_assurance_edge, create_assurance_node, list_assurance_edge_history, list_assurance_node_history, list_assurance_nodes, revise_assurance_edge, revise_assurance_node, ) from govoplan_risk_compliance.backend.db.models import ( RiskAssuranceEdge, RiskAssuranceNode, ) from govoplan_risk_compliance.backend.permissions import ( READ_SCOPE, WRITE_SCOPE, ) from govoplan_risk_compliance.backend.search_source import ( PROVIDER_ID, RESOURCE_TYPE, RiskAssuranceSearchSource, ) NOW = datetime(2026, 8, 1, 12, 0, tzinfo=UTC) TABLES = ( RiskAssuranceNode.__table__, RiskAssuranceEdge.__table__, ) def principal( tenant_id: str = "tenant-1", *, scopes: tuple[str, ...] = (READ_SCOPE, WRITE_SCOPE), ) -> ApiPrincipal: return ApiPrincipal( principal=PrincipalRef( account_id=f"account-{tenant_id}", membership_id=f"membership-{tenant_id}", tenant_id=tenant_id, scopes=frozenset(scopes), ), account=SimpleNamespace(id=f"account-{tenant_id}"), user=SimpleNamespace(id=f"account-{tenant_id}"), ) class AssuranceGraphTests(unittest.TestCase): def setUp(self) -> None: self.engine = create_engine("sqlite:///:memory:") Base.metadata.create_all(self.engine, tables=TABLES) self.session = Session(self.engine) def tearDown(self) -> None: self.session.close() self.engine.dispose() def test_non_sanctions_control_uses_complete_revisioned_graph(self) -> None: node_specs = ( ("retention-obligation", "obligation", "active"), ("customer-register", "governed_object", "active"), ("over-retention-risk", "risk", "identified"), ("retention-control", "control", "implemented"), ("retention-evidence", "evidence", "current"), ("retention-finding", "finding", "open"), ("retention-measure", "corrective_measure", "planned"), ("retention-review", "effectiveness_review", "pending"), ) nodes = {} for stable_id, kind, state in node_specs: nodes[stable_id] = create_assurance_node( self.session, principal(), value=AssuranceNodeInput( stable_id=stable_id, kind=kind, label=stable_id.replace("-", " ").title(), state=state, owner_ref="organizations:function:data-governance", scope_ref="datasource:customer-register", governed_object_ref=( "datasources:customer-register" if kind == "governed_object" else None ), valid_from=NOW, provenance={"fixture": "non-sanctions"}, legal_basis_refs=("law:retention:2026",), ), ) relations = ( ("retention-obligation", "applies_to", "customer-register"), ("customer-register", "exposes_risk", "over-retention-risk"), ("over-retention-risk", "mitigated_by", "retention-control"), ("retention-control", "evidenced_by", "retention-evidence"), ("retention-evidence", "results_in", "retention-finding"), ("retention-finding", "addressed_by", "retention-measure"), ("retention-measure", "reviewed_by", "retention-review"), ) edges = [] for index, (source, relation, target) in enumerate(relations, start=1): edges.append( create_assurance_edge( self.session, principal(), value=AssuranceEdgeInput( stable_id=f"retention-edge-{index}", source_node_ref=source, target_node_ref=target, relation=relation, owner_ref="organizations:function:data-governance", valid_from=NOW, provenance={"fixture": "non-sanctions"}, ), ) ) self.session.commit() graph = assurance_graph( self.session, principal(), root_ref="retention-obligation", max_depth=8, limit=100, ) summary = assurance_summary(self.session, principal()) self.assertEqual(8, len(graph.nodes)) self.assertEqual(7, len(graph.edges)) self.assertFalse(graph.truncated) self.assertEqual(8, summary["node_count"]) self.assertEqual(7, summary["edge_count"]) self.assertEqual(1, summary["by_kind"]["effectiveness_review"]) self.assertEqual("control", nodes["retention-control"].kind) self.assertEqual("reviewed_by", edges[-1].relation) def test_node_and_edge_revisions_are_immutable_and_occ_guarded(self) -> None: risk = self._node("risk-1", "risk", "identified") control = self._node("control-1", "control", "implemented") edge = create_assurance_edge( self.session, principal(), value=AssuranceEdgeInput( stable_id="edge-1", source_node_ref=risk.stable_id, target_node_ref=control.stable_id, relation="mitigated_by", owner_ref="function:risk-owner", valid_from=NOW, ), ) revised = revise_assurance_node( self.session, principal(), stable_id=risk.stable_id, expected_revision=1, value=AssuranceNodeInput( stable_id=risk.stable_id, kind="risk", label="Risk 1", state="assessed", owner_ref="function:risk-owner", valid_from=NOW, evidence_refs=("evidence:assessment-1",), ), ) revised_edge = revise_assurance_edge( self.session, principal(), stable_id=edge.stable_id, expected_revision=1, value=AssuranceEdgeInput( stable_id=edge.stable_id, source_node_ref=risk.stable_id, target_node_ref=control.stable_id, relation="mitigated_by", state="suspended", owner_ref="function:risk-owner", valid_from=NOW, evidence_refs=("evidence:suspension-1",), ), ) self.session.commit() self.assertEqual(2, revised.revision) self.assertEqual(2, revised_edge.revision) self.assertEqual( [2, 1], [ item.revision for item in list_assurance_node_history( self.session, principal(), stable_id="risk-1", ) ], ) self.assertEqual( [2, 1], [ item.revision for item in list_assurance_edge_history( self.session, principal(), stable_id="edge-1", ) ], ) with self.assertRaisesRegex( RiskAssuranceConflictError, "current revision is 2", ): revise_assurance_node( self.session, principal(), stable_id="risk-1", expected_revision=1, value=AssuranceNodeInput( stable_id="risk-1", kind="risk", label="Risk 1", state="closed", owner_ref="function:risk-owner", valid_from=NOW, ), ) def test_access_and_tenant_boundaries_are_enforced(self) -> None: self._node("risk-1", "risk", "identified") self.session.commit() self.assertEqual( (), list_assurance_nodes(self.session, principal("tenant-2")), ) with self.assertRaises(RiskAssuranceNotFoundError): assurance_graph( self.session, principal("tenant-2"), root_ref="risk-1", ) with self.assertRaises(RiskAssuranceAccessError): list_assurance_nodes( self.session, principal(scopes=()), ) def test_programmatic_inputs_enforce_text_and_provenance_bounds(self) -> None: with self.assertRaisesRegex(RiskAssuranceError, "description is limited"): AssuranceNodeInput( stable_id="risk-oversized-description", kind="risk", label="Oversized risk", state="identified", owner_ref="function:risk-owner", valid_from=NOW, description="x" * 20_001, ) with self.assertRaisesRegex(RiskAssuranceError, "provenance is limited"): AssuranceNodeInput( stable_id="risk-oversized-provenance", kind="risk", label="Oversized provenance", state="identified", owner_ref="function:risk-owner", valid_from=NOW, provenance={"payload": "x" * 100_001}, ) def test_relation_shape_is_validated(self) -> None: self._node("risk-1", "risk", "identified") self._node("evidence-1", "evidence", "current") with self.assertRaisesRegex( RiskAssuranceConflictError, "requires risk -> control", ): create_assurance_edge( self.session, principal(), value=AssuranceEdgeInput( stable_id="invalid-edge", source_node_ref="risk-1", target_node_ref="evidence-1", relation="mitigated_by", owner_ref="function:risk-owner", valid_from=NOW, ), ) def test_search_backfill_and_authorization_are_tenant_safe(self) -> None: self._node("risk-1", "risk", "identified") self.session.commit() provider = RiskAssuranceSearchSource() page = provider.backfill( self.session, request=SearchBackfillRequest( tenant_id="tenant-1", provider_id=PROVIDER_ID, resource_type=RESOURCE_TYPE, rebuild_id="rebuild-1", ), ) request = SearchAuthorizationRequest( reference=SearchResourceReference( tenant_id="tenant-1", module_id="risk_compliance", resource_type=RESOURCE_TYPE, resource_id="risk-1", ), source_revision="1", ) self.assertEqual(1, len(page.documents)) self.assertTrue( provider.authorize( self.session, principal(), requests=(request,), )[request.reference.key] ) self.assertFalse( provider.authorize( self.session, principal("tenant-2"), requests=(request,), )[request.reference.key] ) def _node( self, stable_id: str, kind: str, state: str, ) -> RiskAssuranceNode: return create_assurance_node( self.session, principal(), value=AssuranceNodeInput( stable_id=stable_id, kind=kind, label=stable_id.replace("-", " ").title(), state=state, owner_ref="function:risk-owner", valid_from=NOW, ), ) if __name__ == "__main__": unittest.main()