mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-08 17:35:17 +02:00
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
844373e9b3 | ||
|
|
eec823cbd6 | ||
|
|
9eac53db4b |
@@ -631,6 +631,26 @@ class RemoteGraph(PregelProtocol):
|
||||
)
|
||||
return self._get_config(response["checkpoint"])
|
||||
|
||||
def _prepare_run_input(
|
||||
self,
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None,
|
||||
) -> tuple[RunnableConfig, dict[str, Any] | Any, CommandSDK | None, str | None]:
|
||||
"""Prepare input for run calls.
|
||||
|
||||
Returns:
|
||||
Tuple of (sanitized_config, input, command, thread_id)
|
||||
"""
|
||||
merged_config = merge_configs(self.config, config)
|
||||
sanitized_config = self._sanitize_config(merged_config)
|
||||
if isinstance(input, Command):
|
||||
command: CommandSDK | None = cast(CommandSDK, asdict(input))
|
||||
input = None
|
||||
else:
|
||||
command = None
|
||||
thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)
|
||||
return sanitized_config, input, command, thread_id
|
||||
|
||||
def _get_stream_modes(
|
||||
self,
|
||||
stream_mode: StreamMode | list[StreamMode] | None,
|
||||
@@ -715,17 +735,12 @@ class RemoteGraph(PregelProtocol):
|
||||
The output of the graph.
|
||||
"""
|
||||
sync_client = self._validate_sync_client()
|
||||
merged_config = merge_configs(self.config, config)
|
||||
sanitized_config = self._sanitize_config(merged_config)
|
||||
sanitized_config, input, command, thread_id = self._prepare_run_input(
|
||||
input, config
|
||||
)
|
||||
stream_modes, requested, req_single, stream = self._get_stream_modes(
|
||||
stream_mode, config
|
||||
)
|
||||
if isinstance(input, Command):
|
||||
command: CommandSDK | None = cast(CommandSDK, asdict(input))
|
||||
input = None
|
||||
else:
|
||||
command = None
|
||||
thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)
|
||||
|
||||
for chunk in sync_client.runs.stream(
|
||||
thread_id=thread_id,
|
||||
@@ -825,17 +840,12 @@ class RemoteGraph(PregelProtocol):
|
||||
The output of the graph.
|
||||
"""
|
||||
client = self._validate_client()
|
||||
merged_config = merge_configs(self.config, config)
|
||||
sanitized_config = self._sanitize_config(merged_config)
|
||||
sanitized_config, input, command, thread_id = self._prepare_run_input(
|
||||
input, config
|
||||
)
|
||||
stream_modes, requested, req_single, stream = self._get_stream_modes(
|
||||
stream_mode, config
|
||||
)
|
||||
if isinstance(input, Command):
|
||||
command: CommandSDK | None = cast(CommandSDK, asdict(input))
|
||||
input = None
|
||||
else:
|
||||
command = None
|
||||
thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)
|
||||
|
||||
async for chunk in client.runs.stream(
|
||||
thread_id=thread_id,
|
||||
@@ -937,27 +947,31 @@ class RemoteGraph(PregelProtocol):
|
||||
interrupt_before: Interrupt the graph before these nodes.
|
||||
interrupt_after: Interrupt the graph after these nodes.
|
||||
headers: Additional headers to pass to the request.
|
||||
**kwargs: Additional params to pass to RemoteGraph.stream.
|
||||
**kwargs: Additional params to pass to client.runs.wait.
|
||||
|
||||
Returns:
|
||||
The output of the graph.
|
||||
"""
|
||||
for chunk in self.stream(
|
||||
input,
|
||||
config=config,
|
||||
sync_client = self._validate_sync_client()
|
||||
sanitized_config, input, command, thread_id = self._prepare_run_input(
|
||||
input, config
|
||||
)
|
||||
|
||||
return sync_client.runs.wait( # type: ignore
|
||||
thread_id=thread_id,
|
||||
assistant_id=self.assistant_id,
|
||||
input=input,
|
||||
command=command,
|
||||
config=sanitized_config,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
headers=headers,
|
||||
stream_mode="values",
|
||||
if_not_exists="create",
|
||||
headers=(
|
||||
_merge_tracing_headers(headers) if self.distributed_tracing else headers
|
||||
),
|
||||
params=params,
|
||||
**kwargs,
|
||||
):
|
||||
pass
|
||||
try:
|
||||
return chunk
|
||||
except UnboundLocalError:
|
||||
logger.warning("No events received from remote graph")
|
||||
return None
|
||||
)
|
||||
|
||||
async def ainvoke(
|
||||
self,
|
||||
@@ -978,27 +992,31 @@ class RemoteGraph(PregelProtocol):
|
||||
interrupt_before: Interrupt the graph before these nodes.
|
||||
interrupt_after: Interrupt the graph after these nodes.
|
||||
headers: Additional headers to pass to the request.
|
||||
**kwargs: Additional params to pass to RemoteGraph.astream.
|
||||
**kwargs: Additional params to pass to client.runs.wait.
|
||||
|
||||
Returns:
|
||||
The output of the graph.
|
||||
"""
|
||||
async for chunk in self.astream(
|
||||
input,
|
||||
config=config,
|
||||
client = self._validate_client()
|
||||
sanitized_config, input, command, thread_id = self._prepare_run_input(
|
||||
input, config
|
||||
)
|
||||
|
||||
return await client.runs.wait(
|
||||
thread_id=thread_id,
|
||||
assistant_id=self.assistant_id,
|
||||
input=input,
|
||||
command=command,
|
||||
config=sanitized_config,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
headers=headers,
|
||||
stream_mode="values",
|
||||
if_not_exists="create",
|
||||
headers=(
|
||||
_merge_tracing_headers(headers) if self.distributed_tracing else headers
|
||||
),
|
||||
params=params,
|
||||
**kwargs,
|
||||
):
|
||||
pass
|
||||
try:
|
||||
return chunk
|
||||
except UnboundLocalError:
|
||||
logger.warning("No events received from remote graph")
|
||||
return None
|
||||
)
|
||||
|
||||
|
||||
def _merge_tracing_headers(headers: dict[str, str] | None) -> dict[str, str] | None:
|
||||
|
||||
@@ -818,13 +818,9 @@ async def test_astream():
|
||||
def test_invoke():
|
||||
# set up test
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.runs.stream.return_value = [
|
||||
StreamPart(event="values", data={"chunk": "data1"}),
|
||||
StreamPart(event="values", data={"chunk": "data2"}),
|
||||
StreamPart(
|
||||
event="values", data={"messages": [{"type": "human", "content": "world"}]}
|
||||
),
|
||||
]
|
||||
mock_sync_client.runs.wait.return_value = {
|
||||
"messages": [{"type": "human", "content": "world"}]
|
||||
}
|
||||
|
||||
# call method / assertions
|
||||
remote_pregel = RemoteGraph(
|
||||
@@ -838,13 +834,19 @@ def test_invoke():
|
||||
)
|
||||
|
||||
assert result == {"messages": [{"type": "human", "content": "world"}]}
|
||||
# verify runs.wait was called with expected args
|
||||
assert mock_sync_client.runs.wait.called
|
||||
_, kwargs = mock_sync_client.runs.wait.call_args
|
||||
assert kwargs.get("thread_id") == "thread_1"
|
||||
assert kwargs.get("assistant_id") == "test_graph_id"
|
||||
assert kwargs.get("if_not_exists") == "create"
|
||||
|
||||
|
||||
def test_invoke_sanitizes_thread_id():
|
||||
# Ensure that invoking with thread_id passes thread_id as a top-level arg
|
||||
# and removes it from the config body.
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.runs.stream.return_value = []
|
||||
mock_sync_client.runs.wait.return_value = {}
|
||||
remote_pregel = RemoteGraph("test_graph_id", sync_client=mock_sync_client)
|
||||
|
||||
config = {"configurable": {"thread_id": "thread_1"}}
|
||||
@@ -852,8 +854,8 @@ def test_invoke_sanitizes_thread_id():
|
||||
{"input": {"messages": [{"type": "human", "content": "hello"}]}}, config
|
||||
)
|
||||
|
||||
assert mock_sync_client.runs.stream.called
|
||||
_, kwargs = mock_sync_client.runs.stream.call_args
|
||||
assert mock_sync_client.runs.wait.called
|
||||
_, kwargs = mock_sync_client.runs.wait.call_args
|
||||
assert kwargs.get("thread_id") == "thread_1"
|
||||
passed_config = kwargs.get("config") or {}
|
||||
assert "configurable" in passed_config
|
||||
@@ -883,16 +885,10 @@ def test_stream_sanitizes_thread_id():
|
||||
@pytest.mark.anyio
|
||||
async def test_ainvoke():
|
||||
# set up test
|
||||
mock_async_client = MagicMock()
|
||||
async_iter = MagicMock()
|
||||
async_iter.__aiter__.return_value = [
|
||||
StreamPart(event="values", data={"chunk": "data1"}),
|
||||
StreamPart(event="values", data={"chunk": "data2"}),
|
||||
StreamPart(
|
||||
event="values", data={"messages": [{"type": "human", "content": "world"}]}
|
||||
),
|
||||
]
|
||||
mock_async_client.runs.stream.return_value = async_iter
|
||||
mock_async_client = AsyncMock()
|
||||
mock_async_client.runs.wait.return_value = {
|
||||
"messages": [{"type": "human", "content": "world"}]
|
||||
}
|
||||
|
||||
# call method / assertions
|
||||
remote_pregel = RemoteGraph(
|
||||
@@ -906,6 +902,12 @@ async def test_ainvoke():
|
||||
)
|
||||
|
||||
assert result == {"messages": [{"type": "human", "content": "world"}]}
|
||||
# verify runs.wait was called with expected args
|
||||
assert mock_async_client.runs.wait.called
|
||||
_, kwargs = mock_async_client.runs.wait.call_args
|
||||
assert kwargs.get("thread_id") == "thread_1"
|
||||
assert kwargs.get("assistant_id") == "test_graph_id"
|
||||
assert kwargs.get("if_not_exists") == "create"
|
||||
|
||||
|
||||
@pytest.mark.skip(
|
||||
@@ -1241,12 +1243,18 @@ async def test_include_headers(
|
||||
async_iter.__aiter__.return_value = return_value
|
||||
astream_mock = mock_async_client.runs.stream
|
||||
astream_mock.return_value = async_iter
|
||||
# Mock for ainvoke which uses runs.wait
|
||||
await_mock = AsyncMock(return_value={"chunk": "data1"})
|
||||
mock_async_client.runs.wait = await_mock
|
||||
|
||||
mock_sync_client = MagicMock()
|
||||
sync_iter = MagicMock()
|
||||
sync_iter.__iter__.return_value = return_value
|
||||
stream_mock = mock_sync_client.runs.stream
|
||||
stream_mock.return_value = async_iter
|
||||
# Mock for invoke which uses runs.wait
|
||||
wait_mock = MagicMock(return_value={"chunk": "data1"})
|
||||
mock_sync_client.runs.wait = wait_mock
|
||||
|
||||
remote_pregel = RemoteGraph(
|
||||
"test_graph_id",
|
||||
@@ -1279,8 +1287,12 @@ async def test_include_headers(
|
||||
expected["langsmith-trace"] = AnyStr()
|
||||
expected["baggage"] = AnyStr("langsmith-metadata=")
|
||||
|
||||
assert astream_mock.call_args.kwargs["headers"] == expected
|
||||
if stream:
|
||||
assert astream_mock.call_args.kwargs["headers"] == expected
|
||||
else:
|
||||
assert await_mock.call_args.kwargs["headers"] == expected
|
||||
stream_mock.assert_not_called()
|
||||
wait_mock.assert_not_called()
|
||||
|
||||
with ls.tracing_context(enabled=True, client=MagicMock()):
|
||||
with ls.trace("foo"):
|
||||
@@ -1298,4 +1310,7 @@ async def test_include_headers(
|
||||
config,
|
||||
headers=headers,
|
||||
)
|
||||
assert stream_mock.call_args.kwargs["headers"] == expected
|
||||
if stream:
|
||||
assert stream_mock.call_args.kwargs["headers"] == expected
|
||||
else:
|
||||
assert wait_mock.call_args.kwargs["headers"] == expected
|
||||
|
||||
@@ -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