Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
graph.py716 linesDownload Raw Back to entities
1import uuid2from collections.abc import Mapping3from typing import Any, Optional, cast4 5from pydantic import BaseModel, Field6 7from core.workflow.graph_engine.entities.run_condition import RunCondition8from core.workflow.nodes import NodeType9from core.workflow.nodes.answer.answer_stream_generate_router import AnswerStreamGeneratorRouter10from core.workflow.nodes.answer.entities import AnswerStreamGenerateRoute11from core.workflow.nodes.end.end_stream_generate_router import EndStreamGeneratorRouter12from core.workflow.nodes.end.entities import EndStreamParam13 14 15class GraphEdge(BaseModel):16    source_node_id: str = Field(..., description="source node id")17    target_node_id: str = Field(..., description="target node id")18    run_condition: Optional[RunCondition] = None19    """run condition"""20 21 22class GraphParallel(BaseModel):23    id: str = Field(default_factory=lambda: str(uuid.uuid4()), description="random uuid parallel id")24    start_from_node_id: str = Field(..., description="start from node id")25    parent_parallel_id: Optional[str] = None26    """parent parallel id"""27    parent_parallel_start_node_id: Optional[str] = None28    """parent parallel start node id"""29    end_to_node_id: Optional[str] = None30    """end to node id"""31 32 33class Graph(BaseModel):34    root_node_id: str = Field(..., description="root node id of the graph")35    node_ids: list[str] = Field(default_factory=list, description="graph node ids")36    node_id_config_mapping: dict[str, dict] = Field(37        default_factory=list, description="node configs mapping (node id: node config)"38    )39    edge_mapping: dict[str, list[GraphEdge]] = Field(40        default_factory=dict, description="graph edge mapping (source node id: edges)"41    )42    reverse_edge_mapping: dict[str, list[GraphEdge]] = Field(43        default_factory=dict, description="reverse graph edge mapping (target node id: edges)"44    )45    parallel_mapping: dict[str, GraphParallel] = Field(46        default_factory=dict, description="graph parallel mapping (parallel id: parallel)"47    )48    node_parallel_mapping: dict[str, str] = Field(49        default_factory=dict, description="graph node parallel mapping (node id: parallel id)"50    )51    answer_stream_generate_routes: AnswerStreamGenerateRoute = Field(..., description="answer stream generate routes")52    end_stream_param: EndStreamParam = Field(..., description="end stream param")53 54    @classmethod55    def init(cls, graph_config: Mapping[str, Any], root_node_id: Optional[str] = None) -> "Graph":56        """57        Init graph58 59        :param graph_config: graph config60        :param root_node_id: root node id61        :return: graph62        """63        # edge configs64        edge_configs = graph_config.get("edges")65        if edge_configs is None:66            edge_configs = []67 68        edge_configs = cast(list, edge_configs)69 70        # reorganize edges mapping71        edge_mapping: dict[str, list[GraphEdge]] = {}72        reverse_edge_mapping: dict[str, list[GraphEdge]] = {}73        target_edge_ids = set()74        for edge_config in edge_configs:75            source_node_id = edge_config.get("source")76            if not source_node_id:77                continue78 79            if source_node_id not in edge_mapping:80                edge_mapping[source_node_id] = []81 82            target_node_id = edge_config.get("target")83            if not target_node_id:84                continue85 86            if target_node_id not in reverse_edge_mapping:87                reverse_edge_mapping[target_node_id] = []88 89            target_edge_ids.add(target_node_id)90 91            # parse run condition92            run_condition = None93            if edge_config.get("sourceHandle") and edge_config.get("sourceHandle") != "source":94                run_condition = RunCondition(type="branch_identify", branch_identify=edge_config.get("sourceHandle"))95 96            graph_edge = GraphEdge(97                source_node_id=source_node_id, target_node_id=target_node_id, run_condition=run_condition98            )99 100            edge_mapping[source_node_id].append(graph_edge)101            reverse_edge_mapping[target_node_id].append(graph_edge)102 103        # node configs104        node_configs = graph_config.get("nodes")105        if not node_configs:106            raise ValueError("Graph must have at least one node")107 108        node_configs = cast(list, node_configs)109 110        # fetch nodes that have no predecessor node111        root_node_configs = []112        all_node_id_config_mapping: dict[str, dict] = {}113        for node_config in node_configs:114            node_id = node_config.get("id")115            if not node_id:116                continue117 118            if node_id not in target_edge_ids:119                root_node_configs.append(node_config)120 121            all_node_id_config_mapping[node_id] = node_config122 123        root_node_ids = [node_config.get("id") for node_config in root_node_configs]124 125        # fetch root node126        if not root_node_id:127            # if no root node id, use the START type node as root node128            root_node_id = next(129                (130                    node_config.get("id")131                    for node_config in root_node_configs132                    if node_config.get("data", {}).get("type", "") == NodeType.START.value133                ),134                None,135            )136 137        if not root_node_id or root_node_id not in root_node_ids:138            raise ValueError(f"Root node id {root_node_id} not found in the graph")139 140        # Check whether it is connected to the previous node141        cls._check_connected_to_previous_node(route=[root_node_id], edge_mapping=edge_mapping)142 143        # fetch all node ids from root node144        node_ids = [root_node_id]145        cls._recursively_add_node_ids(node_ids=node_ids, edge_mapping=edge_mapping, node_id=root_node_id)146 147        node_id_config_mapping = {node_id: all_node_id_config_mapping[node_id] for node_id in node_ids}148 149        # init parallel mapping150        parallel_mapping: dict[str, GraphParallel] = {}151        node_parallel_mapping: dict[str, str] = {}152        cls._recursively_add_parallels(153            edge_mapping=edge_mapping,154            reverse_edge_mapping=reverse_edge_mapping,155            start_node_id=root_node_id,156            parallel_mapping=parallel_mapping,157            node_parallel_mapping=node_parallel_mapping,158        )159 160        # Check if it exceeds N layers of parallel161        for parallel in parallel_mapping.values():162            if parallel.parent_parallel_id:163                cls._check_exceed_parallel_limit(164                    parallel_mapping=parallel_mapping, level_limit=3, parent_parallel_id=parallel.parent_parallel_id165                )166 167        # init answer stream generate routes168        answer_stream_generate_routes = AnswerStreamGeneratorRouter.init(169            node_id_config_mapping=node_id_config_mapping, reverse_edge_mapping=reverse_edge_mapping170        )171 172        # init end stream param173        end_stream_param = EndStreamGeneratorRouter.init(174            node_id_config_mapping=node_id_config_mapping,175            reverse_edge_mapping=reverse_edge_mapping,176            node_parallel_mapping=node_parallel_mapping,177        )178 179        # init graph180        graph = cls(181            root_node_id=root_node_id,182            node_ids=node_ids,183            node_id_config_mapping=node_id_config_mapping,184            edge_mapping=edge_mapping,185            reverse_edge_mapping=reverse_edge_mapping,186            parallel_mapping=parallel_mapping,187            node_parallel_mapping=node_parallel_mapping,188            answer_stream_generate_routes=answer_stream_generate_routes,189            end_stream_param=end_stream_param,190        )191 192        return graph193 194    def add_extra_edge(195        self, source_node_id: str, target_node_id: str, run_condition: Optional[RunCondition] = None196    ) -> None:197        """198        Add extra edge to the graph199 200        :param source_node_id: source node id201        :param target_node_id: target node id202        :param run_condition: run condition203        """204        if source_node_id not in self.node_ids or target_node_id not in self.node_ids:205            return206 207        if source_node_id not in self.edge_mapping:208            self.edge_mapping[source_node_id] = []209 210        if target_node_id in [graph_edge.target_node_id for graph_edge in self.edge_mapping[source_node_id]]:211            return212 213        graph_edge = GraphEdge(214            source_node_id=source_node_id, target_node_id=target_node_id, run_condition=run_condition215        )216 217        self.edge_mapping[source_node_id].append(graph_edge)218 219    def get_leaf_node_ids(self) -> list[str]:220        """221        Get leaf node ids of the graph222 223        :return: leaf node ids224        """225        leaf_node_ids = []226        for node_id in self.node_ids:227            if node_id not in self.edge_mapping or (228                len(self.edge_mapping[node_id]) == 1229                and self.edge_mapping[node_id][0].target_node_id == self.root_node_id230            ):231                leaf_node_ids.append(node_id)232 233        return leaf_node_ids234 235    @classmethod236    def _recursively_add_node_ids(237        cls, node_ids: list[str], edge_mapping: dict[str, list[GraphEdge]], node_id: str238    ) -> None:239        """240        Recursively add node ids241 242        :param node_ids: node ids243        :param edge_mapping: edge mapping244        :param node_id: node id245        """246        for graph_edge in edge_mapping.get(node_id, []):247            if graph_edge.target_node_id in node_ids:248                continue249 250            node_ids.append(graph_edge.target_node_id)251            cls._recursively_add_node_ids(252                node_ids=node_ids, edge_mapping=edge_mapping, node_id=graph_edge.target_node_id253            )254 255    @classmethod256    def _check_connected_to_previous_node(cls, route: list[str], edge_mapping: dict[str, list[GraphEdge]]) -> None:257        """258        Check whether it is connected to the previous node259        """260        last_node_id = route[-1]261 262        for graph_edge in edge_mapping.get(last_node_id, []):263            if not graph_edge.target_node_id:264                continue265 266            if graph_edge.target_node_id in route:267                raise ValueError(268                    f"Node {graph_edge.source_node_id} is connected to the previous node, please check the graph."269                )270 271            new_route = route.copy()272            new_route.append(graph_edge.target_node_id)273            cls._check_connected_to_previous_node(274                route=new_route,275                edge_mapping=edge_mapping,276            )277 278    @classmethod279    def _recursively_add_parallels(280        cls,281        edge_mapping: dict[str, list[GraphEdge]],282        reverse_edge_mapping: dict[str, list[GraphEdge]],283        start_node_id: str,284        parallel_mapping: dict[str, GraphParallel],285        node_parallel_mapping: dict[str, str],286        parent_parallel: Optional[GraphParallel] = None,287    ) -> None:288        """289        Recursively add parallel ids290 291        :param edge_mapping: edge mapping292        :param start_node_id: start from node id293        :param parallel_mapping: parallel mapping294        :param node_parallel_mapping: node parallel mapping295        :param parent_parallel: parent parallel296        """297        target_node_edges = edge_mapping.get(start_node_id, [])298        parallel = None299        if len(target_node_edges) > 1:300            # fetch all node ids in current parallels301            parallel_branch_node_ids = {}302            condition_edge_mappings = {}303            for graph_edge in target_node_edges:304                if graph_edge.run_condition is None:305                    if "default" not in parallel_branch_node_ids:306                        parallel_branch_node_ids["default"] = []307 308                    parallel_branch_node_ids["default"].append(graph_edge.target_node_id)309                else:310                    condition_hash = graph_edge.run_condition.hash311                    if condition_hash not in condition_edge_mappings:312                        condition_edge_mappings[condition_hash] = []313 314                    condition_edge_mappings[condition_hash].append(graph_edge)315 316            for condition_hash, graph_edges in condition_edge_mappings.items():317                if len(graph_edges) > 1:318                    if condition_hash not in parallel_branch_node_ids:319                        parallel_branch_node_ids[condition_hash] = []320 321                    for graph_edge in graph_edges:322                        parallel_branch_node_ids[condition_hash].append(graph_edge.target_node_id)323 324            condition_parallels = {}325            for condition_hash, condition_parallel_branch_node_ids in parallel_branch_node_ids.items():326                # any target node id in node_parallel_mapping327                parallel = None328                if condition_parallel_branch_node_ids:329                    parent_parallel_id = parent_parallel.id if parent_parallel else None330 331                    parallel = GraphParallel(332                        start_from_node_id=start_node_id,333                        parent_parallel_id=parent_parallel.id if parent_parallel else None,334                        parent_parallel_start_node_id=parent_parallel.start_from_node_id if parent_parallel else None,335                    )336                    parallel_mapping[parallel.id] = parallel337                    condition_parallels[condition_hash] = parallel338 339                    in_branch_node_ids = cls._fetch_all_node_ids_in_parallels(340                        edge_mapping=edge_mapping,341                        reverse_edge_mapping=reverse_edge_mapping,342                        parallel_branch_node_ids=condition_parallel_branch_node_ids,343                    )344 345                    # collect all branches node ids346                    parallel_node_ids = []347                    for _, node_ids in in_branch_node_ids.items():348                        for node_id in node_ids:349                            in_parent_parallel = True350                            if parent_parallel_id:351                                in_parent_parallel = False352                                for parallel_node_id, parallel_id in node_parallel_mapping.items():353                                    if parallel_id == parent_parallel_id and parallel_node_id == node_id:354                                        in_parent_parallel = True355                                        break356 357                            if in_parent_parallel:358                                parallel_node_ids.append(node_id)359                                node_parallel_mapping[node_id] = parallel.id360 361                    outside_parallel_target_node_ids = set()362                    for node_id in parallel_node_ids:363                        if node_id == parallel.start_from_node_id:364                            continue365 366                        node_edges = edge_mapping.get(node_id)367                        if not node_edges:368                            continue369 370                        if len(node_edges) > 1:371                            continue372 373                        target_node_id = node_edges[0].target_node_id374                        if target_node_id in parallel_node_ids:375                            continue376 377                        if parent_parallel_id:378                            parent_parallel = parallel_mapping.get(parent_parallel_id)379                            if not parent_parallel:380                                continue381 382                        if (383                            (384                                node_parallel_mapping.get(target_node_id)385                                and node_parallel_mapping.get(target_node_id) == parent_parallel_id386                            )387                            or (388                                parent_parallel389                                and parent_parallel.end_to_node_id390                                and target_node_id == parent_parallel.end_to_node_id391                            )392                            or (not node_parallel_mapping.get(target_node_id) and not parent_parallel)393                        ):394                            outside_parallel_target_node_ids.add(target_node_id)395 396                    if len(outside_parallel_target_node_ids) == 1:397                        if (398                            parent_parallel399                            and parent_parallel.end_to_node_id400                            and parallel.end_to_node_id == parent_parallel.end_to_node_id401                        ):402                            parallel.end_to_node_id = None403                        else:404                            parallel.end_to_node_id = outside_parallel_target_node_ids.pop()405 406            if condition_edge_mappings:407                for condition_hash, graph_edges in condition_edge_mappings.items():408                    for graph_edge in graph_edges:409                        current_parallel: GraphParallel | None = cls._get_current_parallel(410                            parallel_mapping=parallel_mapping,411                            graph_edge=graph_edge,412                            parallel=condition_parallels.get(condition_hash),413                            parent_parallel=parent_parallel,414                        )415 416                        cls._recursively_add_parallels(417                            edge_mapping=edge_mapping,418                            reverse_edge_mapping=reverse_edge_mapping,419                            start_node_id=graph_edge.target_node_id,420                            parallel_mapping=parallel_mapping,421                            node_parallel_mapping=node_parallel_mapping,422                            parent_parallel=current_parallel,423                        )424            else:425                for graph_edge in target_node_edges:426                    current_parallel = cls._get_current_parallel(427                        parallel_mapping=parallel_mapping,428                        graph_edge=graph_edge,429                        parallel=parallel,430                        parent_parallel=parent_parallel,431                    )432 433                    cls._recursively_add_parallels(434                        edge_mapping=edge_mapping,435                        reverse_edge_mapping=reverse_edge_mapping,436                        start_node_id=graph_edge.target_node_id,437                        parallel_mapping=parallel_mapping,438                        node_parallel_mapping=node_parallel_mapping,439                        parent_parallel=current_parallel,440                    )441        else:442            for graph_edge in target_node_edges:443                current_parallel = cls._get_current_parallel(444                    parallel_mapping=parallel_mapping,445                    graph_edge=graph_edge,446                    parallel=parallel,447                    parent_parallel=parent_parallel,448                )449 450                cls._recursively_add_parallels(451                    edge_mapping=edge_mapping,452                    reverse_edge_mapping=reverse_edge_mapping,453                    start_node_id=graph_edge.target_node_id,454                    parallel_mapping=parallel_mapping,455                    node_parallel_mapping=node_parallel_mapping,456                    parent_parallel=current_parallel,457                )458 459    @classmethod460    def _get_current_parallel(461        cls,462        parallel_mapping: dict[str, GraphParallel],463        graph_edge: GraphEdge,464        parallel: Optional[GraphParallel] = None,465        parent_parallel: Optional[GraphParallel] = None,466    ) -> Optional[GraphParallel]:467        """468        Get current parallel469        """470        current_parallel = None471        if parallel:472            current_parallel = parallel473        elif parent_parallel:474            if not parent_parallel.end_to_node_id or (475                parent_parallel.end_to_node_id and graph_edge.target_node_id != parent_parallel.end_to_node_id476            ):477                current_parallel = parent_parallel478            else:479                # fetch parent parallel's parent parallel480                parent_parallel_parent_parallel_id = parent_parallel.parent_parallel_id481                if parent_parallel_parent_parallel_id:482                    parent_parallel_parent_parallel = parallel_mapping.get(parent_parallel_parent_parallel_id)483                    if parent_parallel_parent_parallel and (484                        not parent_parallel_parent_parallel.end_to_node_id485                        or (486                            parent_parallel_parent_parallel.end_to_node_id487                            and graph_edge.target_node_id != parent_parallel_parent_parallel.end_to_node_id488                        )489                    ):490                        current_parallel = parent_parallel_parent_parallel491 492        return current_parallel493 494    @classmethod495    def _check_exceed_parallel_limit(496        cls,497        parallel_mapping: dict[str, GraphParallel],498        level_limit: int,499        parent_parallel_id: str,500        current_level: int = 1,501    ) -> None:502        """503        Check if it exceeds N layers of parallel504        """505        parent_parallel = parallel_mapping.get(parent_parallel_id)506        if not parent_parallel:507            return508 509        current_level += 1510        if current_level > level_limit:511            raise ValueError(f"Exceeds {level_limit} layers of parallel")512 513        if parent_parallel.parent_parallel_id:514            cls._check_exceed_parallel_limit(515                parallel_mapping=parallel_mapping,516                level_limit=level_limit,517                parent_parallel_id=parent_parallel.parent_parallel_id,518                current_level=current_level,519            )520 521    @classmethod522    def _recursively_add_parallel_node_ids(523        cls,524        branch_node_ids: list[str],525        edge_mapping: dict[str, list[GraphEdge]],526        merge_node_id: str,527        start_node_id: str,528    ) -> None:529        """530        Recursively add node ids531 532        :param branch_node_ids: in branch node ids533        :param edge_mapping: edge mapping534        :param merge_node_id: merge node id535        :param start_node_id: start node id536        """537        for graph_edge in edge_mapping.get(start_node_id, []):538            if graph_edge.target_node_id != merge_node_id and graph_edge.target_node_id not in branch_node_ids:539                branch_node_ids.append(graph_edge.target_node_id)540                cls._recursively_add_parallel_node_ids(541                    branch_node_ids=branch_node_ids,542                    edge_mapping=edge_mapping,543                    merge_node_id=merge_node_id,544                    start_node_id=graph_edge.target_node_id,545                )546 547    @classmethod548    def _fetch_all_node_ids_in_parallels(549        cls,550        edge_mapping: dict[str, list[GraphEdge]],551        reverse_edge_mapping: dict[str, list[GraphEdge]],552        parallel_branch_node_ids: list[str],553    ) -> dict[str, list[str]]:554        """555        Fetch all node ids in parallels556        """557        routes_node_ids: dict[str, list[str]] = {}558        for parallel_branch_node_id in parallel_branch_node_ids:559            routes_node_ids[parallel_branch_node_id] = [parallel_branch_node_id]560 561            # fetch routes node ids562            cls._recursively_fetch_routes(563                edge_mapping=edge_mapping,564                start_node_id=parallel_branch_node_id,565                routes_node_ids=routes_node_ids[parallel_branch_node_id],566            )567 568        # fetch leaf node ids from routes node ids569        leaf_node_ids: dict[str, list[str]] = {}570        merge_branch_node_ids: dict[str, list[str]] = {}571        for branch_node_id, node_ids in routes_node_ids.items():572            for node_id in node_ids:573                if node_id not in edge_mapping or len(edge_mapping[node_id]) == 0:574                    if branch_node_id not in leaf_node_ids:575                        leaf_node_ids[branch_node_id] = []576 577                    leaf_node_ids[branch_node_id].append(node_id)578 579                for branch_node_id2, inner_route2 in routes_node_ids.items():580                    if (581                        branch_node_id != branch_node_id2582                        and node_id in inner_route2583                        and len(reverse_edge_mapping.get(node_id, [])) > 1584                        and cls._is_node_in_routes(585                            reverse_edge_mapping=reverse_edge_mapping,586                            start_node_id=node_id,587                            routes_node_ids=routes_node_ids,588                        )589                    ):590                        if node_id not in merge_branch_node_ids:591                            merge_branch_node_ids[node_id] = []592 593                        if branch_node_id2 not in merge_branch_node_ids[node_id]:594                            merge_branch_node_ids[node_id].append(branch_node_id2)595 596        # sorted merge_branch_node_ids by branch_node_ids length desc597        merge_branch_node_ids = dict(sorted(merge_branch_node_ids.items(), key=lambda x: len(x[1]), reverse=True))598 599        duplicate_end_node_ids = {}600        for node_id, branch_node_ids in merge_branch_node_ids.items():601            for node_id2, branch_node_ids2 in merge_branch_node_ids.items():602                if node_id != node_id2 and set(branch_node_ids) == set(branch_node_ids2):603                    if (node_id, node_id2) not in duplicate_end_node_ids and (604                        node_id2,605                        node_id,606                    ) not in duplicate_end_node_ids:607                        duplicate_end_node_ids[(node_id, node_id2)] = branch_node_ids608 609        for (node_id, node_id2), branch_node_ids in duplicate_end_node_ids.items():610            # check which node is after611            if cls._is_node2_after_node1(node1_id=node_id, node2_id=node_id2, edge_mapping=edge_mapping):612                if node_id in merge_branch_node_ids:613                    del merge_branch_node_ids[node_id2]614            elif cls._is_node2_after_node1(node1_id=node_id2, node2_id=node_id, edge_mapping=edge_mapping):615                if node_id2 in merge_branch_node_ids:616                    del merge_branch_node_ids[node_id]617 618        branches_merge_node_ids: dict[str, str] = {}619        for node_id, branch_node_ids in merge_branch_node_ids.items():620            if len(branch_node_ids) <= 1:621                continue622 623            for branch_node_id in branch_node_ids:624                if branch_node_id in branches_merge_node_ids:625                    continue626 627                branches_merge_node_ids[branch_node_id] = node_id628 629        in_branch_node_ids: dict[str, list[str]] = {}630        for branch_node_id, node_ids in routes_node_ids.items():631            in_branch_node_ids[branch_node_id] = []632            if branch_node_id not in branches_merge_node_ids:633                # all node ids in current branch is in this thread634                in_branch_node_ids[branch_node_id].append(branch_node_id)635                in_branch_node_ids[branch_node_id].extend(node_ids)636            else:637                merge_node_id = branches_merge_node_ids[branch_node_id]638                if merge_node_id != branch_node_id:639                    in_branch_node_ids[branch_node_id].append(branch_node_id)640 641                # fetch all node ids from branch_node_id and merge_node_id642                cls._recursively_add_parallel_node_ids(643                    branch_node_ids=in_branch_node_ids[branch_node_id],644                    edge_mapping=edge_mapping,645                    merge_node_id=merge_node_id,646                    start_node_id=branch_node_id,647                )648 649        return in_branch_node_ids650 651    @classmethod652    def _recursively_fetch_routes(653        cls, edge_mapping: dict[str, list[GraphEdge]], start_node_id: str, routes_node_ids: list[str]654    ) -> None:655        """656        Recursively fetch route657        """658        if start_node_id not in edge_mapping:659            return660 661        for graph_edge in edge_mapping[start_node_id]:662            # find next node ids663            if graph_edge.target_node_id not in routes_node_ids:664                routes_node_ids.append(graph_edge.target_node_id)665 666                cls._recursively_fetch_routes(667                    edge_mapping=edge_mapping, start_node_id=graph_edge.target_node_id, routes_node_ids=routes_node_ids668                )669 670    @classmethod671    def _is_node_in_routes(672        cls, reverse_edge_mapping: dict[str, list[GraphEdge]], start_node_id: str, routes_node_ids: dict[str, list[str]]673    ) -> bool:674        """675        Recursively check if the node is in the routes676        """677        if start_node_id not in reverse_edge_mapping:678            return False679 680        all_routes_node_ids = set()681        parallel_start_node_ids: dict[str, list[str]] = {}682        for branch_node_id, node_ids in routes_node_ids.items():683            all_routes_node_ids.update(node_ids)684 685            if branch_node_id in reverse_edge_mapping:686                for graph_edge in reverse_edge_mapping[branch_node_id]:687                    if graph_edge.source_node_id not in parallel_start_node_ids:688                        parallel_start_node_ids[graph_edge.source_node_id] = []689 690                    parallel_start_node_ids[graph_edge.source_node_id].append(branch_node_id)691 692        for _, branch_node_ids in parallel_start_node_ids.items():693            if set(branch_node_ids) == set(routes_node_ids.keys()):694                return True695 696        return False697 698    @classmethod699    def _is_node2_after_node1(cls, node1_id: str, node2_id: str, edge_mapping: dict[str, list[GraphEdge]]) -> bool:700        """701        is node2 after node1702        """703        if node1_id not in edge_mapping:704            return False705 706        for graph_edge in edge_mapping[node1_id]:707            if graph_edge.target_node_id == node2_id:708                return True709 710            if cls._is_node2_after_node1(711                node1_id=graph_edge.target_node_id, node2_id=node2_id, edge_mapping=edge_mapping712            ):713                return True714 715        return False716