From 5fc0805c69262fc0eebe6509d155d4d076350d45 Mon Sep 17 00:00:00 2001 From: oliver Date: Wed, 6 May 2026 23:19:28 +0800 Subject: [PATCH] feat(memory): improve wiki search observability and retrieval flow Expand memory wiki search with context-rich paginated results and bounded query expansion, and add default three-stage retrieval guidance for agents using memory_wiki_search. Co-authored-by: Cursor --- runtime/extensions/memory-wiki/api.py | 149 +++++++++++++++++++++-- runtime/system_prompt.py | 3 + runtime/workspaces/memory/ROLE_SYSTEM.md | 6 + tests/test_memory_wiki_search_api.py | 98 +++++++++++++++ 4 files changed, 246 insertions(+), 10 deletions(-) create mode 100644 tests/test_memory_wiki_search_api.py diff --git a/runtime/extensions/memory-wiki/api.py b/runtime/extensions/memory-wiki/api.py index d41750ba..11246fdd 100644 --- a/runtime/extensions/memory-wiki/api.py +++ b/runtime/extensions/memory-wiki/api.py @@ -128,6 +128,78 @@ def _read_lines(path: Path) -> list[str]: return text.splitlines() +def _clip_text(s: str, *, max_chars: int = 240) -> str: + t = str(s or "").strip() + if len(t) <= max_chars: + return t + return t[: max(1, max_chars - 15)] + "..." + + +def _tokens_for_query(raw_query: str) -> list[str]: + q = str(raw_query or "").strip() + if not q: + return [] + # Keep CJK runs, alpha words, and numbers. + parts = re.findall(r"[\u4e00-\u9fff]+|[A-Za-z]+|\d+", q) + out: list[str] = [] + seen: set[str] = set() + for p in parts: + tok = str(p or "").strip() + key = tok.lower() + if not tok or key in seen: + continue + seen.add(key) + out.append(tok) + return out + + +def _expand_query_variants(raw_query: str) -> list[str]: + query = str(raw_query or "").strip() + if not query: + return [] + variants: list[str] = [query] + synonym_map: dict[str, list[str]] = { + "用户": ["创作者", "项目所有者", "owner", "profile", "identity"], + "身份": ["角色", "identity", "profile"], + "创建者": ["创作者", "所有者", "owner"], + "owner": ["所有者", "项目所有者", "创作者"], + "creator": ["创作者", "创建者", "项目所有者"], + "name": ["名字", "姓名"], + } + toks = _tokens_for_query(query) + for tok in toks: + for syn in synonym_map.get(tok.lower(), []): + if syn not in variants: + variants.append(syn) + if len(toks) > 1: + joined = " ".join(toks) + if joined not in variants: + variants.append(joined) + return variants + + +def _build_line_hit( + *, + lines: list[str], + line_idx: int, + rel_path: str, + matched_query: str, + context_lines: int, +) -> dict[str, Any]: + before_start = max(0, line_idx - context_lines) + before_lines = lines[before_start:line_idx] + after_end = min(len(lines), line_idx + 1 + context_lines) + after_lines = lines[line_idx + 1 : after_end] + return { + "path": rel_path, + "line": int(line_idx + 1), + "text": _clip_text(lines[line_idx]), + "before": [_clip_text(x) for x in before_lines], + "after": [_clip_text(x) for x in after_lines], + "matched_query": str(matched_query or ""), + } + + def _wiki_status(rt: WikiRuntime, args: dict[str, Any]) -> dict[str, Any]: del args files = _list_md_files(rt.wiki_root) @@ -171,17 +243,69 @@ def _wiki_search(rt: WikiRuntime, args: dict[str, Any]) -> dict[str, Any]: is_regex = bool(args.get("is_regex")) req_limit = int(args.get("limit") or rt.max_search_results) limit = max(1, min(req_limit, rt.max_search_results)) + offset = max(0, int(args.get("offset") or 0)) + context_lines = max(0, min(int(args.get("context_lines") or 0), 5)) + path_prefix = str(args.get("path_prefix") or "").strip().replace("\\", "/").lstrip("/") + expand_query = bool(args.get("expand_query")) and not is_regex + max_rounds = max(1, min(int(args.get("max_rounds") or 2), 5)) + + queries = [query] + if expand_query: + queries = _expand_query_variants(query)[:max_rounds] + flags = 0 if case_sensitive else re.IGNORECASE - pattern = re.compile(query if is_regex else re.escape(query), flags=flags) - hits: list[dict[str, Any]] = [] - for fp in _list_md_files(rt.wiki_root): - rel = str(fp.relative_to(rt.wiki_root)).replace("\\", "/") - for idx, line in enumerate(_read_lines(fp), start=1): - if pattern.search(line): - hits.append({"path": rel, "line": idx, "text": line.strip()}) - if len(hits) >= limit: - return {"ok": True, "query": query, "hits": hits, "truncated": True} - return {"ok": True, "query": query, "hits": hits, "truncated": False} + hits_all: list[dict[str, Any]] = [] + seen_keys: set[tuple[str, int]] = set() + files_scanned = 0 + + for qv in queries: + pattern = re.compile(qv if is_regex else re.escape(qv), flags=flags) + for fp in _list_md_files(rt.wiki_root): + rel = str(fp.relative_to(rt.wiki_root)).replace("\\", "/") + if path_prefix and not rel.startswith(path_prefix): + continue + files_scanned += 1 + lines = _read_lines(fp) + for idx, line in enumerate(lines): + if not pattern.search(line): + continue + key = (rel, int(idx + 1)) + if key in seen_keys: + continue + seen_keys.add(key) + hits_all.append( + _build_line_hit( + lines=lines, + line_idx=idx, + rel_path=rel, + matched_query=qv, + context_lines=context_lines, + ) + ) + + hits_page = hits_all[offset : offset + limit] + next_offset = offset + len(hits_page) + truncated = next_offset < len(hits_all) + query_used = str(hits_page[0].get("matched_query") or query) if hits_page else query + + # Backward-compatible fields remain: ok/query/hits/truncated. + # New fields are additive diagnostics/controls for iterative search flows. + return { + "ok": True, + "query": query, + "hits": hits_page, + "truncated": bool(truncated), + "query_used": query_used, + "queries_attempted": queries, + "total_hits_estimate": len(hits_all), + "offset": offset, + "next_offset": next_offset if truncated else None, + "limit": limit, + "context_lines": context_lines, + "path_prefix": path_prefix, + "files_scanned": files_scanned, + "expanded": bool(expand_query), + } def _wiki_lint(rt: WikiRuntime, args: dict[str, Any]) -> dict[str, Any]: @@ -290,8 +414,13 @@ def build_wiki_tool_specs(api: Any) -> list[dict[str, Any]]: "properties": { "query": {"type": "string"}, "limit": {"type": "integer"}, + "offset": {"type": "integer"}, "is_regex": {"type": "boolean"}, "case_sensitive": {"type": "boolean"}, + "context_lines": {"type": "integer"}, + "path_prefix": {"type": "string"}, + "expand_query": {"type": "boolean"}, + "max_rounds": {"type": "integer"}, }, "required": ["query"], }, diff --git a/runtime/system_prompt.py b/runtime/system_prompt.py index ce751122..78a01240 100644 --- a/runtime/system_prompt.py +++ b/runtime/system_prompt.py @@ -46,6 +46,9 @@ def _unified_skill_policy_guidance() -> str: "- 不要为了“列出技能”而去读取 SKILL.md。只有在你确实需要某个技能的详细使用说明时,才读取对应 SKILL.md。\n" "- 当你需要技能细节时,请按目录中给出的 path 读取对应的 SKILL.md。\n" "- 当对话涉及长期记忆、用户身份/偏好、项目背景延续、复发问题沉淀时,优先启用 wiki-first-autonomy 技能,并优先使用 memory_wiki_search/memory_wiki_get 检索上下文,再执行与回复。\n" + "- 使用 memory_wiki_search 时默认采用三段检索:先 `query` 精确检索;若结果不足再启用 `expand_query=true`;仍不足时加 `path_prefix` 定向到候选目录后重检。\n" + "- 三段检索建议参数:第一段 `limit=5~8`;第二段 `expand_query=true,max_rounds=2~3`;第三段在保留前述参数基础上增加 `path_prefix` 并分页(`offset`/`limit`)。\n" + "- 每段检索后先依据 `queries_attempted`、`total_hits_estimate`、`next_offset` 判断是否继续,避免一次无命中就直接下结论。\n" "- memory wiki 请仅传相对 wiki 根的路径(相对 `data/wiki`),例如 `improvement/learnings.md`;不要传 `data/wiki/...` 或绝对路径。\n" "- 当新增事实会影响后续决策时,完成当前任务后使用 memory_wiki_apply 写入结构化记忆,并用 memory_wiki_lint 做质量检查。\n" "- 技能包由说明文档和可选文件组成。运行时不会自动执行技能 `scripts/` 目录下的文件;\n" diff --git a/runtime/workspaces/memory/ROLE_SYSTEM.md b/runtime/workspaces/memory/ROLE_SYSTEM.md index 4c1fb8c3..1ce9ea5e 100644 --- a/runtime/workspaces/memory/ROLE_SYSTEM.md +++ b/runtime/workspaces/memory/ROLE_SYSTEM.md @@ -15,6 +15,12 @@ - 质量检查:`memory_wiki_lint`。 - 写入/更新:`memory_wiki_apply`(write/append/delete)。 +## wiki 检索默认流程(三段): +- 第1段:先用 `memory_wiki_search` 做 `query` 精确检索(建议 `limit=5~8`)。 +- 第2段:若结果不足,启用 `expand_query=true`(建议 `max_rounds=2~3`)做补搜。 +- 第3段:若仍不足,结合候选目录设置 `path_prefix` 做定向检索,并使用 `offset`/`limit` 分页补全。 +- 每段结束后必须检查 `queries_attempted`、`total_hits_estimate`、`next_offset` 再决定是否继续,不允许“一次没命中就断言不存在”。 + ## 记忆策略(按需记忆): - 满足“稳定、可复用、可检索”时才写入;否则不写入并说明原因。 - 写入前先检索相近条目,优先增量更新,避免重复堆砌。 diff --git a/tests/test_memory_wiki_search_api.py b/tests/test_memory_wiki_search_api.py new file mode 100644 index 00000000..b81a50c3 --- /dev/null +++ b/tests/test_memory_wiki_search_api.py @@ -0,0 +1,98 @@ +from __future__ import annotations + +import importlib.util +from pathlib import Path +import sys +from types import SimpleNamespace + +from oclaw.platform.config.paths import PROJECT_ROOT + + +def _wiki_search_handler(tmp_path: Path): + api_path = (PROJECT_ROOT / "runtime" / "extensions" / "memory-wiki" / "api.py").resolve() + spec = importlib.util.spec_from_file_location("memory_wiki_api_test", str(api_path)) + assert spec is not None and spec.loader is not None + mod = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = mod + spec.loader.exec_module(mod) # type: ignore[assignment] + wiki_root = tmp_path / "wiki" + wiki_root.mkdir(parents=True, exist_ok=True) + rows = mod.build_wiki_tool_specs( + SimpleNamespace(plugin_config={"wiki_root": str(wiki_root), "max_search_results": 10, "max_get_lines": 200}) + ) + for r in rows: + if str(r.get("name") or "") == "wiki_search": + return wiki_root, r["handler"] + raise AssertionError("wiki_search handler missing") + + +def test_wiki_search_returns_context_and_pagination_metadata(tmp_path: Path) -> None: + wiki_root, handler = _wiki_search_handler(tmp_path) + p = wiki_root / "users" / "identity.md" + p.parent.mkdir(parents=True, exist_ok=True) + p.write_text( + "\n".join( + [ + "# Identity", + "creator is Oliver", + "project owner is Oliver", + "session profile line", + ] + ), + encoding="utf-8", + ) + + out = handler({"query": "Oliver", "limit": 1, "offset": 0, "context_lines": 1}) + + assert out.get("ok") is True + assert out.get("query") == "Oliver" + assert out.get("truncated") is True + assert out.get("next_offset") == 1 + hits = out.get("hits") or [] + assert len(hits) == 1 + hit = hits[0] + assert hit.get("path") == "users/identity.md" + assert isinstance(hit.get("before"), list) + assert isinstance(hit.get("after"), list) + assert "matched_query" in hit + + +def test_wiki_search_offset_paginates_deterministically(tmp_path: Path) -> None: + wiki_root, handler = _wiki_search_handler(tmp_path) + p = wiki_root / "identity.md" + p.write_text("Oliver one\nOliver two\nOliver three\n", encoding="utf-8") + + page1 = handler({"query": "Oliver", "limit": 2, "offset": 0}) + page2 = handler({"query": "Oliver", "limit": 2, "offset": 2}) + + hits1 = page1.get("hits") or [] + hits2 = page2.get("hits") or [] + assert [h.get("line") for h in hits1] == [1, 2] + assert [h.get("line") for h in hits2] == [3] + assert page2.get("truncated") is False + assert page2.get("next_offset") is None + + +def test_wiki_search_expand_query_reports_attempted_queries(tmp_path: Path) -> None: + wiki_root, handler = _wiki_search_handler(tmp_path) + p = wiki_root / "users" / "identity.md" + p.parent.mkdir(parents=True, exist_ok=True) + p.write_text("项目所有者:Oliver\n", encoding="utf-8") + + out = handler({"query": "用户", "expand_query": True, "max_rounds": 3, "limit": 5}) + assert out.get("ok") is True + attempted = out.get("queries_attempted") or [] + assert attempted + assert attempted[0] == "用户" + assert len(attempted) <= 3 + hits = out.get("hits") or [] + assert hits + assert any(h.get("matched_query") != "用户" for h in hits) + + +def test_wiki_search_keeps_backward_compatible_keys(tmp_path: Path) -> None: + wiki_root, handler = _wiki_search_handler(tmp_path) + (wiki_root / "a.md").write_text("hello\n", encoding="utf-8") + out = handler({"query": "hello"}) + assert {"ok", "query", "hits", "truncated"}.issubset(set(out.keys())) +