mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-26 17:42:24 +02:00
Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f8006b2fee | ||
|
|
2164b7daa3 | ||
|
|
6d20a0b9c7 | ||
|
|
df8becd5cf | ||
|
|
056ba91a71 | ||
|
|
02300de24c | ||
|
|
ac16bdb795 | ||
|
|
0d4ac836e3 | ||
|
|
201c8015ea |
@@ -129,11 +129,25 @@ REDIRECT_MAP = {
|
||||
"how-tos/human_in_the_loop/edit-graph-state.ipynb": "https://docs.langchain.com/oss/python/langgraph/use-time-travel",
|
||||
|
||||
# LGP mintlify migration redirects
|
||||
"examples/index.md": "https://docs.langchain.com/oss/python/learn",
|
||||
"guides/index.md": "https://docs.langchain.com/oss/python/langgraph/overview",
|
||||
"concepts/index.md": "https://docs.langchain.com/oss/python/langgraph/overview",
|
||||
"tutorials/index.md": "https://docs.langchain.com/oss/python/learn",
|
||||
"llms-txt-overview.md": "https://docs.langchain.com/llms.txt",
|
||||
"tutorials/rag/langgraph_adaptive_rag.md": "https://docs.langchain.com/oss/python/langgraph/agentic-rag",
|
||||
"tutorials/multi_agent/multi-agent-collaboration.ipynb": "https://docs.langchain.com/oss/python/langchain/multi-agent",
|
||||
"how-tos/create-react-agent-manage-message-history.ipynb": "https://docs.langchain.com/oss/python/langgraph/add-memory",
|
||||
"how-tos/many-tools.ipynb": "https://docs.langchain.com/oss/python/langchain/tools",
|
||||
"tutorials/customer-support/customer-support.ipynb": "https://docs.langchain.com/oss/python/langgraph/agentic-rag",
|
||||
"how-tos/react-agent-structured-output.ipynb": "https://docs.langchain.com/oss/python/langchain/agents#structured-output",
|
||||
"tutorials/code_assistant/langgraph_code_assistant.ipynb": "https://docs.langchain.com/oss/python/langgraph/agentic-rag",
|
||||
"tutorials/multi_agent/hierarchical_agent_teams.ipynb": "https://docs.langchain.com/oss/python/langchain/supervisor",
|
||||
"tutorials/auth/getting_started.md": "https://docs.langchain.com/langsmith/auth",
|
||||
"tutorials/auth/resource_auth.md": "https://docs.langchain.com/langsmith/resource-auth",
|
||||
"tutorials/auth/add_auth_server.md": "https://docs.langchain.com/langsmith/add-auth-server",
|
||||
"how-tos/use-remote-graph.md": "https://docs.langchain.com/langsmith/use-remote-graph",
|
||||
"how-tos/autogen-integration.md": "https://docs.langchain.com/langsmith/autogen-integration",
|
||||
"how-tos/human_in_the_loop/wait-user-input.ipynb": "https://docs.langchain.com/oss/python/langgraph/interrupts",
|
||||
"cloud/how-tos/use_stream_react.md": "https://docs.langchain.com/langsmith/use-stream-react",
|
||||
"cloud/how-tos/generative_ui_react.md": "https://docs.langchain.com/langsmith/generative-ui-react",
|
||||
"concepts/langgraph_platform.md": "https://docs.langchain.com/langsmith/deployments",
|
||||
@@ -219,12 +233,15 @@ REDIRECT_MAP = {
|
||||
"tutorials/get-started/6-time-travel.md": "https://docs.langchain.com/oss/python/langgraph/quickstart",
|
||||
"tutorials/langsmith/local-server.md": "https://docs.langchain.com/oss/python/langgraph/local-server",
|
||||
"tutorials/workflows.md": "https://docs.langchain.com/oss/python/langgraph/workflows-agents",
|
||||
"tutorials/plan-and-execute/plan-and-execute.ipynb": "https://docs.langchain.com/oss/python/langchain/middleware/built-in#to-do-list",
|
||||
"tutorials/langgraph-platform/local-server/local-server.md": "https://docs.langchain.com/langsmith/local-server",
|
||||
"concepts/agentic_concepts.md": "https://docs.langchain.com/oss/python/langgraph/workflows-agents",
|
||||
"guides/index.md": "https://docs.langchain.com/oss/python/langchain/overview",
|
||||
"agents/overview.md": "https://docs.langchain.com/oss/python/langchain/agents",
|
||||
"agents/run_agents.md": "https://docs.langchain.com/oss/python/langgraph/quickstart",
|
||||
"concepts/low_level.md": "https://docs.langchain.com/oss/python/langgraph/graph-api",
|
||||
"how-tos/graph-api.md": "https://docs.langchain.com/oss/python/langgraph/graph-api",
|
||||
"how-tos/react-agent-from-scratch.ipynb": "https://docs.langchain.com/oss/python/langchain/quickstart",
|
||||
"concepts/functional_api.md": "https://docs.langchain.com/oss/python/langgraph/functional-api",
|
||||
"how-tos/use-functional-api.md": "https://docs.langchain.com/oss/python/langgraph/functional-api",
|
||||
"concepts/pregel.md": "https://docs.langchain.com/oss/python/langgraph/pregel",
|
||||
@@ -282,8 +299,8 @@ REDIRECT_MAP = {
|
||||
"reference/supervisor.md": "https://reference.langchain.com/python/langgraph/supervisor/",
|
||||
"reference/swarm.md": "https://reference.langchain.com/python/langgraph/swarm/",
|
||||
"reference/mcp.md": "https://reference.langchain.com/python/langgraph/mcp/",
|
||||
"cloud/reference/sdk/python_sdk_ref.md": "https://reference.langchain.com/python/platform/python_sdk/",
|
||||
"reference/remote_graph.md": "https://reference.langchain.com/python/platform/remote_graph/",
|
||||
"cloud/reference/sdk/python_sdk_ref.md": "https://reference.langchain.com/python/langsmith/deployment/sdk/",
|
||||
"reference/remote_graph.md": "https://reference.langchain.com/python/langsmith/deployment/remote_graph/",
|
||||
|
||||
# additional exclude-search entries from mkdocs.yml
|
||||
"additional-resources/index.md": "https://docs.langchain.com/oss/python/langchain/overview",
|
||||
@@ -793,6 +810,17 @@ def on_post_build(config):
|
||||
# Track which paths have explicit redirects
|
||||
redirected_paths = set()
|
||||
|
||||
# Collect all existing HTML files in the site
|
||||
all_html_files = set()
|
||||
for root, dirs, files in os.walk(site_dir):
|
||||
for file in files:
|
||||
if file.endswith(".html"):
|
||||
# Get relative path from site_dir
|
||||
html_path = os.path.relpath(os.path.join(root, file), site_dir)
|
||||
# Normalize path separators to forward slashes
|
||||
html_path = html_path.replace(os.sep, "/")
|
||||
all_html_files.add(html_path)
|
||||
|
||||
# Process explicit redirects from REDIRECT_MAP
|
||||
for page_old, page_new in REDIRECT_MAP.items():
|
||||
# Convert .ipynb to .md for path calculation
|
||||
@@ -847,6 +875,24 @@ def on_post_build(config):
|
||||
|
||||
_write_html(site_dir, old_html_path, new_html_path)
|
||||
|
||||
# Create catch-all redirects for any HTML files not explicitly redirected
|
||||
catchall_url = "https://docs.langchain.com/oss/python/langgraph/overview"
|
||||
for html_file in all_html_files:
|
||||
# Skip if this file is already explicitly redirected
|
||||
if html_file in redirected_paths:
|
||||
continue
|
||||
|
||||
# Skip the root index.html (we handle that separately)
|
||||
if html_file == "index.html":
|
||||
continue
|
||||
|
||||
# Skip reference documentation (keep those accessible)
|
||||
if html_file.startswith("reference/"):
|
||||
continue
|
||||
|
||||
# Create redirect for this unmapped file
|
||||
_write_html(site_dir, html_file, catchall_url)
|
||||
|
||||
# Create root index.html redirect
|
||||
root_redirect_html = """<!doctype html>
|
||||
<html lang="en">
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
{
|
||||
"permissions": {
|
||||
"allow": [
|
||||
"Bash(rg:*)",
|
||||
"Bash(python:*)",
|
||||
"Bash(grep:*)",
|
||||
"Bash(sed:*)",
|
||||
"Bash(awk:*)",
|
||||
"Bash(uv run mypy:*)",
|
||||
"Bash(uv run:*)",
|
||||
"Bash(make test:*)",
|
||||
"Bash(make test_parallel:*)"
|
||||
],
|
||||
"deny": []
|
||||
}
|
||||
}
|
||||
@@ -258,9 +258,11 @@ def apply_writes(
|
||||
next_version = None
|
||||
else:
|
||||
next_version = get_next_version(
|
||||
max(checkpoint["channel_versions"].values())
|
||||
if checkpoint["channel_versions"]
|
||||
else None,
|
||||
(
|
||||
max(checkpoint["channel_versions"].values())
|
||||
if checkpoint["channel_versions"]
|
||||
else None
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
@@ -491,6 +493,11 @@ def prepare_next_tasks(
|
||||
PUSH_TRIGGER = (PUSH,)
|
||||
|
||||
|
||||
class _TaskIDFn(Protocol):
|
||||
def __call__(self, namespace: bytes, *parts: str | bytes) -> str:
|
||||
pass
|
||||
|
||||
|
||||
def prepare_single_task(
|
||||
task_path: tuple[Any, ...],
|
||||
task_id_checksum: str | None,
|
||||
@@ -520,250 +527,50 @@ def prepare_single_task(
|
||||
task_id_func = _xxhash_str if checkpoint["v"] > 1 else _uuid5_str
|
||||
|
||||
if task_path[0] == PUSH and isinstance(task_path[-1], Call):
|
||||
# (PUSH, parent task path, idx of PUSH write, id of parent task, Call)
|
||||
task_path_t = cast(tuple[str, tuple, int, str, Call], task_path)
|
||||
call = task_path_t[-1]
|
||||
proc_ = get_runnable_for_task(call.func)
|
||||
name = proc_.name
|
||||
if name is None:
|
||||
raise ValueError("`call` functions must have a `__name__` attribute")
|
||||
# create task id
|
||||
triggers: Sequence[str] = PUSH_TRIGGER
|
||||
checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name
|
||||
task_id = task_id_func(
|
||||
checkpoint_id_bytes,
|
||||
checkpoint_ns,
|
||||
str(step),
|
||||
name,
|
||||
PUSH,
|
||||
task_path_str(task_path[1]),
|
||||
str(task_path[2]),
|
||||
return prepare_push_task_functional(
|
||||
cast(tuple[str, tuple, int, str, Call], task_path),
|
||||
task_id_checksum,
|
||||
checkpoint=checkpoint,
|
||||
checkpoint_id_bytes=checkpoint_id_bytes,
|
||||
pending_writes=pending_writes,
|
||||
channels=channels,
|
||||
managed=managed,
|
||||
config=config,
|
||||
step=step,
|
||||
stop=stop,
|
||||
for_execution=for_execution,
|
||||
store=store,
|
||||
checkpointer=checkpointer,
|
||||
manager=manager,
|
||||
cache_policy=cache_policy,
|
||||
retry_policy=retry_policy,
|
||||
parent_ns=parent_ns,
|
||||
task_id_func=task_id_func,
|
||||
)
|
||||
|
||||
elif task_path[0] == PUSH:
|
||||
return prepare_push_task_send(
|
||||
cast(tuple[str, tuple], task_path),
|
||||
task_id_checksum,
|
||||
checkpoint=checkpoint,
|
||||
checkpoint_id_bytes=checkpoint_id_bytes,
|
||||
pending_writes=pending_writes,
|
||||
channels=channels,
|
||||
managed=managed,
|
||||
config=config,
|
||||
step=step,
|
||||
processes=processes,
|
||||
stop=stop,
|
||||
for_execution=for_execution,
|
||||
store=store,
|
||||
checkpointer=checkpointer,
|
||||
manager=manager,
|
||||
cache_policy=cache_policy,
|
||||
retry_policy=retry_policy,
|
||||
parent_ns=parent_ns,
|
||||
task_id_func=task_id_func,
|
||||
)
|
||||
task_checkpoint_ns = f"{checkpoint_ns}:{task_id}"
|
||||
# we append True to the task path to indicate that a call is being
|
||||
# made, so we should not return interrupts from this task (responsibility lies with the parent)
|
||||
task_path = (*task_path[:3], True)
|
||||
metadata = {
|
||||
"langgraph_step": step,
|
||||
"langgraph_node": name,
|
||||
"langgraph_triggers": triggers,
|
||||
"langgraph_path": task_path,
|
||||
"langgraph_checkpoint_ns": task_checkpoint_ns,
|
||||
}
|
||||
if task_id_checksum is not None:
|
||||
assert task_id == task_id_checksum, f"{task_id} != {task_id_checksum}"
|
||||
if for_execution:
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
cache_policy = call.cache_policy or cache_policy
|
||||
if cache_policy:
|
||||
args_key = cache_policy.key_func(*call.input[0], **call.input[1])
|
||||
cache_key: CacheKey | None = CacheKey(
|
||||
(
|
||||
CACHE_NS_WRITES,
|
||||
(identifier(call.func) or "__dynamic__"),
|
||||
),
|
||||
xxh3_128_hexdigest(
|
||||
args_key.encode() if isinstance(args_key, str) else args_key,
|
||||
),
|
||||
cache_policy.ttl,
|
||||
)
|
||||
else:
|
||||
cache_key = None
|
||||
scratchpad = _scratchpad(
|
||||
config[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||
pending_writes,
|
||||
task_id,
|
||||
xxh3_128_hexdigest(task_checkpoint_ns.encode()),
|
||||
config[CONF].get(CONFIG_KEY_RESUME_MAP),
|
||||
step,
|
||||
stop,
|
||||
)
|
||||
runtime = cast(
|
||||
Runtime, configurable.get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME)
|
||||
)
|
||||
runtime = runtime.override(store=store)
|
||||
return PregelExecutableTask(
|
||||
name,
|
||||
call.input,
|
||||
proc_,
|
||||
writes,
|
||||
patch_config(
|
||||
merge_configs(config, {"metadata": metadata}),
|
||||
run_name=name,
|
||||
callbacks=call.callbacks
|
||||
or (manager.get_child(f"graph:step:{step}") if manager else None),
|
||||
configurable={
|
||||
CONFIG_KEY_TASK_ID: task_id,
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: writes.extend,
|
||||
CONFIG_KEY_READ: partial(
|
||||
local_read,
|
||||
scratchpad,
|
||||
channels,
|
||||
managed,
|
||||
PregelTaskWrites(task_path, name, writes, triggers),
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINTER: (
|
||||
checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER)
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINT_MAP: {
|
||||
**configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {}),
|
||||
parent_ns: checkpoint["id"],
|
||||
},
|
||||
CONFIG_KEY_CHECKPOINT_ID: None,
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
|
||||
CONFIG_KEY_SCRATCHPAD: scratchpad,
|
||||
CONFIG_KEY_RUNTIME: runtime,
|
||||
},
|
||||
),
|
||||
triggers,
|
||||
call.retry_policy or retry_policy,
|
||||
cache_key,
|
||||
task_id,
|
||||
task_path,
|
||||
)
|
||||
else:
|
||||
return PregelTask(task_id, name, task_path)
|
||||
elif task_path[0] == PUSH:
|
||||
if len(task_path) == 2:
|
||||
# SEND tasks, executed in superstep n+1
|
||||
# (PUSH, idx of pending send)
|
||||
idx = cast(int, task_path[1])
|
||||
if not channels[TASKS].is_available():
|
||||
return
|
||||
sends: Sequence[Send] = channels[TASKS].get()
|
||||
if idx < 0 or idx >= len(sends):
|
||||
return
|
||||
packet = sends[idx]
|
||||
if not isinstance(packet, Send):
|
||||
logger.warning(
|
||||
f"Ignoring invalid packet type {type(packet)} in pending sends"
|
||||
)
|
||||
return
|
||||
|
||||
if packet.node not in processes:
|
||||
logger.warning(
|
||||
f"Ignoring unknown node name {packet.node} in pending sends"
|
||||
)
|
||||
return
|
||||
# find process
|
||||
proc = processes[packet.node]
|
||||
proc_node = proc.node
|
||||
if proc_node is None:
|
||||
return
|
||||
# create task id
|
||||
triggers = PUSH_TRIGGER
|
||||
checkpoint_ns = (
|
||||
f"{parent_ns}{NS_SEP}{packet.node}" if parent_ns else packet.node
|
||||
)
|
||||
task_id = task_id_func(
|
||||
checkpoint_id_bytes,
|
||||
checkpoint_ns,
|
||||
str(step),
|
||||
packet.node,
|
||||
PUSH,
|
||||
str(idx),
|
||||
)
|
||||
else:
|
||||
logger.warning(f"Ignoring invalid PUSH task path {task_path}")
|
||||
return
|
||||
task_checkpoint_ns = f"{checkpoint_ns}:{task_id}"
|
||||
# we append False to the task path to indicate that a call is not being made
|
||||
# so we should return interrupts from this task
|
||||
task_path = (*task_path[:3], False)
|
||||
metadata = {
|
||||
"langgraph_step": step,
|
||||
"langgraph_node": packet.node,
|
||||
"langgraph_triggers": triggers,
|
||||
"langgraph_path": task_path,
|
||||
"langgraph_checkpoint_ns": task_checkpoint_ns,
|
||||
}
|
||||
if task_id_checksum is not None:
|
||||
assert task_id == task_id_checksum, f"{task_id} != {task_id_checksum}"
|
||||
if for_execution:
|
||||
if proc.metadata:
|
||||
metadata.update(proc.metadata)
|
||||
writes = deque()
|
||||
cache_policy = proc.cache_policy or cache_policy
|
||||
if cache_policy:
|
||||
args_key = cache_policy.key_func(packet.arg)
|
||||
cache_key = CacheKey(
|
||||
(
|
||||
CACHE_NS_WRITES,
|
||||
(identifier(proc) or "__dynamic__"),
|
||||
packet.node,
|
||||
),
|
||||
xxh3_128_hexdigest(
|
||||
args_key.encode() if isinstance(args_key, str) else args_key,
|
||||
),
|
||||
cache_policy.ttl,
|
||||
)
|
||||
else:
|
||||
cache_key = None
|
||||
scratchpad = _scratchpad(
|
||||
config[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||
pending_writes,
|
||||
task_id,
|
||||
xxh3_128_hexdigest(task_checkpoint_ns.encode()),
|
||||
config[CONF].get(CONFIG_KEY_RESUME_MAP),
|
||||
step,
|
||||
stop,
|
||||
)
|
||||
runtime = cast(
|
||||
Runtime, configurable.get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME)
|
||||
)
|
||||
runtime = runtime.override(
|
||||
store=store, previous=checkpoint["channel_values"].get(PREVIOUS, None)
|
||||
)
|
||||
additional_config: RunnableConfig = {
|
||||
"metadata": metadata,
|
||||
"tags": proc.tags,
|
||||
}
|
||||
return PregelExecutableTask(
|
||||
packet.node,
|
||||
packet.arg,
|
||||
proc_node,
|
||||
writes,
|
||||
patch_config(
|
||||
merge_configs(config, additional_config),
|
||||
run_name=packet.node,
|
||||
callbacks=(
|
||||
manager.get_child(f"graph:step:{step}") if manager else None
|
||||
),
|
||||
configurable={
|
||||
CONFIG_KEY_TASK_ID: task_id,
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: writes.extend,
|
||||
CONFIG_KEY_READ: partial(
|
||||
local_read,
|
||||
scratchpad,
|
||||
channels,
|
||||
managed,
|
||||
PregelTaskWrites(task_path, packet.node, writes, triggers),
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINTER: (
|
||||
checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER)
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINT_MAP: {
|
||||
**configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {}),
|
||||
parent_ns: checkpoint["id"],
|
||||
},
|
||||
CONFIG_KEY_CHECKPOINT_ID: None,
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
|
||||
CONFIG_KEY_SCRATCHPAD: scratchpad,
|
||||
CONFIG_KEY_RUNTIME: runtime,
|
||||
},
|
||||
),
|
||||
triggers,
|
||||
proc.retry_policy or retry_policy,
|
||||
cache_key,
|
||||
task_id,
|
||||
task_path,
|
||||
writers=proc.flat_writers,
|
||||
subgraphs=proc.subgraphs,
|
||||
)
|
||||
else:
|
||||
return PregelTask(task_id, packet.node, task_path)
|
||||
elif task_path[0] == PULL:
|
||||
# (PULL, node name)
|
||||
name = cast(str, task_path[1])
|
||||
@@ -834,7 +641,7 @@ def prepare_single_task(
|
||||
if node := proc.node:
|
||||
if proc.metadata:
|
||||
metadata.update(proc.metadata)
|
||||
writes = deque()
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
cache_policy = proc.cache_policy or cache_policy
|
||||
if cache_policy:
|
||||
args_key = cache_policy.key_func(val)
|
||||
@@ -845,9 +652,11 @@ def prepare_single_task(
|
||||
name,
|
||||
),
|
||||
xxh3_128_hexdigest(
|
||||
args_key.encode()
|
||||
if isinstance(args_key, str)
|
||||
else args_key,
|
||||
(
|
||||
args_key.encode()
|
||||
if isinstance(args_key, str)
|
||||
else args_key
|
||||
),
|
||||
),
|
||||
cache_policy.ttl,
|
||||
)
|
||||
@@ -870,7 +679,9 @@ def prepare_single_task(
|
||||
node,
|
||||
writes,
|
||||
patch_config(
|
||||
merge_configs(config, additional_config),
|
||||
merge_configs(
|
||||
config, cast(RunnableConfig, additional_config)
|
||||
),
|
||||
run_name=name,
|
||||
callbacks=(
|
||||
manager.get_child(f"graph:step:{step}")
|
||||
@@ -919,6 +730,297 @@ def prepare_single_task(
|
||||
return PregelTask(task_id, name, task_path[:3])
|
||||
|
||||
|
||||
def prepare_push_task_functional(
|
||||
task_path: tuple[str, tuple, int, str, Call],
|
||||
# (PUSH, parent task path, idx of PUSH write, id of parent task, Call)
|
||||
task_id_checksum: str | None,
|
||||
*,
|
||||
checkpoint: Checkpoint,
|
||||
checkpoint_id_bytes: bytes,
|
||||
pending_writes: list[PendingWrite],
|
||||
channels: Mapping[str, BaseChannel],
|
||||
managed: ManagedValueMapping,
|
||||
config: RunnableConfig,
|
||||
step: int,
|
||||
stop: int,
|
||||
for_execution: bool,
|
||||
store: BaseStore | None = None,
|
||||
checkpointer: BaseCheckpointSaver | None = None,
|
||||
manager: None | ParentRunManager | AsyncParentRunManager = None,
|
||||
cache_policy: CachePolicy | None = None,
|
||||
retry_policy: Sequence[RetryPolicy] = (),
|
||||
parent_ns: str,
|
||||
# namespace: bytes, *parts: str | bytes
|
||||
task_id_func: _TaskIDFn,
|
||||
) -> PregelTask | PregelExecutableTask:
|
||||
"""Prepare a push task with an attached caller. Used for the functional API."""
|
||||
configurable = config.get(CONF, {})
|
||||
|
||||
call = task_path[-1]
|
||||
proc_ = get_runnable_for_task(call.func)
|
||||
name = proc_.name
|
||||
if name is None:
|
||||
raise ValueError("`call` functions must have a `__name__` attribute")
|
||||
# create task id
|
||||
triggers: Sequence[str] = PUSH_TRIGGER
|
||||
checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name
|
||||
task_id = task_id_func(
|
||||
checkpoint_id_bytes,
|
||||
checkpoint_ns,
|
||||
str(step),
|
||||
name,
|
||||
PUSH,
|
||||
task_path_str(task_path[1]),
|
||||
str(task_path[2]),
|
||||
)
|
||||
task_checkpoint_ns = f"{checkpoint_ns}:{task_id}"
|
||||
# we append True to the task path to indicate that a call is being
|
||||
# made, so we should not return interrupts from this task (responsibility lies with the parent)
|
||||
in_progress_task_path = (*task_path[:3], True)
|
||||
metadata = {
|
||||
"langgraph_step": step,
|
||||
"langgraph_node": name,
|
||||
"langgraph_triggers": triggers,
|
||||
"langgraph_path": in_progress_task_path,
|
||||
"langgraph_checkpoint_ns": task_checkpoint_ns,
|
||||
}
|
||||
if task_id_checksum is not None:
|
||||
assert task_id == task_id_checksum, f"{task_id} != {task_id_checksum}"
|
||||
if for_execution:
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
cache_policy = call.cache_policy or cache_policy
|
||||
if cache_policy:
|
||||
args_key = cache_policy.key_func(*call.input[0], **call.input[1])
|
||||
cache_key: CacheKey | None = CacheKey(
|
||||
(
|
||||
CACHE_NS_WRITES,
|
||||
(identifier(call.func) or "__dynamic__"),
|
||||
),
|
||||
xxh3_128_hexdigest(
|
||||
args_key.encode() if isinstance(args_key, str) else args_key,
|
||||
),
|
||||
cache_policy.ttl,
|
||||
)
|
||||
else:
|
||||
cache_key = None
|
||||
scratchpad = _scratchpad(
|
||||
configurable.get(CONFIG_KEY_SCRATCHPAD),
|
||||
pending_writes,
|
||||
task_id,
|
||||
xxh3_128_hexdigest(task_checkpoint_ns.encode()),
|
||||
configurable.get(CONFIG_KEY_RESUME_MAP),
|
||||
step,
|
||||
stop,
|
||||
)
|
||||
runtime = cast(Runtime, configurable.get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME))
|
||||
runtime = runtime.override(store=store)
|
||||
return PregelExecutableTask(
|
||||
name,
|
||||
call.input,
|
||||
proc_,
|
||||
writes,
|
||||
patch_config(
|
||||
merge_configs(config, {"metadata": metadata}),
|
||||
run_name=name,
|
||||
callbacks=call.callbacks
|
||||
or (manager.get_child(f"graph:step:{step}") if manager else None),
|
||||
configurable={
|
||||
CONFIG_KEY_TASK_ID: task_id,
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: writes.extend,
|
||||
CONFIG_KEY_READ: partial(
|
||||
local_read,
|
||||
scratchpad,
|
||||
channels,
|
||||
managed,
|
||||
PregelTaskWrites(in_progress_task_path, name, writes, triggers),
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINTER: (
|
||||
checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER)
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINT_MAP: {
|
||||
**configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {}),
|
||||
parent_ns: checkpoint["id"],
|
||||
},
|
||||
CONFIG_KEY_CHECKPOINT_ID: None,
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
|
||||
CONFIG_KEY_SCRATCHPAD: scratchpad,
|
||||
CONFIG_KEY_RUNTIME: runtime,
|
||||
},
|
||||
),
|
||||
triggers,
|
||||
call.retry_policy or retry_policy,
|
||||
cache_key,
|
||||
task_id,
|
||||
in_progress_task_path,
|
||||
)
|
||||
else:
|
||||
return PregelTask(task_id, name, in_progress_task_path)
|
||||
|
||||
|
||||
def prepare_push_task_send(
|
||||
task_path: tuple[str, tuple],
|
||||
# (PUSH, parent task path)
|
||||
task_id_checksum: str | None,
|
||||
*,
|
||||
checkpoint: Checkpoint,
|
||||
checkpoint_id_bytes: bytes,
|
||||
pending_writes: list[PendingWrite],
|
||||
channels: Mapping[str, BaseChannel],
|
||||
managed: ManagedValueMapping,
|
||||
config: RunnableConfig,
|
||||
step: int,
|
||||
stop: int,
|
||||
for_execution: bool,
|
||||
store: BaseStore | None = None,
|
||||
checkpointer: BaseCheckpointSaver | None = None,
|
||||
manager: None | ParentRunManager | AsyncParentRunManager = None,
|
||||
cache_policy: CachePolicy | None = None,
|
||||
retry_policy: Sequence[RetryPolicy] = (),
|
||||
parent_ns: str,
|
||||
task_id_func: _TaskIDFn,
|
||||
processes: Mapping[str, PregelNode],
|
||||
) -> PregelTask | PregelExecutableTask | None:
|
||||
if len(task_path) == 2:
|
||||
# SEND tasks, executed in superstep n+1
|
||||
# (PUSH, idx of pending send)
|
||||
idx = cast(int, task_path[1])
|
||||
if not channels[TASKS].is_available():
|
||||
return
|
||||
sends: Sequence[Send] = channels[TASKS].get()
|
||||
if idx < 0 or idx >= len(sends):
|
||||
return
|
||||
packet = sends[idx]
|
||||
if not isinstance(packet, Send):
|
||||
logger.warning(
|
||||
f"Ignoring invalid packet type {type(packet)} in pending sends"
|
||||
)
|
||||
return
|
||||
|
||||
if packet.node not in processes:
|
||||
logger.warning(f"Ignoring unknown node name {packet.node} in pending sends")
|
||||
return
|
||||
# find process
|
||||
proc = processes[packet.node]
|
||||
proc_node = proc.node
|
||||
if proc_node is None:
|
||||
return
|
||||
# create task id
|
||||
triggers = PUSH_TRIGGER
|
||||
checkpoint_ns = (
|
||||
f"{parent_ns}{NS_SEP}{packet.node}" if parent_ns else packet.node
|
||||
)
|
||||
task_id = task_id_func(
|
||||
checkpoint_id_bytes,
|
||||
checkpoint_ns,
|
||||
str(step),
|
||||
packet.node,
|
||||
PUSH,
|
||||
str(idx),
|
||||
)
|
||||
else:
|
||||
logger.warning(f"Ignoring invalid PUSH task path {task_path}")
|
||||
return
|
||||
configurable = config.get(CONF, {})
|
||||
task_checkpoint_ns = f"{checkpoint_ns}:{task_id}"
|
||||
# we append False to the task path to indicate that a call is not being made
|
||||
# so we should return interrupts from this task
|
||||
translated_task_path = (*task_path[:3], False)
|
||||
metadata = {
|
||||
"langgraph_step": step,
|
||||
"langgraph_node": packet.node,
|
||||
"langgraph_triggers": triggers,
|
||||
"langgraph_path": translated_task_path,
|
||||
"langgraph_checkpoint_ns": task_checkpoint_ns,
|
||||
}
|
||||
if task_id_checksum is not None:
|
||||
assert task_id == task_id_checksum, f"{task_id} != {task_id_checksum}"
|
||||
if for_execution:
|
||||
if proc.metadata:
|
||||
metadata.update(proc.metadata)
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
cache_policy = proc.cache_policy or cache_policy
|
||||
if cache_policy:
|
||||
args_key = cache_policy.key_func(packet.arg)
|
||||
cache_key = CacheKey(
|
||||
(
|
||||
CACHE_NS_WRITES,
|
||||
(identifier(proc) or "__dynamic__"),
|
||||
packet.node,
|
||||
),
|
||||
xxh3_128_hexdigest(
|
||||
args_key.encode() if isinstance(args_key, str) else args_key,
|
||||
),
|
||||
cache_policy.ttl,
|
||||
)
|
||||
else:
|
||||
cache_key = None
|
||||
scratchpad = _scratchpad(
|
||||
config[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||
pending_writes,
|
||||
task_id,
|
||||
xxh3_128_hexdigest(task_checkpoint_ns.encode()),
|
||||
config[CONF].get(CONFIG_KEY_RESUME_MAP),
|
||||
step,
|
||||
stop,
|
||||
)
|
||||
runtime = cast(Runtime, configurable.get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME))
|
||||
runtime = runtime.override(
|
||||
store=store, previous=checkpoint["channel_values"].get(PREVIOUS, None)
|
||||
)
|
||||
additional_config: RunnableConfig = {
|
||||
"metadata": metadata,
|
||||
"tags": proc.tags,
|
||||
}
|
||||
return PregelExecutableTask(
|
||||
packet.node,
|
||||
packet.arg,
|
||||
proc_node,
|
||||
writes,
|
||||
patch_config(
|
||||
merge_configs(config, additional_config),
|
||||
run_name=packet.node,
|
||||
callbacks=(
|
||||
manager.get_child(f"graph:step:{step}") if manager else None
|
||||
),
|
||||
configurable={
|
||||
CONFIG_KEY_TASK_ID: task_id,
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: writes.extend,
|
||||
CONFIG_KEY_READ: partial(
|
||||
local_read,
|
||||
scratchpad,
|
||||
channels,
|
||||
managed,
|
||||
PregelTaskWrites(
|
||||
translated_task_path, packet.node, writes, triggers
|
||||
),
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINTER: (
|
||||
checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER)
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINT_MAP: {
|
||||
**configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {}),
|
||||
parent_ns: checkpoint["id"],
|
||||
},
|
||||
CONFIG_KEY_CHECKPOINT_ID: None,
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
|
||||
CONFIG_KEY_SCRATCHPAD: scratchpad,
|
||||
CONFIG_KEY_RUNTIME: runtime,
|
||||
},
|
||||
),
|
||||
triggers,
|
||||
proc.retry_policy or retry_policy,
|
||||
cache_key,
|
||||
task_id,
|
||||
translated_task_path,
|
||||
writers=proc.flat_writers,
|
||||
subgraphs=proc.subgraphs,
|
||||
)
|
||||
else:
|
||||
return PregelTask(task_id, packet.node, translated_task_path)
|
||||
|
||||
|
||||
def checkpoint_null_version(
|
||||
checkpoint: Checkpoint,
|
||||
) -> V | None:
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph"
|
||||
version = "1.0.2"
|
||||
version = "1.0.3"
|
||||
description = "Building stateful, multi-actor applications with LLMs"
|
||||
authors = []
|
||||
requires-python = ">=3.10"
|
||||
|
||||
Generated
+3
-3
@@ -1,5 +1,5 @@
|
||||
version = 1
|
||||
revision = 3
|
||||
revision = 2
|
||||
requires-python = ">=3.10"
|
||||
resolution-markers = [
|
||||
"python_full_version >= '3.14'",
|
||||
@@ -1345,7 +1345,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "1.0.2"
|
||||
version = "1.0.3"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -1710,7 +1710,7 @@ test = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-prebuilt"
|
||||
version = "1.0.2"
|
||||
version = "1.0.5"
|
||||
source = { editable = "../prebuilt" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
{
|
||||
"permissions": {
|
||||
"allow": [
|
||||
"Bash(make test:*)",
|
||||
"Bash(uv run pytest:*)",
|
||||
"Bash(LANGGRAPH_TEST_FAST=0 make start-services:*)",
|
||||
"Bash(LANGGRAPH_TEST_FAST=0 uv run:*)",
|
||||
"Bash(EXIT_CODE=$?)",
|
||||
"Bash(make stop-services:*)",
|
||||
"Bash(exit $EXIT_CODE)",
|
||||
"Read(//Users/sydney_runkle/oss/langgraph/**)",
|
||||
"Bash(python3:*)",
|
||||
"Bash(find:*)",
|
||||
"Bash(python -m pytest:*)",
|
||||
"Bash(python:*)",
|
||||
"Read(//tmp/**)"
|
||||
],
|
||||
"deny": [],
|
||||
"ask": []
|
||||
}
|
||||
}
|
||||
@@ -63,7 +63,7 @@ class AgentState(TypedDict):
|
||||
|
||||
|
||||
@deprecated(
|
||||
"AgentStatePydantic has been moved to `langchain.agents`. Please update your import to `from langchain.agents import AgentStatePydantic`.",
|
||||
"AgentStatePydantic has been deprecated in favor of AgentState in `langchain.agents`.",
|
||||
category=LangGraphDeprecatedSinceV10,
|
||||
)
|
||||
class AgentStatePydantic(BaseModel):
|
||||
@@ -82,7 +82,7 @@ with warnings.catch_warnings():
|
||||
)
|
||||
|
||||
@deprecated(
|
||||
"AgentStateWithStructuredResponse has been moved to `langchain.agents`. Please update your import to `from langchain.agents import AgentStateWithStructuredResponse`.",
|
||||
"AgentStateWithStructuredResponse has been deprecated in favor of AgentState in `langchain.agents`.",
|
||||
category=LangGraphDeprecatedSinceV10,
|
||||
)
|
||||
class AgentStateWithStructuredResponse(AgentState):
|
||||
@@ -95,11 +95,11 @@ with warnings.catch_warnings():
|
||||
warnings.filterwarnings(
|
||||
"ignore",
|
||||
category=LangGraphDeprecatedSinceV10,
|
||||
message="AgentStatePydantic has been moved to `langchain.agents`.*",
|
||||
message="AgentStatePydantic has been deprecated in favor of AgentState in `langchain.agents`.",
|
||||
)
|
||||
|
||||
@deprecated(
|
||||
"AgentStateWithStructuredResponsePydantic has deprecated. The new `langchain.agents.AgentState` contains the structured response by default.",
|
||||
"AgentStateWithStructuredResponsePydantic has been deprecated in favor of AgentState in `langchain.agents`.",
|
||||
category=LangGraphDeprecatedSinceV10,
|
||||
)
|
||||
class AgentStateWithStructuredResponsePydantic(AgentStatePydantic):
|
||||
|
||||
@@ -142,6 +142,25 @@ class ToolCallRequest:
|
||||
state: Any
|
||||
runtime: ToolRuntime
|
||||
|
||||
def __setattr__(self, name: str, value: Any) -> None:
|
||||
"""Raise deprecation warning when setting attributes directly.
|
||||
|
||||
Direct attribute assignment is deprecated. Use the `override()` method instead.
|
||||
"""
|
||||
import warnings
|
||||
|
||||
# Allow setting attributes during initialization
|
||||
if not hasattr(self, "__dataclass_fields__") or not hasattr(self, name):
|
||||
object.__setattr__(self, name, value)
|
||||
else:
|
||||
warnings.warn(
|
||||
f"Setting attribute '{name}' on ToolCallRequest is deprecated. "
|
||||
"Use the override() method instead to create a new instance with modified values.",
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
object.__setattr__(self, name, value)
|
||||
|
||||
def override(
|
||||
self, **overrides: Unpack[_ToolCallRequestOverrides]
|
||||
) -> ToolCallRequest:
|
||||
@@ -202,8 +221,9 @@ Examples:
|
||||
|
||||
```python
|
||||
def handler(request, execute):
|
||||
request.tool_call["args"]["value"] *= 2
|
||||
return execute(request)
|
||||
modified_call = {**request.tool_call, "args": {**request.tool_call["args"], "value": request.tool_call["args"]["value"] * 2}}
|
||||
modified_request = request.override(tool_call=modified_call)
|
||||
return execute(modified_request)
|
||||
```
|
||||
|
||||
Retry on error (execute multiple times):
|
||||
@@ -479,9 +499,7 @@ def _infer_handled_types(handler: Callable[..., str]) -> tuple[type[Exception],
|
||||
|
||||
def _filter_validation_errors(
|
||||
validation_error: ValidationError,
|
||||
tool_to_state_args: dict[str, str | None],
|
||||
tool_to_store_arg: str | None,
|
||||
tool_to_runtime_arg: str | None,
|
||||
injected_args: _InjectedArgs | None,
|
||||
) -> list[ErrorDetails]:
|
||||
"""Filter validation errors to only include LLM-controlled arguments.
|
||||
|
||||
@@ -496,25 +514,28 @@ def _filter_validation_errors(
|
||||
|
||||
Args:
|
||||
validation_error: The Pydantic ValidationError raised during tool invocation.
|
||||
tool_to_state_args: Mapping of state argument names to state field names.
|
||||
tool_to_store_arg: Name of the store argument, if any.
|
||||
tool_to_runtime_arg: Name of the runtime argument, if any.
|
||||
injected_args: The _InjectedArgs structure containing all injected arguments,
|
||||
or None if there are no injected arguments.
|
||||
|
||||
Returns:
|
||||
List of ErrorDetails containing only errors for LLM-controlled arguments,
|
||||
with system-injected argument values removed from the input field.
|
||||
"""
|
||||
injected_args = set(tool_to_state_args.keys())
|
||||
if tool_to_store_arg:
|
||||
injected_args.add(tool_to_store_arg)
|
||||
if tool_to_runtime_arg:
|
||||
injected_args.add(tool_to_runtime_arg)
|
||||
# Collect all injected argument names
|
||||
injected_arg_names: set[str] = set()
|
||||
if injected_args:
|
||||
if injected_args.state:
|
||||
injected_arg_names.update(injected_args.state.keys())
|
||||
if injected_args.store:
|
||||
injected_arg_names.add(injected_args.store)
|
||||
if injected_args.runtime:
|
||||
injected_arg_names.add(injected_args.runtime)
|
||||
|
||||
filtered_errors: list[ErrorDetails] = []
|
||||
for error in validation_error.errors():
|
||||
# Check if error location contains any injected argument
|
||||
# error['loc'] is a tuple like ('field_name',) or ('field_name', 'nested_field')
|
||||
if error["loc"] and error["loc"][0] not in injected_args:
|
||||
if error["loc"] and error["loc"][0] not in injected_arg_names:
|
||||
# Create a copy of the error dict to avoid mutating the original
|
||||
error_copy: dict[str, Any] = {**error}
|
||||
|
||||
@@ -522,7 +543,7 @@ def _filter_validation_errors(
|
||||
if isinstance(error_copy.get("input"), dict):
|
||||
input_dict = error_copy["input"]
|
||||
input_copy = {
|
||||
k: v for k, v in input_dict.items() if k not in injected_args
|
||||
k: v for k, v in input_dict.items() if k not in injected_arg_names
|
||||
}
|
||||
error_copy["input"] = input_copy
|
||||
|
||||
@@ -532,6 +553,60 @@ def _filter_validation_errors(
|
||||
return filtered_errors
|
||||
|
||||
|
||||
@dataclass
|
||||
class _InjectedArgs:
|
||||
"""Internal structure for tracking injected arguments for a tool.
|
||||
|
||||
This data structure is built once during ToolNode initialization by analyzing
|
||||
the tool's signature and args schema, then reused during execution for efficient
|
||||
injection without repeated reflection.
|
||||
|
||||
The structure maps from tool parameter names to their injection sources, enabling
|
||||
the ToolNode to know exactly which arguments need to be injected and where to
|
||||
get their values from.
|
||||
|
||||
Attributes:
|
||||
state: Mapping from tool parameter names to state field names for injection.
|
||||
Keys are tool parameter names, values are either:
|
||||
- str: Name of the state field to extract and inject
|
||||
- None: Inject the entire state object
|
||||
Empty dict if no state injection is needed.
|
||||
store: Name of the tool parameter where the store should be injected,
|
||||
or None if no store injection is needed.
|
||||
runtime: Name of the tool parameter where the runtime should be injected,
|
||||
or None if no runtime injection is needed.
|
||||
|
||||
Example:
|
||||
For a tool with signature:
|
||||
```python
|
||||
def my_tool(
|
||||
x: int,
|
||||
messages: Annotated[list, InjectedState("messages")],
|
||||
full_state: Annotated[dict, InjectedState()],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
runtime: ToolRuntime,
|
||||
) -> str:
|
||||
...
|
||||
```
|
||||
|
||||
The resulting `_InjectedArgs` would be:
|
||||
```python
|
||||
_InjectedArgs(
|
||||
state={
|
||||
"messages": "messages", # Extract state["messages"]
|
||||
"full_state": None, # Inject entire state
|
||||
},
|
||||
store="store", # Inject into "store" parameter
|
||||
runtime="runtime", # Inject into "runtime" parameter
|
||||
)
|
||||
```
|
||||
"""
|
||||
|
||||
state: dict[str, str | None]
|
||||
store: str | None
|
||||
runtime: str | None
|
||||
|
||||
|
||||
class ToolNode(RunnableCallable):
|
||||
"""A node for executing tools in LangGraph workflows.
|
||||
|
||||
@@ -676,9 +751,7 @@ class ToolNode(RunnableCallable):
|
||||
"""
|
||||
super().__init__(self._func, self._afunc, name=name, tags=tags, trace=False)
|
||||
self._tools_by_name: dict[str, BaseTool] = {}
|
||||
self._tool_to_state_args: dict[str, dict[str, str | None]] = {}
|
||||
self._tool_to_store_arg: dict[str, str | None] = {}
|
||||
self._tool_to_runtime_arg: dict[str, str | None] = {}
|
||||
self._injected_args: dict[str, _InjectedArgs] = {}
|
||||
self._handle_tool_errors = handle_tool_errors
|
||||
self._messages_key = messages_key
|
||||
self._wrap_tool_call = wrap_tool_call
|
||||
@@ -689,9 +762,8 @@ class ToolNode(RunnableCallable):
|
||||
else:
|
||||
tool_ = tool
|
||||
self._tools_by_name[tool_.name] = tool_
|
||||
self._tool_to_state_args[tool_.name] = _get_state_args(tool_)
|
||||
self._tool_to_store_arg[tool_.name] = _get_store_arg(tool_)
|
||||
self._tool_to_runtime_arg[tool_.name] = _get_runtime_arg(tool_)
|
||||
# Build injected args mapping once during initialization in a single pass
|
||||
self._injected_args[tool_.name] = _get_all_injected_args(tool_)
|
||||
|
||||
@property
|
||||
def tools_by_name(self) -> dict[str, BaseTool]:
|
||||
@@ -844,12 +916,8 @@ class ToolNode(RunnableCallable):
|
||||
response = tool.invoke(call_args, config)
|
||||
except ValidationError as exc:
|
||||
# Filter out errors for injected arguments
|
||||
filtered_errors = _filter_validation_errors(
|
||||
exc,
|
||||
self._tool_to_state_args.get(call["name"], {}),
|
||||
self._tool_to_store_arg.get(call["name"]),
|
||||
self._tool_to_runtime_arg.get(call["name"]),
|
||||
)
|
||||
injected = self._injected_args.get(call["name"])
|
||||
filtered_errors = _filter_validation_errors(exc, injected)
|
||||
# Use original call["args"] without injected values for error reporting
|
||||
raise ToolInvocationError(
|
||||
call["name"], exc, call["args"], filtered_errors
|
||||
@@ -1001,12 +1069,8 @@ class ToolNode(RunnableCallable):
|
||||
response = await tool.ainvoke(call_args, config)
|
||||
except ValidationError as exc:
|
||||
# Filter out errors for injected arguments
|
||||
filtered_errors = _filter_validation_errors(
|
||||
exc,
|
||||
self._tool_to_state_args.get(call["name"], {}),
|
||||
self._tool_to_store_arg.get(call["name"]),
|
||||
self._tool_to_runtime_arg.get(call["name"]),
|
||||
)
|
||||
injected = self._injected_args.get(call["name"])
|
||||
filtered_errors = _filter_validation_errors(exc, injected)
|
||||
# Use original call["args"] without injected values for error reporting
|
||||
raise ToolInvocationError(
|
||||
call["name"], exc, call["args"], filtered_errors
|
||||
@@ -1199,86 +1263,6 @@ class ToolNode(RunnableCallable):
|
||||
return input["state"]
|
||||
return input
|
||||
|
||||
def _inject_state(
|
||||
self,
|
||||
tool_call: ToolCall,
|
||||
state: list[AnyMessage] | dict[str, Any] | BaseModel,
|
||||
) -> ToolCall:
|
||||
state_args = self._tool_to_state_args[tool_call["name"]]
|
||||
|
||||
if state_args and isinstance(state, list):
|
||||
required_fields = list(state_args.values())
|
||||
if (
|
||||
len(required_fields) == 1 and required_fields[0] == self._messages_key
|
||||
) or required_fields[0] is None:
|
||||
state = {self._messages_key: state}
|
||||
else:
|
||||
err_msg = (
|
||||
f"Invalid input to ToolNode. Tool {tool_call['name']} requires "
|
||||
f"graph state dict as input."
|
||||
)
|
||||
if any(state_field for state_field in state_args.values()):
|
||||
required_fields_str = ", ".join(f for f in required_fields if f)
|
||||
err_msg += f" State should contain fields {required_fields_str}."
|
||||
raise ValueError(err_msg)
|
||||
|
||||
if isinstance(state, dict):
|
||||
tool_state_args = {
|
||||
tool_arg: state[state_field] if state_field else state
|
||||
for tool_arg, state_field in state_args.items()
|
||||
}
|
||||
else:
|
||||
tool_state_args = {
|
||||
tool_arg: getattr(state, state_field) if state_field else state
|
||||
for tool_arg, state_field in state_args.items()
|
||||
}
|
||||
|
||||
tool_call["args"] = {
|
||||
**tool_call["args"],
|
||||
**tool_state_args,
|
||||
}
|
||||
return tool_call
|
||||
|
||||
def _inject_store(self, tool_call: ToolCall, store: BaseStore | None) -> ToolCall:
|
||||
store_arg = self._tool_to_store_arg[tool_call["name"]]
|
||||
if not store_arg:
|
||||
return tool_call
|
||||
|
||||
if store is None:
|
||||
msg = (
|
||||
"Cannot inject store into tools with InjectedStore annotations - "
|
||||
"please compile your graph with a store."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
|
||||
tool_call["args"] = {
|
||||
**tool_call["args"],
|
||||
store_arg: store,
|
||||
}
|
||||
return tool_call
|
||||
|
||||
def _inject_runtime(
|
||||
self, tool_call: ToolCall, tool_runtime: ToolRuntime
|
||||
) -> ToolCall:
|
||||
"""Inject ToolRuntime into tool call arguments.
|
||||
|
||||
Args:
|
||||
tool_call: The tool call to inject runtime into.
|
||||
tool_runtime: The ToolRuntime instance to inject.
|
||||
|
||||
Returns:
|
||||
The tool call with runtime injected if needed.
|
||||
"""
|
||||
runtime_arg = self._tool_to_runtime_arg.get(tool_call["name"])
|
||||
if not runtime_arg:
|
||||
return tool_call
|
||||
|
||||
tool_call["args"] = {
|
||||
**tool_call["args"],
|
||||
runtime_arg: tool_runtime,
|
||||
}
|
||||
return tool_call
|
||||
|
||||
def _inject_tool_args(
|
||||
self,
|
||||
tool_call: ToolCall,
|
||||
@@ -1317,12 +1301,64 @@ class ToolNode(RunnableCallable):
|
||||
if tool_call["name"] not in self.tools_by_name:
|
||||
return tool_call
|
||||
|
||||
injected = self._injected_args.get(tool_call["name"])
|
||||
if not injected:
|
||||
return tool_call
|
||||
|
||||
tool_call_copy: ToolCall = copy(tool_call)
|
||||
tool_call_with_state = self._inject_state(tool_call_copy, tool_runtime.state)
|
||||
tool_call_with_store = self._inject_store(
|
||||
tool_call_with_state, tool_runtime.store
|
||||
)
|
||||
return self._inject_runtime(tool_call_with_store, tool_runtime)
|
||||
injected_args = {}
|
||||
|
||||
# Inject state
|
||||
if injected.state:
|
||||
state = tool_runtime.state
|
||||
# Handle list state by converting to dict
|
||||
if isinstance(state, list):
|
||||
required_fields = list(injected.state.values())
|
||||
if (
|
||||
len(required_fields) == 1
|
||||
and required_fields[0] == self._messages_key
|
||||
) or required_fields[0] is None:
|
||||
state = {self._messages_key: state}
|
||||
else:
|
||||
err_msg = (
|
||||
f"Invalid input to ToolNode. Tool {tool_call['name']} requires "
|
||||
f"graph state dict as input."
|
||||
)
|
||||
if any(state_field for state_field in injected.state.values()):
|
||||
required_fields_str = ", ".join(f for f in required_fields if f)
|
||||
err_msg += (
|
||||
f" State should contain fields {required_fields_str}."
|
||||
)
|
||||
raise ValueError(err_msg)
|
||||
|
||||
# Extract state values
|
||||
if isinstance(state, dict):
|
||||
for tool_arg, state_field in injected.state.items():
|
||||
injected_args[tool_arg] = (
|
||||
state[state_field] if state_field else state
|
||||
)
|
||||
else:
|
||||
for tool_arg, state_field in injected.state.items():
|
||||
injected_args[tool_arg] = (
|
||||
getattr(state, state_field) if state_field else state
|
||||
)
|
||||
|
||||
# Inject store
|
||||
if injected.store:
|
||||
if tool_runtime.store is None:
|
||||
msg = (
|
||||
"Cannot inject store into tools with InjectedStore annotations - "
|
||||
"please compile your graph with a store."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
injected_args[injected.store] = tool_runtime.store
|
||||
|
||||
# Inject runtime
|
||||
if injected.runtime:
|
||||
injected_args[injected.runtime] = tool_runtime
|
||||
|
||||
tool_call_copy["args"] = {**tool_call_copy["args"], **injected_args}
|
||||
return tool_call_copy
|
||||
|
||||
def _validate_tool_command(
|
||||
self,
|
||||
@@ -1715,120 +1751,86 @@ def _is_injection(
|
||||
return False
|
||||
|
||||
|
||||
def _get_state_args(tool: BaseTool) -> dict[str, str | None]:
|
||||
"""Extract state injection mappings from tool annotations.
|
||||
|
||||
This function analyzes a tool's input schema to identify arguments that should
|
||||
be injected with graph state. It processes InjectedState annotations to build
|
||||
a mapping of tool argument names to state field names.
|
||||
def _get_injection_from_type(
|
||||
type_: Any, injection_type: type[InjectedState | InjectedStore | ToolRuntime]
|
||||
) -> Any | None:
|
||||
"""Extract injection instance from a type annotation.
|
||||
|
||||
Args:
|
||||
tool: The tool to analyze for state injection requirements.
|
||||
type_: The type annotation to check.
|
||||
injection_type: The injection type to look for.
|
||||
|
||||
Returns:
|
||||
A dictionary mapping tool argument names to state field names. If a field
|
||||
name is None, the entire state should be injected for that argument.
|
||||
The injection instance if found, True if injection marker found without instance, None otherwise.
|
||||
"""
|
||||
full_schema = tool.get_input_schema()
|
||||
tool_args_to_state_fields: dict = {}
|
||||
type_args = get_args(type_)
|
||||
matches = [arg for arg in type_args if _is_injection(arg, injection_type)]
|
||||
|
||||
for name, type_ in get_all_basemodel_annotations(full_schema).items():
|
||||
injections = [
|
||||
type_arg
|
||||
for type_arg in get_args(type_)
|
||||
if _is_injection(type_arg, InjectedState)
|
||||
]
|
||||
if len(injections) > 1:
|
||||
msg = (
|
||||
"A tool argument should not be annotated with InjectedState more than "
|
||||
f"once. Received arg {name} with annotations {injections}."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
if len(injections) == 1:
|
||||
injection = injections[0]
|
||||
if isinstance(injection, InjectedState) and injection.field:
|
||||
tool_args_to_state_fields[name] = injection.field
|
||||
else:
|
||||
tool_args_to_state_fields[name] = None
|
||||
else:
|
||||
pass
|
||||
return tool_args_to_state_fields
|
||||
if len(matches) > 1:
|
||||
msg = (
|
||||
f"A tool argument should not be annotated with {injection_type.__name__} "
|
||||
f"more than once. Found: {matches}"
|
||||
)
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
def _get_store_arg(tool: BaseTool) -> str | None:
|
||||
"""Extract store injection argument from tool annotations.
|
||||
|
||||
This function analyzes a tool's input schema to identify the argument that
|
||||
should be injected with the graph store. Only one store argument is supported
|
||||
per tool.
|
||||
|
||||
Args:
|
||||
tool: The tool to analyze for store injection requirements.
|
||||
|
||||
Returns:
|
||||
The name of the argument that should receive the store injection, or None
|
||||
if no store injection is required.
|
||||
|
||||
Raises:
|
||||
ValueError: If a tool argument has multiple InjectedStore annotations.
|
||||
"""
|
||||
full_schema = tool.get_input_schema()
|
||||
for name, type_ in get_all_basemodel_annotations(full_schema).items():
|
||||
injections = [
|
||||
type_arg
|
||||
for type_arg in get_args(type_)
|
||||
if _is_injection(type_arg, InjectedStore)
|
||||
]
|
||||
if len(injections) > 1:
|
||||
msg = (
|
||||
"A tool argument should not be annotated with InjectedStore more than "
|
||||
f"once. Received arg {name} with annotations {injections}."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
if len(injections) == 1:
|
||||
return name
|
||||
if len(matches) == 1:
|
||||
return matches[0]
|
||||
elif _is_injection(type_, injection_type):
|
||||
return True
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _get_runtime_arg(tool: BaseTool) -> str | None:
|
||||
"""Extract runtime injection argument from tool annotations.
|
||||
def _get_all_injected_args(tool: BaseTool) -> _InjectedArgs:
|
||||
"""Extract all injected arguments from tool in a single pass.
|
||||
|
||||
This function analyzes a tool's input schema to identify the argument that
|
||||
should be injected with the ToolRuntime instance. Only one runtime argument
|
||||
is supported per tool.
|
||||
This function analyzes both the tool's input schema and function signature
|
||||
to identify all arguments that should be injected (state, store, runtime).
|
||||
|
||||
Args:
|
||||
tool: The tool to analyze for runtime injection requirements.
|
||||
tool: The tool to analyze for injection requirements.
|
||||
|
||||
Returns:
|
||||
The name of the argument that should receive the runtime injection, or None
|
||||
if no runtime injection is required.
|
||||
|
||||
Raises:
|
||||
ValueError: If a tool argument has multiple ToolRuntime annotations.
|
||||
_InjectedArgs structure containing all detected injections.
|
||||
"""
|
||||
# Get annotations from both schema and function signature
|
||||
full_schema = tool.get_input_schema()
|
||||
for name, type_ in get_all_basemodel_annotations(full_schema).items():
|
||||
# Check if the parameter name is "runtime" (regardless of type)
|
||||
schema_annotations = get_all_basemodel_annotations(full_schema)
|
||||
|
||||
func = getattr(tool, "func", None) or getattr(tool, "coroutine", None)
|
||||
func_annotations = get_type_hints(func, include_extras=True) if func else {}
|
||||
|
||||
# Combine both annotation sources, preferring schema annotations
|
||||
# In the future, we might want to add more restrictions here...
|
||||
all_annotations = {**func_annotations, **schema_annotations}
|
||||
|
||||
# Track injected args
|
||||
state_args: dict[str, str | None] = {}
|
||||
store_arg: str | None = None
|
||||
runtime_arg: str | None = None
|
||||
|
||||
for name, type_ in all_annotations.items():
|
||||
# Check for runtime (special case: parameter named "runtime")
|
||||
if name == "runtime":
|
||||
return name
|
||||
# Check if the type itself is ToolRuntime (direct usage)
|
||||
if _is_injection(type_, ToolRuntime):
|
||||
return name
|
||||
# Check if ToolRuntime is in Annotated args
|
||||
injections = [
|
||||
type_arg
|
||||
for type_arg in get_args(type_)
|
||||
if _is_injection(type_arg, ToolRuntime)
|
||||
]
|
||||
if len(injections) > 1:
|
||||
msg = (
|
||||
"A tool argument should not be annotated with ToolRuntime more than "
|
||||
f"once. Received arg {name} with annotations {injections}."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
if len(injections) == 1:
|
||||
return name
|
||||
runtime_arg = name
|
||||
|
||||
return None
|
||||
# Check for InjectedState
|
||||
if state_inj := _get_injection_from_type(type_, InjectedState):
|
||||
if isinstance(state_inj, InjectedState) and state_inj.field:
|
||||
state_args[name] = state_inj.field
|
||||
else:
|
||||
state_args[name] = None
|
||||
|
||||
# Check for InjectedStore
|
||||
if _get_injection_from_type(type_, InjectedStore):
|
||||
store_arg = name
|
||||
|
||||
# Check for ToolRuntime
|
||||
if _get_injection_from_type(type_, ToolRuntime):
|
||||
runtime_arg = name
|
||||
|
||||
return _InjectedArgs(
|
||||
state=state_args,
|
||||
store=store_arg,
|
||||
runtime=runtime_arg,
|
||||
)
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph-prebuilt"
|
||||
version = "1.0.2"
|
||||
version = "1.0.5"
|
||||
description = "Library with high-level APIs for creating and executing LangGraph agents and tools."
|
||||
authors = []
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -130,11 +130,17 @@ def test_modify_arguments() -> None:
|
||||
execute: Callable[[ToolCallRequest], ToolMessage | Command],
|
||||
) -> ToolMessage | Command:
|
||||
"""Handler that doubles the input arguments."""
|
||||
# Modify the arguments
|
||||
request.tool_call["args"]["a"] *= 2
|
||||
request.tool_call["args"]["b"] *= 2
|
||||
|
||||
return execute(request)
|
||||
# Modify the arguments using override method
|
||||
modified_call = {
|
||||
**request.tool_call,
|
||||
"args": {
|
||||
**request.tool_call["args"],
|
||||
"a": request.tool_call["args"]["a"] * 2,
|
||||
"b": request.tool_call["args"]["b"] * 2,
|
||||
},
|
||||
}
|
||||
modified_request = request.override(tool_call=modified_call)
|
||||
return execute(modified_request)
|
||||
|
||||
tool_node = ToolNode([add], wrap_tool_call=modify_args_handler)
|
||||
|
||||
@@ -362,10 +368,17 @@ async def test_handler_with_async_execution() -> None:
|
||||
execute: Callable[[ToolCallRequest], ToolMessage | Command],
|
||||
) -> ToolMessage | Command:
|
||||
"""Handler that modifies arguments."""
|
||||
# Add 10 to both arguments
|
||||
request.tool_call["args"]["a"] += 10
|
||||
request.tool_call["args"]["b"] += 10
|
||||
return execute(request)
|
||||
# Add 10 to both arguments using override method
|
||||
modified_call = {
|
||||
**request.tool_call,
|
||||
"args": {
|
||||
**request.tool_call["args"],
|
||||
"a": request.tool_call["args"]["a"] + 10,
|
||||
"b": request.tool_call["args"]["b"] + 10,
|
||||
},
|
||||
}
|
||||
modified_request = request.override(tool_call=modified_call)
|
||||
return execute(modified_request)
|
||||
|
||||
tool_node = ToolNode([async_add], wrap_tool_call=modifying_handler)
|
||||
|
||||
@@ -1305,3 +1318,64 @@ async def test_state_extraction_with_tool_call_with_context_async() -> None:
|
||||
assert state_seen[0] == actual_state
|
||||
assert "__type" not in state_seen[0]
|
||||
assert "tool_call" not in state_seen[0]
|
||||
|
||||
|
||||
def test_tool_call_request_is_frozen() -> None:
|
||||
"""Test that ToolCallRequest raises deprecation warnings on direct attribute reassignment."""
|
||||
tool_call: ToolCall = {"name": "add", "args": {"a": 1, "b": 2}, "id": "call_1"}
|
||||
state: dict = {"messages": []}
|
||||
runtime = None
|
||||
|
||||
request = ToolCallRequest(
|
||||
tool_call=tool_call, tool=add, state=state, runtime=runtime
|
||||
) # type: ignore[arg-type]
|
||||
|
||||
# Test that direct attribute reassignment raises DeprecationWarning
|
||||
with pytest.warns(
|
||||
DeprecationWarning,
|
||||
match="Setting attribute 'tool_call' on ToolCallRequest is deprecated",
|
||||
):
|
||||
request.tool_call = {"name": "other", "args": {}, "id": "call_2"} # type: ignore[misc]
|
||||
|
||||
with pytest.warns(
|
||||
DeprecationWarning,
|
||||
match="Setting attribute 'tool' on ToolCallRequest is deprecated",
|
||||
):
|
||||
request.tool = None # type: ignore[misc]
|
||||
|
||||
with pytest.warns(
|
||||
DeprecationWarning,
|
||||
match="Setting attribute 'state' on ToolCallRequest is deprecated",
|
||||
):
|
||||
request.state = {} # type: ignore[misc]
|
||||
|
||||
with pytest.warns(
|
||||
DeprecationWarning,
|
||||
match="Setting attribute 'runtime' on ToolCallRequest is deprecated",
|
||||
):
|
||||
request.runtime = None # type: ignore[misc]
|
||||
|
||||
# Test that override method works correctly
|
||||
new_tool_call: ToolCall = {
|
||||
"name": "multiply",
|
||||
"args": {"x": 5, "y": 10},
|
||||
"id": "call_3",
|
||||
}
|
||||
|
||||
# Original request should be unchanged (note: it was modified by the warnings tests above)
|
||||
# So we create a fresh request to test override properly
|
||||
fresh_request = ToolCallRequest(
|
||||
tool_call=tool_call, tool=add, state=state, runtime=runtime
|
||||
) # type: ignore[arg-type]
|
||||
fresh_new_request = fresh_request.override(tool_call=new_tool_call)
|
||||
|
||||
# Original request should be unchanged
|
||||
assert fresh_request.tool_call == tool_call
|
||||
assert fresh_request.tool_call["name"] == "add"
|
||||
|
||||
# New request should have the updated tool_call
|
||||
assert fresh_new_request.tool_call == new_tool_call
|
||||
assert fresh_new_request.tool_call["name"] == "multiply"
|
||||
assert fresh_new_request.tool == add # Other fields should remain the same
|
||||
assert fresh_new_request.state == state
|
||||
assert fresh_new_request.runtime is None
|
||||
|
||||
@@ -53,7 +53,6 @@ from langgraph.prebuilt.chat_agent_executor import (
|
||||
from langgraph.prebuilt.tool_node import (
|
||||
InjectedState,
|
||||
InjectedStore,
|
||||
_get_state_args,
|
||||
_infer_handled_types,
|
||||
)
|
||||
from tests.any_str import AnyStr
|
||||
@@ -1084,21 +1083,6 @@ async def test_return_direct(version: str) -> None:
|
||||
]
|
||||
|
||||
|
||||
def test__get_state_args() -> None:
|
||||
class Schema1(BaseModel):
|
||||
a: Annotated[str, InjectedState]
|
||||
|
||||
class Schema2(Schema1):
|
||||
b: Annotated[int, InjectedState("bar")]
|
||||
|
||||
@dec_tool(args_schema=Schema2)
|
||||
def foo(a: str, b: int) -> float:
|
||||
"""return"""
|
||||
return 0.0
|
||||
|
||||
assert _get_state_args(foo) == {"a": None, "b": "bar"}
|
||||
|
||||
|
||||
def test_inspect_react() -> None:
|
||||
model = FakeToolCallingModel(tool_calls=[])
|
||||
agent = create_react_agent(model, [])
|
||||
|
||||
@@ -42,6 +42,7 @@ from langgraph.prebuilt import (
|
||||
from langgraph.prebuilt.tool_node import (
|
||||
TOOL_CALL_ERROR_TEMPLATE,
|
||||
ToolInvocationError,
|
||||
ToolRuntime,
|
||||
tools_condition,
|
||||
)
|
||||
|
||||
@@ -1610,3 +1611,252 @@ def test_tool_node_stream_writer() -> None:
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def test_tool_call_request_setattr_deprecation_warning():
|
||||
"""Test that ToolCallRequest raises a deprecation warning on direct attribute modification."""
|
||||
import warnings
|
||||
|
||||
from langgraph.prebuilt.tool_node import ToolCallRequest
|
||||
|
||||
# Create a mock ToolCall
|
||||
tool_call = {"name": "test", "args": {"a": 1}, "id": "call_1", "type": "tool_call"}
|
||||
|
||||
# Create a ToolCallRequest
|
||||
request = ToolCallRequest(
|
||||
tool_call=tool_call,
|
||||
tool=None,
|
||||
state={"messages": []},
|
||||
runtime=None,
|
||||
)
|
||||
|
||||
# Test 1: Direct attribute assignment should raise deprecation warning but still work
|
||||
with pytest.warns(DeprecationWarning, match="deprecated.*override"):
|
||||
request.tool_call = {"name": "other", "args": {}, "id": "call_2"}
|
||||
|
||||
# Verify the attribute was actually modified
|
||||
assert request.tool_call == {"name": "other", "args": {}, "id": "call_2"}
|
||||
|
||||
# Reset for further tests
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore")
|
||||
request.tool_call = tool_call
|
||||
|
||||
# Test 2: override method should work without warnings
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
warnings.simplefilter("always")
|
||||
new_tool_call = {
|
||||
"name": "new_tool",
|
||||
"args": {"b": 2},
|
||||
"id": "call_3",
|
||||
"type": "tool_call",
|
||||
}
|
||||
new_request = request.override(tool_call=new_tool_call)
|
||||
|
||||
# Verify no warning was raised
|
||||
assert len(w) == 0
|
||||
|
||||
# Verify original is unchanged
|
||||
assert request.tool_call == tool_call
|
||||
|
||||
# Verify new request has updated values
|
||||
assert new_request.tool_call == new_tool_call
|
||||
|
||||
# Test 3: Initialization should not trigger warning
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
warnings.simplefilter("always")
|
||||
ToolCallRequest(
|
||||
tool_call=tool_call,
|
||||
tool=None,
|
||||
state={"messages": []},
|
||||
runtime=None,
|
||||
)
|
||||
# Verify no warning was raised during initialization
|
||||
assert len(w) == 0
|
||||
|
||||
|
||||
async def test_tool_node_inject_async_all_types_signature_only() -> None:
|
||||
"""Test all injection types without @tool decorator."""
|
||||
store = InMemoryStore()
|
||||
namespace = ("test",)
|
||||
store.put(namespace, "test_key", {"store_data": "from_store"})
|
||||
|
||||
class TestState(TypedDict):
|
||||
messages: list
|
||||
foo: str
|
||||
bar: int
|
||||
|
||||
async def comprehensive_async_tool(
|
||||
x: int,
|
||||
whole_state: Annotated[TestState, InjectedState],
|
||||
foo_field: Annotated[str, InjectedState("foo")],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
runtime: ToolRuntime,
|
||||
) -> str:
|
||||
"""Async tool that uses all injection types."""
|
||||
bar_from_whole = whole_state["bar"]
|
||||
foo_value = foo_field
|
||||
store_val = store.get(namespace, "test_key").value["store_data"]
|
||||
foo_from_runtime = runtime.state["foo"]
|
||||
tool_call_id = runtime.tool_call_id
|
||||
|
||||
return (
|
||||
f"x={x}, "
|
||||
f"bar_from_whole={bar_from_whole}, "
|
||||
f"foo_field={foo_value}, "
|
||||
f"store={store_val}, "
|
||||
f"foo_from_runtime={foo_from_runtime}, "
|
||||
f"tool_call_id={tool_call_id}"
|
||||
)
|
||||
|
||||
node = ToolNode([comprehensive_async_tool], handle_tool_errors=True)
|
||||
tool_call = {
|
||||
"name": "comprehensive_async_tool",
|
||||
"args": {"x": 42},
|
||||
"id": "test_call_123",
|
||||
"type": "tool_call",
|
||||
}
|
||||
msg = AIMessage("hi?", tool_calls=[tool_call])
|
||||
|
||||
config = _create_config_with_runtime(store=store)
|
||||
result = await node.ainvoke(
|
||||
{"messages": [msg], "foo": "foo_value", "bar": 99}, config=config
|
||||
)
|
||||
|
||||
tool_message = result["messages"][-1]
|
||||
assert tool_message.content == (
|
||||
"x=42, "
|
||||
"bar_from_whole=99, "
|
||||
"foo_field=foo_value, "
|
||||
"store=from_store, "
|
||||
"foo_from_runtime=foo_value, "
|
||||
"tool_call_id=test_call_123"
|
||||
)
|
||||
|
||||
|
||||
async def test_tool_node_inject_async_all_types_with_decorator() -> None:
|
||||
"""Test all injection types with @tool decorator."""
|
||||
store = InMemoryStore()
|
||||
namespace = ("test",)
|
||||
store.put(namespace, "test_key", {"store_data": "from_store"})
|
||||
|
||||
class TestState(TypedDict):
|
||||
messages: list
|
||||
foo: str
|
||||
bar: int
|
||||
|
||||
@dec_tool
|
||||
async def comprehensive_async_tool(
|
||||
x: int,
|
||||
whole_state: Annotated[TestState, InjectedState],
|
||||
foo_field: Annotated[str, InjectedState("foo")],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
runtime: ToolRuntime,
|
||||
) -> str:
|
||||
"""Async tool that uses all injection types."""
|
||||
bar_from_whole = whole_state["bar"]
|
||||
foo_value = foo_field
|
||||
store_val = store.get(namespace, "test_key").value["store_data"]
|
||||
foo_from_runtime = runtime.state["foo"]
|
||||
tool_call_id = runtime.tool_call_id
|
||||
|
||||
return (
|
||||
f"x={x}, "
|
||||
f"bar_from_whole={bar_from_whole}, "
|
||||
f"foo_field={foo_value}, "
|
||||
f"store={store_val}, "
|
||||
f"foo_from_runtime={foo_from_runtime}, "
|
||||
f"tool_call_id={tool_call_id}"
|
||||
)
|
||||
|
||||
node = ToolNode([comprehensive_async_tool], handle_tool_errors=True)
|
||||
tool_call = {
|
||||
"name": "comprehensive_async_tool",
|
||||
"args": {"x": 42},
|
||||
"id": "test_call_456",
|
||||
"type": "tool_call",
|
||||
}
|
||||
msg = AIMessage("hi?", tool_calls=[tool_call])
|
||||
|
||||
config = _create_config_with_runtime(store=store)
|
||||
result = await node.ainvoke(
|
||||
{"messages": [msg], "foo": "foo_value", "bar": 99}, config=config
|
||||
)
|
||||
|
||||
tool_message = result["messages"][-1]
|
||||
assert tool_message.content == (
|
||||
"x=42, "
|
||||
"bar_from_whole=99, "
|
||||
"foo_field=foo_value, "
|
||||
"store=from_store, "
|
||||
"foo_from_runtime=foo_value, "
|
||||
"tool_call_id=test_call_456"
|
||||
)
|
||||
|
||||
|
||||
async def test_tool_node_inject_async_all_types_with_schema() -> None:
|
||||
"""Test all injection types with explicit schema."""
|
||||
store = InMemoryStore()
|
||||
namespace = ("test",)
|
||||
store.put(namespace, "test_key", {"store_data": "from_store"})
|
||||
|
||||
class TestState(TypedDict):
|
||||
messages: list
|
||||
foo: str
|
||||
bar: int
|
||||
|
||||
class ComprehensiveToolSchema(BaseModel):
|
||||
model_config = {"arbitrary_types_allowed": True}
|
||||
x: int
|
||||
whole_state: Annotated[TestState, InjectedState]
|
||||
foo_field: Annotated[str, InjectedState("foo")]
|
||||
store: Annotated[BaseStore, InjectedStore()]
|
||||
runtime: ToolRuntime
|
||||
|
||||
@dec_tool(args_schema=ComprehensiveToolSchema)
|
||||
async def comprehensive_async_tool(
|
||||
x: int,
|
||||
whole_state: Annotated[TestState, InjectedState],
|
||||
foo_field: Annotated[str, InjectedState("foo")],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
runtime: ToolRuntime,
|
||||
) -> str:
|
||||
"""Async tool that uses all injection types."""
|
||||
bar_from_whole = whole_state["bar"]
|
||||
foo_value = foo_field
|
||||
store_val = store.get(namespace, "test_key").value["store_data"]
|
||||
foo_from_runtime = runtime.state["foo"]
|
||||
tool_call_id = runtime.tool_call_id
|
||||
|
||||
return (
|
||||
f"x={x}, "
|
||||
f"bar_from_whole={bar_from_whole}, "
|
||||
f"foo_field={foo_value}, "
|
||||
f"store={store_val}, "
|
||||
f"foo_from_runtime={foo_from_runtime}, "
|
||||
f"tool_call_id={tool_call_id}"
|
||||
)
|
||||
|
||||
node = ToolNode([comprehensive_async_tool], handle_tool_errors=True)
|
||||
tool_call = {
|
||||
"name": "comprehensive_async_tool",
|
||||
"args": {"x": 42},
|
||||
"id": "test_call_789",
|
||||
"type": "tool_call",
|
||||
}
|
||||
msg = AIMessage("hi?", tool_calls=[tool_call])
|
||||
|
||||
config = _create_config_with_runtime(store=store)
|
||||
result = await node.ainvoke(
|
||||
{"messages": [msg], "foo": "foo_value", "bar": 99}, config=config
|
||||
)
|
||||
|
||||
tool_message = result["messages"][-1]
|
||||
assert tool_message.content == (
|
||||
"x=42, "
|
||||
"bar_from_whole=99, "
|
||||
"foo_field=foo_value, "
|
||||
"store=from_store, "
|
||||
"foo_from_runtime=foo_value, "
|
||||
"tool_call_id=test_call_789"
|
||||
)
|
||||
|
||||
Generated
+3
-3
@@ -1,5 +1,5 @@
|
||||
version = 1
|
||||
revision = 3
|
||||
revision = 2
|
||||
requires-python = ">=3.10"
|
||||
|
||||
[[package]]
|
||||
@@ -246,7 +246,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "1.0.2"
|
||||
version = "1.0.3"
|
||||
source = { editable = "../langgraph" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -467,7 +467,7 @@ test = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-prebuilt"
|
||||
version = "1.0.2"
|
||||
version = "1.0.5"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Reference in New Issue
Block a user