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:
Laveesh Rohra
2026-01-14 10:56:11 -08:00
committed by GitHub
Unverified
parent 1e36ba33c4
commit 3df916064c
48 changed files with 4221 additions and 1153 deletions
@@ -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(
+2 -2
View File
@@ -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)