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:
parent
751edcf5dc
commit
bc6e918f49
6 changed files with 222 additions and 55 deletions
111
src/backend/tests/unit/graph/graph/test_cycles.py
Normal file
111
src/backend/tests/unit/graph/graph/test_cycles.py
Normal file
|
|
@ -0,0 +1,111 @@
|
|||
import pytest
|
||||
|
||||
from langflow.components.inputs.ChatInput import ChatInput
|
||||
from langflow.components.outputs.ChatOutput import ChatOutput
|
||||
from langflow.components.outputs.TextOutput import TextOutputComponent
|
||||
from langflow.components.prototypes.ConditionalRouter import ConditionalRouterComponent
|
||||
from langflow.custom.custom_component.component import Component
|
||||
from langflow.graph.graph.base import Graph
|
||||
from langflow.io import MessageTextInput, Output
|
||||
from langflow.schema.message import Message
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client():
|
||||
pass
|
||||
|
||||
|
||||
class Concatenate(Component):
|
||||
display_name = "Concatenate"
|
||||
description = "Concatenates two strings"
|
||||
|
||||
inputs = [
|
||||
MessageTextInput(name="text", display_name="Text", required=True),
|
||||
]
|
||||
outputs = [
|
||||
Output(display_name="Text", name="some_text", method="concatenate"),
|
||||
]
|
||||
|
||||
def concatenate(self) -> Message:
|
||||
return Message(text=f"{self.text}{self.text}" or "test")
|
||||
|
||||
|
||||
def test_cycle_in_graph():
|
||||
chat_input = ChatInput(_id="chat_input")
|
||||
router = ConditionalRouterComponent(_id="router")
|
||||
chat_input.set(input_value=router.false_response)
|
||||
concat_component = Concatenate(_id="concatenate")
|
||||
concat_component.set(text=chat_input.message_response)
|
||||
router.set(
|
||||
input_text=chat_input.message_response,
|
||||
match_text="testtesttesttest",
|
||||
operator="equals",
|
||||
message=concat_component.concatenate,
|
||||
)
|
||||
text_output = TextOutputComponent(_id="text_output")
|
||||
text_output.set(input_value=router.true_response)
|
||||
chat_output = ChatOutput(_id="chat_output")
|
||||
chat_output.set(input_value=text_output.text_response)
|
||||
|
||||
graph = Graph(chat_input, chat_output)
|
||||
assert graph.is_cyclic is True
|
||||
|
||||
# Run queue should contain chat_input and not router
|
||||
assert "chat_input" in graph._run_queue
|
||||
assert "router" not in graph._run_queue
|
||||
results = []
|
||||
max_iterations = 20
|
||||
snapshots = [graph._snapshot()]
|
||||
for result in graph.start(max_iterations=max_iterations, config={"output": {"cache": False}}):
|
||||
snapshots.append(graph._snapshot())
|
||||
results.append(result)
|
||||
results_ids = [result.vertex.id for result in results if hasattr(result, "vertex")]
|
||||
assert results_ids[-2:] == ["text_output", "chat_output"]
|
||||
assert len(results_ids) > len(graph.vertices), snapshots
|
||||
# Check that chat_output and text_output are the last vertices in the results
|
||||
assert results_ids == [
|
||||
"chat_input",
|
||||
"concatenate",
|
||||
"router",
|
||||
"chat_input",
|
||||
"concatenate",
|
||||
"router",
|
||||
"chat_input",
|
||||
"concatenate",
|
||||
"router",
|
||||
"chat_input",
|
||||
"concatenate",
|
||||
"router",
|
||||
"text_output",
|
||||
"chat_output",
|
||||
], f"Results: {results_ids}"
|
||||
|
||||
|
||||
def test_cycle_in_graph_max_iterations():
|
||||
chat_input = ChatInput(_id="chat_input")
|
||||
router = ConditionalRouterComponent(_id="router")
|
||||
chat_input.set(input_value=router.false_response)
|
||||
concat_component = Concatenate(_id="concatenate")
|
||||
concat_component.set(text=chat_input.message_response)
|
||||
router.set(
|
||||
input_text=chat_input.message_response,
|
||||
match_text="testtesttesttest",
|
||||
operator="equals",
|
||||
message=concat_component.concatenate,
|
||||
)
|
||||
text_output = TextOutputComponent(_id="text_output")
|
||||
text_output.set(input_value=router.true_response)
|
||||
chat_output = ChatOutput(_id="chat_output")
|
||||
chat_output.set(input_value=text_output.text_response)
|
||||
|
||||
graph = Graph(chat_input, chat_output)
|
||||
assert graph.is_cyclic is True
|
||||
|
||||
# Run queue should contain chat_input and not router
|
||||
assert "chat_input" in graph._run_queue
|
||||
assert "router" not in graph._run_queue
|
||||
results = []
|
||||
|
||||
with pytest.raises(ValueError, match="Max iterations reached"):
|
||||
for result in graph.start(max_iterations=2, config={"output": {"cache": False}}):
|
||||
results.append(result)
|
||||
|
|
@ -30,23 +30,35 @@ AI: """
|
|||
openai_component.set(
|
||||
input_value=prompt_component.build_prompt, max_tokens=100, temperature=0.1, api_key="test_api_key"
|
||||
)
|
||||
openai_component.get_output("text_output").value = "Mock response"
|
||||
openai_component.set_on_output(name="text_output", value="Mock response", cache=True)
|
||||
|
||||
chat_output = ChatOutput(_id="chat_output")
|
||||
chat_output.set(input_value=openai_component.text_response)
|
||||
|
||||
graph = Graph(chat_input, chat_output)
|
||||
assert graph.in_degree_map == {"chat_output": 1, "prompt": 2, "openai": 1, "chat_input": 0, "chat_memory": 0}
|
||||
return graph
|
||||
|
||||
|
||||
def test_memory_chatbot(memory_chatbot_graph):
|
||||
# Now we run step by step
|
||||
expected_order = deque(["chat_input", "chat_memory", "prompt", "openai", "chat_output"])
|
||||
assert memory_chatbot_graph.in_degree_map == {
|
||||
"chat_output": 1,
|
||||
"prompt": 2,
|
||||
"openai": 1,
|
||||
"chat_input": 0,
|
||||
"chat_memory": 0,
|
||||
}
|
||||
assert memory_chatbot_graph.vertices_layers == [["prompt"], ["openai"], ["chat_output"]]
|
||||
assert memory_chatbot_graph.first_layer == ["chat_input", "chat_memory"]
|
||||
|
||||
for step in expected_order:
|
||||
result = memory_chatbot_graph.step()
|
||||
if isinstance(result, Finish):
|
||||
break
|
||||
assert step == result.vertex.id
|
||||
|
||||
assert step == result.vertex.id, (memory_chatbot_graph.in_degree_map, memory_chatbot_graph.vertices_layers)
|
||||
|
||||
|
||||
def test_memory_chatbot_dump_structure(memory_chatbot_graph: Graph):
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
import copy
|
||||
from collections import Counter, defaultdict
|
||||
from textwrap import dedent
|
||||
|
||||
import pytest
|
||||
|
|
@ -15,7 +14,6 @@ from langflow.components.prompts.Prompt import PromptComponent
|
|||
from langflow.components.vectorstores.AstraDB import AstraVectorStoreComponent
|
||||
from langflow.graph.graph.base import Graph
|
||||
from langflow.graph.graph.constants import Finish
|
||||
from langflow.graph.graph.schema import VertexBuildResult
|
||||
from langflow.schema.data import Data
|
||||
|
||||
|
||||
|
|
@ -29,7 +27,7 @@ def ingestion_graph():
|
|||
# Ingestion Graph
|
||||
file_component = FileComponent(_id="file-123")
|
||||
file_component.set(path="test.txt")
|
||||
file_component.set_on_output("data", value=Data(text="This is a test file."))
|
||||
file_component.set_on_output(name="data", value=Data(text="This is a test file."), cache=True)
|
||||
text_splitter = SplitTextComponent(_id="text-splitter-123")
|
||||
text_splitter.set(data_inputs=file_component.load_file)
|
||||
openai_embeddings = OpenAIEmbeddingsComponent(_id="openai-embeddings-123")
|
||||
|
|
@ -43,9 +41,10 @@ def ingestion_graph():
|
|||
api_endpoint="https://astra.example.com",
|
||||
token="token",
|
||||
)
|
||||
vector_store.set_on_output("vector_store", value="mock_vector_store")
|
||||
vector_store.set_on_output("base_retriever", value="mock_retriever")
|
||||
vector_store.set_on_output("search_results", value=[Data(text="This is a test file.")])
|
||||
vector_store.set_on_output(name="vector_store", value="mock_vector_store", cache=True)
|
||||
vector_store.set_on_output(name="base_retriever", value="mock_retriever", cache=True)
|
||||
vector_store.set_on_output(name="search_results", value=[Data(text="This is a test file.")], cache=True)
|
||||
|
||||
ingestion_graph = Graph(file_component, vector_store)
|
||||
return ingestion_graph
|
||||
|
||||
|
|
@ -65,14 +64,15 @@ def rag_graph():
|
|||
)
|
||||
# Mock search_documents
|
||||
rag_vector_store.set_on_output(
|
||||
"search_results",
|
||||
name="search_results",
|
||||
value=[
|
||||
Data(data={"text": "Hello, world!"}),
|
||||
Data(data={"text": "Goodbye, world!"}),
|
||||
],
|
||||
cache=True,
|
||||
)
|
||||
rag_vector_store.set_on_output("base_retriever", value="mock_retriever")
|
||||
rag_vector_store.set_on_output("vector_store", value="mock_vector_store")
|
||||
rag_vector_store.set_on_output(name="vector_store", value="mock_vector_store", cache=True)
|
||||
rag_vector_store.set_on_output(name="base_retriever", value="mock_retriever", cache=True)
|
||||
parse_data = ParseDataComponent(_id="parse-data-123")
|
||||
parse_data.set(data=rag_vector_store.search_documents)
|
||||
prompt_component = PromptComponent(_id="prompt-123")
|
||||
|
|
@ -88,7 +88,7 @@ def rag_graph():
|
|||
|
||||
openai_component = OpenAIModelComponent(_id="openai-123")
|
||||
openai_component.set(api_key="sk-123", openai_api_base="https://api.openai.com/v1")
|
||||
openai_component.set_on_output("text_output", value="Hello, world!")
|
||||
openai_component.set_on_output(name="text_output", value="Hello, world!", cache=True)
|
||||
openai_component.set(input_value=prompt_component.build_prompt)
|
||||
|
||||
chat_output = ChatOutput(_id="chatoutput-123")
|
||||
|
|
@ -98,7 +98,7 @@ def rag_graph():
|
|||
return graph
|
||||
|
||||
|
||||
def test_vector_store_rag(ingestion_graph: Graph, rag_graph: Graph):
|
||||
def test_vector_store_rag(ingestion_graph, rag_graph):
|
||||
assert ingestion_graph is not None
|
||||
ingestion_ids = [
|
||||
"file-123",
|
||||
|
|
@ -117,17 +117,11 @@ def test_vector_store_rag(ingestion_graph: Graph, rag_graph: Graph):
|
|||
"openai-embeddings-124",
|
||||
]
|
||||
for ids, graph, len_results in zip([ingestion_ids, rag_ids], [ingestion_graph, rag_graph], [5, 8]):
|
||||
results: list[VertexBuildResult] = []
|
||||
ids_count = Counter(ids)
|
||||
results_id_count: dict[str, int] = defaultdict(int)
|
||||
for result in graph.start(config={"output": {"cache": True}}):
|
||||
results = []
|
||||
for result in graph.start():
|
||||
results.append(result)
|
||||
if hasattr(result, "vertex"):
|
||||
results_id_count[result.vertex.id] += 1
|
||||
|
||||
assert (
|
||||
len(results) == len_results
|
||||
), f"Counts: {ids_count} != {results_id_count}, Diff: {set(ids_count.keys()) - set(results_id_count.keys())}"
|
||||
assert len(results) == len_results
|
||||
vids = [result.vertex.id for result in results if hasattr(result, "vertex")]
|
||||
assert all(vid in ids for vid in vids), f"Diff: {set(vids) - set(ids)}"
|
||||
assert results[-1] == Finish()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue