diff --git a/permchain/langgraph/__init__.py b/permchain/langgraph/__init__.py index f968b98b9..6f97bf491 100644 --- a/permchain/langgraph/__init__.py +++ b/permchain/langgraph/__init__.py @@ -1,3 +1,4 @@ +from asyncio import iscoroutinefunction from collections import defaultdict from typing import Any, Callable, Dict, NamedTuple @@ -57,6 +58,8 @@ 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") self.branches[start_key].append(Branch(condition, conditional_edge_mapping)) diff --git a/pyproject.toml b/pyproject.toml index ce2741652..c94b6d9d9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -67,6 +67,6 @@ asyncio_mode = "auto" # # https://github.com/tophat/syrupy # --snapshot-warn-unused Prints a warning on unused snapshots rather than fail the test suite. -addopts = "-x --full-trace --strict-markers --strict-config --durations=5 --snapshot-warn-unused" +addopts = "-x -vv --full-trace --strict-markers --strict-config --durations=5 --snapshot-warn-unused" # Registering custom markers. # https://docs.pytest.org/en/7.1.x/example/markers.html#registering-markers diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 8e5979746..b8253945f 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -2,7 +2,7 @@ import operator import time from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager -from typing import Generator +from typing import Any, Generator import pytest from langchain_core.runnables import RunnablePassthrough @@ -554,3 +554,178 @@ def test_channel_enter_exit_timing(mocker: MockerFixture) -> None: else: assert False, "Expected only two chunks" assert cleanup.call_count == 1, "Expected cleanup to be called once" + + +def test_conditional_graph() -> None: + from copy import deepcopy + + from langchain.llms.fake import FakeStreamingListLLM + from langchain_community.tools import tool + from langchain_core.agents import AgentAction, AgentFinish + from langchain_core.prompts import PromptTemplate + from langchain_core.runnables import RunnablePassthrough + + from permchain.langgraph import END + + # Assemble the tools + @tool() + def search_api(query: str) -> str: + """Searches the API for the query.""" + return f"result for {query}" + + tools = [search_api] + + # Construct the agent + prompt = PromptTemplate.from_template("Hello!") + + llm = FakeStreamingListLLM( + responses=[ + "tool:search_api:query", + "tool:search_api:another", + "finish:answer", + ] + ) + + def agent_parser(input: str) -> AgentFinish | AgentAction: + if input.startswith("finish"): + _, answer = input.split(":") + return AgentFinish(return_values={"answer": answer}, log=input) + else: + _, tool_name, tool_input = input.split(":") + return AgentAction(tool=tool_name, tool_input=tool_input, log=input) + + agent = RunnablePassthrough.assign(agent_outcome=prompt | llm | agent_parser) + + # Define tool execution logic + def execute_tools(data): + agent_action: AgentAction | AgentFinish = data.pop("agent_outcome") + observation = {t.name: t for t in tools}[agent_action.tool].invoke( + agent_action.tool_input + ) + if data.get("intermediate_steps") is None: + data["intermediate_steps"] = [] + data["intermediate_steps"].append((agent_action, observation)) + return data + + # Define decision-making logic + def should_continue(data): + # Logic to decide whether to continue in the loop or exit + if isinstance(data["agent_outcome"], AgentFinish): + return "exit" + else: + return "continue" + + # Define a new graph + workflow = Graph() + + workflow.add_node("agent", agent) + workflow.add_node("tools", execute_tools) + + workflow.set_entry_point("agent") + + workflow.add_conditional_edges( + "agent", should_continue, {"continue": "tools", "exit": END} + ) + + workflow.add_edge("tools", "agent") + + app = workflow.compile() + + assert app.invoke({"input": "what is weather in sf"}) == { + "input": "what is weather in sf", + "intermediate_steps": [ + ( + AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + "result for query", + ), + ( + AgentAction( + tool="search_api", + tool_input="another", + log="tool:search_api:another", + ), + "result for another", + ), + ], + "agent_outcome": AgentFinish( + return_values={"answer": "answer"}, log="finish:answer" + ), + } + + assert [ + deepcopy(c) + for c in app.stream( + {"input": "what is weather in sf"}, output=["agent", "tools"] + ) + ] == [ + { + "tools": { + "input": "what is weather in sf", + "agent_outcome": AgentAction( + tool="search_api", tool_input="query", log="tool:search_api:query" + ), + } + }, + { + "agent": { + "input": "what is weather in sf", + "intermediate_steps": [ + ( + AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + "result for query", + ) + ], + } + }, + { + "tools": { + "input": "what is weather in sf", + "intermediate_steps": [ + ( + AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + "result for query", + ) + ], + "agent_outcome": AgentAction( + tool="search_api", + tool_input="another", + log="tool:search_api:another", + ), + } + }, + { + "agent": { + "input": "what is weather in sf", + "intermediate_steps": [ + ( + AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + "result for query", + ), + ( + AgentAction( + tool="search_api", + tool_input="another", + log="tool:search_api:another", + ), + "result for another", + ), + ], + } + }, + ] diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index abcc0e8b6..760124d5f 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -582,3 +582,178 @@ async def test_channel_enter_exit_timing(mocker: MockerFixture) -> None: assert cleanup_sync.call_count == 0 assert setup_async.call_count == 1, "Expected setup to be called once" assert cleanup_async.call_count == 1, "Expected cleanup to be called once" + + +async def test_conditional_graph() -> None: + from copy import deepcopy + + from langchain.llms.fake import FakeStreamingListLLM + from langchain_community.tools import tool + from langchain_core.agents import AgentAction, AgentFinish + from langchain_core.prompts import PromptTemplate + from langchain_core.runnables import RunnablePassthrough + + from permchain.langgraph import END + + # Assemble the tools + @tool() + def search_api(query: str) -> str: + """Searches the API for the query.""" + return f"result for {query}" + + tools = [search_api] + + # Construct the agent + prompt = PromptTemplate.from_template("Hello!") + + llm = FakeStreamingListLLM( + responses=[ + "tool:search_api:query", + "tool:search_api:another", + "finish:answer", + ] + ) + + async def agent_parser(input: str) -> AgentFinish | AgentAction: + if input.startswith("finish"): + _, answer = input.split(":") + return AgentFinish(return_values={"answer": answer}, log=input) + else: + _, tool_name, tool_input = input.split(":") + return AgentAction(tool=tool_name, tool_input=tool_input, log=input) + + agent = RunnablePassthrough.assign(agent_outcome=prompt | llm | agent_parser) + + # Define tool execution logic + async def execute_tools(data): + agent_action: AgentAction | AgentFinish = data.pop("agent_outcome") + observation = await {t.name: t for t in tools}[agent_action.tool].ainvoke( + agent_action.tool_input + ) + if data.get("intermediate_steps") is None: + data["intermediate_steps"] = [] + data["intermediate_steps"].append((agent_action, observation)) + return data + + # Define decision-making logic + def should_continue(data): + # Logic to decide whether to continue in the loop or exit + if isinstance(data["agent_outcome"], AgentFinish): + return "exit" + else: + return "continue" + + # Define a new graph + workflow = Graph() + + workflow.add_node("agent", agent) + workflow.add_node("tools", execute_tools) + + workflow.set_entry_point("agent") + + workflow.add_conditional_edges( + "agent", should_continue, {"continue": "tools", "exit": END} + ) + + workflow.add_edge("tools", "agent") + + app = workflow.compile() + + assert await app.ainvoke({"input": "what is weather in sf"}) == { + "input": "what is weather in sf", + "intermediate_steps": [ + ( + AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + "result for query", + ), + ( + AgentAction( + tool="search_api", + tool_input="another", + log="tool:search_api:another", + ), + "result for another", + ), + ], + "agent_outcome": AgentFinish( + return_values={"answer": "answer"}, log="finish:answer" + ), + } + + assert [ + deepcopy(c) + async for c in app.astream( + {"input": "what is weather in sf"}, output=["agent", "tools"] + ) + ] == [ + { + "tools": { + "input": "what is weather in sf", + "agent_outcome": AgentAction( + tool="search_api", tool_input="query", log="tool:search_api:query" + ), + } + }, + { + "agent": { + "input": "what is weather in sf", + "intermediate_steps": [ + ( + AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + "result for query", + ) + ], + } + }, + { + "tools": { + "input": "what is weather in sf", + "intermediate_steps": [ + ( + AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + "result for query", + ) + ], + "agent_outcome": AgentAction( + tool="search_api", + tool_input="another", + log="tool:search_api:another", + ), + } + }, + { + "agent": { + "input": "what is weather in sf", + "intermediate_steps": [ + ( + AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + "result for query", + ), + ( + AgentAction( + tool="search_api", + tool_input="another", + log="tool:search_api:another", + ), + "result for another", + ), + ], + } + }, + ]