chore: drop Python 3.9 (and syntax) (#6289)

* `strict=False` is the default, pyupgrade to min version 3.10 adds this
to be explicit w/ behavior
This commit is contained in:
Sydney Runkle
2025-10-16 20:17:46 -04:00
committed by GitHub
parent 06f9142419
commit 2d3121a17c
113 changed files with 730 additions and 1489 deletions
@@ -94,6 +94,7 @@ class PostgresSaver(BasePostgresSaver):
for v, migration in zip(
range(version + 1, len(self.MIGRATIONS)),
self.MIGRATIONS[version + 1 :],
strict=False,
):
cur.execute(migration)
cur.execute(f"INSERT INTO checkpoint_migrations (v) VALUES ({v})")
@@ -2,13 +2,12 @@
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import Union
from psycopg import AsyncConnection
from psycopg.rows import DictRow
from psycopg_pool import AsyncConnectionPool
Conn = Union[AsyncConnection[DictRow], AsyncConnectionPool[AsyncConnection[DictRow]]]
Conn = AsyncConnection[DictRow] | AsyncConnectionPool[AsyncConnection[DictRow]]
@asynccontextmanager
@@ -2,13 +2,12 @@
from collections.abc import Iterator
from contextlib import contextmanager
from typing import Union
from psycopg import Connection
from psycopg.rows import DictRow
from psycopg_pool import ConnectionPool
Conn = Union[Connection[DictRow], ConnectionPool[Connection[DictRow]]]
Conn = Connection[DictRow] | ConnectionPool[Connection[DictRow]]
@contextmanager
@@ -99,6 +99,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
for v, migration in zip(
range(version + 1, len(self.MIGRATIONS)),
self.MIGRATIONS[version + 1 :],
strict=False,
):
await cur.execute(migration)
await cur.execute(f"INSERT INTO checkpoint_migrations (v) VALUES ({v})")
@@ -4,7 +4,7 @@ import random
import warnings
from collections.abc import Sequence
from importlib.metadata import version as get_version
from typing import Any, Optional, cast
from typing import Any, cast
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
@@ -16,7 +16,7 @@ from langgraph.checkpoint.base import (
from langgraph.checkpoint.serde.types import TASKS
from psycopg.types.json import Jsonb
MetadataInput = Optional[dict[str, Any]]
MetadataInput = dict[str, Any] | None
try:
major, minor = get_version("langgraph").split(".")[:2]
@@ -3,7 +3,7 @@ import threading
import warnings
from collections.abc import AsyncIterator, Iterator, Sequence
from contextlib import asynccontextmanager, contextmanager
from typing import Any, Optional
from typing import Any
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
@@ -151,7 +151,7 @@ def _dump_blobs(
checkpoint_ns: str,
values: dict[str, Any],
versions: ChannelVersions,
) -> list[tuple[str, str, str, str, Optional[bytes]]]:
) -> list[tuple[str, str, str, str, bytes | None]]:
if not versions:
return []
@@ -186,8 +186,8 @@ class ShallowPostgresSaver(BasePostgresSaver):
def __init__(
self,
conn: _internal.Conn,
pipe: Optional[Pipeline] = None,
serde: Optional[SerializerProtocol] = None,
pipe: Pipeline | None = None,
serde: SerializerProtocol | None = None,
) -> None:
warnings.warn(
"ShallowPostgresSaver is deprecated as of version 2.0.20 and will be removed in 3.0.0. "
@@ -249,6 +249,7 @@ class ShallowPostgresSaver(BasePostgresSaver):
for v, migration in zip(
range(version + 1, len(self.MIGRATIONS)),
self.MIGRATIONS[version + 1 :],
strict=False,
):
cur.execute(migration)
cur.execute(f"INSERT INTO checkpoint_migrations (v) VALUES ({v})")
@@ -257,11 +258,11 @@ class ShallowPostgresSaver(BasePostgresSaver):
def list(
self,
config: Optional[RunnableConfig],
config: RunnableConfig | None,
*,
filter: Optional[dict[str, Any]] = None,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
filter: dict[str, Any] | None = None,
before: RunnableConfig | None = None,
limit: int | None = None,
) -> Iterator[CheckpointTuple]:
"""List checkpoints from the database.
@@ -299,7 +300,7 @@ class ShallowPostgresSaver(BasePostgresSaver):
pending_writes=self._load_writes(value["pending_writes"]),
)
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
"""Get a checkpoint tuple from the database.
This method retrieves a checkpoint tuple from the Postgres database based on the
@@ -542,8 +543,8 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
def __init__(
self,
conn: _ainternal.Conn,
pipe: Optional[AsyncPipeline] = None,
serde: Optional[SerializerProtocol] = None,
pipe: AsyncPipeline | None = None,
serde: SerializerProtocol | None = None,
) -> None:
warnings.warn(
"AsyncShallowPostgresSaver is deprecated as of version 2.0.20 and will be removed in 3.0.0. "
@@ -570,7 +571,7 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
conn_string: str,
*,
pipeline: bool = False,
serde: Optional[SerializerProtocol] = None,
serde: SerializerProtocol | None = None,
) -> AsyncIterator["AsyncShallowPostgresSaver"]:
"""Create a new AsyncShallowPostgresSaver instance from a connection string.
@@ -610,6 +611,7 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
for v, migration in zip(
range(version + 1, len(self.MIGRATIONS)),
self.MIGRATIONS[version + 1 :],
strict=False,
):
await cur.execute(migration)
await cur.execute(f"INSERT INTO checkpoint_migrations (v) VALUES ({v})")
@@ -618,11 +620,11 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
async def alist(
self,
config: Optional[RunnableConfig],
config: RunnableConfig | None,
*,
filter: Optional[dict[str, Any]] = None,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
filter: dict[str, Any] | None = None,
before: RunnableConfig | None = None,
limit: int | None = None,
) -> AsyncIterator[CheckpointTuple]:
"""List checkpoints from the database asynchronously.
@@ -662,7 +664,7 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
),
)
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
"""Get a checkpoint tuple from the database asynchronously.
This method retrieves a checkpoint tuple from the Postgres database based on the
@@ -861,11 +863,11 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
def list(
self,
config: Optional[RunnableConfig],
config: RunnableConfig | None,
*,
filter: Optional[dict[str, Any]] = None,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
filter: dict[str, Any] | None = None,
before: RunnableConfig | None = None,
limit: int | None = None,
) -> Iterator[CheckpointTuple]:
"""List checkpoints from the database.
@@ -883,7 +885,7 @@ class AsyncShallowPostgresSaver(BasePostgresSaver):
except StopAsyncIteration:
break
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
"""Get a checkpoint tuple from the database.
This method retrieves a checkpoint tuple from the Postgres database based on the
@@ -2,10 +2,10 @@ from __future__ import annotations
import asyncio
import logging
from collections.abc import AsyncIterator, Iterable, Sequence
from collections.abc import AsyncIterator, Callable, Iterable, Sequence
from contextlib import asynccontextmanager
from types import TracebackType
from typing import Any, Callable, cast
from typing import Any, cast
import orjson
from langgraph.store.base import (
@@ -465,7 +465,9 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
query,
[
p
for (ns, k, pathname, _), vector in zip(txt_params, vectors)
for (ns, k, pathname, _), vector in zip(
txt_params, vectors, strict=False
)
for p in (ns, k, pathname, vector)
],
)
@@ -486,13 +488,13 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
vectors = await self.embeddings.aembed_documents(
[query for _, query in embedding_requests]
)
for (idx, _), vector in zip(embedding_requests, vectors):
for (idx, _), vector in zip(embedding_requests, vectors, strict=False):
_paramslist = queries[idx][1]
for i in range(len(_paramslist)):
if _paramslist[i] is PLACEHOLDER:
_paramslist[i] = vector
for (idx, _), (query, params) in zip(search_ops, queries):
for (idx, _), (query, params) in zip(search_ops, queries, strict=False):
await cur.execute(query, params)
rows = cast(list[Row], await cur.fetchall())
items = [
@@ -510,7 +512,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
cur: AsyncCursor[DictRow],
) -> None:
queries = self._get_batch_list_namespaces_queries(list_ops)
for (query, params), (idx, _) in zip(queries, list_ops):
for (query, params), (idx, _) in zip(queries, list_ops, strict=False):
await cur.execute(query, params)
rows = cast(list[dict], await cur.fetchall())
namespaces = [_decode_ns_bytes(row["truncated_prefix"]) for row in rows]
@@ -6,18 +6,16 @@ import json
import logging
import threading
from collections import defaultdict
from collections.abc import Iterable, Iterator, Sequence
from collections.abc import Callable, Iterable, Iterator, Sequence
from contextlib import contextmanager
from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
Callable,
Generic,
Literal,
NamedTuple,
TypeVar,
Union,
cast,
)
@@ -141,7 +139,7 @@ CREATE INDEX CONCURRENTLY IF NOT EXISTS store_vectors_embedding_idx ON store_vec
]
C = TypeVar("C", bound=Union[_pg_internal.Conn, _ainternal.Conn])
C = TypeVar("C", bound=_pg_internal.Conn | _ainternal.Conn)
class PoolConfig(TypedDict, total=False):
@@ -255,7 +253,7 @@ class BasePostgresStore(Generic[C]):
results = []
for namespace, items in namespace_groups.items():
_, keys = zip(*items)
_, keys = zip(*items, strict=False)
this_refresh_ttls = refresh_ttls[namespace]
query = """
@@ -1014,7 +1012,9 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
query,
[
p
for (ns, k, pathname, _), vector in zip(txt_params, vectors)
for (ns, k, pathname, _), vector in zip(
txt_params, vectors, strict=False
)
for p in (ns, k, pathname, vector)
],
)
@@ -1035,13 +1035,15 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
embeddings = self.embeddings.embed_documents(
[query for _, query in embedding_requests]
)
for (idx, _), embedding in zip(embedding_requests, embeddings):
for (idx, _), embedding in zip(
embedding_requests, embeddings, strict=False
):
_paramslist = queries[idx][1]
for i in range(len(_paramslist)):
if _paramslist[i] is PLACEHOLDER:
_paramslist[i] = embedding
for (idx, _), (query, params) in zip(search_ops, queries):
for (idx, _), (query, params) in zip(search_ops, queries, strict=False):
cur.execute(query, params)
rows = cast(list[Row], cur.fetchall())
results[idx] = [
@@ -1058,7 +1060,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
cur: Cursor[DictRow],
) -> None:
for (query, params), (idx, _) in zip(
self._get_batch_list_namespaces_queries(list_ops), list_ops
self._get_batch_list_namespaces_queries(list_ops), list_ops, strict=False
):
cur.execute(query, params)
results[idx] = [_decode_ns_bytes(row["truncated_prefix"]) for row in cur]