mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
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:
committed by
GitHub
Unverified
parent
a5f948c215
commit
cc98d5b6f7
@@ -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.
|
||||
|
||||
+2
-1
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user