diff --git a/libs/langgraph/langgraph/graph/graph.py b/libs/langgraph/langgraph/graph/graph.py index 212987e67..63405912f 100644 --- a/libs/langgraph/langgraph/graph/graph.py +++ b/libs/langgraph/langgraph/graph/graph.py @@ -203,13 +203,18 @@ class Graph: "not be reflected in the compiled graph." ) # coerce path_map to a dictionary - if isinstance(path_map, dict): - path_map = path_map.copy() - elif isinstance(path_map, list): - path_map = {name: name for name in path_map} - elif rtn_type := get_type_hints(path).get("return"): - if get_origin(rtn_type) is Literal: - path_map = {name: name for name in get_args(rtn_type)} + try: + if isinstance(path_map, dict): + path_map = path_map.copy() + elif isinstance(path_map, list): + path_map = {name: name for name in path_map} + elif rtn_type := get_type_hints(path.__call__).get( + "return" + ) or get_type_hints(path).get("return"): + if get_origin(rtn_type) is Literal: + path_map = {name: name for name in get_args(rtn_type)} + except Exception: + pass # find a name for the condition path = coerce_to_runnable(path, name=None, trace=True) name = path.name or "condition" diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 6c4ffb69e..a04bb2b54 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -7023,6 +7023,57 @@ def test_in_one_fan_out_state_graph_waiting_edge_multiple() -> None: ] +def test_callable_in_conditional_edges_with_no_path_map() -> None: + class State(TypedDict, total=False): + query: str + + def rewrite(data: State) -> State: + return {"query": f'query: {data["query"]}'} + + def analyze(data: State) -> State: + return {"query": f'analyzed: {data["query"]}'} + + class ChooseAnalyzer: + def __call__(self, data: State) -> str: + return "analyzer" + + workflow = StateGraph(State) + workflow.add_node("rewriter", rewrite) + workflow.add_node("analyzer", analyze) + workflow.add_conditional_edges("rewriter", ChooseAnalyzer()) + workflow.set_entry_point("rewriter") + app = workflow.compile() + + assert app.invoke({"query": "what is weather in sf"}) == { + "query": "analyzed: query: what is weather in sf", + } + + +def test_function_in_conditional_edges_with_no_path_map() -> None: + class State(TypedDict, total=False): + query: str + + def rewrite(data: State) -> State: + return {"query": f'query: {data["query"]}'} + + def analyze(data: State) -> State: + return {"query": f'analyzed: {data["query"]}'} + + def choose_analyzer(data: State) -> str: + return "analyzer" + + workflow = StateGraph(State) + workflow.add_node("rewriter", rewrite) + workflow.add_node("analyzer", analyze) + workflow.add_conditional_edges("rewriter", choose_analyzer) + workflow.set_entry_point("rewriter") + app = workflow.compile() + + assert app.invoke({"query": "what is weather in sf"}) == { + "query": "analyzed: query: what is weather in sf", + } + + def test_in_one_fan_out_state_graph_waiting_edge_multiple_cond_edge() -> None: def sorted_add( x: list[str], y: Union[list[str], list[tuple[str, str]]]