Skip to content
Open
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
2 changes: 1 addition & 1 deletion temporalio/nexus/_operation_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -814,10 +814,10 @@ def _apply_nexus_context_to_start_activity_request( # pyright: ignore[reportUnu
req.on_conflict_options.attach_completion_callbacks = True
req.on_conflict_options.attach_links = True

req.request_id = nexus_ctx.nexus_context.request_id
request_links = nexus_ctx._get_request_links()

if _in_nexus_backing_start_context():
req.request_id = nexus_ctx.nexus_context.request_id
callbacks = nexus_ctx._get_callbacks(
OperationToken(
type=OperationTokenType.ACTIVITY,
Expand Down
3 changes: 2 additions & 1 deletion tests/nexus/test_link_propagation.py
Original file line number Diff line number Diff line change
Expand Up @@ -497,7 +497,8 @@ async def test_activity_start_forwards_inbound_links() -> None:

assert len(req.links) == 1
assert req.links[0] == _inbound_nexus_link()
assert req.request_id == "req-id"
assert req.request_id
assert req.request_id != "req-id"
assert len(req.completion_callbacks) == 0


Expand Down
75 changes: 75 additions & 0 deletions tests/nexus/test_temporal_operation.py
Original file line number Diff line number Diff line change
Expand Up @@ -1139,6 +1139,81 @@ async def test_temporal_operation_start_activity(
assert result == "test"


async def test_temporal_operation_can_start_multiple_activities_with_same_id(
client: Client, env: WorkflowEnvironment
):
if env.supports_time_skipping:
pytest.skip(
"Standalone Nexus Operation tests don't work with time-skipping server"
)

task_queue = str(uuid.uuid4())
endpoint_name = make_nexus_endpoint_name(task_queue)
activity_id = str(uuid.uuid4())
await env.create_nexus_endpoint(endpoint_name, task_queue)

@service_handler
class MultipleActivitiesHandler:
def __init__(self) -> None:
self.first_activity_run_id: str | None = None
self.start_completed = asyncio.Event()

@nexus.temporal_operation
async def start_activities(
self,
_ctx: nexus.TemporalStartOperationContext,
client: nexus.TemporalNexusClient,
_input: None,
) -> nexus.TemporalOperationResult[str]:
try:
first_handle = await nexus.client().start_activity(
echo_activity,
Input(value="first", task_queue=task_queue),
id=activity_id,
task_queue=task_queue,
schedule_to_close_timeout=timedelta(seconds=5),
)
assert first_handle.run_id
self.first_activity_run_id = first_handle.run_id
await first_handle.result()

return await client.start_activity(
echo_activity,
Input(value="second", task_queue=task_queue),
id=activity_id,
schedule_to_close_timeout=timedelta(seconds=5),
)
finally:
self.start_completed.set()

handler = MultipleActivitiesHandler()
async with Worker(
env.client,
task_queue=task_queue,
nexus_service_handlers=[handler],
activities=[echo_activity],
):
nexus_client = client.create_nexus_client(
MultipleActivitiesHandler, endpoint_name
)
operation_handle = await nexus_client.start_operation(
MultipleActivitiesHandler.start_activities,
None,
id=str(uuid.uuid4()),
)
await asyncio.wait_for(handler.start_completed.wait(), timeout=15)

assert handler.first_activity_run_id
first_result = await client.get_activity_handle(
activity_id, run_id=handler.first_activity_run_id, result_type=str
).result()
second_result = await client.get_activity_handle(
activity_id, result_type=str
).result()
assert [first_result, second_result] == ["first", "second"]
assert await operation_handle.result() == second_result


async def test_temporal_operation_backing_activity_does_not_duplicate_links(
client: Client, env: WorkflowEnvironment
):
Expand Down
Loading