mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: Added thread to AgentRunContext (#1732)
* Added thread to agent run context * Added sample * Update python/samples/getting_started/middleware/thread_behavior_middleware.py Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * Small fix --------- Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
Unverified
parent
a816408cd4
commit
adba312cd6
@@ -56,6 +56,20 @@ class TestAgentRunContext:
|
||||
assert context.is_streaming is True
|
||||
assert context.metadata == metadata
|
||||
|
||||
def test_init_with_thread(self, mock_agent: AgentProtocol) -> None:
|
||||
"""Test AgentRunContext initialization with thread parameter."""
|
||||
from agent_framework import AgentThread
|
||||
|
||||
messages = [ChatMessage(role=Role.USER, text="test")]
|
||||
thread = AgentThread()
|
||||
context = AgentRunContext(agent=mock_agent, messages=messages, thread=thread)
|
||||
|
||||
assert context.agent is mock_agent
|
||||
assert context.messages == messages
|
||||
assert context.thread is thread
|
||||
assert context.is_streaming is False
|
||||
assert context.metadata == {}
|
||||
|
||||
|
||||
class TestFunctionInvocationContext:
|
||||
"""Test cases for FunctionInvocationContext."""
|
||||
@@ -337,6 +351,61 @@ class TestAgentMiddlewarePipeline:
|
||||
assert context.terminate
|
||||
assert execution_order == ["handler_start", "handler_end"]
|
||||
|
||||
async def test_execute_with_thread_in_context(self, mock_agent: AgentProtocol) -> None:
|
||||
"""Test pipeline execution properly passes thread to middleware."""
|
||||
from agent_framework import AgentThread
|
||||
|
||||
captured_thread = None
|
||||
|
||||
class ThreadCapturingMiddleware(AgentMiddleware):
|
||||
async def process(
|
||||
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
|
||||
) -> None:
|
||||
nonlocal captured_thread
|
||||
captured_thread = context.thread
|
||||
await next(context)
|
||||
|
||||
middleware = ThreadCapturingMiddleware()
|
||||
pipeline = AgentMiddlewarePipeline([middleware])
|
||||
messages = [ChatMessage(role=Role.USER, text="test")]
|
||||
thread = AgentThread()
|
||||
context = AgentRunContext(agent=mock_agent, messages=messages, thread=thread)
|
||||
|
||||
expected_response = AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="response")])
|
||||
|
||||
async def final_handler(ctx: AgentRunContext) -> AgentRunResponse:
|
||||
return expected_response
|
||||
|
||||
result = await pipeline.execute(mock_agent, messages, context, final_handler)
|
||||
assert result == expected_response
|
||||
assert captured_thread is thread
|
||||
|
||||
async def test_execute_with_no_thread_in_context(self, mock_agent: AgentProtocol) -> None:
|
||||
"""Test pipeline execution when no thread is provided."""
|
||||
captured_thread = "not_none" # Use string to distinguish from None
|
||||
|
||||
class ThreadCapturingMiddleware(AgentMiddleware):
|
||||
async def process(
|
||||
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
|
||||
) -> None:
|
||||
nonlocal captured_thread
|
||||
captured_thread = context.thread
|
||||
await next(context)
|
||||
|
||||
middleware = ThreadCapturingMiddleware()
|
||||
pipeline = AgentMiddlewarePipeline([middleware])
|
||||
messages = [ChatMessage(role=Role.USER, text="test")]
|
||||
context = AgentRunContext(agent=mock_agent, messages=messages, thread=None)
|
||||
|
||||
expected_response = AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="response")])
|
||||
|
||||
async def final_handler(ctx: AgentRunContext) -> AgentRunResponse:
|
||||
return expected_response
|
||||
|
||||
result = await pipeline.execute(mock_agent, messages, context, final_handler)
|
||||
assert result == expected_response
|
||||
assert captured_thread is None
|
||||
|
||||
|
||||
class TestFunctionMiddlewarePipeline:
|
||||
"""Test cases for FunctionMiddlewarePipeline."""
|
||||
|
||||
@@ -1442,6 +1442,119 @@ class TestMiddlewareDecoratorLogic:
|
||||
assert test_function_middleware._middleware_type == MiddlewareType.FUNCTION # type: ignore[attr-defined]
|
||||
|
||||
|
||||
class TestChatAgentThreadBehavior:
|
||||
"""Test cases for thread behavior in AgentRunContext across multiple runs."""
|
||||
|
||||
async def test_agent_run_context_thread_behavior_across_multiple_runs(self, chat_client: "MockChatClient") -> None:
|
||||
"""Test that AgentRunContext.thread property behaves correctly across multiple agent runs."""
|
||||
thread_states: list[dict[str, Any]] = []
|
||||
|
||||
class ThreadTrackingMiddleware(AgentMiddleware):
|
||||
async def process(
|
||||
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
|
||||
) -> None:
|
||||
# Capture state before next() call
|
||||
thread_messages = []
|
||||
if context.thread and context.thread.message_store:
|
||||
thread_messages = await context.thread.message_store.list_messages()
|
||||
|
||||
before_state = {
|
||||
"before_next": True,
|
||||
"messages_count": len(context.messages),
|
||||
"thread_count": len(thread_messages),
|
||||
"messages_text": [msg.text for msg in context.messages if msg.text],
|
||||
"thread_messages_text": [msg.text for msg in thread_messages if msg.text],
|
||||
}
|
||||
thread_states.append(before_state)
|
||||
|
||||
await next(context)
|
||||
|
||||
# Capture state after next() call
|
||||
thread_messages_after = []
|
||||
if context.thread and context.thread.message_store:
|
||||
thread_messages_after = await context.thread.message_store.list_messages()
|
||||
|
||||
after_state = {
|
||||
"before_next": False,
|
||||
"messages_count": len(context.messages),
|
||||
"thread_count": len(thread_messages_after),
|
||||
"messages_text": [msg.text for msg in context.messages if msg.text],
|
||||
"thread_messages_text": [msg.text for msg in thread_messages_after if msg.text],
|
||||
}
|
||||
thread_states.append(after_state)
|
||||
|
||||
# Import the ChatMessageStore to configure the agent with a message store factory
|
||||
from agent_framework import ChatMessageStore
|
||||
|
||||
# Create ChatAgent with thread tracking middleware and a message store factory
|
||||
middleware = ThreadTrackingMiddleware()
|
||||
agent = ChatAgent(chat_client=chat_client, middleware=[middleware], chat_message_store_factory=ChatMessageStore)
|
||||
|
||||
# Create a thread that will persist messages between runs
|
||||
thread = agent.get_new_thread()
|
||||
|
||||
# First run
|
||||
first_messages = [ChatMessage(role=Role.USER, text="first message")]
|
||||
first_response = await agent.run(first_messages, thread=thread)
|
||||
|
||||
# Verify first response
|
||||
assert first_response is not None
|
||||
assert len(first_response.messages) > 0
|
||||
|
||||
# Second run - use the same thread
|
||||
second_messages = [ChatMessage(role=Role.USER, text="second message")]
|
||||
second_response = await agent.run(second_messages, thread=thread)
|
||||
|
||||
# Verify second response
|
||||
assert second_response is not None
|
||||
assert len(second_response.messages) > 0
|
||||
|
||||
# Verify we captured states for both runs (before and after next() for each)
|
||||
assert len(thread_states) == 4
|
||||
|
||||
# First run - before next()
|
||||
first_before = thread_states[0]
|
||||
assert first_before["before_next"] is True
|
||||
assert first_before["messages_count"] == 1
|
||||
assert first_before["thread_count"] == 0 # Thread is empty before first run
|
||||
assert first_before["messages_text"] == ["first message"]
|
||||
assert first_before["thread_messages_text"] == []
|
||||
|
||||
# First run - after next()
|
||||
first_after = thread_states[1]
|
||||
assert first_after["before_next"] is False
|
||||
assert first_after["messages_count"] == 1 # Input messages unchanged
|
||||
assert first_after["thread_count"] == 2 # Input + response
|
||||
assert first_after["messages_text"] == ["first message"]
|
||||
# Thread should contain input + response
|
||||
assert "first message" in first_after["thread_messages_text"]
|
||||
assert "test response" in " ".join(first_after["thread_messages_text"])
|
||||
|
||||
# Second run - before next()
|
||||
second_before = thread_states[2]
|
||||
assert second_before["before_next"] is True
|
||||
assert second_before["messages_count"] == 1 # Only current run input
|
||||
assert second_before["thread_count"] == 2 # Previous run history (input + response)
|
||||
assert second_before["messages_text"] == ["second message"]
|
||||
# Thread should contain previous run history but not current input yet
|
||||
assert "first message" in second_before["thread_messages_text"]
|
||||
assert "test response" in " ".join(second_before["thread_messages_text"])
|
||||
assert "second message" not in second_before["thread_messages_text"]
|
||||
|
||||
# Second run - after next()
|
||||
second_after = thread_states[3]
|
||||
assert second_after["before_next"] is False
|
||||
assert second_after["messages_count"] == 1 # Input messages unchanged
|
||||
assert second_after["thread_count"] == 4 # Previous history + current input + current response
|
||||
assert second_after["messages_text"] == ["second message"]
|
||||
# Thread should contain: first input + first response + second input + second response
|
||||
assert "first message" in second_after["thread_messages_text"]
|
||||
assert "second message" in second_after["thread_messages_text"]
|
||||
# Should have two "test response" entries (one for each run)
|
||||
response_count = sum(1 for text in second_after["thread_messages_text"] if "test response" in text)
|
||||
assert response_count == 2
|
||||
|
||||
|
||||
class TestChatAgentChatMiddleware:
|
||||
"""Test cases for chat middleware integration with ChatAgent."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user