mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-03 06:55:13 +02:00
Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
18a2a2a832 | ||
|
|
5ae0aff522 | ||
|
|
b7d068677c | ||
|
|
706f3e981e | ||
|
|
344ab65351 | ||
|
|
94698c8a34 | ||
|
|
7d557cb865 |
@@ -67,7 +67,7 @@ While LangGraph can be used standalone, it also integrates seamlessly with any L
|
||||
- [LangChain](https://docs.langchain.com/oss/python/langchain/overview) – Provides integrations and composable components to streamline LLM application development.
|
||||
|
||||
> [!NOTE]
|
||||
> Looking for the JS version of LangGraph? See the [JS repo](https://github.com/langchain-ai/langgraphjs) and the [JS docs](https://langchain-ai.github.io/langgraphjs/).
|
||||
> Looking for the JS version of LangGraph? See the [JS repo](https://github.com/langchain-ai/langgraphjs) and the [JS docs](https://docs.langchain.com/oss/javascript/langgraph/overview).
|
||||
|
||||
## Additional resources
|
||||
|
||||
|
||||
@@ -329,8 +329,9 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
FROM checkpoints
|
||||
{where}
|
||||
ORDER BY checkpoint_id DESC"""
|
||||
if limit:
|
||||
query += f" LIMIT {limit}"
|
||||
if limit is not None:
|
||||
query += " LIMIT ?"
|
||||
param_values = (*param_values, limit)
|
||||
with self.cursor(transaction=False) as cur, closing(self.conn.cursor()) as wcur:
|
||||
cur.execute(query, param_values)
|
||||
for (
|
||||
|
||||
@@ -425,8 +425,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
FROM checkpoints
|
||||
{where}
|
||||
ORDER BY checkpoint_id DESC"""
|
||||
if limit:
|
||||
query += f" LIMIT {limit}"
|
||||
if limit is not None:
|
||||
query += " LIMIT ?"
|
||||
params = (*params, limit)
|
||||
async with (
|
||||
self.lock,
|
||||
self.conn.execute(query, params) as cur,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
@@ -8,6 +9,23 @@ from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.checkpoint.base import get_checkpoint_id
|
||||
|
||||
|
||||
def _validate_filter_key(key: str) -> None:
|
||||
"""Validate that a filter key is safe for use in SQL queries.
|
||||
|
||||
Args:
|
||||
key: The filter key to validate
|
||||
|
||||
Raises:
|
||||
ValueError: If the key contains invalid characters that could enable SQL injection
|
||||
"""
|
||||
# Allow alphanumeric characters, underscores, dots, and hyphens
|
||||
# This covers typical JSON property names while preventing SQL injection
|
||||
if not re.match(r"^[a-zA-Z0-9_.-]+$", key):
|
||||
raise ValueError(
|
||||
f"Invalid filter key: '{key}'. Filter keys must contain only alphanumeric characters, underscores, dots, and hyphens."
|
||||
)
|
||||
|
||||
|
||||
def _metadata_predicate(
|
||||
metadata_filter: dict[str, Any],
|
||||
) -> tuple[Sequence[str], Sequence[Any]]:
|
||||
@@ -43,6 +61,7 @@ def _metadata_predicate(
|
||||
|
||||
# process metadata query
|
||||
for query_key, query_value in metadata_filter.items():
|
||||
_validate_filter_key(query_key)
|
||||
operator, param_value = _where_value(query_value)
|
||||
predicates.append(
|
||||
f"json_extract(CAST(metadata AS TEXT), '$.{query_key}') {operator}"
|
||||
|
||||
@@ -404,12 +404,9 @@ class BaseSqliteStore:
|
||||
# SQLite json_extract returns unquoted string values
|
||||
if isinstance(value, str):
|
||||
filter_conditions.append(
|
||||
"json_extract(value, '$."
|
||||
+ key
|
||||
+ "') = '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'"
|
||||
"json_extract(value, '$." + key + "') = ?"
|
||||
)
|
||||
filter_params.append(value)
|
||||
elif value is None:
|
||||
filter_conditions.append(
|
||||
"json_extract(value, '$." + key + "') IS NULL"
|
||||
@@ -423,9 +420,11 @@ class BaseSqliteStore:
|
||||
+ ("1" if value else "0")
|
||||
)
|
||||
elif isinstance(value, (int, float)):
|
||||
# Use parameterized query to handle special floats and large integers
|
||||
filter_conditions.append(
|
||||
"json_extract(value, '$." + key + "') = " + str(value)
|
||||
"json_extract(value, '$." + key + "') = ?"
|
||||
)
|
||||
filter_params.append(float(value))
|
||||
else:
|
||||
# Complex objects (list, dict, …) – compare JSON text
|
||||
filter_conditions.append(
|
||||
@@ -636,85 +635,66 @@ class BaseSqliteStore:
|
||||
# We need to properly format values for SQLite JSON extraction comparison
|
||||
if op == "$eq":
|
||||
if isinstance(value, str):
|
||||
# Direct string comparison with proper quoting for unquoted json_extract result
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') = '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') = ?", [value]
|
||||
elif value is None:
|
||||
return f"json_extract(value, '$.{key}') IS NULL", []
|
||||
elif isinstance(value, bool):
|
||||
# SQLite JSON stores booleans as integers
|
||||
return f"json_extract(value, '$.{key}') = {1 if value else 0}", []
|
||||
elif isinstance(value, (int, float)):
|
||||
return f"json_extract(value, '$.{key}') = {value}", []
|
||||
# Convert to float to handle inf, -inf, nan, and very large integers
|
||||
# SQLite REAL can handle these cases better than INTEGER
|
||||
return f"json_extract(value, '$.{key}') = ?", [float(value)]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') = ?", [orjson.dumps(value)]
|
||||
elif op == "$gt":
|
||||
# For numeric values, SQLite needs to compare as numbers, not strings
|
||||
if isinstance(value, (int, float)):
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) > {value}", []
|
||||
# Convert to float to handle special values and very large integers
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) > ?", [
|
||||
float(value)
|
||||
]
|
||||
elif isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') > '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') > ?", [value]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') > ?", [orjson.dumps(value)]
|
||||
elif op == "$gte":
|
||||
if isinstance(value, (int, float)):
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) >= {value}", []
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) >= ?", [
|
||||
float(value)
|
||||
]
|
||||
elif isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') >= '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') >= ?", [value]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') >= ?", [orjson.dumps(value)]
|
||||
elif op == "$lt":
|
||||
if isinstance(value, (int, float)):
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) < {value}", []
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) < ?", [
|
||||
float(value)
|
||||
]
|
||||
elif isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') < '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') < ?", [value]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') < ?", [orjson.dumps(value)]
|
||||
elif op == "$lte":
|
||||
if isinstance(value, (int, float)):
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) <= {value}", []
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) <= ?", [
|
||||
float(value)
|
||||
]
|
||||
elif isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') <= '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') <= ?", [value]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') <= ?", [orjson.dumps(value)]
|
||||
elif op == "$ne":
|
||||
if isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') != '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') != ?", [value]
|
||||
elif value is None:
|
||||
return f"json_extract(value, '$.{key}') IS NOT NULL", []
|
||||
elif isinstance(value, bool):
|
||||
return f"json_extract(value, '$.{key}') != {1 if value else 0}", []
|
||||
elif isinstance(value, (int, float)):
|
||||
return f"json_extract(value, '$.{key}') != {value}", []
|
||||
# Convert to float for consistency
|
||||
return f"json_extract(value, '$.{key}') != ?", [float(value)]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') != ?", [orjson.dumps(value)]
|
||||
else:
|
||||
@@ -874,85 +854,66 @@ class SqliteStore(BaseSqliteStore, BaseStore):
|
||||
# We need to properly format values for SQLite JSON extraction comparison
|
||||
if op == "$eq":
|
||||
if isinstance(value, str):
|
||||
# Direct string comparison with proper quoting for unquoted json_extract result
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') = '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') = ?", [value]
|
||||
elif value is None:
|
||||
return f"json_extract(value, '$.{key}') IS NULL", []
|
||||
elif isinstance(value, bool):
|
||||
# SQLite JSON stores booleans as integers
|
||||
return f"json_extract(value, '$.{key}') = {1 if value else 0}", []
|
||||
elif isinstance(value, (int, float)):
|
||||
return f"json_extract(value, '$.{key}') = {value}", []
|
||||
# Convert to float to handle inf, -inf, nan, and very large integers
|
||||
# SQLite REAL can handle these cases better than INTEGER
|
||||
return f"json_extract(value, '$.{key}') = ?", [float(value)]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') = ?", [orjson.dumps(value)]
|
||||
elif op == "$gt":
|
||||
# For numeric values, SQLite needs to compare as numbers, not strings
|
||||
if isinstance(value, (int, float)):
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) > {value}", []
|
||||
# Convert to float to handle special values and very large integers
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) > ?", [
|
||||
float(value)
|
||||
]
|
||||
elif isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') > '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') > ?", [value]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') > ?", [orjson.dumps(value)]
|
||||
elif op == "$gte":
|
||||
if isinstance(value, (int, float)):
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) >= {value}", []
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) >= ?", [
|
||||
float(value)
|
||||
]
|
||||
elif isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') >= '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') >= ?", [value]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') >= ?", [orjson.dumps(value)]
|
||||
elif op == "$lt":
|
||||
if isinstance(value, (int, float)):
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) < {value}", []
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) < ?", [
|
||||
float(value)
|
||||
]
|
||||
elif isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') < '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') < ?", [value]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') < ?", [orjson.dumps(value)]
|
||||
elif op == "$lte":
|
||||
if isinstance(value, (int, float)):
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) <= {value}", []
|
||||
return f"CAST(json_extract(value, '$.{key}') AS REAL) <= ?", [
|
||||
float(value)
|
||||
]
|
||||
elif isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') <= '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') <= ?", [value]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') <= ?", [orjson.dumps(value)]
|
||||
elif op == "$ne":
|
||||
if isinstance(value, str):
|
||||
return (
|
||||
f"json_extract(value, '$.{key}') != '"
|
||||
+ value.replace("'", "''")
|
||||
+ "'",
|
||||
[],
|
||||
)
|
||||
return f"json_extract(value, '$.{key}') != ?", [value]
|
||||
elif value is None:
|
||||
return f"json_extract(value, '$.{key}') IS NOT NULL", []
|
||||
elif isinstance(value, bool):
|
||||
return f"json_extract(value, '$.{key}') != {1 if value else 0}", []
|
||||
elif isinstance(value, (int, float)):
|
||||
return f"json_extract(value, '$.{key}') != {value}", []
|
||||
# Convert to float for consistency
|
||||
return f"json_extract(value, '$.{key}') != ?", [float(value)]
|
||||
else:
|
||||
return f"json_extract(value, '$.{key}') != ?", [orjson.dumps(value)]
|
||||
else:
|
||||
|
||||
@@ -113,4 +113,78 @@ class TestAsyncSqliteSaver:
|
||||
search_results_5[1].config["configurable"]["checkpoint_ns"],
|
||||
} == {"", "inner"}
|
||||
|
||||
# TODO: test before and limit params
|
||||
# Test limit param
|
||||
search_results_6 = [
|
||||
c
|
||||
async for c in saver.alist(
|
||||
{"configurable": {"thread_id": "thread-2"}}, limit=1
|
||||
)
|
||||
]
|
||||
assert len(search_results_6) == 1
|
||||
assert search_results_6[0].config["configurable"]["thread_id"] == "thread-2"
|
||||
|
||||
# Test before param
|
||||
search_results_7 = [
|
||||
c async for c in saver.alist(None, before=search_results_5[1].config)
|
||||
]
|
||||
assert len(search_results_7) == 1
|
||||
assert search_results_7[0].config["configurable"]["thread_id"] == "thread-1"
|
||||
|
||||
async def test_limit_parameter_sql_injection_prevention(self) -> None:
|
||||
"""Test that the limit parameter properly uses parameterized queries to prevent SQL injection."""
|
||||
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
|
||||
# Setup: Create multiple checkpoints
|
||||
for i in range(5):
|
||||
config: RunnableConfig = {
|
||||
"configurable": {
|
||||
"thread_id": f"thread-{i}",
|
||||
"checkpoint_ns": "",
|
||||
}
|
||||
}
|
||||
checkpoint = empty_checkpoint()
|
||||
metadata: CheckpointMetadata = {"index": i}
|
||||
await saver.aput(config, checkpoint, metadata, {})
|
||||
|
||||
# Test that limit works correctly with valid integer
|
||||
results = [c async for c in saver.alist(None, limit=2)]
|
||||
assert len(results) == 2
|
||||
|
||||
# Test that limit=0 returns no results
|
||||
results = [c async for c in saver.alist(None, limit=0)]
|
||||
assert len(results) == 0
|
||||
|
||||
# Test that limit=None returns all results
|
||||
results = [c async for c in saver.alist(None, limit=None)]
|
||||
assert len(results) == 5
|
||||
|
||||
# Test explicit SQL injection attempt via limit parameter
|
||||
# Even if type checking is bypassed and a malicious string is passed,
|
||||
# the parameterized query will treat it as a value, not SQL code
|
||||
# This would cause an error (can't convert string to int for LIMIT),
|
||||
# which is the correct secure behavior
|
||||
malicious_limits = [
|
||||
"1; DROP TABLE checkpoints; --",
|
||||
"1 OR 1=1",
|
||||
"999999 UNION SELECT * FROM checkpoints",
|
||||
]
|
||||
|
||||
for malicious_limit in malicious_limits:
|
||||
# The parameterized query should safely reject non-integer limits
|
||||
# or convert them in a way that prevents SQL injection
|
||||
try:
|
||||
# Bypass type checking by casting
|
||||
results = [
|
||||
c
|
||||
async for c in saver.alist(None, limit=malicious_limit) # type: ignore
|
||||
]
|
||||
# If it doesn't raise an error, it should at least not execute the injection
|
||||
# SQLite's parameter binding will try to convert the string to an integer
|
||||
# which will either fail or treat it as 0
|
||||
except Exception:
|
||||
# Expected: SQLite should reject invalid limit values
|
||||
pass
|
||||
|
||||
# Verify the checkpoints table still exists and has all data
|
||||
# (would have been dropped if injection succeeded)
|
||||
results = [c async for c in saver.alist(None, limit=None)]
|
||||
assert len(results) == 5
|
||||
|
||||
@@ -182,3 +182,120 @@ class TestSqliteSaver:
|
||||
with pytest.raises(NotImplementedError, match="AsyncSqliteSaver"):
|
||||
async for _ in saver.alist(self.config_1):
|
||||
pass
|
||||
|
||||
def test_metadata_predicate_sql_injection_prevention(self) -> None:
|
||||
"""Test that _metadata_predicate rejects malicious filter keys."""
|
||||
# Test various SQL injection payloads
|
||||
malicious_keys = [
|
||||
"x') OR '1'='1", # Boolean-based injection
|
||||
"x') OR 1=1 --", # Comment-based injection
|
||||
"x') UNION SELECT 1,2,3,4,5,6,7 --", # UNION-based injection
|
||||
"access') = 'public' OR '1'='1' OR json_extract(value, '$.", # Complex injection
|
||||
"'; DROP TABLE checkpoints; --", # Destructive injection
|
||||
]
|
||||
|
||||
for malicious_key in malicious_keys:
|
||||
with pytest.raises(ValueError, match="Invalid filter key"):
|
||||
_metadata_predicate({malicious_key: "dummy"})
|
||||
|
||||
def test_checkpoint_search_sql_injection_prevention(self) -> None:
|
||||
"""Test that SQL injection via malicious filter keys is prevented in checkpoint search."""
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
# Setup: Create checkpoints with different metadata
|
||||
config_public: RunnableConfig = {
|
||||
"configurable": {
|
||||
"thread_id": "thread-public",
|
||||
"checkpoint_ns": "",
|
||||
}
|
||||
}
|
||||
config_private: RunnableConfig = {
|
||||
"configurable": {
|
||||
"thread_id": "thread-private",
|
||||
"checkpoint_ns": "",
|
||||
}
|
||||
}
|
||||
|
||||
checkpoint_public = empty_checkpoint()
|
||||
checkpoint_private = empty_checkpoint()
|
||||
|
||||
metadata_public: CheckpointMetadata = {
|
||||
"access": "public",
|
||||
"data": "public information",
|
||||
}
|
||||
metadata_private: CheckpointMetadata = {
|
||||
"access": "private",
|
||||
"data": "secret information",
|
||||
"password": "secret123",
|
||||
}
|
||||
|
||||
saver.put(config_public, checkpoint_public, metadata_public, {})
|
||||
saver.put(config_private, checkpoint_private, metadata_private, {})
|
||||
|
||||
# Normal query - should return only public checkpoint
|
||||
normal_results = list(saver.list(None, filter={"access": "public"}))
|
||||
assert len(normal_results) == 1
|
||||
assert normal_results[0].metadata["access"] == "public"
|
||||
|
||||
# SQL injection attempt should raise ValueError
|
||||
malicious_key = (
|
||||
"access') = 'public' OR '1'='1' OR json_extract(metadata, '$."
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid filter key"):
|
||||
list(saver.list(None, filter={malicious_key: "dummy"}))
|
||||
|
||||
def test_limit_parameter_sql_injection_prevention(self) -> None:
|
||||
"""Test that the limit parameter properly uses parameterized queries to prevent SQL injection."""
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
# Setup: Create multiple checkpoints
|
||||
for i in range(5):
|
||||
config: RunnableConfig = {
|
||||
"configurable": {
|
||||
"thread_id": f"thread-{i}",
|
||||
"checkpoint_ns": "",
|
||||
}
|
||||
}
|
||||
checkpoint = empty_checkpoint()
|
||||
metadata: CheckpointMetadata = {"index": i}
|
||||
saver.put(config, checkpoint, metadata, {})
|
||||
|
||||
# Test that limit works correctly with valid integer
|
||||
results = list(saver.list(None, limit=2))
|
||||
assert len(results) == 2
|
||||
|
||||
# Test that limit=0 returns no results
|
||||
results = list(saver.list(None, limit=0))
|
||||
assert len(results) == 0
|
||||
|
||||
# Test that limit=None returns all results
|
||||
results = list(saver.list(None, limit=None))
|
||||
assert len(results) == 5
|
||||
|
||||
# Test explicit SQL injection attempt via limit parameter
|
||||
# Even if type checking is bypassed and a malicious string is passed,
|
||||
# the parameterized query will treat it as a value, not SQL code
|
||||
# This would cause an error (can't convert string to int for LIMIT),
|
||||
# which is the correct secure behavior
|
||||
malicious_limits = [
|
||||
"1; DROP TABLE checkpoints; --",
|
||||
"1 OR 1=1",
|
||||
"999999 UNION SELECT * FROM checkpoints",
|
||||
]
|
||||
|
||||
for malicious_limit in malicious_limits:
|
||||
# The parameterized query should safely reject non-integer limits
|
||||
# or convert them in a way that prevents SQL injection
|
||||
try:
|
||||
# Bypass type checking by casting
|
||||
results = list(saver.list(None, limit=malicious_limit)) # type: ignore
|
||||
# If it doesn't raise an error, it should at least not execute the injection
|
||||
# SQLite's parameter binding will try to convert the string to an integer
|
||||
# which will either fail or treat it as 0
|
||||
except Exception:
|
||||
# Expected: SQLite should reject invalid limit values
|
||||
pass
|
||||
|
||||
# Verify the checkpoints table still exists and has all data
|
||||
# (would have been dropped if injection succeeded)
|
||||
results = list(saver.list(None, limit=None))
|
||||
assert len(results) == 5
|
||||
|
||||
@@ -1069,6 +1069,110 @@ def test_sql_injection_vulnerability(store: SqliteStore) -> None:
|
||||
store.search(("docs",), filter={malicious_key: "dummy"})
|
||||
|
||||
|
||||
def test_sql_injection_filter_values(store: SqliteStore) -> None:
|
||||
"""Test that SQL injection via malicious filter values is properly escaped."""
|
||||
# Setup: Create documents with different access levels
|
||||
store.put(("docs",), "doc1", {"access": "public", "title": "Public Document"})
|
||||
store.put(("docs",), "doc2", {"access": "private", "title": "Private Document"})
|
||||
store.put(("docs",), "doc3", {"access": "secret", "title": "Secret Document"})
|
||||
|
||||
# Test 1: Basic SQL injection attempt with single quote
|
||||
malicious_value = "public' OR '1'='1"
|
||||
results = store.search(("docs",), filter={"access": malicious_value})
|
||||
# Should return 0 results because the malicious value is escaped and won't match anything
|
||||
assert len(results) == 0, "SQL injection via string value should be blocked"
|
||||
|
||||
# Test 2: SQL injection with comment
|
||||
malicious_value = "public'; --"
|
||||
results = store.search(("docs",), filter={"access": malicious_value})
|
||||
assert len(results) == 0, "SQL comment injection should be blocked"
|
||||
|
||||
# Test 3: UNION injection attempt
|
||||
malicious_value = "public' UNION SELECT * FROM store --"
|
||||
results = store.search(("docs",), filter={"access": malicious_value})
|
||||
assert len(results) == 0, "UNION injection should be blocked"
|
||||
|
||||
# Test 4: Parameterized queries handle strings with null bytes and SQL injection attempts safely
|
||||
malicious_value = "public\x00' OR '1'='1"
|
||||
results = store.search(("docs",), filter={"access": malicious_value})
|
||||
assert len(results) == 0, (
|
||||
"Parameterized queries treat injection attempts as literal strings"
|
||||
)
|
||||
|
||||
# Test 5: Multiple single quotes
|
||||
malicious_value = "''''"
|
||||
results = store.search(("docs",), filter={"access": malicious_value})
|
||||
assert len(results) == 0, "Multiple quotes should be handled safely"
|
||||
|
||||
# Test 6: Legitimate value with single quote should work
|
||||
store.put(("docs",), "doc4", {"title": "O'Brien's Document", "access": "public"})
|
||||
results = store.search(("docs",), filter={"title": "O'Brien's Document"})
|
||||
assert len(results) == 1, "Legitimate single quotes should work"
|
||||
assert results[0].value["title"] == "O'Brien's Document"
|
||||
|
||||
# Test 7: Unicode characters with injection attempt
|
||||
malicious_value = "public' OR 'א'='א"
|
||||
results = store.search(("docs",), filter={"access": malicious_value})
|
||||
assert len(results) == 0, "Unicode-based injection should be blocked"
|
||||
|
||||
|
||||
def test_numeric_filter_safety(store: SqliteStore) -> None:
|
||||
"""Test that numeric filter values are handled safely."""
|
||||
# Setup: Create documents with numeric fields
|
||||
store.put(("items",), "item1", {"price": 10, "quantity": 5})
|
||||
store.put(("items",), "item2", {"price": 20, "quantity": 3})
|
||||
store.put(("items",), "item3", {"price": 30, "quantity": 1})
|
||||
|
||||
# Test 1: Normal numeric comparison
|
||||
results = store.search(("items",), filter={"price": {"$gt": 15}})
|
||||
assert len(results) == 2
|
||||
assert all(r.value["price"] > 15 for r in results)
|
||||
|
||||
# Test 2: Special float values (infinity)
|
||||
results = store.search(("items",), filter={"price": {"$lt": float("inf")}})
|
||||
assert len(results) == 3, "All finite values should be less than infinity"
|
||||
|
||||
# Test 3: Special float values (negative infinity)
|
||||
results = store.search(("items",), filter={"price": {"$gt": float("-inf")}})
|
||||
assert len(results) == 3, (
|
||||
"All finite values should be greater than negative infinity"
|
||||
)
|
||||
|
||||
# Test 4: NaN handling - NaN comparisons should not cause errors
|
||||
try:
|
||||
results = store.search(("items",), filter={"price": {"$eq": float("nan")}})
|
||||
# NaN never equals anything, including itself, so should return 0 results
|
||||
assert len(results) == 0
|
||||
except Exception as e:
|
||||
pytest.fail(f"NaN handling should not raise exception: {e}")
|
||||
|
||||
# Test 5: Very large numbers
|
||||
results = store.search(("items",), filter={"price": {"$lt": 10**100}})
|
||||
assert len(results) == 3, "Very large numbers should be handled safely"
|
||||
|
||||
# Test 6: Negative numbers
|
||||
store.put(("items",), "item4", {"price": -10, "quantity": 0})
|
||||
results = store.search(("items",), filter={"price": {"$lt": 0}})
|
||||
assert len(results) == 1
|
||||
assert results[0].key == "item4"
|
||||
|
||||
|
||||
def test_boolean_filter_safety(store: SqliteStore) -> None:
|
||||
"""Test that boolean filter values are handled safely."""
|
||||
store.put(("flags",), "flag1", {"active": True, "name": "Feature A"})
|
||||
store.put(("flags",), "flag2", {"active": False, "name": "Feature B"})
|
||||
store.put(("flags",), "flag3", {"active": True, "name": "Feature C"})
|
||||
|
||||
# Test boolean filters
|
||||
results = store.search(("flags",), filter={"active": True})
|
||||
assert len(results) == 2
|
||||
assert all(r.value["active"] is True for r in results)
|
||||
|
||||
results = store.search(("flags",), filter={"active": False})
|
||||
assert len(results) == 1
|
||||
assert results[0].value["active"] is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("distance_type", VECTOR_TYPES)
|
||||
def test_non_ascii(
|
||||
fake_embeddings: CharacterEmbeddings,
|
||||
|
||||
@@ -154,6 +154,7 @@ def validate_config(config: Config) -> Config:
|
||||
"env": config.get("env", {}),
|
||||
"store": config.get("store"),
|
||||
"auth": config.get("auth"),
|
||||
"encryption": config.get("encryption"),
|
||||
"http": config.get("http"),
|
||||
"checkpointer": config.get("checkpointer"),
|
||||
"ui": config.get("ui"),
|
||||
@@ -228,6 +229,14 @@ def validate_config(config: Config) -> Config:
|
||||
f"Invalid auth.path format: '{auth_conf['path']}'. "
|
||||
"Must be in format './path/to/file.py:attribute_name'"
|
||||
)
|
||||
# Validate encryption config
|
||||
if encryption_conf := config.get("encryption"):
|
||||
if "path" in encryption_conf:
|
||||
if ":" not in encryption_conf["path"]:
|
||||
raise ValueError(
|
||||
f"Invalid encryption.path format: '{encryption_conf['path']}'. "
|
||||
"Must be in format './path/to/file.py:attribute_name'"
|
||||
)
|
||||
if http_conf := config.get("http"):
|
||||
if "app" in http_conf:
|
||||
if ":" not in http_conf["app"]:
|
||||
@@ -614,6 +623,49 @@ def _update_auth_path(
|
||||
)
|
||||
|
||||
|
||||
def _update_encryption_path(
|
||||
config_path: pathlib.Path, config: Config, local_deps: LocalDeps
|
||||
) -> None:
|
||||
"""Update encryption.path to use Docker container paths."""
|
||||
encryption_conf = config.get("encryption")
|
||||
if not encryption_conf or not (path_str := encryption_conf.get("path")):
|
||||
return
|
||||
|
||||
module_str, sep, attr_str = path_str.partition(":")
|
||||
if not sep or not module_str.startswith("."):
|
||||
return # Already validated or absolute path
|
||||
|
||||
resolved = config_path.parent / module_str
|
||||
if not resolved.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Encryption file not found: {resolved} (from {path_str})"
|
||||
)
|
||||
if not resolved.is_file():
|
||||
raise IsADirectoryError(f"Encryption path must be a file: {resolved}")
|
||||
|
||||
# Check faux packages first (higher priority)
|
||||
for faux_path, (_, destpath) in local_deps.faux_pkgs.items():
|
||||
if resolved.is_relative_to(faux_path):
|
||||
new_path = f"{destpath}/{resolved.relative_to(faux_path)}:{attr_str}"
|
||||
encryption_conf["path"] = new_path
|
||||
return
|
||||
|
||||
# Check real packages
|
||||
for real_path in local_deps.real_pkgs:
|
||||
if resolved.is_relative_to(real_path):
|
||||
new_path = (
|
||||
f"/deps/{real_path.name}/{resolved.relative_to(real_path)}:{attr_str}"
|
||||
)
|
||||
encryption_conf["path"] = new_path
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"Encryption file '{resolved}' not covered by dependencies.\n"
|
||||
"Add its parent directory to the 'dependencies' array in your config.\n"
|
||||
f"Current dependencies: {config['dependencies']}"
|
||||
)
|
||||
|
||||
|
||||
def _update_http_app_path(
|
||||
config_path: pathlib.Path, config: Config, local_deps: LocalDeps
|
||||
) -> None:
|
||||
@@ -809,6 +861,8 @@ def python_config_to_docker(
|
||||
_update_graph_paths(config_path, config, local_deps)
|
||||
# Rewrite auth path, so it points to the correct location in the Docker container
|
||||
_update_auth_path(config_path, config, local_deps)
|
||||
# Rewrite encryption path, so it points to the correct location in the Docker container
|
||||
_update_encryption_path(config_path, config, local_deps)
|
||||
# Rewrite HTTP app path, so it points to the correct location in the Docker container
|
||||
_update_http_app_path(config_path, config, local_deps)
|
||||
|
||||
@@ -899,6 +953,9 @@ ADD {relpath} /deps/{name}
|
||||
if (auth_config := config.get("auth")) is not None:
|
||||
env_vars.append(f"ENV LANGGRAPH_AUTH='{json.dumps(auth_config)}'")
|
||||
|
||||
if (encryption_config := config.get("encryption")) is not None:
|
||||
env_vars.append(f"ENV LANGGRAPH_ENCRYPTION='{json.dumps(encryption_config)}'")
|
||||
|
||||
if (http_config := config.get("http")) is not None:
|
||||
env_vars.append(f"ENV LANGGRAPH_HTTP='{json.dumps(http_config)}'")
|
||||
|
||||
@@ -1022,6 +1079,9 @@ def node_config_to_docker(
|
||||
if (auth_config := config.get("auth")) is not None:
|
||||
env_vars.append(f"ENV LANGGRAPH_AUTH='{json.dumps(auth_config)}'")
|
||||
|
||||
if (encryption_config := config.get("encryption")) is not None:
|
||||
env_vars.append(f"ENV LANGGRAPH_ENCRYPTION='{json.dumps(encryption_config)}'")
|
||||
|
||||
if (http_config := config.get("http")) is not None:
|
||||
env_vars.append(f"ENV LANGGRAPH_HTTP='{json.dumps(http_config)}'")
|
||||
|
||||
|
||||
@@ -302,6 +302,27 @@ class AuthConfig(TypedDict, total=False):
|
||||
"""
|
||||
|
||||
|
||||
class EncryptionConfig(TypedDict, total=False):
|
||||
"""Configuration for custom at-rest encryption logic.
|
||||
|
||||
Allows you to implement custom encryption for sensitive data stored in the database,
|
||||
including metadata fields and checkpoint blobs.
|
||||
"""
|
||||
|
||||
path: str
|
||||
"""Required. Path to an instance of the Encryption() class that implements custom encryption handlers.
|
||||
|
||||
Format: "path/to/file.py:my_encryption"
|
||||
|
||||
Example:
|
||||
{
|
||||
"encryption": {
|
||||
"path": "./encryption.py:my_encryption"
|
||||
}
|
||||
}
|
||||
"""
|
||||
|
||||
|
||||
class CorsConfig(TypedDict, total=False):
|
||||
"""Specifies Cross-Origin Resource Sharing (CORS) rules for your server.
|
||||
|
||||
@@ -577,10 +598,16 @@ class Config(TypedDict, total=False):
|
||||
"""
|
||||
|
||||
auth: AuthConfig | None
|
||||
"""Optional. Custom authentication config, including the path to your Python auth logic and
|
||||
"""Optional. Custom authentication config, including the path to your Python auth logic and
|
||||
the OpenAPI security definitions it uses.
|
||||
"""
|
||||
|
||||
encryption: EncryptionConfig | None
|
||||
"""Optional. Custom at-rest encryption config, including the path to your Python encryption logic.
|
||||
|
||||
Allows you to implement custom encryption for sensitive data stored in the database.
|
||||
"""
|
||||
|
||||
http: HttpConfig | None
|
||||
"""Optional. Configuration for the built-in HTTP server, controlling which custom routes are exposed
|
||||
and how cross-origin requests are handled.
|
||||
@@ -603,6 +630,7 @@ __all__ = [
|
||||
"StoreConfig",
|
||||
"CheckpointerConfig",
|
||||
"AuthConfig",
|
||||
"EncryptionConfig",
|
||||
"HttpConfig",
|
||||
"MiddlewareOrders",
|
||||
"Distros",
|
||||
|
||||
@@ -99,6 +99,17 @@
|
||||
},
|
||||
"description": "Optional. Additional Docker instructions that will be appended to your base Dockerfile.\n\nUseful for installing OS packages, setting environment variables, etc."
|
||||
},
|
||||
"encryption": {
|
||||
"anyOf": [
|
||||
{
|
||||
"$ref": "#/$defs/EncryptionConfig"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Optional. Custom at-rest encryption config, including the path to your Python encryption logic.\n\nAllows you to implement custom encryption for sensitive data stored in the database.\n"
|
||||
},
|
||||
"env": {
|
||||
"anyOf": [
|
||||
{
|
||||
@@ -292,6 +303,17 @@
|
||||
},
|
||||
"description": "Optional. Additional Docker instructions that will be appended to your base Dockerfile.\n\nUseful for installing OS packages, setting environment variables, etc."
|
||||
},
|
||||
"encryption": {
|
||||
"anyOf": [
|
||||
{
|
||||
"$ref": "#/$defs/EncryptionConfig"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Optional. Custom at-rest encryption config, including the path to your Python encryption logic.\n\nAllows you to implement custom encryption for sensitive data stored in the database.\n"
|
||||
},
|
||||
"env": {
|
||||
"anyOf": [
|
||||
{
|
||||
@@ -592,6 +614,17 @@
|
||||
},
|
||||
"required": []
|
||||
},
|
||||
"EncryptionConfig": {
|
||||
"title": "EncryptionConfig",
|
||||
"description": "Configuration for custom at-rest encryption logic.\n\n Allows you to implement custom encryption for sensitive data stored in the database,\n including metadata fields and checkpoint blobs.",
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": []
|
||||
},
|
||||
"HttpConfig": {
|
||||
"title": "HttpConfig",
|
||||
"description": "Configuration for the built-in HTTP server that powers your deployment's routes and endpoints.",
|
||||
|
||||
@@ -99,6 +99,17 @@
|
||||
},
|
||||
"description": "Optional. Additional Docker instructions that will be appended to your base Dockerfile.\n\nUseful for installing OS packages, setting environment variables, etc."
|
||||
},
|
||||
"encryption": {
|
||||
"anyOf": [
|
||||
{
|
||||
"$ref": "#/$defs/EncryptionConfig"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Optional. Custom at-rest encryption config, including the path to your Python encryption logic.\n\nAllows you to implement custom encryption for sensitive data stored in the database.\n"
|
||||
},
|
||||
"env": {
|
||||
"anyOf": [
|
||||
{
|
||||
@@ -292,6 +303,17 @@
|
||||
},
|
||||
"description": "Optional. Additional Docker instructions that will be appended to your base Dockerfile.\n\nUseful for installing OS packages, setting environment variables, etc."
|
||||
},
|
||||
"encryption": {
|
||||
"anyOf": [
|
||||
{
|
||||
"$ref": "#/$defs/EncryptionConfig"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Optional. Custom at-rest encryption config, including the path to your Python encryption logic.\n\nAllows you to implement custom encryption for sensitive data stored in the database.\n"
|
||||
},
|
||||
"env": {
|
||||
"anyOf": [
|
||||
{
|
||||
@@ -592,6 +614,17 @@
|
||||
},
|
||||
"required": []
|
||||
},
|
||||
"EncryptionConfig": {
|
||||
"title": "EncryptionConfig",
|
||||
"description": "Configuration for custom at-rest encryption logic.\n\n Allows you to implement custom encryption for sensitive data stored in the database,\n including metadata fields and checkpoint blobs.",
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": []
|
||||
},
|
||||
"HttpConfig": {
|
||||
"title": "HttpConfig",
|
||||
"description": "Configuration for the built-in HTTP server that powers your deployment's routes and endpoints.",
|
||||
|
||||
@@ -48,6 +48,7 @@ def test_validate_config():
|
||||
"env": {},
|
||||
"store": None,
|
||||
"auth": None,
|
||||
"encryption": None,
|
||||
"checkpointer": None,
|
||||
"http": None,
|
||||
"ui": None,
|
||||
@@ -74,6 +75,7 @@ def test_validate_config():
|
||||
"env": env,
|
||||
"store": None,
|
||||
"auth": None,
|
||||
"encryption": None,
|
||||
"checkpointer": None,
|
||||
"http": None,
|
||||
"ui": None,
|
||||
@@ -748,6 +750,57 @@ RUN (test ! -f /api/langgraph_api/js/build.mts && echo "Prebuild script not foun
|
||||
assert additional_contexts == {}
|
||||
|
||||
|
||||
def test_config_to_docker_python_encryption():
|
||||
# Test that encryption config is included in validation
|
||||
graphs = {"agent": "./agent.py:graph"}
|
||||
validated = validate_config(
|
||||
{
|
||||
"python_version": "3.11",
|
||||
"graphs": graphs,
|
||||
"dependencies": ["."],
|
||||
"encryption": {"path": "./encryption.py:encryption"},
|
||||
}
|
||||
)
|
||||
|
||||
# Verify that encryption config is preserved after validation
|
||||
assert validated.get("encryption") is not None
|
||||
assert validated["encryption"]["path"] == "./encryption.py:encryption"
|
||||
|
||||
|
||||
def test_config_to_docker_python_encryption_bad_path():
|
||||
# Test that invalid encryption path format raises ValueError
|
||||
graphs = {"agent": "./agent.py:graph"}
|
||||
with pytest.raises(ValueError, match="Invalid encryption.path format"):
|
||||
validate_config(
|
||||
{
|
||||
"python_version": "3.11",
|
||||
"graphs": graphs,
|
||||
"dependencies": ["."],
|
||||
"encryption": {"path": "./encryption.py"}, # Missing :attribute
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_config_to_docker_python_encryption_formatted():
|
||||
# Test that encryption config is properly formatted in Docker output
|
||||
graphs = {"agent": "./graphs/agent.py:graph"}
|
||||
actual_docker_stdin, additional_contexts = config_to_docker(
|
||||
PATH_TO_CONFIG,
|
||||
validate_config(
|
||||
{
|
||||
"python_version": "3.11",
|
||||
"dependencies": ["."],
|
||||
"graphs": graphs,
|
||||
"encryption": {"path": "./agent.py:my_encryption"},
|
||||
}
|
||||
),
|
||||
"langchain/langgraph-api",
|
||||
)
|
||||
# Verify that LANGGRAPH_ENCRYPTION is in the docker output with the correct path
|
||||
assert "LANGGRAPH_ENCRYPTION=" in actual_docker_stdin
|
||||
assert "/deps/outer-unit_tests/unit_tests/agent.py:my_encryption" in actual_docker_stdin
|
||||
|
||||
|
||||
def test_config_to_docker_nodejs_internal_docker_tag():
|
||||
graphs = {"agent": "./graphs/agent.js:graph"}
|
||||
actual_docker_stdin, additional_contexts = config_to_docker(
|
||||
|
||||
@@ -67,7 +67,7 @@ While LangGraph can be used standalone, it also integrates seamlessly with any L
|
||||
- [LangChain](https://docs.langchain.com/oss/python/langchain/overview) – Provides integrations and composable components to streamline LLM application development.
|
||||
|
||||
> [!NOTE]
|
||||
> Looking for the JS version of LangGraph? See the [JS repo](https://github.com/langchain-ai/langgraphjs) and the [JS docs](https://langchain-ai.github.io/langgraphjs/).
|
||||
> Looking for the JS version of LangGraph? See the [JS repo](https://github.com/langchain-ai/langgraphjs) and the [JS docs](https://docs.langchain.com/oss/javascript/langgraph/overview).
|
||||
|
||||
## Additional resources
|
||||
|
||||
|
||||
Generated
+2
@@ -1801,12 +1801,14 @@ dev = [
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watch" },
|
||||
{ name = "ruff", specifier = "==0.14.7" },
|
||||
{ name = "starlette" },
|
||||
{ name = "ty", specifier = "==0.0.1a27" },
|
||||
]
|
||||
lint = [
|
||||
{ name = "codespell" },
|
||||
{ name = "mypy", specifier = "==1.19.0" },
|
||||
{ name = "ruff", specifier = "==0.14.7" },
|
||||
{ name = "starlette" },
|
||||
{ name = "ty", specifier = "==0.0.1a27" },
|
||||
]
|
||||
test = [
|
||||
|
||||
Generated
+2
@@ -601,12 +601,14 @@ dev = [
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watch" },
|
||||
{ name = "ruff", specifier = "==0.14.7" },
|
||||
{ name = "starlette" },
|
||||
{ name = "ty", specifier = "==0.0.1a27" },
|
||||
]
|
||||
lint = [
|
||||
{ name = "codespell" },
|
||||
{ name = "mypy", specifier = "==1.19.0" },
|
||||
{ name = "ruff", specifier = "==0.14.7" },
|
||||
{ name = "starlette" },
|
||||
{ name = "ty", specifier = "==0.0.1a27" },
|
||||
]
|
||||
test = [
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from langgraph_sdk.auth import Auth
|
||||
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.2.14"
|
||||
|
||||
__all__ = ["Auth", "get_client", "get_sync_client"]
|
||||
__all__ = ["Auth", "Encryption", "EncryptionContext", "get_client", "get_sync_client"]
|
||||
|
||||
@@ -0,0 +1,554 @@
|
||||
"""Custom encryption support for LangGraph.
|
||||
|
||||
.. warning::
|
||||
This API is in beta and may change in future versions.
|
||||
|
||||
This module provides a framework for implementing custom at-rest encryption
|
||||
in LangGraph applications. Similar to the Auth system, it allows developers
|
||||
to define custom encryption and decryption handlers that are executed
|
||||
server-side.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import inspect
|
||||
import typing
|
||||
import warnings
|
||||
|
||||
from langgraph_sdk.encryption import types
|
||||
|
||||
|
||||
class LangGraphBetaWarning(UserWarning):
|
||||
"""Warning for beta features in LangGraph SDK."""
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _warn_encryption_beta() -> None:
|
||||
warnings.warn(
|
||||
"The Encryption API is in beta and may change in future versions.",
|
||||
LangGraphBetaWarning,
|
||||
stacklevel=4,
|
||||
)
|
||||
|
||||
|
||||
class DuplicateHandlerError(Exception):
|
||||
"""Raised when attempting to register a duplicate encryption/decryption handler."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
def _validate_handler(fn: typing.Callable, handler_type: str) -> None:
|
||||
"""Validate that a handler function has the correct signature.
|
||||
|
||||
Args:
|
||||
fn: The handler function to validate
|
||||
handler_type: Description of the handler for error messages
|
||||
|
||||
Raises:
|
||||
TypeError: If the handler is not an async function or has wrong parameter count
|
||||
"""
|
||||
if not inspect.iscoroutinefunction(fn):
|
||||
raise TypeError(f"{handler_type} must be an async function, got {type(fn)}")
|
||||
|
||||
sig = inspect.signature(fn)
|
||||
params = [
|
||||
p
|
||||
for p in sig.parameters.values()
|
||||
if p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD)
|
||||
]
|
||||
if len(params) != 2:
|
||||
raise TypeError(
|
||||
f"{handler_type} must accept exactly 2 parameters "
|
||||
f"(ctx, data), got {len(params)}"
|
||||
)
|
||||
|
||||
|
||||
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.
|
||||
|
||||
Provides @encryption.encrypt.blob and @encryption.encrypt.json decorators for
|
||||
registering encryption functions.
|
||||
"""
|
||||
|
||||
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.
|
||||
|
||||
The handler will be called to encrypt opaque data like checkpoint blobs.
|
||||
|
||||
Example:
|
||||
```python
|
||||
@encryption.encrypt.blob
|
||||
async def encrypt_blob(ctx: EncryptionContext, blob: bytes) -> bytes:
|
||||
# Encrypt the blob using your encryption service
|
||||
return encrypted_blob
|
||||
```
|
||||
|
||||
Args:
|
||||
fn: The encryption handler function
|
||||
|
||||
Returns:
|
||||
The registered handler function
|
||||
|
||||
Raises:
|
||||
DuplicateHandlerError: If blob encryptor already registered
|
||||
TypeError: If handler has invalid signature
|
||||
"""
|
||||
if self._parent._blob_encryptor is not None:
|
||||
raise DuplicateHandlerError("Blob encryptor already registered")
|
||||
_validate_handler(fn, "Blob encryptor")
|
||||
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
|
||||
|
||||
Example:
|
||||
```python
|
||||
@encryption.encrypt.json
|
||||
async def default_encrypt(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Default encryption for all models
|
||||
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)
|
||||
```
|
||||
"""
|
||||
return self._json
|
||||
|
||||
|
||||
class _DecryptDecorators:
|
||||
"""Decorators for decryption handlers.
|
||||
|
||||
Provides @encryption.decrypt.blob and @encryption.decrypt.json decorators for
|
||||
registering decryption functions.
|
||||
"""
|
||||
|
||||
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.
|
||||
|
||||
The handler will be called to decrypt opaque data like checkpoint blobs.
|
||||
|
||||
Example:
|
||||
```python
|
||||
@encryption.decrypt.blob
|
||||
async def decrypt_blob(ctx: EncryptionContext, blob: bytes) -> bytes:
|
||||
# Decrypt the blob using your encryption service
|
||||
return decrypted_blob
|
||||
```
|
||||
|
||||
Args:
|
||||
fn: The decryption handler function
|
||||
|
||||
Returns:
|
||||
The registered handler function
|
||||
|
||||
Raises:
|
||||
DuplicateHandlerError: If blob decryptor already registered
|
||||
TypeError: If handler has invalid signature
|
||||
"""
|
||||
if self._parent._blob_decryptor is not None:
|
||||
raise DuplicateHandlerError("Blob decryptor already registered")
|
||||
_validate_handler(fn, "Blob decryptor")
|
||||
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
|
||||
|
||||
Example:
|
||||
```python
|
||||
@encryption.decrypt.json
|
||||
async def default_decrypt(ctx: EncryptionContext, data: dict) -> dict:
|
||||
# Default decryption for all models
|
||||
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)
|
||||
```
|
||||
"""
|
||||
return self._json
|
||||
|
||||
|
||||
class Encryption:
|
||||
"""Add custom at-rest encryption to your LangGraph application.
|
||||
|
||||
.. warning::
|
||||
This API is in beta and may change in future versions.
|
||||
|
||||
The Encryption class provides a system for implementing custom encryption
|
||||
of data at rest in LangGraph applications. It supports encryption of
|
||||
both opaque blobs (like checkpoints) and structured JSON data (like
|
||||
metadata, context, kwargs, values, etc.).
|
||||
|
||||
To use, create a separate Python file and add the path to the file to your
|
||||
LangGraph API configuration file (`langgraph.json`). Within that file, create
|
||||
an instance of the Encryption class and register encryption and decryption
|
||||
handlers as needed.
|
||||
|
||||
Example `langgraph.json` file:
|
||||
|
||||
```json
|
||||
{
|
||||
"dependencies": ["."],
|
||||
"graphs": {
|
||||
"agent": "./my_agent/agent.py:graph"
|
||||
},
|
||||
"env": ".env",
|
||||
"encryption": {
|
||||
"path": "./encryption.py:my_encryption"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Then the LangGraph server will load your encryption file and use it to
|
||||
encrypt/decrypt data at rest.
|
||||
|
||||
???+ example "Basic Usage"
|
||||
|
||||
```python
|
||||
from langgraph_sdk import Encryption, EncryptionContext
|
||||
|
||||
my_encryption = Encryption()
|
||||
|
||||
@my_encryption.encrypt.blob
|
||||
async def encrypt_blob(ctx: EncryptionContext, blob: bytes) -> bytes:
|
||||
# Call your encryption service
|
||||
return encrypted_blob
|
||||
|
||||
@my_encryption.decrypt.blob
|
||||
async def decrypt_blob(ctx: EncryptionContext, blob: bytes) -> bytes:
|
||||
# Call your decryption service
|
||||
return decrypted_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)
|
||||
else:
|
||||
# Pass through (including "owner" for search)
|
||||
encrypted[key] = value
|
||||
return encrypted
|
||||
|
||||
@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)
|
||||
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)
|
||||
```
|
||||
|
||||
???+ example "Field-Specific Logic"
|
||||
|
||||
The `ctx.field` attribute tells you which specific field is being encrypted,
|
||||
allowing different logic within the same model:
|
||||
|
||||
```python
|
||||
@my_encryption.encrypt.json.thread
|
||||
async def encrypt_thread(ctx: EncryptionContext, data: dict) -> dict:
|
||||
if ctx.field == "metadata":
|
||||
# Thread metadata - standard encryption
|
||||
return encrypt_standard(data)
|
||||
elif ctx.field == "values":
|
||||
# Thread values - more sensitive, use stronger encryption
|
||||
return encrypt_sensitive(data)
|
||||
else:
|
||||
return encrypt_standard(data)
|
||||
```
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"_blob_decryptor",
|
||||
"_blob_encryptor",
|
||||
"_context_handler",
|
||||
"_json_decryptor",
|
||||
"_json_decryptors",
|
||||
"_json_encryptor",
|
||||
"_json_encryptors",
|
||||
"decrypt",
|
||||
"encrypt",
|
||||
)
|
||||
|
||||
types = types
|
||||
"""Reference to encryption type definitions.
|
||||
|
||||
Provides access to all type definitions used in the encryption system,
|
||||
including EncryptionContext, BlobEncryptor, BlobDecryptor,
|
||||
JsonEncryptor, and JsonDecryptor.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the Encryption instance."""
|
||||
_warn_encryption_beta()
|
||||
self.encrypt = _EncryptDecorators(self)
|
||||
self.decrypt = _DecryptDecorators(self)
|
||||
self._blob_encryptor: types.BlobEncryptor | None = None
|
||||
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:
|
||||
"""Register a context handler to derive encryption context from auth.
|
||||
|
||||
The handler receives the authenticated user and current EncryptionContext,
|
||||
and returns a dict that becomes ctx.metadata for encrypt/decrypt handlers.
|
||||
|
||||
This allows encryption context to be derived from JWT claims or other
|
||||
auth-derived data instead of requiring a separate X-Encryption-Context header.
|
||||
|
||||
Note: The context handler is called once per request in middleware,
|
||||
so ctx.model and ctx.field will be None in the handler.
|
||||
|
||||
Example:
|
||||
```python
|
||||
from langgraph_sdk import Encryption, EncryptionContext
|
||||
from starlette.authentication import BaseUser
|
||||
|
||||
encryption = Encryption()
|
||||
|
||||
@encryption.context
|
||||
async def get_context(user: BaseUser, ctx: EncryptionContext) -> dict:
|
||||
# Derive encryption context from authenticated user
|
||||
return {
|
||||
**ctx.metadata, # preserve X-Encryption-Context header if present
|
||||
"tenant_id": user.tenant_id,
|
||||
}
|
||||
```
|
||||
|
||||
Args:
|
||||
fn: The context handler function
|
||||
|
||||
Returns:
|
||||
The registered handler function
|
||||
"""
|
||||
self._context_handler = fn
|
||||
return fn
|
||||
|
||||
def get_json_encryptor(
|
||||
self, model: str | None = None
|
||||
) -> types.JsonEncryptor | None:
|
||||
"""Get the JSON encryptor for a specific model.
|
||||
|
||||
Args:
|
||||
model: The model type (e.g., "thread", "assistant"). If None, returns default.
|
||||
|
||||
Returns:
|
||||
Model-specific encryptor if registered, otherwise default encryptor, or None.
|
||||
"""
|
||||
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
|
||||
) -> types.JsonDecryptor | None:
|
||||
"""Get the JSON decryptor for a specific model.
|
||||
|
||||
Args:
|
||||
model: The model type (e.g., "thread", "assistant"). If None, returns default.
|
||||
|
||||
Returns:
|
||||
Model-specific decryptor if registered, otherwise default decryptor, or None.
|
||||
"""
|
||||
if model and model in self._json_decryptors:
|
||||
return self._json_decryptors[model]
|
||||
return self._json_decryptor
|
||||
|
||||
def __repr__(self) -> str:
|
||||
handlers = []
|
||||
if self._blob_encryptor:
|
||||
handlers.append("blob_encryptor")
|
||||
if self._blob_decryptor:
|
||||
handlers.append("blob_decryptor")
|
||||
if self._json_encryptor:
|
||||
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)}])"
|
||||
@@ -0,0 +1,147 @@
|
||||
"""Encryption and decryption types for LangGraph.
|
||||
|
||||
This module defines the core types used for custom at-rest encryption
|
||||
in LangGraph. It includes context types and typed dictionaries for
|
||||
encryption operations.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import typing
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
Json = dict[str, typing.Any]
|
||||
"""JSON-serializable dictionary type for structured data encryption."""
|
||||
|
||||
|
||||
class EncryptionContext:
|
||||
"""Context passed to encryption/decryption handlers.
|
||||
|
||||
Contains arbitrary non-secret key-values that will be stored on encrypt.
|
||||
These key-values are intended to be sent to an external service that
|
||||
manages keys and handles the actual encryption and decryption of data.
|
||||
|
||||
Attributes:
|
||||
model: The model type being encrypted (e.g., "assistant", "thread", "run", "checkpoint")
|
||||
field: The specific field being encrypted (e.g., "metadata", "context", "kwargs", "values")
|
||||
metadata: Additional context metadata that can be used for encryption decisions
|
||||
"""
|
||||
|
||||
__slots__ = ("field", "metadata", "model")
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: str | None = None,
|
||||
metadata: dict[str, typing.Any] | None = None,
|
||||
field: str | None = None,
|
||||
):
|
||||
self.model = model
|
||||
self.field = field
|
||||
self.metadata = metadata or {}
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"EncryptionContext(model={self.model!r}, field={self.field!r}, metadata={self.metadata!r})"
|
||||
|
||||
|
||||
BlobEncryptor = Callable[[EncryptionContext, bytes], Awaitable[bytes]]
|
||||
"""Handler for encrypting opaque blob data like checkpoints.
|
||||
|
||||
Note: Must be an async function. Encryption typically involves I/O operations
|
||||
(calling external KMS services), which should be async.
|
||||
|
||||
Args:
|
||||
ctx: Encryption context with model type and metadata
|
||||
blob: The raw bytes to encrypt
|
||||
|
||||
Returns:
|
||||
Awaitable that resolves to encrypted bytes
|
||||
"""
|
||||
|
||||
BlobDecryptor = Callable[[EncryptionContext, bytes], Awaitable[bytes]]
|
||||
"""Handler for decrypting opaque blob data like checkpoints.
|
||||
|
||||
Note: Must be an async function. Decryption typically involves I/O operations
|
||||
(calling external KMS services), which should be async.
|
||||
|
||||
Args:
|
||||
ctx: Encryption context with model type and metadata
|
||||
blob: The encrypted bytes to decrypt
|
||||
|
||||
Returns:
|
||||
Awaitable that resolves to decrypted bytes
|
||||
"""
|
||||
|
||||
JsonEncryptor = Callable[[EncryptionContext, Json], Awaitable[Json]]
|
||||
"""Handler for encrypting structured JSON data.
|
||||
|
||||
Note: Must be an async function. Encryption typically involves I/O operations
|
||||
(calling external KMS services), which should be async.
|
||||
|
||||
Used for encrypting structured data like metadata, context, kwargs, values,
|
||||
and other JSON-serializable fields across different model types.
|
||||
|
||||
Maps plaintext fields to encrypted fields. A practical approach:
|
||||
- Keep "owner" field unencrypted for search/filtering
|
||||
- Encrypt VALUES (not keys) for fields with specific prefix (e.g., "my.customer.org/")
|
||||
- Pass through all other fields unencrypted
|
||||
|
||||
Example:
|
||||
Input: {"owner": "user123", "my.customer.org/email": "john@example.com", "tenant_id": "t-456"}
|
||||
Output: {"owner": "user123", "my.customer.org/email": "ENCRYPTED", "tenant_id": "t-456"}
|
||||
|
||||
Note: Encrypted field VALUES cannot be reliably searched, as most real-world
|
||||
encryption implementations use nonces (non-deterministic encryption).
|
||||
Only unencrypted fields can be used in search queries.
|
||||
|
||||
Args:
|
||||
ctx: Encryption context with model type, field name, and metadata
|
||||
data: The plaintext JSON dictionary
|
||||
|
||||
Returns:
|
||||
Awaitable that resolves to encrypted JSON dictionary
|
||||
"""
|
||||
|
||||
JsonDecryptor = Callable[[EncryptionContext, Json], Awaitable[Json]]
|
||||
"""Handler for decrypting structured JSON data.
|
||||
|
||||
Note: Must be an async function. Decryption typically involves I/O operations
|
||||
(calling external KMS services), which should be async.
|
||||
|
||||
Inverse of JsonEncryptor. Must be able to decrypt data that
|
||||
was encrypted by the corresponding encryptor.
|
||||
|
||||
Args:
|
||||
ctx: Encryption context with model type, field name, and metadata
|
||||
data: The encrypted JSON dictionary
|
||||
|
||||
Returns:
|
||||
Awaitable that resolves to decrypted JSON dictionary
|
||||
"""
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from starlette.authentication import BaseUser
|
||||
|
||||
ContextHandler = Callable[
|
||||
["BaseUser", EncryptionContext], Awaitable[dict[str, typing.Any]]
|
||||
]
|
||||
"""Handler for deriving encryption context from authenticated user info.
|
||||
|
||||
Note: Must be an async function as it may involve I/O operations.
|
||||
|
||||
The context handler is called once per request in middleware (after auth),
|
||||
allowing encryption context to be derived from JWT claims, user properties,
|
||||
or other auth-derived data instead of requiring a separate X-Encryption-Context header.
|
||||
|
||||
The return value becomes ctx.metadata for subsequent encrypt/decrypt operations
|
||||
and is persisted with encrypted data for later decryption.
|
||||
|
||||
Note: ctx.model and ctx.field will be None in context handlers since
|
||||
the handler runs once per request before any specific model/field is known.
|
||||
|
||||
Args:
|
||||
user: The authenticated user (from Starlette's AuthenticationMiddleware)
|
||||
ctx: Current encryption context with metadata from X-Encryption-Context header
|
||||
|
||||
Returns:
|
||||
Awaitable that resolves to dict that becomes the new ctx.metadata
|
||||
"""
|
||||
@@ -34,6 +34,7 @@ lint = [
|
||||
"codespell",
|
||||
"mypy==1.19.0",
|
||||
"ty==0.0.1a27",
|
||||
"starlette",
|
||||
]
|
||||
dev = [
|
||||
{ include-group = "test" },
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
import pytest
|
||||
|
||||
from langgraph_sdk.encryption import DuplicateHandlerError, Encryption
|
||||
|
||||
|
||||
class TestHandlerValidation:
|
||||
"""Test duplicate handler and signature validation."""
|
||||
|
||||
def test_duplicate_handlers_raise_error(self):
|
||||
"""Registering the same handler type twice raises DuplicateHandlerError."""
|
||||
encryption = Encryption()
|
||||
|
||||
@encryption.encrypt.blob
|
||||
async def blob_enc(_ctx, data):
|
||||
return data
|
||||
|
||||
@encryption.decrypt.blob
|
||||
async def blob_dec(_ctx, data):
|
||||
return data
|
||||
|
||||
@encryption.encrypt.json
|
||||
async def json_enc(_ctx, data):
|
||||
return data
|
||||
|
||||
@encryption.decrypt.json
|
||||
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):
|
||||
|
||||
@encryption.encrypt.blob
|
||||
async def dup(_ctx, data):
|
||||
return data
|
||||
|
||||
with pytest.raises(DuplicateHandlerError):
|
||||
|
||||
@encryption.decrypt.blob
|
||||
async def dup(_ctx, data):
|
||||
return data
|
||||
|
||||
with pytest.raises(DuplicateHandlerError):
|
||||
|
||||
@encryption.encrypt.json
|
||||
async def dup(_ctx, data):
|
||||
return data
|
||||
|
||||
with pytest.raises(DuplicateHandlerError):
|
||||
|
||||
@encryption.decrypt.json
|
||||
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()
|
||||
|
||||
with pytest.raises(TypeError, match="must be an async function"):
|
||||
|
||||
@encryption.encrypt.blob
|
||||
def sync_handler(_ctx, data):
|
||||
return data
|
||||
|
||||
def test_handlers_must_have_two_params(self):
|
||||
"""Wrong parameter count raises TypeError."""
|
||||
encryption = Encryption()
|
||||
|
||||
with pytest.raises(TypeError, match="must accept exactly 2 parameters"):
|
||||
|
||||
@encryption.encrypt.blob # type: ignore[arg-type]
|
||||
async def wrong_params(ctx):
|
||||
return ctx
|
||||
Generated
+26
@@ -152,12 +152,14 @@ dev = [
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watch" },
|
||||
{ name = "ruff" },
|
||||
{ name = "starlette" },
|
||||
{ name = "ty" },
|
||||
]
|
||||
lint = [
|
||||
{ name = "codespell" },
|
||||
{ name = "mypy" },
|
||||
{ name = "ruff" },
|
||||
{ name = "starlette" },
|
||||
{ name = "ty" },
|
||||
]
|
||||
test = [
|
||||
@@ -183,12 +185,14 @@ dev = [
|
||||
{ name = "pytest-mock" },
|
||||
{ name = "pytest-watch" },
|
||||
{ name = "ruff", specifier = "==0.14.7" },
|
||||
{ name = "starlette" },
|
||||
{ name = "ty", specifier = "==0.0.1a27" },
|
||||
]
|
||||
lint = [
|
||||
{ name = "codespell" },
|
||||
{ name = "mypy", specifier = "==1.19.0" },
|
||||
{ name = "ruff", specifier = "==0.14.7" },
|
||||
{ name = "starlette" },
|
||||
{ name = "ty", specifier = "==0.0.1a27" },
|
||||
]
|
||||
test = [
|
||||
@@ -658,6 +662,28 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/1d/d2/1637f4360ada6a368d3265bf39f2cf737a0aaab15ab520fc005903e883f8/ruff-0.14.7-py3-none-win_arm64.whl", hash = "sha256:be4d653d3bea1b19742fcc6502354e32f65cd61ff2fbdb365803ef2c2aec6228", size = 13609215, upload-time = "2025-11-28T20:55:15.375Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sniffio"
|
||||
version = "1.3.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/a2/87/a6771e1546d97e7e041b6ae58d80074f81b7d5121207425c964ddf5cfdbd/sniffio-1.3.1.tar.gz", hash = "sha256:f4324edc670a0f49750a81b895f35c3adb843cca46f0530f79fc1babb23789dc", size = 20372, upload-time = "2024-02-25T23:20:04.057Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/e9/44/75a9c9421471a6c4805dbf2356f7c181a29c1879239abab1ea2cc8f38b40/sniffio-1.3.1-py3-none-any.whl", hash = "sha256:2f6da418d1f1e0fddd844478f41680e794e6051915791a034ff65e5f100525a2", size = 10235, upload-time = "2024-02-25T23:20:01.196Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "starlette"
|
||||
version = "0.50.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "anyio" },
|
||||
{ name = "typing-extensions", marker = "python_full_version < '3.13'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/ba/b8/73a0e6a6e079a9d9cfa64113d771e421640b6f679a52eeb9b32f72d871a1/starlette-0.50.0.tar.gz", hash = "sha256:a2a17b22203254bcbc2e1f926d2d55f3f9497f769416b3190768befe598fa3ca", size = 2646985, upload-time = "2025-11-01T15:25:27.516Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/d9/52/1064f510b141bd54025f9b55105e26d1fa970b9be67ad766380a3c9b74b0/starlette-0.50.0-py3-none-any.whl", hash = "sha256:9e5391843ec9b6e472eed1365a78c8098cfceb7a74bfd4d6b1c0c0095efb3bca", size = 74033, upload-time = "2025-11-01T15:25:25.461Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tomli"
|
||||
version = "2.3.0"
|
||||
|
||||
Reference in New Issue
Block a user