Python: [BREAKING] Scope provider state by source_id and standardize source IDs (#3995)

* Initial plan

* Add FoundryMemoryProvider and tests

Co-authored-by: eavanvalkenburg <13749212+eavanvalkenburg@users.noreply.github.com>

* Add sample and documentation for FoundryMemoryProvider

Co-authored-by: eavanvalkenburg <13749212+eavanvalkenburg@users.noreply.github.com>

* Address code review feedback for FoundryMemoryProvider

Co-authored-by: eavanvalkenburg <13749212+eavanvalkenburg@users.noreply.github.com>

* Address PR review comments: Add DEFAULT_SOURCE_ID, use logging.getLogger, move state to session.state

Co-authored-by: eavanvalkenburg <13749212+eavanvalkenburg@users.noreply.github.com>

* Fix Foundry memory ItemParam usage and exports

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Refactor provider hook state and standardize source IDs

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Support endpoint-based Foundry memory init

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Fix core README workflows link

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* updated implementation and sample

* Split out Foundry memory provider changes

Remove FoundryMemoryProvider implementation/tests/sample plus export and docs mentions from this branch so only non-Foundry changes remain.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

* Trigger CI rerun for PR #3995

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

---------

Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
Co-authored-by: eavanvalkenburg <13749212+eavanvalkenburg@users.noreply.github.com>
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
Eduard van Valkenburg
2026-02-17 20:12:28 +01:00
committed by GitHub
Unverified
parent a5f948c215
commit cc98d5b6f7
28 changed files with 359 additions and 148 deletions
@@ -1072,6 +1072,50 @@ Rationale for B1 over B2: Simpler is better. The whole state dict is passed to e
> **Note on trust:** Since all `ContextProvider` instances reason over conversation messages (which may contain sensitive user data), they should be **trusted by default**. This is also why we allow all plugins to see all state - if a plugin is untrusted, it shouldn't be in the pipeline at all. The whole state dict is passed rather than isolated slices because plugins that handle messages already have access to the full conversation context.
### Addendum (2026-02-17): Provider-scoped hook state and default source IDs
This addendum introduces a **breaking change** that supersedes earlier references in this ADR where hooks received the
entire `session.state` object as their `state` parameter.
#### Hook state contract
- `before_run` and `after_run` now receive a **provider-scoped** mutable state dict.
- The framework passes `session.state.setdefault(provider.source_id, {})` to hook `state`.
- Cross-provider/global inspection remains available through `session.state` on `AgentSession`.
#### Session requirement and fallback behavior
- Provider hooks must use session-backed scoped state; there is no ad-hoc `{}` fallback state.
- If providers run without a caller-supplied session, the framework creates an internal run-scoped `AgentSession` and
passes provider-scoped state from that session.
#### Migration guidance
Migrate provider implementations and samples from nested access to scoped access:
- `state[self.source_id]["key"]``state["key"]`
- `state.setdefault(self.source_id, {})["key"]``state["key"]`
#### DEFAULT_SOURCE_ID standardization
Aligned with and extending [PR #3944](https://github.com/microsoft/agent-framework/pull/3944), all built-in/connector
providers in this surface now define a `DEFAULT_SOURCE_ID` and allow constructor override via `source_id`.
Naming convention:
- snake_case
- close to the provider class name
- history providers may use `*_memory` where differentiation is useful
Defaults introduced by this change:
- `InMemoryHistoryProvider.DEFAULT_SOURCE_ID = "in_memory"`
- `Mem0ContextProvider.DEFAULT_SOURCE_ID = "mem0"`
- `RedisContextProvider.DEFAULT_SOURCE_ID = "redis"`
- `RedisHistoryProvider.DEFAULT_SOURCE_ID = "redis_memory"`
- `AzureAISearchContextProvider.DEFAULT_SOURCE_ID = "azure_ai_search"`
## Comparison to .NET Implementation
The .NET Agent Framework provides equivalent functionality through a different structure. Both implementations achieve the same goals using idioms natural to their respective languages.
@@ -139,10 +139,11 @@ class AzureAISearchContextProvider(BaseContextProvider):
"""
_DEFAULT_SEARCH_CONTEXT_PROMPT: ClassVar[str] = "Use the following context to answer the question:"
DEFAULT_SOURCE_ID: ClassVar[str] = "azure_ai_search"
def __init__(
self,
source_id: str,
source_id: str = DEFAULT_SOURCE_ID,
endpoint: str | None = None,
index_name: str | None = None,
api_key: str | AzureKeyCredential | None = None,
@@ -59,7 +59,7 @@ def mock_search_client_empty() -> AsyncMock:
def _make_provider(**overrides) -> AzureAISearchContextProvider:
"""Create a semantic-mode provider with mocked internals (skips auto-discovery)."""
defaults = {
"source_id": "aisearch",
"source_id": AzureAISearchContextProvider.DEFAULT_SOURCE_ID,
"endpoint": "https://test.search.windows.net",
"index_name": "test-index",
"api_key": "test-key",
@@ -78,7 +78,7 @@ class TestInitSemantic:
def test_valid_init(self) -> None:
provider = _make_provider()
assert provider.source_id == "aisearch"
assert provider.source_id == AzureAISearchContextProvider.DEFAULT_SOURCE_ID
assert provider.endpoint == "https://test.search.windows.net"
assert provider.index_name == "test-index"
assert provider.mode == "semantic"
@@ -182,10 +182,12 @@ class TestBeforeRunSemantic:
input_messages=[Message(role="user", contents=["test query"])],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_search_client.search.assert_awaited_once()
msgs = ctx.context_messages.get("aisearch", [])
msgs = ctx.context_messages.get(provider.source_id, [])
assert len(msgs) >= 2 # context_prompt + at least one result
assert msgs[0].text == provider.context_prompt
@@ -195,10 +197,12 @@ class TestBeforeRunSemantic:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_search_client.search.assert_not_awaited()
assert ctx.context_messages.get("aisearch") is None
assert ctx.context_messages.get(provider.source_id) is None
async def test_no_results_no_messages(self, mock_search_client_empty: AsyncMock) -> None:
provider = _make_provider()
@@ -209,10 +213,12 @@ class TestBeforeRunSemantic:
input_messages=[Message(role="user", contents=["test query"])],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_search_client_empty.search.assert_awaited_once()
assert ctx.context_messages.get("aisearch") is None
assert ctx.context_messages.get(provider.source_id) is None
async def test_context_prompt_prepended(self, mock_search_client: AsyncMock) -> None:
custom_prompt = "Custom search context:"
@@ -224,9 +230,11 @@ class TestBeforeRunSemantic:
input_messages=[Message(role="user", contents=["test query"])],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
msgs = ctx.context_messages["aisearch"]
msgs = ctx.context_messages[provider.source_id]
assert msgs[0].text == custom_prompt
@@ -248,7 +256,9 @@ class TestBeforeRunFiltering:
],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_search_client.search.assert_awaited_once()
call_kwargs = mock_search_client.search.call_args[1]
@@ -265,7 +275,9 @@ class TestBeforeRunFiltering:
input_messages=[Message(role="system", contents=["system prompt"])],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_search_client.search.assert_not_awaited()
+1 -1
View File
@@ -220,7 +220,7 @@ if __name__ == "__main__":
- [Getting Started with Agents](../../samples/02-agents): Basic agent creation and tool usage
- [Chat Client Examples](../../samples/02-agents/chat_client): Direct chat client usage patterns
- [Azure AI Integration](https://github.com/microsoft/agent-framework/tree/main/python/packages/azure-ai): Azure AI integration
- [.NET Workflows Samples](https://github.com/microsoft/agent-framework/tree/main/dotnet/samples/GettingStarted/Workflows): Advanced multi-agent patterns (.NET)
- [.NET Workflows Samples](../../../dotnet/samples/GettingStarted/Workflows): Advanced multi-agent patterns (.NET)
## Agent Framework Documentation
+29 -13
View File
@@ -420,13 +420,18 @@ class BaseAgent(SerializationMixin):
session: The conversation session.
context: The invocation context with response populated.
"""
state = session.state if session else {}
provider_session = session
if provider_session is None and self.context_providers:
provider_session = AgentSession()
for provider in reversed(self.context_providers):
if provider_session is None:
raise RuntimeError("Provider session must be available when context providers are configured.")
await provider.after_run(
agent=self, # type: ignore[arg-type]
session=session, # type: ignore[arg-type]
session=provider_session,
context=context,
state=state,
state=provider_session.state.setdefault(provider.source_id, {}),
)
def as_tool(
@@ -988,10 +993,14 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
and not opts.get("store")
and not (getattr(self.client, "STORES_BY_DEFAULT", False) and opts.get("store") is not False)
):
self.context_providers.append(InMemoryHistoryProvider("memory"))
self.context_providers.append(InMemoryHistoryProvider())
active_session = session
if active_session is None and self.context_providers:
active_session = AgentSession()
session_context, chat_options = await self._prepare_session_and_messages(
session=session,
session=active_session,
input_messages=input_messages,
options=opts,
)
@@ -1018,7 +1027,9 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
# Build options dict from run() options merged with provided options
run_opts: dict[str, Any] = {
"model_id": opts.pop("model_id", None),
"conversation_id": session.service_session_id if session else opts.pop("conversation_id", None),
"conversation_id": active_session.service_session_id
if active_session
else opts.pop("conversation_id", None),
"allow_multiple_tool_calls": opts.pop("allow_multiple_tool_calls", None),
"additional_function_arguments": opts.pop("additional_function_arguments", None),
"frequency_penalty": opts.pop("frequency_penalty", None),
@@ -1046,12 +1057,12 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
# Ensure session is forwarded in kwargs for tool invocation
finalize_kwargs = dict(kwargs)
finalize_kwargs["session"] = session
finalize_kwargs["session"] = active_session
# Filter chat_options from kwargs to prevent duplicate keyword argument
filtered_kwargs = {k: v for k, v in finalize_kwargs.items() if k != "chat_options"}
return {
"session": session,
"session": active_session,
"session_context": session_context,
"input_messages": input_messages,
"session_messages": session_messages,
@@ -1129,23 +1140,28 @@ class RawAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
else:
chat_options = {}
provider_session = session
if provider_session is None and self.context_providers:
provider_session = AgentSession()
session_context = SessionContext(
session_id=session.session_id if session else None,
service_session_id=session.service_session_id if session else None,
session_id=provider_session.session_id if provider_session else None,
service_session_id=provider_session.service_session_id if provider_session else None,
input_messages=input_messages or [],
options=options or {},
)
# Run before_run providers (forward order, skip BaseHistoryProvider with load_messages=False)
state = session.state if session else {}
for provider in self.context_providers:
if isinstance(provider, BaseHistoryProvider) and not provider.load_messages:
continue
if provider_session is None:
raise RuntimeError("Provider session must be available when context providers are configured.")
await provider.before_run(
agent=self, # type: ignore[arg-type]
session=session, # type: ignore[arg-type]
session=provider_session,
context=session_context,
state=state,
state=provider_session.state.setdefault(provider.source_id, {}),
)
# Merge provider-contributed tools into chat_options
@@ -16,7 +16,7 @@ import copy
import uuid
from abc import abstractmethod
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, ClassVar
from ._types import AgentResponse, Message
@@ -310,7 +310,8 @@ class BaseContextProvider:
agent: The agent running this invocation.
session: The current session.
context: The invocation context - add messages/instructions/tools here.
state: The session's mutable state dict.
state: The provider-scoped mutable state dict for this provider.
Full cross-provider state remains available at ``session.state``.
"""
async def after_run(
@@ -330,7 +331,8 @@ class BaseContextProvider:
agent: The agent that ran this invocation.
session: The current session.
context: The invocation context with response populated.
state: The session's mutable state dict.
state: The provider-scoped mutable state dict for this provider.
Full cross-provider state remains available at ``session.state``.
"""
@@ -520,25 +522,56 @@ class AgentSession:
class InMemoryHistoryProvider(BaseHistoryProvider):
"""Built-in history provider that stores messages in session.state.
Messages are stored in ``state[source_id]["messages"]`` as a list of
Messages are stored in ``state["messages"]`` as a list of
``Message`` objects. Serialization to/from dicts is handled by
``AgentSession.to_dict()``/``from_dict()`` using ``SerializationProtocol``.
This provider holds no instance state — all data lives in the session's
state dict, passed as a named ``state`` parameter to ``get_messages``/``save_messages``.
This is the default provider auto-added by the agent when no providers
are configured and ``conversation_id`` or ``store=True`` is set.
This is the default provider auto-added by the agent for local sessions
when no providers are configured and service-side storage is not requested.
"""
DEFAULT_SOURCE_ID: ClassVar[str] = "in_memory"
def __init__(
self,
source_id: str | None = None,
*,
load_messages: bool = True,
store_inputs: bool = True,
store_context_messages: bool = False,
store_context_from: set[str] | None = None,
store_outputs: bool = True,
) -> None:
"""Initialize the in-memory history provider.
Args:
source_id: Unique identifier for this provider instance.
Defaults to DEFAULT_SOURCE_ID when not provided.
load_messages: Whether to load messages before invocation.
store_inputs: Whether to store input messages.
store_context_messages: Whether to store context from other providers.
store_context_from: If set, only store context from these source_ids.
store_outputs: Whether to store response messages.
"""
super().__init__(
source_id=source_id or self.DEFAULT_SOURCE_ID,
load_messages=load_messages,
store_inputs=store_inputs,
store_context_messages=store_context_messages,
store_context_from=store_context_from,
store_outputs=store_outputs,
)
async def get_messages(
self, session_id: str | None, *, state: dict[str, Any] | None = None, **kwargs: Any
) -> list[Message]:
"""Retrieve messages from session state."""
if state is None:
return []
my_state = state.get(self.source_id, {})
return list(my_state.get("messages", []))
return list(state.get("messages", []))
async def save_messages(
self,
@@ -551,6 +584,5 @@ class InMemoryHistoryProvider(BaseHistoryProvider):
"""Persist messages to session state."""
if state is None:
return
my_state = state.setdefault(self.source_id, {})
existing = my_state.get("messages", [])
my_state["messages"] = [*existing, *messages]
existing = state.get("messages", [])
state["messages"] = [*existing, *messages]
@@ -121,7 +121,7 @@ class WorkflowAgent(BaseAgent):
resolved_context_providers = list(context_providers) if context_providers is not None else []
if not resolved_context_providers:
resolved_context_providers.append(InMemoryHistoryProvider("memory"))
resolved_context_providers.append(InMemoryHistoryProvider())
super().__init__(
id=id,
@@ -237,23 +237,27 @@ class WorkflowAgent(BaseAgent):
An AgentResponse representing the workflow execution results.
"""
input_messages = normalize_messages_input(messages)
provider_session = session
if provider_session is None and self.context_providers:
provider_session = AgentSession()
# run the context providers with the session
session_context = SessionContext(
session_id=session.session_id if session else None,
service_session_id=session.service_session_id if session else None,
session_id=provider_session.session_id if provider_session else None,
service_session_id=provider_session.service_session_id if provider_session else None,
input_messages=input_messages or [],
options={},
)
state = session.state if session else {}
for provider in self.context_providers:
if isinstance(provider, BaseHistoryProvider) and not provider.load_messages:
continue
if provider_session is None:
raise RuntimeError("Provider session must be available when context providers are configured.")
await provider.before_run(
agent=self, # type: ignore[arg-type]
session=session, # type: ignore[arg-type]
session=provider_session,
context=session_context,
state=state,
state=provider_session.state.setdefault(provider.source_id, {}),
)
# combine the messages
session_messages: list[Message] = session_context.get_messages(include_input=True)
@@ -266,7 +270,7 @@ class WorkflowAgent(BaseAgent):
output_events.append(event)
result = self._convert_workflow_events_to_agent_response(response_id, output_events)
await self._run_after_providers(session=session, context=session_context)
await self._run_after_providers(session=provider_session, context=session_context)
return result
async def _run_stream_impl(
@@ -293,23 +297,27 @@ class WorkflowAgent(BaseAgent):
AgentResponseUpdate objects representing the workflow execution progress.
"""
input_messages = normalize_messages_input(messages)
provider_session = session
if provider_session is None and self.context_providers:
provider_session = AgentSession()
# run the context providers with the session
session_context = SessionContext(
session_id=session.session_id if session else None,
service_session_id=session.service_session_id if session else None,
session_id=provider_session.session_id if provider_session else None,
service_session_id=provider_session.service_session_id if provider_session else None,
input_messages=input_messages or [],
options={},
)
state = session.state if session else {}
for provider in self.context_providers:
if isinstance(provider, BaseHistoryProvider) and not provider.load_messages:
continue
if provider_session is None:
raise RuntimeError("Provider session must be available when context providers are configured.")
await provider.before_run(
agent=self, # type: ignore[arg-type]
session=session, # type: ignore[arg-type]
session=provider_session,
context=session_context,
state=state,
state=provider_session.state.setdefault(provider.source_id, {}),
)
# combine the messages
@@ -320,7 +328,7 @@ class WorkflowAgent(BaseAgent):
updates = self._convert_workflow_event_to_agent_response_updates(response_id, event)
for update in updates:
yield update
await self._run_after_providers(session=session, context=session_context)
await self._run_after_providers(session=provider_session, context=session_context)
async def _run_core(
self,
@@ -107,10 +107,10 @@ async def test_chat_client_agent_create_session(client: SupportsChatGetResponse)
async def test_chat_client_agent_prepare_session_and_messages(client: SupportsChatGetResponse) -> None:
from agent_framework._sessions import InMemoryHistoryProvider
agent = Agent(client=client, context_providers=[InMemoryHistoryProvider("memory")])
agent = Agent(client=client, context_providers=[InMemoryHistoryProvider()])
message = Message(role="user", text="Hello")
session = AgentSession()
session.state["memory"] = {"messages": [message]}
session.state[InMemoryHistoryProvider.DEFAULT_SOURCE_ID] = {"messages": [message]}
session_context, _ = await agent._prepare_session_and_messages( # type: ignore[reportPrivateUsage]
session=session,
@@ -267,6 +267,8 @@ async def test_chat_client_agent_update_session_id_streaming_does_not_use_respon
async def test_chat_client_agent_update_session_messages(client: SupportsChatGetResponse) -> None:
from agent_framework._sessions import InMemoryHistoryProvider
agent = Agent(client=client)
session = agent.create_session()
@@ -275,7 +277,7 @@ async def test_chat_client_agent_update_session_messages(client: SupportsChatGet
assert session.service_session_id is None
chat_messages: list[Message] = session.state.get("memory", {}).get("messages", [])
chat_messages: list[Message] = session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {}).get("messages", [])
assert chat_messages is not None
assert len(chat_messages) == 2
@@ -27,6 +27,7 @@ from agent_framework import (
chat_middleware,
function_middleware,
)
from agent_framework._sessions import InMemoryHistoryProvider
from .conftest import MockBaseChatClient, MockChatClient
@@ -1416,8 +1417,10 @@ class TestChatAgentSessionBehavior:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Capture state before next() call
thread_messages = []
if context.session and context.session.state.get("memory"):
thread_messages = context.session.state.get("memory", {}).get("messages", [])
if context.session and context.session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID):
thread_messages = context.session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {}).get(
"messages", []
)
before_state = {
"before_next": True,
@@ -1432,8 +1435,10 @@ class TestChatAgentSessionBehavior:
# Capture state after next() call
thread_messages_after = []
if context.session and context.session.state.get("memory"):
thread_messages_after = context.session.state.get("memory", {}).get("messages", [])
if context.session and context.session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID):
thread_messages_after = context.session.state.get(
InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {}
).get("messages", [])
after_state = {
"before_next": False,
@@ -359,30 +359,50 @@ class TestAgentSession:
class TestInMemoryHistoryProvider:
async def test_empty_state_returns_no_messages(self) -> None:
provider = InMemoryHistoryProvider("memory")
provider = InMemoryHistoryProvider()
session = AgentSession()
ctx = SessionContext(session_id="s1", input_messages=[])
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
assert ctx.context_messages.get("memory", []) == []
await provider.before_run( # type: ignore[arg-type]
agent=None,
session=session,
context=ctx,
state=session.state.setdefault(provider.source_id, {}),
)
assert ctx.context_messages.get(provider.source_id, []) == []
async def test_stores_and_loads_messages(self) -> None:
from agent_framework import AgentResponse
provider = InMemoryHistoryProvider("memory")
provider = InMemoryHistoryProvider()
session = AgentSession()
# First run: send input, get response
input_msg = Message(role="user", contents=["hello"])
resp_msg = Message(role="assistant", contents=["hi there"])
ctx1 = SessionContext(session_id="s1", input_messages=[input_msg])
await provider.before_run(agent=None, session=session, context=ctx1, state=session.state) # type: ignore[arg-type]
await provider.before_run( # type: ignore[arg-type]
agent=None,
session=session,
context=ctx1,
state=session.state.setdefault(provider.source_id, {}),
)
ctx1._response = AgentResponse(messages=[resp_msg])
await provider.after_run(agent=None, session=session, context=ctx1, state=session.state) # type: ignore[arg-type]
await provider.after_run( # type: ignore[arg-type]
agent=None,
session=session,
context=ctx1,
state=session.state.setdefault(provider.source_id, {}),
)
# Second run: should load previous messages
ctx2 = SessionContext(session_id="s1", input_messages=[Message(role="user", contents=["again"])])
await provider.before_run(agent=None, session=session, context=ctx2, state=session.state) # type: ignore[arg-type]
loaded = ctx2.context_messages.get("memory", [])
await provider.before_run( # type: ignore[arg-type]
agent=None,
session=session,
context=ctx2,
state=session.state.setdefault(provider.source_id, {}),
)
loaded = ctx2.context_messages.get(provider.source_id, [])
assert len(loaded) == 2
assert loaded[0].text == "hello"
assert loaded[1].text == "hi there"
@@ -390,17 +410,27 @@ class TestInMemoryHistoryProvider:
async def test_state_is_serializable(self) -> None:
from agent_framework import AgentResponse
provider = InMemoryHistoryProvider("memory")
provider = InMemoryHistoryProvider()
session = AgentSession()
input_msg = Message(role="user", contents=["test"])
ctx = SessionContext(session_id="s1", input_messages=[input_msg])
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run( # type: ignore[arg-type]
agent=None,
session=session,
context=ctx,
state=session.state.setdefault(provider.source_id, {}),
)
ctx._response = AgentResponse(messages=[Message(role="assistant", contents=["reply"])])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run( # type: ignore[arg-type]
agent=None,
session=session,
context=ctx,
state=session.state.setdefault(provider.source_id, {}),
)
# State contains Message objects (not dicts)
assert isinstance(session.state["memory"]["messages"][0], Message)
assert isinstance(session.state[provider.source_id]["messages"][0], Message)
# to_dict() serializes them via SerializationProtocol
session_dict = session.to_dict()
@@ -409,9 +439,9 @@ class TestInMemoryHistoryProvider:
# Round-trip through session serialization restores Message objects
restored = AgentSession.from_dict(json.loads(json_str))
assert isinstance(restored.state["memory"]["messages"][0], Message)
assert restored.state["memory"]["messages"][0].text == "test"
assert restored.state["memory"]["messages"][1].text == "reply"
assert isinstance(restored.state[provider.source_id]["messages"][0], Message)
assert restored.state[provider.source_id]["messages"][0].text == "test"
assert restored.state[provider.source_id]["messages"][1].text == "reply"
async def test_source_id_attribution(self) -> None:
provider = InMemoryHistoryProvider("custom-source")
@@ -18,7 +18,7 @@ class SlidingWindowHistoryProvider(InMemoryHistoryProvider):
def __init__(
self,
source_id: str = "memory",
source_id: str = InMemoryHistoryProvider.DEFAULT_SOURCE_ID,
*,
max_tokens: int = 3800,
system_message: str | None = None,
@@ -12,6 +12,7 @@ from agent_framework import (
AgentExecutorResponse,
AgentResponse,
FunctionExecutor,
InMemoryHistoryProvider,
Message,
SupportsChatGetResponse,
Workflow,
@@ -359,7 +360,9 @@ class TaskRunner:
# 2. The assistant's session state (full history, not just the truncated window)
# 3. The final user message (if any)
session_state: dict[str, Any] = self._assistant_executor._session.state # type: ignore
all_messages: list[Message] = list(session_state.get("memory", {}).get("messages", [])) # type: ignore
all_messages: list[Message] = list(
session_state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {}).get("messages", [])
) # type: ignore
full_conversation = [first_message, *all_messages]
if self._final_user_message is not None:
full_conversation.extend(self._final_user_message)
@@ -4,6 +4,7 @@
from unittest.mock import patch
from agent_framework import InMemoryHistoryProvider
from agent_framework._types import Content, Message
from agent_framework_lab_tau2._sliding_window import SlidingWindowHistoryProvider
@@ -12,7 +13,7 @@ def _make_state(provider: SlidingWindowHistoryProvider, messages: list[Message]
"""Helper to create a session state dict with messages pre-loaded."""
state: dict = {}
if messages:
state[provider.source_id] = {"messages": list(messages)}
state["messages"] = list(messages)
return state
@@ -27,7 +28,7 @@ def test_initialization():
assert provider.max_tokens == 2000
assert provider.system_message == "You are a helpful assistant"
assert provider.tool_definitions == [{"name": "test_tool"}]
assert provider.source_id == "memory"
assert provider.source_id == InMemoryHistoryProvider.DEFAULT_SOURCE_ID
async def test_get_messages_empty():
@@ -66,7 +67,7 @@ async def test_save_and_get_messages():
# get_messages returns truncated
truncated = await provider.get_messages(None, state=state)
# Full history is in session state
all_msgs = state[provider.source_id]["messages"]
all_msgs = state["messages"]
assert len(all_msgs) == 10
assert len(truncated) < len(all_msgs)
@@ -196,7 +197,7 @@ async def test_real_world_scenario():
await provider.save_messages(None, conversation, state=state)
truncated = await provider.get_messages(None, state=state)
all_msgs = state[provider.source_id]["messages"]
all_msgs = state["messages"]
assert len(all_msgs) == 6
assert len(truncated) <= 6
@@ -10,7 +10,7 @@ from __future__ import annotations
import sys
from contextlib import AbstractAsyncContextManager
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, ClassVar
from agent_framework import Message
from agent_framework._sessions import AgentSession, BaseContextProvider, SessionContext
@@ -42,10 +42,11 @@ class Mem0ContextProvider(BaseContextProvider):
"""
DEFAULT_CONTEXT_PROMPT = "## Memories\nConsider the following memories when answering user questions:"
DEFAULT_SOURCE_ID: ClassVar[str] = "mem0"
def __init__(
self,
source_id: str,
source_id: str = DEFAULT_SOURCE_ID,
mem0_client: AsyncMemory | AsyncMemoryClient | None = None,
api_key: str | None = None,
application_id: str | None = None,
@@ -97,7 +97,9 @@ class TestBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="Hello")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_mem0_client.search.assert_awaited_once()
assert "mem0" in ctx.context_messages
@@ -113,7 +115,9 @@ class TestBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_mem0_client.search.assert_not_awaited()
assert "mem0" not in ctx.context_messages
@@ -125,7 +129,9 @@ class TestBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="test")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
assert "mem0" not in ctx.context_messages
@@ -136,7 +142,9 @@ class TestBeforeRun:
ctx = SessionContext(input_messages=[Message(role="user", text="test")], session_id="s1")
with pytest.raises(ServiceInitializationError, match="At least one of the filters"):
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
async def test_v1_1_response_format(self, mock_mem0_client: AsyncMock) -> None:
"""Search response in v1.1 dict format with 'results' key."""
@@ -145,7 +153,9 @@ class TestBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="test")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
added = ctx.context_messages["mem0"]
assert "remembered fact" in added[0].text # type: ignore[operator]
@@ -163,7 +173,9 @@ class TestBeforeRun:
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
call_kwargs = mock_mem0_client.search.call_args.kwargs
assert call_kwargs["query"] == "Hello\nWorld"
@@ -175,7 +187,9 @@ class TestBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="Hello")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
call_kwargs = mock_oss_mem0_client.search.call_args.kwargs
assert call_kwargs["query"] == "Hello"
@@ -191,7 +205,9 @@ class TestBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="Hello")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
call_kwargs = mock_oss_mem0_client.search.call_args.kwargs
assert call_kwargs["user_id"] == "u1"
@@ -205,7 +221,9 @@ class TestBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="Hello")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
call_kwargs = mock_mem0_client.search.call_args.kwargs
assert call_kwargs["query"] == "Hello"
@@ -226,7 +244,9 @@ class TestAfterRun:
ctx = SessionContext(input_messages=[Message(role="user", text="question")], session_id="s1")
ctx._response = AgentResponse(messages=[Message(role="assistant", text="answer")])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_mem0_client.add.assert_awaited_once()
call_kwargs = mock_mem0_client.add.call_args.kwargs
@@ -250,7 +270,9 @@ class TestAfterRun:
)
ctx._response = AgentResponse(messages=[Message(role="assistant", text="reply")])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
call_kwargs = mock_mem0_client.add.call_args.kwargs
roles = [m["role"] for m in call_kwargs["messages"]]
@@ -270,7 +292,9 @@ class TestAfterRun:
)
ctx._response = AgentResponse(messages=[])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_mem0_client.add.assert_not_awaited()
@@ -281,7 +305,9 @@ class TestAfterRun:
ctx = SessionContext(input_messages=[Message(role="user", text="hi")], session_id="my-session")
ctx._response = AgentResponse(messages=[Message(role="assistant", text="hey")])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
assert mock_mem0_client.add.call_args.kwargs["run_id"] == "my-session"
@@ -293,7 +319,9 @@ class TestAfterRun:
ctx._response = AgentResponse(messages=[Message(role="assistant", text="hey")])
with pytest.raises(ServiceInitializationError, match="At least one of the filters"):
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
async def test_stores_with_application_id_metadata(self, mock_mem0_client: AsyncMock) -> None:
"""application_id is passed in metadata."""
@@ -304,7 +332,9 @@ class TestAfterRun:
ctx = SessionContext(input_messages=[Message(role="user", text="hi")], session_id="s1")
ctx._response = AgentResponse(messages=[])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
assert mock_mem0_client.add.call_args.kwargs["metadata"] == {"application_id": "app1"}
@@ -12,7 +12,7 @@ import json
import sys
from functools import reduce
from operator import and_
from typing import TYPE_CHECKING, Any, Literal, cast
from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast
import numpy as np
from agent_framework import Message
@@ -50,10 +50,11 @@ class RedisContextProvider(BaseContextProvider):
"""
DEFAULT_CONTEXT_PROMPT = "## Memories\nConsider the following memories when answering user questions:"
DEFAULT_SOURCE_ID: ClassVar[str] = "redis"
def __init__(
self,
source_id: str,
source_id: str = DEFAULT_SOURCE_ID,
redis_url: str = "redis://localhost:6379",
index_name: str = "context",
prefix: str = "context",
@@ -9,7 +9,7 @@ This module provides ``RedisHistoryProvider``, built on the new
from __future__ import annotations
from collections.abc import Sequence
from typing import Any
from typing import Any, ClassVar
import redis.asyncio as redis
from agent_framework import Message
@@ -24,9 +24,11 @@ class RedisHistoryProvider(BaseHistoryProvider):
unique Redis key.
"""
DEFAULT_SOURCE_ID: ClassVar[str] = "redis_memory"
def __init__(
self,
source_id: str,
source_id: str = DEFAULT_SOURCE_ID,
redis_url: str | None = None,
credential_provider: CredentialProvider | None = None,
host: str | None = None,
+27 -9
View File
@@ -144,7 +144,9 @@ class TestRedisContextProviderBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=["test query"])], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
assert "ctx" in ctx.context_messages
msgs = ctx.context_messages["ctx"]
@@ -161,7 +163,9 @@ class TestRedisContextProviderBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=[" "])], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_index.query.assert_not_called()
assert "ctx" not in ctx.context_messages
@@ -176,7 +180,9 @@ class TestRedisContextProviderBeforeRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=["hello"])], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
assert "ctx" not in ctx.context_messages
@@ -193,7 +199,9 @@ class TestRedisContextProviderAfterRun:
ctx = SessionContext(input_messages=[Message(role="user", contents=["user input"])], session_id="s1")
ctx._response = response
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_index.load.assert_called_once()
loaded = mock_index.load.call_args[0][0]
@@ -210,7 +218,9 @@ class TestRedisContextProviderAfterRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=[" "])], session_id="s1")
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_index.load.assert_not_called()
@@ -223,7 +233,9 @@ class TestRedisContextProviderAfterRun:
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", contents=["hello"])], session_id="s1")
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
loaded = mock_index.load.call_args[0][0]
doc = loaded[0]
@@ -419,7 +431,9 @@ class TestRedisHistoryProviderBeforeAfterRun:
session = AgentSession(session_id="test")
ctx = SessionContext(input_messages=[Message(role="user", contents=["new msg"])], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.before_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
assert "mem" in ctx.context_messages
assert len(ctx.context_messages["mem"]) == 1
@@ -434,7 +448,9 @@ class TestRedisHistoryProviderBeforeAfterRun:
ctx = SessionContext(input_messages=[Message(role="user", contents=["hi"])], session_id="s1")
ctx._response = AgentResponse(messages=[Message(role="assistant", contents=["hello"])])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
pipeline = mock_redis_client.pipeline.return_value.__aenter__.return_value
assert pipeline.rpush.call_count == 2
@@ -450,6 +466,8 @@ class TestRedisHistoryProviderBeforeAfterRun:
session = AgentSession(session_id="test")
ctx = SessionContext(input_messages=[Message(role="user", contents=["hi"])], session_id="s1")
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
await provider.after_run(
agent=None, session=session, context=ctx, state=session.state.setdefault(provider.source_id, {})
) # type: ignore[arg-type]
mock_redis_client.pipeline.assert_not_called()
@@ -13,6 +13,7 @@ from agent_framework import (
ChatResponseUpdate,
Content,
FunctionInvocationLayer,
InMemoryHistoryProvider,
Message,
ResponseStream,
)
@@ -195,7 +196,7 @@ async def main() -> None:
print(f"Agent: {result.messages[0].text}\n")
# Check conversation history
memory_state = session.state.get("memory", {})
memory_state = session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
session_messages = memory_state.get("messages", [])
if session_messages:
print(f"Session contains {len(session_messages)} messages")
@@ -17,7 +17,9 @@ class UserInfo(BaseModel):
class UserInfoMemory(BaseContextProvider):
def __init__(self, source_id: str = "user-info-memory", *, client: SupportsChatGetResponse, **kwargs: Any):
DEFAULT_SOURCE_ID = "user_info_memory"
def __init__(self, source_id: str = DEFAULT_SOURCE_ID, *, client: SupportsChatGetResponse, **kwargs: Any):
"""Create the memory.
If you pass in kwargs, they will be attempted to be used to create a UserInfo object.
@@ -39,9 +41,7 @@ class UserInfoMemory(BaseContextProvider):
# Check if we need to extract user info from user messages
user_messages = [msg for msg in request_messages if hasattr(msg, "role") and msg.role == "user"] # type: ignore
if (
state[self.source_id]["user_info"].name is None or state[self.source_id]["user_info"].age is None
) and user_messages:
if (state["user_info"].name is None or state["user_info"].age is None) and user_messages:
with suppress(Exception):
# Use the chat client to extract structured information
result = await self._chat_client.get_response(
@@ -54,10 +54,10 @@ class UserInfoMemory(BaseContextProvider):
# Update user info with extracted data
with suppress(Exception):
extracted = result.value
if state[self.source_id]["user_info"].name is None and extracted.name:
state[self.source_id]["user_info"].name = extracted.name
if state[self.source_id]["user_info"].age is None and extracted.age:
state[self.source_id]["user_info"].age = extracted.age
if state["user_info"].name is None and extracted.name:
state["user_info"].name = extracted.name
if state["user_info"].age is None and extracted.age:
state["user_info"].age = extracted.age
async def before_run(
self,
@@ -68,20 +68,19 @@ class UserInfoMemory(BaseContextProvider):
state: dict[str, Any],
) -> None:
"""Provide user information context before each agent call."""
if state.setdefault(self.source_id, None) is None:
state[self.source_id] = {"user_info": UserInfo()}
state.setdefault("user_info", UserInfo())
context.extend_instructions(
self.source_id,
"Ask the user for their name and politely decline to answer any questions until they provide it."
if state[self.source_id]["user_info"].name is None
else f"The user's name is {state[self.source_id]['user_info'].name}.",
if state["user_info"].name is None
else f"The user's name is {state['user_info'].name}.",
)
context.extend_instructions(
self.source_id,
"Ask the user for their age and politely decline to answer any questions until they provide it."
if state[self.source_id]["user_info"].age is None
else f"The user's age is {state[self.source_id]['user_info'].age}.",
if state["user_info"].age is None
else f"The user's age is {state['user_info'].age}.",
)
@@ -92,7 +91,7 @@ async def main():
credential=AzureCliCredential(),
)
context_name = "user-info-memory"
context_name = UserInfoMemory.DEFAULT_SOURCE_ID
# Create the memory provider
memory_provider = UserInfoMemory(context_name, client=client)
@@ -6,6 +6,7 @@ from typing import Annotated
from agent_framework import (
AgentContext,
InMemoryHistoryProvider,
tool,
)
from agent_framework.azure import AzureOpenAIChatClient
@@ -50,7 +51,7 @@ async def thread_tracking_middleware(
"""MiddlewareTypes that tracks and logs session behavior across runs."""
session_message_count = 0
if context.session:
memory_state = context.session.state.get("memory", {})
memory_state = context.session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
session_message_count = len(memory_state.get("messages", []))
print(f"[MiddlewareTypes pre-execution] Current input messages: {len(context.messages)}")
@@ -62,7 +63,7 @@ async def thread_tracking_middleware(
# Check session state after agent execution
updated_session_message_count = 0
if context.session:
memory_state = context.session.state.get("memory", {})
memory_state = context.session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
updated_session_message_count = len(memory_state.get("messages", []))
print(f"[MiddlewareTypes post-execution] Updated session messages: {updated_session_message_count}")
@@ -4,7 +4,7 @@ import asyncio
from random import randint
from typing import Annotated
from agent_framework import Agent, AgentSession, tool
from agent_framework import Agent, AgentSession, InMemoryHistoryProvider, tool
from agent_framework.azure import AzureOpenAIChatClient
from azure.identity import AzureCliCredential
from pydantic import Field
@@ -112,7 +112,7 @@ async def example_with_existing_session_messages() -> None:
print(f"Agent: {result1.text}")
# The session now contains the conversation history in state
memory_state = session.state.get("memory", {})
memory_state = session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
messages = memory_state.get("messages", [])
if messages:
print(f"Session contains {len(messages)} messages")
@@ -10,6 +10,7 @@ from agent_framework import (
AgentSession,
BaseAgent,
Content,
InMemoryHistoryProvider,
Message,
Role,
normalize_messages,
@@ -93,7 +94,9 @@ class EchoAgent(BaseAgent):
if not normalized_messages:
response_message = Message(
role=Role.ASSISTANT,
contents=[Content.from_text(text="Hello! I'm a custom echo agent. Send me a message and I'll echo it back.")],
contents=[
Content.from_text(text="Hello! I'm a custom echo agent. Send me a message and I'll echo it back.")
],
)
else:
# For simplicity, echo the last user message
@@ -199,7 +202,7 @@ async def main() -> None:
print(f"Agent: {result2.messages[0].text}")
# Check conversation history
memory_state = session.state.get("memory", {})
memory_state = session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
messages = memory_state.get("messages", [])
if messages:
print(f"\nSession contains {len(messages)} messages in history")
@@ -4,7 +4,7 @@ import asyncio
from random import randint
from typing import Annotated
from agent_framework import Agent, AgentSession, tool
from agent_framework import Agent, AgentSession, InMemoryHistoryProvider, tool
from agent_framework.openai import OpenAIChatClient
from pydantic import Field
@@ -105,7 +105,7 @@ async def example_with_existing_session_messages() -> None:
print(f"Agent: {result1.text}")
# The session now contains the conversation history in state
memory_state = session.state.get("memory", {})
memory_state = session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
messages = memory_state.get("messages", [])
if messages:
print(f"Session contains {len(messages)} messages")
@@ -59,17 +59,17 @@ async def main() -> None:
credential=AzureCliCredential(),
)
# set the same context provider, with the same source_id, for both agents to share the thread
# set the same context provider (same default source_id) for both agents to share the thread
writer = client.as_agent(
instructions=("You are a concise copywriter. Provide a single, punchy marketing sentence based on the prompt."),
name="writer",
context_providers=[InMemoryHistoryProvider("memory")],
context_providers=[InMemoryHistoryProvider()],
)
reviewer = client.as_agent(
instructions=("You are a thoughtful reviewer. Give brief feedback on the previous assistant message."),
name="reviewer",
context_providers=[InMemoryHistoryProvider("memory")],
context_providers=[InMemoryHistoryProvider()],
)
# Create the shared session
@@ -96,7 +96,7 @@ async def main() -> None:
# The shared session now contains the conversation between the writer and reviewer. Print it out.
print("=== Shared Session Conversation ===")
memory_state = shared_session.state.get("memory", {})
memory_state = shared_session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
for message in memory_state.get("messages", []):
print(f"{message.author_name or message.role}: {message.text}")
@@ -3,7 +3,7 @@
import asyncio
import os
from agent_framework import AgentSession
from agent_framework import AgentSession, InMemoryHistoryProvider
from agent_framework.azure import AzureOpenAIResponsesClient
from agent_framework.orchestrations import SequentialBuilder
from azure.identity import AzureCliCredential
@@ -109,7 +109,7 @@ async def main() -> None:
print("\n" + "=" * 60)
print("Full Session History")
print("=" * 60)
memory_state = session.state.get("memory", {})
memory_state = session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
history = memory_state.get("messages", [])
for i, msg in enumerate(history, start=1):
role = msg.role if hasattr(msg.role, "value") else str(msg.role)
@@ -29,6 +29,7 @@ import os
from agent_framework import (
InMemoryCheckpointStorage,
InMemoryHistoryProvider,
)
from agent_framework.azure import AzureOpenAIResponsesClient
from agent_framework.orchestrations import SequentialBuilder
@@ -122,7 +123,7 @@ async def checkpointing_with_thread() -> None:
checkpoints = await checkpoint_storage.list_checkpoints(workflow_name=workflow.name)
print(f"\nTotal checkpoints across both turns: {len(checkpoints)}")
memory_state = session.state.get("memory", {})
memory_state = session.state.get(InMemoryHistoryProvider.DEFAULT_SOURCE_ID, {})
history = memory_state.get("messages", [])
print(f"Messages in session history: {len(history)}")
@@ -23,12 +23,14 @@ class UserInfo(BaseModel):
class UserInfoMemory(BaseContextProvider):
"""Context provider that extracts and remembers user info (name, age).
State is stored in ``session.state["user-info-memory"]`` so it survives
State is stored in ``session.state["user_info_memory"]`` so it survives
serialization via ``session.to_dict()`` / ``AgentSession.from_dict()``.
"""
DEFAULT_SOURCE_ID = "user_info_memory"
def __init__(self, client: SupportsChatGetResponse):
super().__init__("user-info-memory")
super().__init__(self.DEFAULT_SOURCE_ID)
self._chat_client = client
async def before_run(
@@ -40,8 +42,7 @@ class UserInfoMemory(BaseContextProvider):
state: dict[str, Any],
) -> None:
"""Provide user information context before each agent call."""
my_state = state.setdefault(self.source_id, {})
user_info = my_state.setdefault("user_info", UserInfo())
user_info = state.setdefault("user_info", UserInfo())
instructions: list[str] = []
@@ -70,8 +71,7 @@ class UserInfoMemory(BaseContextProvider):
state: dict[str, Any],
) -> None:
"""Extract user information from messages after each agent call."""
my_state = state.setdefault(self.source_id, {})
user_info = my_state.setdefault("user_info", UserInfo())
user_info = state.setdefault("user_info", UserInfo())
if user_info.name is not None and user_info.age is not None:
return # Already have everything
@@ -92,7 +92,7 @@ class UserInfoMemory(BaseContextProvider):
user_info.name = extracted.name
if extracted and user_info.age is None and extracted.age:
user_info.age = extracted.age
state.setdefault(self.source_id, {})["user_info"] = user_info
state["user_info"] = user_info
except Exception:
pass # Failed to extract, continue without updating
@@ -113,7 +113,7 @@ async def main():
print(await agent.run("I am 20 years old", session=session))
# Inspect extracted user info from session state
user_info = session.state.get("user-info-memory", {}).get("user_info", UserInfo())
user_info = session.state.get(UserInfoMemory.DEFAULT_SOURCE_ID, {}).get("user_info", UserInfo())
print()
print(f"MEMORY - User Name: {user_info.name}")
print(f"MEMORY - User Age: {user_info.age}")