Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 34 additions & 28 deletions apps/application/workflow/common.py
Original file line number Diff line number Diff line change
@@ -1,19 +1,19 @@
# coding=utf-8
"""
@project: MaxKB
@Author:虎虎虎
@file: workflow.py
@date:2026/6/29 10:58
@desc:
@project: MaxKB
@Author:虎虎虎
@file: workflow.py
@date:2026/6/29 10:58
@desc:
"""

from enum import Enum
from typing import List, Dict

from common.utils.common import group_by


class Node:

def __init__(self, _id: str, _type: str, x: int, y: int, properties: dict, **kwargs):
"""

Expand Down Expand Up @@ -64,29 +64,29 @@ def init_fields(workflow):
result = []
for node in workflow.nodes:
properties = node.properties
node_name = properties.get('stepName')
node_name = properties.get("stepName")
node_id = node.id
node_config = properties.get('config')
result.append(NodeField(node_id, node_name, '异常信息', 'exception_message'))
node_config = properties.get("config")
result.append(NodeField(node_id, node_name, "异常信息", "exception_message"))
if node_config is not None:
fields = node_config.get('fields')
fields = node_config.get("fields")
if fields is not None:
for field in fields:
result.append(NodeField(node_id, node_name, field.get('label'), field.get('value')))
global_fields = node_config.get('globalFields')
result.append(NodeField(node_id, node_name, field.get("label"), field.get("value")))
global_fields = node_config.get("globalFields")
if global_fields is not None:
for global_field in global_fields:
result.append(NodeField('global', '全局变量', global_field.get('label'), global_field.get('value')))
chat_fields = node_config.get('chatFields')
result.append(NodeField("global", "全局变量", global_field.get("label"), global_field.get("value")))
chat_fields = node_config.get("chatFields")
if chat_fields is not None:
for chat_field in chat_fields:
result.append(NodeField('chat', 'chat', chat_field.get('label'), chat_field.get('value')))
result.append(NodeField("chat", "chat", chat_field.get("label"), chat_field.get("value")))
result.sort(key=lambda f: len(f.node_name + f.value), reverse=True)
return result


def get_node_parameters(node):
return node.properties.get('node_data', {})
return node.properties.get("node_data", {})


class NodeField:
Expand All @@ -99,13 +99,18 @@ def __init__(self, node_id, node_name, label, value):
def reset_variable(self, prompt: str):
userVariable = self.node_name + "." + self.value
systemVariable = f"context.get('{self.node_id}').get('{self.value}','')"
return prompt.replace(userVariable, systemVariable)
prompt = prompt.replace(userVariable, systemVariable)
# 全局变量:前端用 global.xxx 引用,也要能解析到 context['global']
if self.node_id == "global":
prompt = prompt.replace(f"global.{self.value}", systemVariable)
return prompt


class Workflow:
"""
节点列表
"""

nodes: List[Node]
"""
线列表
Expand Down Expand Up @@ -133,14 +138,15 @@ def __init__(self, nodes: List[Node], edges: List[Edge]):
self.edges = edges
self.node_map = {node.id: node for node in nodes}

self.up_node_map = {key: [EdgeNode(edge, self.node_map.get(edge.sourceNodeId)) for
edge in edges] for
key, edges in
group_by(edges, key=lambda edge: edge.targetNodeId).items()}
self.up_node_map = {
key: [EdgeNode(edge, self.node_map.get(edge.sourceNodeId)) for edge in edges]
for key, edges in group_by(edges, key=lambda edge: edge.targetNodeId).items()
}

self.next_node_map = {key: [EdgeNode(edge, self.node_map.get(edge.targetNodeId)) for edge in edges] for
key, edges in
group_by(edges, key=lambda edge: edge.sourceNodeId).items()}
self.next_node_map = {
key: [EdgeNode(edge, self.node_map.get(edge.targetNodeId)) for edge in edges]
for key, edges in group_by(edges, key=lambda edge: edge.sourceNodeId).items()
}
self.node_field_list = init_fields(self)

def get_node(self, node_id):
Expand Down Expand Up @@ -199,8 +205,8 @@ class WorkflowType(Enum):


def new_instance(flow_obj: Dict, workflow_type: WorkflowType = WorkflowType.APPLICATION):
nodes = flow_obj.get('nodes')
edges = flow_obj.get('edges')
nodes = [Node(node.get('id'), node.get('type'), **node) for node in nodes]
edges = [Edge(edge.get('id'), edge.get('type'), **edge) for edge in edges]
nodes = flow_obj.get("nodes")
edges = flow_obj.get("edges")
nodes = [Node(node.get("id"), node.get("type"), **node) for node in nodes]
edges = [Edge(edge.get("id"), edge.get("type"), **edge) for edge in edges]
return Workflow(nodes, edges)
5 changes: 3 additions & 2 deletions apps/application/workflow/nodes/start_node/start_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,15 +93,16 @@ def execute(self):
for key, value in node_variable.items():
self.write_context(key, value)

# 全局变量统一放进 context['global'],与 reset_variable / get_reference_field 的引用约定一致
for key, value in workflow_variable.items():
self.workflow_manage.context[key] = value
self.workflow_manage.write_context("global", key, value)

config = self.node.properties.get("config", {})
if config:
for field in config.get("globalFields", []):
key = field.get("value")
if key:
self.workflow_manage.context[key] = workflow_variable.get(key, "")
self.workflow_manage.write_context("global", key, workflow_variable.get(key, ""))

def get_details(self, index: int = 0, position: dict = None, old_details: dict = None, **kwargs):
details = super().get_details(index, position, old_details, **kwargs)
Expand Down
Loading