mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
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:
parent
43a407bc4f
commit
7369342115
7 changed files with 151 additions and 10 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -54,5 +54,6 @@ vault/
|
|||
docs/_build/
|
||||
site/
|
||||
|
||||
longmemeval/
|
||||
evaluation/
|
||||
datasets/
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue