Merge pull request #267 from langchain-ai/nc/2apr/optimize-tracing-output

Optimize tracing output of Graph/StateGraph/MessageGraph
This commit is contained in:
Nuno Campos
2024-04-02 15:32:24 -07:00
committed by GitHub
17 changed files with 677 additions and 2650 deletions
+10 -10
View File
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+3
View File
@@ -1,3 +1,6 @@
CONFIG_KEY_SEND = "__pregel_send"
CONFIG_KEY_READ = "__pregel_read"
INTERRUPT = "__interrupt__"
TAG_HIDDEN = "langsmith:hidden"
+28 -57
View File
@@ -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
View File
@@ -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
+19 -21
View File
@@ -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
+12 -19
View File
@@ -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
+2
View File
@@ -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,
+2 -3
View File
@@ -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
View File
@@ -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
View File
@@ -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
+59
View File
@@ -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
View File
@@ -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
View File
@@ -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]
File diff suppressed because it is too large Load Diff
+18 -8
View File
@@ -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")},
+14 -4
View File
@@ -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())]}},
]