Files
2026-07-31 12:27:41 +08:00

243 lines
8.2 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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