Skip to content

Commit 462195d

Browse files
authored
fix(task-create): handle WorkflowAlreadyStartedError gracefully (#489)
1 parent f2b1808 commit 462195d

4 files changed

Lines changed: 174 additions & 2 deletions

File tree

‎src/agentex/lib/core/clients/temporal/temporal_client.py‎

Lines changed: 23 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,11 @@
55
from collections.abc import Callable
66

77
from temporalio.client import Client, WorkflowExecutionStatus
8-
from temporalio.common import RetryPolicy as TemporalRetryPolicy, WorkflowIDReusePolicy
8+
from temporalio.common import (
9+
RetryPolicy as TemporalRetryPolicy,
10+
WorkflowIDReusePolicy,
11+
WorkflowIDConflictPolicy,
12+
)
913
from temporalio.service import RPCError, RPCStatusCode
1014
from temporalio.converter import PayloadCodec, DataConverter
1115

@@ -15,6 +19,7 @@
1519
TaskStatus,
1620
RetryPolicy,
1721
WorkflowState,
22+
ConflictWorkflowPolicy,
1823
DuplicateWorkflowPolicy,
1924
)
2025
from agentex.lib.core.clients.temporal.utils import get_temporal_client
@@ -75,6 +80,13 @@
7580
DuplicateWorkflowPolicy.TERMINATE_IF_RUNNING: WorkflowIDReusePolicy.TERMINATE_IF_RUNNING,
7681
}
7782

83+
CONFLICT_POLICY_TO_ID_CONFLICT_POLICY = {
84+
ConflictWorkflowPolicy.UNSPECIFIED: WorkflowIDConflictPolicy.UNSPECIFIED,
85+
ConflictWorkflowPolicy.FAIL: WorkflowIDConflictPolicy.FAIL,
86+
ConflictWorkflowPolicy.USE_EXISTING: WorkflowIDConflictPolicy.USE_EXISTING,
87+
ConflictWorkflowPolicy.TERMINATE_EXISTING: WorkflowIDConflictPolicy.TERMINATE_EXISTING,
88+
}
89+
7890

7991
class TemporalClient:
8092
def __init__(
@@ -151,18 +163,28 @@ async def start_workflow(
151163
self,
152164
*args: Any,
153165
duplicate_policy: DuplicateWorkflowPolicy = DuplicateWorkflowPolicy.ALLOW_DUPLICATE,
166+
conflict_policy: ConflictWorkflowPolicy = ConflictWorkflowPolicy.UNSPECIFIED,
154167
retry_policy: RetryPolicy = DEFAULT_RETRY_POLICY,
155168
task_timeout: timedelta = timedelta(seconds=10),
156169
execution_timeout: timedelta | None = None,
157170
**kwargs: Any,
158171
) -> str:
172+
if (
173+
duplicate_policy == DuplicateWorkflowPolicy.TERMINATE_IF_RUNNING
174+
and conflict_policy != ConflictWorkflowPolicy.UNSPECIFIED
175+
):
176+
raise ValueError(
177+
"conflict_policy cannot be set when duplicate_policy is TERMINATE_IF_RUNNING; "
178+
"use ConflictWorkflowPolicy.TERMINATE_EXISTING instead"
179+
)
159180
temporal_retry_policy = TemporalRetryPolicy(**retry_policy.model_dump(exclude_unset=True))
160181
workflow_handle = await self.client.start_workflow(
161182
*args,
162183
retry_policy=temporal_retry_policy,
163184
task_timeout=task_timeout,
164185
execution_timeout=execution_timeout,
165186
id_reuse_policy=DUPLICATE_POLICY_TO_ID_REUSE_POLICY[duplicate_policy],
187+
id_conflict_policy=CONFLICT_POLICY_TO_ID_CONFLICT_POLICY[conflict_policy],
166188
**kwargs,
167189
)
168190
return workflow_handle.id

‎src/agentex/lib/core/clients/temporal/types.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,13 @@ class DuplicateWorkflowPolicy(str, Enum):
4040
TERMINATE_IF_RUNNING = "TERMINATE_IF_RUNNING"
4141

4242

43+
class ConflictWorkflowPolicy(str, Enum):
44+
UNSPECIFIED = "UNSPECIFIED"
45+
FAIL = "FAIL"
46+
USE_EXISTING = "USE_EXISTING"
47+
TERMINATE_EXISTING = "TERMINATE_EXISTING"
48+
49+
4350
class TaskStatus(str, Enum):
4451
CANCELED = "CANCELED"
4552
COMPLETED = "COMPLETED"

‎src/agentex/lib/core/temporal/services/temporal_task_service.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
from agentex.types.event import Event
1212
from agentex.protocol.acp import SendEventParams, CreateTaskParams, InterruptTaskParams
1313
from agentex.lib.environment_variables import EnvironmentVariables
14-
from agentex.lib.core.clients.temporal.types import WorkflowState
14+
from agentex.lib.core.clients.temporal.types import WorkflowState, ConflictWorkflowPolicy
1515
from agentex.lib.core.temporal.types.workflow import SignalName
1616
from agentex.lib.core.clients.temporal.temporal_client import TemporalClient
1717

@@ -89,6 +89,8 @@ async def submit_task(self, agent: Agent, task: Task, params: dict[str, Any] | N
8989
# value bounds the whole continue-as-new chain's wall-clock lifetime.
9090
timeout_seconds = self._env_vars.WORKFLOW_EXECUTION_TIMEOUT_SECONDS
9191
execution_timeout = timedelta(seconds=timeout_seconds) if timeout_seconds and timeout_seconds > 0 else None
92+
# USE_EXISTING makes task/create idempotent
93+
# If same task ID is already running Temporal returns a handle to the existing run instead of raising WorkflowAlreadyStarted
9294
with _acp_dispatch_span("acp.task_create", task_id=task.id):
9395
return await self._temporal_client.start_workflow(
9496
workflow=self._env_vars.WORKFLOW_NAME,
@@ -100,6 +102,7 @@ async def submit_task(self, agent: Agent, task: Task, params: dict[str, Any] | N
100102
id=task.id,
101103
task_queue=self._env_vars.WORKFLOW_TASK_QUEUE,
102104
execution_timeout=execution_timeout,
105+
conflict_policy=ConflictWorkflowPolicy.USE_EXISTING,
103106
)
104107

105108
async def get_state(self, task_id: str) -> WorkflowState:
Lines changed: 140 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,140 @@
1+
"""Unit tests for TemporalTaskService idempotency behavior.
2+
3+
Covers the ``task/create`` idempotency guarantee: duplicate submits for the
4+
same task ID must not raise ``WorkflowAlreadyStartedError``. The service
5+
achieves this by passing ``ConflictWorkflowPolicy.USE_EXISTING`` through the
6+
``TemporalClient`` wrapper, which maps to ``WorkflowIDConflictPolicy.USE_EXISTING``
7+
on the underlying temporalio client so Temporal returns a handle to the
8+
existing run instead of erroring.
9+
"""
10+
11+
from __future__ import annotations
12+
13+
from unittest.mock import Mock, AsyncMock
14+
15+
import pytest
16+
from temporalio.common import WorkflowIDConflictPolicy
17+
18+
from agentex.types.task import Task
19+
from agentex.types.agent import Agent
20+
from agentex.lib.core.clients.temporal.types import (
21+
ConflictWorkflowPolicy,
22+
DuplicateWorkflowPolicy,
23+
)
24+
from agentex.lib.core.clients.temporal.temporal_client import TemporalClient
25+
from agentex.lib.core.temporal.services.temporal_task_service import TemporalTaskService
26+
27+
28+
def _agent() -> Agent:
29+
return Agent(
30+
id="test-agent-456",
31+
name="test-agent",
32+
description="test-agent",
33+
acp_type="async",
34+
created_at="2023-01-01T00:00:00Z",
35+
updated_at="2023-01-01T00:00:00Z",
36+
)
37+
38+
39+
def _task() -> Task:
40+
return Task(id="test-task-123", status="RUNNING")
41+
42+
43+
def _env_vars() -> Mock:
44+
env_vars = Mock()
45+
env_vars.WORKFLOW_NAME = "test-workflow"
46+
env_vars.WORKFLOW_TASK_QUEUE = "test-queue"
47+
env_vars.WORKFLOW_EXECUTION_TIMEOUT_SECONDS = 0
48+
return env_vars
49+
50+
51+
class TestSubmitTaskIdempotency:
52+
async def test_submit_task_uses_use_existing_conflict_policy(self) -> None:
53+
"""Duplicate task/create must be idempotent.
54+
55+
Passing ``ConflictWorkflowPolicy.USE_EXISTING`` tells Temporal to
56+
return the existing workflow handle instead of raising
57+
``WorkflowAlreadyStartedError`` when a run with that ID is already
58+
active. Without this, load-balanced agentex-agent replicas racing on
59+
the same task ID surface Temporal's start conflict as an error log.
60+
"""
61+
temporal_client = Mock()
62+
temporal_client.start_workflow = AsyncMock(return_value="test-task-123")
63+
64+
service = TemporalTaskService(temporal_client=temporal_client, env_vars=_env_vars())
65+
66+
result = await service.submit_task(agent=_agent(), task=_task(), params=None)
67+
68+
temporal_client.start_workflow.assert_awaited_once()
69+
kwargs = temporal_client.start_workflow.await_args.kwargs
70+
assert kwargs["conflict_policy"] == ConflictWorkflowPolicy.USE_EXISTING
71+
assert kwargs["id"] == "test-task-123"
72+
assert result == "test-task-123"
73+
74+
75+
class TestTemporalClientConflictPolicyPlumbing:
76+
"""Boundary tests: ``TemporalClient.start_workflow`` wraps
77+
``WorkflowIDConflictPolicy`` in a local ``ConflictWorkflowPolicy`` enum
78+
(mirroring the existing ``DuplicateWorkflowPolicy`` pattern) so SDK users
79+
don't have to import ``temporalio.common`` to opt into non-default behavior.
80+
Also guards the incompatible ``TERMINATE_IF_RUNNING`` + explicit-conflict
81+
combo client-side rather than letting it round-trip to the frontend as
82+
``InvalidArgument``.
83+
"""
84+
85+
async def test_forwards_conflict_policy_when_set(self) -> None:
86+
inner_client = Mock()
87+
inner_handle = Mock()
88+
inner_handle.id = "wf-1"
89+
inner_client.start_workflow = AsyncMock(return_value=inner_handle)
90+
91+
tc = TemporalClient(temporal_client=inner_client)
92+
93+
await tc.start_workflow(
94+
workflow="w",
95+
arg={},
96+
id="id-1",
97+
task_queue="q",
98+
conflict_policy=ConflictWorkflowPolicy.USE_EXISTING,
99+
)
100+
101+
kwargs = inner_client.start_workflow.await_args.kwargs
102+
assert kwargs["id_conflict_policy"] == WorkflowIDConflictPolicy.USE_EXISTING
103+
104+
async def test_default_conflict_policy_is_unspecified(self) -> None:
105+
inner_client = Mock()
106+
inner_handle = Mock()
107+
inner_handle.id = "wf-1"
108+
inner_client.start_workflow = AsyncMock(return_value=inner_handle)
109+
110+
tc = TemporalClient(temporal_client=inner_client)
111+
112+
await tc.start_workflow(workflow="w", arg={}, id="id-1", task_queue="q")
113+
114+
kwargs = inner_client.start_workflow.await_args.kwargs
115+
assert kwargs["id_conflict_policy"] == WorkflowIDConflictPolicy.UNSPECIFIED
116+
117+
async def test_terminate_if_running_with_explicit_conflict_policy_raises(self) -> None:
118+
"""temporalio rejects this combo at the frontend as InvalidArgument;
119+
we fail fast client-side with a clearer message.
120+
"""
121+
inner_client = Mock()
122+
inner_client.start_workflow = AsyncMock()
123+
124+
tc = TemporalClient(temporal_client=inner_client)
125+
126+
with pytest.raises(ValueError, match="TERMINATE_EXISTING"):
127+
await tc.start_workflow(
128+
workflow="w",
129+
arg={},
130+
id="id-1",
131+
task_queue="q",
132+
duplicate_policy=DuplicateWorkflowPolicy.TERMINATE_IF_RUNNING,
133+
conflict_policy=ConflictWorkflowPolicy.USE_EXISTING,
134+
)
135+
136+
inner_client.start_workflow.assert_not_awaited()
137+
138+
139+
if __name__ == "__main__": # pragma: no cover
140+
raise SystemExit(pytest.main([__file__, "-v"]))

0 commit comments

Comments
 (0)