diff --git a/README.md b/README.md index d981382..e85c102 100644 --- a/README.md +++ b/README.md @@ -65,6 +65,11 @@ The SDK's organizational architecture strictly mirrors the Rust version: * **Facet Aggregation & Grouping**: Out-of-the-box support for multi-dimensional facet aggregations, group-bys, and hierarchical data processing. * **Provider Support**: Highly extensible asynchronous database connectivity (integrating third-party async drivers like `aiosqlite` through a unified Transport layer). * **Context & Logging Management**: Built-in support for lifecycle context passing, end-to-end tracing, and SQL execution log interception and dispatch. +* **Governed Mutation Policy**: An application-owned policy can review an + immutable whole-graph plan after Checker/Fix and before the first provider + mutation. Exact policy identity and approval state are retained with audit + evidence. Missing customer policy or approval emits stable warnings without + changing persistence semantics; an explicit denial fails closed. * **TeaQL Federal Protocol Client**: `TeaQLFederalClient` and `TfpHttpProvider` execute governed canonical TFP v1 queries and audited mutations against a remote TeaQL endpoint such as Rust. Direct query execution returns @@ -104,5 +109,25 @@ development-only acknowledgement are maintained in the canonical Opaque tokens never replace the backend's authorization, tenant, ownership, role, or optimistic-version checks. +### Mutation Policy installation + +Policy implementations are installed from trusted application startup through +`UserContext`; request JSON cannot select or replace them. Built-in SQL and TFP +providers enter the same governed boundary. + +```python +context = ( + UserContext.new() + .with_mutation_policy_registry(policy_registry) + .with_mutation_policy_approval_provider(approval_provider) + .with_mutation_governance_sink(warning_sink) +) +``` + +Generated graph saves call `preflight_mutation(...)` for every operation before +the first provider write. See the repeatable +[`examples/mutation-policy`](examples/mutation-policy) example for allow, +approval, audit propagation, and zero-write denial evidence. + --- To run test validations and business logic simulations locally, simply run `pytest` in the project root. diff --git a/examples/mutation-policy/README.md b/examples/mutation-policy/README.md new file mode 100644 index 0000000..1d30287 --- /dev/null +++ b/examples/mutation-policy/README.md @@ -0,0 +1,16 @@ +# Mutation Policy example + +This focused example installs an application-owned policy and exact approval +through `UserContext`, preflights an entire two-entity graph after Checker/Fix, +and proves that a denied graph reaches no persistent provider mutation. The +successful audit events retain the same governance snapshot. + +Run it against the local runtime under development: + +```bash +PYTHONPATH=src python examples/mutation-policy/main.py +``` + +The deterministic in-memory transaction keeps the example independent of a +database. Built-in SQL and TFP providers exercise the same policy boundary in +their focused runtime tests. diff --git a/examples/mutation-policy/main.py b/examples/mutation-policy/main.py new file mode 100644 index 0000000..008ed06 --- /dev/null +++ b/examples/mutation-policy/main.py @@ -0,0 +1,174 @@ +"""Focused Mutation Policy example with an atomic in-memory transaction.""" + +import asyncio +from datetime import datetime, timezone + +from teaql.core.mutation import InsertCommand, MutationRequest, TraceNode +from teaql.data_service import DataServiceOperation, ExecutionMetadata, MutationResult +from teaql.runtime import ( + DelegatingMutationPolicyApprovalProvider, + DelegatingMutationPolicyRegistry, + MutationDecision, + MutationPolicyApproval, + MutationPolicyApprovalStatus, + MutationPolicyError, + MutationPolicyIdentity, + UserContext, +) +from teaql.runtime.audit import MutationAuditKind, RawAuditEvent + + +class OrderPolicy: + identity = MutationPolicyIdentity( + "order-submission", "1", "sha256:order-submission-v1" + ) + + def review(self, context, plan): + if any( + operation.changed_values.get("name").try_text() == "DENIED" + for operation in plan.operations + if operation.changed_values.get("name") is not None + ): + return MutationDecision.denied( + "ORDER_DENIED", "the order policy rejected this graph", "Order.name" + ) + return MutationDecision.allowed() + + +class AuditRecorder: + def __init__(self): + self.events = [] + + async def on_safe_event(self, context, event): + self.events.append(event) + + +class MemoryTransaction: + def __init__(self, provider): + self.provider = provider + self.pending = [] + + async def mutate(self, context, request): + with context.mutation_policy_execution(request): + command = request._data + self.pending.append(command.entity) + await context.send_audit_event( + RawAuditEvent( + MutationAuditKind.CREATED, + command.entity, + command.values.get("id"), + (), + tuple(request.trace_chain()), + context.user_identifier(), + "mutation-policy-example", + context.current_mutation_governance(), + ) + ) + now = datetime.now(timezone.utc) + return MutationResult( + affected_rows=1, + generated_values={}, + persisted_record={ + key: value.val for key, value in command.values.items() + }, + metadata=ExecutionMetadata( + backend="memory-example", + operation=DataServiceOperation.Insert, + started_at=now, + ended_at=now, + affected_rows=1, + ), + ) + + async def commit(self, context): + self.provider.persisted.extend(self.pending) + + async def rollback(self, context): + self.pending.clear() + + +class MemoryProvider: + def __init__(self): + self.persisted = [] + + async def begin(self, context): + return MemoryTransaction(self) + + +def insert(entity, entity_id, name): + command = ( + InsertCommand.new(entity) + .value("id", entity_id) + .value("version", 1) + .value("name", name) + ) + command.trace_chain.append(TraceNode(entity, entity_id, f"create {entity}")) + return command + + +def context_for(provider, audit): + policy = OrderPolicy() + return ( + UserContext.new() + .insert_resource("dataService", provider) + .with_trace_id("python-mutation-policy-example") + .with_app_audit_event_sink(audit) + .with_mutation_policy_registry( + DelegatingMutationPolicyRegistry(lambda request_key: policy) + ) + .with_mutation_policy_approval_provider( + DelegatingMutationPolicyApprovalProvider( + lambda identity: MutationPolicyApproval( + identity, "security-owner", datetime.now(timezone.utc) + ) + ) + ) + ) + + +async def save(context, *commands): + async def graph(): + transaction = context.require_resource("dataService") + for command in commands: + context.preflight_mutation(command) + for command in commands: + await transaction.mutate(context, MutationRequest(command)) + + await context.execute_graph_save(graph) + + +async def main(): + provider = MemoryProvider() + audit = AuditRecorder() + allowed = context_for(provider, audit) + await save( + allowed, + insert("Order", 42, "SUBMITTED"), + insert("OrderLine", 99, "LINE-1"), + ) + + denied_provider = MemoryProvider() + denied = context_for(denied_provider, AuditRecorder()) + try: + await save(denied, insert("Order", 43, "DENIED")) + except MutationPolicyError as error: + assert "ORDER_DENIED" in str(error) + else: + raise AssertionError("denied graph unexpectedly persisted") + + assert provider.persisted == ["Order", "OrderLine"] + assert denied_provider.persisted == [] + assert len(audit.events) == 2 + assert all( + event.mutation_governance.approval_status + == MutationPolicyApprovalStatus.APPROVED + for event in audit.events + ) + print( + "PYTHON_MUTATION_POLICY_PASS " + "allowed_operations=2 denied_provider_mutations=0 audit_events=2" + ) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/scripts/verify-examples.sh b/scripts/verify-examples.sh index 0dc7274..bcb0e4d 100755 --- a/scripts/verify-examples.sh +++ b/scripts/verify-examples.sh @@ -2,7 +2,7 @@ set -euo pipefail repo="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" -expected=(conformance order-management school-management task_board) +expected=(conformance mutation-policy order-management school-management task_board) mapfile -t actual < <(find "$repo/examples" -mindepth 1 -maxdepth 1 -type d -printf '%f\n' | sort) if [[ "${actual[*]}" != "${expected[*]}" ]]; then echo "example inventory changed; update scripts/verify-examples.sh: ${actual[*]}" >&2 @@ -24,6 +24,7 @@ PYTHONPATH="$repo/examples/conformance:$repo/src" python -m app.main PYTHONPATH="$repo/src" python -m unittest discover -s "$repo/examples/conformance" -p 'test_sql_log_intent.py' -v PYTHONPATH="$repo/examples/school-management:$repo/src" python -m app.main PYTHONPATH="$repo/src" python -m unittest discover -s "$repo/examples/school-management" -p 'test_sql_log_intent.py' -v +PYTHONPATH="$repo/src" python "$repo/examples/mutation-policy/main.py" order_management_tmp="$(mktemp -d)" task_board_tmp="$(mktemp -d)" trap 'rm -rf "$order_management_tmp" "$task_board_tmp"' EXIT diff --git a/src/teaql/provider/tfp_client/__init__.py b/src/teaql/provider/tfp_client/__init__.py index 95a9a41..2926837 100644 --- a/src/teaql/provider/tfp_client/__init__.py +++ b/src/teaql/provider/tfp_client/__init__.py @@ -1,5 +1,6 @@ from __future__ import annotations +from contextlib import nullcontext from dataclasses import dataclass, field from datetime import datetime from typing import Any, Awaitable, Callable, Dict, Mapping, Optional @@ -164,27 +165,31 @@ async def query(self, context: Any, request: QueryRequest) -> QueryResult: ) async def mutate(self, context: Any, request: MutationRequest) -> MutationResult: - started_at = datetime.now() - federal = _mutation_request(request) - data = await self.federal_client.execute_mutation(federal) - records = data.get("data") or [] - generated = records[0] if records else {} - affected = int(data.get("affectedRows", 0)) - operation = { - "Create": DataServiceOperation.Insert, - "Update": DataServiceOperation.Update, - "Delete": DataServiceOperation.Delete, - "Recover": DataServiceOperation.Recover, - }[federal.action] - return MutationResult( - affected_rows=affected, generated_values=generated, - persisted_record=generated or None, - metadata=ExecutionMetadata( - backend="teaql-federal", operation=operation, - started_at=started_at, ended_at=datetime.now(), affected_rows=affected, - comment=federal.comment, - ), - ) + if context is not None and not context.consume_mutation_checked(request._data): + context.check_and_fix_mutation(request._data) + scope = context.mutation_policy_execution(request) if context is not None else nullcontext() + with scope: + started_at = datetime.now() + federal = _mutation_request(request) + data = await self.federal_client.execute_mutation(federal) + records = data.get("data") or [] + generated = records[0] if records else {} + affected = int(data.get("affectedRows", 0)) + operation = { + "Create": DataServiceOperation.Insert, + "Update": DataServiceOperation.Update, + "Delete": DataServiceOperation.Delete, + "Recover": DataServiceOperation.Recover, + }[federal.action] + return MutationResult( + affected_rows=affected, generated_values=generated, + persisted_record=generated or None, + metadata=ExecutionMetadata( + backend="teaql-federal", operation=operation, + started_at=started_at, ended_at=datetime.now(), affected_rows=affected, + comment=federal.comment, + ), + ) def _federal_query_payload(query: FederalQuery) -> Dict[str, Any]: diff --git a/src/teaql/runtime/__init__.py b/src/teaql/runtime/__init__.py index e0c46fa..2421648 100644 --- a/src/teaql/runtime/__init__.py +++ b/src/teaql/runtime/__init__.py @@ -10,12 +10,44 @@ from .module import RuntimeModule, DefaultEntityDataServiceBehavior from .store import DataStore from .audit import RawAuditEvent, SafeAuditEvent, MutationAuditKind +from .mutation_policy import ( + MISSING_APPROVAL, + MISSING_POLICY, + DelegatingMutationGovernanceSink, + DelegatingMutationPolicyApprovalProvider, + DelegatingMutationPolicyRegistry, + MutationDecision, + MutationGovernanceEvent, + MutationGovernanceSnapshot, + MutationOperation, + MutationOperationKind, + MutationOperationSummary, + MutationPlan, + MutationPolicyApproval, + MutationPolicyApprovalStatus, + MutationPolicyError, + MutationPolicyIdentity, + MutationPolicySource, + MutationVerdict, +) from .i18n import CheckException, CheckResult, I18nCatalog, JsonFieldNamingProfile, Locale, ObjectLocation, UnsupportedLocaleError from .wire_fields import NormalizedWireInput, WireEntityMetadata, WireFieldMetadata, WireInputError, create_wire_entity_metadata, encode_wire_output, normalize_wire_input, retain_submitted_paths from teaql.core.entity import EntityKey, EntityChangeSet, EntityRoot __all__ = ["WireFieldMetadata", "WireEntityMetadata", "NormalizedWireInput", "WireInputError", "create_wire_entity_metadata", "normalize_wire_input", "encode_wire_output", "retain_submitted_paths", "EntityKey", "EntityChangeSet", "EntityRoot", "ContextEntityRef", "ContextRootError", "CheckException", "CheckResult", "I18nCatalog", "JsonFieldNamingProfile", "Locale", "ObjectLocation", "UnsupportedLocaleError", "UserContext", "TeaqlRuntime", "SqlLogEntry", "SqlLogOperation", "DiagnosticSqlLogSink", "TextDiagnosticSqlLogSink", "ServiceRuntimeFromEnv", "RuntimeModule", "DataStore", "RawAuditEvent", "SafeAuditEvent", "MutationAuditKind", "ContextTools", "ExecutableHttpTool", "HTTP_TOOL", "HttpIntentPhase", "HttpTool", "HttpToolProvider", "ToolDeniedError", "ToolError", "ToolPolicy", "ToolRisk", "Tools", "ToolToken", "ToolUnavailableError"] +__all__ += [ + "MISSING_APPROVAL", "MISSING_POLICY", + "DelegatingMutationGovernanceSink", + "DelegatingMutationPolicyApprovalProvider", + "DelegatingMutationPolicyRegistry", "MutationDecision", + "MutationGovernanceEvent", "MutationGovernanceSnapshot", + "MutationOperation", "MutationOperationKind", "MutationOperationSummary", + "MutationPlan", "MutationPolicyApproval", "MutationPolicyApprovalStatus", + "MutationPolicyError", "MutationPolicyIdentity", "MutationPolicySource", + "MutationVerdict", +] + def __getattr__(name): # Keep provider construction lazy: importing a SQL provider loads runtime diff --git a/src/teaql/runtime/audit.py b/src/teaql/runtime/audit.py index 0d8096a..1097a9a 100644 --- a/src/teaql/runtime/audit.py +++ b/src/teaql/runtime/audit.py @@ -27,6 +27,7 @@ class RawAuditEvent: trace_chain: tuple[Any, ...] = field(default_factory=tuple) actor: Optional[str] = None category: Optional[str] = None + mutation_governance: Any = None def safe(self, mask_fields: List[str], max_length: Optional[int]) -> "SafeAuditEvent": from .log_privacy import REDACTED, credential_name, payload_has_credentials, plaintext_enabled, scrub, value_strings @@ -51,7 +52,7 @@ def safe(self, mask_fields: List[str], max_length: Optional[int]) -> "SafeAuditE intent_values = secrets + value_strings(self.entity_id) return SafeAuditEvent( self.kind, self.entity, self.entity_id, scrub(tuple(fields), secrets), scrub(self.trace_chain, intent_values), - scrub(self.actor, intent_values), self.category, + scrub(self.actor, intent_values), self.category, self.mutation_governance, ) @@ -72,6 +73,7 @@ class SafeAuditEvent: trace_chain: tuple[Any, ...] = field(default_factory=tuple) actor: Optional[str] = None category: Optional[str] = None + mutation_governance: Any = None def _mask(value: str) -> str: diff --git a/src/teaql/runtime/context.py b/src/teaql/runtime/context.py index 5b299d4..0c3d0ca 100644 --- a/src/teaql/runtime/context.py +++ b/src/teaql/runtime/context.py @@ -105,6 +105,8 @@ def __init__(self): self._fix_evidence_current: List[FixEvidence] = [] self._fix_evidence_last: List[FixEvidence] = [] self._checked_mutations = set() + from .mutation_policy import MutationPolicyRuntimeState + self._mutation_policy = MutationPolicyRuntimeState() def begin_fix_evidence(self): self._fix_evidence_current = [] @@ -159,6 +161,7 @@ async def execute_graph_save(self, work): transaction = await begin() owner_token = self._graph_save_owner.set(object()) self._graph_save_active = True + self._mutation_policy.begin_graph() self._graph_commit_actions = [] self._graph_rollback_actions = [] self.insert_resource("fix_time", datetime.now()) @@ -175,6 +178,7 @@ async def execute_graph_save(self, work): raise else: try: + self._mutation_policy.ensure_graph_complete() await self._finish_graph_transaction(transaction, "commit") except BaseException: try: @@ -189,6 +193,7 @@ async def execute_graph_save(self, work): finally: self.insert_resource("dataService", provider) self._graph_save_active = False + self._mutation_policy.end_graph() self._graph_commit_actions = [] self._graph_rollback_actions = [] self._resources.pop("fix_time", None) @@ -639,6 +644,41 @@ def with_app_audit_event_sink(self, sink: Any) -> 'UserContext': self._app_audit_sink = sink return self + def with_mutation_policy_registry(self, registry: Any) -> 'UserContext': + if registry is None or not callable(getattr(registry, "resolve", None)): + raise TypeError("mutation policy registry must expose resolve(request_key)") + self._mutation_policy.registry = registry + return self + + def with_mutation_policy_approval_provider(self, provider: Any) -> 'UserContext': + if provider is None or not callable(getattr(provider, "find_approval", None)): + raise TypeError( + "mutation policy approval provider must expose find_approval(identity)" + ) + self._mutation_policy.approval_provider = provider + return self + + def with_mutation_governance_sink(self, sink: Any) -> 'UserContext': + if sink is None or not callable(getattr(sink, "on_warning", None)): + raise TypeError("mutation governance sink must expose on_warning(context, event)") + self._mutation_policy.warning_sink = sink + return self + + def current_mutation_governance(self): + return self._mutation_policy.current + + def review_mutation_plan(self, plan: Any): + return self._mutation_policy.review(self, plan) + + def preflight_mutation(self, mutation: Any) -> None: + """Check/fix and snapshot one operation for whole-graph policy review.""" + self.check_and_fix_mutation(mutation) + self._mutation_policy.record_preflight(mutation) + + def mutation_policy_execution(self, request: Any): + """Provider boundary scope; entered after validation and before mutation.""" + return self._mutation_policy.enter_mutation(self, request) + def with_sql_log_options(self, options: 'SqlLogOptions') -> 'UserContext': self.insert_resource("sql_log_options", options) return self diff --git a/src/teaql/runtime/mutation_policy.py b/src/teaql/runtime/mutation_policy.py new file mode 100644 index 0000000..78c68f0 --- /dev/null +++ b/src/teaql/runtime/mutation_policy.py @@ -0,0 +1,607 @@ +"""Governed, whole-graph mutation policy contracts. + +The policy receives a detached snapshot after Checker/Fix and before the first +provider mutation. Runtime-specific transaction setup may already have happened; +the portable guarantee is that no provider mutation has executed. +""" + +from __future__ import annotations + +import contextvars +import logging +import threading +from collections import Counter +from copy import deepcopy +from dataclasses import dataclass +from datetime import date, datetime +from decimal import Decimal +from enum import Enum +from types import MappingProxyType +from typing import Any, Callable, Mapping, Optional, Protocol, Sequence + +from teaql.core.mutation import ( + DeleteCommand, + InsertCommand, + MutationRequest, + RecoverCommand, + UpdateCommand, +) +from teaql.core.value import Timestamp, Value + + +MISSING_POLICY = "MUTATION-POLICY-001" +MISSING_APPROVAL = "MUTATION-POLICY-002" + + +class MutationPolicyError(RuntimeError): + pass + + +class MutationOperationKind(str, Enum): + CREATE = "create" + UPDATE = "update" + DELETE = "delete" + RECOVER = "recover" + + +@dataclass(frozen=True) +class MutationPolicyIdentity: + policy_id: str + version: str + fingerprint: str + + def __post_init__(self) -> None: + for name in ("policy_id", "version", "fingerprint"): + value = getattr(self, name) + if not isinstance(value, str) or not value.strip(): + raise ValueError("mutation policy identity values must not be blank") + object.__setattr__(self, name, value.strip()) + + +@dataclass(frozen=True) +class MutationOperation: + kind: MutationOperationKind + entity: str + entity_id: Optional[Value] + original_version: Optional[int] + changed_values: Mapping[str, Value] + + +@dataclass(frozen=True) +class MutationPlan: + execution_id: str + request_key: str + root_entity_type: str + audit_reason: Optional[str] + operations: tuple[MutationOperation, ...] + + +class MutationVerdict(str, Enum): + ALLOW = "allow" + DENY = "deny" + + +@dataclass(frozen=True) +class MutationDecision: + verdict: MutationVerdict + code: Optional[str] = None + message: Optional[str] = None + field_paths: tuple[str, ...] = () + + @classmethod + def allowed(cls) -> "MutationDecision": + return cls(MutationVerdict.ALLOW) + + @classmethod + def denied( + cls, code: str, message: str, *field_paths: str + ) -> "MutationDecision": + if not isinstance(code, str) or not code.strip(): + raise ValueError("a mutation policy denial code is required") + return cls(MutationVerdict.DENY, code.strip(), message, tuple(field_paths)) + + +class MutationPolicy(Protocol): + identity: MutationPolicyIdentity + + def review(self, context: Any, plan: MutationPlan) -> MutationDecision: + ... + + +class MutationPolicyRegistry(Protocol): + def resolve(self, request_key: str) -> Optional[MutationPolicy]: + ... + + +class DelegatingMutationPolicyRegistry: + def __init__(self, resolve: Callable[[str], Optional[MutationPolicy]]): + self._resolve = resolve + + def resolve(self, request_key: str) -> Optional[MutationPolicy]: + return self._resolve(request_key) + + +@dataclass(frozen=True) +class MutationPolicyApproval: + policy: MutationPolicyIdentity + approved_by: str + approved_at: datetime + + def is_valid_for(self, identity: MutationPolicyIdentity) -> bool: + return ( + self.policy == identity + and isinstance(self.approved_by, str) + and bool(self.approved_by.strip()) + and isinstance(self.approved_at, datetime) + and self.approved_at.replace(tzinfo=None) != datetime.min + ) + + +class MutationPolicyApprovalProvider(Protocol): + def find_approval( + self, identity: MutationPolicyIdentity + ) -> Optional[MutationPolicyApproval]: + ... + + +class DelegatingMutationPolicyApprovalProvider: + def __init__( + self, + find: Callable[[MutationPolicyIdentity], Optional[MutationPolicyApproval]], + ): + self._find = find + + def find_approval( + self, identity: MutationPolicyIdentity + ) -> Optional[MutationPolicyApproval]: + return self._find(identity) + + +class MutationPolicySource(str, Enum): + GENERATED_DEFAULT = "generated_default" + CUSTOMER = "customer" + + +class MutationPolicyApprovalStatus(str, Enum): + NOT_APPLICABLE = "not_applicable" + MISSING = "missing" + APPROVED = "approved" + + +@dataclass(frozen=True) +class MutationOperationSummary: + kind: MutationOperationKind + entity: str + entity_id: Optional[Value] + changed_fields: tuple[str, ...] + + +@dataclass(frozen=True) +class MutationGovernanceSnapshot: + execution_id: str + request_key: str + source: MutationPolicySource + policy: Optional[MutationPolicyIdentity] + approval_status: MutationPolicyApprovalStatus + warning_codes: tuple[str, ...] + operations: tuple[MutationOperationSummary, ...] + + +@dataclass(frozen=True) +class MutationGovernanceEvent: + snapshot: MutationGovernanceSnapshot + warning_code: str + first_occurrence: bool + + +class MutationGovernanceSink(Protocol): + def on_warning(self, context: Any, warning: MutationGovernanceEvent) -> None: + ... + + +class DelegatingMutationGovernanceSink: + def __init__(self, warning: Callable[[Any, MutationGovernanceEvent], None]): + self._warning = warning + + def on_warning(self, context: Any, warning: MutationGovernanceEvent) -> None: + self._warning(context, warning) + + +class _TextMutationGovernanceSink: + def on_warning(self, context: Any, warning: MutationGovernanceEvent) -> None: + if not warning.first_occurrence: + return + logging.getLogger("teaql.mutation_policy").warning( + "TeaQL mutation policy warning code=%s request_key=%s source=%s approval=%s", + warning.warning_code, + warning.snapshot.request_key, + warning.snapshot.source.value, + warning.snapshot.approval_status.value, + ) + + +class MutationPolicyRuntimeState: + _sequence = 0 + _sequence_lock = threading.Lock() + + def __init__(self) -> None: + self.registry: Optional[MutationPolicyRegistry] = None + self.approval_provider: Optional[MutationPolicyApprovalProvider] = None + self.warning_sink: MutationGovernanceSink = _TextMutationGovernanceSink() + self._emitted_warnings: set[str] = set() + self._warning_lock = threading.Lock() + self._active = contextvars.ContextVar( + f"teaql_mutation_policy_{id(self)}", default=None + ) + self._graph_active = False + self._graph_reviewed = False + self._preflight: list[MutationOperation] = [] + self._root_entity: Optional[str] = None + self._audit_reason: Optional[str] = None + self._remaining: Counter[Any] = Counter() + + @property + def current(self) -> Optional[MutationGovernanceSnapshot]: + return self._active.get() + + def begin_graph(self) -> None: + self._graph_active = True + self._graph_reviewed = False + self._preflight.clear() + self._root_entity = None + self._audit_reason = None + self._remaining.clear() + self._active.set(None) + + def end_graph(self) -> None: + self._graph_active = False + self._graph_reviewed = False + self._preflight.clear() + self._root_entity = None + self._audit_reason = None + self._remaining.clear() + self._active.set(None) + + def record_preflight(self, command: Any) -> None: + if not self._graph_active: + return + if self._graph_reviewed: + raise MutationPolicyError( + "mutation preflight cannot add operations after policy review" + ) + operations = _operations_from_data(command) + if not operations: + raise MutationPolicyError("mutation preflight must contain an operation") + self._root_entity = self._root_entity or operations[0].entity + self._audit_reason = self._audit_reason or _comment_from_data(command) + self._preflight.extend(operations) + + def enter_mutation(self, context: Any, request: MutationRequest): + operations = _operations_from_data(request._data) + if not operations: + raise MutationPolicyError("mutation request must contain an operation") + # Legacy/manual graph transactions cannot describe the complete graph + # before their first provider call. Keep those writes governed by the + # generated-default policy one operation at a time. Installing a + # customer policy still requires a complete preflight and therefore + # remains fail-closed. + if self._graph_active and self.registry is None and not self._preflight: + plan = self._build_plan( + context, operations[0].entity, _request_comment(request), tuple(operations) + ) + token = self._active.set(self.review(context, plan)) + return _ResetScope(self._active, token) + if self._graph_active: + if not self._graph_reviewed: + if self.registry is not None and not self._preflight: + raise MutationPolicyError( + "customer mutation policy requires complete graph preflight " + "before provider mutation" + ) + planned = tuple(self._preflight or operations) + root = self._root_entity or planned[0].entity + reason = self._audit_reason or _request_comment(request) + snapshot = self.review(context, self._build_plan(context, root, reason, planned)) + self._active.set(snapshot) + self._remaining = Counter(_operation_signature(item) for item in planned) + self._graph_reviewed = True + self._consume_planned(operations) + return _NoopScope() + + plan = self._build_plan( + context, operations[0].entity, _request_comment(request), tuple(operations) + ) + token = self._active.set(self.review(context, plan)) + return _ResetScope(self._active, token) + + def ensure_graph_complete(self) -> None: + if self._graph_reviewed and self._remaining: + raise MutationPolicyError( + "reviewed mutation plan contains operations that were not executed" + ) + + def review(self, context: Any, plan: MutationPlan) -> MutationGovernanceSnapshot: + _validate_plan(plan) + policy = self.registry.resolve(plan.request_key) if self.registry else None + if policy is None: + source = MutationPolicySource.GENERATED_DEFAULT + identity = None + approval = MutationPolicyApprovalStatus.NOT_APPLICABLE + warnings = (MISSING_POLICY,) + else: + identity = policy.identity + if not isinstance(identity, MutationPolicyIdentity): + raise MutationPolicyError("customer mutation policy identity is invalid") + decision = policy.review(context, _clone_plan(plan)) + if not isinstance(decision, MutationDecision): + raise MutationPolicyError("customer mutation policy returned an invalid decision") + if decision.verdict == MutationVerdict.DENY: + raise MutationPolicyError( + f"[MUTATION POLICY DENIED] " + f"{decision.code or 'MUTATION-POLICY-DENIED'}: " + f"{decision.message or 'mutation rejected'}" + ) + if decision.verdict != MutationVerdict.ALLOW: + raise MutationPolicyError("customer mutation policy returned an invalid verdict") + source = MutationPolicySource.CUSTOMER + found = ( + self.approval_provider.find_approval(identity) + if self.approval_provider + else None + ) + approval = ( + MutationPolicyApprovalStatus.APPROVED + if found is not None and found.is_valid_for(identity) + else MutationPolicyApprovalStatus.MISSING + ) + warnings = () if approval == MutationPolicyApprovalStatus.APPROVED else (MISSING_APPROVAL,) + + snapshot = MutationGovernanceSnapshot( + execution_id=plan.execution_id, + request_key=plan.request_key, + source=source, + policy=identity, + approval_status=approval, + warning_codes=warnings, + operations=tuple( + MutationOperationSummary( + item.kind, + item.entity, + _clone_value(item.entity_id) if item.entity_id else None, + tuple(sorted(item.changed_values)), + ) + for item in plan.operations + ), + ) + for warning in warnings: + self._emit_warning(context, snapshot, warning) + return snapshot + + def _consume_planned(self, operations: Sequence[MutationOperation]) -> None: + for operation in operations: + signature = _operation_signature(operation) + if self._remaining[signature] <= 0: + raise MutationPolicyError( + "provider mutation is not present in the reviewed graph plan" + ) + self._remaining[signature] -= 1 + if self._remaining[signature] == 0: + del self._remaining[signature] + + def _build_plan( + self, + context: Any, + root: str, + reason: Optional[str], + operations: tuple[MutationOperation, ...], + ) -> MutationPlan: + with self._sequence_lock: + type(self)._sequence += 1 + sequence = type(self)._sequence + trace_id = getattr(context, "trace_id", lambda: "")() + return MutationPlan( + execution_id=f"{trace_id or 'teaql'}-mutation-{sequence}", + request_key=f"{root}.saveGraph", + root_entity_type=root, + audit_reason=reason, + operations=tuple(_clone_operation(item) for item in operations), + ) + + def _emit_warning( + self, + context: Any, + snapshot: MutationGovernanceSnapshot, + warning_code: str, + ) -> None: + identity = ( + "none" + if snapshot.policy is None + else f"{snapshot.policy.policy_id}:{snapshot.policy.version}:" + f"{snapshot.policy.fingerprint}" + ) + key = f"{snapshot.request_key}|{identity}|{warning_code}" + with self._warning_lock: + first = key not in self._emitted_warnings + self._emitted_warnings.add(key) + try: + self.warning_sink.on_warning( + context, + MutationGovernanceEvent(snapshot, warning_code, first), + ) + except BaseException: + # Warning delivery must not change business persistence semantics. + pass + + +class _NoopScope: + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + return False + + +class _ResetScope: + def __init__(self, variable: contextvars.ContextVar, token: contextvars.Token): + self._variable = variable + self._token = token + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + self._variable.reset(self._token) + return False + + +def _operations_from_data(data: Any) -> list[MutationOperation]: + if isinstance(data, MutationRequest): + return _operations_from_data(data._data) + if isinstance(data, list): + operations: list[MutationOperation] = [] + for item in data: + operations.extend(_operations_from_data(item)) + return operations + if isinstance(data, InsertCommand): + return [MutationOperation( + MutationOperationKind.CREATE, + data.entity, + _clone_value(data.values.get("id")), + None, + _readonly_values(data.values), + )] + if isinstance(data, UpdateCommand): + return [MutationOperation( + MutationOperationKind.UPDATE, + data.entity, + _clone_value(data.id), + data.expected_version_val, + _readonly_values(data.values), + )] + if isinstance(data, DeleteCommand): + return [MutationOperation( + MutationOperationKind.DELETE, + data.entity, + _clone_value(data.id), + data.expected_version_val, + MappingProxyType({}), + )] + if isinstance(data, RecoverCommand): + return [MutationOperation( + MutationOperationKind.RECOVER, + data.entity, + _clone_value(data.id), + data.expected_version_val, + MappingProxyType({}), + )] + raise MutationPolicyError(f"unsupported mutation command {type(data).__name__}") + + +def _comment_from_data(data: Any) -> Optional[str]: + if isinstance(data, MutationRequest): + return _request_comment(data) + if isinstance(data, list): + for item in data: + comment = _comment_from_data(item) + if comment: + return comment + return None + traces = getattr(data, "trace_chain", ()) + return traces[-1].comment if traces else None + + +def _request_comment(request: MutationRequest) -> Optional[str]: + comment = getattr(request, "comment", None) + return comment() if callable(comment) else comment + + +def _readonly_values(values: Mapping[str, Value]) -> Mapping[str, Value]: + return MappingProxyType({key: _clone_value(value) for key, value in values.items()}) + + +def _clone_value(value: Optional[Value]) -> Optional[Value]: + if value is None: + return None + raw = value.val + if hasattr(raw, "id"): + raw = getattr(raw, "id") + else: + try: + raw = deepcopy(raw) + except BaseException: + raw = repr(raw) + return Value(raw, getattr(value, "_type_hint", None)) + + +def _clone_operation(operation: MutationOperation) -> MutationOperation: + return MutationOperation( + operation.kind, + operation.entity, + _clone_value(operation.entity_id), + operation.original_version, + _readonly_values(operation.changed_values), + ) + + +def _clone_plan(plan: MutationPlan) -> MutationPlan: + return MutationPlan( + plan.execution_id, + plan.request_key, + plan.root_entity_type, + plan.audit_reason, + tuple(_clone_operation(item) for item in plan.operations), + ) + + +def _validate_plan(plan: MutationPlan) -> None: + if not isinstance(plan.execution_id, str) or not plan.execution_id.strip(): + raise MutationPolicyError("mutation plan execution id is required") + if not isinstance(plan.request_key, str) or not plan.request_key.strip(): + raise MutationPolicyError("mutation plan request key is required") + if not isinstance(plan.root_entity_type, str) or not plan.root_entity_type.strip(): + raise MutationPolicyError("mutation plan root entity type is required") + if not plan.operations: + raise MutationPolicyError("mutation plan must contain an operation") + if any(not item.entity for item in plan.operations): + raise MutationPolicyError("mutation operation entity type is required") + + +def _operation_signature(operation: MutationOperation): + return ( + operation.kind.value, + operation.entity, + _freeze_value(operation.entity_id), + operation.original_version, + tuple( + (key, _freeze_value(value)) + for key, value in sorted(operation.changed_values.items()) + ), + ) + + +def _freeze_value(value: Optional[Value]): + if value is None: + return None + raw = value.val + hint = getattr(getattr(value, "_type_hint", None), "name", None) + return (hint, _freeze_raw(raw)) + + +def _freeze_raw(value: Any): + if isinstance(value, Value): + return _freeze_value(value) + if isinstance(value, Mapping): + return tuple((str(key), _freeze_raw(item)) for key, item in sorted(value.items())) + if isinstance(value, (list, tuple)): + return tuple(_freeze_raw(item) for item in value) + if isinstance(value, Timestamp): + return ("timestamp", value.millis) + if isinstance(value, (date, datetime)): + return (type(value).__name__, value.isoformat()) + if isinstance(value, Decimal): + return ("decimal", str(value)) + if hasattr(value, "id"): + return (type(value).__name__, _freeze_raw(getattr(value, "id"))) + if isinstance(value, (str, int, float, bool, type(None))): + return value + return (type(value).__name__, repr(value)) diff --git a/src/teaql/sql/executor.py b/src/teaql/sql/executor.py index a8450b1..59b00b6 100644 --- a/src/teaql/sql/executor.py +++ b/src/teaql/sql/executor.py @@ -53,6 +53,13 @@ def _intent_bindings(compiled, request): sql_origin='generated') +class _NoopContextManager: + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + return False + def _canonical_id_set_value(value): if isinstance(value, Value): return ("Value", str(value._type_hint), _canonical_id_set_value(value.val)) @@ -670,16 +677,23 @@ async def mutate(self, context: 'UserContext', request: MutationRequest) -> Muta entity = getattr(request._data, "entity", "unknown") kind = type(request._data).__name__.replace("Command", "").lower() if context is not None: - context.check_and_fix_mutation(request._data) + if not context.consume_mutation_checked(request._data): + context.check_and_fix_mutation(request._data) telemetry = context.runtime_telemetry() if context is not None else None - return await observe_runtime_operation( - telemetry, - RuntimeOperation("mutation", f"{entity}.{kind}", { - "teaql.entity.type": entity, - "teaql.mutation.kind": kind, - }), - lambda: self._mutate(context, request), + scope = ( + context.mutation_policy_execution(request) + if context is not None + else _NoopContextManager() ) + with scope: + return await observe_runtime_operation( + telemetry, + RuntimeOperation("mutation", f"{entity}.{kind}", { + "teaql.entity.type": entity, + "teaql.mutation.kind": kind, + }), + lambda: self._mutate(context, request), + ) async def _mutate(self, context: 'UserContext', request: MutationRequest) -> MutationResult: if isinstance(self.transport, SqlTransactionTransport): @@ -855,6 +869,7 @@ async def _mutate(self, context: 'UserContext', request: MutationRequest) -> Mut tuple(request.trace_chain()), context.user_identifier(), context.get_resource("bootstrapCategory"), + context.current_mutation_governance(), )) return MutationResult( affected_rows=affected_rows, diff --git a/tests/provider/test_tfp_client.py b/tests/provider/test_tfp_client.py index ac2ad34..2709941 100644 --- a/tests/provider/test_tfp_client.py +++ b/tests/provider/test_tfp_client.py @@ -12,6 +12,13 @@ FederalMutation, FederalQuery, TeaQLFederalClient, TfpError, TfpHttpProvider, _reject_trusted_fields, ) +from teaql.runtime import ( + DelegatingMutationPolicyRegistry, + MutationDecision, + MutationPolicyError, + MutationPolicyIdentity, + UserContext, +) class RecordingTelemetry: @@ -121,6 +128,45 @@ async def handler(request): assert payloads[1]["expectedVersion"] == 3 +@pytest.mark.asyncio +async def test_provider_policy_denial_prevents_remote_mutation_request(): + calls = 0 + + async def handler(_): + nonlocal calls + calls += 1 + return httpx.Response(200, json={"affectedRows": 1}) + + class DenyAllPolicy: + identity = MutationPolicyIdentity( + "tfp-order-policy", "1", "sha256:tfp-order-policy-v1" + ) + + def review(self, context, plan): + assert plan.request_key == "CustomerOrder.saveGraph" + return MutationDecision.denied( + "REMOTE_MUTATION_DENIED", "remote mutation is disabled" + ) + + http = httpx.AsyncClient( + transport=httpx.MockTransport(handler), base_url="https://tfp.test" + ) + provider = TfpHttpProvider("https://tfp.test", client=http) + context = UserContext.new().with_mutation_policy_registry( + DelegatingMutationPolicyRegistry(lambda _: DenyAllPolicy()) + ) + command = ( + UpdateCommand.new("CustomerOrder", 42) + .expected_version(3) + .value("status", "PAID") + ) + command.trace_chain.append(TraceNode(comment="Mark paid")) + + with pytest.raises(MutationPolicyError, match="REMOTE_MUTATION_DENIED"): + await provider.mutate(context, MutationRequest.Update(command)) + assert calls == 0 + + @pytest.mark.asyncio async def test_fails_closed_for_trusted_fields_errors_and_streaming(): async def handler(_): diff --git a/tests/runtime/test_mutation_policy.py b/tests/runtime/test_mutation_policy.py new file mode 100644 index 0000000..ef516b6 --- /dev/null +++ b/tests/runtime/test_mutation_policy.py @@ -0,0 +1,371 @@ +from datetime import datetime, timezone + +import pytest + +from teaql.core.mutation import InsertCommand, MutationRequest, TraceNode +from teaql.core.value import Value +from teaql.data_service import ( + DataServiceOperation, + ExecutionMetadata, + MutationResult, +) +from teaql.runtime import ( + MISSING_APPROVAL, + MISSING_POLICY, + DelegatingMutationPolicyApprovalProvider, + DelegatingMutationPolicyRegistry, + MutationDecision, + MutationOperation, + MutationOperationKind, + MutationPlan, + MutationPolicyApproval, + MutationPolicyApprovalStatus, + MutationPolicyError, + MutationPolicyIdentity, + MutationPolicySource, + UserContext, +) +from teaql.runtime.audit import MutationAuditKind, RawAuditEvent +from teaql.sql.executor import SqlDataServiceExecutor + + +class TestPolicy: + __test__ = False + + def __init__(self, identity, review): + self.identity = identity + self._review = review + + def review(self, context, plan): + return self._review(plan) + + +class RecordingWarnings: + def __init__(self, failure=None): + self.events = [] + self.failure = failure + + def on_warning(self, context, warning): + self.events.append(warning) + if self.failure: + raise self.failure + + +class RecordingAudit: + def __init__(self): + self.events = [] + + async def on_safe_event(self, context, event): + self.events.append(event) + + +class RecordingTransaction: + """Runs the real SqlDataServiceExecutor.mutate boundary without a database.""" + + def __init__(self, owner): + self.owner = owner + self.executor = object.__new__(SqlDataServiceExecutor) + self.executor._sync_generated_schema = lambda context: None + self.executor._mutate = self._mutate + + async def mutate(self, context, request): + return await self.executor.mutate(context, request) + + async def _mutate(self, context, request): + self.owner.mutations += 1 + command = request._data + await context.send_audit_event(RawAuditEvent( + MutationAuditKind.CREATED, + command.entity, + command.values.get("id"), + (), + tuple(request.trace_chain()), + context.user_identifier(), + "mutation-policy-test", + context.current_mutation_governance(), + )) + now = datetime.now(timezone.utc) + return MutationResult( + affected_rows=1, + generated_values={}, + persisted_record={ + key: value.val for key, value in command.values.items() + }, + metadata=ExecutionMetadata( + backend="test", + operation=DataServiceOperation.Insert, + started_at=now, + ended_at=now, + affected_rows=1, + ), + ) + + async def commit(self, context): + self.owner.commits += 1 + + async def rollback(self, context): + self.owner.rollbacks += 1 + + +class RecordingProvider: + def __init__(self): + self.begins = 0 + self.mutations = 0 + self.commits = 0 + self.rollbacks = 0 + + async def begin(self, context): + self.begins += 1 + return RecordingTransaction(self) + + +def test_generated_default_and_exact_approval_warning_semantics(): + warnings = RecordingWarnings() + context = UserContext.new().with_mutation_governance_sink(warnings) + + first = context.review_mutation_plan(_plan("one")) + second = context.review_mutation_plan(_plan("two")) + assert first.source == MutationPolicySource.GENERATED_DEFAULT + assert first.warning_codes == (MISSING_POLICY,) + assert second.approval_status == MutationPolicyApprovalStatus.NOT_APPLICABLE + assert [event.first_occurrence for event in warnings.events] == [True, False] + + identity = MutationPolicyIdentity("orders", "3", "sha256:orders-v3") + policy = TestPolicy(identity, lambda plan: MutationDecision.allowed()) + customer = (UserContext.new() + .with_mutation_policy_registry( + DelegatingMutationPolicyRegistry(lambda _: policy)) + .with_mutation_governance_sink(warnings)) + missing = customer.review_mutation_plan(_plan("missing")) + assert missing.warning_codes == (MISSING_APPROVAL,) + + wrong = MutationPolicyIdentity("orders", "3", "sha256:wrong") + customer.with_mutation_policy_approval_provider( + DelegatingMutationPolicyApprovalProvider( + lambda _: MutationPolicyApproval( + wrong, "security-owner", datetime.now(timezone.utc)))) + mismatched = customer.review_mutation_plan(_plan("mismatched")) + assert mismatched.approval_status == MutationPolicyApprovalStatus.MISSING + + customer.with_mutation_policy_approval_provider( + DelegatingMutationPolicyApprovalProvider( + lambda candidate: MutationPolicyApproval( + candidate, "security-owner", datetime.min))) + default_time = customer.review_mutation_plan(_plan("default-time")) + assert default_time.approval_status == MutationPolicyApprovalStatus.MISSING + + customer.with_mutation_policy_approval_provider( + DelegatingMutationPolicyApprovalProvider( + lambda candidate: MutationPolicyApproval( + candidate, "security-owner", datetime.now(timezone.utc)))) + approved = customer.review_mutation_plan(_plan("approved")) + assert approved.approval_status == MutationPolicyApprovalStatus.APPROVED + assert approved.warning_codes == () + + +@pytest.mark.asyncio +async def test_complete_graph_policy_audit_and_immutable_preflight(): + observed = {} + identity = MutationPolicyIdentity("orders", "1", "sha256:orders") + + def review(plan): + observed["count"] = len(plan.operations) + observed["name"] = plan.operations[0].changed_values["name"].try_text() + return MutationDecision.allowed() + + audit = RecordingAudit() + provider = RecordingProvider() + context = _context(provider, TestPolicy(identity, review), audit=audit, approval=True) + order = _insert("Order", 42, "DRAFT") + line = _insert("OrderLine", 99, "LINE") + + async def save_graph(): + context.preflight_mutation(order) + context.preflight_mutation(line) + order.values["name"] = Value.Text("APPROVED") + # The policy still receives DRAFT, and an execution may only consume the + # original reviewed operation rather than the changed command. + transaction = context.require_resource("dataService") + await transaction.mutate( + context, MutationRequest(_insert("Order", 42, "DRAFT"))) + return await transaction.mutate(context, MutationRequest(line)) + + await context.execute_graph_save(save_graph) + assert observed == {"count": 2, "name": "DRAFT"} + assert provider.mutations == 2 + assert provider.commits == 1 + assert provider.rollbacks == 0 + assert len(audit.events) == 2 + assert all(event.mutation_governance.policy == identity for event in audit.events) + assert all( + event.mutation_governance.approval_status + == MutationPolicyApprovalStatus.APPROVED + for event in audit.events + ) + + +@pytest.mark.asyncio +async def test_denial_and_missing_preflight_leave_zero_provider_mutations(): + identity = MutationPolicyIdentity("orders", "1", "sha256:deny") + denied_provider = RecordingProvider() + denied = _context( + denied_provider, + TestPolicy( + identity, + lambda plan: MutationDecision.denied( + "ORDER_DENIED", "orders disabled", "Order.name")), + ) + order = _insert("Order", 43, "DENIED") + line = _insert("OrderLine", 100, "DENIED-LINE") + + async def denied_graph(): + denied.preflight_mutation(order) + denied.preflight_mutation(line) + return await denied.require_resource("dataService").mutate( + denied, MutationRequest(order)) + + with pytest.raises(MutationPolicyError, match="ORDER_DENIED"): + await denied.execute_graph_save(denied_graph) + assert denied_provider.begins == 1 + assert denied_provider.mutations == 0 + assert denied_provider.rollbacks == 1 + + missing_provider = RecordingProvider() + missing = _context( + missing_provider, + TestPolicy(identity, lambda plan: MutationDecision.allowed()), + ) + with pytest.raises(MutationPolicyError, match="complete graph preflight"): + await missing.execute_graph_save( + lambda: missing.require_resource("dataService").mutate( + missing, MutationRequest(_insert("Order", 44, "MISSING")))) + assert missing_provider.mutations == 0 + assert missing_provider.rollbacks == 1 + + +@pytest.mark.asyncio +async def test_generated_default_allows_legacy_graph_without_complete_preflight(): + provider = RecordingProvider() + warnings = RecordingWarnings() + context = (UserContext.new() + .insert_resource("dataService", provider) + .with_mutation_governance_sink(warnings)) + + async def legacy_graph(): + transaction = context.require_resource("dataService") + await transaction.mutate( + context, MutationRequest(_insert("Order", 47, "FIRST"))) + await transaction.mutate( + context, MutationRequest(_insert("OrderLine", 102, "SECOND"))) + + await context.execute_graph_save(legacy_graph) + assert provider.mutations == 2 + assert provider.commits == 1 + assert provider.rollbacks == 0 + assert [event.warning_code for event in warnings.events] == [MISSING_POLICY, MISSING_POLICY] + + +@pytest.mark.asyncio +async def test_unplanned_and_incomplete_operations_fail_closed(): + identity = MutationPolicyIdentity("orders", "1", "sha256:strict") + unplanned_provider = RecordingProvider() + unplanned = _context( + unplanned_provider, + TestPolicy(identity, lambda plan: MutationDecision.allowed()), + ) + + async def unplanned_graph(): + unplanned.preflight_mutation(_insert("Order", 45, "PLANNED")) + return await unplanned.require_resource("dataService").mutate( + unplanned, MutationRequest(_insert("Order", 45, "DIFFERENT"))) + + with pytest.raises(MutationPolicyError, match="not present"): + await unplanned.execute_graph_save(unplanned_graph) + assert unplanned_provider.mutations == 0 + assert unplanned_provider.rollbacks == 1 + + incomplete_provider = RecordingProvider() + incomplete = _context( + incomplete_provider, + TestPolicy(identity, lambda plan: MutationDecision.allowed()), + ) + first = _insert("Order", 46, "FIRST") + second = _insert("OrderLine", 101, "SECOND") + + async def incomplete_graph(): + incomplete.preflight_mutation(first) + incomplete.preflight_mutation(second) + return await incomplete.require_resource("dataService").mutate( + incomplete, MutationRequest(first)) + + with pytest.raises(MutationPolicyError, match="were not executed"): + await incomplete.execute_graph_save(incomplete_graph) + assert incomplete_provider.mutations == 1 + assert incomplete_provider.commits == 0 + assert incomplete_provider.rollbacks == 1 + + +@pytest.mark.asyncio +async def test_warning_sink_failure_is_fail_open(): + identity = MutationPolicyIdentity("orders", "1", "sha256:warning") + provider = RecordingProvider() + warnings = RecordingWarnings(RuntimeError("warning sink unavailable")) + context = _context( + provider, + TestPolicy(identity, lambda plan: MutationDecision.allowed()), + warning_sink=warnings, + ) + order = _insert("Order", 47, "ALLOWED") + + async def graph(): + context.preflight_mutation(order) + return await context.require_resource("dataService").mutate( + context, MutationRequest(order)) + + await context.execute_graph_save(graph) + assert warnings.events[0].warning_code == MISSING_APPROVAL + assert provider.mutations == 1 + assert provider.commits == 1 + + +def _context(provider, policy, audit=None, approval=False, warning_sink=None): + context = (UserContext.new() + .insert_resource("dataService", provider) + .with_trace_id("python-mutation-policy") + .with_mutation_policy_registry( + DelegatingMutationPolicyRegistry(lambda _: policy))) + if audit: + context.with_app_audit_event_sink(audit) + if approval: + context.with_mutation_policy_approval_provider( + DelegatingMutationPolicyApprovalProvider( + lambda candidate: MutationPolicyApproval( + candidate, "security-owner", datetime.now(timezone.utc)))) + if warning_sink: + context.with_mutation_governance_sink(warning_sink) + return context + + +def _plan(execution_id): + return MutationPlan( + execution_id, + "Order.saveGraph", + "Order", + "submit order", + ( + MutationOperation( + MutationOperationKind.UPDATE, + "Order", + Value.I64(42), + 7, + {"name": Value.Text("Updated")}, + ), + ), + ) + + +def _insert(entity, entity_id, name): + command = InsertCommand.new(entity).value("id", entity_id).value("name", name) + command.value("version", 1) + command.trace_chain.append(TraceNode(entity, entity_id, f"create {entity}")) + return command