mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING] Simplify API: ChatAgent -> Agent, ChatMessage -> Message (#3747)
* [BREAKING] Rename ChatAgent -> Agent, ChatMessage -> Message, ChatClientProtocol -> SupportsChatGetResponse Simplify the public API by removing redundant 'Chat' prefix from core types: - ChatAgent -> Agent - RawChatAgent -> RawAgent - ChatMessage -> Message - ChatClientProtocol -> SupportsChatGetResponse Also renamed internal WorkflowMessage (was Message in _runner_context) to avoid collision. No backward compatibility aliases - this is a clean breaking change. * [BREAKING] Rename Agent chat_client parameter to client * Fix rebase issues: WorkflowMessage references and broken markdown links * Fix formatting and lint issues from code quality checks * Fix import ordering in workflow sample files * fixed rebase * Fix test failures: use WorkflowMessage and A2AMessage after ChatMessage→Message rename - Replace Message(data=..., source_id=...) with WorkflowMessage(...) in workflow tests - Fix isinstance check in A2A agent to use A2AMessage instead of Message - Fix import in test_workflow_observability.py (Message→WorkflowMessage) * Fix lint, fmt, and sample errors after ChatMessage→Message rename - Auto-fix 70+ ruff lint issues across samples (ChatMessage→Message refs) - Fix HostedVectorStoreContent→Content.from_hosted_vector_store in file search sample - Fix _normalize_messages→normalize_messages in custom agent sample - Fix context.terminate→raise MiddlewareTermination in middleware samples - Fix with_update_hook→with_transform_hook in override middleware sample - Add TOptions_co import back to custom_chat_client sample - Add noqa for FastAPI File() default in chatkit sample - Fix B023 loop variable capture in weather agent sample * fix: update Agent constructor calls from chat_client to client in declaration-only tool tests * fix: add register_cleanup to devui lazy-loading proxy and type stub * fixed tests and updated new pieces * fix agui typevar * fix merge errors * fix merge conflicts * fiux merge * Remove unused links --------- Co-authored-by: Evan Mattson <evan.mattson@microsoft.com>
This commit is contained in:
committed by
GitHub
Unverified
parent
a4c9e43afb
commit
0521f5bed8
@@ -12,7 +12,7 @@ Integration with Mem0 for agent memory management.
|
||||
from agent_framework.mem0 import Mem0Provider
|
||||
|
||||
provider = Mem0Provider(api_key="your-key")
|
||||
agent = ChatAgent(..., context_provider=provider)
|
||||
agent = Agent(..., context_provider=provider)
|
||||
```
|
||||
|
||||
## Import Path
|
||||
|
||||
@@ -13,7 +13,7 @@ import sys
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework import Message
|
||||
from agent_framework._sessions import AgentSession, BaseContextProvider, SessionContext
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
from mem0 import AsyncMemory, AsyncMemoryClient
|
||||
@@ -131,7 +131,7 @@ class _Mem0ContextProvider(BaseContextProvider):
|
||||
if line_separated_memories:
|
||||
context.extend_messages(
|
||||
self.source_id,
|
||||
[ChatMessage(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")],
|
||||
[Message(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")],
|
||||
)
|
||||
|
||||
async def after_run(
|
||||
@@ -145,7 +145,7 @@ class _Mem0ContextProvider(BaseContextProvider):
|
||||
"""Store request/response messages to Mem0 for future retrieval."""
|
||||
self._validate_filters()
|
||||
|
||||
messages_to_store: list[ChatMessage] = list(context.input_messages)
|
||||
messages_to_store: list[Message] = list(context.input_messages)
|
||||
if context.response and context.response.messages:
|
||||
messages_to_store.extend(context.response.messages)
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ from collections.abc import MutableSequence, Sequence
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from typing import Any
|
||||
|
||||
from agent_framework import ChatMessage, Context, ContextProvider
|
||||
from agent_framework import Context, ContextProvider, Message
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
from mem0 import AsyncMemory, AsyncMemoryClient
|
||||
|
||||
@@ -103,19 +103,17 @@ class Mem0Provider(ContextProvider):
|
||||
@override
|
||||
async def invoked(
|
||||
self,
|
||||
request_messages: ChatMessage | Sequence[ChatMessage],
|
||||
response_messages: ChatMessage | Sequence[ChatMessage] | None = None,
|
||||
request_messages: Message | Sequence[Message],
|
||||
response_messages: Message | Sequence[Message] | None = None,
|
||||
invoke_exception: Exception | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
self._validate_filters()
|
||||
|
||||
request_messages_list = (
|
||||
[request_messages] if isinstance(request_messages, ChatMessage) else list(request_messages)
|
||||
)
|
||||
request_messages_list = [request_messages] if isinstance(request_messages, Message) else list(request_messages)
|
||||
response_messages_list = (
|
||||
[response_messages]
|
||||
if isinstance(response_messages, ChatMessage)
|
||||
if isinstance(response_messages, Message)
|
||||
else list(response_messages)
|
||||
if response_messages
|
||||
else []
|
||||
@@ -142,7 +140,7 @@ class Mem0Provider(ContextProvider):
|
||||
)
|
||||
|
||||
@override
|
||||
async def invoking(self, messages: ChatMessage | MutableSequence[ChatMessage], **kwargs: Any) -> Context:
|
||||
async def invoking(self, messages: Message | MutableSequence[Message], **kwargs: Any) -> Context:
|
||||
"""Called before invoking the AI model to provide context.
|
||||
|
||||
Args:
|
||||
@@ -155,7 +153,7 @@ class Mem0Provider(ContextProvider):
|
||||
Context: Context object containing instructions with memories.
|
||||
"""
|
||||
self._validate_filters()
|
||||
messages_list = [messages] if isinstance(messages, ChatMessage) else list(messages)
|
||||
messages_list = [messages] if isinstance(messages, Message) else list(messages)
|
||||
input_text = "\n".join(msg.text for msg in messages_list if msg and msg.text and msg.text.strip())
|
||||
|
||||
# Validate input text is not empty before searching (possible for function approval responses)
|
||||
@@ -182,7 +180,7 @@ class Mem0Provider(ContextProvider):
|
||||
line_separated_memories = "\n".join(memory.get("memory", "") for memory in memories)
|
||||
|
||||
return Context(
|
||||
messages=[ChatMessage(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")]
|
||||
messages=[Message(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")]
|
||||
if line_separated_memories
|
||||
else None
|
||||
)
|
||||
|
||||
@@ -7,7 +7,7 @@ import sys
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatMessage, Content, Context
|
||||
from agent_framework import Content, Context, Message
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
from agent_framework.mem0 import Mem0Provider
|
||||
|
||||
@@ -33,12 +33,12 @@ def mock_mem0_client() -> AsyncMock:
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_messages() -> list[ChatMessage]:
|
||||
def sample_messages() -> list[Message]:
|
||||
"""Create sample chat messages for testing."""
|
||||
return [
|
||||
ChatMessage(role="user", text="Hello, how are you?"),
|
||||
ChatMessage(role="assistant", text="I'm doing well, thank you!"),
|
||||
ChatMessage(role="system", text="You are a helpful assistant"),
|
||||
Message(role="user", text="Hello, how are you?"),
|
||||
Message(role="assistant", text="I'm doing well, thank you!"),
|
||||
Message(role="system", text="You are a helpful assistant"),
|
||||
]
|
||||
|
||||
|
||||
@@ -157,7 +157,7 @@ class TestMem0ProviderMessagesAdding:
|
||||
async def test_messages_adding_fails_without_filters(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Test that invoked fails when no filters are provided."""
|
||||
provider = Mem0Provider(mem0_client=mock_mem0_client)
|
||||
message = ChatMessage(role="user", text="Hello!")
|
||||
message = Message(role="user", text="Hello!")
|
||||
|
||||
with pytest.raises(ServiceInitializationError) as exc_info:
|
||||
await provider.invoked(message)
|
||||
@@ -167,7 +167,7 @@ class TestMem0ProviderMessagesAdding:
|
||||
async def test_messages_adding_single_message(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Test adding a single message."""
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
message = ChatMessage(role="user", text="Hello!")
|
||||
message = Message(role="user", text="Hello!")
|
||||
|
||||
await provider.invoked(message)
|
||||
|
||||
@@ -177,7 +177,7 @@ class TestMem0ProviderMessagesAdding:
|
||||
assert call_args.kwargs["user_id"] == "user123"
|
||||
|
||||
async def test_messages_adding_multiple_messages(
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[ChatMessage]
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
|
||||
) -> None:
|
||||
"""Test adding multiple messages."""
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
@@ -194,7 +194,7 @@ class TestMem0ProviderMessagesAdding:
|
||||
assert call_args.kwargs["messages"] == expected_messages
|
||||
|
||||
async def test_messages_adding_with_agent_id(
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[ChatMessage]
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
|
||||
) -> None:
|
||||
"""Test adding messages with agent_id."""
|
||||
provider = Mem0Provider(agent_id="agent123", mem0_client=mock_mem0_client)
|
||||
@@ -206,7 +206,7 @@ class TestMem0ProviderMessagesAdding:
|
||||
assert call_args.kwargs["user_id"] is None
|
||||
|
||||
async def test_messages_adding_with_application_id(
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[ChatMessage]
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
|
||||
) -> None:
|
||||
"""Test adding messages with application_id in metadata."""
|
||||
provider = Mem0Provider(user_id="user123", application_id="app123", mem0_client=mock_mem0_client)
|
||||
@@ -217,7 +217,7 @@ class TestMem0ProviderMessagesAdding:
|
||||
assert call_args.kwargs["metadata"] == {"application_id": "app123"}
|
||||
|
||||
async def test_messages_adding_with_scope_to_per_operation_thread_id(
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[ChatMessage]
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
|
||||
) -> None:
|
||||
"""Test adding messages with scope_to_per_operation_thread_id enabled."""
|
||||
provider = Mem0Provider(
|
||||
@@ -235,7 +235,7 @@ class TestMem0ProviderMessagesAdding:
|
||||
assert call_args.kwargs["run_id"] == "operation_thread"
|
||||
|
||||
async def test_messages_adding_without_scope_uses_base_thread_id(
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[ChatMessage]
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
|
||||
) -> None:
|
||||
"""Test adding messages without scope uses base thread_id."""
|
||||
provider = Mem0Provider(
|
||||
@@ -254,9 +254,9 @@ class TestMem0ProviderMessagesAdding:
|
||||
"""Test that empty or invalid messages are filtered out."""
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
messages = [
|
||||
ChatMessage(role="user", text=""), # Empty text
|
||||
ChatMessage(role="user", text=" "), # Whitespace only
|
||||
ChatMessage(role="user", text="Valid message"),
|
||||
Message(role="user", text=""), # Empty text
|
||||
Message(role="user", text=" "), # Whitespace only
|
||||
Message(role="user", text="Valid message"),
|
||||
]
|
||||
|
||||
await provider.invoked(messages)
|
||||
@@ -269,8 +269,8 @@ class TestMem0ProviderMessagesAdding:
|
||||
"""Test that mem0 client is not called when no valid messages exist."""
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
messages = [
|
||||
ChatMessage(role="user", text=""),
|
||||
ChatMessage(role="user", text=" "),
|
||||
Message(role="user", text=""),
|
||||
Message(role="user", text=" "),
|
||||
]
|
||||
|
||||
await provider.invoked(messages)
|
||||
@@ -284,7 +284,7 @@ class TestMem0ProviderModelInvoking:
|
||||
async def test_model_invoking_fails_without_filters(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Test that invoking fails when no filters are provided."""
|
||||
provider = Mem0Provider(mem0_client=mock_mem0_client)
|
||||
message = ChatMessage(role="user", text="What's the weather?")
|
||||
message = Message(role="user", text="What's the weather?")
|
||||
|
||||
with pytest.raises(ServiceInitializationError) as exc_info:
|
||||
await provider.invoking(message)
|
||||
@@ -294,7 +294,7 @@ class TestMem0ProviderModelInvoking:
|
||||
async def test_model_invoking_single_message(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Test invoking with a single message."""
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
message = ChatMessage(role="user", text="What's the weather?")
|
||||
message = Message(role="user", text="What's the weather?")
|
||||
|
||||
# Mock search results
|
||||
mock_mem0_client.search.return_value = [
|
||||
@@ -319,7 +319,7 @@ class TestMem0ProviderModelInvoking:
|
||||
assert context.messages[0].text == expected_instructions
|
||||
|
||||
async def test_model_invoking_multiple_messages(
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[ChatMessage]
|
||||
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
|
||||
) -> None:
|
||||
"""Test invoking with multiple messages."""
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
@@ -335,7 +335,7 @@ class TestMem0ProviderModelInvoking:
|
||||
async def test_model_invoking_with_agent_id(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Test invoking with agent_id."""
|
||||
provider = Mem0Provider(agent_id="agent123", mem0_client=mock_mem0_client)
|
||||
message = ChatMessage(role="user", text="Hello")
|
||||
message = Message(role="user", text="Hello")
|
||||
|
||||
mock_mem0_client.search.return_value = []
|
||||
|
||||
@@ -353,7 +353,7 @@ class TestMem0ProviderModelInvoking:
|
||||
mem0_client=mock_mem0_client,
|
||||
)
|
||||
provider._per_operation_thread_id = "operation_thread"
|
||||
message = ChatMessage(role="user", text="Hello")
|
||||
message = Message(role="user", text="Hello")
|
||||
|
||||
mock_mem0_client.search.return_value = []
|
||||
|
||||
@@ -365,7 +365,7 @@ class TestMem0ProviderModelInvoking:
|
||||
async def test_model_invoking_no_memories_returns_none_instructions(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Test that no memories returns context with None instructions."""
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
message = ChatMessage(role="user", text="Hello")
|
||||
message = Message(role="user", text="Hello")
|
||||
|
||||
mock_mem0_client.search.return_value = []
|
||||
|
||||
@@ -381,7 +381,7 @@ class TestMem0ProviderModelInvoking:
|
||||
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
function_call = Content.from_function_call(call_id="1", name="test_func", arguments='{"arg1": "value1"}')
|
||||
message = ChatMessage(
|
||||
message = Message(
|
||||
role="user",
|
||||
contents=[
|
||||
Content.from_function_approval_response(
|
||||
@@ -403,9 +403,9 @@ class TestMem0ProviderModelInvoking:
|
||||
"""Test that empty message text is filtered out from query."""
|
||||
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
|
||||
messages = [
|
||||
ChatMessage(role="user", text=""),
|
||||
ChatMessage(role="user", text="Valid message"),
|
||||
ChatMessage(role="user", text=" "),
|
||||
Message(role="user", text=""),
|
||||
Message(role="user", text="Valid message"),
|
||||
Message(role="user", text=" "),
|
||||
]
|
||||
|
||||
mock_mem0_client.search.return_value = []
|
||||
@@ -423,7 +423,7 @@ class TestMem0ProviderModelInvoking:
|
||||
context_prompt=custom_prompt,
|
||||
mem0_client=mock_mem0_client,
|
||||
)
|
||||
message = ChatMessage(role="user", text="Hello")
|
||||
message = Message(role="user", text="Hello")
|
||||
|
||||
mock_mem0_client.search.return_value = [{"memory": "Test memory"}]
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from __future__ import annotations
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import AgentResponse, ChatMessage
|
||||
from agent_framework import AgentResponse, Message
|
||||
from agent_framework._sessions import AgentSession, SessionContext
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
|
||||
@@ -84,7 +84,7 @@ class TestBeforeRun:
|
||||
]
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="Hello")], session_id="s1")
|
||||
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]
|
||||
|
||||
@@ -100,7 +100,7 @@ class TestBeforeRun:
|
||||
"""Empty input messages → no search performed."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="")], session_id="s1")
|
||||
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]
|
||||
|
||||
@@ -112,7 +112,7 @@ class TestBeforeRun:
|
||||
mock_mem0_client.search.return_value = []
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
|
||||
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]
|
||||
|
||||
@@ -122,7 +122,7 @@ class TestBeforeRun:
|
||||
"""Raises ServiceInitializationError when no filters."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
|
||||
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]
|
||||
@@ -132,7 +132,7 @@ class TestBeforeRun:
|
||||
mock_mem0_client.search.return_value = {"results": [{"memory": "remembered fact"}]}
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
|
||||
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]
|
||||
|
||||
@@ -146,8 +146,8 @@ class TestBeforeRun:
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(
|
||||
input_messages=[
|
||||
ChatMessage(role="user", text="Hello"),
|
||||
ChatMessage(role="user", text="World"),
|
||||
Message(role="user", text="Hello"),
|
||||
Message(role="user", text="World"),
|
||||
],
|
||||
session_id="s1",
|
||||
)
|
||||
@@ -168,8 +168,8 @@ class TestAfterRun:
|
||||
"""Stores input+response messages to mem0 via client.add."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="question")], session_id="s1")
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="answer")])
|
||||
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]
|
||||
|
||||
@@ -188,12 +188,12 @@ class TestAfterRun:
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(
|
||||
input_messages=[
|
||||
ChatMessage(role="user", text="hello"),
|
||||
ChatMessage(role="tool", text="tool output"),
|
||||
Message(role="user", text="hello"),
|
||||
Message(role="tool", text="tool output"),
|
||||
],
|
||||
session_id="s1",
|
||||
)
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="reply")])
|
||||
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]
|
||||
|
||||
@@ -208,8 +208,8 @@ class TestAfterRun:
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(
|
||||
input_messages=[
|
||||
ChatMessage(role="user", text=""),
|
||||
ChatMessage(role="user", text=" "),
|
||||
Message(role="user", text=""),
|
||||
Message(role="user", text=" "),
|
||||
],
|
||||
session_id="s1",
|
||||
)
|
||||
@@ -223,8 +223,8 @@ class TestAfterRun:
|
||||
"""Uses session_id as run_id."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="my-session")
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="hey")])
|
||||
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]
|
||||
|
||||
@@ -234,8 +234,8 @@ class TestAfterRun:
|
||||
"""Raises ServiceInitializationError when no filters."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="s1")
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="hey")])
|
||||
ctx = SessionContext(input_messages=[Message(role="user", text="hi")], session_id="s1")
|
||||
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]
|
||||
@@ -246,7 +246,7 @@ class TestAfterRun:
|
||||
source_id="mem0", mem0_client=mock_mem0_client, user_id="u1", application_id="app1"
|
||||
)
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="s1")
|
||||
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]
|
||||
|
||||
Reference in New Issue
Block a user