mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 20:45:05 +02:00
Merge pull request #267 from langchain-ai/nc/2apr/optimize-tracing-output
Optimize tracing output of Graph/StateGraph/MessageGraph
This commit is contained in:
+10
-10
File diff suppressed because one or more lines are too long
+249
-62
File diff suppressed because one or more lines are too long
@@ -1,3 +1,6 @@
|
||||
CONFIG_KEY_SEND = "__pregel_send"
|
||||
CONFIG_KEY_READ = "__pregel_read"
|
||||
|
||||
INTERRUPT = "__interrupt__"
|
||||
|
||||
TAG_HIDDEN = "langsmith:hidden"
|
||||
|
||||
+28
-57
@@ -1,16 +1,15 @@
|
||||
import asyncio
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from typing import (
|
||||
Any,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Coroutine,
|
||||
Dict,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Sequence,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
from langchain_core.runnables import Runnable
|
||||
@@ -25,9 +24,11 @@ from langchain_core.runnables.graph import (
|
||||
|
||||
from langgraph.channels.ephemeral_value import EphemeralValue
|
||||
from langgraph.checkpoint import BaseCheckpointSaver
|
||||
from langgraph.constants import TAG_HIDDEN
|
||||
from langgraph.pregel import Channel, Pregel
|
||||
from langgraph.pregel.read import PregelNode
|
||||
from langgraph.pregel.write import ChannelWrite
|
||||
from langgraph.utils import RunnableCallable
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -35,36 +36,8 @@ START = "__start__"
|
||||
END = "__end__"
|
||||
|
||||
|
||||
class RunnableCallable(Runnable):
|
||||
def __init__(
|
||||
self,
|
||||
func: Callable[..., Optional[Runnable]],
|
||||
afunc: Callable[..., Awaitable[Optional[Runnable]]],
|
||||
name: str,
|
||||
writer: Callable[[str], Optional[Runnable]],
|
||||
) -> None:
|
||||
self.name = name
|
||||
self.func = func
|
||||
self.afunc = afunc
|
||||
self.writer = writer
|
||||
|
||||
def invoke(self, input: Any, config: Optional[RunnableConfig] = None) -> Any:
|
||||
ret = self._call_with_config(self.func, input, config, writer=self.writer)
|
||||
if isinstance(ret, Runnable):
|
||||
return ret.invoke(input, config)
|
||||
return ret
|
||||
|
||||
async def ainvoke(self, input: Any, config: Optional[RunnableConfig] = None) -> Any:
|
||||
ret = await self._acall_with_config(
|
||||
self.afunc, input, config, writer=self.writer
|
||||
)
|
||||
if isinstance(ret, Runnable):
|
||||
return await ret.ainvoke(input, config)
|
||||
return ret
|
||||
|
||||
|
||||
class Branch(NamedTuple):
|
||||
condition: Union[Runnable[Any, str], Callable[..., str], Coroutine[Any, Any, str]]
|
||||
condition: Runnable[Any, str]
|
||||
ends: Optional[dict[str, str]]
|
||||
|
||||
def run(self, writer: Callable[[str], Optional[Runnable]]) -> None:
|
||||
@@ -73,19 +46,19 @@ class Branch(NamedTuple):
|
||||
func=self._route,
|
||||
afunc=self._aroute,
|
||||
writer=writer,
|
||||
name=self.condition.name
|
||||
if isinstance(self.condition, Runnable)
|
||||
else self.condition.__name__,
|
||||
name=None,
|
||||
trace=False,
|
||||
)
|
||||
)
|
||||
|
||||
def _route(
|
||||
self, input: Any, *, writer: Callable[[str], Optional[Runnable]]
|
||||
self,
|
||||
input: Any,
|
||||
config: RunnableConfig,
|
||||
*,
|
||||
writer: Callable[[str], Optional[Runnable]],
|
||||
) -> Runnable:
|
||||
if isinstance(self.condition, Runnable):
|
||||
result = self.condition.invoke(input, {"run_name": "condition"})
|
||||
else:
|
||||
result = self.condition(input)
|
||||
result = self.condition.invoke(input, config)
|
||||
if self.ends:
|
||||
destination = self.ends[result]
|
||||
else:
|
||||
@@ -93,14 +66,13 @@ class Branch(NamedTuple):
|
||||
return writer(destination)
|
||||
|
||||
async def _aroute(
|
||||
self, input: Any, *, writer: Callable[[str], Optional[Runnable]]
|
||||
self,
|
||||
input: Any,
|
||||
config: RunnableConfig,
|
||||
*,
|
||||
writer: Callable[[str], Optional[Runnable]],
|
||||
) -> Runnable:
|
||||
if isinstance(self.condition, Runnable):
|
||||
result = await self.condition.ainvoke(input, {"run_name": "condition"})
|
||||
elif asyncio.iscoroutinefunction(self.condition):
|
||||
result = await self.condition(input)
|
||||
else:
|
||||
result = self.condition(input)
|
||||
result = await self.condition.ainvoke(input, config)
|
||||
if self.ends:
|
||||
destination = self.ends[result]
|
||||
else:
|
||||
@@ -172,12 +144,8 @@ class Graph:
|
||||
"not be reflected in the compiled graph."
|
||||
)
|
||||
# find a name for the condition
|
||||
try:
|
||||
name = (
|
||||
condition.__name__ if condition.__name__ != "<lambda>" else "condition"
|
||||
)
|
||||
except AttributeError:
|
||||
name = "condition"
|
||||
condition = coerce_to_runnable(condition)
|
||||
name = condition.name or "condition"
|
||||
# validate the condition
|
||||
if start_key not in self.nodes and start_key != START:
|
||||
raise ValueError(f"Need to add_node `{start_key}` first")
|
||||
@@ -293,14 +261,16 @@ class CompiledGraph(Pregel):
|
||||
def attach_node(self, key: str, node: Runnable) -> None:
|
||||
self.channels[key] = EphemeralValue(Any)
|
||||
self.nodes[key] = (
|
||||
PregelNode(channels=[], triggers=[]) | node | Channel.write_to(key)
|
||||
PregelNode(channels=[], triggers=[])
|
||||
| node
|
||||
| Channel.write_to(key, tags=[TAG_HIDDEN])
|
||||
)
|
||||
self.stream_channels.append(key)
|
||||
cast(list[str], self.stream_channels).append(key)
|
||||
|
||||
def attach_edge(self, start: str, end: str) -> None:
|
||||
if end == END:
|
||||
# publish to end channel
|
||||
self.nodes[start].writers.append(Channel.write_to(END))
|
||||
self.nodes[start].writers.append(Channel.write_to(END, tags=[TAG_HIDDEN]))
|
||||
else:
|
||||
# subscribe to start channel
|
||||
self.nodes[end].triggers.append(start)
|
||||
@@ -309,12 +279,13 @@ class CompiledGraph(Pregel):
|
||||
def attach_branch(self, start: str, name: str, branch: Branch) -> None:
|
||||
def branch_writer(end: str) -> Optional[ChannelWrite]:
|
||||
return Channel.write_to(
|
||||
f"branch:{start}:{name}:{end}" if end != END else END
|
||||
f"branch:{start}:{name}:{end}" if end != END else END,
|
||||
tags=[TAG_HIDDEN],
|
||||
)
|
||||
|
||||
# add hidden start node
|
||||
if start == START and start not in self.nodes:
|
||||
self.nodes[start] = Channel.subscribe_to(START, tags=["langsmith:hidden"])
|
||||
self.nodes[start] = Channel.subscribe_to(START, tags=[TAG_HIDDEN])
|
||||
|
||||
# attach branch writer
|
||||
self.nodes[start] |= branch.run(branch_writer)
|
||||
|
||||
+40
-26
@@ -3,7 +3,7 @@ from functools import partial
|
||||
from inspect import signature
|
||||
from typing import Any, Optional, Sequence, Type, Union
|
||||
|
||||
from langchain_core.runnables import Runnable, RunnableLambda
|
||||
from langchain_core.runnables import Runnable, RunnableConfig
|
||||
from langchain_core.runnables.base import RunnableLike
|
||||
|
||||
from langgraph.channels.base import BaseChannel, InvalidUpdateError
|
||||
@@ -12,10 +12,11 @@ from langgraph.channels.ephemeral_value import EphemeralValue
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.channels.named_barrier_value import NamedBarrierValue
|
||||
from langgraph.checkpoint import BaseCheckpointSaver
|
||||
from langgraph.constants import TAG_HIDDEN
|
||||
from langgraph.graph.graph import END, START, Branch, CompiledGraph, Graph
|
||||
from langgraph.pregel import Channel
|
||||
from langgraph.pregel.read import ChannelRead, PregelNode
|
||||
from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry
|
||||
from langgraph.utils import RunnableCallable
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -123,7 +124,7 @@ class CompiledStateGraph(CompiledGraph):
|
||||
graph: StateGraph
|
||||
|
||||
def attach_node(self, key: str, node: Optional[Runnable]) -> None:
|
||||
def _get_state_key(key: str, input: dict) -> Any:
|
||||
def _get_state_key(input: dict, config: RunnableConfig, *, key: str) -> Any:
|
||||
if input is None:
|
||||
return SKIP_WRITE
|
||||
elif not isinstance(input, dict):
|
||||
@@ -136,15 +137,36 @@ class CompiledStateGraph(CompiledGraph):
|
||||
state_write_entries = [
|
||||
ChannelWriteEntry(key, None, skip_none=True)
|
||||
if key == "__root__"
|
||||
else ChannelWriteEntry(key, RunnableLambda(partial(_get_state_key, key)))
|
||||
else ChannelWriteEntry(
|
||||
key, RunnableCallable(_get_state_key, key=key, trace=False)
|
||||
)
|
||||
for key in state_keys
|
||||
]
|
||||
# node that reads current state with (this node's) updates applied
|
||||
state_reader = ChannelRead(
|
||||
state_keys[0] if state_keys == ["__root__"] else state_keys,
|
||||
tags=[TAG_HIDDEN],
|
||||
fresh=True,
|
||||
# coerce state dict to schema class (eg. pydantic model)
|
||||
mapper=(
|
||||
None
|
||||
if state_keys == ["__root__"]
|
||||
else partial(_coerce_state, self.graph.schema)
|
||||
),
|
||||
)
|
||||
|
||||
# add node and output channel
|
||||
if key == START:
|
||||
self.nodes[key] = Channel.subscribe_to(
|
||||
START, tags=["langsmith:hidden"]
|
||||
).pipe(ChannelWrite(state_write_entries))
|
||||
self.nodes[key] = PregelNode(
|
||||
tags=[TAG_HIDDEN],
|
||||
triggers=[START],
|
||||
channels=[START],
|
||||
writers=[
|
||||
ChannelWrite(state_write_entries, tags=[TAG_HIDDEN]),
|
||||
# read back state with updates applied
|
||||
state_reader,
|
||||
],
|
||||
)
|
||||
else:
|
||||
self.channels[key] = EphemeralValue(Any)
|
||||
self.nodes[key] = PregelNode(
|
||||
@@ -156,24 +178,15 @@ class CompiledStateGraph(CompiledGraph):
|
||||
else {chan: chan for chan in state_keys}
|
||||
),
|
||||
# coerce state dict to schema class (eg. pydantic model)
|
||||
mapper=(
|
||||
None
|
||||
if state_keys == ["__root__"]
|
||||
else partial(_coerce_state, self.graph.schema)
|
||||
),
|
||||
# publish to this channel and state keys
|
||||
mapper=state_reader.mapper,
|
||||
writers=[
|
||||
ChannelWrite([ChannelWriteEntry(key)] + state_write_entries),
|
||||
# read back state with updates applied
|
||||
ChannelRead(
|
||||
state_keys[0] if state_keys == ["__root__"] else state_keys,
|
||||
fresh=True,
|
||||
mapper=(
|
||||
None
|
||||
if state_keys == ["__root__"]
|
||||
else partial(_coerce_state, self.graph.schema)
|
||||
),
|
||||
# publish to this channel and state keys
|
||||
ChannelWrite(
|
||||
[ChannelWriteEntry(key)] + state_write_entries,
|
||||
tags=[TAG_HIDDEN],
|
||||
),
|
||||
# read back state with updates applied
|
||||
state_reader,
|
||||
],
|
||||
).pipe(node)
|
||||
|
||||
@@ -187,7 +200,7 @@ class CompiledStateGraph(CompiledGraph):
|
||||
self.nodes[end].triggers.append(channel_name)
|
||||
# publish to channel
|
||||
self.nodes[START] |= ChannelWrite(
|
||||
[ChannelWriteEntry(channel_name, START)]
|
||||
[ChannelWriteEntry(channel_name, START)], tags=[TAG_HIDDEN]
|
||||
)
|
||||
elif end != END:
|
||||
# subscribe to start channel
|
||||
@@ -201,14 +214,15 @@ class CompiledStateGraph(CompiledGraph):
|
||||
# publish to channel
|
||||
for start in starts:
|
||||
self.nodes[start] |= ChannelWrite(
|
||||
[ChannelWriteEntry(channel_name, start)]
|
||||
[ChannelWriteEntry(channel_name, start)], tags=[TAG_HIDDEN]
|
||||
)
|
||||
|
||||
def attach_branch(self, start: str, name: str, branch: Branch) -> None:
|
||||
def branch_writer(end: str) -> Optional[ChannelWrite]:
|
||||
if end != END:
|
||||
return ChannelWrite(
|
||||
[ChannelWriteEntry(f"branch:{start}:{name}:{end}", start)]
|
||||
[ChannelWriteEntry(f"branch:{start}:{name}:{end}", start)],
|
||||
tags=[TAG_HIDDEN],
|
||||
)
|
||||
|
||||
# attach branch publisher
|
||||
|
||||
@@ -3,10 +3,10 @@ from typing import Annotated, Sequence, TypedDict, Union
|
||||
|
||||
from langchain_core.agents import AgentAction, AgentFinish
|
||||
from langchain_core.messages import BaseMessage
|
||||
from langchain_core.runnables import RunnableLambda
|
||||
|
||||
from langgraph.graph import END, StateGraph
|
||||
from langgraph.prebuilt.tool_executor import ToolExecutor
|
||||
from langgraph.utils import RunnableCallable
|
||||
|
||||
|
||||
def _get_agent_state(input_schema=None):
|
||||
@@ -72,35 +72,33 @@ def create_agent_executor(agent_runnable, tools, input_schema=None):
|
||||
def execute_tools(data):
|
||||
# Get the most recent agent_outcome - this is the key added in the `agent` above
|
||||
agent_action = data["agent_outcome"]
|
||||
if isinstance(agent_action, list):
|
||||
output = tool_executor.batch(agent_action, return_exceptions=True)
|
||||
return {
|
||||
"intermediate_steps": [
|
||||
(action, str(out)) for action, out in zip(agent_action, output)
|
||||
]
|
||||
}
|
||||
output = tool_executor.invoke(agent_action)
|
||||
return {"intermediate_steps": [(agent_action, str(output))]}
|
||||
if not isinstance(agent_action, list):
|
||||
agent_action = [agent_action]
|
||||
output = tool_executor.batch(agent_action, return_exceptions=True)
|
||||
return {
|
||||
"intermediate_steps": [
|
||||
(action, str(out)) for action, out in zip(agent_action, output)
|
||||
]
|
||||
}
|
||||
|
||||
async def aexecute_tools(data):
|
||||
# Get the most recent agent_outcome - this is the key added in the `agent` above
|
||||
agent_action = data["agent_outcome"]
|
||||
if isinstance(agent_action, list):
|
||||
output = await tool_executor.abatch(agent_action, return_exceptions=True)
|
||||
return {
|
||||
"intermediate_steps": [
|
||||
(action, str(out)) for action, out in zip(agent_action, output)
|
||||
]
|
||||
}
|
||||
output = await tool_executor.ainvoke(agent_action)
|
||||
return {"intermediate_steps": [(agent_action, str(output))]}
|
||||
if not isinstance(agent_action, list):
|
||||
agent_action = [agent_action]
|
||||
output = await tool_executor.abatch(agent_action, return_exceptions=True)
|
||||
return {
|
||||
"intermediate_steps": [
|
||||
(action, str(out)) for action, out in zip(agent_action, output)
|
||||
]
|
||||
}
|
||||
|
||||
# Define a new graph
|
||||
workflow = StateGraph(state)
|
||||
|
||||
# Define the two nodes we will cycle between
|
||||
workflow.add_node("agent", RunnableLambda(run_agent, arun_agent))
|
||||
workflow.add_node("action", RunnableLambda(execute_tools, aexecute_tools))
|
||||
workflow.add_node("agent", RunnableCallable(run_agent, arun_agent))
|
||||
workflow.add_node("action", RunnableCallable(execute_tools, aexecute_tools))
|
||||
|
||||
# Set the entrypoint as `agent`
|
||||
# This means that this node is the first one called
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
from typing import Any, Sequence, Union
|
||||
|
||||
from langchain_core.load.serializable import Serializable
|
||||
from langchain_core.runnables import RunnableBinding, RunnableConfig, RunnableLambda
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
from langgraph.utils import RunnableCallable
|
||||
|
||||
INVALID_TOOL_MSG_TEMPLATE = (
|
||||
"{requested_tool_name} is not a valid tool, "
|
||||
"try one of [{available_tool_names_str}]."
|
||||
@@ -26,29 +28,20 @@ class ToolInvocation(Serializable):
|
||||
"""The input to pass in to the Tool."""
|
||||
|
||||
|
||||
class ToolExecutor(RunnableBinding):
|
||||
tools: Sequence[BaseTool]
|
||||
tool_map: dict
|
||||
invalid_tool_msg_template: str
|
||||
|
||||
class ToolExecutor(RunnableCallable):
|
||||
def __init__(
|
||||
self,
|
||||
tools: Sequence[BaseTool],
|
||||
*,
|
||||
invalid_tool_msg_template: str = INVALID_TOOL_MSG_TEMPLATE,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
bound = RunnableLambda(self._execute, afunc=self._aexecute)
|
||||
super().__init__(
|
||||
bound=bound,
|
||||
tools=tools,
|
||||
tool_map={t.name: t for t in tools},
|
||||
invalid_tool_msg_template=invalid_tool_msg_template,
|
||||
**kwargs,
|
||||
)
|
||||
super().__init__(self._execute, afunc=self._aexecute, trace=False)
|
||||
self.tools = tools
|
||||
self.tool_map = {t.name: t for t in tools}
|
||||
self.invalid_tool_msg_template = invalid_tool_msg_template
|
||||
|
||||
def _execute(
|
||||
self, tool_invocation: ToolInvocationInterface, *, config: RunnableConfig
|
||||
self, tool_invocation: ToolInvocationInterface, config: RunnableConfig
|
||||
) -> Any:
|
||||
if tool_invocation.tool not in self.tool_map:
|
||||
return self.invalid_tool_msg_template.format(
|
||||
@@ -57,11 +50,11 @@ class ToolExecutor(RunnableBinding):
|
||||
)
|
||||
else:
|
||||
tool = self.tool_map[tool_invocation.tool]
|
||||
output = tool.invoke(tool_invocation.tool_input, config=config)
|
||||
output = tool.invoke(tool_invocation.tool_input, config)
|
||||
return output
|
||||
|
||||
async def _aexecute(
|
||||
self, tool_invocation: ToolInvocationInterface, *, config: RunnableConfig
|
||||
self, tool_invocation: ToolInvocationInterface, config: RunnableConfig
|
||||
) -> Any:
|
||||
if tool_invocation.tool not in self.tool_map:
|
||||
return self.invalid_tool_msg_template.format(
|
||||
@@ -70,5 +63,5 @@ class ToolExecutor(RunnableBinding):
|
||||
)
|
||||
else:
|
||||
tool = self.tool_map[tool_invocation.tool]
|
||||
output = await tool.ainvoke(tool_invocation.tool_input, config=config)
|
||||
output = await tool.ainvoke(tool_invocation.tool_input, config)
|
||||
return output
|
||||
|
||||
@@ -430,6 +430,7 @@ class Pregel(
|
||||
task.input,
|
||||
patch_config(
|
||||
config,
|
||||
run_name=self.name + "UpdateState",
|
||||
configurable={
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: task.writes.extend,
|
||||
@@ -489,6 +490,7 @@ class Pregel(
|
||||
task.input,
|
||||
patch_config(
|
||||
config,
|
||||
run_name=self.name + "UpdateState",
|
||||
configurable={
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: task.writes.extend,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from typing import Any, Iterator, Mapping, Optional, Sequence, Union
|
||||
|
||||
from langgraph.channels.base import BaseChannel, EmptyChannelError
|
||||
from langgraph.constants import TAG_HIDDEN
|
||||
from langgraph.pregel.log import logger
|
||||
from langgraph.pregel.types import PregelExecutableTask
|
||||
|
||||
@@ -81,9 +82,7 @@ def map_output_updates(
|
||||
) -> Optional[dict[str, Union[Any, dict[str, Any]]]]:
|
||||
"""Map pending writes (a sequence of tuples (channel, value)) to output chunk."""
|
||||
output_tasks = [
|
||||
t
|
||||
for t in tasks
|
||||
if not t.config or "langsmith:hidden" not in t.config.get("tags")
|
||||
t for t in tasks if not t.config or TAG_HIDDEN not in t.config.get("tags")
|
||||
]
|
||||
if isinstance(output_channels, str):
|
||||
if updated := {
|
||||
|
||||
+22
-11
@@ -6,7 +6,6 @@ from langchain_core.pydantic_v1 import Field
|
||||
from langchain_core.runnables import (
|
||||
Runnable,
|
||||
RunnableConfig,
|
||||
RunnableLambda,
|
||||
RunnablePassthrough,
|
||||
RunnableSequence,
|
||||
RunnableSerializable,
|
||||
@@ -17,11 +16,12 @@ from langchain_core.runnables.utils import ConfigurableFieldSpec
|
||||
|
||||
from langgraph.constants import CONFIG_KEY_READ
|
||||
from langgraph.pregel.write import ChannelWrite
|
||||
from langgraph.utils import RunnableCallable
|
||||
|
||||
READ_TYPE = Callable[[str, bool], Union[Any, dict[str, Any]]]
|
||||
|
||||
|
||||
class ChannelRead(RunnableLambda):
|
||||
class ChannelRead(RunnableCallable):
|
||||
channel: Union[str, list[str]]
|
||||
|
||||
fresh: bool = False
|
||||
@@ -46,12 +46,23 @@ class ChannelRead(RunnableLambda):
|
||||
*,
|
||||
fresh: bool = False,
|
||||
mapper: Optional[Callable[[Any], Any]] = None,
|
||||
tags: Optional[list[str]] = None,
|
||||
) -> None:
|
||||
super().__init__(func=self._read, afunc=self._aread)
|
||||
super().__init__(func=self._read, afunc=self._aread, tags=tags, name=None)
|
||||
self.fresh = fresh
|
||||
self.mapper = mapper
|
||||
self.channel = channel
|
||||
self.name = f"ChannelRead<{channel}>"
|
||||
|
||||
def get_name(
|
||||
self, suffix: Optional[str] = None, *, name: Optional[str] = None
|
||||
) -> str:
|
||||
if name:
|
||||
pass
|
||||
elif isinstance(self.channel, str):
|
||||
name = f"ChannelRead<{self.channel}>"
|
||||
else:
|
||||
name = f"ChannelRead<{','.join(self.channel)}>"
|
||||
return super().get_name(suffix, name=name)
|
||||
|
||||
def _read(self, _: Any, config: RunnableConfig) -> Any:
|
||||
try:
|
||||
@@ -80,7 +91,7 @@ class ChannelRead(RunnableLambda):
|
||||
return read(self.channel, self.fresh)
|
||||
|
||||
|
||||
default_bound: RunnablePassthrough = RunnablePassthrough()
|
||||
DEFAULT_BOUND: RunnablePassthrough = RunnablePassthrough()
|
||||
|
||||
|
||||
class PregelNode(RunnableBindingBase):
|
||||
@@ -92,7 +103,7 @@ class PregelNode(RunnableBindingBase):
|
||||
|
||||
writers: list[Runnable] = Field(default_factory=list)
|
||||
|
||||
bound: Runnable[Any, Any] = Field(default=default_bound)
|
||||
bound: Runnable[Any, Any] = Field(default=DEFAULT_BOUND)
|
||||
|
||||
kwargs: Mapping[str, Any] = Field(default_factory=dict)
|
||||
|
||||
@@ -125,11 +136,11 @@ class PregelNode(RunnableBindingBase):
|
||||
|
||||
def get_node(self) -> Optional[Runnable[Any, Any]]:
|
||||
writers = self.get_writers()
|
||||
if self.bound is default_bound and not writers:
|
||||
if self.bound is DEFAULT_BOUND and not writers:
|
||||
return None
|
||||
elif self.bound is default_bound and len(writers) == 1:
|
||||
elif self.bound is DEFAULT_BOUND and len(writers) == 1:
|
||||
return writers[0]
|
||||
elif self.bound is default_bound:
|
||||
elif self.bound is DEFAULT_BOUND:
|
||||
return RunnableSequence(*writers)
|
||||
elif writers:
|
||||
return RunnableSequence(self.bound, *writers)
|
||||
@@ -154,7 +165,7 @@ class PregelNode(RunnableBindingBase):
|
||||
triggers=triggers,
|
||||
mapper=mapper,
|
||||
writers=writers or [],
|
||||
bound=bound or default_bound,
|
||||
bound=bound or DEFAULT_BOUND,
|
||||
kwargs=kwargs or {},
|
||||
config=merge_configs(config, {"tags": tags or []}),
|
||||
**other_kwargs,
|
||||
@@ -201,7 +212,7 @@ class PregelNode(RunnableBindingBase):
|
||||
kwargs=self.kwargs,
|
||||
config=self.config,
|
||||
)
|
||||
elif self.bound is default_bound:
|
||||
elif self.bound is DEFAULT_BOUND:
|
||||
return PregelNode(
|
||||
channels=self.channels,
|
||||
triggers=self.triggers,
|
||||
|
||||
+15
-22
@@ -3,14 +3,11 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
from typing import Any, Callable, NamedTuple, Optional, Sequence, TypeVar, Union
|
||||
|
||||
from langchain_core.runnables import (
|
||||
Runnable,
|
||||
RunnableConfig,
|
||||
RunnablePassthrough,
|
||||
)
|
||||
from langchain_core.runnables import Runnable, RunnableConfig
|
||||
from langchain_core.runnables.utils import ConfigurableFieldSpec
|
||||
|
||||
from langgraph.constants import CONFIG_KEY_SEND
|
||||
from langgraph.utils import RunnableCallable
|
||||
|
||||
TYPE_SEND = Callable[[Sequence[tuple[str, Any]]], None]
|
||||
R = TypeVar("R", bound=Runnable)
|
||||
@@ -25,7 +22,7 @@ class ChannelWriteEntry(NamedTuple):
|
||||
skip_none: bool = False
|
||||
|
||||
|
||||
class ChannelWrite(RunnablePassthrough):
|
||||
class ChannelWrite(RunnableCallable):
|
||||
writes: Sequence[ChannelWriteEntry]
|
||||
"""
|
||||
Sequence of write entries, each of which is a tuple of:
|
||||
@@ -34,11 +31,11 @@ class ChannelWrite(RunnablePassthrough):
|
||||
- whether to skip writing if the mapped value is None
|
||||
"""
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
def __init__(self, writes: Sequence[ChannelWriteEntry]):
|
||||
super().__init__(func=self._write, afunc=self._awrite, writes=writes)
|
||||
def __init__(
|
||||
self, writes: Sequence[ChannelWriteEntry], *, tags: Optional[list[str]] = None
|
||||
):
|
||||
super().__init__(func=self._write, afunc=self._awrite, name=None, tags=tags)
|
||||
self.writes = writes
|
||||
|
||||
def __repr_args__(self) -> Any:
|
||||
return [("writes", self.writes)]
|
||||
@@ -46,15 +43,9 @@ class ChannelWrite(RunnablePassthrough):
|
||||
def get_name(
|
||||
self, suffix: Optional[str] = None, *, name: Optional[str] = None
|
||||
) -> str:
|
||||
return super().get_name(
|
||||
suffix,
|
||||
name=name
|
||||
or f"ChannelWrite<{','.join(chan for chan, _, _ in self.writes)}>",
|
||||
)
|
||||
|
||||
@property
|
||||
def is_channel_writer(self) -> bool:
|
||||
return True
|
||||
if not name:
|
||||
name = f"ChannelWrite<{','.join(chan for chan, _, _ in self.writes)}>"
|
||||
return super().get_name(suffix, name=name)
|
||||
|
||||
@property
|
||||
def config_specs(self) -> list[ConfigurableFieldSpec]:
|
||||
@@ -85,8 +76,8 @@ class ChannelWrite(RunnablePassthrough):
|
||||
for write, (_, _, skip_none) in zip(values, self.writes)
|
||||
if not skip_none or write[1] is not None
|
||||
]
|
||||
|
||||
self.do_write(config, **dict(values))
|
||||
return input
|
||||
|
||||
async def _awrite(self, input: Any, config: RunnableConfig) -> None:
|
||||
values = await asyncio.gather(
|
||||
@@ -104,8 +95,8 @@ class ChannelWrite(RunnablePassthrough):
|
||||
for val, (chan, _, skip_none) in zip(values, self.writes)
|
||||
if not skip_none or val is not None
|
||||
]
|
||||
|
||||
self.do_write(config, **dict(values))
|
||||
return input
|
||||
|
||||
@staticmethod
|
||||
def do_write(config: RunnableConfig, **values: Any) -> None:
|
||||
@@ -121,6 +112,8 @@ class ChannelWrite(RunnablePassthrough):
|
||||
|
||||
@staticmethod
|
||||
def register_writer(runnable: R) -> R:
|
||||
# using object.__setattr__ to work around objects that override __setattr__
|
||||
# eg. pydantic models and dataclasses
|
||||
object.__setattr__(runnable, "_is_channel_writer", True)
|
||||
return runnable
|
||||
|
||||
|
||||
@@ -1,4 +1,8 @@
|
||||
import enum
|
||||
from typing import Any, Awaitable, Callable, Optional
|
||||
|
||||
from langchain_core.runnables import Runnable, RunnableConfig
|
||||
from langchain_core.runnables.config import merge_configs
|
||||
|
||||
|
||||
# Before Python 3.11 native StrEnum is not available
|
||||
@@ -6,3 +10,58 @@ class StrEnum(str, enum.Enum):
|
||||
"""A string enum."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class RunnableCallable(Runnable):
|
||||
"""A much simpler version of RunnableLambda that requires sync and async functions."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
func: Callable[..., Optional[Runnable]],
|
||||
afunc: Optional[Callable[..., Awaitable[Optional[Runnable]]]] = None,
|
||||
*,
|
||||
name: Optional[str] = None,
|
||||
tags: Optional[list[str]] = None,
|
||||
trace: bool = True,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
self.name = name or func.__name__
|
||||
self.func = func
|
||||
self.afunc = afunc
|
||||
self.config = {"tags": tags} if tags else None
|
||||
self.kwargs = kwargs
|
||||
self.trace = trace
|
||||
|
||||
def __repr__(self) -> str:
|
||||
repr_args = {
|
||||
k: v
|
||||
for k, v in self.__dict__.items()
|
||||
if k not in {"name", "func", "afunc", "config", "kwargs", "trace"}
|
||||
}
|
||||
return f"{self.get_name()}({', '.join(f'{k}={v!r}' for k, v in repr_args.items())})"
|
||||
|
||||
def invoke(self, input: Any, config: Optional[RunnableConfig] = None) -> Any:
|
||||
if self.trace:
|
||||
ret = self._call_with_config(
|
||||
self.func, input, merge_configs(self.config, config), **self.kwargs
|
||||
)
|
||||
else:
|
||||
ret = self.func(input, merge_configs(self.config, config), **self.kwargs)
|
||||
if isinstance(ret, Runnable):
|
||||
return ret.invoke(input, config)
|
||||
return ret
|
||||
|
||||
async def ainvoke(self, input: Any, config: Optional[RunnableConfig] = None) -> Any:
|
||||
if not self.afunc:
|
||||
return self.invoke(input, config)
|
||||
if self.trace:
|
||||
ret = await self._acall_with_config(
|
||||
self.afunc, input, merge_configs(self.config, config), **self.kwargs
|
||||
)
|
||||
else:
|
||||
ret = await self.afunc(
|
||||
input, merge_configs(self.config, config), **self.kwargs
|
||||
)
|
||||
if isinstance(ret, Runnable):
|
||||
return await ret.ainvoke(input, config)
|
||||
return ret
|
||||
|
||||
Generated
+5
-6
@@ -1587,17 +1587,16 @@ extended-testing = ["aiosqlite (>=0.19.0,<0.20.0)", "aleph-alpha-client (>=2.15.
|
||||
|
||||
[[package]]
|
||||
name = "langchain-core"
|
||||
version = "0.1.30"
|
||||
version = "0.1.38"
|
||||
description = "Building applications with LLMs through composability"
|
||||
optional = false
|
||||
python-versions = ">=3.8.1,<4.0"
|
||||
python-versions = "<4.0,>=3.8.1"
|
||||
files = [
|
||||
{file = "langchain_core-0.1.30-py3-none-any.whl", hash = "sha256:c9643505e41d25ba8f20a2e8bf083d0f0d50b9a098d901511fff8df79f831ada"},
|
||||
{file = "langchain_core-0.1.30.tar.gz", hash = "sha256:e13a016e55e7f082ff3eeeda2d0cb505b89a8830e3a23c1d134d0a89d7871894"},
|
||||
{file = "langchain_core-0.1.38-py3-none-any.whl", hash = "sha256:d881b2754254cb4bdb0d5bb56e5c138d032b6e75e5cb21f151b01224b322e02b"},
|
||||
{file = "langchain_core-0.1.38.tar.gz", hash = "sha256:ee8da6d061c06cce7dc22fec224b6ecbc3a8de106d6dd9f409c7fe448ea41861"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
anyio = ">=3,<5"
|
||||
jsonpatch = ">=1.33,<2.0"
|
||||
langsmith = ">=0.1.0,<0.2.0"
|
||||
packaging = ">=23.2,<24.0"
|
||||
@@ -3860,4 +3859,4 @@ testing = ["big-O", "jaraco.functools", "jaraco.itertools", "more-itertools", "p
|
||||
[metadata]
|
||||
lock-version = "2.0"
|
||||
python-versions = ">=3.9.0,<4.0"
|
||||
content-hash = "2d35e923bf3902e0e11a305f58d17b0efc3fbb444dff8d6cb92e070a993115c9"
|
||||
content-hash = "3f31fdccb53a66dc294a53d63bcef0233f38c9909170113c44b8dc215fc77d7c"
|
||||
|
||||
+1
-1
@@ -9,7 +9,7 @@ repository = "https://www.github.com/langchain-ai/langgraph"
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
python = ">=3.9.0,<4.0"
|
||||
langchain-core = "^0.1.25"
|
||||
langchain-core = "^0.1.38"
|
||||
|
||||
|
||||
[tool.poetry.group.test.dependencies]
|
||||
|
||||
+178
-2400
File diff suppressed because it is too large
Load Diff
+18
-8
@@ -1773,6 +1773,7 @@ def test_conditional_entrypoint_graph_state(snapshot: SnapshotAssertion) -> None
|
||||
class AgentState(TypedDict, total=False):
|
||||
input: str
|
||||
output: str
|
||||
steps: Annotated[list[str], operator.add]
|
||||
|
||||
def left(data: AgentState) -> AgentState:
|
||||
return {"output": data["input"] + "->left"}
|
||||
@@ -1781,6 +1782,7 @@ def test_conditional_entrypoint_graph_state(snapshot: SnapshotAssertion) -> None
|
||||
return {"output": data["input"] + "->right"}
|
||||
|
||||
def should_start(data: AgentState) -> str:
|
||||
assert data["steps"] == [], "Expected input to be read from the state"
|
||||
# Logic to decide where to start
|
||||
if len(data["input"]) > 10:
|
||||
return "go-right"
|
||||
@@ -1891,6 +1893,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
"messages": [
|
||||
HumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"tool_calls": [
|
||||
@@ -1907,6 +1910,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
),
|
||||
ToolMessage(content="result for query", tool_call_id="tool_call123"),
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"tool_calls": [
|
||||
@@ -1931,7 +1935,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
),
|
||||
ToolMessage(content="result for another", tool_call_id="tool_call234"),
|
||||
ToolMessage(content="result for a third one", tool_call_id="tool_call567"),
|
||||
AIMessage(content="answer"),
|
||||
AIMessage(content="answer", id=AnyStr()),
|
||||
]
|
||||
}
|
||||
|
||||
@@ -1942,6 +1946,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"tool_calls": [
|
||||
@@ -1970,6 +1975,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"tool_calls": [
|
||||
@@ -2007,7 +2013,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
]
|
||||
}
|
||||
},
|
||||
{"agent": {"messages": [AIMessage(content="answer")]}},
|
||||
{"agent": {"messages": [AIMessage(content="answer", id=AnyStr())]}},
|
||||
]
|
||||
|
||||
|
||||
@@ -2065,6 +2071,7 @@ def test_prebuilt_chat(snapshot: SnapshotAssertion) -> None:
|
||||
"messages": [
|
||||
HumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"function_call": {"name": "search_api", "arguments": '"query"'}
|
||||
@@ -2072,13 +2079,14 @@ def test_prebuilt_chat(snapshot: SnapshotAssertion) -> None:
|
||||
),
|
||||
FunctionMessage(content="result for query", name="search_api"),
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"function_call": {"name": "search_api", "arguments": '"another"'}
|
||||
},
|
||||
),
|
||||
FunctionMessage(content="result for another", name="search_api"),
|
||||
AIMessage(content="answer"),
|
||||
AIMessage(content="answer", id=AnyStr()),
|
||||
]
|
||||
}
|
||||
|
||||
@@ -2089,6 +2097,7 @@ def test_prebuilt_chat(snapshot: SnapshotAssertion) -> None:
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"function_call": {
|
||||
@@ -2111,6 +2120,7 @@ def test_prebuilt_chat(snapshot: SnapshotAssertion) -> None:
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"function_call": {
|
||||
@@ -2129,7 +2139,7 @@ def test_prebuilt_chat(snapshot: SnapshotAssertion) -> None:
|
||||
]
|
||||
}
|
||||
},
|
||||
{"agent": {"messages": [AIMessage(content="answer")]}},
|
||||
{"agent": {"messages": [AIMessage(content="answer", id=AnyStr())]}},
|
||||
]
|
||||
|
||||
|
||||
@@ -2297,7 +2307,7 @@ def test_message_graph(
|
||||
FunctionMessage(
|
||||
content="result for query",
|
||||
name="search_api",
|
||||
id="00000000-0000-4000-8000-000000000014",
|
||||
id="00000000-0000-4000-8000-000000000013",
|
||||
),
|
||||
AIMessage(
|
||||
content="",
|
||||
@@ -2309,7 +2319,7 @@ def test_message_graph(
|
||||
FunctionMessage(
|
||||
content="result for another",
|
||||
name="search_api",
|
||||
id="00000000-0000-4000-8000-000000000026",
|
||||
id="00000000-0000-4000-8000-000000000024",
|
||||
),
|
||||
AIMessage(content="answer", id="ai3"),
|
||||
]
|
||||
@@ -2328,7 +2338,7 @@ def test_message_graph(
|
||||
"action": FunctionMessage(
|
||||
content="result for query",
|
||||
name="search_api",
|
||||
id="00000000-0000-4000-8000-000000000046",
|
||||
id="00000000-0000-4000-8000-000000000043",
|
||||
)
|
||||
},
|
||||
{
|
||||
@@ -2344,7 +2354,7 @@ def test_message_graph(
|
||||
"action": FunctionMessage(
|
||||
content="result for another",
|
||||
name="search_api",
|
||||
id="00000000-0000-4000-8000-000000000058",
|
||||
id="00000000-0000-4000-8000-000000000054",
|
||||
)
|
||||
},
|
||||
{"agent": AIMessage(content="answer", id="ai3")},
|
||||
|
||||
@@ -1814,6 +1814,7 @@ async def test_conditional_entrypoint_graph_state() -> None:
|
||||
class AgentState(TypedDict, total=False):
|
||||
input: str
|
||||
output: str
|
||||
steps: Annotated[list[str], operator.add]
|
||||
|
||||
async def left(data: AgentState) -> AgentState:
|
||||
return {"output": data["input"] + "->left"}
|
||||
@@ -1822,6 +1823,7 @@ async def test_conditional_entrypoint_graph_state() -> None:
|
||||
return {"output": data["input"] + "->right"}
|
||||
|
||||
def should_start(data: AgentState) -> str:
|
||||
assert data["steps"] == [], "Expected input to be read from the state"
|
||||
# Logic to decide where to start
|
||||
if len(data["input"]) > 10:
|
||||
return "go-right"
|
||||
@@ -1922,6 +1924,7 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
"messages": [
|
||||
HumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"tool_calls": [
|
||||
@@ -1938,6 +1941,7 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
),
|
||||
ToolMessage(content="result for query", tool_call_id="tool_call123"),
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"tool_calls": [
|
||||
@@ -1962,7 +1966,7 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
),
|
||||
ToolMessage(content="result for another", tool_call_id="tool_call234"),
|
||||
ToolMessage(content="result for a third one", tool_call_id="tool_call567"),
|
||||
AIMessage(content="answer"),
|
||||
AIMessage(content="answer", id=AnyStr()),
|
||||
]
|
||||
}
|
||||
|
||||
@@ -1976,6 +1980,7 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"tool_calls": [
|
||||
@@ -2004,6 +2009,7 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"tool_calls": [
|
||||
@@ -2041,7 +2047,7 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
]
|
||||
}
|
||||
},
|
||||
{"agent": {"messages": [AIMessage(content="answer")]}},
|
||||
{"agent": {"messages": [AIMessage(content="answer", id=AnyStr())]}},
|
||||
]
|
||||
|
||||
|
||||
@@ -2094,6 +2100,7 @@ async def test_prebuilt_chat() -> None:
|
||||
"messages": [
|
||||
HumanMessage(content="what is weather in sf"),
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"function_call": {"name": "search_api", "arguments": '"query"'}
|
||||
@@ -2101,13 +2108,14 @@ async def test_prebuilt_chat() -> None:
|
||||
),
|
||||
FunctionMessage(content="result for query", name="search_api"),
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"function_call": {"name": "search_api", "arguments": '"another"'}
|
||||
},
|
||||
),
|
||||
FunctionMessage(content="result for another", name="search_api"),
|
||||
AIMessage(content="answer"),
|
||||
AIMessage(content="answer", id=AnyStr()),
|
||||
]
|
||||
}
|
||||
|
||||
@@ -2121,6 +2129,7 @@ async def test_prebuilt_chat() -> None:
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"function_call": {
|
||||
@@ -2143,6 +2152,7 @@ async def test_prebuilt_chat() -> None:
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=AnyStr(),
|
||||
content="",
|
||||
additional_kwargs={
|
||||
"function_call": {
|
||||
@@ -2161,7 +2171,7 @@ async def test_prebuilt_chat() -> None:
|
||||
]
|
||||
}
|
||||
},
|
||||
{"agent": {"messages": [AIMessage(content="answer")]}},
|
||||
{"agent": {"messages": [AIMessage(content="answer", id=AnyStr())]}},
|
||||
]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user