26-7-31-1

This commit is contained in:
shuai
2026-07-31 12:27:41 +08:00
commit 9b8b8b57b3
142 changed files with 22408 additions and 0 deletions
+242
View File
@@ -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