浅谈LangChain的用法
前言:
简单聊一下使用LangChain的一些心得,个人感觉就是操作大模型调用自己的向量数据库,然后触发自定义工具包的一个生产框架,主要包括RAG(向量数据库)和Agent(工具包)两部分核心组成。
RAG一般来源于企业生产中的文档资料,需要转换成大模型识别的向量数据库,一般要先把Word\Excel\Pdf这种文件资料,先进行整理,变成向量数据库:
import json
from langchain_core.documents import Document
from langchain_community.vectorstores import Chroma
from langchain_huggingface import HuggingFaceEmbeddings
import os
os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
def build_vector_store():
with open("faq.json", "r", encoding="utf-8") as f:
faq_list = json.load(f)
documents = []
for item in faq_list:
page_content = f"问题:{item['question']}\n客服答:{item['answer']}"
metadata = {
"id": item["id"],
"category": item["category"]
}
doc = Document(page_content=page_content, metadata=metadata)
documents.append(doc)
embedding_model = HuggingFaceEmbeddings(
model_name="BAAI/bge-small-zh-v1.5"
)
vectorstore = Chroma.from_documents(
documents=documents,
embedding=embedding_model,
persist_directory="./chroma_db"
)
print("Done")
if __name__ == "__main__":
build_vector_store()经过这个操作后,我们拿到了chroma_db,然后就可以写向量工具类了:
from typing import List
from langchain_chroma import Chroma
from langchain_core.tools import tool
from langchain_huggingface import HuggingFaceEmbeddings
from app.core.reranker import LocalReranker
embedding_model = HuggingFaceEmbeddings(
model_name="./models/BAAI--bge-small-zh-v1.5"
)
vectorstore = Chroma(
persist_directory="./chroma_db",
embedding_function=embedding_model
)
retriever = vectorstore.as_retriever(search_kwargs={"k": 10})
reranker = LocalReranker(model_name="./models/BAAI--bge-small-zh-v1.5")
def format_docs(docs: List) -> str:
formatted = []
for i, doc in enumerate(docs, 1):
score = doc.metadata.get("rerank_score", 0.0)
formatted.append(f"[参考资料 {i}] (置信度得分: {score:.4f}):\n{doc.page_content}")
return "\n\n".join(formatted)
@tool
async def search_knowledge_base(query: str) -> str:
raw_docs = await retriever.ainvoke(query)
if not raw_docs:
return "知识库中未检索到相关内容。"
reranked_docs = reranker.rerank(query, raw_docs, top_n=2)
if reranked_docs[0].metadata["rerank_score"] < -2.0: # bge-reranker 得分区间
return "知识库检索结果置信度过低,未匹配到有效文档。"
return format_docs(reranked_docs)准备好了RAG和Agent就可以接入框架了:
接口层:
发送请求并获取流式输出结果:
import json
from fastapi import APIRouter, HTTPException, status
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from app.services.agent_service import AgentService
# 定义路由组
router = APIRouter(prefix="/agent", tags=["Agent 智能助手"])
# 初始化单例 Agent 服务
agent_service = AgentService(redis_url="redis://localhost:6379/0")
# 显式校验 message 和 session_id
class AgentRequest(BaseModel):
message: str = Field(..., min_length=1, description="用户发送的问题文本", example="今天天气怎么样?")
session_id: str = Field(..., min_length=1, description="唯一会话ID,用于隔离与持久化历史记忆", example="sess_123456")
@router.post("/chat/stream", summary="Agent SSE 流式对话接口")
async def agent_chat_stream_endpoint(request: AgentRequest):
try:
return StreamingResponse(
agent_service.get_stream_response(request.message, request.session_id),
media_type="text/event-stream"
)
except Exception as e:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Agent 流式服务内部异常: {str(e)}"
)
@router.post("/chat", summary="Agent 普通同步对话接口")
async def agent_chat_sync_endpoint(request: AgentRequest):
try:
full_content = ""
# 复用异步流生成器,在内部拼接所有 chunk 结果
async for chunk in agent_service.get_stream_response(request.message, request.session_id):
if chunk.startswith("data: "):
data_str = chunk.replace("data: ", "").strip()
if data_str and data_str != "[DONE]":
try:
payload = json.loads(data_str)
if "content" in payload:
full_content += payload["content"]
except json.JSONDecodeError:
continue
return {"response": full_content, "session_id": request.session_id}
except Exception as e:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Agent 同步服务内部异常: {str(e)}"
)服务层:
import json
from typing import AsyncGenerator, List
from langchain_community.chat_message_histories import RedisChatMessageHistory
from langchain_core.messages import SystemMessage, HumanMessage, ToolMessage, AIMessage, BaseMessage
from app.core.llm import get_llm
from app.tools import ALL_TOOLS, TOOLS_MAP
class AgentService:
def __init__(self, redis_url: str = "redis://localhost:6379/0"):
self.llm = get_llm()
self.tools_map = TOOLS_MAP
self.redis_url = redis_url
self.llm_with_tools = self.llm.bind_tools(ALL_TOOLS)
def _get_history(self, session_id: str) -> RedisChatMessageHistory:
return RedisChatMessageHistory(
session_id=session_id,
url=self.redis_url,
ttl=86400
)
async def get_stream_response(self, question: str, session_id: str) -> AsyncGenerator[str, None]:
# 1.Redis
history_store = self._get_history(session_id)
# 2.Prompt
system_prompt = SystemMessage(
content=(
"你是一个严谨且高效的智能助手。\n"
"规则要求:\n"
"1. 如果调用的工具返回了错误或‘未收录’等提示,必须如实告知用户无法提供该信息,绝对严禁凭空捏造数据!\n"
"2. 严格基于工具返回的真实结果进行回答。"
)
)
saved_messages: List[BaseMessage] = history_store.messages
current_human_msg = HumanMessage(content=question)
messages = [system_prompt] + saved_messages + [current_human_msg]
new_messages_to_save: List[BaseMessage] = [current_human_msg]
first_response = await self.llm_with_tools.ainvoke(messages)
messages.append(first_response)
new_messages_to_save.append(first_response)
# 触发Tool
if first_response.tool_calls:
for tool_call in first_response.tool_calls:
tool_name = tool_call["name"]
tool_args = tool_call["args"]
tool_func = self.tools_map.get(tool_name)
if tool_func:
# 使用异步,消除阻塞
tool_result = await tool_func.ainvoke(tool_args)
print(f"[Agent Log] 异步执行工具 [{tool_name}],输入: {tool_args},输出: {tool_result}")
tool_msg = ToolMessage(
content=str(tool_result),
tool_call_id=tool_call["id"]
)
messages.append(tool_msg)
new_messages_to_save.append(tool_msg)
full_ai_content = ""
async for chunk in self.llm_with_tools.astream(messages):
if chunk.content:
full_ai_content += chunk.content
payload = json.dumps({"content": chunk.content}, ensure_ascii=False)
yield f"data: {payload}\n\n"
if full_ai_content:
new_messages_to_save.append(AIMessage(content=full_ai_content))
else:
# 普通场景
if first_response.content:
payload = json.dumps({"content": first_response.content}, ensure_ascii=False)
yield f"data: {payload}\n\n"
# 持久化至 Redis
history_store.add_messages(new_messages_to_save)Reranker 工具类
from typing import List
from langchain_core.documents import Document
from sentence_transformers import CrossEncoder
class LocalReranker:
def __init__(self, model_name: str = "BAAI/bge-reranker-base"):
# 加载模型
self.model = CrossEncoder(model_name)
def rerank(self, query: str, docs: List[Document], top_n: int = 2) -> List[Document]:
if not docs:
return []
pairs = [[query, doc.page_content] for doc in docs]
scores = self.model.predict(pairs)
for doc, score in zip(docs, scores):
doc.metadata["rerank_score"] = float(score)
sorted_docs = sorted(docs, key=lambda x: x.metadata["rerank_score"], reverse=True)
return sorted_docs[:top_n]工具层:
from app.tools.rag_tool import search_knowledge_base
ALL_TOOLS = [
search_knowledge_base,
]
TOOLS_MAP = {t.name: t for t in ALL_TOOLS}这里调用工具类,这里直接调用就行。
许可协议:
CC BY-NC 4.0
本文同步发布于个人博客 Neo·元,转载请注明出处。