mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
Add exact totals to search pagination
Signed-off-by: phernandez <paul@basicmachines.co>
This commit is contained in:
@@ -65,7 +65,6 @@ async def search(
|
||||
has_filters=bool(query.note_types or query.entity_types or query.metadata_filters),
|
||||
):
|
||||
offset = (page - 1) * page_size
|
||||
fetch_limit = page_size + 1
|
||||
try:
|
||||
with logfire.span(
|
||||
"api.search.search.execute_query",
|
||||
@@ -75,7 +74,8 @@ async def search(
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
):
|
||||
results = await search_service.search(query, limit=fetch_limit, offset=offset)
|
||||
results = await search_service.search(query, limit=page_size, offset=offset)
|
||||
total = await search_service.count(query)
|
||||
except SemanticSearchDisabledError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except SemanticDependenciesMissingError as exc:
|
||||
@@ -90,9 +90,7 @@ async def search(
|
||||
phase="paginate_results",
|
||||
result_count=len(results),
|
||||
):
|
||||
has_more = len(results) > page_size
|
||||
if has_more:
|
||||
results = results[:page_size]
|
||||
has_more = offset + len(results) < total
|
||||
|
||||
with logfire.span(
|
||||
"api.search.search.hydrate_results",
|
||||
@@ -113,6 +111,7 @@ async def search(
|
||||
results=search_results,
|
||||
current_page=page,
|
||||
page_size=page_size,
|
||||
total=total,
|
||||
has_more=has_more,
|
||||
)
|
||||
|
||||
|
||||
@@ -686,7 +686,17 @@ class PostgresSearchRepository(SearchRepositoryBase):
|
||||
# FTS search (Postgres-specific)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def search(
|
||||
@staticmethod
|
||||
def _is_tsquery_syntax_error(exc: Exception) -> bool:
|
||||
msg = str(exc).lower()
|
||||
return (
|
||||
"syntax error in tsquery" in msg
|
||||
or "invalid input syntax for type tsquery" in msg
|
||||
or "no operand in tsquery" in msg
|
||||
or "no operator in tsquery" in msg
|
||||
)
|
||||
|
||||
async def _build_fts_query_parts(
|
||||
self,
|
||||
search_text: Optional[str] = None,
|
||||
permalink: Optional[str] = None,
|
||||
@@ -696,31 +706,8 @@ class PostgresSearchRepository(SearchRepositoryBase):
|
||||
after_date: Optional[datetime] = None,
|
||||
search_item_types: Optional[List[SearchItemType]] = None,
|
||||
metadata_filters: Optional[dict] = None,
|
||||
retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS,
|
||||
min_similarity: Optional[float] = None,
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
) -> List[SearchIndexRow]:
|
||||
"""Search across all indexed content using PostgreSQL tsvector."""
|
||||
# --- Dispatch vector / hybrid modes (shared logic) ---
|
||||
dispatched = await self._dispatch_retrieval_mode(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=min_similarity,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
if dispatched is not None:
|
||||
return dispatched
|
||||
|
||||
# --- FTS mode (Postgres-specific) ---
|
||||
) -> tuple[str, str, dict, str, str]:
|
||||
"""Build Postgres FTS FROM/WHERE params shared by search and count."""
|
||||
conditions = []
|
||||
params = {}
|
||||
order_by_clause = ""
|
||||
@@ -868,10 +855,6 @@ class PostgresSearchRepository(SearchRepositoryBase):
|
||||
params["project_id"] = self.project_id
|
||||
conditions.append("search_index.project_id = :project_id")
|
||||
|
||||
# set limit and offset
|
||||
params["limit"] = limit
|
||||
params["offset"] = offset
|
||||
|
||||
# Build WHERE clause
|
||||
where_clause = " AND ".join(conditions) if conditions else "1=1"
|
||||
|
||||
@@ -884,6 +867,64 @@ class PostgresSearchRepository(SearchRepositoryBase):
|
||||
else:
|
||||
score_expr = "0"
|
||||
|
||||
return from_clause, where_clause, params, order_by_clause, score_expr
|
||||
|
||||
async def search(
|
||||
self,
|
||||
search_text: Optional[str] = None,
|
||||
permalink: Optional[str] = None,
|
||||
permalink_match: Optional[str] = None,
|
||||
title: Optional[str] = None,
|
||||
note_types: Optional[List[str]] = None,
|
||||
after_date: Optional[datetime] = None,
|
||||
search_item_types: Optional[List[SearchItemType]] = None,
|
||||
metadata_filters: Optional[dict] = None,
|
||||
retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS,
|
||||
min_similarity: Optional[float] = None,
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
) -> List[SearchIndexRow]:
|
||||
"""Search across all indexed content using PostgreSQL tsvector."""
|
||||
# --- Dispatch vector / hybrid modes (shared logic) ---
|
||||
dispatched = await self._dispatch_retrieval_mode(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=min_similarity,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
if dispatched is not None:
|
||||
return dispatched
|
||||
|
||||
# --- FTS mode (Postgres-specific) ---
|
||||
(
|
||||
from_clause,
|
||||
where_clause,
|
||||
params,
|
||||
order_by_clause,
|
||||
score_expr,
|
||||
) = await self._build_fts_query_parts(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
)
|
||||
|
||||
# set limit and offset
|
||||
params["limit"] = limit
|
||||
params["offset"] = offset
|
||||
|
||||
sql = f"""
|
||||
SELECT
|
||||
search_index.project_id,
|
||||
@@ -915,17 +956,7 @@ class PostgresSearchRepository(SearchRepositoryBase):
|
||||
result = await session.execute(text(sql), params)
|
||||
rows = result.fetchall()
|
||||
except Exception as e:
|
||||
# Handle tsquery syntax errors (and only those).
|
||||
#
|
||||
# Important: Postgres errors for other failures (e.g. missing table) will still mention
|
||||
# `to_tsquery(...)` in the SQL text, so checking for the substring "tsquery" is too broad.
|
||||
msg = str(e).lower()
|
||||
if (
|
||||
"syntax error in tsquery" in msg
|
||||
or "invalid input syntax for type tsquery" in msg
|
||||
or "no operand in tsquery" in msg
|
||||
or "no operator in tsquery" in msg
|
||||
):
|
||||
if self._is_tsquery_syntax_error(e):
|
||||
logger.warning(f"tsquery syntax error for search term: {search_text}, error: {e}")
|
||||
return []
|
||||
|
||||
@@ -966,3 +997,65 @@ class PostgresSearchRepository(SearchRepositoryBase):
|
||||
)
|
||||
|
||||
return results
|
||||
|
||||
async def count(
|
||||
self,
|
||||
search_text: Optional[str] = None,
|
||||
permalink: Optional[str] = None,
|
||||
permalink_match: Optional[str] = None,
|
||||
title: Optional[str] = None,
|
||||
note_types: Optional[List[str]] = None,
|
||||
after_date: Optional[datetime] = None,
|
||||
search_item_types: Optional[List[SearchItemType]] = None,
|
||||
metadata_filters: Optional[dict] = None,
|
||||
retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS,
|
||||
min_similarity: Optional[float] = None,
|
||||
) -> int:
|
||||
"""Count indexed content matching the Postgres FTS query."""
|
||||
mode = (
|
||||
retrieval_mode.value
|
||||
if isinstance(retrieval_mode, SearchRetrievalMode)
|
||||
else str(retrieval_mode)
|
||||
)
|
||||
if mode != SearchRetrievalMode.FTS.value:
|
||||
return await super().count(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=min_similarity,
|
||||
)
|
||||
|
||||
(
|
||||
from_clause,
|
||||
where_clause,
|
||||
params,
|
||||
_order_by_clause,
|
||||
_score_expr,
|
||||
) = await self._build_fts_query_parts(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
)
|
||||
sql = f"SELECT COUNT(*) FROM {from_clause} WHERE {where_clause}"
|
||||
logger.trace(f"Count {sql} params: {params}")
|
||||
try:
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
result = await session.execute(text(sql), params)
|
||||
return int(result.scalar_one())
|
||||
except Exception as e:
|
||||
if self._is_tsquery_syntax_error(e):
|
||||
logger.warning(f"tsquery syntax error for search term: {search_text}, error: {e}")
|
||||
return 0
|
||||
logger.error(f"Database error during search count: {e}")
|
||||
raise
|
||||
|
||||
@@ -50,6 +50,22 @@ class SearchRepository(Protocol):
|
||||
"""Search across indexed content."""
|
||||
...
|
||||
|
||||
async def count(
|
||||
self,
|
||||
search_text: Optional[str] = None,
|
||||
permalink: Optional[str] = None,
|
||||
permalink_match: Optional[str] = None,
|
||||
title: Optional[str] = None,
|
||||
note_types: Optional[List[str]] = None,
|
||||
after_date: Optional[datetime] = None,
|
||||
search_item_types: Optional[List[SearchItemType]] = None,
|
||||
metadata_filters: Optional[dict] = None,
|
||||
retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS,
|
||||
min_similarity: Optional[float] = None,
|
||||
) -> int:
|
||||
"""Count indexed content matching the same filters as search."""
|
||||
...
|
||||
|
||||
async def index_item(self, search_index_row: SearchIndexRow) -> None:
|
||||
"""Index a single item."""
|
||||
...
|
||||
|
||||
@@ -247,6 +247,36 @@ class SearchRepositoryBase(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
async def count(
|
||||
self,
|
||||
search_text: Optional[str] = None,
|
||||
permalink: Optional[str] = None,
|
||||
permalink_match: Optional[str] = None,
|
||||
title: Optional[str] = None,
|
||||
note_types: Optional[List[str]] = None,
|
||||
after_date: Optional[datetime] = None,
|
||||
search_item_types: Optional[List[SearchItemType]] = None,
|
||||
metadata_filters: Optional[Dict[str, Any]] = None,
|
||||
retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS,
|
||||
min_similarity: Optional[float] = None,
|
||||
) -> int:
|
||||
"""Count results for retrieval modes that cannot use a backend COUNT query."""
|
||||
results = await self.search(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=min_similarity,
|
||||
limit=VECTOR_FILTER_SCAN_LIMIT,
|
||||
offset=0,
|
||||
)
|
||||
return len(results)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Abstract methods — semantic search (backend-specific DB operations)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -701,7 +701,11 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
# FTS search (backend-specific)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def search(
|
||||
@staticmethod
|
||||
def _is_fts5_syntax_error(exc: Exception) -> bool:
|
||||
return "fts5: syntax error" in str(exc).lower()
|
||||
|
||||
async def _build_fts_query_parts(
|
||||
self,
|
||||
search_text: Optional[str] = None,
|
||||
permalink: Optional[str] = None,
|
||||
@@ -711,31 +715,8 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
after_date: Optional[datetime] = None,
|
||||
search_item_types: Optional[List[SearchItemType]] = None,
|
||||
metadata_filters: Optional[dict] = None,
|
||||
retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS,
|
||||
min_similarity: Optional[float] = None,
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
) -> List[SearchIndexRow]:
|
||||
"""Search across all indexed content using SQLite FTS5."""
|
||||
# --- Dispatch vector / hybrid modes (shared logic) ---
|
||||
dispatched = await self._dispatch_retrieval_mode(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=min_similarity,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
if dispatched is not None:
|
||||
return dispatched
|
||||
|
||||
# --- FTS mode (SQLite-specific) ---
|
||||
) -> tuple[str, str, dict, str]:
|
||||
"""Build SQLite FTS FROM/WHERE params shared by search and count."""
|
||||
conditions = []
|
||||
match_conditions = []
|
||||
params = {}
|
||||
@@ -911,13 +892,60 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
params["project_id"] = self.project_id
|
||||
conditions.append("search_index.project_id = :project_id")
|
||||
|
||||
# Build WHERE clause
|
||||
where_clause = " AND ".join(conditions) if conditions else "1=1"
|
||||
return from_clause, where_clause, params, order_by_clause
|
||||
|
||||
async def search(
|
||||
self,
|
||||
search_text: Optional[str] = None,
|
||||
permalink: Optional[str] = None,
|
||||
permalink_match: Optional[str] = None,
|
||||
title: Optional[str] = None,
|
||||
note_types: Optional[List[str]] = None,
|
||||
after_date: Optional[datetime] = None,
|
||||
search_item_types: Optional[List[SearchItemType]] = None,
|
||||
metadata_filters: Optional[dict] = None,
|
||||
retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS,
|
||||
min_similarity: Optional[float] = None,
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
) -> List[SearchIndexRow]:
|
||||
"""Search across all indexed content using SQLite FTS5."""
|
||||
# --- Dispatch vector / hybrid modes (shared logic) ---
|
||||
dispatched = await self._dispatch_retrieval_mode(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=min_similarity,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
if dispatched is not None:
|
||||
return dispatched
|
||||
|
||||
# --- FTS mode (SQLite-specific) ---
|
||||
from_clause, where_clause, params, order_by_clause = await self._build_fts_query_parts(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
)
|
||||
|
||||
# set limit on search query
|
||||
params["limit"] = limit
|
||||
params["offset"] = offset
|
||||
|
||||
# Build WHERE clause
|
||||
where_clause = " AND ".join(conditions) if conditions else "1=1"
|
||||
|
||||
sql = f"""
|
||||
SELECT
|
||||
search_index.project_id,
|
||||
@@ -950,7 +978,7 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
rows = result.fetchall()
|
||||
except Exception as e:
|
||||
# Handle FTS5 syntax errors and provide user-friendly feedback
|
||||
if "fts5: syntax error" in str(e).lower(): # pragma: no cover
|
||||
if self._is_fts5_syntax_error(e): # pragma: no cover
|
||||
logger.warning(f"FTS5 syntax error for search term: {search_text}, error: {e}")
|
||||
# Return empty results rather than crashing
|
||||
return []
|
||||
@@ -988,3 +1016,59 @@ class SQLiteSearchRepository(SearchRepositoryBase):
|
||||
)
|
||||
|
||||
return results
|
||||
|
||||
async def count(
|
||||
self,
|
||||
search_text: Optional[str] = None,
|
||||
permalink: Optional[str] = None,
|
||||
permalink_match: Optional[str] = None,
|
||||
title: Optional[str] = None,
|
||||
note_types: Optional[List[str]] = None,
|
||||
after_date: Optional[datetime] = None,
|
||||
search_item_types: Optional[List[SearchItemType]] = None,
|
||||
metadata_filters: Optional[dict] = None,
|
||||
retrieval_mode: SearchRetrievalMode = SearchRetrievalMode.FTS,
|
||||
min_similarity: Optional[float] = None,
|
||||
) -> int:
|
||||
"""Count indexed content matching the SQLite FTS query."""
|
||||
mode = (
|
||||
retrieval_mode.value
|
||||
if isinstance(retrieval_mode, SearchRetrievalMode)
|
||||
else str(retrieval_mode)
|
||||
)
|
||||
if mode != SearchRetrievalMode.FTS.value:
|
||||
return await super().count(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=min_similarity,
|
||||
)
|
||||
|
||||
from_clause, where_clause, params, _order_by_clause = await self._build_fts_query_parts(
|
||||
search_text=search_text,
|
||||
permalink=permalink,
|
||||
permalink_match=permalink_match,
|
||||
title=title,
|
||||
note_types=note_types,
|
||||
after_date=after_date,
|
||||
search_item_types=search_item_types,
|
||||
metadata_filters=metadata_filters,
|
||||
)
|
||||
sql = f"SELECT COUNT(*) FROM {from_clause} WHERE {where_clause}"
|
||||
logger.trace(f"Count {sql} params: {params}")
|
||||
try:
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
result = await session.execute(text(sql), params)
|
||||
return int(result.scalar_one())
|
||||
except Exception as e:
|
||||
if self._is_fts5_syntax_error(e): # pragma: no cover
|
||||
logger.warning(f"FTS5 syntax error for search term: {search_text}, error: {e}")
|
||||
return 0
|
||||
logger.error(f"Database error during search count: {e}")
|
||||
raise
|
||||
|
||||
@@ -142,4 +142,5 @@ class SearchResponse(BaseModel):
|
||||
results: List[SearchResult]
|
||||
current_page: int
|
||||
page_size: int
|
||||
total: int = 0
|
||||
has_more: bool = False
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
import asyncio
|
||||
import ast
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import List, Optional, Set, Dict, Any
|
||||
|
||||
@@ -64,6 +65,22 @@ FTS_RELAXED_STOPWORDS = {
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _PreparedSearchQuery:
|
||||
"""Normalized query inputs shared by search and count."""
|
||||
|
||||
search_text: str | None
|
||||
permalink: str | None
|
||||
permalink_match: str | None
|
||||
title: str | None
|
||||
note_types: list[str] | None
|
||||
search_item_types: list[SearchItemType] | None
|
||||
after_date: datetime | None
|
||||
metadata_filters: dict[str, Any] | None
|
||||
retrieval_mode: SearchRetrievalMode
|
||||
min_similarity: float | None
|
||||
|
||||
|
||||
def _strip_nul(value: str) -> str:
|
||||
"""Strip NUL bytes that PostgreSQL text columns cannot store.
|
||||
|
||||
@@ -132,27 +149,20 @@ class SearchService:
|
||||
|
||||
logger.info("Reindex complete")
|
||||
|
||||
async def search(self, query: SearchQuery, limit=10, offset=0) -> List[SearchIndexRow]:
|
||||
"""Search across all indexed content.
|
||||
def _prepare_query(self, query: SearchQuery) -> _PreparedSearchQuery | None:
|
||||
"""Normalize a SearchQuery into repository arguments."""
|
||||
search_text = query.text
|
||||
tags = query.tags
|
||||
|
||||
Supports three modes:
|
||||
1. Exact permalink: finds direct matches for a specific path
|
||||
2. Pattern match: handles * wildcards in paths
|
||||
3. Text search: full-text search across title/content
|
||||
"""
|
||||
# Support tag:<tag> shorthand by mapping to tags filter
|
||||
if query.text:
|
||||
text = query.text.strip()
|
||||
# Support tag:<tag> shorthand by mapping to tags filter.
|
||||
if search_text:
|
||||
text = search_text.strip()
|
||||
if text.lower().startswith("tag:"):
|
||||
tag_values = re.split(r"[,\s]+", text[4:].strip())
|
||||
tags = [t for t in tag_values if t]
|
||||
if tags:
|
||||
query.tags = tags
|
||||
query.text = None
|
||||
|
||||
if query.no_criteria():
|
||||
logger.debug("no criteria passed to query")
|
||||
return []
|
||||
parsed_tags = [t for t in tag_values if t]
|
||||
if parsed_tags:
|
||||
tags = parsed_tags
|
||||
search_text = None
|
||||
|
||||
after_date = (
|
||||
(
|
||||
@@ -164,50 +174,124 @@ class SearchService:
|
||||
else None
|
||||
)
|
||||
|
||||
# Merge structured metadata filters (explicit + convenience fields)
|
||||
# Merge structured metadata filters (explicit + convenience fields).
|
||||
metadata_filters: Optional[Dict[str, Any]] = None
|
||||
if query.metadata_filters or query.tags or query.status:
|
||||
if query.metadata_filters or tags or query.status:
|
||||
metadata_filters = dict(query.metadata_filters or {})
|
||||
if query.tags:
|
||||
metadata_filters.setdefault("tags", query.tags)
|
||||
if tags:
|
||||
metadata_filters.setdefault("tags", tags)
|
||||
if query.status:
|
||||
metadata_filters.setdefault("status", query.status)
|
||||
|
||||
retrieval_mode = query.retrieval_mode or SearchRetrievalMode.FTS
|
||||
strict_search_text = query.text
|
||||
prepared = _PreparedSearchQuery(
|
||||
search_text=search_text,
|
||||
permalink=query.permalink,
|
||||
permalink_match=query.permalink_match,
|
||||
title=query.title,
|
||||
note_types=query.note_types,
|
||||
search_item_types=query.entity_types,
|
||||
after_date=after_date,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=query.retrieval_mode or SearchRetrievalMode.FTS,
|
||||
min_similarity=query.min_similarity,
|
||||
)
|
||||
|
||||
has_criteria = bool(
|
||||
prepared.search_text
|
||||
or prepared.permalink
|
||||
or prepared.permalink_match
|
||||
or prepared.title
|
||||
or prepared.note_types
|
||||
or prepared.search_item_types
|
||||
or prepared.after_date
|
||||
or prepared.metadata_filters
|
||||
)
|
||||
if not has_criteria:
|
||||
logger.debug("no criteria passed to query")
|
||||
return None
|
||||
return prepared
|
||||
|
||||
@staticmethod
|
||||
def _prepared_has_filters(prepared: _PreparedSearchQuery) -> bool:
|
||||
return bool(
|
||||
prepared.metadata_filters
|
||||
or prepared.note_types
|
||||
or prepared.search_item_types
|
||||
or prepared.after_date
|
||||
)
|
||||
|
||||
async def _search_repository(
|
||||
self,
|
||||
prepared: _PreparedSearchQuery,
|
||||
*,
|
||||
search_text: str | None,
|
||||
limit: int,
|
||||
offset: int,
|
||||
) -> List[SearchIndexRow]:
|
||||
return await self.repository.search(
|
||||
search_text=search_text,
|
||||
permalink=prepared.permalink,
|
||||
permalink_match=prepared.permalink_match,
|
||||
title=prepared.title,
|
||||
note_types=prepared.note_types,
|
||||
search_item_types=prepared.search_item_types,
|
||||
after_date=prepared.after_date,
|
||||
metadata_filters=prepared.metadata_filters,
|
||||
retrieval_mode=prepared.retrieval_mode,
|
||||
min_similarity=prepared.min_similarity,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
async def _count_repository(
|
||||
self,
|
||||
prepared: _PreparedSearchQuery,
|
||||
*,
|
||||
search_text: str | None,
|
||||
) -> int:
|
||||
return await self.repository.count(
|
||||
search_text=search_text,
|
||||
permalink=prepared.permalink,
|
||||
permalink_match=prepared.permalink_match,
|
||||
title=prepared.title,
|
||||
note_types=prepared.note_types,
|
||||
search_item_types=prepared.search_item_types,
|
||||
after_date=prepared.after_date,
|
||||
metadata_filters=prepared.metadata_filters,
|
||||
retrieval_mode=prepared.retrieval_mode,
|
||||
min_similarity=prepared.min_similarity,
|
||||
)
|
||||
|
||||
async def search(self, query: SearchQuery, limit=10, offset=0) -> List[SearchIndexRow]:
|
||||
"""Search across all indexed content.
|
||||
|
||||
Supports three modes:
|
||||
1. Exact permalink: finds direct matches for a specific path
|
||||
2. Pattern match: handles * wildcards in paths
|
||||
3. Text search: full-text search across title/content
|
||||
"""
|
||||
prepared = self._prepare_query(query)
|
||||
if prepared is None:
|
||||
return []
|
||||
|
||||
strict_search_text = prepared.search_text
|
||||
has_query = bool(
|
||||
strict_search_text or query.title or query.permalink or query.permalink_match
|
||||
)
|
||||
has_filters = bool(
|
||||
metadata_filters
|
||||
or query.note_types
|
||||
or query.entity_types
|
||||
or after_date
|
||||
or query.tags
|
||||
or query.status
|
||||
strict_search_text or prepared.title or prepared.permalink or prepared.permalink_match
|
||||
)
|
||||
has_filters = self._prepared_has_filters(prepared)
|
||||
|
||||
with logfire.span(
|
||||
"search.execute",
|
||||
retrieval_mode=retrieval_mode.value,
|
||||
retrieval_mode=prepared.retrieval_mode.value,
|
||||
has_query=has_query,
|
||||
has_filters=has_filters,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
):
|
||||
logger.trace(f"Searching with query: {query}")
|
||||
# First pass: preserve existing strict search behavior.
|
||||
results = await self.repository.search(
|
||||
results = await self._search_repository(
|
||||
prepared,
|
||||
search_text=strict_search_text,
|
||||
permalink=query.permalink,
|
||||
permalink_match=query.permalink_match,
|
||||
title=query.title,
|
||||
note_types=query.note_types,
|
||||
search_item_types=query.entity_types,
|
||||
after_date=after_date,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=query.min_similarity,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
@@ -217,7 +301,9 @@ class SearchService:
|
||||
# Outcome: retry once with relaxed OR terms while preserving explicit boolean intent.
|
||||
if results:
|
||||
return results
|
||||
if not self._is_relaxed_fts_fallback_eligible(query, strict_search_text, retrieval_mode):
|
||||
if not self._is_relaxed_fts_fallback_eligible(
|
||||
query, strict_search_text, prepared.retrieval_mode
|
||||
):
|
||||
return results
|
||||
|
||||
assert strict_search_text is not None
|
||||
@@ -231,26 +317,57 @@ class SearchService:
|
||||
)
|
||||
with logfire.span(
|
||||
"search.relaxed_fts_retry",
|
||||
retrieval_mode=retrieval_mode.value,
|
||||
retrieval_mode=prepared.retrieval_mode.value,
|
||||
token_count=len(self._tokenize_fts_text(strict_search_text)),
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
):
|
||||
return await self.repository.search(
|
||||
return await self._search_repository(
|
||||
prepared,
|
||||
search_text=relaxed_search_text,
|
||||
permalink=query.permalink,
|
||||
permalink_match=query.permalink_match,
|
||||
title=query.title,
|
||||
note_types=query.note_types,
|
||||
search_item_types=query.entity_types,
|
||||
after_date=after_date,
|
||||
metadata_filters=metadata_filters,
|
||||
retrieval_mode=retrieval_mode,
|
||||
min_similarity=query.min_similarity,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
async def count(self, query: SearchQuery) -> int:
|
||||
"""Count all indexed rows matching a query."""
|
||||
prepared = self._prepare_query(query)
|
||||
if prepared is None:
|
||||
return 0
|
||||
|
||||
strict_search_text = prepared.search_text
|
||||
has_query = bool(
|
||||
strict_search_text or prepared.title or prepared.permalink or prepared.permalink_match
|
||||
)
|
||||
has_filters = self._prepared_has_filters(prepared)
|
||||
|
||||
with logfire.span(
|
||||
"search.count",
|
||||
retrieval_mode=prepared.retrieval_mode.value,
|
||||
has_query=has_query,
|
||||
has_filters=has_filters,
|
||||
):
|
||||
total = await self._count_repository(prepared, search_text=strict_search_text)
|
||||
|
||||
if total > 0:
|
||||
return total
|
||||
if not self._is_relaxed_fts_fallback_eligible(
|
||||
query, strict_search_text, prepared.retrieval_mode
|
||||
):
|
||||
return total
|
||||
|
||||
assert strict_search_text is not None
|
||||
relaxed_search_text = self._build_relaxed_fts_query(strict_search_text)
|
||||
if relaxed_search_text == strict_search_text:
|
||||
return total
|
||||
|
||||
with logfire.span(
|
||||
"search.count.relaxed_fts_retry",
|
||||
retrieval_mode=prepared.retrieval_mode.value,
|
||||
token_count=len(self._tokenize_fts_text(strict_search_text)),
|
||||
):
|
||||
return await self._count_repository(prepared, search_text=relaxed_search_text)
|
||||
|
||||
@staticmethod
|
||||
def _tokenize_fts_text(search_text: str) -> list[str]:
|
||||
"""Tokenize text into alphanumeric terms for relaxed FTS fallback."""
|
||||
|
||||
Reference in New Issue
Block a user