Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
graph.py740 linesDownload Raw Back to runnables
1"""Graph used in `Runnable` objects."""2 3from __future__ import annotations4 5import inspect6from collections import defaultdict7from dataclasses import dataclass, field8from enum import Enum9from typing import (10    TYPE_CHECKING,11    Any,12    NamedTuple,13    Protocol,14    TypedDict,15    overload,16)17from uuid import UUID, uuid418 19from langchain_core.load.serializable import to_json_not_implemented20from langchain_core.runnables.base import Runnable, RunnableSerializable21from langchain_core.utils.pydantic import _IgnoreUnserializable, is_basemodel_subclass22 23if TYPE_CHECKING:24    from collections.abc import Callable, Sequence25 26    from pydantic import BaseModel27 28    from langchain_core.runnables.base import Runnable as RunnableType29 30 31class Stringifiable(Protocol):32    """Protocol for objects that can be converted to a string."""33 34    def __str__(self) -> str:35        """Convert the object to a string."""36 37 38class LabelsDict(TypedDict):39    """Dictionary of labels for nodes and edges in a graph."""40 41    nodes: dict[str, str]42    """Labels for nodes."""43    edges: dict[str, str]44    """Labels for edges."""45 46 47def is_uuid(value: str) -> bool:48    """Check if a string is a valid UUID.49 50    Args:51        value: The string to check.52 53    Returns:54        `True` if the string is a valid UUID, `False` otherwise.55    """56    try:57        UUID(value)58    except ValueError:59        return False60    return True61 62 63class Edge(NamedTuple):64    """Edge in a graph."""65 66    source: str67    """The source node id."""68    target: str69    """The target node id."""70    data: Stringifiable | None = None71    """Optional data associated with the edge. """72    conditional: bool = False73    """Whether the edge is conditional."""74 75    def copy(self, *, source: str | None = None, target: str | None = None) -> Edge:76        """Return a copy of the edge with optional new source and target nodes.77 78        Args:79            source: The new source node id.80            target: The new target node id.81 82        Returns:83            A copy of the edge with the new source and target nodes.84        """85        return Edge(86            source=source or self.source,87            target=target or self.target,88            data=self.data,89            conditional=self.conditional,90        )91 92 93class Node(NamedTuple):94    """Node in a graph."""95 96    id: str97    """The unique identifier of the node."""98    name: str99    """The name of the node."""100    data: type[BaseModel] | RunnableType | None101    """The data of the node."""102    metadata: dict[str, Any] | None103    """Optional metadata for the node. """104 105    def copy(106        self,107        *,108        id: str | None = None,109        name: str | None = None,110    ) -> Node:111        """Return a copy of the node with optional new id and name.112 113        Args:114            id: The new node id.115            name: The new node name.116 117        Returns:118            A copy of the node with the new id and name.119        """120        return Node(121            id=id or self.id,122            name=name or self.name,123            data=self.data,124            metadata=self.metadata,125        )126 127 128class Branch(NamedTuple):129    """Branch in a graph."""130 131    condition: Callable[..., str]132    """A callable that returns a string representation of the condition."""133    ends: dict[str, str] | None134    """Optional dictionary of end node IDs for the branches. """135 136 137class CurveStyle(Enum):138    """Enum for different curve styles supported by Mermaid."""139 140    BASIS = "basis"141    BUMP_X = "bumpX"142    BUMP_Y = "bumpY"143    CARDINAL = "cardinal"144    CATMULL_ROM = "catmullRom"145    LINEAR = "linear"146    MONOTONE_X = "monotoneX"147    MONOTONE_Y = "monotoneY"148    NATURAL = "natural"149    STEP = "step"150    STEP_AFTER = "stepAfter"151    STEP_BEFORE = "stepBefore"152 153 154@dataclass155class NodeStyles:156    """Schema for Hexadecimal color codes for different node types.157 158    Args:159        default: The default color code.160        first: The color code for the first node.161        last: The color code for the last node.162    """163 164    default: str = "fill:#f2f0ff,line-height:1.2"165    first: str = "fill-opacity:0"166    last: str = "fill:#bfb6fc"167 168 169class MermaidDrawMethod(Enum):170    """Enum for different draw methods supported by Mermaid."""171 172    PYPPETEER = "pyppeteer"173    """Uses Pyppeteer to render the graph"""174    API = "api"175    """Uses Mermaid.INK API to render the graph"""176 177 178def node_data_str(179    id: str,180    data: type[BaseModel] | RunnableType | None,181) -> str:182    """Convert the data of a node to a string.183 184    Args:185        id: The node id.186        data: The node data.187 188    Returns:189        A string representation of the data.190    """191    if not is_uuid(id) or data is None:192        return id193    data_str = data.get_name() if isinstance(data, Runnable) else data.__name__194    return data_str if not data_str.startswith("Runnable") else data_str[8:]195 196 197def node_data_json(198    node: Node, *, with_schemas: bool = False199) -> dict[str, str | dict[str, Any]]:200    """Convert the data of a node to a JSON-serializable format.201 202    Args:203        node: The `Node` to convert.204        with_schemas: Whether to include the schema of the data if it is a Pydantic205            model.206 207    Returns:208        A dictionary with the type of the data and the data itself.209    """210    if node.data is None:211        json: dict[str, Any] = {}212    elif isinstance(node.data, RunnableSerializable):213        json = {214            "type": "runnable",215            "data": {216                "id": node.data.lc_id(),217                "name": node_data_str(node.id, node.data),218            },219        }220    elif isinstance(node.data, Runnable):221        json = {222            "type": "runnable",223            "data": {224                "id": to_json_not_implemented(node.data)["id"],225                "name": node_data_str(node.id, node.data),226            },227        }228    elif inspect.isclass(node.data) and is_basemodel_subclass(node.data):229        json = (230            {231                "type": "schema",232                "data": node.data.model_json_schema(233                    schema_generator=_IgnoreUnserializable234                ),235            }236            if with_schemas237            else {238                "type": "schema",239                "data": node_data_str(node.id, node.data),240            }241        )242    else:243        json = {244            "type": "unknown",245            "data": node_data_str(node.id, node.data),246        }247    if node.metadata is not None:248        json["metadata"] = node.metadata249    return json250 251 252@dataclass253class Graph:254    """Graph of nodes and edges.255 256    Args:257        nodes: Dictionary of nodes in the graph. Defaults to an empty dictionary.258        edges: List of edges in the graph. Defaults to an empty list.259    """260 261    nodes: dict[str, Node] = field(default_factory=dict)262    edges: list[Edge] = field(default_factory=list)263 264    def to_json(self, *, with_schemas: bool = False) -> dict[str, list[dict[str, Any]]]:265        """Convert the graph to a JSON-serializable format.266 267        Args:268            with_schemas: Whether to include the schemas of the nodes if they are269                Pydantic models.270 271        Returns:272            A dictionary with the nodes and edges of the graph.273        """274        stable_node_ids = {275            node.id: i if is_uuid(node.id) else node.id276            for i, node in enumerate(self.nodes.values())277        }278        edges: list[dict[str, Any]] = []279        for edge in self.edges:280            edge_dict = {281                "source": stable_node_ids[edge.source],282                "target": stable_node_ids[edge.target],283            }284            if edge.data is not None:285                edge_dict["data"] = edge.data  # type: ignore[assignment]286            if edge.conditional:287                edge_dict["conditional"] = True288            edges.append(edge_dict)289 290        return {291            "nodes": [292                {293                    "id": stable_node_ids[node.id],294                    **node_data_json(node, with_schemas=with_schemas),295                }296                for node in self.nodes.values()297            ],298            "edges": edges,299        }300 301    def __bool__(self) -> bool:302        """Return whether the graph has any nodes."""303        return bool(self.nodes)304 305    def next_id(self) -> str:306        """Return a new unique node identifier.307 308        It that can be used to add a node to the graph.309        """310        return uuid4().hex311 312    def add_node(313        self,314        data: type[BaseModel] | RunnableType | None,315        id: str | None = None,316        *,317        metadata: dict[str, Any] | None = None,318    ) -> Node:319        """Add a node to the graph and return it.320 321        Args:322            data: The data of the node.323            id: The id of the node.324            metadata: Optional metadata for the node.325 326        Returns:327            The node that was added to the graph.328 329        Raises:330            ValueError: If a node with the same id already exists.331        """332        if id is not None and id in self.nodes:333            msg = f"Node with id {id} already exists"334            raise ValueError(msg)335        id_ = id or self.next_id()336        node = Node(id=id_, data=data, metadata=metadata, name=node_data_str(id_, data))337        self.nodes[node.id] = node338        return node339 340    def remove_node(self, node: Node) -> None:341        """Remove a node from the graph and all edges connected to it.342 343        Args:344            node: The node to remove.345        """346        self.nodes.pop(node.id)347        self.edges = [348            edge for edge in self.edges if node.id not in {edge.source, edge.target}349        ]350 351    def add_edge(352        self,353        source: Node,354        target: Node,355        data: Stringifiable | None = None,356        conditional: bool = False,  # noqa: FBT001,FBT002357    ) -> Edge:358        """Add an edge to the graph and return it.359 360        Args:361            source: The source node of the edge.362            target: The target node of the edge.363            data: Optional data associated with the edge.364            conditional: Whether the edge is conditional.365 366        Returns:367            The edge that was added to the graph.368 369        Raises:370            ValueError: If the source or target node is not in the graph.371        """372        if source.id not in self.nodes:373            msg = f"Source node {source.id} not in graph"374            raise ValueError(msg)375        if target.id not in self.nodes:376            msg = f"Target node {target.id} not in graph"377            raise ValueError(msg)378        edge = Edge(379            source=source.id, target=target.id, data=data, conditional=conditional380        )381        self.edges.append(edge)382        return edge383 384    def extend(385        self, graph: Graph, *, prefix: str = ""386    ) -> tuple[Node | None, Node | None]:387        """Add all nodes and edges from another graph.388 389        Note this doesn't check for duplicates, nor does it connect the graphs.390 391        Args:392            graph: The graph to add.393            prefix: The prefix to add to the node ids.394 395        Returns:396            A tuple of the first and last nodes of the subgraph.397        """398        if all(is_uuid(node.id) for node in graph.nodes.values()):399            prefix = ""400 401        def prefixed(id_: str) -> str:402            return f"{prefix}:{id_}" if prefix else id_403 404        # prefix each node405        self.nodes.update(406            {prefixed(k): v.copy(id=prefixed(k)) for k, v in graph.nodes.items()}407        )408        # prefix each edge's source and target409        self.edges.extend(410            [411                edge.copy(source=prefixed(edge.source), target=prefixed(edge.target))412                for edge in graph.edges413            ]414        )415        # return (prefixed) first and last nodes of the subgraph416        first, last = graph.first_node(), graph.last_node()417        return (418            first.copy(id=prefixed(first.id)) if first else None,419            last.copy(id=prefixed(last.id)) if last else None,420        )421 422    def reid(self) -> Graph:423        """Return a new graph with all nodes re-identified.424 425        Uses their unique, readable names where possible.426        """427        node_name_to_ids = defaultdict(list)428        for node in self.nodes.values():429            node_name_to_ids[node.name].append(node.id)430 431        unique_labels = {432            node_id: node_name if len(node_ids) == 1 else f"{node_name}_{i + 1}"433            for node_name, node_ids in node_name_to_ids.items()434            for i, node_id in enumerate(node_ids)435        }436 437        def _get_node_id(node_id: str) -> str:438            label = unique_labels[node_id]439            if is_uuid(node_id):440                return label441            return node_id442 443        return Graph(444            nodes={445                _get_node_id(id_): node.copy(id=_get_node_id(id_))446                for id_, node in self.nodes.items()447            },448            edges=[449                edge.copy(450                    source=_get_node_id(edge.source),451                    target=_get_node_id(edge.target),452                )453                for edge in self.edges454            ],455        )456 457    def first_node(self) -> Node | None:458        """Find the single node that is not a target of any edge.459 460        If there is no such node, or there are multiple, return `None`.461        When drawing the graph, this node would be the origin.462 463        Returns:464            The first node, or None if there is no such node or multiple465            candidates.466        """467        return _first_node(self)468 469    def last_node(self) -> Node | None:470        """Find the single node that is not a source of any edge.471 472        If there is no such node, or there are multiple, return `None`.473        When drawing the graph, this node would be the destination.474 475        Returns:476            The last node, or None if there is no such node or multiple477            candidates.478        """479        return _last_node(self)480 481    def trim_first_node(self) -> None:482        """Remove the first node if it exists and has a single outgoing edge.483 484        i.e., if removing it would not leave the graph without a "first" node.485        """486        first_node = self.first_node()487        if (488            first_node489            and _first_node(self, exclude=[first_node.id])490            and len({e for e in self.edges if e.source == first_node.id}) == 1491        ):492            self.remove_node(first_node)493 494    def trim_last_node(self) -> None:495        """Remove the last node if it exists and has a single incoming edge.496 497        i.e., if removing it would not leave the graph without a "last" node.498        """499        last_node = self.last_node()500        if (501            last_node502            and _last_node(self, exclude=[last_node.id])503            and len({e for e in self.edges if e.target == last_node.id}) == 1504        ):505            self.remove_node(last_node)506 507    def draw_ascii(self) -> str:508        """Draw the graph as an ASCII art string.509 510        Returns:511            The ASCII art string.512        """513        # Import locally to prevent circular import514        from langchain_core.runnables.graph_ascii import draw_ascii  # noqa: PLC0415515 516        return draw_ascii(517            {node.id: node.name for node in self.nodes.values()},518            self.edges,519        )520 521    def print_ascii(self) -> None:522        """Print the graph as an ASCII art string."""523        print(self.draw_ascii())  # noqa: T201524 525    @overload526    def draw_png(527        self,528        output_file_path: str,529        fontname: str | None = None,530        labels: LabelsDict | None = None,531    ) -> None: ...532 533    @overload534    def draw_png(535        self,536        output_file_path: None,537        fontname: str | None = None,538        labels: LabelsDict | None = None,539    ) -> bytes: ...540 541    def draw_png(542        self,543        output_file_path: str | None = None,544        fontname: str | None = None,545        labels: LabelsDict | None = None,546    ) -> bytes | None:547        """Draw the graph as a PNG image.548 549        Args:550            output_file_path: The path to save the image to. If `None`, the image551                is not saved.552            fontname: The name of the font to use.553            labels: Optional labels for nodes and edges in the graph. Defaults to554                `None`.555 556        Returns:557            The PNG image as bytes if output_file_path is None, None otherwise.558        """559        # Import locally to prevent circular import560        from langchain_core.runnables.graph_png import PngDrawer  # noqa: PLC0415561 562        default_node_labels = {node.id: node.name for node in self.nodes.values()}563 564        return PngDrawer(565            fontname,566            LabelsDict(567                nodes={568                    **default_node_labels,569                    **(labels["nodes"] if labels is not None else {}),570                },571                edges=labels["edges"] if labels is not None else {},572            ),573        ).draw(self, output_file_path)574 575    def draw_mermaid(576        self,577        *,578        with_styles: bool = True,579        curve_style: CurveStyle = CurveStyle.LINEAR,580        node_colors: NodeStyles | None = None,581        wrap_label_n_words: int = 9,582        frontmatter_config: dict[str, Any] | None = None,583    ) -> str:584        """Draw the graph as a Mermaid syntax string.585 586        Args:587            with_styles: Whether to include styles in the syntax.588            curve_style: The style of the edges.589            node_colors: The colors of the nodes.590            wrap_label_n_words: The number of words to wrap the node labels at.591            frontmatter_config: Mermaid frontmatter config.592                Can be used to customize theme and styles. Will be converted to YAML and593                added to the beginning of the mermaid graph.594 595                See more here: https://mermaid.js.org/config/configuration.html.596 597                Example config:598 599                ```python600                {601                    "config": {602                        "theme": "neutral",603                        "look": "handDrawn",604                        "themeVariables": {"primaryColor": "#e2e2e2"},605                    }606                }607                ```608        Returns:609            The Mermaid syntax string.610        """611        # Import locally to prevent circular import612        from langchain_core.runnables.graph_mermaid import draw_mermaid  # noqa: PLC0415613 614        graph = self.reid()615        first_node = graph.first_node()616        last_node = graph.last_node()617 618        return draw_mermaid(619            nodes=graph.nodes,620            edges=graph.edges,621            first_node=first_node.id if first_node else None,622            last_node=last_node.id if last_node else None,623            with_styles=with_styles,624            curve_style=curve_style,625            node_styles=node_colors,626            wrap_label_n_words=wrap_label_n_words,627            frontmatter_config=frontmatter_config,628        )629 630    def draw_mermaid_png(631        self,632        *,633        curve_style: CurveStyle = CurveStyle.LINEAR,634        node_colors: NodeStyles | None = None,635        wrap_label_n_words: int = 9,636        output_file_path: str | None = None,637        draw_method: MermaidDrawMethod = MermaidDrawMethod.API,638        background_color: str = "white",639        padding: int = 10,640        max_retries: int = 1,641        retry_delay: float = 1.0,642        frontmatter_config: dict[str, Any] | None = None,643        base_url: str | None = None,644        proxies: dict[str, str] | None = None,645    ) -> bytes:646        """Draw the graph as a PNG image using Mermaid.647 648        Args:649            curve_style: The style of the edges.650            node_colors: The colors of the nodes.651            wrap_label_n_words: The number of words to wrap the node labels at.652            output_file_path: The path to save the image to. If `None`, the image653                is not saved.654            draw_method: The method to use to draw the graph.655            background_color: The color of the background.656            padding: The padding around the graph.657            max_retries: The maximum number of retries (`MermaidDrawMethod.API`).658            retry_delay: The delay between retries (`MermaidDrawMethod.API`).659            frontmatter_config: Mermaid frontmatter config.660                Can be used to customize theme and styles. Will be converted to YAML and661                added to the beginning of the mermaid graph.662 663                See more here: https://mermaid.js.org/config/configuration.html.664 665                Example config:666 667                ```python668                {669                    "config": {670                        "theme": "neutral",671                        "look": "handDrawn",672                        "themeVariables": {"primaryColor": "#e2e2e2"},673                    }674                }675                ```676            base_url: The base URL of the Mermaid server for rendering via API.677            proxies: HTTP/HTTPS proxies for requests (e.g. `{"http": "http://127.0.0.1:7890"}`).678 679        Returns:680            The PNG image as bytes.681        """682        # Import locally to prevent circular import683        from langchain_core.runnables.graph_mermaid import (  # noqa: PLC0415684            draw_mermaid_png,685        )686 687        mermaid_syntax = self.draw_mermaid(688            curve_style=curve_style,689            node_colors=node_colors,690            wrap_label_n_words=wrap_label_n_words,691            frontmatter_config=frontmatter_config,692        )693        return draw_mermaid_png(694            mermaid_syntax=mermaid_syntax,695            output_file_path=output_file_path,696            draw_method=draw_method,697            background_color=background_color,698            padding=padding,699            max_retries=max_retries,700            retry_delay=retry_delay,701            proxies=proxies,702            base_url=base_url,703        )704 705 706def _first_node(graph: Graph, exclude: Sequence[str] = ()) -> Node | None:707    """Find the single node that is not a target of any edge.708 709    Exclude nodes/sources with IDs in the exclude list.710 711    If there is no such node, or there are multiple, return `None`.712 713    When drawing the graph, this node would be the origin.714    """715    targets = {edge.target for edge in graph.edges if edge.source not in exclude}716    found: list[Node] = [717        node718        for node in graph.nodes.values()719        if node.id not in exclude and node.id not in targets720    ]721    return found[0] if len(found) == 1 else None722 723 724def _last_node(graph: Graph, exclude: Sequence[str] = ()) -> Node | None:725    """Find the single node that is not a source of any edge.726 727    Exclude nodes/targets with IDs in the exclude list.728 729    If there is no such node, or there are multiple, return `None`.730 731    When drawing the graph, this node would be the destination.732    """733    sources = {edge.source for edge in graph.edges if edge.target not in exclude}734    found: list[Node] = [735        node736        for node in graph.nodes.values()737        if node.id not in exclude and node.id not in sources738    ]739    return found[0] if len(found) == 1 else None740 
codekingpro/portable-devtools · Team Ai