diff --git a/reme/steps/index/bm25_search.py b/reme/steps/index/bm25_search.py index b9fc47f3..1ea71fb2 100644 --- a/reme/steps/index/bm25_search.py +++ b/reme/steps/index/bm25_search.py @@ -47,6 +47,7 @@ class Bm25SearchStep(BaseStep): assert self.context is not None query: str = (self.context.get("query", "") or "").strip() limit: int = int(self.context.get("limit") or 5) + min_score: float = float(self.context.get("min_score") or 0.0) tool_context_id: str = (self.context.get("tool_context_id", "") or "").strip() if not query: @@ -59,6 +60,9 @@ class Bm25SearchStep(BaseStep): results = await self.file_store.keyword_search(query, candidates, {}) self.logger.info(f"[{self.name}] query={query!r} candidates={candidates} hits={len(results)}") + if min_score > 0.0: + results = [chunk for chunk in results if chunk.score >= min_score] + if tool_context_id: results = self._dedupe_tool_context(results, tool_context_id, limit) else: diff --git a/reme/steps/index/vector_search.py b/reme/steps/index/vector_search.py index 673252c2..492bac0e 100644 --- a/reme/steps/index/vector_search.py +++ b/reme/steps/index/vector_search.py @@ -47,6 +47,7 @@ class VectorSearchStep(BaseStep): assert self.context is not None query: str = (self.context.get("query", "") or "").strip() limit: int = int(self.context.get("limit") or 5) + min_score: float = float(self.context.get("min_score") or 0.0) tool_context_id: str = (self.context.get("tool_context_id", "") or "").strip() if not query: @@ -59,6 +60,9 @@ class VectorSearchStep(BaseStep): results = await self.file_store.vector_search(query, candidates, {}) self.logger.info(f"[{self.name}] query={query!r} candidates={candidates} hits={len(results)}") + if min_score > 0.0: + results = [chunk for chunk in results if chunk.score >= min_score] + if tool_context_id: results = self._dedupe_tool_context(results, tool_context_id, limit) else: diff --git a/tests/unit/test_search_step.py b/tests/unit/test_search_step.py index 69b06594..ec38e63a 100644 --- a/tests/unit/test_search_step.py +++ b/tests/unit/test_search_step.py @@ -7,7 +7,7 @@ from reme.components import ApplicationContext from reme.components.runtime_context import RuntimeContext from reme.enumeration import LinkScopeEnum from reme.schema import FileChunk, FileLink, FileNode -from reme.steps.index import AddDraftStep, ReadAllDraftStep, SearchStep +from reme.steps.index import AddDraftStep, Bm25SearchStep, ReadAllDraftStep, SearchStep, VectorSearchStep class FakeSearchStore(BaseFileStore): @@ -151,6 +151,38 @@ def test_search_step_keyword_only_uses_keyword_scores_and_min_score(): asyncio.run(run()) +def test_plain_search_steps_apply_min_score_before_truncation(): + """Vector-only and BM25-only tools should not return hits below ``min_score``.""" + + async def run(): + vector_store = FakeSearchStore( + vector_results=[ + _chunk("vector-high", "daily/high.md", "strong vector hit", "vector", 0.9), + _chunk("vector-low", "daily/low.md", "weak vector hit", "vector", 0.2), + ], + ) + keyword_store = FakeSearchStore( + keyword_results=[ + _chunk("keyword-high", "daily/high.md", "strong keyword hit", "keyword", 4.0), + _chunk("keyword-low", "daily/low.md", "weak keyword hit", "keyword", 0.2), + ], + ) + + vector = await VectorSearchStep(file_store=vector_store)( + RuntimeContext(query="alpha", limit=5, min_score=0.5), + ) + keyword = await Bm25SearchStep(file_store=keyword_store)( + RuntimeContext(query="alpha", limit=5, min_score=1.0), + ) + + assert [result["id"] for result in vector.metadata["results"]] == ["vector-high"] + assert [result["id"] for result in keyword.metadata["results"]] == ["keyword-high"] + assert vector.answer == "strong vector hit" + assert keyword.answer == "strong keyword hit" + + asyncio.run(run()) + + def test_search_step_tool_context_deduplicates_returned_chunks_only(): """When tool_context_id is supplied, repeated searches skip previously returned chunks."""