"""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