Files
basicmachines-co-basic-memory/tests/test_relation_repository.py
T
2024-12-11 12:27:31 -06:00

190 lines
6.5 KiB
Python

"""Tests for RelationRepository."""
import pytest
import pytest_asyncio
from basic_memory.models import Entity, Relation
from basic_memory.repository.relation_repository import RelationRepository
pytestmark = pytest.mark.asyncio
class TestRelationRepository:
@pytest_asyncio.fixture(scope="function")
async def related_entity(self, entity_repository):
"""Create a second entity for testing relations"""
entity_data = {
'id': '20240102-related',
'name': 'Related Entity',
'entity_type': 'test',
'description': 'A related test entity',
'references': ''
}
return await entity_repository.create(entity_data)
@pytest_asyncio.fixture(scope="function")
async def sample_relation(
self,
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity
):
"""Create a sample relation for testing"""
relation_data = {
'from_id': sample_entity.id,
'to_id': related_entity.id,
'relation_type': 'test_relation',
'context': 'test-context'
}
return await relation_repository.create(relation_data)
@pytest_asyncio.fixture(scope="function")
async def multiple_relations(
self,
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity
):
"""Create multiple relations for testing"""
relations_data = [
{
'from_id': sample_entity.id,
'to_id': related_entity.id,
'relation_type': 'relation_one',
'context': 'context_one'
},
{
'from_id': sample_entity.id,
'to_id': related_entity.id,
'relation_type': 'relation_two',
'context': 'context_two'
},
{
'from_id': related_entity.id,
'to_id': sample_entity.id,
'relation_type': 'relation_one',
'context': 'context_three'
}
]
return [await relation_repository.create(data) for data in relations_data]
async def test_create_relation(
self,
relation_repository: RelationRepository,
sample_entity: Entity,
related_entity: Entity
):
"""Test creating a new relation"""
relation_data = {
'from_id': sample_entity.id,
'to_id': related_entity.id,
'relation_type': 'test_relation',
'context': 'test-context'
}
relation = await relation_repository.create(relation_data)
assert relation.from_id == sample_entity.id
assert relation.to_id == related_entity.id
assert relation.relation_type == 'test_relation'
assert relation.id is not None # Should be auto-generated
async def test_find_by_entities(
self,
relation_repository: RelationRepository,
sample_relation: Relation,
sample_entity: Entity,
related_entity: Entity
):
"""Test finding relations between specific entities"""
relations = await relation_repository.find_by_entities(
sample_entity.id,
related_entity.id
)
assert len(relations) == 1
assert relations[0].id == sample_relation.id
assert relations[0].relation_type == sample_relation.relation_type
async def test_find_by_type(
self,
relation_repository: RelationRepository,
sample_relation: Relation
):
"""Test finding relations by type"""
relations = await relation_repository.find_by_type('test_relation')
assert len(relations) == 1
assert relations[0].id == sample_relation.id
async def test_delete_by_fields_single_field(
self,
relation_repository: RelationRepository,
multiple_relations: list[Relation]
):
"""Test deleting relations by a single field."""
# Delete all relations of type 'relation_one'
result = await relation_repository.delete_by_fields(relation_type='relation_one')
assert result is True
# Verify deletion
remaining = await relation_repository.find_by_type('relation_one')
assert len(remaining) == 0
# Other relations should still exist
others = await relation_repository.find_by_type('relation_two')
assert len(others) == 1
async def test_delete_by_fields_multiple_fields(
self,
relation_repository: RelationRepository,
multiple_relations: list[Relation],
sample_entity: Entity,
related_entity: Entity
):
"""Test deleting relations by multiple fields."""
# Delete specific relation matching both from_id and relation_type
result = await relation_repository.delete_by_fields(
from_id=sample_entity.id,
relation_type='relation_one'
)
assert result is True
# Verify correct relation was deleted
remaining = await relation_repository.find_by_entities(
sample_entity.id,
related_entity.id
)
assert len(remaining) == 1 # Only relation_two should remain
assert remaining[0].relation_type == 'relation_two'
async def test_delete_by_fields_no_match(
self,
relation_repository: RelationRepository,
multiple_relations: list[Relation]
):
"""Test delete_by_fields when no relations match."""
result = await relation_repository.delete_by_fields(
relation_type='nonexistent_type'
)
assert result is False
async def test_delete_by_fields_all_fields(
self,
relation_repository: RelationRepository,
multiple_relations: list[Relation],
sample_entity: Entity,
related_entity: Entity
):
"""Test deleting relation by matching all fields."""
# Get first relation's data
relation = multiple_relations[0]
# Delete using all fields
result = await relation_repository.delete_by_fields(
from_id=relation.from_id,
to_id=relation.to_id,
relation_type=relation.relation_type,
context=relation.context
)
assert result is True
# Verify only exact match was deleted
remaining = await relation_repository.find_by_type(relation.relation_type)
assert len(remaining) == 1 # One other relation_one should remain
assert remaining[0].context != relation.context