codekingpro/portable-devtools
114k
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 