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
@@ -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"}