diff --git a/src/basic_memory/repository/search_repository.py b/src/basic_memory/repository/search_repository.py index 0e5116db..88f117b0 100644 --- a/src/basic_memory/repository/search_repository.py +++ b/src/basic_memory/repository/search_repository.py @@ -8,9 +8,10 @@ from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from basic_memory import db from basic_memory.repository.repository import Repository -from basic_memory.schemas.search import SearchQuery, SearchResult +from basic_memory.schemas.search import SearchQuery, SearchResult, SearchItemType from basic_memory.models.search import CREATE_SEARCH_INDEX + class SearchRepository(): """Repository for search index operations.""" @@ -38,7 +39,8 @@ class SearchRepository(): # Handle type filter if query.types: - type_list = ", ".join(f"'{t}'" for t in query.types) + # Get string values from enums + type_list = ", ".join(f"'{t.value}'" for t in query.types) conditions.append(f"type IN ({type_list})") # Handle entity type filter @@ -78,7 +80,7 @@ class SearchRepository(): SearchResult( path_id=row.path_id, file_path=row.file_path, - type=row.type, + type=SearchItemType(row.type), # Convert string to enum score=row.score, metadata=json.loads(row.metadata) ) @@ -90,7 +92,7 @@ class SearchRepository(): content: str, path_id: str, file_path: str, - type: str, + type: SearchItemType, # Now accepts enum metadata: dict ): """Index or update a single item.""" @@ -114,7 +116,7 @@ class SearchRepository(): "content": content, "path_id": path_id, "file_path": file_path, - "type": type, + "type": type.value, # Store the string value "metadata": json.dumps(metadata) } ) diff --git a/src/basic_memory/schemas/search.py b/src/basic_memory/schemas/search.py index 07ed749d..ac530b89 100644 --- a/src/basic_memory/schemas/search.py +++ b/src/basic_memory/schemas/search.py @@ -1,16 +1,29 @@ +"""Search schemas for Basic Memory.""" + from typing import Optional, List from datetime import datetime +from enum import Enum from pydantic import BaseModel + +class SearchItemType(str, Enum): + """Types of searchable items.""" + DOCUMENT = "document" + ENTITY = "entity" + + class SearchQuery(BaseModel): + """Search query parameters.""" text: str - types: Optional[List[str]] = None + types: Optional[List[SearchItemType]] = None entity_types: Optional[List[str]] = None after_date: Optional[datetime] = None + class SearchResult(BaseModel): + """Search result item.""" path_id: str file_path: str - type: str + type: SearchItemType score: float metadata: dict \ No newline at end of file diff --git a/src/basic_memory/services/search_service.py b/src/basic_memory/services/search_service.py index 85518486..465e5e04 100644 --- a/src/basic_memory/services/search_service.py +++ b/src/basic_memory/services/search_service.py @@ -3,7 +3,7 @@ from typing import List, Optional from basic_memory.repository.search_repository import SearchRepository -from basic_memory.schemas.search import SearchQuery, SearchResult +from basic_memory.schemas.search import SearchQuery, SearchResult, SearchItemType class SearchService: @@ -52,7 +52,7 @@ class SearchService: content=content, path_id=entity.path_id, file_path=entity.file_path, - type="entity", + type=SearchItemType.ENTITY, metadata=metadata, ) else: @@ -60,10 +60,37 @@ class SearchService: content=content, path_id=entity.path_id, file_path=entity.file_path, - type="entity", + type=SearchItemType.ENTITY, + metadata=metadata, + ) + + async def index_document(self, document, content: str, background_tasks=None): + """Index a document and its content.""" + metadata = { + **document.doc_metadata, + "created_at": document.created_at.isoformat(), + "updated_at": document.updated_at.isoformat(), + } + + # Queue indexing if background_tasks provided + if background_tasks: + background_tasks.add_task( + self._do_index, + content=content, + path_id=document.path_id, + file_path=document.file_path, + type=SearchItemType.DOCUMENT, + metadata=metadata, + ) + else: + await self._do_index( + content=content, + path_id=document.path_id, + file_path=document.file_path, + type=SearchItemType.DOCUMENT, metadata=metadata, ) async def _do_index(self, **kwargs): """Actually perform the indexing.""" - await self.repository.index_item(**kwargs) + await self.repository.index_item(**kwargs) \ No newline at end of file diff --git a/tests/services/test_search_service.py b/tests/services/test_search_service.py index fa25e79d..e12b44cb 100644 --- a/tests/services/test_search_service.py +++ b/tests/services/test_search_service.py @@ -6,7 +6,7 @@ from sqlalchemy import text from basic_memory import db from basic_memory.repository.search_repository import SearchRepository -from basic_memory.schemas.search import SearchQuery +from basic_memory.schemas.search import SearchQuery, SearchItemType from basic_memory.services.search_service import SearchService @@ -26,6 +26,23 @@ def test_entity(): relations = [] return Entity() + +@pytest.fixture +def test_document(): + """Create a test document""" + class Document: + id = 1 + path_id = "docs/test_doc.md" + file_path = "docs/test_doc.md" + doc_metadata = { + "title": "Test Document", + "type": "technical" + } + created_at = datetime.now(timezone.utc) + updated_at = datetime.now(timezone.utc) + return Document() + + @pytest.mark.asyncio async def test_init_search_index(search_service, session_maker): """Test search index initialization""" @@ -36,6 +53,7 @@ async def test_init_search_index(search_service, session_maker): )) assert result.scalar() == "search_index" + @pytest.mark.asyncio async def test_index_entity(search_service, test_entity): """Test indexing an entity""" @@ -45,6 +63,8 @@ async def test_index_entity(search_service, test_entity): results = await search_service.search(SearchQuery(text="test component")) assert len(results) == 1 assert results[0].path_id == test_entity.path_id + assert results[0].type == SearchItemType.ENTITY + @pytest.mark.asyncio async def test_search_filtering(search_service, test_entity): @@ -55,7 +75,7 @@ async def test_search_filtering(search_service, test_entity): results = await search_service.search( SearchQuery( text="test", - types=["entity"], + types=[SearchItemType.ENTITY], entity_types=["component"] ) ) @@ -65,11 +85,12 @@ async def test_search_filtering(search_service, test_entity): results = await search_service.search( SearchQuery( text="test", - types=["document"] + types=[SearchItemType.DOCUMENT] ) ) assert len(results) == 0 + @pytest.mark.asyncio async def test_update_index(search_service, test_entity): """Test updating indexed content""" @@ -83,6 +104,7 @@ async def test_update_index(search_service, test_entity): results = await search_service.search(SearchQuery(text="new terms")) assert len(results) == 1 + @pytest.mark.asyncio async def test_search_date_filter(search_service, test_entity): """Test searching with date filter""" @@ -96,4 +118,39 @@ async def test_search_date_filter(search_service, test_entity): after_date=future ) ) - assert len(results) == 0 \ No newline at end of file + assert len(results) == 0 + + +@pytest.mark.asyncio +async def test_index_document(search_service, test_document): + """Test indexing a document""" + content = """# Test Document + +This is a test document with some searchable content. +It contains technical information about implementation.""" + + await search_service.index_document(test_document, content) + + # Search for document content + results = await search_service.search(SearchQuery(text="searchable content")) + assert len(results) == 1 + assert results[0].path_id == test_document.path_id + assert results[0].type == SearchItemType.DOCUMENT + + # Verify metadata + assert results[0].metadata["title"] == "Test Document" + assert results[0].metadata["type"] == "technical" + + +@pytest.mark.asyncio +async def test_update_document_index(search_service, test_document): + """Test updating an indexed document""" + # Initial indexing + await search_service.index_document(test_document, "Initial content") + + # Update with new content + await search_service.index_document(test_document, "Updated content with new terms") + + # Search for new terms + results = await search_service.search(SearchQuery(text="new terms")) + assert len(results) == 1 \ No newline at end of file