mirror of
https://github.com/countbot-ai/CountBot.git
synced 2026-09-14 20:46:47 +08:00
Merge pull request #115 from ddggkkcc/pr/rag-crag-refusal
[Feature] Wiki 块级问答增加评估路由与诚实拒答(CRAG 式):知识库没有答案时不再硬答
This commit is contained in:
+136
-10
@@ -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)
|
||||
|
||||
@@ -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 provider;Exception 实例则抛出"""
|
||||
|
||||
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_grade:LLM 输出 → (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]
|
||||
|
||||
Reference in New Issue
Block a user