Merge branch 'main' into tool_call

This commit is contained in:
midas8181919
2024-02-10 23:25:32 +00:00
2 changed files with 42 additions and 8 deletions
+34 -4
View File
@@ -1,3 +1,4 @@
import logging
from asyncio import iscoroutinefunction
from collections import defaultdict
from typing import Any, Callable, Dict, NamedTuple, Optional, Sequence
@@ -12,6 +13,8 @@ from langchain_core.runnables.base import (
from langgraph.checkpoint import BaseCheckpointSaver
from langgraph.pregel import Channel, Pregel
logger = logging.getLogger(__name__)
END = "__end__"
@@ -34,8 +37,14 @@ class Graph:
self.edges = set[tuple[str, str]]()
self.branches: defaultdict[str, list[Branch]] = defaultdict(list)
self.support_multiple_edges = False
self.compiled = False
def add_node(self, key: str, action: RunnableLike) -> None:
if self.compiled:
logger.warning(
"Adding a node to a graph that has already been compiled. This will "
"not be reflected in the compiled graph."
)
if key in self.nodes:
raise ValueError(f"Node `{key}` already present.")
if key == END:
@@ -44,6 +53,11 @@ class Graph:
self.nodes[key] = coerce_to_runnable(action)
def add_edge(self, start_key: str, end_key: str) -> None:
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."
)
if start_key == END:
raise ValueError("END cannot be a start node")
if start_key not in self.nodes:
@@ -64,6 +78,11 @@ class Graph:
condition: Callable[..., str],
conditional_edge_mapping: Optional[Dict[str, str]] = None,
) -> None:
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."
)
if start_key not in self.nodes:
raise ValueError(f"Need to add_node `{start_key}` first")
if iscoroutinefunction(condition):
@@ -81,6 +100,11 @@ class Graph:
self.branches[start_key].append(Branch(condition, conditional_edge_mapping))
def set_entry_point(self, key: str) -> None:
if self.compiled:
logger.warning(
"Setting the entry point of a graph that has already been compiled. "
"This will not be reflected in the compiled graph."
)
if key not in self.nodes:
raise ValueError(f"Need to add_node `{key}` first")
self.entry_point = key
@@ -119,12 +143,17 @@ class Graph:
if node not in self.nodes:
raise ValueError(f"Node `{node}` is not present")
self.compiled = True
def compile(
self,
checkpointer: Optional[BaseCheckpointSaver] = None,
interrupt_before: Optional[Sequence[str]] = None,
interrupt_after: Optional[Sequence[str]] = None,
) -> Pregel:
self.validate(interrupt=interrupt_before)
interrupt_before = interrupt_before or []
interrupt_after = interrupt_after or []
self.validate(interrupt=interrupt_before + interrupt_after)
outgoing_edges = defaultdict(list)
for start, end in self.edges:
@@ -154,7 +183,8 @@ class Graph:
output=END,
hidden=[f"{node}:inbox" for node in self.nodes],
checkpointer=checkpointer,
interrupt=[f"{node}:inbox" for node in interrupt_before]
if interrupt_before
else [],
interrupt=(
[f"{node}:inbox" for node in interrupt_before]
+ [node for node in interrupt_after]
),
)
+8 -4
View File
@@ -38,8 +38,11 @@ class StateGraph(Graph):
self,
checkpointer: Optional[BaseCheckpointSaver] = None,
interrupt_before: Optional[Sequence[str]] = None,
interrupt_after: Optional[Sequence[str]] = None,
) -> Pregel:
self.validate(interrupt=interrupt_before)
interrupt_before = interrupt_before or []
interrupt_after = interrupt_after or []
self.validate(interrupt=interrupt_before + interrupt_after)
state_keys = list(self.channels)
state_keys_read = state_keys[0] if state_keys == ["__root__"] else state_keys
@@ -102,9 +105,10 @@ class StateGraph(Graph):
output=END,
hidden=[f"{node}:inbox" for node in self.nodes] + [START] + state_keys,
checkpointer=checkpointer,
interrupt=[f"{node}:inbox" for node in interrupt_before]
if interrupt_before
else [],
interrupt=(
[f"{node}:inbox" for node in interrupt_before]
+ [node for node in interrupt_after]
),
)