mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-02 14:35:18 +02:00
Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9aab70bd73 | ||
|
|
aaea76b475 | ||
|
|
16b363fbb0 | ||
|
|
68a75135b0 | ||
|
|
5c0c0fb186 | ||
|
|
4571b708d9 | ||
|
|
e365b2b8bd | ||
|
|
b5504506a7 | ||
|
|
c6ae8d25b9 | ||
|
|
0bd7dd2c52 | ||
|
|
82978a8dd8 |
Generated
+1
@@ -329,6 +329,7 @@ dev = [
|
||||
{ name = "pytest-asyncio" },
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
]
|
||||
|
||||
|
||||
Generated
+1
@@ -341,6 +341,7 @@ dev = [
|
||||
{ name = "pytest-asyncio" },
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
]
|
||||
|
||||
|
||||
+144
@@ -0,0 +1,144 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langgraph.cache.base import BaseCache, FullKey, Namespace, ValueT
|
||||
from langgraph.checkpoint.serde.base import SerializerProtocol
|
||||
|
||||
|
||||
class RedisCache(BaseCache[ValueT]):
|
||||
"""Redis-based cache implementation with TTL support."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
redis: Any,
|
||||
*,
|
||||
serde: SerializerProtocol | None = None,
|
||||
prefix: str = "langgraph:cache:",
|
||||
) -> None:
|
||||
"""Initialize the cache with a Redis client.
|
||||
|
||||
Args:
|
||||
redis: Redis client instance (sync or async)
|
||||
serde: Serializer to use for values
|
||||
prefix: Key prefix for all cached values
|
||||
"""
|
||||
super().__init__(serde=serde)
|
||||
self.redis = redis
|
||||
self.prefix = prefix
|
||||
|
||||
def _make_key(self, ns: Namespace, key: str) -> str:
|
||||
"""Create a Redis key from namespace and key."""
|
||||
ns_str = ":".join(ns) if ns else ""
|
||||
return f"{self.prefix}{ns_str}:{key}" if ns_str else f"{self.prefix}{key}"
|
||||
|
||||
def _parse_key(self, redis_key: str) -> tuple[Namespace, str]:
|
||||
"""Parse a Redis key back to namespace and key."""
|
||||
if not redis_key.startswith(self.prefix):
|
||||
raise ValueError(
|
||||
f"Key {redis_key} does not start with prefix {self.prefix}"
|
||||
)
|
||||
|
||||
remaining = redis_key[len(self.prefix) :]
|
||||
if ":" in remaining:
|
||||
parts = remaining.split(":")
|
||||
key = parts[-1]
|
||||
ns_parts = parts[:-1]
|
||||
return (tuple(ns_parts), key)
|
||||
else:
|
||||
return (tuple(), remaining)
|
||||
|
||||
def get(self, keys: Sequence[FullKey]) -> dict[FullKey, ValueT]:
|
||||
"""Get the cached values for the given keys."""
|
||||
if not keys:
|
||||
return {}
|
||||
|
||||
# Build Redis keys
|
||||
redis_keys = [self._make_key(ns, key) for ns, key in keys]
|
||||
|
||||
# Get values from Redis using MGET
|
||||
try:
|
||||
raw_values = self.redis.mget(redis_keys)
|
||||
except Exception:
|
||||
# If Redis is unavailable, return empty dict
|
||||
return {}
|
||||
|
||||
values: dict[FullKey, ValueT] = {}
|
||||
for i, raw_value in enumerate(raw_values):
|
||||
if raw_value is not None:
|
||||
try:
|
||||
# Deserialize the value
|
||||
encoding, data = raw_value.split(b":", 1)
|
||||
values[keys[i]] = self.serde.loads_typed((encoding.decode(), data))
|
||||
except Exception:
|
||||
# Skip corrupted entries
|
||||
continue
|
||||
|
||||
return values
|
||||
|
||||
async def aget(self, keys: Sequence[FullKey]) -> dict[FullKey, ValueT]:
|
||||
"""Asynchronously get the cached values for the given keys."""
|
||||
return self.get(keys)
|
||||
|
||||
def set(self, mapping: Mapping[FullKey, tuple[ValueT, int | None]]) -> None:
|
||||
"""Set the cached values for the given keys and TTLs."""
|
||||
if not mapping:
|
||||
return
|
||||
|
||||
# Use pipeline for efficient batch operations
|
||||
pipe = self.redis.pipeline()
|
||||
|
||||
for (ns, key), (value, ttl) in mapping.items():
|
||||
redis_key = self._make_key(ns, key)
|
||||
encoding, data = self.serde.dumps_typed(value)
|
||||
|
||||
# Store as "encoding:data" format
|
||||
serialized_value = f"{encoding}:".encode() + data
|
||||
|
||||
if ttl is not None:
|
||||
pipe.setex(redis_key, ttl, serialized_value)
|
||||
else:
|
||||
pipe.set(redis_key, serialized_value)
|
||||
|
||||
try:
|
||||
pipe.execute()
|
||||
except Exception:
|
||||
# Silently fail if Redis is unavailable
|
||||
pass
|
||||
|
||||
async def aset(self, mapping: Mapping[FullKey, tuple[ValueT, int | None]]) -> None:
|
||||
"""Asynchronously set the cached values for the given keys and TTLs."""
|
||||
self.set(mapping)
|
||||
|
||||
def clear(self, namespaces: Sequence[Namespace] | None = None) -> None:
|
||||
"""Delete the cached values for the given namespaces.
|
||||
If no namespaces are provided, clear all cached values."""
|
||||
try:
|
||||
if namespaces is None:
|
||||
# Clear all keys with our prefix
|
||||
pattern = f"{self.prefix}*"
|
||||
keys = self.redis.keys(pattern)
|
||||
if keys:
|
||||
self.redis.delete(*keys)
|
||||
else:
|
||||
# Clear specific namespaces
|
||||
keys_to_delete = []
|
||||
for ns in namespaces:
|
||||
ns_str = ":".join(ns) if ns else ""
|
||||
pattern = (
|
||||
f"{self.prefix}{ns_str}:*" if ns_str else f"{self.prefix}*"
|
||||
)
|
||||
keys = self.redis.keys(pattern)
|
||||
keys_to_delete.extend(keys)
|
||||
|
||||
if keys_to_delete:
|
||||
self.redis.delete(*keys_to_delete)
|
||||
except Exception:
|
||||
# Silently fail if Redis is unavailable
|
||||
pass
|
||||
|
||||
async def aclear(self, namespaces: Sequence[Namespace] | None = None) -> None:
|
||||
"""Asynchronously delete the cached values for the given namespaces.
|
||||
If no namespaces are provided, clear all cached values."""
|
||||
self.clear(namespaces)
|
||||
@@ -81,6 +81,9 @@ class Checkpoint(TypedDict):
|
||||
This keeps track of the versions of the channels that each node has seen.
|
||||
Used to determine which nodes to execute next.
|
||||
"""
|
||||
updated_channels: list[str] | None
|
||||
"""The channels that were updated in this checkpoint.
|
||||
"""
|
||||
|
||||
|
||||
def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
|
||||
@@ -92,6 +95,7 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
|
||||
channel_versions=checkpoint["channel_versions"].copy(),
|
||||
versions_seen={k: v.copy() for k, v in checkpoint["versions_seen"].items()},
|
||||
pending_sends=checkpoint.get("pending_sends", []).copy(),
|
||||
updated_channels=checkpoint.get("updated_channels", None),
|
||||
)
|
||||
|
||||
|
||||
@@ -437,6 +441,7 @@ def empty_checkpoint() -> Checkpoint:
|
||||
channel_versions={},
|
||||
versions_seen={},
|
||||
pending_sends=[],
|
||||
updated_channels=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -470,4 +475,5 @@ def create_checkpoint(
|
||||
channel_versions=checkpoint["channel_versions"],
|
||||
versions_seen=checkpoint["versions_seen"],
|
||||
pending_sends=checkpoint.get("pending_sends", []),
|
||||
updated_channels=None,
|
||||
)
|
||||
|
||||
@@ -64,14 +64,21 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
super().__init__()
|
||||
self._loop = asyncio.get_running_loop()
|
||||
self._aqueue: asyncio.Queue[tuple[asyncio.Future, Op]] = asyncio.Queue()
|
||||
self._task = self._loop.create_task(_run(self._aqueue, weakref.ref(self)))
|
||||
self._task: asyncio.Task | None = None
|
||||
self._ensure_task()
|
||||
|
||||
def __del__(self) -> None:
|
||||
try:
|
||||
self._task.cancel()
|
||||
if self._task:
|
||||
self._task.cancel()
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
def _ensure_task(self) -> None:
|
||||
"""Ensure the background processing loop is running."""
|
||||
if self._task is None or self._task.done():
|
||||
self._task = self._loop.create_task(_run(self._aqueue, weakref.ref(self)))
|
||||
|
||||
async def aget(
|
||||
self,
|
||||
namespace: tuple[str, ...],
|
||||
@@ -79,7 +86,7 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
*,
|
||||
refresh_ttl: bool | None = None,
|
||||
) -> Item | None:
|
||||
assert not self._task.done()
|
||||
self._ensure_task()
|
||||
fut = self._loop.create_future()
|
||||
self._aqueue.put_nowait(
|
||||
(
|
||||
@@ -104,7 +111,7 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
offset: int = 0,
|
||||
refresh_ttl: bool | None = None,
|
||||
) -> list[SearchItem]:
|
||||
assert not self._task.done()
|
||||
self._ensure_task()
|
||||
fut = self._loop.create_future()
|
||||
self._aqueue.put_nowait(
|
||||
(
|
||||
@@ -130,7 +137,7 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
*,
|
||||
ttl: float | None | NotProvided = NOT_PROVIDED,
|
||||
) -> None:
|
||||
assert not self._task.done()
|
||||
self._ensure_task()
|
||||
_validate_namespace(namespace)
|
||||
fut = self._loop.create_future()
|
||||
self._aqueue.put_nowait(
|
||||
@@ -148,7 +155,7 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
namespace: tuple[str, ...],
|
||||
key: str,
|
||||
) -> None:
|
||||
assert not self._task.done()
|
||||
self._ensure_task()
|
||||
fut = self._loop.create_future()
|
||||
self._aqueue.put_nowait((fut, PutOp(namespace, key, None)))
|
||||
return await fut
|
||||
@@ -162,7 +169,7 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
) -> list[tuple[str, ...]]:
|
||||
assert not self._task.done()
|
||||
self._ensure_task()
|
||||
fut = self._loop.create_future()
|
||||
match_conditions = []
|
||||
if prefix:
|
||||
|
||||
@@ -32,6 +32,7 @@ dev = [
|
||||
"numpy",
|
||||
"pandas",
|
||||
"pandas-stubs>=2.2.2.240807",
|
||||
"redis",
|
||||
]
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
|
||||
@@ -0,0 +1,313 @@
|
||||
"""Unit tests for Redis cache implementation."""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
import redis
|
||||
|
||||
from langgraph.cache.redis import RedisCache
|
||||
|
||||
|
||||
class TestRedisCache:
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup(self):
|
||||
"""Set up test Redis client and cache."""
|
||||
self.client = redis.Redis(
|
||||
host="localhost", port=6379, db=0, decode_responses=False
|
||||
)
|
||||
try:
|
||||
self.client.ping()
|
||||
except redis.ConnectionError:
|
||||
pytest.skip("Redis server not available")
|
||||
|
||||
self.cache = RedisCache(self.client, prefix="test:cache:")
|
||||
|
||||
# Clean up before each test
|
||||
self.client.flushdb()
|
||||
|
||||
def teardown_method(self):
|
||||
"""Clean up after each test."""
|
||||
try:
|
||||
self.client.flushdb()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def test_basic_set_and_get(self):
|
||||
"""Test basic set and get operations."""
|
||||
keys = [(("graph", "node"), "key1")]
|
||||
values = {keys[0]: ({"result": 42}, None)}
|
||||
|
||||
# Set value
|
||||
self.cache.set(values)
|
||||
|
||||
# Get value
|
||||
result = self.cache.get(keys)
|
||||
assert len(result) == 1
|
||||
assert result[keys[0]] == {"result": 42}
|
||||
|
||||
def test_batch_operations(self):
|
||||
"""Test batch set and get operations."""
|
||||
keys = [
|
||||
(("graph", "node1"), "key1"),
|
||||
(("graph", "node2"), "key2"),
|
||||
(("other", "node"), "key3"),
|
||||
]
|
||||
values = {
|
||||
keys[0]: ({"result": 1}, None),
|
||||
keys[1]: ({"result": 2}, 60), # With TTL
|
||||
keys[2]: ({"result": 3}, None),
|
||||
}
|
||||
|
||||
# Set values
|
||||
self.cache.set(values)
|
||||
|
||||
# Get all values
|
||||
result = self.cache.get(keys)
|
||||
assert len(result) == 3
|
||||
assert result[keys[0]] == {"result": 1}
|
||||
assert result[keys[1]] == {"result": 2}
|
||||
assert result[keys[2]] == {"result": 3}
|
||||
|
||||
def test_ttl_behavior(self):
|
||||
"""Test TTL (time-to-live) functionality."""
|
||||
key = (("graph", "node"), "ttl_key")
|
||||
values = {key: ({"data": "expires_soon"}, 1)} # 1 second TTL
|
||||
|
||||
# Set with TTL
|
||||
self.cache.set(values)
|
||||
|
||||
# Should be available immediately
|
||||
result = self.cache.get([key])
|
||||
assert len(result) == 1
|
||||
assert result[key] == {"data": "expires_soon"}
|
||||
|
||||
# Wait for expiration
|
||||
time.sleep(1.1)
|
||||
|
||||
# Should be expired
|
||||
result = self.cache.get([key])
|
||||
assert len(result) == 0
|
||||
|
||||
def test_namespace_isolation(self):
|
||||
"""Test that different namespaces are isolated."""
|
||||
key1 = (("graph1", "node"), "same_key")
|
||||
key2 = (("graph2", "node"), "same_key")
|
||||
|
||||
values = {key1: ({"graph": 1}, None), key2: ({"graph": 2}, None)}
|
||||
|
||||
self.cache.set(values)
|
||||
|
||||
result = self.cache.get([key1, key2])
|
||||
assert result[key1] == {"graph": 1}
|
||||
assert result[key2] == {"graph": 2}
|
||||
|
||||
def test_clear_all(self):
|
||||
"""Test clearing all cached values."""
|
||||
keys = [(("graph", "node1"), "key1"), (("graph", "node2"), "key2")]
|
||||
values = {keys[0]: ({"result": 1}, None), keys[1]: ({"result": 2}, None)}
|
||||
|
||||
self.cache.set(values)
|
||||
|
||||
# Verify data exists
|
||||
result = self.cache.get(keys)
|
||||
assert len(result) == 2
|
||||
|
||||
# Clear all
|
||||
self.cache.clear()
|
||||
|
||||
# Verify data is gone
|
||||
result = self.cache.get(keys)
|
||||
assert len(result) == 0
|
||||
|
||||
def test_clear_by_namespace(self):
|
||||
"""Test clearing cached values by namespace."""
|
||||
keys = [
|
||||
(("graph1", "node"), "key1"),
|
||||
(("graph2", "node"), "key2"),
|
||||
(("graph1", "other"), "key3"),
|
||||
]
|
||||
values = {
|
||||
keys[0]: ({"result": 1}, None),
|
||||
keys[1]: ({"result": 2}, None),
|
||||
keys[2]: ({"result": 3}, None),
|
||||
}
|
||||
|
||||
self.cache.set(values)
|
||||
|
||||
# Clear only graph1 namespace
|
||||
self.cache.clear([("graph1", "node"), ("graph1", "other")])
|
||||
|
||||
# graph1 should be cleared, graph2 should remain
|
||||
result = self.cache.get(keys)
|
||||
assert len(result) == 1
|
||||
assert result[keys[1]] == {"result": 2}
|
||||
|
||||
def test_empty_operations(self):
|
||||
"""Test behavior with empty keys/values."""
|
||||
# Empty get
|
||||
result = self.cache.get([])
|
||||
assert result == {}
|
||||
|
||||
# Empty set
|
||||
self.cache.set({}) # Should not raise error
|
||||
|
||||
def test_nonexistent_keys(self):
|
||||
"""Test getting keys that don't exist."""
|
||||
keys = [(("graph", "node"), "nonexistent")]
|
||||
result = self.cache.get(keys)
|
||||
assert len(result) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_operations(self):
|
||||
"""Test async set and get operations with sync Redis client."""
|
||||
# Create sync Redis client and cache (like main integration tests)
|
||||
client = redis.Redis(
|
||||
host="localhost", port=6379, db=1, decode_responses=False
|
||||
)
|
||||
try:
|
||||
client.ping()
|
||||
except Exception:
|
||||
pytest.skip("Redis not available")
|
||||
|
||||
cache = RedisCache(client, prefix="test:async:")
|
||||
|
||||
keys = [(("graph", "node"), "async_key")]
|
||||
values = {keys[0]: ({"async": True}, None)}
|
||||
|
||||
# Async set (delegates to sync)
|
||||
await cache.aset(values)
|
||||
|
||||
# Async get (delegates to sync)
|
||||
result = await cache.aget(keys)
|
||||
assert len(result) == 1
|
||||
assert result[keys[0]] == {"async": True}
|
||||
|
||||
# Cleanup
|
||||
client.flushdb()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_clear(self):
|
||||
"""Test async clear operations with sync Redis client."""
|
||||
# Create sync Redis client and cache (like main integration tests)
|
||||
client = redis.Redis(
|
||||
host="localhost", port=6379, db=1, decode_responses=False
|
||||
)
|
||||
try:
|
||||
client.ping()
|
||||
except Exception:
|
||||
pytest.skip("Redis not available")
|
||||
|
||||
cache = RedisCache(client, prefix="test:async:")
|
||||
|
||||
keys = [(("graph", "node"), "key")]
|
||||
values = {keys[0]: ({"data": "test"}, None)}
|
||||
|
||||
await cache.aset(values)
|
||||
|
||||
# Verify data exists
|
||||
result = await cache.aget(keys)
|
||||
assert len(result) == 1
|
||||
|
||||
# Clear all (delegates to sync)
|
||||
await cache.aclear()
|
||||
|
||||
# Verify data is gone
|
||||
result = await cache.aget(keys)
|
||||
assert len(result) == 0
|
||||
|
||||
# Cleanup
|
||||
client.flushdb()
|
||||
|
||||
def test_redis_unavailable_get(self):
|
||||
"""Test behavior when Redis is unavailable during get operations."""
|
||||
# Create cache with non-existent Redis server
|
||||
bad_client = redis.Redis(
|
||||
host="nonexistent", port=9999, socket_connect_timeout=0.1
|
||||
)
|
||||
cache = RedisCache(bad_client, prefix="test:cache:")
|
||||
|
||||
keys = [(("graph", "node"), "key")]
|
||||
result = cache.get(keys)
|
||||
|
||||
# Should return empty dict when Redis unavailable
|
||||
assert result == {}
|
||||
|
||||
def test_redis_unavailable_set(self):
|
||||
"""Test behavior when Redis is unavailable during set operations."""
|
||||
# Create cache with non-existent Redis server
|
||||
bad_client = redis.Redis(
|
||||
host="nonexistent", port=9999, socket_connect_timeout=0.1
|
||||
)
|
||||
cache = RedisCache(bad_client, prefix="test:cache:")
|
||||
|
||||
keys = [(("graph", "node"), "key")]
|
||||
values = {keys[0]: ({"data": "test"}, None)}
|
||||
|
||||
# Should not raise exception when Redis unavailable
|
||||
cache.set(values) # Should silently fail
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_unavailable_async(self):
|
||||
"""Test async behavior when Redis is unavailable."""
|
||||
# Create sync cache with non-existent Redis server (like main integration tests)
|
||||
bad_client = redis.Redis(
|
||||
host="nonexistent", port=9999, socket_connect_timeout=0.1
|
||||
)
|
||||
cache = RedisCache(bad_client, prefix="test:cache:")
|
||||
|
||||
keys = [(("graph", "node"), "key")]
|
||||
values = {keys[0]: ({"data": "test"}, None)}
|
||||
|
||||
# Should return empty dict for get (delegates to sync)
|
||||
result = await cache.aget(keys)
|
||||
assert result == {}
|
||||
|
||||
# Should not raise exception for set (delegates to sync)
|
||||
await cache.aset(values) # Should silently fail
|
||||
|
||||
def test_corrupted_data_handling(self):
|
||||
"""Test handling of corrupted data in Redis."""
|
||||
# Set some valid data first
|
||||
keys = [(("graph", "node"), "valid_key")]
|
||||
values = {keys[0]: ({"data": "valid"}, None)}
|
||||
self.cache.set(values)
|
||||
|
||||
# Manually insert corrupted data
|
||||
corrupted_key = self.cache._make_key(("graph", "node"), "corrupted_key")
|
||||
self.client.set(corrupted_key, b"invalid:data:format:too:many:colons")
|
||||
|
||||
# Should skip corrupted entry and return only valid ones
|
||||
all_keys = [keys[0], (("graph", "node"), "corrupted_key")]
|
||||
result = self.cache.get(all_keys)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[keys[0]] == {"data": "valid"}
|
||||
|
||||
def test_key_parsing_edge_cases(self):
|
||||
"""Test key parsing with edge cases."""
|
||||
# Test empty namespace
|
||||
key1 = ((), "empty_ns")
|
||||
values = {key1: ({"data": "empty_ns"}, None)}
|
||||
self.cache.set(values)
|
||||
result = self.cache.get([key1])
|
||||
assert result[key1] == {"data": "empty_ns"}
|
||||
|
||||
# Test namespace with special characters
|
||||
key2 = (("graph:with:colons", "node-with-dashes"), "key_with_underscores")
|
||||
values = {key2: ({"data": "special_chars"}, None)}
|
||||
self.cache.set(values)
|
||||
result = self.cache.get([key2])
|
||||
assert result[key2] == {"data": "special_chars"}
|
||||
|
||||
def test_large_data_serialization(self):
|
||||
"""Test handling of large data objects."""
|
||||
# Create a large data structure
|
||||
large_data = {"large_list": list(range(1000)), "nested": {"data": "x" * 1000}}
|
||||
key = (("graph", "node"), "large_key")
|
||||
values = {key: (large_data, None)}
|
||||
|
||||
self.cache.set(values)
|
||||
result = self.cache.get([key])
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[key] == large_data
|
||||
@@ -34,6 +34,42 @@ class MockAsyncBatchedStore(AsyncBatchedBaseStore):
|
||||
return self._store.batch(ops)
|
||||
|
||||
|
||||
async def test_async_batch_store_resilience() -> None:
|
||||
"""Test that AsyncBatchedBaseStore recovers gracefully from task cancellation."""
|
||||
doc = {"foo": "bar"}
|
||||
async_store = MockAsyncBatchedStore()
|
||||
|
||||
await async_store.aput(("foo", "langgraph", "foo"), "bar", doc)
|
||||
|
||||
# Store the original task reference
|
||||
original_task = async_store._task
|
||||
assert original_task is not None
|
||||
assert not original_task.done()
|
||||
|
||||
# Cancel the background task
|
||||
original_task.cancel()
|
||||
await asyncio.sleep(0.01)
|
||||
assert original_task.cancelled()
|
||||
|
||||
# Perform a new operation - this should trigger _ensure_task() to create a new task
|
||||
result = await async_store.asearch(("foo", "langgraph", "foo"))
|
||||
assert len(result) > 0
|
||||
assert result[0].value == doc
|
||||
|
||||
# Verify a new task was created
|
||||
new_task = async_store._task
|
||||
assert new_task is not None
|
||||
assert new_task is not original_task
|
||||
assert not new_task.done()
|
||||
|
||||
# Test that operations continue to work with the new task
|
||||
doc2 = {"baz": "qux"}
|
||||
await async_store.aput(("test", "namespace"), "key", doc2)
|
||||
result2 = await async_store.aget(("test", "namespace"), "key")
|
||||
assert result2 is not None
|
||||
assert result2.value == doc2
|
||||
|
||||
|
||||
def test_get_text_at_path() -> None:
|
||||
nested_data = {
|
||||
"name": "test",
|
||||
|
||||
Generated
+23
@@ -32,6 +32,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/a1/ee/48ca1a7c89ffec8b6a0c5d02b89c305671d5ffd8d3c94acf8b8c408575bb/anyio-4.9.0-py3-none-any.whl", hash = "sha256:9f76d541cad6e36af7beb62e978876f3b41e3e04f2c1fbf0884604c0a9c4d93c", size = 100916, upload-time = "2025-03-17T00:02:52.713Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "async-timeout"
|
||||
version = "5.0.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/a5/ae/136395dfbfe00dfc94da3f3e136d0b13f394cba8f4841120e34226265780/async_timeout-5.0.1.tar.gz", hash = "sha256:d9321a7a3d5a6a5e187e824d2fa0793ce379a202935782d555d6e9d2735677d3", size = 9274, upload-time = "2024-11-06T16:41:39.6Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/fe/ba/e2081de779ca30d473f21f5b30e0e737c438205440784c7dfc81efc2b029/async_timeout-5.0.1-py3-none-any.whl", hash = "sha256:39e3809566ff85354557ec2398b55e096c8364bacac9405a7a1fa429e77fe76c", size = 6233, upload-time = "2024-11-06T16:41:37.9Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "certifi"
|
||||
version = "2025.7.9"
|
||||
@@ -345,6 +354,7 @@ dev = [
|
||||
{ name = "pytest-asyncio" },
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
]
|
||||
|
||||
@@ -366,6 +376,7 @@ dev = [
|
||||
{ name = "pytest-asyncio" },
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
]
|
||||
|
||||
@@ -1153,6 +1164,18 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/19/87/5124b1c1f2412bb95c59ec481eaf936cd32f0fe2a7b16b97b81c4c017a6a/PyYAML-6.0.2-cp39-cp39-win_amd64.whl", hash = "sha256:39693e1f8320ae4f43943590b49779ffb98acb81f788220ea932a6b6c51004d8", size = 162312, upload-time = "2024-08-06T20:33:49.073Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "redis"
|
||||
version = "6.3.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "async-timeout", marker = "python_full_version < '3.11.3'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/21/cd/030274634a1a052b708756016283ea3d84e91ae45f74d7f5dcf55d753a0f/redis-6.3.0.tar.gz", hash = "sha256:3000dbe532babfb0999cdab7b3e5744bcb23e51923febcfaeb52c8cfb29632ef", size = 4647275, upload-time = "2025-08-05T08:12:31.648Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/df/a7/2fe45801534a187543fc45d28b3844d84559c1589255bc2ece30d92dc205/redis-6.3.0-py3-none-any.whl", hash = "sha256:92f079d656ded871535e099080f70fab8e75273c0236797126ac60242d638e9b", size = 280018, upload-time = "2025-08-05T08:12:30.093Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "requests"
|
||||
version = "2.32.4"
|
||||
|
||||
+10
-10
@@ -37,11 +37,11 @@ coverage:
|
||||
--cov-report xml \
|
||||
--cov-report term-missing:skip-covered
|
||||
|
||||
start-postgres:
|
||||
docker compose -f tests/compose-postgres.yml up -V --force-recreate --wait --remove-orphans
|
||||
start-services:
|
||||
docker compose -f tests/compose-postgres.yml -f tests/compose-redis.yml up -V --force-recreate --wait --remove-orphans
|
||||
|
||||
stop-postgres:
|
||||
docker compose -f tests/compose-postgres.yml down -v
|
||||
stop-services:
|
||||
docker compose -f tests/compose-postgres.yml -f tests/compose-redis.yml down -v
|
||||
|
||||
start-dev-server:
|
||||
LOG_LEVEL=warning uv run langgraph dev --config tests/example_app/langgraph.json --no-browser & echo "$$!" > .devserver.pid
|
||||
@@ -60,11 +60,11 @@ NO_DOCKER ?= $(sh command -v docker >/dev/null 2>&1 && echo "false" || echo "tru
|
||||
|
||||
test:
|
||||
if [ "$(NO_DOCKER)" = "false" ]; then \
|
||||
make start-postgres &&\
|
||||
make start-services &&\
|
||||
make start-dev-server &&\
|
||||
uv run pytest $(TEST); \
|
||||
EXIT_CODE=$$?; \
|
||||
make stop-postgres; \
|
||||
make stop-services; \
|
||||
make stop-dev-server; \
|
||||
exit $$EXIT_CODE; \
|
||||
else \
|
||||
@@ -74,11 +74,11 @@ test:
|
||||
fi
|
||||
|
||||
test_parallel:
|
||||
make start-postgres &&\
|
||||
make start-services &&\
|
||||
make start-dev-server &&\
|
||||
uv run pytest -n auto --dist worksteal $(TEST); \
|
||||
EXIT_CODE=$$?; \
|
||||
make stop-postgres; \
|
||||
make stop-services; \
|
||||
make stop-dev-server; \
|
||||
exit $$EXIT_CODE
|
||||
|
||||
@@ -93,11 +93,11 @@ MAXFAIL_ARGS := $(if $(MAXFAIL),--maxfail $(MAXFAIL),)
|
||||
XDIST_ARGS := $(if $(WORKERS),-x $(XDIST_ARGS),)
|
||||
|
||||
test_watch:
|
||||
make start-postgres &&\
|
||||
make start-services &&\
|
||||
make start-dev-server &&\
|
||||
uv run ptw . -- --ff -vv $(XDIST_ARGS) $(MAXFAIL_ARGS) $(TEST); \
|
||||
EXIT_CODE=$$?; \
|
||||
make stop-postgres; \
|
||||
make stop-services; \
|
||||
make stop-dev-server; \
|
||||
exit $$EXIT_CODE
|
||||
|
||||
|
||||
@@ -41,16 +41,16 @@ _Writer = Callable[
|
||||
|
||||
|
||||
def _get_branch_path_input_schema(
|
||||
path: Callable[..., Hashable | list[Hashable]]
|
||||
| Callable[..., Awaitable[Hashable | list[Hashable]]]
|
||||
| Runnable[Any, Hashable | list[Hashable]],
|
||||
path: Callable[..., Hashable | Sequence[Hashable]]
|
||||
| Callable[..., Awaitable[Hashable | Sequence[Hashable]]]
|
||||
| Runnable[Any, Hashable | Sequence[Hashable]],
|
||||
) -> type[Any] | None:
|
||||
input = None
|
||||
# detect input schema annotation in the branch callable
|
||||
try:
|
||||
callable_: (
|
||||
Callable[..., Hashable | list[Hashable]]
|
||||
| Callable[..., Awaitable[Hashable | list[Hashable]]]
|
||||
Callable[..., Hashable | Sequence[Hashable]]
|
||||
| Callable[..., Awaitable[Hashable | Sequence[Hashable]]]
|
||||
| None
|
||||
) = None
|
||||
if isinstance(path, (RunnableCallable, RunnableLambda)):
|
||||
|
||||
@@ -22,10 +22,11 @@ from langchain_core.messages import (
|
||||
convert_to_messages,
|
||||
message_chunk_to_message,
|
||||
)
|
||||
from typing_extensions import TypedDict
|
||||
from typing_extensions import TypedDict, deprecated
|
||||
|
||||
from langgraph._internal._constants import CONF, CONFIG_KEY_SEND, NS_SEP
|
||||
from langgraph.graph.state import StateGraph
|
||||
from langgraph.warnings import LangGraphDeprecatedSinceV10
|
||||
|
||||
__all__ = (
|
||||
"add_messages",
|
||||
@@ -233,9 +234,16 @@ def add_messages(
|
||||
return merged
|
||||
|
||||
|
||||
@deprecated(
|
||||
"MessageGraph is deprecated in LangGraph v1.0.0, to be removed in v2.0.0. Please use StateGraph with a `messages` key instead.",
|
||||
category=None,
|
||||
)
|
||||
class MessageGraph(StateGraph):
|
||||
"""A StateGraph where every node receives a list of messages as input and returns one or more messages as output.
|
||||
|
||||
!!! warning "Deprecation"
|
||||
MessageGraph is deprecated in LangGraph v1.0.0, to be removed in v2.0.0. Please use StateGraph with a `messages` key instead.
|
||||
|
||||
MessageGraph is a subclass of StateGraph whose entire state is a single, append-only* list of messages.
|
||||
Each node in a MessageGraph takes a list of messages as input and returns zero or more
|
||||
messages as output. The `add_messages` function is used to merge the output messages from each node
|
||||
@@ -281,6 +289,11 @@ class MessageGraph(StateGraph):
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
warnings.warn(
|
||||
"MessageGraph is deprecated in LangGraph v1.0.0, to be removed in v2.0.0. Please use StateGraph with a `messages` key instead.",
|
||||
category=LangGraphDeprecatedSinceV10,
|
||||
stacklevel=2,
|
||||
)
|
||||
super().__init__(Annotated[list[AnyMessage], add_messages]) # type: ignore[arg-type]
|
||||
|
||||
|
||||
|
||||
@@ -607,9 +607,9 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
|
||||
def add_conditional_edges(
|
||||
self,
|
||||
source: str,
|
||||
path: Callable[..., Hashable | list[Hashable]]
|
||||
| Callable[..., Awaitable[Hashable | list[Hashable]]]
|
||||
| Runnable[Any, Hashable | list[Hashable]],
|
||||
path: Callable[..., Hashable | Sequence[Hashable]]
|
||||
| Callable[..., Awaitable[Hashable | Sequence[Hashable]]]
|
||||
| Runnable[Any, Hashable | Sequence[Hashable]],
|
||||
path_map: dict[Hashable, str] | list[str] | None = None,
|
||||
) -> Self:
|
||||
"""Add a conditional edge from the starting node to any number of destination nodes.
|
||||
@@ -710,9 +710,9 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
|
||||
|
||||
def set_conditional_entry_point(
|
||||
self,
|
||||
path: Callable[..., Hashable | list[Hashable]]
|
||||
| Callable[..., Awaitable[Hashable | list[Hashable]]]
|
||||
| Runnable[Any, Hashable | list[Hashable]],
|
||||
path: Callable[..., Hashable | Sequence[Hashable]]
|
||||
| Callable[..., Awaitable[Hashable | Sequence[Hashable]]]
|
||||
| Runnable[Any, Hashable | Sequence[Hashable]],
|
||||
path_map: dict[Hashable, str] | list[str] | None = None,
|
||||
) -> Self:
|
||||
"""Sets a conditional entry point in the graph.
|
||||
|
||||
@@ -29,6 +29,7 @@ def create_checkpoint(
|
||||
step: int,
|
||||
*,
|
||||
id: str | None = None,
|
||||
updated_channels: set[str] | None = None,
|
||||
) -> Checkpoint:
|
||||
"""Create a checkpoint for the given channels."""
|
||||
ts = datetime.now(timezone.utc).isoformat()
|
||||
@@ -49,6 +50,7 @@ def create_checkpoint(
|
||||
channel_values=values,
|
||||
channel_versions=checkpoint["channel_versions"],
|
||||
versions_seen=checkpoint["versions_seen"],
|
||||
updated_channels=None if updated_channels is None else sorted(updated_channels),
|
||||
)
|
||||
|
||||
|
||||
@@ -81,4 +83,5 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
|
||||
channel_values=checkpoint["channel_values"].copy(),
|
||||
channel_versions=checkpoint["channel_versions"].copy(),
|
||||
versions_seen={k: v.copy() for k, v in checkpoint["versions_seen"].items()},
|
||||
updated_channels=checkpoint.get("updated_channels", None),
|
||||
)
|
||||
|
||||
@@ -568,7 +568,9 @@ class PregelLoop:
|
||||
if task := tasks.get(tid):
|
||||
task.writes.append((k, v))
|
||||
|
||||
def _first(self, *, input_keys: str | Sequence[str]) -> set[str] | None:
|
||||
def _first(
|
||||
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
|
||||
) -> set[str] | None:
|
||||
# resuming from previous checkpoint requires
|
||||
# - finding a previous checkpoint
|
||||
# - receiving None input (outer graph) or RESUMING flag (subgraph)
|
||||
@@ -585,8 +587,6 @@ class PregelLoop:
|
||||
),
|
||||
)
|
||||
)
|
||||
# this can be set only when there are input_writes
|
||||
updated_channels: set[str] | None = None
|
||||
|
||||
# map command to writes
|
||||
if isinstance(self.input, Command):
|
||||
@@ -614,13 +614,15 @@ class PregelLoop:
|
||||
if null_writes := [
|
||||
w[1:] for w in self.checkpoint_pending_writes if w[0] == NULL_TASK_ID
|
||||
]:
|
||||
apply_writes(
|
||||
null_updated_channels = apply_writes(
|
||||
self.checkpoint,
|
||||
self.channels,
|
||||
[PregelTaskWrites((), INPUT, null_writes, [])],
|
||||
self.checkpointer_get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
if updated_channels is not None:
|
||||
updated_channels.update(null_updated_channels)
|
||||
# proceed past previous checkpoint
|
||||
if is_resuming:
|
||||
self.checkpoint["versions_seen"].setdefault(INTERRUPT, {})
|
||||
@@ -648,6 +650,7 @@ class PregelLoop:
|
||||
store=None,
|
||||
checkpointer=None,
|
||||
manager=None,
|
||||
updated_channels=updated_channels,
|
||||
)
|
||||
# apply input writes
|
||||
updated_channels = apply_writes(
|
||||
@@ -661,6 +664,7 @@ class PregelLoop:
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
# save input checkpoint
|
||||
self.updated_channels = updated_channels
|
||||
self._put_checkpoint({"source": "input"})
|
||||
elif CONFIG_KEY_RESUMING not in configurable:
|
||||
raise EmptyInputError(f"Received no input for {input_keys}")
|
||||
@@ -693,6 +697,7 @@ class PregelLoop:
|
||||
self.channels if do_checkpoint else None,
|
||||
self.step,
|
||||
id=self.checkpoint["id"] if exiting else None,
|
||||
updated_channels=self.updated_channels,
|
||||
)
|
||||
# bail if no checkpointer
|
||||
if do_checkpoint and self._checkpointer_put_after_previous is not None:
|
||||
@@ -1036,7 +1041,12 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
self.step = self.checkpoint_metadata["step"] + 1
|
||||
self.stop = self.step + self.config["recursion_limit"] + 1
|
||||
self.checkpoint_previous_versions = self.checkpoint["channel_versions"].copy()
|
||||
self.updated_channels = self._first(input_keys=self.input_keys)
|
||||
self.updated_channels = self._first(
|
||||
input_keys=self.input_keys,
|
||||
updated_channels=set(self.checkpoint.get("updated_channels")) # type: ignore[arg-type]
|
||||
if self.checkpoint.get("updated_channels")
|
||||
else None,
|
||||
)
|
||||
|
||||
return self
|
||||
|
||||
@@ -1212,7 +1222,12 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
self.step = self.checkpoint_metadata["step"] + 1
|
||||
self.stop = self.step + self.config["recursion_limit"] + 1
|
||||
self.checkpoint_previous_versions = self.checkpoint["channel_versions"].copy()
|
||||
self.updated_channels = self._first(input_keys=self.input_keys)
|
||||
self.updated_channels = self._first(
|
||||
input_keys=self.input_keys,
|
||||
updated_channels=set(self.checkpoint.get("updated_channels")) # type: ignore[arg-type]
|
||||
if self.checkpoint.get("updated_channels")
|
||||
else None,
|
||||
)
|
||||
|
||||
return self
|
||||
|
||||
|
||||
@@ -29,16 +29,51 @@ Meta = tuple[tuple[str, ...], dict[str, Any]]
|
||||
|
||||
class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
|
||||
"""A callback handler that implements stream_mode=messages.
|
||||
Collects messages from (1) chat model stream events and (2) node outputs."""
|
||||
|
||||
Collects messages from:
|
||||
(1) chat model stream events; and
|
||||
(2) node outputs.
|
||||
"""
|
||||
|
||||
run_inline = True
|
||||
"""We want this callback to run in the main thread, to avoid order/locking issues."""
|
||||
"""We want this callback to run in the main thread to avoid order/locking issues."""
|
||||
|
||||
def __init__(self, stream: Callable[[StreamChunk], None], subgraphs: bool):
|
||||
def __init__(
|
||||
self,
|
||||
stream: Callable[[StreamChunk], None],
|
||||
subgraphs: bool,
|
||||
*,
|
||||
parent_ns: tuple[str, ...] | None = None,
|
||||
) -> None:
|
||||
"""Configure the handler to stream messages from LLMs and nodes.
|
||||
|
||||
Args:
|
||||
stream: A callable that takes a StreamChunk and emits it.
|
||||
subgraphs: Whether to emit messages from subgraphs.
|
||||
parent_ns: The namespace where the handler was created.
|
||||
We keep track of this namespace to allow calls to subgraphs that
|
||||
were explicitly requested as a stream with `messages` mode
|
||||
configured.
|
||||
|
||||
Example:
|
||||
parent_ns is used to handle scenarios where the subgraph is explicitly
|
||||
streamed with `stream_mode="messages"`.
|
||||
|
||||
```python
|
||||
def parent_graph_node():
|
||||
# This node is in the parent graph.
|
||||
async for event in some_subgraph(..., stream_mode="messages"):
|
||||
do something with event # <-- these events will be emitted
|
||||
return ...
|
||||
|
||||
parent_graph.invoke(subgraphs=False)
|
||||
```
|
||||
"""
|
||||
self.stream = stream
|
||||
self.subgraphs = subgraphs
|
||||
self.metadata: dict[UUID, Meta] = {}
|
||||
self.seen: set[int | str] = set()
|
||||
self.parent_ns = parent_ns
|
||||
|
||||
def _emit(self, meta: Meta, message: BaseMessage, *, dedupe: bool = False) -> None:
|
||||
if dedupe and message.id in self.seen:
|
||||
@@ -100,7 +135,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
|
||||
ns = tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP))[
|
||||
:-1
|
||||
]
|
||||
if not self.subgraphs and len(ns) > 0:
|
||||
if not self.subgraphs and len(ns) > 0 and ns != self.parent_ns:
|
||||
return
|
||||
if tags:
|
||||
if filtered_tags := [t for t in tags if not t.startswith("seq:step")]:
|
||||
|
||||
@@ -11,7 +11,7 @@ from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
|
||||
from dataclasses import is_dataclass
|
||||
from functools import partial
|
||||
from inspect import isclass
|
||||
from typing import Any, Callable, Generic, Union, cast, get_type_hints
|
||||
from typing import Any, Callable, Generic, Optional, Union, cast, get_type_hints
|
||||
from uuid import UUID, uuid5
|
||||
|
||||
from langchain_core.globals import get_debug
|
||||
@@ -2534,8 +2534,13 @@ class Pregel(
|
||||
config[CONF][CONFIG_KEY_CHECKPOINT_NS] = recast_checkpoint_ns(ns)
|
||||
# set up messages stream mode
|
||||
if "messages" in stream_modes:
|
||||
ns_ = cast(Optional[str], config[CONF].get(CONFIG_KEY_CHECKPOINT_NS))
|
||||
run_manager.inheritable_handlers.append(
|
||||
StreamMessagesHandler(stream.put, subgraphs)
|
||||
StreamMessagesHandler(
|
||||
stream.put,
|
||||
subgraphs,
|
||||
parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None,
|
||||
)
|
||||
)
|
||||
|
||||
# set up custom stream mode
|
||||
@@ -2814,8 +2819,14 @@ class Pregel(
|
||||
config[CONF][CONFIG_KEY_CHECKPOINT_NS] = recast_checkpoint_ns(ns)
|
||||
# set up messages stream mode
|
||||
if "messages" in stream_modes:
|
||||
# namespace can be None in a root level graph?
|
||||
ns_ = cast(Optional[str], config[CONF].get(CONFIG_KEY_CHECKPOINT_NS))
|
||||
run_manager.inheritable_handlers.append(
|
||||
StreamMessagesHandler(stream_put, subgraphs)
|
||||
StreamMessagesHandler(
|
||||
stream_put,
|
||||
subgraphs,
|
||||
parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None,
|
||||
)
|
||||
)
|
||||
|
||||
# set up custom stream mode
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph"
|
||||
version = "0.6.3"
|
||||
version = "0.6.4"
|
||||
description = "Building stateful, multi-actor applications with LLMs"
|
||||
authors = []
|
||||
requires-python = ">=3.9"
|
||||
@@ -49,6 +49,7 @@ dev = [
|
||||
"types-requests",
|
||||
"pycryptodome",
|
||||
"langgraph-cli[inmem]",
|
||||
"redis",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
name: langgraph-tests
|
||||
services:
|
||||
redis-test:
|
||||
image: redis:7-alpine
|
||||
ports:
|
||||
- "6379:6379"
|
||||
command: redis-server --maxmemory 256mb --maxmemory-policy allkeys-lru
|
||||
healthcheck:
|
||||
test: redis-cli ping
|
||||
start_period: 10s
|
||||
timeout: 1s
|
||||
retries: 5
|
||||
interval: 5s
|
||||
start_interval: 1s
|
||||
tmpfs:
|
||||
- /data # Use tmpfs for faster testing
|
||||
@@ -3,10 +3,12 @@ from collections.abc import AsyncIterator, Iterator
|
||||
from uuid import UUID
|
||||
|
||||
import pytest
|
||||
import redis
|
||||
from pytest_mock import MockerFixture
|
||||
|
||||
from langgraph.cache.base import BaseCache
|
||||
from langgraph.cache.memory import InMemoryCache
|
||||
from langgraph.cache.redis import RedisCache
|
||||
from langgraph.cache.sqlite import SqliteCache
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.store.base import BaseStore
|
||||
@@ -55,12 +57,34 @@ def durability(request: pytest.FixtureRequest) -> Durability:
|
||||
return request.param
|
||||
|
||||
|
||||
@pytest.fixture(scope="function", params=["sqlite", "memory"])
|
||||
@pytest.fixture(
|
||||
scope="function",
|
||||
params=["sqlite", "memory"] if NO_DOCKER else ["sqlite", "memory", "redis"],
|
||||
)
|
||||
def cache(request: pytest.FixtureRequest) -> Iterator[BaseCache]:
|
||||
if request.param == "sqlite":
|
||||
yield SqliteCache(path=":memory:")
|
||||
elif request.param == "memory":
|
||||
yield InMemoryCache()
|
||||
elif request.param == "redis":
|
||||
# Get worker ID for parallel test isolation
|
||||
worker_id = getattr(request.config, "workerinput", {}).get("workerid", "master")
|
||||
|
||||
redis_client = redis.Redis(
|
||||
host="localhost", port=6379, db=0, decode_responses=False
|
||||
)
|
||||
# Use worker-specific prefix to avoid cache pollution between parallel tests
|
||||
cache = RedisCache(redis_client, prefix=f"test:cache:{worker_id}:")
|
||||
yield cache
|
||||
|
||||
try:
|
||||
# Only clear keys with our specific prefix
|
||||
pattern = f"test:cache:{worker_id}:*"
|
||||
keys = redis_client.keys(pattern)
|
||||
if keys:
|
||||
redis_client.delete(*keys)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
raise ValueError(f"Unknown cache type: {request.param}")
|
||||
|
||||
|
||||
@@ -330,6 +330,7 @@ SAVED_CHECKPOINTS = {
|
||||
"docs": ["doc1", "doc2", "doc3", "doc4"],
|
||||
"answer": "doc1,doc2,doc3,doc4",
|
||||
},
|
||||
"updated_channels": None,
|
||||
},
|
||||
metadata={
|
||||
"source": "loop",
|
||||
@@ -390,6 +391,7 @@ SAVED_CHECKPOINTS = {
|
||||
"docs": ["doc1", "doc2", "doc3", "doc4"],
|
||||
"branch:to:qa": None,
|
||||
},
|
||||
"updated_channels": None,
|
||||
},
|
||||
metadata={
|
||||
"source": "loop",
|
||||
@@ -465,6 +467,7 @@ SAVED_CHECKPOINTS = {
|
||||
"branch:to:retriever_one": None,
|
||||
"docs": ["doc3", "doc4"],
|
||||
},
|
||||
"updated_channels": None,
|
||||
},
|
||||
metadata={
|
||||
"source": "loop",
|
||||
@@ -516,6 +519,7 @@ SAVED_CHECKPOINTS = {
|
||||
"branch:to:analyzer_one": None,
|
||||
"branch:to:retriever_two": None,
|
||||
},
|
||||
"updated_channels": None,
|
||||
},
|
||||
metadata={
|
||||
"source": "loop",
|
||||
@@ -570,6 +574,7 @@ SAVED_CHECKPOINTS = {
|
||||
"query": "what is weather in sf",
|
||||
"branch:to:rewrite_query": None,
|
||||
},
|
||||
"updated_channels": None,
|
||||
},
|
||||
metadata={
|
||||
"source": "loop",
|
||||
@@ -618,6 +623,7 @@ SAVED_CHECKPOINTS = {
|
||||
},
|
||||
"versions_seen": {"__input__": {}},
|
||||
"channel_values": {"__start__": {"query": "what is weather in sf"}},
|
||||
"updated_channels": None,
|
||||
},
|
||||
metadata={
|
||||
"source": "input",
|
||||
|
||||
@@ -12,6 +12,7 @@ from langgraph.channels.last_value import LastValue
|
||||
from langgraph.errors import NodeInterrupt
|
||||
from langgraph.func import entrypoint, task
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.graph.message import MessageGraph
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
from langgraph.types import Interrupt, RetryPolicy
|
||||
from langgraph.warnings import LangGraphDeprecatedSinceV05, LangGraphDeprecatedSinceV10
|
||||
@@ -332,3 +333,11 @@ def test_config_parameter_incorrect_typing() -> None:
|
||||
|
||||
builder.add_node(async_node_with_untyped_config)
|
||||
assert len(w) == 0
|
||||
|
||||
|
||||
def test_message_graph_deprecation() -> None:
|
||||
with pytest.warns(
|
||||
LangGraphDeprecatedSinceV10,
|
||||
match="MessageGraph is deprecated in LangGraph v1.0.0, to be removed in v2.0.0. Please use StateGraph with a `messages` key instead.",
|
||||
):
|
||||
MessageGraph()
|
||||
|
||||
@@ -6,7 +6,9 @@ from dataclasses import replace
|
||||
from typing import Annotated, Any, Literal, Optional, Union, cast
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, AnyMessage, ToolCall
|
||||
from langchain_core.runnables import RunnableConfig, RunnableMap, RunnablePick
|
||||
from langchain_core.tools import tool
|
||||
from pytest_mock import MockerFixture
|
||||
from syrupy import SnapshotAssertion
|
||||
from typing_extensions import TypedDict
|
||||
@@ -18,7 +20,7 @@ from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.constants import END, START
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.graph.message import MessageGraph, MessagesState, add_messages
|
||||
from langgraph.graph.message import MessagesState, add_messages
|
||||
from langgraph.prebuilt.chat_agent_executor import create_react_agent
|
||||
from langgraph.prebuilt.tool_node import ToolNode
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
@@ -2441,7 +2443,7 @@ def test_message_graph(
|
||||
return "continue"
|
||||
|
||||
# Define a new graph
|
||||
workflow = MessageGraph()
|
||||
workflow = StateGraph(state_schema=Annotated[list[AnyMessage], add_messages]) # type: ignore[arg-type]
|
||||
|
||||
# Define the two nodes we will cycle between
|
||||
workflow.add_node("agent", model)
|
||||
@@ -2487,7 +2489,7 @@ def test_message_graph(
|
||||
assert json.dumps(app.get_graph().to_json(), indent=2) == snapshot
|
||||
assert app.get_graph().draw_mermaid(with_styles=False) == snapshot
|
||||
|
||||
assert app.invoke(HumanMessage(content="what is weather in sf")) == [
|
||||
assert app.invoke([HumanMessage(content="what is weather in sf")]) == [
|
||||
_AnyIdHumanMessage(
|
||||
content="what is weather in sf",
|
||||
),
|
||||
@@ -6435,10 +6437,6 @@ def test_weather_subgraph(
|
||||
from langchain_core.language_models.fake_chat_models import (
|
||||
FakeMessagesListChatModel,
|
||||
)
|
||||
from langchain_core.messages import AIMessage, ToolCall
|
||||
from langchain_core.tools import tool
|
||||
|
||||
from langgraph.graph import MessagesState
|
||||
|
||||
# setup subgraph
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ from typing import (
|
||||
)
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import ToolCall
|
||||
from langchain_core.messages import AnyMessage, ToolCall
|
||||
from langchain_core.runnables import RunnableConfig, RunnablePick
|
||||
from pytest_mock import MockerFixture
|
||||
from typing_extensions import TypedDict
|
||||
@@ -21,7 +21,7 @@ from langgraph.channels.last_value import LastValue
|
||||
from langgraph.channels.untracked_value import UntrackedValue
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.constants import END, START
|
||||
from langgraph.graph.message import MessageGraph, add_messages
|
||||
from langgraph.graph.message import add_messages
|
||||
from langgraph.graph.state import StateGraph
|
||||
from langgraph.prebuilt.chat_agent_executor import create_react_agent
|
||||
from langgraph.prebuilt.tool_node import ToolNode
|
||||
@@ -2117,7 +2117,7 @@ async def test_message_graph(async_checkpointer: BaseCheckpointSaver) -> None:
|
||||
return "continue"
|
||||
|
||||
# Define a new graph
|
||||
workflow = MessageGraph()
|
||||
workflow = StateGraph(state_schema=Annotated[list[AnyMessage], add_messages]) # type: ignore[arg-type]
|
||||
|
||||
# Define the two nodes we will cycle between
|
||||
workflow.add_node("agent", model)
|
||||
@@ -2157,7 +2157,7 @@ async def test_message_graph(async_checkpointer: BaseCheckpointSaver) -> None:
|
||||
# meaning you can use it as you would any other runnable
|
||||
app = workflow.compile()
|
||||
|
||||
assert await app.ainvoke(HumanMessage(content="what is weather in sf")) == [
|
||||
assert await app.ainvoke([HumanMessage(content="what is weather in sf")]) == [
|
||||
_AnyIdHumanMessage(
|
||||
content="what is weather in sf",
|
||||
),
|
||||
|
||||
@@ -16,6 +16,7 @@ from typing import Annotated, Any, Literal, Optional, Union, get_type_hints
|
||||
|
||||
import pytest
|
||||
from langchain_core.language_models import GenericFakeChatModel
|
||||
from langchain_core.messages import AnyMessage
|
||||
from langchain_core.runnables import (
|
||||
RunnableConfig,
|
||||
RunnableLambda,
|
||||
@@ -26,7 +27,7 @@ from langsmith import traceable
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
from pytest_mock import MockerFixture
|
||||
from syrupy import SnapshotAssertion
|
||||
from typing_extensions import TypedDict
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
from langgraph._internal._constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL
|
||||
from langgraph.cache.base import BaseCache
|
||||
@@ -45,7 +46,7 @@ from langgraph.config import get_stream_writer
|
||||
from langgraph.errors import GraphRecursionError, InvalidUpdateError, ParentCommand
|
||||
from langgraph.func import entrypoint, task
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from langgraph.graph.message import MessageGraph, MessagesState, add_messages
|
||||
from langgraph.graph.message import MessagesState, add_messages
|
||||
from langgraph.prebuilt.tool_node import ToolNode
|
||||
from langgraph.pregel import (
|
||||
NodeBuilder,
|
||||
@@ -967,6 +968,7 @@ def test_pending_writes_resume(
|
||||
"branch:to:two": AnyVersion(),
|
||||
},
|
||||
"channel_values": {"value": 6},
|
||||
"updated_channels": ["value"],
|
||||
},
|
||||
metadata={
|
||||
"parents": {},
|
||||
@@ -1014,6 +1016,7 @@ def test_pending_writes_resume(
|
||||
"branch:to:one": None,
|
||||
"branch:to:two": None,
|
||||
},
|
||||
"updated_channels": ["branch:to:one", "branch:to:two", "value"],
|
||||
},
|
||||
metadata={
|
||||
"parents": {},
|
||||
@@ -1065,6 +1068,7 @@ def test_pending_writes_resume(
|
||||
"__start__": AnyVersion(),
|
||||
},
|
||||
"channel_values": {"__start__": {"value": 1}},
|
||||
"updated_channels": ["__start__"],
|
||||
},
|
||||
metadata={
|
||||
"parents": {},
|
||||
@@ -3907,7 +3911,7 @@ def test_remove_message_via_state_update(
|
||||
) -> None:
|
||||
from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage
|
||||
|
||||
workflow = MessageGraph()
|
||||
workflow = StateGraph(state_schema=Annotated[list[AnyMessage], add_messages]) # type: ignore[arg-type]
|
||||
workflow.add_node(
|
||||
"chatbot",
|
||||
lambda state: [
|
||||
@@ -3940,7 +3944,7 @@ def test_remove_message_via_state_update(
|
||||
def test_remove_message_from_node():
|
||||
from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage
|
||||
|
||||
workflow = MessageGraph()
|
||||
workflow = StateGraph(state_schema=Annotated[list[AnyMessage], add_messages]) # type: ignore[arg-type]
|
||||
workflow.add_node(
|
||||
"chatbot",
|
||||
lambda state: [
|
||||
@@ -8262,3 +8266,53 @@ def test_fork_and_update_task_results(sync_checkpointer: BaseCheckpointSaver) ->
|
||||
],
|
||||
],
|
||||
]
|
||||
|
||||
|
||||
def test_subgraph_streaming_sync() -> None:
|
||||
"""Test subgraph streaming when used as a node in sync version"""
|
||||
|
||||
# Create a fake chat model that returns a simple response
|
||||
model = GenericFakeChatModel(messages=iter(["The weather is sunny today."]))
|
||||
|
||||
# Create a subgraph that uses the fake chat model
|
||||
def call_model_node(state: MessagesState, config: RunnableConfig) -> MessagesState:
|
||||
"""Node that calls the model with the last message."""
|
||||
messages = state["messages"]
|
||||
last_message = messages[-1].content if messages else ""
|
||||
response = model.invoke([("user", last_message)], config)
|
||||
return {"messages": [response]}
|
||||
|
||||
# Build the subgraph
|
||||
subgraph = StateGraph(MessagesState)
|
||||
subgraph.add_node("call_model", call_model_node)
|
||||
subgraph.add_edge(START, "call_model")
|
||||
compiled_subgraph = subgraph.compile()
|
||||
|
||||
class SomeCustomState(TypedDict):
|
||||
last_chunk: NotRequired[str]
|
||||
num_chunks: NotRequired[int]
|
||||
|
||||
# Will invoke a subgraph as a function
|
||||
def parent_node(state: SomeCustomState, config: RunnableConfig) -> dict:
|
||||
"""Node that runs the subgraph."""
|
||||
msgs = {"messages": [("user", "What is the weather in Tokyo?")]}
|
||||
events = []
|
||||
for event in compiled_subgraph.stream(msgs, config, stream_mode="messages"):
|
||||
events.append(event)
|
||||
ai_msg_chunks = [ai_msg_chunk for ai_msg_chunk, _ in events]
|
||||
return {
|
||||
"last_chunk": ai_msg_chunks[-1],
|
||||
"num_chunks": len(ai_msg_chunks),
|
||||
}
|
||||
|
||||
# Build the main workflow
|
||||
workflow = StateGraph(SomeCustomState)
|
||||
workflow.add_node("subgraph", parent_node)
|
||||
workflow.add_edge(START, "subgraph")
|
||||
compiled_workflow = workflow.compile()
|
||||
|
||||
# Test the basic functionality
|
||||
result = compiled_workflow.invoke({})
|
||||
|
||||
assert result["last_chunk"].content == "today."
|
||||
assert result["num_chunks"] == 9
|
||||
|
||||
@@ -26,7 +26,7 @@ from langchain_core.utils.aiter import aclosing
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
from pytest_mock import MockerFixture
|
||||
from syrupy import SnapshotAssertion
|
||||
from typing_extensions import TypedDict
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
from langgraph._internal._constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL
|
||||
from langgraph.cache.base import BaseCache
|
||||
@@ -1908,6 +1908,7 @@ async def test_pending_writes_resume(
|
||||
"branch:to:two": AnyVersion(),
|
||||
},
|
||||
"channel_values": {"value": 6},
|
||||
"updated_channels": ["value"],
|
||||
},
|
||||
metadata={
|
||||
"parents": {},
|
||||
@@ -1955,6 +1956,7 @@ async def test_pending_writes_resume(
|
||||
"branch:to:one": None,
|
||||
"branch:to:two": None,
|
||||
},
|
||||
"updated_channels": ["branch:to:one", "branch:to:two", "value"],
|
||||
},
|
||||
metadata={
|
||||
"parents": {},
|
||||
@@ -2002,6 +2004,7 @@ async def test_pending_writes_resume(
|
||||
"__start__": AnyVersion(),
|
||||
},
|
||||
"channel_values": {"__start__": {"value": 1}},
|
||||
"updated_channels": ["__start__"],
|
||||
},
|
||||
metadata={
|
||||
"parents": {},
|
||||
@@ -9050,3 +9053,57 @@ async def test_fork_and_update_task_results(
|
||||
],
|
||||
],
|
||||
]
|
||||
|
||||
|
||||
async def test_subgraph_streaming_async() -> None:
|
||||
"""Test subgraph streaming when used as a node in async version"""
|
||||
|
||||
# Create a fake chat model that returns a simple response
|
||||
model = GenericFakeChatModel(messages=iter(["The weather is sunny today."]))
|
||||
|
||||
# Create a subgraph that uses the fake chat model
|
||||
async def call_model_node(
|
||||
state: MessagesState, config: RunnableConfig
|
||||
) -> MessagesState:
|
||||
"""Node that calls the model with the last message."""
|
||||
messages = state["messages"]
|
||||
last_message = messages[-1].content if messages else ""
|
||||
response = await model.ainvoke([("user", last_message)], config)
|
||||
return {"messages": [response]}
|
||||
|
||||
# Build the subgraph
|
||||
subgraph = StateGraph(MessagesState)
|
||||
subgraph.add_node("call_model", call_model_node)
|
||||
subgraph.add_edge(START, "call_model")
|
||||
compiled_subgraph = subgraph.compile()
|
||||
|
||||
class SomeCustomState(TypedDict):
|
||||
last_chunk: NotRequired[str]
|
||||
num_chunks: NotRequired[int]
|
||||
|
||||
# Will invoke a subgraph as a function
|
||||
async def parent_node(state: SomeCustomState, config: RunnableConfig) -> dict:
|
||||
"""Node that runs the subgraph."""
|
||||
msgs = {"messages": [("user", "What is the weather in Tokyo?")]}
|
||||
events = []
|
||||
async for event in compiled_subgraph.astream(
|
||||
msgs, config, stream_mode="messages"
|
||||
):
|
||||
events.append(event)
|
||||
ai_msg_chunks = [ai_msg_chunk for ai_msg_chunk, _ in events]
|
||||
return {
|
||||
"last_chunk": ai_msg_chunks[-1],
|
||||
"num_chunks": len(ai_msg_chunks),
|
||||
}
|
||||
|
||||
# Build the main workflow
|
||||
workflow = StateGraph(SomeCustomState)
|
||||
workflow.add_node("subgraph", parent_node)
|
||||
workflow.add_edge(START, "subgraph")
|
||||
compiled_workflow = workflow.compile()
|
||||
|
||||
# Test the basic functionality
|
||||
result = await compiled_workflow.ainvoke({})
|
||||
|
||||
assert result["last_chunk"].content == "today."
|
||||
assert result["num_chunks"] == 9
|
||||
|
||||
Generated
+26
-2
@@ -119,6 +119,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/03/49/d10027df9fce941cb8184e78a02857af36360d33e1721df81c5ed2179a1a/async_lru-2.0.5-py3-none-any.whl", hash = "sha256:ab95404d8d2605310d345932697371a5f40def0487c03d6d0ad9138de52c9943", size = 6069, upload-time = "2025-03-16T17:25:35.422Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "async-timeout"
|
||||
version = "5.0.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/a5/ae/136395dfbfe00dfc94da3f3e136d0b13f394cba8f4841120e34226265780/async_timeout-5.0.1.tar.gz", hash = "sha256:d9321a7a3d5a6a5e187e824d2fa0793ce379a202935782d555d6e9d2735677d3", size = 9274, upload-time = "2024-11-06T16:41:39.6Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/fe/ba/e2081de779ca30d473f21f5b30e0e737c438205440784c7dfc81efc2b029/async_timeout-5.0.1-py3-none-any.whl", hash = "sha256:39e3809566ff85354557ec2398b55e096c8364bacac9405a7a1fa429e77fe76c", size = 6233, upload-time = "2024-11-06T16:41:37.9Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "attrs"
|
||||
version = "25.3.0"
|
||||
@@ -1192,7 +1201,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "0.6.3"
|
||||
version = "0.6.4"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -1225,6 +1234,7 @@ dev = [
|
||||
{ name = "pytest-repeat" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "pytest-xdist", extra = ["psutil"] },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
{ name = "syrupy" },
|
||||
{ name = "types-requests" },
|
||||
@@ -1263,6 +1273,7 @@ dev = [
|
||||
{ name = "pytest-repeat" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "pytest-xdist", extras = ["psutil"] },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
{ name = "syrupy" },
|
||||
{ name = "types-requests" },
|
||||
@@ -1326,6 +1337,7 @@ dev = [
|
||||
{ name = "pytest-asyncio" },
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
]
|
||||
|
||||
@@ -1433,7 +1445,7 @@ dev = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-prebuilt"
|
||||
version = "0.6.3"
|
||||
version = "0.6.4"
|
||||
source = { editable = "../prebuilt" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -2628,6 +2640,18 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/51/8b/619a9ee2fa4d3c724fbadde946427735ade64da03894b071bbdc3b789d83/pyzmq-27.0.0-pp39-pypy39_pp73-win_amd64.whl", hash = "sha256:096af9e133fec3a72108ddefba1e42985cb3639e9de52cfd336b6fc23aa083e9", size = 544715, upload-time = "2025-06-13T14:09:05.579Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "redis"
|
||||
version = "6.3.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "async-timeout", marker = "python_full_version < '3.11.3'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/21/cd/030274634a1a052b708756016283ea3d84e91ae45f74d7f5dcf55d753a0f/redis-6.3.0.tar.gz", hash = "sha256:3000dbe532babfb0999cdab7b3e5744bcb23e51923febcfaeb52c8cfb29632ef", size = 4647275, upload-time = "2025-08-05T08:12:31.648Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/df/a7/2fe45801534a187543fc45d28b3844d84559c1589255bc2ece30d92dc205/redis-6.3.0-py3-none-any.whl", hash = "sha256:92f079d656ded871535e099080f70fab8e75273c0236797126ac60242d638e9b", size = 280018, upload-time = "2025-08-05T08:12:30.093Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "referencing"
|
||||
version = "0.36.2"
|
||||
|
||||
@@ -7,11 +7,11 @@ all: help
|
||||
# TESTING AND COVERAGE
|
||||
######################
|
||||
|
||||
start-postgres:
|
||||
docker compose -f tests/compose-postgres.yml up -V --force-recreate --wait --remove-orphans
|
||||
start-services:
|
||||
docker compose -f tests/compose-postgres.yml -f tests/compose-redis.yml up -V --force-recreate --wait --remove-orphans
|
||||
|
||||
stop-postgres:
|
||||
docker compose -f tests/compose-postgres.yml down -v
|
||||
stop-services:
|
||||
docker compose -f tests/compose-postgres.yml -f tests/compose-redis.yml down -v
|
||||
|
||||
TEST ?= .
|
||||
|
||||
@@ -19,15 +19,15 @@ test-fast:
|
||||
LANGGRAPH_TEST_FAST=1 uv run pytest $(TEST)
|
||||
|
||||
test:
|
||||
make start-postgres && LANGGRAPH_TEST_FAST=0 uv run pytest $(TEST); \
|
||||
make start-services && LANGGRAPH_TEST_FAST=0 uv run pytest $(TEST); \
|
||||
EXIT_CODE=$$?; \
|
||||
make stop-postgres; \
|
||||
make stop-services; \
|
||||
exit $$EXIT_CODE
|
||||
|
||||
test_watch:
|
||||
make start-postgres && LANGGRAPH_TEST_FAST=0 uv run ptw $(TEST); \
|
||||
make start-services && LANGGRAPH_TEST_FAST=0 uv run ptw $(TEST); \
|
||||
EXIT_CODE=$$?; \
|
||||
make stop-postgres; \
|
||||
make stop-services; \
|
||||
exit $$EXIT_CODE
|
||||
|
||||
######################
|
||||
|
||||
@@ -1,26 +0,0 @@
|
||||
from typing import Any, Literal, TypedDict
|
||||
|
||||
from langchain_core.messages import ToolCall
|
||||
|
||||
|
||||
class ToolCallWithContext(TypedDict):
|
||||
"""ToolCall with additional context for graph state.
|
||||
|
||||
This is an internal data-structure meant to help the ToolNode accept
|
||||
tools calls with additional context (e.g. state) when dispatched using the
|
||||
`Send` API.
|
||||
|
||||
The Send API is used in create_react_agent to be able to distribute the tool
|
||||
calls in parallel and support human-in-the-loop workflows where graph execution
|
||||
may be paused for an indefinite time.
|
||||
"""
|
||||
|
||||
tool_call: ToolCall
|
||||
__type: Literal["tool_call_with_context"]
|
||||
"""Type to parameterize the payload.
|
||||
|
||||
Using "__" as a prefix to be defensive against potential name collisions with
|
||||
regular user state.
|
||||
"""
|
||||
state: Any
|
||||
"""The state is provided as additional context."""
|
||||
@@ -12,6 +12,7 @@ from typing import (
|
||||
cast,
|
||||
get_type_hints,
|
||||
)
|
||||
from operator import add
|
||||
from warnings import warn
|
||||
|
||||
from langchain_core.language_models import (
|
||||
@@ -43,11 +44,10 @@ from langgraph.graph import END, StateGraph
|
||||
from langgraph.graph.message import add_messages
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
from langgraph.managed import RemainingSteps
|
||||
from langgraph.prebuilt._internal import ToolCallWithContext
|
||||
from langgraph.prebuilt.tool_node import ToolNode
|
||||
from langgraph.runtime import Runtime
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.types import Checkpointer, Send
|
||||
from langgraph.types import Checkpointer, Command, Send
|
||||
from langgraph.typing import ContextT
|
||||
from langgraph.warnings import LangGraphDeprecatedSinceV10
|
||||
|
||||
@@ -67,6 +67,8 @@ class AgentState(TypedDict):
|
||||
|
||||
remaining_steps: NotRequired[RemainingSteps]
|
||||
|
||||
model_calls: Annotated[NotRequired[int], add]
|
||||
|
||||
|
||||
class AgentStatePydantic(BaseModel):
|
||||
"""The state of the agent."""
|
||||
@@ -192,29 +194,6 @@ def _should_bind_tools(
|
||||
return False
|
||||
|
||||
|
||||
def _get_model(model: LanguageModelLike) -> BaseChatModel:
|
||||
"""Get the underlying model from a RunnableBinding or return the model itself."""
|
||||
if isinstance(model, RunnableSequence):
|
||||
model = next(
|
||||
(
|
||||
step
|
||||
for step in model.steps
|
||||
if isinstance(step, (RunnableBinding, BaseChatModel))
|
||||
),
|
||||
model,
|
||||
)
|
||||
|
||||
if isinstance(model, RunnableBinding):
|
||||
model = model.bound
|
||||
|
||||
if not isinstance(model, BaseChatModel):
|
||||
raise TypeError(
|
||||
f"Expected `model` to be a ChatModel or RunnableBinding (e.g. model.bind_tools(...)), got {type(model)}"
|
||||
)
|
||||
|
||||
return model
|
||||
|
||||
|
||||
def _validate_chat_history(
|
||||
messages: Sequence[BaseMessage],
|
||||
) -> None:
|
||||
@@ -246,6 +225,10 @@ def _validate_chat_history(
|
||||
raise ValueError(error_message)
|
||||
|
||||
|
||||
class StepCountIs(BaseModel):
|
||||
count: int
|
||||
|
||||
|
||||
def create_react_agent(
|
||||
model: Union[
|
||||
str,
|
||||
@@ -277,6 +260,7 @@ def create_react_agent(
|
||||
debug: bool = False,
|
||||
version: Literal["v1", "v2"] = "v2",
|
||||
name: Optional[str] = None,
|
||||
stop_when: Optional[Callable[[StateSchema], bool]] | StepCountIs = None,
|
||||
**deprecated_kwargs: Any,
|
||||
) -> CompiledStateGraph:
|
||||
"""Creates an agent graph that calls tools in a loop until a stopping condition is met.
|
||||
@@ -474,6 +458,11 @@ def create_react_agent(
|
||||
if context_schema is None:
|
||||
context_schema = config_schema
|
||||
|
||||
if len(deprecated_kwargs) > 0:
|
||||
raise TypeError(
|
||||
f"create_react_agent() got unexpected keyword arguments: {deprecated_kwargs}"
|
||||
)
|
||||
|
||||
if version not in ("v1", "v2"):
|
||||
raise ValueError(
|
||||
f"Invalid version {version}. Supported versions are 'v1' and 'v2'."
|
||||
@@ -495,13 +484,19 @@ def create_react_agent(
|
||||
else AgentState
|
||||
)
|
||||
|
||||
structured_output_tools: list[type] = []
|
||||
llm_builtin_tools: list[dict] = []
|
||||
if isinstance(tools, ToolNode):
|
||||
tool_classes = list(tools.tools_by_name.values())
|
||||
tool_node = tools
|
||||
else:
|
||||
llm_builtin_tools = [t for t in tools if isinstance(t, dict)]
|
||||
tool_node = ToolNode([t for t in tools if not isinstance(t, dict)])
|
||||
structured_output_tools = (
|
||||
[response_format] if response_format is not None else []
|
||||
)
|
||||
tool_node = ToolNode(
|
||||
[t for t in [*tools, *structured_output_tools] if not isinstance(t, dict)]
|
||||
)
|
||||
tool_classes = list(tool_node.tools_by_name.values())
|
||||
|
||||
is_dynamic_model = not isinstance(model, (str, Runnable)) and callable(model)
|
||||
@@ -523,12 +518,19 @@ def create_react_agent(
|
||||
|
||||
model = cast(BaseChatModel, init_chat_model(model))
|
||||
|
||||
# Add structured output tool if response_format is provided
|
||||
structured_output_tools = (
|
||||
[response_format] if response_format is not None else []
|
||||
)
|
||||
|
||||
if (
|
||||
_should_bind_tools(model, tool_classes, num_builtin=len(llm_builtin_tools)) # type: ignore[arg-type]
|
||||
and len(tool_classes + llm_builtin_tools) > 0
|
||||
):
|
||||
model = cast(BaseChatModel, model).bind_tools(
|
||||
tool_classes + llm_builtin_tools # type: ignore[operator]
|
||||
tool_classes + llm_builtin_tools + structured_output_tools, # type: ignore[operator],
|
||||
tool_choice="any",
|
||||
parallel_tool_calls=False,
|
||||
)
|
||||
|
||||
static_model: Optional[Runnable] = _get_prompt_runnable(prompt) | model # type: ignore[operator]
|
||||
@@ -538,7 +540,11 @@ def create_react_agent(
|
||||
|
||||
# If any of the tools are configured to return_directly after running,
|
||||
# our graph needs to check if these were called
|
||||
should_return_direct = {t.name for t in tool_classes if t.return_direct}
|
||||
should_return_direct = {
|
||||
t.name for t in tool_classes if getattr(t, "return_direct", False)
|
||||
}
|
||||
if response_format is not None:
|
||||
should_return_direct.add(response_format.__name__)
|
||||
|
||||
def _resolve_model(
|
||||
state: StateSchema, runtime: Runtime[ContextT]
|
||||
@@ -604,7 +610,7 @@ def create_react_agent(
|
||||
# Define the function that calls the model
|
||||
def call_model(
|
||||
state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig
|
||||
) -> StateSchema:
|
||||
) -> dict[str, list[AIMessage]] | Command:
|
||||
if is_async_dynamic_model:
|
||||
msg = (
|
||||
"Async model callable provided but agent invoked synchronously. "
|
||||
@@ -615,11 +621,25 @@ def create_react_agent(
|
||||
|
||||
model_input = _get_model_input_state(state)
|
||||
|
||||
if stop_when is not None:
|
||||
post_model_node = "post_model_hook" if post_model_hook is not None else END
|
||||
if isinstance(stop_when, StepCountIs):
|
||||
|
||||
if (model_calls := _get_state_value(state, "model_calls", 0)) == (stop_when.count - 1):
|
||||
# set tool_choice to structured output tool if response_format is provided
|
||||
# though we don't currently expose support for that binding here.
|
||||
...
|
||||
elif model_calls == stop_when.count:
|
||||
return Command(goto=post_model_node)
|
||||
else:
|
||||
if stop_when(state):
|
||||
return Command(goto=post_model_node)
|
||||
|
||||
if is_dynamic_model:
|
||||
# Resolve dynamic model at runtime and apply prompt
|
||||
dynamic_model = _resolve_model(state, runtime)
|
||||
response = cast(AIMessage, dynamic_model.invoke(model_input, config)) # type: ignore[arg-type]
|
||||
else:
|
||||
else:
|
||||
response = cast(AIMessage, static_model.invoke(model_input, config)) # type: ignore[union-attr]
|
||||
|
||||
# add agent name to the AIMessage
|
||||
@@ -634,12 +654,12 @@ def create_react_agent(
|
||||
)
|
||||
]
|
||||
}
|
||||
# We return a list, because this will get added to the existing list
|
||||
return {"messages": [response]}
|
||||
|
||||
return {"messages": [response], "model_calls": 1}
|
||||
|
||||
async def acall_model(
|
||||
state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig
|
||||
) -> StateSchema:
|
||||
) -> dict[str, list[AIMessage]]:
|
||||
model_input = _get_model_input_state(state)
|
||||
|
||||
if is_dynamic_model:
|
||||
@@ -685,49 +705,6 @@ def create_react_agent(
|
||||
else:
|
||||
input_schema = state_schema
|
||||
|
||||
def generate_structured_response(
|
||||
state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig
|
||||
) -> StateSchema:
|
||||
if is_async_dynamic_model:
|
||||
msg = (
|
||||
"Async model callable provided but agent invoked synchronously. "
|
||||
"Use agent.ainvoke() or agent.astream(), or provide a sync model callable."
|
||||
)
|
||||
raise RuntimeError(msg)
|
||||
|
||||
messages = _get_state_value(state, "messages")
|
||||
structured_response_schema = response_format
|
||||
if isinstance(response_format, tuple):
|
||||
system_prompt, structured_response_schema = response_format
|
||||
messages = [SystemMessage(content=system_prompt)] + list(messages)
|
||||
|
||||
resolved_model = _resolve_model(state, runtime)
|
||||
model_with_structured_output = _get_model(
|
||||
resolved_model
|
||||
).with_structured_output(
|
||||
cast(StructuredResponseSchema, structured_response_schema)
|
||||
)
|
||||
response = model_with_structured_output.invoke(messages, config)
|
||||
return {"structured_response": response}
|
||||
|
||||
async def agenerate_structured_response(
|
||||
state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig
|
||||
) -> StateSchema:
|
||||
messages = _get_state_value(state, "messages")
|
||||
structured_response_schema = response_format
|
||||
if isinstance(response_format, tuple):
|
||||
system_prompt, structured_response_schema = response_format
|
||||
messages = [SystemMessage(content=system_prompt)] + list(messages)
|
||||
|
||||
resolved_model = await _aresolve_model(state, runtime)
|
||||
model_with_structured_output = _get_model(
|
||||
resolved_model
|
||||
).with_structured_output(
|
||||
cast(StructuredResponseSchema, structured_response_schema)
|
||||
)
|
||||
response = await model_with_structured_output.ainvoke(messages, config)
|
||||
return {"structured_response": response}
|
||||
|
||||
if not tool_calling_enabled:
|
||||
# Define a new graph
|
||||
workflow = StateGraph(state_schema=state_schema, context_schema=context_schema)
|
||||
@@ -749,19 +726,6 @@ def create_react_agent(
|
||||
workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type]
|
||||
workflow.add_edge("agent", "post_model_hook")
|
||||
|
||||
if response_format is not None:
|
||||
workflow.add_node(
|
||||
"generate_structured_response",
|
||||
RunnableCallable(
|
||||
generate_structured_response,
|
||||
agenerate_structured_response,
|
||||
),
|
||||
)
|
||||
if post_model_hook is not None:
|
||||
workflow.add_edge("post_model_hook", "generate_structured_response")
|
||||
else:
|
||||
workflow.add_edge("agent", "generate_structured_response")
|
||||
|
||||
return workflow.compile(
|
||||
checkpointer=checkpointer,
|
||||
store=store,
|
||||
@@ -773,16 +737,14 @@ def create_react_agent(
|
||||
|
||||
# Define the function that determines whether to continue or not
|
||||
def should_continue(state: StateSchema) -> Union[str, list[Send]]:
|
||||
|
||||
post_model_node = "post_model_hook" if post_model_hook is not None else END
|
||||
|
||||
messages = _get_state_value(state, "messages")
|
||||
last_message = messages[-1]
|
||||
# If there is no function call, then we finish
|
||||
if not isinstance(last_message, AIMessage) or not last_message.tool_calls:
|
||||
if post_model_hook is not None:
|
||||
return "post_model_hook"
|
||||
elif response_format is not None:
|
||||
return "generate_structured_response"
|
||||
else:
|
||||
return END
|
||||
return post_model_node
|
||||
# Otherwise if there is, we continue
|
||||
else:
|
||||
if version == "v1":
|
||||
@@ -790,17 +752,11 @@ def create_react_agent(
|
||||
elif version == "v2":
|
||||
if post_model_hook is not None:
|
||||
return "post_model_hook"
|
||||
return [
|
||||
Send(
|
||||
"tools",
|
||||
ToolCallWithContext(
|
||||
__type="tool_call_with_context",
|
||||
tool_call=tool_call,
|
||||
state=state,
|
||||
),
|
||||
)
|
||||
for tool_call in last_message.tool_calls
|
||||
tool_calls = [
|
||||
tool_node.inject_tool_args(call, state, store) # type: ignore[arg-type]
|
||||
for call in last_message.tool_calls
|
||||
]
|
||||
return [Send("tools", [tool_call]) for tool_call in tool_calls]
|
||||
|
||||
# Define a new graph
|
||||
workflow = StateGraph(
|
||||
@@ -827,45 +783,19 @@ def create_react_agent(
|
||||
# Set the entrypoint as `agent`
|
||||
# This means that this node is the first one called
|
||||
workflow.set_entry_point(entrypoint)
|
||||
|
||||
agent_paths = []
|
||||
post_model_hook_paths = [entrypoint, "tools"]
|
||||
|
||||
# Add a post model hook node if post_model_hook is provided
|
||||
if post_model_hook is not None:
|
||||
workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type]
|
||||
agent_paths.append("post_model_hook")
|
||||
workflow.add_edge("agent", "post_model_hook")
|
||||
else:
|
||||
agent_paths.append("tools")
|
||||
|
||||
# Add a structured output node if response_format is provided
|
||||
if response_format is not None:
|
||||
workflow.add_node(
|
||||
"generate_structured_response",
|
||||
RunnableCallable(
|
||||
generate_structured_response,
|
||||
agenerate_structured_response,
|
||||
),
|
||||
)
|
||||
if post_model_hook is not None:
|
||||
post_model_hook_paths.append("generate_structured_response")
|
||||
else:
|
||||
agent_paths.append("generate_structured_response")
|
||||
else:
|
||||
if post_model_hook is not None:
|
||||
post_model_hook_paths.append(END)
|
||||
else:
|
||||
agent_paths.append(END)
|
||||
|
||||
if post_model_hook is not None:
|
||||
|
||||
def post_model_hook_router(state: StateSchema) -> Union[str, list[Send]]:
|
||||
"""Route to the next node after post_model_hook.
|
||||
|
||||
Routes to one of:
|
||||
* "tools": if there are pending tool calls without a corresponding message.
|
||||
* "generate_structured_response": if no pending tool calls exist and response_format is specified.
|
||||
* END: if no pending tool calls exist and no response_format is specified.
|
||||
"""
|
||||
|
||||
@@ -881,37 +811,32 @@ def create_react_agent(
|
||||
]
|
||||
|
||||
if pending_tool_calls:
|
||||
return [
|
||||
Send(
|
||||
"tools",
|
||||
ToolCallWithContext(
|
||||
__type="tool_call_with_context",
|
||||
tool_call=tool_call,
|
||||
state=state,
|
||||
),
|
||||
)
|
||||
for tool_call in pending_tool_calls
|
||||
pending_tool_calls = [
|
||||
tool_node.inject_tool_args(call, state, store) # type: ignore[arg-type]
|
||||
for call in pending_tool_calls
|
||||
]
|
||||
return [Send("tools", [tool_call]) for tool_call in pending_tool_calls]
|
||||
elif isinstance(messages[-1], ToolMessage):
|
||||
return entrypoint
|
||||
elif response_format is not None:
|
||||
return "generate_structured_response"
|
||||
else:
|
||||
return END
|
||||
|
||||
workflow.add_conditional_edges(
|
||||
"post_model_hook",
|
||||
post_model_hook_router, # type: ignore[arg-type]
|
||||
path_map=post_model_hook_paths,
|
||||
post_model_hook_router,
|
||||
path_map=[entrypoint, "tools", END],
|
||||
)
|
||||
else:
|
||||
agent_paths.append("tools")
|
||||
agent_paths.append(END)
|
||||
|
||||
workflow.add_conditional_edges(
|
||||
"agent",
|
||||
should_continue, # type: ignore[arg-type]
|
||||
should_continue,
|
||||
path_map=agent_paths,
|
||||
)
|
||||
|
||||
def route_tool_responses(state: StateSchema) -> str:
|
||||
def route_tool_responses(state: StateSchema) -> str | Command:
|
||||
for m in reversed(_get_state_value(state, "messages")):
|
||||
if not isinstance(m, ToolMessage):
|
||||
break
|
||||
|
||||
@@ -74,7 +74,6 @@ from typing_extensions import Annotated, get_args, get_origin
|
||||
from langgraph._internal._runnable import RunnableCallable
|
||||
from langgraph.errors import GraphBubbleUp
|
||||
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
||||
from langgraph.prebuilt._internal import ToolCallWithContext
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.types import Command, Send
|
||||
|
||||
@@ -341,11 +340,19 @@ class ToolNode(RunnableCallable):
|
||||
self.tools_by_name: dict[str, BaseTool] = {}
|
||||
self.tool_to_state_args: dict[str, dict[str, Optional[str]]] = {}
|
||||
self.tool_to_store_arg: dict[str, Optional[str]] = {}
|
||||
self.structured_output_tools: list[str] = []
|
||||
self.handle_tool_errors = handle_tool_errors
|
||||
self.messages_key = messages_key
|
||||
for tool_ in tools:
|
||||
if not isinstance(tool_, BaseTool):
|
||||
if issubclass(tool_, BaseModel):
|
||||
self.tools_by_name[tool_.__name__] = tool_
|
||||
self.tool_to_state_args[tool_.__name__] = {}
|
||||
self.tool_to_store_arg[tool_.__name__] = None
|
||||
self.structured_output_tools.append(tool_.__name__)
|
||||
continue
|
||||
elif not isinstance(tool_, BaseTool):
|
||||
tool_ = create_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_)
|
||||
@@ -361,8 +368,7 @@ class ToolNode(RunnableCallable):
|
||||
*,
|
||||
store: Optional[BaseStore],
|
||||
) -> Any:
|
||||
tool_calls, input_type = self._parse_input(input)
|
||||
tool_calls = [self.inject_tool_args(call, input, store) for call in tool_calls]
|
||||
tool_calls, input_type = self._parse_input(input, store)
|
||||
config_list = get_config_list(config, len(tool_calls))
|
||||
input_types = [input_type] * len(tool_calls)
|
||||
with get_executor_for_config(config) as executor:
|
||||
@@ -383,8 +389,7 @@ class ToolNode(RunnableCallable):
|
||||
*,
|
||||
store: Optional[BaseStore],
|
||||
) -> Any:
|
||||
tool_calls, input_type = self._parse_input(input)
|
||||
tool_calls = [self.inject_tool_args(call, input, store) for call in tool_calls]
|
||||
tool_calls, input_type = self._parse_input(input, store)
|
||||
outputs = await asyncio.gather(
|
||||
*(self._arun_one(call, input_type, config) for call in tool_calls)
|
||||
)
|
||||
@@ -440,11 +445,26 @@ class ToolNode(RunnableCallable):
|
||||
call: ToolCall,
|
||||
input_type: Literal["list", "dict", "tool_calls"],
|
||||
config: RunnableConfig,
|
||||
) -> ToolMessage:
|
||||
) -> ToolMessage | Command:
|
||||
"""Run a single tool call synchronously."""
|
||||
if invalid_tool_message := self._validate_tool_call(call):
|
||||
return invalid_tool_message
|
||||
try:
|
||||
if call["name"] in self.structured_output_tools:
|
||||
response_schema = self.tools_by_name[call["name"]]
|
||||
return Command(
|
||||
update={
|
||||
"messages": [
|
||||
ToolMessage(
|
||||
content="structured output generated",
|
||||
name="structured_output",
|
||||
tool_call_id=call["id"],
|
||||
status="success",
|
||||
),
|
||||
],
|
||||
"structured_response": response_schema(**call["args"]),
|
||||
}
|
||||
)
|
||||
call_args = {**call, **{"type": "tool_call"}}
|
||||
response = self.tools_by_name[call["name"]].invoke(call_args, config)
|
||||
|
||||
@@ -502,13 +522,14 @@ class ToolNode(RunnableCallable):
|
||||
return invalid_tool_message
|
||||
|
||||
try:
|
||||
input = {**call, **{"type": "tool_call"}}
|
||||
response = await self.tools_by_name[call["name"]].ainvoke(input, config)
|
||||
call_args = {**call, **{"type": "tool_call"}}
|
||||
response = await self.tools_by_name[call["name"]].ainvoke(call_args, config)
|
||||
|
||||
# GraphInterrupt is a special exception that will always be raised.
|
||||
# It can be triggered in the following scenarios:
|
||||
# (1) a NodeInterrupt is raised inside a tool
|
||||
# (2) a NodeInterrupt is raised inside a graph node for a graph called as a tool
|
||||
# It can be triggered in the following scenarios,
|
||||
# Where GraphInterrupt(GraphBubbleUp) is raised from an `interrupt` invocation most commonly:
|
||||
# (1) a GraphInterrupt is raised inside a tool
|
||||
# (2) a GraphInterrupt is raised inside a graph node for a graph called as a tool
|
||||
# (3) a GraphInterrupt is raised when a subgraph is interrupted inside a graph called as a tool
|
||||
# (2 and 3 can happen in a "supervisor w/ tools" multi-agent architecture)
|
||||
except GraphBubbleUp as e:
|
||||
@@ -555,6 +576,7 @@ class ToolNode(RunnableCallable):
|
||||
dict[str, Any],
|
||||
BaseModel,
|
||||
],
|
||||
store: Optional[BaseStore],
|
||||
) -> Tuple[list[ToolCall], Literal["list", "dict", "tool_calls"]]:
|
||||
input_type: Literal["list", "dict", "tool_calls"]
|
||||
if isinstance(input, list):
|
||||
@@ -565,15 +587,6 @@ class ToolNode(RunnableCallable):
|
||||
else:
|
||||
input_type = "list"
|
||||
messages = input
|
||||
elif (
|
||||
isinstance(input, dict) and input.get("__type") == "tool_call_with_context"
|
||||
):
|
||||
# mypy will not be able to type narrow correctly since the signature
|
||||
# for input contains dict[str, Any]. We'd need to type dict[str, Any]
|
||||
# before we can apply correct typing.
|
||||
input = cast(ToolCallWithContext, input) # type: ignore[assignment]
|
||||
input_type = "tool_calls"
|
||||
return [input["tool_call"]], input_type
|
||||
elif isinstance(input, dict) and (messages := input.get(self.messages_key, [])):
|
||||
input_type = "dict"
|
||||
elif messages := getattr(input, self.messages_key, []):
|
||||
@@ -589,7 +602,10 @@ class ToolNode(RunnableCallable):
|
||||
except StopIteration:
|
||||
raise ValueError("No AIMessage found in input")
|
||||
|
||||
tool_calls = [call for call in latest_ai_message.tool_calls]
|
||||
tool_calls = [
|
||||
self.inject_tool_args(call, input, store)
|
||||
for call in latest_ai_message.tool_calls
|
||||
]
|
||||
return tool_calls, input_type
|
||||
|
||||
def _validate_tool_call(self, call: ToolCall) -> Optional[ToolMessage]:
|
||||
@@ -632,19 +648,14 @@ class ToolNode(RunnableCallable):
|
||||
err_msg += f" State should contain fields {required_fields_str}."
|
||||
raise ValueError(err_msg)
|
||||
|
||||
if isinstance(input, dict) and input.get("__type") == "tool_call_with_context":
|
||||
state = input["state"]
|
||||
else:
|
||||
state = input
|
||||
|
||||
if isinstance(state, dict):
|
||||
if isinstance(input, dict):
|
||||
tool_state_args = {
|
||||
tool_arg: state[state_field] if state_field else state
|
||||
tool_arg: input[state_field] if state_field else input
|
||||
for tool_arg, state_field in state_args.items()
|
||||
}
|
||||
else:
|
||||
tool_state_args = {
|
||||
tool_arg: getattr(state, state_field) if state_field else state
|
||||
tool_arg: getattr(input, state_field) if state_field else input
|
||||
for tool_arg, state_field in state_args.items()
|
||||
}
|
||||
|
||||
@@ -802,7 +813,6 @@ def tools_condition(
|
||||
|
||||
Args:
|
||||
state: The current graph state to examine for tool calls. Supported formats:
|
||||
- List of messages (for MessageGraph)
|
||||
- Dictionary containing a messages key (for StateGraph)
|
||||
- BaseModel instance with a messages attribute
|
||||
messages_key: The key or attribute name containing the message list in the state.
|
||||
|
||||
@@ -2,8 +2,7 @@
|
||||
in a langchain graph. It applies a pydantic schema to tool_calls in the models' outputs,
|
||||
and returns a ToolMessage with the validated content. If the schema is not valid, it
|
||||
returns a ToolMessage with the error message. The ValidationNode can be used in a
|
||||
StateGraph with a "messages" key or in a MessageGraph. If multiple tool calls are
|
||||
requested, they will be run in parallel.
|
||||
StateGraph with a "messages" key. If multiple tool calls are requested, they will be run in parallel.
|
||||
"""
|
||||
|
||||
from typing import (
|
||||
@@ -49,7 +48,7 @@ def _default_format_error(
|
||||
class ValidationNode(RunnableCallable):
|
||||
"""A node that validates all tools requests from the last AIMessage.
|
||||
|
||||
It can be used either in StateGraph with a "messages" key or in MessageGraph.
|
||||
It can be used either in StateGraph with a "messages" key.
|
||||
|
||||
!!! note
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph-prebuilt"
|
||||
version = "0.6.3"
|
||||
version = "0.6.4"
|
||||
description = "Library with high-level APIs for creating and executing LangGraph agents and tools."
|
||||
authors = []
|
||||
requires-python = ">=3.9"
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
name: langgraph-tests-redis
|
||||
services:
|
||||
redis-test:
|
||||
image: redis:7-alpine
|
||||
ports:
|
||||
- "6379:6379"
|
||||
command: redis-server --maxmemory 256mb --maxmemory-policy allkeys-lru
|
||||
healthcheck:
|
||||
test: redis-cli ping
|
||||
start_period: 10s
|
||||
timeout: 1s
|
||||
retries: 5
|
||||
interval: 5s
|
||||
start_interval: 1s
|
||||
tmpfs:
|
||||
- /data # Use tmpfs for faster testing
|
||||
@@ -31,3 +31,11 @@ def test_config_schema_deprecation() -> None:
|
||||
match="`get_config_jsonschema` is deprecated. Use `get_context_jsonschema` instead.",
|
||||
):
|
||||
assert agent.get_config_jsonschema() is not None
|
||||
|
||||
|
||||
def test_extra_kwargs_deprecation() -> None:
|
||||
with pytest.raises(
|
||||
TypeError,
|
||||
match="create_react_agent\(\) got unexpected keyword arguments: \{'extra': 'extra'\}",
|
||||
):
|
||||
create_react_agent(FakeToolCallingModel(), [], extra="extra")
|
||||
|
||||
@@ -24,7 +24,7 @@ from langchain_core.messages import (
|
||||
ToolCall,
|
||||
ToolMessage,
|
||||
)
|
||||
from langchain_core.runnables import RunnableLambda
|
||||
from langchain_core.runnables import RunnableConfig, RunnableLambda
|
||||
from langchain_core.tools import InjectedToolCallId, ToolException
|
||||
from langchain_core.tools import tool as dec_tool
|
||||
from pydantic import BaseModel, Field
|
||||
@@ -1236,6 +1236,190 @@ def test_tool_node_stream_writer() -> None:
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_react_agent_subgraph_streaming_sync(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test React agent streaming when used as a subgraph node sync version"""
|
||||
|
||||
@dec_tool
|
||||
def get_weather(city: str) -> str:
|
||||
"""Get the weather of a city."""
|
||||
return f"The weather of {city} is sunny."
|
||||
|
||||
# Create a React agent
|
||||
model = FakeToolCallingModel(
|
||||
tool_calls=[
|
||||
[{"args": {"city": "Tokyo"}, "id": "1", "name": "get_weather"}],
|
||||
[],
|
||||
]
|
||||
)
|
||||
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
tools=[get_weather],
|
||||
prompt="You are a helpful travel assistant.",
|
||||
version=version,
|
||||
)
|
||||
|
||||
# Create a subgraph that uses the React agent as a node
|
||||
def react_agent_node(state: MessagesState, config: RunnableConfig) -> MessagesState:
|
||||
"""Node that runs the React agent and collects streaming output."""
|
||||
collected_content = ""
|
||||
|
||||
# Stream the agent output and collect content
|
||||
for msg_chunk, msg_metadata in agent.stream(
|
||||
{"messages": [("user", state["messages"][-1].content)]},
|
||||
config,
|
||||
stream_mode="messages",
|
||||
):
|
||||
if hasattr(msg_chunk, "content") and msg_chunk.content:
|
||||
collected_content += msg_chunk.content
|
||||
|
||||
return {"messages": [("assistant", collected_content)]}
|
||||
|
||||
# Create the main workflow with the React agent as a subgraph node
|
||||
workflow = StateGraph(MessagesState)
|
||||
workflow.add_node("react_agent", react_agent_node)
|
||||
workflow.add_edge(START, "react_agent")
|
||||
workflow.add_edge("react_agent", "__end__")
|
||||
compiled_workflow = workflow.compile()
|
||||
|
||||
# Test the streaming functionality
|
||||
result = compiled_workflow.invoke(
|
||||
{"messages": [("user", "What is the weather in Tokyo?")]}
|
||||
)
|
||||
|
||||
# Verify the result contains expected structure
|
||||
assert len(result["messages"]) == 2
|
||||
assert result["messages"][0].content == "What is the weather in Tokyo?"
|
||||
assert "assistant" in str(result["messages"][1])
|
||||
|
||||
# Test streaming with subgraphs = True
|
||||
result = compiled_workflow.invoke(
|
||||
{"messages": [("user", "What is the weather in Tokyo?")]},
|
||||
subgraphs=True,
|
||||
)
|
||||
assert len(result["messages"]) == 2
|
||||
|
||||
events = []
|
||||
for event in compiled_workflow.stream(
|
||||
{"messages": [("user", "What is the weather in Tokyo?")]},
|
||||
stream_mode="messages",
|
||||
subgraphs=False,
|
||||
):
|
||||
events.append(event)
|
||||
|
||||
assert len(events) == 0
|
||||
|
||||
events = []
|
||||
for event in compiled_workflow.stream(
|
||||
{"messages": [("user", "What is the weather in Tokyo?")]},
|
||||
stream_mode="messages",
|
||||
subgraphs=True,
|
||||
):
|
||||
events.append(event)
|
||||
|
||||
assert len(events) == 3
|
||||
namespace, (msg, metadata) = events[0]
|
||||
# FakeToolCallingModel returns a single AIMessage with tool calls
|
||||
# The content of the AIMessage reflects the input message
|
||||
assert msg.content.startswith("You are a helpful travel assistant")
|
||||
namespace, (msg, metadata) = events[1] # ToolMessage
|
||||
assert msg.content.startswith("The weather of Tokyo is sunny.")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
async def test_react_agent_subgraph_streaming(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test React agent streaming when used as a subgraph node."""
|
||||
|
||||
@dec_tool
|
||||
def get_weather(city: str) -> str:
|
||||
"""Get the weather of a city."""
|
||||
return f"The weather of {city} is sunny."
|
||||
|
||||
# Create a React agent
|
||||
model = FakeToolCallingModel(
|
||||
tool_calls=[
|
||||
[{"args": {"city": "Tokyo"}, "id": "1", "name": "get_weather"}],
|
||||
[],
|
||||
]
|
||||
)
|
||||
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
tools=[get_weather],
|
||||
prompt="You are a helpful travel assistant.",
|
||||
version=version,
|
||||
)
|
||||
|
||||
# Create a subgraph that uses the React agent as a node
|
||||
async def react_agent_node(
|
||||
state: MessagesState, config: RunnableConfig
|
||||
) -> MessagesState:
|
||||
"""Node that runs the React agent and collects streaming output."""
|
||||
collected_content = ""
|
||||
|
||||
# Stream the agent output and collect content
|
||||
async for msg_chunk, msg_metadata in agent.astream(
|
||||
{"messages": [("user", state["messages"][-1].content)]},
|
||||
config,
|
||||
stream_mode="messages",
|
||||
):
|
||||
if hasattr(msg_chunk, "content") and msg_chunk.content:
|
||||
collected_content += msg_chunk.content
|
||||
|
||||
return {"messages": [("assistant", collected_content)]}
|
||||
|
||||
# Create the main workflow with the React agent as a subgraph node
|
||||
workflow = StateGraph(MessagesState)
|
||||
workflow.add_node("react_agent", react_agent_node)
|
||||
workflow.add_edge(START, "react_agent")
|
||||
workflow.add_edge("react_agent", "__end__")
|
||||
compiled_workflow = workflow.compile()
|
||||
|
||||
# Test the streaming functionality
|
||||
result = await compiled_workflow.ainvoke(
|
||||
{"messages": [("user", "What is the weather in Tokyo?")]}
|
||||
)
|
||||
|
||||
# Verify the result contains expected structure
|
||||
assert len(result["messages"]) == 2
|
||||
assert result["messages"][0].content == "What is the weather in Tokyo?"
|
||||
assert "assistant" in str(result["messages"][1])
|
||||
|
||||
# Test streaming with subgraphs = True
|
||||
result = await compiled_workflow.ainvoke(
|
||||
{"messages": [("user", "What is the weather in Tokyo?")]},
|
||||
subgraphs=True,
|
||||
)
|
||||
assert len(result["messages"]) == 2
|
||||
|
||||
events = []
|
||||
async for event in compiled_workflow.astream(
|
||||
{"messages": [("user", "What is the weather in Tokyo?")]},
|
||||
stream_mode="messages",
|
||||
subgraphs=False,
|
||||
):
|
||||
events.append(event)
|
||||
|
||||
assert len(events) == 0
|
||||
|
||||
events = []
|
||||
async for event in compiled_workflow.astream(
|
||||
{"messages": [("user", "What is the weather in Tokyo?")]},
|
||||
stream_mode="messages",
|
||||
subgraphs=True,
|
||||
):
|
||||
events.append(event)
|
||||
|
||||
assert len(events) == 3
|
||||
namespace, (msg, metadata) = events[0]
|
||||
# FakeToolCallingModel returns a single AIMessage with tool calls
|
||||
# The content of the AIMessage reflects the input message
|
||||
assert msg.content.startswith("You are a helpful travel assistant")
|
||||
namespace, (msg, metadata) = events[1] # ToolMessage
|
||||
assert msg.content.startswith("The weather of Tokyo is sunny.")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_tool_node_node_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver, version: str
|
||||
|
||||
Generated
+4
-2
@@ -316,7 +316,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "0.6.3"
|
||||
version = "0.6.4"
|
||||
source = { editable = "../langgraph" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -359,6 +359,7 @@ dev = [
|
||||
{ name = "pytest-repeat" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "pytest-xdist", extras = ["psutil"] },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
{ name = "syrupy" },
|
||||
{ name = "types-requests" },
|
||||
@@ -392,6 +393,7 @@ dev = [
|
||||
{ name = "pytest-asyncio" },
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watcher" },
|
||||
{ name = "redis" },
|
||||
{ name = "ruff" },
|
||||
]
|
||||
|
||||
@@ -460,7 +462,7 @@ dev = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-prebuilt"
|
||||
version = "0.6.3"
|
||||
version = "0.6.4"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Reference in New Issue
Block a user