Compare commits

..
Author SHA1 Message Date
William Fu-Hinthorn fba8665718 merge 2024-11-27 14:54:03 -08:00
William Fu-Hinthorn 70b42812b5 Merge branch 'wfh/store/base/add_vector_earch' into wfh/store/add_vector_search 2024-11-27 14:46:44 -08:00
William FHandGitHub 471b2fad51 Merge branch 'main' into wfh/store/base/add_vector_earch 2024-11-27 14:41:37 -08:00
William FHandGitHub 98db2c08d7 Merge branch 'main' into wfh/store/base/add_vector_earch 2024-11-27 14:36:01 -08:00
William Fu-Hinthorn 1f5cab1869 Format 2024-11-27 14:34:31 -08:00
William Fu-Hinthorn 33dc6c84bb Merge branch 'wfh/store/base/add_vector_earch' into wfh/store/add_vector_search 2024-11-27 14:16:08 -08:00
William Fu-Hinthorn 4455e4185c Improve docs 2024-11-27 14:16:00 -08:00
William Fu-Hinthorn d3b2e2dc95 merge 2024-11-27 14:13:43 -08:00
William Fu-Hinthorn 2bff330bd7 Merge branch 'wfh/store/base/add_vector_earch' into wfh/store/add_vector_search 2024-11-27 14:08:05 -08:00
William Fu-Hinthorn c2a1f3a30a Improve docs 2024-11-27 14:06:54 -08:00
William Fu-Hinthorn 112a4f6c12 Test put-time fields 2024-11-27 13:19:04 -08:00
William Fu-Hinthorn 34a4ca3eaf Rename file 2024-11-27 11:52:11 -08:00
William Fu-Hinthorn d75492c330 Rename config 2024-11-27 11:20:03 -08:00
William Fu-Hinthorn 0ed26a3af9 handle no emb situation 2024-11-27 06:40:21 -08:00
William Fu-Hinthorn 74229c03dc merge 2024-11-26 21:30:41 -08:00
William Fu-Hinthorn 9a83dd7900 Add test 2024-11-26 19:23:48 -08:00
William Fu-Hinthorn b03235f92d Add in-mem vector search 2024-11-26 18:59:44 -08:00
William Fu-Hinthorn 6ec2afe958 Handle conflicts 2024-11-26 18:56:42 -08:00
William Fu-Hinthorn 11dc4d2691 merge 2024-11-26 18:46:21 -08:00
William Fu-Hinthorn 01b4d15aca Add in-mem vector search 2024-11-26 18:43:16 -08:00
William Fu-Hinthorn d81e350105 feat: Add vector search 2024-11-25 17:09:32 -08:00
34 changed files with 975 additions and 2192 deletions
+9 -1
View File
@@ -19,7 +19,7 @@
"\n",
"## Setup\n",
"\n",
"First, install the required packages and configure your environment:"
"First, install the required packages:"
]
},
{
@@ -33,6 +33,14 @@
"%pip install -U langgraph langsmith langchain_anthropic"
]
},
{
"cell_type": "markdown",
"id": "a6d1e870-1bc0-4d44-86c0-96681ccf6113",
"metadata": {},
"source": [
"In this tutorial, we'll be "
]
},
{
"cell_type": "code",
"execution_count": 2,
@@ -35,10 +35,10 @@ Create a new app from the `react-agent` template. This template is a simple agen
## Install Dependencies
In the root of your new LangGraph app, install the dependencies in `edit` mode so your local changes are used by the server:
In the root of your new LangGraph app, install the dependencies:
```shell
pip install -e .
pip install .
```
## Create a `.env` file
@@ -5,6 +5,7 @@ from typing import Any, Optional
from langchain_core.runnables import RunnableConfig
from psycopg import Capabilities, Connection, Cursor, Pipeline
from psycopg.errors import UndefinedTable
from psycopg.rows import DictRow, dict_row
from psycopg.types.json import Jsonb
from psycopg_pool import ConnectionPool
@@ -75,15 +76,16 @@ class PostgresSaver(BasePostgresSaver):
the first time checkpointer is used.
"""
with self._cursor() as cur:
cur.execute(self.MIGRATIONS[0])
results = cur.execute(
"SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1"
)
row = results.fetchone()
if row is None:
try:
row = cur.execute(
"SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1"
).fetchone()
if row is None:
version = -1
else:
version = row["v"]
except UndefinedTable:
version = -1
else:
version = row["v"]
for v, migration in zip(
range(version + 1, len(self.MIGRATIONS)),
self.MIGRATIONS[version + 1 :],
@@ -5,6 +5,7 @@ from typing import Any, Optional
from langchain_core.runnables import RunnableConfig
from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline, Capabilities
from psycopg.errors import UndefinedTable
from psycopg.rows import DictRow, dict_row
from psycopg.types.json import Jsonb
from psycopg_pool import AsyncConnectionPool
@@ -80,15 +81,17 @@ class AsyncPostgresSaver(BasePostgresSaver):
the first time checkpointer is used.
"""
async with self._cursor() as cur:
await cur.execute(self.MIGRATIONS[0])
results = await cur.execute(
"SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1"
)
row = await results.fetchone()
if row is None:
try:
results = await cur.execute(
"SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1"
)
row = await results.fetchone()
if row is None:
version = -1
else:
version = row["v"]
except UndefinedTable:
version = -1
else:
version = row["v"]
for v, migration in zip(
range(version + 1, len(self.MIGRATIONS)),
self.MIGRATIONS[version + 1 :],
@@ -21,7 +21,6 @@ from langgraph.store.base import (
)
from langgraph.store.base.batch import AsyncBatchedBaseStore
from langgraph.store.postgres.base import (
_PLACEHOLDER,
BasePostgresStore,
PoolConfig,
PostgresIndexConfig,
@@ -148,10 +147,11 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
already exist and runs database migrations. It MUST be called directly by the user
the first time the store is used.
"""
async def _get_version(cur: AsyncCursor[DictRow], table: str) -> int:
async with self._cursor() as cur:
try:
await cur.execute(f"SELECT v FROM {table} ORDER BY v DESC LIMIT 1")
await cur.execute(
"SELECT v FROM store_migrations ORDER BY v DESC LIMIT 1"
)
row = await cur.fetchone()
if row is None:
version = -1
@@ -160,25 +160,22 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
except UndefinedTable:
version = -1
await cur.execute(
f"""
CREATE TABLE IF NOT EXISTS {table} (
"""
CREATE TABLE IF NOT EXISTS store_migrations (
v INTEGER PRIMARY KEY
)
"""
)
return version
async with self._cursor() as cur:
version = await _get_version(cur, table="store_migrations")
for v, sql in enumerate(self.MIGRATIONS[version + 1 :], start=version + 1):
await cur.execute(sql)
await cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,))
for v, migration in enumerate(
self.MIGRATIONS[version + 1 :], start=version + 1
):
if isinstance(migration, str):
sql = migration
else:
if migration.condition and not migration.condition(self):
continue
if self.index_config:
version = await _get_version(cur, table="vector_migrations")
for v, migration in enumerate(
self.VECTOR_MIGRATIONS[version + 1 :], start=version + 1
):
sql = migration.sql
if migration.params:
params = {
@@ -186,10 +183,9 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
for k, v in migration.params.items()
}
sql = sql % params
await cur.execute(sql)
await cur.execute(
"INSERT INTO vector_migrations (v) VALUES (%s)", (v,)
)
await cur.execute(sql)
await cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,))
async def _execute_batch(
self,
@@ -293,10 +289,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
[query for _, query in embedding_requests]
)
for (idx, _), vector in zip(embedding_requests, vectors):
_paramslist = queries[idx][1]
for i in range(len(_paramslist)):
if _paramslist[i] is _PLACEHOLDER:
_paramslist[i] = vector
queries[idx][1][0] = vector
for (idx, _), (query, params) in zip(search_ops, queries):
await cur.execute(query, params)
@@ -55,11 +55,16 @@ class Migration(NamedTuple):
"""A database migration with optional conditions and parameters."""
sql: str
condition: Optional[Callable[[Any], bool]] = None
params: Optional[dict[str, Any]] = None
condition: Optional[Callable[["BasePostgresStore"], bool]] = None
MIGRATIONS: Sequence[str] = [
def _embedding_requested(store: Any) -> bool:
"""Check if vector operations are available in the database."""
return bool(store.index_config)
MIGRATIONS: Sequence[Union[str, Migration]] = [
"""
CREATE TABLE IF NOT EXISTS store (
-- 'prefix' represents the doc's 'namespace'
@@ -75,13 +80,11 @@ CREATE TABLE IF NOT EXISTS store (
-- For faster lookups by prefix
CREATE INDEX IF NOT EXISTS store_prefix_idx ON store USING btree (prefix text_pattern_ops);
""",
]
VECTOR_MIGRATIONS: Sequence[Migration] = [
Migration(
"""
CREATE EXTENSION IF NOT EXISTS vector;
""",
condition=_embedding_requested,
),
Migration(
"""
@@ -96,11 +99,12 @@ CREATE TABLE IF NOT EXISTS store_vectors (
FOREIGN KEY (prefix, key) REFERENCES store(prefix, key) ON DELETE CASCADE
);
""",
condition=_embedding_requested,
params={
"dims": lambda store: store.index_config["dims"],
"vector_type": lambda store: (
cast(PostgresIndexConfig, store.index_config)
.get("ann_index_config", {})
.get("db_index_config", {})
.get("vector_type", "vector")
),
},
@@ -110,9 +114,7 @@ CREATE TABLE IF NOT EXISTS store_vectors (
CREATE INDEX IF NOT EXISTS store_vectors_embedding_idx ON store_vectors
USING %(index_type)s (embedding %(ops)s)%(index_params)s;
""",
condition=lambda store: bool(
store.index_config and _get_index_params(store)[0] != "flat"
),
condition=_embedding_requested,
params={
"index_type": lambda store: _get_index_params(store)[0],
"ops": lambda store: _get_vector_type_ops(store),
@@ -156,10 +158,10 @@ class PoolConfig(TypedDict, total=False):
"""
class ANNIndexConfig(TypedDict, total=False):
class DBIndexConfig(TypedDict, total=False):
"""Configuration for vector index in PostgreSQL store."""
kind: Literal["hnsw", "ivfflat", "flat"]
kind: Literal["hnsw", "ivfflat"]
"""Type of index to use: 'hnsw' for Hierarchical Navigable Small World, or 'ivfflat' for Inverted File Flat."""
vector_type: Literal["vector", "halfvec"]
"""Type of vector storage to use.
@@ -169,7 +171,7 @@ class ANNIndexConfig(TypedDict, total=False):
"""
class HNSWConfig(ANNIndexConfig, total=False):
class HNSWConfig(DBIndexConfig, total=False):
"""Configuration for HNSW (Hierarchical Navigable Small World) index."""
kind: Literal["hnsw"] # type: ignore[misc]
@@ -179,7 +181,7 @@ class HNSWConfig(ANNIndexConfig, total=False):
"""Size of dynamic candidate list for index construction. Default is 64."""
class IVFFlatConfig(ANNIndexConfig, total=False):
class IVFFlatConfig(DBIndexConfig, total=False):
"""IVFFlat index divides vectors into lists, and then searches a subset of those lists that are closest to the query vector. It has faster build times and uses less memory than HNSW, but has lower query performance (in terms of speed-recall tradeoff).
Three keys to achieving good recall are:
@@ -204,7 +206,7 @@ class PostgresIndexConfig(IndexConfig, total=False):
Extends EmbeddingConfig with additional configuration for pgvector index and vector types.
"""
ann_index_config: ANNIndexConfig
db_index_config: Union[HNSWConfig, IVFFlatConfig]
"""Specific configuration for the chosen index type (HNSW or IVF Flat)."""
distance_type: Literal["l2", "inner_product", "cosine"]
"""Distance metric to use for vector similarity search:
@@ -216,11 +218,17 @@ class PostgresIndexConfig(IndexConfig, total=False):
class BasePostgresStore(Generic[C]):
MIGRATIONS = MIGRATIONS
VECTOR_MIGRATIONS = VECTOR_MIGRATIONS
conn: C
_deserializer: Optional[Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]]
index_config: Optional[PostgresIndexConfig]
@staticmethod
def _get_default_index_config() -> IndexConfig:
return HNSWConfig(
kind="hnsw",
vector_type="vector",
)
def _get_batch_GET_ops_queries(
self,
get_ops: Sequence[tuple[int, GetOp]],
@@ -290,12 +298,13 @@ class BasePostgresStore(Generic[C]):
[
_namespace_to_text(op.namespace),
op.key,
Jsonb(cast(dict, op.value)),
Jsonb(cast(dict, op.value).copy()),
]
)
# Then handle embeddings if configured
if self.index_config:
paths = self.index_config["__tokenized_fields"]
for op in inserts:
if op.index is False:
continue
@@ -303,11 +312,6 @@ class BasePostgresStore(Generic[C]):
ns = _namespace_to_text(op.namespace)
k = op.key
if op.index is None:
paths = self.index_config["__tokenized_fields"]
else:
paths = [(ix, tokenize_path(ix)) for ix in op.index]
for path, tokenized_path in paths:
texts = get_text_at_path(value, tokenized_path)
for i, text in enumerate(texts):
@@ -351,30 +355,22 @@ class BasePostgresStore(Generic[C]):
embedding_requests = []
for idx, (_, op) in enumerate(search_ops):
# Build filter conditions first
filter_params = []
filter_conditions = []
if op.filter:
for key, value in op.filter.items():
if isinstance(value, dict):
for op_name, val in value.items():
condition, filter_params_ = self._get_filter_condition(
key, op_name, val
)
filter_conditions.append(condition)
filter_params.extend(filter_params_)
else:
filter_conditions.append("value->%s = %s::jsonb")
filter_params.extend([key, json.dumps(value)])
base_query = """
SELECT prefix, key, value, created_at, updated_at
FROM store
WHERE prefix LIKE %s
"""
params: list = [f"{_namespace_to_text(op.namespace_prefix)}%"]
needs_vector_search = False
# Vector search branch
if op.query and self.index_config:
needs_vector_search = True
embedding_requests.append((idx, op.query))
score_operator, post_operator = _get_distance_operator(self)
score_expr = _get_distance_operator(self)
vector_type = (
cast(PostgresIndexConfig, self.index_config)
.get("ann_index_config", {})
.get("db_index_config", self._get_default_index_config())
.get("vector_type", "vector")
)
@@ -382,81 +378,61 @@ class BasePostgresStore(Generic[C]):
vector_type == "bit"
and self.index_config.get("distance_type") == "hamming"
):
score_operator = score_operator % (
"%s",
self.index_config["dims"],
)
score_expr = score_expr % ("%s", self.index_config["dims"])
else:
score_operator = score_operator % (
"%s",
vector_type,
)
score_expr = score_expr % ("%s", vector_type)
vectors_per_doc_estimate = self.index_config["__estimated_num_vectors"]
expanded_limit = (op.limit * vectors_per_doc_estimate * 2) + 1
# Vector search with CTE for proper score handling
filter_str = (
""
if not filter_conditions
else " AND " + " AND ".join(filter_conditions)
)
if op.namespace_prefix:
prefix_filter_str = f"WHERE s.prefix LIKE %s {filter_str} "
ns_args: Sequence = (f"{_namespace_to_text(op.namespace_prefix)}%",)
else:
ns_args = ()
if filter_str:
prefix_filter_str = f"WHERE {filter_str} "
else:
prefix_filter_str = ""
# Direct query with DISTINCT ON to get best score per document
base_query = f"""
WITH scored AS (
SELECT s.prefix, s.key, s.value, s.created_at, s.updated_at, {score_operator} AS neg_score
with scored as (
SELECT DISTINCT ON (s.prefix, s.key)
s.prefix, s.key, s.value, s.created_at, s.updated_at,
{score_expr} as score
FROM store s
JOIN store_vectors sv ON s.prefix = sv.prefix AND s.key = sv.key
{prefix_filter_str}
ORDER BY {score_operator} ASC
WHERE s.prefix LIKE %s
ORDER BY s.prefix, s.key, score DESC
LIMIT %s
)
SELECT * FROM (
SELECT DISTINCT ON (prefix, key)
prefix, key, value, created_at, updated_at, {post_operator} as score
FROM scored
ORDER BY prefix, key, score DESC
) AS unique_docs
ORDER BY score DESC
LIMIT %s
OFFSET %s
SELECT * FROM scored
"""
params = [
_PLACEHOLDER, # Vector placeholder
*ns_args,
*filter_params,
_PLACEHOLDER,
None, # Vector placeholder
f"{_namespace_to_text(op.namespace_prefix)}%",
expanded_limit,
op.limit,
op.offset,
]
# Regular search branch
else:
base_query = """
SELECT prefix, key, value, created_at, updated_at
FROM store
WHERE prefix LIKE %s
"""
params = [f"{_namespace_to_text(op.namespace_prefix)}%"]
if op.filter:
filter_conditions = []
for key, value in op.filter.items():
if isinstance(value, dict):
for op_name, val in value.items():
condition, filter_params = self._get_filter_condition(
key, op_name, val
)
filter_conditions.append(condition)
params.extend(filter_params)
else:
filter_conditions.append("value->%s = %s::jsonb")
params.extend([key, json.dumps(value)])
if filter_conditions:
params.extend(filter_params)
base_query += " AND " + " AND ".join(filter_conditions)
if needs_vector_search:
base_query += " WHERE " + " AND ".join(filter_conditions)
else:
base_query += " AND " + " AND ".join(filter_conditions)
if needs_vector_search:
base_query += " ORDER BY score DESC"
else:
base_query += " ORDER BY updated_at DESC"
base_query += " LIMIT %s OFFSET %s"
params.extend([op.limit, op.offset])
base_query += " LIMIT %s OFFSET %s"
params.extend([op.limit, op.offset])
queries.append((base_query, params))
return queries, embedding_requests
@@ -583,7 +559,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
pool_config (Optional[PoolArgs]): Configuration for the connection pool.
If provided, will create a connection pool and use it instead of a single connection.
This overrides the `pipeline` argument.
index (Optional[PostgresIndexConfig]): The index configuration for the store.
embedding (Optional[PostgresIndexConfig]): The embedding config.
Returns:
PostgresStore: A new PostgresStore instance.
@@ -729,6 +705,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
vectors = self.embeddings.embed_documents(
[param[-1] for param in txt_params]
)
queries.append(
(
query,
@@ -756,10 +733,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
[query for _, query in embedding_requests]
)
for (idx, _), embedding in zip(embedding_requests, embeddings):
_paramslist = queries[idx][1]
for i in range(len(_paramslist)):
if _paramslist[i] is _PLACEHOLDER:
_paramslist[i] = embedding
queries[idx][1][0] = embedding
for (idx, _), (query, params) in zip(search_ops, queries):
cur.execute(query, params)
@@ -793,10 +767,9 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
already exist and runs database migrations. It MUST be called directly by the user
the first time the store is used.
"""
def _get_version(cur: Cursor[dict[str, Any]], table: str) -> int:
with self._cursor() as cur:
try:
cur.execute(f"SELECT v FROM {table} ORDER BY v DESC LIMIT 1")
cur.execute("SELECT v FROM store_migrations ORDER BY v DESC LIMIT 1")
row = cast(dict, cur.fetchone())
if row is None:
version = -1
@@ -805,27 +778,21 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
except UndefinedTable:
version = -1
cur.execute(
f"""
CREATE TABLE IF NOT EXISTS {table} (
"""
CREATE TABLE IF NOT EXISTS store_migrations (
v INTEGER PRIMARY KEY
)
"""
)
return version
with self._cursor() as cur:
version = _get_version(cur, table="store_migrations")
for v, sql in enumerate(self.MIGRATIONS[version + 1 :], start=version + 1):
cur.execute(sql)
cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,))
if self.index_config:
version = _get_version(cur, table="vector_migrations")
for v, migration in enumerate(
self.VECTOR_MIGRATIONS[version + 1 :], start=version + 1
):
for v, migration in enumerate(
self.MIGRATIONS[version + 1 :], start=version + 1
):
if isinstance(migration, str):
sql = migration
else:
if migration.condition and not migration.condition(self):
continue
sql = migration.sql
if migration.params:
params = {
@@ -833,8 +800,8 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
for k, v in migration.params.items()
}
sql = sql % params
cur.execute(sql)
cur.execute("INSERT INTO vector_migrations (v) VALUES (%s)", (v,))
cur.execute(sql)
cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,))
class Row(TypedDict):
@@ -847,10 +814,6 @@ class Row(TypedDict):
# Private utilities
_DEFAULT_ANN_CONFIG = ANNIndexConfig(
vector_type="vector",
)
def _get_vector_type_ops(store: BasePostgresStore) -> str:
"""Get the vector type operator class based on config."""
@@ -858,7 +821,9 @@ def _get_vector_type_ops(store: BasePostgresStore) -> str:
return "vector_cosine_ops"
config = cast(PostgresIndexConfig, store.index_config)
index_config = config.get("ann_index_config", _DEFAULT_ANN_CONFIG).copy()
index_config = config.get(
"db_index_config", BasePostgresStore._get_default_index_config()
)
vector_type = cast(str, index_config.get("vector_type", "vector"))
if vector_type not in ("vector", "halfvec"):
raise ValueError(
@@ -890,7 +855,8 @@ def _get_index_params(store: Any) -> tuple[str, dict[str, Any]]:
return "hnsw", {}
config = cast(PostgresIndexConfig, store.index_config)
index_config = config.get("ann_index_config", _DEFAULT_ANN_CONFIG).copy()
default_config = BasePostgresStore._get_default_index_config()
index_config = config.get("db_index_config", default_config).copy()
kind = index_config.pop("kind", "hnsw")
index_config.pop("vector_type", None)
return kind, index_config
@@ -959,6 +925,32 @@ def _row_to_search_item(
)
def _row_to_search_item(
namespace: tuple[str, ...],
row: Row,
*,
loader: Optional[Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]] = None,
) -> SearchItem:
"""Convert a row from the database into an Item."""
loader = loader or _json_loads
val = row["value"]
score = row.get("score")
if score is not None:
try:
score = float(score) # type: ignore[arg-type]
except ValueError:
logger.warning("Invalid score: %s", score)
score = None
return SearchItem(
value=val if isinstance(val, dict) else loader(val),
key=row["key"],
namespace=namespace,
created_at=row["created_at"],
updated_at=row["updated_at"],
score=score,
)
def _group_ops(ops: Iterable[Op]) -> tuple[dict[type, list[tuple[int, Op]]], int]:
grouped_ops: dict[type, list[tuple[int, Op]]] = defaultdict(list)
tot = 0
@@ -988,17 +980,8 @@ def _decode_ns_bytes(namespace: Union[str, bytes, list]) -> tuple[str, ...]:
return tuple(namespace.split("."))
def _get_distance_operator(store: Any) -> tuple[str, str]:
def _get_distance_operator(store: Any) -> str:
"""Get the distance operator and score expression based on config."""
# Note: Today, we are not using ANN indices due to restrictions
# on PGVector's support for mixing vector and non-vector filters
# To use the index, PGVector expects:
# - ORDER BY the operator NOT an expression (even negation blocks it)
# - ASCENDING order
# - Any WHERE clause should be over a partial index.
# If we violate any of these, it will use a sequential scan
# See https://github.com/pgvector/pgvector/issues/216 and the
# pgvector documentation for more details.
if not store.index_config:
raise ValueError(
"Embedding configuration is required for vector operations "
@@ -1009,22 +992,12 @@ def _get_distance_operator(store: Any) -> tuple[str, str]:
config = cast(PostgresIndexConfig, store.index_config)
distance_type = config.get("distance_type", "cosine")
# Return the operator and the score expression
# The operator is used in the CTE and will be compatible with an ASCENDING ORDER
# sort clause.
# The score expression is used in the final query and will be compatible with
# a DESCENDING ORDER sort clause and the user's expectations of what the similarity score
# should be.
if distance_type == "l2":
# Final: "-(sv.embedding <-> %s::%s)"
# We return the "l2 similarity" so that the sorting order is the same
return "sv.embedding <-> %s::%s", "-scored.neg_score"
return "1 - (sv.embedding <-> %s::%s)"
elif distance_type == "inner_product":
# Final: "-(sv.embedding <#> %s::%s)"
return "sv.embedding <#> %s::%s", "-(scored.neg_score)"
else: # cosine similarity
# Final: "1 - (sv.embedding <=> %s::%s)"
return "sv.embedding <=> %s::%s", "1 - scored.neg_score"
return "-(sv.embedding <#> %s::%s)"
else: # cosine
return "1 - (sv.embedding <=> %s::%s)"
def _ensure_index_config(
@@ -1052,6 +1025,3 @@ def _ensure_index_config(
index_config.get("embed"),
)
return embeddings, index_config
_PLACEHOLDER = object()
+426 -529
View File
File diff suppressed because it is too large Load Diff
+4 -4
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "langgraph-checkpoint-postgres"
version = "2.0.7"
version = "2.0.4"
description = "Library with a Postgres implementation of LangGraph checkpoint saver."
authors = []
license = "MIT"
@@ -10,10 +10,10 @@ packages = [{ include = "langgraph" }]
[tool.poetry.dependencies]
python = "^3.9.0,<4.0"
langgraph-checkpoint = "^2.0.7"
langgraph-checkpoint = "^2.0.2"
orjson = ">=3.10.1"
psycopg = "^3.2.0"
psycopg-pool = "^3.2.0"
psycopg = "^3.0.0"
psycopg-pool = "^3.0.0"
[tool.poetry.group.dev.dependencies]
ruff = "^0.6.2"
+1 -1
View File
@@ -7,7 +7,6 @@ from psycopg.rows import DictRow, dict_row
from tests.embed_test_utils import CharacterEmbeddings
DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5441/"
DEFAULT_URI = "postgres://postgres:postgres@localhost:5441/postgres?sslmode=disable"
@@ -41,4 +40,5 @@ def fake_embeddings() -> CharacterEmbeddings:
return CharacterEmbeddings(dims=500)
INDEX_TYPES = ["hnsw", "ivfflat"]
VECTOR_TYPES = ["vector", "halfvec"]
+85 -200
View File
@@ -1,14 +1,7 @@
# type: ignore
from contextlib import asynccontextmanager
from typing import Any
from uuid import uuid4
import pytest
from langchain_core.runnables import RunnableConfig
from psycopg import AsyncConnection
from psycopg.rows import dict_row
from psycopg_pool import AsyncConnectionPool
from langgraph.checkpoint.base import (
Checkpoint,
@@ -17,212 +10,104 @@ from langgraph.checkpoint.base import (
empty_checkpoint,
)
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
from tests.conftest import DEFAULT_POSTGRES_URI
from tests.conftest import DEFAULT_URI
@asynccontextmanager
async def _pool_saver():
"""Fixture for pool mode testing."""
database = f"test_{uuid4().hex[:16]}"
# create unique db
async with await AsyncConnection.connect(
DEFAULT_POSTGRES_URI, autocommit=True
) as conn:
await conn.execute(f"CREATE DATABASE {database}")
try:
# yield checkpointer
async with AsyncConnectionPool(
DEFAULT_POSTGRES_URI + database,
max_size=10,
kwargs={"autocommit": True, "row_factory": dict_row},
) as pool:
checkpointer = AsyncPostgresSaver(pool)
await checkpointer.setup()
yield checkpointer
finally:
# drop unique db
async with await AsyncConnection.connect(
DEFAULT_POSTGRES_URI, autocommit=True
) as conn:
await conn.execute(f"DROP DATABASE {database}")
@asynccontextmanager
async def _pipe_saver():
"""Fixture for pipeline mode testing."""
database = f"test_{uuid4().hex[:16]}"
# create unique db
async with await AsyncConnection.connect(
DEFAULT_POSTGRES_URI, autocommit=True
) as conn:
await conn.execute(f"CREATE DATABASE {database}")
try:
async with await AsyncConnection.connect(
DEFAULT_POSTGRES_URI + database,
autocommit=True,
prepare_threshold=0,
row_factory=dict_row,
) as conn:
async with conn.pipeline() as pipe:
checkpointer = AsyncPostgresSaver(conn, pipe=pipe)
await checkpointer.setup()
async with conn.pipeline() as pipe:
checkpointer = AsyncPostgresSaver(conn, pipe=pipe)
yield checkpointer
finally:
# drop unique db
async with await AsyncConnection.connect(
DEFAULT_POSTGRES_URI, autocommit=True
) as conn:
await conn.execute(f"DROP DATABASE {database}")
@asynccontextmanager
async def _base_saver():
"""Fixture for regular connection mode testing."""
database = f"test_{uuid4().hex[:16]}"
# create unique db
async with await AsyncConnection.connect(
DEFAULT_POSTGRES_URI, autocommit=True
) as conn:
await conn.execute(f"CREATE DATABASE {database}")
try:
async with await AsyncConnection.connect(
DEFAULT_POSTGRES_URI + database,
autocommit=True,
prepare_threshold=0,
row_factory=dict_row,
) as conn:
checkpointer = AsyncPostgresSaver(conn)
await checkpointer.setup()
yield checkpointer
finally:
# drop unique db
async with await AsyncConnection.connect(
DEFAULT_POSTGRES_URI, autocommit=True
) as conn:
await conn.execute(f"DROP DATABASE {database}")
@asynccontextmanager
async def _saver(name: str):
if name == "base":
async with _base_saver() as saver:
yield saver
elif name == "pool":
async with _pool_saver() as saver:
yield saver
elif name == "pipe":
async with _pipe_saver() as saver:
yield saver
@pytest.fixture
def test_data():
"""Fixture providing test data for checkpoint tests."""
config_1: RunnableConfig = {
"configurable": {
"thread_id": "thread-1",
# for backwards compatibility testing
"thread_ts": "1",
"checkpoint_ns": "",
class TestAsyncPostgresSaver:
@pytest.fixture(autouse=True)
async def setup(self) -> None:
# objects for test setup
self.config_1: RunnableConfig = {
"configurable": {
"thread_id": "thread-1",
# for backwards compatibility testing
"thread_ts": "1",
"checkpoint_ns": "",
}
}
}
config_2: RunnableConfig = {
"configurable": {
"thread_id": "thread-2",
"checkpoint_id": "2",
"checkpoint_ns": "",
self.config_2: RunnableConfig = {
"configurable": {
"thread_id": "thread-2",
"checkpoint_id": "2",
"checkpoint_ns": "",
}
}
}
config_3: RunnableConfig = {
"configurable": {
"thread_id": "thread-2",
"checkpoint_id": "2-inner",
"checkpoint_ns": "inner",
self.config_3: RunnableConfig = {
"configurable": {
"thread_id": "thread-2",
"checkpoint_id": "2-inner",
"checkpoint_ns": "inner",
}
}
}
chkpnt_1: Checkpoint = empty_checkpoint()
chkpnt_2: Checkpoint = create_checkpoint(chkpnt_1, {}, 1)
chkpnt_3: Checkpoint = empty_checkpoint()
self.chkpnt_1: Checkpoint = empty_checkpoint()
self.chkpnt_2: Checkpoint = create_checkpoint(self.chkpnt_1, {}, 1)
self.chkpnt_3: Checkpoint = empty_checkpoint()
metadata_1: CheckpointMetadata = {
"source": "input",
"step": 2,
"writes": {},
"score": 1,
}
metadata_2: CheckpointMetadata = {
"source": "loop",
"step": 1,
"writes": {"foo": "bar"},
"score": None,
}
metadata_3: CheckpointMetadata = {}
return {
"configs": [config_1, config_2, config_3],
"checkpoints": [chkpnt_1, chkpnt_2, chkpnt_3],
"metadata": [metadata_1, metadata_2, metadata_3],
}
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
async def test_asearch(request, saver_name: str, test_data) -> None:
async with _saver(saver_name) as saver:
configs = test_data["configs"]
checkpoints = test_data["checkpoints"]
metadata = test_data["metadata"]
await saver.aput(configs[0], checkpoints[0], metadata[0], {})
await saver.aput(configs[1], checkpoints[1], metadata[1], {})
await saver.aput(configs[2], checkpoints[2], metadata[2], {})
# call method / assertions
query_1 = {"source": "input"} # search by 1 key
query_2 = {
self.metadata_1: CheckpointMetadata = {
"source": "input",
"step": 2,
"writes": {},
"score": 1,
}
self.metadata_2: CheckpointMetadata = {
"source": "loop",
"step": 1,
"writes": {"foo": "bar"},
} # search by multiple keys
query_3: dict[str, Any] = {} # search by no keys, return all checkpoints
query_4 = {"source": "update", "step": 1} # no match
"score": None,
}
self.metadata_3: CheckpointMetadata = {}
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
await saver.setup()
search_results_1 = [c async for c in saver.alist(None, filter=query_1)]
assert len(search_results_1) == 1
assert search_results_1[0].metadata == metadata[0]
async def test_asearch(self) -> None:
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
await saver.aput(self.config_1, self.chkpnt_1, self.metadata_1, {})
await saver.aput(self.config_2, self.chkpnt_2, self.metadata_2, {})
await saver.aput(self.config_3, self.chkpnt_3, self.metadata_3, {})
search_results_2 = [c async for c in saver.alist(None, filter=query_2)]
assert len(search_results_2) == 1
assert search_results_2[0].metadata == metadata[1]
# call method / assertions
query_1 = {"source": "input"} # search by 1 key
query_2 = {
"step": 1,
"writes": {"foo": "bar"},
} # search by multiple keys
query_3: dict[str, Any] = {} # search by no keys, return all checkpoints
query_4 = {"source": "update", "step": 1} # no match
search_results_3 = [c async for c in saver.alist(None, filter=query_3)]
assert len(search_results_3) == 3
search_results_1 = [c async for c in saver.alist(None, filter=query_1)]
assert len(search_results_1) == 1
assert search_results_1[0].metadata == self.metadata_1
search_results_4 = [c async for c in saver.alist(None, filter=query_4)]
assert len(search_results_4) == 0
search_results_2 = [c async for c in saver.alist(None, filter=query_2)]
assert len(search_results_2) == 1
assert search_results_2[0].metadata == self.metadata_2
# search by config (defaults to checkpoints across all namespaces)
search_results_5 = [
c async for c in saver.alist({"configurable": {"thread_id": "thread-2"}})
]
assert len(search_results_5) == 2
assert {
search_results_5[0].config["configurable"]["checkpoint_ns"],
search_results_5[1].config["configurable"]["checkpoint_ns"],
} == {"", "inner"}
search_results_3 = [c async for c in saver.alist(None, filter=query_3)]
assert len(search_results_3) == 3
search_results_4 = [c async for c in saver.alist(None, filter=query_4)]
assert len(search_results_4) == 0
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
async def test_null_chars(request, saver_name: str, test_data) -> None:
async with _saver(saver_name) as saver:
config = await saver.aput(
test_data["configs"][0],
test_data["checkpoints"][0],
{"my_key": "\x00abc"},
{},
)
assert (await saver.aget_tuple(config)).metadata["my_key"] == "abc" # type: ignore
assert [c async for c in saver.alist(None, filter={"my_key": "abc"})][
0
].metadata["my_key"] == "abc"
# search by config (defaults to checkpoints across all namespaces)
search_results_5 = [
c
async for c in saver.alist({"configurable": {"thread_id": "thread-2"}})
]
assert len(search_results_5) == 2
assert {
search_results_5[0].config["configurable"]["checkpoint_ns"],
search_results_5[1].config["configurable"]["checkpoint_ns"],
} == {"", "inner"}
# TODO: test before and limit params
async def test_null_chars(self) -> None:
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
config = await saver.aput(
self.config_1, self.chkpnt_1, {"my_key": "\x00abc"}, {}
)
assert (await saver.aget_tuple(config)).metadata["my_key"] == "abc" # type: ignore
assert [c async for c in saver.alist(None, filter={"my_key": "abc"})][
0
].metadata["my_key"] == "abc"
@@ -1,5 +1,4 @@
# type: ignore
import itertools
import sys
import uuid
from collections.abc import AsyncIterator
@@ -14,6 +13,7 @@ from langgraph.store.base import GetOp, Item, ListNamespacesOp, PutOp, SearchOp
from langgraph.store.postgres import AsyncPostgresStore
from tests.conftest import (
DEFAULT_URI,
INDEX_TYPES,
VECTOR_TYPES,
CharacterEmbeddings,
)
@@ -191,6 +191,7 @@ async def test_batch_list_namespaces_ops(store: AsyncPostgresStore) -> None:
@asynccontextmanager
async def _create_vector_store(
index_type: str,
vector_type: str,
distance_type: str,
fake_embeddings: CharacterEmbeddings,
@@ -214,7 +215,8 @@ async def _create_vector_store(
index_config = {
"dims": fake_embeddings.dims,
"embed": fake_embeddings,
"ann_index_config": {
"db_index_config": {
"kind": index_type,
"vector_type": vector_type,
},
"distance_type": distance_type,
@@ -242,22 +244,25 @@ async def _create_vector_store(
@pytest.fixture(
scope="function",
params=[
(vector_type, distance_type)
(index_type, vector_type, distance_type)
for index_type in INDEX_TYPES
for vector_type in VECTOR_TYPES
for distance_type in (
["hamming"] if vector_type == "bit" else ["l2", "inner_product", "cosine"]
(["hamming"] if index_type == "ivfflat" else ["hamming", "jaccard"])
if vector_type == "bit"
else ["l2", "inner_product", "cosine"]
)
],
ids=lambda p: f"{p[0]}_{p[1]}",
ids=lambda p: f"{p[0]}_{p[1]}_{p[2]}",
)
async def vector_store(
request,
fake_embeddings: CharacterEmbeddings,
) -> AsyncIterator[AsyncPostgresStore]:
"""Create a store with vector search enabled."""
vector_type, distance_type = request.param
index_type, vector_type, distance_type = request.param
async with _create_vector_store(
vector_type, distance_type, fake_embeddings
index_type, vector_type, distance_type, fake_embeddings
) as store:
yield store
@@ -412,19 +417,24 @@ async def test_vector_search_edge_cases(vector_store: AsyncPostgresStore) -> Non
@pytest.mark.parametrize(
"vector_type,distance_type",
"index_type,vector_type,distance_type",
[
*itertools.product(["vector", "halfvec"], ["cosine", "inner_product", "l2"]),
("ivfflat", "vector", "cosine"),
("hnsw", "vector", "cosine"),
("hnsw", "halfvec", "cosine"),
("hnsw", "halfvec", "inner_product"),
],
)
async def test_embed_with_path(
request: Any,
fake_embeddings: CharacterEmbeddings,
index_type: str,
vector_type: str,
distance_type: str,
) -> None:
"""Test vector search with specific text fields in Postgres store."""
async with _create_vector_store(
index_type,
vector_type,
distance_type,
fake_embeddings,
@@ -468,40 +478,3 @@ async def test_embed_with_path(
assert len(results) == 2
assert results[0].score < ascore
assert results[1].score < ascore
@pytest.mark.parametrize(
"vector_type,distance_type",
[
*itertools.product(["vector", "halfvec"], ["cosine", "inner_product", "l2"]),
],
)
async def test_search_sorting(
request: Any,
fake_embeddings: CharacterEmbeddings,
vector_type: str,
distance_type: str,
) -> None:
"""Test operation-level field configuration for vector search."""
async with _create_vector_store(
vector_type,
distance_type,
fake_embeddings,
text_fields=["key1"], # Default fields that won't match our test data
) as store:
amatch = {
"key1": "mmm",
}
await store.aput(("test", "M"), "M", amatch)
N = 100
for i in range(N):
await store.aput(("test", "A"), f"A{i}", {"key1": "no"})
for i in range(N):
await store.aput(("test", "Z"), f"Z{i}", {"key1": "no"})
results = await store.asearch(("test",), query="mmm", limit=10)
assert len(results) == 10
assert len(set(r.key for r in results)) == 10
assert results[0].key == "M"
assert results[0].score > results[1].score
+21 -177
View File
@@ -19,6 +19,7 @@ from langgraph.store.base import (
from langgraph.store.postgres import PostgresStore
from tests.conftest import (
DEFAULT_URI,
INDEX_TYPES,
VECTOR_TYPES,
CharacterEmbeddings,
)
@@ -351,6 +352,7 @@ class TestPostgresStore:
@contextmanager
def _create_vector_store(
index_type: str,
vector_type: str,
distance_type: str,
fake_embeddings: Embeddings,
@@ -371,7 +373,8 @@ def _create_vector_store(
index_config = {
"dims": fake_embeddings.dims,
"embed": fake_embeddings,
"ann_index_config": {
"db_index_config": {
"kind": index_type,
"vector_type": vector_type,
},
"distance_type": distance_type,
@@ -395,21 +398,26 @@ def _create_vector_store(
@pytest.fixture(
scope="function",
params=[
(vector_type, distance_type)
(index_type, vector_type, distance_type)
for index_type in INDEX_TYPES
for vector_type in VECTOR_TYPES
for distance_type in (
["hamming"] if vector_type == "bit" else ["l2", "inner_product", "cosine"]
(["hamming"] if index_type == "ivfflat" else ["hamming", "jaccard"])
if vector_type == "bit"
else ["l2", "inner_product", "cosine"]
)
],
ids=lambda p: f"{p[0]}_{p[1]}",
ids=lambda p: f"{p[0]}_{p[1]}_{p[2]}",
)
def vector_store(
request,
fake_embeddings: Embeddings,
) -> PostgresStore:
"""Create a store with vector search enabled."""
vector_type, distance_type = request.param
with _create_vector_store(vector_type, distance_type, fake_embeddings) as store:
index_type, vector_type, distance_type = request.param
with _create_vector_store(
index_type, vector_type, distance_type, fake_embeddings
) as store:
yield store
@@ -547,22 +555,24 @@ def test_vector_search_edge_cases(vector_store: PostgresStore) -> None:
@pytest.mark.parametrize(
"vector_type,distance_type",
"index_type,vector_type,distance_type",
[
("vector", "cosine"),
("vector", "inner_product"),
("halfvec", "cosine"),
("halfvec", "inner_product"),
("ivfflat", "vector", "cosine"),
("hnsw", "vector", "cosine"),
("hnsw", "halfvec", "cosine"),
("hnsw", "halfvec", "inner_product"),
],
)
def test_embed_with_path_sync(
request: Any,
fake_embeddings: CharacterEmbeddings,
index_type: str,
vector_type: str,
distance_type: str,
) -> None:
"""Test vector search with specific text fields in Postgres store."""
with _create_vector_store(
index_type,
vector_type,
distance_type,
fake_embeddings,
@@ -616,169 +626,3 @@ def test_embed_with_path_sync(
assert results[0].key != results[1].key
assert results[0].score < ascore
assert results[1].score < ascore
@pytest.mark.parametrize(
"vector_type,distance_type",
[
("vector", "cosine"),
("vector", "inner_product"),
("halfvec", "cosine"),
("halfvec", "inner_product"),
],
)
def test_embed_with_path_operation_config(
request: Any,
fake_embeddings: CharacterEmbeddings,
vector_type: str,
distance_type: str,
) -> None:
"""Test operation-level field configuration for vector search."""
with _create_vector_store(
vector_type,
distance_type,
fake_embeddings,
text_fields=["key17"], # Default fields that won't match our test data
) as store:
doc3 = {
"key0": "aaa",
"key1": "bbb",
"key2": "ccc",
"key3": "ddd",
}
doc4 = {
"key0": "eee",
"key1": "bbb", # Same as doc3.key1
"key2": "fff",
"key3": "ggg",
}
store.put(("test",), "doc3", doc3, index=["key0", "key1"])
store.put(("test",), "doc4", doc4, index=["key1", "key3"])
results = store.search(("test",), query="aaa")
assert len(results) == 2
assert results[0].key == "doc3"
assert len(set(r.key for r in results)) == 2
assert results[0].score > results[1].score
results = store.search(("test",), query="ggg")
assert len(results) == 2
assert results[0].key == "doc4"
assert results[0].score > results[1].score
results = store.search(("test",), query="bbb")
assert len(results) == 2
assert results[0].key != results[1].key
assert results[0].score == pytest.approx(results[1].score, abs=1e-3)
results = store.search(("test",), query="ccc")
assert len(results) == 2
assert all(
r.score < 0.9 for r in results
) # Unindexed field should have low scores
# Test index=False behavior
doc5 = {
"key0": "hhh",
"key1": "iii",
}
store.put(("test",), "doc5", doc5, index=False)
results = store.search(("test",))
assert len(results) == 3
assert all(r.score is None for r in results)
assert any(r.key == "doc5" for r in results)
results = store.search(("test",), query="hhh")
# TODO: We don't currently fill in additional results if there are not enough
# returned during vector search.
# assert len(results) == 3
# doc5_result = next(r for r in results if r.key == "doc5")
# assert doc5_result.score is None
def _cosine_similarity(X: list[float], Y: list[list[float]]) -> list[float]:
"""
Compute cosine similarity between a vector X and a matrix Y.
Lazy import numpy for efficiency.
"""
similarities = []
for y in Y:
dot_product = sum(a * b for a, b in zip(X, y))
norm1 = sum(a * a for a in X) ** 0.5
norm2 = sum(a * a for a in y) ** 0.5
similarity = dot_product / (norm1 * norm2) if norm1 > 0 and norm2 > 0 else 0.0
similarities.append(similarity)
return similarities
def _inner_product(X: list[float], Y: list[list[float]]) -> list[float]:
"""
Compute inner product between a vector X and a matrix Y.
Lazy import numpy for efficiency.
"""
similarities = []
for y in Y:
similarity = sum(a * b for a, b in zip(X, y))
similarities.append(similarity)
return similarities
def _neg_l2_distance(X: list[float], Y: list[list[float]]) -> list[float]:
"""
Compute l2 distance between a vector X and a matrix Y.
Lazy import numpy for efficiency.
"""
similarities = []
for y in Y:
similarity = sum((a - b) ** 2 for a, b in zip(X, y)) ** 0.5
similarities.append(-similarity)
return similarities
@pytest.mark.parametrize(
"vector_type,distance_type",
[
("vector", "cosine"),
("vector", "inner_product"),
("halfvec", "l2"),
],
)
@pytest.mark.parametrize("query", ["aaa", "bbb", "ccc", "abcd", "poisson"])
def test_scores(
fake_embeddings: CharacterEmbeddings,
vector_type: str,
distance_type: str,
query: str,
) -> None:
"""Test operation-level field configuration for vector search."""
with _create_vector_store(
vector_type,
distance_type,
fake_embeddings,
text_fields=["key0"],
) as store:
doc = {
"key0": "aaa",
}
store.put(("test",), "doc", doc, index=["key0", "key1"])
results = store.search((), query=query)
vec0 = fake_embeddings.embed_query(doc["key0"])
vec1 = fake_embeddings.embed_query(query)
if distance_type == "cosine":
similarities = _cosine_similarity(vec1, [vec0])
elif distance_type == "inner_product":
similarities = _inner_product(vec1, [vec0])
elif distance_type == "l2":
similarities = _neg_l2_distance(vec1, [vec0])
assert len(results) == 1
assert results[0].score == pytest.approx(similarities[0], abs=1e-3)
+84 -187
View File
@@ -1,14 +1,7 @@
# type: ignore
from contextlib import contextmanager
from typing import Any
from uuid import uuid4
import pytest
from langchain_core.runnables import RunnableConfig
from psycopg import Connection
from psycopg.rows import dict_row
from psycopg_pool import ConnectionPool
from langgraph.checkpoint.base import (
Checkpoint,
@@ -17,199 +10,103 @@ from langgraph.checkpoint.base import (
empty_checkpoint,
)
from langgraph.checkpoint.postgres import PostgresSaver
from tests.conftest import DEFAULT_POSTGRES_URI
from tests.conftest import DEFAULT_URI
@contextmanager
def _pool_saver():
"""Fixture for pool mode testing."""
database = f"test_{uuid4().hex[:16]}"
# create unique db
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
conn.execute(f"CREATE DATABASE {database}")
try:
# yield checkpointer
with ConnectionPool(
DEFAULT_POSTGRES_URI + database,
max_size=10,
kwargs={"autocommit": True, "row_factory": dict_row},
) as pool:
checkpointer = PostgresSaver(pool)
checkpointer.setup()
yield checkpointer
finally:
# drop unique db
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
conn.execute(f"DROP DATABASE {database}")
@contextmanager
def _pipe_saver():
"""Fixture for pipeline mode testing."""
database = f"test_{uuid4().hex[:16]}"
# create unique db
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
conn.execute(f"CREATE DATABASE {database}")
try:
with Connection.connect(
DEFAULT_POSTGRES_URI + database,
autocommit=True,
prepare_threshold=0,
row_factory=dict_row,
) as conn:
with conn.pipeline() as pipe:
checkpointer = PostgresSaver(conn, pipe=pipe)
checkpointer.setup()
with conn.pipeline() as pipe:
checkpointer = PostgresSaver(conn, pipe=pipe)
yield checkpointer
finally:
# drop unique db
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
conn.execute(f"DROP DATABASE {database}")
@contextmanager
def _base_saver():
"""Fixture for regular connection mode testing."""
database = f"test_{uuid4().hex[:16]}"
# create unique db
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
conn.execute(f"CREATE DATABASE {database}")
try:
with Connection.connect(
DEFAULT_POSTGRES_URI + database,
autocommit=True,
prepare_threshold=0,
row_factory=dict_row,
) as conn:
checkpointer = PostgresSaver(conn)
checkpointer.setup()
yield checkpointer
finally:
# drop unique db
with Connection.connect(DEFAULT_POSTGRES_URI, autocommit=True) as conn:
conn.execute(f"DROP DATABASE {database}")
@contextmanager
def _saver(name: str):
if name == "base":
with _base_saver() as saver:
yield saver
elif name == "pool":
with _pool_saver() as saver:
yield saver
elif name == "pipe":
with _pipe_saver() as saver:
yield saver
@pytest.fixture
def test_data():
"""Fixture providing test data for checkpoint tests."""
config_1: RunnableConfig = {
"configurable": {
"thread_id": "thread-1",
# for backwards compatibility testing
"thread_ts": "1",
"checkpoint_ns": "",
class TestPostgresSaver:
@pytest.fixture(autouse=True)
def setup(self) -> None:
# objects for test setup
self.config_1: RunnableConfig = {
"configurable": {
"thread_id": "thread-1",
# for backwards compatibility testing
"thread_ts": "1",
"checkpoint_ns": "",
}
}
}
config_2: RunnableConfig = {
"configurable": {
"thread_id": "thread-2",
"checkpoint_id": "2",
"checkpoint_ns": "",
self.config_2: RunnableConfig = {
"configurable": {
"thread_id": "thread-2",
"checkpoint_id": "2",
"checkpoint_ns": "",
}
}
}
config_3: RunnableConfig = {
"configurable": {
"thread_id": "thread-2",
"checkpoint_id": "2-inner",
"checkpoint_ns": "inner",
self.config_3: RunnableConfig = {
"configurable": {
"thread_id": "thread-2",
"checkpoint_id": "2-inner",
"checkpoint_ns": "inner",
}
}
}
chkpnt_1: Checkpoint = empty_checkpoint()
chkpnt_2: Checkpoint = create_checkpoint(chkpnt_1, {}, 1)
chkpnt_3: Checkpoint = empty_checkpoint()
self.chkpnt_1: Checkpoint = empty_checkpoint()
self.chkpnt_2: Checkpoint = create_checkpoint(self.chkpnt_1, {}, 1)
self.chkpnt_3: Checkpoint = empty_checkpoint()
metadata_1: CheckpointMetadata = {
"source": "input",
"step": 2,
"writes": {},
"score": 1,
}
metadata_2: CheckpointMetadata = {
"source": "loop",
"step": 1,
"writes": {"foo": "bar"},
"score": None,
}
metadata_3: CheckpointMetadata = {}
return {
"configs": [config_1, config_2, config_3],
"checkpoints": [chkpnt_1, chkpnt_2, chkpnt_3],
"metadata": [metadata_1, metadata_2, metadata_3],
}
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
def test_search(saver_name: str, test_data) -> None:
with _saver(saver_name) as saver:
configs = test_data["configs"]
checkpoints = test_data["checkpoints"]
metadata = test_data["metadata"]
saver.put(configs[0], checkpoints[0], metadata[0], {})
saver.put(configs[1], checkpoints[1], metadata[1], {})
saver.put(configs[2], checkpoints[2], metadata[2], {})
# call method / assertions
query_1 = {"source": "input"} # search by 1 key
query_2 = {
self.metadata_1: CheckpointMetadata = {
"source": "input",
"step": 2,
"writes": {},
"score": 1,
}
self.metadata_2: CheckpointMetadata = {
"source": "loop",
"step": 1,
"writes": {"foo": "bar"},
} # search by multiple keys
query_3: dict[str, Any] = {} # search by no keys, return all checkpoints
query_4 = {"source": "update", "step": 1} # no match
"score": None,
}
self.metadata_3: CheckpointMetadata = {}
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
saver.setup()
search_results_1 = list(saver.list(None, filter=query_1))
assert len(search_results_1) == 1
assert search_results_1[0].metadata == metadata[0]
def test_search(self) -> None:
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
# save checkpoints
saver.put(self.config_1, self.chkpnt_1, self.metadata_1, {})
saver.put(self.config_2, self.chkpnt_2, self.metadata_2, {})
saver.put(self.config_3, self.chkpnt_3, self.metadata_3, {})
search_results_2 = list(saver.list(None, filter=query_2))
assert len(search_results_2) == 1
assert search_results_2[0].metadata == metadata[1]
# call method / assertions
query_1 = {"source": "input"} # search by 1 key
query_2 = {
"step": 1,
"writes": {"foo": "bar"},
} # search by multiple keys
query_3: dict[str, Any] = {} # search by no keys, return all checkpoints
query_4 = {"source": "update", "step": 1} # no match
search_results_3 = list(saver.list(None, filter=query_3))
assert len(search_results_3) == 3
search_results_1 = list(saver.list(None, filter=query_1))
assert len(search_results_1) == 1
assert search_results_1[0].metadata == self.metadata_1
search_results_4 = list(saver.list(None, filter=query_4))
assert len(search_results_4) == 0
search_results_2 = list(saver.list(None, filter=query_2))
assert len(search_results_2) == 1
assert search_results_2[0].metadata == self.metadata_2
# search by config (defaults to checkpoints across all namespaces)
search_results_5 = list(saver.list({"configurable": {"thread_id": "thread-2"}}))
assert len(search_results_5) == 2
assert {
search_results_5[0].config["configurable"]["checkpoint_ns"],
search_results_5[1].config["configurable"]["checkpoint_ns"],
} == {"", "inner"}
search_results_3 = list(saver.list(None, filter=query_3))
assert len(search_results_3) == 3
search_results_4 = list(saver.list(None, filter=query_4))
assert len(search_results_4) == 0
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
def test_null_chars(saver_name: str, test_data) -> None:
with _saver(saver_name) as saver:
config = saver.put(
test_data["configs"][0],
test_data["checkpoints"][0],
{"my_key": "\x00abc"},
{},
)
assert saver.get_tuple(config).metadata["my_key"] == "abc" # type: ignore
assert (
list(saver.list(None, filter={"my_key": "abc"}))[0].metadata["my_key"]
== "abc"
)
# search by config (defaults to checkpoints across all namespaces)
search_results_5 = list(
saver.list({"configurable": {"thread_id": "thread-2"}})
)
assert len(search_results_5) == 2
assert {
search_results_5[0].config["configurable"]["checkpoint_ns"],
search_results_5[1].config["configurable"]["checkpoint_ns"],
} == {"", "inner"}
# TODO: test before and limit params
def test_null_chars(self) -> None:
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
config = saver.put(self.config_1, self.chkpnt_1, {"my_key": "\x00abc"}, {})
assert saver.get_tuple(config).metadata["my_key"] == "abc" # type: ignore
assert (
list(saver.list(None, filter={"my_key": "abc"}))[0].metadata["my_key"] # type: ignore
== "abc"
)
@@ -413,8 +413,6 @@ def _cosine_similarity(X: list[float], Y: list[list[float]]) -> list[float]:
Compute cosine similarity between a vector X and a matrix Y.
Lazy import numpy for efficiency.
"""
if not Y:
return []
if _check_numpy():
import numpy as np # type: ignore
+1 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "langgraph-checkpoint"
version = "2.0.8"
version = "2.0.5"
description = "Library with base interfaces for LangGraph checkpoint savers."
authors = []
license = "MIT"
+14 -24
View File
@@ -511,6 +511,19 @@ def dockerfile(save_path: str, config: pathlib.Path, add_docker_compose: bool) -
)
@click.argument("path", required=False)
@click.option(
"--template",
type=str,
help=TEMPLATE_HELP_STRING,
)
@cli.command("new", help="🌱 Create a new LangGraph project from a template.")
@log_command
def new(path: Optional[str], template: Optional[str]) -> None:
"""Create a new LangGraph project from a template."""
return create_new(path, template)
@click.option(
"--host",
default="127.0.0.1",
@@ -550,12 +563,6 @@ def dockerfile(save_path: str, config: pathlib.Path, add_docker_compose: bool) -
type=int,
help="Enable remote debugging by listening on specified port. Requires debugpy to be installed",
)
@click.option(
"--wait-for-client",
is_flag=True,
help="Wait for a debugger client to connect to the debug port before starting the server",
default=False,
)
@cli.command(
"dev",
help="🏃‍♀️‍➡️ Run LangGraph API server in development mode with hot reloading and debugging support",
@@ -569,7 +576,6 @@ def dev(
n_jobs_per_worker: Optional[int],
no_browser: bool,
debug_port: Optional[int],
wait_for_client: bool,
):
"""CLI entrypoint for running the LangGraph API server."""
try:
@@ -602,7 +608,6 @@ def dev(
sys.path.append(str(dep_path))
graphs = config_json.get("graphs", {})
run_server(
host,
port,
@@ -611,25 +616,10 @@ def dev(
n_jobs_per_worker=n_jobs_per_worker,
open_browser=not no_browser,
debug_port=debug_port,
env=config_json.get("env"),
store=config_json.get("store"),
wait_for_client=wait_for_client,
env=config_json.get("env", None),
)
@click.argument("path", required=False)
@click.option(
"--template",
type=str,
help=TEMPLATE_HELP_STRING,
)
@cli.command("new", help="🌱 Create a new LangGraph project from a template.")
@log_command
def new(path: Optional[str], template: Optional[str]) -> None:
"""Create a new LangGraph project from a template."""
return create_new(path, template)
def prepare_args_and_stdin(
*,
capabilities: DockerCapabilities,
+5 -59
View File
@@ -10,44 +10,7 @@ MIN_NODE_VERSION = "20"
MIN_PYTHON_VERSION = "3.11"
class IndexConfig(TypedDict, total=False):
"""Configuration for indexing documents for semantic search in the store."""
dims: int
"""Number of dimensions in the embedding vectors.
Common embedding models have the following dimensions:
- OpenAI text-embedding-3-large: 256, 1024, or 3072
- OpenAI text-embedding-3-small: 512 or 1536
- OpenAI text-embedding-ada-002: 1536
- Cohere embed-english-v3.0: 1024
- Cohere embed-english-light-v3.0: 384
- Cohere embed-multilingual-v3.0: 1024
- Cohere embed-multilingual-light-v3.0: 384
"""
embed: str
"""Optional model (string) to generate embeddings from text or path to model or function.
Examples:
- "openai:text-embedding-3-large"
- "cohere:embed-multilingual-v3.0"
- "src/app.py:embeddings
"""
fields: Optional[list[str]]
"""Fields to extract text from for embedding generation.
Defaults to the root ["$"], which embeds the json object as a whole.
"""
class StoreConfig(TypedDict, total=False):
embed: Optional[IndexConfig]
"""Configuration for vector embeddings in store."""
class Config(TypedDict, total=False):
class Config(TypedDict):
python_version: str
node_version: Optional[str]
pip_config_file: Optional[str]
@@ -55,7 +18,6 @@ class Config(TypedDict, total=False):
dependencies: list[str]
graphs: dict[str, str]
env: Union[dict[str, str], str]
store: Optional[StoreConfig]
def _parse_version(version_str: str) -> tuple[int, int]:
@@ -87,7 +49,6 @@ def validate_config(config: Config) -> Config:
"dockerfile_lines": config.get("dockerfile_lines", []),
"graphs": config.get("graphs", {}),
"env": config.get("env", {}),
"store": config.get("store"),
}
if config.get("node_version")
else {
@@ -97,7 +58,6 @@ def validate_config(config: Config) -> Config:
"dependencies": config.get("dependencies", []),
"graphs": config.get("graphs", {}),
"env": config.get("env", {}),
"store": config.get("store"),
}
)
@@ -392,14 +352,7 @@ RUN set -ex && \\
],
)
)
store_config = config.get("store")
env_additional_config = (
""
if not store_config
else f"""
ENV LANGGRAPH_STORE='{json.dumps(store_config)}'
"""
)
return f"""FROM {base_image}:{config['python_version']}
{os.linesep.join(config["dockerfile_lines"])}
@@ -407,7 +360,7 @@ ENV LANGGRAPH_STORE='{json.dumps(store_config)}'
{installs}
RUN {pip_install} -e /deps/*
{env_additional_config}
ENV LANGSERVE_GRAPHS='{json.dumps(config["graphs"])}'
{f"WORKDIR {local_deps.working_dir}" if local_deps.working_dir else ""}"""
@@ -437,14 +390,7 @@ def node_config_to_docker(config_path: pathlib.Path, config: Config, base_image:
install_cmd = "npm ci"
else:
install_cmd = "npm i"
store_config = config.get("store")
env_additional_config = (
""
if not store_config
else f"""
ENV LANGGRAPH_STORE='{json.dumps(store_config)}'
"""
)
return f"""FROM {base_image}:{config['node_version']}
{os.linesep.join(config["dockerfile_lines"])}
@@ -452,7 +398,7 @@ ENV LANGGRAPH_STORE='{json.dumps(store_config)}'
ADD . {faux_path}
RUN cd {faux_path} && {install_cmd}
{env_additional_config}
ENV LANGSERVE_GRAPHS='{json.dumps(config["graphs"])}'
WORKDIR {faux_path}
+21 -54
View File
@@ -565,13 +565,13 @@ langgraph-sdk = ">=0.1.32,<0.2.0"
[[package]]
name = "langgraph-api"
version = "0.0.6"
version = "0.0.2"
description = ""
optional = true
python-versions = "<4.0,>=3.11.0"
files = [
{file = "langgraph_api-0.0.6-py3-none-any.whl", hash = "sha256:f64b13959d721143f6a023af5b9ffc9aa054064af98d21d5d8090cda7e7bffd2"},
{file = "langgraph_api-0.0.6.tar.gz", hash = "sha256:badac44fa1ec979509e56fc0da57eeb5f278ee5871f27803f73ea6d8822c21b9"},
{file = "langgraph_api-0.0.2-py3-none-any.whl", hash = "sha256:7a30fb21987572eacc93dd1c69c2155c17957afed71dde18d6f47992b3124d65"},
{file = "langgraph_api-0.0.2.tar.gz", hash = "sha256:b751afca96cb6db67fe2f48e798ada27a4df068f0df86b36d2b8eee52344bbf0"},
]
[package.dependencies]
@@ -579,8 +579,8 @@ cryptography = ">=43.0.3,<44.0.0"
httpx = ">=0.27.0"
jsonschema-rs = ">=0.25.0,<0.26.0"
langchain-core = ">=0.2.38,<0.4.0"
langgraph = ">=0.2.52,<0.3.0"
langgraph-checkpoint = ">=2.0.7,<3.0"
langgraph = ">=0.2.52"
langgraph-checkpoint = ">=2.0.5,<3.0"
langsmith = ">=0.1.63,<0.2.0"
orjson = ">=3.10.1"
pyjwt = ">=2.9.0,<3.0.0"
@@ -593,13 +593,13 @@ watchfiles = ">=0.13"
[[package]]
name = "langgraph-checkpoint"
version = "2.0.7"
version = "2.0.6"
description = "Library with base interfaces for LangGraph checkpoint savers."
optional = true
python-versions = "<4.0.0,>=3.9.0"
files = [
{file = "langgraph_checkpoint-2.0.7-py3-none-any.whl", hash = "sha256:9709f672e1c5a47e13352067c2ffa114dd91d443967b7ce8a1d36d6fc170370e"},
{file = "langgraph_checkpoint-2.0.7.tar.gz", hash = "sha256:88d648a331d20aa8ce65280de34a34a9190380b004f6afcc5f9894fe3abeed08"},
{file = "langgraph_checkpoint-2.0.6-py3-none-any.whl", hash = "sha256:2878283c3ee2519bf180df9b7b7155b73fa05eb63b1af9600a03e03a930d8c53"},
{file = "langgraph_checkpoint-2.0.6.tar.gz", hash = "sha256:69ab9c61c4e2992264671f55579c24070b7b6cedc105a33da3fba6526df248cf"},
]
[package.dependencies]
@@ -608,13 +608,13 @@ msgpack = ">=1.1.0,<2.0.0"
[[package]]
name = "langgraph-sdk"
version = "0.1.40"
version = "0.1.36"
description = "SDK for interacting with LangGraph API"
optional = true
python-versions = "<4.0.0,>=3.9.0"
files = [
{file = "langgraph_sdk-0.1.40-py3-none-any.whl", hash = "sha256:8810cca5e4144cf3a5441fc76b4ee6e658ec95f932d3a0bf9ad63de117e925b9"},
{file = "langgraph_sdk-0.1.40.tar.gz", hash = "sha256:ab2719ac7274612a791a7a0ad9395d250357106cba8ba81bca9968fc91009af2"},
{file = "langgraph_sdk-0.1.36-py3-none-any.whl", hash = "sha256:b11e1f0bc67631134d09d50c812dc73f9eb30394764ae1144d7d2a786a715355"},
{file = "langgraph_sdk-0.1.36.tar.gz", hash = "sha256:2a2c651b7851ba15aeaab7e4e3ea7fd8357ef1cb0b592f264916fa990cdda6e7"},
]
[package.dependencies]
@@ -624,13 +624,13 @@ orjson = ">=3.10.1"
[[package]]
name = "langsmith"
version = "0.1.147"
version = "0.1.146"
description = "Client library to connect to the LangSmith LLM Tracing and Evaluation Platform."
optional = true
python-versions = "<4.0,>=3.8.1"
files = [
{file = "langsmith-0.1.147-py3-none-any.whl", hash = "sha256:7166fc23b965ccf839d64945a78e9f1157757add228b086141eb03a60d699a15"},
{file = "langsmith-0.1.147.tar.gz", hash = "sha256:2e933220318a4e73034657103b3b1a3a6109cc5db3566a7e8e03be8d6d7def7a"},
{file = "langsmith-0.1.146-py3-none-any.whl", hash = "sha256:9d062222f1a32c9b047dab0149b24958f988989cd8d4a5f9139ff959a51e59d8"},
{file = "langsmith-0.1.146.tar.gz", hash = "sha256:ead8b0b9d5b6cd3ac42937ec48bdf09d4afe7ca1bba22dc05eb65591a18106f8"},
]
[package.dependencies]
@@ -643,9 +643,6 @@ pydantic = [
requests = ">=2,<3"
requests-toolbelt = ">=1.0.0,<2.0.0"
[package.extras]
langsmith-pyo3 = ["langsmith-pyo3 (>=0.1.0rc2,<0.2.0)"]
[[package]]
name = "msgpack"
version = "1.1.0"
@@ -1038,13 +1035,13 @@ typing-extensions = ">=4.6.0,<4.7.0 || >4.7.0"
[[package]]
name = "pyjwt"
version = "2.10.1"
version = "2.10.0"
description = "JSON Web Token implementation in Python"
optional = true
python-versions = ">=3.9"
files = [
{file = "PyJWT-2.10.1-py3-none-any.whl", hash = "sha256:dcdd193e30abefd5debf142f9adfcdd2b58004e644f25406ffaebd50bd98dacb"},
{file = "pyjwt-2.10.1.tar.gz", hash = "sha256:3cc5772eb20009233caf06e9d8a0577824723b44e6648ee0a2aedb6cf9381953"},
{file = "PyJWT-2.10.0-py3-none-any.whl", hash = "sha256:543b77207db656de204372350926bed5a86201c4cbff159f623f79c7bb487a15"},
{file = "pyjwt-2.10.0.tar.gz", hash = "sha256:7628a7eb7938959ac1b26e819a1df0fd3259505627b575e4bad6d08f76db695c"},
]
[package.extras]
@@ -1345,43 +1342,13 @@ test = ["pytest", "tornado (>=4.5)", "typeguard"]
[[package]]
name = "tomli"
version = "2.2.1"
version = "2.1.0"
description = "A lil' TOML parser"
optional = false
python-versions = ">=3.8"
files = [
{file = "tomli-2.2.1-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:678e4fa69e4575eb77d103de3df8a895e1591b48e740211bd1067378c69e8249"},
{file = "tomli-2.2.1-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:023aa114dd824ade0100497eb2318602af309e5a55595f76b626d6d9f3b7b0a6"},
{file = "tomli-2.2.1-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:ece47d672db52ac607a3d9599a9d48dcb2f2f735c6c2d1f34130085bb12b112a"},
{file = "tomli-2.2.1-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:6972ca9c9cc9f0acaa56a8ca1ff51e7af152a9f87fb64623e31d5c83700080ee"},
{file = "tomli-2.2.1-cp311-cp311-manylinux_2_5_i686.manylinux1_i686.manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:c954d2250168d28797dd4e3ac5cf812a406cd5a92674ee4c8f123c889786aa8e"},
{file = "tomli-2.2.1-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:8dd28b3e155b80f4d54beb40a441d366adcfe740969820caf156c019fb5c7ec4"},
{file = "tomli-2.2.1-cp311-cp311-musllinux_1_2_i686.whl", hash = "sha256:e59e304978767a54663af13c07b3d1af22ddee3bb2fb0618ca1593e4f593a106"},
{file = "tomli-2.2.1-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:33580bccab0338d00994d7f16f4c4ec25b776af3ffaac1ed74e0b3fc95e885a8"},
{file = "tomli-2.2.1-cp311-cp311-win32.whl", hash = "sha256:465af0e0875402f1d226519c9904f37254b3045fc5084697cefb9bdde1ff99ff"},
{file = "tomli-2.2.1-cp311-cp311-win_amd64.whl", hash = "sha256:2d0f2fdd22b02c6d81637a3c95f8cd77f995846af7414c5c4b8d0545afa1bc4b"},
{file = "tomli-2.2.1-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:4a8f6e44de52d5e6c657c9fe83b562f5f4256d8ebbfe4ff922c495620a7f6cea"},
{file = "tomli-2.2.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:8d57ca8095a641b8237d5b079147646153d22552f1c637fd3ba7f4b0b29167a8"},
{file = "tomli-2.2.1-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:4e340144ad7ae1533cb897d406382b4b6fede8890a03738ff1683af800d54192"},
{file = "tomli-2.2.1-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:db2b95f9de79181805df90bedc5a5ab4c165e6ec3fe99f970d0e302f384ad222"},
{file = "tomli-2.2.1-cp312-cp312-manylinux_2_5_i686.manylinux1_i686.manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:40741994320b232529c802f8bc86da4e1aa9f413db394617b9a256ae0f9a7f77"},
{file = "tomli-2.2.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:400e720fe168c0f8521520190686ef8ef033fb19fc493da09779e592861b78c6"},
{file = "tomli-2.2.1-cp312-cp312-musllinux_1_2_i686.whl", hash = "sha256:02abe224de6ae62c19f090f68da4e27b10af2b93213d36cf44e6e1c5abd19fdd"},
{file = "tomli-2.2.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:b82ebccc8c8a36f2094e969560a1b836758481f3dc360ce9a3277c65f374285e"},
{file = "tomli-2.2.1-cp312-cp312-win32.whl", hash = "sha256:889f80ef92701b9dbb224e49ec87c645ce5df3fa2cc548664eb8a25e03127a98"},
{file = "tomli-2.2.1-cp312-cp312-win_amd64.whl", hash = "sha256:7fc04e92e1d624a4a63c76474610238576942d6b8950a2d7f908a340494e67e4"},
{file = "tomli-2.2.1-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:f4039b9cbc3048b2416cc57ab3bda989a6fcf9b36cf8937f01a6e731b64f80d7"},
{file = "tomli-2.2.1-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:286f0ca2ffeeb5b9bd4fcc8d6c330534323ec51b2f52da063b11c502da16f30c"},
{file = "tomli-2.2.1-cp313-cp313-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:a92ef1a44547e894e2a17d24e7557a5e85a9e1d0048b0b5e7541f76c5032cb13"},
{file = "tomli-2.2.1-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:9316dc65bed1684c9a98ee68759ceaed29d229e985297003e494aa825ebb0281"},
{file = "tomli-2.2.1-cp313-cp313-manylinux_2_5_i686.manylinux1_i686.manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:e85e99945e688e32d5a35c1ff38ed0b3f41f43fad8df0bdf79f72b2ba7bc5272"},
{file = "tomli-2.2.1-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:ac065718db92ca818f8d6141b5f66369833d4a80a9d74435a268c52bdfa73140"},
{file = "tomli-2.2.1-cp313-cp313-musllinux_1_2_i686.whl", hash = "sha256:d920f33822747519673ee656a4b6ac33e382eca9d331c87770faa3eef562aeb2"},
{file = "tomli-2.2.1-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:a198f10c4d1b1375d7687bc25294306e551bf1abfa4eace6650070a5c1ae2744"},
{file = "tomli-2.2.1-cp313-cp313-win32.whl", hash = "sha256:d3f5614314d758649ab2ab3a62d4f2004c825922f9e370b29416484086b264ec"},
{file = "tomli-2.2.1-cp313-cp313-win_amd64.whl", hash = "sha256:a38aa0308e754b0e3c67e344754dff64999ff9b513e691d0e786265c93583c69"},
{file = "tomli-2.2.1-py3-none-any.whl", hash = "sha256:cb55c73c5f4408779d0cf3eef9f762b9c9f147a77de7b258bef0a5628adc85cc"},
{file = "tomli-2.2.1.tar.gz", hash = "sha256:cd45e1dc79c835ce60f7404ec8119f2eb06d38b1deba146f07ced3bbc44505ff"},
{file = "tomli-2.1.0-py3-none-any.whl", hash = "sha256:a5c57c3d1c56f5ccdf89f6523458f60ef716e210fc47c4cfb188c5ba473e0391"},
{file = "tomli-2.1.0.tar.gz", hash = "sha256:3f646cae2aec94e17d04973e4249548320197cfabdf130015d023de4b74d8ab8"},
]
[[package]]
@@ -1561,4 +1528,4 @@ inmem = ["langgraph-api", "python-dotenv"]
[metadata]
lock-version = "2.0"
python-versions = "^3.9.0,<4.0"
content-hash = "8eaaa66d9e6e447699e3bcee336dfe779b58c956f8c2ad6678008a07be935838"
content-hash = "3d655bb578e20219e19152d4a3d86be370fe3be61b5559847f0204dfff499b4a"
+2 -2
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "langgraph-cli"
version = "0.1.61"
version = "0.1.59"
description = "CLI for interacting with LangGraph API"
authors = []
license = "MIT"
@@ -14,7 +14,7 @@ langgraph = "langgraph_cli.cli:cli"
[tool.poetry.dependencies]
python = "^3.9.0,<4.0"
click = "^8.1.7"
langgraph-api = { version = ">=0.0.6,<0.1.0", optional = true, python = ">=3.11,<4.0" }
langgraph-api = { version = ">=0.0.2,<0.1.0", optional = true, python = ">=3.11,<4.0" }
python-dotenv = { version = ">=0.8.0", optional = true }
[tool.poetry.group.dev.dependencies]
-2
View File
@@ -30,7 +30,6 @@ def test_validate_config():
"pip_config_file": None,
"dockerfile_lines": [],
"env": {},
"store": None,
**expected_config,
}
actual_config = validate_config(expected_config)
@@ -47,7 +46,6 @@ def test_validate_config():
"agent": "./agent.py:graph",
},
"env": env,
"store": None,
}
actual_config = validate_config(expected_config)
assert actual_config == expected_config
+4 -5
View File
@@ -374,11 +374,6 @@ class Graph:
if source not in self.nodes and source != START:
raise ValueError(f"Found edge starting at unknown node '{source}'")
if START not in all_sources:
raise ValueError(
"Graph must have an entrypoint: add at least one edge from START to another node"
)
# assemble targets
all_targets = {end for _, end in self._all_edges}
for start, branches in self.branches.items():
@@ -400,6 +395,10 @@ class Graph:
for name, spec in self.nodes.items():
if spec.ends:
all_targets.update(spec.ends)
# validate targets
for node in self.nodes:
if node not in all_targets:
raise ValueError(f"Node `{node}` is not reachable")
for target in all_targets:
if target not in self.nodes and target != END:
raise ValueError(f"Found edge ending at unknown node `{target}`")
+2 -2
View File
@@ -933,12 +933,12 @@ def _is_field_binop(typ: Type[Any]) -> Optional[BinaryOperatorAggregate]:
if hasattr(typ, "__metadata__"):
meta = typ.__metadata__
if len(meta) >= 1 and callable(meta[-1]):
sig = signature(meta[-1])
sig = signature(meta[0])
params = list(sig.parameters.values())
if len(params) == 2 and all(
p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD) for p in params
):
return BinaryOperatorAggregate(typ, meta[-1])
return BinaryOperatorAggregate(typ, meta[0])
else:
raise ValueError(
f"Invalid reducer signature. Expected (a, b) -> c. Got {sig}"
+2 -2
View File
@@ -595,7 +595,7 @@ def prepare_single_task(
for tid, c, v in pending_writes
if tid in (NULL_TASK_ID, task_id) and c == RESUME
),
configurable.get(CONFIG_KEY_RESUME_VALUE, MISSING),
MISSING,
),
},
),
@@ -720,7 +720,7 @@ def prepare_single_task(
if tid in (NULL_TASK_ID, task_id)
and c == RESUME
),
configurable.get(CONFIG_KEY_RESUME_VALUE, MISSING),
MISSING,
),
},
),
+9 -27
View File
@@ -1,4 +1,3 @@
from dataclasses import asdict
from typing import (
Any,
AsyncIterator,
@@ -28,7 +27,6 @@ from langgraph_sdk.client import (
get_sync_client,
)
from langgraph_sdk.schema import Checkpoint, ThreadState
from langgraph_sdk.schema import Command as CommandSDK
from langgraph_sdk.schema import StreamMode as StreamModeSDK
from typing_extensions import Self
@@ -43,7 +41,7 @@ from langgraph.constants import (
from langgraph.errors import GraphInterrupt
from langgraph.pregel.protocol import PregelProtocol
from langgraph.pregel.types import All, PregelTask, StateSnapshot, StreamMode
from langgraph.types import Command, Interrupt, StreamProtocol
from langgraph.types import Interrupt, StreamProtocol
from langgraph.utils.config import merge_configs
@@ -575,7 +573,6 @@ class RemoteGraph(PregelProtocol):
interrupt_before: Optional[Union[All, Sequence[str]]] = None,
interrupt_after: Optional[Union[All, Sequence[str]]] = None,
subgraphs: bool = False,
**kwargs: Any,
) -> Iterator[Union[dict[str, Any], Any]]:
"""Create a run and stream the results.
@@ -590,7 +587,6 @@ class RemoteGraph(PregelProtocol):
interrupt_before: Interrupt the graph before these nodes.
interrupt_after: Interrupt the graph after these nodes.
subgraphs: Stream from subgraphs.
**kwargs: Additional params to pass to client.runs.stream.
Yields:
The output of the graph.
@@ -601,24 +597,17 @@ class RemoteGraph(PregelProtocol):
stream_modes, requested, req_single, stream = self._get_stream_modes(
stream_mode, config
)
if isinstance(input, Command):
command: Optional[CommandSDK] = cast(CommandSDK, asdict(input))
input = None
else:
command = None
for chunk in sync_client.runs.stream(
thread_id=sanitized_config["configurable"].get("thread_id"),
assistant_id=self.name,
input=input,
command=command,
config=sanitized_config,
stream_mode=stream_modes,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
stream_subgraphs=subgraphs or stream is not None,
if_not_exists="create",
**kwargs,
):
# split mode and ns
if NS_SEP in chunk.event:
@@ -667,7 +656,6 @@ class RemoteGraph(PregelProtocol):
interrupt_before: Optional[Union[All, Sequence[str]]] = None,
interrupt_after: Optional[Union[All, Sequence[str]]] = None,
subgraphs: bool = False,
**kwargs: Any,
) -> AsyncIterator[Union[dict[str, Any], Any]]:
"""Create a run and stream the results.
@@ -682,7 +670,6 @@ class RemoteGraph(PregelProtocol):
interrupt_before: Interrupt the graph before these nodes.
interrupt_after: Interrupt the graph after these nodes.
subgraphs: Stream from subgraphs.
**kwargs: Additional params to pass to client.runs.stream.
Yields:
The output of the graph.
@@ -693,24 +680,17 @@ class RemoteGraph(PregelProtocol):
stream_modes, requested, req_single, stream = self._get_stream_modes(
stream_mode, config
)
if isinstance(input, Command):
command: Optional[CommandSDK] = cast(CommandSDK, asdict(input))
input = None
else:
command = None
async for chunk in client.runs.stream(
thread_id=sanitized_config["configurable"].get("thread_id"),
assistant_id=self.name,
input=input,
command=command,
config=sanitized_config,
stream_mode=stream_modes,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
stream_subgraphs=subgraphs or stream is not None,
if_not_exists="create",
**kwargs,
):
# split mode and ns
if NS_SEP in chunk.event:
@@ -773,16 +753,18 @@ class RemoteGraph(PregelProtocol):
*,
interrupt_before: Optional[Union[All, Sequence[str]]] = None,
interrupt_after: Optional[Union[All, Sequence[str]]] = None,
**kwargs: Any,
) -> Union[dict[str, Any], Any]:
"""Create a run, wait until it finishes and return the final state.
This method calls `POST /threads/{thread_id}/runs/wait` if a `thread_id`
is speciffed in the `configurable` field of the config or
`POST /runs/wait` otherwise.
Args:
input: Input to the graph.
config: A `RunnableConfig` for graph invocation.
interrupt_before: Interrupt the graph before these nodes.
interrupt_after: Interrupt the graph after these nodes.
**kwargs: Additional params to pass to RemoteGraph.stream.
Returns:
The output of the graph.
@@ -793,7 +775,6 @@ class RemoteGraph(PregelProtocol):
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
stream_mode="values",
**kwargs,
):
pass
try:
@@ -808,16 +789,18 @@ class RemoteGraph(PregelProtocol):
*,
interrupt_before: Optional[Union[All, Sequence[str]]] = None,
interrupt_after: Optional[Union[All, Sequence[str]]] = None,
**kwargs: Any,
) -> Union[dict[str, Any], Any]:
"""Create a run, wait until it finishes and return the final state.
This method calls `POST /threads/{thread_id}/runs/wait` if a `thread_id`
is speciffed in the `configurable` field of the config or
`POST /runs/wait` otherwise.
Args:
input: Input to the graph.
config: A `RunnableConfig` for graph invocation.
interrupt_before: Interrupt the graph before these nodes.
interrupt_after: Interrupt the graph after these nodes.
**kwargs: Additional params to pass to RemoteGraph.astream.
Returns:
The output of the graph.
@@ -828,7 +811,6 @@ class RemoteGraph(PregelProtocol):
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
stream_mode="values",
**kwargs,
):
pass
try:
+13 -222
View File
@@ -160,7 +160,7 @@ def test_graph_validation() -> None:
workflow = Graph()
workflow.add_node("agent", logic)
workflow.set_finish_point("agent")
with pytest.raises(ValueError, match="must have an entrypoint"):
with pytest.raises(ValueError, match="not reachable"):
workflow.compile()
workflow = Graph()
@@ -207,6 +207,18 @@ def test_graph_validation() -> None:
with pytest.raises(ValueError, match="unknown"): # extra is not defined
workflow.compile()
workflow = Graph()
workflow.add_node("agent", logic)
workflow.add_node("tools", logic)
workflow.add_node("extra", logic)
workflow.set_entry_point("agent")
workflow.add_conditional_edges("agent", logic, {"continue": "tools", "exit": END})
workflow.add_edge("tools", "agent")
with pytest.raises(
ValueError, match="Node `extra` is not reachable"
): # extra is not reachable
workflow.compile()
workflow = Graph()
workflow.add_node("agent", logic)
workflow.add_node("tools", logic)
@@ -264,25 +276,6 @@ def test_graph_validation() -> None:
graph.invoke({"hello": "there"})
def test_graph_validation_with_command() -> None:
class State(TypedDict):
foo: str
bar: str
def node_a(state: State):
return GraphCommand(goto="b", update={"foo": "bar"})
def node_b(state: State):
return GraphCommand(goto=END, update={"bar": "baz"})
builder = StateGraph(State)
builder.add_node("a", node_a)
builder.add_node("b", node_b)
builder.add_edge(START, "a")
graph = builder.compile()
assert graph.invoke({"foo": ""}) == {"foo": "bar", "bar": "baz"}
def test_checkpoint_errors() -> None:
class FaultyGetCheckpointer(MemorySaver):
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
@@ -8735,176 +8728,6 @@ def test_copy_checkpoint(
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_dynamic_interrupt_subgraph(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
class SubgraphState(TypedDict):
my_key: str
market: str
tool_two_node_count = 0
def tool_two_node(s: SubgraphState) -> SubgraphState:
nonlocal tool_two_node_count
tool_two_node_count += 1
if s["market"] == "DE":
answer = interrupt("Just because...")
else:
answer = " all good"
return {"my_key": answer}
subgraph = StateGraph(SubgraphState)
subgraph.add_node("do", tool_two_node, retry=RetryPolicy())
subgraph.add_edge(START, "do")
class State(TypedDict):
my_key: Annotated[str, operator.add]
market: str
tool_two_graph = StateGraph(State)
tool_two_graph.add_node("tool_two", subgraph.compile())
tool_two_graph.add_edge(START, "tool_two")
tool_two = tool_two_graph.compile()
tracer = FakeTracer()
assert tool_two.invoke(
{"my_key": "value", "market": "DE"}, {"callbacks": [tracer]}
) == {
"my_key": "value",
"market": "DE",
}
assert tool_two_node_count == 1, "interrupts aren't retried"
assert len(tracer.runs) == 1
run = tracer.runs[0]
assert run.end_time is not None
assert run.error is None
assert run.outputs == {"market": "DE", "my_key": "value"}
assert tool_two.invoke({"my_key": "value", "market": "US"}) == {
"my_key": "value all good",
"market": "US",
}
tool_two = tool_two_graph.compile(checkpointer=checkpointer)
# missing thread_id
with pytest.raises(ValueError, match="thread_id"):
tool_two.invoke({"my_key": "value", "market": "DE"})
# flow: interrupt -> resume with answer
thread2 = {"configurable": {"thread_id": "2"}}
# stop when about to enter node
assert [
c for c in tool_two.stream({"my_key": "value ⛰️", "market": "DE"}, thread2)
] == [
{
"__interrupt__": (
Interrupt(
value="Just because...",
resumable=True,
ns=[AnyStr("tool_two:"), AnyStr("do:")],
),
)
},
]
# resume with answer
assert [c for c in tool_two.stream(Command(resume=" my answer"), thread2)] == [
{"tool_two": {"my_key": " my answer", "market": "DE"}},
]
# flow: interrupt -> clear tasks
thread1 = {"configurable": {"thread_id": "1"}}
# stop when about to enter node
assert tool_two.invoke({"my_key": "value ⛰️", "market": "DE"}, thread1) == {
"my_key": "value ⛰️",
"market": "DE",
}
assert [
c.metadata
for c in tool_two.checkpointer.list(
{"configurable": {"thread_id": "1", "checkpoint_ns": ""}}
)
] == [
{
"parents": {},
"source": "loop",
"step": 0,
"writes": None,
"thread_id": "1",
},
{
"parents": {},
"source": "input",
"step": -1,
"writes": {"__start__": {"my_key": "value ⛰️", "market": "DE"}},
"thread_id": "1",
},
]
assert tool_two.get_state(thread1) == StateSnapshot(
values={"my_key": "value ⛰️", "market": "DE"},
next=("tool_two",),
tasks=(
PregelTask(
AnyStr(),
"tool_two",
(PULL, "tool_two"),
interrupts=(
Interrupt(
value="Just because...",
resumable=True,
ns=[AnyStr("tool_two:"), AnyStr("do:")],
),
),
state={
"configurable": {
"thread_id": "1",
"checkpoint_ns": AnyStr("tool_two:"),
}
},
),
),
config=tool_two.checkpointer.get_tuple(thread1).config,
created_at=tool_two.checkpointer.get_tuple(thread1).checkpoint["ts"],
metadata={
"parents": {},
"source": "loop",
"step": 0,
"writes": None,
"thread_id": "1",
},
parent_config=[
*tool_two.checkpointer.list(
{"configurable": {"thread_id": "1", "checkpoint_ns": ""}}, limit=2
)
][-1].config,
)
# clear the interrupt and next tasks
tool_two.update_state(thread1, None, as_node=END)
# interrupt and next tasks are cleared
assert tool_two.get_state(thread1) == StateSnapshot(
values={"my_key": "value ⛰️", "market": "DE"},
next=(),
tasks=(),
config=tool_two.checkpointer.get_tuple(thread1).config,
created_at=tool_two.checkpointer.get_tuple(thread1).checkpoint["ts"],
metadata={
"parents": {},
"source": "update",
"step": 1,
"writes": {},
"thread_id": "1",
},
parent_config=[
*tool_two.checkpointer.list(
{"configurable": {"thread_id": "1", "checkpoint_ns": ""}}, limit=2
)
][-1].config,
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_start_branch_then(
snapshot: SnapshotAssertion, request: pytest.FixtureRequest, checkpointer_name: str
@@ -14648,35 +14471,3 @@ def test_parent_command(request: pytest.FixtureRequest, checkpointer_name: str)
},
tasks=(),
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_interrupt_subgraph(request: pytest.FixtureRequest, checkpointer_name: str):
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
class State(TypedDict):
baz: str
def foo(state):
return {"baz": "foo"}
def bar(state):
value = interrupt("Please provide baz value:")
return {"baz": value}
child_builder = StateGraph(State)
child_builder.add_node(bar)
child_builder.add_edge(START, "bar")
builder = StateGraph(State)
builder.add_node(foo)
builder.add_node("bar", child_builder.compile())
builder.add_edge(START, "foo")
builder.add_edge("foo", "bar")
graph = builder.compile(checkpointer=checkpointer)
thread1 = {"configurable": {"thread_id": "1"}}
# First run, interrupted at bar
assert graph.invoke({"baz": ""}, thread1)
# Resume with answer
assert graph.invoke(Command(resume="bar"), thread1)
-219
View File
@@ -429,189 +429,6 @@ async def test_dynamic_interrupt(checkpointer_name: str) -> None:
)
@pytest.mark.skipif(
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_dynamic_interrupt_subgraph(checkpointer_name: str) -> None:
class SubgraphState(TypedDict):
my_key: str
market: str
tool_two_node_count = 0
def tool_two_node(s: SubgraphState) -> SubgraphState:
nonlocal tool_two_node_count
tool_two_node_count += 1
if s["market"] == "DE":
answer = interrupt("Just because...")
else:
answer = " all good"
return {"my_key": answer}
subgraph = StateGraph(SubgraphState)
subgraph.add_node("do", tool_two_node, retry=RetryPolicy())
subgraph.add_edge(START, "do")
class State(TypedDict):
my_key: Annotated[str, operator.add]
market: str
tool_two_graph = StateGraph(State)
tool_two_graph.add_node("tool_two", subgraph.compile())
tool_two_graph.add_edge(START, "tool_two")
tool_two = tool_two_graph.compile()
tracer = FakeTracer()
assert await tool_two.ainvoke(
{"my_key": "value", "market": "DE"}, {"callbacks": [tracer]}
) == {
"my_key": "value",
"market": "DE",
}
assert tool_two_node_count == 1, "interrupts aren't retried"
assert len(tracer.runs) == 1
run = tracer.runs[0]
assert run.end_time is not None
assert run.error is None
assert run.outputs == {"market": "DE", "my_key": "value"}
assert await tool_two.ainvoke({"my_key": "value", "market": "US"}) == {
"my_key": "value all good",
"market": "US",
}
async with awith_checkpointer(checkpointer_name) as checkpointer:
tool_two = tool_two_graph.compile(checkpointer=checkpointer)
# missing thread_id
with pytest.raises(ValueError, match="thread_id"):
await tool_two.ainvoke({"my_key": "value", "market": "DE"})
# flow: interrupt -> resume with answer
thread2 = {"configurable": {"thread_id": "2"}}
# stop when about to enter node
assert [
c
async for c in tool_two.astream(
{"my_key": "value ⛰️", "market": "DE"}, thread2
)
] == [
{
"__interrupt__": (
Interrupt(
value="Just because...",
resumable=True,
ns=[AnyStr("tool_two:"), AnyStr("do:")],
),
)
},
]
# resume with answer
assert [
c async for c in tool_two.astream(Command(resume=" my answer"), thread2)
] == [
{"tool_two": {"my_key": " my answer", "market": "DE"}},
]
# flow: interrupt -> clear
thread1 = {"configurable": {"thread_id": "1"}}
thread1root = {"configurable": {"thread_id": "1", "checkpoint_ns": ""}}
# stop when about to enter node
assert [
c
async for c in tool_two.astream(
{"my_key": "value ⛰️", "market": "DE"}, thread1
)
] == [
{
"__interrupt__": (
Interrupt(
value="Just because...",
resumable=True,
ns=[AnyStr("tool_two:"), AnyStr("do:")],
),
)
},
]
assert [c.metadata async for c in tool_two.checkpointer.alist(thread1root)] == [
{
"parents": {},
"source": "loop",
"step": 0,
"writes": None,
"thread_id": "1",
},
{
"parents": {},
"source": "input",
"step": -1,
"writes": {"__start__": {"my_key": "value ⛰️", "market": "DE"}},
"thread_id": "1",
},
]
tup = await tool_two.checkpointer.aget_tuple(thread1)
assert await tool_two.aget_state(thread1) == StateSnapshot(
values={"my_key": "value ⛰️", "market": "DE"},
next=("tool_two",),
tasks=(
PregelTask(
AnyStr(),
"tool_two",
(PULL, "tool_two"),
interrupts=(
Interrupt(
value="Just because...",
resumable=True,
ns=[AnyStr("tool_two:"), AnyStr("do:")],
),
),
state={
"configurable": {
"thread_id": "1",
"checkpoint_ns": AnyStr("tool_two:"),
}
},
),
),
config=tup.config,
created_at=tup.checkpoint["ts"],
metadata={
"parents": {},
"source": "loop",
"step": 0,
"writes": None,
"thread_id": "1",
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread1root, limit=2)
][-1].config,
)
# clear the interrupt and next tasks
await tool_two.aupdate_state(thread1, None, as_node=END)
# interrupt is cleared, as well as the next tasks
tup = await tool_two.checkpointer.aget_tuple(thread1)
assert await tool_two.aget_state(thread1) == StateSnapshot(
values={"my_key": "value ⛰️", "market": "DE"},
next=(),
tasks=(),
config=tup.config,
created_at=tup.checkpoint["ts"],
metadata={
"parents": {},
"source": "update",
"step": 1,
"writes": {},
"thread_id": "1",
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread1root, limit=2)
][-1].config,
)
@pytest.mark.skipif(not FF_SEND_V2, reason="send v2 is not enabled")
@pytest.mark.skipif(
sys.version_info < (3, 11),
@@ -12860,39 +12677,3 @@ async def test_parent_command(checkpointer_name: str) -> None:
},
tasks=(),
)
@pytest.mark.skipif(
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_interrupt_subgraph(checkpointer_name: str):
class State(TypedDict):
baz: str
def foo(state):
return {"baz": "foo"}
def bar(state):
value = interrupt("Please provide baz value:")
return {"baz": value}
child_builder = StateGraph(State)
child_builder.add_node(bar)
child_builder.add_edge(START, "bar")
builder = StateGraph(State)
builder.add_node(foo)
builder.add_node("bar", child_builder.compile())
builder.add_edge(START, "foo")
builder.add_edge("foo", "bar")
async with awith_checkpointer(checkpointer_name) as checkpointer:
graph = builder.compile(checkpointer=checkpointer)
thread1 = {"configurable": {"thread_id": "1"}}
# First run, interrupted at bar
assert await graph.ainvoke({"baz": ""}, thread1)
# Resume with answer
assert await graph.ainvoke(Command(resume="bar"), thread1)
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@langchain/langgraph-sdk",
"version": "0.0.30",
"version": "0.0.29",
"description": "Client library for interacting with the LangGraph API",
"type": "module",
"packageManager": "yarn@1.22.19",
+4 -20
View File
@@ -7,7 +7,6 @@ import {
GraphSchema,
Metadata,
Run,
RunStatus,
Thread,
ThreadState,
Cron,
@@ -945,18 +944,12 @@ export class RunsClient extends BaseClient {
* Defaults to 0.
*/
offset?: number;
/**
* Status of the run to filter by.
*/
status?: RunStatus;
},
): Promise<Run[]> {
return this.fetch<Run[]>(`/threads/${threadId}/runs`, {
params: {
limit: options?.limit ?? 10,
offset: options?.offset ?? 0,
status: options?.status ?? undefined,
},
});
}
@@ -1021,28 +1014,19 @@ export class RunsClient extends BaseClient {
*
* @param threadId The ID of the thread.
* @param runId The ID of the run.
* @param signal An optional abort signal.
* @returns An async generator yielding stream parts.
*/
async *joinStream(
threadId: string,
runId: string,
options?:
| { signal?: AbortSignal; cancelOnDisconnect?: boolean }
| AbortSignal,
signal?: AbortSignal,
): AsyncGenerator<{ event: StreamEvent; data: any }> {
const opts =
typeof options === "object" &&
options != null &&
options instanceof AbortSignal
? { signal: options }
: options;
const response = await this.asyncCaller.fetch(
...this.prepareFetchOptions(`/threads/${threadId}/runs/${runId}/stream`, {
method: "GET",
timeoutMs: null,
signal: opts?.signal,
params: { cancel_on_disconnect: opts?.cancelOnDisconnect ? "1" : "0" },
signal,
}),
);
@@ -1057,7 +1041,7 @@ export class RunsClient extends BaseClient {
async start(ctrl) {
parser = createParser((event) => {
if (
(opts?.signal && opts.signal.aborted) ||
(signal && signal.aborted) ||
(event.type === "event" && event.data === "[DONE]")
) {
ctrl.terminate();
+1 -1
View File
@@ -2,7 +2,7 @@ import type { JSONSchema7 } from "json-schema";
type Optional<T> = T | null | undefined;
export type RunStatus =
type RunStatus =
| "pending"
| "running"
| "error"
+32 -101
View File
@@ -18,7 +18,6 @@ from typing import (
Dict,
Iterator,
List,
Literal,
Optional,
Sequence,
Union,
@@ -26,6 +25,7 @@ from typing import (
)
import httpx
import httpx_sse
import orjson
from httpx._types import QueryParamTypes
@@ -50,7 +50,6 @@ from langgraph_sdk.schema import (
OnConflictBehavior,
Run,
RunCreate,
RunStatus,
SearchItemsResponse,
StreamMode,
StreamPart,
@@ -60,7 +59,7 @@ from langgraph_sdk.schema import (
ThreadStatus,
ThreadUpdateStateResponse,
)
from langgraph_sdk.sse import SSEDecoder, aiter_lines_raw, iter_lines_raw
from langgraph_sdk.sse import EventSource
logger = logging.getLogger(__name__)
@@ -282,35 +281,22 @@ class HttpClient:
) -> AsyncIterator[StreamPart]:
"""Stream results using SSE."""
headers, content = await aencode_json(json)
headers["Accept"] = "text/event-stream"
headers["Cache-Control"] = "no-store"
async with self.client.stream(
method, path, headers=headers, content=content
) as res:
# check status
async with httpx_sse.aconnect_sse(
self.client, method, path, headers=headers, content=content
) as sse:
try:
res.raise_for_status()
sse.response.raise_for_status()
except httpx.HTTPStatusError as e:
body = (await res.aread()).decode()
body = (await sse.response.aread()).decode()
if sys.version_info >= (3, 11):
e.add_note(body)
else:
logger.error(f"Error from langgraph-api: {body}", exc_info=e)
raise e
# check content type
content_type = res.headers.get("content-type", "").partition(";")[0]
if "text/event-stream" not in content_type:
raise httpx.TransportError(
"Expected response header Content-Type to contain 'text/event-stream', "
f"got {content_type!r}"
async for event in EventSource(sse.response).aiter_sse():
yield StreamPart(
event.event, orjson.loads(event.data) if event.data else None
)
# parse SSE
decoder = SSEDecoder()
async for line in aiter_lines_raw(res):
sse = decoder.decode(line=line.rstrip(b"\n"))
if sse is not None:
yield sse
async def aencode_json(json: Any) -> tuple[dict[str, str], bytes]:
@@ -1698,12 +1684,7 @@ class RunsClient:
return response
async def list(
self,
thread_id: str,
*,
limit: int = 10,
offset: int = 0,
status: Optional[RunStatus] = None,
self, thread_id: str, *, limit: int = 10, offset: int = 0
) -> List[Run]:
"""List runs.
@@ -1711,7 +1692,6 @@ class RunsClient:
thread_id: The thread ID to list runs for.
limit: The maximum number of results to return.
offset: The number of results to skip.
status: The status of the run to filter by.
Returns:
List[Run]: The runs for the thread.
@@ -1725,13 +1705,9 @@ class RunsClient:
)
""" # noqa: E501
params = {
"limit": limit,
"offset": offset,
}
if status is not None:
params["status"] = status
return await self.http.get(f"/threads/{thread_id}/runs", params=params)
return await self.http.get(
f"/threads/{thread_id}/runs?limit={limit}&offset={offset}"
)
async def get(self, thread_id: str, run_id: str) -> Run:
"""Get a run.
@@ -1809,9 +1785,7 @@ class RunsClient:
""" # noqa: E501
return await self.http.get(f"/threads/{thread_id}/runs/{run_id}/join")
def join_stream(
self, thread_id: str, run_id: str, *, cancel_on_disconnect: bool = False
) -> AsyncIterator[StreamPart]:
def join_stream(self, thread_id: str, run_id: str) -> AsyncIterator[StreamPart]:
"""Stream output from a run in real-time, until the run is done.
Output is not buffered, so any output produced before this call will
not be received here.
@@ -1819,7 +1793,6 @@ class RunsClient:
Args:
thread_id: The thread ID to join.
run_id: The run ID to join.
cancel_on_disconnect: Whether to cancel the run when the stream is disconnected.
Returns:
None
@@ -1832,11 +1805,7 @@ class RunsClient:
)
""" # noqa: E501
return self.http.stream(
f"/threads/{thread_id}/runs/{run_id}/stream",
"GET",
params={"cancel_on_disconnect": cancel_on_disconnect},
)
return self.http.stream(f"/threads/{thread_id}/runs/{run_id}/stream", "GET")
async def delete(self, thread_id: str, run_id: str) -> None:
"""Delete a run.
@@ -1978,7 +1947,7 @@ class CronClient:
Example Usage:
cron_run = client.crons.create(
cron_run = await client.crons.create(
assistant_id="agent",
schedule="27 15 * * *",
input={"messages": [{"role": "user", "content": "hello!"}]},
@@ -2102,12 +2071,7 @@ class StoreClient:
self.http = http
async def put_item(
self,
namespace: Sequence[str],
/,
key: str,
value: dict[str, Any],
index: Optional[Union[Literal[False], list[str]]] = None,
self, namespace: Sequence[str], /, key: str, value: dict[str, Any]
) -> None:
"""Store or update an item.
@@ -2115,7 +2079,6 @@ class StoreClient:
namespace: A list of strings representing the namespace path.
key: The unique identifier for the item within the namespace.
value: A dictionary containing the item's data.
index: Controls search indexing - None (use defaults), False (disable), or list of field paths to index.
Returns:
None
@@ -2133,7 +2096,11 @@ class StoreClient:
raise ValueError(
f"Invalid namespace label '{label}'. Namespace labels cannot contain periods ('.')."
)
payload = {"namespace": namespace, "key": key, "value": value, "index": index}
payload = {
"namespace": namespace,
"key": key,
"value": value,
}
await self.http.put("/store/items", json=payload)
async def get_item(self, namespace: Sequence[str], /, key: str) -> Item:
@@ -2201,7 +2168,6 @@ class StoreClient:
filter: Optional[dict[str, Any]] = None,
limit: int = 10,
offset: int = 0,
query: Optional[str] = None,
) -> SearchItemsResponse:
"""Search for items within a namespace prefix.
@@ -2210,7 +2176,6 @@ class StoreClient:
filter: Optional dictionary of key-value pairs to filter results.
limit: Maximum number of items to return (default is 10).
offset: Number of items to skip before returning results (default is 0).
query: Optional query for natural language search.
Returns:
List[Item]: A list of items matching the search criteria.
@@ -2248,7 +2213,6 @@ class StoreClient:
"filter": filter,
"limit": limit,
"offset": offset,
"query": query,
}
return await self.http.post("/store/items/search", json=_provided_vals(payload))
@@ -2451,30 +2415,22 @@ class SyncHttpClient:
) -> Iterator[StreamPart]:
"""Stream the results of a request using SSE."""
headers, content = encode_json(json)
with self.client.stream(method, path, headers=headers, content=content) as res:
# check status
with httpx_sse.connect_sse(
self.client, method, path, headers=headers, content=content
) as sse:
try:
res.raise_for_status()
sse.response.raise_for_status()
except httpx.HTTPStatusError as e:
body = (res.read()).decode()
body = sse.response.read().decode()
if sys.version_info >= (3, 11):
e.add_note(body)
else:
logger.error(f"Error from langgraph-api: {body}", exc_info=e)
raise e
# check content type
content_type = res.headers.get("content-type", "").partition(";")[0]
if "text/event-stream" not in content_type:
raise httpx.TransportError(
"Expected response header Content-Type to contain 'text/event-stream', "
f"got {content_type!r}"
for event in EventSource(sse.response).iter_sse():
yield StreamPart(
event.event, orjson.loads(event.data) if event.data else None
)
# parse SSE
decoder = SSEDecoder()
for line in iter_lines_raw(res):
sse = decoder.decode(line.rstrip(b"\n"))
if sse is not None:
yield sse
def encode_json(json: Any) -> tuple[dict[str, str], bytes]:
@@ -3342,7 +3298,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values",
stream_subgraphs: bool = False,
metadata: Optional[dict] = None,
@@ -3366,7 +3321,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values",
stream_subgraphs: bool = False,
metadata: Optional[dict] = None,
@@ -3387,7 +3341,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values",
stream_subgraphs: bool = False,
metadata: Optional[dict] = None,
@@ -3412,7 +3365,6 @@ class SyncRunsClient:
assistant_id: The assistant ID or graph name to stream from.
If using graph name, will default to first assistant created from that graph.
input: The input to the graph.
command: The command to execute.
stream_mode: The stream mode(s) to use.
stream_subgraphs: Whether to stream output from subgraphs.
metadata: Metadata to assign to the run.
@@ -3463,7 +3415,6 @@ class SyncRunsClient:
""" # noqa: E501
payload = {
"input": input,
"command": command,
"config": config,
"metadata": metadata,
"stream_mode": stream_mode,
@@ -3497,7 +3448,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values",
stream_subgraphs: bool = False,
metadata: Optional[dict] = None,
@@ -3517,7 +3467,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values",
stream_subgraphs: bool = False,
metadata: Optional[dict] = None,
@@ -3538,7 +3487,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values",
stream_subgraphs: bool = False,
metadata: Optional[dict] = None,
@@ -3561,7 +3509,6 @@ class SyncRunsClient:
assistant_id: The assistant ID or graph name to stream from.
If using graph name, will default to first assistant created from that graph.
input: The input to the graph.
command: The command to execute.
stream_mode: The stream mode(s) to use.
stream_subgraphs: Whether to stream output from subgraphs.
metadata: Metadata to assign to the run.
@@ -3648,7 +3595,6 @@ class SyncRunsClient:
""" # noqa: E501
payload = {
"input": input,
"command": command,
"stream_mode": stream_mode,
"stream_subgraphs": stream_subgraphs,
"config": config,
@@ -3686,7 +3632,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
metadata: Optional[dict] = None,
config: Optional[Config] = None,
checkpoint: Optional[Checkpoint] = None,
@@ -3707,7 +3652,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
metadata: Optional[dict] = None,
config: Optional[Config] = None,
interrupt_before: Optional[Union[All, Sequence[str]]] = None,
@@ -3725,7 +3669,6 @@ class SyncRunsClient:
assistant_id: str,
*,
input: Optional[dict] = None,
command: Optional[Command] = None,
metadata: Optional[dict] = None,
config: Optional[Config] = None,
checkpoint: Optional[Checkpoint] = None,
@@ -3747,7 +3690,6 @@ class SyncRunsClient:
assistant_id: The assistant ID or graph name to run.
If using graph name, will default to first assistant created from that graph.
input: The input to the graph.
command: The command to execute.
metadata: Metadata to assign to the run.
config: The configuration for the assistant.
checkpoint: The checkpoint to resume from.
@@ -3814,7 +3756,6 @@ class SyncRunsClient:
""" # noqa: E501
payload = {
"input": input,
"command": command,
"config": config,
"metadata": metadata,
"assistant_id": assistant_id,
@@ -4214,12 +4155,7 @@ class SyncStoreClient:
self.http = http
def put_item(
self,
namespace: Sequence[str],
/,
key: str,
value: dict[str, Any],
index: Optional[Union[Literal[False], list[str]]] = None,
self, namespace: Sequence[str], /, key: str, value: dict[str, Any]
) -> None:
"""Store or update an item.
@@ -4227,7 +4163,6 @@ class SyncStoreClient:
namespace: A list of strings representing the namespace path.
key: The unique identifier for the item within the namespace.
value: A dictionary containing the item's data.
index: Controls search indexing - None (use defaults), False (disable), or list of field paths to index.
Returns:
None
@@ -4249,7 +4184,6 @@ class SyncStoreClient:
"namespace": namespace,
"key": key,
"value": value,
"index": index,
}
self.http.put("/store/items", json=payload)
@@ -4317,7 +4251,6 @@ class SyncStoreClient:
filter: Optional[dict[str, Any]] = None,
limit: int = 10,
offset: int = 0,
query: Optional[str] = None,
) -> SearchItemsResponse:
"""Search for items within a namespace prefix.
@@ -4326,7 +4259,6 @@ class SyncStoreClient:
filter: Optional dictionary of key-value pairs to filter results.
limit: Maximum number of items to return (default is 10).
offset: Number of items to skip before returning results (default is 0).
query: Optional query for natural language search.
Returns:
List[Item]: A list of items matching the search criteria.
@@ -4364,7 +4296,6 @@ class SyncStoreClient:
"filter": filter,
"limit": limit,
"offset": offset,
"query": query,
}
return self.http.post("/store/items/search", json=_provided_vals(payload))
+3 -29
View File
@@ -1,7 +1,7 @@
"""Data models for interacting with the LangGraph API."""
from datetime import datetime
from typing import Any, Dict, Literal, NamedTuple, Optional, Sequence, TypedDict, Union
from typing import Any, Literal, NamedTuple, Optional, Sequence, TypedDict, Union
Json = Optional[dict[str, Any]]
"""Represents a JSON-like structure, which can be None or a dictionary with string keys and any values."""
@@ -176,19 +176,6 @@ class Assistant(AssistantBase):
"""The name of the assistant"""
class Interrupt(TypedDict, total=False):
"""Represents an interruption in the execution flow."""
value: Any
"""The value associated with the interrupt."""
when: Literal["during"]
"""When the interrupt occurred."""
resumable: bool
"""Whether the interrupt can be resumed."""
ns: Optional[list[str]]
"""Optional namespace for the interrupt."""
class Thread(TypedDict):
"""Represents a conversation thread."""
@@ -204,8 +191,6 @@ class Thread(TypedDict):
"""The status of the thread, one of 'idle', 'busy', 'interrupted'."""
values: Json
"""The current state of the thread."""
interrupts: Dict[str, list[Interrupt]]
"""Interrupts which were thrown in this thread"""
class ThreadTask(TypedDict):
@@ -214,7 +199,7 @@ class ThreadTask(TypedDict):
id: str
name: str
error: Optional[str]
interrupts: list[Interrupt]
interrupts: list[dict]
checkpoint: Optional[Checkpoint]
state: Optional["ThreadState"]
result: Optional[dict[str, Any]]
@@ -340,21 +325,10 @@ class ListNamespaceResponse(TypedDict):
"""A list of namespace paths, where each path is a list of strings."""
class SearchItem(Item, total=False):
"""Item with an optional relevance score from search operations.
Attributes:
score (Optional[float]): Relevance/similarity score. Included when
searching a compatible store with a natural language query.
"""
score: Optional[float]
class SearchItemsResponse(TypedDict):
"""Response structure for searching items."""
items: list[SearchItem]
items: list[Item]
"""A list of items matching the search criteria."""
+35 -77
View File
@@ -1,13 +1,11 @@
"""Adapted from httpx_sse to split lines on \n, \r, \r\n per the SSE spec."""
from typing import AsyncIterator, Iterator, Optional, Union
import io
from typing import AsyncIterator, Iterator
import httpx
import orjson
from langgraph_sdk.schema import StreamPart
BytesLike = Union[bytes, bytearray, memoryview]
import httpx_sse
import httpx_sse._decoders
class BytesLineDecoder:
@@ -19,10 +17,10 @@ class BytesLineDecoder:
"""
def __init__(self) -> None:
self.buffer = bytearray()
self.buffer = io.BytesIO()
self.trailing_cr: bool = False
def decode(self, text: bytes) -> list[BytesLike]:
def decode(self, text: bytes) -> list[bytes]:
# See https://docs.python.org/3/glossary.html#term-universal-newlines
NEWLINE_CHARS = b"\n\r"
@@ -44,93 +42,33 @@ class BytesLineDecoder:
if len(lines) == 1 and not trailing_newline:
# No new lines, buffer the input and continue.
self.buffer.extend(lines[0])
self.buffer.append(lines[0])
return []
if self.buffer:
# Include any existing buffer in the first portion of the
# splitlines result.
self.buffer.extend(lines[0])
lines = [self.buffer] + lines[1:]
self.buffer = bytearray()
lines = [self.buffer.getvalue() + lines[0]] + lines[1:]
self.buffer.truncate(0)
if not trailing_newline:
# If the last segment of splitlines is not newline terminated,
# then drop it from our output and start a new buffer.
self.buffer.extend(lines.pop())
self.buffer.write(lines.pop())
return lines
def flush(self) -> list[BytesLike]:
def flush(self) -> list[bytes]:
if not self.buffer and not self.trailing_cr:
return []
lines = [self.buffer]
self.buffer = bytearray()
lines = [self.buffer.getvalue()] if self.buffer else []
self.buffer.truncate(0)
self.trailing_cr = False
return lines
class SSEDecoder:
def __init__(self) -> None:
self._event = ""
self._data = bytearray()
self._last_event_id = ""
self._retry: Optional[int] = None
def decode(self, line: bytes) -> Optional[StreamPart]:
# See: https://html.spec.whatwg.org/multipage/server-sent-events.html#event-stream-interpretation # noqa: E501
if not line:
if (
not self._event
and not self._data
and not self._last_event_id
and self._retry is None
):
return None
sse = StreamPart(
event=self._event,
data=orjson.loads(self._data) if self._data else None,
)
# NOTE: as per the SSE spec, do not reset last_event_id.
self._event = ""
self._data = bytearray()
self._retry = None
return sse
if line.startswith(b":"):
return None
fieldname, _, value = line.partition(b":")
if value.startswith(b" "):
value = value[1:]
if fieldname == b"event":
self._event = value.decode()
elif fieldname == b"data":
self._data.extend(value)
elif fieldname == b"id":
if b"\0" in value:
pass
else:
self._last_event_id = value.decode()
elif fieldname == b"retry":
try:
self._retry = int(value)
except (TypeError, ValueError):
pass
else:
pass # Field is ignored.
return None
async def aiter_lines_raw(response: httpx.Response) -> AsyncIterator[BytesLike]:
async def aiter_lines_raw(response: httpx.Response) -> AsyncIterator[bytes]:
decoder = BytesLineDecoder()
async for chunk in response.aiter_bytes():
for line in decoder.decode(chunk):
@@ -139,10 +77,30 @@ async def aiter_lines_raw(response: httpx.Response) -> AsyncIterator[BytesLike]:
yield line
def iter_lines_raw(response: httpx.Response) -> Iterator[BytesLike]:
def iter_lines_raw(response: httpx.Response) -> Iterator[bytes]:
decoder = BytesLineDecoder()
for chunk in response.iter_bytes():
for line in decoder.decode(chunk):
yield line
for line in decoder.flush():
yield line
class EventSource(httpx_sse.EventSource):
async def aiter_sse(self) -> AsyncIterator[httpx_sse.ServerSentEvent]:
self._check_content_type()
decoder = httpx_sse._decoders.SSEDecoder()
async for line in aiter_lines_raw(self._response):
line = line.rstrip(b"\n")
sse = decoder.decode(line.decode())
if sse is not None:
yield sse
def iter_sse(self) -> Iterator[httpx_sse.ServerSentEvent]:
self._check_content_type()
decoder = httpx_sse._decoders.SSEDecoder()
for line in iter_lines_raw(self._response):
line = line.rstrip(b"\n")
sse = decoder.decode(line.decode())
if sse is not None:
yield sse
+12 -1
View File
@@ -141,6 +141,17 @@ cli = ["click (==8.*)", "pygments (==2.*)", "rich (>=10,<14)"]
http2 = ["h2 (>=3,<5)"]
socks = ["socksio (==1.*)"]
[[package]]
name = "httpx-sse"
version = "0.4.0"
description = "Consume Server-Sent Event (SSE) messages with HTTPX."
optional = false
python-versions = ">=3.8"
files = [
{file = "httpx-sse-0.4.0.tar.gz", hash = "sha256:1e81a3a3070ce322add1d3529ed42eb5f70817f45ed6ec915ab753f961139721"},
{file = "httpx_sse-0.4.0-py3-none-any.whl", hash = "sha256:f329af6eae57eaa2bdfd962b42524764af68075ea87370a2de920af5341e318f"},
]
[[package]]
name = "idna"
version = "3.7"
@@ -479,4 +490,4 @@ watchmedo = ["PyYAML (>=3.10)"]
[metadata]
lock-version = "2.0"
python-versions = "^3.9.0,<4.0"
content-hash = "1262a6148df18cc44ade00466b6e0f8305897a460eea370c8de649d8d20cd7a2"
content-hash = "832acea0ad21ce71ae74edef225a1ad6f8fb166f6bf1531d876fe80fac7495f0"
+2 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "langgraph-sdk"
version = "0.1.42"
version = "0.1.37"
description = "SDK for interacting with LangGraph API"
authors = []
license = "MIT"
@@ -11,6 +11,7 @@ packages = [{ include = "langgraph_sdk" }]
[tool.poetry.dependencies]
python = "^3.9.0,<4.0"
httpx = ">=0.25.2"
httpx-sse = ">=0.4.0"
orjson = ">=3.10.1"
[tool.poetry.group.dev.dependencies]