Merge pull request #115 from ddggkkcc/pr/rag-crag-refusal

[Feature] Wiki 块级问答增加评估路由与诚实拒答(CRAG 式):知识库没有答案时不再硬答
This commit is contained in:
countbot-ai
2026-09-04 12:02:02 +08:00
committed by GitHub
2 changed files with 297 additions and 10 deletions
+136 -10
View File
@@ -1,8 +1,10 @@
"""Wiki Tool - Agent工具接口"""
import json
import os
import re
from pathlib import Path
from typing import Any, Dict, Optional
from typing import Any, Dict, List, Optional, Tuple
from loguru import logger
from backend.modules.tools.base import Tool
@@ -258,11 +260,141 @@ class WikiTool(Tool):
return "\n".join(lines)
async def _rag_ask(self, question: str) -> str:
"""块级问答:注入 top-K 相关块(而非整篇),带溯源"""
"""块级问答:检索 → LLM 评估相关性 → 三段路由(直接生成 / 过滤后生成 / 改写重试后拒答)。
评估不可用(无 provider / 调用失败 / 输出不可解析)时回退为
"全部注入直接生成"(即评估前的行为),保证零破坏。
"""
chunks = self._rag.search_chunks(question, top_k=6)
if not chunks:
return "Wiki 知识库为空或没有找到相关内容。"
grade, relevant = await self._grade_chunks(question, chunks)
if grade == "none":
# 检索结果全不相关:改写问题重试一次,仍不相关则如实拒答
rewritten = await self._rewrite_query(question)
if rewritten and rewritten != question:
chunks2 = self._rag.search_chunks(rewritten, top_k=6)
if chunks2:
grade2, relevant2 = await self._grade_chunks(rewritten, chunks2)
if grade2 == "all":
return await self._generate_from_chunks(rewritten, chunks2)
if grade2 == "partial":
return await self._generate_from_chunks(
rewritten, self._filter_chunks(chunks2, relevant2) or chunks2
)
return ("Wiki 知识库中没有找到与该问题相关的内容。"
"(已检索并逐条校验相关性,结果均与问题无关;"
"可以换个问法重试,或确认知识库中是否已有相关条目。)")
if grade == "partial" and relevant:
chunks = self._filter_chunks(chunks, relevant) or chunks
return await self._generate_from_chunks(question, chunks)
# ---------- 块级问答的评估与路由组件(仅 COUNTBOT_RAG_CHUNKS=1 路径使用) ----------
@staticmethod
def _get_provider():
"""获取共享 LLM provider不可用时返回 None各调用方自行回退"""
try:
from backend.app import get_shared_provider
return get_shared_provider()
except Exception:
return None
async def _grade_chunks(self, question: str, chunks: List[dict]) -> Tuple[Optional[str], Optional[List[int]]]:
"""LLM 评估检索块与问题的相关性(单次轻量调用,约 200-400 token
Returns:
("all"|"partial"|"none", relevant_ids);评估不可用时 (None, None)。
纯分数阈值无法做拒答:实测负样本 top1 分数与正样本重叠率 7/10。
"""
provider = self._get_provider()
if provider is None:
return None, None
lines = [
f"{i}. {c['doc_title']} {c['section']}{c['content'].strip()[:120]}"
for i, c in enumerate(chunks, 1)
]
prompt = (
"你是知识库检索质量评估器。判断下面的检索结果能否支撑回答问题。\n\n"
f"问题:{question}\n\n"
"检索结果(编号. 文档 章节:内容摘录):\n" + "\n".join(lines) + "\n\n"
'只输出一行 JSON不要输出其他内容\n'
'{"grade": "all|partial|none", "relevant": [相关编号]}\n'
"- all检索结果基本都与问题相关足以回答\n"
"- partial有任何一条结果可能包含与问题相关的信息"
"(哪怕只覆盖问题的一部分、或只提供部分线索),"
"relevant 列出这些结果的编号\n"
"- none仅当所有结果谈论的都是与问题完全无关的主题时使用"
"拿不准时优先 partial不要轻易判 none"
)
try:
resp = await provider.chat_completion(prompt, max_tokens=200, temperature=0.0)
return self._parse_grade(resp, len(chunks))
except Exception as e:
logger.warning(f"Chunk grading failed, falling back to plain generation: {e}")
return None, None
@staticmethod
def _parse_grade(text: str, n_chunks: int) -> Tuple[Optional[str], Optional[List[int]]]:
"""解析评估输出为 (grade, relevant_ids);不可解析时返回 (None, None)"""
if not text:
return None, None
m = re.search(r"\{[^{}]*\}", text, re.S)
if not m:
return None, None
try:
data = json.loads(m.group(0))
except Exception:
return None, None
grade = str(data.get("grade", "")).strip().lower()
if grade not in ("all", "partial", "none"):
return None, None
ids = set()
for i in (data.get("relevant") or []):
try:
n = int(i)
except (TypeError, ValueError):
continue
if 1 <= n <= n_chunks:
ids.add(n)
return grade, sorted(ids)
async def _rewrite_query(self, question: str) -> Optional[str]:
"""全不相关时,让 LLM 把问题改写为更贴近知识库术语的检索词(一次机会)"""
provider = self._get_provider()
if provider is None:
return None
prompt = (
"把下面的问题改写成更适合知识库关键词检索的问法"
"(保留原意,尽量使用文档中可能出现的术语,不要回答问题本身),"
"只输出改写后的问题:\n" + question
)
try:
resp = await provider.chat_completion(prompt, max_tokens=100, temperature=0.0)
text = (resp or "").strip().strip('"').strip("\u201c\u201d")
return text or None
except Exception as e:
logger.warning(f"Query rewrite failed: {e}")
return None
@staticmethod
def _filter_chunks(chunks: List[dict], relevant_ids: List[int]) -> List[dict]:
"""按评估给出的编号保留相关块(编号从 1 开始)"""
if not relevant_ids:
return []
return [c for i, c in enumerate(chunks, 1) if i in set(relevant_ids)]
async def _generate_from_chunks(self, question: str, chunks: List[dict]) -> str:
"""用给定块组装上下文并生成回答;无 provider 或失败时回退块级搜索结果"""
provider = self._get_provider()
if not provider:
return self._rag_search(question, top_k=6)
context_parts = []
for c in chunks:
context_parts.append(
@@ -270,14 +402,7 @@ class WikiTool(Tool):
)
context = "\n\n".join(context_parts)
try:
from backend.app import get_shared_provider
provider = get_shared_provider()
if not provider:
return self._rag_search(question, top_k=6)
prompt = f"""请根据以下 Wiki 知识库内容回答问题。引用时注明来源 [slug#section]。
prompt = f"""请根据以下 Wiki 知识库内容回答问题。引用时注明来源 [slug#section]。
问题:{question}
@@ -288,6 +413,7 @@ class WikiTool(Tool):
---
如果知识库中没有相关内容,请如实告知。"""
try:
return await provider.chat_completion(prompt, max_tokens=2000, temperature=0.3)
except Exception:
return self._rag_search(question, top_k=6)
+161
View File
@@ -214,3 +214,164 @@ class TestRagServiceInternals:
assert all("docker-compose" not in r["content"]
for r in rag.search_chunks("docker-compose 一键启动"))
assert any("helm" in r["content"] for r in rag.search_chunks("k8s helm"))
class ScriptedProvider:
"""按调用顺序返回预设响应的假 LLM providerException 实例则抛出"""
def __init__(self, replies):
self.replies = list(replies)
self.calls = []
async def chat_completion(self, prompt, **kwargs):
self.calls.append(prompt)
if not self.replies:
return "DEFAULT"
reply = self.replies.pop(0)
if isinstance(reply, Exception):
raise reply
return reply
def _install_provider(monkeypatch, provider):
fake_app = types.ModuleType("backend.app")
fake_app.get_shared_provider = lambda: provider
monkeypatch.setitem(sys.modules, "backend.app", fake_app)
GRADING_MARK = "检索质量评估器"
REWRITE_MARK = "改写成更适合知识库关键词检索"
GEN_MARK = "请根据以下 Wiki 知识库内容回答问题"
class TestRagAskGrading:
"""块级问答的评估-路由CRAG行为三段路由 + 拒答 + 全链路回退"""
def test_grade_all_generates_directly(self, wiki_dir, rag_env, monkeypatch):
provider = ScriptedProvider([
'{"grade": "all", "relevant": [1, 2]}',
"ANSWER_OK",
])
_install_provider(monkeypatch, provider)
tool = WikiTool(wiki_dir)
out = asyncio.run(tool._handle_ask("如何用 Docker 部署"))
assert out == "ANSWER_OK"
assert len(provider.calls) == 2 # 评估 1 次 + 生成 1 次,无改写
assert GRADING_MARK in provider.calls[0]
assert GEN_MARK in provider.calls[1]
assert "docker-compose" in provider.calls[1] # 生成上下文含检索块
def test_grade_partial_filters_unrelated_chunks(self, wiki_dir, rag_env, monkeypatch):
provider = ScriptedProvider([
'{"grade": "partial", "relevant": [1]}',
"ANSWER_OK",
])
_install_provider(monkeypatch, provider)
tool = WikiTool(wiki_dir)
out = asyncio.run(tool._handle_ask("如何用 Docker 部署"))
assert out == "ANSWER_OK"
gen_prompt = provider.calls[1]
# 只注入编号 1 的块无关文档memory不得进入生成上下文
assert "memory.md" not in gen_prompt
assert "检索" not in gen_prompt.split("问题")[0] or "Docker" in gen_prompt
def test_grade_none_refuses_after_one_rewrite(self, wiki_dir, rag_env, monkeypatch):
provider = ScriptedProvider([
'{"grade": "none", "relevant": []}', # 首次评估:全不相关
"Docker 部署方法", # 改写
'{"grade": "none", "relevant": []}', # 重试评估:仍不相关
])
_install_provider(monkeypatch, provider)
tool = WikiTool(wiki_dir)
out = asyncio.run(tool._handle_ask("如何用 Docker 部署"))
assert "没有找到" in out
# 评估 2 次 + 改写 1 次,绝不发起生成
assert len(provider.calls) == 3
assert not any(GEN_MARK in c for c in provider.calls)
def test_grade_none_retry_then_succeeds(self, wiki_dir, rag_env, monkeypatch):
provider = ScriptedProvider([
'{"grade": "none", "relevant": []}', # 首次评估:不相关
"docker compose 部署", # 改写
'{"grade": "all", "relevant": [1]}', # 重试评估:相关
"ANSWER_AFTER_RETRY",
])
_install_provider(monkeypatch, provider)
tool = WikiTool(wiki_dir)
out = asyncio.run(tool._handle_ask("如何用 Docker 部署"))
assert out == "ANSWER_AFTER_RETRY"
assert len(provider.calls) == 4
def test_grading_failure_falls_back_to_generation(self, wiki_dir, rag_env, monkeypatch):
provider = ScriptedProvider([
RuntimeError("grading exploded"), # 评估失败 → 回退旧行为:直接生成
"ANSWER_FALLBACK",
])
_install_provider(monkeypatch, provider)
tool = WikiTool(wiki_dir)
out = asyncio.run(tool._handle_ask("如何用 Docker 部署"))
assert out == "ANSWER_FALLBACK"
assert len(provider.calls) == 2
def test_unparseable_grade_falls_back(self, wiki_dir, rag_env, monkeypatch):
provider = ScriptedProvider([
"我觉得这些结果看起来都挺好的!", # 非 JSON → 评估不可用
"ANSWER_UNPARSED",
])
_install_provider(monkeypatch, provider)
tool = WikiTool(wiki_dir)
out = asyncio.run(tool._handle_ask("如何用 Docker 部署"))
assert out == "ANSWER_UNPARSED"
assert len(provider.calls) == 2
def test_no_provider_keeps_legacy_fallback(self, wiki_dir, rag_env, monkeypatch):
"""无 provider评估/改写/生成全部跳过,回退块级搜索(与旧行为一致)"""
provider = ScriptedProvider([])
_install_provider(monkeypatch, None)
tool = WikiTool(wiki_dir)
out = asyncio.run(tool._handle_ask("如何用 Docker 部署"))
assert "matching sections" in out
assert len(provider.calls) == 0
class TestParseGrade:
"""_parse_gradeLLM 输出 → (grade, relevant_ids) 的解析契约"""
@pytest.mark.parametrize("text,grade,rel", [
('{"grade": "all", "relevant": [1, 2]}', "all", [1, 2]),
('好的,这是我的判断:\n{"grade": "partial", "relevant": [2]}', "partial", [2]),
('{"grade": "NONE", "relevant": []}', "none", []),
('{"grade": "none"}', "none", []),
])
def test_valid(self, text, grade, rel):
assert WikiTool._parse_grade(text, 6) == (grade, rel)
@pytest.mark.parametrize("text", [
"",
"完全没有 JSON",
'{"grade": "maybe"}',
'{"grade": "partial", "relevant": "not-a-list"}', # relevant 非列表 → 空编号
])
def test_invalid_returns_none(self, text):
if "not-a-list" in text:
grade, rel = WikiTool._parse_grade(text, 6)
assert grade == "partial" and rel == []
else:
assert WikiTool._parse_grade(text, 6) == (None, None)
def test_out_of_range_ids_dropped(self):
grade, rel = WikiTool._parse_grade('{"grade": "partial", "relevant": [0, 1, 9, "2"]}', 6)
assert grade == "partial"
assert rel == [1, 2]