26-7-31-1
This commit is contained in:
@@ -0,0 +1,242 @@
|
||||
"""LangChain 风格 RAG:文档切分 + 关键词检索 + 可选 LLM 生成。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from app.config import Settings, get_settings
|
||||
from app.schemas.models import SourceRef
|
||||
from app.services.llm_client import get_llm
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_LC_OK = False
|
||||
try:
|
||||
from langchain_core.documents import Document
|
||||
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||
|
||||
_LC_OK = True
|
||||
except Exception: # pragma: no cover
|
||||
Document = None # type: ignore
|
||||
RecursiveCharacterTextSplitter = None # type: ignore
|
||||
logger.info("LangChain 未完全安装,将使用内置简易检索")
|
||||
|
||||
|
||||
def langchain_available() -> bool:
|
||||
return _LC_OK
|
||||
|
||||
|
||||
@dataclass
|
||||
class Chunk:
|
||||
title: str
|
||||
category: str
|
||||
content: str
|
||||
source: str
|
||||
|
||||
|
||||
class RagPipeline:
|
||||
def __init__(self, settings: Settings | None = None):
|
||||
self.settings = settings or get_settings()
|
||||
self.chunks: list[Chunk] = []
|
||||
self._load_knowledge()
|
||||
|
||||
def _load_knowledge(self) -> None:
|
||||
path = self.settings.knowledge_path
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
files = sorted(list(path.glob("*.md")) + list(path.glob("*.txt")))
|
||||
raw_docs: list[tuple[str, str, str]] = []
|
||||
for f in files:
|
||||
try:
|
||||
text = f.read_text(encoding="utf-8")
|
||||
except Exception:
|
||||
continue
|
||||
title, category, body = self._parse_doc(f.stem, text)
|
||||
raw_docs.append((title, category, body))
|
||||
|
||||
if not raw_docs:
|
||||
logger.warning("知识库目录为空: %s", path)
|
||||
self.chunks = []
|
||||
return
|
||||
|
||||
if _LC_OK:
|
||||
splitter = RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=80)
|
||||
for title, category, body in raw_docs:
|
||||
docs = splitter.split_documents(
|
||||
[Document(page_content=body, metadata={"title": title, "category": category})]
|
||||
)
|
||||
for d in docs:
|
||||
self.chunks.append(
|
||||
Chunk(
|
||||
title=title,
|
||||
category=category,
|
||||
content=d.page_content,
|
||||
source=title,
|
||||
)
|
||||
)
|
||||
else:
|
||||
for title, category, body in raw_docs:
|
||||
for part in self._simple_split(body, 500):
|
||||
self.chunks.append(
|
||||
Chunk(title=title, category=category, content=part, source=title)
|
||||
)
|
||||
|
||||
logger.info("知识库已加载 %d 个文档片段", len(self.chunks))
|
||||
|
||||
@staticmethod
|
||||
def _parse_doc(stem: str, text: str) -> tuple[str, str, str]:
|
||||
title = stem
|
||||
category = "临床指南"
|
||||
body = text.strip()
|
||||
lines = body.splitlines()
|
||||
if lines and lines[0].startswith("#"):
|
||||
title = lines[0].lstrip("#").strip() or stem
|
||||
body = "\n".join(lines[1:]).strip()
|
||||
m = re.search(r"category:\s*(.+)", body, re.I)
|
||||
if m:
|
||||
category = m.group(1).strip()
|
||||
return title, category, body
|
||||
|
||||
@staticmethod
|
||||
def _simple_split(text: str, size: int) -> list[str]:
|
||||
if len(text) <= size:
|
||||
return [text]
|
||||
parts: list[str] = []
|
||||
i = 0
|
||||
while i < len(text):
|
||||
parts.append(text[i : i + size])
|
||||
i += max(1, size - 50)
|
||||
return parts
|
||||
|
||||
def reload(self) -> int:
|
||||
self.chunks = []
|
||||
self._load_knowledge()
|
||||
return len(self.chunks)
|
||||
|
||||
def doc_count(self) -> int:
|
||||
return len(self.chunks)
|
||||
|
||||
def retrieve(self, query: str, top_k: int = 4) -> list[SourceRef]:
|
||||
if not query or not self.chunks:
|
||||
return []
|
||||
tokens = self._tokenize(query)
|
||||
scored: list[tuple[float, Chunk]] = []
|
||||
q = query.lower()
|
||||
for ch in self.chunks:
|
||||
score = self._score(ch, tokens, q)
|
||||
if score > 0:
|
||||
scored.append((score, ch))
|
||||
scored.sort(key=lambda x: x[0], reverse=True)
|
||||
results: list[SourceRef] = []
|
||||
seen: set[tuple[str, str]] = set()
|
||||
for score, ch in scored:
|
||||
key = (ch.title, ch.content[:80])
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
results.append(
|
||||
SourceRef(
|
||||
title=ch.title,
|
||||
category=ch.category,
|
||||
snippet=ch.content[:220].replace("\n", " "),
|
||||
score=round(score, 2),
|
||||
)
|
||||
)
|
||||
if len(results) >= top_k:
|
||||
break
|
||||
return results
|
||||
|
||||
def build_context(self, sources: list[SourceRef]) -> str:
|
||||
if not sources:
|
||||
return ""
|
||||
parts = []
|
||||
for i, s in enumerate(sources, 1):
|
||||
parts.append(f"【资料{i}】{s.title}\n{s.snippet}")
|
||||
return "\n\n".join(parts)
|
||||
|
||||
def query(self, query: str, top_k: int = 4, extra_context: str = "") -> dict[str, Any]:
|
||||
sources = self.retrieve(query, top_k=top_k)
|
||||
context = self.build_context(sources)
|
||||
if extra_context:
|
||||
context = (extra_context.strip() + "\n\n" + context).strip()
|
||||
|
||||
llm = get_llm()
|
||||
if llm.enabled and (context or query):
|
||||
try:
|
||||
answer = llm.chat(
|
||||
[
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
"你是医院临床辅助决策助手。请仅依据给定资料与问题作答,"
|
||||
"语言专业简洁,并提醒需医师审核。不要编造未提供的检查数据。"
|
||||
),
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"参考资料:\n{context or '(无)'}\n\n问题:{query}",
|
||||
},
|
||||
]
|
||||
)
|
||||
return {"answer": answer, "sources": sources, "engine": "langchain-rag+llm"}
|
||||
except Exception as e:
|
||||
logger.warning("RAG LLM 失败: %s", e)
|
||||
|
||||
if sources:
|
||||
answer = (
|
||||
f"基于本地知识库检索(LangChain 文档切分 + 关键词排序),与「{query}」相关的要点:\n"
|
||||
+ "\n".join(f"- {s.title}:{s.snippet[:120]}" for s in sources)
|
||||
+ "\n\n(未配置大模型 API Key 或调用失败时展示检索摘要,仅供参考。)"
|
||||
)
|
||||
return {"answer": answer, "sources": sources, "engine": "langchain-rag-local"}
|
||||
|
||||
return {
|
||||
"answer": f"知识库中未检索到与「{query}」高度相关的条目,建议补充临床指南或完善病历描述。",
|
||||
"sources": [],
|
||||
"engine": "langchain-rag-empty",
|
||||
}
|
||||
|
||||
def ingest_text(self, title: str, content: str, category: str = "自定义") -> None:
|
||||
path = self.settings.knowledge_path
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
safe = re.sub(r"[^\w\u4e00-\u9fff\-]+", "_", title)[:40] or "doc"
|
||||
file_path = path / f"{safe}.md"
|
||||
file_path.write_text(
|
||||
f"# {title}\n\ncategory: {category}\n\n{content}\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
self.reload()
|
||||
|
||||
@staticmethod
|
||||
def _tokenize(text: str) -> list[str]:
|
||||
parts = re.split(r"[\s,,。.!!??;;::、/\\|_\-—()()\[\]{}]+", text.lower())
|
||||
return [p for p in parts if len(p) >= 2]
|
||||
|
||||
@staticmethod
|
||||
def _score(ch: Chunk, tokens: list[str], q: str) -> float:
|
||||
title = ch.title.lower()
|
||||
content = ch.content.lower()
|
||||
score = 0.0
|
||||
if q and q in title:
|
||||
score += 30
|
||||
if q and q in content:
|
||||
score += 15
|
||||
for t in tokens:
|
||||
if t in title:
|
||||
score += 8
|
||||
if t in content:
|
||||
score += 3
|
||||
if t in ch.category.lower():
|
||||
score += 2
|
||||
return score
|
||||
|
||||
|
||||
_rag: RagPipeline | None = None
|
||||
|
||||
|
||||
def get_rag() -> RagPipeline:
|
||||
global _rag
|
||||
if _rag is None:
|
||||
_rag = RagPipeline()
|
||||
return _rag
|
||||
Reference in New Issue
Block a user