fix(rag): route exercise titles to knowledge base
This commit is contained in:
@@ -8,7 +8,7 @@ import re
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from time import perf_counter
|
from time import perf_counter
|
||||||
|
|
||||||
from sqlalchemy import select
|
from sqlalchemy import or_, select
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from app.core.config import get_settings
|
from app.core.config import get_settings
|
||||||
@@ -40,7 +40,13 @@ MAX_LEXICAL_CANDIDATES = 12
|
|||||||
MAX_SELECTED_SECTIONS = 4
|
MAX_SELECTED_SECTIONS = 4
|
||||||
MAX_OVERVIEW_LEXICAL_CANDIDATES = 48
|
MAX_OVERVIEW_LEXICAL_CANDIDATES = 48
|
||||||
MAX_OVERVIEW_SELECTED_SECTIONS = 24
|
MAX_OVERVIEW_SELECTED_SECTIONS = 24
|
||||||
BUSINESS_MARKERS = {"课程", "大本营", "训练营", "老师", "卢慧", "功课", "学员", "课堂", "练习", "觉察", "内在"}
|
BUSINESS_MARKERS = {
|
||||||
|
"课程", "大本营", "训练营", "老师", "卢慧", "功课", "作业", "学员",
|
||||||
|
"课堂", "练习", "静心", "觉察", "内在",
|
||||||
|
}
|
||||||
|
_ROUTING_STOP_TERMS = {
|
||||||
|
"课程", "练习", "功课", "作业", "静心", "具体", "详细", "内容", "怎么", "操作", "什么", "方法",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -101,6 +107,14 @@ class KnowledgeAgentService:
|
|||||||
)
|
)
|
||||||
trace.append(rewrite_trace)
|
trace.append(rewrite_trace)
|
||||||
need_knowledge, selected_ids, reason = cls._decide(retrieval_question, catalog)
|
need_knowledge, selected_ids, reason = cls._decide(retrieval_question, catalog)
|
||||||
|
terms = cls._query_terms(retrieval_question)
|
||||||
|
title_routes = cls._route_knowledge_by_titles(db, terms, catalog)
|
||||||
|
if title_routes:
|
||||||
|
need_knowledge = True
|
||||||
|
selected_ids = list(dict.fromkeys(
|
||||||
|
[item["knowledgeId"] for item in title_routes] + selected_ids
|
||||||
|
))[:4]
|
||||||
|
reason = "章节标题直接命中知识库"
|
||||||
trace.append(cls._trace(
|
trace.append(cls._trace(
|
||||||
"agent_decision",
|
"agent_decision",
|
||||||
len(trace) + 1,
|
len(trace) + 1,
|
||||||
@@ -108,10 +122,16 @@ class KnowledgeAgentService:
|
|||||||
{"needKnowledge": need_knowledge, "selectedKnowledgeIds": selected_ids, "reason": reason},
|
{"needKnowledge": need_knowledge, "selectedKnowledgeIds": selected_ids, "reason": reason},
|
||||||
started,
|
started,
|
||||||
))
|
))
|
||||||
|
trace.append(cls._trace(
|
||||||
|
"route_knowledge_by_titles",
|
||||||
|
len(trace) + 1,
|
||||||
|
{"queryTerms": terms},
|
||||||
|
{"count": len(title_routes), "matches": title_routes},
|
||||||
|
started,
|
||||||
|
))
|
||||||
chunks: list[RetrievedChunk] = []
|
chunks: list[RetrievedChunk] = []
|
||||||
try:
|
try:
|
||||||
if need_knowledge and selected_ids:
|
if need_knowledge and selected_ids:
|
||||||
terms = cls._query_terms(retrieval_question)
|
|
||||||
practice_overview = cls._is_practice_overview(retrieval_question)
|
practice_overview = cls._is_practice_overview(retrieval_question)
|
||||||
candidate_limit = cls._candidate_limit(retrieval_question)
|
candidate_limit = cls._candidate_limit(retrieval_question)
|
||||||
candidates = cls.search_knowledge(
|
candidates = cls.search_knowledge(
|
||||||
@@ -557,6 +577,41 @@ class KnowledgeAgentService:
|
|||||||
return True, selected, "问题与知识目录主题相关"
|
return True, selected, "问题与知识目录主题相关"
|
||||||
return False, [], "问题与当前知识目录无关,无需调用知识库"
|
return False, [], "问题与当前知识目录无关,无需调用知识库"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _route_knowledge_by_titles(db: Session, terms: list[str], catalog: list[dict]) -> list[dict]:
|
||||||
|
"""Map a named exercise/chapter back to its knowledge base before retrieval."""
|
||||||
|
version_ids = [int(item["versionId"]) for item in catalog if item.get("versionId")]
|
||||||
|
if not version_ids:
|
||||||
|
return []
|
||||||
|
distinctive_terms = list(dict.fromkeys(
|
||||||
|
term.lower() for term in terms
|
||||||
|
if len(term) >= 2 and term not in _ROUTING_STOP_TERMS
|
||||||
|
))
|
||||||
|
if not distinctive_terms:
|
||||||
|
return []
|
||||||
|
distinctive_terms = sorted(distinctive_terms, key=len, reverse=True)[:12]
|
||||||
|
|
||||||
|
rows = db.execute(
|
||||||
|
select(KnowledgeChunk.knowledge_id, KnowledgeChunk.title)
|
||||||
|
.where(
|
||||||
|
KnowledgeChunk.version_id.in_(version_ids),
|
||||||
|
or_(*(KnowledgeChunk.title.contains(term) for term in distinctive_terms)),
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
best_by_knowledge: dict[int, tuple[float, str]] = {}
|
||||||
|
for knowledge_id, title in rows:
|
||||||
|
normalized_title = title.lower()
|
||||||
|
score = sum(min(len(term), 4) ** 2 for term in distinctive_terms if term in normalized_title)
|
||||||
|
current = best_by_knowledge.get(knowledge_id)
|
||||||
|
if score > 0 and (current is None or score > current[0]):
|
||||||
|
best_by_knowledge[knowledge_id] = (float(score), title)
|
||||||
|
|
||||||
|
ranked = sorted(best_by_knowledge.items(), key=lambda item: (item[1][0], item[0]), reverse=True)
|
||||||
|
return [
|
||||||
|
{"knowledgeId": knowledge_id, "matchedTitle": title, "score": score}
|
||||||
|
for knowledge_id, (score, title) in ranked[:4]
|
||||||
|
]
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _query_terms(question: str) -> list[str]:
|
def _query_terms(question: str) -> list[str]:
|
||||||
base = extract_terms(question)
|
base = extract_terms(question)
|
||||||
|
|||||||
@@ -341,11 +341,14 @@ def test_same_level_practice_headings_are_read_as_one_semantic_unit():
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
db.commit()
|
db.commit()
|
||||||
|
for knowledge_id in range(2, 6):
|
||||||
|
_add_published_knowledge(db, knowledge_id=knowledge_id, name=f"其他课程{knowledge_id}")
|
||||||
|
|
||||||
content = KnowledgeAgentService._read_complete_section(db, detail)
|
content = KnowledgeAgentService._read_complete_section(db, detail)
|
||||||
result = asyncio.run(
|
result = asyncio.run(
|
||||||
KnowledgeAgentService.build_result(db, question="金色光欧姆静心具体怎么练习?")
|
KnowledgeAgentService.build_result(db, question="金色光欧姆是什么?")
|
||||||
)
|
)
|
||||||
|
title_route = next(item for item in result.tool_trace if item["tool"] == "route_knowledge_by_titles")
|
||||||
|
|
||||||
assert "金色光欧姆静心" in content
|
assert "金色光欧姆静心" in content
|
||||||
assert "保持自然呼吸" in content
|
assert "保持自然呼吸" in content
|
||||||
@@ -353,6 +356,7 @@ def test_same_level_practice_headings_are_read_as_one_semantic_unit():
|
|||||||
assert "下一项练习内容" not in content
|
assert "下一项练习内容" not in content
|
||||||
assert result.chunks
|
assert result.chunks
|
||||||
assert "保持自然呼吸" in result.chunks[0].content
|
assert "保持自然呼吸" in result.chunks[0].content
|
||||||
|
assert title_route["matches"][0]["knowledgeId"] == knowledge.id
|
||||||
|
|
||||||
|
|
||||||
def test_parser_merges_practice_title_and_same_level_detail_headings():
|
def test_parser_merges_practice_title_and_same_level_detail_headings():
|
||||||
|
|||||||
Reference in New Issue
Block a user