From 5ba6d0d1f52aaff7f515bccd014bff67348bb9c5 Mon Sep 17 00:00:00 2001 From: phernandez Date: Mon, 20 Jan 2025 20:26:16 -0600 Subject: [PATCH] add fuzzy matching to notes --- .../api/routers/resource_router.py | 16 +++--- src/basic_memory/deps.py | 10 +++- src/basic_memory/mcp/tools/notes.py | 29 ++--------- src/basic_memory/services/relation_service.py | 13 +++-- tests/api/test_knowledge_router.py | 14 ++--- tests/api/test_resource_router.py | 23 +++++++++ tests/conftest.py | 11 ++-- tests/mcp/test_tool_create_relations.py | 26 +++------- tests/mcp/test_tool_notes.py | 10 ++-- tests/services/test_relation_service.py | 51 +++++++++++++------ 10 files changed, 111 insertions(+), 92 deletions(-) diff --git a/src/basic_memory/api/routers/resource_router.py b/src/basic_memory/api/routers/resource_router.py index 3e03b566..cc313cda 100644 --- a/src/basic_memory/api/routers/resource_router.py +++ b/src/basic_memory/api/routers/resource_router.py @@ -6,24 +6,24 @@ from fastapi import APIRouter, HTTPException from fastapi.responses import FileResponse from loguru import logger -from basic_memory.deps import EntityRepositoryDep, ProjectConfigDep +from basic_memory.deps import EntityRepositoryDep, ProjectConfigDep, LinkResolverDep router = APIRouter(prefix="/resource", tags=["resources"]) -@router.get("/{permalink:path}") +@router.get("/{identifier:path}") async def get_resource_content( config: ProjectConfigDep, - entity_repository: EntityRepositoryDep, - permalink: str, + link_resolver: LinkResolverDep, + identifier: str, ) -> FileResponse: - """Get resource content by permalink.""" - logger.debug(f"Getting content for permalink: {permalink}") + """Get resource content by identifier: name or permalink.""" + logger.debug(f"Getting content for permalink: {identifier}") # Find entity by permalink - entity = await entity_repository.get_by_permalink(permalink) + entity = await link_resolver.resolve_link(identifier) if not entity: - raise HTTPException(status_code=404, detail=f"Entity not found: {permalink}") + raise HTTPException(status_code=404, detail=f"Entity not found: {identifier}") file_path = Path(f"{config.home}/{entity.file_path}") if not file_path.exists(): diff --git a/src/basic_memory/deps.py b/src/basic_memory/deps.py index 4937a2f8..d43d776b 100644 --- a/src/basic_memory/deps.py +++ b/src/basic_memory/deps.py @@ -23,6 +23,7 @@ from basic_memory.services import ( ) from basic_memory.services.context_service import ContextService from basic_memory.services.file_service import FileService +from basic_memory.services.link_resolver import LinkResolver from basic_memory.services.search_service import SearchService @@ -105,7 +106,6 @@ SearchRepositoryDep = Annotated[SearchRepository, Depends(get_search_repository) ## services - async def get_file_service(project_config: ProjectConfigDep) -> FileService: return FileService(project_config.home, KnowledgeWriter()) @@ -143,12 +143,14 @@ async def get_relation_service( relation_repository: RelationRepositoryDep, entity_repository: EntityRepositoryDep, file_service: FileServiceDep, + link_resolver: "LinkResolverDep", ) -> RelationService: """Create RelationService with repository.""" return RelationService( relation_repository=relation_repository, entity_repository=entity_repository, file_service=file_service, + link_resolver=link_resolver, ) @@ -171,6 +173,12 @@ async def get_knowledge_writer() -> KnowledgeWriter: KnowledgeWriterDep = Annotated[KnowledgeWriter, Depends(get_knowledge_writer)] +async def get_link_resolver(entity_repository: EntityRepositoryDep, + search_service: SearchServiceDep) -> LinkResolver: + return LinkResolver(entity_repository=entity_repository, + search_service=search_service) + +LinkResolverDep = Annotated[LinkResolver, Depends(get_link_resolver)] async def get_context_service( search_repository: SearchRepositoryDep, entity_repository: EntityRepositoryDep diff --git a/src/basic_memory/mcp/tools/notes.py b/src/basic_memory/mcp/tools/notes.py index 50214ece..8a7c9838 100644 --- a/src/basic_memory/mcp/tools/notes.py +++ b/src/basic_memory/mcp/tools/notes.py @@ -7,17 +7,14 @@ while leveraging the underlying knowledge graph structure. from typing import Optional, List from loguru import logger -from mcp.server.fastmcp.exceptions import ToolError from basic_memory.mcp.server import mcp from basic_memory.mcp.async_client import client -from basic_memory.mcp.tools.search import search from basic_memory.schemas.request import CreateEntityRequest from basic_memory.schemas.base import Entity, Relation from basic_memory.schemas.request import CreateRelationsRequest from basic_memory.mcp.tools.knowledge import create_entities, create_relations from basic_memory.mcp.tools.utils import call_get -from basic_memory.schemas.search import SearchQuery @mcp.tool( @@ -95,25 +92,8 @@ async def read_note(identifier: str) -> str: Raises: ValueError: If the note cannot be found """ - try: - # Try as permalink first - response = await call_get(client, f"/resource/{identifier}") - return response.text - except ToolError as e: - if "404" in str(e): - # If not found, try searching by title - search_response = await search(SearchQuery(text=identifier, entity_types=["note"])) - - if not search_response.results: - raise ValueError(f"Note not found: {identifier}") - - # if we found results, return the first one - response = await call_get(client, f"/resource/{search_response.results[0].permalink}") - return response.text - - raise ValueError(f"Error reading note: {e}") - except Exception as e: - raise ValueError(f"Unexpected error reading note: {e}") + response = await call_get(client, f"/resource/{identifier}") + return response.text @mcp.tool(description="Create a semantic link between two notes") @@ -122,7 +102,7 @@ async def link_notes( to_note: str, relationship: str = "relates_to", context: Optional[str] = None, -) -> None: +) -> str: """Create a semantic link between two notes. Args: @@ -156,4 +136,5 @@ async def link_notes( ) ] ) - await create_relations(request) + response = await create_relations(request) + return response.entities[0].permalink diff --git a/src/basic_memory/services/relation_service.py b/src/basic_memory/services/relation_service.py index 1b02f98f..41fb46d3 100644 --- a/src/basic_memory/services/relation_service.py +++ b/src/basic_memory/services/relation_service.py @@ -9,6 +9,7 @@ from basic_memory.models import Entity as EntityModel, Relation as RelationModel from basic_memory.repository.relation_repository import RelationRepository from . import FileService from .exceptions import EntityNotFoundError +from .link_resolver import LinkResolver from .service import BaseService from ..repository import EntityRepository @@ -24,10 +25,12 @@ class RelationService(BaseService[RelationRepository]): relation_repository: RelationRepository, entity_repository: EntityRepository, file_service: FileService, + link_resolver: LinkResolver, ): super().__init__(relation_repository) self.entity_repository = entity_repository self.file_service = file_service + self.link_resolver = link_resolver async def create_relations(self, relations: List[RelationSchema]) -> Sequence[EntityModel]: """Create relations and return updated entities.""" @@ -37,9 +40,10 @@ class RelationService(BaseService[RelationRepository]): for rs in relations: try: - from_entity = await self.entity_repository.get_by_permalink(rs.from_id) - to_entity = await self.entity_repository.get_by_permalink(rs.to_id) - + # Use link resolver instead of direct permalink lookup + from_entity = await self.link_resolver.resolve_link(rs.from_id) + to_entity = await self.link_resolver.resolve_link(rs.to_id) + relation = RelationModel( from_id=from_entity.id, to_id=to_entity.id, @@ -50,8 +54,7 @@ class RelationService(BaseService[RelationRepository]): await self.repository.add(relation) # Keep track of entities we need to update - entities_to_update.add(rs.from_id) - entities_to_update.add(rs.to_id) + entities_to_update.add(from_entity.permalink) except Exception as e: logger.error(f"Failed to create relation: {e}") diff --git a/tests/api/test_knowledge_router.py b/tests/api/test_knowledge_router.py index 235371a1..5cfc023a 100644 --- a/tests/api/test_knowledge_router.py +++ b/tests/api/test_knowledge_router.py @@ -91,22 +91,16 @@ async def create_related_results(client) -> List[RelationResponse]: # pyright: data = response.json() relation_response = EntityListResponse.model_validate(data) - assert len(relation_response.entities) == 2 + assert len(relation_response.entities) == 1 source_entity = relation_response.entities[0] - target_entity = relation_response.entities[1] assert len(source_entity.relations) == 1 source_relation = source_entity.relations[0] assert source_relation.from_id == source_permalink assert source_relation.to_id == target_permalink - assert len(target_entity.relations) == 1 - target_relation = target_entity.relations[0] - assert target_relation.from_id == source_permalink - assert target_relation.to_id == target_permalink - - return source_entity.relations + target_entity.relations + return source_entity.relations @pytest.mark.asyncio @@ -275,7 +269,7 @@ async def test_delete_observations(client, observation_repository): async def test_delete_relations(client, relation_repository): """Test deleting relations between entities.""" relations = await create_related_results(client) - assert len(relations) == 2 + assert len(relations) == 1 relation = relations[0] # Delete relation @@ -380,7 +374,7 @@ async def test_full_knowledge_flow(client: AsyncClient): ) assert relations_response.status_code == 200 relations_entities = relations_response.json() - assert len(relations_entities["entities"]) == 3 + assert len(relations_entities["entities"]) == 1 # 4. Add observations to main entity await client.post( diff --git a/tests/api/test_resource_router.py b/tests/api/test_resource_router.py index a84df3e6..222d16b1 100644 --- a/tests/api/test_resource_router.py +++ b/tests/api/test_resource_router.py @@ -30,6 +30,29 @@ async def test_get_resource_content(client, test_config, entity_repository): assert response.headers["content-type"] == "text/markdown; charset=utf-8" assert response.text == content +async def test_get_resource_by_title(client, test_config, entity_repository): + """Test getting content by permalink.""" + # Create a test file + content = "# Test Content\n\nThis is a test file." + test_file = Path(test_config.home) / "test" / "test.md" + test_file.parent.mkdir(parents=True, exist_ok=True) + test_file.write_text(content) + + # Create entity referencing the file + entity = await entity_repository.create( + { + "title": "Test Entity", + "entity_type": "test", + "permalink": "test/test", + "file_path": "test/test.md", # Relative to config.home + "content_type": "text/markdown", + } + ) + + # Test getting the content + response = await client.get(f"/resource/{entity.title}") + assert response.status_code == 200 + @pytest.mark.asyncio async def test_get_resource_missing_entity(client): diff --git a/tests/conftest.py b/tests/conftest.py index 37023d79..e548c8dc 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -108,12 +108,14 @@ async def relation_service( relation_repository: RelationRepository, entity_repository: EntityRepository, file_service: FileService, + link_resolver: LinkResolver, ) -> RelationService: """Create RelationService with repository.""" return RelationService( relation_repository=relation_repository, entity_repository=entity_repository, file_service=file_service, + link_resolver=link_resolver, ) @@ -157,7 +159,6 @@ def file_change_scanner(entity_repository) -> FileChangeScanner: return FileChangeScanner(entity_repository) - @pytest_asyncio.fixture async def entity_sync_service( entity_repository: EntityRepository, @@ -254,7 +255,9 @@ async def full_entity(sample_entity, entity_repository): @pytest_asyncio.fixture -async def test_graph(entity_repository, relation_repository, observation_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 = [ @@ -361,8 +364,8 @@ async def test_graph(entity_repository, relation_repository, observation_reposit # Save relations related_entities = await relation_repository.add_all(relations) - - # get latest + + # get latest entities = await entity_repository.find_all() # Index everything for search for entity in entities: diff --git a/tests/mcp/test_tool_create_relations.py b/tests/mcp/test_tool_create_relations.py index 8ad62c12..4c4ba6e3 100644 --- a/tests/mcp/test_tool_create_relations.py +++ b/tests/mcp/test_tool_create_relations.py @@ -27,15 +27,13 @@ async def test_create_basic_relation(client): ) result = await create_relations(relation_request) - assert len(result.entities) == 2 + assert len(result.entities) == 1 # Find source and target entities - source = next(e for e in result.entities if e.permalink == "source-entity") - target = next(e for e in result.entities if e.permalink == "target-entity") + source = result.entities[0] # Both entities should have the relation for bi-directional navigation assert len(source.relations) == 1 - assert len(target.relations) == 1 # Source's relation shows it depends_on target source_relation = source.relations[0] @@ -43,12 +41,6 @@ async def test_create_basic_relation(client): assert source_relation.to_id == "target-entity" assert source_relation.relation_type == "depends_on" - # Target's relation is the same, allowing backwards traversal - target_relation = target.relations[0] - assert target_relation.from_id == "source-entity" - assert target_relation.to_id == "target-entity" - assert target_relation.relation_type == "depends_on" - @pytest.mark.asyncio async def test_create_relation_with_context(client): @@ -74,14 +66,11 @@ async def test_create_relation_with_context(client): ) result = await create_relations(relation_request) - source = next(e for e in result.entities if e.permalink == "source") - target = next(e for e in result.entities if e.permalink == "target") + source = result.entities[0] # Both entities should have the relation with context assert len(source.relations) == 1 - assert len(target.relations) == 1 assert source.relations[0].context == "Implementation details" - assert target.relations[0].context == "Implementation details" @pytest.mark.asyncio @@ -105,24 +94,21 @@ async def test_create_multiple_relations(client): ) result = await create_relations(relation_request) - # Should return all involved entities - assert len(result.entities) == 3 + # Should return all source entities + assert len(result.entities) == 2 # Get entities entity1 = next(e for e in result.entities if e.permalink == "entity1") entity2 = next(e for e in result.entities if e.permalink == "entity2") - entity3 = next(e for e in result.entities if e.permalink == "entity3") # Entity1 and Entity2 should share the connects_to relation assert len(entity1.relations) == 1 assert len(entity2.relations) == 2 # Has both relations - assert len(entity3.relations) == 1 # Verify relation types assert any(r.relation_type == "connects_to" for r in entity1.relations) assert any(r.relation_type == "connects_to" for r in entity2.relations) assert any(r.relation_type == "depends_on" for r in entity2.relations) - assert any(r.relation_type == "depends_on" for r in entity3.relations) @pytest.mark.asyncio @@ -198,7 +184,7 @@ async def test_create_duplicate_relation(client): # Create first relation first_result = await create_relations(relation_request) - assert len(first_result.entities) == 2 + assert len(first_result.entities) == 1 assert len(first_result.entities[0].relations) == 1 # Attempt to create same relation again diff --git a/tests/mcp/test_tool_notes.py b/tests/mcp/test_tool_notes.py index f4658e58..862c6349 100644 --- a/tests/mcp/test_tool_notes.py +++ b/tests/mcp/test_tool_notes.py @@ -1,6 +1,7 @@ """Tests for note tools that exercise the full stack with SQLite.""" import pytest +from mcp.server.fastmcp.exceptions import ToolError from basic_memory.mcp.tools import notes @@ -48,7 +49,7 @@ async def test_write_note_no_tags(app): @pytest.mark.asyncio async def test_read_note_not_found(app): """Test trying to read a non-existent note.""" - with pytest.raises(ValueError, match="Note not found"): + with pytest.raises(ToolError, match="Error calling tool: Client error '404 Not Found'"): await notes.read_note("notes/does-not-exist") @@ -81,16 +82,15 @@ async def test_link_notes(app): ) # Link them - await notes.link_notes( + permalink = await notes.link_notes( from_note=note1, to_note=note2, relationship="inspires", context="Design informs implementation" ) - # TODO: Add verification of the link - # We might want to add a get_note_links() tool - # or use the existing knowledge tools to verify + content = await notes.read_note(permalink) + assert "- inspires [[Implementation]] (Design informs implementation)" in content @pytest.mark.asyncio diff --git a/tests/services/test_relation_service.py b/tests/services/test_relation_service.py index c4e71961..0fe8a6d3 100644 --- a/tests/services/test_relation_service.py +++ b/tests/services/test_relation_service.py @@ -64,7 +64,7 @@ async def test_create_relations( entities = await relation_service.create_relations(relation_data) - assert len(entities) == 2 + assert len(entities) == 1 # verify relations on e0 relations_e0 = entities[0].outgoing_relations @@ -78,17 +78,18 @@ async def test_create_relations( assert relations_e0[1].to_id == entity2.id assert relations_e0[1].relation_type == "type_1" - # verify relations on e1 - relations_e1 = entities[1].incoming_relations - assert len(relations_e1) == 2 + # verify relations on entity2 + e2 = await entity_service.get_by_permalink(entity2.permalink) + relations_e2 = e2.incoming_relations + assert len(relations_e2) == 2 - assert relations_e1[0].from_id == entity1.id - assert relations_e1[0].to_id == entity2.id - assert relations_e1[0].relation_type == "type_0" + assert relations_e2[0].from_id == entity1.id + assert relations_e2[0].to_id == entity2.id + assert relations_e2[0].relation_type == "type_0" - assert relations_e1[1].from_id == entity1.id - assert relations_e1[1].to_id == entity2.id - assert relations_e1[1].relation_type == "type_1" + assert relations_e2[1].from_id == entity1.id + assert relations_e2[1].to_id == entity2.id + assert relations_e2[1].relation_type == "type_1" # Verify outgoing relation is updated found = await entity_service.get_by_permalink(entity1.permalink) @@ -99,11 +100,31 @@ async def test_create_relations( assert f"- type_0 [[{entity2.title}]] (context_0)" in content assert f"- type_1 [[{entity2.title}]] (context_1)" in content - # Verify other entity file is not updated - found = await entity_service.get_by_permalink(entity2.permalink) - file_path = file_service.get_entity_path(found) - content, _ = await file_service.read_file(file_path) - assert "type_0" not in content + +@pytest.mark.asyncio +async def test_create_relations_resolve_links( + relation_service: RelationService, + entity_service: EntityService, + file_service: FileService, + test_entities: tuple[EntityModel, EntityModel], +): + """Test creating a basic relation between two entities.""" + entity1, entity2 = test_entities + + relation_data = [ + RelationSchema( + from_id=entity1.title, + to_id=entity2.title, + relation_type="type_0", + context="context_0", + ), + ] + + entities = await relation_service.create_relations(relation_data) + assert len(entities) == 1 + + assert entities[0].outgoing_relations[0].from_id == entity1.id + assert entities[0].outgoing_relations[0].to_id == entity2.id @pytest.mark.asyncio