feat(search): add tool context deduplication and improve search configuration (#321)

* feat(search): add tool context deduplication and improve search configuration

- Modify _make_tool methods to accept and inject tool_context_id parameter
- Add tool_context_id handling in AS and CC agent wrappers
- Increase search candidate multiplier from 3.0 to 5.0 in default config
- Extend HTTP client timeout from 30s to 3600s
- Add tool context deduplication logic to prevent duplicate search results
- Implement TTL-based expiration for seen chunks in tool contexts
- Add comprehensive unit tests for tool context deduplication behavior
- Update .gitignore to exclude longmemeval directory
- Add time import for timestamp functionality in search step

* refactor(search): replace time module with datetime for timestamp generation

- Removed unused time import
- Added static method _now_ts using datetime.timestamp
- Updated clock parameter to use _now_ts method instead of time.time
- Maintained same timestamp precision and functionality
This commit is contained in:
jinliyl 2026-07-06 15:18:39 +08:00 • committed by GitHub
parent 43a407bc4f
commit 7369342115
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 151 additions and 10 deletions

1
.gitignore vendored
View file

@ -54,5 +54,6 @@ vault/
docs/_build/
site/
longmemeval/
evaluation/
datasets/

View file

@ -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,
)

View file

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

View file

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

View file

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

View file

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

View file

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