rag_agent.py 5.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159
  1. import os
  2. import chromadb
  3. from dotenv import load_dotenv
  4. from langchain_community.document_loaders import PyMuPDFLoader
  5. from langchain_community.embeddings import DashScopeEmbeddings
  6. from langchain_community.vectorstores import Chroma
  7. from langchain_core.output_parsers import StrOutputParser
  8. from langchain_core.prompts import ChatPromptTemplate
  9. from langchain_openai import ChatOpenAI
  10. from langchain_classic.chains.combine_documents import create_stuff_documents_chain
  11. from langchain_text_splitters import RecursiveCharacterTextSplitter
  12. # ── 配置 ──────────────────────────────────────────────────────────
  13. BASE_DIR = os.path.dirname(os.path.abspath(__file__))
  14. PDF_PATH = os.path.join(BASE_DIR, "docs", "浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf")
  15. CHROMA_DB_PATH = os.path.join(BASE_DIR, "chroma")
  16. COLLECTION_NAME = "shanghai_bank_policy"
  17. PROMPT_TEMPLATE = """
  18. 你是一个专业的知识库助手。请根据以下上下文回答问题。
  19. **规则:**
  20. - 只基于提供的上下文回答,不要编造
  21. - 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
  22. - 回答要简洁直接,引用原文时用引号
  23. **上下文:**
  24. {context}
  25. **问题:**
  26. {question}
  27. """
  28. def load_config() -> dict:
  29. """加载 .env 中的配置项"""
  30. load_dotenv()
  31. config = {
  32. "api_key": os.getenv("ALIYUN_API_KEY"),
  33. "base_url": os.getenv("ALIYUN_BASE_URL"),
  34. "chat_model": os.getenv("ALIYUN_CHAT_MODEL"),
  35. "embedding_model": os.getenv("ALIYUN_EMBEDDING_MODEL", "text-embedding-v3"),
  36. }
  37. missing = [k for k, v in config.items() if not v and k != "embedding_model"]
  38. if missing:
  39. raise EnvironmentError(f"缺少环境变量: {missing},请检查 .env 文件")
  40. return config
  41. def init_llm(config: dict) -> ChatOpenAI:
  42. return ChatOpenAI(
  43. base_url=config["base_url"],
  44. api_key=config["api_key"],
  45. model=config["chat_model"],
  46. temperature=0,
  47. timeout=10,
  48. max_retries=1,
  49. )
  50. def init_embeddings(config: dict) -> DashScopeEmbeddings:
  51. return DashScopeEmbeddings(
  52. model=config["embedding_model"],
  53. dashscope_api_key=config["api_key"],
  54. )
  55. def clean_pdf_text(text: str) -> str:
  56. """清洗 PDF 解析出的文本,去除常见噪声"""
  57. import re
  58. # 删除非中文字符之间的换行符
  59. text = re.sub(r'[^一](\n)[^一]',
  60. lambda m: m.group(0).replace('\n', ''), text)
  61. # 删除项目符号和多余空格
  62. text = text.replace('•', '').replace(' ', ' ')
  63. # 删除连续的换行符(保留一个)
  64. text = re.sub(r'\n{2,}', '\n', text)
  65. return text.strip()
  66. def build_or_load_vectorstore(embeddings: DashScopeEmbeddings) -> Chroma:
  67. """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
  68. chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
  69. existing_collections = [col.name for col in chroma_client.list_collections()]
  70. if COLLECTION_NAME in existing_collections:
  71. print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
  72. return Chroma(
  73. embedding_function=embeddings,
  74. collection_name=COLLECTION_NAME,
  75. client=chroma_client,
  76. )
  77. print(f"ℹ️ 未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
  78. loader = PyMuPDFLoader(PDF_PATH)
  79. pages = loader.load()
  80. # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
  81. clean_pages = [clean_pdf_text(page) for page in pages]
  82. text_splitter = RecursiveCharacterTextSplitter(
  83. chunk_size=500,
  84. chunk_overlap=50,
  85. separators=["\n\n", "\n", "。", ";", ",", " ", ""],
  86. )
  87. docs = text_splitter.split_documents(pages)
  88. vectorstore = Chroma.from_documents(
  89. embedding=embeddings,
  90. collection_name=COLLECTION_NAME,
  91. client=chroma_client,
  92. documents=docs,
  93. )
  94. print(f"✅ 建库完成,共 {len(docs)} 个分块")
  95. return vectorstore
  96. def ask(query: str, retriever, llm) -> str:
  97. """检索 + 生成回答"""
  98. relevant_docs = retriever.invoke(query)
  99. context = "\n\n---\n\n".join(d.page_content for d in relevant_docs)
  100. prompt = ChatPromptTemplate.from_template(PROMPT_TEMPLATE)
  101. chain = prompt | llm | StrOutputParser()
  102. return chain.invoke({"context": context, "question": query})
  103. def main():
  104. config = load_config()
  105. llm = init_llm(config)
  106. embeddings = init_embeddings(config)
  107. print("✅ 模型初始化完成")
  108. vectorstore = build_or_load_vectorstore(embeddings)
  109. retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
  110. query = "什么是工作质量考核标准?"
  111. answer = ask(query, retriever, llm)
  112. print("\n" + "=" * 60)
  113. print(f"问题:{query}")
  114. print("=" * 60)
  115. print(answer)
  116. if __name__ == "__main__":
  117. main()