Rename to add_edge(string[], string)

This commit is contained in:
Nuno Campos
2024-03-21 17:56:35 -07:00
parent e11153cf09
commit db2ec5b2de
3 changed files with 16 additions and 13 deletions
+10 -7
View File
@@ -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,
+3 -3
View File
@@ -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")
+3 -3
View File
@@ -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")