架构对照表
| 层级 | 文件 | 职责 | 技术 | 核心导出 |
|---|
| L1 基础设施层 | llm.py | LLM/Embedding 统一入口,单例+重试+超时 | ChatOpenAI / OpenAIEmbeddings | get_llm() get_embeddings() |
| L2 数据检索层 | rag.py | 向量知识库 + RAG 检索 + 阈值过滤 | langchain-chroma + Retriever | rag_service.search() as_retriever() |
| L3 能力工具层 | tools.py | Agent 工具定义,规范入参 | schema @tool + Pydantic | ALL_TOOLS |
| L4 智能体编排层 | agent.py | 工作流 + 状态流转 + 异常兜底 + 会话持久化 | LangGraph StateGraph | agent_instance.chat() stream_chat() |
1.基础设施层
作用:封装底层大模型和向量模型的客户端实例,提供统一、可靠的调用接口
| 技术 | 作用 | 代码示例 |
|---|
| 单例模式 ( @lru_cache ) | 确保全局只有一个 ChatOpenAI 和 Embeddings 实例,避免重复初始化开销 | @lru_cache(maxsize=1) |
| 内置重试 ( max_retries ) | 自动处理网络抖动和 API 限流,无需上层写 try-catch | ChatOpenAI(max_retries=3) |
| 超时控制 ( timeout ) | 防止 LLM 调用无限挂起,保障接口响应时间 | ChatOpenAI(timeout=30) |
| 环境变量驱动 | 敏感信息(API Key)和配置参数通过环境变量注入,便于部署和切换 | os.getenv(“AI_LLM_MODEL”) |
"""
L1 基础设施层:底层模型客户端封装,统一 LLM/Embedding 入口
"""
import os
from functools import lru_cache
from langchain_openai import ChatOpenAI, OpenAIEmbeddings
@lru_cache(maxsize=1)
def get_llm() -> ChatOpenAI:
"""全局 LLM 单例(内置重试 + 超时)"""
return ChatOpenAI(
model=os.getenv("AI_LLM_MODEL", "qwen-plus"),
api_key=os.getenv("AI_LLM_API_KEY"),
base_url=os.getenv("AI_LLM_BASE_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1"),
temperature=float(os.getenv("AI_LLM_TEMPERATURE", "0.7")),
max_tokens=int(os.getenv("AI_LLM_MAX_TOKENS", "2000")),
timeout=int(os.getenv("AI_LLM_TIMEOUT", "30")),
max_retries=int(os.getenv("AI_LLM_MAX_RETRIES", "3")),
)
@lru_cache(maxsize=1)
def get_embeddings() -> OpenAIEmbeddings:
"""全局 Embedding 单例"""
return OpenAIEmbeddings(
model=os.getenv("AI_EMBEDDING_MODEL", "text-embedding-v3"),
api_key=os.getenv("AI_LLM_API_KEY"),
base_url=os.getenv("AI_LLM_BASE_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1"),
)
2.rag数据检索层
作用:管理向量知识库,为智能体提供文档召回能力
| 技术 | 作用 | 代码示例 |
|---|
| 向量数据库 ( Chroma ) | 轻量级嵌入式向量数据库,持久化存储文档及其向量表示 | Chroma(collection_name=“knowledge”, …) |
| 语义检索 ( similarity_search ) | 根据用户问题的语义相似度快速召回最相关的文档 | vectorstore.similarity_search_with_relevance_scores(query) |
| 阈值过滤 (Threshold) | 过滤掉相似度低于设定值(如 0.6)的文档,确保召回质量 | if score >= SIMILARITY_THRESHOLD |
| LangChain 集成 ( as_retriever ) | 暴露为标准的 LangChain Retriever 接口,可直接接入其他 Chain 或 Agent | vectorstore.as_retriever() |
| 延迟初始化 (Lazy Init) | 向量库在首次使用时才初始化,加快应用启动速度 | @property def vectorstore |
"""
L2 数据检索层:向量知识库 + RAG 检索逻辑,提供文档召回能力
支持:文档传入 → 自动分块 → 向量化 → 持久化存储 → 语义检索
"""
import os
import logging
from typing import Optional
from langchain_chroma import Chroma
from langchain_core.documents import Document
from langchain_text_splitters import RecursiveCharacterTextSplitter
from ai_module.llm import get_embeddings
logger = logging.getLogger(__name__)
SIMILARITY_THRESHOLD = float(os.getenv("AI_RAG_THRESHOLD", "0.6"))
SEARCH_TOP_K = int(os.getenv("AI_RAG_TOP_K", "5"))
CHUNK_SIZE = int(os.getenv("AI_RAG_CHUNK_SIZE", "500"))
CHUNK_OVERLAP = int(os.getenv("AI_RAG_CHUNK_OVERLAP", "50"))
class RAGService:
"""RAG 检索服务(支持文档传入 + 检索召回)"""
def __init__(self):
self._vectorstore: Optional[Chroma] = None
self._splitter: Optional[RecursiveCharacterTextSplitter] = None
@property
def vectorstore(self) -> Chroma:
"""延迟初始化向量库"""
if self._vectorstore is None:
self._vectorstore = Chroma(
collection_name="knowledge",
persist_directory=os.getenv("AI_CHROMA_DIR", "./chroma_db"),
embedding_function=get_embeddings(),
)
return self._vectorstore
@property
def splitter(self) -> RecursiveCharacterTextSplitter:
"""延迟初始化分块器"""
if self._splitter is None:
self._splitter = RecursiveCharacterTextSplitter(
chunk_size=CHUNK_SIZE,
chunk_overlap=CHUNK_OVERLAP,
separators=["\n\n", "\n", "。", "!", "?", ";", ",", " ", ""],
)
return self._splitter
def add_text(self, text: str, metadata: dict = None) -> int:
"""
传入纯文本,自动分块后写入向量库
Args:
text: 纯文本内容
metadata: 元数据(如 {"title": "xx", "source": "manual"})
Returns:
写入的文档块数量
"""
if not text or not text.strip():
return 0
chunks = self.splitter.split_text(text)
documents = [
Document(page_content=chunk, metadata=metadata or {})
for chunk in chunks
]
self.vectorstore.add_documents(documents)
logger.info(f"add_text: 写入 {len(documents)} 个文档块")
return len(documents)
def add_file(self, file_path: str, metadata: dict = None) -> int:
"""
传入文件路径,自动加载 + 分块 + 写入
支持 .txt / .md 文件
Args:
file_path: 文件路径
metadata: 元数据(自动补充 {"source": file_path})
Returns:
写入的文档块数量
"""
if not os.path.exists(file_path):
logger.error(f"文件不存在: {file_path}")
return 0
try:
with open(file_path, "r", encoding="utf-8") as f:
text = f.read()
except UnicodeDecodeError:
with open(file_path, "r", encoding="gbk", errors="ignore") as f:
text = f.read()
file_meta = {"source": file_path, "filename": os.path.basename(file_path)}
if metadata:
file_meta.update(metadata)
return self.add_text(text, file_meta)
def add_documents(self, documents: list) -> int:
"""
传入结构化文档列表,批量分块写入
Args:
documents: [{"content": str, "metadata": dict}, ...]
Returns:
写入的文档块总数
"""
total = 0
for doc in documents:
content = doc.get("content", "")
metadata = doc.get("metadata", {})
total += self.add_text(content, metadata)
logger.info(f"add_documents: 批量写入完成,共 {total} 个文档块")
return total
def search(self, query: str, top_k: int = None) -> list:
"""
向量检索 + 相似度阈值过滤
Args:
query: 查询文本
top_k: 返回数量,默认用配置值
Returns:
[{"content": str, "metadata": dict, "score": float}, ...]
"""
try:
raw = self.vectorstore.similarity_search_with_relevance_scores(
query, k=top_k or SEARCH_TOP_K
)
return [
{"content": doc.page_content, "metadata": doc.metadata, "score": score}
for doc, score in raw
if score >= SIMILARITY_THRESHOLD
]
except Exception as e:
logger.error(f"RAG 检索失败: {e}")
return []
def as_retriever(self, top_k: int = None):
"""暴露为 LangChain Retriever(可接入 LCEL chain)"""
return self.vectorstore.as_retriever(
search_kwargs={"k": top_k or SEARCH_TOP_K}
)
def clear(self) -> bool:
"""清空向量库所有文档"""
try:
collection = self.vectorstore._collection
all_ids = collection.get()["ids"]
if all_ids:
collection.delete(ids=all_ids)
logger.info("向量库已清空")
return True
except Exception as e:
logger.error(f"清空向量库失败: {e}")
return False
def get_stats(self) -> dict:
"""获取知识库统计信息"""
try:
count = self.vectorstore._collection.count()
return {
"status": "ok",
"document_count": count,
"chunk_size": CHUNK_SIZE,
"chunk_overlap": CHUNK_OVERLAP,
"similarity_threshold": SIMILARITY_THRESHOLD,
}
except Exception as e:
return {"status": "error", "document_count": 0, "error": str(e)}
rag_service = RAGService()
3.能力工具层
| 技术 | 作用 | 代码示例 |
|---|
| 工具装饰器 ( @tool ) | LangChain 提供的装饰器,自动将 Python 函数转换为 LLM 可识别的工具 | @tool(args_schema=KnowledgeQueryArgs) |
| Pydantic Schema | 使用 Pydantic 定义工具的输入参数结构,LLM 会严格遵循此结构传参 | class KnowledgeQueryArgs(BaseModel) |
| 工具描述 ( description ) | 工具的 Docstring 会传给 LLM,LLM 根据描述判断何时调用此工具 | “”“查询内部知识库…”“” |
| 异常吞噬 | 工具内部捕获所有异常,返回友好的错误提示给 LLM,防止 Agent 崩溃 | except Exception: return “查询失败” |
| 注册表模式 | 所有工具集中注册到 ALL_TOOLS 列表,便于 Agent 统一绑定 | ALL_TOOLS = [query_knowledge_base, …] |
"""
L3 能力工具层:Agent 可用函数工具,规范入参 schema
基于 @tool + Pydantic
"""
from langchain_core.tools import tool
from pydantic import BaseModel, Field
from typing import Optional
from ai_module.rag import rag_service
class KnowledgeQueryArgs(BaseModel):
query: str = Field(description="查询内容,用于检索知识库")
class UserDataQueryArgs(BaseModel):
user_id: Optional[int] = Field(default=None, description="用户ID,系统自动注入")
query: str = Field(description="查询内容", default="")
@tool(args_schema=KnowledgeQueryArgs)
def query_knowledge_base(query: str) -> str:
"""
查询内部知识库。
当用户询问产品信息、技术细节、使用规则等专业知识时使用此工具。
"""
results = rag_service.search(query)
if not results:
return "知识库中未找到相关信息。"
return "\n\n".join(
[f"【{r['metadata'].get('title', '')}】\n{r['content']}" for r in results]
)
@tool(args_schema=UserDataQueryArgs)
def query_user_data(user_id: int, query: str = "") -> str:
"""
查询用户个人数据(订单、预约、积分等)。
仅当用户明确询问自身相关数据时使用。
"""
if not user_id or user_id <= 0:
return "无法识别您的身份,请重新登录后再试。"
return f"用户 {user_id} 的数据查询结果(占位)。"
ALL_TOOLS = [query_knowledge_base, query_user_data]
4.智能体编排层
| 技术 | 作用 | 代码示例 |
|---|
| 状态图 ( StateGraph ) | 定义节点(Node)和有向边(Edge),构建可预测的工作流 | g = StateGraph(AgentState) |
| 状态持久化 ( MemorySaver ) | 保存每次对话的上下文历史,实现多轮对话记忆 | g.compile(checkpointer=MemorySaver()) |
| 最大迭代限制 ( recursion_limit ) | 防止 Agent 陷入无限循环(如工具反复调用失败) | config={“recursion_limit”: 25} |
| 异常兜底节点 (Fallback Node) | 任何节点报错都会最终汇聚到 Fallback 节点,返回友好提示 | g.add_node(“fallback”, self._fallback_node) |
| Function Calling ( bind_tools ) | 让 LLM 自主决定何时调用工具、调用哪个工具、传什么参数 | self.llm_with_tools = self.llm.bind_tools(ALL_TOOLS) |
| 流式输出 ( astream_events ) | 支持打字机效果的流式响应,提升用户体验 | async for event in self.graph.astream_events(…) |
| TypedDict 状态 | 用 TypedDict 定义流转数据结构,类型安全 | class AgentState(TypedDict) |
"""
L4 智能体编排层:工作流、分支、状态流转、节点调度
基于 LangGraph StateGraph
"""
import os
import time
import logging
from typing import Annotated, AsyncIterator, Optional, TypedDict
import operator
from langgraph.graph import StateGraph, END
from langgraph.checkpoint.memory import MemorySaver
from langchain_core.messages import HumanMessage, ToolMessage
from pydantic import BaseModel, Field
from ai_module.llm import get_llm
from ai_module.tools import ALL_TOOLS, query_user_data
logger = logging.getLogger(__name__)
MAX_ITERATIONS = int(os.getenv("AI_MAX_ITERATIONS", "25"))
INTENT_PROMPT = """判断用户问题的意图,只返回以下类别之一:
- NEED_TOOL: 需要查询数据
- CHAT: 通用闲聊
问题:{question}
只返回 NEED_TOOL 或 CHAT。"""
SYSTEM_PROMPT = "你是一个专业的 AI 助手,请用中文简洁准确地回答用户问题。"
TOOL_SYSTEM_PROMPT = """你是数据查询助手。根据用户问题选择合适的工具查询数据。
注意:用户ID由系统自动处理,你只需传递问题相关的查询内容。"""
FALLBACK_MSG = "抱歉,AI 服务暂时不可用,请稍后重试。"
class IntentResult(BaseModel):
intent: str = Field(description="NEED_TOOL 或 CHAT")
class AgentState(TypedDict):
messages: Annotated[list, operator.add]
user_id: Optional[int]
intent: str
final_answer: str
error: Optional[str]
TOOL_MAP = {t.name: t for t in ALL_TOOLS}
class UniversalAgent:
def __init__(self):
self.llm = get_llm()
self.llm_with_tools = self.llm.bind_tools(ALL_TOOLS)
self.llm_router = self.llm.with_structured_output(IntentResult)
self.checkpointer = MemorySaver()
self.graph = self._build_graph()
def _build_graph(self):
g = StateGraph(AgentState)
g.add_node("route", self._route_node)
g.add_node("tool_call", self._tool_call_node)
g.add_node("answer", self._answer_node)
g.add_node("fallback", self._fallback_node)
g.set_entry_point("route")
g.add_conditional_edges("route", self._route_by_intent,
{"need_tool": "tool_call", "chat": "answer"})
g.add_edge("tool_call", "answer")
g.add_edge("answer", END)
g.add_edge("fallback", END)
return g.compile(checkpointer=self.checkpointer)
def _route_node(self, state: AgentState) -> dict:
"""意图识别"""
question = state["messages"][-1].content
try:
result = self.llm_router.invoke(
[{"role": "user", "content": INTENT_PROMPT.format(question=question)}]
)
intent = result.intent.strip().upper()
if intent not in ("NEED_TOOL", "CHAT"):
intent = "CHAT"
except Exception as e:
logger.warning(f"意图识别失败,默认 CHAT: {e}")
intent = "CHAT"
logger.info(f"意图: {intent}")
return {"intent": intent}
def _tool_call_node(self, state: AgentState) -> dict:
"""
工具调用节点 — 核心:user_id 在这里被注入
"""
messages = state["messages"]
real_user_id = state.get("user_id")
try:
sys_msg = {"role": "system", "content": TOOL_SYSTEM_PROMPT}
resp = self.llm_with_tools.invoke([sys_msg] + messages)
new_messages = [resp]
if resp.tool_calls:
for tc in resp.tool_calls:
tool_name = tc["name"]
tool = TOOL_MAP.get(tool_name)
if not tool:
logger.warning(f"未知工具: {tool_name}")
continue
args = dict(tc["args"])
if tool is query_user_data:
args["user_id"] = real_user_id
logger.info(f"注入 user_id={real_user_id} 到 {tool_name}")
logger.info(f"执行: {tool_name} | args={args}")
try:
result = tool.invoke(args)
except Exception as e:
result = f"工具执行失败: {e}"
logger.error(f"工具 {tool_name} 异常: {e}")
new_messages.append(
ToolMessage(content=str(result), tool_call_id=tc["id"])
)
return {"messages": new_messages}
except Exception as e:
logger.error(f"工具调用节点异常: {e}", exc_info=True)
return {"error": str(e)}
def _answer_node(self, state: AgentState) -> dict:
"""答案生成"""
messages = [{"role": "system", "content": SYSTEM_PROMPT}] + state["messages"]
try:
resp = self.llm.invoke(messages)
return {"messages": [resp], "final_answer": resp.content or ""}
except Exception as e:
logger.error(f"答案生成失败: {e}", exc_info=True)
return {"error": str(e)}
def _fallback_node(self, state: AgentState) -> dict:
"""异常兜底"""
logger.warning(f"兜底: {state.get('error')}")
return {"final_answer": FALLBACK_MSG}
def _route_by_intent(self, state: AgentState) -> str:
return "need_tool" if state.get("intent") == "NEED_TOOL" else "chat"
def chat(self, session_id: str, question: str, user_id: int = None) -> dict:
"""
同步对话
Args:
session_id: 会话ID
question: 用户问题
user_id: 登录用户ID(可选,用于工具注入)
"""
start = time.time()
cfg = {
"configurable": {"thread_id": session_id},
"recursion_limit": MAX_ITERATIONS,
}
try:
result = self.graph.invoke(
{
"messages": [HumanMessage(content=question)],
"user_id": user_id,
},
config=cfg,
)
return {
"success": not result.get("error"),
"answer": result.get("final_answer", FALLBACK_MSG),
"session_id": session_id,
"response_time": round(time.time() - start, 2),
}
except RecursionError:
return {"success": False, "answer": "思考过于复杂,请简化问题。",
"session_id": session_id, "response_time": round(time.time() - start, 2)}
except Exception as e:
logger.error(f"对话异常: {e}", exc_info=True)
return {"success": False, "answer": FALLBACK_MSG,
"session_id": session_id, "response_time": round(time.time() - start, 2)}
async def stream_chat(self, session_id: str, question: str,
user_id: int = None) -> AsyncIterator[str]:
"""流式对话"""
cfg = {"configurable": {"thread_id": session_id},
"recursion_limit": MAX_ITERATIONS}
try:
async for event in self.graph.astream_events(
{"messages": [HumanMessage(content=question)],
"user_id": user_id},
config=cfg, version="v1",
):
if event["event"] == "on_chat_model_stream":
content = event["data"]["chunk"].content
if content:
yield content
except Exception as e:
yield f"\n[异常: {str(e)}]"
def clear_session(self, session_id: str) -> bool:
try:
self.checkpointer.delete_thread({"configurable": {"thread_id": session_id}})
return True
except Exception:
return False
agent_instance = UniversalAgent()
5.django适配
import json
import asyncio
from django.http import JsonResponse, StreamingHttpResponse
from rest_framework.views import APIView
from ai_module import agent_instance
class AIChatView(APIView):
def post(self, request):
session_id = request.data.get("session_id", "anonymous")
question = request.data.get("question", "").strip()
if not question:
return JsonResponse({"error": "问题不能为空"}, status=400)
user_id = self._get_user_id(request)
if request.data.get("stream"):
return self._stream(session_id, question, user_id)
result = agent_instance.chat(session_id, question, user_id=user_id)
return JsonResponse(result, status=200 if result["success"] else 500)
def _stream(self, session_id, question, user_id):
def gen():
loop = asyncio.new_event_loop()
try:
agen = agent_instance.stream_chat(session_id, question, user_id=user_id)
while True:
try:
token = loop.run_until_complete(agen.__anext__())
yield f"data: {json.dumps({'token': token})}\n\n"
except StopAsyncIteration:
break
yield f"data: {json.dumps({'done': True})}\n\n"
finally:
loop.close()
resp = StreamingHttpResponse(gen(), content_type="text/event-stream")
resp["Cache-Control"] = "no-cache"
return resp
@staticmethod
def _get_user_id(request):
"""
获取当前请求的用户 ID
- 已登录:返回 request.user.id
- 未登录:返回 None(匿名用户)
"""
if request.user.is_authenticated:
return request.user.id
return None
class AISessionView(APIView):
def delete(self, request, session_id):
return JsonResponse({"success": agent_instance.clear_session(session_id)})
6.fastapi适配
import json
from fastapi import APIRouter, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel
from typing import Optional
from ai_module import agent_instance
router = APIRouter(prefix="/api/ai", tags=["AI"])
class ChatReq(BaseModel):
session_id: Optional[str] = None
question: str
stream: bool = False
user_id: Optional[int] = None
@router.post("/chat")
async def chat(req: ChatReq, request: Request):
sid = req.session_id or "anonymous"
user_id = req.user_id or _get_user_id(request)
if req.stream:
async def gen():
async for t in agent_instance.stream_chat(sid, req.question, user_id=user_id):
yield f"data: {json.dumps({'token': t})}\n\n"
yield f"data: {json.dumps({'done': True})}\n\n"
return StreamingResponse(gen(), media_type="text/event-stream",
headers={"Cache-Control": "no-cache"})
return agent_instance.chat(sid, req.question, user_id=user_id)
@router.delete("/sessions/{session_id}")
async def clear_session(session_id: str):
return {"success": agent_instance.clear_session(session_id)}
def _get_user_id(request: Request):
"""
FastAPI 获取用户 ID(根据你的认证方式调整)
- JWT: 从 request.state.user 或中间件解析
- Session: 从 request.session 获取
- 匿名: 返回 None
"""
if hasattr(request.state, "user"):
return request.state.user.id
return None
所有评论(0)