From 3d20c8dc38e186b87e72f780391e19976d46cd73 Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:50:38 -0300 Subject: [PATCH 1/9] refactor: removes most of the circular dependencies in the Graph --- src/backend/langflow/graph/edge/base.py | 38 ++-- src/backend/langflow/graph/graph/base.py | 220 ++++++++++++---------- src/backend/langflow/graph/vertex/base.py | 90 ++++----- 3 files changed, 170 insertions(+), 178 deletions(-) diff --git a/src/backend/langflow/graph/edge/base.py b/src/backend/langflow/graph/edge/base.py index f9d77741b..c9c1077f6 100644 --- a/src/backend/langflow/graph/edge/base.py +++ b/src/backend/langflow/graph/edge/base.py @@ -1,7 +1,7 @@ +from typing import TYPE_CHECKING, List, Optional + from loguru import logger -from typing import TYPE_CHECKING from pydantic import BaseModel, Field -from typing import List, Optional if TYPE_CHECKING: from langflow.graph.vertex.base import Vertex @@ -22,8 +22,8 @@ class TargetHandle(BaseModel): class Edge: def __init__(self, source: "Vertex", target: "Vertex", edge: dict): - self.source: "Vertex" = source - self.target: "Vertex" = target + self.source_id: str = source.id + self.target_id: str = target.id if data := edge.get("data", {}): self._source_handle = data.get("sourceHandle", {}) self._target_handle = data.get("targetHandle", {}) @@ -31,7 +31,7 @@ class Edge: self.target_handle: TargetHandle = TargetHandle(**self._target_handle) self.target_param = self.target_handle.fieldName # validate handles - self.validate_handles() + self.validate_handles(source, target) else: # Logging here because this is a breaking change logger.error("Edge data is empty") @@ -41,9 +41,9 @@ class Edge: # target_param is documents self.target_param = self._target_handle.split("|")[1] # Validate in __init__ to fail fast - self.validate_edge() + self.validate_edge(source, target) - def validate_handles(self) -> None: + def validate_handles(self, source, target) -> None: if self.target_handle.inputTypes is None: self.valid_handles = self.target_handle.type in self.source_handle.baseClasses else: @@ -54,26 +54,20 @@ class Edge: if not self.valid_handles: logger.debug(self.source_handle) logger.debug(self.target_handle) - raise ValueError( - f"Edge between {self.source.vertex_type} and {self.target.vertex_type} " f"has invalid handles" - ) + raise ValueError(f"Edge between {source.vertex_type} and {target.vertex_type} " f"has invalid handles") def __setstate__(self, state): - self.source = state["source"] - self.target = state["target"] + self.source_id = state["source_id"] + self.target_id = state["target_id"] self.target_param = state["target_param"] self.source_handle = state.get("source_handle") self.target_handle = state.get("target_handle") - def reset(self) -> None: - self.source._build_params() - self.target._build_params() - - def validate_edge(self) -> None: + def validate_edge(self, source, target) -> None: # Validate that the outputs of the source node are valid inputs # for the target node - self.source_types = self.source.output - self.target_reqs = self.target.required_inputs + self.target.optional_inputs + self.source_types = source.output + self.target_reqs = target.required_inputs + target.optional_inputs # Both lists contain strings and sometimes a string contains the value we are # looking for e.g. comgin_out=["Chain"] and target_reqs=["LLMChain"] # so we need to check if any of the strings in source_types is in target_reqs @@ -88,13 +82,11 @@ class Edge: if no_matched_type: logger.debug(self.source_types) logger.debug(self.target_reqs) - raise ValueError( - f"Edge between {self.source.vertex_type} and {self.target.vertex_type} " f"has no matched type" - ) + raise ValueError(f"Edge between {source.vertex_type} and {target.vertex_type} " f"has no matched type") def __repr__(self) -> str: return ( - f"Edge(source={self.source.id}, target={self.target.id}, target_param={self.target_param}" + f"Edge(source={self.source_id}, target={self.target_id}, target_param={self.target_param}" f", matched_type={self.matched_type})" ) diff --git a/src/backend/langflow/graph/graph/base.py b/src/backend/langflow/graph/graph/base.py index 48e2bac30..9f80907a8 100644 --- a/src/backend/langflow/graph/graph/base.py +++ b/src/backend/langflow/graph/graph/base.py @@ -13,24 +13,24 @@ from langflow.utils import payload class Graph: - """A class representing a graph of nodes and edges.""" + """A class representing a graph of vertices and edges.""" def __init__( self, nodes: List[Dict], edges: List[Dict[str, str]], ) -> None: - self._nodes = nodes + self._vertices = nodes self._edges = edges self.raw_graph_data = {"nodes": nodes, "edges": edges} - self.top_level_nodes = [] - for node in self._nodes: - if node_id := node.get("id"): - self.top_level_nodes.append(node_id) + self.top_level_vertices = [] + for vertex in self._vertices: + if vertex_id := vertex.get("id"): + self.top_level_vertices.append(vertex_id) self._graph_data = process_flow(self.raw_graph_data) - self._nodes = self._graph_data["nodes"] + self._vertices = self._graph_data["nodes"] self._edges = self._graph_data["edges"] self._build_graph() @@ -54,9 +54,9 @@ class Graph: if "data" in payload: payload = payload["data"] try: - nodes = payload["nodes"] + vertices = payload["nodes"] edges = payload["edges"] - return cls(nodes, edges) + return cls(vertices, edges) except KeyError as exc: logger.exception(exc) raise ValueError( @@ -69,61 +69,69 @@ class Graph: return self.__repr__() == other.__repr__() def _build_graph(self) -> None: - """Builds the graph from the nodes and edges.""" - self.nodes = self._build_vertices() + """Builds the graph from the vertices and edges.""" + self.vertices = self._build_vertices() + self.vertex_ids = [vertex.id for vertex in self.vertices] self.edges = self._build_edges() - for edge in self.edges: - edge.source.add_edge(edge) - edge.target.add_edge(edge) - # This is a hack to make sure that the LLM node is sent to - # the toolkit node - self._build_node_params() - # remove invalid nodes - self._validate_nodes() + # This is a hack to make sure that the LLM vertex is sent to + # the toolkit vertex + self._build_vertex_params() + # remove invalid vertices + self._validate_vertices() - def _build_node_params(self) -> None: - """Identifies and handles the LLM node within the graph.""" - llm_node = None - for node in self.nodes: - node._build_params() - if isinstance(node, LLMVertex): - llm_node = node + def _build_vertex_params(self) -> None: + """Identifies and handles the LLM vertex within the graph.""" + llm_vertex = None + for vertex in self.vertices: + vertex._build_params() + if isinstance(vertex, LLMVertex): + llm_vertex = vertex - if llm_node: - for node in self.nodes: - if isinstance(node, ToolkitVertex): - node.params["llm"] = llm_node + if llm_vertex: + for vertex in self.vertices: + if isinstance(vertex, ToolkitVertex): + vertex.params["llm"] = llm_vertex - def _validate_nodes(self) -> None: - """Check that all nodes have edges""" - if len(self.nodes) == 1: + def _validate_vertices(self) -> None: + """Check that all vertices have edges""" + if len(self.vertices) == 1: return - for node in self.nodes: - if not self._validate_node(node): - raise ValueError(f"{node.vertex_type} is not connected to any other components") + for vertex in self.vertices: + if not self._validate_vertex(vertex): + raise ValueError(f"{vertex.vertex_type} is not connected to any other components") - def _validate_node(self, node: Vertex) -> bool: - """Validates a node.""" - # All nodes that do not have edges are invalid - return len(node.edges) > 0 + def _validate_vertex(self, vertex: Vertex) -> bool: + """Validates a vertex.""" + # All vertices that do not have edges are invalid + return len(self.get_vertex_edges(vertex.id)) > 0 - def get_node(self, node_id: str) -> Union[None, Vertex]: - """Returns a node by id.""" - return next((node for node in self.nodes if node.id == node_id), None) + def get_vertex(self, vertex_id: str) -> Union[None, Vertex]: + """Returns a vertex by id.""" + return next((vertex for vertex in self.vertices if vertex.id == vertex_id), None) - def get_nodes_with_target(self, node: Vertex) -> List[Vertex]: - """Returns the nodes connected to a node.""" - connected_nodes: List[Vertex] = [edge.source for edge in self.edges if edge.target == node] - return connected_nodes + def get_vertex_edges(self, vertex_id: str) -> List[Edge]: + """Returns a list of edges for a given vertex.""" + return [edge for edge in self.edges if edge.source_id == vertex_id or edge.target_id == vertex_id] + + def get_vertices_with_target(self, vertex_id: str) -> List[Vertex]: + """Returns the vertices connected to a vertex.""" + vertices: List[Vertex] = [] + for edge in self.edges: + if edge.target_id == vertex_id: + vertex = self.get_vertex(edge.source_id) + if vertex is None: + continue + vertices.append(vertex) + return vertices async def build(self) -> Chain: """Builds the graph.""" - # Get root node - root_node = payload.get_root_node(self) - if root_node is None: - raise ValueError("No root node found") - return await root_node.build() + # Get root vertex + root_vertex = payload.get_root_vertex(self) + if root_vertex is None: + raise ValueError("No root vertex found") + return await root_vertex.build() def topological_sort(self) -> List[Vertex]: """ @@ -136,25 +144,25 @@ class Graph: ValueError: If the graph contains a cycle. """ # States: 0 = unvisited, 1 = visiting, 2 = visited - state = {node: 0 for node in self.nodes} + state = {vertex: 0 for vertex in self.vertices} sorted_vertices = [] - def dfs(node): - if state[node] == 1: + def dfs(vertex): + if state[vertex] == 1: # We have a cycle raise ValueError("Graph contains a cycle, cannot perform topological sort") - if state[node] == 0: - state[node] = 1 - for edge in node.edges: - if edge.source == node: + if state[vertex] == 0: + state[vertex] = 1 + for edge in vertex.edges: + if edge.source == vertex: dfs(edge.target) - state[node] = 2 - sorted_vertices.append(node) + state[vertex] = 2 + sorted_vertices.append(vertex) - # Visit each node - for node in self.nodes: - if state[node] == 0: - dfs(node) + # Visit each vertex + for vertex in self.vertices: + if state[vertex] == 0: + dfs(vertex) return list(reversed(sorted_vertices)) @@ -164,17 +172,21 @@ class Graph: logger.debug("There are %s vertices in the graph", len(sorted_vertices)) yield from sorted_vertices - def get_node_neighbors(self, node: Vertex) -> Dict[Vertex, int]: - """Returns the neighbors of a node.""" + def get_vertex_neighbors(self, vertex: Vertex) -> Dict[Vertex, int]: + """Returns the neighbors of a vertex.""" neighbors: Dict[Vertex, int] = {} for edge in self.edges: - if edge.source == node: - neighbor = edge.target + if edge.source_id == vertex.id: + neighbor = self.get_vertex(edge.target_id) + if neighbor is None: + continue if neighbor not in neighbors: neighbors[neighbor] = 0 neighbors[neighbor] += 1 - elif edge.target == node: - neighbor = edge.source + elif edge.target_id == vertex.id: + neighbor = self.get_vertex(edge.source_id) + if neighbor is None: + continue if neighbor not in neighbors: neighbors[neighbor] = 0 neighbors[neighbor] += 1 @@ -182,59 +194,59 @@ class Graph: def _build_edges(self) -> List[Edge]: """Builds the edges of the graph.""" - # Edge takes two nodes as arguments, so we need to build the nodes first + # Edge takes two vertices as arguments, so we need to build the vertices first # and then build the edges - # if we can't find a node, we raise an error + # if we can't find a vertex, we raise an error edges: List[Edge] = [] for edge in self._edges: - source = self.get_node(edge["source"]) - target = self.get_node(edge["target"]) + source = self.get_vertex(edge["source"]) + target = self.get_vertex(edge["target"]) if source is None: - raise ValueError(f"Source node {edge['source']} not found") + raise ValueError(f"Source vertex {edge['source']} not found") if target is None: - raise ValueError(f"Target node {edge['target']} not found") + raise ValueError(f"Target vertex {edge['target']} not found") edges.append(Edge(source, target, edge)) return edges - def _get_vertex_class(self, node_type: str, node_lc_type: str) -> Type[Vertex]: - """Returns the node class based on the node type.""" - if node_type in FILE_TOOLS: + def _get_vertex_class(self, vertex_type: str, vertex_lc_type: str) -> Type[Vertex]: + """Returns the vertex class based on the vertex type.""" + if vertex_type in FILE_TOOLS: return FileToolVertex - if node_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP: - return lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_type] + if vertex_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP: + return lazy_load_vertex_dict.VERTEX_TYPE_MAP[vertex_type] return ( - lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_lc_type] - if node_lc_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP + lazy_load_vertex_dict.VERTEX_TYPE_MAP[vertex_lc_type] + if vertex_lc_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP else Vertex ) def _build_vertices(self) -> List[Vertex]: """Builds the vertices of the graph.""" - nodes: List[Vertex] = [] - for node in self._nodes: - node_data = node["data"] - node_type: str = node_data["type"] # type: ignore - node_lc_type: str = node_data["node"]["template"]["_type"] # type: ignore + vertices: List[Vertex] = [] + for vertex in self._vertices: + vertex_data = vertex["data"] + vertex_type: str = vertex_data["type"] # type: ignore + vertex_lc_type: str = vertex_data["node"]["template"]["_type"] # type: ignore - VertexClass = self._get_vertex_class(node_type, node_lc_type) - vertex = VertexClass(node) - vertex.set_top_level(self.top_level_nodes) - nodes.append(vertex) + VertexClass = self._get_vertex_class(vertex_type, vertex_lc_type) + vertex = VertexClass(vertex, graph=self) + vertex.set_top_level(self.top_level_vertices) + vertices.append(vertex) - return nodes + return vertices - def get_children_by_node_type(self, node: Vertex, node_type: str) -> List[Vertex]: - """Returns the children of a node based on the node type.""" + def get_children_by_vertex_type(self, vertex: Vertex, vertex_type: str) -> List[Vertex]: + """Returns the children of a vertex based on the vertex type.""" children = [] - node_types = [node.data["type"]] - if "node" in node.data: - node_types += node.data["node"]["base_classes"] - if node_type in node_types: - children.append(node) + vertex_types = [vertex.data["type"]] + if "node" in vertex.data: + vertex_types += vertex.data["node"]["base_classes"] + if vertex_type in vertex_types: + children.append(vertex) return children def __repr__(self): - node_ids = [node.id for node in self.nodes] - edges_repr = "\n".join([f"{edge.source.id} --> {edge.target.id}" for edge in self.edges]) - return f"Graph:\nNodes: {node_ids}\nConnections:\n{edges_repr}" + vertex_ids = [vertex.id for vertex in self.vertices] + edges_repr = "\n".join([f"{edge.source_id} --> {edge.target_id}" for edge in self.edges]) + return f"Graph:\nNodes: {vertex_ids}\nConnections:\n{edges_repr}" diff --git a/src/backend/langflow/graph/vertex/base.py b/src/backend/langflow/graph/vertex/base.py index 5ea645980..dc9e76f8a 100644 --- a/src/backend/langflow/graph/vertex/base.py +++ b/src/backend/langflow/graph/vertex/base.py @@ -1,33 +1,32 @@ import ast import inspect -import pickle import types from typing import TYPE_CHECKING, Any, Dict, List, Optional -from loguru import logger - from langflow.graph.utils import UnbuiltObject -from langflow.graph.vertex.utils import is_basic_type from langflow.interface.initialize import loading from langflow.interface.listing import lazy_load_dict from langflow.utils.constants import DIRECT_TYPES from langflow.utils.util import sync_to_async +from loguru import logger if TYPE_CHECKING: from langflow.graph.edge.base import Edge + from langflow.graph.graph.base import Graph class Vertex: def __init__( self, data: Dict, + graph: "Graph", base_type: Optional[str] = None, is_task: bool = False, params: Optional[Dict] = None, ) -> None: + self.graph = graph self.id: str = data["id"] self._data = data - self.edges: List["Edge"] = [] self.base_type: Optional[str] = base_type self._parse_data() self._built_object = UnbuiltObject() @@ -39,43 +38,28 @@ class Vertex: self.parent_node_id: Optional[str] = self._data.get("parent_node_id") self.parent_is_top_level = False - def reset_params(self): - for edge in self.edges: - if edge.source != self: - target_param = edge.target_param - if target_param in ["document", "texts"]: - # this means they got data and have already ingested it - # so we continue after removing the param - self.params.pop(target_param, None) - continue - - if target_param in self.params and not is_basic_type(self.params[target_param]): - # edge.source.params = {} - edge.source._build_params() - edge.source._built_object = UnbuiltObject() - edge.source._built = False - - self.params[target_param] = edge.source + @property + def edges(self) -> List["Edge"]: + return self.graph.get_vertex_edges(self.id) def __getstate__(self): - state_dict = self.__dict__.copy() - try: - # try pickling the built object - # if it fails, then we need to delete it - # and build it again - pickle.dumps(state_dict["_built_object"]) - except Exception: - self.reset_params() - del state_dict["_built_object"] - del state_dict["_built"] - return state_dict + return { + "_data": self._data, + "params": {}, + "base_type": self.base_type, + "is_task": self.is_task, + "id": self.id, + "_built_object": UnbuiltObject(), + "_built": False, + "parent_node_id": self.parent_node_id, + "parent_is_top_level": self.parent_is_top_level, + } def __setstate__(self, state): self._data = state["_data"] self.params = state["params"] self.base_type = state["base_type"] self.is_task = state["is_task"] - self.edges = state["edges"] self.id = state["id"] self._parse_data() if "_built_object" in state: @@ -144,6 +128,10 @@ class Vertex: # and use that as the value for the param # If the type is "str", then we need to get the value of the "value" key # and use that as the value for the param + + if self.graph is None: + raise ValueError("Graph not found") + template_dict = {key: value for key, value in self.data["node"]["template"].items() if isinstance(value, dict)} params = self.params.copy() if self.params else {} @@ -155,9 +143,9 @@ class Vertex: if template_dict[param_key]["list"]: if param_key not in params: params[param_key] = [] - params[param_key].append(edge.source) - elif edge.target.id == self.id: - params[param_key] = edge.source + params[param_key].append(self.graph.get_vertex(edge.source_id)) + elif edge.target_id == self.id: + params[param_key] = self.graph.get_vertex(edge.source_id) for key, value in template_dict.items(): if key in params: @@ -177,33 +165,33 @@ class Vertex: else: raise ValueError(f"File path not found for {self.vertex_type}") elif value.get("type") in DIRECT_TYPES and params.get(key) is None: + val = value.get("value") if value.get("type") == "code": try: - params[key] = ast.literal_eval(value.get("value")) + params[key] = ast.literal_eval(val) if val else None except Exception as exc: logger.debug(f"Error parsing code: {exc}") - params[key] = value.get("value") + params[key] = val elif value.get("type") in ["dict", "NestedDict"]: # When dict comes from the frontend it comes as a # list of dicts, so we need to convert it to a dict # before passing it to the build method - _value = value.get("value") - if isinstance(_value, list): + if isinstance(val, list): params[key] = {k: v for item in value.get("value", []) for k, v in item.items()} - elif isinstance(_value, dict): - params[key] = _value - elif value.get("type") == "int" and value.get("value") is not None: + elif isinstance(val, dict): + params[key] = val + elif value.get("type") == "int" and val is not None: try: - params[key] = int(value.get("value")) + params[key] = int(val) except ValueError: - params[key] = value.get("value") - elif value.get("type") == "float" and value.get("value") is not None: + params[key] = val + elif value.get("type") == "float" and val is not None: try: - params[key] = float(value.get("value")) + params[key] = float(val) except ValueError: - params[key] = value.get("value") + params[key] = val else: - params[key] = value.get("value") + params[key] = val if not value.get("required") and params.get(key) is None: if value.get("default"): @@ -266,7 +254,7 @@ class Vertex: pass # If there's no task_id, build the vertex locally - await self.build(user_id) + await self.build(user_id=user_id) return self._built_object async def _build_node_and_update_params(self, key, node, user_id=None): From 1facfefb193fedfca75a4a4836206d0fced2dd6d Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:50:43 -0300 Subject: [PATCH 2/9] =?UTF-8?q?=F0=9F=90=9B=20fix(types.py):=20pass=20grap?= =?UTF-8?q?h=20parameter=20to=20Vertex=20constructors=20to=20fix=20missing?= =?UTF-8?q?=20graph=20reference=20=E2=9C=A8=20feat(types.py):=20add=20supp?= =?UTF-8?q?ort=20for=20passing=20graph=20parameter=20to=20Vertex=20constru?= =?UTF-8?q?ctors=20to=20ensure=20proper=20graph=20reference?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/backend/langflow/graph/vertex/types.py | 113 ++++++++++----------- 1 file changed, 56 insertions(+), 57 deletions(-) diff --git a/src/backend/langflow/graph/vertex/types.py b/src/backend/langflow/graph/vertex/types.py index c288a4b0a..92c920d35 100644 --- a/src/backend/langflow/graph/vertex/types.py +++ b/src/backend/langflow/graph/vertex/types.py @@ -1,14 +1,14 @@ import ast from typing import Any, Dict, List, Optional, Union -from langflow.graph.utils import flatten_list +from langflow.graph.utils import UnbuiltObject, flatten_list from langflow.graph.vertex.base import Vertex from langflow.interface.utils import extract_input_variables_from_prompt class AgentVertex(Vertex): - def __init__(self, data: Dict, params: Optional[Dict] = None): - super().__init__(data, base_type="agents", params=params) + def __init__(self, data: Dict, graph, params: Optional[Dict] = None): + super().__init__(data, graph=graph, base_type="agents", params=params) self.tools: List[Union[ToolkitVertex, ToolVertex]] = [] self.chains: List[ChainVertex] = [] @@ -28,7 +28,7 @@ class AgentVertex(Vertex): for edge in self.edges: if not hasattr(edge, "source"): continue - source_node = edge.source + source_node = self.graph.get_vertex(edge.source_id) if isinstance(source_node, (ToolVertex, ToolkitVertex)): self.tools.append(source_node) elif isinstance(source_node, ChainVertex): @@ -51,16 +51,21 @@ class AgentVertex(Vertex): class ToolVertex(Vertex): - def __init__(self, data: Dict, params: Optional[Dict] = None): - super().__init__(data, base_type="tools", params=params) + def __init__( + self, + data: Dict, + graph, + params: Optional[Dict] = None, + ): + super().__init__(data, graph=graph, base_type="tools", params=params) class LLMVertex(Vertex): built_node_type = None class_built_object = None - def __init__(self, data: Dict, params: Optional[Dict] = None): - super().__init__(data, base_type="llms", params=params) + def __init__(self, data: Dict, graph, params: Optional[Dict] = None): + super().__init__(data, graph=graph, base_type="llms", params=params) async def build(self, force: bool = False, user_id=None, *args, **kwargs) -> Any: # LLM is different because some models might take up too much memory @@ -77,18 +82,18 @@ class LLMVertex(Vertex): class ToolkitVertex(Vertex): - def __init__(self, data: Dict, params=None): - super().__init__(data, base_type="toolkits", params=params) + def __init__(self, data: Dict, graph, params=None): + super().__init__(data, graph=graph, base_type="toolkits", params=params) class FileToolVertex(ToolVertex): - def __init__(self, data: Dict, params=None): - super().__init__(data, params=params) + def __init__(self, data: Dict, graph, params=None): + super().__init__(data, graph=graph, params=params) class WrapperVertex(Vertex): - def __init__(self, data: Dict): - super().__init__(data, base_type="wrappers") + def __init__(self, data: Dict, graph): + super().__init__(data, graph=graph, base_type="wrappers") async def build(self, force: bool = False, user_id=None, *args, **kwargs) -> Any: if not self._built or force: @@ -99,14 +104,14 @@ class WrapperVertex(Vertex): class DocumentLoaderVertex(Vertex): - def __init__(self, data: Dict, params: Optional[Dict] = None): - super().__init__(data, base_type="documentloaders", params=params) + def __init__(self, data: Dict, graph, params: Optional[Dict] = None): + super().__init__(data, graph=graph, base_type="documentloaders", params=params) def _built_object_repr(self): # This built_object is a list of documents. Maybe we should # show how many documents are in the list? - if self._built_object: + if self._built_object and not isinstance(self._built_object, UnbuiltObject): avg_length = sum(len(doc.page_content) for doc in self._built_object if hasattr(doc, "page_content")) / len( self._built_object ) @@ -117,28 +122,19 @@ class DocumentLoaderVertex(Vertex): class EmbeddingVertex(Vertex): - def __init__(self, data: Dict, params: Optional[Dict] = None): - super().__init__(data, base_type="embeddings", params=params) + def __init__(self, data: Dict, graph, params: Optional[Dict] = None): + super().__init__(data, graph=graph, base_type="embeddings", params=params) class VectorStoreVertex(Vertex): - def __init__(self, data: Dict, params=None): - super().__init__(data, base_type="vectorstores") + def __init__(self, data: Dict, graph, params=None): + super().__init__(data, graph=graph, base_type="vectorstores") self.params = params or {} # VectorStores may contain databse connections # so we need to define the __reduce__ method and the __setstate__ method # to avoid pickling errors - def clean_edges_for_pickling(self): - # for each edge that has self as source - # we need to clear the _built_object of the target - # so that we don't try to pickle a database connection - for edge in self.edges: - if edge.source == self: - edge.target._built_object = None - edge.target._built = False - edge.target.params[edge.target_param] = self def remove_docs_and_texts_from_params(self): # remove documents and texts from params @@ -146,17 +142,16 @@ class VectorStoreVertex(Vertex): self.params.pop("documents", None) self.params.pop("texts", None) - def __getstate__(self): - # We want to save the params attribute - # and if "documents" or "texts" are in the params - # we want to remove them because they have already - # been processed. - params = self.params.copy() - params.pop("documents", None) - params.pop("texts", None) - self.clean_edges_for_pickling() + # def __getstate__(self): + # # We want to save the params attribute + # # and if "documents" or "texts" are in the params + # # we want to remove them because they have already + # # been processed. + # params = self.params.copy() + # params.pop("documents", None) + # params.pop("texts", None) - return super().__getstate__() + # return super().__getstate__() def __setstate__(self, state): super().__setstate__(state) @@ -164,24 +159,24 @@ class VectorStoreVertex(Vertex): class MemoryVertex(Vertex): - def __init__(self, data: Dict): - super().__init__(data, base_type="memory") + def __init__(self, data: Dict, graph): + super().__init__(data, graph=graph, base_type="memory") class RetrieverVertex(Vertex): - def __init__(self, data: Dict): - super().__init__(data, base_type="retrievers") + def __init__(self, data: Dict, graph): + super().__init__(data, graph=graph, base_type="retrievers") class TextSplitterVertex(Vertex): - def __init__(self, data: Dict, params: Optional[Dict] = None): - super().__init__(data, base_type="textsplitters", params=params) + def __init__(self, data: Dict, graph, params: Optional[Dict] = None): + super().__init__(data, graph=graph, base_type="textsplitters", params=params) def _built_object_repr(self): # This built_object is a list of documents. Maybe we should # show how many documents are in the list? - if self._built_object: + if self._built_object and not isinstance(self._built_object, UnbuiltObject): avg_length = sum(len(doc.page_content) for doc in self._built_object) / len(self._built_object) return f"""{self.vertex_type}({len(self._built_object)} documents) \nAvg. Document Length (characters): {int(avg_length)} @@ -190,8 +185,8 @@ class TextSplitterVertex(Vertex): class ChainVertex(Vertex): - def __init__(self, data: Dict): - super().__init__(data, base_type="chains") + def __init__(self, data: Dict, graph): + super().__init__(data, graph=graph, base_type="chains") async def build( self, @@ -220,8 +215,8 @@ class ChainVertex(Vertex): class PromptVertex(Vertex): - def __init__(self, data: Dict): - super().__init__(data, base_type="prompts") + def __init__(self, data: Dict, graph): + super().__init__(data, graph=graph, base_type="prompts") async def build( self, @@ -271,9 +266,13 @@ class PromptVertex(Vertex): # so the prompt format doesn't break artifacts.pop("handle_keys", None) try: - if not hasattr(self._built_object, "template") and hasattr(self._built_object, "prompt"): + if ( + not hasattr(self._built_object, "template") + and hasattr(self._built_object, "prompt") + and not isinstance(self._built_object, UnbuiltObject) + ): template = self._built_object.prompt.template - else: + elif not isinstance(self._built_object, UnbuiltObject) and hasattr(self._built_object, "template"): template = self._built_object.template for key, value in artifacts.items(): if value: @@ -285,13 +284,13 @@ class PromptVertex(Vertex): class OutputParserVertex(Vertex): - def __init__(self, data: Dict): - super().__init__(data, base_type="output_parsers") + def __init__(self, data: Dict, graph): + super().__init__(data, graph=graph, base_type="output_parsers") class CustomComponentVertex(Vertex): - def __init__(self, data: Dict): - super().__init__(data, base_type="custom_components", is_task=True) + def __init__(self, data: Dict, graph): + super().__init__(data, graph=graph, base_type="custom_components", is_task=True) def _built_object_repr(self): if self.task_id and self.is_task: From c8d84490299c5eaf1d3c4215021357d0636c0c9e Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:50:55 -0300 Subject: [PATCH 3/9] Refactor API imports and variable names --- src/backend/langflow/api/v1/chat.py | 22 +++++----------------- 1 file changed, 5 insertions(+), 17 deletions(-) diff --git a/src/backend/langflow/api/v1/chat.py b/src/backend/langflow/api/v1/chat.py index 66e23aa5a..877ee07bd 100644 --- a/src/backend/langflow/api/v1/chat.py +++ b/src/backend/langflow/api/v1/chat.py @@ -1,27 +1,15 @@ -from fastapi import ( - APIRouter, - Depends, - HTTPException, - Query, - WebSocket, - WebSocketException, - status, -) +from fastapi import APIRouter, Depends, HTTPException, Query, WebSocket, WebSocketException, status from fastapi.responses import StreamingResponse -from loguru import logger -from sqlmodel import Session - from langflow.api.utils import build_input_keys_response from langflow.api.v1.schemas import BuildStatus, BuiltResponse, InitResponse, StreamData from langflow.graph.graph.base import Graph -from langflow.services.auth.utils import ( - get_current_active_user, - get_current_user_by_jwt, -) +from langflow.services.auth.utils import get_current_active_user, get_current_user_by_jwt from langflow.services.cache.service import BaseCacheService from langflow.services.cache.utils import update_build_status from langflow.services.chat.service import ChatService from langflow.services.deps import get_cache_service, get_chat_service, get_session +from loguru import logger +from sqlmodel import Session router = APIRouter(tags=["Chat"]) @@ -148,7 +136,7 @@ async def stream_build( # Some error could happen when building the graph graph = Graph.from_payload(graph_data) - number_of_nodes = len(graph.nodes) + number_of_nodes = len(graph.vertices) update_build_status(cache_service, flow_id, BuildStatus.IN_PROGRESS) try: From 73e0cdba16e3d2af2a835304669d74cbf0886587 Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:51:01 -0300 Subject: [PATCH 4/9] Refactor API endpoints and reload custom component --- src/backend/langflow/api/v1/endpoints.py | 19 +++++-------------- 1 file changed, 5 insertions(+), 14 deletions(-) diff --git a/src/backend/langflow/api/v1/endpoints.py b/src/backend/langflow/api/v1/endpoints.py index 18e241133..51d3e8bad 100644 --- a/src/backend/langflow/api/v1/endpoints.py +++ b/src/backend/langflow/api/v1/endpoints.py @@ -3,8 +3,6 @@ from typing import Annotated, Optional, Union import sqlalchemy as sa from fastapi import APIRouter, Body, Depends, HTTPException, UploadFile, status -from loguru import logger - from langflow.api.v1.schemas import ( CustomComponentCode, ProcessResponse, @@ -20,12 +18,8 @@ from langflow.services.auth.utils import api_key_security, get_current_active_us from langflow.services.cache.utils import save_uploaded_file from langflow.services.database.models.flow import Flow from langflow.services.database.models.user.model import User -from langflow.services.deps import ( - get_session, - get_session_service, - get_settings_service, - get_task_service, -) +from langflow.services.deps import get_session, get_session_service, get_settings_service, get_task_service +from loguru import logger try: from langflow.worker import process_graph_cached_task @@ -35,9 +29,8 @@ except ImportError: raise NotImplementedError("Celery is not installed") -from sqlmodel import Session - from langflow.services.task.service import TaskService +from sqlmodel import Session # build router router = APIRouter(tags=["Base"]) @@ -218,10 +211,8 @@ async def custom_component( @router.post("/custom_component/reload", status_code=HTTPStatus.OK) -async def reload_custom_component(path: str): - from langflow.interface.types import ( - build_langchain_template_custom_component, - ) +async def reload_custom_component(path: str, user: User = Depends(get_current_active_user)): + from langflow.interface.types import build_langchain_template_custom_component try: reader = DirectoryReader("") From e9871976eff8fe906b48c4eef253ce94e4e682ba Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:51:06 -0300 Subject: [PATCH 5/9] Fix user_id immutability issue --- src/backend/langflow/interface/custom/component.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/backend/langflow/interface/custom/component.py b/src/backend/langflow/interface/custom/component.py index dfd7c0b00..594ec982f 100644 --- a/src/backend/langflow/interface/custom/component.py +++ b/src/backend/langflow/interface/custom/component.py @@ -25,6 +25,7 @@ class Component: code: Optional[str] = None _function_entrypoint_name: str = "build" field_config: dict = {} + _user_id: Optional[str] def __init__(self, **data): self.cache = TTLCache(maxsize=1024, ttl=60) @@ -36,9 +37,8 @@ class Component: def __setattr__(self, key, value): if key == "_user_id" and hasattr(self, "_user_id"): - warnings.warn("Modification of user_id is not allowed") - else: - super().__setattr__(key, value) + warnings.warn("user_id is immutable and cannot be changed.") + super().__setattr__(key, value) @cachedmethod(cache=operator.attrgetter("cache")) def get_code_tree(self, code: str): From 5c90a9013c31f9ca4beadaca201fd86b9961e22e Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:51:10 -0300 Subject: [PATCH 6/9] Fix user_id attribute in list_flows method --- src/backend/langflow/interface/custom/custom_component.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/backend/langflow/interface/custom/custom_component.py b/src/backend/langflow/interface/custom/custom_component.py index 98be2299e..f887ed025 100644 --- a/src/backend/langflow/interface/custom/custom_component.py +++ b/src/backend/langflow/interface/custom/custom_component.py @@ -5,7 +5,6 @@ from uuid import UUID import yaml from cachetools import TTLCache, cachedmethod from fastapi import HTTPException - from langflow.field_typing.constants import CUSTOM_COMPONENT_SUPPORTED_TYPES from langflow.interface.custom.component import Component from langflow.interface.custom.directory_reader import DirectoryReader @@ -232,7 +231,7 @@ class CustomComponent(Component): return await build_sorted_vertices(graph_data, self.user_id) def list_flows(self, *, get_session: Optional[Callable] = None) -> List[Flow]: - if not self.user_id: + if not self._user_id: raise ValueError("Session is invalid") try: get_session = get_session or session_getter From 2c79494d8b752693e08b841239718ffb4f8e8d94 Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:51:16 -0300 Subject: [PATCH 7/9] Fix get_root_node function to use vertices instead of nodes --- src/backend/langflow/utils/payload.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/src/backend/langflow/utils/payload.py b/src/backend/langflow/utils/payload.py index 02cca5c71..0e2f0fc7a 100644 --- a/src/backend/langflow/utils/payload.py +++ b/src/backend/langflow/utils/payload.py @@ -28,16 +28,16 @@ def extract_input_variables(nodes): return nodes -def get_root_node(graph): +def get_root_vertex(graph): """ Returns the root node of the template. """ - incoming_edges = {edge.source for edge in graph.edges} + incoming_edges = {edge.source_id for edge in graph.edges} - if not incoming_edges and len(graph.nodes) == 1: - return graph.nodes[0] + if not incoming_edges and len(graph.vertices) == 1: + return graph.vertices[0] - return next((node for node in graph.nodes if node not in incoming_edges), None) + return next((node for node in graph.vertices if node.id not in incoming_edges), None) def build_json(root, graph) -> Dict: From 3ac18c82565a3892cb421a71d9b2f9e1e6d1e04a Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:51:21 -0300 Subject: [PATCH 8/9] Fix import statement in validate.py --- src/backend/langflow/utils/validate.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/backend/langflow/utils/validate.py b/src/backend/langflow/utils/validate.py index 51c4894d5..0a33dfa01 100644 --- a/src/backend/langflow/utils/validate.py +++ b/src/backend/langflow/utils/validate.py @@ -1,7 +1,7 @@ import ast import contextlib import importlib -import types +from types import FunctionType from typing import Dict @@ -61,7 +61,7 @@ def eval_function(function_string: str): ( obj for name, obj in namespace.items() - if isinstance(obj, types.FunctionType) and obj.__code__.co_filename == "" + if isinstance(obj, FunctionType) and obj.__code__.co_filename == "" ), None, ) From 5ad879bd9b0c5fdaea7a32fbb3007175b8ef90b2 Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Mon, 27 Nov 2023 21:51:31 -0300 Subject: [PATCH 9/9] Fix inconsistencies in test cases --- tests/test_cache.py | 4 +-- tests/test_custom_component.py | 7 ++--- tests/test_graph.py | 55 ++++++++++++++++++---------------- 3 files changed, 33 insertions(+), 33 deletions(-) diff --git a/tests/test_cache.py b/tests/test_cache.py index c2c706ee9..925402769 100644 --- a/tests/test_cache.py +++ b/tests/test_cache.py @@ -1,7 +1,7 @@ import json -from langflow.graph import Graph import pytest +from langflow.graph import Graph def get_graph(_type="basic"): @@ -41,5 +41,5 @@ def langchain_objects_are_equal(obj1, obj2): def test_build_graph(client, basic_data_graph): graph = Graph.from_payload(basic_data_graph) assert graph is not None - assert len(graph.nodes) == len(basic_data_graph["nodes"]) + assert len(graph.vertices) == len(basic_data_graph["nodes"]) assert len(graph.edges) == len(basic_data_graph["edges"]) diff --git a/tests/test_custom_component.py b/tests/test_custom_component.py index b07753b8d..636bd63b1 100644 --- a/tests/test_custom_component.py +++ b/tests/test_custom_component.py @@ -7,10 +7,7 @@ from fastapi import HTTPException from langflow.field_typing.constants import Data from langflow.interface.custom.base import CustomComponent from langflow.interface.custom.code_parser import CodeParser, CodeSyntaxError -from langflow.interface.custom.component import ( - Component, - ComponentCodeNullError, -) +from langflow.interface.custom.component import Component, ComponentCodeNullError from langflow.services.database.models.flow import Flow, FlowCreate code_default = """ @@ -445,7 +442,7 @@ def test_custom_component_build_not_implemented(): def test_build_config_no_code(): component = CustomComponent(code=None) - assert component.get_function_entrypoint_args == "" + assert component.get_function_entrypoint_args == [] assert component.get_function_entrypoint_return_type == [] diff --git a/tests/test_graph.py b/tests/test_graph.py index cb69d79d5..020642798 100644 --- a/tests/test_graph.py +++ b/tests/test_graph.py @@ -24,7 +24,7 @@ from langflow.graph.utils import UnbuiltObject from langflow.graph.vertex.base import Vertex from langflow.graph.vertex.types import FileToolVertex, LLMVertex, ToolkitVertex from langflow.processing.process import get_result_and_thought -from langflow.utils.payload import get_root_node +from langflow.utils.payload import get_root_vertex # Test cases for the graph module @@ -70,19 +70,19 @@ def sample_nodes(): def get_node_by_type(graph, node_type: Type[Vertex]) -> Union[Vertex, None]: """Get a node by type""" - return next((node for node in graph.nodes if isinstance(node, node_type)), None) + return next((node for node in graph.vertices if isinstance(node, node_type)), None) def test_graph_structure(basic_graph): assert isinstance(basic_graph, Graph) - assert len(basic_graph.nodes) > 0 + assert len(basic_graph.vertices) > 0 assert len(basic_graph.edges) > 0 - for node in basic_graph.nodes: + for node in basic_graph.vertices: assert isinstance(node, Vertex) for edge in basic_graph.edges: assert isinstance(edge, Edge) - assert edge.source in basic_graph.nodes - assert edge.target in basic_graph.nodes + assert edge.source_id in basic_graph.vertex_ids + assert edge.target_id in basic_graph.vertex_ids def test_circular_dependencies(basic_graph): @@ -90,7 +90,7 @@ def test_circular_dependencies(basic_graph): def check_circular(node, visited): visited.add(node) - neighbors = basic_graph.get_nodes_with_target(node) + neighbors = basic_graph.get_vertices_with_target(node) for neighbor in neighbors: if neighbor in visited: return True @@ -98,7 +98,7 @@ def test_circular_dependencies(basic_graph): return True return False - for node in basic_graph.nodes: + for node in basic_graph.vertices: assert not check_circular(node, set()) @@ -123,13 +123,13 @@ def test_invalid_node_types(): Graph(graph_data["nodes"], graph_data["edges"]) -def test_get_nodes_with_target(basic_graph): +def test_get_vertices_with_target(basic_graph): """Test getting connected nodes""" assert isinstance(basic_graph, Graph) # Get root node - root = get_root_node(basic_graph) + root = get_root_vertex(basic_graph) assert root is not None - connected_nodes = basic_graph.get_nodes_with_target(root) + connected_nodes = basic_graph.get_vertices_with_target(root.id) assert connected_nodes is not None @@ -138,9 +138,9 @@ def test_get_node_neighbors_basic(basic_graph): assert isinstance(basic_graph, Graph) # Get root node - root = get_root_node(basic_graph) + root = get_root_vertex(basic_graph) assert root is not None - neighbors = basic_graph.get_node_neighbors(root) + neighbors = basic_graph.get_vertex_neighbors(root) assert neighbors is not None assert isinstance(neighbors, dict) # Root Node is an Agent, it requires an LLMChain and tools @@ -153,8 +153,8 @@ def test_get_node_neighbors_basic(basic_graph): def test_get_node(basic_graph): """Test getting a single node""" - node_id = basic_graph.nodes[0].id - node = basic_graph.get_node(node_id) + node_id = basic_graph.vertices[0].id + node = basic_graph.get_vertex(node_id) assert isinstance(node, Vertex) assert node.id == node_id @@ -162,8 +162,8 @@ def test_get_node(basic_graph): def test_build_nodes(basic_graph): """Test building nodes""" - assert len(basic_graph.nodes) == len(basic_graph._nodes) - for node in basic_graph.nodes: + assert len(basic_graph.vertices) == len(basic_graph._vertices) + for node in basic_graph.vertices: assert isinstance(node, Vertex) @@ -172,20 +172,21 @@ def test_build_edges(basic_graph): assert len(basic_graph.edges) == len(basic_graph._edges) for edge in basic_graph.edges: assert isinstance(edge, Edge) - assert isinstance(edge.source, Vertex) - assert isinstance(edge.target, Vertex) + + assert isinstance(edge.source_id, str) + assert isinstance(edge.target_id, str) -def test_get_root_node(client, basic_graph, complex_graph): +def test_get_root_vertex(client, basic_graph, complex_graph): """Test getting root node""" assert isinstance(basic_graph, Graph) - root = get_root_node(basic_graph) + root = get_root_vertex(basic_graph) assert root is not None assert isinstance(root, Vertex) assert root.data["type"] == "TimeTravelGuideChain" # For complex example, the root node is a ZeroShotAgent too assert isinstance(complex_graph, Graph) - root = get_root_node(complex_graph) + root = get_root_vertex(complex_graph) assert root is not None assert isinstance(root, Vertex) assert root.data["type"] == "ZeroShotAgent" @@ -221,7 +222,7 @@ def test_build_params(basic_graph): # The matched_type attribute should be in the source_types attr assert all(edge.matched_type in edge.source_types for edge in basic_graph.edges) # Get the root node - root = get_root_node(basic_graph) + root = get_root_vertex(basic_graph) # Root node is a TimeTravelGuideChain # which requires an llm and memory assert root is not None @@ -278,7 +279,7 @@ async def test_file_tool_node_build(client, openapi_graph): assert Path(file_path).exists() file_tool_node = get_node_by_type(openapi_graph, FileToolVertex) - assert file_tool_node is not UnbuiltObject + assert file_tool_node is not UnbuiltObject and file_tool_node is not None built_object = await file_tool_node.build() assert built_object is not UnbuiltObject # Remove the file @@ -301,7 +302,7 @@ async def test_get_result_and_thought(basic_graph): llm_node._built = True langchain_object = await basic_graph.build() # assert all nodes are built - assert all(node._built for node in basic_graph.nodes) + assert all(node._built for node in basic_graph.vertices) # now build again and check if FakeListLLM was used # Get the result and thought @@ -420,10 +421,12 @@ def test_update_template(sample_template, sample_nodes): node2_updated = next((n for n in nodes_copy if n["id"] == "node2"), None) node3_updated = next((n for n in nodes_copy if n["id"] == "node3"), None) + assert node1_updated is not None assert node1_updated["data"]["node"]["template"]["some_field"]["show"] is True assert node1_updated["data"]["node"]["template"]["some_field"]["advanced"] is False assert node1_updated["data"]["node"]["template"]["some_field"]["display_name"] == "Name1" + assert node2_updated is not None assert node2_updated["data"]["node"]["template"]["other_field"]["show"] is False assert node2_updated["data"]["node"]["template"]["other_field"]["advanced"] is True assert node2_updated["data"]["node"]["template"]["other_field"]["display_name"] == "DisplayName2" @@ -502,7 +505,7 @@ async def test_pickle_each_vertex(json_vector_store): loaded_json = json.loads(json_vector_store) graph = Graph.from_payload(loaded_json) assert isinstance(graph, Graph) - for vertex in graph.nodes: + for vertex in graph.vertices: await vertex.build() pickled = pickle.dumps(vertex) assert pickled is not UnbuiltObject