架构对照表

层级文件职责技术核心导出
L1 基础设施层llm.pyLLM/Embedding 统一入口,单例+重试+超时ChatOpenAI / OpenAIEmbeddingsget_llm() get_embeddings()
L2 数据检索层rag.py向量知识库 + RAG 检索 + 阈值过滤langchain-chroma + Retrieverrag_service.search() as_retriever()
L3 能力工具层tools.pyAgent 工具定义,规范入参schema @tool + PydanticALL_TOOLS
L4 智能体编排层agent.py工作流 + 状态流转 + 异常兜底 + 会话持久化LangGraph StateGraphagent_instance.chat() stream_chat()

1.基础设施层

作用:封装底层大模型和向量模型的客户端实例,提供统一、可靠的调用接口

技术作用代码示例
单例模式 ( @lru_cache )确保全局只有一个 ChatOpenAI 和 Embeddings 实例,避免重复初始化开销@lru_cache(maxsize=1)
内置重试 ( max_retries )自动处理网络抖动和 API 限流,无需上层写 try-catchChatOpenAI(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 或 Agentvectorstore.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:
            # 尝试 GBK 编码
            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


# ========== 工具入参 Schema ==========
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:
    """
    查询用户个人数据(订单、预约、积分等)。
    仅当用户明确询问自身相关数据时使用。
    """
    # 防御:user_id 必须有效
    if not user_id or user_id <= 0:
        return "无法识别您的身份,请重新登录后再试。"
    
    # TODO: 对接业务服务
    # from apps.order.models import Order
    # orders = Order.objects.filter(user_id=user_id).order_by('-created_at')[:5]
    
    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"))

# ========== Prompt ==========
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}


# ========== Agent ==========
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")  # ✅ 从 state 取真实用户 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"])

                    # ✅===== 关键:user_id 注入 =====#
                    # 只有当工具是 query_user_data 时,才注入 user_id
                    if tool is query_user_data:
                        args["user_id"] = real_user_id  # ✅ 强制覆盖,不信任 LLM
                        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

    # 从 request 获取用户 ID
    @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适配

# fastapi_router.py
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):  #  注入 request 用于获取用户
    sid = req.session_id or "anonymous"
    
    # 优先用请求体的 user_id,否则从认证信息获取
    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)}


# FastAPI 获取用户 ID 的辅助函数
def _get_user_id(request: Request):
    """
    FastAPI 获取用户 ID(根据你的认证方式调整)
    - JWT: 从 request.state.user 或中间件解析
    - Session: 从 request.session 获取
    - 匿名: 返回 None
    """
    # 方式 1: 从 JWT 中间件解析后注入
    if hasattr(request.state, "user"):
        return request.state.user.id
    # 方式 2: 从请求头解析 Authorization
    # 方式 3: 默认匿名
    return None
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐