mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING] update context provider APIs, middleware, and per-service-call history persistence (#4992)
* Rename provider base APIs Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Allow provider-added chat and function middleware Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Simulate service-stored history per model call Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Fix typing regressions in CI Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Fix response ID suppression review feedback Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Rename per-service-call history persistence APIs Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Address context persistence review feedback Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Stabilize markdown sample docs Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Persist service continuation state per call Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --------- Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
Unverified
parent
38de991481
commit
b065a4ce51
@@ -3,8 +3,8 @@
|
||||
import contextlib
|
||||
import inspect
|
||||
import json
|
||||
from collections.abc import AsyncIterable, MutableSequence
|
||||
from typing import Any
|
||||
from collections.abc import AsyncIterable, Awaitable, Callable, MutableSequence, Sequence
|
||||
from typing import Any, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
@@ -18,22 +18,29 @@ from agent_framework import (
|
||||
AgentResponse,
|
||||
AgentResponseUpdate,
|
||||
AgentSession,
|
||||
BaseContextProvider,
|
||||
ChatContext,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
ContextProvider,
|
||||
FunctionTool,
|
||||
HistoryProvider,
|
||||
InMemoryHistoryProvider,
|
||||
Message,
|
||||
ResponseStream,
|
||||
SessionContext,
|
||||
SlidingWindowStrategy,
|
||||
SupportsAgentRun,
|
||||
SupportsChatGetResponse,
|
||||
TruncationStrategy,
|
||||
chat_middleware,
|
||||
tool,
|
||||
)
|
||||
from agent_framework._agents import _get_tool_name, _merge_options, _sanitize_agent_name
|
||||
from agent_framework._mcp import MCPTool, _build_prefixed_mcp_name, _normalize_mcp_name
|
||||
from agent_framework._middleware import FunctionInvocationContext
|
||||
from agent_framework.exceptions import AgentInvalidRequestException, ChatClientInvalidResponseException
|
||||
|
||||
|
||||
class _FixedTokenizer:
|
||||
@@ -68,6 +75,49 @@ class _ConnectedMCPTool(MCPTool):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class _RecordingHistoryProvider(HistoryProvider):
|
||||
def __init__(self, source_id: str = "recording_history") -> None:
|
||||
super().__init__(source_id=source_id)
|
||||
|
||||
async def get_messages(
|
||||
self,
|
||||
session_id: str | None,
|
||||
*,
|
||||
state: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> list[Message]:
|
||||
if state is None:
|
||||
return []
|
||||
state["get_call_count"] = state.get("get_call_count", 0) + 1
|
||||
return list(cast(list[Message], state.get("messages", [])))
|
||||
|
||||
async def save_messages(
|
||||
self,
|
||||
session_id: str | None,
|
||||
messages: Sequence[Message],
|
||||
*,
|
||||
state: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
if state is None:
|
||||
return
|
||||
state["save_call_count"] = state.get("save_call_count", 0) + 1
|
||||
state.setdefault("messages", []).extend(messages)
|
||||
|
||||
|
||||
class _ResponseIdRecordingHistoryProvider(_RecordingHistoryProvider):
|
||||
async def after_run(
|
||||
self,
|
||||
*,
|
||||
agent: SupportsAgentRun,
|
||||
session: AgentSession,
|
||||
context: SessionContext,
|
||||
state: dict[str, Any],
|
||||
) -> None:
|
||||
state.setdefault("response_ids", []).append(context.response.response_id if context.response else None)
|
||||
await super().after_run(agent=agent, session=session, context=context, state=state)
|
||||
|
||||
|
||||
def test_agent_session_type(agent_session: AgentSession) -> None:
|
||||
assert isinstance(agent_session, AgentSession)
|
||||
|
||||
@@ -314,6 +364,413 @@ async def test_prepare_run_context_handles_function_kwargs(
|
||||
assert ctx["client_kwargs"]["session"] is session
|
||||
|
||||
|
||||
async def test_chat_agent_persists_history_per_service_call(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
provider = _RecordingHistoryProvider()
|
||||
|
||||
@tool(name="lookup_weather", approval_mode="never_require")
|
||||
def lookup_weather(location: str) -> str:
|
||||
return f"Weather in {location}: sunny"
|
||||
|
||||
session = AgentSession()
|
||||
session.state[provider.source_id] = {
|
||||
"messages": [
|
||||
Message(role="user", text="Earlier question"),
|
||||
Message(role="assistant", text="Earlier answer"),
|
||||
]
|
||||
}
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_1",
|
||||
name="lookup_weather",
|
||||
arguments='{"location": "Seattle"}',
|
||||
)
|
||||
],
|
||||
),
|
||||
response_id="resp_call_1",
|
||||
),
|
||||
ChatResponse(messages=Message(role="assistant", text="It is sunny in Seattle."), response_id="resp_call_2"),
|
||||
]
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
tools=[lookup_weather],
|
||||
context_providers=[provider],
|
||||
require_per_service_call_history_persistence=True,
|
||||
)
|
||||
|
||||
result = await agent.run("What's the weather in Seattle?", session=session)
|
||||
|
||||
provider_state = session.state[provider.source_id]
|
||||
stored_messages = cast(list[Message], provider_state["messages"])
|
||||
|
||||
assert result.text == "It is sunny in Seattle."
|
||||
assert result.response_id is None
|
||||
assert chat_client_base.call_count == 2
|
||||
assert provider_state["get_call_count"] == 2
|
||||
assert provider_state["save_call_count"] == 2
|
||||
assert stored_messages[-1].text == "It is sunny in Seattle."
|
||||
assert session.service_session_id is None
|
||||
|
||||
|
||||
async def test_chat_agent_persists_history_per_service_call_streaming(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
provider = _RecordingHistoryProvider()
|
||||
|
||||
@tool(name="lookup_weather", approval_mode="never_require")
|
||||
def lookup_weather(location: str) -> str:
|
||||
return f"Weather in {location}: sunny"
|
||||
|
||||
session = AgentSession()
|
||||
session.state[provider.source_id] = {
|
||||
"messages": [
|
||||
Message(role="user", text="Earlier question"),
|
||||
Message(role="assistant", text="Earlier answer"),
|
||||
]
|
||||
}
|
||||
chat_client_base.streaming_responses = [
|
||||
[
|
||||
ChatResponseUpdate(
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_1",
|
||||
name="lookup_weather",
|
||||
arguments='{"location": "Seattle"}',
|
||||
)
|
||||
],
|
||||
role="assistant",
|
||||
finish_reason="stop",
|
||||
response_id="resp_call_1",
|
||||
)
|
||||
],
|
||||
[
|
||||
ChatResponseUpdate(
|
||||
contents=[Content.from_text("It is sunny in Seattle.")],
|
||||
role="assistant",
|
||||
finish_reason="stop",
|
||||
response_id="resp_call_2",
|
||||
)
|
||||
],
|
||||
]
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
tools=[lookup_weather],
|
||||
context_providers=[provider],
|
||||
require_per_service_call_history_persistence=True,
|
||||
)
|
||||
|
||||
stream = agent.run("What's the weather in Seattle?", session=session, stream=True)
|
||||
async for _ in stream:
|
||||
pass
|
||||
result = await stream.get_final_response()
|
||||
|
||||
provider_state = session.state[provider.source_id]
|
||||
stored_messages = cast(list[Message], provider_state["messages"])
|
||||
|
||||
assert result.text == "It is sunny in Seattle."
|
||||
assert result.response_id is None
|
||||
assert chat_client_base.call_count == 2
|
||||
assert provider_state["get_call_count"] == 2
|
||||
assert provider_state["save_call_count"] == 2
|
||||
assert stored_messages[-1].text == "It is sunny in Seattle."
|
||||
assert session.service_session_id is None
|
||||
|
||||
|
||||
async def test_streaming_per_service_call_persistence_hides_response_id_from_after_run(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
provider = _ResponseIdRecordingHistoryProvider()
|
||||
|
||||
@tool(name="lookup_weather", approval_mode="never_require")
|
||||
def lookup_weather(location: str) -> str:
|
||||
return f"Weather in {location}: sunny"
|
||||
|
||||
session = AgentSession()
|
||||
session.state[provider.source_id] = {"messages": []}
|
||||
chat_client_base.streaming_responses = [
|
||||
[
|
||||
ChatResponseUpdate(
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_1",
|
||||
name="lookup_weather",
|
||||
arguments='{"location": "Seattle"}',
|
||||
)
|
||||
],
|
||||
role="assistant",
|
||||
finish_reason="stop",
|
||||
response_id="resp_call_1",
|
||||
)
|
||||
],
|
||||
[
|
||||
ChatResponseUpdate(
|
||||
contents=[Content.from_text("It is sunny in Seattle.")],
|
||||
role="assistant",
|
||||
finish_reason="stop",
|
||||
response_id="resp_call_2",
|
||||
)
|
||||
],
|
||||
]
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
tools=[lookup_weather],
|
||||
context_providers=[provider],
|
||||
require_per_service_call_history_persistence=True,
|
||||
)
|
||||
|
||||
stream = agent.run("What's the weather in Seattle?", session=session, stream=True)
|
||||
async for _ in stream:
|
||||
pass
|
||||
result = await stream.get_final_response()
|
||||
|
||||
provider_state = session.state[provider.source_id]
|
||||
|
||||
assert result.response_id is None
|
||||
assert provider_state["response_ids"] == [None, None]
|
||||
|
||||
|
||||
async def test_per_service_call_persistence_uses_real_service_storage_when_client_stores_by_default(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
provider = _RecordingHistoryProvider()
|
||||
|
||||
@tool(name="lookup_weather", approval_mode="never_require")
|
||||
def lookup_weather(location: str) -> str:
|
||||
return f"Weather in {location}: sunny"
|
||||
|
||||
chat_client_base.STORES_BY_DEFAULT = True # type: ignore[attr-defined]
|
||||
|
||||
session = AgentSession()
|
||||
session.state[provider.source_id] = {"messages": []}
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_1",
|
||||
name="lookup_weather",
|
||||
arguments='{"location": "Seattle"}',
|
||||
)
|
||||
],
|
||||
),
|
||||
conversation_id="resp_service_managed",
|
||||
response_id="resp_call_1",
|
||||
),
|
||||
ChatResponse(
|
||||
messages=Message(role="assistant", text="It is sunny in Seattle."),
|
||||
conversation_id="resp_service_managed",
|
||||
response_id="resp_call_2",
|
||||
),
|
||||
]
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
tools=[lookup_weather],
|
||||
context_providers=[provider],
|
||||
require_per_service_call_history_persistence=True,
|
||||
)
|
||||
|
||||
result = await agent.run("What's the weather in Seattle?", session=session)
|
||||
|
||||
provider_state = session.state[provider.source_id]
|
||||
|
||||
assert result.text == "It is sunny in Seattle."
|
||||
assert result.response_id == "resp_call_2"
|
||||
assert chat_client_base.call_count == 2
|
||||
assert "get_call_count" not in provider_state
|
||||
assert "save_call_count" not in provider_state
|
||||
assert session.service_session_id == "resp_service_managed"
|
||||
|
||||
|
||||
async def test_service_storage_updates_session_handle_per_service_call_before_non_streaming_failure(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
provider = _RecordingHistoryProvider()
|
||||
|
||||
@tool(name="lookup_weather", approval_mode="never_require")
|
||||
def lookup_weather(location: str) -> str:
|
||||
return f"Weather in {location}: sunny"
|
||||
|
||||
chat_client_base.STORES_BY_DEFAULT = True # type: ignore[attr-defined]
|
||||
|
||||
session = AgentSession()
|
||||
session.state[provider.source_id] = {"messages": []}
|
||||
first_response = ChatResponse(
|
||||
messages=Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_1",
|
||||
name="lookup_weather",
|
||||
arguments='{"location": "Seattle"}',
|
||||
)
|
||||
],
|
||||
),
|
||||
conversation_id="resp_call_1",
|
||||
response_id="resp_call_1",
|
||||
)
|
||||
mock_get_non_streaming_response = AsyncMock(
|
||||
side_effect=[first_response, RuntimeError("service down")],
|
||||
)
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
tools=[lookup_weather],
|
||||
context_providers=[provider],
|
||||
require_per_service_call_history_persistence=True,
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(chat_client_base, "_get_non_streaming_response", new=mock_get_non_streaming_response),
|
||||
pytest.raises(RuntimeError, match="service down"),
|
||||
):
|
||||
await agent.run("What's the weather in Seattle?", session=session)
|
||||
|
||||
assert mock_get_non_streaming_response.await_count == 2
|
||||
assert session.service_session_id == "resp_call_1"
|
||||
|
||||
|
||||
async def test_service_storage_updates_session_handle_per_service_call_before_streaming_failure(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
provider = _RecordingHistoryProvider()
|
||||
|
||||
@tool(name="lookup_weather", approval_mode="never_require")
|
||||
def lookup_weather(location: str) -> str:
|
||||
return f"Weather in {location}: sunny"
|
||||
|
||||
chat_client_base.STORES_BY_DEFAULT = True # type: ignore[attr-defined]
|
||||
|
||||
session = AgentSession()
|
||||
session.state[provider.source_id] = {"messages": []}
|
||||
|
||||
async def _first_stream_updates() -> AsyncIterable[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_1",
|
||||
name="lookup_weather",
|
||||
arguments='{"location": "Seattle"}',
|
||||
)
|
||||
],
|
||||
role="assistant",
|
||||
finish_reason="stop",
|
||||
)
|
||||
|
||||
def _finalize_first_stream(_updates: Sequence[ChatResponseUpdate]) -> ChatResponse[Any]:
|
||||
return ChatResponse(
|
||||
messages=Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_1",
|
||||
name="lookup_weather",
|
||||
arguments='{"location": "Seattle"}',
|
||||
)
|
||||
],
|
||||
),
|
||||
conversation_id="resp_call_1",
|
||||
response_id="resp_call_1",
|
||||
)
|
||||
|
||||
first_stream = ResponseStream(_first_stream_updates(), finalizer=_finalize_first_stream)
|
||||
mock_get_streaming_response = MagicMock(side_effect=[first_stream, RuntimeError("service down")])
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
tools=[lookup_weather],
|
||||
context_providers=[provider],
|
||||
require_per_service_call_history_persistence=True,
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(chat_client_base, "_get_streaming_response", new=mock_get_streaming_response),
|
||||
pytest.raises(RuntimeError, match="service down"),
|
||||
):
|
||||
stream = agent.run("What's the weather in Seattle?", session=session, stream=True)
|
||||
async for _ in stream:
|
||||
pass
|
||||
|
||||
assert mock_get_streaming_response.call_count == 2
|
||||
assert session.service_session_id == "resp_call_1"
|
||||
|
||||
|
||||
async def test_chat_agent_without_per_service_call_persistence_preserves_response_id(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(role="assistant", text="Hello"),
|
||||
response_id="resp_call_1",
|
||||
)
|
||||
]
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
context_providers=[InMemoryHistoryProvider()],
|
||||
)
|
||||
|
||||
result = await agent.run("Hello", session=AgentSession(), options={"store": False})
|
||||
|
||||
assert result.response_id == "resp_call_1"
|
||||
|
||||
|
||||
async def test_per_service_call_persistence_rejects_real_service_conversation_id(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
provider = _RecordingHistoryProvider()
|
||||
chat_client_base.STORES_BY_DEFAULT = True # type: ignore[attr-defined]
|
||||
session = AgentSession()
|
||||
session.state[provider.source_id] = {"messages": []}
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=Message(role="assistant", text="Hello"),
|
||||
conversation_id="resp_service_managed",
|
||||
)
|
||||
]
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
context_providers=[provider],
|
||||
require_per_service_call_history_persistence=True,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
ChatClientInvalidResponseException,
|
||||
match="require_per_service_call_history_persistence cannot be used",
|
||||
):
|
||||
await agent.run("Hello", session=session, options={"store": False})
|
||||
|
||||
|
||||
async def test_per_service_call_persistence_rejects_existing_conversation_id_when_service_not_storing_history(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
provider = _RecordingHistoryProvider()
|
||||
session = AgentSession()
|
||||
session.state[provider.source_id] = {"messages": []}
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
context_providers=[provider],
|
||||
require_per_service_call_history_persistence=True,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
AgentInvalidRequestException,
|
||||
match="require_per_service_call_history_persistence cannot be used",
|
||||
):
|
||||
await agent.run("Hello", session=session, options={"store": False, "conversation_id": "existing_conversation"})
|
||||
|
||||
|
||||
async def test_chat_client_agent_run_with_session(chat_client_base: SupportsChatGetResponse) -> None:
|
||||
mock_response = ChatResponse(
|
||||
messages=[Message(role="assistant", contents=[Content.from_text("test response")])],
|
||||
@@ -586,7 +1043,7 @@ async def test_chat_client_agent_author_name_is_used_from_response(
|
||||
|
||||
|
||||
# Mock context provider for testing
|
||||
class MockContextProvider(BaseContextProvider):
|
||||
class MockContextProvider(ContextProvider):
|
||||
def __init__(self, messages: list[Message] | None = None) -> None:
|
||||
super().__init__(source_id="mock")
|
||||
self.context_messages = messages
|
||||
@@ -1723,7 +2180,7 @@ async def test_agent_create_session_with_context_providers(
|
||||
):
|
||||
"""Test that create_session works when context_providers are set on the agent."""
|
||||
|
||||
class TestContextProvider(BaseContextProvider):
|
||||
class TestContextProvider(ContextProvider):
|
||||
def __init__(self):
|
||||
super().__init__(source_id="test")
|
||||
|
||||
@@ -1798,7 +2255,7 @@ async def test_chat_agent_context_provider_adds_tools_when_agent_has_none(
|
||||
"""A tool provided by context."""
|
||||
return text
|
||||
|
||||
class ToolContextProvider(BaseContextProvider):
|
||||
class ToolContextProvider(ContextProvider):
|
||||
def __init__(self):
|
||||
super().__init__(source_id="tool-context")
|
||||
|
||||
@@ -1827,7 +2284,7 @@ async def test_chat_agent_context_provider_adds_instructions_when_agent_has_none
|
||||
):
|
||||
"""Test that context provider instructions are used when agent has no default instructions."""
|
||||
|
||||
class InstructionContextProvider(BaseContextProvider):
|
||||
class InstructionContextProvider(ContextProvider):
|
||||
def __init__(self):
|
||||
super().__init__(source_id="instruction-context")
|
||||
|
||||
@@ -1849,6 +2306,33 @@ async def test_chat_agent_context_provider_adds_instructions_when_agent_has_none
|
||||
assert options.get("instructions") == "Context-provided instructions"
|
||||
|
||||
|
||||
async def test_chat_agent_context_provider_adds_middleware_when_agent_has_none(
|
||||
chat_client_base: SupportsChatGetResponse,
|
||||
) -> None:
|
||||
"""Test that context provider middleware is collected during preparation."""
|
||||
|
||||
@chat_middleware
|
||||
async def context_chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
await call_next()
|
||||
|
||||
class MiddlewareContextProvider(ContextProvider):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(source_id="middleware-context")
|
||||
|
||||
async def before_run(self, *, agent, session, context, state) -> None:
|
||||
context.extend_middleware("middleware-context", context_chat_middleware)
|
||||
|
||||
agent = Agent(client=chat_client_base, context_providers=[MiddlewareContextProvider()])
|
||||
|
||||
session_context, _ = await agent._prepare_session_and_messages( # type: ignore[reportPrivateUsage]
|
||||
session=None,
|
||||
input_messages=[Message(role="user", text="Hello")],
|
||||
)
|
||||
|
||||
assert session_context.middleware["middleware-context"] == [context_chat_middleware]
|
||||
assert session_context.get_middleware() == [context_chat_middleware]
|
||||
|
||||
|
||||
# region STORES_BY_DEFAULT tests
|
||||
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ from agent_framework import (
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
ContextProvider,
|
||||
FunctionInvocationContext,
|
||||
FunctionMiddleware,
|
||||
FunctionTool,
|
||||
@@ -464,6 +465,31 @@ class TestChatAgentMultipleMiddlewareOrdering:
|
||||
expected_order = ["class_agent_before", "function_agent_before", "function_agent_after", "class_agent_after"]
|
||||
assert execution_order == expected_order
|
||||
|
||||
async def test_provider_added_agent_middleware_is_rejected(self, chat_client_base: "MockBaseChatClient") -> None:
|
||||
"""Test provider-added agent middleware is rejected explicitly."""
|
||||
|
||||
@agent_middleware
|
||||
async def provider_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
await call_next()
|
||||
|
||||
class ProviderMiddlewareContextProvider(ContextProvider):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(source_id="provider-middleware")
|
||||
|
||||
async def before_run(self, *, agent, session, context, state) -> None:
|
||||
context.extend_middleware(self.source_id, provider_middleware)
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
context_providers=[ProviderMiddlewareContextProvider()],
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
MiddlewareException,
|
||||
match="Context providers may only add chat or function middleware",
|
||||
):
|
||||
await agent.run([Message(role="user", text="test message")])
|
||||
|
||||
|
||||
# region Tool Functions for Testing
|
||||
|
||||
@@ -2066,6 +2092,121 @@ class TestChatAgentChatMiddleware:
|
||||
"agent_middleware_after",
|
||||
]
|
||||
|
||||
async def test_provider_added_chat_and_function_middleware_are_forwarded(
|
||||
self, chat_client_base: "MockBaseChatClient"
|
||||
) -> None:
|
||||
"""Test provider-added chat and function middleware forwarding and ordering."""
|
||||
execution_order: list[str] = []
|
||||
|
||||
@chat_middleware
|
||||
async def constructor_chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
execution_order.append("constructor_chat_before")
|
||||
await call_next()
|
||||
execution_order.append("constructor_chat_after")
|
||||
|
||||
@chat_middleware
|
||||
async def provider_chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
execution_order.append("provider_chat_before")
|
||||
await call_next()
|
||||
execution_order.append("provider_chat_after")
|
||||
|
||||
@chat_middleware
|
||||
async def run_chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
execution_order.append("run_chat_before")
|
||||
await call_next()
|
||||
execution_order.append("run_chat_after")
|
||||
|
||||
@function_middleware
|
||||
async def constructor_function_middleware(
|
||||
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
|
||||
) -> None:
|
||||
execution_order.append("constructor_function_before")
|
||||
await call_next()
|
||||
execution_order.append("constructor_function_after")
|
||||
|
||||
@function_middleware
|
||||
async def provider_function_middleware(
|
||||
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
|
||||
) -> None:
|
||||
execution_order.append("provider_function_before")
|
||||
await call_next()
|
||||
execution_order.append("provider_function_after")
|
||||
|
||||
@function_middleware
|
||||
async def run_function_middleware(
|
||||
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
|
||||
) -> None:
|
||||
execution_order.append("run_function_before")
|
||||
await call_next()
|
||||
execution_order.append("run_function_after")
|
||||
|
||||
class ProviderMiddlewareContextProvider(ContextProvider):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(source_id="provider-middleware")
|
||||
|
||||
async def before_run(self, *, agent, session, context, state) -> None:
|
||||
context.extend_middleware(
|
||||
self.source_id,
|
||||
[
|
||||
provider_chat_middleware,
|
||||
provider_function_middleware,
|
||||
],
|
||||
)
|
||||
|
||||
chat_client_base.run_responses = [
|
||||
ChatResponse(
|
||||
messages=[
|
||||
Message(
|
||||
role="assistant",
|
||||
contents=[
|
||||
Content.from_function_call(
|
||||
call_id="call_provider",
|
||||
name="sample_tool_function",
|
||||
arguments='{"location": "Seattle"}',
|
||||
)
|
||||
],
|
||||
)
|
||||
]
|
||||
),
|
||||
ChatResponse(messages=[Message(role="assistant", text="Final response")]),
|
||||
]
|
||||
|
||||
agent = Agent(
|
||||
client=chat_client_base,
|
||||
middleware=[constructor_chat_middleware, constructor_function_middleware],
|
||||
context_providers=[ProviderMiddlewareContextProvider()],
|
||||
tools=[sample_tool_function],
|
||||
)
|
||||
|
||||
response = await agent.run(
|
||||
[Message(role="user", text="Get weather for Seattle")],
|
||||
middleware=[run_chat_middleware, run_function_middleware],
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert chat_client_base.call_count == 2
|
||||
assert response.messages[-1].text == "Final response"
|
||||
assert execution_order == [
|
||||
"constructor_chat_before",
|
||||
"run_chat_before",
|
||||
"provider_chat_before",
|
||||
"provider_chat_after",
|
||||
"run_chat_after",
|
||||
"constructor_chat_after",
|
||||
"constructor_function_before",
|
||||
"run_function_before",
|
||||
"provider_function_before",
|
||||
"provider_function_after",
|
||||
"run_function_after",
|
||||
"constructor_function_after",
|
||||
"constructor_chat_before",
|
||||
"run_chat_before",
|
||||
"provider_chat_before",
|
||||
"provider_chat_after",
|
||||
"run_chat_after",
|
||||
"constructor_chat_after",
|
||||
]
|
||||
|
||||
async def test_agent_middleware_can_access_and_override_options(self) -> None:
|
||||
"""Test that agent middleware can access and override runtime options."""
|
||||
captured_options: dict[str, Any] = {}
|
||||
|
||||
@@ -1,16 +1,26 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
|
||||
from agent_framework import Message
|
||||
from agent_framework._sessions import (
|
||||
import pytest
|
||||
|
||||
from agent_framework import (
|
||||
AgentContext,
|
||||
AgentSession,
|
||||
BaseContextProvider,
|
||||
BaseHistoryProvider,
|
||||
ChatContext,
|
||||
ContextProvider,
|
||||
HistoryProvider,
|
||||
InMemoryHistoryProvider,
|
||||
Message,
|
||||
SessionContext,
|
||||
agent_middleware,
|
||||
chat_middleware,
|
||||
)
|
||||
from agent_framework._sessions import LOCAL_HISTORY_CONVERSATION_ID, is_local_history_conversation_id
|
||||
from agent_framework.exceptions import MiddlewareException
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SessionContext tests
|
||||
@@ -102,6 +112,50 @@ class TestSessionContext:
|
||||
ctx.extend_instructions("sys", ["Be helpful", "Be concise"])
|
||||
assert ctx.instructions == ["Be helpful", "Be concise"]
|
||||
|
||||
def test_extend_middleware_creates_key_and_appends(self) -> None:
|
||||
ctx = SessionContext(input_messages=[])
|
||||
|
||||
@chat_middleware
|
||||
async def first_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
await call_next()
|
||||
|
||||
@chat_middleware
|
||||
async def second_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
await call_next()
|
||||
|
||||
ctx.extend_middleware("rag", first_middleware)
|
||||
ctx.extend_middleware("rag", [second_middleware])
|
||||
|
||||
assert ctx.middleware["rag"] == [first_middleware, second_middleware]
|
||||
assert ctx.get_middleware() == [first_middleware, second_middleware]
|
||||
|
||||
def test_extend_middleware_preserves_source_order(self) -> None:
|
||||
ctx = SessionContext(input_messages=[])
|
||||
|
||||
@chat_middleware
|
||||
async def first_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
await call_next()
|
||||
|
||||
@chat_middleware
|
||||
async def second_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
await call_next()
|
||||
|
||||
ctx.extend_middleware("a", first_middleware)
|
||||
ctx.extend_middleware("b", second_middleware)
|
||||
|
||||
assert list(ctx.middleware.keys()) == ["a", "b"]
|
||||
assert ctx.get_middleware() == [first_middleware, second_middleware]
|
||||
|
||||
def test_extend_middleware_rejects_agent_middleware(self) -> None:
|
||||
ctx = SessionContext(input_messages=[])
|
||||
|
||||
@agent_middleware
|
||||
async def provider_agent_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
|
||||
await call_next()
|
||||
|
||||
with pytest.raises(MiddlewareException, match="Context providers may only add chat or function middleware"):
|
||||
ctx.extend_middleware("rag", provider_agent_middleware)
|
||||
|
||||
def test_get_messages_all(self) -> None:
|
||||
ctx = SessionContext(input_messages=[])
|
||||
ctx.extend_messages("a", [Message(role="user", contents=["a"])])
|
||||
@@ -154,37 +208,58 @@ class TestSessionContext:
|
||||
ctx._response = resp
|
||||
assert ctx.response is resp
|
||||
|
||||
def test_local_history_conversation_id_sentinel(self) -> None:
|
||||
assert is_local_history_conversation_id(LOCAL_HISTORY_CONVERSATION_ID) is True
|
||||
assert is_local_history_conversation_id("some_other_id") is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BaseContextProvider tests
|
||||
# ContextProvider tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestContextProviderBase:
|
||||
class TestContextProvider:
|
||||
def test_source_id_required(self) -> None:
|
||||
provider = BaseContextProvider(source_id="test")
|
||||
provider = ContextProvider(source_id="test")
|
||||
assert provider.source_id == "test"
|
||||
|
||||
async def test_before_run_is_noop(self) -> None:
|
||||
provider = BaseContextProvider(source_id="test")
|
||||
provider = ContextProvider(source_id="test")
|
||||
session = AgentSession()
|
||||
ctx = SessionContext(input_messages=[])
|
||||
# Should not raise
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state={}) # type: ignore[arg-type]
|
||||
|
||||
async def test_after_run_is_noop(self) -> None:
|
||||
provider = BaseContextProvider(source_id="test")
|
||||
provider = ContextProvider(source_id="test")
|
||||
session = AgentSession()
|
||||
ctx = SessionContext(input_messages=[])
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state={}) # type: ignore[arg-type]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BaseHistoryProvider tests
|
||||
# Deprecated provider alias tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ConcreteHistoryProvider(BaseHistoryProvider):
|
||||
class TestDeprecatedProviderAliases:
|
||||
def test_base_context_provider_warns_and_is_compatible(self) -> None:
|
||||
with pytest.warns(DeprecationWarning, match="BaseContextProvider is deprecated. Use ContextProvider instead."):
|
||||
provider = BaseContextProvider(source_id="test")
|
||||
|
||||
assert isinstance(provider, ContextProvider)
|
||||
|
||||
def test_base_provider_aliases_preserve_subtyping(self) -> None:
|
||||
assert issubclass(BaseContextProvider, ContextProvider)
|
||||
assert issubclass(BaseHistoryProvider, HistoryProvider)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HistoryProvider tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ConcreteHistoryProvider(HistoryProvider):
|
||||
"""Concrete test implementation."""
|
||||
|
||||
def __init__(self, source_id: str, stored_messages: list[Message] | None = None, **kwargs) -> None:
|
||||
|
||||
Reference in New Issue
Block a user