Merge pull request #624 from langchain-ai/vb/conditional-edges-typing

allow path_map keys to be non-strings in add_conditional_edges
This commit is contained in:
Nuno Campos
2024-06-07 16:09:38 -07:00
committed by GitHub
3 changed files with 115 additions and 10 deletions
+12 -10
View File
@@ -4,7 +4,7 @@ from typing import (
Any,
Awaitable,
Callable,
Dict,
Hashable,
Literal,
NamedTuple,
Optional,
@@ -38,8 +38,8 @@ logger = logging.getLogger(__name__)
class Branch(NamedTuple):
path: Runnable[Any, Union[str, list[str]]]
ends: Optional[dict[str, str]]
path: Runnable[Any, Union[Hashable, list[Hashable]]]
ends: Optional[dict[Hashable, str]]
then: Optional[str] = None
def run(
@@ -174,11 +174,11 @@ class Graph:
self,
source: str,
path: Union[
Callable[..., Union[str, list[str]]],
Callable[..., Awaitable[Union[str, list[str]]]],
Runnable[Any, Union[str, list[str]]],
Callable[..., Union[Hashable, list[Hashable]]],
Callable[..., Awaitable[Union[Hashable, list[Hashable]]]],
Runnable[Any, Union[Hashable, list[Hashable]]],
],
path_map: Optional[Union[dict[str, str], list[str]]] = None,
path_map: Optional[Union[dict[Hashable, str], list[str]]] = None,
then: Optional[str] = None,
) -> None:
"""Add a conditional edge from the starting node to any number of destination nodes.
@@ -189,7 +189,7 @@ class Graph:
path (Union[Callable, Runnable]): The callable that determines the next
node or nodes. If not specifying `path_map` it should return one or
more nodes. If it returns END, the graph will stop execution.
path_map (Optional[dict[str, str]]): Optional mapping of paths to node
path_map (Optional[dict[Hashable, str]]): Optional mapping of paths to node
names. If omitted the paths returned by `path` should be node names.
then (Optional[str]): The name of a node to execute after the nodes
selected by `path`.
@@ -235,9 +235,11 @@ class Graph:
def set_conditional_entry_point(
self,
path: Union[
Callable[..., str], Callable[..., Awaitable[str]], Runnable[Any, str]
Callable[..., Union[Hashable, list[Hashable]]],
Callable[..., Awaitable[Union[Hashable, list[Hashable]]]],
Runnable[Any, Union[Hashable, list[Hashable]]],
],
path_map: Optional[Dict[str, str]] = None,
path_map: Optional[Union[dict[Hashable, str], list[str]]] = None,
then: Optional[str] = None,
) -> None:
"""Sets a conditional entry point in the graph.
+61
View File
@@ -114,6 +114,67 @@
'''
# ---
# name: test_conditional_entrypoint_to_multiple_state_graph
'{"title": "LangGraphInput", "$ref": "#/definitions/OverallState", "definitions": {"OverallState": {"title": "OverallState", "type": "object", "properties": {"locations": {"title": "Locations", "type": "array", "items": {"type": "string"}}, "results": {"title": "Results", "type": "array", "items": {"type": "string"}}}, "required": ["locations", "results"]}}}'
# ---
# name: test_conditional_entrypoint_to_multiple_state_graph.1
'{"title": "LangGraphOutput", "type": "object", "properties": {"locations": {"title": "Locations", "type": "array", "items": {"type": "string"}}, "results": {"title": "Results", "type": "array", "items": {"type": "string"}}}}'
# ---
# name: test_conditional_entrypoint_to_multiple_state_graph.2
'''
{
"nodes": [
{
"id": "__start__",
"type": "schema",
"data": "__start__"
},
{
"id": "__end__",
"type": "schema",
"data": "__end__"
},
{
"id": "get_weather",
"type": "runnable",
"data": {
"id": [
"langgraph",
"utils",
"RunnableCallable"
],
"name": "get_weather"
}
}
],
"edges": [
{
"source": "get_weather",
"target": "__end__"
},
{
"source": "__start__",
"target": "get_weather",
"conditional": true
},
{
"source": "__start__",
"target": "__end__",
"conditional": true
}
]
}
'''
# ---
# name: test_conditional_entrypoint_to_multiple_state_graph.3
'''
graph TD;
get_weather --> __end__;
__start__ -.-> get_weather;
__start__ -.-> __end__;
'''
# ---
# name: test_conditional_entrypoint_graph_state
'{"title": "LangGraphInput", "$ref": "#/definitions/AgentState", "definitions": {"AgentState": {"title": "AgentState", "type": "object", "properties": {"input": {"title": "Input", "type": "string"}, "output": {"title": "Output", "type": "string"}, "steps": {"title": "Steps", "type": "array", "items": {"type": "string"}}}}}}'
# ---
+42
View File
@@ -1934,6 +1934,48 @@ def test_conditional_entrypoint_graph(snapshot: SnapshotAssertion) -> None:
]
def test_conditional_entrypoint_to_multiple_state_graph(
snapshot: SnapshotAssertion,
) -> None:
class OverallState(TypedDict):
locations: list[str]
results: Annotated[list[str], operator.add]
def get_weather(state: OverallState) -> OverallState:
location = state["location"]
weather = "sunny" if len(location) > 2 else "cloudy"
return {"results": [f"It's {weather} in {location}"]}
def continue_to_weather(state: OverallState) -> list[Send]:
return [
Send("get_weather", {"location": location})
for location in state["locations"]
]
workflow = StateGraph(OverallState)
workflow.add_node("get_weather", get_weather)
workflow.add_edge("get_weather", END)
workflow.set_conditional_entry_point(continue_to_weather)
app = workflow.compile()
assert app.get_input_schema().schema_json() == snapshot
assert app.get_output_schema().schema_json() == snapshot
assert json.dumps(app.get_graph().to_json(), indent=2) == snapshot
assert app.get_graph().draw_mermaid(with_styles=False) == snapshot
assert app.invoke({"locations": ["sf", "nyc"]}, debug=True) == {
"locations": ["sf", "nyc"],
"results": ["It's cloudy in sf", "It's sunny in nyc"],
}
assert [*app.stream({"locations": ["sf", "nyc"]}, stream_mode="values")][-1] == {
"locations": ["sf", "nyc"],
"results": ["It's cloudy in sf", "It's sunny in nyc"],
}
def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
from langchain_core.agents import AgentAction, AgentFinish
from langchain_core.language_models.fake import FakeStreamingListLLM