diff --git a/src/basic_memory/repository/postgres_search_repository.py b/src/basic_memory/repository/postgres_search_repository.py index 8e44f371..e5cecaed 100644 --- a/src/basic_memory/repository/postgres_search_repository.py +++ b/src/basic_memory/repository/postgres_search_repository.py @@ -705,6 +705,8 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, ) -> tuple[str, str, dict, str, str]: """Build Postgres FTS FROM/WHERE params shared by search and count.""" @@ -753,14 +755,28 @@ class PostgresSearchRepository(SearchRepositoryBase): else: conditions.append("search_index.permalink = :permalink") - # Handle search item type filter (parameterized for defense-in-depth) - if search_item_types: - type_placeholders = [] - for idx, t in enumerate(search_item_types): - param_name = f"search_type_{idx}" - params[param_name] = t.value - type_placeholders.append(f":{param_name}") - conditions.append(f"search_index.type IN ({', '.join(type_placeholders)})") + # Handle typed search row filters (parameterized for defense-in-depth) + self._append_in_filter( + conditions, + params, + column="search_index.type", + values=search_item_types, + param_prefix="search_type", + ) + self._append_in_filter( + conditions, + params, + column="search_index.category", + values=observation_categories, + param_prefix="observation_category", + ) + self._append_in_filter( + conditions, + params, + column="search_index.relation_type", + values=relation_types, + param_prefix="relation_type", + ) # Handle note type filter using JSONB containment (parameterized) if note_types: @@ -878,6 +894,8 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -894,6 +912,8 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, retrieval_mode=retrieval_mode, min_similarity=min_similarity, @@ -918,6 +938,8 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, ) @@ -1007,6 +1029,8 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -1021,6 +1045,8 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, retrieval_mode=retrieval_mode, min_similarity=min_similarity, @@ -1040,6 +1066,8 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, ) sql = f"SELECT COUNT(*) FROM {from_clause} WHERE {where_clause}" diff --git a/src/basic_memory/repository/search_repository.py b/src/basic_memory/repository/search_repository.py index 6bfab46e..a4c57eb4 100644 --- a/src/basic_memory/repository/search_repository.py +++ b/src/basic_memory/repository/search_repository.py @@ -41,6 +41,8 @@ class SearchRepository(Protocol): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -59,6 +61,8 @@ class SearchRepository(Protocol): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, diff --git a/src/basic_memory/repository/search_repository_base.py b/src/basic_memory/repository/search_repository_base.py index c5367295..0da5b4d2 100644 --- a/src/basic_memory/repository/search_repository_base.py +++ b/src/basic_memory/repository/search_repository_base.py @@ -218,6 +218,8 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[Dict[str, Any]] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -234,6 +236,8 @@ class SearchRepositoryBase(ABC): note_types: Filter by note types (from metadata.note_type) after_date: Filter by created_at > after_date search_item_types: Filter by SearchItemType (ENTITY, OBSERVATION, RELATION) + observation_categories: Filter observation rows by category + relation_types: Filter relation rows by relation_type metadata_filters: Structured frontmatter metadata filters limit: Maximum results to return offset: Number of results to skip @@ -256,6 +260,8 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[Dict[str, Any]] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -473,6 +479,26 @@ class SearchRepositoryBase(ABC): logger.debug(f"Query executed successfully in {elapsed_time:.2f}s.") return result + @staticmethod + def _append_in_filter( + conditions: list[str], + params: dict[str, Any], + *, + column: str, + values: list[Any] | None, + param_prefix: str, + ) -> None: + """Append a parameterized SQL IN clause for controlled column names.""" + if not values: + return + + placeholders: list[str] = [] + for idx, value in enumerate(values): + param_name = f"{param_prefix}_{idx}" + params[param_name] = value.value if isinstance(value, SearchItemType) else value + placeholders.append(f":{param_name}") + conditions.append(f"{column} IN ({', '.join(placeholders)})") + async def delete_entity_vector_rows(self, entity_id: int) -> None: """Delete one entity's derived vector rows using the backend's cleanup path.""" await self._ensure_vector_tables() @@ -1777,6 +1803,8 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]], after_date: Optional[datetime], search_item_types: Optional[List[SearchItemType]], + observation_categories: Optional[List[str]], + relation_types: Optional[List[str]], metadata_filters: Optional[dict], retrieval_mode: SearchRetrievalMode, min_similarity: Optional[float] = None, @@ -1810,6 +1838,8 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, min_similarity=min_similarity, limit=limit, @@ -1829,6 +1859,8 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, min_similarity=min_similarity, limit=limit, @@ -1858,6 +1890,8 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]], after_date: Optional[datetime], search_item_types: Optional[List[SearchItemType]], + observation_categories: Optional[List[str]], + relation_types: Optional[List[str]], metadata_filters: Optional[dict], min_similarity: Optional[float] = None, limit: int, @@ -1968,6 +2002,8 @@ class SearchRepositoryBase(ABC): note_types, after_date, search_item_types, + observation_categories, + relation_types, metadata_filters, ] ) @@ -1981,6 +2017,8 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, retrieval_mode=SearchRetrievalMode.FTS, limit=VECTOR_FILTER_SCAN_LIMIT, @@ -2131,6 +2169,8 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]], after_date: Optional[datetime], search_item_types: Optional[List[SearchItemType]], + observation_categories: Optional[List[str]], + relation_types: Optional[List[str]], metadata_filters: Optional[dict], min_similarity: Optional[float] = None, limit: int, @@ -2155,6 +2195,8 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, retrieval_mode=SearchRetrievalMode.FTS, limit=candidate_limit, @@ -2170,6 +2212,8 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, min_similarity=min_similarity, limit=candidate_limit, diff --git a/src/basic_memory/repository/sqlite_search_repository.py b/src/basic_memory/repository/sqlite_search_repository.py index 2eaec06c..3acb3aa4 100644 --- a/src/basic_memory/repository/sqlite_search_repository.py +++ b/src/basic_memory/repository/sqlite_search_repository.py @@ -714,6 +714,8 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, ) -> tuple[str, str, dict, str]: """Build SQLite FTS FROM/WHERE params shared by search and count.""" @@ -766,14 +768,28 @@ class SQLiteSearchRepository(SearchRepositoryBase): params["permalink"] = permalink_text match_conditions.append("search_index.permalink MATCH :permalink") - # Handle entity type filter (parameterized for defense-in-depth) - if search_item_types: - type_placeholders = [] - for idx, t in enumerate(search_item_types): - param_name = f"search_type_{idx}" - params[param_name] = t.value - type_placeholders.append(f":{param_name}") - conditions.append(f"search_index.type IN ({', '.join(type_placeholders)})") + # Handle typed search row filters (parameterized for defense-in-depth) + self._append_in_filter( + conditions, + params, + column="search_index.type", + values=search_item_types, + param_prefix="search_type", + ) + self._append_in_filter( + conditions, + params, + column="search_index.category", + values=observation_categories, + param_prefix="observation_category", + ) + self._append_in_filter( + conditions, + params, + column="search_index.relation_type", + values=relation_types, + param_prefix="relation_type", + ) # Handle note type filter (frontmatter type field, parameterized) if note_types: @@ -905,6 +921,8 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -921,6 +939,8 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, retrieval_mode=retrieval_mode, min_similarity=min_similarity, @@ -939,6 +959,8 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, ) @@ -1026,6 +1048,8 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + observation_categories: Optional[List[str]] = None, + relation_types: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -1040,6 +1064,8 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, retrieval_mode=retrieval_mode, min_similarity=min_similarity, @@ -1053,6 +1079,8 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + observation_categories=observation_categories, + relation_types=relation_types, metadata_filters=metadata_filters, ) sql = f"SELECT COUNT(*) FROM {from_clause} WHERE {where_clause}" diff --git a/src/basic_memory/schemas/search.py b/src/basic_memory/schemas/search.py index 51a00d83..7ee29897 100644 --- a/src/basic_memory/schemas/search.py +++ b/src/basic_memory/schemas/search.py @@ -46,6 +46,8 @@ class SearchQuery(BaseModel): - metadata_filters: Structured frontmatter filters (field -> value) - tags: Convenience frontmatter tag filter - status: Convenience frontmatter status filter + - observation_categories: Limit observation results to categories + - relation_types: Limit relation results to relationship types Boolean search examples: - "python AND flask" - Find items with both terms @@ -67,6 +69,8 @@ class SearchQuery(BaseModel): metadata_filters: Optional[dict[str, Any]] = None # Structured frontmatter filters tags: Optional[List[str]] = None # Convenience tag filter status: Optional[str] = None # Convenience status filter + observation_categories: Optional[List[str]] = None # Filter observations by category + relation_types: Optional[List[str]] = None # Filter relations by relation_type retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS min_similarity: Optional[float] = None # Per-query override for semantic_min_similarity @@ -85,6 +89,8 @@ class SearchQuery(BaseModel): status_is_empty = self.status is None or (isinstance(self.status, str) and not self.status) note_types_is_empty = not self.note_types entity_types_is_empty = not self.entity_types + observation_categories_is_empty = not self.observation_categories + relation_types_is_empty = not self.relation_types return ( self.permalink is None and self.permalink_match is None @@ -96,6 +102,8 @@ class SearchQuery(BaseModel): and metadata_is_empty and tags_is_empty and status_is_empty + and observation_categories_is_empty + and relation_types_is_empty ) def has_boolean_operators(self) -> bool: diff --git a/src/basic_memory/services/search_service.py b/src/basic_memory/services/search_service.py index 46b12ba0..0573a683 100644 --- a/src/basic_memory/services/search_service.py +++ b/src/basic_memory/services/search_service.py @@ -75,6 +75,8 @@ class _PreparedSearchQuery: title: str | None note_types: list[str] | None search_item_types: list[SearchItemType] | None + observation_categories: list[str] | None + relation_types: list[str] | None after_date: datetime | None metadata_filters: dict[str, Any] | None retrieval_mode: SearchRetrievalMode @@ -190,6 +192,8 @@ class SearchService: title=query.title, note_types=query.note_types, search_item_types=query.entity_types, + observation_categories=query.observation_categories, + relation_types=query.relation_types, after_date=after_date, metadata_filters=metadata_filters, retrieval_mode=query.retrieval_mode or SearchRetrievalMode.FTS, @@ -203,6 +207,8 @@ class SearchService: or prepared.title or prepared.note_types or prepared.search_item_types + or prepared.observation_categories + or prepared.relation_types or prepared.after_date or prepared.metadata_filters ) @@ -217,6 +223,8 @@ class SearchService: prepared.metadata_filters or prepared.note_types or prepared.search_item_types + or prepared.observation_categories + or prepared.relation_types or prepared.after_date ) @@ -235,6 +243,8 @@ class SearchService: title=prepared.title, note_types=prepared.note_types, search_item_types=prepared.search_item_types, + observation_categories=prepared.observation_categories, + relation_types=prepared.relation_types, after_date=prepared.after_date, metadata_filters=prepared.metadata_filters, retrieval_mode=prepared.retrieval_mode, @@ -256,6 +266,8 @@ class SearchService: title=prepared.title, note_types=prepared.note_types, search_item_types=prepared.search_item_types, + observation_categories=prepared.observation_categories, + relation_types=prepared.relation_types, after_date=prepared.after_date, metadata_filters=prepared.metadata_filters, retrieval_mode=prepared.retrieval_mode, diff --git a/tests/api/v2/test_search_router.py b/tests/api/v2/test_search_router.py index a7f93ce2..5dd39936 100644 --- a/tests/api/v2/test_search_router.py +++ b/tests/api/v2/test_search_router.py @@ -155,6 +155,43 @@ async def test_search_with_item_type_filter_returns_total( assert len(data["results"]) == 3 +@pytest.mark.asyncio +async def test_search_accepts_typed_facet_filters( + client: AsyncClient, + app, + v2_project_url: str, +): + """The v2 API contract should preserve observation and relation typed filters.""" + captured_queries = [] + + class CapturingSearchService: + async def search(self, query, *, limit, offset): + captured_queries.append(("search", query, limit, offset)) + return [] + + async def count(self, query): + captured_queries.append(("count", query)) + return 0 + + app.dependency_overrides[get_search_service_v2_external] = lambda: CapturingSearchService() + try: + response = await client.post( + f"{v2_project_url}/search/", + json={ + "entity_types": ["observation"], + "observation_categories": ["appearance"], + "relation_types": ["associated_with"], + }, + ) + finally: + app.dependency_overrides.pop(get_search_service_v2_external, None) + + assert response.status_code == 200 + _, query, _, _ = captured_queries[0] + assert query.observation_categories == ["appearance"] + assert query.relation_types == ["associated_with"] + + @pytest.mark.asyncio async def test_search_by_permalink( client: AsyncClient, diff --git a/tests/repository/test_hybrid_fusion.py b/tests/repository/test_hybrid_fusion.py index c0226f0d..bce0f549 100644 --- a/tests/repository/test_hybrid_fusion.py +++ b/tests/repository/test_hybrid_fusion.py @@ -71,6 +71,8 @@ class ConcreteSearchRepo(SearchRepositoryBase): note_types: Optional[list[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[list[SearchItemType]] = None, + observation_categories: Optional[list[str]] = None, + relation_types: Optional[list[str]] = None, metadata_filters: Optional[dict[str, Any]] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -126,6 +128,8 @@ HYBRID_KWARGS: dict[str, Any] = dict( note_types=None, after_date=None, search_item_types=None, + observation_categories=None, + relation_types=None, metadata_filters=None, limit=10, offset=0, diff --git a/tests/repository/test_search_repository.py b/tests/repository/test_search_repository.py index 27e8b8da..41d82b76 100644 --- a/tests/repository/test_search_repository.py +++ b/tests/repository/test_search_repository.py @@ -1044,3 +1044,111 @@ async def test_search_item_types_parameterized(search_repository): results = await search_repository.search(search_item_types=[SearchItemType.ENTITY]) # Should not raise — parameterized query handles enum values safely assert isinstance(results, list) + + +@pytest.mark.asyncio +async def test_search_filters_observation_categories(search_repository, search_entity): + """Observation category facets should filter observation search rows directly.""" + now = datetime.now(timezone.utc) + await search_repository.bulk_index_items( + [ + SearchIndexRow( + id=search_entity.id + 100, + type=SearchItemType.OBSERVATION.value, + title="appearance: Broad shoulders", + content_stems="shared observation content", + content_snippet="Broad shoulders", + permalink=f"{search_entity.permalink}/observations/appearance", + file_path=search_entity.file_path, + entity_id=search_entity.id, + category="appearance", + metadata={"facet_test": True}, + created_at=now, + updated_at=now, + project_id=search_repository.project_id, + ), + SearchIndexRow( + id=search_entity.id + 101, + type=SearchItemType.OBSERVATION.value, + title="trait: Stoic resolve", + content_stems="shared observation content", + content_snippet="Stoic resolve", + permalink=f"{search_entity.permalink}/observations/trait", + file_path=search_entity.file_path, + entity_id=search_entity.id, + category="trait", + metadata={"facet_test": True}, + created_at=now, + updated_at=now, + project_id=search_repository.project_id, + ), + ] + ) + + results = await search_repository.search( + search_item_types=[SearchItemType.OBSERVATION], + observation_categories=["appearance"], + ) + total = await search_repository.count( + search_item_types=[SearchItemType.OBSERVATION], + observation_categories=["appearance"], + ) + + assert total == 1 + assert [(result.category, result.title) for result in results] == [ + ("appearance", "appearance: Broad shoulders") + ] + + +@pytest.mark.asyncio +async def test_search_filters_relation_types(search_repository, search_entity): + """Relationship type facets should filter relation search rows directly.""" + now = datetime.now(timezone.utc) + await search_repository.bulk_index_items( + [ + SearchIndexRow( + id=search_entity.id + 200, + type=SearchItemType.RELATION.value, + title="Starbuck -> The Pequod", + content_stems="shared relation content", + content_snippet="Starbuck serves on The Pequod", + permalink=f"{search_entity.permalink}/relations/serves-on", + file_path=search_entity.file_path, + entity_id=search_entity.id, + relation_type="serves_on", + metadata={"facet_test": True}, + created_at=now, + updated_at=now, + project_id=search_repository.project_id, + ), + SearchIndexRow( + id=search_entity.id + 201, + type=SearchItemType.RELATION.value, + title="Starbuck -> Ahab", + content_stems="shared relation content", + content_snippet="Starbuck contrasts with Ahab", + permalink=f"{search_entity.permalink}/relations/contrasts-with", + file_path=search_entity.file_path, + entity_id=search_entity.id, + relation_type="contrasts_with", + metadata={"facet_test": True}, + created_at=now, + updated_at=now, + project_id=search_repository.project_id, + ), + ] + ) + + results = await search_repository.search( + search_item_types=[SearchItemType.RELATION], + relation_types=["contrasts_with"], + ) + total = await search_repository.count( + search_item_types=[SearchItemType.RELATION], + relation_types=["contrasts_with"], + ) + + assert total == 1 + assert [(result.relation_type, result.title) for result in results] == [ + ("contrasts_with", "Starbuck -> Ahab") + ] diff --git a/tests/repository/test_semantic_search_base.py b/tests/repository/test_semantic_search_base.py index ba4e08ef..2aabd036 100644 --- a/tests/repository/test_semantic_search_base.py +++ b/tests/repository/test_semantic_search_base.py @@ -83,6 +83,8 @@ class _ConcreteRepo(SearchRepositoryBase): note_types: list[str] | None = None, after_date: datetime | None = None, search_item_types: list[SearchItemType] | None = None, + observation_categories: list[str] | None = None, + relation_types: list[str] | None = None, metadata_filters: dict[str, Any] | None = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: float | None = None, diff --git a/tests/repository/test_vector_pagination.py b/tests/repository/test_vector_pagination.py index 821ce84a..71f57df5 100644 --- a/tests/repository/test_vector_pagination.py +++ b/tests/repository/test_vector_pagination.py @@ -56,6 +56,8 @@ class ConcreteSearchRepo(SearchRepositoryBase): note_types: list[str] | None = None, after_date: datetime | None = None, search_item_types: list[SearchItemType] | None = None, + observation_categories: list[str] | None = None, + relation_types: list[str] | None = None, metadata_filters: dict[str, Any] | None = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: float | None = None, @@ -158,6 +160,8 @@ async def test_page1_scores_gte_page2_scores(): note_types=None, after_date=None, search_item_types=None, + observation_categories=None, + relation_types=None, metadata_filters=None, limit=limit, offset=offset, diff --git a/tests/repository/test_vector_threshold.py b/tests/repository/test_vector_threshold.py index f4da09ed..071f57e8 100644 --- a/tests/repository/test_vector_threshold.py +++ b/tests/repository/test_vector_threshold.py @@ -60,6 +60,8 @@ class ConcreteSearchRepo(SearchRepositoryBase): note_types: Optional[list[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[list[SearchItemType]] = None, + observation_categories: Optional[list[str]] = None, + relation_types: Optional[list[str]] = None, metadata_filters: Optional[dict[str, Any]] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -130,6 +132,8 @@ COMMON_SEARCH_KWARGS: dict[str, Any] = dict( note_types=None, after_date=None, search_item_types=None, + observation_categories=None, + relation_types=None, metadata_filters=None, limit=10, offset=0,