From 4fe6fe09c8bd36fd206f01b1900208d63c05433d Mon Sep 17 00:00:00 2001 From: Paul Hernandez <60959+phernandez@users.noreply.github.com> Date: Sun, 7 Jun 2026 19:57:19 -0500 Subject: [PATCH] feat(core): add observation category filter to search (#908) Signed-off-by: phernandez Co-authored-by: Claude Opus 4.8 (1M context) --- .../api/v2/routers/search_router.py | 4 +- src/basic_memory/mcp/tools/search.py | 55 ++++++++++-- .../repository/postgres_search_repository.py | 21 +++++ .../repository/search_repository.py | 2 + .../repository/search_repository_base.py | 12 +++ .../repository/sqlite_search_repository.py | 21 +++++ src/basic_memory/schemas/search.py | 4 + src/basic_memory/services/search_service.py | 6 ++ tests/mcp/test_tool_contracts.py | 1 + tests/mcp/test_tool_search.py | 87 +++++++++++++++++++ tests/mcp/test_tool_telemetry.py | 1 + tests/repository/test_hybrid_fusion.py | 2 + .../test_postgres_search_repository.py | 69 +++++++++++++++ tests/repository/test_search_repository.py | 80 +++++++++++++++++ tests/repository/test_semantic_search_base.py | 1 + tests/repository/test_vector_pagination.py | 2 + tests/repository/test_vector_threshold.py | 2 + tests/services/test_search_service.py | 33 +++++++ 18 files changed, 395 insertions(+), 8 deletions(-) diff --git a/src/basic_memory/api/v2/routers/search_router.py b/src/basic_memory/api/v2/routers/search_router.py index be80d4b4..cd48c0e0 100644 --- a/src/basic_memory/api/v2/routers/search_router.py +++ b/src/basic_memory/api/v2/routers/search_router.py @@ -64,7 +64,9 @@ async def search( or query.permalink or query.permalink_match ), - has_filters=bool(query.note_types or query.entity_types or query.metadata_filters), + has_filters=bool( + query.note_types or query.entity_types or query.categories or query.metadata_filters + ), ): offset = (page - 1) * page_size exact_count_available = query.retrieval_mode == SearchRetrievalMode.FTS diff --git a/src/basic_memory/mcp/tools/search.py b/src/basic_memory/mcp/tools/search.py index e47c687f..17af0b3f 100644 --- a/src/basic_memory/mcp/tools/search.py +++ b/src/basic_memory/mcp/tools/search.py @@ -155,7 +155,7 @@ def _format_search_error_response( - Boolean NOT: `project NOT archived` - Grouped: `(project OR planning) AND notes` - Exact phrases: `"weekly standup meeting"` - - Content-specific: `tag:example` or `category:observation` + - Content-specific: `tag:example` ## Try again with: ``` @@ -222,7 +222,7 @@ def _format_search_error_response( 6. **Try advanced search patterns**: - Tag search: `search_notes("{project}","tag:your-tag")` - - Category search: `search_notes("{project}","category:observation")` + - Observation category: `search_notes("{project}","{query}", entity_types=["observation"], categories=["requirement"])` - Pattern matching: `search_notes("{project}","*{query}*", search_type="permalink")` ## Explore what content exists: @@ -299,7 +299,8 @@ Error searching for '{query}': {error_message} - **Boolean**: `term1 AND term2`, `term1 OR term2`, `term1 NOT term2` - **Phrases**: `"exact phrase"` - **Grouping**: `(term1 OR term2) AND term3` -- **Patterns**: `tag:example`, `category:observation`""" +- **Tags**: `tag:example` +- **Observation categories**: `entity_types=["observation"], categories=["requirement"]`""" def _format_search_markdown(result: SearchResponse, project: str, query: str | None) -> str: @@ -497,6 +498,7 @@ async def _search_all_projects( output_format: Literal["text", "json"], note_types: list[str], entity_types: list[str], + categories: list[str], after_date: str | None, metadata_filters: dict[str, Any] | None, tags: list[str] | None, @@ -555,6 +557,7 @@ async def _search_all_projects( output_format="json", note_types=note_types or None, entity_types=entity_types or None, + categories=categories or None, after_date=after_date, metadata_filters=metadata_filters, tags=tags, @@ -653,6 +656,14 @@ async def search_notes( "'relation'. Defaults to 'entity'. Do NOT pass schema/frontmatter types like " "'Chapter' here — use note_types instead.", ] = None, + categories: Annotated[ + List[str] | None, + BeforeValidator(coerce_list), + Field(default=None, validation_alias=AliasChoices("categories", "category")), + "Filter observation results to these exact categories (e.g. ['requirement']). " + "Pair with entity_types=['observation'] to return only observations whose " + "category matches exactly — not every row mentioning the word.", + ] = None, # Time-filter naming varies wildly across APIs. after_date: Annotated[ Optional[str], @@ -707,7 +718,8 @@ async def search_notes( ### Content-Specific Searches - `search_notes("research", "tag:example")` - Search within specific tags (if supported by content) - - `search_notes("work-project", "category:observation")` - Filter by observation categories + - `search_notes("work-project", "req", entity_types=["observation"], categories=["requirement"])` + - Return only observations whose category is exactly "requirement" - `search_notes("team-docs", "author:username")` - Find content by author (if metadata available) **Note:** `tag:` shorthand is automatically converted to a `tags` filter, so it works @@ -725,6 +737,8 @@ async def search_notes( - `search_notes("my-project", "query", note_types=["note"])` - Search only notes - `search_notes("work-docs", "query", note_types=["note", "person"])` - Multiple note types - `search_notes("research", "query", entity_types=["observation"])` - Filter by entity type + - `search_notes("research", "query", entity_types=["observation"], categories=["requirement"])` + - Filter observations to an exact category - `search_notes("team-docs", "query", after_date="2024-01-01")` - Recent content only - `search_notes("my-project", "query", after_date="1 week")` - Relative date filtering - `search_notes("my-project", "query", tags=["security"])` - Filter by frontmatter tags @@ -776,6 +790,9 @@ async def search_notes( "json" returns a machine-readable dictionary payload. note_types: Optional list of note types to search (e.g., ["note", "person"]) entity_types: Optional list of entity types to filter by (e.g., ["entity", "observation"]) + categories: Optional list of observation categories for exact matching (e.g., + ["requirement"]). Pair with entity_types=["observation"] to return only + observations whose category matches exactly. after_date: Optional date filter for recent content (e.g., "1 week", "2d", "2024-01-01") metadata_filters: Optional structured frontmatter filters (e.g., {"status": "in-progress"}) tags: Optional tag filter (frontmatter tags); shorthand for metadata_filters["tags"] @@ -853,6 +870,9 @@ async def search_notes( # Lowercase note_types so "Chapter" matches the stored "chapter". note_types = [t.lower() for t in note_types] if note_types else [] entity_types = entity_types or [] + # Categories are matched exactly against the indexed observation category, + # so preserve their original casing (unlike the lowercased note_types). + categories = categories or [] # Parse tag: shorthand at tool level so it works with all search modes. # Handles "tag:security", "tag:coffee tag:brewing", "tag:coffee AND tag:brewing". @@ -894,6 +914,7 @@ async def search_notes( output_format=output_format, note_types=note_types, entity_types=entity_types, + categories=categories, after_date=after_date, metadata_filters=metadata_filters, tags=tags, @@ -917,8 +938,15 @@ async def search_notes( has_query=bool(query and query.strip()), note_type_filter_count=len(note_types), entity_type_filter_count=len(entity_types), + category_filter_count=len(categories), has_filters=bool( - metadata_filters or tags or status or note_types or entity_types or after_date + metadata_filters + or tags + or status + or note_types + or entity_types + or categories + or after_date ), has_tags_filter=bool(tags), has_status_filter=bool(status), @@ -981,6 +1009,8 @@ async def search_notes( # Add optional filters if provided (empty lists are treated as no filter) if entity_types: search_query.entity_types = [SearchItemType(t) for t in entity_types] + if categories: + search_query.categories = categories if note_types: search_query.note_types = note_types if after_date: @@ -1006,7 +1036,8 @@ async def search_notes( return ( "# No Search Criteria\n\n" "Please provide at least one of: `query`, `metadata_filters`, " - "`tags`, `status`, `note_types`, `entity_types`, or `after_date`." + "`tags`, `status`, `note_types`, `entity_types`, `categories`, " + "or `after_date`." ) # Default to entity-level results to avoid returning individual @@ -1014,7 +1045,17 @@ async def search_notes( # Applied after no_criteria() so that the implicit default doesn't # mask a truly empty search request. if not search_query.entity_types: - search_query.entity_types = [SearchItemType("entity")] + # Trigger: a category filter was supplied without an explicit + # entity_types. + # Why: categories only exist on observations — defaulting to "entity" + # (whose rows have NULL category) would AND a category filter against + # entity rows and return nothing, defeating a category-only search. + # Outcome: scope the implicit default to observations so + # search_notes(categories=[...]) returns the matching bullets. + if search_query.categories: + search_query.entity_types = [SearchItemType("observation")] + else: + search_query.entity_types = [SearchItemType("entity")] logger.debug( f"Search request: project={active_project.name} " diff --git a/src/basic_memory/repository/postgres_search_repository.py b/src/basic_memory/repository/postgres_search_repository.py index fd83217f..ce259521 100644 --- a/src/basic_memory/repository/postgres_search_repository.py +++ b/src/basic_memory/repository/postgres_search_repository.py @@ -705,6 +705,7 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: 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.""" @@ -762,6 +763,20 @@ class PostgresSearchRepository(SearchRepositoryBase): type_placeholders.append(f":{param_name}") conditions.append(f"search_index.type IN ({', '.join(type_placeholders)})") + # Handle observation category filter (parameterized for defense-in-depth). + # Trigger: caller passed `categories` to scope observation results. + # Why: `entity_types=["observation"]` only narrows to the observation row type; + # callers expect exact-category matching, not incidental text matches. + # Outcome: only rows whose indexed category exactly equals a requested value + # survive (entities/relations have NULL category and are excluded). + if categories: + category_placeholders = [] + for idx, category in enumerate(categories): + param_name = f"category_{idx}" + params[param_name] = category + category_placeholders.append(f":{param_name}") + conditions.append(f"search_index.category IN ({', '.join(category_placeholders)})") + # Handle note type filter using JSONB containment (parameterized) if note_types: type_conditions = [] @@ -879,6 +894,7 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -895,6 +911,7 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, retrieval_mode=retrieval_mode, min_similarity=min_similarity, @@ -919,6 +936,7 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, ) @@ -1008,6 +1026,7 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -1022,6 +1041,7 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, retrieval_mode=retrieval_mode, min_similarity=min_similarity, @@ -1041,6 +1061,7 @@ class PostgresSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, 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 663c0122..3428f428 100644 --- a/src/basic_memory/repository/search_repository.py +++ b/src/basic_memory/repository/search_repository.py @@ -42,6 +42,7 @@ class SearchRepository(Protocol): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -60,6 +61,7 @@ class SearchRepository(Protocol): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: 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 9e13cb4a..39c03227 100644 --- a/src/basic_memory/repository/search_repository_base.py +++ b/src/basic_memory/repository/search_repository_base.py @@ -218,6 +218,7 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: Optional[List[str]] = None, metadata_filters: Optional[Dict[str, Any]] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -234,6 +235,7 @@ 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) + categories: Filter observations by exact category (e.g. "requirement") metadata_filters: Structured frontmatter metadata filters limit: Maximum results to return offset: Number of results to skip @@ -256,6 +258,7 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: Optional[List[str]] = None, metadata_filters: Optional[Dict[str, Any]] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -1785,6 +1788,7 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]], after_date: Optional[datetime], search_item_types: Optional[List[SearchItemType]], + categories: Optional[List[str]], metadata_filters: Optional[dict], retrieval_mode: SearchRetrievalMode, min_similarity: Optional[float] = None, @@ -1818,6 +1822,7 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, min_similarity=min_similarity, limit=limit, @@ -1837,6 +1842,7 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, min_similarity=min_similarity, limit=limit, @@ -1866,6 +1872,7 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]], after_date: Optional[datetime], search_item_types: Optional[List[SearchItemType]], + categories: Optional[List[str]], metadata_filters: Optional[dict], min_similarity: Optional[float] = None, limit: int, @@ -1976,6 +1983,7 @@ class SearchRepositoryBase(ABC): note_types, after_date, search_item_types, + categories, metadata_filters, ] ) @@ -1989,6 +1997,7 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, retrieval_mode=SearchRetrievalMode.FTS, limit=VECTOR_FILTER_SCAN_LIMIT, @@ -2139,6 +2148,7 @@ class SearchRepositoryBase(ABC): note_types: Optional[List[str]], after_date: Optional[datetime], search_item_types: Optional[List[SearchItemType]], + categories: Optional[List[str]], metadata_filters: Optional[dict], min_similarity: Optional[float] = None, limit: int, @@ -2163,6 +2173,7 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, retrieval_mode=SearchRetrievalMode.FTS, limit=candidate_limit, @@ -2178,6 +2189,7 @@ class SearchRepositoryBase(ABC): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, 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 9ffc73f0..745c7f23 100644 --- a/src/basic_memory/repository/sqlite_search_repository.py +++ b/src/basic_memory/repository/sqlite_search_repository.py @@ -733,6 +733,7 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: 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.""" @@ -794,6 +795,20 @@ class SQLiteSearchRepository(SearchRepositoryBase): type_placeholders.append(f":{param_name}") conditions.append(f"search_index.type IN ({', '.join(type_placeholders)})") + # Handle observation category filter (parameterized for defense-in-depth). + # Trigger: caller passed `categories` to scope observation results. + # Why: `entity_types=["observation"]` only narrows to the observation row type; + # callers expect exact-category matching, not incidental text matches. + # Outcome: only rows whose indexed category exactly equals a requested value + # survive (entities/relations have NULL category and are excluded). + if categories: + category_placeholders = [] + for idx, category in enumerate(categories): + param_name = f"category_{idx}" + params[param_name] = category + category_placeholders.append(f":{param_name}") + conditions.append(f"search_index.category IN ({', '.join(category_placeholders)})") + # Handle note type filter (frontmatter type field, parameterized) if note_types: type_placeholders = [] @@ -925,6 +940,7 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -941,6 +957,7 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, retrieval_mode=retrieval_mode, min_similarity=min_similarity, @@ -959,6 +976,7 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, ) @@ -1046,6 +1064,7 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types: Optional[List[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[List[SearchItemType]] = None, + categories: Optional[List[str]] = None, metadata_filters: Optional[dict] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -1060,6 +1079,7 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, metadata_filters=metadata_filters, retrieval_mode=retrieval_mode, min_similarity=min_similarity, @@ -1073,6 +1093,7 @@ class SQLiteSearchRepository(SearchRepositoryBase): note_types=note_types, after_date=after_date, search_item_types=search_item_types, + categories=categories, 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..cae201ea 100644 --- a/src/basic_memory/schemas/search.py +++ b/src/basic_memory/schemas/search.py @@ -42,6 +42,7 @@ class SearchQuery(BaseModel): Optionally filter results by: - note_types: Limit to specific note types (frontmatter "type") - entity_types: Limit to search item types (entity/observation/relation) + - categories: Limit observation results to exact category matches (e.g. "requirement") - after_date: Only items after date - metadata_filters: Structured frontmatter filters (field -> value) - tags: Convenience frontmatter tag filter @@ -63,6 +64,7 @@ class SearchQuery(BaseModel): # Optional filters note_types: Optional[List[str]] = None # Filter by note type (frontmatter "type") entity_types: Optional[List[SearchItemType]] = None # Filter by entity type + categories: Optional[List[str]] = None # Filter observations by exact category after_date: Optional[Union[datetime, str]] = None # Time-based filter metadata_filters: Optional[dict[str, Any]] = None # Structured frontmatter filters tags: Optional[List[str]] = None # Convenience tag filter @@ -85,6 +87,7 @@ 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 + categories_is_empty = not self.categories return ( self.permalink is None and self.permalink_match is None @@ -93,6 +96,7 @@ class SearchQuery(BaseModel): and self.after_date is None and note_types_is_empty and entity_types_is_empty + and categories_is_empty and metadata_is_empty and tags_is_empty and status_is_empty diff --git a/src/basic_memory/services/search_service.py b/src/basic_memory/services/search_service.py index 46264088..c1f28c2a 100644 --- a/src/basic_memory/services/search_service.py +++ b/src/basic_memory/services/search_service.py @@ -75,6 +75,7 @@ class _PreparedSearchQuery: title: str | None note_types: list[str] | None search_item_types: list[SearchItemType] | None + categories: list[str] | None after_date: datetime | None metadata_filters: dict[str, Any] | None retrieval_mode: SearchRetrievalMode @@ -194,6 +195,7 @@ class SearchService: title=query.title, note_types=query.note_types, search_item_types=query.entity_types, + categories=query.categories, after_date=after_date, metadata_filters=metadata_filters, retrieval_mode=query.retrieval_mode or SearchRetrievalMode.FTS, @@ -207,6 +209,7 @@ class SearchService: or prepared.title or prepared.note_types or prepared.search_item_types + or prepared.categories or prepared.after_date or prepared.metadata_filters ) @@ -221,6 +224,7 @@ class SearchService: prepared.metadata_filters or prepared.note_types or prepared.search_item_types + or prepared.categories or prepared.after_date ) @@ -239,6 +243,7 @@ class SearchService: title=prepared.title, note_types=prepared.note_types, search_item_types=prepared.search_item_types, + categories=prepared.categories, after_date=prepared.after_date, metadata_filters=prepared.metadata_filters, retrieval_mode=prepared.retrieval_mode, @@ -260,6 +265,7 @@ class SearchService: title=prepared.title, note_types=prepared.note_types, search_item_types=prepared.search_item_types, + categories=prepared.categories, after_date=prepared.after_date, metadata_filters=prepared.metadata_filters, retrieval_mode=prepared.retrieval_mode, diff --git a/tests/mcp/test_tool_contracts.py b/tests/mcp/test_tool_contracts.py index c1665d19..4e1be403 100644 --- a/tests/mcp/test_tool_contracts.py +++ b/tests/mcp/test_tool_contracts.py @@ -91,6 +91,7 @@ EXPECTED_TOOL_SIGNATURES: dict[str, list[str]] = { "output_format", "note_types", "entity_types", + "categories", "after_date", "metadata_filters", "tags", diff --git a/tests/mcp/test_tool_search.py b/tests/mcp/test_tool_search.py index 2c2bc535..1660af40 100644 --- a/tests/mcp/test_tool_search.py +++ b/tests/mcp/test_tool_search.py @@ -370,6 +370,93 @@ async def test_search_with_entity_type_filter(client, test_project): pytest.fail(f"Search failed with error: {response}") +@pytest.mark.asyncio +async def test_search_with_categories_filter(client, test_project): + """Observation category filter returns only the exact category (#430). + + Writes a note whose body has a [requirement] observation and a [decision] + observation that also mentions the word "requirement". The categories filter + must return only the requirement observation. + """ + await write_note( + project=test_project.name, + title="Category Filter Note", + directory="test", + content=( + "# Category Filter Note\n" + "- [requirement] The system must enforce auth on every request\n" + "- [decision] We deferred the auth requirement to next sprint\n" + ), + ) + + response = await search_notes( + project=test_project.name, + query="requirement", + search_type="text", + entity_types=["observation"], + categories=["requirement"], + output_format="json", + ) + + assert isinstance(response, dict), f"Search failed with error: {response}" + results = response["results"] + assert len(results) > 0 + # Every result is a requirement observation; the [decision] row is excluded + # even though its text contains the word "requirement". + assert all(r["type"] == "observation" for r in results) + assert all(r["category"] == "requirement" for r in results) + + # A non-matching category yields no results for the same text query. + decision = await search_notes( + project=test_project.name, + query="requirement", + search_type="text", + entity_types=["observation"], + categories=["decision"], + output_format="json", + ) + assert isinstance(decision, dict), f"Search failed with error: {decision}" + assert all(r["category"] == "decision" for r in decision["results"]) + # The requirement observation must not leak into a decision-scoped search. + assert all("requirement" != r.get("category") for r in decision["results"]) + + +@pytest.mark.asyncio +async def test_search_categories_without_entity_types_returns_observations(client, test_project): + """categories=[...] WITHOUT entity_types must return the matching observations (#908). + + search_notes defaults entity_types to "entity" when unset, but categories only exist on + observations — so a category filter without an explicit entity_types would AND the + category against entity rows (which have NULL category) and return nothing. The implicit + default must scope to observations when categories is supplied. + """ + await write_note( + project=test_project.name, + title="Category Default Note", + directory="test", + content=( + "# Category Default Note\n" + "- [requirement] Auth tokens must rotate every 24 hours\n" + "- [decision] We chose JWT for the auth token format\n" + ), + ) + + # Note: no entity_types passed — exercises the implicit default. + response = await search_notes( + project=test_project.name, + query="auth", + search_type="text", + categories=["requirement"], + output_format="json", + ) + + assert isinstance(response, dict), f"Search failed with error: {response}" + results = response["results"] + assert len(results) > 0, "category-only search must return matching observations" + assert all(r["type"] == "observation" for r in results) + assert all(r["category"] == "requirement" for r in results) + + @pytest.mark.asyncio async def test_search_with_date_filter(client, test_project): """Test search with date filter.""" diff --git a/tests/mcp/test_tool_telemetry.py b/tests/mcp/test_tool_telemetry.py index fa85ee12..a9a304f8 100644 --- a/tests/mcp/test_tool_telemetry.py +++ b/tests/mcp/test_tool_telemetry.py @@ -165,6 +165,7 @@ async def test_search_notes_emits_root_operation_and_project_context( "has_query": True, "note_type_filter_count": 0, "entity_type_filter_count": 0, + "category_filter_count": 0, "has_filters": True, "has_tags_filter": True, "has_status_filter": False, diff --git a/tests/repository/test_hybrid_fusion.py b/tests/repository/test_hybrid_fusion.py index c0226f0d..7ae8a9e0 100644 --- a/tests/repository/test_hybrid_fusion.py +++ b/tests/repository/test_hybrid_fusion.py @@ -71,6 +71,7 @@ class ConcreteSearchRepo(SearchRepositoryBase): note_types: Optional[list[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[list[SearchItemType]] = None, + categories: Optional[list[str]] = None, metadata_filters: Optional[dict[str, Any]] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -126,6 +127,7 @@ HYBRID_KWARGS: dict[str, Any] = dict( note_types=None, after_date=None, search_item_types=None, + categories=None, metadata_filters=None, limit=10, offset=0, diff --git a/tests/repository/test_postgres_search_repository.py b/tests/repository/test_postgres_search_repository.py index fb4aa134..7c022014 100644 --- a/tests/repository/test_postgres_search_repository.py +++ b/tests/repository/test_postgres_search_repository.py @@ -932,3 +932,72 @@ async def test_postgres_metadata_filters_path_parameterized(session_maker, test_ # Nested path should work without SQL injection risk results = await repo.search(metadata_filters={"schema.confidence": {"$gt": 0.5}}) assert isinstance(results, list) + + +@pytest.mark.asyncio +async def test_postgres_search_categories_exact_match(session_maker, test_project): + """categories filter matches the observation category exactly (mirror of #430). + + A [decision] observation that merely mentions "requirement" must be excluded + when categories=["requirement"] is requested. + """ + repo = PostgresSearchRepository(session_maker, project_id=test_project.id) + now = datetime.now(timezone.utc) + + await repo.bulk_index_items( + [ + SearchIndexRow( + project_id=test_project.id, + id=70101, + type=SearchItemType.OBSERVATION.value, + content_stems="the auth requirement must be enforced on every call", + content_snippet="the auth requirement must be enforced on every call", + permalink="test/obs/requirement/70101", + file_path="test/obs.md", + entity_id=1, + category="requirement", + metadata={"note_type": "note"}, + created_at=now, + updated_at=now, + ), + SearchIndexRow( + project_id=test_project.id, + id=70102, + type=SearchItemType.OBSERVATION.value, + content_stems="we deferred the auth requirement to next sprint", + content_snippet="we deferred the auth requirement to next sprint", + permalink="test/obs/decision/70102", + file_path="test/obs.md", + entity_id=1, + category="decision", + metadata={"note_type": "note"}, + created_at=now, + updated_at=now, + ), + ] + ) + + # Without the category filter, a text search for "requirement" matches both. + text_results = await repo.search( + search_text="requirement", + search_item_types=[SearchItemType.OBSERVATION], + ) + assert {r.id for r in text_results} == {70101, 70102} + + # With categories=["requirement"], only the requirement observation survives. + filtered = await repo.search( + search_text="requirement", + search_item_types=[SearchItemType.OBSERVATION], + categories=["requirement"], + ) + assert {r.id for r in filtered} == {70101} + assert filtered[0].category == "requirement" + + # Standalone filter and count both honor the exact category. + filtered_only = await repo.search(categories=["requirement"]) + assert {r.id for r in filtered_only} == {70101} + assert await repo.count(categories=["requirement"]) == 1 + + # Multiple categories union. + multi = await repo.search(categories=["requirement", "decision"]) + assert {r.id for r in multi} == {70101, 70102} diff --git a/tests/repository/test_search_repository.py b/tests/repository/test_search_repository.py index 27e8b8da..0fb6d630 100644 --- a/tests/repository/test_search_repository.py +++ b/tests/repository/test_search_repository.py @@ -1044,3 +1044,83 @@ 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) + + +async def _index_observation( + search_repository, + *, + row_id: int, + entity_id: int, + category: str, + content: str, +) -> None: + """Index a single observation row with an explicit category for filter tests.""" + now = datetime.now(timezone.utc) + search_row = SearchIndexRow( + id=row_id, + type=SearchItemType.OBSERVATION.value, + content_stems=content, + content_snippet=content, + permalink=f"test/obs/{category}/{row_id}", + file_path="test/obs.md", + entity_id=entity_id, + category=category, + metadata={"note_type": "note"}, + created_at=now, + updated_at=now, + project_id=search_repository.project_id, + ) + await search_repository.index_item(search_row) + + +@pytest.mark.asyncio +async def test_search_categories_exact_match(search_repository, search_entity): + """categories must match the observation category exactly, not by text. + + Regression for #430: searching observations for "requirement" used to also + return a [decision] observation that merely mentions the word. The categories + filter scopes results to the exact indexed category. + """ + await _index_observation( + search_repository, + row_id=70001, + entity_id=search_entity.id, + category="requirement", + content="The auth requirement must be enforced on every call", + ) + # A decision observation whose text mentions "requirement" but whose category + # is NOT requirement — it must be excluded by an exact-category filter. + await _index_observation( + search_repository, + row_id=70002, + entity_id=search_entity.id, + category="decision", + content="We deferred the auth requirement to next sprint", + ) + + # Without the category filter, a text search for "requirement" matches both. + text_results = await search_repository.search( + search_text="requirement", + search_item_types=[SearchItemType.OBSERVATION], + ) + assert {r.id for r in text_results} == {70001, 70002} + + # With categories=["requirement"], only the requirement observation survives. + filtered = await search_repository.search( + search_text="requirement", + search_item_types=[SearchItemType.OBSERVATION], + categories=["requirement"], + ) + assert {r.id for r in filtered} == {70001} + assert filtered[0].category == "requirement" + + # categories also works as a standalone filter (no text query). + filtered_only = await search_repository.search(categories=["requirement"]) + assert {r.id for r in filtered_only} == {70001} + + # count mirrors the filtered search. + assert await search_repository.count(categories=["requirement"]) == 1 + + # Multiple categories union: both observations come back. + multi = await search_repository.search(categories=["requirement", "decision"]) + assert {r.id for r in multi} == {70001, 70002} diff --git a/tests/repository/test_semantic_search_base.py b/tests/repository/test_semantic_search_base.py index ba4e08ef..b75f8871 100644 --- a/tests/repository/test_semantic_search_base.py +++ b/tests/repository/test_semantic_search_base.py @@ -83,6 +83,7 @@ class _ConcreteRepo(SearchRepositoryBase): note_types: list[str] | None = None, after_date: datetime | None = None, search_item_types: list[SearchItemType] | None = None, + categories: 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..21bc98a8 100644 --- a/tests/repository/test_vector_pagination.py +++ b/tests/repository/test_vector_pagination.py @@ -56,6 +56,7 @@ class ConcreteSearchRepo(SearchRepositoryBase): note_types: list[str] | None = None, after_date: datetime | None = None, search_item_types: list[SearchItemType] | None = None, + categories: list[str] | None = None, metadata_filters: dict[str, Any] | None = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: float | None = None, @@ -158,6 +159,7 @@ async def test_page1_scores_gte_page2_scores(): note_types=None, after_date=None, search_item_types=None, + categories=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..266ee9ae 100644 --- a/tests/repository/test_vector_threshold.py +++ b/tests/repository/test_vector_threshold.py @@ -60,6 +60,7 @@ class ConcreteSearchRepo(SearchRepositoryBase): note_types: Optional[list[str]] = None, after_date: Optional[datetime] = None, search_item_types: Optional[list[SearchItemType]] = None, + categories: Optional[list[str]] = None, metadata_filters: Optional[dict[str, Any]] = None, retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS, min_similarity: Optional[float] = None, @@ -130,6 +131,7 @@ COMMON_SEARCH_KWARGS: dict[str, Any] = dict( note_types=None, after_date=None, search_item_types=None, + categories=None, metadata_filters=None, limit=10, offset=0, diff --git a/tests/services/test_search_service.py b/tests/services/test_search_service.py index 7ee913b8..32fec22a 100644 --- a/tests/services/test_search_service.py +++ b/tests/services/test_search_service.py @@ -277,6 +277,39 @@ async def test_search_entity_type(search_service, test_graph): assert r.type == SearchItemType.ENTITY +@pytest.mark.asyncio +async def test_search_categories_filter(search_service, test_graph): + """categories propagates through _prepare_query/has_criteria to scope results. + + The test_graph fixture indexes observations with categories "note" and "tech". + A categories filter must return only matching observation categories. + """ + # categories alone is enough criteria to run a query (has_criteria True). + note_results = await search_service.search(SearchQuery(categories=["note"])) + assert len(note_results) > 0 + assert all(r.type == SearchItemType.OBSERVATION for r in note_results) + assert all(r.category == "note" for r in note_results) + + # A different category yields a disjoint, non-empty result set. + tech_results = await search_service.search(SearchQuery(categories=["tech"])) + assert len(tech_results) > 0 + assert all(r.category == "tech" for r in tech_results) + + note_ids = {r.id for r in note_results} + tech_ids = {r.id for r in tech_results} + assert note_ids.isdisjoint(tech_ids) + + # count() must agree with the filtered search via the same prepared query. + assert await search_service.count(SearchQuery(categories=["note"])) == len(note_results) + + +@pytest.mark.asyncio +async def test_search_categories_only_is_not_no_criteria(): + """A SearchQuery carrying only categories must not be treated as empty.""" + assert SearchQuery(categories=["requirement"]).no_criteria() is False + assert SearchQuery().no_criteria() is True + + @pytest.mark.asyncio async def test_extract_entity_tags_exception_handling(search_service): """Test the _extract_entity_tags method exception handling (lines 147-151)."""