Merge branch 'better_graph' into feature/store

This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-11-27 22:59:41 -03:00
commit 5b4d0c62c8
13 changed files with 280 additions and 311 deletions

View file

@ -1,27 +1,15 @@
from fastapi import ( from fastapi import APIRouter, Depends, HTTPException, Query, WebSocket, WebSocketException, status
APIRouter,
Depends,
HTTPException,
Query,
WebSocket,
WebSocketException,
status,
)
from fastapi.responses import StreamingResponse 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.utils import build_input_keys_response
from langflow.api.v1.schemas import BuildStatus, BuiltResponse, InitResponse, StreamData from langflow.api.v1.schemas import BuildStatus, BuiltResponse, InitResponse, StreamData
from langflow.graph.graph.base import Graph from langflow.graph.graph.base import Graph
from langflow.services.auth.utils import ( from langflow.services.auth.utils import get_current_active_user, get_current_user_by_jwt
get_current_active_user,
get_current_user_by_jwt,
)
from langflow.services.cache.service import BaseCacheService from langflow.services.cache.service import BaseCacheService
from langflow.services.cache.utils import update_build_status from langflow.services.cache.utils import update_build_status
from langflow.services.chat.service import ChatService from langflow.services.chat.service import ChatService
from langflow.services.deps import get_cache_service, get_chat_service, get_session 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"]) router = APIRouter(tags=["Chat"])
@ -148,7 +136,7 @@ async def stream_build(
# Some error could happen when building the graph # Some error could happen when building the graph
graph = Graph.from_payload(graph_data) 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) update_build_status(cache_service, flow_id, BuildStatus.IN_PROGRESS)
try: try:

View file

@ -3,8 +3,6 @@ from typing import Annotated, Optional, Union
import sqlalchemy as sa import sqlalchemy as sa
from fastapi import APIRouter, Body, Depends, HTTPException, UploadFile, status from fastapi import APIRouter, Body, Depends, HTTPException, UploadFile, status
from loguru import logger
from langflow.api.v1.schemas import ( from langflow.api.v1.schemas import (
CustomComponentCode, CustomComponentCode,
ProcessResponse, 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.cache.utils import save_uploaded_file
from langflow.services.database.models.flow import Flow from langflow.services.database.models.flow import Flow
from langflow.services.database.models.user.model import User from langflow.services.database.models.user.model import User
from langflow.services.deps import ( from langflow.services.deps import get_session, get_session_service, get_settings_service, get_task_service
get_session, from loguru import logger
get_session_service,
get_settings_service,
get_task_service,
)
try: try:
from langflow.worker import process_graph_cached_task from langflow.worker import process_graph_cached_task
@ -35,9 +29,8 @@ except ImportError:
raise NotImplementedError("Celery is not installed") raise NotImplementedError("Celery is not installed")
from sqlmodel import Session
from langflow.services.task.service import TaskService from langflow.services.task.service import TaskService
from sqlmodel import Session
# build router # build router
router = APIRouter(tags=["Base"]) router = APIRouter(tags=["Base"])
@ -218,10 +211,8 @@ async def custom_component(
@router.post("/custom_component/reload", status_code=HTTPStatus.OK) @router.post("/custom_component/reload", status_code=HTTPStatus.OK)
async def reload_custom_component(path: str): async def reload_custom_component(path: str, user: User = Depends(get_current_active_user)):
from langflow.interface.types import ( from langflow.interface.types import build_langchain_template_custom_component
build_langchain_template_custom_component,
)
try: try:
reader = DirectoryReader("") reader = DirectoryReader("")

View file

@ -1,7 +1,7 @@
from typing import TYPE_CHECKING, List, Optional
from loguru import logger from loguru import logger
from typing import TYPE_CHECKING
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from typing import List, Optional
if TYPE_CHECKING: if TYPE_CHECKING:
from langflow.graph.vertex.base import Vertex from langflow.graph.vertex.base import Vertex
@ -22,8 +22,8 @@ class TargetHandle(BaseModel):
class Edge: class Edge:
def __init__(self, source: "Vertex", target: "Vertex", edge: dict): def __init__(self, source: "Vertex", target: "Vertex", edge: dict):
self.source: "Vertex" = source self.source_id: str = source.id
self.target: "Vertex" = target self.target_id: str = target.id
if data := edge.get("data", {}): if data := edge.get("data", {}):
self._source_handle = data.get("sourceHandle", {}) self._source_handle = data.get("sourceHandle", {})
self._target_handle = data.get("targetHandle", {}) self._target_handle = data.get("targetHandle", {})
@ -31,7 +31,7 @@ class Edge:
self.target_handle: TargetHandle = TargetHandle(**self._target_handle) self.target_handle: TargetHandle = TargetHandle(**self._target_handle)
self.target_param = self.target_handle.fieldName self.target_param = self.target_handle.fieldName
# validate handles # validate handles
self.validate_handles() self.validate_handles(source, target)
else: else:
# Logging here because this is a breaking change # Logging here because this is a breaking change
logger.error("Edge data is empty") logger.error("Edge data is empty")
@ -41,9 +41,9 @@ class Edge:
# target_param is documents # target_param is documents
self.target_param = self._target_handle.split("|")[1] self.target_param = self._target_handle.split("|")[1]
# Validate in __init__ to fail fast # 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: if self.target_handle.inputTypes is None:
self.valid_handles = self.target_handle.type in self.source_handle.baseClasses self.valid_handles = self.target_handle.type in self.source_handle.baseClasses
else: else:
@ -54,26 +54,20 @@ class Edge:
if not self.valid_handles: if not self.valid_handles:
logger.debug(self.source_handle) logger.debug(self.source_handle)
logger.debug(self.target_handle) logger.debug(self.target_handle)
raise ValueError( raise ValueError(f"Edge between {source.vertex_type} and {target.vertex_type} " f"has invalid handles")
f"Edge between {self.source.vertex_type} and {self.target.vertex_type} " f"has invalid handles"
)
def __setstate__(self, state): def __setstate__(self, state):
self.source = state["source"] self.source_id = state["source_id"]
self.target = state["target"] self.target_id = state["target_id"]
self.target_param = state["target_param"] self.target_param = state["target_param"]
self.source_handle = state.get("source_handle") self.source_handle = state.get("source_handle")
self.target_handle = state.get("target_handle") self.target_handle = state.get("target_handle")
def reset(self) -> None: def validate_edge(self, source, target) -> None:
self.source._build_params()
self.target._build_params()
def validate_edge(self) -> None:
# Validate that the outputs of the source node are valid inputs # Validate that the outputs of the source node are valid inputs
# for the target node # for the target node
self.source_types = self.source.output self.source_types = source.output
self.target_reqs = self.target.required_inputs + self.target.optional_inputs self.target_reqs = target.required_inputs + target.optional_inputs
# Both lists contain strings and sometimes a string contains the value we are # Both lists contain strings and sometimes a string contains the value we are
# looking for e.g. comgin_out=["Chain"] and target_reqs=["LLMChain"] # 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 # 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: if no_matched_type:
logger.debug(self.source_types) logger.debug(self.source_types)
logger.debug(self.target_reqs) logger.debug(self.target_reqs)
raise ValueError( raise ValueError(f"Edge between {source.vertex_type} and {target.vertex_type} " f"has no matched type")
f"Edge between {self.source.vertex_type} and {self.target.vertex_type} " f"has no matched type"
)
def __repr__(self) -> str: def __repr__(self) -> str:
return ( 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})" f", matched_type={self.matched_type})"
) )

View file

@ -13,24 +13,24 @@ from langflow.utils import payload
class Graph: class Graph:
"""A class representing a graph of nodes and edges.""" """A class representing a graph of vertices and edges."""
def __init__( def __init__(
self, self,
nodes: List[Dict], nodes: List[Dict],
edges: List[Dict[str, str]], edges: List[Dict[str, str]],
) -> None: ) -> None:
self._nodes = nodes self._vertices = nodes
self._edges = edges self._edges = edges
self.raw_graph_data = {"nodes": nodes, "edges": edges} self.raw_graph_data = {"nodes": nodes, "edges": edges}
self.top_level_nodes = [] self.top_level_vertices = []
for node in self._nodes: for vertex in self._vertices:
if node_id := node.get("id"): if vertex_id := vertex.get("id"):
self.top_level_nodes.append(node_id) self.top_level_vertices.append(vertex_id)
self._graph_data = process_flow(self.raw_graph_data) 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._edges = self._graph_data["edges"]
self._build_graph() self._build_graph()
@ -54,9 +54,9 @@ class Graph:
if "data" in payload: if "data" in payload:
payload = payload["data"] payload = payload["data"]
try: try:
nodes = payload["nodes"] vertices = payload["nodes"]
edges = payload["edges"] edges = payload["edges"]
return cls(nodes, edges) return cls(vertices, edges)
except KeyError as exc: except KeyError as exc:
logger.exception(exc) logger.exception(exc)
raise ValueError( raise ValueError(
@ -69,61 +69,69 @@ class Graph:
return self.__repr__() == other.__repr__() return self.__repr__() == other.__repr__()
def _build_graph(self) -> None: def _build_graph(self) -> None:
"""Builds the graph from the nodes and edges.""" """Builds the graph from the vertices and edges."""
self.nodes = self._build_vertices() self.vertices = self._build_vertices()
self.vertex_ids = [vertex.id for vertex in self.vertices]
self.edges = self._build_edges() 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 # This is a hack to make sure that the LLM vertex is sent to
# the toolkit node # the toolkit vertex
self._build_node_params() self._build_vertex_params()
# remove invalid nodes # remove invalid vertices
self._validate_nodes() self._validate_vertices()
def _build_node_params(self) -> None: def _build_vertex_params(self) -> None:
"""Identifies and handles the LLM node within the graph.""" """Identifies and handles the LLM vertex within the graph."""
llm_node = None llm_vertex = None
for node in self.nodes: for vertex in self.vertices:
node._build_params() vertex._build_params()
if isinstance(node, LLMVertex): if isinstance(vertex, LLMVertex):
llm_node = node llm_vertex = vertex
if llm_node: if llm_vertex:
for node in self.nodes: for vertex in self.vertices:
if isinstance(node, ToolkitVertex): if isinstance(vertex, ToolkitVertex):
node.params["llm"] = llm_node vertex.params["llm"] = llm_vertex
def _validate_nodes(self) -> None: def _validate_vertices(self) -> None:
"""Check that all nodes have edges""" """Check that all vertices have edges"""
if len(self.nodes) == 1: if len(self.vertices) == 1:
return return
for node in self.nodes: for vertex in self.vertices:
if not self._validate_node(node): if not self._validate_vertex(vertex):
raise ValueError(f"{node.vertex_type} is not connected to any other components") raise ValueError(f"{vertex.vertex_type} is not connected to any other components")
def _validate_node(self, node: Vertex) -> bool: def _validate_vertex(self, vertex: Vertex) -> bool:
"""Validates a node.""" """Validates a vertex."""
# All nodes that do not have edges are invalid # All vertices that do not have edges are invalid
return len(node.edges) > 0 return len(self.get_vertex_edges(vertex.id)) > 0
def get_node(self, node_id: str) -> Union[None, Vertex]: def get_vertex(self, vertex_id: str) -> Union[None, Vertex]:
"""Returns a node by id.""" """Returns a vertex by id."""
return next((node for node in self.nodes if node.id == node_id), None) return next((vertex for vertex in self.vertices if vertex.id == vertex_id), None)
def get_nodes_with_target(self, node: Vertex) -> List[Vertex]: def get_vertex_edges(self, vertex_id: str) -> List[Edge]:
"""Returns the nodes connected to a node.""" """Returns a list of edges for a given vertex."""
connected_nodes: List[Vertex] = [edge.source for edge in self.edges if edge.target == node] return [edge for edge in self.edges if edge.source_id == vertex_id or edge.target_id == vertex_id]
return connected_nodes
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: async def build(self) -> Chain:
"""Builds the graph.""" """Builds the graph."""
# Get root node # Get root vertex
root_node = payload.get_root_node(self) root_vertex = payload.get_root_vertex(self)
if root_node is None: if root_vertex is None:
raise ValueError("No root node found") raise ValueError("No root vertex found")
return await root_node.build() return await root_vertex.build()
def topological_sort(self) -> List[Vertex]: def topological_sort(self) -> List[Vertex]:
""" """
@ -136,25 +144,25 @@ class Graph:
ValueError: If the graph contains a cycle. ValueError: If the graph contains a cycle.
""" """
# States: 0 = unvisited, 1 = visiting, 2 = visited # 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 = [] sorted_vertices = []
def dfs(node): def dfs(vertex):
if state[node] == 1: if state[vertex] == 1:
# We have a cycle # We have a cycle
raise ValueError("Graph contains a cycle, cannot perform topological sort") raise ValueError("Graph contains a cycle, cannot perform topological sort")
if state[node] == 0: if state[vertex] == 0:
state[node] = 1 state[vertex] = 1
for edge in node.edges: for edge in vertex.edges:
if edge.source == node: if edge.source == vertex:
dfs(edge.target) dfs(edge.target)
state[node] = 2 state[vertex] = 2
sorted_vertices.append(node) sorted_vertices.append(vertex)
# Visit each node # Visit each vertex
for node in self.nodes: for vertex in self.vertices:
if state[node] == 0: if state[vertex] == 0:
dfs(node) dfs(vertex)
return list(reversed(sorted_vertices)) return list(reversed(sorted_vertices))
@ -164,17 +172,21 @@ class Graph:
logger.debug("There are %s vertices in the graph", len(sorted_vertices)) logger.debug("There are %s vertices in the graph", len(sorted_vertices))
yield from sorted_vertices yield from sorted_vertices
def get_node_neighbors(self, node: Vertex) -> Dict[Vertex, int]: def get_vertex_neighbors(self, vertex: Vertex) -> Dict[Vertex, int]:
"""Returns the neighbors of a node.""" """Returns the neighbors of a vertex."""
neighbors: Dict[Vertex, int] = {} neighbors: Dict[Vertex, int] = {}
for edge in self.edges: for edge in self.edges:
if edge.source == node: if edge.source_id == vertex.id:
neighbor = edge.target neighbor = self.get_vertex(edge.target_id)
if neighbor is None:
continue
if neighbor not in neighbors: if neighbor not in neighbors:
neighbors[neighbor] = 0 neighbors[neighbor] = 0
neighbors[neighbor] += 1 neighbors[neighbor] += 1
elif edge.target == node: elif edge.target_id == vertex.id:
neighbor = edge.source neighbor = self.get_vertex(edge.source_id)
if neighbor is None:
continue
if neighbor not in neighbors: if neighbor not in neighbors:
neighbors[neighbor] = 0 neighbors[neighbor] = 0
neighbors[neighbor] += 1 neighbors[neighbor] += 1
@ -182,59 +194,59 @@ class Graph:
def _build_edges(self) -> List[Edge]: def _build_edges(self) -> List[Edge]:
"""Builds the edges of the graph.""" """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 # 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] = [] edges: List[Edge] = []
for edge in self._edges: for edge in self._edges:
source = self.get_node(edge["source"]) source = self.get_vertex(edge["source"])
target = self.get_node(edge["target"]) target = self.get_vertex(edge["target"])
if source is None: 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: 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)) edges.append(Edge(source, target, edge))
return edges return edges
def _get_vertex_class(self, node_type: str, node_lc_type: str) -> Type[Vertex]: def _get_vertex_class(self, vertex_type: str, vertex_lc_type: str) -> Type[Vertex]:
"""Returns the node class based on the node type.""" """Returns the vertex class based on the vertex type."""
if node_type in FILE_TOOLS: if vertex_type in FILE_TOOLS:
return FileToolVertex return FileToolVertex
if node_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP: if vertex_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP:
return lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_type] return lazy_load_vertex_dict.VERTEX_TYPE_MAP[vertex_type]
return ( return (
lazy_load_vertex_dict.VERTEX_TYPE_MAP[node_lc_type] lazy_load_vertex_dict.VERTEX_TYPE_MAP[vertex_lc_type]
if node_lc_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP if vertex_lc_type in lazy_load_vertex_dict.VERTEX_TYPE_MAP
else Vertex else Vertex
) )
def _build_vertices(self) -> List[Vertex]: def _build_vertices(self) -> List[Vertex]:
"""Builds the vertices of the graph.""" """Builds the vertices of the graph."""
nodes: List[Vertex] = [] vertices: List[Vertex] = []
for node in self._nodes: for vertex in self._vertices:
node_data = node["data"] vertex_data = vertex["data"]
node_type: str = node_data["type"] # type: ignore vertex_type: str = vertex_data["type"] # type: ignore
node_lc_type: str = node_data["node"]["template"]["_type"] # type: ignore vertex_lc_type: str = vertex_data["node"]["template"]["_type"] # type: ignore
VertexClass = self._get_vertex_class(node_type, node_lc_type) VertexClass = self._get_vertex_class(vertex_type, vertex_lc_type)
vertex = VertexClass(node) vertex = VertexClass(vertex, graph=self)
vertex.set_top_level(self.top_level_nodes) vertex.set_top_level(self.top_level_vertices)
nodes.append(vertex) vertices.append(vertex)
return nodes return vertices
def get_children_by_node_type(self, node: Vertex, node_type: str) -> List[Vertex]: def get_children_by_vertex_type(self, vertex: Vertex, vertex_type: str) -> List[Vertex]:
"""Returns the children of a node based on the node type.""" """Returns the children of a vertex based on the vertex type."""
children = [] children = []
node_types = [node.data["type"]] vertex_types = [vertex.data["type"]]
if "node" in node.data: if "node" in vertex.data:
node_types += node.data["node"]["base_classes"] vertex_types += vertex.data["node"]["base_classes"]
if node_type in node_types: if vertex_type in vertex_types:
children.append(node) children.append(vertex)
return children return children
def __repr__(self): def __repr__(self):
node_ids = [node.id for node in self.nodes] 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]) 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}" return f"Graph:\nNodes: {vertex_ids}\nConnections:\n{edges_repr}"

View file

@ -1,33 +1,32 @@
import ast import ast
import inspect import inspect
import pickle
import types import types
from typing import TYPE_CHECKING, Any, Dict, List, Optional from typing import TYPE_CHECKING, Any, Dict, List, Optional
from loguru import logger
from langflow.graph.utils import UnbuiltObject from langflow.graph.utils import UnbuiltObject
from langflow.graph.vertex.utils import is_basic_type
from langflow.interface.initialize import loading from langflow.interface.initialize import loading
from langflow.interface.listing import lazy_load_dict from langflow.interface.listing import lazy_load_dict
from langflow.utils.constants import DIRECT_TYPES from langflow.utils.constants import DIRECT_TYPES
from langflow.utils.util import sync_to_async from langflow.utils.util import sync_to_async
from loguru import logger
if TYPE_CHECKING: if TYPE_CHECKING:
from langflow.graph.edge.base import Edge from langflow.graph.edge.base import Edge
from langflow.graph.graph.base import Graph
class Vertex: class Vertex:
def __init__( def __init__(
self, self,
data: Dict, data: Dict,
graph: "Graph",
base_type: Optional[str] = None, base_type: Optional[str] = None,
is_task: bool = False, is_task: bool = False,
params: Optional[Dict] = None, params: Optional[Dict] = None,
) -> None: ) -> None:
self.graph = graph
self.id: str = data["id"] self.id: str = data["id"]
self._data = data self._data = data
self.edges: List["Edge"] = []
self.base_type: Optional[str] = base_type self.base_type: Optional[str] = base_type
self._parse_data() self._parse_data()
self._built_object = UnbuiltObject() self._built_object = UnbuiltObject()
@ -39,43 +38,28 @@ class Vertex:
self.parent_node_id: Optional[str] = self._data.get("parent_node_id") self.parent_node_id: Optional[str] = self._data.get("parent_node_id")
self.parent_is_top_level = False self.parent_is_top_level = False
def reset_params(self): @property
for edge in self.edges: def edges(self) -> List["Edge"]:
if edge.source != self: return self.graph.get_vertex_edges(self.id)
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
def __getstate__(self): def __getstate__(self):
state_dict = self.__dict__.copy() return {
try: "_data": self._data,
# try pickling the built object "params": {},
# if it fails, then we need to delete it "base_type": self.base_type,
# and build it again "is_task": self.is_task,
pickle.dumps(state_dict["_built_object"]) "id": self.id,
except Exception: "_built_object": UnbuiltObject(),
self.reset_params() "_built": False,
del state_dict["_built_object"] "parent_node_id": self.parent_node_id,
del state_dict["_built"] "parent_is_top_level": self.parent_is_top_level,
return state_dict }
def __setstate__(self, state): def __setstate__(self, state):
self._data = state["_data"] self._data = state["_data"]
self.params = state["params"] self.params = state["params"]
self.base_type = state["base_type"] self.base_type = state["base_type"]
self.is_task = state["is_task"] self.is_task = state["is_task"]
self.edges = state["edges"]
self.id = state["id"] self.id = state["id"]
self._parse_data() self._parse_data()
if "_built_object" in state: if "_built_object" in state:
@ -144,6 +128,10 @@ class Vertex:
# and use that as the value for the param # 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 # 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 # 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)} 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 {} params = self.params.copy() if self.params else {}
@ -155,9 +143,9 @@ class Vertex:
if template_dict[param_key]["list"]: if template_dict[param_key]["list"]:
if param_key not in params: if param_key not in params:
params[param_key] = [] params[param_key] = []
params[param_key].append(edge.source) params[param_key].append(self.graph.get_vertex(edge.source_id))
elif edge.target.id == self.id: elif edge.target_id == self.id:
params[param_key] = edge.source params[param_key] = self.graph.get_vertex(edge.source_id)
for key, value in template_dict.items(): for key, value in template_dict.items():
if key in params: if key in params:
@ -177,33 +165,33 @@ class Vertex:
else: else:
raise ValueError(f"File path not found for {self.vertex_type}") raise ValueError(f"File path not found for {self.vertex_type}")
elif value.get("type") in DIRECT_TYPES and params.get(key) is None: elif value.get("type") in DIRECT_TYPES and params.get(key) is None:
val = value.get("value")
if value.get("type") == "code": if value.get("type") == "code":
try: try:
params[key] = ast.literal_eval(value.get("value")) params[key] = ast.literal_eval(val) if val else None
except Exception as exc: except Exception as exc:
logger.debug(f"Error parsing code: {exc}") logger.debug(f"Error parsing code: {exc}")
params[key] = value.get("value") params[key] = val
elif value.get("type") in ["dict", "NestedDict"]: elif value.get("type") in ["dict", "NestedDict"]:
# When dict comes from the frontend it comes as a # When dict comes from the frontend it comes as a
# list of dicts, so we need to convert it to a dict # list of dicts, so we need to convert it to a dict
# before passing it to the build method # before passing it to the build method
_value = value.get("value") if isinstance(val, list):
if isinstance(_value, list):
params[key] = {k: v for item in value.get("value", []) for k, v in item.items()} params[key] = {k: v for item in value.get("value", []) for k, v in item.items()}
elif isinstance(_value, dict): elif isinstance(val, dict):
params[key] = _value params[key] = val
elif value.get("type") == "int" and value.get("value") is not None: elif value.get("type") == "int" and val is not None:
try: try:
params[key] = int(value.get("value")) params[key] = int(val)
except ValueError: except ValueError:
params[key] = value.get("value") params[key] = val
elif value.get("type") == "float" and value.get("value") is not None: elif value.get("type") == "float" and val is not None:
try: try:
params[key] = float(value.get("value")) params[key] = float(val)
except ValueError: except ValueError:
params[key] = value.get("value") params[key] = val
else: else:
params[key] = value.get("value") params[key] = val
if not value.get("required") and params.get(key) is None: if not value.get("required") and params.get(key) is None:
if value.get("default"): if value.get("default"):
@ -266,7 +254,7 @@ class Vertex:
pass pass
# If there's no task_id, build the vertex locally # 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 return self._built_object
async def _build_node_and_update_params(self, key, node, user_id=None): async def _build_node_and_update_params(self, key, node, user_id=None):

View file

@ -1,14 +1,14 @@
import ast import ast
from typing import Any, Dict, List, Optional, Union 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.graph.vertex.base import Vertex
from langflow.interface.utils import extract_input_variables_from_prompt from langflow.interface.utils import extract_input_variables_from_prompt
class AgentVertex(Vertex): class AgentVertex(Vertex):
def __init__(self, data: Dict, params: Optional[Dict] = None): def __init__(self, data: Dict, graph, params: Optional[Dict] = None):
super().__init__(data, base_type="agents", params=params) super().__init__(data, graph=graph, base_type="agents", params=params)
self.tools: List[Union[ToolkitVertex, ToolVertex]] = [] self.tools: List[Union[ToolkitVertex, ToolVertex]] = []
self.chains: List[ChainVertex] = [] self.chains: List[ChainVertex] = []
@ -28,7 +28,7 @@ class AgentVertex(Vertex):
for edge in self.edges: for edge in self.edges:
if not hasattr(edge, "source"): if not hasattr(edge, "source"):
continue continue
source_node = edge.source source_node = self.graph.get_vertex(edge.source_id)
if isinstance(source_node, (ToolVertex, ToolkitVertex)): if isinstance(source_node, (ToolVertex, ToolkitVertex)):
self.tools.append(source_node) self.tools.append(source_node)
elif isinstance(source_node, ChainVertex): elif isinstance(source_node, ChainVertex):
@ -51,16 +51,21 @@ class AgentVertex(Vertex):
class ToolVertex(Vertex): class ToolVertex(Vertex):
def __init__(self, data: Dict, params: Optional[Dict] = None): def __init__(
super().__init__(data, base_type="tools", params=params) self,
data: Dict,
graph,
params: Optional[Dict] = None,
):
super().__init__(data, graph=graph, base_type="tools", params=params)
class LLMVertex(Vertex): class LLMVertex(Vertex):
built_node_type = None built_node_type = None
class_built_object = None class_built_object = None
def __init__(self, data: Dict, params: Optional[Dict] = None): def __init__(self, data: Dict, graph, params: Optional[Dict] = None):
super().__init__(data, base_type="llms", params=params) super().__init__(data, graph=graph, base_type="llms", params=params)
async def build(self, force: bool = False, user_id=None, *args, **kwargs) -> Any: 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 # LLM is different because some models might take up too much memory
@ -77,18 +82,18 @@ class LLMVertex(Vertex):
class ToolkitVertex(Vertex): class ToolkitVertex(Vertex):
def __init__(self, data: Dict, params=None): def __init__(self, data: Dict, graph, params=None):
super().__init__(data, base_type="toolkits", params=params) super().__init__(data, graph=graph, base_type="toolkits", params=params)
class FileToolVertex(ToolVertex): class FileToolVertex(ToolVertex):
def __init__(self, data: Dict, params=None): def __init__(self, data: Dict, graph, params=None):
super().__init__(data, params=params) super().__init__(data, graph=graph, params=params)
class WrapperVertex(Vertex): class WrapperVertex(Vertex):
def __init__(self, data: Dict): def __init__(self, data: Dict, graph):
super().__init__(data, base_type="wrappers") super().__init__(data, graph=graph, base_type="wrappers")
async def build(self, force: bool = False, user_id=None, *args, **kwargs) -> Any: async def build(self, force: bool = False, user_id=None, *args, **kwargs) -> Any:
if not self._built or force: if not self._built or force:
@ -99,14 +104,14 @@ class WrapperVertex(Vertex):
class DocumentLoaderVertex(Vertex): class DocumentLoaderVertex(Vertex):
def __init__(self, data: Dict, params: Optional[Dict] = None): def __init__(self, data: Dict, graph, params: Optional[Dict] = None):
super().__init__(data, base_type="documentloaders", params=params) super().__init__(data, graph=graph, base_type="documentloaders", params=params)
def _built_object_repr(self): def _built_object_repr(self):
# This built_object is a list of documents. Maybe we should # This built_object is a list of documents. Maybe we should
# show how many documents are in the list? # 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( avg_length = sum(len(doc.page_content) for doc in self._built_object if hasattr(doc, "page_content")) / len(
self._built_object self._built_object
) )
@ -117,28 +122,19 @@ class DocumentLoaderVertex(Vertex):
class EmbeddingVertex(Vertex): class EmbeddingVertex(Vertex):
def __init__(self, data: Dict, params: Optional[Dict] = None): def __init__(self, data: Dict, graph, params: Optional[Dict] = None):
super().__init__(data, base_type="embeddings", params=params) super().__init__(data, graph=graph, base_type="embeddings", params=params)
class VectorStoreVertex(Vertex): class VectorStoreVertex(Vertex):
def __init__(self, data: Dict, params=None): def __init__(self, data: Dict, graph, params=None):
super().__init__(data, base_type="vectorstores") super().__init__(data, graph=graph, base_type="vectorstores")
self.params = params or {} self.params = params or {}
# VectorStores may contain databse connections # VectorStores may contain databse connections
# so we need to define the __reduce__ method and the __setstate__ method # so we need to define the __reduce__ method and the __setstate__ method
# to avoid pickling errors # 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): def remove_docs_and_texts_from_params(self):
# remove documents and texts from params # remove documents and texts from params
@ -146,17 +142,16 @@ class VectorStoreVertex(Vertex):
self.params.pop("documents", None) self.params.pop("documents", None)
self.params.pop("texts", None) self.params.pop("texts", None)
def __getstate__(self): # def __getstate__(self):
# We want to save the params attribute # # We want to save the params attribute
# and if "documents" or "texts" are in the params # # and if "documents" or "texts" are in the params
# we want to remove them because they have already # # we want to remove them because they have already
# been processed. # # been processed.
params = self.params.copy() # params = self.params.copy()
params.pop("documents", None) # params.pop("documents", None)
params.pop("texts", None) # params.pop("texts", None)
self.clean_edges_for_pickling()
return super().__getstate__() # return super().__getstate__()
def __setstate__(self, state): def __setstate__(self, state):
super().__setstate__(state) super().__setstate__(state)
@ -164,24 +159,24 @@ class VectorStoreVertex(Vertex):
class MemoryVertex(Vertex): class MemoryVertex(Vertex):
def __init__(self, data: Dict): def __init__(self, data: Dict, graph):
super().__init__(data, base_type="memory") super().__init__(data, graph=graph, base_type="memory")
class RetrieverVertex(Vertex): class RetrieverVertex(Vertex):
def __init__(self, data: Dict): def __init__(self, data: Dict, graph):
super().__init__(data, base_type="retrievers") super().__init__(data, graph=graph, base_type="retrievers")
class TextSplitterVertex(Vertex): class TextSplitterVertex(Vertex):
def __init__(self, data: Dict, params: Optional[Dict] = None): def __init__(self, data: Dict, graph, params: Optional[Dict] = None):
super().__init__(data, base_type="textsplitters", params=params) super().__init__(data, graph=graph, base_type="textsplitters", params=params)
def _built_object_repr(self): def _built_object_repr(self):
# This built_object is a list of documents. Maybe we should # This built_object is a list of documents. Maybe we should
# show how many documents are in the list? # 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) 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) return f"""{self.vertex_type}({len(self._built_object)} documents)
\nAvg. Document Length (characters): {int(avg_length)} \nAvg. Document Length (characters): {int(avg_length)}
@ -190,8 +185,8 @@ class TextSplitterVertex(Vertex):
class ChainVertex(Vertex): class ChainVertex(Vertex):
def __init__(self, data: Dict): def __init__(self, data: Dict, graph):
super().__init__(data, base_type="chains") super().__init__(data, graph=graph, base_type="chains")
async def build( async def build(
self, self,
@ -220,8 +215,8 @@ class ChainVertex(Vertex):
class PromptVertex(Vertex): class PromptVertex(Vertex):
def __init__(self, data: Dict): def __init__(self, data: Dict, graph):
super().__init__(data, base_type="prompts") super().__init__(data, graph=graph, base_type="prompts")
async def build( async def build(
self, self,
@ -271,9 +266,13 @@ class PromptVertex(Vertex):
# so the prompt format doesn't break # so the prompt format doesn't break
artifacts.pop("handle_keys", None) artifacts.pop("handle_keys", None)
try: 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 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 template = self._built_object.template
for key, value in artifacts.items(): for key, value in artifacts.items():
if value: if value:
@ -285,13 +284,13 @@ class PromptVertex(Vertex):
class OutputParserVertex(Vertex): class OutputParserVertex(Vertex):
def __init__(self, data: Dict): def __init__(self, data: Dict, graph):
super().__init__(data, base_type="output_parsers") super().__init__(data, graph=graph, base_type="output_parsers")
class CustomComponentVertex(Vertex): class CustomComponentVertex(Vertex):
def __init__(self, data: Dict): def __init__(self, data: Dict, graph):
super().__init__(data, base_type="custom_components", is_task=True) super().__init__(data, graph=graph, base_type="custom_components", is_task=True)
def _built_object_repr(self): def _built_object_repr(self):
if self.task_id and self.is_task: if self.task_id and self.is_task:

View file

@ -25,6 +25,7 @@ class Component:
code: Optional[str] = None code: Optional[str] = None
_function_entrypoint_name: str = "build" _function_entrypoint_name: str = "build"
field_config: dict = {} field_config: dict = {}
_user_id: Optional[str]
def __init__(self, **data): def __init__(self, **data):
self.cache = TTLCache(maxsize=1024, ttl=60) self.cache = TTLCache(maxsize=1024, ttl=60)
@ -36,8 +37,7 @@ class Component:
def __setattr__(self, key, value): def __setattr__(self, key, value):
if key == "_user_id" and hasattr(self, "_user_id"): if key == "_user_id" and hasattr(self, "_user_id"):
warnings.warn("Modification of user_id is not allowed") warnings.warn("user_id is immutable and cannot be changed.")
else:
super().__setattr__(key, value) super().__setattr__(key, value)
@cachedmethod(cache=operator.attrgetter("cache")) @cachedmethod(cache=operator.attrgetter("cache"))

View file

@ -5,7 +5,6 @@ from uuid import UUID
import yaml import yaml
from cachetools import TTLCache, cachedmethod from cachetools import TTLCache, cachedmethod
from fastapi import HTTPException from fastapi import HTTPException
from langflow.field_typing.constants import CUSTOM_COMPONENT_SUPPORTED_TYPES from langflow.field_typing.constants import CUSTOM_COMPONENT_SUPPORTED_TYPES
from langflow.interface.custom.component import Component from langflow.interface.custom.component import Component
from langflow.interface.custom.directory_reader import DirectoryReader 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) return await build_sorted_vertices(graph_data, self.user_id)
def list_flows(self, *, get_session: Optional[Callable] = None) -> List[Flow]: 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") raise ValueError("Session is invalid")
try: try:
get_session = get_session or session_getter get_session = get_session or session_getter

View file

@ -28,16 +28,16 @@ def extract_input_variables(nodes):
return nodes return nodes
def get_root_node(graph): def get_root_vertex(graph):
""" """
Returns the root node of the template. 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: if not incoming_edges and len(graph.vertices) == 1:
return graph.nodes[0] 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: def build_json(root, graph) -> Dict:

View file

@ -1,7 +1,7 @@
import ast import ast
import contextlib import contextlib
import importlib import importlib
import types from types import FunctionType
from typing import Dict from typing import Dict
@ -61,7 +61,7 @@ def eval_function(function_string: str):
( (
obj obj
for name, obj in namespace.items() for name, obj in namespace.items()
if isinstance(obj, types.FunctionType) and obj.__code__.co_filename == "<string>" if isinstance(obj, FunctionType) and obj.__code__.co_filename == "<string>"
), ),
None, None,
) )

View file

@ -1,7 +1,7 @@
import json import json
from langflow.graph import Graph
import pytest import pytest
from langflow.graph import Graph
def get_graph(_type="basic"): def get_graph(_type="basic"):
@ -41,5 +41,5 @@ def langchain_objects_are_equal(obj1, obj2):
def test_build_graph(client, basic_data_graph): def test_build_graph(client, basic_data_graph):
graph = Graph.from_payload(basic_data_graph) graph = Graph.from_payload(basic_data_graph)
assert graph is not None 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"]) assert len(graph.edges) == len(basic_data_graph["edges"])

View file

@ -7,10 +7,7 @@ from fastapi import HTTPException
from langflow.field_typing.constants import Data from langflow.field_typing.constants import Data
from langflow.interface.custom.base import CustomComponent from langflow.interface.custom.base import CustomComponent
from langflow.interface.custom.code_parser import CodeParser, CodeSyntaxError from langflow.interface.custom.code_parser import CodeParser, CodeSyntaxError
from langflow.interface.custom.component import ( from langflow.interface.custom.component import Component, ComponentCodeNullError
Component,
ComponentCodeNullError,
)
from langflow.services.database.models.flow import Flow, FlowCreate from langflow.services.database.models.flow import Flow, FlowCreate
code_default = """ code_default = """
@ -445,7 +442,7 @@ def test_custom_component_build_not_implemented():
def test_build_config_no_code(): def test_build_config_no_code():
component = CustomComponent(code=None) component = CustomComponent(code=None)
assert component.get_function_entrypoint_args == "" assert component.get_function_entrypoint_args == []
assert component.get_function_entrypoint_return_type == [] assert component.get_function_entrypoint_return_type == []

View file

@ -24,7 +24,7 @@ from langflow.graph.utils import UnbuiltObject
from langflow.graph.vertex.base import Vertex from langflow.graph.vertex.base import Vertex
from langflow.graph.vertex.types import FileToolVertex, LLMVertex, ToolkitVertex from langflow.graph.vertex.types import FileToolVertex, LLMVertex, ToolkitVertex
from langflow.processing.process import get_result_and_thought 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 # 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]: def get_node_by_type(graph, node_type: Type[Vertex]) -> Union[Vertex, None]:
"""Get a node by type""" """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): def test_graph_structure(basic_graph):
assert isinstance(basic_graph, Graph) assert isinstance(basic_graph, Graph)
assert len(basic_graph.nodes) > 0 assert len(basic_graph.vertices) > 0
assert len(basic_graph.edges) > 0 assert len(basic_graph.edges) > 0
for node in basic_graph.nodes: for node in basic_graph.vertices:
assert isinstance(node, Vertex) assert isinstance(node, Vertex)
for edge in basic_graph.edges: for edge in basic_graph.edges:
assert isinstance(edge, Edge) assert isinstance(edge, Edge)
assert edge.source in basic_graph.nodes assert edge.source_id in basic_graph.vertex_ids
assert edge.target in basic_graph.nodes assert edge.target_id in basic_graph.vertex_ids
def test_circular_dependencies(basic_graph): def test_circular_dependencies(basic_graph):
@ -90,7 +90,7 @@ def test_circular_dependencies(basic_graph):
def check_circular(node, visited): def check_circular(node, visited):
visited.add(node) visited.add(node)
neighbors = basic_graph.get_nodes_with_target(node) neighbors = basic_graph.get_vertices_with_target(node)
for neighbor in neighbors: for neighbor in neighbors:
if neighbor in visited: if neighbor in visited:
return True return True
@ -98,7 +98,7 @@ def test_circular_dependencies(basic_graph):
return True return True
return False return False
for node in basic_graph.nodes: for node in basic_graph.vertices:
assert not check_circular(node, set()) assert not check_circular(node, set())
@ -123,13 +123,13 @@ def test_invalid_node_types():
Graph(graph_data["nodes"], graph_data["edges"]) 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""" """Test getting connected nodes"""
assert isinstance(basic_graph, Graph) assert isinstance(basic_graph, Graph)
# Get root node # Get root node
root = get_root_node(basic_graph) root = get_root_vertex(basic_graph)
assert root is not None 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 assert connected_nodes is not None
@ -138,9 +138,9 @@ def test_get_node_neighbors_basic(basic_graph):
assert isinstance(basic_graph, Graph) assert isinstance(basic_graph, Graph)
# Get root node # Get root node
root = get_root_node(basic_graph) root = get_root_vertex(basic_graph)
assert root is not None 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 neighbors is not None
assert isinstance(neighbors, dict) assert isinstance(neighbors, dict)
# Root Node is an Agent, it requires an LLMChain and tools # 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): def test_get_node(basic_graph):
"""Test getting a single node""" """Test getting a single node"""
node_id = basic_graph.nodes[0].id node_id = basic_graph.vertices[0].id
node = basic_graph.get_node(node_id) node = basic_graph.get_vertex(node_id)
assert isinstance(node, Vertex) assert isinstance(node, Vertex)
assert node.id == node_id assert node.id == node_id
@ -162,8 +162,8 @@ def test_get_node(basic_graph):
def test_build_nodes(basic_graph): def test_build_nodes(basic_graph):
"""Test building nodes""" """Test building nodes"""
assert len(basic_graph.nodes) == len(basic_graph._nodes) assert len(basic_graph.vertices) == len(basic_graph._vertices)
for node in basic_graph.nodes: for node in basic_graph.vertices:
assert isinstance(node, Vertex) assert isinstance(node, Vertex)
@ -172,20 +172,21 @@ def test_build_edges(basic_graph):
assert len(basic_graph.edges) == len(basic_graph._edges) assert len(basic_graph.edges) == len(basic_graph._edges)
for edge in basic_graph.edges: for edge in basic_graph.edges:
assert isinstance(edge, Edge) 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""" """Test getting root node"""
assert isinstance(basic_graph, Graph) assert isinstance(basic_graph, Graph)
root = get_root_node(basic_graph) root = get_root_vertex(basic_graph)
assert root is not None assert root is not None
assert isinstance(root, Vertex) assert isinstance(root, Vertex)
assert root.data["type"] == "TimeTravelGuideChain" assert root.data["type"] == "TimeTravelGuideChain"
# For complex example, the root node is a ZeroShotAgent too # For complex example, the root node is a ZeroShotAgent too
assert isinstance(complex_graph, Graph) assert isinstance(complex_graph, Graph)
root = get_root_node(complex_graph) root = get_root_vertex(complex_graph)
assert root is not None assert root is not None
assert isinstance(root, Vertex) assert isinstance(root, Vertex)
assert root.data["type"] == "ZeroShotAgent" 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 # 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) assert all(edge.matched_type in edge.source_types for edge in basic_graph.edges)
# Get the root node # Get the root node
root = get_root_node(basic_graph) root = get_root_vertex(basic_graph)
# Root node is a TimeTravelGuideChain # Root node is a TimeTravelGuideChain
# which requires an llm and memory # which requires an llm and memory
assert root is not None assert root is not None
@ -278,7 +279,7 @@ async def test_file_tool_node_build(client, openapi_graph):
assert Path(file_path).exists() assert Path(file_path).exists()
file_tool_node = get_node_by_type(openapi_graph, FileToolVertex) 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() built_object = await file_tool_node.build()
assert built_object is not UnbuiltObject assert built_object is not UnbuiltObject
# Remove the file # Remove the file
@ -301,7 +302,7 @@ async def test_get_result_and_thought(basic_graph):
llm_node._built = True llm_node._built = True
langchain_object = await basic_graph.build() langchain_object = await basic_graph.build()
# assert all nodes are built # 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 # now build again and check if FakeListLLM was used
# Get the result and thought # 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) 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) 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"]["show"] is True
assert node1_updated["data"]["node"]["template"]["some_field"]["advanced"] is False assert node1_updated["data"]["node"]["template"]["some_field"]["advanced"] is False
assert node1_updated["data"]["node"]["template"]["some_field"]["display_name"] == "Name1" 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"]["show"] is False
assert node2_updated["data"]["node"]["template"]["other_field"]["advanced"] is True assert node2_updated["data"]["node"]["template"]["other_field"]["advanced"] is True
assert node2_updated["data"]["node"]["template"]["other_field"]["display_name"] == "DisplayName2" 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) loaded_json = json.loads(json_vector_store)
graph = Graph.from_payload(loaded_json) graph = Graph.from_payload(loaded_json)
assert isinstance(graph, Graph) assert isinstance(graph, Graph)
for vertex in graph.nodes: for vertex in graph.vertices:
await vertex.build() await vertex.build()
pickled = pickle.dumps(vertex) pickled = pickle.dumps(vertex)
assert pickled is not UnbuiltObject assert pickled is not UnbuiltObject