883 lines
28 KiB
Python
883 lines
28 KiB
Python
|
|
"""Tests for the Agent Memory plugin — service and route layers.
|
||
|
|
|
||
|
|
Uses AsyncMock for all DB operations and embedding generation.
|
||
|
|
No real DB connections required.
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import uuid
|
||
|
|
from datetime import UTC, datetime
|
||
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
|
|
||
|
|
import pytest
|
||
|
|
from fastapi import FastAPI
|
||
|
|
from httpx import ASGITransport, AsyncClient
|
||
|
|
|
||
|
|
from app.plugins.builtins.agent_memory.models import AgentMemory
|
||
|
|
from app.plugins.builtins.agent_memory.routes import router as agent_memory_router
|
||
|
|
from app.plugins.builtins.agent_memory.services import (
|
||
|
|
delete_memory,
|
||
|
|
retrieve_relevant_memories,
|
||
|
|
store_memory,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
# Override conftest DB fixtures — these tests use mocks, no real DB needed
|
||
|
|
@pytest.fixture(autouse=True, scope="session")
|
||
|
|
def db_setup():
|
||
|
|
"""No-op override of conftest db_setup."""
|
||
|
|
yield
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.fixture(autouse=True)
|
||
|
|
def clean_tables(db_setup):
|
||
|
|
"""No-op override of conftest clean_tables."""
|
||
|
|
yield
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.fixture(autouse=True)
|
||
|
|
def _mock_check_permission():
|
||
|
|
"""Patch check_permission to always return True for route tests."""
|
||
|
|
with patch("app.core.permissions.check_permission", return_value=True):
|
||
|
|
yield
|
||
|
|
|
||
|
|
|
||
|
|
# ─── Helpers ───
|
||
|
|
|
||
|
|
|
||
|
|
def _make_memory(
|
||
|
|
*,
|
||
|
|
tenant_id: uuid.UUID | None = None,
|
||
|
|
agent_id: uuid.UUID | None = None,
|
||
|
|
memory_type: str = "fact",
|
||
|
|
content: str = "test content",
|
||
|
|
owner_id: uuid.UUID | None = None,
|
||
|
|
) -> AgentMemory:
|
||
|
|
"""Create an AgentMemory instance with defaults."""
|
||
|
|
return AgentMemory(
|
||
|
|
id=uuid.uuid4(),
|
||
|
|
tenant_id=tenant_id or uuid.uuid4(),
|
||
|
|
agent_id=agent_id or uuid.uuid4(),
|
||
|
|
memory_type=memory_type,
|
||
|
|
content=content,
|
||
|
|
owner_id=owner_id,
|
||
|
|
created_at=datetime.now(UTC),
|
||
|
|
updated_at=datetime.now(UTC),
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def _mock_session() -> AsyncMock:
|
||
|
|
"""Create a mock AsyncSession with common patterns."""
|
||
|
|
session = AsyncMock()
|
||
|
|
session.add = MagicMock()
|
||
|
|
session.flush = AsyncMock()
|
||
|
|
session.refresh = AsyncMock()
|
||
|
|
session.delete = AsyncMock()
|
||
|
|
session.execute = AsyncMock()
|
||
|
|
session.commit = AsyncMock()
|
||
|
|
session.rollback = AsyncMock()
|
||
|
|
return session
|
||
|
|
|
||
|
|
|
||
|
|
# ─── Service-Layer Tests ───
|
||
|
|
|
||
|
|
|
||
|
|
class TestStoreMemory:
|
||
|
|
"""Tests for store_memory() service function."""
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_store_memory_creates_memory_with_embedding(self):
|
||
|
|
"""store_memory creates a memory and stores the embedding."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
owner_id = uuid.uuid4()
|
||
|
|
fake_embedding = [0.1] * 768
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
# generate_embedding returns a list of floats
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=fake_embedding,
|
||
|
|
) as mock_gen:
|
||
|
|
# After flush+refresh, the memory object should have id/created_at
|
||
|
|
def _refresh_side_effect(obj):
|
||
|
|
obj.id = uuid.uuid4()
|
||
|
|
obj.created_at = datetime.now(UTC)
|
||
|
|
obj.updated_at = datetime.now(UTC)
|
||
|
|
|
||
|
|
db.refresh.side_effect = _refresh_side_effect
|
||
|
|
|
||
|
|
result = await store_memory(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_id,
|
||
|
|
agent_id=agent_id,
|
||
|
|
content="The sky is blue",
|
||
|
|
memory_type="fact",
|
||
|
|
owner_id=owner_id,
|
||
|
|
)
|
||
|
|
|
||
|
|
# Verify embedding was generated
|
||
|
|
mock_gen.assert_awaited_once()
|
||
|
|
# Verify memory was added to session
|
||
|
|
db.add.assert_called_once()
|
||
|
|
db.flush.assert_awaited()
|
||
|
|
# Verify embedding SQL was executed
|
||
|
|
assert db.execute.await_count >= 1
|
||
|
|
# Verify result shape
|
||
|
|
assert "id" in result
|
||
|
|
assert result["agent_id"] == str(agent_id)
|
||
|
|
assert result["memory_type"] == "fact"
|
||
|
|
assert result["content"] == "The sky is blue"
|
||
|
|
assert result["owner_id"] == str(owner_id)
|
||
|
|
assert result["created_at"] is not None
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_store_memory_without_embedding(self):
|
||
|
|
"""store_memory still creates memory when embedding generation fails (returns None/empty)."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=None,
|
||
|
|
):
|
||
|
|
def _refresh_side_effect(obj):
|
||
|
|
obj.id = uuid.uuid4()
|
||
|
|
obj.created_at = datetime.now(UTC)
|
||
|
|
obj.updated_at = datetime.now(UTC)
|
||
|
|
|
||
|
|
db.refresh.side_effect = _refresh_side_effect
|
||
|
|
|
||
|
|
result = await store_memory(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_id,
|
||
|
|
agent_id=agent_id,
|
||
|
|
content="No embedding for this",
|
||
|
|
)
|
||
|
|
|
||
|
|
assert result["content"] == "No embedding for this"
|
||
|
|
assert result["owner_id"] is None
|
||
|
|
# No embedding SQL should be executed when embedding is None
|
||
|
|
db.execute.assert_not_awaited()
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_store_memory_empty_embedding_list(self):
|
||
|
|
"""store_memory handles empty embedding list gracefully."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=[],
|
||
|
|
):
|
||
|
|
def _refresh_side_effect(obj):
|
||
|
|
obj.id = uuid.uuid4()
|
||
|
|
obj.created_at = datetime.now(UTC)
|
||
|
|
|
||
|
|
db.refresh.side_effect = _refresh_side_effect
|
||
|
|
|
||
|
|
result = await store_memory(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_id,
|
||
|
|
agent_id=agent_id,
|
||
|
|
content="Empty embedding",
|
||
|
|
)
|
||
|
|
|
||
|
|
assert result["content"] == "Empty embedding"
|
||
|
|
db.execute.assert_not_awaited()
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_store_memory_default_memory_type(self):
|
||
|
|
"""store_memory uses 'fact' as default memory_type."""
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=[0.1] * 768,
|
||
|
|
):
|
||
|
|
def _refresh_side_effect(obj):
|
||
|
|
obj.id = uuid.uuid4()
|
||
|
|
obj.created_at = datetime.now(UTC)
|
||
|
|
|
||
|
|
db.refresh.side_effect = _refresh_side_effect
|
||
|
|
|
||
|
|
result = await store_memory(
|
||
|
|
db=db,
|
||
|
|
tenant_id=uuid.uuid4(),
|
||
|
|
agent_id=uuid.uuid4(),
|
||
|
|
content="Default type test",
|
||
|
|
)
|
||
|
|
|
||
|
|
assert result["memory_type"] == "fact"
|
||
|
|
|
||
|
|
|
||
|
|
class TestRetrieveRelevantMemories:
|
||
|
|
"""Tests for retrieve_relevant_memories() service function."""
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_retrieve_with_embedding_semantic_search(self):
|
||
|
|
"""retrieve_relevant_memories performs vector similarity search when embedding is available."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
fake_embedding = [0.2] * 768
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
# Mock the SQL query result with rows
|
||
|
|
mock_row = {
|
||
|
|
"id": uuid.uuid4(),
|
||
|
|
"agent_id": agent_id,
|
||
|
|
"memory_type": "fact",
|
||
|
|
"content": "Paris is in France",
|
||
|
|
"score": 0.95,
|
||
|
|
"created_at": datetime.now(UTC),
|
||
|
|
}
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.mappings.return_value.all.return_value = [mock_row]
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=fake_embedding,
|
||
|
|
):
|
||
|
|
results = await retrieve_relevant_memories(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_id,
|
||
|
|
agent_id=agent_id,
|
||
|
|
query="capital of France",
|
||
|
|
limit=5,
|
||
|
|
)
|
||
|
|
|
||
|
|
assert len(results) == 1
|
||
|
|
assert results[0]["content"] == "Paris is in France"
|
||
|
|
assert results[0]["score"] == 0.95
|
||
|
|
assert results[0]["agent_id"] == str(agent_id)
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_retrieve_fallback_on_embedding_failure(self):
|
||
|
|
"""retrieve_relevant_memories falls back to recent memories when embedding fails."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
# Create mock memory objects for fallback query
|
||
|
|
mem1 = _make_memory(
|
||
|
|
tenant_id=tenant_id, agent_id=agent_id, content="Recent memory 1"
|
||
|
|
)
|
||
|
|
mem2 = _make_memory(
|
||
|
|
tenant_id=tenant_id, agent_id=agent_id, content="Recent memory 2"
|
||
|
|
)
|
||
|
|
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.scalars.return_value.all.return_value = [mem1, mem2]
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=None,
|
||
|
|
):
|
||
|
|
results = await retrieve_relevant_memories(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_id,
|
||
|
|
agent_id=agent_id,
|
||
|
|
query="some query",
|
||
|
|
limit=10,
|
||
|
|
)
|
||
|
|
|
||
|
|
assert len(results) == 2
|
||
|
|
assert results[0]["content"] == "Recent memory 1"
|
||
|
|
assert results[0]["score"] == 0.0 # fallback has no score
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_retrieve_with_memory_type_filter(self):
|
||
|
|
"""retrieve_relevant_memories filters by memory_type when provided."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
mock_row = {
|
||
|
|
"id": uuid.uuid4(),
|
||
|
|
"agent_id": agent_id,
|
||
|
|
"memory_type": "instruction",
|
||
|
|
"content": "Always be polite",
|
||
|
|
"score": 0.88,
|
||
|
|
"created_at": datetime.now(UTC),
|
||
|
|
}
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.mappings.return_value.all.return_value = [mock_row]
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=[0.3] * 768,
|
||
|
|
):
|
||
|
|
results = await retrieve_relevant_memories(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_id,
|
||
|
|
agent_id=agent_id,
|
||
|
|
query="be polite",
|
||
|
|
memory_type="instruction",
|
||
|
|
)
|
||
|
|
|
||
|
|
assert len(results) == 1
|
||
|
|
assert results[0]["memory_type"] == "instruction"
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_retrieve_min_score_filter(self):
|
||
|
|
"""retrieve_relevant_memories filters out results below min_score."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
# Two rows: one above threshold, one below
|
||
|
|
mock_row_high = {
|
||
|
|
"id": uuid.uuid4(),
|
||
|
|
"agent_id": agent_id,
|
||
|
|
"memory_type": "fact",
|
||
|
|
"content": "High score memory",
|
||
|
|
"score": 0.9,
|
||
|
|
"created_at": datetime.now(UTC),
|
||
|
|
}
|
||
|
|
mock_row_low = {
|
||
|
|
"id": uuid.uuid4(),
|
||
|
|
"agent_id": agent_id,
|
||
|
|
"memory_type": "fact",
|
||
|
|
"content": "Low score memory",
|
||
|
|
"score": 0.3,
|
||
|
|
"created_at": datetime.now(UTC),
|
||
|
|
}
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.mappings.return_value.all.return_value = [mock_row_high, mock_row_low]
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=[0.5] * 768,
|
||
|
|
):
|
||
|
|
results = await retrieve_relevant_memories(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_id,
|
||
|
|
agent_id=agent_id,
|
||
|
|
query="test",
|
||
|
|
min_score=0.5,
|
||
|
|
)
|
||
|
|
|
||
|
|
# Only the high-score row should pass the min_score filter
|
||
|
|
assert len(results) == 1
|
||
|
|
assert results[0]["content"] == "High score memory"
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_retrieve_fallback_with_memory_type_filter(self):
|
||
|
|
"""Fallback query also respects memory_type filter."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
mem = _make_memory(
|
||
|
|
tenant_id=tenant_id, agent_id=agent_id, memory_type="pattern", content="Pattern A"
|
||
|
|
)
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.scalars.return_value.all.return_value = [mem]
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=None,
|
||
|
|
):
|
||
|
|
results = await retrieve_relevant_memories(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_id,
|
||
|
|
agent_id=agent_id,
|
||
|
|
query="test",
|
||
|
|
memory_type="pattern",
|
||
|
|
)
|
||
|
|
|
||
|
|
assert len(results) == 1
|
||
|
|
assert results[0]["memory_type"] == "pattern"
|
||
|
|
|
||
|
|
|
||
|
|
class TestDeleteMemory:
|
||
|
|
"""Tests for delete_memory() service function."""
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_delete_memory_success(self):
|
||
|
|
"""delete_memory returns True when memory exists."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
memory_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
memory = _make_memory(tenant_id=tenant_id)
|
||
|
|
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.scalar_one_or_none.return_value = memory
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
result = await delete_memory(db, tenant_id, memory_id)
|
||
|
|
|
||
|
|
assert result is True
|
||
|
|
db.delete.assert_awaited_once_with(memory)
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_delete_memory_not_found(self):
|
||
|
|
"""delete_memory returns False when memory does not exist."""
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
memory_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.scalar_one_or_none.return_value = None
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
result = await delete_memory(db, tenant_id, memory_id)
|
||
|
|
|
||
|
|
assert result is False
|
||
|
|
db.delete.assert_not_awaited()
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_delete_memory_tenant_isolation(self):
|
||
|
|
"""delete_memory only finds memories within the same tenant."""
|
||
|
|
tenant_a = uuid.uuid4()
|
||
|
|
tenant_b = uuid.uuid4()
|
||
|
|
memory_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
mock_result = MagicMock()
|
||
|
|
# Memory from tenant B should not be found when querying with tenant A
|
||
|
|
mock_result.scalar_one_or_none.return_value = None
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
result = await delete_memory(db, tenant_a, memory_id)
|
||
|
|
|
||
|
|
assert result is False
|
||
|
|
|
||
|
|
|
||
|
|
# ─── Model Tests ───
|
||
|
|
|
||
|
|
|
||
|
|
class TestAgentMemoryModel:
|
||
|
|
"""Tests for the AgentMemory model attributes."""
|
||
|
|
|
||
|
|
def test_model_has_required_fields(self):
|
||
|
|
"""AgentMemory model has tenant_id, agent_id, memory_type, content."""
|
||
|
|
mem = AgentMemory(
|
||
|
|
tenant_id=uuid.uuid4(),
|
||
|
|
agent_id=uuid.uuid4(),
|
||
|
|
memory_type="fact",
|
||
|
|
content="test",
|
||
|
|
)
|
||
|
|
assert mem.tenant_id is not None
|
||
|
|
assert mem.agent_id is not None
|
||
|
|
assert mem.memory_type == "fact"
|
||
|
|
assert mem.content == "test"
|
||
|
|
|
||
|
|
def test_model_default_memory_type(self):
|
||
|
|
"""AgentMemory has default='fact' for memory_type column."""
|
||
|
|
mem = AgentMemory(
|
||
|
|
tenant_id=uuid.uuid4(),
|
||
|
|
agent_id=uuid.uuid4(),
|
||
|
|
content="test",
|
||
|
|
)
|
||
|
|
# The default is set at DB level (default="fact"), so we check the column default
|
||
|
|
col = AgentMemory.__table__.c.memory_type
|
||
|
|
assert col.default.arg == "fact"
|
||
|
|
|
||
|
|
def test_model_table_name(self):
|
||
|
|
"""AgentMemory uses correct table name."""
|
||
|
|
assert AgentMemory.__tablename__ == "agent_memories"
|
||
|
|
|
||
|
|
|
||
|
|
# ─── Route-Layer Tests ───
|
||
|
|
|
||
|
|
|
||
|
|
def _create_agent_memory_app() -> FastAPI:
|
||
|
|
"""Create a minimal FastAPI app with agent_memory router and mocked dependencies."""
|
||
|
|
app = FastAPI()
|
||
|
|
app.include_router(agent_memory_router)
|
||
|
|
|
||
|
|
async def _mock_get_db():
|
||
|
|
db = _mock_session()
|
||
|
|
yield db
|
||
|
|
|
||
|
|
async def _mock_get_current_user():
|
||
|
|
return {
|
||
|
|
"user_id": str(uuid.uuid4()),
|
||
|
|
"tenant_id": str(uuid.uuid4()),
|
||
|
|
"is_system_admin": True,
|
||
|
|
"role": "admin",
|
||
|
|
"permissions": [],
|
||
|
|
}
|
||
|
|
|
||
|
|
async def _mock_require_permission(permission: str):
|
||
|
|
async def _check():
|
||
|
|
return {
|
||
|
|
"user_id": str(uuid.uuid4()),
|
||
|
|
"tenant_id": str(uuid.uuid4()),
|
||
|
|
"is_system_admin": True,
|
||
|
|
}
|
||
|
|
return _check
|
||
|
|
|
||
|
|
from app.deps import get_current_user
|
||
|
|
from app.core.db import get_db
|
||
|
|
|
||
|
|
app.dependency_overrides[get_db] = _mock_get_db
|
||
|
|
app.dependency_overrides[get_current_user] = _mock_get_current_user
|
||
|
|
|
||
|
|
return app
|
||
|
|
|
||
|
|
|
||
|
|
class TestAgentMemoryRoutes:
|
||
|
|
"""Tests for agent memory API routes."""
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_create_memory_route(self):
|
||
|
|
"""POST /api/v1/agent-memory creates a memory."""
|
||
|
|
app = _create_agent_memory_app()
|
||
|
|
|
||
|
|
# Override get_db to return our mock
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
user_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
async def _mock_get_db():
|
||
|
|
yield db
|
||
|
|
|
||
|
|
async def _mock_get_current_user():
|
||
|
|
return {
|
||
|
|
"user_id": str(user_id),
|
||
|
|
"tenant_id": str(tenant_id),
|
||
|
|
"is_system_admin": True,
|
||
|
|
}
|
||
|
|
|
||
|
|
from app.deps import get_current_user
|
||
|
|
from app.core.db import get_db
|
||
|
|
app.dependency_overrides[get_db] = _mock_get_db
|
||
|
|
app.dependency_overrides[get_current_user] = _mock_get_current_user
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=[0.1] * 768,
|
||
|
|
):
|
||
|
|
def _refresh_side_effect(obj):
|
||
|
|
obj.id = uuid.uuid4()
|
||
|
|
obj.created_at = datetime.now(UTC)
|
||
|
|
obj.updated_at = datetime.now(UTC)
|
||
|
|
|
||
|
|
db.refresh.side_effect = _refresh_side_effect
|
||
|
|
|
||
|
|
transport = ASGITransport(app=app)
|
||
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||
|
|
resp = await client.post(
|
||
|
|
"/api/v1/agent-memory",
|
||
|
|
json={
|
||
|
|
"agent_id": str(agent_id),
|
||
|
|
"memory_type": "fact",
|
||
|
|
"content": "Route test memory",
|
||
|
|
},
|
||
|
|
)
|
||
|
|
|
||
|
|
assert resp.status_code == 201
|
||
|
|
data = resp.json()
|
||
|
|
assert data["content"] == "Route test memory"
|
||
|
|
assert data["agent_id"] == str(agent_id)
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_create_memory_invalid_agent_id(self):
|
||
|
|
"""POST /api/v1/agent-memory returns 400 for invalid agent_id UUID."""
|
||
|
|
app = _create_agent_memory_app()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
async def _mock_get_db():
|
||
|
|
yield db
|
||
|
|
|
||
|
|
async def _mock_get_current_user():
|
||
|
|
return {
|
||
|
|
"user_id": str(uuid.uuid4()),
|
||
|
|
"tenant_id": str(uuid.uuid4()),
|
||
|
|
"is_system_admin": True,
|
||
|
|
}
|
||
|
|
|
||
|
|
from app.deps import get_current_user
|
||
|
|
from app.core.db import get_db
|
||
|
|
app.dependency_overrides[get_db] = _mock_get_db
|
||
|
|
app.dependency_overrides[get_current_user] = _mock_get_current_user
|
||
|
|
|
||
|
|
transport = ASGITransport(app=app)
|
||
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||
|
|
resp = await client.post(
|
||
|
|
"/api/v1/agent-memory",
|
||
|
|
json={
|
||
|
|
"agent_id": "not-a-uuid",
|
||
|
|
"content": "test",
|
||
|
|
},
|
||
|
|
)
|
||
|
|
|
||
|
|
assert resp.status_code == 400
|
||
|
|
assert resp.json()["detail"]["code"] == "invalid_id"
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_search_memories_route(self):
|
||
|
|
"""GET /api/v1/agent-memory/search returns search results."""
|
||
|
|
app = _create_agent_memory_app()
|
||
|
|
tenant_id = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
async def _mock_get_db():
|
||
|
|
yield db
|
||
|
|
|
||
|
|
async def _mock_get_current_user():
|
||
|
|
return {
|
||
|
|
"user_id": str(uuid.uuid4()),
|
||
|
|
"tenant_id": str(tenant_id),
|
||
|
|
"is_system_admin": True,
|
||
|
|
}
|
||
|
|
|
||
|
|
from app.deps import get_current_user
|
||
|
|
from app.core.db import get_db
|
||
|
|
app.dependency_overrides[get_db] = _mock_get_db
|
||
|
|
app.dependency_overrides[get_current_user] = _mock_get_current_user
|
||
|
|
|
||
|
|
mock_row = {
|
||
|
|
"id": uuid.uuid4(),
|
||
|
|
"agent_id": agent_id,
|
||
|
|
"memory_type": "fact",
|
||
|
|
"content": "Found memory",
|
||
|
|
"score": 0.92,
|
||
|
|
"created_at": datetime.now(UTC),
|
||
|
|
}
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.mappings.return_value.all.return_value = [mock_row]
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=[0.4] * 768,
|
||
|
|
):
|
||
|
|
transport = ASGITransport(app=app)
|
||
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||
|
|
resp = await client.get(
|
||
|
|
"/api/v1/agent-memory/search",
|
||
|
|
params={
|
||
|
|
"agent_id": str(agent_id),
|
||
|
|
"query": "test query",
|
||
|
|
},
|
||
|
|
)
|
||
|
|
|
||
|
|
assert resp.status_code == 200
|
||
|
|
data = resp.json()
|
||
|
|
assert data["total"] == 1
|
||
|
|
assert data["items"][0]["content"] == "Found memory"
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_delete_memory_route_success(self):
|
||
|
|
"""DELETE /api/v1/agent-memory/{id} returns 204 on success."""
|
||
|
|
app = _create_agent_memory_app()
|
||
|
|
memory_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
memory = _make_memory()
|
||
|
|
|
||
|
|
async def _mock_get_db():
|
||
|
|
yield db
|
||
|
|
|
||
|
|
async def _mock_get_current_user():
|
||
|
|
return {
|
||
|
|
"user_id": str(uuid.uuid4()),
|
||
|
|
"tenant_id": str(uuid.uuid4()),
|
||
|
|
"is_system_admin": True,
|
||
|
|
}
|
||
|
|
|
||
|
|
from app.deps import get_current_user
|
||
|
|
from app.core.db import get_db
|
||
|
|
app.dependency_overrides[get_db] = _mock_get_db
|
||
|
|
app.dependency_overrides[get_current_user] = _mock_get_current_user
|
||
|
|
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.scalar_one_or_none.return_value = memory
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
transport = ASGITransport(app=app)
|
||
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||
|
|
resp = await client.delete(f"/api/v1/agent-memory/{memory_id}")
|
||
|
|
|
||
|
|
assert resp.status_code == 204
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_delete_memory_route_not_found(self):
|
||
|
|
"""DELETE /api/v1/agent-memory/{id} returns 404 when not found."""
|
||
|
|
app = _create_agent_memory_app()
|
||
|
|
memory_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
async def _mock_get_db():
|
||
|
|
yield db
|
||
|
|
|
||
|
|
async def _mock_get_current_user():
|
||
|
|
return {
|
||
|
|
"user_id": str(uuid.uuid4()),
|
||
|
|
"tenant_id": str(uuid.uuid4()),
|
||
|
|
"is_system_admin": True,
|
||
|
|
}
|
||
|
|
|
||
|
|
from app.deps import get_current_user
|
||
|
|
from app.core.db import get_db
|
||
|
|
app.dependency_overrides[get_db] = _mock_get_db
|
||
|
|
app.dependency_overrides[get_current_user] = _mock_get_current_user
|
||
|
|
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.scalar_one_or_none.return_value = None
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
transport = ASGITransport(app=app)
|
||
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||
|
|
resp = await client.delete(f"/api/v1/agent-memory/{memory_id}")
|
||
|
|
|
||
|
|
assert resp.status_code == 404
|
||
|
|
assert resp.json()["detail"]["code"] == "not_found"
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_list_memories_route(self):
|
||
|
|
"""GET /api/v1/agent-memory lists memories with pagination."""
|
||
|
|
app = _create_agent_memory_app()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
async def _mock_get_db():
|
||
|
|
yield db
|
||
|
|
|
||
|
|
async def _mock_get_current_user():
|
||
|
|
return {
|
||
|
|
"user_id": str(uuid.uuid4()),
|
||
|
|
"tenant_id": str(uuid.uuid4()),
|
||
|
|
"is_system_admin": True,
|
||
|
|
}
|
||
|
|
|
||
|
|
from app.deps import get_current_user
|
||
|
|
from app.core.db import get_db
|
||
|
|
app.dependency_overrides[get_db] = _mock_get_db
|
||
|
|
app.dependency_overrides[get_current_user] = _mock_get_current_user
|
||
|
|
|
||
|
|
# Mock count query
|
||
|
|
count_result = MagicMock()
|
||
|
|
count_result.scalar_one.return_value = 1
|
||
|
|
|
||
|
|
# Mock paginated query
|
||
|
|
mem = _make_memory(agent_id=agent_id, content="Listed memory")
|
||
|
|
list_result = MagicMock()
|
||
|
|
list_result.scalars.return_value.all.return_value = [mem]
|
||
|
|
|
||
|
|
# execute is called twice: once for count, once for list
|
||
|
|
db.execute.side_effect = [count_result, list_result]
|
||
|
|
|
||
|
|
transport = ASGITransport(app=app)
|
||
|
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||
|
|
resp = await client.get(
|
||
|
|
"/api/v1/agent-memory",
|
||
|
|
params={"agent_id": str(agent_id)},
|
||
|
|
)
|
||
|
|
|
||
|
|
assert resp.status_code == 200
|
||
|
|
data = resp.json()
|
||
|
|
assert data["total"] == 1
|
||
|
|
assert len(data["items"]) == 1
|
||
|
|
assert data["items"][0]["content"] == "Listed memory"
|
||
|
|
|
||
|
|
|
||
|
|
# ─── Tenant Isolation Tests ───
|
||
|
|
|
||
|
|
|
||
|
|
class TestTenantIsolation:
|
||
|
|
"""Tests for tenant_id isolation in agent memory."""
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_store_memory_uses_correct_tenant_id(self):
|
||
|
|
"""store_memory creates memory with the provided tenant_id."""
|
||
|
|
tenant_a = uuid.uuid4()
|
||
|
|
tenant_b = uuid.uuid4()
|
||
|
|
agent_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
captured_tenant_ids = []
|
||
|
|
|
||
|
|
def _track_add(obj):
|
||
|
|
captured_tenant_ids.append(obj.tenant_id)
|
||
|
|
|
||
|
|
db.add.side_effect = _track_add
|
||
|
|
|
||
|
|
with patch(
|
||
|
|
"app.plugins.builtins.agent_memory.services.generate_embedding",
|
||
|
|
new_callable=AsyncMock,
|
||
|
|
return_value=[0.1] * 768,
|
||
|
|
):
|
||
|
|
def _refresh_side_effect(obj):
|
||
|
|
obj.id = uuid.uuid4()
|
||
|
|
obj.created_at = datetime.now(UTC)
|
||
|
|
|
||
|
|
db.refresh.side_effect = _refresh_side_effect
|
||
|
|
|
||
|
|
await store_memory(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_a,
|
||
|
|
agent_id=agent_id,
|
||
|
|
content="Tenant A memory",
|
||
|
|
)
|
||
|
|
|
||
|
|
await store_memory(
|
||
|
|
db=db,
|
||
|
|
tenant_id=tenant_b,
|
||
|
|
agent_id=agent_id,
|
||
|
|
content="Tenant B memory",
|
||
|
|
)
|
||
|
|
|
||
|
|
assert tenant_a in captured_tenant_ids
|
||
|
|
assert tenant_b in captured_tenant_ids
|
||
|
|
assert len(captured_tenant_ids) == 2
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_delete_memory_respects_tenant_boundary(self):
|
||
|
|
"""delete_memory does not delete memories from other tenants."""
|
||
|
|
tenant_a = uuid.uuid4()
|
||
|
|
tenant_b = uuid.uuid4()
|
||
|
|
memory_id = uuid.uuid4()
|
||
|
|
|
||
|
|
db = _mock_session()
|
||
|
|
|
||
|
|
# Simulate that the memory belongs to tenant_b, not tenant_a
|
||
|
|
mock_result = MagicMock()
|
||
|
|
mock_result.scalar_one_or_none.return_value = None
|
||
|
|
db.execute.return_value = mock_result
|
||
|
|
|
||
|
|
result = await delete_memory(db, tenant_a, memory_id)
|
||
|
|
|
||
|
|
# Should return False because the memory is not in tenant_a
|
||
|
|
assert result is False
|
||
|
|
db.delete.assert_not_awaited()
|