diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 6a171a3..b3db688 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -58,7 +58,7 @@ jobs: uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: repository: fruwehq/determa-state-conformance - ref: 263644f951f342b0eeaa3aceef4877293d2d7c67 + ref: 5ba78c7ef90b8556e76de6481a18b23a3d0c2378 path: .pinned/determa-state-conformance - name: Check out pinned specification uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 diff --git a/README.md b/README.md index 5561916..5565318 100644 --- a/README.md +++ b/README.md @@ -5,8 +5,8 @@ a language-agnostic statechart engine with a shared normative conformance suite. This implementation supports Determa State `format: 1` at specification commit `cc4b0d734aa1c5953de75fb53b63e390a3b72761`. Correctness is determined by the -111-case core suite, persistence profiles, and 85-vector execution-checkpoint profile -at conformance commit `263644f951f342b0eeaa3aceef4877293d2d7c67`. +111-case core suite, persistence profiles, and 91-vector execution-checkpoint profile +at conformance commit `5ba78c7ef90b8556e76de6481a18b23a3d0c2378`. Version `0.1.0` is the published synchronized release of the specification, conformance suite, Python engine, and Rust engine. @@ -132,6 +132,10 @@ the exact supplied state object. semantic validation path. Native values must satisfy the same portable Unicode and numeric domain as source documents. +Category-specific `StrEnum` definitions are used by production emitters. +`PORTABLE_CODE_SETS` is the immutable category-to-string mapping derived from those +definitions. + ## Persist And Migrate `serialize_aggregate` produces the canonical ยง16 aggregate artifact. Restoration diff --git a/conformance/execution_checkpoint.py b/conformance/execution_checkpoint.py index 2f096b0..b3b4320 100644 --- a/conformance/execution_checkpoint.py +++ b/conformance/execution_checkpoint.py @@ -3,7 +3,10 @@ from __future__ import annotations import copy +import hashlib import json +from collections.abc import Iterator +from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path from typing import Any @@ -18,6 +21,7 @@ ExecutionStore, ExecutionStoreError, ExecutionStoreRegistry, + ExecutionStoreTransaction, MemoryArtifactResolver, MemoryExecutionStore, load_bundle, @@ -25,6 +29,7 @@ restore_execution_checkpoint, serialize_execution_checkpoint, ) +from determa.state.wire import hash_value from .harness import conformance_root @@ -246,6 +251,54 @@ def health(self) -> dict[str, Any]: return {"healthy": True} +class _ObservedTransaction(ExecutionStoreTransaction): + def __init__( + self, transaction: ExecutionStoreTransaction, calls: list[str] + ) -> None: + self._transaction = transaction + self._calls = calls + + @property + def root_instance_id(self) -> str: + return self._transaction.root_instance_id + + def load(self) -> bytes | None: + self._calls.append("load_checkpoint") + return self._transaction.load() + + def insert(self, checkpoint: bytes) -> bool: + self._calls.append("insert_checkpoint") + return self._transaction.insert(checkpoint) + + def replace( + self, + expected_revision: str, + expected_checkpoint_digest: str, + checkpoint: bytes, + ) -> bool: + self._calls.append("compare_and_swap_checkpoint") + return self._transaction.replace( + expected_revision, expected_checkpoint_digest, checkpoint + ) + + +class _ObservedMemoryStore(MemoryExecutionStore): + def __init__(self, initial: dict[str, bytes], calls: list[str]) -> None: + super().__init__(initial) + self._calls = calls + + @property + def root_instance_ids(self) -> frozenset[str]: + return frozenset(self._records) + + @contextmanager + def transaction( + self, root_instance_id: str + ) -> Iterator[ExecutionStoreTransaction]: + with super().transaction(root_instance_id) as transaction: + yield _ObservedTransaction(transaction, self._calls) + + def _adapter_operation(vector: dict[str, Any]) -> dict[str, Any]: operation = vector["operation"] if operation == "inject_execution_store": @@ -387,10 +440,177 @@ def _expected_response( raise AssertionError(f"no exact response projection for {operation}") +def _scope_state( + case: ExecutionCheckpointCase, reference: dict[str, str] +) -> dict[str, Any]: + return copy.deepcopy( + _pointer(_json(case.path / reference["file"]), reference["pointer"]) + ) + + +def _outbox_records(checkpoint: dict[str, Any]) -> list[dict[str, str]]: + records = [ + ("pending", record) for record in checkpoint["pending_outbox_intents"] + ] + records.extend( + ("terminal", record) for record in checkpoint["terminal_outbox_records"] + ) + records.extend( + ("tombstone", record) for record in checkpoint["outbox_effect_tombstones"] + ) + return [ + { + "effect_id": ( + record["effect_id"] + if kind == "tombstone" + else record["intent"]["effect_id"] + ), + "record_kind": kind, + "source_digest": hash_value(record), + } + for kind, record in records + ] + + +def _expected_scope_maps( + case: ExecutionCheckpointCase, + expected: dict[str, Any], +) -> dict[str, dict[str, Any]]: + result = {} + for scope in expected["scopes"]: + checkpoints = {} + outbox_records = {} + for root_instance_id, binding in scope["checkpoints"].items(): + expected_checkpoint = _json(case.path / binding["file"]) + canonical = serialize_execution_checkpoint(expected_checkpoint) + assert ( + f"sha256:{hashlib.sha256(canonical).hexdigest()}" + == binding["serialization_digest"] + ) + assert ( + expected_checkpoint["execution_checkpoint_digest"] + == binding["execution_checkpoint_digest"] + ) + checkpoints[root_instance_id] = canonical + outbox_records[root_instance_id] = _outbox_records(expected_checkpoint) + assert [ + record + for root_records in outbox_records.values() + for record in root_records + ] == scope["outbox_records"] + result[scope["logical_scope_id"]] = { + "checkpoints": checkpoints, + "outbox_records": outbox_records, + } + return result + + +def _actual_scope_maps( + hosts: dict[str, ExecutionHost], + stores: dict[str, _ObservedMemoryStore], +) -> dict[str, dict[str, Any]]: + result = {} + for scope_id, store in stores.items(): + checkpoints = {} + outbox_records = {} + for root_instance_id in sorted(store.root_instance_ids): + restored = hosts[scope_id].read_checkpoint(root_instance_id) + assert restored is not None + checkpoints[root_instance_id] = restored.canonical_bytes + outbox_records[root_instance_id] = _outbox_records(restored.document) + result[scope_id] = { + "checkpoints": checkpoints, + "outbox_records": outbox_records, + } + return result + + +def _assert_complete_scope_maps( + actual: dict[str, dict[str, Any]], + expected: dict[str, dict[str, Any]], +) -> None: + assert actual == expected + + +def _scope_hosts( + case: ExecutionCheckpointCase, + state: dict[str, Any], + store_calls: list[str], +) -> tuple[dict[str, ExecutionHost], dict[str, _ObservedMemoryStore]]: + hosts = {} + stores = {} + for scope in state["scopes"]: + initial = { + root_instance_id: (case.path / binding["file"]).read_bytes() + for root_instance_id, binding in scope["checkpoints"].items() + } + scope_id = scope["logical_scope_id"] + store = _ObservedMemoryStore(initial, store_calls) + stores[scope_id] = store + hosts[scope_id] = ExecutionHost(store, _resolver(case)) + return hosts, stores + + +def _run_scope_vector(item: ExecutionCheckpointVector) -> None: + case = item.case + vector = item.vector + expected = vector["expect"] + before = _scope_state(case, vector["scope_state_before"]) + after = _scope_state(case, vector["scope_state_after"]) + store_calls: list[str] = [] + hosts, stores = _scope_hosts(case, before, store_calls) + + resolver_calls = ["resolve_execution_store_scope"] + selection = vector["scope_selection"] + requested = selection["requested_scope_id"] + candidates = selection["candidates"] + selected = ( + requested + if requested is not None + and len(candidates) == 1 + and candidates[0]["logical_scope_id"] == requested + and candidates[0]["authorized"] + else None + ) + host_calls = [] + if selected is not None: + host_calls.append("update_pending_outbox") + selected_scope = next( + scope for scope in before["scopes"] if scope["logical_scope_id"] == selected + ) + root_instance_id, binding = next(iter(selected_scope["checkpoints"].items())) + checkpoint = _json(case.path / binding["file"]) + hosts[selected].update_pending_outbox( + root_instance_id, + vector["effect_id"], + vector["desired_pending_state"], + expected_revision=checkpoint["revision"], + expected_checkpoint_digest=checkpoint["execution_checkpoint_digest"], + ) + + assert expected["selection_result"] == ( + "selected" if selected is not None else "rejected" + ) + assert expected["selected_scope_id"] == selected + assert expected["calls"] == { + "resolver": resolver_calls, + "execution_host": host_calls, + "store": store_calls, + "core": [], + } + _assert_complete_scope_maps( + _actual_scope_maps(hosts, stores), + _expected_scope_maps(case, after), + ) + + def run_execution_checkpoint_vector(item: ExecutionCheckpointVector) -> None: case = item.case vector = item.vector expected = vector["expect"] + if "scope_state_before" in vector: + _run_scope_vector(item) + return before_name = vector.get("checkpoint_before") after_name = expected["checkpoint_after"] if vector["operation"] in { diff --git a/conformance/pins.py b/conformance/pins.py index 4e3cbaf..bd9024f 100644 --- a/conformance/pins.py +++ b/conformance/pins.py @@ -4,7 +4,7 @@ from pathlib import Path -CONFORMANCE_COMMIT = "263644f951f342b0eeaa3aceef4877293d2d7c67" +CONFORMANCE_COMMIT = "5ba78c7ef90b8556e76de6481a18b23a3d0c2378" SPEC_COMMIT = "cc4b0d734aa1c5953de75fb53b63e390a3b72761" ROOT = Path(__file__).resolve().parent.parent diff --git a/conformance/test_conformance.py b/conformance/test_conformance.py index 71895bb..e883955 100644 --- a/conformance/test_conformance.py +++ b/conformance/test_conformance.py @@ -9,17 +9,22 @@ import pytest from jsonschema import Draft202012Validator -from determa.state import load_bundle +from determa.state import PORTABLE_CODE_SETS, load_bundle from determa.state.validator import schema as bundled_schema from determa.state.wire import artifact_schema from .execution_checkpoint import ( + _actual_scope_maps, + _assert_complete_scope_maps, + _expected_scope_maps, + _scope_hosts, + _scope_state, execution_checkpoint_cases, execution_checkpoint_vectors, run_execution_checkpoint_vector, validate_execution_checkpoint_artifact, ) -from .harness import CORE_DIR, CoreCase, core_cases, run_case +from .harness import CORE_DIR, CoreCase, conformance_root, core_cases, run_case from .persistence import persistence_vector_cases, run_persistence_vectors from .persistence_profiles import ( persistence_profile_cases, @@ -46,7 +51,102 @@ def _spec_root() -> Path | None: def test_suite_present() -> None: assert CORE_DIR.exists(), "pinned conformance suite is unavailable" assert len(core_cases()) == 111 - assert len(execution_checkpoint_vectors()) == 85 + assert len(execution_checkpoint_vectors()) == 91 + + +def test_portable_code_sets_match_authoritative_registry() -> None: + vector_path = ( + conformance_root() + / "conformance" + / "closed-code-registry" + / "vectors.generated.json" + ) + vectors = json.loads(vector_path.read_text(encoding="utf-8")) + expected = { + category["id"]: frozenset(category["codes"]) + for category in vectors["categories"] + if category["id"] != "execution_store_failure" + } + + missing_categories = sorted(set(expected) - set(PORTABLE_CODE_SETS)) + extra_categories = sorted(set(PORTABLE_CODE_SETS) - set(expected)) + mismatches = [] + for category in sorted(set(expected) & set(PORTABLE_CODE_SETS)): + missing = sorted(expected[category] - PORTABLE_CODE_SETS[category]) + extra = sorted(PORTABLE_CODE_SETS[category] - expected[category]) + if missing or extra: + mismatches.append(f"{category}: missing={missing}, extra={extra}") + + assert not (missing_categories or extra_categories or mismatches), "\n".join( + [ + f"categories: missing={missing_categories}, extra={extra_categories}", + *mismatches, + ] + ) + + +def _scope_map_pair() -> tuple[dict, dict]: + item = next( + item + for item in execution_checkpoint_vectors() + if item.vector["name"] == "pending_outbox_update_in_scope_a" + ) + before = _scope_state(item.case, item.vector["scope_state_before"]) + hosts, stores = _scope_hosts(item.case, before, []) + selected_scope = next( + scope + for scope in before["scopes"] + if scope["logical_scope_id"] == "scope-a" + ) + root_instance_id, binding = next(iter(selected_scope["checkpoints"].items())) + checkpoint = json.loads( + (item.case.path / binding["file"]).read_text(encoding="utf-8") + ) + hosts["scope-a"].update_pending_outbox( + root_instance_id, + item.vector["effect_id"], + item.vector["desired_pending_state"], + expected_revision=checkpoint["revision"], + expected_checkpoint_digest=checkpoint["execution_checkpoint_digest"], + ) + after = _scope_state(item.case, item.vector["scope_state_after"]) + return ( + _actual_scope_maps(hosts, stores), + _expected_scope_maps(item.case, after), + ) + + +def test_scope_map_rejects_omitted_unchanged_scope_records() -> None: + actual, expected = _scope_map_pair() + root_instance_id = next(iter(expected["scope-b"]["checkpoints"])) + expected["scope-b"]["checkpoints"].pop(root_instance_id) + expected["scope-b"]["outbox_records"].pop(root_instance_id) + + with pytest.raises(AssertionError): + _assert_complete_scope_maps(actual, expected) + + +def test_scope_map_rejects_unexpected_checkpoint() -> None: + actual, expected = _scope_map_pair() + actual["scope-a"]["checkpoints"]["unexpected-root"] = b"unexpected" + + with pytest.raises(AssertionError): + _assert_complete_scope_maps(actual, expected) + + +def test_scope_map_rejects_unexpected_outbox_record() -> None: + actual, expected = _scope_map_pair() + root_instance_id = next(iter(actual["scope-a"]["outbox_records"])) + actual["scope-a"]["outbox_records"][root_instance_id].append( + { + "effect_id": "unexpected-effect", + "record_kind": "pending", + "source_digest": "sha256:unexpected", + } + ) + + with pytest.raises(AssertionError): + _assert_complete_scope_maps(actual, expected) def test_bundled_schema_matches_pinned_spec() -> None: diff --git a/src/determa/state/__init__.py b/src/determa/state/__init__.py index db8c3c2..a708faa 100644 --- a/src/determa/state/__init__.py +++ b/src/determa/state/__init__.py @@ -14,6 +14,19 @@ validate_execution_checkpoint_member, validate_execution_checkpoint_semantics, ) +from .codes import ( + PORTABLE_CODE_SETS, + CheckpointArtifactFailureCode, + CheckpointHostFailureCode, + CheckpointPreAcceptanceFailureCode, + CreationRejectionCode, + DispatchRejectionCode, + DispositionCode, + EngineFaultCode, + ExecutionStoreAdapterFailureCode, + MachineLoadFailureCode, + PersistenceFailureCode, +) from .definition import Bundle, BundleSource, load_bundle from .engine import Delivery, Result, create, dispatch from .errors import ( @@ -92,14 +105,21 @@ "Bundle", "BundleSource", "CelError", + "CheckpointArtifactFailureCode", + "CheckpointHostFailureCode", + "CheckpointPreAcceptanceFailureCode", "COMPACT_EFFECT_IDENTITY_RETENTION", "DURABLE_CONCURRENT", "DURABLE_SINGLE_WRITER", "DetermaError", "DefinitionResolver", "Delivery", + "CreationRejectionCode", + "DispatchRejectionCode", + "DispositionCode", "ErrorRecord", "EPHEMERAL", + "EngineFaultCode", "ExecutionHost", "ExecutionHostError", "ExecutionStore", @@ -107,9 +127,11 @@ "ExecutionStoreFactory", "ExecutionStoreRegistry", "ExecutionStoreTransaction", + "ExecutionStoreAdapterFailureCode", "FileExecutionStore", "MemoryArtifactResolver", "MemoryExecutionStore", + "MachineLoadFailureCode", "MigrationDescriptorResolver", "MigrationDispatchResult", "MigrationFailure", @@ -118,6 +140,8 @@ "PERMANENT_OUTBOX_TERMINAL_RETENTION", "PERMANENT_RECEIPT_RETENTION", "PostgreSQLExecutionStore", + "PORTABLE_CODE_SETS", + "PersistenceFailureCode", "RESTART_PERSISTENT", "ROOT_IDENTITY_RETENTION", "Result", diff --git a/src/determa/state/checkpoint.py b/src/determa/state/checkpoint.py index de034fe..b7fcab4 100644 --- a/src/determa/state/checkpoint.py +++ b/src/determa/state/checkpoint.py @@ -9,6 +9,12 @@ from functools import cache from typing import Any +from .codes import ( + CheckpointArtifactFailureCode as CheckpointCode, +) +from .codes import ( + PersistenceFailureCode as PersistenceCode, +) from .errors import ArtifactError from .wire import ( ArtifactSource, @@ -57,7 +63,7 @@ def serialize_execution_checkpoint(document: Mapping[str, Any]) -> bytes: def _invalid() -> ArtifactError: - return ArtifactError("invalid_execution_checkpoint") + return ArtifactError(CheckpointCode.INVALID_EXECUTION_CHECKPOINT) @cache @@ -520,14 +526,14 @@ def restore_execution_checkpoint( ) except ArtifactError as exc: if exc.code in { - "source_definition_unavailable", - "definition_untrusted", - "definition_fingerprint_mismatch", + PersistenceCode.SOURCE_DEFINITION_UNAVAILABLE, + PersistenceCode.DEFINITION_UNTRUSTED, + PersistenceCode.DEFINITION_FINGERPRINT_MISMATCH, }: raise raise _invalid() from exc if execution_checkpoint_digest(document) != document["execution_checkpoint_digest"]: - raise ArtifactError("execution_checkpoint_digest_mismatch") + raise ArtifactError(CheckpointCode.EXECUTION_CHECKPOINT_DIGEST_MISMATCH) validate_execution_checkpoint_semantics(document) return RestoredExecutionCheckpoint( document=copy.deepcopy(document), diff --git a/src/determa/state/codes.py b/src/determa/state/codes.py new file mode 100644 index 0000000..290a6d9 --- /dev/null +++ b/src/determa/state/codes.py @@ -0,0 +1,153 @@ +"""Closed portable code sets implemented by Determa State Python.""" + +from __future__ import annotations + +from collections.abc import Mapping +from enum import StrEnum +from types import MappingProxyType + + +class CheckpointArtifactFailureCode(StrEnum): + EXECUTION_CHECKPOINT_DIGEST_MISMATCH = "execution_checkpoint_digest_mismatch" + INVALID_EXECUTION_CHECKPOINT = "invalid_execution_checkpoint" + UNSUPPORTED_EXECUTION_CHECKPOINT_FORMAT = "unsupported_execution_checkpoint_format" + UNSUPPORTED_EXECUTION_CHECKPOINT_SCHEMA_VERSION = ( + "unsupported_execution_checkpoint_schema_version" + ) + + +class CheckpointHostFailureCode(StrEnum): + CHECKPOINT_REVISION_CONFLICT = "checkpoint_revision_conflict" + CREATION_ID_CONFLICT = "creation_id_conflict" + CREATION_REJECTED = "creation_rejected" + EFFECT_ID_CONFLICT = "effect_id_conflict" + EVENT_ID_CONFLICT = "event_id_conflict" + INJECTED_PRE_COMMIT_FAILURE = "injected_pre_commit_failure" + INVALID_EXECUTION_CHECKPOINT = "invalid_execution_checkpoint" + OPERATION_ID_CONFLICT = "operation_id_conflict" + PHYSICAL_DELETION_UNSUPPORTED = "physical_deletion_unsupported" + RESPONSE_LOST_AFTER_COMMIT = "response_lost_after_commit" + + +class CheckpointPreAcceptanceFailureCode(StrEnum): + DELIVERY_DIGEST_MISMATCH = "delivery_digest_mismatch" + EVENT_ID_CONFLICT = "event_id_conflict" + INVALID_DELIVERY_MODE = "invalid_delivery_mode" + INVALID_DELIVERY_ORIGIN = "invalid_delivery_origin" + MALFORMED_DELIVERY = "malformed_delivery" + TOMBSTONED_ROOT = "tombstoned_root" + WRONG_ROOT = "wrong_root" + + +class CreationRejectionCode(StrEnum): + INVALID_BINDING = "invalid_binding" + INVALID_CREATION_REQUEST = "invalid_creation_request" + INVALID_MACHINE_TARGET = "invalid_machine_target" + + +class DispatchRejectionCode(StrEnum): + INACTIVE_COMPONENT_TARGET = "inactive_component_target" + INCOMPATIBLE_BUNDLE = "incompatible_bundle" + INVALID_CORRELATION = "invalid_correlation" + INVALID_EVENT = "invalid_event" + INVALID_INSTANCE_TARGET = "invalid_instance_target" + INVALID_PAYLOAD = "invalid_payload" + INVALID_PRIOR_STATE = "invalid_prior_state" + + +class DispositionCode(StrEnum): + FAULTED = "faulted" + HANDLED = "handled" + REJECTED = "rejected" + UNHANDLED = "unhandled" + + +class EngineFaultCode(StrEnum): + ACTION_FAULT = "action_fault" + BINDING_NOT_EMPTY = "binding_not_empty" + CASCADE_FAULT = "cascade_fault" + CONTAINED_RUNTIME_FAULT = "contained_runtime_fault" + GUARD_FAULT = "guard_fault" + INACTIVE_COMPONENT_TARGET = "inactive_component_target" + INVALID_INSTANCE_TARGET = "invalid_instance_target" + INVARIANT_FAULT = "invariant_fault" + + +class ExecutionStoreAdapterFailureCode(StrEnum): + ADAPTER_CAPABILITY_MISMATCH = "adapter_capability_mismatch" + DUPLICATE_ADAPTER_REGISTRATION = "duplicate_adapter_registration" + INVALID_ADAPTER_CONFIGURATION = "invalid_adapter_configuration" + UNKNOWN_ADAPTER = "unknown_adapter" + + +class MachineLoadFailureCode(StrEnum): + CEL_PROFILE_ERROR = "cel_profile_error" + DESTROYED_REFERENCE_BINDING = "destroyed_reference_binding" + DESTROYED_VARIABLE_WRITE = "destroyed_variable_write" + DUPLICATE_KEY = "duplicate_key" + INVALID_BINDING = "invalid_binding" + INVALID_BOOLEAN_SYNTAX = "invalid_boolean_syntax" + INVALID_NULL_SYNTAX = "invalid_null_syntax" + INVALID_NUMERIC_SYNTAX = "invalid_numeric_syntax" + INVALID_UNICODE = "invalid_unicode" + NON_JSON_VALUE = "non_json_value" + NON_STRING_MAP_KEY = "non_string_map_key" + NUMERIC_VALUE_OUT_OF_RANGE = "numeric_value_out_of_range" + ROOT_LOCAL_TRANSITION = "root_local_transition" + ROOT_REENTRY = "root_reentry" + SEMANTIC_VALIDATION = "semantic_validation" + UNSUPPORTED_FORMAT = "unsupported_format" + UNSUPPORTED_YAML_FEATURE = "unsupported_yaml_feature" + + +class PersistenceFailureCode(StrEnum): + AGGREGATE_STATE_DIGEST_MISMATCH = "aggregate_state_digest_mismatch" + DEFINITION_FINGERPRINT_MISMATCH = "definition_fingerprint_mismatch" + DEFINITION_UNTRUSTED = "definition_untrusted" + INVALID_AGGREGATE_STATE = "invalid_aggregate_state" + INVALID_AGGREGATE_STATE_PACKAGE = "invalid_aggregate_state_package" + INVALID_MIGRATION_DESCRIPTOR = "invalid_migration_descriptor" + INVALID_MIGRATION_REQUEST = "invalid_migration_request" + MIGRATION_DESCRIPTOR_UNTRUSTED = "migration_descriptor_untrusted" + MIGRATION_RESOURCE_LIMIT_EXCEEDED = "migration_resource_limit_exceeded" + MIGRATION_ROUTE_MISMATCH = "migration_route_mismatch" + MIGRATION_ROUTE_MISSING = "migration_route_missing" + MIGRATION_TOTALITY_FAILURE = "migration_totality_failure" + MIGRATION_TRANSFORM_FAULT = "migration_transform_fault" + SOURCE_DEFINITION_UNAVAILABLE = "source_definition_unavailable" + TARGET_DEFINITION_UNAVAILABLE = "target_definition_unavailable" + TERMINAL_MIGRATION_REJECTED = "terminal_migration_rejected" + TERMINAL_MIGRATION_REQUIRES_MAINTENANCE = "terminal_migration_requires_maintenance" + UNSUPPORTED_AGGREGATE_STATE_FORMAT = "unsupported_aggregate_state_format" + UNSUPPORTED_AGGREGATE_STATE_PACKAGE_FORMAT = "unsupported_aggregate_state_package_format" + UNSUPPORTED_AGGREGATE_STATE_PACKAGE_SCHEMA_VERSION = ( + "unsupported_aggregate_state_package_schema_version" + ) + UNSUPPORTED_AGGREGATE_STATE_SCHEMA_VERSION = "unsupported_aggregate_state_schema_version" + UNSUPPORTED_MIGRATION_DESCRIPTOR_FORMAT = "unsupported_migration_descriptor_format" + UNSUPPORTED_MIGRATION_DESCRIPTOR_SCHEMA_VERSION = ( + "unsupported_migration_descriptor_schema_version" + ) + + +_CATEGORY_TYPES: Mapping[str, type[StrEnum]] = MappingProxyType( + { + "checkpoint_artifact_failure": CheckpointArtifactFailureCode, + "checkpoint_host_failure": CheckpointHostFailureCode, + "checkpoint_pre_acceptance_failure": CheckpointPreAcceptanceFailureCode, + "creation_rejection": CreationRejectionCode, + "dispatch_rejection": DispatchRejectionCode, + "disposition": DispositionCode, + "engine_fault": EngineFaultCode, + "execution_store_adapter_failure": ExecutionStoreAdapterFailureCode, + "machine_load_failure": MachineLoadFailureCode, + "persistence_failure": PersistenceFailureCode, + } +) + +PORTABLE_CODE_SETS: Mapping[str, frozenset[str]] = MappingProxyType( + { + category: frozenset(code.value for code in code_type) + for category, code_type in _CATEGORY_TYPES.items() + } +) diff --git a/src/determa/state/definition.py b/src/determa/state/definition.py index bb342ee..b631cbf 100644 --- a/src/determa/state/definition.py +++ b/src/determa/state/definition.py @@ -12,6 +12,7 @@ from typing import Any from . import yaml12 +from .codes import MachineLoadFailureCode as LoadCode from .errors import ValidationError BundleSource = str | Mapping[str, Any] @@ -33,7 +34,7 @@ def _normalize_typed_literal(declaration: dict[str, Any], member: str) -> None: value = float(value) if isinstance(value, float): if not math.isfinite(value): - raise ValidationError("numeric_value_out_of_range") + raise ValidationError(LoadCode.NUMERIC_VALUE_OUT_OF_RANGE) value = 0.0 if value == 0.0 else value declaration[member] = value @@ -129,7 +130,7 @@ def _typed_value(value: Any) -> list[Any]: if isinstance(value, dict): entries = [[key, _typed_value(value[key])] for key in sorted(value, key=_utf8_key)] return ["map", entries] - raise ValidationError("non_json_value") + raise ValidationError(LoadCode.NON_JSON_VALUE) def canonical_json(value: Any) -> str: @@ -173,13 +174,13 @@ def load_bundle(source: BundleSource) -> Bundle: document = copy.deepcopy(dict(source)) yaml12.validate_portable_values(document) if not yaml12.validate_unicode(document): - raise ValidationError("invalid_unicode") + raise ValidationError(LoadCode.INVALID_UNICODE) else: - raise ValidationError("non_json_value") + raise ValidationError(LoadCode.NON_JSON_VALUE) if not isinstance(document, dict): raise ValidationError("structural_validation") if document.get("format") != 1 or isinstance(document.get("format"), bool): - raise ValidationError("unsupported_format") + raise ValidationError(LoadCode.UNSUPPORTED_FORMAT) from .validator import validate validate(document) diff --git a/src/determa/state/engine.py b/src/determa/state/engine.py index 9026dda..1d41e25 100644 --- a/src/determa/state/engine.py +++ b/src/determa/state/engine.py @@ -9,6 +9,21 @@ from typing import Any, Literal, cast from . import cel +from .codes import ( + CreationRejectionCode as CreationCode, +) +from .codes import ( + DispatchRejectionCode as DispatchCode, +) +from .codes import ( + DispositionCode as Disposition, +) +from .codes import ( + EngineFaultCode as FaultCode, +) +from .codes import ( + MachineLoadFailureCode as LoadCode, +) from .definition import Bundle, BundleSource, _escape_pointer, hash_identity, load_bundle from .errors import CelError, StepFault, ValidationError from .model import BundleModel, MachineModel, StateNode @@ -279,19 +294,19 @@ def create( or not validate_unicode([machine_id, root_instance_id, creation_id]) ): result = _empty_result(status="rejected", state=None, disposition=None) - result["rejection"] = {"code": "invalid_creation_request"} + result["rejection"] = {"code": CreationCode.INVALID_CREATION_REQUEST.value} return result models = BundleModel(validated) if machine_id not in models.machines: result = _empty_result(status="rejected", state=None, disposition=None) - result["rejection"] = {"code": "invalid_machine_target"} + result["rejection"] = {"code": CreationCode.INVALID_MACHINE_TARGET.value} return result machine = models.machine(machine_id) try: root_bindings = _creation_bindings(machine, bindings or {}) except ValueError: result = _empty_result(status="rejected", state=None, disposition=None) - result["rejection"] = {"code": "invalid_binding"} + result["rejection"] = {"code": CreationCode.INVALID_BINDING.value} return result root_id = _root_runtime_id(validated, machine.raw, root_instance_id) state: dict[str, Any] = { @@ -366,15 +381,17 @@ def dispatch( if isinstance(prior_state, dict) else "faulted", state=prior_state, - disposition="rejected", + disposition=Disposition.REJECTED.value, ) - result["rejection"] = {"code": "invalid_prior_state"} + result["rejection"] = {"code": DispatchCode.INVALID_PRIOR_STATE.value} return result if prior_state["validated_bundle_fingerprint"] != validated.fingerprint: result = _empty_result( - status=prior_state["status"], state=prior_state, disposition="rejected" + status=prior_state["status"], + state=prior_state, + disposition=Disposition.REJECTED.value, ) - result["rejection"] = {"code": "incompatible_bundle"} + result["rejection"] = {"code": DispatchCode.INCOMPATIBLE_BUNDLE.value} result["fault"] = copy.deepcopy(prior_state.get("fault")) return result if delivery is None: @@ -382,7 +399,7 @@ def dispatch( result["fault"] = copy.deepcopy(prior_state.get("fault")) return result if not isinstance(delivery, dict) or set(delivery) not in ({"input"}, {"internal"}): - return _rejected(prior_state, "invalid_event") + return _rejected(prior_state, DispatchCode.INVALID_EVENT) mode = next(iter(delivery)) envelope = delivery[mode] models = BundleModel(validated) @@ -431,13 +448,17 @@ def dispatch( state["fault"] = copy.deepcopy(runtime["fault"]) else: execution.emit_failure(runtime, str(envelope["event_id"])) - result = _empty_result(status=state["status"], state=state, disposition="faulted") + result = _empty_result( + status=state["status"], state=state, disposition=Disposition.FAULTED.value + ) result["fault"] = copy.deepcopy(runtime["fault"]) result["emissions"] = execution.emissions return result if not handled: result = _empty_result( - status=prior_state["status"], state=prior_state, disposition="unhandled" + status=prior_state["status"], + state=prior_state, + disposition=Disposition.UNHANDLED.value, ) result["fault"] = copy.deepcopy(prior_state.get("fault")) return result @@ -445,16 +466,22 @@ def dispatch( root = state["runtimes"][state["root_runtime_id"]] state["status"] = root["status"] state["fault"] = copy.deepcopy(root.get("fault")) - result = _empty_result(status=state["status"], state=state, disposition="handled") + result = _empty_result( + status=state["status"], state=state, disposition=Disposition.HANDLED.value + ) result["emissions"] = execution.emissions result["fault"] = copy.deepcopy(root.get("fault")) if state["status"] == "faulted" else None return result -def _rejected(prior_state: dict[str, Any], code: str) -> Result: - result = _empty_result(status=prior_state["status"], state=prior_state, disposition="rejected") +def _rejected(prior_state: dict[str, Any], code: DispatchCode) -> Result: + result = _empty_result( + status=prior_state["status"], + state=prior_state, + disposition=Disposition.REJECTED.value, + ) result["fault"] = copy.deepcopy(prior_state.get("fault")) - result["rejection"] = {"code": code} + result["rejection"] = {"code": code.value} return result @@ -500,12 +527,12 @@ def _validate_prior_state_values(state: dict[str, Any]) -> None: def visit(value: Any, path: tuple[str | int, ...], ancestors: set[int]) -> None: if _is_prior_counter_path(path): if not _logical_counter(value): - raise ValidationError("numeric_value_out_of_range") + raise ValidationError(LoadCode.NUMERIC_VALUE_OUT_OF_RANGE) return if isinstance(value, list): identity = id(value) if identity in ancestors: - raise ValidationError("non_json_value") + raise ValidationError(LoadCode.NON_JSON_VALUE) ancestors.add(identity) for index, item in enumerate(value): visit(item, (*path, index), ancestors) @@ -514,11 +541,11 @@ def visit(value: Any, path: tuple[str | int, ...], ancestors: set[int]) -> None: if isinstance(value, dict): identity = id(value) if identity in ancestors: - raise ValidationError("non_json_value") + raise ValidationError(LoadCode.NON_JSON_VALUE) ancestors.add(identity) for key, item in value.items(): if not isinstance(key, str): - raise ValidationError("non_string_map_key") + raise ValidationError(LoadCode.NON_STRING_MAP_KEY) visit(item, (*path, key), ancestors) ancestors.remove(identity) return @@ -1068,16 +1095,16 @@ def _valid_fault( next_logical_step_sequence: int, ) -> bool: pointer_codes = { - "guard_fault", - "action_fault", - "invalid_instance_target", - "inactive_component_target", - "binding_not_empty", + FaultCode.GUARD_FAULT.value, + FaultCode.ACTION_FAULT.value, + FaultCode.INVALID_INSTANCE_TARGET.value, + FaultCode.INACTIVE_COMPONENT_TARGET.value, + FaultCode.BINDING_NOT_EMPTY.value, } system_locators = { - "contained_runtime_fault": "system:unhandled_contained_failure", - "cascade_fault": "system:cascade_cleanup", - "invariant_fault": "system:invariant", + FaultCode.CONTAINED_RUNTIME_FAULT.value: "system:unhandled_contained_failure", + FaultCode.CASCADE_FAULT.value: "system:cascade_cleanup", + FaultCode.INVARIANT_FAULT.value: "system:invariant", } code = fault.get("code") if isinstance(fault, dict) else None locator = fault.get("source_locator") if isinstance(fault, dict) else None @@ -1171,15 +1198,15 @@ def _validate_envelope( state: dict[str, Any], envelope: Any, mode: Literal["input", "internal"], -) -> str | None: +) -> DispatchCode | None: del models if state["status"] == "faulted": - return "invalid_instance_target" + return DispatchCode.INVALID_INSTANCE_TARGET if not isinstance(envelope, dict): - return "invalid_event" + return DispatchCode.INVALID_EVENT allowed_members = {"event", "event_id", "target", "payload", "correlation_id"} if set(envelope) - allowed_members: - return "invalid_event" + return DispatchCode.INVALID_EVENT event = envelope.get("event") event_id = envelope.get("event_id") if ( @@ -1189,7 +1216,7 @@ def _validate_envelope( or not event_id or not validate_unicode([event, event_id]) ): - return "invalid_event" + return DispatchCode.INVALID_EVENT target = envelope.get("target") target_code, runtime = _locate_target(state, target) if target_code is not None: @@ -1198,17 +1225,17 @@ def _validate_envelope( try: validate_portable_values(target) except ValidationError: - return "invalid_instance_target" + return DispatchCode.INVALID_INSTANCE_TARGET if not validate_unicode(target): - return "invalid_instance_target" + return DispatchCode.INVALID_INSTANCE_TARGET if runtime["status"] != "running": return ( - "inactive_component_target" + DispatchCode.INACTIVE_COMPONENT_TARGET if runtime["role"] == "component" - else "invalid_instance_target" + else DispatchCode.INVALID_INSTANCE_TARGET ) if mode == "input" and runtime["role"] == "component": - return "invalid_instance_target" + return DispatchCode.INVALID_INSTANCE_TARGET machine = next( item for item in bundle.raw["machines"] if item["machine_id"] == runtime["machine_id"] ) @@ -1219,21 +1246,21 @@ def _validate_envelope( (mode == "input" and runtime["role"] in {"root", "spawned"}) or (mode == "internal" and runtime["role"] == "component") ): - return "invalid_event" + return DispatchCode.INVALID_EVENT if "correlation_id" in envelope: - return "invalid_correlation" + return DispatchCode.INVALID_CORRELATION payload = envelope.get("payload") if not isinstance(payload, dict) or set(payload) != {"changed"}: - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD changed = payload["changed"] if not isinstance(changed, dict) or not changed: - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD try: validate_portable_values(changed) except ValidationError: - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD if not validate_unicode(changed): - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD runtime_root = _pointer_get(bundle.raw, runtime["root_pointer"]) variables = runtime_root.get("variables") or {} external = { @@ -1242,32 +1269,32 @@ def _validate_envelope( if declaration.get("external") is True } if set(changed) - set(external): - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD try: for name, value in changed.items(): _normalize_value(value, str(external[name]["type"])) except ValueError: - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD return None declaration = declarations.get(event) if declaration is None: if mode == "internal" and event in _reserved_events(): return _validate_reserved_payload(event, envelope) - return "invalid_event" + return DispatchCode.INVALID_EVENT expected_direction = "input" if mode == "input" else "internal" if declaration["direction"] != expected_direction: - return "invalid_event" + return DispatchCode.INVALID_EVENT correlation = envelope.get("correlation_id") if correlation is not None and ( not isinstance(correlation, str) or not correlation or not validate_unicode(correlation) ): - return "invalid_correlation" + return DispatchCode.INVALID_CORRELATION if declaration.get("correlates_to") and correlation is None: - return "invalid_correlation" + return DispatchCode.INVALID_CORRELATION if "payload" not in envelope: - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD if _normalize_payload(declaration, envelope.get("payload")) is None: - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD return None @@ -1280,16 +1307,18 @@ def _reserved_events() -> set[str]: } -def _validate_reserved_payload(event: str, envelope: dict[str, Any]) -> str | None: +def _validate_reserved_payload( + event: str, envelope: dict[str, Any] +) -> DispatchCode | None: payload = envelope.get("payload") if not isinstance(payload, dict): - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD try: validate_portable_values(payload) except ValidationError: - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD if not validate_unicode(payload): - return "invalid_payload" + return DispatchCode.INVALID_PAYLOAD if event == "determa.component_completed": valid = set(payload) == {"component_id", "component_runtime_id"} and all( isinstance(payload[name], str) and bool(payload[name]) @@ -1349,7 +1378,7 @@ def _validate_reserved_payload(event: str, envelope: dict[str, Any]) -> str | No ) else: valid = False - return None if valid else "invalid_payload" + return None if valid else DispatchCode.INVALID_PAYLOAD def _valid_public_fault(value: Any) -> bool: @@ -1372,9 +1401,11 @@ def _valid_public_fault(value: Any) -> bool: ) -def _locate_target(state: dict[str, Any], target: Any) -> tuple[str | None, dict[str, Any] | None]: +def _locate_target( + state: dict[str, Any], target: Any +) -> tuple[DispatchCode | None, dict[str, Any] | None]: if not isinstance(target, dict) or len(target) != 1: - return "invalid_instance_target", None + return DispatchCode.INVALID_INSTANCE_TARGET, None runtimes = state["runtimes"] if "root" in target: value = target["root"] @@ -1384,33 +1415,35 @@ def _locate_target(state: dict[str, Any], target: Any) -> tuple[str | None, dict or value.get("root_instance_id") != state["root_instance_id"] or value.get("root_runtime_id") != state["root_runtime_id"] ): - return "invalid_instance_target", None + return DispatchCode.INVALID_INSTANCE_TARGET, None runtime = runtimes[state["root_runtime_id"]] return _target_eligibility(state, runtime), runtime if "spawned_instance" in target: reference = target["spawned_instance"] if not _is_instance_reference(reference): - return "invalid_instance_target", None + return DispatchCode.INVALID_INSTANCE_TARGET, None runtime = runtimes.get(reference["instance_id"]) if runtime is None or runtime.get("instance_reference") != reference: - return "invalid_instance_target", None + return DispatchCode.INVALID_INSTANCE_TARGET, None return _target_eligibility(state, runtime), runtime if "component" in target: value = target["component"] if not isinstance(value, dict): - return "inactive_component_target", None + return DispatchCode.INACTIVE_COMPONENT_TARGET, None runtime = runtimes.get(value.get("component_runtime_id")) if runtime is None or runtime.get("target") != target: - return "inactive_component_target", None + return DispatchCode.INACTIVE_COMPONENT_TARGET, None return _target_eligibility(state, runtime), runtime - return "invalid_instance_target", None + return DispatchCode.INVALID_INSTANCE_TARGET, None -def _target_eligibility(state: dict[str, Any], runtime: dict[str, Any]) -> str | None: +def _target_eligibility( + state: dict[str, Any], runtime: dict[str, Any] +) -> DispatchCode | None: code = ( - "inactive_component_target" + DispatchCode.INACTIVE_COMPONENT_TARGET if runtime["role"] == "component" - else "invalid_instance_target" + else DispatchCode.INVALID_INSTANCE_TARGET ) if runtime["status"] != "running": return code @@ -1486,7 +1519,7 @@ def new_runtime( **copy.deepcopy(metadata), } if not replace and runtime_id in self.state["runtimes"]: - raise StepFault("invariant_fault", "system:invariant") + raise StepFault(FaultCode.INVARIANT_FAULT, "system:invariant") self.state["runtimes"][runtime_id] = runtime return runtime @@ -1507,7 +1540,12 @@ def model_for(self, runtime: dict[str, Any]) -> MachineModel: def runtime_for_target(self, target: dict[str, Any]) -> dict[str, Any]: code, runtime = _locate_target(self.state, target) if code is not None or runtime is None: - raise StepFault(code or "invalid_instance_target", "system:invariant") + fault_code = ( + FaultCode(code.value) + if code is not None + else FaultCode.INVALID_INSTANCE_TARGET + ) + raise StepFault(fault_code, "system:invariant") return runtime def event_declaration(self, runtime: dict[str, Any], event_name: str) -> dict[str, Any] | None: @@ -1624,7 +1662,7 @@ def initialize_variables( if selected is None and "init" in declaration: selected = declaration["init"] if selected is None and declaration["type"] != "instance_reference": - raise StepFault("invariant_fault", "system:invariant") + raise StepFault(FaultCode.INVARIANT_FAULT, "system:invariant") values[name] = copy.deepcopy(selected) return values @@ -1645,7 +1683,7 @@ def variable_slot( if name in declarations and current.path in runtime["scopes"]: return current.path, declarations[name] current = current.parent - raise StepFault("invariant_fault", "system:invariant") + raise StepFault(FaultCode.INVARIANT_FAULT, "system:invariant") def activation( self, @@ -1674,7 +1712,8 @@ def evaluate( try: return cel.evaluate(expression, activation) except CelError as exc: - raise StepFault("guard_fault" if guard else "action_fault", pointer) from exc + code = FaultCode.GUARD_FAULT if guard else FaultCode.ACTION_FAULT + raise StepFault(code, pointer) from exc def allocate_components( self, runtime: dict[str, Any], machine: MachineModel, state: StateNode @@ -1760,7 +1799,7 @@ def evaluate_author_bindings( result[kind][name] = _normalize_value(value, str(declaration["type"])) except ValueError as exc: raise StepFault( - "action_fault", f"{pointer}/with/{kind}/{_escape_pointer(name)}" + FaultCode.ACTION_FAULT, f"{pointer}/with/{kind}/{_escape_pointer(name)}" ) from exc for name, declaration in declarations.items(): if declaration.get(kind) and name not in result[kind]: @@ -1806,7 +1845,10 @@ def process(self, runtime: dict[str, Any], envelope: dict[str, Any]) -> bool: "determa.component_failed", "determa.spawned_instance_failed", }: - raise StepFault("contained_runtime_fault", "system:unhandled_contained_failure") + raise StepFault( + FaultCode.CONTAINED_RUNTIME_FAULT, + "system:unhandled_contained_failure", + ) return False source, transition, pointer = selected try: @@ -1859,7 +1901,7 @@ def resolve_compound_transition( seen: set[str] = set() while target.is_choice: if target.path in seen: - raise StepFault("invariant_fault", "system:invariant") + raise StepFault(FaultCode.INVARIANT_FAULT, "system:invariant") seen.add(target.path) branch_selected = None for index, branch in enumerate(target.raw["choice"]): @@ -1878,7 +1920,7 @@ def resolve_compound_transition( branch_selected = (branch, branch_pointer) break if branch_selected is None: - raise StepFault("invariant_fault", "system:invariant") + raise StepFault(FaultCode.INVARIANT_FAULT, "system:invariant") branch, branch_pointer = branch_selected self.run_actions( runtime, @@ -1921,7 +1963,7 @@ def run_actions( ) except ValueError as exc: raise StepFault( - "action_fault", + FaultCode.ACTION_FAULT, f"{action_pointer}/assign/{_escape_pointer(name)}", ) from exc elif "send" in action: @@ -1997,7 +2039,7 @@ def send( evaluated_targets.append((target_spec, value)) if send["event"] == "env": if not isinstance(payload_values["changed"], dict): - raise StepFault("action_fault", f"{pointer}/payload/changed") + raise StepFault(FaultCode.ACTION_FAULT, f"{pointer}/payload/changed") normalized_payload = {"changed": copy.deepcopy(payload_values["changed"])} else: assert declaration is not None @@ -2009,7 +2051,7 @@ def send( if supplied else f"{pointer}/payload" ) - raise StepFault("action_fault", locator) + raise StepFault(FaultCode.ACTION_FAULT, locator) normalized_payload = payload_result resolved = [ self.resolve_send_target(runtime, target_spec, value, pointer, index, "targets" in send) @@ -2071,24 +2113,24 @@ def resolve_send_target( if target_spec.get("owner") is True: owner_id = runtime.get("owner_runtime_id") if owner_id is None or owner_id not in self.state["runtimes"]: - raise StepFault("invalid_instance_target", f"{pointer}{suffix}") + raise StepFault(FaultCode.INVALID_INSTANCE_TARGET, f"{pointer}{suffix}") return self.target_for(self.state["runtimes"][owner_id]) if "component" in target_spec: child_id = runtime["components"].get(target_spec["component"]) child = self.state["runtimes"].get(child_id) if child is None or _target_eligibility(self.state, child) is not None: - raise StepFault("inactive_component_target", f"{pointer}{suffix}") + raise StepFault(FaultCode.INACTIVE_COMPONENT_TARGET, f"{pointer}{suffix}") return cast(dict[str, Any], copy.deepcopy(child["target"])) if "instance" in target_spec: if not _is_instance_reference(evaluated): - raise StepFault("invalid_instance_target", f"{pointer}{suffix}/instance") + raise StepFault(FaultCode.INVALID_INSTANCE_TARGET, f"{pointer}{suffix}/instance") child = self.state["runtimes"].get(evaluated["instance_id"]) if child is None or _target_eligibility(self.state, child) is not None: - raise StepFault("invalid_instance_target", f"{pointer}{suffix}/instance") + raise StepFault(FaultCode.INVALID_INSTANCE_TARGET, f"{pointer}{suffix}/instance") return {"spawned_instance": copy.deepcopy(evaluated)} if target_spec.get("external") is True: return "external" - raise StepFault("invalid_instance_target", f"{pointer}{suffix}") + raise StepFault(FaultCode.INVALID_INSTANCE_TARGET, f"{pointer}{suffix}") def target_for(self, runtime: dict[str, Any]) -> dict[str, Any]: if runtime["role"] == "root": @@ -2115,7 +2157,7 @@ def refresh( selected = refresh.get("only", list(changed)) for index, name in enumerate(selected): if name not in changed: - raise StepFault("action_fault", f"{pointer}/refresh/only/{index}") + raise StepFault(FaultCode.ACTION_FAULT, f"{pointer}/refresh/only/{index}") for name in selected: scope_path, declaration = self.variable_slot(runtime, state, name) runtime["scopes"][scope_path][name] = _normalize_value( @@ -2161,7 +2203,7 @@ def spawn( name = spawn["bind_to"] scope_path, declaration = self.variable_slot(runtime, state, name) if runtime["scopes"][scope_path][name] is not None: - raise StepFault("binding_not_empty", f"{pointer}/bind_to") + raise StepFault(FaultCode.BINDING_NOT_EMPTY, f"{pointer}/bind_to") runtime["scopes"][scope_path][name] = copy.deepcopy(reference) holder_state = machine.states[scope_path] holder = { @@ -2398,7 +2440,7 @@ def cleanup_descendant( try: self.cleanup_runtime(runtime, dispose=True, frozen=frozen) except StepFault as exc: - raise StepFault("cascade_fault", "system:cascade_cleanup") from exc + raise StepFault(FaultCode.CASCADE_FAULT, "system:cascade_cleanup") from exc def cleanup_runtime( self, diff --git a/src/determa/state/errors.py b/src/determa/state/errors.py index a0d0ebd..361df0f 100644 --- a/src/determa/state/errors.py +++ b/src/determa/state/errors.py @@ -4,6 +4,8 @@ from dataclasses import dataclass +from .codes import EngineFaultCode + class DetermaError(Exception): """Base class for Determa State errors.""" @@ -22,10 +24,13 @@ class ValidationError(DetermaError): """A source, schema, or semantic validation failure.""" def __init__(self, code: str, path: str = "", message: str = "") -> None: - self.code = code + normalized_code = str(code) + self.code = normalized_code self.path = path - self.message = message or code - self.errors = [ErrorRecord(code=code, path=path, message=self.message)] + self.message = message or normalized_code + self.errors = [ + ErrorRecord(code=normalized_code, path=path, message=self.message) + ] super().__init__(self.message) @@ -41,16 +46,17 @@ class ArtifactError(DetermaError): """A portable persistence artifact is invalid or unsupported.""" def __init__(self, code: str, path: str = "", message: str = "") -> None: - self.code = code + normalized_code = str(code) + self.code = normalized_code self.path = path - self.message = message or code + self.message = message or normalized_code super().__init__(self.message) class StepFault(DetermaError): """Internal control flow for one atomic RTC fault.""" - def __init__(self, code: str, source_locator: str) -> None: - self.code = code + def __init__(self, code: EngineFaultCode, source_locator: str) -> None: + self.code = code.value self.source_locator = source_locator super().__init__(f"{code} at {source_locator}") diff --git a/src/determa/state/host.py b/src/determa/state/host.py index 80f6822..3110661 100644 --- a/src/determa/state/host.py +++ b/src/determa/state/host.py @@ -15,6 +15,18 @@ serialize_execution_checkpoint, validate_execution_checkpoint_member, ) +from .codes import ( + CheckpointHostFailureCode as HostCode, +) +from .codes import ( + CheckpointPreAcceptanceFailureCode as PreAcceptanceCode, +) +from .codes import ( + ExecutionStoreAdapterFailureCode as AdapterCode, +) +from .codes import ( + PersistenceFailureCode as PersistenceCode, +) from .definition import Bundle, BundleSource, load_bundle from .engine import create as core_create from .engine import dispatch as core_dispatch @@ -48,8 +60,8 @@ class ExecutionHostError(DetermaError): """A closed host-layer failure.""" def __init__(self, code: str, message: str = "") -> None: - self.code = code - self.message = message or code + self.code = str(code) + self.message = message or self.code super().__init__(self.message) @@ -67,17 +79,17 @@ def _checkpoint_number(value: Any) -> int: ) ) ): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) try: return int(value) except ValueError as exc: - raise ExecutionHostError("invalid_execution_checkpoint") from exc + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) from exc def _increment_checkpoint_number(value: Any) -> str: result = str(_checkpoint_number(value) + 1) if len(result) > _MAX_DECIMAL_DIGITS: - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) return result @@ -147,7 +159,7 @@ def portable_envelope( if correlation_id is not None: result["correlation_id"] = correlation_id if not validate_execution_checkpoint_member("envelope", result): - raise ExecutionHostError("malformed_delivery") + raise ExecutionHostError(PreAcceptanceCode.MALFORMED_DELIVERY) return result @@ -248,7 +260,7 @@ def validate_host_profile( and "native_shared_application_transaction" in host_features ) if not valid: - raise ExecutionHostError("adapter_capability_mismatch") + raise ExecutionHostError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) @dataclass(frozen=True) @@ -269,7 +281,7 @@ def _project_fault( candidate = runtime["fault"] if candidate is not None and candidate["runtime_id"] == fault["runtime_id"]: return cast(dict[str, Any], copy.deepcopy(candidate)) - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) def _project_emission(emission: Mapping[str, Any]) -> dict[str, Any]: @@ -466,7 +478,7 @@ def __init__( fault_injector: FaultInjector | None = None, ) -> None: if not required_capabilities.issubset(store.capabilities): - raise ExecutionHostError("adapter_capability_mismatch") + raise ExecutionHostError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) if required_capabilities or profile is not None: store.validate_schema() if profile is not None: @@ -525,12 +537,12 @@ def _restore( PERMANENT_RECEIPT_RETENTION in self.store.capabilities and restored.document["replay_retention"]["mode"] != "permanent" ): - raise ExecutionHostError("adapter_capability_mismatch") + raise ExecutionHostError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) if ( PERMANENT_OUTBOX_TERMINAL_RETENTION in self.store.capabilities and restored.document["outbox_effect_tombstones"] ): - raise ExecutionHostError("adapter_capability_mismatch") + raise ExecutionHostError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) return restored def _transaction( @@ -555,7 +567,7 @@ def run_shared_transaction( ) -> dict[str, Any]: """Commit application writes and exactly one staged host operation together.""" if SHARED_APPLICATION_TRANSACTION not in self.store.capabilities: - raise ExecutionHostError("adapter_capability_mismatch") + raise ExecutionHostError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) with self.store.shared_transaction(root_instance_id) as ( native_transaction, store_transaction, @@ -584,7 +596,7 @@ def _check_expected( or checkpoint["execution_checkpoint_digest"] != expected_checkpoint_digest ): - raise ExecutionHostError("checkpoint_revision_conflict") + raise ExecutionHostError(HostCode.CHECKPOINT_REVISION_CONFLICT) def _stage_insert( self, transaction: ExecutionStoreTransaction, candidate: dict[str, Any] @@ -592,7 +604,7 @@ def _stage_insert( restore_execution_checkpoint(candidate, self.artifact_resolver) self._fault("before_commit") if not transaction.insert(serialize_execution_checkpoint(candidate)): - raise ExecutionHostError("checkpoint_revision_conflict") + raise ExecutionHostError(HostCode.CHECKPOINT_REVISION_CONFLICT) def _stage_replace( self, @@ -607,7 +619,7 @@ def _stage_replace( previous["execution_checkpoint_digest"], serialize_execution_checkpoint(candidate), ): - raise ExecutionHostError("checkpoint_revision_conflict") + raise ExecutionHostError(HostCode.CHECKPOINT_REVISION_CONFLICT) def read_checkpoint( self, @@ -651,7 +663,7 @@ def create( and receipt["request_digest"] == request_digest ): return {"result": "committed", "receipt": copy.deepcopy(receipt)} - raise ExecutionHostError("creation_id_conflict") + raise ExecutionHostError(HostCode.CREATION_ID_CONFLICT) result = core_create( validated, machine_id, @@ -662,7 +674,7 @@ def create( projected = _project_core_result(validated, result) aggregate = projected["aggregate_state"] if aggregate is None: - raise ExecutionHostError("creation_rejected") + raise ExecutionHostError(HostCode.CREATION_REJECTED) candidate = _new_checkpoint(aggregate, request_digest, projected) self._stage_insert(transaction, candidate) receipt = copy.deepcopy(candidate["operation_receipts"][0]) @@ -706,8 +718,10 @@ def _delivery_candidate( supplied_digest, ) - def _not_accepted(self, code: str) -> dict[str, Any]: - return {"result": "not_accepted", "failure": {"code": code}} + def _not_accepted( + self, code: PreAcceptanceCode + ) -> dict[str, Any]: + return {"result": "not_accepted", "failure": {"code": code.value}} def _delivery_replay( self, @@ -718,7 +732,7 @@ def _delivery_replay( for pending in checkpoint["pending_deliveries"]: if pending["envelope"]["event_id"] == event_id: if pending["envelope_digest"] != digest: - return self._not_accepted("event_id_conflict") + return self._not_accepted(PreAcceptanceCode.EVENT_ID_CONFLICT) return { "result": "pending", "event_id": event_id, @@ -728,7 +742,7 @@ def _delivery_replay( for receipt in checkpoint["operation_receipts"]: if receipt["operation_kind"] == "delivery" and receipt["event_id"] == event_id: if receipt["request_digest"] != digest: - return self._not_accepted("event_id_conflict") + return self._not_accepted(PreAcceptanceCode.EVENT_ID_CONFLICT) return {"result": "committed", "receipt": copy.deepcopy(receipt)} return None @@ -740,9 +754,9 @@ def _prepare_acceptance( parsed = self._delivery_candidate(candidate) root_instance_id, mode, origin, envelope, supplied_digest = parsed if root_instance_id is None or mode is None or envelope is None: - return None, self._not_accepted("malformed_delivery") + return None, self._not_accepted(PreAcceptanceCode.MALFORMED_DELIVERY) if root_instance_id != checkpoint["root_instance_id"]: - return None, self._not_accepted("wrong_root") + return None, self._not_accepted(PreAcceptanceCode.WRONG_ROOT) digest = delivery_request_digest(root_instance_id, mode, envelope) replay = self._delivery_replay( @@ -751,7 +765,7 @@ def _prepare_acceptance( if replay is not None: return None, replay if checkpoint["root_record"]["status"] == "tombstone": - return None, self._not_accepted("tombstoned_root") + return None, self._not_accepted(PreAcceptanceCode.TOMBSTONED_ROOT) valid_mode = mode in {"input", "internal"} valid_origin = validate_execution_checkpoint_member("deliveryOrigin", origin) @@ -763,16 +777,18 @@ def _prepare_acceptance( and origin.get("kind") == "internal_emission" ) if not valid_mode: - return None, self._not_accepted("invalid_delivery_mode") + return None, self._not_accepted(PreAcceptanceCode.INVALID_DELIVERY_MODE) if not valid_origin or not valid_pair: - return None, self._not_accepted("invalid_delivery_origin") + return None, self._not_accepted(PreAcceptanceCode.INVALID_DELIVERY_ORIGIN) if supplied_digest is not None and supplied_digest != digest: - return None, self._not_accepted("delivery_digest_mismatch") + return None, self._not_accepted( + PreAcceptanceCode.DELIVERY_DIGEST_MISMATCH + ) if ( _target_root_instance_id(envelope["target"]) != checkpoint["root_instance_id"] ): - return None, self._not_accepted("wrong_root") + return None, self._not_accepted(PreAcceptanceCode.WRONG_ROOT) return { "root_instance_id": root_instance_id, "delivery_mode": mode, @@ -792,7 +808,7 @@ def accept_delivery( with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - return self._not_accepted("wrong_root") + return self._not_accepted(PreAcceptanceCode.WRONG_ROOT) checkpoint = self._restore(source, root_instance_id).document prepared, result = self._prepare_acceptance(checkpoint, candidate) if result is not None: @@ -864,7 +880,7 @@ def _commit_delivery( ) aggregate = copy.deepcopy(projected["aggregate_state"]) if aggregate is None or restored.aggregate is None: - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) candidate["root_record"]["aggregate_state"] = aggregate receipt = { "operation_kind": "delivery", @@ -902,7 +918,7 @@ def process_pending_delivery( with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) restored = self._restore(source, root_instance_id) checkpoint = restored.document parsed = self._delivery_candidate(candidate) @@ -913,10 +929,10 @@ def process_pending_delivery( or origin is None or envelope is None ): - raise ExecutionHostError("malformed_delivery") + raise ExecutionHostError(PreAcceptanceCode.MALFORMED_DELIVERY) digest = delivery_request_digest(root_instance_id, mode, envelope) if supplied_digest is not None and supplied_digest != digest: - raise ExecutionHostError("delivery_digest_mismatch") + raise ExecutionHostError(PreAcceptanceCode.DELIVERY_DIGEST_MISMATCH) replay = self._delivery_replay( checkpoint, envelope["event_id"], digest ) @@ -933,12 +949,12 @@ def process_pending_delivery( None, ) if pending is None or pending["envelope_digest"] != digest: - raise ExecutionHostError("event_id_conflict") + raise ExecutionHostError(HostCode.EVENT_ID_CONFLICT) self._check_expected( checkpoint, expected_revision, expected_checkpoint_digest ) if restored.aggregate is None: - raise ExecutionHostError("tombstoned_root") + raise ExecutionHostError(PreAcceptanceCode.TOMBSTONED_ROOT) result = core_dispatch( restored.aggregate.bundle, restored.aggregate.state, @@ -971,7 +987,7 @@ def foreground_process_delivery( with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) restored = self._restore(source, root_instance_id) checkpoint = restored.document prepared, replay = self._prepare_acceptance(checkpoint, candidate) @@ -982,7 +998,7 @@ def foreground_process_delivery( checkpoint, expected_revision, expected_checkpoint_digest ) if restored.aggregate is None: - raise ExecutionHostError("tombstoned_root") + raise ExecutionHostError(PreAcceptanceCode.TOMBSTONED_ROOT) result = core_dispatch( restored.aggregate.bundle, restored.aggregate.state, @@ -1022,7 +1038,7 @@ def maintenance_migration( "sha256", source_aggregate_state_digest ) ): - raise ExecutionHostError("invalid_migration_request") + raise ExecutionHostError(PersistenceCode.INVALID_MIGRATION_REQUEST) request_digest = maintenance_migration_request_digest( root_instance_id, operation_id, @@ -1034,7 +1050,7 @@ def maintenance_migration( with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) restored = self._restore(source, root_instance_id) checkpoint = restored.document for receipt in checkpoint["operation_receipts"]: @@ -1047,14 +1063,14 @@ def maintenance_migration( "result": "committed", "receipt": copy.deepcopy(receipt), } - raise ExecutionHostError("operation_id_conflict") + raise ExecutionHostError(HostCode.OPERATION_ID_CONFLICT) if restored.aggregate is None: - raise ExecutionHostError("tombstoned_root") + raise ExecutionHostError(PreAcceptanceCode.TOMBSTONED_ROOT) current_source_digest = restored.aggregate.aggregate_envelope[ "aggregate_state_digest" ] if source_aggregate_state_digest != current_source_digest: - raise ExecutionHostError("invalid_migration_request") + raise ExecutionHostError(PersistenceCode.INVALID_MIGRATION_REQUEST) self._check_expected( checkpoint, expected_revision, expected_checkpoint_digest ) @@ -1122,11 +1138,11 @@ def update_pending_outbox( ) -> dict[str, Any]: desired = copy.deepcopy(dict(desired_pending_state)) if not validate_execution_checkpoint_member("pendingOutboxState", desired): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) checkpoint = self._restore(source, root_instance_id).document item = next( ( @@ -1137,7 +1153,7 @@ def update_pending_outbox( None, ) if item is None: - raise ExecutionHostError("effect_id_conflict") + raise ExecutionHostError(HostCode.EFFECT_ID_CONFLICT) if item["delivery_state"] == desired: return {"result": "committed", "record": copy.deepcopy(item)} self._check_expected( @@ -1168,11 +1184,11 @@ def terminalize_outbox( ) -> dict[str, Any]: outcome = copy.deepcopy(dict(terminal_outcome)) if not validate_execution_checkpoint_member("terminalOutboxOutcome", outcome): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) checkpoint = self._restore(source, root_instance_id).document for record in checkpoint["terminal_outbox_records"]: if record["intent"]["effect_id"] == effect_id: @@ -1181,7 +1197,7 @@ def terminalize_outbox( "result": "committed", "record": copy.deepcopy(record), } - raise ExecutionHostError("effect_id_conflict") + raise ExecutionHostError(HostCode.EFFECT_ID_CONFLICT) for record in checkpoint["outbox_effect_tombstones"]: if record["effect_id"] == effect_id: if record["outcome"] == outcome: @@ -1189,7 +1205,7 @@ def terminalize_outbox( "result": "committed", "record": copy.deepcopy(record), } - raise ExecutionHostError("effect_id_conflict") + raise ExecutionHostError(HostCode.EFFECT_ID_CONFLICT) pending = next( ( value @@ -1199,7 +1215,7 @@ def terminalize_outbox( None, ) if pending is None: - raise ExecutionHostError("effect_id_conflict") + raise ExecutionHostError(HostCode.EFFECT_ID_CONFLICT) self._check_expected( checkpoint, expected_revision, expected_checkpoint_digest ) @@ -1236,11 +1252,11 @@ def compact_outbox( expected_checkpoint_digest: str, ) -> dict[str, Any]: if PERMANENT_OUTBOX_TERMINAL_RETENTION in self.store.capabilities: - raise ExecutionHostError("adapter_capability_mismatch") + raise ExecutionHostError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) checkpoint = self._restore(source, root_instance_id).document existing = next( ( @@ -1261,7 +1277,7 @@ def compact_outbox( None, ) if terminal is None: - raise ExecutionHostError("effect_id_conflict") + raise ExecutionHostError(HostCode.EFFECT_ID_CONFLICT) self._check_expected( checkpoint, expected_revision, expected_checkpoint_digest ) @@ -1303,11 +1319,11 @@ def delete_outbox_record( PERMANENT_OUTBOX_TERMINAL_RETENTION, COMPACT_EFFECT_IDENTITY_RETENTION, }.intersection(self.store.capabilities): - raise ExecutionHostError("adapter_capability_mismatch") + raise ExecutionHostError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) checkpoint = self._restore(source, root_instance_id).document if any( emission.get("kind") == "external_outbox" @@ -1315,7 +1331,7 @@ def delete_outbox_record( for receipt in checkpoint["operation_receipts"] for emission in receipt.get("emission_references", []) ): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) self._check_expected( checkpoint, expected_revision, expected_checkpoint_digest ) @@ -1336,7 +1352,7 @@ def delete_outbox_record( if prior_count == len(candidate["terminal_outbox_records"]) + len( candidate["outbox_effect_tombstones"] ): - raise ExecutionHostError("effect_id_conflict") + raise ExecutionHostError(HostCode.EFFECT_ID_CONFLICT) candidate = seal_execution_checkpoint(candidate) self._stage_replace(transaction, checkpoint, candidate) self._after_commit() @@ -1352,22 +1368,22 @@ def update_replay_retention( ) -> dict[str, Any]: target = copy.deepcopy(dict(target_replay_retention)) if not validate_execution_checkpoint_member("replayRetention", target): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) if ( PERMANENT_RECEIPT_RETENTION in self.store.capabilities and target["mode"] != "permanent" ): - raise ExecutionHostError("adapter_capability_mismatch") + raise ExecutionHostError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) checkpoint = self._restore(source, root_instance_id).document current = checkpoint["replay_retention"] if current == target: return {"result": "committed", "replay_retention": copy.deepcopy(current)} if current["mode"] == "bounded" and target["mode"] == "permanent": - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) current_cutoff = current["pruned_through_receipt_sequence"] target_cutoff = target["pruned_through_receipt_sequence"] if target["mode"] == "bounded": @@ -1376,7 +1392,7 @@ def update_replay_retention( and current["policy_identifier"] != target["policy_identifier"] ): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) if ( current_cutoff is not None and ( @@ -1385,7 +1401,7 @@ def update_replay_retention( < _checkpoint_number(current_cutoff) ) ): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) if ( target_cutoff is not None and _checkpoint_number(target_cutoff) @@ -1393,7 +1409,7 @@ def update_replay_retention( checkpoint["next_operation_receipt_sequence"] ) ): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) self._check_expected( checkpoint, expected_revision, expected_checkpoint_digest ) @@ -1438,8 +1454,8 @@ def update_replay_retention( try: self._stage_replace(transaction, checkpoint, candidate) except Exception as exc: - if getattr(exc, "code", None) == "invalid_execution_checkpoint": - raise ExecutionHostError("invalid_execution_checkpoint") from exc + if getattr(exc, "code", None) == HostCode.INVALID_EXECUTION_CHECKPOINT: + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) from exc raise response = { "result": "committed", @@ -1457,11 +1473,11 @@ def tombstone_root( expected_checkpoint_digest: str, ) -> dict[str, Any]: if not operation_id: - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) with self._transaction(root_instance_id) as transaction: source = transaction.load() if source is None: - raise ExecutionHostError("wrong_root") + raise ExecutionHostError(PreAcceptanceCode.WRONG_ROOT) restored = self._restore(source, root_instance_id) checkpoint = restored.document root_record = checkpoint["root_record"] @@ -1471,12 +1487,12 @@ def tombstone_root( "result": "tombstoned", "tombstone": copy.deepcopy(root_record), } - raise ExecutionHostError("operation_id_conflict") + raise ExecutionHostError(HostCode.OPERATION_ID_CONFLICT) self._check_expected( checkpoint, expected_revision, expected_checkpoint_digest ) if restored.aggregate is None: - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) root_runtime = restored.aggregate.state["runtimes"][ restored.aggregate.state["root_runtime_id"] ] @@ -1485,7 +1501,7 @@ def tombstone_root( or checkpoint["pending_deliveries"] or checkpoint["pending_outbox_intents"] ): - raise ExecutionHostError("invalid_execution_checkpoint") + raise ExecutionHostError(HostCode.INVALID_EXECUTION_CHECKPOINT) aggregate = root_record["aggregate_state"] candidate = _mutate(checkpoint) tombstone = { @@ -1518,7 +1534,7 @@ def delete_checkpoint( del root_instance_id, expected_revision, expected_checkpoint_digest return { "result": "unsupported", - "failure": {"code": "physical_deletion_unsupported"}, + "failure": {"code": HostCode.PHYSICAL_DELETION_UNSUPPORTED.value}, } diff --git a/src/determa/state/migration.py b/src/determa/state/migration.py index ef1fca9..cad9f67 100644 --- a/src/determa/state/migration.py +++ b/src/determa/state/migration.py @@ -8,6 +8,7 @@ from typing import Any, cast from . import cel +from .codes import PersistenceFailureCode as PersistenceCode from .definition import Bundle, _escape_pointer from .engine import Delivery, dispatch from .errors import ArtifactError, CelError @@ -54,11 +55,11 @@ class MigrationLimits: def from_mapping(cls, value: dict[str, Any]) -> MigrationLimits: expected = set(cls.__dataclass_fields__) if set(value) != expected: - raise ArtifactError("invalid_migration_request") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_REQUEST) try: parsed = {name: decimal(value[name]) for name in expected} except ArtifactError as exc: - raise ArtifactError("invalid_migration_request") from exc + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_REQUEST) from exc return cls(**parsed) @@ -101,7 +102,7 @@ def succeeded(self) -> bool: def _failure(code: str) -> MigrationResult: - return MigrationResult(None, None, (), MigrationFailure(code)) + return MigrationResult(None, None, (), MigrationFailure(str(code))) def _dispatch_failure(code: str) -> MigrationDispatchResult: @@ -151,17 +152,17 @@ def _check_shape_limits( limits: MigrationLimits, ) -> None: if len(canonical_bytes(aggregate)) > limits.maximum_aggregate_bytes: - raise ArtifactError("migration_resource_limit_exceeded") + raise ArtifactError(PersistenceCode.MIGRATION_RESOURCE_LIMIT_EXCEEDED) if any( len(canonical_bytes(typed_value(bundle.raw))) > limits.maximum_definition_bytes for bundle in definitions ): - raise ArtifactError("migration_resource_limit_exceeded") + raise ArtifactError(PersistenceCode.MIGRATION_RESOURCE_LIMIT_EXCEEDED) if any( len(canonical_bytes(descriptor)) > limits.maximum_descriptor_bytes for descriptor in descriptors ): - raise ArtifactError("migration_resource_limit_exceeded") + raise ArtifactError(PersistenceCode.MIGRATION_RESOURCE_LIMIT_EXCEEDED) values: list[Any] = [aggregate, *[bundle.raw for bundle in definitions], *descriptors] metrics = [_resource_metrics(value) for value in values] if ( @@ -180,7 +181,7 @@ def _check_shape_limits( for runtime in aggregate["runtimes"] ) ): - raise ArtifactError("migration_resource_limit_exceeded") + raise ArtifactError(PersistenceCode.MIGRATION_RESOURCE_LIMIT_EXCEEDED) def _ast_nodes(value: Any) -> int: @@ -195,7 +196,7 @@ def _descriptor_static_requirements( ) -> tuple[int, int]: rule_count = sum(len(items) for items in descriptor["mappings"].values()) if rule_count > limits.maximum_descriptor_rules: - raise ArtifactError("migration_resource_limit_exceeded") + raise ArtifactError(PersistenceCode.MIGRATION_RESOURCE_LIMIT_EXCEEDED) expressions = { rule["expression"] for rule in descriptor["mappings"]["variables"] @@ -210,13 +211,13 @@ def _descriptor_static_requirements( or expression_bytes > decimal(requirements["maximum_cel_expression_length"]) or ast_nodes > decimal(requirements["maximum_cel_ast_nodes"]) ): - raise ArtifactError("migration_resource_limit_exceeded") + raise ArtifactError(PersistenceCode.MIGRATION_RESOURCE_LIMIT_EXCEEDED) return expression_bytes, ast_nodes def _pointer_parts(pointer: str) -> list[str]: if not pointer.startswith("/"): - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) return [ item.replace("~1", "/").replace("~0", "~") for item in pointer[1:].split("/") @@ -230,22 +231,22 @@ def _pointer_get(document: Any, pointer: str) -> Any: try: current = current[int(part)] except (IndexError, ValueError) as exc: - raise ArtifactError("invalid_migration_descriptor") from exc + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) from exc elif isinstance(current, dict) and part in current: current = current[part] else: - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) return current def _machine_identity_for_pointer(bundle: Bundle, root_pointer: str) -> dict[str, Any]: parts = _pointer_parts(root_pointer) if len(parts) < 3 or parts[0] != "machines": - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) try: machine = bundle.raw["machines"][int(parts[1])] except (IndexError, ValueError) as exc: - raise ArtifactError("migration_totality_failure") from exc + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) from exc return { "namespace": bundle.namespace, "machine_id": machine["machine_id"], @@ -285,14 +286,14 @@ def _state_nodes(machine: MachineModel) -> dict[str, StateNode]: def _variable_declaration(bundle: Bundle, pointer: str) -> dict[str, Any]: declaration = _pointer_get(bundle.raw, pointer) if not isinstance(declaration, dict) or "type" not in declaration: - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) return declaration def _state_pointer_for_variable(pointer: str) -> str: marker = "/variables/" if marker not in pointer: - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) return pointer.split(marker, 1)[0] @@ -313,7 +314,7 @@ def _active_ancestor_pointers(machine: MachineModel, leaves: list[str]) -> list[ for pointer in leaves: node = nodes.get(pointer) if node is None: - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) result.update(item.pointer for item in node.ancestors(include_self=True)) return sorted(result, key=lambda item: item.encode("utf-8")) @@ -327,7 +328,7 @@ def _unique_mapping( source = rule[source_member] target = rule[target_member] if source in result or target in targets: - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) result[source] = target targets.add(target) return result @@ -347,7 +348,7 @@ def _validate_descriptor_semantics( or descriptor["target_aggregate_shape_fingerprint"] != aggregate_shape_fingerprint(target_bundle) ): - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) mappings = descriptor["mappings"] if descriptor["mode"] == "compatible": if ( @@ -355,7 +356,7 @@ def _validate_descriptor_semantics( != descriptor["target_aggregate_shape_fingerprint"] or any(mappings.values()) ): - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) _unique_mapping( mappings["machines"], "source_definition_pointer", "target_definition_pointer" ) @@ -380,7 +381,7 @@ def _validate_descriptor_semantics( source = rule["source_leaf_state_definition_pointer"] targets = rule["target_leaf_state_definition_pointers"] if source in active_sources or any(target in active_targets for target in targets): - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) active_sources.add(source) active_targets.update(targets) consumed: set[str] = set() @@ -398,7 +399,7 @@ def _validate_descriptor_semantics( if any(source in consumed for source in sources) or ( target is not None and target in produced ): - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) consumed.update(sources) if target is not None: produced.add(target) @@ -416,7 +417,7 @@ def _validate_descriptor_semantics( if target_declaration["type"] == "instance_reference" or any( declaration.kind == "instance_reference" for declaration in scope.values() ): - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) try: cel.check_expression( rule["expression"], @@ -426,7 +427,7 @@ def _validate_descriptor_semantics( owner_fields=None, ) except CelError as exc: - raise ArtifactError("invalid_migration_descriptor") from exc + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) from exc _descriptor_static_requirements(descriptor, limits) @@ -435,13 +436,13 @@ def _resolve_descriptor( ) -> tuple[dict[str, Any], bytes]: source = resolver.resolve_migration_descriptor(digest) if source is None: - raise ArtifactError("migration_route_mismatch") + raise ArtifactError(PersistenceCode.MIGRATION_ROUTE_MISMATCH) if not resolver.migration_descriptor_is_trusted(digest): - raise ArtifactError("migration_descriptor_untrusted") + raise ArtifactError(PersistenceCode.MIGRATION_DESCRIPTOR_UNTRUSTED) document, _raw = load_json_artifact(source, "migration_descriptor") encoded = canonical_bytes(document) if migration_descriptor_digest(document) != digest: - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) return document, encoded @@ -473,7 +474,7 @@ def _counter_transform( operation = rule["operation"] target = rule["target_definition_pointer"] if target in targets: - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) targets.add(target) if operation == "map": pointer = rule["source_definition_pointer"] @@ -491,7 +492,7 @@ def _counter_transform( consumed.update(pointers) result.append({"definition_pointer": target, "next_sequence": str(value)}) if consumed != set(source): - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) result.sort(key=lambda item: item["definition_pointer"].encode("utf-8")) return result @@ -512,13 +513,13 @@ def _mapped_active( ] if len(matching) != 1: raise ArtifactError( - "invalid_migration_descriptor" + PersistenceCode.INVALID_MIGRATION_DESCRIPTOR if len(matching) > 1 - else "migration_totality_failure" + else PersistenceCode.MIGRATION_TOTALITY_FAILURE ) targets.extend(matching[0]["target_leaf_state_definition_pointers"]) if len(set(targets)) != len(targets): - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) targets.sort(key=lambda item: item.encode("utf-8")) source_activations = { item["state_definition_pointer"]: item["activation_sequence"] @@ -540,7 +541,7 @@ def _mapped_active( if mapped == target and source in source_activations ] if len(sources) != 1: - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) activations.append( { "state_definition_pointer": target, @@ -575,7 +576,7 @@ def _transform_variables( for occurrence in runtime["variables"]: pointer = occurrence["variable_declaration_pointer"] if pointer in source_values: - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) source_values[pointer] = occurrence target_required = _target_variable_pointers(target_machine, target_activations) produced: dict[str, dict[str, Any]] = {} @@ -601,11 +602,11 @@ def _transform_variables( if target not in target_required: continue if target in produced: - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) if operation in {"copy", "transform"} and any( value is None for value in applicable_sources ): - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) if operation == "copy": value = copy.deepcopy(cast(dict[str, Any], applicable_sources[0])["value"]) else: @@ -622,7 +623,7 @@ def _transform_variables( try: evaluated = cel.evaluate(expression, activation) except CelError as exc: - raise ArtifactError("migration_transform_fault") from exc + raise ArtifactError(PersistenceCode.MIGRATION_TRANSFORM_FAULT) from exc value = typed_value(evaluated) transformed_bytes += len(canonical_bytes(value)) if operation == "copy": @@ -633,7 +634,7 @@ def _transform_variables( "value": value, } if consumed != set(source_values) or set(produced) != set(target_required): - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) result = sorted( produced.values(), key=lambda item: ( @@ -678,7 +679,7 @@ def _transform_history( } if recorded is not None: if any(pointer not in mapping for pointer in recorded): - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) recorded = sorted( (mapping[pointer] for pointer in recorded), key=lambda item: item.encode("utf-8"), @@ -688,7 +689,7 @@ def _transform_history( "recorded_state_definition_pointers": recorded, } if consumed != set(source) or set(produced) != set(_history_pointers(target_machine)): - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) return sorted( produced.values(), key=lambda item: item["history_declaration_pointer"].encode("utf-8"), @@ -708,7 +709,7 @@ def _component_counter_transform( for item in items: source = item["definition_pointer"] if source not in mapping: - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) result.append( { "definition_pointer": mapping[source], @@ -722,21 +723,21 @@ def _component_counter_transform( def _target_root_for_component(bundle: Bundle, pointer: str) -> str: placement = _pointer_get(bundle.raw, pointer) if not isinstance(placement, dict): - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) if "root" in placement: return f"{pointer}/root" machine_id = placement.get("machine_id") for index, machine in enumerate(bundle.raw["machines"]): if machine["machine_id"] == machine_id: return f"/machines/{index}/root" - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) def _target_root_for_machine(bundle: Bundle, machine_id: str) -> str: for index, machine in enumerate(bundle.raw["machines"]): if machine["machine_id"] == machine_id: return f"/machines/{index}/root" - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) def _transform_candidate( @@ -772,12 +773,12 @@ def _transform_candidate( if relation["kind"] == "root": target_root = machine_mapping.get(source_root) if target_root is None: - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) elif relation["kind"] == "component": source_component = relation["current_component_definition_pointer"] rule = component_mapping.get(source_component) if rule is None: - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) target_component = rule["target_component_definition_pointer"] relation["current_component_definition_pointer"] = target_component relation["component_id"] = rule["target_component_id"] @@ -789,14 +790,14 @@ def _transform_candidate( source_machine = runtime["current_definition"]["machine"]["machine_id"] rule = owned_mapping.get((source_spawn, source_machine)) if rule is None: - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) relation["current_spawn_action_pointer"] = rule["target_spawn_action_pointer"] target_root = _target_root_for_machine(target_bundle, rule["target_machine_id"]) holder = relation["lifetime_holder"] if holder is not None: source_holder = holder["variable_declaration_pointer"] if source_holder not in holder_mapping: - raise ArtifactError("migration_totality_failure") + raise ArtifactError(PersistenceCode.MIGRATION_TOTALITY_FAILURE) holder["variable_declaration_pointer"] = holder_mapping[source_holder] source_machine_model = _machine_model_for_root(source_bundle, source_root) target_machine_model = _machine_model_for_root(target_bundle, target_root) @@ -839,7 +840,7 @@ def _transform_candidate( decimal(requirements["maximum_cel_evaluation_steps"]), ) ): - raise ArtifactError("migration_resource_limit_exceeded") + raise ArtifactError(PersistenceCode.MIGRATION_RESOURCE_LIMIT_EXCEEDED) root = next( runtime for runtime in candidate["runtimes"] @@ -872,9 +873,9 @@ def _apply_descriptor( ) if root["status"] in {"completed", "faulted"}: if not maintenance_mode: - raise ArtifactError("terminal_migration_requires_maintenance") + raise ArtifactError(PersistenceCode.TERMINAL_MIGRATION_REQUIRES_MAINTENANCE) if descriptor["terminal_policy"][root["status"]] != "preserve": - raise ArtifactError("terminal_migration_rejected") + raise ArtifactError(PersistenceCode.TERMINAL_MIGRATION_REJECTED) _validate_descriptor_semantics(descriptor, source_bundle, target_bundle, limits) candidate = ( _compatible_candidate(source, target_bundle) @@ -908,22 +909,22 @@ def migrate_aggregate( or not all(isinstance(item, str) and item for item in migration_route) or not isinstance(maintenance_mode, bool) ): - raise ArtifactError("invalid_migration_request") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_REQUEST) route = list(migration_route) if len(route) > limits.maximum_chain_length: - raise ArtifactError("migration_resource_limit_exceeded") + raise ArtifactError(PersistenceCode.MIGRATION_RESOURCE_LIMIT_EXCEEDED) restored = restore_aggregate(source_copy, artifact_resolver) if not route: if ( restored.aggregate_envelope["validated_bundle_fingerprint"] != target_validated_bundle_fingerprint ): - raise ArtifactError("migration_route_missing") + raise ArtifactError(PersistenceCode.MIGRATION_ROUTE_MISSING) return MigrationResult( copy.deepcopy(restored.aggregate_envelope), source_copy, (), None ) if len(set(route)) != len(route): - raise ArtifactError("migration_route_mismatch") + raise ArtifactError(PersistenceCode.MIGRATION_ROUTE_MISMATCH) descriptors_with_bytes = [ _resolve_descriptor(artifact_resolver, digest) for digest in route ] @@ -939,13 +940,13 @@ def migrate_aggregate( for left, right in zip(descriptors, descriptors[1:], strict=False) ) ): - raise ArtifactError("migration_route_mismatch") + raise ArtifactError(PersistenceCode.MIGRATION_ROUTE_MISMATCH) fingerprints = [descriptors[0]["source_validated_bundle_fingerprint"]] + [ descriptor["target_validated_bundle_fingerprint"] for descriptor in descriptors ] if len(set(fingerprints)) != len(fingerprints): - raise ArtifactError("migration_route_mismatch") + raise ArtifactError(PersistenceCode.MIGRATION_ROUTE_MISMATCH) definitions = [ _bundle_from_resolver( artifact_resolver, @@ -999,7 +1000,7 @@ def migrate_aggregate( return _failure(exc.code) except (CelError, KeyError, TypeError, ValueError) as exc: del exc - return _failure("invalid_migration_descriptor") + return _failure(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) def migrate_and_dispatch( @@ -1029,7 +1030,7 @@ def migrate_and_dispatch( core = dispatch(restored.bundle, restored.state, delivery) state = core["state"] if state is None: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) from .wire import aggregate_envelope envelope = aggregate_envelope(restored.bundle, state) diff --git a/src/determa/state/model.py b/src/determa/state/model.py index b9d6b16..aa95716 100644 --- a/src/determa/state/model.py +++ b/src/determa/state/model.py @@ -5,6 +5,7 @@ from dataclasses import dataclass, field from typing import Any +from .codes import MachineLoadFailureCode as LoadCode from .definition import Bundle, _escape_pointer from .errors import ValidationError @@ -103,7 +104,7 @@ def _build( def _index(self, state: StateNode) -> None: if state.path in self.states: - raise ValidationError("semantic_validation", path=state.pointer) + raise ValidationError(LoadCode.SEMANTIC_VALIDATION, path=state.pointer) self.states[state.path] = state for child in state.children.values(): self._index(child) @@ -129,7 +130,9 @@ def resolve(self, target: str | dict[str, str], source: StateNode | None = None) try: return self.states[path] except KeyError as exc: - raise ValidationError("semantic_validation", message=f"unknown state {path}") from exc + raise ValidationError( + LoadCode.SEMANTIC_VALIDATION, message=f"unknown state {path}" + ) from exc def leaves_under(self, state: StateNode) -> list[StateNode]: return [ @@ -157,7 +160,7 @@ def __init__(self, bundle: Bundle) -> None: for index, raw in enumerate(bundle.raw["machines"]): machine_id = str(raw["machine_id"]) if machine_id in self.machines: - raise ValidationError("semantic_validation", message="duplicate machine_id") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION, message="duplicate machine_id") self.machines[machine_id] = MachineModel(bundle, raw, machine_index=index) def machine(self, machine_id: str) -> MachineModel: @@ -165,7 +168,7 @@ def machine(self, machine_id: str) -> MachineModel: return self.machines[machine_id] except KeyError as exc: raise ValidationError( - "semantic_validation", message=f"unknown machine {machine_id}" + LoadCode.SEMANTIC_VALIDATION, message=f"unknown machine {machine_id}" ) from exc def inline_component( diff --git a/src/determa/state/stores/base.py b/src/determa/state/stores/base.py index 8266df7..9d4d9b4 100644 --- a/src/determa/state/stores/base.py +++ b/src/determa/state/stores/base.py @@ -7,6 +7,7 @@ from contextlib import AbstractContextManager from typing import Any +from ..codes import ExecutionStoreAdapterFailureCode as AdapterCode from ..errors import ArtifactError, DetermaError from ..wire import strict_json @@ -39,8 +40,8 @@ class ExecutionStoreError(DetermaError): """A closed execution-store or adapter failure.""" def __init__(self, code: str, message: str = "") -> None: - self.code = code - self.message = message or code + self.code = str(code) + self.message = message or self.code super().__init__(self.message) @@ -96,7 +97,7 @@ def shared_transaction( ) -> AbstractContextManager[tuple[Any, ExecutionStoreTransaction]]: """Open one host-owned native transaction for application composition.""" del root_instance_id - raise ExecutionStoreError("adapter_capability_mismatch") + raise ExecutionStoreError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) @abstractmethod def setup_schema(self) -> None: diff --git a/src/determa/state/stores/file.py b/src/determa/state/stores/file.py index 0f4270c..e29b56c 100644 --- a/src/determa/state/stores/file.py +++ b/src/determa/state/stores/file.py @@ -11,6 +11,7 @@ from typing import Any, BinaryIO from urllib.parse import unquote, urlsplit +from ..codes import ExecutionStoreAdapterFailureCode as AdapterCode from .base import ( RESTART_PERSISTENT, ExecutionStore, @@ -172,13 +173,13 @@ def file_execution_store_factory( """Create the ordinary bundled file adapter.""" parsed = urlsplit(uri) if parsed.scheme != "file" or parsed.netloc not in {"", "localhost"}: - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) if parsed.query or parsed.fragment or set(configuration) - {"directory"}: - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) configured = configuration.get("directory") if configured is not None and not isinstance(configured, str): - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) directory = configured if configured is not None else unquote(parsed.path) if not directory or not Path(directory).is_absolute(): - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) return FileExecutionStore(directory) diff --git a/src/determa/state/stores/memory.py b/src/determa/state/stores/memory.py index fc9176f..8c34a97 100644 --- a/src/determa/state/stores/memory.py +++ b/src/determa/state/stores/memory.py @@ -7,6 +7,7 @@ from contextlib import contextmanager from typing import Any +from ..codes import ExecutionStoreAdapterFailureCode as AdapterCode from .base import ( EPHEMERAL, ExecutionStore, @@ -108,5 +109,5 @@ def memory_execution_store_factory( if uri != "memory:" or configuration: from .base import ExecutionStoreError - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) return MemoryExecutionStore() diff --git a/src/determa/state/stores/postgresql.py b/src/determa/state/stores/postgresql.py index 6a0c996..f9aa179 100644 --- a/src/determa/state/stores/postgresql.py +++ b/src/determa/state/stores/postgresql.py @@ -9,6 +9,7 @@ from typing import Any from urllib.parse import urlsplit +from ..codes import ExecutionStoreAdapterFailureCode as AdapterCode from .base import ( COMPACT_EFFECT_IDENTITY_RETENTION, DURABLE_CONCURRENT, @@ -133,7 +134,7 @@ def __init__( or replay_retention not in _REPLAY_RETENTION_MODES or outbox_retention not in _OUTBOX_RETENTION_MODES ): - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) self.conninfo = conninfo self.table_name = table_name self.metadata_table = metadata_table @@ -430,13 +431,13 @@ def postgresql_execution_store_factory( """Create the ordinary bundled PostgreSQL adapter without importing Psycopg.""" parsed = urlsplit(uri) if parsed.scheme != "postgresql" or parsed.fragment: - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) if set(configuration) - { "table_name", "replay_retention", "outbox_retention", }: - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) table_name = configuration.get( "table_name", "determa_execution_checkpoints" ) @@ -446,7 +447,7 @@ def postgresql_execution_store_factory( isinstance(value, str) for value in (table_name, replay_retention, outbox_retention) ): - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) return PostgreSQLExecutionStore( uri, table_name=table_name, diff --git a/src/determa/state/stores/registry.py b/src/determa/state/stores/registry.py index d646f8f..e467c1f 100644 --- a/src/determa/state/stores/registry.py +++ b/src/determa/state/stores/registry.py @@ -7,6 +7,7 @@ from typing import Any from urllib.parse import urlsplit +from ..codes import ExecutionStoreAdapterFailureCode as AdapterCode from .base import ExecutionStore, ExecutionStoreError ExecutionStoreFactory = Callable[[str, Mapping[str, Any]], ExecutionStore] @@ -25,9 +26,9 @@ def identifiers(self) -> tuple[str, ...]: def register(self, identifier: str, factory: ExecutionStoreFactory) -> None: if _IDENTIFIER.fullmatch(identifier) is None: - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) if identifier in self._factories: - raise ExecutionStoreError("duplicate_adapter_registration") + raise ExecutionStoreError(AdapterCode.DUPLICATE_ADAPTER_REGISTRATION) self._factories[identifier] = factory def resolve( @@ -38,21 +39,21 @@ def resolve( required_capabilities: set[str] | frozenset[str] = frozenset(), ) -> ExecutionStore: if not isinstance(uri, str): - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) scheme = urlsplit(uri).scheme factory = self._factories.get(scheme) if factory is None: - raise ExecutionStoreError("unknown_adapter") + raise ExecutionStoreError(AdapterCode.UNKNOWN_ADAPTER) try: store = factory(uri, dict(configuration or {})) except ExecutionStoreError: raise except (TypeError, ValueError) as exc: - raise ExecutionStoreError("invalid_adapter_configuration") from exc + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) from exc if required_capabilities and not required_capabilities.issubset( store.capabilities ): - raise ExecutionStoreError("adapter_capability_mismatch") + raise ExecutionStoreError(AdapterCode.ADAPTER_CAPABILITY_MISMATCH) return store diff --git a/src/determa/state/stores/sqlite.py b/src/determa/state/stores/sqlite.py index c5d475f..8f0b010 100644 --- a/src/determa/state/stores/sqlite.py +++ b/src/determa/state/stores/sqlite.py @@ -10,6 +10,7 @@ from typing import Any from urllib.parse import parse_qs, unquote, urlsplit +from ..codes import ExecutionStoreAdapterFailureCode as AdapterCode from .base import ( COMPACT_EFFECT_IDENTITY_RETENTION, DURABLE_SINGLE_WRITER, @@ -135,7 +136,7 @@ def __init__( or replay_retention not in _REPLAY_RETENTION_MODES or outbox_retention not in _OUTBOX_RETENTION_MODES ): - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) @property def capabilities(self) -> frozenset[str]: @@ -188,7 +189,7 @@ def _connect(self) -> sqlite3.Connection: or int(actual_synchronous[0]) != expected_synchronous ): connection.close() - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) return connection def _validate_table( @@ -452,7 +453,7 @@ def _single_query(query: Mapping[str, list[str]], key: str, default: str) -> str if values is None: return default if len(values) != 1: - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) return values[0] @@ -467,10 +468,10 @@ def sqlite_execution_store_factory( or parsed.fragment or configuration ): - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) path = unquote(parsed.path) if not path or not Path(path).is_absolute(): - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) query = parse_qs(parsed.query, keep_blank_values=True) if set(query) - { "journal_mode", @@ -479,7 +480,7 @@ def sqlite_execution_store_factory( "replay_retention", "outbox_retention", }: - raise ExecutionStoreError("invalid_adapter_configuration") + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) journal_mode = _single_query(query, "journal_mode", "WAL") synchronous = _single_query(query, "synchronous", "FULL") timeout_text = _single_query(query, "timeout", "30") @@ -488,7 +489,7 @@ def sqlite_execution_store_factory( try: timeout = float(timeout_text) except ValueError as exc: - raise ExecutionStoreError("invalid_adapter_configuration") from exc + raise ExecutionStoreError(AdapterCode.INVALID_ADAPTER_CONFIGURATION) from exc return SQLiteExecutionStore( path, journal_mode=journal_mode, diff --git a/src/determa/state/validator.py b/src/determa/state/validator.py index 822c9a7..6f43c8d 100644 --- a/src/determa/state/validator.py +++ b/src/determa/state/validator.py @@ -8,6 +8,7 @@ from typing import Any, cast from . import cel +from .codes import MachineLoadFailureCode as LoadCode from .definition import Bundle, _escape_pointer, normalize_bundle from .errors import CelError, ErrorRecord, ValidationError from .model import BundleModel, MachineModel, StateNode @@ -195,25 +196,25 @@ def _check_expression( owner_fields=owner_fields if allow_owner else None, ) except cel.CelProfileError as exc: - raise ValidationError("cel_profile_error", message=str(exc)) from exc + raise ValidationError(LoadCode.CEL_PROFILE_ERROR, message=str(exc)) from exc except CelError as exc: - raise ValidationError("semantic_validation", message=str(exc)) from exc + raise ValidationError(LoadCode.SEMANTIC_VALIDATION, message=str(exc)) from exc def _validate_semantics(bundle: Bundle, model: BundleModel) -> None: if not isinstance(bundle.raw.get("format"), int) or isinstance(bundle.raw.get("format"), bool): - raise ValidationError("unsupported_format") + raise ValidationError(LoadCode.UNSUPPORTED_FORMAT) events = bundle.raw.get("events") or {} for _name, declaration in events.items(): correlation = declaration.get("correlates_to") if correlation is not None: target = events.get(correlation) if declaration["direction"] != "input" or not target or target["direction"] != "output": - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) graph: dict[str, set[str]] = {machine_id: set() for machine_id in model.machines} for machine in model.machines.values(): if not isinstance(machine.raw["version"], int) or isinstance(machine.raw["version"], bool): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) _validate_machine(bundle, model, machine, graph) _reject_initialization_cycles(graph) @@ -227,7 +228,7 @@ def _validate_machine( declarations = _event_declarations(bundle, machine) for name, declaration in (machine.raw.get("events") or {}).items(): if name in _RESERVED_EVENTS or declaration["direction"] != "internal": - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) component_ids: set[str] = set() for state in machine.states.values(): _validate_variable_literals(state) @@ -241,17 +242,17 @@ def _validate_variable_literals(state: StateNode) -> None: for declaration in (state.raw.get("variables") or {}).values(): expected = str(declaration["type"]) if "init" in declaration and not _literal_matches(declaration["init"], expected): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) if expected == "int" and isinstance(declaration.get("init"), float): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) def _validate_payload_literals(declaration: dict[str, Any]) -> None: for field in (declaration.get("payload") or {}).values(): if "default" in field and not _literal_matches(field["default"], str(field["type"])): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) if field["type"] == "int" and isinstance(field.get("default"), float): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) def _validate_state_structure( @@ -301,7 +302,7 @@ def _validate_state_structure( if isinstance(initial, dict): target = machine.resolve(initial["transition_to"], state) if not state.is_ancestor_of(target, strict=True): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) _validate_transition( initial, bundle=bundle, @@ -319,13 +320,13 @@ def _validate_state_structure( for event_name, transition_or_list in (state.raw.get("on_events") or {}).items(): declaration = events.get(event_name) if declaration is None and event_name not in _RESERVED_EVENTS: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) transitions = ( transition_or_list if isinstance(transition_or_list, list) else [transition_or_list] ) for index, transition in enumerate(transitions): if index < len(transitions) - 1 and "guard" not in transition: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) suffix = f"/{index}" if isinstance(transition_or_list, list) else "" _validate_transition( transition, @@ -345,7 +346,7 @@ def _validate_state_structure( branches = state.raw["choice"] defaults = [index for index, branch in enumerate(branches) if "guard" not in branch] if defaults != [len(branches) - 1]: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) for index, branch in enumerate(branches): _validate_transition( branch, @@ -364,7 +365,7 @@ def _validate_state_structure( for index, placement in enumerate(state.raw.get("components") or []): component_id = str(placement["component_id"]) if component_id in component_ids: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) component_ids.add(component_id) pointer = f"{state.pointer}/components/{index}" if "machine_id" in placement: @@ -470,18 +471,18 @@ def _validate_target_shape( transition: dict[str, Any], ) -> None: if target is machine.root: - raise ValidationError("root_reentry") + raise ValidationError(LoadCode.ROOT_REENTRY) local = transition.get("local") is True if local: if source is machine.root: - raise ValidationError("root_local_transition") + raise ValidationError(LoadCode.ROOT_LOCAL_TRANSITION) if source.type != "composite" or not source.is_ancestor_of(target, strict=True): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) if isinstance(target_spec, dict): if target.type != "composite" or target.raw.get("history", "none") == "none": - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) if target.is_ancestor_of(source, strict=True): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) def _transition_boundary( @@ -517,14 +518,14 @@ def _validate_destroyed_destinations( name = next(iter(action["assign"])) declaration_state = declarations[name][0] if boundary.is_ancestor_of(declaration_state, strict=True): - raise ValidationError("destroyed_variable_write") + raise ValidationError(LoadCode.DESTROYED_VARIABLE_WRITE) if "refresh" in action and target.type == "final" and target.parent is machine.root: - raise ValidationError("destroyed_variable_write") + raise ValidationError(LoadCode.DESTROYED_VARIABLE_WRITE) if "spawn" in action and "bind_to" in action["spawn"]: name = action["spawn"]["bind_to"] declaration_state = declarations[name][0] if boundary.is_ancestor_of(declaration_state, strict=True): - raise ValidationError("destroyed_reference_binding") + raise ValidationError(LoadCode.DESTROYED_REFERENCE_BINDING) def _validate_actions( @@ -547,7 +548,7 @@ def _validate_actions( if "assign" in action: name, expression = next(iter(action["assign"].items())) if name not in scope_declarations or scope_declarations[name][1].get("external"): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) _check_expression( expression, scope=scope, @@ -584,13 +585,13 @@ def _validate_actions( bind_to = spawn.get("bind_to") if bind_to is not None: if bind_to not in scope_declarations: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) declaration = scope_declarations[bind_to][1] if declaration["type"] != "instance_reference": - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) constraint = declaration.get("machine_id") if constraint is not None and constraint != target.machine_id: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) elif "cancel" in action: _check_expression( action["cancel"]["instance"], @@ -603,7 +604,7 @@ def _validate_actions( ) elif "refresh" in action: if event_name != "env": - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) def _validate_send( @@ -624,12 +625,12 @@ def _validate_send( external = any(target.get("external") is True for target in targets) if event_name == "env": if len(targets) != 1 or "component" not in targets[0]: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) changed = (send.get("payload") or {}).get("changed") if not isinstance(changed, str) or not changed.strip().startswith("{"): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) if changed.strip() == "{}" or "correlation_id" in send: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) component_id = targets[0]["component"] placement = next( ( @@ -640,7 +641,7 @@ def _validate_send( None, ) if placement is None: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) pointer = ( f"{state.pointer}/components/{(state.raw.get('components') or []).index(placement)}" ) @@ -662,28 +663,28 @@ def _validate_send( event_fields=event_fields, ) except cel.CelProfileError as exc: - raise ValidationError("cel_profile_error", message=str(exc)) from exc + raise ValidationError(LoadCode.CEL_PROFILE_ERROR, message=str(exc)) from exc except CelError as exc: - raise ValidationError("semantic_validation", message=str(exc)) from exc + raise ValidationError(LoadCode.SEMANTIC_VALIDATION, message=str(exc)) from exc return if declaration is None or event_name in _RESERVED_EVENTS: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) expected_direction = "output" if external else "internal" if declaration["direction"] != expected_direction: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) if external and "correlation_id" not in send: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) payload_types = _payload_types(declaration, event_name) supplied = send.get("payload") or {} if set(supplied) - set(payload_types): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) required = { name for name, field in (declaration.get("payload") or {}).items() if field.get("required") is True and "default" not in field } if required - set(supplied): - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) for name, expression in supplied.items(): _check_expression( expression, @@ -733,14 +734,14 @@ def _validate_bindings( } supplied = bindings.get(kind) or {} if set(supplied) - set(expected): - raise ValidationError("invalid_binding") + raise ValidationError(LoadCode.INVALID_BINDING) missing = { name for name, declaration in expected.items() if name not in supplied and "init" not in declaration } if missing: - raise ValidationError("invalid_binding") + raise ValidationError(LoadCode.INVALID_BINDING) for name, expression in supplied.items(): _check_expression( expression, @@ -786,7 +787,7 @@ def enter(state: StateNode) -> None: changed = len(reachable) != before unreachable = [state for state in machine.states.values() if state.path not in reachable] if unreachable: - raise ValidationError("semantic_validation", path=unreachable[0].pointer) + raise ValidationError(LoadCode.SEMANTIC_VALIDATION, path=unreachable[0].pointer) def _reject_initialization_cycles(graph: dict[str, set[str]]) -> None: @@ -795,7 +796,7 @@ def _reject_initialization_cycles(graph: dict[str, set[str]]) -> None: def visit(machine_id: str) -> None: if machine_id in visiting: - raise ValidationError("semantic_validation") + raise ValidationError(LoadCode.SEMANTIC_VALIDATION) if machine_id in visited: return visiting.add(machine_id) diff --git a/src/determa/state/wire.py b/src/determa/state/wire.py index 0a0690f..2baa2bd 100644 --- a/src/determa/state/wire.py +++ b/src/determa/state/wire.py @@ -15,6 +15,15 @@ import rfc8785 +from .codes import ( + CheckpointArtifactFailureCode as CheckpointCode, +) +from .codes import ( + MachineLoadFailureCode as LoadCode, +) +from .codes import ( + PersistenceFailureCode as PersistenceCode, +) from .definition import Bundle, BundleSource, _escape_pointer, load_bundle from .errors import ArtifactError, ValidationError from .model import BundleModel, MachineModel, StateNode @@ -107,14 +116,14 @@ def put_definition( ) -> None: bundle = definition if isinstance(definition, Bundle) else load_bundle(definition) if bundle.fingerprint != fingerprint: - raise ArtifactError("definition_fingerprint_mismatch") + raise ArtifactError(PersistenceCode.DEFINITION_FINGERPRINT_MISMATCH) existing = self._definitions.get(fingerprint) if existing is not None: current = existing if isinstance(existing, Bundle) else load_bundle(existing) if canonical_bytes(typed_value(current.raw)) != canonical_bytes( typed_value(bundle.raw) ): - raise ArtifactError("definition_fingerprint_mismatch") + raise ArtifactError(PersistenceCode.DEFINITION_FINGERPRINT_MISMATCH) else: self._definitions[fingerprint] = bundle if trusted: @@ -125,12 +134,12 @@ def put_migration_descriptor( ) -> None: document, _ = load_json_artifact(descriptor, "migration_descriptor") if migration_descriptor_digest(document) != digest: - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) existing = self._migration_descriptors.get(digest) if existing is not None: current, _ = load_json_artifact(existing, "migration_descriptor") if canonical_bytes(current) != canonical_bytes(document): - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) else: self._migration_descriptors[digest] = copy.deepcopy(document) if trusted: @@ -166,7 +175,7 @@ def _source_bytes(source: ArtifactSource) -> bytes: except ValidationError as exc: raise ArtifactError("invalid_json_value") from exc if not validate_unicode(source): - raise ArtifactError("invalid_unicode") + raise ArtifactError(LoadCode.INVALID_UNICODE) return canonical_bytes(copy.deepcopy(dict(source))) @@ -176,7 +185,7 @@ def strict_json(source: ArtifactSource) -> tuple[Any, bytes]: try: text = raw.decode("utf-8", errors="strict") except UnicodeDecodeError as exc: - raise ArtifactError("invalid_unicode") from exc + raise ArtifactError(LoadCode.INVALID_UNICODE) from exc try: value = json.loads( text, @@ -192,7 +201,7 @@ def strict_json(source: ArtifactSource) -> tuple[Any, bytes]: except ValidationError as exc: raise ArtifactError("invalid_json_value") from exc if not validate_unicode(value): - raise ArtifactError("invalid_unicode") + raise ArtifactError(LoadCode.INVALID_UNICODE) return value, raw @@ -216,38 +225,38 @@ def typed_value(value: Any) -> list[Any]: return ["boolean", value] if isinstance(value, str): if not validate_unicode(value): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) return ["string", value] if isinstance(value, int): if not _INT_MIN <= value <= _INT_MAX: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) return ["integer", str(value)] if isinstance(value, float): if not math.isfinite(value): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) return ["float", struct.pack("!d", 0.0 if value == 0.0 else value).hex()] if isinstance(value, list): return ["list", [typed_value(item) for item in value]] if isinstance(value, dict): if not all(isinstance(key, str) for key in value): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) entries = [ [key, typed_value(value[key])] for key in sorted(value, key=lambda item: item.encode("utf-8")) ] return ["map", entries] - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) def decoded_typed_value(value: Any) -> Any: """Decode one exact typed value, rejecting noncanonical representations.""" if not isinstance(value, list) or not value or not isinstance(value[0], str): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) tag = value[0] if tag == "null" and value == ["null"]: return None if len(value) != 2: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) payload = value[1] if tag == "boolean" and isinstance(payload, bool): return payload @@ -277,14 +286,14 @@ def decoded_typed_value(value: Any) -> Any: or not isinstance(entry[0], str) or not validate_unicode(entry[0]) ): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) encoded = entry[0].encode("utf-8") if previous is not None and encoded <= previous: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) previous = encoded result[entry[0]] = decoded_typed_value(entry[1]) return result - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) def _signed_decimal(value: str) -> int: @@ -299,19 +308,19 @@ def _signed_decimal(value: str) -> int: or not digits.isdigit() or digits.startswith("0") ): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) try: return int(value) except ValueError as exc: - raise ArtifactError("invalid_aggregate_state") from exc + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) from exc def decimal(value: Any, *, positive: bool = False) -> int: if not isinstance(value, str): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) number = _signed_decimal(value) if number < 0 or (positive and number == 0): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) return number @@ -355,32 +364,32 @@ def _format_code(document: Any, kind: str) -> str | None: "determa.aggregate_state", "aggregate_state_schema_version", 1, - "unsupported_aggregate_state_format", - "unsupported_aggregate_state_schema_version", + PersistenceCode.UNSUPPORTED_AGGREGATE_STATE_FORMAT, + PersistenceCode.UNSUPPORTED_AGGREGATE_STATE_SCHEMA_VERSION, ), "migration_descriptor": ( "migration_descriptor_format", "determa.aggregate_migration", "migration_descriptor_schema_version", 1, - "unsupported_migration_descriptor_format", - "unsupported_migration_descriptor_schema_version", + PersistenceCode.UNSUPPORTED_MIGRATION_DESCRIPTOR_FORMAT, + PersistenceCode.UNSUPPORTED_MIGRATION_DESCRIPTOR_SCHEMA_VERSION, ), "aggregate_state_package": ( "aggregate_state_package_format", "determa.aggregate_state_package", "aggregate_state_package_schema_version", 1, - "unsupported_aggregate_state_package_format", - "unsupported_aggregate_state_package_schema_version", + PersistenceCode.UNSUPPORTED_AGGREGATE_STATE_PACKAGE_FORMAT, + PersistenceCode.UNSUPPORTED_AGGREGATE_STATE_PACKAGE_SCHEMA_VERSION, ), "execution_checkpoint": ( "execution_checkpoint_format", "determa.execution_checkpoint", "execution_checkpoint_schema_version", 1, - "unsupported_execution_checkpoint_format", - "unsupported_execution_checkpoint_schema_version", + CheckpointCode.UNSUPPORTED_EXECUTION_CHECKPOINT_FORMAT, + CheckpointCode.UNSUPPORTED_EXECUTION_CHECKPOINT_SCHEMA_VERSION, ), } format_member, expected_format, version_member, expected_version, format_code, version_code = ( @@ -401,10 +410,10 @@ def load_json_artifact( document, raw = strict_json(source) except ArtifactError as exc: code = { - "aggregate_state": "invalid_aggregate_state", - "migration_descriptor": "invalid_migration_descriptor", - "aggregate_state_package": "invalid_aggregate_state_package", - "execution_checkpoint": "invalid_execution_checkpoint", + "aggregate_state": PersistenceCode.INVALID_AGGREGATE_STATE, + "migration_descriptor": PersistenceCode.INVALID_MIGRATION_DESCRIPTOR, + "aggregate_state_package": PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE, + "execution_checkpoint": CheckpointCode.INVALID_EXECUTION_CHECKPOINT, }[kind] raise ArtifactError(code) from exc unsupported = _format_code(document, kind) @@ -417,10 +426,10 @@ def load_json_artifact( ) if not isinstance(document, dict) or next(validator.iter_errors(document), None) is not None: code = { - "aggregate_state": "invalid_aggregate_state", - "migration_descriptor": "invalid_migration_descriptor", - "aggregate_state_package": "invalid_aggregate_state_package", - "execution_checkpoint": "invalid_execution_checkpoint", + "aggregate_state": PersistenceCode.INVALID_AGGREGATE_STATE, + "migration_descriptor": PersistenceCode.INVALID_MIGRATION_DESCRIPTOR, + "aggregate_state_package": PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE, + "execution_checkpoint": CheckpointCode.INVALID_EXECUTION_CHECKPOINT, }[kind] raise ArtifactError(code) return document, raw @@ -450,12 +459,12 @@ def bundle_from_attachment(attachment: Mapping[str, Any]) -> Bundle: raw = decoded_typed_value(attachment["normalized_bundle"]) fingerprint = attachment["validated_bundle_fingerprint"] except (ArtifactError, KeyError, TypeError) as exc: - raise ArtifactError("invalid_aggregate_state_package") from exc + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE) from exc if not isinstance(raw, dict) or not isinstance(fingerprint, str): - raise ArtifactError("invalid_aggregate_state_package") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE) bundle = load_bundle(raw) if bundle.fingerprint != fingerprint: - raise ArtifactError("invalid_aggregate_state_package") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE) return bundle @@ -469,16 +478,18 @@ def _bundle_from_resolver( definition = resolver.resolve_definition(fingerprint) if definition is None: raise ArtifactError( - "source_definition_unavailable" if source else "target_definition_unavailable" + PersistenceCode.SOURCE_DEFINITION_UNAVAILABLE + if source + else PersistenceCode.TARGET_DEFINITION_UNAVAILABLE ) if require_trust and not resolver.definition_is_trusted(fingerprint): - raise ArtifactError("definition_untrusted") + raise ArtifactError(PersistenceCode.DEFINITION_UNTRUSTED) try: bundle = definition if isinstance(definition, Bundle) else load_bundle(definition) except ValidationError as exc: - raise ArtifactError("definition_fingerprint_mismatch") from exc + raise ArtifactError(PersistenceCode.DEFINITION_FINGERPRINT_MISMATCH) from exc if bundle.fingerprint != fingerprint: - raise ArtifactError("definition_fingerprint_mismatch") + raise ArtifactError(PersistenceCode.DEFINITION_FINGERPRINT_MISMATCH) return bundle @@ -581,7 +592,7 @@ def _node_for_runtime(machine: MachineModel, path: str) -> StateNode: try: return machine.states[path] except KeyError as exc: - raise ArtifactError("invalid_aggregate_state") from exc + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) from exc def _runtime_wire( @@ -717,7 +728,7 @@ def aggregate_envelope(bundle: Bundle | BundleSource, state: dict[str, Any]) -> from .engine import _valid_prior_state if not _valid_prior_state(state, validated): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) models = BundleModel(validated) root = state["runtimes"][state["root_runtime_id"]] restored_order = state.get("_wire_runtime_order") @@ -768,7 +779,7 @@ def _state_path_for_pointer(machine: MachineModel, pointer: str) -> str: for path, node in machine.states.items(): if node.pointer == pointer: return path - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) def _variable_for_pointer( @@ -780,14 +791,14 @@ def _variable_for_pointer( for name, declaration in (node.raw.get("variables") or {}).items(): if pointer == f"{prefix}{_escape_pointer(name)}": return path, name, declaration - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) def _history_path_for_pointer(machine: MachineModel, pointer: str) -> str: for path, node in machine.states.items(): if pointer == f"{node.pointer}/history": return "$root" if node is machine.root else path - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) def _machine_for_binding( @@ -800,7 +811,7 @@ def _machine_for_binding( machine_data["namespace"] != bundle.namespace or decimal(machine_data["machine_version"], positive=True) != base.version ): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) root_pointer = machine_data["root_definition_pointer"] if root_pointer == base.root_pointer: return base @@ -822,10 +833,10 @@ def _origin_machine( ) -> tuple[Bundle, MachineModel]: definition = origin.get("definition") if not isinstance(definition, Mapping): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) fingerprint = definition.get("validated_bundle_fingerprint") if not isinstance(fingerprint, str): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) bundle = _bundle_from_resolver(resolver, fingerprint, source=True, require_trust=True) return bundle, _machine_for_binding(bundle, BundleModel(bundle), definition) @@ -843,7 +854,7 @@ def _validate_immutable_identity( "component": "component", "spawned": "owned_spawned_instance", }[role]: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) origin_bundle, origin_machine = _origin_machine(resolver, origin) root_instance_id = aggregate["root_instance_id"] runtime_id = document["runtime_id"] @@ -858,22 +869,22 @@ def _validate_immutable_identity( } } ): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) return if role == "component": component = target.get("component") if not isinstance(component, Mapping): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) pointer = origin.get("component_definition_pointer") if not isinstance(pointer, str): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) try: from .engine import _pointer_get placement = _pointer_get(origin_bundle.raw, pointer) declaration_index = int(pointer.rsplit("/", 1)[1]) except (IndexError, KeyError, TypeError, ValueError): - raise ArtifactError("invalid_aggregate_state") from None + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) from None if ( not isinstance(placement, dict) or decimal(origin.get("declaration_index")) != declaration_index @@ -886,7 +897,7 @@ def _validate_immutable_identity( "activation_sequence": decimal(origin.get("activation_sequence")), } ): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) return spawned = target.get("spawned_instance") if ( @@ -899,7 +910,7 @@ def _validate_immutable_identity( "machine_version": origin_machine.version, } ): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) def _runtime_from_wire( @@ -911,7 +922,7 @@ def _runtime_from_wire( ) -> dict[str, Any]: current = document["current_definition"] if current["validated_bundle_fingerprint"] != bundle.fingerprint: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) machine = _machine_for_binding(bundle, models, current) relation = document["relation"] role = { @@ -923,41 +934,41 @@ def _runtime_from_wire( for item in document["active_state_activations"]: path = _state_path_for_pointer(machine, item["state_definition_pointer"]) if path in activations: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) activations[path] = decimal(item["activation_sequence"]) leaf_pointers = document["active_leaf_state_definition_pointers"] if len(leaf_pointers) > 1: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) active: list[str] = [] if leaf_pointers: leaf = machine.states[_state_path_for_pointer(machine, leaf_pointers[0])] active = [node.path for node in reversed(leaf.ancestors(include_self=True))] if set(active) != set(activations): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) scopes: dict[str, dict[str, Any]] = {path: {} for path in active} for item in document["variables"]: path, name, declaration = _variable_for_pointer( machine, item["variable_declaration_pointer"] ) if path not in scopes or name in scopes[path]: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) if decimal(item["declaring_state_activation_sequence"]) != activations[path]: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) value = decoded_typed_value(item["value"]) from .engine import _value_matches if not _value_matches(value, str(declaration["type"])): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) scopes[path][name] = value for path in active: declarations = machine.states[path].raw.get("variables") or {} if set(scopes[path]) != set(declarations): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) history: dict[str, list[str] | None] = {} for item in document["history"]: path = _history_path_for_pointer(machine, item["history_declaration_pointer"]) if path in history: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) recorded = item["recorded_state_definition_pointers"] history[path] = ( None @@ -968,13 +979,13 @@ def _runtime_from_wire( for item in document["next_state_activation_sequences"]: path = _state_path_for_pointer(machine, item["definition_pointer"]) if path in state_counters: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) state_counters[path] = decimal(item["next_sequence"]) component_counters: dict[str, int] = {} for item in document["next_component_activation_sequences"]: pointer = item["definition_pointer"] if pointer in component_counters: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) component_counters[pointer] = decimal(item["next_sequence"]) target = copy.deepcopy(document["target_identity"]) if "component" in target: @@ -1064,7 +1075,7 @@ def _finish_relationships( if runtime["role"] == "component": owner = state["runtimes"].get(runtime["owner_runtime_id"]) if owner is None: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) relation = runtime["_relation"] owner_machine = _runtime_model(bundle, models, owner) pointer = relation["current_component_definition_pointer"] @@ -1077,25 +1088,25 @@ def _finish_relationships( None, ) if owning is None or owning.path not in owner["state_activation_sequence"]: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) runtime["owning_state_path"] = owning.path runtime["owning_state_activation_sequence"] = owner[ "state_activation_sequence" ][owning.path] if runtime["component_id"] in owner["components"]: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) owner["components"][runtime["component_id"]] = runtime["runtime_id"] elif runtime["role"] == "spawned" and runtime["holder"] is not None: owner = state["runtimes"].get(runtime["owner_runtime_id"]) if owner is None: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) owner_machine = _runtime_model(bundle, models, owner) holder_pointer = runtime["holder"]["pointer"] path, _name, _declaration = _variable_for_pointer( owner_machine, holder_pointer ) if path not in owner["state_activation_sequence"]: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) runtime["holder"]["state_path"] = path @@ -1105,13 +1116,13 @@ def restore_aggregate( """Verify and restore one portable aggregate without changing the source.""" document, raw = load_json_artifact(source, "aggregate_state") if aggregate_state_digest(document) != document["aggregate_state_digest"]: - raise ArtifactError("aggregate_state_digest_mismatch") + raise ArtifactError(PersistenceCode.AGGREGATE_STATE_DIGEST_MISMATCH) fingerprint = document["validated_bundle_fingerprint"] bundle = _bundle_from_resolver( definition_resolver, fingerprint, source=True, require_trust=True ) if document["namespace"] != bundle.namespace: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) models = BundleModel(bundle) state: dict[str, Any] = { "validated_bundle_fingerprint": fingerprint, @@ -1135,12 +1146,12 @@ def restore_aggregate( bundle, models, definition_resolver, document, runtime_document ) if runtime["runtime_id"] in state["runtimes"]: - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) state["runtimes"][runtime["runtime_id"]] = runtime _finish_relationships(bundle, state, models) root = state["runtimes"].get(state["root_runtime_id"]) if root is None or root["role"] != "root": - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) if ( document["root_machine_id"] != root["machine_id"] or decimal(document["root_machine_version"], positive=True) @@ -1148,13 +1159,13 @@ def restore_aggregate( or root["_current_definition"]["machine"]["root_definition_pointer"] != root["root_pointer"] ): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) state["status"] = root["status"] state["fault"] = copy.deepcopy(root["fault"]) from .engine import _valid_prior_state if not _valid_prior_state(state, bundle): - raise ArtifactError("invalid_aggregate_state") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE) canonical = canonical_bytes(document) return RestoredAggregate( bundle=bundle, @@ -1176,16 +1187,16 @@ def restore_aggregate_package( for attachment in document["normalized_definitions"]: bundle = bundle_from_attachment(attachment) if bundle.fingerprint in definitions: - raise ArtifactError("invalid_aggregate_state_package") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE) definitions[bundle.fingerprint] = bundle for descriptor in document["migration_descriptors"]: digest = descriptor["migration_descriptor_digest"] if digest in descriptors or migration_descriptor_digest(descriptor) != digest: - raise ArtifactError("invalid_aggregate_state_package") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE) descriptors[digest] = copy.deepcopy(descriptor) route = tuple(document["migration_route"]) if len(set(route)) != len(route): - raise ArtifactError("invalid_aggregate_state_package") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE) for fingerprint, bundle in definitions.items(): existing_definition = artifact_resolver.resolve_definition(fingerprint) if existing_definition is not None: @@ -1199,7 +1210,7 @@ def restore_aggregate_package( or canonical_bytes(typed_value(current_bundle.raw)) != canonical_bytes(typed_value(bundle.raw)) ): - raise ArtifactError("definition_fingerprint_mismatch") + raise ArtifactError(PersistenceCode.DEFINITION_FINGERPRINT_MISMATCH) for digest, descriptor in descriptors.items(): existing_descriptor = artifact_resolver.resolve_migration_descriptor(digest) if existing_descriptor is not None: @@ -1207,21 +1218,24 @@ def restore_aggregate_package( existing_descriptor, "migration_descriptor" ) if canonical_bytes(current_descriptor) != canonical_bytes(descriptor): - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) except (ArtifactError, KeyError, TypeError, ValidationError) as exc: if isinstance(exc, ArtifactError) and ( exc.code.startswith("unsupported_") or exc.code - in {"definition_fingerprint_mismatch", "invalid_migration_descriptor"} + in { + PersistenceCode.DEFINITION_FINGERPRINT_MISMATCH, + PersistenceCode.INVALID_MIGRATION_DESCRIPTOR, + } ): raise exc - raise ArtifactError("invalid_aggregate_state_package") from exc + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE) from exc put_definition = getattr(artifact_resolver, "put_definition", None) put_descriptor = getattr(artifact_resolver, "put_migration_descriptor", None) if (definitions and not callable(put_definition)) or ( descriptors and not callable(put_descriptor) ): - raise ArtifactError("invalid_aggregate_state_package") + raise ArtifactError(PersistenceCode.INVALID_AGGREGATE_STATE_PACKAGE) overlay = _PackageResolver(artifact_resolver, definitions, descriptors) try: aggregate = restore_aggregate(document["aggregate_state"], overlay) @@ -1421,4 +1435,4 @@ def _resolve_variable_pointer( if name in (current.raw.get("variables") or {}): return f"{current.pointer}/variables/{_escape_pointer(name)}" current = current.parent - raise ArtifactError("invalid_migration_descriptor") + raise ArtifactError(PersistenceCode.INVALID_MIGRATION_DESCRIPTOR) diff --git a/src/determa/state/yaml12.py b/src/determa/state/yaml12.py index 018b4ad..e0797e4 100644 --- a/src/determa/state/yaml12.py +++ b/src/determa/state/yaml12.py @@ -8,6 +8,7 @@ import yaml +from .codes import MachineLoadFailureCode as LoadCode from .errors import ValidationError _JSON_NUMBER = re.compile(r"-?(?:0|[1-9][0-9]*)(?:\.[0-9]+)?(?:[eE][+-]?[0-9]+)?\Z") @@ -97,16 +98,16 @@ def _validate_portable_values(value: Any, ancestors: set[int]) -> None: return if isinstance(value, int): if not _INT_MIN <= value <= _INT_MAX: - raise ValidationError("numeric_value_out_of_range") + raise ValidationError(LoadCode.NUMERIC_VALUE_OUT_OF_RANGE) return if isinstance(value, float): if not math.isfinite(value): - raise ValidationError("numeric_value_out_of_range") + raise ValidationError(LoadCode.NUMERIC_VALUE_OUT_OF_RANGE) return if isinstance(value, list): identity = id(value) if identity in ancestors: - raise ValidationError("non_json_value") + raise ValidationError(LoadCode.NON_JSON_VALUE) ancestors.add(identity) for item in value: _validate_portable_values(item, ancestors) @@ -115,15 +116,15 @@ def _validate_portable_values(value: Any, ancestors: set[int]) -> None: if isinstance(value, dict): identity = id(value) if identity in ancestors: - raise ValidationError("non_json_value") + raise ValidationError(LoadCode.NON_JSON_VALUE) ancestors.add(identity) for key, item in value.items(): if not isinstance(key, str): - raise ValidationError("non_string_map_key") + raise ValidationError(LoadCode.NON_STRING_MAP_KEY) _validate_portable_values(item, ancestors) ancestors.remove(identity) return - raise ValidationError("non_json_value") + raise ValidationError(LoadCode.NON_JSON_VALUE) def _resolve_plain(value: str) -> Any: @@ -132,28 +133,28 @@ def _resolve_plain(value: str) -> Any: if value == "false": return False if value in _INVALID_BOOLEAN: - raise ValidationError("invalid_boolean_syntax") + raise ValidationError(LoadCode.INVALID_BOOLEAN_SYNTAX) if value == "null": return None if value in _INVALID_NULL: - raise ValidationError("invalid_null_syntax") + raise ValidationError(LoadCode.INVALID_NULL_SYNTAX) if value in _STRING_BOOLEAN_LIKE: return value if _JSON_NUMBER.fullmatch(value): if _INTEGER.fullmatch(value): integer = int(value, 10) if not _INT_MIN <= integer <= _INT_MAX: - raise ValidationError("numeric_value_out_of_range") + raise ValidationError(LoadCode.NUMERIC_VALUE_OUT_OF_RANGE) return integer try: double = float(value) except ValueError as exc: - raise ValidationError("invalid_numeric_syntax") from exc + raise ValidationError(LoadCode.INVALID_NUMERIC_SYNTAX) from exc if not math.isfinite(double): - raise ValidationError("numeric_value_out_of_range") + raise ValidationError(LoadCode.NUMERIC_VALUE_OUT_OF_RANGE) return 0.0 if double == 0.0 else double if _NONPORTABLE_YAML_NUMBER.fullmatch(value): - raise ValidationError("invalid_numeric_syntax") + raise ValidationError(LoadCode.INVALID_NUMERIC_SYNTAX) return value @@ -164,7 +165,7 @@ class _PortableLoader(yaml.BaseLoader): def _construct_scalar(loader: _PortableLoader, node: yaml.ScalarNode) -> Any: value = loader.construct_scalar(node) if _has_invalid_unicode(value): - raise ValidationError("invalid_unicode") + raise ValidationError(LoadCode.INVALID_UNICODE) if node.style is None: return _resolve_plain(value) return value @@ -177,9 +178,9 @@ def _construct_mapping( for key_node, value_node in node.value: key = loader.construct_object(key_node, deep=deep) if not isinstance(key, str): - raise ValidationError("non_string_map_key") + raise ValidationError(LoadCode.NON_STRING_MAP_KEY) if key in result: - raise ValidationError("duplicate_key") + raise ValidationError(LoadCode.DUPLICATE_KEY) result[key] = loader.construct_object(value_node, deep=deep) return result @@ -202,27 +203,30 @@ def _reject_yaml_features(text: str) -> None: if isinstance( token, (yaml.tokens.AliasToken, yaml.tokens.AnchorToken, yaml.tokens.TagToken) ): - raise ValidationError("unsupported_yaml_feature") + raise ValidationError(LoadCode.UNSUPPORTED_YAML_FEATURE) except ValidationError: raise except (yaml.YAMLError, UnicodeError) as exc: - raise ValidationError("non_json_value", message=str(exc)) from exc + raise ValidationError(LoadCode.NON_JSON_VALUE, message=str(exc)) from exc def load(text: str) -> Any: """Parse exactly one portable format-1 source document.""" if _has_invalid_unicode(text): - raise ValidationError("invalid_unicode") + raise ValidationError(LoadCode.INVALID_UNICODE) _reject_yaml_features(text) try: documents = list(yaml.load_all(text, Loader=_PortableLoader)) except ValidationError: raise except (yaml.YAMLError, UnicodeError) as exc: - raise ValidationError("non_json_value", message=str(exc)) from exc + raise ValidationError(LoadCode.NON_JSON_VALUE, message=str(exc)) from exc if len(documents) != 1: - raise ValidationError("non_json_value", message="source must contain exactly one document") + raise ValidationError( + LoadCode.NON_JSON_VALUE, + message="source must contain exactly one document", + ) document = documents[0] if not validate_unicode(document): - raise ValidationError("invalid_unicode") + raise ValidationError(LoadCode.INVALID_UNICODE) return document diff --git a/tests/test_portable_codes.py b/tests/test_portable_codes.py new file mode 100644 index 0000000..1a61614 --- /dev/null +++ b/tests/test_portable_codes.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, cast + +import pytest + +import determa.state as ds + + +def test_portable_code_sets_are_deeply_immutable() -> None: + assert isinstance(ds.PORTABLE_CODE_SETS, Mapping) + assert all(isinstance(codes, frozenset) for codes in ds.PORTABLE_CODE_SETS.values()) + + mutable_view = cast(Any, ds.PORTABLE_CODE_SETS) + with pytest.raises(TypeError): + mutable_view["disposition"] = frozenset() + with pytest.raises(AttributeError): + mutable_view["disposition"].add("unknown") + + +def test_nonportable_categories_are_not_exported() -> None: + assert "execution_store_failure" not in ds.PORTABLE_CODE_SETS + assert "structural_validation" not in ds.PORTABLE_CODE_SETS + + +def test_public_code_definitions_are_immutable_strings() -> None: + code_types = [ + getattr(ds, name) for name in ds.__all__ if name.endswith("Code") + ] + assert code_types + for code_type in code_types: + member = next(iter(code_type)) + assert isinstance(member.value, str) + with pytest.raises(AttributeError): + cast(Any, member).value = "changed"