diff --git a/src/basic_memory/api/routers/__init__.py b/src/basic_memory/api/routers/__init__.py index 0f6c91cb..48abda72 100644 --- a/src/basic_memory/api/routers/__init__.py +++ b/src/basic_memory/api/routers/__init__.py @@ -2,6 +2,6 @@ from . import knowledge_router as knowledge from . import discovery_router as discovery -from . import discovery_router as memory +from . import memory_router as memory __all__ = ["knowledge", "discovery", "memory"] diff --git a/src/basic_memory/api/routers/memory_router.py b/src/basic_memory/api/routers/memory_router.py index e102726b..c4faecbe 100644 --- a/src/basic_memory/api/routers/memory_router.py +++ b/src/basic_memory/api/routers/memory_router.py @@ -1,13 +1,15 @@ """Routes for memory:// URI operations.""" - +from dataclasses import asdict from typing import List, Optional from datetime import datetime, timedelta from fastapi import APIRouter +from basic_memory.config import config from basic_memory.schemas.memory import MemoryUrl, GraphContext from basic_memory.deps import ContextServiceDep +from basic_memory.schemas.search import SearchResult, RelatedResult -router = APIRouter(prefix="/memory") +router = APIRouter(prefix="/memory", tags=["memory"]) def parse_timeframe(timeframe: str) -> Optional[datetime]: @@ -36,8 +38,9 @@ async def get_memory_context( timeframe: str = "7d", ) -> GraphContext: """Get rich context from memory:// URI.""" + # add the project name from the config to the url as the "host # Parse URI - memory_url = MemoryUrl.parse(f"memory://{uri}") + memory_url = MemoryUrl.parse(f"memory://{config.project}/{uri}") # Parse timeframe since = parse_timeframe(timeframe) @@ -45,8 +48,11 @@ async def get_memory_context( # Build context context = await context_service.build_context(str(memory_url), depth=depth, since=since) + primary_entities = [SearchResult(**asdict(r)) for r in context["primary_entities"]] + related_entities = [RelatedResult(**asdict(r)) for r in context["related_entities"]] + metadata = context["metadata"] # Transform to GraphContext - return GraphContext.model_validate(context) + return GraphContext(primary_entities=primary_entities, related_entities=related_entities, metadata=metadata) @router.get("/related/{permalink}", response_model=GraphContext) diff --git a/src/basic_memory/api/routers/search_router.py b/src/basic_memory/api/routers/search_router.py index f8e2159c..0e85b98b 100644 --- a/src/basic_memory/api/routers/search_router.py +++ b/src/basic_memory/api/routers/search_router.py @@ -1,4 +1,5 @@ """Router for search operations.""" +from dataclasses import asdict from fastapi import APIRouter, Depends, BackgroundTasks from typing import List @@ -17,7 +18,8 @@ async def search( ): """Search across all knowledge and documents.""" results = await search_service.search(query) - return SearchResponse(results=results) + search_results = [SearchResult.model_validate(asdict(r)) for r in results] + return SearchResponse(results=search_results) @router.post("/reindex") async def reindex( diff --git a/src/basic_memory/config.py b/src/basic_memory/config.py index 04275597..0197e29e 100644 --- a/src/basic_memory/config.py +++ b/src/basic_memory/config.py @@ -18,6 +18,10 @@ class ProjectConfig(BaseSettings): description="Base path for basic-memory files", ) + # Name of the project + project: str = Field(default="default", description="Project name") + + model_config = SettingsConfigDict( env_prefix="BASIC_MEMORY_", extra="ignore", @@ -25,7 +29,6 @@ class ProjectConfig(BaseSettings): env_file_encoding="utf-8", ) - @property def database_path(self) -> Path: """Get SQLite database path.""" diff --git a/src/basic_memory/repository/search_repository.py b/src/basic_memory/repository/search_repository.py index 32383309..5b3a4198 100644 --- a/src/basic_memory/repository/search_repository.py +++ b/src/basic_memory/repository/search_repository.py @@ -1,6 +1,7 @@ """Repository for search operations.""" import json +from dataclasses import dataclass from typing import List, Optional, Any, Dict from loguru import logger @@ -10,7 +11,26 @@ from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from basic_memory import db from basic_memory.models.search import CREATE_SEARCH_INDEX from basic_memory.repository.repository import Repository -from basic_memory.schemas.search import SearchQuery, SearchResult, SearchItemType +from basic_memory.schemas.search import SearchQuery, SearchItemType + +@dataclass +class SearchResultRow(): + """Search result with score and metadata.""" + id: int + type: str + score: float + metadata: dict + + # Common fields + permalink: Optional[str] = None + file_path: Optional[str] = None + + # Type-specific fields + entity_id: Optional[int] = None # For observations + category: Optional[str] = None # For observations + from_id: Optional[int] = None # For relations + to_id: Optional[int] = None # For relations + relation_type: Optional[str] = None # For relations class SearchRepository: @@ -31,7 +51,7 @@ class SearchRepository: return f'"{term}"' return term - async def search(self, query: SearchQuery) -> List[SearchResult]: + async def search(self, query: SearchQuery) -> List[SearchResultRow]: """Search across all indexed content with fuzzy matching.""" conditions = [] params = {} @@ -88,20 +108,18 @@ class SearchRepository: WHERE {where_clause} ORDER BY score ASC """ - - logger.debug(f"Search query: {sql}") + logger.debug(f"Search params: {params}") - async with db.scoped_session(self.session_maker) as session: result = await session.execute(text(sql), params) rows = result.fetchall() - return [ - SearchResult( + results = [ + SearchResultRow( id=row.id, permalink=row.permalink, file_path=row.file_path, - type=SearchItemType(row.type), + type=row.type, score=row.score, metadata=json.loads(row.metadata), from_id=row.from_id, @@ -112,6 +130,10 @@ class SearchRepository: ) for row in rows ] + + logger.debug(f"Search results: {results}") + return results + async def index_item( self, @@ -157,7 +179,7 @@ class SearchRepository: "content": content, "permalink": permalink, "file_path": file_path, - "type": type.value, + "type": type, "metadata": json.dumps(metadata), "from_id": from_id, "to_id": to_id, @@ -168,7 +190,7 @@ class SearchRepository: "updated_at": metadata.get("updated_at") }, ) - logger.debug(f"indexed {permalink}") + logger.debug(f"indexed permalink {permalink}") await session.commit() async def delete_by_permalink(self, permalink: str): @@ -187,7 +209,7 @@ class SearchRepository: use_query_options:bool = True ) -> Result[Any]: """Execute a query asynchronously.""" - logger.debug(f"Executing query: {query}") + #logger.debug(f"Executing query: {query}") async with db.scoped_session(self.session_maker) as session: if params: result = await session.execute(query, params) diff --git a/src/basic_memory/schemas/memory.py b/src/basic_memory/schemas/memory.py index bb1841fc..aa1e8b6f 100644 --- a/src/basic_memory/schemas/memory.py +++ b/src/basic_memory/schemas/memory.py @@ -1,10 +1,11 @@ """Schemas for memory context.""" from typing import Dict, List, Optional, Any + from pydantic import BaseModel, Field from pydantic import field_validator -from basic_memory.schemas.search import SearchResult +from basic_memory.schemas.search import SearchResult, RelatedResult """Memory URL schema for knowledge addressing. @@ -101,7 +102,7 @@ class GraphContext(BaseModel): primary_entities: List[SearchResult] = Field(description="Entities directly matching URI") # Related entities - related_entities: List[SearchResult] = Field(description="Entities found via relations") + related_entities: List[RelatedResult] = Field(description="Entities found via relations") # Context metadata metadata: Dict[str, Any] = Field( diff --git a/src/basic_memory/schemas/search.py b/src/basic_memory/schemas/search.py index 4f7a3538..1dbd172b 100644 --- a/src/basic_memory/schemas/search.py +++ b/src/basic_memory/schemas/search.py @@ -57,8 +57,8 @@ class SearchResult(BaseModel): """Search result with score and metadata.""" id: int type: SearchItemType - score: float - metadata: dict + score: Optional[float] = None + metadata: Optional[dict] = None # Common fields permalink: Optional[str] = None @@ -71,6 +71,21 @@ class SearchResult(BaseModel): to_id: Optional[int] = None # For relations relation_type: Optional[str] = None # For relations +class RelatedResult(BaseModel): + type: SearchItemType + id: int + title: str + permalink: str + depth:int + root_id: int + created_at: datetime + from_id: Optional[int] = None + to_id: Optional[int] = None + relation_type: Optional[str] = None + category: Optional[str] = None + entity_id: Optional[int] = None + content: Optional[str] = None + class SearchResponse(BaseModel): """Wrapper for search results.""" diff --git a/src/basic_memory/services/context_service.py b/src/basic_memory/services/context_service.py index 8236b7d6..1cc940ef 100644 --- a/src/basic_memory/services/context_service.py +++ b/src/basic_memory/services/context_service.py @@ -1,5 +1,5 @@ """Service for building rich context from the knowledge graph.""" - +from dataclasses import dataclass from datetime import datetime, UTC, timezone from typing import List, Optional, Tuple from loguru import logger @@ -10,6 +10,22 @@ from basic_memory.repository.entity_repository import EntityRepository from basic_memory.schemas.memory import MemoryUrl from basic_memory.schemas.search import SearchQuery, SearchItemType +@dataclass +class ContextResultRow: + type: str + id: int + title: str + permalink: str + depth:int + root_id: int + created_at: datetime + from_id: Optional[int] = None + to_id: Optional[int] = None + relation_type: Optional[str] = None + category: Optional[str] = None + entity_id: Optional[int] = None + content: Optional[str] = None + class ContextService: """Service for building rich context from memory:// URIs. @@ -44,16 +60,24 @@ class ContextService: if memory_url.params.get("type") == "related": # Special mode for finding related content target = memory_url.params["target"] - primary = await self.find_related(target) + logger.debug(f"Finding related content for '{target}'") + primary = await self.find_related_1(target) elif memory_url.pattern: # Pattern matching with * + logger.debug(f"Pattern matching for '{memory_url.pattern}'") primary = await self.find_by_pattern(memory_url.pattern) else: # Direct permalink lookup + logger.debug(f"Direct permalink lookup for '{memory_url.relative_path()}'") primary = await self.find_by_permalink(memory_url.relative_path()) + logger.debug(f"Found {len(primary)} primary entities") + for p in primary: + logger.debug(f"Found primary entity: {p}") + # Get type_id pairs for traversal type_id_pairs = [(r.type, r.id) for r in primary] if primary else [] + logger.debug(f"type_id_pairs: {type_id_pairs}") # Find connected content related = await self.find_connected( @@ -61,6 +85,10 @@ class ContextService: max_depth=depth, since=since ) + logger.debug(f"Found {len(related)} related entities") + for r in related: + logger.debug(f"Found related entity: {r}") + # Build response return { @@ -87,15 +115,15 @@ class ContextService: query = SearchQuery(permalink=permalink) return await self.search_repository.search(query) - async def find_related(self, permalink: str): + async def find_related_1(self, permalink: str): """Find entities related to a given permalink.""" # First find the target entity target = await self.find_by_permalink(permalink) if not target: return [] - # Use find_connected to get related items - type_id_pairs = [(r.type.value, r.id) for r in target] + # Use find_connected to get related items at depth=1 + type_id_pairs = [(r.type, r.id) for r in target] return await self.find_connected( type_id_pairs, max_depth=1 # Only immediate relations @@ -177,7 +205,7 @@ class ContextService: r1.type = 'relation' AND (r1.from_id = cg.id OR r1.to_id = cg.id) {r1_date_filter} - ) + ) -- Then join to ALL related items at the same depth JOIN search_index related ON ( -- The found relation @@ -217,5 +245,25 @@ class ContextService: ORDER BY depth, type, id """) - results = await self.search_repository.execute_query(query, params=params) - return results.all() \ No newline at end of file + result = await self.search_repository.execute_query(query, params=params) + rows = result.all() + + context_rows = [ + ContextResultRow( + type=row.type, + id=row.id, + title=row.title, + permalink=row.permalink, + from_id=row.from_id, + to_id=row.to_id, + relation_type=row.relation_type, + category=row.category, + entity_id=row.entity_id, + content=row.content, + depth=row.depth, + root_id=row.root_id, + created_at=row.created_at, + ) + for row in rows + ] + return context_rows \ 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 8afa9701..36551aae 100644 --- a/src/basic_memory/services/search_service.py +++ b/src/basic_memory/services/search_service.py @@ -155,7 +155,7 @@ class SearchService: content=entity_content, permalink=entity.permalink, file_path=entity.file_path, - type=SearchItemType.ENTITY, + type=SearchItemType.ENTITY.value, metadata={ "entity_type": entity.entity_type, "created_at": entity.created_at.isoformat(), @@ -169,7 +169,7 @@ class SearchService: content=entity_content, permalink=entity.permalink, file_path=entity.file_path, - type=SearchItemType.ENTITY, + type=SearchItemType.ENTITY.value, metadata={ "entity_type": entity.entity_type, "created_at": entity.created_at.isoformat(), @@ -193,7 +193,7 @@ class SearchService: content=obs.content, permalink=observation_permalink, file_path=entity.file_path, - type=SearchItemType.OBSERVATION, + type=SearchItemType.OBSERVATION.value, category=obs.category, entity_id=entity.id, metadata={ @@ -209,7 +209,7 @@ class SearchService: content=obs.content, permalink=observation_permalink, file_path=entity.file_path, - type=SearchItemType.OBSERVATION, + type=SearchItemType.OBSERVATION.value, category=obs.category, entity_id=entity.id, metadata={ @@ -237,7 +237,7 @@ class SearchService: content=rel.context or "", permalink=relation_permalink, file_path=entity.file_path, - type=SearchItemType.RELATION, + type=SearchItemType.RELATION.value, from_id=rel.from_id, to_id=rel.to_id, relation_type=rel.relation_type, @@ -253,7 +253,7 @@ class SearchService: content=rel.context or "", permalink=relation_permalink, file_path=entity.file_path, - type=SearchItemType.RELATION, + type=SearchItemType.RELATION.value, from_id=rel.from_id, to_id=rel.to_id, relation_type=rel.relation_type, diff --git a/tests/api/test_memory_router.py b/tests/api/test_memory_router.py new file mode 100644 index 00000000..aef0d127 --- /dev/null +++ b/tests/api/test_memory_router.py @@ -0,0 +1,103 @@ +"""Tests for memory router endpoints.""" + +import pytest + +from basic_memory.schemas.memory import GraphContext + + +@pytest.mark.asyncio +async def test_get_memory_context(client, test_graph): + """Test getting context from memory URL.""" + response = await client.get("/memory/test/root") + assert response.status_code == 200 + + context = GraphContext(**response.json()) + assert len(context.primary_entities) == 1 + assert context.primary_entities[0].permalink == "test/root" + assert len(context.related_entities) > 0 + + # Verify metadata + assert context.metadata["uri"] == "memory://default/test/root" + assert context.metadata["depth"] == 1 # default depth + #assert context.metadata["timeframe"] == "7d" # default timeframe + assert isinstance(context.metadata["generated_at"], str) + assert context.metadata["matched_entities"] == 1 + + +@pytest.mark.asyncio +async def test_get_memory_context_pattern(client, test_graph): + """Test getting context with pattern matching.""" + response = await client.get("/memory/test/*") + assert response.status_code == 200 + + context = GraphContext(**response.json()) + assert len(context.primary_entities) > 1 # Should match multiple test/* paths + assert all("test/" in e.permalink for e in context.primary_entities) + + +@pytest.mark.asyncio +async def test_get_memory_context_depth(client, test_graph): + """Test depth parameter affects relation traversal.""" + # With depth=1, should only get immediate connections + response = await client.get("/memory/test/root?depth=1") + assert response.status_code == 200 + context1 = GraphContext(**response.json()) + + # With depth=2, should get deeper connections + response = await client.get("/memory/test/root?depth=2") + assert response.status_code == 200 + context2 = GraphContext(**response.json()) + + assert len(context2.related_entities) > len(context1.related_entities) + + +@pytest.mark.asyncio +async def test_get_memory_context_timeframe(client, test_graph): + """Test timeframe parameter filters by date.""" + # Recent timeframe + response = await client.get("/memory/test/root?timeframe=1d") + assert response.status_code == 200 + recent = GraphContext(**response.json()) + + # Longer timeframe + response = await client.get("/memory/test/root?timeframe=30d") + assert response.status_code == 200 + older = GraphContext(**response.json()) + + assert len(older.related_entities) >= len(recent.related_entities) + + +@pytest.mark.asyncio +async def test_get_related_context(client, test_graph): + """Test getting related content.""" + response = await client.get("/memory/related/test/root") + assert response.status_code == 200 + + context = GraphContext(**response.json()) + assert len(context.primary_entities) > 0 + assert any("connected1" in e.permalink for e in context.related_entities) + assert any("connected2" in e.permalink for e in context.related_entities) + +@pytest.mark.asyncio +async def test_get_related_context_filters(client, test_graph): + """Test filtering related content by relation type.""" + response = await client.get("/memory/related/test/root?relation_types=connects_to") + assert response.status_code == 200 + + context = GraphContext(**response.json()) + for relation in context.related_entities: + if relation.type == "relation": + assert relation.relation_type == "connects_to" + + + + +@pytest.mark.asyncio +async def test_not_found(client): + """Test handling of non-existent paths.""" + response = await client.get("/memory/test/does-not-exist") + assert response.status_code == 200 + + context = GraphContext(**response.json()) + assert len(context.primary_entities) == 0 + assert len(context.related_entities) == 0 diff --git a/tests/conftest.py b/tests/conftest.py index 7346421d..36504d15 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -252,3 +252,81 @@ async def full_entity(sample_entity, entity_repository): search_entity.observations = observations search_entity.outgoing_relations = relations return await entity_repository.add(search_entity) + +@pytest_asyncio.fixture +async def test_graph(entity_repository, search_service): + """Create a test knowledge graph with entities, relations and observations.""" + # Create some test entities + entities = [ + Entity( + title="Root Entity", + entity_type="test", + permalink="test/root", + file_path="test/root.md", + content_type="text/markdown", + ), + Entity( + title="Connected Entity 1", + entity_type="test", + permalink="test/connected1", + file_path="test/connected1.md", + content_type="text/markdown", + ), + Entity( + title="Connected Entity 2", + entity_type="test", + permalink="test/connected2", + file_path="test/connected2.md", + content_type="text/markdown", + ), + Entity( + title="Deep Entity", + entity_type="test", + permalink="test/deep", + file_path="test/deep.md", + content_type="text/markdown", + ), + ] + entities = await entity_repository.add_all(entities) + root, conn1, conn2, deep = entities + + # Add some observations + root.observations = [ + Observation(content="Root note 1", category=ObservationCategory.NOTE), + Observation(content="Root tech note", category=ObservationCategory.TECH), + ] + + conn1.observations = [ + Observation(content="Connected 1 note", category=ObservationCategory.NOTE) + ] + + # Add relations + relations = [ + # Direct connections to root + Relation(from_id=root.id, to_id=conn1.id, relation_type="connects_to"), + Relation(from_id=conn2.id, to_id=root.id, relation_type="connected_from"), + # Deep connection + Relation(from_id=conn1.id, to_id=deep.id, relation_type="deep_connection"), + ] + + root.outgoing_relations = [relations[0]] + conn1.outgoing_relations = [relations[2]] + conn2.outgoing_relations = [relations[1]] + + # Save relations + root = await entity_repository.add(root) + conn1 = await entity_repository.add(conn1) + conn2 = await entity_repository.add(conn2) + + # Index everything for search + for entity in entities: + await search_service.index_entity(entity) + + return { + "root": root, + "connected1": conn1, + "connected2": conn2, + "deep": deep, + "observations": root.observations + conn1.observations, + "relations": relations, + } diff --git a/tests/schemas/test_memory_url.py b/tests/schemas/test_memory_url.py index f0d36452..62fd9fca 100644 --- a/tests/schemas/test_memory_url.py +++ b/tests/schemas/test_memory_url.py @@ -2,7 +2,7 @@ import pytest from pydantic import ValidationError -from basic_memory.schemas.memory_url import MemoryUrl +from basic_memory.schemas.memory import MemoryUrl def test_basic_permalink(): diff --git a/tests/services/test_context_service.py b/tests/services/test_context_service.py index fb4b336a..8482f0d1 100644 --- a/tests/services/test_context_service.py +++ b/tests/services/test_context_service.py @@ -5,7 +5,6 @@ from datetime import datetime, timedelta, UTC import pytest import pytest_asyncio -from basic_memory.models import Entity, Relation, Observation, ObservationCategory from basic_memory.schemas.search import SearchItemType from basic_memory.services.context_service import ContextService @@ -16,85 +15,6 @@ async def context_service(search_repository, entity_repository): return ContextService(search_repository, entity_repository) -@pytest_asyncio.fixture -async def test_graph(entity_repository, search_service): - """Create a test knowledge graph with entities, relations and observations.""" - # Create some test entities - entities = [ - Entity( - title="Root Entity", - entity_type="test", - permalink="test/root", - file_path="test/root.md", - content_type="text/markdown", - ), - Entity( - title="Connected Entity 1", - entity_type="test", - permalink="test/connected1", - file_path="test/connected1.md", - content_type="text/markdown", - ), - Entity( - title="Connected Entity 2", - entity_type="test", - permalink="test/connected2", - file_path="test/connected2.md", - content_type="text/markdown", - ), - Entity( - title="Deep Entity", - entity_type="test", - permalink="test/deep", - file_path="test/deep.md", - content_type="text/markdown", - ), - ] - entities = await entity_repository.add_all(entities) - root, conn1, conn2, deep = entities - - # Add some observations - root.observations = [ - Observation(content="Root note 1", category=ObservationCategory.NOTE), - Observation(content="Root tech note", category=ObservationCategory.TECH), - ] - - conn1.observations = [ - Observation(content="Connected 1 note", category=ObservationCategory.NOTE) - ] - - # Add relations - relations = [ - # Direct connections to root - Relation(from_id=root.id, to_id=conn1.id, relation_type="connects_to"), - Relation(from_id=conn2.id, to_id=root.id, relation_type="connected_from"), - # Deep connection - Relation(from_id=conn1.id, to_id=deep.id, relation_type="deep_connection"), - ] - - root.outgoing_relations = [relations[0]] - conn1.outgoing_relations = [relations[2]] - conn2.outgoing_relations = [relations[1]] - - # Save relations - root = await entity_repository.add(root) - conn1 = await entity_repository.add(conn1) - conn2 = await entity_repository.add(conn2) - - # Index everything for search - for entity in entities: - await search_service.index_entity(entity) - - return { - "root": root, - "connected1": conn1, - "connected2": conn2, - "deep": deep, - "observations": root.observations + conn1.observations, - "relations": relations, - } - - @pytest.mark.asyncio async def test_find_by_pattern(context_service, test_graph): """Test pattern matching.""" @@ -114,7 +34,7 @@ async def test_find_by_permalink(context_service, test_graph): @pytest.mark.asyncio async def test_find_related(context_service, test_graph): """Test finding related content.""" - results = await context_service.find_related("test/root") + results = await context_service.find_related_1("test/root") # Should get immediate connections assert any("connected1" in r.permalink for r in results) assert any("connected2" in r.permalink for r in results) @@ -229,6 +149,7 @@ async def test_find_connected_timeframe(context_service, test_graph, search_repo entity_ids = {r.id for r in results if r.type == "entity"} assert len(entity_ids) == 0 # No accessible entities within timeframe + @pytest.mark.asyncio async def test_build_context_pattern(context_service, test_graph): """Test building context from pattern.""" @@ -237,35 +158,29 @@ async def test_build_context_pattern(context_service, test_graph): assert "uri" in context["metadata"] assert "total_entities" in context["metadata"] + @pytest.mark.asyncio async def test_build_context_related(context_service, test_graph): """Test building context from related mode.""" - context = await context_service.build_context( - "memory://project/related/test/root" - ) + context = await context_service.build_context("memory://project/related/test/root") assert len(context["primary_entities"]) > 0 assert len(context["related_entities"]) > 0 - - + + @pytest.mark.asyncio async def test_build_context_not_found(context_service): """Test handling non-existent permalinks.""" - context = await context_service.build_context( - "memory://project/does/not/exist" - ) + context = await context_service.build_context("memory://project/does/not/exist") assert len(context["primary_entities"]) == 0 assert len(context["related_entities"]) == 0 - - + + @pytest.mark.asyncio async def test_context_metadata(context_service, test_graph): """Test metadata is correctly populated.""" - context = await context_service.build_context( - "memory://project/test/root", - depth=2 - ) + context = await context_service.build_context("memory://project/test/root", depth=2) metadata = context["metadata"] assert metadata["uri"] == "memory://project/test/root" assert metadata["depth"] == 2 assert metadata["generated_at"] is not None - assert metadata["matched_entities"] > 0 \ No newline at end of file + assert metadata["matched_entities"] > 0 diff --git a/tests/services/test_search_service.py b/tests/services/test_search_service.py index b0a7a326..91da43ec 100644 --- a/tests/services/test_search_service.py +++ b/tests/services/test_search_service.py @@ -1,4 +1,5 @@ """Tests for search service.""" +from datetime import datetime, timezone import pytest import pytest_asyncio @@ -165,6 +166,30 @@ async def test_filters(indexed_search): assert all(r.metadata.get("entity_type") == "component" for r in results) +@pytest.mark.asyncio +async def test_after_date(indexed_search): + """Test search filters.""" + + # Should find with past date + past_date = datetime(2020, 1, 1) + results = await indexed_search.search( + SearchQuery( + text="service", + after_date=past_date.isoformat(), + ) + ) + assert all(datetime.fromisoformat(r.metadata['created_at']) > past_date for r in results) + + # Should not find with future date + future_date = datetime(2030, 1, 1) + results = await indexed_search.search( + SearchQuery( + text="service", + after_date=future_date.isoformat(), + ) + ) + assert len(results) == 0 + @pytest.mark.asyncio async def test_no_criteria(indexed_search): """Test search with no criteria returns empty list."""