Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion src/hermes_workflows/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,12 @@
from .approvals import ApprovalDecision, ApprovalDecisionInput, ApprovalReceipt, ApprovalView, OperatorResponseReceipt
from .domain import CommandType, WorkflowStatus, decode_command_row, decode_event_row, make_command, make_event
from .input_parsing import coerce_workflow_input
from .runtime_services import EmptyRuntimeServicesV1, RuntimeOnlyServiceRegistry, RuntimeServiceRegistry
from .runtime_services import (
EmptyRuntimeServicesV1,
RuntimeOnlyServiceRegistry,
RuntimeServiceRegistry,
validate_runtime_service_resolution,
)
from .status_projection import StatusProjection
from .types import to_json_value
from .workflow_values import Workflow
Expand Down Expand Up @@ -127,6 +132,7 @@ def __init__(
self._init_db()

def resolve_runtime_service(self, service_id: str, contract_version: int) -> object | None:
validate_runtime_service_resolution(service_id, contract_version)
return self.runtime_services.resolve(service_id, contract_version)

def _ensure_writable(self, operation: str) -> None:
Expand Down
6 changes: 3 additions & 3 deletions src/hermes_workflows/runtime_services.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ def __post_init__(self) -> None:
object.__setattr__(self, "services", MappingProxyType(validated))

def resolve(self, service_id: str, contract_version: int) -> object | None:
_validate_resolution(service_id, contract_version)
validate_runtime_service_resolution(service_id, contract_version)
return self.services.get(service_id)

def __reduce_ex__(self, protocol: SupportsIndex):
Expand All @@ -55,7 +55,7 @@ def __reduce_ex__(self, protocol: SupportsIndex):
class EmptyRuntimeServicesV1(RuntimeOnlyServiceRegistry):

def resolve(self, service_id: str, contract_version: int) -> object | None:
_validate_resolution(service_id, contract_version)
validate_runtime_service_resolution(service_id, contract_version)
return None

def __reduce_ex__(self, protocol: SupportsIndex):
Expand All @@ -67,7 +67,7 @@ def _validate_service_id(service_id: object) -> None:
raise ValueError("service_id must match ^[a-z][a-z0-9_.-]{0,63}$")


def _validate_resolution(service_id: object, contract_version: object) -> None:
def validate_runtime_service_resolution(service_id: object, contract_version: object) -> None:
_validate_service_id(service_id)
if type(contract_version) is not int or contract_version < 1:
raise ValueError("contract_version must be an integer >= 1")
12 changes: 12 additions & 0 deletions tests/test_runtime_services_contract.py
Original file line number Diff line number Diff line change
Expand Up @@ -348,6 +348,18 @@ def test_engine_preserves_marked_registry_identity_without_mutation_or_wrapping(
assert engine.resolve_runtime_service("integration.any", 1) is None


@pytest.mark.parametrize(
("service_id", "contract_version"),
[("", 1), ("Upper", 1), ("valid.service", 0), ("valid.service", True)],
)
def test_engine_validates_resolution_for_integration_owned_registry(tmp_path, service_id, contract_version):
registry = IntegrationOwnedRuntimeServices(marker="accepted-integration-registry")
engine = WorkflowEngine(tmp_path / "workflow.sqlite", runtime_services=registry)

with pytest.raises(ValueError):
engine.resolve_runtime_service(service_id, contract_version)


def test_engine_rejects_unmarked_structural_registry(tmp_path):
registry = UnmarkedStructuralRuntimeServices()

Expand Down