fix(search): honor min_score in plain search steps (#338)

This commit is contained in:
Ziyang Guo 2026-07-13 16:32:16 +08:00 • committed by GitHub
parent c5eefe4da3
commit e41b1673ad
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 41 additions and 1 deletions

View file

@ -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:

View file

@ -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:

View file

@ -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."""