From db2ec5b2de78bddaa32a623f9badb9fae4cf10b1 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Thu, 21 Mar 2024 17:56:35 -0700 Subject: [PATCH] Rename to `add_edge(string[], string)` --- langgraph/graph/state.py | 17 ++++++++++------- tests/test_pregel.py | 6 +++--- tests/test_pregel_async.py | 6 +++--- 3 files changed, 16 insertions(+), 13 deletions(-) diff --git a/langgraph/graph/state.py b/langgraph/graph/state.py index 2c766c410..3bb0adf5c 100644 --- a/langgraph/graph/state.py +++ b/langgraph/graph/state.py @@ -2,7 +2,7 @@ import logging from collections import defaultdict from functools import partial from inspect import signature -from typing import Any, Optional, Sequence, Type +from typing import Any, Optional, Sequence, Type, Union from langchain_core.runnables import RunnableLambda from langchain_core.runnables.base import RunnableLike @@ -45,23 +45,26 @@ class StateGraph(Graph): ) return super().add_node(key, action) - def add_waiting_edge(self, starts: Sequence[str], end: str) -> None: + def add_edge(self, start_key: Union[str, list[str]], end_key: str) -> None: + if isinstance(start_key, str): + return super().add_edge(start_key, end_key) + if self.compiled: logger.warning( "Adding an edge to a graph that has already been compiled. This will " "not be reflected in the compiled graph." ) - for start in starts: + for start in start_key: if start == END: raise ValueError("END cannot be a start node") if start not in self.nodes: raise ValueError(f"Need to add_node `{start}` first") - if end == END: + if end_key == END: raise ValueError("END cannot be an end node") - if end not in self.nodes: - raise ValueError(f"Need to add_node `{end}` first") + if end_key not in self.nodes: + raise ValueError(f"Need to add_node `{end_key}` first") - self.w_edges.add((tuple(starts), end)) + self.w_edges.add((tuple(start_key), end_key)) def compile( self, diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 0cd238c66..03e46379e 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -2827,7 +2827,7 @@ def test_in_one_fan_out_state_graph_waiting_edge() -> None: workflow.add_edge("rewrite_query", "analyzer_one") workflow.add_edge("analyzer_one", "retriever_one") workflow.add_edge("rewrite_query", "retriever_two") - workflow.add_waiting_edge(["retriever_one", "retriever_two"], "qa") + workflow.add_edge(["retriever_one", "retriever_two"], "qa") workflow.set_finish_point("qa") app = workflow.compile() @@ -2961,7 +2961,7 @@ def test_in_one_fan_out_state_graph_waiting_edge_plus_regular() -> None: workflow.add_edge("rewrite_query", "analyzer_one") workflow.add_edge("analyzer_one", "retriever_one") workflow.add_edge("rewrite_query", "retriever_two") - workflow.add_waiting_edge(["retriever_one", "retriever_two"], "qa") + workflow.add_edge(["retriever_one", "retriever_two"], "qa") workflow.set_finish_point("qa") # silly edge, to make sure having been triggered before doesn't break @@ -3089,7 +3089,7 @@ def test_in_one_fan_out_state_graph_waiting_edge_multiple() -> None: workflow.add_edge("rewrite_query", "analyzer_one") workflow.add_edge("analyzer_one", "retriever_one") workflow.add_edge("rewrite_query", "retriever_two") - workflow.add_waiting_edge(["retriever_one", "retriever_two"], "decider") + workflow.add_edge(["retriever_one", "retriever_two"], "decider") workflow.add_conditional_edges("decider", decider_cond) workflow.set_finish_point("qa") diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 4cca085e1..5ada8b3c3 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -2856,7 +2856,7 @@ async def test_in_one_fan_out_state_graph_waiting_edge() -> None: workflow.add_edge("rewrite_query", "analyzer_one") workflow.add_edge("analyzer_one", "retriever_one") workflow.add_edge("rewrite_query", "retriever_two") - workflow.add_waiting_edge(["retriever_one", "retriever_two"], "qa") + workflow.add_edge(["retriever_one", "retriever_two"], "qa") workflow.set_finish_point("qa") app = workflow.compile() @@ -2957,7 +2957,7 @@ async def test_in_one_fan_out_state_graph_waiting_edge_plus_regular() -> None: workflow.add_edge("rewrite_query", "analyzer_one") workflow.add_edge("analyzer_one", "retriever_one") workflow.add_edge("rewrite_query", "retriever_two") - workflow.add_waiting_edge(["retriever_one", "retriever_two"], "qa") + workflow.add_edge(["retriever_one", "retriever_two"], "qa") workflow.set_finish_point("qa") # silly edge, to make sure having been triggered before doesn't break @@ -3088,7 +3088,7 @@ async def test_in_one_fan_out_state_graph_waiting_edge_multiple() -> None: workflow.add_edge("rewrite_query", "analyzer_one") workflow.add_edge("analyzer_one", "retriever_one") workflow.add_edge("rewrite_query", "retriever_two") - workflow.add_waiting_edge(["retriever_one", "retriever_two"], "decider") + workflow.add_edge(["retriever_one", "retriever_two"], "decider") workflow.add_conditional_edges("decider", decider_cond) workflow.set_finish_point("qa")