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
34 changes: 31 additions & 3 deletions src/durable_workflow/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from __future__ import annotations

import asyncio
import copy
import hashlib
import json as json_module
import math
Expand Down Expand Up @@ -2672,6 +2673,26 @@ def get_workflow_handle(
"""
return WorkflowHandle(self, workflow_id=workflow_id, run_id=run_id, workflow_type=workflow_type)

def _nexus_target_client(self, target_namespace: str | None) -> Client:
if target_namespace is not None and not target_namespace.strip():
raise ValueError("target_namespace must be nonempty")
if target_namespace is None or target_namespace == self.namespace:
return self
self._auth_token(worker=False)
scoped = copy.copy(self)
scoped.namespace = target_namespace
# Share the HTTP pool, never the caller's namespace discovery or
# provider credentials. The worker keeps using the original client.
scoped.worker_token = None
scoped._cluster_info = None
scoped._cluster_info_lock = asyncio.Lock()
scoped._runtime_external_payload_transport_cache = None
scoped._runtime_external_payload_transport_resolved = False
scoped.external_storage = None
scoped.external_storage_threshold_bytes = None
scoped.external_storage_cache = ExternalPayloadCache()
return scoped

async def execute_nexus_operation(
self,
endpoint_name: str,
Expand All @@ -2685,6 +2706,7 @@ async def execute_nexus_operation(
idempotency_key: str | None = None,
payload_codec: str | None = None,
caller_namespace: str | None = None,
target_namespace: str | None = None,
caller_workflow_instance_id: str | None = None,
caller_workflow_run_id: str | None = None,
service_sdk_language: str | None = None,
Expand All @@ -2708,7 +2730,11 @@ async def execute_nexus_operation(
the worker uses this client method to make the single durable
execute request and records the returned service-call surface into
workflow history.

``target_namespace`` selects the service catalog and its payload
storage. The caller identity and worker namespace are preserved.
"""
execution_client = self._nexus_target_client(target_namespace)
request_body = nexus_request_payload(
arguments=list(arguments) if arguments is not None else [],
payload_codec=payload_codec,
Expand Down Expand Up @@ -2737,9 +2763,9 @@ async def execute_nexus_operation(
context = f"{endpoint_name}/{service_name}/{operation_name}"

wire_body = request_body
transport = await self._runtime_external_payload_transport()
transport = await execution_client._runtime_external_payload_transport()
if transport is not None:
arguments_envelope = await self._externalize_runtime_payloads(
arguments_envelope = await execution_client._externalize_runtime_payloads(
{"codec": request_body["payload_codec"], "blob": request_body["arguments"]},
worker=False,
transport=transport,
Expand All @@ -2750,7 +2776,7 @@ async def execute_nexus_operation(
wire_body = {**request_body, "arguments": arguments_envelope}

try:
data = await self._request("POST", path, json=wire_body, context=context)
data = await execution_client._request("POST", path, json=wire_body, context=context)
except ServerError as exc:
if exc.status != 409 or not isinstance(exc.body, dict):
raise
Expand All @@ -2760,6 +2786,7 @@ async def execute_nexus_operation(
service_name=service_name,
operation_name=operation_name,
request_payload=request_body,
target_namespace=target_namespace,
service_sdk_language=service_sdk_language,
artifact_tuple=artifact_tuple,
published_artifact_worker_execution=published_artifact_worker_execution,
Expand All @@ -2783,6 +2810,7 @@ async def execute_nexus_operation(
service_name=service_name,
operation_name=operation_name,
request_payload=request_body,
target_namespace=target_namespace,
service_sdk_language=service_sdk_language,
artifact_tuple=artifact_tuple,
published_artifact_worker_execution=published_artifact_worker_execution,
Expand Down
11 changes: 11 additions & 0 deletions src/durable_workflow/nexus.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ def stable_nexus_idempotency_key(
service_name: str,
operation_name: str,
arguments: Sequence[Any],
target_namespace: str | None = None,
) -> str:
"""Return a deterministic idempotency key for one workflow-side Nexus call."""
source = {
Expand All @@ -62,6 +63,8 @@ def stable_nexus_idempotency_key(
"operation_name": operation_name,
"arguments": list(arguments),
}
if target_namespace is not None:
source["target_namespace"] = target_namespace
encoded = json.dumps(source, sort_keys=True, separators=(",", ":"), default=_json_default)
return f"dw-py-nexus-{hashlib.sha256(encoded.encode()).hexdigest()}"

Expand Down Expand Up @@ -136,6 +139,7 @@ class NexusOperationResult:
caller_observed_error_type: str | None = None
typed_error_message: str | None = None
result: Any = None
target_namespace: str | None = None

@classmethod
def from_response(
Expand All @@ -149,6 +153,7 @@ def from_response(
service_sdk_language: str | None = None,
artifact_tuple: Mapping[str, Any] | None = None,
published_artifact_worker_execution: bool | None = None,
target_namespace: str | None = None,
) -> NexusOperationResult:
call = data.get("service_call")
source = call if isinstance(call, Mapping) else data
Expand Down Expand Up @@ -190,6 +195,7 @@ def from_response(
typed_error_message=_optional_str(source.get("typed_error_message"))
or _optional_str(data.get("typed_error_message")),
result=response_value,
target_namespace=target_namespace,
)

@classmethod
Expand All @@ -201,6 +207,9 @@ def from_recorded_payload(cls, value: Any, *, expected: Mapping[str, Any]) -> Ne
if value.get("version") != NEXUS_OPERATION_RESULT_VERSION:
raise ValueError("recorded Nexus operation result version is not supported")

if value.get("target_namespace") != expected.get("target_namespace"):
raise ValueError("recorded Nexus target namespace does not match the current workflow command")

for key in ("endpoint_name", "service_name", "operation_name", "idempotency_key"):
expected_value = expected.get(key)
recorded_value = value.get(key)
Expand Down Expand Up @@ -237,6 +246,7 @@ def from_recorded_payload(cls, value: Any, *, expected: Mapping[str, Any]) -> Ne
caller_observed_error_type=_optional_str(value.get("caller_observed_error_type")),
typed_error_message=_optional_str(value.get("typed_error_message")),
result=value.get("result"),
target_namespace=_optional_str(value.get("target_namespace")),
)
if result.is_failure:
raise result.to_failure()
Expand Down Expand Up @@ -289,6 +299,7 @@ def to_recorded_payload(self) -> dict[str, Any]:
"endpoint_name": self.endpoint_name,
"service_name": self.service_name,
"operation_name": self.operation_name,
"target_namespace": self.target_namespace,
"idempotency_key": self.request_payload.get("idempotency_key"),
"caller_workflow_instance_id": self.caller_workflow_instance_id,
"caller_workflow_run_id": self.caller_workflow_run_id,
Expand Down
1 change: 1 addition & 0 deletions src/durable_workflow/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -1462,6 +1462,7 @@ async def _resolve_workflow_nexus_commands(
idempotency_key=command.idempotency_key,
payload_codec=command.payload_codec,
caller_namespace=command.caller_namespace,
target_namespace=command.target_namespace,
caller_workflow_instance_id=_string_or_none(task.get("workflow_id")),
caller_workflow_run_id=_string_or_none(task.get("run_id")),
service_sdk_language=command.service_sdk_language,
Expand Down
18 changes: 18 additions & 0 deletions src/durable_workflow/workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -1196,6 +1196,7 @@ class NexusServiceCall:
memo: dict[str, Any] | None = None
search_attributes: dict[str, Any] | None = None
duplicate_start_policy: str | None = None
target_namespace: str | None = None

def to_server_command(self, *args: Any, **kwargs: Any) -> dict[str, Any]:
raise RuntimeError(
Expand All @@ -1209,6 +1210,7 @@ def expected_recorded_metadata(self) -> dict[str, Any]:
"service_name": self.service_name,
"operation_name": self.operation_name,
"idempotency_key": self.idempotency_key,
"target_namespace": self.target_namespace,
}

def recorded_result(self, value: Any) -> NexusOperationResult:
Expand Down Expand Up @@ -2402,6 +2404,7 @@ def call_nexus_service(
wait_for: str | None = "completed",
wait_timeout_seconds: int | None = None,
caller_namespace: str | None = None,
target_namespace: str | None = None,
service_sdk_language: str | None = None,
artifact_tuple: Mapping[str, Any] | None = None,
published_artifact_worker_execution: bool | None = None,
Expand All @@ -2420,7 +2423,12 @@ def call_nexus_service(
If the service reports a typed failure, replay raises
:class:`~durable_workflow.errors.NexusOperationFailed` at the yield
point so workflow code can compensate or let the workflow fail.

``target_namespace`` addresses a service in another namespace and
participates in the durable call identity and replay checks.
"""
if target_namespace is not None and not target_namespace.strip():
raise ValueError("target_namespace must be nonempty")
args = list(arguments) if arguments is not None else []
if idempotency_key is None:
idempotency_key = stable_nexus_idempotency_key(
Expand All @@ -2431,6 +2439,7 @@ def call_nexus_service(
service_name=service_name,
operation_name=operation_name,
arguments=args,
target_namespace=target_namespace,
)
self._nexus_call_counter += 1
return NexusServiceCall(
Expand All @@ -2444,6 +2453,7 @@ def call_nexus_service(
wait_for=wait_for,
wait_timeout_seconds=wait_timeout_seconds,
caller_namespace=caller_namespace,
target_namespace=target_namespace,
service_sdk_language=service_sdk_language,
artifact_tuple=dict(artifact_tuple) if artifact_tuple is not None else None,
published_artifact_worker_execution=published_artifact_worker_execution,
Expand Down Expand Up @@ -6589,6 +6599,14 @@ def _advance_selection_clock(base: int, size: int, failure: BaseException | None
if result_cursor < len(resolved_results):
_assert_next_step_matches(cmd)
recorded_value = resolved_results[result_cursor]
if (
isinstance(recorded_value, Mapping)
and recorded_value.get("target_namespace") != cmd.target_namespace
):
raise NonDeterministicReplayError(
current_call_sequence, "NexusServiceCall", ["SideEffectRecorded"],
detail="nexus_target_namespace_changed: recorded target differs from the workflow command",
)
result_cursor += 1
try:
next_value = cmd.recorded_result(recorded_value)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
{
"$schema": "https://raw.githubusercontent.com/durable-workflow/.github/main/regression-corpus/evidence-schema.json",
"fixture_schema": "durable-workflow.replay-regression/v1",
"id": "nexus-target-namespace-changed",
"protocol_version": "1.19",
"bindings": ["python"],
"workflow": {"type": "tests.replay.nexus-target", "input": [], "payload_codec": "avro"},
"history": [{"event_type": "SideEffectRecorded", "payload": {
"sequence": 1,
"payload_codec": "avro",
"result": "wwHioz3/VYAiNw4WDHNjaGVtYQpqZHVyYWJsZS13b3JrZmxvdy52Mi5zZGstcHl0aG9uLm5leHVzLW9wZXJhdGlvbi1yZXN1bHQOdmVyc2lvbgQCEGFjY2VwdGVkAgEec2VydmljZV9jYWxsX2lkCiBvcmlnaW5hbC1jYWxsLWlkGmVuZHBvaW50X25hbWUKDmdyZWV0ZXIYc2VydmljZV9uYW1lCgxzaGFyZWQcb3BlcmF0aW9uX25hbWUKCmdyZWV0HmlkZW1wb3RlbmN5X2tleQoab3JpZ2luYWwtY2FsbCB0YXJnZXRfbmFtZXNwYWNlCgxzaGFyZWQMc3RhdHVzChJjb21wbGV0ZWQMcmVzdWx0ChpzZXJ2aWNlLXZhbHVlAA=="
}}],
"expected": {"command_sequence": []},
"expected_replay_error": {
"type": "NonDeterministicReplayError",
"message_contains": "nexus_target_namespace_changed",
"workflow_sequence": 1
}
}
Loading
Loading