fix: add tests to cycles in Graph and improve error handling (#3628)

* Add cycle detection and handling in graph edge building process

- Introduced `cycles` property to detect cycles in the graph.
- Modified `_build_edges` and `build_edge` methods to differentiate between `CycleEdge` and `Edge`.
- Updated imports and type hints to support new functionality.

* Add cycle detection and handling in graph processing

- Introduced `is_cyclic` property to check for cycles in the graph.
- Added `_snapshot` method for capturing the current state of the graph.
- Modified `layered_topological_sort` to handle cyclic graphs by starting from a specified start component.
- Updated imports and type hints for better code clarity and functionality.

* Refactor tests and components for improved caching and data handling

- Updated `test_vector_store_rag.py` to use `set_on_output` with `cache=True` and simplified assertions.
- Enhanced `test_memory_chatbot.py` with additional assertions for graph structure and caching.
- Simplified `to_data` method in `base.py` to directly return `_data` without JSON serialization.

* Add unit tests for detecting cycles in graph

- Introduce `test_cycle_in_graph` to verify cyclic behavior in the graph.
- Add `test_cycle_in_graph_max_iterations` to ensure max iterations limit is respected.
- Implement `Concatenate` component for testing purposes.

* Disable output cache in graph tests to allow loops to work

* Refactor: Update VertexStates enum values to uppercase and optimize imports in base.py

* Refactor type hints and improve error handling in `Vertex` class

- Replace `ValueError` with `NoComponentInstance` exception for missing component instances.
- Add `target_handle_name` parameter to `_get_result` method for better result retrieval.
- Refactor type hints to use `collections.abc` for `AsyncIterator`, `Generator`, and `Iterator`.
- Update type hints for `extract_messages_from_artifacts` and `successors_ids` methods to use generic `dict` and `list`.
This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-09-02 08:45:37 -03:00 • committed by GitHub
commit bc6e918f49
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 222 additions and 55 deletions

View file

@ -4,24 +4,31 @@ import json
import uuid
import warnings
from collections import defaultdict, deque
from collections.abc import Generator, Iterable
from datetime import datetime, timezone
from functools import partial
from itertools import chain
from typing import TYPE_CHECKING, Any, Optional
from collections.abc import Generator
from typing import TYPE_CHECKING, Any, Optional, cast
import nest_asyncio
from loguru import logger
from langflow.exceptions.component import ComponentBuildException
from langflow.graph.edge.base import CycleEdge
from langflow.graph.edge.base import CycleEdge, Edge
from langflow.graph.edge.schema import EdgeData
from langflow.graph.graph.constants import Finish, lazy_load_vertex_dict
from langflow.graph.graph.runnable_vertices_manager import RunnableVerticesManager
from langflow.graph.graph.schema import GraphData, GraphDump, StartConfigDict, VertexBuildResult
from langflow.graph.graph.state_manager import GraphStateManager
from langflow.graph.graph.state_model import create_state_model_from_graph
from langflow.graph.graph.utils import find_start_component_id, process_flow, should_continue, sort_up_to_vertex
from langflow.graph.graph.utils import (
find_all_cycle_edges,
find_start_component_id,
has_cycle,
process_flow,
should_continue,
sort_up_to_vertex,
)
from langflow.graph.schema import InterfaceComponentTypes, RunOutputs
from langflow.graph.vertex.base import Vertex, VertexStates
from langflow.graph.vertex.schema import NodeData
@ -279,6 +286,15 @@ class Graph:
raise ValueError("Max iterations reached")
def _snapshot(self):
return {
"_run_queue": self._run_queue.copy(),
"_first_layer": self._first_layer.copy(),
"vertices_layers": copy.deepcopy(self.vertices_layers),
"vertices_to_run": copy.deepcopy(self.vertices_to_run),
"run_manager": copy.deepcopy(self.run_manager.to_dict()),
}
def __apply_config(self, config: StartConfigDict):
for vertex in self.vertices:
if vertex._custom_component is None:
@ -460,6 +476,23 @@ class Graph:
raise ValueError("Graph not prepared. Call prepare() first.")
return self._first_layer
@property
def is_cyclic(self):
"""
Check if the graph has any cycles.
Returns:
bool: True if the graph has any cycles, False otherwise.
"""
if self._is_cyclic is None:
vertices = [vertex.id for vertex in self.vertices]
try:
edges = [(e["data"]["sourceHandle"]["id"], e["data"]["targetHandle"]["id"]) for e in self._edges]
except KeyError:
edges = [(e["source"], e["target"]) for e in self._edges]
self._is_cyclic = has_cycle(vertices, edges)
return self._is_cyclic
@property
def run_id(self):
"""
@ -1380,7 +1413,7 @@ class Graph:
def find_next_runnable_vertices(self, vertex_id: str, vertex_successors_ids: list[str]) -> list[str]:
next_runnable_vertices = set()
for v_id in vertex_successors_ids:
for v_id in sorted(vertex_successors_ids):
if not self.is_vertex_runnable(v_id):
next_runnable_vertices.update(self.find_runnable_predecessors_for_successor(v_id))
else:
@ -1536,21 +1569,31 @@ class Graph:
neighbors[neighbor] += 1
return neighbors
@property
def cycles(self):
if self._cycles is None:
if self._start is None:
self._cycles = []
else:
entry_vertex = self._start._id
edges = [(e["data"]["sourceHandle"]["id"], e["data"]["targetHandle"]["id"]) for e in self._edges]
self._cycles = find_all_cycle_edges(entry_vertex, edges)
return self._cycles
def _build_edges(self) -> list[CycleEdge]:
"""Builds the edges of the graph."""
# 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 vertex, we raise an error
edges: set[CycleEdge] = set()
edges: set[CycleEdge | Edge] = set()
for edge in self._edges:
new_edge = self.build_edge(edge)
edges.add(new_edge)
if self.vertices and not edges:
warnings.warn("Graph has vertices but no edges")
return list(edges)
return list(cast(Iterable[CycleEdge], edges))
def build_edge(self, edge: EdgeData) -> CycleEdge:
def build_edge(self, edge: EdgeData) -> CycleEdge | Edge:
source = self.get_vertex(edge["source"])
target = self.get_vertex(edge["target"])
@ -1558,7 +1601,10 @@ class Graph:
raise ValueError(f"Source vertex {edge['source']} not found")
if target is None:
raise ValueError(f"Target vertex {edge['target']} not found")
new_edge = CycleEdge(source, target, edge)
if (source.id, target.id) in self.cycles:
new_edge: CycleEdge | Edge = CycleEdge(source, target, edge)
else:
new_edge = Edge(source, target, edge)
return new_edge
def _get_vertex_class(self, node_type: str, node_base_type: str, node_id: str) -> type["Vertex"]:
@ -1608,7 +1654,6 @@ class Graph:
if stop_component_id and start_component_id:
raise ValueError("You can only provide one of stop_component_id or start_component_id")
self.validate_stream()
self.edges = self._build_edges()
if stop_component_id or start_component_id:
try:
@ -1658,12 +1703,25 @@ class Graph:
"""Performs a layered topological sort of the vertices in the graph."""
vertices_ids = {vertex.id for vertex in vertices}
# Queue for vertices with no incoming edges
queue = deque(
vertex.id
for vertex in vertices
# if filter_graphs then only vertex.is_input will be considered
if self.in_degree_map[vertex.id] == 0 and (not filter_graphs or vertex.is_input)
)
in_degree_map = self.in_degree_map.copy()
if self.is_cyclic and all(in_degree_map.values()):
# This means we have a cycle because all vertex have in_degree_map > 0
# because of this we set the queue to start on the ._start if it exists
if self._start is not None:
queue = deque([self._start._id])
else:
# Find the chat input component
chat_input = find_start_component_id(vertices_ids)
if chat_input is None:
raise ValueError("No input component found and no start component provided")
queue = deque([chat_input])
else:
queue = deque(
vertex.id
for vertex in vertices
# if filter_graphs then only vertex.is_input will be considered
if in_degree_map[vertex.id] == 0 and (not filter_graphs or vertex.is_input)
)
layers: list[list[str]] = []
visited = set(queue)
@ -1684,13 +1742,13 @@ class Graph:
if neighbor not in vertices_ids:
continue
self.in_degree_map[neighbor] -= 1 # 'remove' edge
if self.in_degree_map[neighbor] == 0 and neighbor not in visited:
in_degree_map[neighbor] -= 1 # 'remove' edge
if in_degree_map[neighbor] == 0 and neighbor not in visited:
queue.append(neighbor)
# if > 0 it might mean not all predecessors have added to the queue
# so we should process the neighbors predecessors
elif self.in_degree_map[neighbor] > 0:
elif in_degree_map[neighbor] > 0:
for predecessor in self.predecessor_map[neighbor]:
if predecessor not in queue and predecessor not in visited:
queue.append(predecessor)

View file

@ -4,10 +4,9 @@ import inspect
import os
import traceback
import types
import json
from collections.abc import AsyncIterator, Callable, Iterator, Mapping
from enum import Enum
from typing import TYPE_CHECKING, Any, Optional
from collections.abc import AsyncIterator, Callable, Iterator, Mapping
import pandas as pd
from loguru import logger
@ -37,9 +36,9 @@ if TYPE_CHECKING:
class VertexStates(str, Enum):
"""Vertex are related to it being active, inactive, or in an error state."""
ACTIVE = "active"
INACTIVE = "inactive"
ERROR = "error"
ACTIVE = "ACTIVE"
INACTIVE = "INACTIVE"
ERROR = "ERROR"
class Vertex:
@ -105,12 +104,7 @@ class Vertex:
self._custom_component._set_input_value(name, value)
def to_data(self):
try:
data = json.loads(json.dumps(self._data, default=str))
except TypeError:
data = self._data
return data
return self._data
def add_component_instance(self, component_instance: "Component"):
component_instance.set_vertex(self)

View file

@ -102,8 +102,6 @@ class ComponentVertex(Vertex):
)
for edge in self.get_edge_with_target(requester.id):
# We need to check if the edge is a normal edge
# or a contract edge
if edge.is_cycle and edge.target_param:
return requester.get_value_from_template_dict(edge.target_param)