Compare commits

..
Author SHA1 Message Date
Sydney Runkle 9aab70bd73 experimental stop when 2025-08-12 16:19:16 -04:00
Sydney Runkle aaea76b475 hacky solution for now 2025-08-11 19:44:02 -04:00
Sam CrowderandGitHub 16b363fbb0 feat(langgraph): implement redis node level cache (#5834)
###   Description

Adds Redis as a supported cache backend for LangGraph node-level
caching, enabling distributed caching across multiple processes/servers.
This implementation follows the same patterns as existing InMemoryCache
and SqliteCache.

###  Key changes
  - New RedisCache class implementing the BaseCache interface
  - Support for TTL-based expiration and batch operations
  - Worker-specific cache prefixes for parallel test isolation

###  Dependencies

  - redis package (already included in dev dependencies)

### Test Plan

- Unit tests: Added Redis cache tests covering basic operations, TTL,
batch operations, and error handling
- Integration tests: Redis cache integrated into existing LangGraph test
suite, tested with all checkpointer combinations
2025-08-11 09:19:34 -07:00
16 changed files with 667 additions and 145 deletions
+1
View File
@@ -329,6 +329,7 @@ dev = [
{ name = "pytest-asyncio" },
{ name = "pytest-mock" },
{ name = "pytest-watcher" },
{ name = "redis" },
{ name = "ruff" },
]
+1
View File
@@ -341,6 +341,7 @@ dev = [
{ name = "pytest-asyncio" },
{ name = "pytest-mock" },
{ name = "pytest-watcher" },
{ name = "redis" },
{ name = "ruff" },
]
+144
View File
@@ -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)
+1
View File
@@ -32,6 +32,7 @@ dev = [
"numpy",
"pandas",
"pandas-stubs>=2.2.2.240807",
"redis",
]
[tool.hatch.build.targets.wheel]
+313
View File
@@ -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
+23
View File
@@ -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
View File
@@ -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
+1
View File
@@ -49,6 +49,7 @@ dev = [
"types-requests",
"pycryptodome",
"langgraph-cli[inmem]",
"redis",
]
[tool.uv]
+16
View File
@@ -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
+25 -1
View File
@@ -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}")
+24
View File
@@ -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"
@@ -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" },
]
@@ -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"
+8 -8
View File
@@ -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
######################
@@ -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 (
@@ -46,7 +47,7 @@ from langgraph.managed import RemainingSteps
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
@@ -66,6 +67,8 @@ class AgentState(TypedDict):
remaining_steps: NotRequired[RemainingSteps]
model_calls: Annotated[NotRequired[int], add]
class AgentStatePydantic(BaseModel):
"""The state of the agent."""
@@ -191,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:
@@ -245,6 +225,10 @@ def _validate_chat_history(
raise ValueError(error_message)
class StepCountIs(BaseModel):
count: int
def create_react_agent(
model: Union[
str,
@@ -276,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.
@@ -499,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)
@@ -527,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]
@@ -542,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]
@@ -608,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. "
@@ -619,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
@@ -638,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:
@@ -689,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)
@@ -753,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,
@@ -777,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":
@@ -825,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.
"""
@@ -886,16 +818,17 @@ def create_react_agent(
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,
path_map=post_model_hook_paths,
path_map=[entrypoint, "tools", END],
)
else:
agent_paths.append("tools")
agent_paths.append(END)
workflow.add_conditional_edges(
"agent",
@@ -903,7 +836,7 @@ def create_react_agent(
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
+25 -2
View File
@@ -340,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_)
@@ -437,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)
+16
View File
@@ -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
+2
View File
@@ -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" },
]