diff --git a/src/basic_memory/services/context_service.py b/src/basic_memory/services/context_service.py index 8c192315..db4c0e7e 100644 --- a/src/basic_memory/services/context_service.py +++ b/src/basic_memory/services/context_service.py @@ -202,7 +202,7 @@ WITH RECURSIVE context_graph AS ( cg.type = 'relation' AND e.type = 'entity' AND e.id = CASE - WHEN cg.from_id = cg.root_id THEN cg.to_id + WHEN cg.from_id = cg.id THEN cg.to_id ELSE cg.from_id END {related_date_filter} diff --git a/tests/conftest.py b/tests/conftest.py index 0e7ad2f1..125c8a55 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -260,7 +260,7 @@ async def full_entity(sample_entity, entity_repository): @pytest_asyncio.fixture -async def test_graph(entity_repository, search_service): +async def test_graph(entity_repository, relation_repository, observation_repository, search_service): """Create a test knowledge graph with entities, relations and observations.""" # Create some test entities entities = [ @@ -327,6 +327,8 @@ async def test_graph(entity_repository, search_service): updated_at=datetime.now(timezone.utc), ) ] + await observation_repository.add_all(root.observations) + await observation_repository.add_all(conn1.observations) # Add relations relations = [ @@ -339,33 +341,35 @@ async def test_graph(entity_repository, search_service): updated_at=datetime.now(timezone.utc), ), Relation( - from_id=conn2.id, - to_id=root.id, - relation_type="connected_from", + from_id=conn1.id, + to_id=conn2.id, + relation_type="connected_to", created_at=datetime.now(timezone.utc), updated_at=datetime.now(timezone.utc), ), # Deep connection Relation( - from_id=conn1.id, + from_id=conn2.id, to_id=deep.id, relation_type="deep_connection", created_at=datetime.now(timezone.utc), updated_at=datetime.now(timezone.utc), ), - # Deep connection + # Deeper connection Relation( from_id=deep.id, to_id=deeper.id, - relation_type="deep_connection", + relation_type="deeper_connection", created_at=datetime.now(timezone.utc), updated_at=datetime.now(timezone.utc), ), ] # Save relations - root = await entity_repository.add_all(relations) - + related_entities = await relation_repository.add_all(relations) + + # get latest + entities = await entity_repository.find_all() # Index everything for search for entity in entities: await search_service.index_entity(entity) diff --git a/tests/services/test_context_service.py b/tests/services/test_context_service.py index de586fb2..e0ca9efe 100644 --- a/tests/services/test_context_service.py +++ b/tests/services/test_context_service.py @@ -58,8 +58,8 @@ async def test_find_connected_depth_limit(context_service, test_graph): assert (test_graph["deep"].id, "entity") not in shallow_entities - # With depth=2, we get the next level - deep_results = await context_service.find_related(type_id_pairs, max_depth=4, max_results=10) + # search deeper + deep_results = await context_service.find_related(type_id_pairs, max_depth=4, max_results=100) deep_entities = {(r.id, r.type) for r in deep_results if r.type == "entity"} # Should now include Deep entity assert (test_graph["deep"].id, "entity") in deep_entities @@ -148,7 +148,7 @@ async def test_build_context(context_service, test_graph): assert results["metadata"]["depth"] == 1 assert matched_results == 1 assert len(primary_results) == 1 - assert len(related_results) == 8 + assert len(related_results) == 2 assert total_results == len(primary_results) + len(related_results)