mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-03 15:05:06 +02:00
Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ac9c0f511d | ||
|
|
8ccead9560 | ||
|
|
3a9749a0ed | ||
|
|
196fbf2631 | ||
|
|
0acd5decf8 | ||
|
|
eec823cbd6 | ||
|
|
9eac53db4b |
@@ -151,6 +151,7 @@ class PoolConfig(TypedDict, total=False):
|
||||
"""Connection pool settings for PostgreSQL connections.
|
||||
|
||||
Controls connection lifecycle and resource utilization:
|
||||
|
||||
- Small pools (1-5) suit low-concurrency workloads
|
||||
- Larger pools handle concurrent requests but consume more resources
|
||||
- Setting max_size prevents resource exhaustion under load
|
||||
@@ -166,6 +167,7 @@ class PoolConfig(TypedDict, total=False):
|
||||
"""Additional connection arguments passed to each connection in the pool.
|
||||
|
||||
Default kwargs set automatically:
|
||||
|
||||
- autocommit: True
|
||||
- prepare_threshold: 0
|
||||
- row_factory: dict_row
|
||||
@@ -656,7 +658,8 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
|
||||
item = store.get(("users", "123"), "prefs")
|
||||
```
|
||||
|
||||
Or using the convenient from_conn_string helper:
|
||||
Or using the convenient `from_conn_string` helper:
|
||||
|
||||
```python
|
||||
from langgraph.store.postgres import PostgresStore
|
||||
|
||||
|
||||
@@ -722,7 +722,8 @@ class SqliteStore(BaseSqliteStore, BaseStore):
|
||||
item = store.get(("users", "123"), "prefs")
|
||||
```
|
||||
|
||||
Or using the convenient from_conn_string helper:
|
||||
Or using the convenient `from_conn_string` helper:
|
||||
|
||||
```python
|
||||
from langgraph.store.sqlite import SqliteStore
|
||||
|
||||
|
||||
@@ -76,6 +76,7 @@ from langgraph.types import (
|
||||
CachePolicy,
|
||||
Checkpointer,
|
||||
Command,
|
||||
OnInterruptHook,
|
||||
RetryPolicy,
|
||||
Send,
|
||||
ensure_valid_checkpointer,
|
||||
@@ -831,6 +832,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
|
||||
interrupt_after: All | list[str] | None = None,
|
||||
debug: bool = False,
|
||||
name: str | None = None,
|
||||
on_interrupt: OnInterruptHook | None = None,
|
||||
) -> CompiledStateGraph[StateT, ContextT, InputT, OutputT]:
|
||||
"""Compiles the `StateGraph` into a `CompiledStateGraph` object.
|
||||
|
||||
@@ -850,6 +852,10 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
|
||||
interrupt_after: An optional list of node names to interrupt after.
|
||||
debug: A flag indicating whether to enable debug mode.
|
||||
name: The name to use for the compiled graph.
|
||||
on_interrupt: An optional callback that is invoked whenever the graph
|
||||
execution is interrupted. Called with the list of `Interrupt` objects.
|
||||
|
||||
May be a sync function or an async coroutine function.
|
||||
|
||||
Returns:
|
||||
CompiledStateGraph: The compiled `StateGraph`.
|
||||
@@ -910,6 +916,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
|
||||
store=store,
|
||||
cache=cache,
|
||||
name=name or "LangGraph",
|
||||
on_interrupt=on_interrupt,
|
||||
)
|
||||
|
||||
compiled.attach_node(START, None)
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import binascii
|
||||
import concurrent.futures
|
||||
import warnings
|
||||
from collections import defaultdict, deque
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from contextlib import (
|
||||
@@ -12,7 +13,7 @@ from contextlib import (
|
||||
ExitStack,
|
||||
)
|
||||
from datetime import datetime, timezone
|
||||
from inspect import signature
|
||||
from inspect import iscoroutinefunction, signature
|
||||
from types import TracebackType
|
||||
from typing import (
|
||||
Any,
|
||||
@@ -115,6 +116,8 @@ from langgraph.types import (
|
||||
CachePolicy,
|
||||
Command,
|
||||
Durability,
|
||||
Interrupt,
|
||||
OnInterruptHook,
|
||||
PregelExecutableTask,
|
||||
RetryPolicy,
|
||||
Send,
|
||||
@@ -157,6 +160,7 @@ class PregelLoop:
|
||||
manager: None | AsyncParentRunManager | ParentRunManager
|
||||
interrupt_after: All | Sequence[str]
|
||||
interrupt_before: All | Sequence[str]
|
||||
on_interrupt: OnInterruptHook | None
|
||||
durability: Durability
|
||||
retry_policy: Sequence[RetryPolicy]
|
||||
cache_policy: CachePolicy | None
|
||||
@@ -226,6 +230,7 @@ class PregelLoop:
|
||||
migrate_checkpoint: Callable[[Checkpoint], None] | None = None,
|
||||
retry_policy: Sequence[RetryPolicy] = (),
|
||||
cache_policy: CachePolicy | None = None,
|
||||
on_interrupt: OnInterruptHook | None = None,
|
||||
) -> None:
|
||||
self.stream = stream
|
||||
self.config = config
|
||||
@@ -242,6 +247,7 @@ class PregelLoop:
|
||||
self.stream_keys = stream_keys
|
||||
self.interrupt_after = interrupt_after
|
||||
self.interrupt_before = interrupt_before
|
||||
self.on_interrupt = on_interrupt
|
||||
self.manager = manager
|
||||
self.is_nested = CONFIG_KEY_TASK_ID in self.config.get(CONF, {})
|
||||
self.skip_done_tasks = CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
|
||||
@@ -865,14 +871,36 @@ class PregelLoop:
|
||||
[{INTERRUPT: cast(GraphInterrupt, exc_value).args[0]}]
|
||||
),
|
||||
)
|
||||
# save final output
|
||||
# save final output first, so graph state is consistent even
|
||||
# if the on_interrupt hook raises
|
||||
self.output = read_channels(self.channels, self.output_keys)
|
||||
# call on_interrupt hook
|
||||
if self.on_interrupt is not None:
|
||||
interrupts: list[Interrupt] = (
|
||||
list(cast(GraphInterrupt, exc_value).args[0])
|
||||
if exc_value is not None and exc_value.args and exc_value.args[0]
|
||||
else []
|
||||
)
|
||||
self._call_on_interrupt(interrupts)
|
||||
# suppress interrupt
|
||||
return True
|
||||
elif exc_type is None:
|
||||
# save final output
|
||||
self.output = read_channels(self.channels, self.output_keys)
|
||||
|
||||
def _call_on_interrupt(self, interrupts: list[Interrupt]) -> None:
|
||||
"""Call the on_interrupt hook synchronously."""
|
||||
if self.on_interrupt is None:
|
||||
return
|
||||
if iscoroutinefunction(self.on_interrupt):
|
||||
warnings.warn(
|
||||
"Async on_interrupt hook cannot be called from sync graph execution. "
|
||||
"Use a sync function or run the graph with astream/ainvoke.",
|
||||
stacklevel=2,
|
||||
)
|
||||
return
|
||||
self.on_interrupt(interrupts)
|
||||
|
||||
def _emit(
|
||||
self,
|
||||
mode: StreamMode,
|
||||
@@ -985,6 +1013,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
migrate_checkpoint: Callable[[Checkpoint], None] | None = None,
|
||||
retry_policy: Sequence[RetryPolicy] = (),
|
||||
cache_policy: CachePolicy | None = None,
|
||||
on_interrupt: OnInterruptHook | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
input,
|
||||
@@ -1006,6 +1035,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
retry_policy=retry_policy,
|
||||
cache_policy=cache_policy,
|
||||
durability=durability,
|
||||
on_interrupt=on_interrupt,
|
||||
)
|
||||
self.stack = ExitStack()
|
||||
if checkpointer:
|
||||
@@ -1161,6 +1191,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
migrate_checkpoint: Callable[[Checkpoint], None] | None = None,
|
||||
retry_policy: Sequence[RetryPolicy] = (),
|
||||
cache_policy: CachePolicy | None = None,
|
||||
on_interrupt: OnInterruptHook | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
input,
|
||||
@@ -1182,6 +1213,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
retry_policy=retry_policy,
|
||||
cache_policy=cache_policy,
|
||||
durability=durability,
|
||||
on_interrupt=on_interrupt,
|
||||
)
|
||||
self.stack = AsyncExitStack()
|
||||
if checkpointer:
|
||||
@@ -1257,6 +1289,18 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
},
|
||||
)
|
||||
|
||||
_deferred_on_interrupt_args: list[Interrupt] | None = None
|
||||
|
||||
def _call_on_interrupt(self, interrupts: list[Interrupt]) -> None:
|
||||
"""Override for async loop: defer async hooks to __aexit__."""
|
||||
if self.on_interrupt is None:
|
||||
return
|
||||
if iscoroutinefunction(self.on_interrupt):
|
||||
# Defer async hooks — they will be awaited in __aexit__
|
||||
self._deferred_on_interrupt_args = interrupts
|
||||
else:
|
||||
self.on_interrupt(interrupts)
|
||||
|
||||
# context manager
|
||||
|
||||
async def __aenter__(self) -> Self:
|
||||
@@ -1315,14 +1359,25 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
exc_value: BaseException | None,
|
||||
traceback: TracebackType | None,
|
||||
) -> bool | None:
|
||||
# unwind stack
|
||||
# unwind stack (calls _suppress_interrupt synchronously)
|
||||
exit_task = asyncio.create_task(
|
||||
self.stack.__aexit__(exc_type, exc_value, traceback)
|
||||
)
|
||||
try:
|
||||
return await exit_task
|
||||
result = await exit_task
|
||||
except asyncio.CancelledError as e:
|
||||
# Bubble up the exit task upon cancellation to permit the API
|
||||
# consumer to await it before e.g., reusing the DB connection.
|
||||
e.args = (*e.args, exit_task)
|
||||
raise
|
||||
# Await deferred async on_interrupt hook (set by _call_on_interrupt)
|
||||
if (
|
||||
self._deferred_on_interrupt_args is not None
|
||||
and self.on_interrupt is not None
|
||||
):
|
||||
interrupts = self._deferred_on_interrupt_args
|
||||
self._deferred_on_interrupt_args = None
|
||||
coro = self.on_interrupt(interrupts)
|
||||
if coro is not None:
|
||||
await coro
|
||||
return result
|
||||
|
||||
@@ -138,6 +138,7 @@ from langgraph.types import (
|
||||
Command,
|
||||
Durability,
|
||||
Interrupt,
|
||||
OnInterruptHook,
|
||||
Send,
|
||||
StateSnapshot,
|
||||
StateUpdate,
|
||||
@@ -622,6 +623,12 @@ class Pregel(
|
||||
context_schema: type[ContextT] | None = None
|
||||
"""Specifies the schema for the context object that will be passed to the workflow."""
|
||||
|
||||
on_interrupt: OnInterruptHook | None = None
|
||||
"""Optional callback invoked when the graph execution is interrupted.
|
||||
|
||||
Called with the list of `Interrupt` objects whenever the graph pauses.
|
||||
May be a sync or async callable."""
|
||||
|
||||
config: RunnableConfig | None = None
|
||||
|
||||
name: str = "LangGraph"
|
||||
@@ -652,6 +659,7 @@ class Pregel(
|
||||
config: RunnableConfig | None = None,
|
||||
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
|
||||
name: str = "LangGraph",
|
||||
on_interrupt: OnInterruptHook | None = None,
|
||||
**deprecated_kwargs: Unpack[DeprecatedKwargs],
|
||||
) -> None:
|
||||
if (
|
||||
@@ -695,6 +703,7 @@ class Pregel(
|
||||
)
|
||||
self.cache_policy = cache_policy
|
||||
self.context_schema = context_schema
|
||||
self.on_interrupt = on_interrupt
|
||||
self.config = config
|
||||
self.trigger_to_nodes = trigger_to_nodes or {}
|
||||
self.name = name
|
||||
@@ -2599,6 +2608,7 @@ class Pregel(
|
||||
migrate_checkpoint=self._migrate_checkpoint,
|
||||
retry_policy=self.retry_policy,
|
||||
cache_policy=self.cache_policy,
|
||||
on_interrupt=self.on_interrupt,
|
||||
) as loop:
|
||||
# create runner
|
||||
runner = PregelRunner(
|
||||
@@ -2908,6 +2918,7 @@ class Pregel(
|
||||
migrate_checkpoint=self._migrate_checkpoint,
|
||||
retry_policy=self.retry_policy,
|
||||
cache_policy=self.cache_policy,
|
||||
on_interrupt=self.on_interrupt,
|
||||
) as loop:
|
||||
# create runner
|
||||
runner = PregelRunner(
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from collections import deque
|
||||
from collections.abc import Callable, Hashable, Sequence
|
||||
from collections.abc import Awaitable, Callable, Hashable, Sequence
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
@@ -56,6 +56,7 @@ __all__ = (
|
||||
"Durability",
|
||||
"interrupt",
|
||||
"Overwrite",
|
||||
"OnInterruptHook",
|
||||
"ensure_valid_checkpointer",
|
||||
)
|
||||
|
||||
@@ -109,6 +110,18 @@ StreamWriter = Callable[[Any], None]
|
||||
Always injected into nodes if requested as a keyword argument, but it's a no-op
|
||||
when not using `stream_mode="custom"`."""
|
||||
|
||||
OnInterruptHook = (
|
||||
Callable[[list["Interrupt"]], None] | Callable[[list["Interrupt"]], Awaitable[None]]
|
||||
)
|
||||
"""Callback invoked when a graph execution is interrupted.
|
||||
|
||||
Called with the list of `Interrupt` objects whenever the graph pauses due to
|
||||
an `interrupt()` call or `interrupt_before`/`interrupt_after` configuration.
|
||||
|
||||
May be a regular function or an async coroutine function. Async hooks are
|
||||
awaited in async graph execution; in sync execution only sync hooks are called.
|
||||
"""
|
||||
|
||||
_DC_KWARGS = {"kw_only": True, "slots": True, "frozen": True}
|
||||
|
||||
|
||||
|
||||
@@ -3038,7 +3038,7 @@ def test_message_graph(
|
||||
# add an extra message as if it came from "tools" node
|
||||
app_w_interrupt.update_state(config, ("ai", "an extra message"), as_node="tools")
|
||||
|
||||
# extra message is coerced BaseMessge and appended
|
||||
# extra message is coerced BaseMessage and appended
|
||||
# now the next node is "agent" per the graph edges
|
||||
assert app_w_interrupt.get_state(config) == StateSnapshot(
|
||||
values=[
|
||||
@@ -3762,7 +3762,7 @@ def test_root_graph(
|
||||
# add an extra message as if it came from "tools" node
|
||||
app_w_interrupt.update_state(config, ("ai", "an extra message"), as_node="tools")
|
||||
|
||||
# extra message is coerced BaseMessge and appended
|
||||
# extra message is coerced BaseMessage and appended
|
||||
# now the next node is "agent" per the graph edges
|
||||
assert app_w_interrupt.get_state(config) == StateSnapshot(
|
||||
values=[
|
||||
|
||||
@@ -8893,3 +8893,151 @@ def test_fork_does_not_apply_pending_writes(
|
||||
|
||||
# Should be: 1 (input) + 20 (forked node_a) + 100 (node_b) = 121
|
||||
assert result == {"value": 121}
|
||||
|
||||
|
||||
def test_on_interrupt_hook_with_interrupt_call(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Test that on_interrupt hook fires when interrupt() is called in a node."""
|
||||
hook_calls: list[list[Interrupt]] = []
|
||||
|
||||
def my_on_interrupt(interrupts: list[Interrupt]) -> None:
|
||||
hook_calls.append(interrupts)
|
||||
|
||||
class State(TypedDict):
|
||||
value: str
|
||||
|
||||
def ask_human(state: State) -> dict:
|
||||
answer = interrupt("what should I do?")
|
||||
return {"value": answer}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("ask", ask_human)
|
||||
builder.add_edge(START, "ask")
|
||||
|
||||
graph = builder.compile(
|
||||
checkpointer=sync_checkpointer,
|
||||
on_interrupt=my_on_interrupt,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# First invocation: should trigger interrupt and call the hook
|
||||
result = list(graph.stream({"value": ""}, config))
|
||||
assert len(result) == 1
|
||||
assert "__interrupt__" in result[0]
|
||||
|
||||
# Hook should have been called once with the interrupt data
|
||||
assert len(hook_calls) == 1
|
||||
assert len(hook_calls[0]) == 1
|
||||
assert hook_calls[0][0].value == "what should I do?"
|
||||
|
||||
# Resume — no new interrupt, hook should not fire again
|
||||
hook_calls.clear()
|
||||
result = list(graph.stream(Command(resume="do this"), config))
|
||||
assert any("ask" in chunk for chunk in result)
|
||||
assert len(hook_calls) == 0
|
||||
|
||||
|
||||
def test_on_interrupt_hook_with_interrupt_before(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Test that on_interrupt hook fires for interrupt_before config."""
|
||||
hook_calls: list[list[Interrupt]] = []
|
||||
|
||||
def my_on_interrupt(interrupts: list[Interrupt]) -> None:
|
||||
hook_calls.append(interrupts)
|
||||
|
||||
class State(TypedDict):
|
||||
value: int
|
||||
|
||||
def add_one(state: State) -> dict:
|
||||
return {"value": state["value"] + 1}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("add_one", add_one)
|
||||
builder.add_edge(START, "add_one")
|
||||
|
||||
graph = builder.compile(
|
||||
checkpointer=sync_checkpointer,
|
||||
interrupt_before=["add_one"],
|
||||
on_interrupt=my_on_interrupt,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Should interrupt before add_one runs
|
||||
result = list(graph.stream({"value": 0}, config))
|
||||
assert any("__interrupt__" in chunk for chunk in result)
|
||||
|
||||
# Hook should have been called (empty interrupt list for config-level interrupts)
|
||||
assert len(hook_calls) == 1
|
||||
assert hook_calls[0] == []
|
||||
|
||||
|
||||
def test_on_interrupt_hook_not_called_without_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Test that on_interrupt hook is NOT called when graph completes normally."""
|
||||
hook_calls: list[list[Interrupt]] = []
|
||||
|
||||
def my_on_interrupt(interrupts: list[Interrupt]) -> None:
|
||||
hook_calls.append(interrupts)
|
||||
|
||||
class State(TypedDict):
|
||||
value: int
|
||||
|
||||
def add_one(state: State) -> dict:
|
||||
return {"value": state["value"] + 1}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("add_one", add_one)
|
||||
builder.add_edge(START, "add_one")
|
||||
|
||||
graph = builder.compile(
|
||||
checkpointer=sync_checkpointer,
|
||||
on_interrupt=my_on_interrupt,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
result = graph.invoke({"value": 0}, config)
|
||||
assert result == {"value": 1}
|
||||
|
||||
# Hook should NOT have been called
|
||||
assert len(hook_calls) == 0
|
||||
|
||||
|
||||
def test_on_interrupt_hook_exception_propagates(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Test that exceptions in the on_interrupt hook propagate to the caller."""
|
||||
|
||||
def bad_hook(interrupts: list[Interrupt]) -> None:
|
||||
raise RuntimeError("hook exploded")
|
||||
|
||||
class State(TypedDict):
|
||||
value: str
|
||||
|
||||
def ask(state: State) -> dict:
|
||||
answer = interrupt("question")
|
||||
return {"value": answer}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("ask", ask)
|
||||
builder.add_edge(START, "ask")
|
||||
|
||||
graph = builder.compile(
|
||||
checkpointer=sync_checkpointer,
|
||||
on_interrupt=bad_hook,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Hook error should propagate
|
||||
with pytest.raises(RuntimeError, match="hook exploded"):
|
||||
list(graph.stream({"value": ""}, config))
|
||||
|
||||
# Graph state should still be checkpointed and resumable despite the hook error
|
||||
result = list(graph.stream(Command(resume="answer"), config))
|
||||
assert any("ask" in chunk for chunk in result)
|
||||
|
||||
@@ -9345,3 +9345,117 @@ async def test_fork_does_not_apply_pending_writes(
|
||||
|
||||
# 1 (input) + 20 (forked node_a) + 100 (node_b) = 121
|
||||
assert result == {"value": 121}
|
||||
|
||||
|
||||
async def test_on_interrupt_hook_async_with_interrupt_call(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Test that an async on_interrupt hook fires when interrupt() is called."""
|
||||
hook_calls: list[list[Interrupt]] = []
|
||||
|
||||
async def my_on_interrupt(interrupts: list[Interrupt]) -> None:
|
||||
hook_calls.append(interrupts)
|
||||
|
||||
class State(TypedDict):
|
||||
value: str
|
||||
|
||||
def ask_human(state: State) -> dict:
|
||||
answer = interrupt("what should I do?")
|
||||
return {"value": answer}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("ask", ask_human)
|
||||
builder.add_edge(START, "ask")
|
||||
|
||||
graph = builder.compile(
|
||||
checkpointer=async_checkpointer,
|
||||
on_interrupt=my_on_interrupt,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# First invocation: should trigger interrupt and call the async hook
|
||||
result = [chunk async for chunk in graph.astream({"value": ""}, config)]
|
||||
assert len(result) == 1
|
||||
assert "__interrupt__" in result[0]
|
||||
|
||||
# Hook should have been called once with the interrupt data
|
||||
assert len(hook_calls) == 1
|
||||
assert len(hook_calls[0]) == 1
|
||||
assert hook_calls[0][0].value == "what should I do?"
|
||||
|
||||
# Resume — no new interrupt, hook should not fire again
|
||||
hook_calls.clear()
|
||||
result = [chunk async for chunk in graph.astream(Command(resume="do this"), config)]
|
||||
assert any("ask" in chunk for chunk in result)
|
||||
assert len(hook_calls) == 0
|
||||
|
||||
|
||||
async def test_on_interrupt_hook_sync_in_async_graph(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Test that a sync on_interrupt hook works in async graph execution."""
|
||||
hook_calls: list[list[Interrupt]] = []
|
||||
|
||||
def my_sync_hook(interrupts: list[Interrupt]) -> None:
|
||||
hook_calls.append(interrupts)
|
||||
|
||||
class State(TypedDict):
|
||||
value: str
|
||||
|
||||
def ask_human(state: State) -> dict:
|
||||
answer = interrupt("question?")
|
||||
return {"value": answer}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("ask", ask_human)
|
||||
builder.add_edge(START, "ask")
|
||||
|
||||
graph = builder.compile(
|
||||
checkpointer=async_checkpointer,
|
||||
on_interrupt=my_sync_hook,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
result = [chunk async for chunk in graph.astream({"value": ""}, config)]
|
||||
assert "__interrupt__" in result[0]
|
||||
|
||||
# Sync hook should work fine in async execution
|
||||
assert len(hook_calls) == 1
|
||||
assert hook_calls[0][0].value == "question?"
|
||||
|
||||
|
||||
async def test_on_interrupt_hook_async_exception_propagates(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Test that exceptions in the async on_interrupt hook propagate."""
|
||||
|
||||
async def bad_hook(interrupts: list[Interrupt]) -> None:
|
||||
raise RuntimeError("async hook exploded")
|
||||
|
||||
class State(TypedDict):
|
||||
value: str
|
||||
|
||||
def ask(state: State) -> dict:
|
||||
answer = interrupt("question")
|
||||
return {"value": answer}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("ask", ask)
|
||||
builder.add_edge(START, "ask")
|
||||
|
||||
graph = builder.compile(
|
||||
checkpointer=async_checkpointer,
|
||||
on_interrupt=bad_hook,
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# Hook error should propagate
|
||||
with pytest.raises(RuntimeError, match="async hook exploded"):
|
||||
[chunk async for chunk in graph.astream({"value": ""}, config)]
|
||||
|
||||
# Graph state should still be checkpointed and resumable despite the hook error
|
||||
result = [chunk async for chunk in graph.astream(Command(resume="answer"), config)]
|
||||
assert any("ask" in chunk for chunk in result)
|
||||
|
||||
@@ -15,12 +15,12 @@ The module implements design patterns for:
|
||||
|
||||
Key Components:
|
||||
|
||||
- `ToolNode`: Main class for executing tools in LangGraph workflows
|
||||
- `InjectedState`: Annotation for injecting graph state into tools
|
||||
- `InjectedStore`: Annotation for injecting persistent store into tools
|
||||
- `ToolRuntime`: Runtime information for tools, bundling together `state`, `context`,
|
||||
- [`ToolNode`][langgraph.prebuilt.ToolNode]: Main class for executing tools in LangGraph workflows
|
||||
- [`InjectedState`][langgraph.prebuilt.InjectedState]: Annotation for injecting graph state into tools
|
||||
- [`InjectedStore`][langgraph.prebuilt.InjectedStore]: Annotation for injecting persistent store into tools
|
||||
- [`ToolRuntime`][langgraph.prebuilt.ToolRuntime]: Runtime information for tools, bundling together `state`, `context`,
|
||||
`config`, `stream_writer`, `tool_call_id`, and `store`
|
||||
- `tools_condition`: Utility function for conditional routing based on tool calls
|
||||
- [`tools_condition`][langgraph.prebuilt.tools_condition]: Utility function for conditional routing based on tool calls
|
||||
|
||||
Typical Usage:
|
||||
```python
|
||||
@@ -614,8 +614,16 @@ class ToolNode(RunnableCallable):
|
||||
persistent storage, and control flow. Manages parallel execution,
|
||||
error handling.
|
||||
|
||||
Use `ToolNode` when building custom workflows that require fine-grained control over
|
||||
tool execution—for example, custom routing logic, specialized error handling, or
|
||||
non-standard agent architectures.
|
||||
|
||||
For standard ReAct-style agents, use [`create_agent`][langchain.agents.create_agent]
|
||||
instead. It uses `ToolNode` internally with sensible defaults for the agent loop,
|
||||
conditional routing, and error handling.
|
||||
|
||||
Input Formats:
|
||||
1. Graph state with `messages` key that has a list of messages:
|
||||
1. **Graph state** with `messages` key that has a list of messages:
|
||||
- Common representation for agentic workflows
|
||||
- Supports custom messages key via `messages_key` parameter
|
||||
|
||||
|
||||
@@ -3,6 +3,6 @@ from langgraph_sdk.client import get_client, get_sync_client
|
||||
from langgraph_sdk.encryption import Encryption
|
||||
from langgraph_sdk.encryption.types import EncryptionContext
|
||||
|
||||
__version__ = "0.3.0"
|
||||
__version__ = "0.3.1"
|
||||
|
||||
__all__ = ["Auth", "Encryption", "EncryptionContext", "get_client", "get_sync_client"]
|
||||
|
||||
@@ -64,114 +64,6 @@ def _validate_handler(fn: typing.Callable, handler_type: str) -> None:
|
||||
)
|
||||
|
||||
|
||||
class _JsonEncryptDecorators:
|
||||
"""Dynamic decorator factory for JSON encryption handlers.
|
||||
|
||||
Supports both default and model-specific handlers:
|
||||
- @encrypt.json - default handler for all models
|
||||
- @encrypt.json.thread - handler for thread model
|
||||
"""
|
||||
|
||||
def __init__(self, parent: Encryption):
|
||||
self._parent = parent
|
||||
|
||||
def __call__(self, fn: types.JsonEncryptor) -> types.JsonEncryptor:
|
||||
"""Register the default JSON encryption handler.
|
||||
|
||||
Args:
|
||||
fn: The handler function
|
||||
|
||||
Returns:
|
||||
The registered handler function
|
||||
|
||||
Raises:
|
||||
DuplicateHandlerError: If handler already registered
|
||||
TypeError: If handler has invalid signature
|
||||
"""
|
||||
if self._parent._json_encryptor is not None:
|
||||
raise DuplicateHandlerError("Default JSON encryptor already registered")
|
||||
_validate_handler(fn, "Default JSON encryptor")
|
||||
self._parent._json_encryptor = fn
|
||||
return fn
|
||||
|
||||
def __getattr__(
|
||||
self, model: str
|
||||
) -> typing.Callable[[types.JsonEncryptor], types.JsonEncryptor]:
|
||||
"""Dynamic attribute access for model-specific handlers.
|
||||
|
||||
Allows @encryption.encrypt.json.thread, @encryption.encrypt.json.assistant, etc.
|
||||
|
||||
Raises:
|
||||
DuplicateHandlerError: If handler already registered for this model
|
||||
TypeError: If handler has invalid signature
|
||||
"""
|
||||
|
||||
def decorator(fn: types.JsonEncryptor) -> types.JsonEncryptor:
|
||||
if model in self._parent._json_encryptors:
|
||||
raise DuplicateHandlerError(
|
||||
f"JSON encryptor for model '{model}' already registered"
|
||||
)
|
||||
_validate_handler(fn, f"JSON encryptor for model '{model}'")
|
||||
self._parent._json_encryptors[model] = fn
|
||||
return fn
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
class _JsonDecryptDecorators:
|
||||
"""Dynamic decorator factory for JSON decryption handlers.
|
||||
|
||||
Supports both default and model-specific handlers:
|
||||
- @encryption.decrypt.json - default handler for all models
|
||||
- @encryption.decrypt.json.thread - handler for thread model
|
||||
"""
|
||||
|
||||
def __init__(self, parent: Encryption):
|
||||
self._parent = parent
|
||||
|
||||
def __call__(self, fn: types.JsonDecryptor) -> types.JsonDecryptor:
|
||||
"""Register the default JSON decryption handler.
|
||||
|
||||
Args:
|
||||
fn: The handler function
|
||||
|
||||
Returns:
|
||||
The registered handler function
|
||||
|
||||
Raises:
|
||||
DuplicateHandlerError: If handler already registered
|
||||
TypeError: If handler has invalid signature
|
||||
"""
|
||||
if self._parent._json_decryptor is not None:
|
||||
raise DuplicateHandlerError("Default JSON decryptor already registered")
|
||||
_validate_handler(fn, "Default JSON decryptor")
|
||||
self._parent._json_decryptor = fn
|
||||
return fn
|
||||
|
||||
def __getattr__(
|
||||
self, model: str
|
||||
) -> typing.Callable[[types.JsonDecryptor], types.JsonDecryptor]:
|
||||
"""Dynamic attribute access for model-specific handlers.
|
||||
|
||||
Allows @encryption.decrypt.json.thread, @encryption.decrypt.json.assistant, etc.
|
||||
|
||||
Raises:
|
||||
DuplicateHandlerError: If handler already registered for this model
|
||||
TypeError: If handler has invalid signature
|
||||
"""
|
||||
|
||||
def decorator(fn: types.JsonDecryptor) -> types.JsonDecryptor:
|
||||
if model in self._parent._json_decryptors:
|
||||
raise DuplicateHandlerError(
|
||||
f"JSON decryptor for model '{model}' already registered"
|
||||
)
|
||||
_validate_handler(fn, f"JSON decryptor for model '{model}'")
|
||||
self._parent._json_decryptors[model] = fn
|
||||
return fn
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
class _EncryptDecorators:
|
||||
"""Decorators for encryption handlers.
|
||||
|
||||
@@ -181,7 +73,6 @@ class _EncryptDecorators:
|
||||
|
||||
def __init__(self, parent: Encryption):
|
||||
self._parent = parent
|
||||
self._json = _JsonEncryptDecorators(parent)
|
||||
|
||||
def blob(self, fn: types.BlobEncryptor) -> types.BlobEncryptor:
|
||||
"""Register a blob encryption handler.
|
||||
@@ -212,29 +103,32 @@ class _EncryptDecorators:
|
||||
self._parent._blob_encryptor = fn
|
||||
return fn
|
||||
|
||||
@property
|
||||
def json(self) -> _JsonEncryptDecorators:
|
||||
"""Access JSON encryption decorators.
|
||||
|
||||
Supports model-specific handlers:
|
||||
- @encryption.encrypt.json - default handler for all models
|
||||
- @encryption.encrypt.json.thread - handler for thread model only
|
||||
- @encryption.encrypt.json.assistant - handler for assistant model only
|
||||
def json(self, fn: types.JsonEncryptor) -> types.JsonEncryptor:
|
||||
"""Register the JSON encryption handler.
|
||||
|
||||
Example:
|
||||
```python
|
||||
@encryption.encrypt.json
|
||||
async def default_encrypt(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Default encryption for all models
|
||||
async def encrypt_json(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Encrypt the data
|
||||
return encrypt_data(data)
|
||||
|
||||
@encryption.encrypt.json.thread
|
||||
async def encrypt_thread(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Special encryption for thread model only
|
||||
return encrypt_thread_data(data)
|
||||
```
|
||||
|
||||
Args:
|
||||
fn: The encryption handler function
|
||||
|
||||
Returns:
|
||||
The registered handler function
|
||||
|
||||
Raises:
|
||||
DuplicateHandlerError: If JSON encryptor already registered
|
||||
TypeError: If handler has invalid signature
|
||||
"""
|
||||
return self._json
|
||||
if self._parent._json_encryptor is not None:
|
||||
raise DuplicateHandlerError("JSON encryptor already registered")
|
||||
_validate_handler(fn, "JSON encryptor")
|
||||
self._parent._json_encryptor = fn
|
||||
return fn
|
||||
|
||||
|
||||
class _DecryptDecorators:
|
||||
@@ -246,7 +140,6 @@ class _DecryptDecorators:
|
||||
|
||||
def __init__(self, parent: Encryption):
|
||||
self._parent = parent
|
||||
self._json = _JsonDecryptDecorators(parent)
|
||||
|
||||
def blob(self, fn: types.BlobDecryptor) -> types.BlobDecryptor:
|
||||
"""Register a blob decryption handler.
|
||||
@@ -277,29 +170,32 @@ class _DecryptDecorators:
|
||||
self._parent._blob_decryptor = fn
|
||||
return fn
|
||||
|
||||
@property
|
||||
def json(self) -> _JsonDecryptDecorators:
|
||||
"""Access JSON decryption decorators.
|
||||
|
||||
Supports model-specific handlers:
|
||||
- @encryption.decrypt.json - default handler for all models
|
||||
- @encryption.decrypt.json.thread - handler for thread model only
|
||||
- @encryption.decrypt.json.assistant - handler for assistant model only
|
||||
def json(self, fn: types.JsonDecryptor) -> types.JsonDecryptor:
|
||||
"""Register the JSON decryption handler.
|
||||
|
||||
Example:
|
||||
```python
|
||||
@encryption.decrypt.json
|
||||
async def default_decrypt(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Default decryption for all models
|
||||
async def decrypt_json(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Decrypt the data
|
||||
return decrypt_data(data)
|
||||
|
||||
@encryption.decrypt.json.thread
|
||||
async def decrypt_thread(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Special decryption for thread model only
|
||||
return decrypt_thread_data(data)
|
||||
```
|
||||
|
||||
Args:
|
||||
fn: The decryption handler function
|
||||
|
||||
Returns:
|
||||
The registered handler function
|
||||
|
||||
Raises:
|
||||
DuplicateHandlerError: If JSON decryptor already registered
|
||||
TypeError: If handler has invalid signature
|
||||
"""
|
||||
return self._json
|
||||
if self._parent._json_decryptor is not None:
|
||||
raise DuplicateHandlerError("JSON decryptor already registered")
|
||||
_validate_handler(fn, "JSON decryptor")
|
||||
self._parent._json_decryptor = fn
|
||||
return fn
|
||||
|
||||
|
||||
class Encryption:
|
||||
@@ -336,6 +232,28 @@ class Encryption:
|
||||
Then the LangGraph server will load your encryption file and use it to
|
||||
encrypt/decrypt data at rest.
|
||||
|
||||
!!! warning "JSON Encryptors Must Preserve Keys"
|
||||
|
||||
JSON encryptors **must not add or remove keys** from the input dict.
|
||||
Only values may be transformed. This constraint is **enforced at runtime
|
||||
by the server** and exists because SQL JSONB merge operations (used for
|
||||
partial updates) work at the key level.
|
||||
|
||||
**Correct (per-key encryption):**
|
||||
```python
|
||||
# Input: {"secret": "value", "plain": "x"}
|
||||
# Output: {"secret": "<encrypted>", "plain": "x"} ✓ Keys preserved
|
||||
```
|
||||
|
||||
**Incorrect (key consolidation):**
|
||||
```python
|
||||
# Input: {"secret": "value", "plain": "x"}
|
||||
# Output: {"__encrypted__": "<blob>", "plain": "x"} ✗ Key changed
|
||||
```
|
||||
|
||||
If your encryptor needs to store auxiliary data (DEK, IV, etc.), embed it
|
||||
within the encrypted value itself, not as separate keys.
|
||||
|
||||
???+ example "Basic Usage"
|
||||
|
||||
```python
|
||||
@@ -343,89 +261,48 @@ class Encryption:
|
||||
|
||||
my_encryption = Encryption()
|
||||
|
||||
SKIP_FIELDS = {"tenant_id", "owner", "thread_id", "assistant_id"}
|
||||
ENCRYPTED_PREFIX = "encrypted:"
|
||||
|
||||
@my_encryption.encrypt.blob
|
||||
async def encrypt_blob(ctx: EncryptionContext, blob: bytes) -> bytes:
|
||||
# Call your encryption service
|
||||
return encrypted_blob
|
||||
return your_encrypt_bytes(blob)
|
||||
|
||||
@my_encryption.decrypt.blob
|
||||
async def decrypt_blob(ctx: EncryptionContext, blob: bytes) -> bytes:
|
||||
# Call your decryption service
|
||||
return decrypted_blob
|
||||
return your_decrypt_bytes(blob)
|
||||
|
||||
@my_encryption.encrypt.json
|
||||
async def encrypt_json(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Practical encryption strategy:
|
||||
# - "owner" field: unencrypted (for search/filtering)
|
||||
# - "my.customer.org/" prefixed fields: encrypt VALUES only
|
||||
# - All other fields: pass through unencrypted
|
||||
encrypted = {}
|
||||
for key, value in data.items():
|
||||
if key.startswith("my.customer.org/"):
|
||||
# Encrypt VALUE for sensitive customer data
|
||||
encrypted[key] = encrypt_value(value)
|
||||
result = {}
|
||||
for k, v in data.items():
|
||||
if k in SKIP_FIELDS or v is None:
|
||||
result[k] = v
|
||||
else:
|
||||
# Pass through (including "owner" for search)
|
||||
encrypted[key] = value
|
||||
return encrypted
|
||||
result[k] = ENCRYPTED_PREFIX + your_encrypt_string(v)
|
||||
return result
|
||||
|
||||
@my_encryption.decrypt.json
|
||||
async def decrypt_json(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Decrypt VALUES for "my.customer.org/" prefixed fields
|
||||
decrypted = {}
|
||||
for key, value in data.items():
|
||||
if key.startswith("my.customer.org/"):
|
||||
decrypted[key] = decrypt_value(value)
|
||||
result = {}
|
||||
for k, v in data.items():
|
||||
if isinstance(v, str) and v.startswith(ENCRYPTED_PREFIX):
|
||||
result[k] = your_decrypt_string(v[len(ENCRYPTED_PREFIX):])
|
||||
else:
|
||||
decrypted[key] = value
|
||||
return decrypted
|
||||
```
|
||||
|
||||
???+ example "Model-Specific Handlers"
|
||||
|
||||
You can register different encryption handlers for different model types
|
||||
(thread, assistant, run, cron, checkpoint, etc.):
|
||||
|
||||
```python
|
||||
from langgraph_sdk import Encryption, EncryptionContext
|
||||
|
||||
my_encryption = Encryption()
|
||||
|
||||
# Default handler for models without specific handlers
|
||||
@my_encryption.encrypt.json
|
||||
async def default_encrypt(ctx: EncryptionContext, data: dict) -> dict:
|
||||
return standard_encrypt(data)
|
||||
|
||||
# Thread-specific handler (uses different KMS key)
|
||||
@my_encryption.encrypt.json.thread
|
||||
async def encrypt_thread(ctx: EncryptionContext, data: dict) -> dict:
|
||||
return encrypt_with_thread_key(data)
|
||||
|
||||
# Assistant-specific handler
|
||||
@my_encryption.encrypt.json.assistant
|
||||
async def encrypt_assistant(ctx: EncryptionContext, data: dict) -> dict:
|
||||
return encrypt_with_assistant_key(data)
|
||||
|
||||
# Same pattern for decryption
|
||||
@my_encryption.decrypt.json
|
||||
async def default_decrypt(ctx: EncryptionContext, data: dict) -> dict:
|
||||
return standard_decrypt(data)
|
||||
|
||||
@my_encryption.decrypt.json.thread
|
||||
async def decrypt_thread(ctx: EncryptionContext, data: dict) -> dict:
|
||||
return decrypt_with_thread_key(data)
|
||||
result[k] = v
|
||||
return result
|
||||
```
|
||||
|
||||
???+ example "Field-Specific Logic"
|
||||
|
||||
The `ctx.field` attribute tells you which specific field is being encrypted,
|
||||
allowing different logic within the same model:
|
||||
The `ctx.model` and `ctx.field` attributes tell you which model type and
|
||||
specific field is being encrypted, allowing different logic:
|
||||
|
||||
```python
|
||||
@my_encryption.encrypt.json.thread
|
||||
async def encrypt_thread(ctx: EncryptionContext, data: dict) -> dict:
|
||||
@my_encryption.encrypt.json
|
||||
async def encrypt_json(ctx: EncryptionContext, data: dict) -> dict:
|
||||
if ctx.field == "metadata":
|
||||
# Thread metadata - standard encryption
|
||||
# Metadata - standard encryption
|
||||
return encrypt_standard(data)
|
||||
elif ctx.field == "values":
|
||||
# Thread values - more sensitive, use stronger encryption
|
||||
@@ -433,6 +310,49 @@ class Encryption:
|
||||
else:
|
||||
return encrypt_standard(data)
|
||||
```
|
||||
|
||||
!!! warning "Model/Field May Differ Between Encrypt and Decrypt"
|
||||
|
||||
Data encrypted with one `(model, field)` pair is **not guaranteed**
|
||||
to be decrypted with the same pair. The server performs SQL JSONB
|
||||
merges that can move encrypted values between models (e.g., cron
|
||||
metadata → run metadata). Your decryption logic must handle data
|
||||
regardless of the `ctx.model` or `ctx.field` values at decrypt time.
|
||||
|
||||
**Safe:** Use `ctx.model`/`ctx.field` for logging or metrics only.
|
||||
|
||||
**Safe:** Encrypt different keys based on `ctx.field`, but use a
|
||||
single decrypt handler that decrypts any value with the encrypted
|
||||
prefix (and passes through plaintext unchanged):
|
||||
|
||||
```python
|
||||
ENCRYPTED_PREFIX = "enc:"
|
||||
|
||||
@my_encryption.encrypt.json
|
||||
async def encrypt_json(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Encrypt different keys depending on the field
|
||||
if ctx.field == "context":
|
||||
keys_to_encrypt = {"api_key", "secret_token"}
|
||||
else:
|
||||
keys_to_encrypt = {"email", "ssn"}
|
||||
return {
|
||||
k: ENCRYPTED_PREFIX + encrypt(v) if k in keys_to_encrypt else v
|
||||
for k, v in data.items()
|
||||
}
|
||||
|
||||
@my_encryption.decrypt.json
|
||||
async def decrypt_json(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Decrypt ANY value with the prefix, regardless of model/field
|
||||
return {
|
||||
k: decrypt(v[len(ENCRYPTED_PREFIX):])
|
||||
if isinstance(v, str) and v.startswith(ENCRYPTED_PREFIX)
|
||||
else v
|
||||
for k, v in data.items()
|
||||
}
|
||||
```
|
||||
|
||||
**Unsafe:** Using different encryption keys or algorithms based on
|
||||
`ctx.model`/`ctx.field` will cause decryption failures.
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
@@ -440,9 +360,7 @@ class Encryption:
|
||||
"_blob_encryptor",
|
||||
"_context_handler",
|
||||
"_json_decryptor",
|
||||
"_json_decryptors",
|
||||
"_json_encryptor",
|
||||
"_json_encryptors",
|
||||
"decrypt",
|
||||
"encrypt",
|
||||
)
|
||||
@@ -464,8 +382,6 @@ class Encryption:
|
||||
self._blob_decryptor: types.BlobDecryptor | None = None
|
||||
self._json_encryptor: types.JsonEncryptor | None = None
|
||||
self._json_decryptor: types.JsonDecryptor | None = None
|
||||
self._json_encryptors: dict[str, types.JsonEncryptor] = {}
|
||||
self._json_decryptors: dict[str, types.JsonDecryptor] = {}
|
||||
self._context_handler: types.ContextHandler | None = None
|
||||
|
||||
def context(self, fn: types.ContextHandler) -> types.ContextHandler:
|
||||
@@ -506,33 +422,33 @@ class Encryption:
|
||||
return fn
|
||||
|
||||
def get_json_encryptor(
|
||||
self, model: str | None = None
|
||||
self,
|
||||
_model: str | None = None, # kept for langgraph-api compat
|
||||
) -> types.JsonEncryptor | None:
|
||||
"""Get the JSON encryptor for a specific model.
|
||||
"""Get the JSON encryptor.
|
||||
|
||||
Args:
|
||||
model: The model type (e.g., "thread", "assistant"). If None, returns default.
|
||||
_model: Ignored. Kept for backwards compatibility with langgraph-api
|
||||
which passes model_type to this method.
|
||||
|
||||
Returns:
|
||||
Model-specific encryptor if registered, otherwise default encryptor, or None.
|
||||
The JSON encryptor, or None if not registered.
|
||||
"""
|
||||
if model and model in self._json_encryptors:
|
||||
return self._json_encryptors[model]
|
||||
return self._json_encryptor
|
||||
|
||||
def get_json_decryptor(
|
||||
self, model: str | None = None
|
||||
self,
|
||||
_model: str | None = None, # kept for langgraph-api compat
|
||||
) -> types.JsonDecryptor | None:
|
||||
"""Get the JSON decryptor for a specific model.
|
||||
"""Get the JSON decryptor.
|
||||
|
||||
Args:
|
||||
model: The model type (e.g., "thread", "assistant"). If None, returns default.
|
||||
_model: Ignored. Kept for backwards compatibility with langgraph-api
|
||||
which passes model_type to this method.
|
||||
|
||||
Returns:
|
||||
Model-specific decryptor if registered, otherwise default decryptor, or None.
|
||||
The JSON decryptor, or None if not registered.
|
||||
"""
|
||||
if model and model in self._json_decryptors:
|
||||
return self._json_decryptors[model]
|
||||
return self._json_decryptor
|
||||
|
||||
def __repr__(self) -> str:
|
||||
@@ -545,10 +461,6 @@ class Encryption:
|
||||
handlers.append("json_encryptor")
|
||||
if self._json_decryptor:
|
||||
handlers.append("json_decryptor")
|
||||
if self._json_encryptors:
|
||||
handlers.append(f"json_encryptors({list(self._json_encryptors.keys())})")
|
||||
if self._json_decryptors:
|
||||
handlers.append(f"json_decryptors({list(self._json_decryptors.keys())})")
|
||||
if self._context_handler:
|
||||
handlers.append("context_handler")
|
||||
return f"Encryption(handlers=[{', '.join(handlers)}])"
|
||||
|
||||
@@ -26,14 +26,6 @@ class TestHandlerValidation:
|
||||
async def json_dec(_ctx, data):
|
||||
return data
|
||||
|
||||
@encryption.encrypt.json.thread
|
||||
async def thread_enc(_ctx, data):
|
||||
return data
|
||||
|
||||
@encryption.decrypt.json.custom
|
||||
async def custom_dec(_ctx, data):
|
||||
return data
|
||||
|
||||
# All duplicates should raise
|
||||
with pytest.raises(DuplicateHandlerError):
|
||||
|
||||
@@ -59,18 +51,6 @@ class TestHandlerValidation:
|
||||
async def dup(_ctx, data):
|
||||
return data
|
||||
|
||||
with pytest.raises(DuplicateHandlerError):
|
||||
|
||||
@encryption.encrypt.json.thread
|
||||
async def dup(_ctx, data):
|
||||
return data
|
||||
|
||||
with pytest.raises(DuplicateHandlerError):
|
||||
|
||||
@encryption.decrypt.json.custom
|
||||
async def dup(_ctx, data):
|
||||
return data
|
||||
|
||||
def test_handlers_must_be_async(self):
|
||||
"""Sync functions raise TypeError."""
|
||||
encryption = Encryption()
|
||||
|
||||
Reference in New Issue
Block a user