diff --git a/.gitignore b/.gitignore index badafa78..8510c30a 100644 --- a/.gitignore +++ b/.gitignore @@ -54,5 +54,6 @@ vault/ docs/_build/ site/ +longmemeval/ evaluation/ datasets/ diff --git a/reme/components/agent_wrapper/as_agent_wrapper.py b/reme/components/agent_wrapper/as_agent_wrapper.py index e4b7d683..37dd73c5 100644 --- a/reme/components/agent_wrapper/as_agent_wrapper.py +++ b/reme/components/agent_wrapper/as_agent_wrapper.py @@ -96,9 +96,12 @@ class AsAgentWrapper(BaseAgentWrapper): self.session_retention_days = int(session_retention_days) self._session_cleanup_done = False - @staticmethod - def _make_tool(job: "BaseJob") -> FunctionTool: + @classmethod + def _make_tool(cls, job: "BaseJob", tool_context_id: str | None = None) -> FunctionTool: async def run_job(**kwargs) -> ToolChunk: + if tool_context_id: + assert "tool_context_id" not in kwargs, "tool_context_id is injected by agent_wrapper" + kwargs["tool_context_id"] = tool_context_id response = await job(**kwargs) state = ToolResultState.SUCCESS if response.success else ToolResultState.ERROR return ToolChunk(content=[TextBlock(text=str(response.answer))], state=state) @@ -216,8 +219,9 @@ class AsAgentWrapper(BaseAgentWrapper): job_tools: list[str] = kwargs.get("job_tools", []) resolved_jobs = self._resolve_job_tools(job_tools) skills = self._resolve_skills(kwargs.get("skills")) + tool_context_id = kwargs.get("tool_context_id") toolkit = kwargs.get("toolkit") or Toolkit( - tools=[*self._builtin_tools(), *(self._make_tool(job) for job in resolved_jobs)], + tools=[*self._builtin_tools(), *(self._make_tool(job, tool_context_id) for job in resolved_jobs)], skills_or_loaders=skills, ) diff --git a/reme/components/agent_wrapper/cc_agent_wrapper.py b/reme/components/agent_wrapper/cc_agent_wrapper.py index a67e50c0..525dd2cd 100644 --- a/reme/components/agent_wrapper/cc_agent_wrapper.py +++ b/reme/components/agent_wrapper/cc_agent_wrapper.py @@ -216,11 +216,14 @@ class CcAgentWrapper(BaseAgentWrapper): except OSError as exc: self.logger.warning(f"Failed to link Claude Code skills directory {target}: {exc}") - @staticmethod - def _make_tool(job: "BaseJob"): + @classmethod + def _make_tool(cls, job: "BaseJob", tool_context_id: str | None = None): from claude_agent_sdk import SdkMcpTool async def run_job(args): + if tool_context_id: + assert "tool_context_id" not in args, "tool_context_id is injected by agent_wrapper" + args["tool_context_id"] = tool_context_id response = await job(**args) return {"content": [{"type": "text", "text": str(response.answer)}], "is_error": not response.success} @@ -284,7 +287,7 @@ class CcAgentWrapper(BaseAgentWrapper): job_tools: list[str] = kwargs.get("job_tools", []) resolved_jobs = self._resolve_job_tools(job_tools) if resolved_jobs: - sdk_tools = [self._make_tool(job) for job in resolved_jobs] + sdk_tools = [self._make_tool(job, kwargs.get("tool_context_id")) for job in resolved_jobs] server = create_sdk_mcp_server(name="mcp_server", tools=sdk_tools) opts.mcp_servers = opts.mcp_servers if isinstance(opts.mcp_servers, dict) else {} opts.mcp_servers["mcp_server"] = server diff --git a/reme/components/client/http_client.py b/reme/components/client/http_client.py index 871c0d3a..34aac1f4 100644 --- a/reme/components/client/http_client.py +++ b/reme/components/client/http_client.py @@ -21,7 +21,7 @@ class HttpClient(BaseClient): self, host: str | None = None, port: int | None = None, - timeout: float = 30.0, + timeout: float = 3600.0, **kwargs, ): super().__init__(**kwargs) diff --git a/reme/config/default.yaml b/reme/config/default.yaml index 7def83b7..83b29f1f 100644 --- a/reme/config/default.yaml +++ b/reme/config/default.yaml @@ -293,7 +293,7 @@ jobs: steps: - backend: search_step vector_weight: 0.7 - candidate_multiplier: 3.0 + candidate_multiplier: 5.0 expand_links: true max_links_per_direction: 10 diff --git a/reme/steps/index/search.py b/reme/steps/index/search.py index 6090c3c1..a22942d7 100644 --- a/reme/steps/index/search.py +++ b/reme/steps/index/search.py @@ -17,6 +17,19 @@ _MAX_CANDIDATES = 200 class SearchStep(BaseStep): """Hybrid search: run vector + keyword in parallel, fuse via RRF, filter, truncate.""" + def __init__( + self, + *args, + tool_contexts_key: str = "tool_contexts", + search_seen_key: str = "search_seen_chunk_ids", + tool_context_chunk_ttl_hours: float = 24, + **kwargs, + ): + super().__init__(*args, **kwargs) + self.tool_contexts_key = tool_contexts_key + self.search_seen_key = search_seen_key + self.tool_context_chunk_ttl_hours = tool_context_chunk_ttl_hours + @staticmethod def _rrf_merge( vector: list[FileChunk], @@ -61,15 +74,63 @@ class SearchStep(BaseStep): 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.tool_context_chunk_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 execute(self): 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) vector_weight: float = float(self.kwargs.get("vector_weight", 0.7)) - candidate_multiplier: float = float(self.kwargs.get("candidate_multiplier", 3.0)) + 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() strict_date_filter: bool = bool( self.context.get("strict_date_filter") or self.kwargs.get("strict_date_filter", False), ) @@ -143,7 +204,12 @@ class SearchStep(BaseStep): if min_score > 0.0: fused = [c for c in fused if c.score >= min_score] - fused = fused[:limit] + + 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] = ( @@ -169,4 +235,6 @@ class SearchStep(BaseStep): "returned": len(fused), "hybrid": hybrid, } + if dedup is not None: + self.context.response.metadata["dedup"] = dedup return self.context.response diff --git a/tests/unit/test_search_step.py b/tests/unit/test_search_step.py index c5bf1e03..7d347255 100644 --- a/tests/unit/test_search_step.py +++ b/tests/unit/test_search_step.py @@ -130,6 +130,71 @@ def test_search_step_keyword_only_uses_keyword_scores_and_min_score(): 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.""" + + async def run(): + chunks = [ + _chunk("a", "daily/a.md", "first", "keyword", 5.0), + _chunk("b", "daily/b.md", "second", "keyword", 4.0), + _chunk("c", "daily/c.md", "third", "keyword", 3.0), + ] + store = FakeSearchStore(keyword_results=chunks) + step = SearchStep(file_store=store, expand_links=False) + + first = await step(RuntimeContext(query="alpha", limit=2, tool_context_id="ctx-1")) + second = await step(RuntimeContext(query="alpha", limit=2, tool_context_id="ctx-1")) + third = await step(RuntimeContext(query="alpha", limit=2)) + + assert [r["id"] for r in first.metadata["results"]] == ["a", "b"] + assert first.metadata["dedup"] == { + "tool_context_id": "ctx-1", + "seen_before": 0, + "skipped_seen": 0, + "seen_after": 2, + "expired": 0, + "ttl_seconds": 86400.0, + } + assert [r["id"] for r in second.metadata["results"]] == ["c"] + assert second.metadata["dedup"]["seen_before"] == 2 + assert second.metadata["dedup"]["skipped_seen"] == 2 + assert second.metadata["dedup"]["seen_after"] == 3 + assert [r["id"] for r in third.metadata["results"]] == ["a", "b"] + assert "dedup" not in third.metadata + + asyncio.run(run()) + + +def test_search_step_tool_context_seen_chunks_expire_after_ttl(): + """Seen chunk ids under a tool_context_id are reusable after the configured TTL.""" + + async def run(): + now = 1000.0 + chunks = [ + _chunk("a", "daily/a.md", "first", "keyword", 5.0), + _chunk("b", "daily/b.md", "second", "keyword", 4.0), + ] + store = FakeSearchStore(keyword_results=chunks) + step = SearchStep( + file_store=store, + expand_links=False, + tool_context_chunk_ttl_hours=1, + clock=lambda: now, + ) + + first = await step(RuntimeContext(query="alpha", limit=1, tool_context_id="ctx-1")) + now = 4601.0 + second = await step(RuntimeContext(query="alpha", limit=1, tool_context_id="ctx-1")) + + assert [r["id"] for r in first.metadata["results"]] == ["a"] + assert [r["id"] for r in second.metadata["results"]] == ["a"] + assert second.metadata["dedup"]["expired"] == 1 + assert second.metadata["dedup"]["seen_before"] == 0 + assert second.metadata["dedup"]["ttl_seconds"] == 3600.0 + + asyncio.run(run()) + + def test_search_step_empty_query_fails_before_store_calls(): """Empty queries fail fast and do not call file_store search methods."""