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:
Eduard van Valkenburg
2026-04-01 18:13:11 +02:00
committed by GitHub
Unverified
parent 38de991481
commit b065a4ce51
37 changed files with 1836 additions and 396 deletions
+491 -7
View File
@@ -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: