Files

379 lines
13 KiB
Python

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()