Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions tests/agent_memory/integration/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,9 @@
Required for SUBSCRIBER tests. When absent those tests are skipped.
"""

from collections.abc import Generator
import os
import uuid
from pathlib import Path

import pytest
Expand Down Expand Up @@ -84,3 +86,23 @@ def subscriber_tenant() -> str:
)

return tenant


def _delete_all_memories(client: AgentMemoryClient, agent_id: str) -> None:
offset = 0
limit = 50
while True:
memories = client.list_memories(agent_id=agent_id, limit=limit, offset=offset)
for memory in memories:
client.delete_memory(memory.id)
if len(memories) < limit:
break
offset += limit


@pytest.fixture(scope="session")
def run_agent_id(agent_memory_client: AgentMemoryClient) -> Generator[str, None, None]:
"""Return a unique agent ID for this test run and clean up all its data afterwards."""
agent_id = f"test-agent-{uuid.uuid4().hex[:8]}"
yield agent_id
_delete_all_memories(agent_memory_client, agent_id)
40 changes: 21 additions & 19 deletions tests/agent_memory/integration/test_agentmemory_bdd.py
Original file line number Diff line number Diff line change
Expand Up @@ -161,10 +161,12 @@ def test_filter_messages_by_metadata_subscriber():


@pytest.fixture
def context():
def context(run_agent_id):
return {
"access_strategy": AccessStrategy.PROVIDER,
"tenant": None,
"agent_id": run_agent_id,
"invoker_id": "test-user",
}


Expand All @@ -190,7 +192,7 @@ def use_configured_subscriber_tenant(context, subscriber_tenant):
def memory_exists(context, agent_memory_client, agent_id, invoker_id, content):
context["client"] = agent_memory_client
context["memory"] = agent_memory_client.add_memory(
agent_id, invoker_id, content,
context["agent_id"], context["invoker_id"], content,
)


Expand All @@ -202,7 +204,7 @@ def memory_exists(context, agent_memory_client, agent_id, invoker_id, content):
def message_exists(context, agent_memory_client, agent_id, invoker_id, group, role, content):
context["client"] = agent_memory_client
context["message"] = agent_memory_client.add_message(
agent_id, invoker_id, group, role, content,
context["agent_id"], context["invoker_id"], group, role, content,
)


Expand All @@ -214,7 +216,7 @@ def message_exists(context, agent_memory_client, agent_id, invoker_id, group, ro
def message_exists_with_metadata(context, agent_memory_client, agent_id, invoker_id, group, role, content, metadata_value):
context["client"] = agent_memory_client
context["message"] = agent_memory_client.add_message(
agent_id, invoker_id, group, role, content,
context["agent_id"], context["invoker_id"], group, role, content,
metadata={"tag": metadata_value},
)

Expand All @@ -230,7 +232,7 @@ def message_exists_with_metadata(context, agent_memory_client, agent_id, invoker
def add_memory(context, agent_id, invoker_id, content):
client: AgentMemoryClient = context["client"]
context["memory"] = client.add_memory(
agent_id, invoker_id, content,
context["agent_id"], context["invoker_id"], content,
)


Expand All @@ -257,10 +259,10 @@ def update_memory(context, content):
def list_memories(context, agent_id):
client: AgentMemoryClient = context["client"]
context["memories"] = client.list_memories(
agent_id=agent_id,
agent_id=context["agent_id"],
)
context["total"] = client.count_memories(
agent_id=agent_id,
agent_id=context["agent_id"],
)


Expand All @@ -277,8 +279,8 @@ def delete_memory(context):
def search_memories(context, query):
client: AgentMemoryClient = context["client"]
context["search_results"] = client.search_memories(
agent_id="test-agent",
invoker_id="test-user",
agent_id=context["agent_id"],
invoker_id=context["invoker_id"],
query=query,
threshold=0.5,
limit=10,
Expand All @@ -293,7 +295,7 @@ def search_memories(context, query):
def add_message(context, agent_id, invoker_id, group, role, content):
client: AgentMemoryClient = context["client"]
context["message"] = client.add_message(
agent_id, invoker_id, group, MessageRole(role), content,
context["agent_id"], context["invoker_id"], group, MessageRole(role), content,
)


Expand All @@ -305,7 +307,7 @@ def add_message(context, agent_id, invoker_id, group, role, content):
def list_messages(context, agent_id, group):
client: AgentMemoryClient = context["client"]
context["messages"] = client.list_messages(
agent_id=agent_id,
agent_id=context["agent_id"],
message_group=group,
)
context["total"] = len(context["messages"])
Expand Down Expand Up @@ -345,17 +347,17 @@ def update_retention_config(context):
def count_memories(context, agent_id, invoker_id):
client: AgentMemoryClient = context["client"]
context["memory_count"] = client.count_memories(
agent_id=agent_id,
invoker_id=invoker_id,
agent_id=context["agent_id"],
invoker_id=context["invoker_id"],
)


@when(parsers.parse('I list memories filtered by content containing "{substring}"'))
def list_memories_by_content(context, substring):
client: AgentMemoryClient = context["client"]
context["memories"] = client.list_memories(
agent_id="test-agent",
invoker_id="test-user",
agent_id=context["agent_id"],
invoker_id=context["invoker_id"],
filters=[FilterDefinition(target="content", contains=substring)],
)

Expand All @@ -364,8 +366,8 @@ def list_memories_by_content(context, substring):
def list_messages_by_metadata(context, substring):
client: AgentMemoryClient = context["client"]
context["messages"] = client.list_messages(
agent_id="test-agent",
invoker_id="test-user",
agent_id=context["agent_id"],
invoker_id=context["invoker_id"],
message_group="conv-filter",
filters=[FilterDefinition(target="metadata", contains=substring)],
)
Expand All @@ -381,12 +383,12 @@ def check_memory_id(context):

@then(parsers.parse('the memory should have agent_id "{agent_id}"'))
def check_memory_agent_id(context, agent_id):
assert context["memory"].agent_id == agent_id
assert context["memory"].agent_id == context["agent_id"]


@then(parsers.parse('the memory should have invoker_id "{invoker_id}"'))
def check_memory_invoker_id(context, invoker_id):
assert context["memory"].invoker_id == invoker_id
assert context["memory"].invoker_id == context["invoker_id"]


@then(parsers.parse('the memory should have content "{content}"'))
Expand Down
Loading