"""Hybrid search over file_store using RRF fusion of vector + keyword results.""" import asyncio import datetime import os from typing import Final from ..base_step import BaseStep from ..file_io import extract_daily_date from ...components import R from ...schema import FileChunk from ...utils import expand_links, render_expansion_lines _RRF_K: Final = 60 _MAX_CANDIDATES: Final = 200 _DEFAULT_LIMIT_ENV: Final = "REME_SEARCH_LIMIT" _DEFAULT_LIMIT: Final = 5 def _default_limit() -> int: value = os.getenv(_DEFAULT_LIMIT_ENV) if value is None: return _DEFAULT_LIMIT try: return int(value) except ValueError: return _DEFAULT_LIMIT def _filter_values(value: object) -> set: """Normalize a scalar or collection-valued search filter.""" if isinstance(value, (list, tuple, set, frozenset)): return set(value) return {value} @R.register("search_step") class SearchStep(BaseStep): """Hybrid search: run vector + keyword in parallel, fuse via RRF, filter, truncate.""" TOOL_CONTEXTS_KEY: Final[str] = "tool_contexts" SEARCH_SEEN_KEY: Final[str] = "search_seen_chunk_ids" def __init__( self, *args, seen_ttl_hours: float = 24, **kwargs, ): super().__init__(*args, **kwargs) self.seen_ttl_hours = seen_ttl_hours @staticmethod def _rrf_merge( vector: list[FileChunk], keyword: list[FileChunk], vector_weight: float, ) -> list[FileChunk]: """Fuse two ranked lists with Reciprocal Rank Fusion, keyed by chunk.id.""" text_weight = 1.0 - vector_weight merged: dict[str, FileChunk] = {} for rank, chunk in enumerate(vector, start=1): contrib = vector_weight / (_RRF_K + rank) c = chunk.model_copy(deep=False) c.scores = {**chunk.scores, "vector": chunk.scores.get("vector", chunk.score), "score": contrib} merged[c.id] = c for rank, chunk in enumerate(keyword, start=1): contrib = text_weight / (_RRF_K + rank) existing = merged.get(chunk.id) if existing is not None: existing.scores = { **existing.scores, "keyword": chunk.scores.get("keyword", chunk.score), "score": existing.scores["score"] + contrib, } else: c = chunk.model_copy(deep=False) c.scores = {**chunk.scores, "keyword": chunk.scores.get("keyword", chunk.score), "score": contrib} merged[c.id] = c results = list(merged.values()) results.sort(key=lambda r: r.score, reverse=True) return results @staticmethod def _format_scores(scores: dict[str, float], hybrid: bool) -> str: """Format scores for the answer line: always show fused; show per-branch when hybrid.""" parts = [f"score={scores.get('score', 0.0):.4f}"] if hybrid: for k in ("vector", "keyword"): v = scores.get(k) parts.append(f"{k}={v:.4f}" if v is not None else f"{k}=-") return " ".join(parts) @staticmethod def _now_ts() -> float: return datetime.datetime.now().timestamp() def _tool_context_store(self, tool_context_id: str) -> dict: """Return the mutable state bucket for a tool context.""" if self.app_context is not None: contexts = self.app_context.metadata.setdefault(self.TOOL_CONTEXTS_KEY, {}) else: contexts = self.kwargs.setdefault(self.TOOL_CONTEXTS_KEY, {}) return contexts.setdefault(tool_context_id, {}) def _dedupe_tool_context( self, chunks: list[FileChunk], tool_context_id: str, limit: int, ) -> tuple[list[FileChunk], dict]: now = self.kwargs.get("clock", self._now_ts)() ttl = float( self.kwargs.get( "tool_context_chunk_ttl_seconds", float(self.seen_ttl_hours) * 60 * 60, ), ) store = self._tool_context_store(tool_context_id) seen = store.get(self.SEARCH_SEEN_KEY, {}) if not isinstance(seen, dict): seen = dict.fromkeys(seen, now) before_expire = len(seen) store[self.SEARCH_SEEN_KEY] = seen = {chunk_id: ts for chunk_id, ts in seen.items() if now - float(ts) < ttl} seen_before = len(seen) unvisited = [chunk for chunk in chunks if chunk.id not in seen] returned = unvisited[:limit] for chunk in returned: seen[chunk.id] = now return returned, { "tool_context_id": tool_context_id, "seen_before": seen_before, "skipped_seen": len(chunks) - len(unvisited), "seen_after": len(seen), "expired": before_expire - seen_before, "ttl_seconds": ttl, } async def _resolve_tag_filter( self, raw_tags: object, search_filter: dict, ) -> tuple[dict, dict | None, str | None]: """Merge tag-derived paths into the ordinary file-store filter.""" if not raw_tags: return search_filter, None, None if not self.file_store.tag_index_enabled: self.logger.warning( f"[{self.name}] tags requested but tag_index_unavailable", ) return ( search_filter, {"requested": True, "applied": False, "reason": "tag_index_unavailable"}, None, ) tag_index = self.file_store.require_tag_index() if not tag_index.is_healthy: self.logger.warning(f"[{self.name}] tags requested but tag_index_unavailable") return ( search_filter, {"requested": True, "applied": False, "reason": "tag_index_unavailable"}, None, ) normalized_tags = tag_index.normalize_query_tags(raw_tags) if not normalized_tags: return search_filter, None, "Error: tags contained no valid values" allowed_paths = set(await tag_index.paths_for_tags(normalized_tags, match_all=False)) matched_path_count = len(allowed_paths) if not tag_index.is_healthy: self.logger.warning( f"[{self.name}] tag index became unhealthy during lookup", ) return ( search_filter, {"requested": True, "applied": False, "reason": "tag_index_unavailable"}, None, ) exact_paths = set() has_exact_path_filter = False for key in ("path", "paths"): if key in search_filter: has_exact_path_filter = True exact_paths.update(_filter_values(search_filter.pop(key))) if has_exact_path_filter: allowed_paths.intersection_update(exact_paths) search_filter["paths"] = sorted(allowed_paths) metadata = { "requested": True, "applied": True, "tags": normalized_tags, "matched_paths": matched_path_count, } return search_filter, metadata, None async def execute(self): assert self.context is not None query: str = (self.context.get("query", "") or "").strip() limit: int = int(self.context.get("limit") or _default_limit()) min_score: float = float(self.context.get("min_score") or 0.0) # vector_weight: prefer agent-supplied context value; fallback to YAML kwargs / default 0.7. # Convertible numeric inputs are clipped to [0.0, 1.0]; non-numeric inputs are silently ignored. raw_vw = self.context.get("vector_weight") vector_weight: float | None = None if raw_vw is not None: try: vector_weight = float(raw_vw) except (TypeError, ValueError): self.logger.warning( f"[{self.name}] non-numeric vector_weight={raw_vw!r}; ignoring and using default 0.7", ) vector_weight = None if vector_weight is None: vector_weight = float(self.kwargs.get("vector_weight", 0.7)) vector_weight = max(0.0, min(1.0, vector_weight)) candidate_multiplier: float = float(self.kwargs.get("candidate_multiplier", 5.0)) expand_links_enabled: bool = bool(self.kwargs.get("expand_links", True)) max_links_per_direction: int = int(self.kwargs.get("max_links_per_direction", 10)) tool_context_id: str = (self.context.get("tool_context_id", "") or "").strip() max_search_calls = self.context.get("max_search_calls") strict_date_filter: bool = bool( self.context.get("strict_date_filter") or self.kwargs.get("strict_date_filter", False), ) if not query: self.context.response.success = False self.context.response.answer = "Error: query cannot be empty" return self.context.response assert limit > 0, f"limit must be positive, got {limit}" if max_search_calls is not None: if not tool_context_id or self.app_context is None: raise ValueError("max_search_calls requires an application and tool_context_id") maximum = int(max_search_calls) if maximum < 1: raise ValueError("max_search_calls must be positive") budgets = self.app_context.metadata.setdefault("__search_call_budgets", {}) budget = budgets.setdefault(tool_context_id, {"count": 0, "lock": asyncio.Lock()}) async with budget["lock"]: if budget["count"] >= maximum: self.logger.warning( f"[{self.name}] search budget exhausted context={tool_context_id} limit={maximum}", ) self.context.response.success = False self.context.response.answer = f"Error: search call limit of {maximum} reached" return self.context.response budget["count"] += 1 self.logger.info( f"[{self.name}] search budget context={tool_context_id} " f"call={budget['count']}/{maximum}", ) candidates = min(_MAX_CANDIDATES, max(1, int(limit * candidate_multiplier))) search_filter: dict = dict(self.context.get("search_filter", {}) or {}) raw_tags = self.context.get("tags", []) or [] # Promote top-level date parameters into search_filter for file_store. for date_key in ("start_date", "end_date"): value = self.context.get(date_key) if value and date_key not in search_filter: search_filter[date_key] = value # Validate and normalize date filters before they reach file_store. # _matches_search_filter does lexicographic string comparison against # path_date (always a canonical YYYY-MM-DD), so raw caller values like # "2026-2-28" or "abc" would produce silently wrong results. for date_key in ("start_date", "end_date"): raw = search_filter.get(date_key) if raw is None: continue normalized = extract_daily_date(raw) if normalized is None: # Fallback: accept non-zero-padded dates like "2024-1-5". try: normalized = ( datetime.datetime.strptime( str(raw).strip(), "%Y-%m-%d", ) .date() .isoformat() ) except ValueError: self.logger.warning( f"Ignoring invalid {date_key}={raw!r}; " f"expected a valid YYYY-MM-DD date.", ) del search_filter[date_key] continue search_filter[date_key] = normalized if strict_date_filter: search_filter["strict_date_filter"] = True search_filter, tag_filter_metadata, tag_error = await self._resolve_tag_filter(raw_tags, search_filter) if tag_error is not None: self.context.response.success = False self.context.response.answer = tag_error if tag_filter_metadata is not None: self.context.response.metadata["tag_filter"] = tag_filter_metadata return self.context.response text_weight = 1.0 - vector_weight use_vector = vector_weight > 0.0 use_keyword = text_weight > 0.0 if use_vector and use_keyword: vector_results, keyword_results = await asyncio.gather( self.file_store.vector_search(query, candidates, search_filter), self.file_store.keyword_search(query, candidates, search_filter), ) elif use_vector: vector_results = await self.file_store.vector_search(query, candidates, search_filter) keyword_results = [] else: vector_results = [] keyword_results = await self.file_store.keyword_search(query, candidates, search_filter) self.logger.info( f"[{self.name}] query={query!r} candidates={candidates} " f"vector_hits={len(vector_results)} keyword_hits={len(keyword_results)}", ) hybrid = bool(vector_results) and bool(keyword_results) if not vector_results and not keyword_results: fused: list[FileChunk] = [] elif not keyword_results: fused = vector_results elif not vector_results: fused = keyword_results else: fused = self._rrf_merge(vector_results, keyword_results, vector_weight) if min_score > 0.0: fused = [c for c in fused if c.score >= min_score] dedup: dict | None = None if tool_context_id: fused, dedup = self._dedupe_tool_context(fused, tool_context_id, limit) else: fused = fused[:limit] unique_paths = list(dict.fromkeys(c.path for c in fused)) link_expansion: dict[str, dict] = ( await expand_links(self.file_store, unique_paths, max_links_per_direction) if expand_links_enabled else {} ) answer_lines: list[str] = [] for c in fused: answer_lines.append( f"========== {c.path}:{c.start_line}-{c.end_line} " f"[{self._format_scores(c.scores, hybrid)}] ==========\n{c.text}", ) answer_lines.extend(render_expansion_lines(link_expansion.get(c.path, {}))) self.context.response.answer = "\n".join(answer_lines) self.context.response.metadata["results"] = [ c.model_dump(exclude_none=True, exclude={"embedding"}) for c in fused ] self.context.response.metadata["link_expansion"] = link_expansion self.context.response.metadata["counts"] = { "vector": len(vector_results), "keyword": len(keyword_results), "returned": len(fused), "hybrid": hybrid, } if tag_filter_metadata is not None: self.context.response.metadata["tag_filter"] = tag_filter_metadata if dedup is not None: self.context.response.metadata["dedup"] = dedup return self.context.response