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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user