Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
workflow_entry.py295 linesDownload Raw Back to workflow
1import logging2import time3import uuid4from collections.abc import Generator, Mapping, Sequence5from typing import Any, Optional, cast6 7from configs import dify_config8from core.app.app_config.entities import FileExtraConfig9from core.app.apps.base_app_queue_manager import GenerateTaskStoppedError10from core.app.entities.app_invoke_entities import InvokeFrom11from core.file.models import File, FileTransferMethod, FileType, ImageConfig12from core.workflow.callbacks import WorkflowCallback13from core.workflow.entities.variable_pool import VariablePool14from core.workflow.errors import WorkflowNodeRunFailedError15from core.workflow.graph_engine.entities.event import GraphEngineEvent, GraphRunFailedEvent, InNodeEvent16from core.workflow.graph_engine.entities.graph import Graph17from core.workflow.graph_engine.entities.graph_init_params import GraphInitParams18from core.workflow.graph_engine.entities.graph_runtime_state import GraphRuntimeState19from core.workflow.graph_engine.graph_engine import GraphEngine20from core.workflow.nodes import NodeType21from core.workflow.nodes.base import BaseNode, BaseNodeData22from core.workflow.nodes.event import NodeEvent23from core.workflow.nodes.llm import LLMNodeData24from core.workflow.nodes.node_mapping import node_type_classes_mapping25from models.enums import UserFrom26from models.workflow import (27    Workflow,28    WorkflowType,29)30 31logger = logging.getLogger(__name__)32 33 34class WorkflowEntry:35    def __init__(36        self,37        tenant_id: str,38        app_id: str,39        workflow_id: str,40        workflow_type: WorkflowType,41        graph_config: Mapping[str, Any],42        graph: Graph,43        user_id: str,44        user_from: UserFrom,45        invoke_from: InvokeFrom,46        call_depth: int,47        variable_pool: VariablePool,48        thread_pool_id: Optional[str] = None,49    ) -> None:50        """51        Init workflow entry52        :param tenant_id: tenant id53        :param app_id: app id54        :param workflow_id: workflow id55        :param workflow_type: workflow type56        :param graph_config: workflow graph config57        :param graph: workflow graph58        :param user_id: user id59        :param user_from: user from60        :param invoke_from: invoke from61        :param call_depth: call depth62        :param variable_pool: variable pool63        :param thread_pool_id: thread pool id64        """65        # check call depth66        workflow_call_max_depth = dify_config.WORKFLOW_CALL_MAX_DEPTH67        if call_depth > workflow_call_max_depth:68            raise ValueError("Max workflow call depth {} reached.".format(workflow_call_max_depth))69 70        # init workflow run state71        self.graph_engine = GraphEngine(72            tenant_id=tenant_id,73            app_id=app_id,74            workflow_type=workflow_type,75            workflow_id=workflow_id,76            user_id=user_id,77            user_from=user_from,78            invoke_from=invoke_from,79            call_depth=call_depth,80            graph=graph,81            graph_config=graph_config,82            variable_pool=variable_pool,83            max_execution_steps=dify_config.WORKFLOW_MAX_EXECUTION_STEPS,84            max_execution_time=dify_config.WORKFLOW_MAX_EXECUTION_TIME,85            thread_pool_id=thread_pool_id,86        )87 88    def run(89        self,90        *,91        callbacks: Sequence[WorkflowCallback],92    ) -> Generator[GraphEngineEvent, None, None]:93        """94        :param callbacks: workflow callbacks95        """96        graph_engine = self.graph_engine97 98        try:99            # run workflow100            generator = graph_engine.run()101            for event in generator:102                if callbacks:103                    for callback in callbacks:104                        callback.on_event(event=event)105                yield event106        except GenerateTaskStoppedError:107            pass108        except Exception as e:109            logger.exception("Unknown Error when workflow entry running")110            if callbacks:111                for callback in callbacks:112                    callback.on_event(event=GraphRunFailedEvent(error=str(e)))113            return114 115    @classmethod116    def single_step_run(117        cls, workflow: Workflow, node_id: str, user_id: str, user_inputs: dict118    ) -> tuple[BaseNode, Generator[NodeEvent | InNodeEvent, None, None]]:119        """120        Single step run workflow node121        :param workflow: Workflow instance122        :param node_id: node id123        :param user_id: user id124        :param user_inputs: user inputs125        :return:126        """127        # fetch node info from workflow graph128        graph = workflow.graph_dict129        if not graph:130            raise ValueError("workflow graph not found")131 132        nodes = graph.get("nodes")133        if not nodes:134            raise ValueError("nodes not found in workflow graph")135 136        # fetch node config from node id137        node_config = None138        for node in nodes:139            if node.get("id") == node_id:140                node_config = node141                break142 143        if not node_config:144            raise ValueError("node id not found in workflow graph")145 146        # Get node class147        node_type = NodeType(node_config.get("data", {}).get("type"))148        node_cls = node_type_classes_mapping.get(node_type)149        node_cls = cast(type[BaseNode], node_cls)150 151        if not node_cls:152            raise ValueError(f"Node class not found for node type {node_type}")153 154        # init variable pool155        variable_pool = VariablePool(156            system_variables={},157            user_inputs={},158            environment_variables=workflow.environment_variables,159        )160 161        # init graph162        graph = Graph.init(graph_config=workflow.graph_dict)163 164        # init workflow run state165        node_instance = node_cls(166            id=str(uuid.uuid4()),167            config=node_config,168            graph_init_params=GraphInitParams(169                tenant_id=workflow.tenant_id,170                app_id=workflow.app_id,171                workflow_type=WorkflowType.value_of(workflow.type),172                workflow_id=workflow.id,173                graph_config=workflow.graph_dict,174                user_id=user_id,175                user_from=UserFrom.ACCOUNT,176                invoke_from=InvokeFrom.DEBUGGER,177                call_depth=0,178            ),179            graph=graph,180            graph_runtime_state=GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter()),181        )182 183        try:184            # variable selector to variable mapping185            try:186                variable_mapping = node_cls.extract_variable_selector_to_variable_mapping(187                    graph_config=workflow.graph_dict, config=node_config188                )189            except NotImplementedError:190                variable_mapping = {}191 192            cls.mapping_user_inputs_to_variable_pool(193                variable_mapping=variable_mapping,194                user_inputs=user_inputs,195                variable_pool=variable_pool,196                tenant_id=workflow.tenant_id,197                node_type=node_type,198                node_data=node_instance.node_data,199            )200 201            # run node202            generator = node_instance.run()203 204            return node_instance, generator205        except Exception as e:206            raise WorkflowNodeRunFailedError(node_instance=node_instance, error=str(e))207 208    @staticmethod209    def handle_special_values(value: Optional[Mapping[str, Any]]) -> Mapping[str, Any] | None:210        return WorkflowEntry._handle_special_values(value)211 212    @staticmethod213    def _handle_special_values(value: Any) -> Any:214        if value is None:215            return value216        if isinstance(value, dict):217            res = {}218            for k, v in value.items():219                res[k] = WorkflowEntry._handle_special_values(v)220            return res221        if isinstance(value, list):222            res = []223            for item in value:224                res.append(WorkflowEntry._handle_special_values(item))225            return res226        if isinstance(value, File):227            return value.to_dict()228        return value229 230    @classmethod231    def mapping_user_inputs_to_variable_pool(232        cls,233        variable_mapping: Mapping[str, Sequence[str]],234        user_inputs: dict,235        variable_pool: VariablePool,236        tenant_id: str,237        node_type: NodeType,238        node_data: BaseNodeData,239    ) -> None:240        for node_variable, variable_selector in variable_mapping.items():241            # fetch node id and variable key from node_variable242            node_variable_list = node_variable.split(".")243            if len(node_variable_list) < 1:244                raise ValueError(f"Invalid node variable {node_variable}")245 246            node_variable_key = ".".join(node_variable_list[1:])247 248            if (node_variable_key not in user_inputs and node_variable not in user_inputs) and not variable_pool.get(249                variable_selector250            ):251                raise ValueError(f"Variable key {node_variable} not found in user inputs.")252 253            # fetch variable node id from variable selector254            variable_node_id = variable_selector[0]255            variable_key_list = variable_selector[1:]256            variable_key_list = cast(list[str], variable_key_list)257 258            # get input value259            input_value = user_inputs.get(node_variable)260            if not input_value:261                input_value = user_inputs.get(node_variable_key)262 263            # FIXME: temp fix for image type264            if node_type == NodeType.LLM:265                new_value = []266                if isinstance(input_value, list):267                    node_data = cast(LLMNodeData, node_data)268 269                    detail = node_data.vision.configs.detail if node_data.vision.configs else None270 271                    for item in input_value:272                        if isinstance(item, dict) and "type" in item and item["type"] == "image":273                            transfer_method = FileTransferMethod.value_of(item.get("transfer_method"))274                            file = File(275                                tenant_id=tenant_id,276                                type=FileType.IMAGE,277                                transfer_method=transfer_method,278                                remote_url=item.get("url")279                                if transfer_method == FileTransferMethod.REMOTE_URL280                                else None,281                                related_id=item.get("upload_file_id")282                                if transfer_method == FileTransferMethod.LOCAL_FILE283                                else None,284                                _extra_config=FileExtraConfig(285                                    image_config=ImageConfig(detail=detail) if detail else None286                                ),287                            )288                            new_value.append(file)289 290                if new_value:291                    input_value = new_value292 293            # append variable and value to variable pool294            variable_pool.add([variable_node_id] + variable_key_list, input_value)295