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
@@ -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