mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: Add Durabletask samples and minor fixes (#3157)
* Add samples and minor fixes * Add redis sample and wait-for-completion * Add wait-for-completion support * ADd missing docs
This commit is contained in:
committed by
GitHub
Unverified
parent
1e36ba33c4
commit
3df916064c
@@ -142,7 +142,7 @@ class AgentEntity:
|
||||
response_format = run_request.response_format
|
||||
enable_tool_calls = run_request.enable_tool_calls
|
||||
|
||||
logger.debug("[AgentEntity.run] Received Message: %s", run_request)
|
||||
logger.debug("[AgentEntity.run] Received ThreadId %s Message: %s", thread_id, run_request)
|
||||
|
||||
state_request = DurableAgentStateRequest.from_run_request(run_request)
|
||||
self.state.data.conversation_history.append(state_request)
|
||||
|
||||
@@ -16,10 +16,10 @@ from abc import ABC, abstractmethod
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Generic, TypeVar
|
||||
|
||||
from agent_framework import AgentRunResponse, AgentThread, ChatMessage, ErrorContent, Role, get_logger
|
||||
from agent_framework import AgentRunResponse, AgentThread, ChatMessage, ErrorContent, Role, TextContent, get_logger
|
||||
from durabletask.client import TaskHubGrpcClient
|
||||
from durabletask.entities import EntityInstanceId
|
||||
from durabletask.task import CompositeTask, OrchestrationContext, Task
|
||||
from durabletask.task import CompletableTask, CompositeTask, OrchestrationContext, Task
|
||||
from pydantic import BaseModel
|
||||
|
||||
from ._constants import DEFAULT_MAX_POLL_RETRIES, DEFAULT_POLL_INTERVAL_SECONDS
|
||||
@@ -33,16 +33,19 @@ logger = get_logger("agent_framework.durabletask.executors")
|
||||
TaskT = TypeVar("TaskT")
|
||||
|
||||
|
||||
class DurableAgentTask(CompositeTask[AgentRunResponse]):
|
||||
class DurableAgentTask(CompositeTask[AgentRunResponse], CompletableTask[AgentRunResponse]):
|
||||
"""A custom Task that wraps entity calls and provides typed AgentRunResponse results.
|
||||
|
||||
This task wraps the underlying entity call task and intercepts its completion
|
||||
to convert the raw result into a typed AgentRunResponse object.
|
||||
|
||||
When yielded in an orchestration, this task returns an AgentRunResponse:
|
||||
response: AgentRunResponse = yield durable_agent_task
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
entity_task: Task[Any],
|
||||
entity_task: CompletableTask[Any],
|
||||
response_format: type[BaseModel] | None,
|
||||
correlation_id: str,
|
||||
):
|
||||
@@ -55,7 +58,7 @@ class DurableAgentTask(CompositeTask[AgentRunResponse]):
|
||||
"""
|
||||
self._response_format = response_format
|
||||
self._correlation_id = correlation_id
|
||||
super().__init__([entity_task]) # type: ignore[misc]
|
||||
super().__init__([entity_task]) # type: ignore
|
||||
|
||||
def on_child_completed(self, task: Task[Any]) -> None:
|
||||
"""Handle completion of the underlying entity task.
|
||||
@@ -69,11 +72,8 @@ class DurableAgentTask(CompositeTask[AgentRunResponse]):
|
||||
return
|
||||
|
||||
if task.is_failed:
|
||||
# Propagate the failure
|
||||
self._exception = task.get_exception()
|
||||
self._is_complete = True
|
||||
if self._parent is not None:
|
||||
self._parent.on_child_completed(self)
|
||||
# Propagate the failure - pass the original exception directly
|
||||
self.fail("call_entity Task failed", task.get_exception())
|
||||
return
|
||||
|
||||
# Task succeeded - transform the raw result
|
||||
@@ -94,18 +94,12 @@ class DurableAgentTask(CompositeTask[AgentRunResponse]):
|
||||
)
|
||||
|
||||
# Set the typed AgentRunResponse as this task's result
|
||||
self._result = response
|
||||
self._is_complete = True
|
||||
self.complete(response)
|
||||
|
||||
if self._parent is not None:
|
||||
self._parent.on_child_completed(self)
|
||||
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"[DurableAgentTask] Failed to convert result for correlation_id: %s",
|
||||
self._correlation_id,
|
||||
)
|
||||
raise
|
||||
except Exception as ex:
|
||||
err_msg = "[DurableAgentTask] Failed to convert result for correlation_id: " + self._correlation_id
|
||||
logger.exception(err_msg)
|
||||
self.fail(err_msg, ex)
|
||||
|
||||
|
||||
class DurableAgentExecutor(ABC, Generic[TaskT]):
|
||||
@@ -155,6 +149,7 @@ class DurableAgentExecutor(ABC, Generic[TaskT]):
|
||||
message: str,
|
||||
response_format: type[BaseModel] | None,
|
||||
enable_tool_calls: bool,
|
||||
wait_for_response: bool = True,
|
||||
) -> RunRequest:
|
||||
"""Create a RunRequest for the given parameters."""
|
||||
correlation_id = self.generate_unique_id()
|
||||
@@ -162,9 +157,34 @@ class DurableAgentExecutor(ABC, Generic[TaskT]):
|
||||
message=message,
|
||||
response_format=response_format,
|
||||
enable_tool_calls=enable_tool_calls,
|
||||
wait_for_response=wait_for_response,
|
||||
correlation_id=correlation_id,
|
||||
)
|
||||
|
||||
def _create_acceptance_response(self, correlation_id: str) -> AgentRunResponse:
|
||||
"""Create an acceptance response for fire-and-forget mode.
|
||||
|
||||
Args:
|
||||
correlation_id: Correlation ID for tracking the request
|
||||
|
||||
Returns:
|
||||
AgentRunResponse: Acceptance response with correlation ID
|
||||
"""
|
||||
acceptance_message = ChatMessage(
|
||||
role=Role.SYSTEM,
|
||||
contents=[
|
||||
TextContent(
|
||||
f"Request accepted for processing (correlation_id: {correlation_id}). "
|
||||
f"Agent is executing in the background. "
|
||||
f"Retrieve response via your configured streaming or callback mechanism."
|
||||
)
|
||||
],
|
||||
)
|
||||
return AgentRunResponse(
|
||||
messages=[acceptance_message],
|
||||
created_at=datetime.now(timezone.utc).isoformat(),
|
||||
)
|
||||
|
||||
|
||||
class ClientAgentExecutor(DurableAgentExecutor[AgentRunResponse]):
|
||||
"""Execution strategy for external clients.
|
||||
@@ -205,11 +225,20 @@ class ClientAgentExecutor(DurableAgentExecutor[AgentRunResponse]):
|
||||
thread: Optional conversation thread (creates new if not provided)
|
||||
|
||||
Returns:
|
||||
AgentRunResponse: The agent's response after execution completes
|
||||
AgentRunResponse: The agent's response after execution completes, or an immediate
|
||||
acknowledgement if wait_for_response is False
|
||||
"""
|
||||
# Signal the entity with the request
|
||||
entity_id = self._signal_agent_entity(agent_name, run_request, thread)
|
||||
|
||||
# If fire-and-forget mode, return immediately without polling
|
||||
if not run_request.wait_for_response:
|
||||
logger.info(
|
||||
"[ClientAgentExecutor] Fire-and-forget mode: request signaled (correlation: %s)",
|
||||
run_request.correlation_id,
|
||||
)
|
||||
return self._create_acceptance_response(run_request.correlation_id)
|
||||
|
||||
# Poll for the response
|
||||
agent_response = self._poll_for_agent_response(entity_id, run_request.correlation_id)
|
||||
|
||||
@@ -395,11 +424,16 @@ class OrchestrationAgentExecutor(DurableAgentExecutor[DurableAgentTask]):
|
||||
self._context = context
|
||||
logger.debug("[OrchestrationAgentExecutor] Initialized")
|
||||
|
||||
def generate_unique_id(self) -> str:
|
||||
"""Create a new UUID that is safe for replay within an orchestration or operation."""
|
||||
return self._context.new_uuid()
|
||||
|
||||
def get_run_request(
|
||||
self,
|
||||
message: str,
|
||||
response_format: type[BaseModel] | None,
|
||||
enable_tool_calls: bool,
|
||||
wait_for_response: bool = True,
|
||||
) -> RunRequest:
|
||||
"""Get the current run request from the orchestration context.
|
||||
|
||||
@@ -410,6 +444,7 @@ class OrchestrationAgentExecutor(DurableAgentExecutor[DurableAgentTask]):
|
||||
message,
|
||||
response_format,
|
||||
enable_tool_calls,
|
||||
wait_for_response,
|
||||
)
|
||||
request.orchestration_id = self._context.instance_id
|
||||
return request
|
||||
@@ -449,8 +484,22 @@ class OrchestrationAgentExecutor(DurableAgentExecutor[DurableAgentTask]):
|
||||
session_id,
|
||||
)
|
||||
|
||||
# Call the entity and get the underlying task
|
||||
entity_task: Task[Any] = self._context.call_entity(entity_id, "run", run_request.to_dict()) # type: ignore
|
||||
# Branch based on wait_for_response
|
||||
if not run_request.wait_for_response:
|
||||
# Fire-and-forget mode: signal entity and return pre-completed task
|
||||
logger.info(
|
||||
"[OrchestrationAgentExecutor] Fire-and-forget mode: signaling entity (correlation: %s)",
|
||||
run_request.correlation_id,
|
||||
)
|
||||
self._context.signal_entity(entity_id, "run", run_request.to_dict())
|
||||
|
||||
# Create a pre-completed task with acceptance response
|
||||
acceptance_response = self._create_acceptance_response(run_request.correlation_id)
|
||||
entity_task: CompletableTask[AgentRunResponse] = CompletableTask()
|
||||
entity_task.complete(acceptance_response)
|
||||
else:
|
||||
# Blocking mode: call entity and wait for response
|
||||
entity_task = self._context.call_entity(entity_id, "run", run_request.to_dict()) # type: ignore
|
||||
|
||||
# Wrap in DurableAgentTask for response transformation
|
||||
return DurableAgentTask(
|
||||
|
||||
@@ -104,6 +104,8 @@ class RunRequest:
|
||||
role: The role of the message sender (user, system, or assistant)
|
||||
response_format: Optional Pydantic BaseModel type describing the structured response format
|
||||
enable_tool_calls: Whether to enable tool calls for this request
|
||||
wait_for_response: If True (default), caller will wait for agent response. If False,
|
||||
returns immediately after signaling (fire-and-forget mode)
|
||||
correlation_id: Correlation ID for tracking the response to this specific request
|
||||
created_at: Optional timestamp when the request was created
|
||||
orchestration_id: Optional ID of the orchestration that initiated this request
|
||||
@@ -115,6 +117,7 @@ class RunRequest:
|
||||
role: Role = Role.USER
|
||||
response_format: type[BaseModel] | None = None
|
||||
enable_tool_calls: bool = True
|
||||
wait_for_response: bool = True
|
||||
created_at: datetime | None = None
|
||||
orchestration_id: str | None = None
|
||||
|
||||
@@ -126,6 +129,7 @@ class RunRequest:
|
||||
role: Role | str | None = Role.USER,
|
||||
response_format: type[BaseModel] | None = None,
|
||||
enable_tool_calls: bool = True,
|
||||
wait_for_response: bool = True,
|
||||
created_at: datetime | None = None,
|
||||
orchestration_id: str | None = None,
|
||||
) -> None:
|
||||
@@ -135,6 +139,7 @@ class RunRequest:
|
||||
self.response_format = response_format
|
||||
self.request_response_format = request_response_format
|
||||
self.enable_tool_calls = enable_tool_calls
|
||||
self.wait_for_response = wait_for_response
|
||||
self.created_at = created_at if created_at is not None else datetime.now(tz=timezone.utc)
|
||||
self.orchestration_id = orchestration_id
|
||||
|
||||
@@ -155,6 +160,7 @@ class RunRequest:
|
||||
result = {
|
||||
"message": self.message,
|
||||
"enable_tool_calls": self.enable_tool_calls,
|
||||
"wait_for_response": self.wait_for_response,
|
||||
"role": self.role.value,
|
||||
"request_response_format": self.request_response_format,
|
||||
"correlationId": self.correlation_id,
|
||||
@@ -198,6 +204,7 @@ class RunRequest:
|
||||
request_response_format=data.get("request_response_format", REQUEST_RESPONSE_FORMAT_TEXT),
|
||||
role=cls.coerce_role(data.get("role")),
|
||||
response_format=_deserialize_response_format(data.get("response_format")),
|
||||
wait_for_response=data.get("wait_for_response", True),
|
||||
enable_tool_calls=data.get("enable_tool_calls", True),
|
||||
created_at=created_at,
|
||||
orchestration_id=data.get("orchestrationId"),
|
||||
|
||||
@@ -51,10 +51,20 @@ def ensure_response_format(
|
||||
response_format: Optional Pydantic model class to parse the response value into
|
||||
correlation_id: Correlation ID for logging purposes
|
||||
response: The AgentRunResponse object to validate and parse
|
||||
|
||||
Raises:
|
||||
ValueError: If response_format is specified but response.value cannot be parsed
|
||||
"""
|
||||
if response_format is not None and not isinstance(response.value, response_format):
|
||||
response.try_parse_value(response_format)
|
||||
|
||||
# Validate that parsing succeeded
|
||||
if not isinstance(response.value, response_format):
|
||||
raise ValueError(
|
||||
f"Response value could not be parsed into required format {response_format.__name__} "
|
||||
f"for correlation_id {correlation_id}"
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
"[ensure_response_format] Loaded AgentRunResponse.value for correlation_id %s with type: %s",
|
||||
correlation_id,
|
||||
|
||||
@@ -108,9 +108,20 @@ class DurableAIAgent(AgentProtocol, Generic[TaskT]):
|
||||
thread: AgentThread | None = None,
|
||||
response_format: type[BaseModel] | None = None,
|
||||
enable_tool_calls: bool = True,
|
||||
wait_for_response: bool = True,
|
||||
) -> TaskT:
|
||||
"""Execute the agent via the injected provider.
|
||||
|
||||
Args:
|
||||
messages: The message(s) to send to the agent
|
||||
thread: Optional agent thread for conversation context
|
||||
response_format: Optional Pydantic model for structured response
|
||||
enable_tool_calls: Whether to enable tool calls for this request
|
||||
wait_for_response: If True (default), waits for agent response.
|
||||
If False, returns immediately (fire-and-forget mode).
|
||||
|
||||
**Only supported for DurableAIAgentClient contexts.**
|
||||
|
||||
Note:
|
||||
This method overrides AgentProtocol.run() with a different return type:
|
||||
- AgentProtocol.run() returns Coroutine[Any, Any, AgentRunResponse] (async)
|
||||
@@ -121,6 +132,9 @@ class DurableAIAgent(AgentProtocol, Generic[TaskT]):
|
||||
|
||||
Returns:
|
||||
TaskT: The task type specific to the executor
|
||||
|
||||
Raises:
|
||||
ValueError: If wait_for_response=False is used in an unsupported context
|
||||
"""
|
||||
message_str = self._normalize_messages(messages)
|
||||
|
||||
@@ -128,6 +142,7 @@ class DurableAIAgent(AgentProtocol, Generic[TaskT]):
|
||||
message=message_str,
|
||||
response_format=response_format,
|
||||
enable_tool_calls=enable_tool_calls,
|
||||
wait_for_response=wait_for_response,
|
||||
)
|
||||
|
||||
return self._executor.run_durable_agent(
|
||||
|
||||
@@ -23,8 +23,8 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core",
|
||||
"durabletask>=1.1.0",
|
||||
"durabletask-azuremanaged>=1.1.0"
|
||||
"durabletask>=1.3.0",
|
||||
"durabletask-azuremanaged>=1.3.0"
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
|
||||
@@ -39,7 +39,10 @@ def mock_client() -> Mock:
|
||||
@pytest.fixture
|
||||
def mock_entity_task() -> Mock:
|
||||
"""Provide a mock entity task."""
|
||||
return Mock(spec=Task)
|
||||
task = Mock(spec=Task)
|
||||
task.is_complete = False
|
||||
task.is_failed = False
|
||||
return task
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -77,6 +80,32 @@ def successful_agent_response() -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def configure_successful_entity_task(mock_entity_task: Mock) -> Any:
|
||||
"""Provide a helper to configure mock_entity_task with a successful response."""
|
||||
|
||||
def _configure(response: dict[str, Any]) -> Mock:
|
||||
mock_entity_task.is_failed = False
|
||||
mock_entity_task.is_complete = False
|
||||
mock_entity_task.get_result = Mock(return_value=response)
|
||||
return mock_entity_task
|
||||
|
||||
return _configure
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def configure_failed_entity_task(mock_entity_task: Mock) -> Any:
|
||||
"""Provide a helper to configure mock_entity_task with a failure."""
|
||||
|
||||
def _configure(exception: Exception) -> Mock:
|
||||
mock_entity_task.is_failed = True
|
||||
mock_entity_task.is_complete = True
|
||||
mock_entity_task.get_exception = Mock(return_value=exception)
|
||||
return mock_entity_task
|
||||
|
||||
return _configure
|
||||
|
||||
|
||||
class TestExecutorThreadCreation:
|
||||
"""Test that executors properly create DurableAgentThread with parameters."""
|
||||
|
||||
@@ -176,6 +205,115 @@ class TestClientAgentExecutorPollingConfiguration:
|
||||
assert isinstance(result, AgentRunResponse)
|
||||
|
||||
|
||||
class TestClientAgentExecutorFireAndForget:
|
||||
"""Test fire-and-forget mode (wait_for_response=False) for ClientAgentExecutor."""
|
||||
|
||||
def test_fire_and_forget_returns_immediately(self, mock_client: Mock) -> None:
|
||||
"""Verify wait_for_response=False returns immediately without polling."""
|
||||
executor = ClientAgentExecutor(mock_client, max_poll_retries=10, poll_interval_seconds=0.1)
|
||||
|
||||
# Create a request with wait_for_response=False
|
||||
request = RunRequest(message="test message", correlation_id="test-123", wait_for_response=False)
|
||||
|
||||
# Measure time taken
|
||||
start = time.time()
|
||||
result = executor.run_durable_agent("test_agent", request)
|
||||
elapsed = time.time() - start
|
||||
|
||||
# Should return immediately without polling (elapsed time should be very small)
|
||||
assert elapsed < 0.1 # Much faster than any polling would take
|
||||
|
||||
# Should return an AgentRunResponse
|
||||
assert isinstance(result, AgentRunResponse)
|
||||
|
||||
# Should have signaled the entity but not polled
|
||||
assert mock_client.signal_entity.call_count == 1
|
||||
assert mock_client.get_entity.call_count == 0 # No polling occurred
|
||||
|
||||
def test_fire_and_forget_returns_empty_response(self, mock_client: Mock) -> None:
|
||||
"""Verify wait_for_response=False returns an acceptance message with correlation ID."""
|
||||
executor = ClientAgentExecutor(mock_client)
|
||||
|
||||
request = RunRequest(message="test message", correlation_id="test-456", wait_for_response=False)
|
||||
|
||||
result = executor.run_durable_agent("test_agent", request)
|
||||
|
||||
# Verify it contains an acceptance message
|
||||
assert isinstance(result, AgentRunResponse)
|
||||
assert len(result.messages) == 1
|
||||
assert result.messages[0].role == Role.SYSTEM
|
||||
# Check message contains key information
|
||||
message_text = result.messages[0].text
|
||||
assert "accepted" in message_text.lower()
|
||||
assert "test-456" in message_text # Contains correlation ID
|
||||
assert "background" in message_text.lower()
|
||||
|
||||
|
||||
class TestOrchestrationAgentExecutorFireAndForget:
|
||||
"""Test fire-and-forget mode for OrchestrationAgentExecutor."""
|
||||
|
||||
def test_orchestration_fire_and_forget_calls_signal_entity(self, mock_orchestration_context: Mock) -> None:
|
||||
"""Verify wait_for_response=False calls signal_entity instead of call_entity."""
|
||||
executor = OrchestrationAgentExecutor(mock_orchestration_context)
|
||||
mock_orchestration_context.signal_entity = Mock()
|
||||
|
||||
request = RunRequest(message="test", correlation_id="test-123", wait_for_response=False)
|
||||
|
||||
result = executor.run_durable_agent("test_agent", request)
|
||||
|
||||
# Verify signal_entity was called and call_entity was not
|
||||
assert mock_orchestration_context.signal_entity.call_count == 1
|
||||
assert mock_orchestration_context.call_entity.call_count == 0
|
||||
|
||||
# Should still return a DurableAgentTask
|
||||
assert isinstance(result, DurableAgentTask)
|
||||
|
||||
def test_orchestration_fire_and_forget_returns_completed_task(self, mock_orchestration_context: Mock) -> None:
|
||||
"""Verify wait_for_response=False returns pre-completed DurableAgentTask."""
|
||||
executor = OrchestrationAgentExecutor(mock_orchestration_context)
|
||||
mock_orchestration_context.signal_entity = Mock()
|
||||
|
||||
request = RunRequest(message="test", correlation_id="test-456", wait_for_response=False)
|
||||
|
||||
result = executor.run_durable_agent("test_agent", request)
|
||||
|
||||
# Task should be immediately complete
|
||||
assert isinstance(result, DurableAgentTask)
|
||||
assert result.is_complete
|
||||
|
||||
def test_orchestration_fire_and_forget_returns_acceptance_response(self, mock_orchestration_context: Mock) -> None:
|
||||
"""Verify wait_for_response=False returns acceptance response."""
|
||||
executor = OrchestrationAgentExecutor(mock_orchestration_context)
|
||||
mock_orchestration_context.signal_entity = Mock()
|
||||
|
||||
request = RunRequest(message="test", correlation_id="test-789", wait_for_response=False)
|
||||
|
||||
result = executor.run_durable_agent("test_agent", request)
|
||||
|
||||
# Get the result
|
||||
response = result.get_result()
|
||||
assert isinstance(response, AgentRunResponse)
|
||||
assert len(response.messages) == 1
|
||||
assert response.messages[0].role == Role.SYSTEM
|
||||
assert "test-789" in response.messages[0].text
|
||||
|
||||
def test_orchestration_blocking_mode_calls_call_entity(self, mock_orchestration_context: Mock) -> None:
|
||||
"""Verify wait_for_response=True uses call_entity as before."""
|
||||
executor = OrchestrationAgentExecutor(mock_orchestration_context)
|
||||
mock_orchestration_context.signal_entity = Mock()
|
||||
|
||||
request = RunRequest(message="test", correlation_id="test-abc", wait_for_response=True)
|
||||
|
||||
result = executor.run_durable_agent("test_agent", request)
|
||||
|
||||
# Verify call_entity was called and signal_entity was not
|
||||
assert mock_orchestration_context.call_entity.call_count == 1
|
||||
assert mock_orchestration_context.signal_entity.call_count == 0
|
||||
|
||||
# Should return a DurableAgentTask
|
||||
assert isinstance(result, DurableAgentTask)
|
||||
|
||||
|
||||
class TestOrchestrationAgentExecutorRun:
|
||||
"""Test OrchestrationAgentExecutor.run_durable_agent implementation."""
|
||||
|
||||
@@ -240,11 +378,10 @@ class TestDurableAgentTask:
|
||||
"""Test DurableAgentTask completion and response transformation."""
|
||||
|
||||
def test_durable_agent_task_transforms_successful_result(
|
||||
self, mock_entity_task: Mock, successful_agent_response: dict[str, Any]
|
||||
self, configure_successful_entity_task: Any, successful_agent_response: dict[str, Any]
|
||||
) -> None:
|
||||
"""Verify DurableAgentTask converts successful entity result to AgentRunResponse."""
|
||||
mock_entity_task.is_failed = False
|
||||
mock_entity_task.get_result = Mock(return_value=successful_agent_response)
|
||||
mock_entity_task = configure_successful_entity_task(successful_agent_response)
|
||||
|
||||
task = DurableAgentTask(entity_task=mock_entity_task, response_format=None, correlation_id="test-123")
|
||||
|
||||
@@ -257,10 +394,9 @@ class TestDurableAgentTask:
|
||||
assert len(result.messages) == 1
|
||||
assert result.messages[0].role == Role.ASSISTANT
|
||||
|
||||
def test_durable_agent_task_propagates_failure(self, mock_entity_task: Mock) -> None:
|
||||
def test_durable_agent_task_propagates_failure(self, configure_failed_entity_task: Any) -> None:
|
||||
"""Verify DurableAgentTask propagates task failures."""
|
||||
mock_entity_task.is_failed = True
|
||||
mock_entity_task.get_exception = Mock(return_value=ValueError("Entity error"))
|
||||
mock_entity_task = configure_failed_entity_task(ValueError("Entity error"))
|
||||
|
||||
task = DurableAgentTask(entity_task=mock_entity_task, response_format=None, correlation_id="test-123")
|
||||
|
||||
@@ -269,19 +405,17 @@ class TestDurableAgentTask:
|
||||
|
||||
assert task.is_complete
|
||||
assert task.is_failed
|
||||
# The exception is wrapped in TaskFailedError by the durabletask library
|
||||
exception = task.get_exception()
|
||||
assert isinstance(exception, ValueError)
|
||||
assert str(exception) == "Entity error"
|
||||
assert exception is not None
|
||||
|
||||
def test_durable_agent_task_validates_response_format(self, mock_entity_task: Mock) -> None:
|
||||
def test_durable_agent_task_validates_response_format(self, configure_successful_entity_task: Any) -> None:
|
||||
"""Verify DurableAgentTask validates response format when provided."""
|
||||
mock_entity_task.is_failed = False
|
||||
mock_entity_task.get_result = Mock(
|
||||
return_value={
|
||||
"messages": [{"role": "assistant", "contents": [{"type": "text", "text": '{"answer": "42"}'}]}],
|
||||
"created_at": "2025-12-30T10:00:00Z",
|
||||
}
|
||||
)
|
||||
response = {
|
||||
"messages": [{"role": "assistant", "contents": [{"type": "text", "text": '{"answer": "42"}'}]}],
|
||||
"created_at": "2025-12-30T10:00:00Z",
|
||||
}
|
||||
mock_entity_task = configure_successful_entity_task(response)
|
||||
|
||||
class TestResponse(BaseModel):
|
||||
answer: str
|
||||
@@ -296,11 +430,10 @@ class TestDurableAgentTask:
|
||||
assert isinstance(result, AgentRunResponse)
|
||||
|
||||
def test_durable_agent_task_ignores_duplicate_completion(
|
||||
self, mock_entity_task: Mock, successful_agent_response: dict[str, Any]
|
||||
self, configure_successful_entity_task: Any, successful_agent_response: dict[str, Any]
|
||||
) -> None:
|
||||
"""Verify DurableAgentTask ignores duplicate completion calls."""
|
||||
mock_entity_task.is_failed = False
|
||||
mock_entity_task.get_result = Mock(return_value=successful_agent_response)
|
||||
mock_entity_task = configure_successful_entity_task(successful_agent_response)
|
||||
|
||||
task = DurableAgentTask(entity_task=mock_entity_task, response_format=None, correlation_id="test-123")
|
||||
|
||||
@@ -315,6 +448,124 @@ class TestDurableAgentTask:
|
||||
assert first_result is second_result
|
||||
assert mock_entity_task.get_result.call_count == 1
|
||||
|
||||
def test_durable_agent_task_fails_on_malformed_response(self, configure_successful_entity_task: Any) -> None:
|
||||
"""Verify DurableAgentTask fails when entity returns malformed response data."""
|
||||
# Use data that will cause AgentRunResponse.from_dict to fail
|
||||
# Using a list instead of dict, or other invalid structure
|
||||
mock_entity_task = configure_successful_entity_task("invalid string response")
|
||||
|
||||
task = DurableAgentTask(entity_task=mock_entity_task, response_format=None, correlation_id="test-123")
|
||||
|
||||
# Simulate child task completion with malformed data
|
||||
task.on_child_completed(mock_entity_task)
|
||||
|
||||
assert task.is_complete
|
||||
assert task.is_failed
|
||||
|
||||
def test_durable_agent_task_fails_on_invalid_response_format(self, configure_successful_entity_task: Any) -> None:
|
||||
"""Verify DurableAgentTask fails when response doesn't match required format."""
|
||||
response = {
|
||||
"messages": [{"role": "assistant", "contents": [{"type": "text", "text": '{"wrong": "field"}'}]}],
|
||||
"created_at": "2025-12-30T10:00:00Z",
|
||||
}
|
||||
mock_entity_task = configure_successful_entity_task(response)
|
||||
|
||||
class StrictResponse(BaseModel):
|
||||
required_field: str
|
||||
|
||||
task = DurableAgentTask(entity_task=mock_entity_task, response_format=StrictResponse, correlation_id="test-123")
|
||||
|
||||
# Simulate child task completion with wrong format
|
||||
task.on_child_completed(mock_entity_task)
|
||||
|
||||
assert task.is_complete
|
||||
assert task.is_failed
|
||||
|
||||
def test_durable_agent_task_handles_empty_response(self, configure_successful_entity_task: Any) -> None:
|
||||
"""Verify DurableAgentTask handles response with empty messages list."""
|
||||
response: dict[str, str | list[Any]] = {
|
||||
"messages": [],
|
||||
"created_at": "2025-12-30T10:00:00Z",
|
||||
}
|
||||
mock_entity_task = configure_successful_entity_task(response)
|
||||
|
||||
task = DurableAgentTask(entity_task=mock_entity_task, response_format=None, correlation_id="test-123")
|
||||
|
||||
# Simulate child task completion
|
||||
task.on_child_completed(mock_entity_task)
|
||||
|
||||
assert task.is_complete
|
||||
result = task.get_result()
|
||||
assert isinstance(result, AgentRunResponse)
|
||||
assert len(result.messages) == 0
|
||||
|
||||
def test_durable_agent_task_handles_multiple_messages(self, configure_successful_entity_task: Any) -> None:
|
||||
"""Verify DurableAgentTask correctly processes response with multiple messages."""
|
||||
response = {
|
||||
"messages": [
|
||||
{"role": "assistant", "contents": [{"type": "text", "text": "First message"}]},
|
||||
{"role": "assistant", "contents": [{"type": "text", "text": "Second message"}]},
|
||||
],
|
||||
"created_at": "2025-12-30T10:00:00Z",
|
||||
}
|
||||
mock_entity_task = configure_successful_entity_task(response)
|
||||
|
||||
task = DurableAgentTask(entity_task=mock_entity_task, response_format=None, correlation_id="test-123")
|
||||
|
||||
# Simulate child task completion
|
||||
task.on_child_completed(mock_entity_task)
|
||||
|
||||
assert task.is_complete
|
||||
result = task.get_result()
|
||||
assert isinstance(result, AgentRunResponse)
|
||||
assert len(result.messages) == 2
|
||||
assert result.messages[0].role == Role.ASSISTANT
|
||||
assert result.messages[1].role == Role.ASSISTANT
|
||||
|
||||
def test_durable_agent_task_is_not_complete_initially(self, mock_entity_task: Mock) -> None:
|
||||
"""Verify DurableAgentTask is not complete when first created."""
|
||||
task = DurableAgentTask(entity_task=mock_entity_task, response_format=None, correlation_id="test-123")
|
||||
|
||||
assert not task.is_complete
|
||||
assert not task.is_failed
|
||||
|
||||
def test_durable_agent_task_completes_with_complex_response_format(
|
||||
self, configure_successful_entity_task: Any
|
||||
) -> None:
|
||||
"""Verify DurableAgentTask validates complex nested response formats correctly."""
|
||||
response = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"contents": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": '{"name": "test", "count": 42, "items": ["a", "b", "c"]}',
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
"created_at": "2025-12-30T10:00:00Z",
|
||||
}
|
||||
mock_entity_task = configure_successful_entity_task(response)
|
||||
|
||||
class ComplexResponse(BaseModel):
|
||||
name: str
|
||||
count: int
|
||||
items: list[str]
|
||||
|
||||
task = DurableAgentTask(
|
||||
entity_task=mock_entity_task, response_format=ComplexResponse, correlation_id="test-123"
|
||||
)
|
||||
|
||||
# Simulate child task completion
|
||||
task.on_child_completed(mock_entity_task)
|
||||
|
||||
assert task.is_complete
|
||||
assert not task.is_failed
|
||||
result = task.get_result()
|
||||
assert isinstance(result, AgentRunResponse)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v", "--tb=short"])
|
||||
|
||||
@@ -25,6 +25,7 @@ class TestRunRequest:
|
||||
assert request.role == Role.USER
|
||||
assert request.response_format is None
|
||||
assert request.enable_tool_calls is True
|
||||
assert request.wait_for_response is True
|
||||
|
||||
def test_init_with_all_fields(self) -> None:
|
||||
"""Test RunRequest initialization with all fields."""
|
||||
@@ -35,6 +36,7 @@ class TestRunRequest:
|
||||
role=Role.SYSTEM,
|
||||
response_format=schema,
|
||||
enable_tool_calls=False,
|
||||
wait_for_response=False,
|
||||
)
|
||||
|
||||
assert request.message == "Hello"
|
||||
@@ -42,6 +44,7 @@ class TestRunRequest:
|
||||
assert request.role == Role.SYSTEM
|
||||
assert request.response_format is schema
|
||||
assert request.enable_tool_calls is False
|
||||
assert request.wait_for_response is False
|
||||
|
||||
def test_init_coerces_string_role(self) -> None:
|
||||
"""Ensure string role values are coerced into Role instances."""
|
||||
@@ -56,6 +59,7 @@ class TestRunRequest:
|
||||
|
||||
assert data["message"] == "Test message"
|
||||
assert data["enable_tool_calls"] is True
|
||||
assert data["wait_for_response"] is True
|
||||
assert data["role"] == "user"
|
||||
assert data["correlationId"] == "corr-004"
|
||||
assert "response_format" not in data or data["response_format"] is None
|
||||
@@ -70,6 +74,7 @@ class TestRunRequest:
|
||||
role=Role.ASSISTANT,
|
||||
response_format=schema,
|
||||
enable_tool_calls=False,
|
||||
wait_for_response=False,
|
||||
)
|
||||
data = request.to_dict()
|
||||
|
||||
@@ -80,6 +85,7 @@ class TestRunRequest:
|
||||
assert data["response_format"]["module"] == schema.__module__
|
||||
assert data["response_format"]["qualname"] == schema.__qualname__
|
||||
assert data["enable_tool_calls"] is False
|
||||
assert data["wait_for_response"] is False
|
||||
assert "thread_id" not in data
|
||||
|
||||
def test_from_dict_with_defaults(self) -> None:
|
||||
@@ -91,6 +97,7 @@ class TestRunRequest:
|
||||
assert request.correlation_id == "corr-006"
|
||||
assert request.role == Role.USER
|
||||
assert request.enable_tool_calls is True
|
||||
assert request.wait_for_response is True
|
||||
|
||||
def test_from_dict_ignores_thread_id_field(self) -> None:
|
||||
"""Ensure legacy thread_id input does not break RunRequest parsing."""
|
||||
|
||||
@@ -34,7 +34,10 @@ def mock_executor() -> Mock:
|
||||
|
||||
# Mock get_run_request to create actual RunRequest objects
|
||||
def create_run_request(
|
||||
message: str, response_format: type[BaseModel] | None = None, enable_tool_calls: bool = True
|
||||
message: str,
|
||||
response_format: type[BaseModel] | None = None,
|
||||
enable_tool_calls: bool = True,
|
||||
wait_for_response: bool = True,
|
||||
) -> RunRequest:
|
||||
import uuid
|
||||
|
||||
@@ -43,6 +46,7 @@ def mock_executor() -> Mock:
|
||||
correlation_id=str(uuid.uuid4()),
|
||||
response_format=response_format,
|
||||
enable_tool_calls=enable_tool_calls,
|
||||
wait_for_response=wait_for_response,
|
||||
)
|
||||
|
||||
mock.get_run_request = Mock(side_effect=create_run_request)
|
||||
|
||||
Reference in New Issue
Block a user