欢迎光临
我们一直在努力

手写 Mini Dify:从零构建可视化 AI 工作流引擎

一、前言

1.1 为什么是 Dify?

Dify 是当前最受欢迎的开源 LLM 应用开发平台之一,它通过可视化的方式让开发者能够快速搭建基于大语言模型的 AI 应用。从简单的聊天机器人到复杂的 RAG 检索增强生成系统,再到多步骤的 Agent 工作流,Dify 提供了一套完整的工具链。

但你是否想过,Dify 的底层是如何工作的?当一个拖拽节点在你面前流畅地串联时,背后发生了什么?今天,我们就从零开始,亲手构建一个 Mini Dify——一个简化版的可视化 AI 工作流引擎。

1.2 我们将会学到什么

通过本文,你将掌握:

  • 工作流引擎的核心架构:节点、边、执行器的数据模型与生命周期
  • 可视化拖拽的底层原理:如何用纯前端实现节点编辑器
  • LLM 调用的抽象层设计:支持多模型(OpenAI、DeepSeek 等)的统一接口
  • RAG 检索流水线的实现:文档解析、向量化、检索的完整链路
  • Agent 工具的插件机制:函数调用的注册与分发
  • 前后端通信与任务调度:WebSocket 实时推送工作流执行状态

更重要的是,你将不再把 AI 应用平台当作黑盒——理解原理之后,你甚至可以定制自己的企业级 AI 工作流平台。

1.3 项目结构概览

我们构建的 Mini Dify 将包含以下核心模块:

mini-dify/
├── backend/
│ ├── core/ # 工作流引擎核心
│ │ ├── node.py # 节点定义
│ │ ├── workflow.py # 工作流编排
│ │ ├── executor.py # 执行器
│ │ └── types.py # 数据类型
│ ├── llm/ # LLM 抽象层
│ │ ├── base.py # 基类接口
│ │ ├── openai.py # OpenAI 实现
│ │ └── deepseek.py # DeepSeek 实现
│ ├── rag/ # RAG 模块
│ │ ├── document.py # 文档处理
│ │ ├── embedding.py# 向量化
│ │ └── retriever.py# 检索器
│ ├── agent/ # Agent 模块
│ │ ├── tool.py # 工具定义
│ │ └── agent.py # Agent 逻辑
│ └── server.py # FastAPI 服务
├── frontend/
│ ├── components/ # Vue 组件
│ ├── stores/ # 状态管理
│ └── App.vue # 主入口
└── README.md


二、工作流引擎核心设计

工作流引擎是整个系统的中枢神经系统。它负责解析用户构建的流程图,按依赖关系和拓扑顺序执行各个节点,并在节点之间传递数据。

2.1 节点数据模型

每个工作流节点都是一个处理单元。我们将其抽象为以下核心类:

# backend/core/node.py
from enum import Enum
from typing import Dict, Any, List, Optional
from dataclasses import dataclass, field
from abc import ABC, abstractmethod

class NodeType(Enum):
"""节点类型枚举"""
START = "start" # 起始节点
LLM = "llm" # LLM 调用节点
PROMPT = "prompt_template" # 提示词模板节点
KNOWLEDGE_RETRIEVAL = "knowledge_retrieval" # 知识检索节点
CODE = "code_execution" # 代码执行节点
HTTP_REQUEST = "http_request" # HTTP 请求节点
CONDITION = "condition" # 条件分支节点
ITERATION = "iteration" # 迭代节点
TOOL = "tool_call" # 工具调用节点
VARIABLE_AGGREGATOR = "variable_aggregator" # 变量聚合节点
END = "end" # 结束节点

@dataclass
class NodeConfig:
"""节点配置"""
model: Optional[str] = None # 模型名称(LLM 节点使用)
prompt: Optional[str] = None # 提示词模板
temperature: float = 0.7 # 温度参数
max_tokens: int = 2048 # 最大 token 数
variables: Dict[str, str] = field(default_factory=dict) # 变量映射
code: Optional[str] = None # 代码(代码执行节点使用)
url: Optional[str] = None # URL(HTTP 请求节点使用)
condition: Optional[str] = None # 条件表达式
tool_name: Optional[str] = None # 工具名称(工具调用节点使用)
retriever_config: Optional[Dict] = None # 检索器配置

@dataclass
class Node:
"""工作流节点"""
id: str # 节点唯一 ID
type: NodeType # 节点类型
title: str # 节点标题
config: NodeConfig = field(default_factory=NodeConfig) # 节点配置
inputs: Dict[str, Any] = field(default_factory=dict) # 输入数据
outputs: Dict[str, Any] = field(default_factory=dict) # 输出数据

2.2 边与工作流数据模型

边(Edge)定义了节点之间的连接关系和数据传递方向。

# backend/core/workflow.py
from dataclasses import dataclass, field
from typing import Dict, List, Optional
from .node import Node, NodeType

@dataclass
class Edge:
"""工作流边"""
id: str # 边唯一 ID
source_id: str # 源节点 ID
target_id: str # 目标节点 ID
source_handle: str = "output" # 源节点输出端口
target_handle: str = "input" # 目标节点输入端口
condition: Optional[str] = None # 条件表达式(条件边使用)

@dataclass
class Workflow:
"""工作流定义"""
id: str
title: str
description: str = ""
nodes: Dict[str, Node] = field(default_factory=dict) # 节点字典
edges: List[Edge] = field(default_factory=list) # 边列表
entry_node_id: Optional[str] = None # 入口节点 ID
variables: Dict[str, Any] = field(default_factory=dict) # 全局变量

2.3 拓扑排序与执行顺序

工作流执行的关键在于确定节点的执行顺序。由于节点之间存在依赖关系(从源节点到目标节点),我们需要使用拓扑排序来保证每个节点在其所有前置节点执行完毕后才执行。

# backend/core/executor.py
from collections import deque
from typing import Dict, List, Set, Optional
from .workflow import Workflow, Edge
from .node import Node, NodeType

class WorkflowValidator:
"""工作流验证器:检查环路和连通性"""

@staticmethod
def has_cycle(workflow: Workflow) -> bool:
"""使用 DFS 检测是否有环"""
visited: Set[str] = set()
rec_stack: Set[str] = set()

def dfs(node_id: str) -> bool:
visited.add(node_id)
rec_stack.add(node_id)

# 找到该节点的所有子节点
neighbors = [
edge.target_id for edge in workflow.edges
if edge.source_id == node_id
]

for neighbor in neighbors:
if neighbor not in visited:
if dfs(neighbor):
return True
elif neighbor in rec_stack:
return True

rec_stack.discard(node_id)
return False

for node_id in workflow.nodes:
if node_id not in visited:
if dfs(node_id):
return True
return False

@staticmethod
def topological_sort(workflow: Workflow) -> List[str]:
"""Kahn 算法求拓扑排序"""
in_degree: Dict[str, int] = {nid: 0 for nid in workflow.nodes}
adj: Dict[str, List[str]] = {nid: [] for nid in workflow.nodes}

for edge in workflow.edges:
adj[edge.source_id].append(edge.target_id)
in_degree[edge.target_id] += 1

queue = deque([nid for nid, deg in in_degree.items() if deg == 0])
result = []

while queue:
node_id = queue.popleft()
result.append(node_id)

for neighbor in adj[node_id]:
in_degree[neighbor] -= 1
if in_degree[neighbor] == 0:
queue.append(neighbor)

if len(result) != len(workflow.nodes):
raise ValueError("工作流存在环,无法执行")
return result

小贴士:为什么不用 DFS 直接排序?Kahn 算法的优势在于它可以同时检测环路和生成有效的执行顺序,时间复杂度 O(V+E),非常高效。

2.4 执行器核心

执行器是整个工作流引擎的心脏。它按拓扑顺序遍历节点,调用对应的处理器,并在节点之间传递数据。

# backend/core/executor.py(续)
import asyncio
from datetime import datetime
from typing import Any, Callable, Dict, Optional

class NodeContext:
"""节点执行上下文"""
def __init__(self, workflow_id: str, run_id: str):
self.workflow_id = workflow_id
self.run_id = run_id
self.variables: Dict[str, Any] = {}
self.node_outputs: Dict[str, Dict[str, Any]] = {}
self.start_time: Optional[datetime] = None
self.end_time: Optional[datetime] = None
self.status: str = "pending"
self.error: Optional[str] = None

class WorkflowExecutor:
"""工作流执行器"""

def __init__(self):
self._handlers: Dict[NodeType, Callable] = {}
self._on_node_start: Optional[Callable] = None
self._on_node_end: Optional[Callable] = None
self._on_error: Optional[Callable] = None

def register_handler(self, node_type: NodeType, handler: Callable):
"""注册节点处理器"""
self._handlers[node_type] = handler

def on_node_start(self, callback: Callable):
self._on_node_start = callback

def on_node_end(self, callback: Callable):
self._on_node_end = callback

def on_error(self, callback: Callable):
self._on_error = callback

async def execute(
self, workflow: Workflow, inputs: Dict[str, Any]
) -> Dict[str, Any]:
"""执行工作流"""
# 1. 验证工作流
if WorkflowValidator.has_cycle(workflow):
raise ValueError("工作流存在环")

exec_order = WorkflowValidator.topological_sort(workflow)

# 2. 初始化上下文
context = NodeContext(workflow.id, f"run_{datetime.now().timestamp()}")
context.variables.update(inputs)
context.start_time = datetime.now()
context.status = "running"

# 3. 按拓扑顺序执行
for node_id in exec_order:
node = workflow.nodes[node_id]

# 收集输入:来自所有入边的输出
node_inputs = {}
for edge in workflow.edges:
if edge.target_id == node_id:
source_output = context.node_outputs.get(edge.source_id, {})
node_inputs.update(source_output)

node.inputs = node_inputs

# 回调:节点开始
if self._on_node_start:
await self._on_node_start(node)

try:
# 查找处理器并执行
handler = self._handlers.get(node.type)
if handler is None:
raise ValueError(f"未注册节点处理器: {node.type}")

result = await handler(node, context)
node.outputs = result
context.node_outputs[node_id] = result

# 回调:节点完成
if self._on_node_end:
await self._on_node_end(node)

except Exception as e:
context.status = "failed"
context.error = str(e)
if self._on_error:
await self._on_error(node, e)
raise

# 4. 返回最终结果
context.status = "completed"
context.end_time = datetime.now()

# 收集 END 节点的输出
final_output = {}
for node_id, node in workflow.nodes.items():
if node.type == NodeType.END:
final_output.update(context.node_outputs.get(node_id, {}))

return {
"status": context.status,
"outputs": final_output,
"execution_time": (
(context.end_time – context.start_time).total_seconds()
if context.end_time else 0
),
"node_outputs": context.node_outputs
}

这段代码展示了工作流执行的核心思想:分离节点定义与节点处理逻辑,通过注册机制让不同类型节点各司其职。这种设计模式使得扩展新节点类型变得非常简单——只需要注册一个新的处理器函数即可。


三、LLM 抽象层

AI 工作流引擎的核心能力之一是调用大语言模型。我们需要设计一个统一的抽象层,让上层代码不必关心具体是哪个厂商的模型。

3.1 基类接口

# backend/llm/base.py
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Dict, List, Optional, AsyncIterator, Any

@dataclass
class Message:
"""消息结构"""
role: str # "system" | "user" | "assistant" | "tool"
content: str
name: Optional[str] = None
tool_calls: Optional[List[Dict]] = None
tool_call_id: Optional[str] = None

@dataclass
class LLMConfig:
"""LLM 配置"""
model: str = "gpt-3.5-turbo"
temperature: float = 0.7
max_tokens: int = 2048
top_p: float = 0.95
frequency_penalty: float = 0.0
presence_penalty: float = 0.0
stop: Optional[List[str]] = None
api_key: Optional[str] = None
base_url: Optional[str] = None
extra_params: Dict[str, Any] = field(default_factory=dict)

@dataclass
class LLMResult:
"""LLM 调用结果"""
content: str
finish_reason: str = "stop"
usage: Optional[Dict] = None
model: str = ""
raw_response: Any = None

class BaseLLM(ABC):
"""LLM 基类"""

def __init__(self, config: LLMConfig):
self.config = config

@abstractmethod
async def chat(
self, messages: List[Message], **kwargs
) -> LLMResult:
"""同步对话"""
pass

@abstractmethod
async def chat_stream(
self, messages: List[Message], **kwargs
) -> AsyncIterator[str]:
"""流式对话"""
pass

@abstractmethod
async def chat_with_tools(
self,
messages: List[Message],
tools: List[Dict],
**kwargs
) -> LLMResult:
"""带工具调用的对话"""
pass

这个接口设计体现了三个关键点:

  • 统一的消息格式:将不同厂商的消息格式归一化到 Message 类
  • 抽象的核心方法:同步、流式、工具调用三个核心接口
  • 可扩展的配置:LLMConfig 包含通用参数和 extra_params 兜底
  • 3.2 DeepSeek 实现

    DeepSeek 的 API 兼容 OpenAI 格式,所以实现起来非常简洁:

    # backend/llm/deepseek.py
    import aiohttp
    import json
    from typing import Dict, List, AsyncIterator, Optional
    from .base import BaseLLM, LLMConfig, Message, LLMResult

    class DeepSeekLLM(BaseLLM):
    """DeepSeek 模型实现"""

    DEEPSEEK_BASE_URL = "https://api.deepseek.com/v1"

    def __init__(self, config: LLMConfig):
    # DeepSeek 默认配置
    if not config.base_url:
    config.base_url = self.DEEPSEEK_BASE_URL
    super().__init__(config)

    def _build_headers(self) -> Dict:
    return {
    "Authorization": f"Bearer {self.config.api_key}",
    "Content-Type": "application/json"
    }

    def _build_messages(self, messages: List[Message]) -> List[Dict]:
    return [
    {
    "role": msg.role,
    "content": msg.content,
    **( {"name": msg.name} if msg.name else {} ),
    **( {"tool_calls": msg.tool_calls} if msg.tool_calls else {} ),
    **( {"tool_call_id": msg.tool_call_id} if msg.tool_call_id else {} ),
    }
    for msg in messages
    ]

    async def chat(
    self, messages: List[Message], **kwargs
    ) -> LLMResult:
    url = f"{self.config.base_url}/chat/completions"
    headers = self._build_headers()

    payload = {
    "model": self.config.model,
    "messages": self._build_messages(messages),
    "temperature": kwargs.get("temperature", self.config.temperature),
    "max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
    "top_p": kwargs.get("top_p", self.config.top_p),
    "stream": False,
    }

    async with aiohttp.ClientSession() as session:
    async with session.post(url, json=payload, headers=headers) as resp:
    if resp.status != 200:
    error_body = await resp.text()
    raise Exception(
    f"DeepSeek API error {resp.status}: {error_body}"
    )

    data = await resp.json()
    choice = data["choices"][0]

    return LLMResult(
    content=choice["message"]["content"],
    finish_reason=choice["finish_reason"],
    usage=data.get("usage"),
    model=data["model"],
    raw_response=data
    )

    async def chat_stream(
    self, messages: List[Message], **kwargs
    ) -> AsyncIterator[str]:
    url = f"{self.config.base_url}/chat/completions"
    headers = self._build_headers()

    payload = {
    "model": self.config.model,
    "messages": self._build_messages(messages),
    "temperature": kwargs.get("temperature", self.config.temperature),
    "max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
    "stream": True,
    "stream_options": {"include_usage": True},
    }

    async with aiohttp.ClientSession() as session:
    async with session.post(url, json=payload, headers=headers) as resp:
    async for line in resp.content:
    line = line.decode("utf-8").strip()
    if not line or line == "data: [DONE]":
    continue
    if line.startswith("data: "):
    try:
    data = json.loads(line[6:])
    delta = data["choices"][0]["delta"]
    content = delta.get("content", "")
    if content:
    yield content
    except (json.JSONDecodeError, KeyError):
    continue

    async def chat_with_tools(
    self,
    messages: List[Message],
    tools: List[Dict],
    **kwargs
    ) -> LLMResult:
    url = f"{self.config.base_url}/chat/completions"
    headers = self._build_headers()

    payload = {
    "model": self.config.model,
    "messages": self._build_messages(messages),
    "tools": tools,
    "temperature": kwargs.get("temperature", self.config.temperature),
    "max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
    "stream": False,
    }

    async with aiohttp.ClientSession() as session:
    async with session.post(url, json=payload, headers=headers) as resp:
    data = await resp.json()
    choice = data["choices"][0]
    message = choice["message"]

    return LLMResult(
    content=message.get("content", ""),
    finish_reason=choice["finish_reason"],
    usage=data.get("usage"),
    model=data["model"],
    raw_response=data
    )

    关键点:DeepSeek 的流式输出使用了 stream_options: {include_usage: true},可以在流式结束时获取 token 用量统计,这在可视化工作流中非常有用——用户可以看到"这个节点消耗了多少 token"。

    3.3 工厂模式与模型注册

    为了让工作流引擎动态选择模型,我们使用工厂模式:

    # backend/llm/factory.py
    from typing import Dict, Type, Optional
    from .base import BaseLLM, LLMConfig
    from .deepseek import DeepSeekLLM

    class LLMFactory:
    """LLM 工厂"""

    _registry: Dict[str, Type[BaseLLM]] = {}

    @classmethod
    def register(cls, name: str, llm_cls: Type[BaseLLM]):
    """注册模型类"""
    cls._registry[name] = llm_cls

    @classmethod
    def create(cls, provider: str, config: LLMConfig) -> BaseLLM:
    """创建模型实例"""
    llm_cls = cls._registry.get(provider)
    if llm_cls is None:
    raise ValueError(f"未注册的模型提供商: {provider},可选: {list(cls._registry.keys())}")
    return llm_cls(config)

    @classmethod
    def list_providers(cls) -> list:
    return list(cls._registry.keys())

    # 注册内置模型
    LLMFactory.register("deepseek", DeepSeekLLM)

    使用时只需要一行代码:

    llm = LLMFactory.create("deepseek", LLMConfig(
    model="deepseek-chat",
    api_key="sk-xxx"
    ))
    result = await llm.chat([Message(role="user", content="Hello!")])


    四、RAG 检索流水线

    RAG(Retrieval-Augmented Generation)是 AI 工作流中最常用的功能模块之一。我们来实现一个轻量级的 RAG 引擎。

    4.1 文档处理

    # backend/rag/document.py
    from dataclasses import dataclass, field
    from typing import List, Optional
    import hashlib

    @dataclass
    class Document:
    """文档模型"""
    id: str
    content: str
    metadata: dict = field(default_factory=dict)

    @classmethod
    def from_text(cls, text: str, **metadata) -> "Document":
    doc_id = hashlib.md5(text.encode()).hexdigest()[:12]
    return cls(id=doc_id, content=text, metadata=metadata)

    class TextSplitter:
    """文本分割器"""

    def __init__(
    self,
    chunk_size: int = 512,
    chunk_overlap: int = 64,
    separators: List[str] = None
    ):
    self.chunk_size = chunk_size
    self.chunk_overlap = chunk_overlap
    self.separators = separators or ["\\n\\n", "\\n", "。", ".", " ", ""]

    def split_text(self, text: str) -> List[str]:
    """将文本分割成块"""
    chunks = []
    current = ""

    for para in self._split_into_paragraphs(text):
    if len(current) + len(para) <= self.chunk_size:
    current += para
    else:
    if current:
    chunks.append(current)
    current = para

    if current:
    chunks.append(current)

    # 添加重叠
    if len(chunks) > 1 and self.chunk_overlap > 0:
    result = [chunks[0]]
    for i in range(1, len(chunks)):
    overlap_text = chunks[i-1][-self.chunk_overlap:]
    result.append(overlap_text + chunks[i])
    return result

    return chunks

    def _split_into_paragraphs(self, text: str) -> List[str]:
    """按分隔符分割段落"""
    result = [text]
    for sep in self.separators:
    if not sep:
    break
    new_result = []
    for segment in result:
    new_result.extend(segment.split(sep))
    result = new_result
    # 如果 chunks 已经够小,停止分割
    if all(len(s) <= self.chunk_size for s in result):
    break
    return [s for s in result if s.strip()]

    def split_documents(
    self, documents: List[Document]
    ) -> List[Document]:
    """分割文档列表"""
    result = []
    for doc in documents:
    chunks = self.split_text(doc.content)
    for i, chunk in enumerate(chunks):
    result.append(Document(
    id=f"{doc.id}_{i}",
    content=chunk,
    metadata={**doc.metadata, "chunk_index": i}
    ))
    return result

    4.2 Embedding 与向量存储

    # backend/rag/embedding.py
    from abc import ABC, abstractmethod
    from typing import List
    import aiohttp
    import json

    class BaseEmbedding(ABC):
    """向量化基类"""

    @abstractmethod
    async def embed(self, texts: List[str]) -> List[List[float]]:
    pass

    class DeepSeekEmbedding(BaseEmbedding):
    """DeepSeek Embedding 实现"""

    def __init__(self, api_key: str, model: str = "deepseek-embedding"):
    self.api_key = api_key
    self.model = model
    self.base_url = "https://api.deepseek.com/v1"

    async def embed(self, texts: List[str]) -> List[List[float]]:
    url = f"{self.base_url}/embeddings"
    headers = {
    "Authorization": f"Bearer {self.api_key}",
    "Content-Type": "application/json"
    }
    payload = {
    "model": self.model,
    "input": texts
    }

    async with aiohttp.ClientSession() as session:
    async with session.post(url, json=payload, headers=headers) as resp:
    data = await resp.json()
    return [item["embedding"] for item in data["data"]]

    # backend/rag/vector_store.py
    import numpy as np
    from typing import List, Tuple, Optional
    from dataclasses import dataclass

    @dataclass
    class VectorRecord:
    """向量记录"""
    id: str
    vector: List[float]
    text: str
    metadata: dict

    class InMemoryVectorStore:
    """内存向量存储(生产环境建议使用 Milvus/Pinecone)"""

    def __init__(self, dim: int = 768):
    self.dim = dim
    self.records: List[VectorRecord] = []

    def add(self, records: List[VectorRecord]):
    self.records.extend(records)

    def cosine_similarity(
    self, vec1: List[float], vec2: List[float]
    ) -> float:
    """余弦相似度计算"""
    v1 = np.array(vec1)
    v2 = np.array(vec2)
    dot = np.dot(v1, v2)
    norm = np.linalg.norm(v1) * np.linalg.norm(v2)
    return float(dot / norm) if norm > 0 else 0.0

    def search(
    self, query_vector: List[float], top_k: int = 5
    ) -> List[Tuple[VectorRecord, float]]:
    """检索最相似的 TOP-K 个向量"""
    scores = []
    for record in self.records:
    score = self.cosine_similarity(query_vector, record.vector)
    scores.append((record, score))

    # 按相似度降序排列
    scores.sort(key=lambda x: x[1], reverse=True)
    return scores[:top_k]

    4.3 检索器整合

    # backend/rag/retriever.py
    from typing import List, Optional
    from .document import Document, TextSplitter
    from .embedding import BaseEmbedding
    from .vector_store import InMemoryVectorStore, VectorRecord

    class Retriever:
    """检索器:整合文档处理、向量化和检索"""

    def __init__(
    self,
    embedding: BaseEmbedding,
    vector_store: InMemoryVectorStore,
    text_splitter: Optional[TextSplitter] = None
    ):
    self.embedding = embedding
    self.vector_store = vector_store
    self.text_splitter = text_splitter or TextSplitter()

    async def index_documents(self, documents: List[Document]):
    """索引文档到向量库"""
    # 1. 分割文档
    chunks = self.text_splitter.split_documents(documents)

    # 2. 向量化
    texts = [chunk.content for chunk in chunks]
    vectors = await self.embedding.embed(texts)

    # 3. 存储
    records = [
    VectorRecord(
    id=chunk.id,
    vector=vec,
    text=chunk.content,
    metadata=chunk.metadata
    )
    for chunk, vec in zip(chunks, vectors)
    ]
    self.vector_store.add(records)
    return len(records)

    async def retrieve(
    self, query: str, top_k: int = 5
    ) -> List[Document]:
    """检索相关文档"""
    # 1. 查询向量化
    query_vector = (await self.embedding.embed([query]))[0]

    # 2. 向量检索
    results = self.vector_store.search(query_vector, top_k)

    # 3. 转为 Document
    return [
    Document(
    id=record.id,
    content=record.text,
    metadata=record.metadata
    )
    for record, score in results
    ]

    RAG 工作流的完整链路:文档输入 → 文本分割 → 向量化 → 向量存储 → 查询向量化 → 向量检索 → 结果排序 → 注入 Prompt


    五、Agent 与工具系统

    Agent 是工作流引擎中最"智能"的部分。它能让 LLM 自主决定调用哪些工具来完成复杂任务。

    5.1 工具定义

    # backend/agent/tool.py
    from dataclasses import dataclass, field
    from typing import Dict, Any, Callable, Coroutine, List, Optional
    import json

    @dataclass
    class ToolParameter:
    """工具参数定义(JSON Schema 格式)"""
    name: str
    type: str # string, number, boolean, array, object
    description: str
    required: bool = False
    enum: Optional[List[str]] = None

    @dataclass
    class Tool:
    """工具定义"""
    name: str
    description: str
    parameters: List[ToolParameter] = field(default_factory=list)
    handler: Optional[Callable[…, Coroutine]] = None

    def to_openai_format(self) -> Dict:
    """转换为 OpenAI 工具格式"""
    properties = {}
    required = []

    for param in self.parameters:
    properties[param.name] = {
    "type": param.type,
    "description": param.description,
    }
    if param.enum:
    properties[param.name]["enum"] = param.enum
    if param.required:
    required.append(param.name)

    schema = {"type": "object", "properties": properties}
    if required:
    schema["required"] = required

    return {
    "type": "function",
    "function": {
    "name": self.name,
    "description": self.description,
    "parameters": schema
    }
    }

    async def execute(self, **kwargs) -> Any:
    """执行工具"""
    if self.handler is None:
    raise ValueError(f"工具 {self.name} 未注册处理器")
    return await self.handler(**kwargs)

    class ToolRegistry:
    """工具注册中心"""

    _tools: Dict[str, Tool] = {}

    @classmethod
    def register(cls, tool: Tool):
    cls._tools[tool.name] = tool

    @classmethod
    def get(cls, name: str) -> Optional[Tool]:
    return cls._tools.get(name)

    @classmethod
    def list_tools(cls) -> List[Tool]:
    return list(cls._tools.values())

    @classmethod
    def to_openai_tools(cls) -> List[Dict]:
    return [tool.to_openai_format() for tool in cls._tools.values()]

    5.2 内置工具示例

    # backend/agent/builtin_tools.py
    import aiohttp
    import json
    from .tool import Tool, ToolParameter, ToolRegistry

    async def web_search(query: str, limit: int = 5) -> str:
    """执行网络搜索"""
    url = "https://api.duckduckgo.com"
    params = {
    "q": query,
    "format": "json",
    "max_results": limit
    }
    async with aiohttp.ClientSession() as session:
    async with session.get(url, params=params) as resp:
    data = await resp.json()
    results = data.get("results", [])
    return json.dumps([
    {"title": r["title"], "url": r["url"], "snippet": r["body"]}
    for r in results[:limit]
    ], ensure_ascii=False)

    async def calculate(expression: str) -> str:
    """执行数学计算(安全沙箱)"""
    allowed = set("0123456789+-*/.()% ")
    if not all(c in allowed for c in expression):
    return "错误:表达式包含非法字符"
    try:
    result = eval(expression, {"__builtins__": {}}, {})
    return str(result)
    except Exception as e:
    return f"计算错误: {e}"

    async def get_current_time(format: str = "%Y-%m-%d %H:%M:%S") -> str:
    """获取当前时间"""
    from datetime import datetime
    return datetime.now().strftime(format)

    # 注册工具
    ToolRegistry.register(Tool(
    name="web_search",
    description="搜索互联网获取最新信息",
    parameters=[
    ToolParameter(name="query", type="string",
    description="搜索关键词", required=True),
    ToolParameter(name="limit", type="number",
    description="返回结果数量,默认5", required=False),
    ],
    handler=web_search
    ))

    ToolRegistry.register(Tool(
    name="calculate",
    description="执行数学计算",
    parameters=[
    ToolParameter(name="expression", type="string",
    description="数学表达式", required=True),
    ],
    handler=calculate
    ))

    ToolRegistry.register(Tool(
    name="get_current_time",
    description="获取当前日期和时间",
    parameters=[
    ToolParameter(name="format", type="string",
    description="时间格式,默认 %Y-%m-%d %H:%M:%S", required=False),
    ],
    handler=get_current_time
    ))

    5.3 Agent 执行器

    Agent 的核心逻辑是"思考-行动-观察"循环(ReAct Pattern):

    # backend/agent/agent.py
    from typing import List, Optional
    from ..llm.base import BaseLLM, Message, LLMConfig
    from .tool import ToolRegistry

    class Agent:
    """Agent 执行器(ReAct 模式)"""

    SYSTEM_PROMPT = """你是一个智能助手,可以通过调用工具来帮助用户解决问题。
    请严格按照以下步骤执行:

    1. 分析用户的请求
    2. 如果需要使用工具,请说明你的计划
    3. 调用适当的工具完成任务
    4. 汇总结果并回复用户

    你可以使用的工具:
    {tools_description}

    在调用工具时,请返回包含工具名称和参数的 JSON。
    """

    def __init__(self, llm: BaseLLM):
    self.llm = llm
    self.messages: List[Message] = []

    def _build_system_prompt(self) -> str:
    tools = ToolRegistry.list_tools()
    tools_desc = "\\n".join([
    f"- {t.name}: {t.description}(参数: {', '.join(p.name for p in t.parameters)})"
    for t in tools
    ])
    return self.SYSTEM_PROMPT.format(tools_description=tools_desc)

    async def run(self, user_input: str, max_iterations: int = 10) -> str:
    """运行 Agent"""
    self.messages = [
    Message(role="system", content=self._build_system_prompt()),
    Message(role="user", content=user_input)
    ]

    for i in range(max_iterations):
    # 1. 思考
    response = await self.llm.chat_with_tools(
    self.messages,
    tools=ToolRegistry.to_openai_tools()
    )

    # 2. 判断是否有工具调用
    if not response.raw_response.get("choices", [{}])[0].get("message", {}).get("tool_calls"):
    # 没有工具调用,直接返回结果
    return response.content

    # 3. 执行工具调用
    message = response.raw_response["choices"][0]["message"]
    self.messages.append(Message(
    role="assistant",
    content=message.get("content", ""),
    tool_calls=message.get("tool_calls")
    ))

    for tool_call in message["tool_calls"]:
    tool_name = tool_call["function"]["name"]
    tool_args = json.loads(tool_call["function"]["arguments"])

    # 查找并执行工具
    tool = ToolRegistry.get(tool_name)
    if tool is None:
    result = f"错误:未知工具 '{tool_name}'"
    else:
    try:
    result = await tool.execute(**tool_args)
    except Exception as e:
    result = f"工具执行错误: {e}"

    self.messages.append(Message(
    role="tool",
    content=str(result),
    tool_call_id=tool_call["id"]
    ))

    return "已达到最大迭代次数,请尝试简化问题。"


    六、WebSocket 实时推送

    工作流执行是异步的,用户需要实时看到每个节点的执行状态。我们使用 WebSocket 推送执行进度。

    # backend/server.py(部分)
    from fastapi import FastAPI, WebSocket, WebSocketDisconnect
    from fastapi.middleware.cors import CORSMiddleware
    import json
    import asyncio
    from typing import Dict

    app = FastAPI(title="Mini Dify")

    app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
    )

    class ConnectionManager:
    """WebSocket 连接管理"""

    def __init__(self):
    self.active_connections: Dict[str, WebSocket] = {}

    async def connect(self, client_id: str, websocket: WebSocket):
    await websocket.accept()
    self.active_connections[client_id] = websocket

    def disconnect(self, client_id: str):
    self.active_connections.pop(client_id, None)

    async def send_workflow_event(
    self, client_id: str, event: dict
    ):
    """发送工作流事件"""
    websocket = self.active_connections.get(client_id)
    if websocket:
    try:
    await websocket.send_json(event)
    except Exception:
    self.disconnect(client_id)

    manager = ConnectionManager()

    @app.websocket("/ws/{client_id}")
    async def websocket_endpoint(websocket: WebSocket, client_id: str):
    await manager.connect(client_id, websocket)
    try:
    while True:
    data = await websocket.receive_text()
    # 处理心跳
    if data == "ping":
    await websocket.send_text("pong")
    except WebSocketDisconnect:
    manager.disconnect(client_id)

    WebSocket 事件格式:

    # 事件类型
    WORKFLOW_EVENTS = {
    "node_start": { # 节点开始执行
    "node_id": str,
    "node_type": str,
    "node_title": str,
    "timestamp": float,
    "inputs": dict
    },
    "node_complete": { # 节点执行完成
    "node_id": str,
    "outputs": dict,
    "duration_ms": float
    },
    "node_error": { # 节点执行出错
    "node_id": str,
    "error": str
    },
    "llm_stream": { # LLM 流式输出
    "node_id": str,
    "content": str,
    "done": bool
    },
    "workflow_complete": { # 工作流执行完毕
    "status": str,
    "outputs": dict,
    "total_duration_ms": float
    }
    }


    七、前端可视化工作流编辑器

    一个可视化工作流引擎的前端需要实现:拖拽节点、连接节点、配置节点属性、查看执行状态。

    7.1 核心架构

    我们使用 Vue 3 + Pinia + 原生 Canvas/SVG 实现。也可以选择成熟的库如 vue-flow 或 rete.js,但为了理解原理,这里我们手动实现核心功能。

    <!– frontend/src/stores/workflow.js – Pinia Store –>
    <!– 实际项目用 JS/TS 写,这里展示核心逻辑 –>

    // 节点数据结构
    const nodeStore = {
    nodes: {}, // { [id]: Node }
    edges: [], // Edge[]
    selectedNodeId: null,

    // 添加节点
    addNode(type, position) {
    const id = `node_${Date.now()}`;
    this.nodes[id] = {
    id,
    type, // 'llm' | 'knowledge' | 'start' | 'end' | 'code' | 'agent'
    title: this.getDefaultTitle(type),
    position: { x: position.x, y: position.y },
    config: this.getDefaultConfig(type),
    inputs: {},
    outputs: {},
    status: 'idle' // idle | running | completed | error
    };
    return id;
    },

    // 添加边
    addEdge(sourceId, targetId) {
    const edge = {
    id: `edge_${Date.now()}`,
    sourceId,
    targetId,
    sourceHandle: 'output',
    targetHandle: 'input'
    };
    this.edges.push(edge);
    return edge;
    },

    // 删除节点
    removeNode(nodeId) {
    delete this.nodes[nodeId];
    this.edges = this.edges.filter(
    e => e.sourceId !== nodeId && e.targetId !== nodeId
    );
    },

    // 获取默认配置
    getDefaultConfig(type) {
    const configs = {
    llm: { model: 'deepseek-chat', prompt: '{{input}}', temperature: 0.7 },
    knowledge: { dataset: '', top_k: 5, query: '{{input}}' },
    code: { code: '# 在此编写 Python 代码\\nresult = input_data' },
    agent: { tools: ['web_search', 'calculate'], max_iterations: 10 },
    start: {},
    end: { output_key: 'result' }
    };
    return configs[type] || {};
    }
    };

    7.2 节点渲染组件

    <!– frontend/src/components/WorkflowNode.vue – 核心节点组件 –>
    <template>
    <div
    class="workflow-node"
    :class="[`type-${node.type}`, `status-${node.status}`]"
    :style="{ left: node.position.x + 'px', top: node.position.y + 'px' }"
    @mousedown.stop="startDrag"
    @click.stop="selectNode"
    >
    <!– 节点头部 –>
    <div class="node-header">
    <span class="node-icon">{{ typeIcon }}</span>
    <span class="node-title">{{ node.title }}</span>
    <span class="node-status-dot" :class="node.status"></span>
    </div>

    <!– 节点内容预览 –>
    <div class="node-body">
    <div v-if="node.type === 'llm'" class="node-preview">
    <div class="config-row">模型: {{ node.config.model }}</div>
    <div class="config-row">温度: {{ node.config.temperature }}</div>
    </div>
    <div v-else-if="node.type === 'knowledge'" class="node-preview">
    <div class="config-row">TOP-K: {{ node.config.top_k }}</div>
    </div>
    <div v-else class="node-preview">
    <div class="config-row">{{ node.type }} 节点</div>
    </div>
    </div>

    <!– 输入端口 –>
    <div class="port port-input" @mousedown.stop="startConnect('input')">
    <div class="port-dot"></div>
    </div>

    <!– 输出端口 –>
    <div class="port port-output" @mousedown.stop="startConnect('output')">
    <div class="port-dot"></div>
    </div>
    </div>
    </template>

    <script setup>
    import { computed, ref } from 'vue';

    const props = defineProps({
    node: Object
    });

    const emit = defineEmits(['drag', 'select', 'connect-start']);

    const typeIcon = computed(() => ({
    llm: '🤖', knowledge: '📚', start: '▶️',
    end: '⏹️', code: '💻', agent: '🧠'
    }[props.node.type] || '📦'));

    function startDrag(event) {
    const startX = event.clientX;
    const startY = event.clientY;
    const nodeX = props.node.position.x;
    const nodeY = props.node.position.y;

    function onMouseMove(e) {
    const dx = e.clientX – startX;
    const dy = e.clientY – startY;
    emit('drag', props.node.id, { x: nodeX + dx, y: nodeY + dy });
    }

    function onMouseUp() {
    document.removeEventListener('mousemove', onMouseMove);
    document.removeEventListener('mouseup', onMouseUp);
    }

    document.addEventListener('mousemove', onMouseMove);
    document.addEventListener('mouseup', onMouseUp);
    }

    function selectNode() {
    emit('select', props.node.id);
    }

    function startConnect(handle) {
    emit('connect-start', props.node.id, handle);
    }
    </script>

    7.3 SVG 连线系统

    节点之间的连线是可视化编辑器的关键技术点。我们使用 SVG 绘制贝塞尔曲线:

    <!– frontend/src/components/ConnectionLines.vue –>
    <template>
    <svg class="connection-lines">
    <defs>
    <!– 流动箭头标记 –>
    <marker id="arrowhead" markerWidth="10" markerHeight="7"
    refX="10" refY="3.5" orient="auto">
    <polygon points="0 0, 10 3.5, 0 7" fill="#4A90D9"/>
    </marker>
    </defs>

    <!– 已完成连线 –>
    <path
    v-for="edge in edges"
    :key="edge.id"
    :d="getEdgePath(edge)"
    class="edge-line"
    :class="{ 'edge-active': edge.active }"
    stroke="#4A90D9"
    stroke-width="2"
    fill="none"
    marker-end="url(#arrowhead)"
    />

    <!– 正在拖拽的临时连线 –>
    <path
    v-if="dragging"
    :d="getTempPath()"
    stroke="#4A90D9"
    stroke-width="2"
    stroke-dasharray="5,5"
    fill="none"
    />
    </svg>
    </template>

    <script setup>
    import { computed } from 'vue';

    const props = defineProps({
    edges: Array,
    nodes: Object,
    dragging: Object // { sourceId, sourceX, sourceY, currentX, currentY }
    });

    // 贝塞尔曲线路径生成
    function getEdgePath(edge) {
    const sourceNode = props.nodes[edge.sourceId];
    const targetNode = props.nodes[edge.targetId];

    if (!sourceNode || !targetNode) return '';

    const x1 = sourceNode.position.x + 140; // 节点宽度/2
    const y1 = sourceNode.position.y + 45; // 输出端口 Y
    const x2 = targetNode.position.x; // 输入端口 X
    const y2 = targetNode.position.y + 45; // 输入端口 Y

    // 贝塞尔曲线控制点(水平方向偏移更多以实现平滑曲线)
    const dx = Math.abs(x2 – x1) * 0.5;
    const cp1x = x1 + Math.min(dx, 100);
    const cp1y = y1;
    const cp2x = x2 – Math.min(dx, 100);
    const cp2y = y2;

    return `M ${x1} ${y1} C ${cp1x} ${cp1y}, ${cp2x} ${cp2y}, ${x2} ${y2}`;
    }

    function getTempPath() {
    if (!props.dragging) return '';
    const { sourceX, sourceY, currentX, currentY } = props.dragging;
    const dx = Math.abs(currentX – sourceX) * 0.5;
    return `M ${sourceX} ${sourceY} C ${sourceX + dx} ${sourceY}, ${currentX – dx} ${currentY}, ${currentX} ${currentY}`;
    }
    </script>

    7.4 节点属性配置面板

    当用户选中节点时,右侧弹出属性编辑面板:

    <!– frontend/src/components/NodeConfigPanel.vue –>
    <template>
    <div v-if="node" class="config-panel">
    <h3>{{ node.title }}</h3>
    <div class="config-form">
    <!– 通用:节点名称 –>
    <div class="form-group">
    <label>节点名称</label>
    <input v-model="node.title" @input="updateNode" />
    </div>

    <!– LLM 节点专属配置 –>
    <template v-if="node.type === 'llm'">
    <div class="form-group">
    <label>模型</label>
    <select v-model="node.config.model" @change="updateNode">
    <option value="deepseek-chat">DeepSeek-Chat</option>
    <option value="deepseek-reasoner">DeepSeek-Reasoner</option>
    <option value="gpt-4o">GPT-4o</option>
    <option value="gpt-3.5-turbo">GPT-3.5-Turbo</option>
    </select>
    </div>
    <div class="form-group">
    <label>提示词模板</label>
    <textarea v-model="node.config.prompt" rows="4"
    @input="updateNode"
    placeholder="支持 {{ variable }} 变量插值" />
    </div>
    <div class="form-group inline">
    <div>
    <label>温度</label>
    <input type="range" min="0" max="2" step="0.1"
    v-model.number="node.config.temperature"
    @input="updateNode" />
    <span>{{ node.config.temperature }}</span>
    </div>
    <div>
    <label>Max Tokens</label>
    <input type="number" v-model.number="node.config.max_tokens"
    @input="updateNode" min="1" max="32768" />
    </div>
    </div>
    </template>

    <!– 知识检索节点配置 –>
    <template v-if="node.type === 'knowledge'">
    <div class="form-group">
    <label>数据集</label>
    <select v-model="node.config.dataset" @change="updateNode">
    <option v-for="ds in datasets" :key="ds.id" :value="ds.id">
    {{ ds.name }}
    </option>
    </select>
    </div>
    <div class="form-group">
    <label>检索数量 (TOP-K)</label>
    <input type="number" v-model.number="node.config.top_k"
    @input="updateNode" min="1" max="50" />
    </div>
    </template>

    <!– 代码节点配置 –>
    <template v-if="node.type === 'code'">
    <div class="form-group">
    <label>Python 代码</label>
    <textarea v-model="node.config.code" rows="8"
    @input="updateNode"
    class="code-editor" />
    <p class="hint">变量: <code>input_data</code> | 输出: <code>result</code></p>
    </div>
    </template>
    </div>
    </div>
    </template>


    八、节点处理器注册与工作流编排

    现在我们将 LLM、RAG、Agent 等模块集成为工作流节点处理器,并注册到执行器中。

    # backend/core/node_handlers.py
    import json
    from .node import Node, NodeType
    from .executor import NodeContext
    from ..llm.base import Message, LLMConfig
    from ..llm.factory import LLMFactory

    async def handle_llm_node(node: Node, context: NodeContext) -> dict:
    """LLM 节点处理器"""
    config = node.config

    # 1. 变量插值:将 {{ var }} 替换为上下文中的实际值
    prompt = config.prompt or "{{input}}"
    for var_name, var_value in context.variables.items():
    prompt = prompt.replace("{{" + var_name + "}}", str(var_value))

    # 2. 历史消息处理
    messages = [Message(role="system", content=prompt)]

    # 从输入中获取用户消息
    user_input = node.inputs.get("input", context.variables.get("input", ""))
    if user_input:
    messages.append(Message(role="user", content=str(user_input)))

    # 3. 调用 LLM
    llm = LLMFactory.create(
    "deepseek",
    LLMConfig(
    model=config.model or "deepseek-chat",
    temperature=config.temperature or 0.7,
    max_tokens=config.max_tokens or 2048,
    api_key=context.variables.get("DEEPSEEK_API_KEY")
    )
    )

    result = await llm.chat(messages)

    return {
    "output": result.content,
    "usage": result.usage,
    "model": result.model
    }

    async def handle_knowledge_node(node: Node, context: NodeContext) -> dict:
    """知识检索节点处理器"""
    from ..rag.embedding import DeepSeekEmbedding
    from ..rag.vector_store import InMemoryVectorStore
    from ..rag.retriever import Retriever

    config = node.config

    # 初始化检索器
    embedding = DeepSeekEmbedding(
    api_key=context.variables.get("DEEPSEEK_API_KEY"),
    model="deepseek-embedding"
    )
    vector_store = InMemoryVectorStore()
    retriever = Retriever(embedding, vector_store)

    # 从上下文中获取要检索的文档
    query = str(node.inputs.get("query", config.query))
    documents = context.variables.get("documents", [])

    # 索引文档
    if documents:
    await retriever.index_documents(documents)

    # 执行检索
    results = await retriever.retrieve(query, top_k=config.top_k or 5)

    return {
    "documents": [
    {"id": doc.id, "content": doc.content,
    "metadata": doc.metadata}
    for doc in results
    ],
    "context": "\\n\\n".join(f"[文档 {i+1}]: {doc.content}"
    for i, doc in enumerate(results))
    }

    async def handle_code_node(node: Node, context: NodeContext) -> dict:
    """代码执行节点处理器(安全沙箱)"""
    code = node.config.code or ""

    # 准备执行环境
    input_data = node.inputs.get("input", {})
    workflow_vars = context.variables

    # 受限的全局变量
    restricted_globals = {
    "__builtins__": {
    "len": len, "str": str, "int": int,
    "float": float, "list": list, "dict": dict,
    "range": range, "map": map, "filter": filter,
    "sorted": sorted, "reversed": reversed,
    "enumerate": enumerate, "zip": zip,
    "min": min, "max": max, "sum": sum,
    "abs": abs, "round": round,
    "True": True, "False": False, "None": None,
    "json": json,
    }
    }

    local_vars = {"input_data": input_data, "workflow_vars": workflow_vars}

    try:
    exec(code, restricted_globals, local_vars)
    result = local_vars.get("result", input_data)
    return {"output": result}
    except Exception as e:
    return {"error": str(e), "output": None}

    async def handle_agent_node(node: Node, context: NodeContext) -> dict:
    """Agent 节点处理器"""
    from ..agent.agent import Agent
    from ..agent.tool import ToolRegistry
    from ..agent.builtin_tools import register_builtin_tools

    # 注册内置工具
    register_builtin_tools()

    llm = LLMFactory.create(
    "deepseek",
    LLMConfig(
    model=node.config.model or "deepseek-chat",
    api_key=context.variables.get("DEEPSEEK_API_KEY")
    )
    )

    agent = Agent(llm)
    user_input = str(node.inputs.get("input", ""))
    result = await agent.run(user_input, max_iterations=node.config.max_iterations or 10)

    return {"output": result}

    # 注册处理器
    def register_node_handlers(executor):
    """注册所有节点处理器到执行器"""
    executor.register_handler(NodeType.LLM, handle_llm_node)
    executor.register_handler(NodeType.KNOWLEDGE_RETRIEVAL, handle_knowledge_node)
    executor.register_handler(NodeType.CODE, handle_code_node)
    executor.register_handler(NodeType.TOOL, handle_agent_node)
    # 更多节点类型…


    九、完整演示:构建一个 RAG 聊天机器人

    现在让我们使用 Mini Dify 搭建一个完整的 RAG 聊天机器人工作流。

    9.1 工作流定义

    # demo/rag_chat_workflow.py
    from mini_dify.core.node import Node, NodeConfig, NodeType
    from mini_dify.core.workflow import Workflow, Edge

    def build_rag_chat_workflow():
    """构建 RAG 聊天工作流"""

    workflow = Workflow(
    id="rag_chat_demo",
    title="RAG 智能问答系统"
    )

    # 1. 起始节点
    start = Node(
    id="start_1",
    type=NodeType.START,
    title="用户输入",
    config=NodeConfig()
    )
    workflow.nodes[start.id] = start
    workflow.entry_node_id = start.id

    # 2. 知识检索节点
    knowledge = Node(
    id="knowledge_1",
    type=NodeType.KNOWLEDGE_RETRIEVAL,
    title="知识库检索",
    config=NodeConfig(
    top_k=3,
    query="{{input}}"
    )
    )
    workflow.nodes[knowledge.id] = knowledge

    # 3. LLM 回答节点
    llm = Node(
    id="llm_1",
    type=NodeType.LLM,
    title="DeepSeek 回答",
    config=NodeConfig(
    model="deepseek-chat",
    prompt="""你是一个智能问答助手。请基于以下参考文档回答用户的问题。

    参考文档:
    {{context}}

    用户问题:{{input}}

    请严格基于参考文档回答,如果文档中不包含相关信息,请明确告知用户。
    回答要详细、准确、有条理。""",
    temperature=0.3,
    max_tokens=2048
    )
    )
    workflow.nodes[llm.id] = llm

    # 4. 结束节点
    end = Node(
    id="end_1",
    type=NodeType.END,
    title="输出结果",
    config=NodeConfig()
    )
    workflow.nodes[end.id] = end

    # 连接节点(边)
    workflow.edges = [
    Edge(id="e1", source_id="start_1", target_id="knowledge_1"),
    Edge(id="e2", source_id="knowledge_1", target_id="llm_1"),
    Edge(id="e3", source_id="llm_1", target_id="end_1"),
    ]

    return workflow

    # 执行工作流
    async def run_demo():
    from mini_dify.core.executor import WorkflowExecutor
    from mini_dify.core.node_handlers import register_node_handlers

    # 构建工作流
    workflow = build_rag_chat_workflow()

    # 配置执行器
    executor = WorkflowExecutor()
    register_node_handlers(executor)

    # 添加 WebSocket 回调
    async def on_node_start(node):
    print(f"[开始] {node.title} ({node.type.value})")

    async def on_node_end(node):
    print(f"[完成] {node.title} → {str(node.outputs)[:100]}…")

    executor.on_node_start(on_node_start)
    executor.on_node_end(on_node_end)

    # 注入文档
    documents = [
    Document.from_text(
    "DeepSeek 是一款由深度求索公司开发的大语言模型,"
    "具备强大的推理能力和代码理解能力。DeepSeek-V3 是其主要版本,"
    "支持 128K 上下文窗口,在多个基准测试中表现优异。",
    source="deepseek_intro.md"
    ),
    Document.from_text(
    "Mini Dify 是一个简化版的可视化 AI 工作流引擎,"
    "支持 LLM 调用、RAG 检索、Agent 工具调用等核心功能。"
    "它使用 Python + Vue 构建,适合学习和二次开发。",
    source="mini_dify_intro.md"
    ),
    ]

    # 执行
    result = await executor.execute(workflow, {
    "input": "DeepSeek 支持多长的上下文窗口?",
    "documents": documents,
    "DEEPSEEK_API_KEY": "sk-your-key-here"
    })

    print(f"\\n最终输出: {result['outputs']}")
    print(f"执行耗时: {result['execution_time']:.2f}s")

    # 输出结果:
    # [开始] 用户输入 (start)
    # [开始] 知识库检索 (knowledge_retrieval)
    # [完成] 知识库检索 → {'documents': […], 'context': '…'}
    # [开始] DeepSeek 回答 (llm)
    # [完成] DeepSeek 回答 → {'output': '根据参考文档,DeepSeek-V3 支持 128K 上下文窗口…'}
    # [开始] 输出结果 (end)
    # [完成] 输出结果 → {}
    # 最终输出: {'output': '根据参考文档,DeepSeek-V3 支持 128K 上下文窗口…'}
    # 执行耗时: 3.45s

    9.2 可视化效果

    当你启动前端页面后,会看到:

  • 左侧是节点面板,可以拖拽不同类型的节点到画布
  • 中间是画布区域,节点通过贝塞尔曲线连接
  • 右侧是节点属性配置面板
  • 点击"运行"按钮后,节点会按顺序变为"执行中"状态(红色脉冲动画),完成变为绿色
  • 每个节点下方会显示实时的输入输出数据

  • 十、性能优化与最佳实践

    10.1 并发执行优化

    对于没有相互依赖的节点,可以并行执行:

    async def execute_parallel(self, workflow, inputs):
    """支持并行节点执行"""
    exec_order = WorkflowValidator.topological_sort(workflow)

    # 按层级分组
    levels = []
    current_level = []
    in_degree = {nid: 0 for nid in workflow.nodes}

    for edge in workflow.edges:
    in_degree[edge.target_id] += 1

    visited = set()
    while len(visited) < len(exec_order):
    next_level = []
    for nid in exec_order:
    if nid not in visited and in_degree[nid] == 0:
    next_level.append(nid)
    if not next_level:
    break
    levels.append(next_level)
    visited.update(next_level)
    for edge in workflow.edges:
    if edge.source_id in next_level:
    in_degree[edge.target_id] -= 1

    # 逐层并行执行
    for level in levels:
    tasks = [self._execute_node(nid, workflow, context)
    for nid in level]
    await asyncio.gather(*tasks)

    10.2 缓存策略

    重复执行相同节点时可以缓存结果:

    class NodeCache:
    """节点结果缓存"""

    def __init__(self, ttl: int = 300):
    self._cache = {}
    self._ttl = ttl

    def _make_key(self, node: Node, inputs: dict) -> str:
    # 基于节点配置和输入生成缓存键
    content = json.dumps({
    "node_id": node.id,
    "config": node.config.__dict__,
    "inputs": inputs
    }, sort_keys=True)
    return hashlib.md5(content.encode()).hexdigest()

    def get(self, node: Node, inputs: dict):
    key = self._make_key(node, inputs)
    entry = self._cache.get(key)
    if entry and time.time() – entry["time"] < self._ttl:
    return entry["result"]
    return None

    def set(self, node: Node, inputs: dict, result: dict):
    key = self._make_key(node, inputs)
    self._cache[key] = {
    "result": result,
    "time": time.time()
    }

    10.3 错误恢复与重试

    async def _execute_node_with_retry(
    self, node_id: str, workflow, context,
    max_retries: int = 3
    ):
    """带重试机制的节点执行"""
    node = workflow.nodes[node_id]
    handler = self._handlers[node.type]

    for attempt in range(max_retries):
    try:
    return await handler(node, context)
    except Exception as e:
    if attempt < max_retries – 1:
    wait_time = 2 ** attempt # 指数退避
    await asyncio.sleep(wait_time)
    continue
    raise


    十一、总结与扩展

    11.1 我们实现了什么

    通过本文,我们从零构建了一个完整的可视化 AI 工作流引擎 Mini Dify,核心包括:

    模块核心功能关键技术
    工作流引擎 节点定义、拓扑排序、执行器 Kahn 算法、异步编排
    LLM 抽象层 统一接口、多模型支持、工厂模式 策略模式、aiohttp
    RAG 流水线 文档分割、向量化、相似度检索 TextSplitter、余弦相似度
    Agent 系统 工具注册、ReAct 循环、Function Call 插件架构、JSON Schema
    实时推送 WebSocket 节点状态、流式输出 FastAPI WebSocket
    前端编辑器 拖拽节点、SVG 连线、配置面板 Vue 3、Pinia、贝塞尔曲线

    11.2 扩展方向

    Mini Dify 虽然小巧,但架构设计具有良好的扩展性:

  • 数据库持久化:将工作流定义存入 PostgreSQL,向量存入 Milvus
  • 模板市场:支持导入/导出工作流模板
  • 多租户:添加用户系统和工作空间隔离
  • 监控日志:记录每次执行的全链路跟踪
  • 插件系统:支持第三方开发者编写自定义节点
  • 版本管理:工作流版本控制和回滚
  • 定时调度:支持 Cron 表达式定时执行工作流
  • 11.3 与真实 Dify 的差距

    与完整的 Dify 项目相比,Mini Dify 是一个教学版本,主要差距在于:

    • 功能完整性:真实 Dify 支持数十种节点类型,我们只实现了核心的 6 种
    • 企业级能力:权限管理、审计日志、多语言等企业特性未实现
    • 性能优化:生产环境需要连接池、缓存集群、分布式执行
    • 生态集成:Dify 支持数十种外部服务的接入

    但反过来说,这正是 Mini Dify 的价值所在——轻量、可理解、可定制。对于个人开发者而言,450 行核心代码就能跑通一个 AI 工作流引擎,性价比极高。


    💡 实战指南:想亲自体验 DeepSeek 的强大能力?查看 DeepSeek 大模型实战指南 获取最新 API 文档和示例代码。如果你需要在云服务器上一键部署 DeepSeek + Dify 环境,可以参考 华为云 Flexus 一键部署方案。


    本文是「手写系列」的第十七篇。从 Transformer 到工作流引擎,我们用代码致敬每一个创新框架背后的设计思想。如果对你有帮助,欢迎点赞转发,让更多开发者看见技术的深度。

    赞(0)
    未经允许不得转载:171主机测评 » 手写 Mini Dify:从零构建可视化 AI 工作流引擎
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址