Merge pull request #236 from langchain-ai/jacob/async_edges

Adds async conditional edge support
This commit is contained in:
Nuno Campos
2024-03-28 15:42:34 -07:00
committed by GitHub
5 changed files with 155 additions and 154 deletions
+33 -28
View File
@@ -1,7 +1,6 @@
import logging
from asyncio import iscoroutinefunction
from collections import defaultdict
from typing import Any, Callable, Dict, NamedTuple, Optional, Sequence
from typing import Any, Awaitable, Callable, Dict, NamedTuple, Optional, Sequence, Union
from langchain_core.runnables import Runnable
from langchain_core.runnables.base import (
@@ -28,11 +27,23 @@ END = "__end__"
class Branch(NamedTuple):
condition: Callable[..., str]
condition: Runnable[Any, str]
ends: Optional[dict[str, str]]
def runnable(self, input: Any) -> Runnable:
result = self.condition(input)
@property
def runnable(self):
return RunnableLambda(self._route, self._aroute, name=self.condition.name)
def _route(self, input: Any) -> Runnable:
result = self.condition.invoke(input, {"run_name": "condition"})
if self.ends:
destination = self.ends[result]
else:
destination = result
return Channel.write_to(f"{destination}:inbox" if destination != END else END)
async def _aroute(self, input: Any) -> Runnable:
result = await self.condition.ainvoke(input, {"run_name": "condition"})
if self.ends:
destination = self.ends[result]
else:
@@ -90,7 +101,9 @@ class Graph:
def add_conditional_edges(
self,
start_key: str,
condition: Callable[..., str],
condition: Union[
Callable[..., str], Callable[..., Awaitable[str]], Runnable[Any, str]
],
conditional_edge_mapping: Optional[Dict[str, str]] = None,
) -> None:
if self.compiled:
@@ -100,8 +113,6 @@ class Graph:
)
if start_key not in self.nodes:
raise ValueError(f"Need to add_node `{start_key}` first")
if iscoroutinefunction(condition):
raise ValueError("Condition cannot be a coroutine function")
if conditional_edge_mapping and set(
conditional_edge_mapping.values()
).difference([END]).difference(self.nodes):
@@ -111,6 +122,8 @@ class Graph:
f"{list(conditional_edge_mapping.values())}. Possible nodes are "
f"{list(self.nodes.keys())}."
)
if not isinstance(condition, Runnable):
condition = RunnableLambda(condition)
self.branches[start_key].append(Branch(condition, conditional_edge_mapping))
@@ -126,7 +139,9 @@ class Graph:
def set_conditional_entry_point(
self,
condition: Callable[..., str],
condition: Union[
Callable[..., str], Callable[..., Awaitable[str]], Runnable[Any, str]
],
conditional_edge_mapping: Optional[Dict[str, str]] = None,
) -> None:
if self.compiled:
@@ -134,8 +149,6 @@ class Graph:
"Setting the entry point of a graph that has already been compiled. "
"This will not be reflected in the compiled graph."
)
if iscoroutinefunction(condition):
raise ValueError("Condition cannot be a coroutine function")
if conditional_edge_mapping and set(
conditional_edge_mapping.values()
).difference([END]).difference(self.nodes):
@@ -145,6 +158,8 @@ class Graph:
f"{list(conditional_edge_mapping.values())}. Possible nodes are "
f"{list(self.nodes.keys())}."
)
if not isinstance(condition, Runnable):
condition = RunnableLambda(condition)
self.entry_point_branch = Branch(condition, conditional_edge_mapping)
def set_finish_point(self, key: str) -> None:
@@ -218,15 +233,12 @@ class Graph:
nodes[edges_key] |= Channel.write_to(*[dest for dest in outgoing])
if key in self.branches:
for branch in self.branches[key]:
nodes[edges_key] |= RunnableLambda(
branch.runnable, name=f"{key}_condition"
)
nodes[edges_key] |= branch.runnable
if self.entry_point_branch:
nodes[f"{START}:edges"] = Channel.subscribe_to(
START, tags=["langsmith:hidden"]
) | RunnableLambda(
self.entry_point_branch.runnable, name=f"{START}_condition"
nodes[f"{START}:edges"] = (
Channel.subscribe_to(START, tags=["langsmith:hidden"])
| self.entry_point_branch.runnable
)
elif self.entry_point is None:
raise ValueError("No entry point set")
@@ -285,13 +297,10 @@ class CompiledGraph(Pregel):
graph.add_edge(start_nodes[start], end_nodes[end])
for start, branches in self.graph.branches.items():
for i, branch in enumerate(branches):
name = f"{start}_{branch.condition.__name__}"
name = f"{start}_{branch.runnable.name}"
if i > 0:
name += f"_{i}"
cond = graph.add_node(
RunnableLambda(branch.runnable, name=branch.condition.__name__),
name,
)
cond = graph.add_node(branch.runnable, name)
graph.add_edge(start_nodes[start], cond)
ends = branch.ends or {
**{k: k for k in self.graph.nodes},
@@ -301,11 +310,7 @@ class CompiledGraph(Pregel):
graph.add_edge(cond, end_nodes[end], label)
if self.graph.entry_point_branch:
cond = graph.add_node(
RunnableLambda(
self.graph.entry_point_branch.runnable,
name=self.graph.entry_point_branch.condition.__name__,
),
f"{START}_condition",
self.graph.entry_point_branch.runnable, f"{START}_condition"
)
graph.add_edge(start_nodes[START], cond)
ends = self.graph.entry_point_branch.ends or {
+2 -6
View File
@@ -159,9 +159,7 @@ class StateGraph(Graph):
)
if key in self.branches:
for branch in self.branches[key]:
nodes[edges_key] |= RunnableLambda(
branch.runnable, name=f"{key}_condition"
)
nodes[edges_key] |= branch.runnable
nodes[START] = Channel.subscribe_to(
f"{START}:inbox", tags=["langsmith:hidden"]
@@ -174,9 +172,7 @@ class StateGraph(Graph):
if self.entry_point:
nodes[f"{START}:edges"] |= Channel.write_to(f"{self.entry_point}:inbox")
elif self.entry_point_branch:
nodes[f"{START}:edges"] |= RunnableLambda(
self.entry_point_branch.runnable, name=f"{START}_condition"
)
nodes[f"{START}:edges"] |= self.entry_point_branch.runnable
else:
raise ValueError("No entry point set")
+78 -78
View File
@@ -50,7 +50,7 @@
}
},
{
"id": "left_<lambda>",
"id": "left__route",
"type": "runnable",
"data": {
"id": [
@@ -59,7 +59,7 @@
"base",
"RunnableLambda"
],
"name": "<lambda>"
"name": "_route"
}
},
{
@@ -83,20 +83,20 @@
},
{
"source": "left",
"target": "left_<lambda>"
"target": "left__route"
},
{
"source": "left_<lambda>",
"source": "left__route",
"target": "left",
"data": "left"
},
{
"source": "left_<lambda>",
"source": "left__route",
"target": "right",
"data": "right"
},
{
"source": "left_<lambda>",
"source": "left__route",
"target": "__end__",
"data": "__end__"
},
@@ -120,39 +120,39 @@
# ---
# name: test_conditional_entrypoint_graph.3
'''
+-----------+
| __start__ |
+-----------+
*
*
*
+---------------------+
| __start___condition |
+---------------------+
*** ***
* *
** ***
+------+ *
| left | *
+------+ *
* *
* *
* *
+---------------+ *
| left_<lambda> | *
+---------------+* *
* ***** *
* *** *
* *** *
** +-------+
* | right |
*** +-------+
* ***
*** *
* **
+---------+
| __end__ |
+---------+
+-----------+
| __start__ |
+-----------+
*
*
*
+---------------------+
| __start___condition |
+---------------------+
*** ***
* *
** ***
+------+ *
| left | *
+------+ *
* *
* *
* *
+-------------+ *
| left__route | *
+-------------+** *
* **** *
* **** *
* ** *
* +-------+
** | right |
** +-------+
** ***
** *
* **
+---------+
| __end__ |
+---------+
'''
# ---
# name: test_conditional_entrypoint_graph_state
@@ -240,7 +240,7 @@
}
},
{
"id": "left_<lambda>",
"id": "left__route",
"type": "runnable",
"data": {
"id": [
@@ -249,7 +249,7 @@
"base",
"RunnableLambda"
],
"name": "<lambda>"
"name": "_route"
}
},
{
@@ -273,20 +273,20 @@
},
{
"source": "left",
"target": "left_<lambda>"
"target": "left__route"
},
{
"source": "left_<lambda>",
"source": "left__route",
"target": "left",
"data": "left"
},
{
"source": "left_<lambda>",
"source": "left__route",
"target": "right",
"data": "right"
},
{
"source": "left_<lambda>",
"source": "left__route",
"target": "__end__",
"data": "__end__"
},
@@ -310,39 +310,39 @@
# ---
# name: test_conditional_entrypoint_graph_state.3
'''
+-----------+
| __start__ |
+-----------+
*
*
*
+---------------------+
| __start___condition |
+---------------------+
*** ***
* *
** ***
+------+ *
| left | *
+------+ *
* *
* *
* *
+---------------+ *
| left_<lambda> | *
+---------------+* *
* ***** *
* *** *
* *** *
** +-------+
* | right |
*** +-------+
* ***
*** *
* **
+---------+
| __end__ |
+---------+
+-----------+
| __start__ |
+-----------+
*
*
*
+---------------------+
| __start___condition |
+---------------------+
*** ***
* *
** ***
+------+ *
| left | *
+------+ *
* *
* *
* *
+-------------+ *
| left__route | *
+-------------+** *
* **** *
* **** *
* ** *
* +-------+
** | right |
** +-------+
** ***
** *
* **
+---------+
| __end__ |
+---------+
'''
# ---
# name: test_conditional_graph
+25 -25
View File
@@ -2480,7 +2480,7 @@ def test_message_graph(
FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000014",
id="00000000-0000-4000-8000-000000000015",
),
AIMessage(
content="",
@@ -2492,7 +2492,7 @@ def test_message_graph(
FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000026",
id="00000000-0000-4000-8000-000000000028",
),
AIMessage(content="answer", id="ai3"),
]
@@ -2511,7 +2511,7 @@ def test_message_graph(
"action": FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000047",
id="00000000-0000-4000-8000-000000000051",
)
},
{
@@ -2527,7 +2527,7 @@ def test_message_graph(
"action": FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000059",
id="00000000-0000-4000-8000-000000000064",
)
},
{"agent": AIMessage(content="answer", id="ai3")},
@@ -2535,7 +2535,7 @@ def test_message_graph(
"__end__": [
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000035",
id="00000000-0000-4000-8000-000000000038",
),
AIMessage(
content="",
@@ -2547,7 +2547,7 @@ def test_message_graph(
FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000047",
id="00000000-0000-4000-8000-000000000051",
),
AIMessage(
content="",
@@ -2562,7 +2562,7 @@ def test_message_graph(
FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000059",
id="00000000-0000-4000-8000-000000000064",
),
AIMessage(content="answer", id="ai3"),
]
@@ -2595,7 +2595,7 @@ def test_message_graph(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2619,7 +2619,7 @@ def test_message_graph(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2641,7 +2641,7 @@ def test_message_graph(
"action": FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
id="00000000-0000-4000-8000-000000000088",
)
},
{
@@ -2659,7 +2659,7 @@ def test_message_graph(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2674,7 +2674,7 @@ def test_message_graph(
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
id="00000000-0000-4000-8000-000000000088",
),
AIMessage(
content="",
@@ -2698,7 +2698,7 @@ def test_message_graph(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2713,7 +2713,7 @@ def test_message_graph(
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
id="00000000-0000-4000-8000-000000000088",
),
AIMessage(content="answer", id="ai2"),
],
@@ -2726,7 +2726,7 @@ def test_message_graph(
"__end__": [
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2741,7 +2741,7 @@ def test_message_graph(
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
id="00000000-0000-4000-8000-000000000088",
),
AIMessage(content="answer", id="ai2"),
]
@@ -2775,7 +2775,7 @@ def test_message_graph(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000091",
id="00000000-0000-4000-8000-000000000099",
),
AIMessage(
content="",
@@ -2799,7 +2799,7 @@ def test_message_graph(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000091",
id="00000000-0000-4000-8000-000000000099",
),
AIMessage(
content="",
@@ -2821,7 +2821,7 @@ def test_message_graph(
"action": FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000106",
id="00000000-0000-4000-8000-000000000116",
)
},
{
@@ -2839,7 +2839,7 @@ def test_message_graph(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000091",
id="00000000-0000-4000-8000-000000000099",
),
AIMessage(
content="",
@@ -2854,7 +2854,7 @@ def test_message_graph(
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000106",
id="00000000-0000-4000-8000-000000000116",
),
AIMessage(
content="",
@@ -2878,7 +2878,7 @@ def test_message_graph(
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000091",
id="00000000-0000-4000-8000-000000000099",
),
AIMessage(
content="",
@@ -2893,7 +2893,7 @@ def test_message_graph(
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000106",
id="00000000-0000-4000-8000-000000000116",
),
AIMessage(content="answer", id="ai2"),
],
@@ -2906,7 +2906,7 @@ def test_message_graph(
"__end__": [
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000091",
id="00000000-0000-4000-8000-000000000099",
),
AIMessage(
content="",
@@ -2921,7 +2921,7 @@ def test_message_graph(
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000106",
id="00000000-0000-4000-8000-000000000116",
),
AIMessage(content="answer", id="ai2"),
]
+17 -17
View File
@@ -820,7 +820,7 @@ async def test_conditional_graph() -> None:
return data
# Define decision-making logic
def should_continue(data: dict) -> str:
async def should_continue(data: dict) -> str:
# Logic to decide whether to continue in the loop or exit
if isinstance(data["agent_outcome"], AgentFinish):
return "exit"
@@ -2483,7 +2483,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000014",
id="00000000-0000-4000-8000-000000000015",
),
AIMessage(
content="",
@@ -2495,7 +2495,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000026",
id="00000000-0000-4000-8000-000000000028",
),
AIMessage(content="answer", id="ai3"),
]
@@ -2516,7 +2516,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
"action": FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000047",
id="00000000-0000-4000-8000-000000000051",
)
},
{
@@ -2532,7 +2532,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
"action": FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000059",
id="00000000-0000-4000-8000-000000000064",
)
},
{"agent": AIMessage(content="answer", id="ai3")},
@@ -2540,7 +2540,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
"__end__": [
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000035",
id="00000000-0000-4000-8000-000000000038",
),
AIMessage(
content="",
@@ -2552,7 +2552,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
FunctionMessage(
content="result for query",
name="search_api",
id="00000000-0000-4000-8000-000000000047",
id="00000000-0000-4000-8000-000000000051",
),
AIMessage(
content="",
@@ -2567,7 +2567,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
FunctionMessage(
content="result for another",
name="search_api",
id="00000000-0000-4000-8000-000000000059",
id="00000000-0000-4000-8000-000000000064",
),
AIMessage(content="answer", id="ai3"),
]
@@ -2600,7 +2600,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2624,7 +2624,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2646,7 +2646,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
"action": FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
id="00000000-0000-4000-8000-000000000088",
)
},
{
@@ -2664,7 +2664,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2679,7 +2679,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
id="00000000-0000-4000-8000-000000000088",
),
AIMessage(
content="",
@@ -2703,7 +2703,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
values=[
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2718,7 +2718,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
id="00000000-0000-4000-8000-000000000088",
),
AIMessage(content="answer", id="ai2"),
],
@@ -2731,7 +2731,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
"__end__": [
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000068",
id="00000000-0000-4000-8000-000000000074",
),
AIMessage(
content="",
@@ -2746,7 +2746,7 @@ async def test_message_graph(deterministic_uuids: MockerFixture) -> None:
FunctionMessage(
content="result for a different query",
name="search_api",
id="00000000-0000-4000-8000-000000000081",
id="00000000-0000-4000-8000-000000000088",
),
AIMessage(content="answer", id="ai2"),
]