mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: workflow instance should not be invoked concurrently (#948)
* workflow instance should not be invoked concurrently * address comments
This commit is contained in:
committed by
GitHub
Unverified
parent
0d0231b512
commit
a9608465d9
@@ -227,12 +227,25 @@ class Workflow(AFBaseModel):
|
||||
workflow_id=id,
|
||||
)
|
||||
|
||||
# Flag to prevent concurrent workflow executions
|
||||
self._is_running = False
|
||||
|
||||
# Capture a canonical fingerprint of the workflow graph so checkpoints
|
||||
# can assert they are resumed with an equivalent topology.
|
||||
self._graph_signature = self._compute_graph_signature()
|
||||
self._graph_signature_hash = self._hash_graph_signature(self._graph_signature)
|
||||
self._runner.graph_signature_hash = self._graph_signature_hash
|
||||
|
||||
def _ensure_not_running(self) -> None:
|
||||
"""Ensure the workflow is not already running."""
|
||||
if self._is_running:
|
||||
raise RuntimeError("Workflow is already running. Concurrent executions are not allowed.")
|
||||
self._is_running = True
|
||||
|
||||
def _reset_running_flag(self) -> None:
|
||||
"""Reset the running flag."""
|
||||
self._is_running = False
|
||||
|
||||
def model_dump(self, **kwargs: Any) -> dict[str, Any]:
|
||||
"""Custom serialization that properly handles WorkflowExecutor nested workflows."""
|
||||
data = super().model_dump(**kwargs)
|
||||
@@ -370,20 +383,26 @@ class Workflow(AFBaseModel):
|
||||
Yields:
|
||||
WorkflowEvent: The events generated during the workflow execution.
|
||||
"""
|
||||
self._ensure_not_running()
|
||||
try:
|
||||
|
||||
async def initial_execution() -> None:
|
||||
executor = self.get_start_executor()
|
||||
await executor.execute(
|
||||
message,
|
||||
[self.__class__.__name__], # source_executor_ids
|
||||
self._shared_state, # shared_state
|
||||
self._runner.context, # runner_context
|
||||
trace_contexts=None, # No parent trace context for workflow start
|
||||
source_span_ids=None, # No source span for workflow start
|
||||
)
|
||||
async def initial_execution() -> None:
|
||||
executor = self.get_start_executor()
|
||||
await executor.execute(
|
||||
message,
|
||||
[self.__class__.__name__], # source_executor_ids
|
||||
self._shared_state, # shared_state
|
||||
self._runner.context, # runner_context
|
||||
trace_contexts=None, # No parent trace context for workflow start
|
||||
source_span_ids=None, # No source span for workflow start
|
||||
)
|
||||
|
||||
async for event in self._run_workflow_with_tracing(initial_executor_fn=initial_execution, reset_context=True):
|
||||
yield event
|
||||
async for event in self._run_workflow_with_tracing(
|
||||
initial_executor_fn=initial_execution, reset_context=True
|
||||
):
|
||||
yield event
|
||||
finally:
|
||||
self._reset_running_flag()
|
||||
|
||||
async def run_stream_from_checkpoint(
|
||||
self,
|
||||
@@ -407,60 +426,64 @@ class Workflow(AFBaseModel):
|
||||
ValueError: If neither checkpoint_storage is provided nor checkpointing is enabled.
|
||||
RuntimeError: If checkpoint restoration fails.
|
||||
"""
|
||||
self._ensure_not_running()
|
||||
try:
|
||||
|
||||
async def checkpoint_restoration() -> None:
|
||||
has_checkpointing = self._runner.context.has_checkpointing()
|
||||
async def checkpoint_restoration() -> None:
|
||||
has_checkpointing = self._runner.context.has_checkpointing()
|
||||
|
||||
if not has_checkpointing and checkpoint_storage is None:
|
||||
raise ValueError(
|
||||
"Cannot restore from checkpoint: either provide checkpoint_storage parameter "
|
||||
"or build workflow with WorkflowBuilder.with_checkpointing(checkpoint_storage)."
|
||||
)
|
||||
if not has_checkpointing and checkpoint_storage is None:
|
||||
raise ValueError(
|
||||
"Cannot restore from checkpoint: either provide checkpoint_storage parameter "
|
||||
"or build workflow with WorkflowBuilder.with_checkpointing(checkpoint_storage)."
|
||||
)
|
||||
|
||||
restored = await self._runner.restore_from_checkpoint(checkpoint_id, checkpoint_storage)
|
||||
restored = await self._runner.restore_from_checkpoint(checkpoint_id, checkpoint_storage)
|
||||
|
||||
if not restored:
|
||||
raise RuntimeError(f"Failed to restore from checkpoint: {checkpoint_id}")
|
||||
if not restored:
|
||||
raise RuntimeError(f"Failed to restore from checkpoint: {checkpoint_id}")
|
||||
|
||||
# Process any pending messages from the checkpoint first
|
||||
# This ensures that RequestInfoExecutor state is properly populated
|
||||
# before we try to handle responses
|
||||
if await self._runner.context.has_messages():
|
||||
# Run one iteration to process pending messages
|
||||
# This will populate RequestInfoExecutor._request_events properly
|
||||
await self._runner._run_iteration() # type: ignore
|
||||
# Process any pending messages from the checkpoint first
|
||||
# This ensures that RequestInfoExecutor state is properly populated
|
||||
# before we try to handle responses
|
||||
if await self._runner.context.has_messages():
|
||||
# Run one iteration to process pending messages
|
||||
# This will populate RequestInfoExecutor._request_events properly
|
||||
await self._runner._run_iteration() # type: ignore
|
||||
|
||||
if responses:
|
||||
request_info_executor = self._find_request_info_executor()
|
||||
if request_info_executor:
|
||||
for request_id, response_data in responses.items():
|
||||
ctx: WorkflowContext[Any] = WorkflowContext(
|
||||
request_info_executor.id,
|
||||
[self.__class__.__name__],
|
||||
self._shared_state,
|
||||
self._runner.context,
|
||||
trace_contexts=None, # No parent trace context for new workflow span
|
||||
source_span_ids=None, # No source span for response handling
|
||||
)
|
||||
|
||||
if not await request_info_executor.has_pending_request(request_id, ctx):
|
||||
logger.debug(
|
||||
f"Skipping pre-supplied response for request {request_id}; no pending request found "
|
||||
f"after checkpoint restoration."
|
||||
if responses:
|
||||
request_info_executor = self._find_request_info_executor()
|
||||
if request_info_executor:
|
||||
for request_id, response_data in responses.items():
|
||||
ctx: WorkflowContext[Any] = WorkflowContext(
|
||||
request_info_executor.id,
|
||||
[self.__class__.__name__],
|
||||
self._shared_state,
|
||||
self._runner.context,
|
||||
trace_contexts=None, # No parent trace context for new workflow span
|
||||
source_span_ids=None, # No source span for response handling
|
||||
)
|
||||
continue
|
||||
|
||||
await request_info_executor.handle_response(
|
||||
response_data,
|
||||
request_id,
|
||||
ctx,
|
||||
)
|
||||
if not await request_info_executor.has_pending_request(request_id, ctx):
|
||||
logger.debug(
|
||||
f"Skipping pre-supplied response for request {request_id}; "
|
||||
f"no pending request found after checkpoint restoration."
|
||||
)
|
||||
continue
|
||||
|
||||
async for event in self._run_workflow_with_tracing(
|
||||
initial_executor_fn=checkpoint_restoration,
|
||||
reset_context=False, # Don't reset context when resuming from checkpoint
|
||||
):
|
||||
yield event
|
||||
await request_info_executor.handle_response(
|
||||
response_data,
|
||||
request_id,
|
||||
ctx,
|
||||
)
|
||||
|
||||
async for event in self._run_workflow_with_tracing(
|
||||
initial_executor_fn=checkpoint_restoration,
|
||||
reset_context=False, # Don't reset context when resuming from checkpoint
|
||||
):
|
||||
yield event
|
||||
finally:
|
||||
self._reset_running_flag()
|
||||
|
||||
async def send_responses_streaming(self, responses: dict[str, Any]) -> AsyncIterable[WorkflowEvent]:
|
||||
"""Send responses back to the workflow and stream the events generated by the workflow.
|
||||
@@ -472,36 +495,40 @@ class Workflow(AFBaseModel):
|
||||
Yields:
|
||||
WorkflowEvent: The events generated during the workflow execution after sending the responses.
|
||||
"""
|
||||
self._ensure_not_running()
|
||||
try:
|
||||
|
||||
async def send_responses() -> None:
|
||||
request_info_executor = self._find_request_info_executor()
|
||||
if not request_info_executor:
|
||||
raise ValueError("No RequestInfoExecutor found in workflow.")
|
||||
async def send_responses() -> None:
|
||||
request_info_executor = self._find_request_info_executor()
|
||||
if not request_info_executor:
|
||||
raise ValueError("No RequestInfoExecutor found in workflow.")
|
||||
|
||||
async def _handle_response(response: Any, request_id: str) -> None:
|
||||
"""Handle the response from the RequestInfoExecutor."""
|
||||
await request_info_executor.handle_response(
|
||||
response,
|
||||
request_id,
|
||||
WorkflowContext(
|
||||
request_info_executor.id,
|
||||
[self.__class__.__name__],
|
||||
self._shared_state,
|
||||
self._runner.context,
|
||||
trace_contexts=None, # No parent trace context for new workflow span
|
||||
source_span_ids=None, # No source span for response handling
|
||||
),
|
||||
)
|
||||
async def _handle_response(response: Any, request_id: str) -> None:
|
||||
"""Handle the response from the RequestInfoExecutor."""
|
||||
await request_info_executor.handle_response(
|
||||
response,
|
||||
request_id,
|
||||
WorkflowContext(
|
||||
request_info_executor.id,
|
||||
[self.__class__.__name__],
|
||||
self._shared_state,
|
||||
self._runner.context,
|
||||
trace_contexts=None, # No parent trace context for new workflow span
|
||||
source_span_ids=None, # No source span for response handling
|
||||
),
|
||||
)
|
||||
|
||||
await asyncio.gather(*[
|
||||
_handle_response(response, request_id) for request_id, response in responses.items()
|
||||
])
|
||||
await asyncio.gather(*[
|
||||
_handle_response(response, request_id) for request_id, response in responses.items()
|
||||
])
|
||||
|
||||
async for event in self._run_workflow_with_tracing(
|
||||
initial_executor_fn=send_responses,
|
||||
reset_context=False, # Don't reset context when sending responses
|
||||
):
|
||||
yield event
|
||||
async for event in self._run_workflow_with_tracing(
|
||||
initial_executor_fn=send_responses,
|
||||
reset_context=False, # Don't reset context when sending responses
|
||||
):
|
||||
yield event
|
||||
finally:
|
||||
self._reset_running_flag()
|
||||
|
||||
async def run(self, message: Any, *, include_status_events: bool = False) -> WorkflowRunResult:
|
||||
"""Run the workflow with the given message.
|
||||
@@ -513,11 +540,31 @@ class Workflow(AFBaseModel):
|
||||
Returns:
|
||||
A WorkflowRunResult instance containing a list of events generated during the workflow execution.
|
||||
"""
|
||||
from agent_framework import AgentRunResponse, AgentRunResponseUpdate
|
||||
self._ensure_not_running()
|
||||
try:
|
||||
from agent_framework import AgentRunResponse, AgentRunResponseUpdate
|
||||
|
||||
from ._events import AgentRunEvent, AgentRunUpdateEvent # Local import to avoid cycles
|
||||
from ._events import AgentRunEvent, AgentRunUpdateEvent # Local import to avoid cycles
|
||||
|
||||
raw_events = [event async for event in self.run_stream(message)]
|
||||
async def initial_execution() -> None:
|
||||
executor = self.get_start_executor()
|
||||
await executor.execute(
|
||||
message,
|
||||
[self.__class__.__name__], # source_executor_ids
|
||||
self._shared_state, # shared_state
|
||||
self._runner.context, # runner_context
|
||||
trace_contexts=None, # No parent trace context for workflow start
|
||||
source_span_ids=None, # No source span for workflow start
|
||||
)
|
||||
|
||||
raw_events = [
|
||||
event
|
||||
async for event in self._run_workflow_with_tracing(
|
||||
initial_executor_fn=initial_execution, reset_context=True
|
||||
)
|
||||
]
|
||||
finally:
|
||||
self._reset_running_flag()
|
||||
|
||||
# Coalesce streaming update events into a single AgentRunEvent per executor sequence.
|
||||
coalesced: list[WorkflowEvent] = [] # type: ignore[name-defined]
|
||||
@@ -589,12 +636,69 @@ class Workflow(AFBaseModel):
|
||||
ValueError: If neither checkpoint_storage is provided nor checkpointing is enabled.
|
||||
RuntimeError: If checkpoint restoration fails.
|
||||
"""
|
||||
events = [
|
||||
event async for event in self.run_stream_from_checkpoint(checkpoint_id, checkpoint_storage, responses)
|
||||
]
|
||||
status_events = [e for e in events if isinstance(e, WorkflowStatusEvent)]
|
||||
filtered_events = [e for e in events if not isinstance(e, (WorkflowStatusEvent, WorkflowStartedEvent))]
|
||||
return WorkflowRunResult(filtered_events, status_events)
|
||||
self._ensure_not_running()
|
||||
try:
|
||||
|
||||
async def checkpoint_restoration() -> None:
|
||||
has_checkpointing = self._runner.context.has_checkpointing()
|
||||
|
||||
if not has_checkpointing and checkpoint_storage is None:
|
||||
raise ValueError(
|
||||
"Cannot restore from checkpoint: either provide checkpoint_storage parameter "
|
||||
"or build workflow with WorkflowBuilder.with_checkpointing(checkpoint_storage)."
|
||||
)
|
||||
|
||||
restored = await self._runner.restore_from_checkpoint(checkpoint_id, checkpoint_storage)
|
||||
|
||||
if not restored:
|
||||
raise RuntimeError(f"Failed to restore from checkpoint: {checkpoint_id}")
|
||||
|
||||
# Process any pending messages from the checkpoint first
|
||||
# This ensures that RequestInfoExecutor state is properly populated
|
||||
# before we try to handle responses
|
||||
if await self._runner.context.has_messages():
|
||||
# Run one iteration to process pending messages
|
||||
# This will populate RequestInfoExecutor._request_events properly
|
||||
await self._runner._run_iteration() # type: ignore
|
||||
|
||||
if responses:
|
||||
request_info_executor = self._find_request_info_executor()
|
||||
if request_info_executor:
|
||||
for request_id, response_data in responses.items():
|
||||
ctx: WorkflowContext[Any] = WorkflowContext(
|
||||
request_info_executor.id,
|
||||
[self.__class__.__name__],
|
||||
self._shared_state,
|
||||
self._runner.context,
|
||||
trace_contexts=None, # No parent trace context for new workflow span
|
||||
source_span_ids=None, # No source span for response handling
|
||||
)
|
||||
|
||||
if not await request_info_executor.has_pending_request(request_id, ctx):
|
||||
logger.debug(
|
||||
f"Skipping pre-supplied response for request {request_id}; "
|
||||
f"no pending request found after checkpoint restoration."
|
||||
)
|
||||
continue
|
||||
|
||||
await request_info_executor.handle_response(
|
||||
response_data,
|
||||
request_id,
|
||||
ctx,
|
||||
)
|
||||
|
||||
events = [
|
||||
event
|
||||
async for event in self._run_workflow_with_tracing(
|
||||
initial_executor_fn=checkpoint_restoration,
|
||||
reset_context=False, # Don't reset context when resuming from checkpoint
|
||||
)
|
||||
]
|
||||
status_events = [e for e in events if isinstance(e, WorkflowStatusEvent)]
|
||||
filtered_events = [e for e in events if not isinstance(e, (WorkflowStatusEvent, WorkflowStartedEvent))]
|
||||
return WorkflowRunResult(filtered_events, status_events)
|
||||
finally:
|
||||
self._reset_running_flag()
|
||||
|
||||
async def send_responses(self, responses: dict[str, Any]) -> WorkflowRunResult:
|
||||
"""Send responses back to the workflow.
|
||||
@@ -605,10 +709,45 @@ class Workflow(AFBaseModel):
|
||||
Returns:
|
||||
A WorkflowRunResult instance containing a list of events generated during the workflow execution.
|
||||
"""
|
||||
events = [event async for event in self.send_responses_streaming(responses)]
|
||||
status_events = [e for e in events if isinstance(e, WorkflowStatusEvent)]
|
||||
filtered_events = [e for e in events if not isinstance(e, (WorkflowStatusEvent, WorkflowStartedEvent))]
|
||||
return WorkflowRunResult(filtered_events, status_events)
|
||||
self._ensure_not_running()
|
||||
try:
|
||||
|
||||
async def send_responses_internal() -> None:
|
||||
request_info_executor = self._find_request_info_executor()
|
||||
if not request_info_executor:
|
||||
raise ValueError("No RequestInfoExecutor found in workflow.")
|
||||
|
||||
async def _handle_response(response: Any, request_id: str) -> None:
|
||||
"""Handle the response from the RequestInfoExecutor."""
|
||||
await request_info_executor.handle_response(
|
||||
response,
|
||||
request_id,
|
||||
WorkflowContext(
|
||||
request_info_executor.id,
|
||||
[self.__class__.__name__],
|
||||
self._shared_state,
|
||||
self._runner.context,
|
||||
trace_contexts=None, # No parent trace context for new workflow span
|
||||
source_span_ids=None, # No source span for response handling
|
||||
),
|
||||
)
|
||||
|
||||
await asyncio.gather(*[
|
||||
_handle_response(response, request_id) for request_id, response in responses.items()
|
||||
])
|
||||
|
||||
events = [
|
||||
event
|
||||
async for event in self._run_workflow_with_tracing(
|
||||
initial_executor_fn=send_responses_internal,
|
||||
reset_context=False, # Don't reset context when sending responses
|
||||
)
|
||||
]
|
||||
status_events = [e for e in events if isinstance(e, WorkflowStatusEvent)]
|
||||
filtered_events = [e for e in events if not isinstance(e, (WorkflowStatusEvent, WorkflowStartedEvent))]
|
||||
return WorkflowRunResult(filtered_events, status_events)
|
||||
finally:
|
||||
self._reset_running_flag()
|
||||
|
||||
def _get_executor_by_id(self, executor_id: str) -> Executor:
|
||||
"""Get an executor by its ID.
|
||||
|
||||
Reference in New Issue
Block a user