From 5e7566f4a37d663cf666240f60fd0e1d61a43446 Mon Sep 17 00:00:00 2001 From: Sydney Runkle <54324534+sydney-runkle@users.noreply.github.com> Date: Wed, 4 Jun 2025 21:50:16 -0400 Subject: [PATCH] lint: use `pep 604` union syntax and `pep 585` generic syntax (#4963) * new union syntax * fix test * second round of conversions by injecting future annotations * format + add top level makefile --- Makefile | 58 + .../langgraph/checkpoint/postgres/__init__.py | 20 +- .../langgraph/checkpoint/postgres/aio.py | 32 +- .../langgraph/checkpoint/postgres/base.py | 10 +- .../langgraph/store/postgres/aio.py | 36 +- .../langgraph/store/postgres/base.py | 59 +- libs/checkpoint-postgres/pyproject.toml | 2 +- .../tests/checkpoint_utils.py | 10 +- .../tests/test_async_store.py | 6 +- libs/checkpoint-postgres/tests/test_store.py | 5 +- .../langgraph/checkpoint/sqlite/__init__.py | 30 +- .../langgraph/checkpoint/sqlite/aio.py | 30 +- .../langgraph/checkpoint/sqlite/utils.py | 10 +- .../langgraph/store/sqlite/aio.py | 33 +- .../langgraph/store/sqlite/base.py | 49 +- libs/checkpoint-sqlite/pyproject.toml | 2 +- .../tests/checkpoint_utils.py | 10 +- .../langgraph/checkpoint/base/__init__.py | 38 +- .../langgraph/checkpoint/base/id.py | 17 +- .../langgraph/checkpoint/memory/__init__.py | 56 +- .../langgraph/checkpoint/serde/jsonplus.py | 18 +- .../langgraph/store/base/__init__.py | 77 +- libs/checkpoint/langgraph/store/base/batch.py | 48 +- libs/checkpoint/langgraph/store/base/embed.py | 14 +- .../langgraph/store/memory/__init__.py | 10 +- libs/checkpoint/pyproject.toml | 2 +- libs/checkpoint/tests/checkpoint_utils.py | 10 +- libs/langgraph/langgraph/channels/topic.py | 6 +- libs/langgraph/langgraph/func/__init__.py | 55 +- libs/langgraph/langgraph/graph/branch.py | 50 +- libs/langgraph/langgraph/graph/message.py | 13 +- libs/langgraph/langgraph/graph/state.py | 155 ++- libs/langgraph/langgraph/graph/ui.py | 16 +- libs/langgraph/langgraph/pregel/__init__.py | 8 +- libs/langgraph/langgraph/pregel/algo.py | 93 +- libs/langgraph/langgraph/pregel/call.py | 16 +- libs/langgraph/langgraph/pregel/checkpoint.py | 9 +- libs/langgraph/langgraph/pregel/debug.py | 33 +- libs/langgraph/langgraph/pregel/draw.py | 26 +- libs/langgraph/langgraph/pregel/executor.py | 25 +- libs/langgraph/langgraph/pregel/io.py | 22 +- libs/langgraph/langgraph/pregel/loop.py | 178 +-- libs/langgraph/langgraph/pregel/messages.py | 32 +- libs/langgraph/langgraph/pregel/protocol.py | 68 +- libs/langgraph/langgraph/pregel/remote.py | 130 +- libs/langgraph/langgraph/pregel/retry.py | 17 +- libs/langgraph/langgraph/pregel/runner.py | 86 +- libs/langgraph/langgraph/pregel/utils.py | 6 +- libs/langgraph/langgraph/pregel/validate.py | 16 +- libs/langgraph/langgraph/pregel/write.py | 10 +- libs/langgraph/langgraph/types.py | 49 +- libs/langgraph/langgraph/utils/config.py | 26 +- libs/langgraph/langgraph/utils/fields.py | 6 +- libs/langgraph/langgraph/utils/future.py | 14 +- libs/langgraph/langgraph/utils/pydantic.py | 14 +- libs/langgraph/langgraph/utils/queue.py | 4 +- libs/langgraph/langgraph/utils/runnable.py | 42 +- libs/langgraph/pyproject.toml | 5 +- libs/langgraph/tests/test_runnable.py | 6 +- libs/sdk-py/langgraph_sdk/auth/__init__.py | 94 +- libs/sdk-py/langgraph_sdk/auth/exceptions.py | 8 +- libs/sdk-py/langgraph_sdk/auth/types.py | 78 +- libs/sdk-py/langgraph_sdk/client.py | 1076 ++++++++--------- libs/sdk-py/langgraph_sdk/schema.py | 65 +- libs/sdk-py/langgraph_sdk/sse.py | 8 +- libs/sdk-py/pyproject.toml | 2 +- 66 files changed, 1671 insertions(+), 1588 deletions(-) create mode 100644 Makefile diff --git a/Makefile b/Makefile new file mode 100644 index 000000000..aa4eb9945 --- /dev/null +++ b/Makefile @@ -0,0 +1,58 @@ +# Define the directories containing projects +LIBS_DIRS := $(wildcard libs/*) + +# Default target +.PHONY: all +all: lint format lock test + +# Install dependencies for all projects +.PHONY: install +install: + @echo "Creating virtual environment..." + @uv venv + @for dir in $(LIBS_DIRS); do \ + if [ -f $$dir/pyproject.toml ]; then \ + echo "Installing dependencies for $$dir"; \ + uv pip install -e $$dir; \ + fi; \ + done + +# Lint all projects +.PHONY: lint +lint: + @for dir in $(LIBS_DIRS); do \ + if [ -f $$dir/Makefile ]; then \ + echo "Running lint in $$dir"; \ + $(MAKE) -C $$dir lint; \ + fi; \ + done + +# Format all projects +.PHONY: format +format: + @for dir in $(LIBS_DIRS); do \ + if [ -f $$dir/Makefile ]; then \ + echo "Running format in $$dir"; \ + $(MAKE) -C $$dir format; \ + fi; \ + done + +# Lock all projects +.PHONY: lock +lock: + @for dir in $(LIBS_DIRS); do \ + if [ -f $$dir/Makefile ]; then \ + echo "Running lock in $$dir"; \ + (cd $$dir && uv lock); \ + fi; \ + done + +# Test all projects +.PHONY: test +test: + @for dir in $(LIBS_DIRS); do \ + if [ -f $$dir/Makefile ]; then \ + echo "Running test in $$dir"; \ + $(MAKE) -C $$dir test; \ + fi; \ + done \ No newline at end of file diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index be6f3c137..d8e6947d3 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -1,8 +1,10 @@ +from __future__ import annotations + import threading from collections import defaultdict from collections.abc import Iterator, Sequence from contextlib import contextmanager -from typing import Any, Optional +from typing import Any from langchain_core.runnables import RunnableConfig from psycopg import Capabilities, Connection, Cursor, Pipeline @@ -34,8 +36,8 @@ class PostgresSaver(BasePostgresSaver): def __init__( self, conn: _internal.Conn, - pipe: Optional[Pipeline] = None, - serde: Optional[SerializerProtocol] = None, + pipe: Pipeline | None = None, + serde: SerializerProtocol | None = None, ) -> None: super().__init__(serde=serde) if isinstance(conn, ConnectionPool) and pipe is not None: @@ -52,7 +54,7 @@ class PostgresSaver(BasePostgresSaver): @contextmanager def from_conn_string( cls, conn_string: str, *, pipeline: bool = False - ) -> Iterator["PostgresSaver"]: + ) -> Iterator[PostgresSaver]: """Create a new PostgresSaver instance from a connection string. Args: @@ -99,11 +101,11 @@ class PostgresSaver(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. @@ -200,7 +202,7 @@ class PostgresSaver(BasePostgresSaver): 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 diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py index 8253d4d8f..13ee08044 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py @@ -1,8 +1,10 @@ +from __future__ import annotations + import asyncio from collections import defaultdict from collections.abc import AsyncIterator, Iterator, Sequence from contextlib import asynccontextmanager -from typing import Any, Optional +from typing import Any from langchain_core.runnables import RunnableConfig from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline, Capabilities @@ -34,8 +36,8 @@ class AsyncPostgresSaver(BasePostgresSaver): def __init__( self, conn: _ainternal.Conn, - pipe: Optional[AsyncPipeline] = None, - serde: Optional[SerializerProtocol] = None, + pipe: AsyncPipeline | None = None, + serde: SerializerProtocol | None = None, ) -> None: super().__init__(serde=serde) if isinstance(conn, AsyncConnectionPool) and pipe is not None: @@ -56,8 +58,8 @@ class AsyncPostgresSaver(BasePostgresSaver): conn_string: str, *, pipeline: bool = False, - serde: Optional[SerializerProtocol] = None, - ) -> AsyncIterator["AsyncPostgresSaver"]: + serde: SerializerProtocol | None = None, + ) -> AsyncIterator[AsyncPostgresSaver]: """Create a new AsyncPostgresSaver instance from a connection string. Args: @@ -104,11 +106,11 @@ class AsyncPostgresSaver(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. @@ -187,7 +189,7 @@ class AsyncPostgresSaver(BasePostgresSaver): await asyncio.to_thread(self._load_writes, value["pending_writes"]), ) - 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 @@ -424,11 +426,11 @@ class AsyncPostgresSaver(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. @@ -466,7 +468,7 @@ class AsyncPostgresSaver(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 diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py index e4b9b2e65..8c502a7fe 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import random from collections.abc import Sequence from typing import Any, Optional, cast @@ -186,7 +188,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): checkpoint_ns: str, values: dict[str, Any], versions: ChannelVersions, - ) -> list[tuple[str, str, str, str, str, Optional[bytes]]]: + ) -> list[tuple[str, str, str, str, str, bytes | None]]: if not versions: return [] @@ -244,7 +246,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): for idx, (channel, value) in enumerate(writes) ] - def get_next_version(self, current: Optional[str]) -> str: + def get_next_version(self, current: str | None) -> str: if current is None: current_v = 0 elif isinstance(current, int): @@ -257,9 +259,9 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): def _search_where( self, - config: Optional[RunnableConfig], + config: RunnableConfig | None, filter: MetadataInput, - before: Optional[RunnableConfig] = None, + before: RunnableConfig | None = None, ) -> tuple[str, list[Any]]: """Return WHERE clause predicates for alist() given config, filter, before. diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py index 22d15ee5f..65b316d13 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py @@ -1,9 +1,11 @@ +from __future__ import annotations + import asyncio import logging from collections.abc import AsyncIterator, Iterable, Sequence from contextlib import asynccontextmanager from types import TracebackType -from typing import Any, Callable, Optional, Union, cast +from typing import Any, Callable, cast import orjson from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline, Capabilities @@ -132,12 +134,10 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con self, conn: _ainternal.Conn, *, - pipe: Optional[AsyncPipeline] = None, - deserializer: Optional[ - Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]] - ] = None, - index: Optional[PostgresIndexConfig] = None, - ttl: Optional[TTLConfig] = None, + pipe: AsyncPipeline | None = None, + deserializer: Callable[[bytes | orjson.Fragment], dict[str, Any]] | None = None, + index: PostgresIndexConfig | None = None, + ttl: TTLConfig | None = None, ) -> None: if isinstance(conn, AsyncConnectionPool) and pipe is not None: raise ValueError( @@ -157,7 +157,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con self.embeddings = None self.ttl_config = ttl - self._ttl_sweeper_task: Optional[asyncio.Task[None]] = None + self._ttl_sweeper_task: asyncio.Task[None] | None = None self._ttl_stop_event = asyncio.Event() async def abatch(self, ops: Iterable[Op]) -> list[Result]: @@ -180,10 +180,10 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con conn_string: str, *, pipeline: bool = False, - pool_config: Optional[PoolConfig] = None, - index: Optional[PostgresIndexConfig] = None, - ttl: Optional[TTLConfig] = None, - ) -> AsyncIterator["AsyncPostgresStore"]: + pool_config: PoolConfig | None = None, + index: PostgresIndexConfig | None = None, + ttl: TTLConfig | None = None, + ) -> AsyncIterator[AsyncPostgresStore]: """Create a new AsyncPostgresStore instance from a connection string. Args: @@ -289,7 +289,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con return deleted_count async def start_ttl_sweeper( - self, sweep_interval_minutes: Optional[int] = None + self, sweep_interval_minutes: int | None = None ) -> asyncio.Task[None]: """Periodically delete expired store items based on TTL. @@ -334,7 +334,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con self._ttl_sweeper_task = task return task - async def stop_ttl_sweeper(self, timeout: Optional[float] = None) -> bool: + async def stop_ttl_sweeper(self, timeout: float | None = None) -> bool: """Stop the TTL sweeper task if it's running. Args: @@ -369,14 +369,14 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con return success - async def __aenter__(self) -> "AsyncPostgresStore": + async def __aenter__(self) -> AsyncPostgresStore: return self async def __aexit__( self, - exc_type: Optional[type[BaseException]], - exc_val: Optional[BaseException], - exc_tb: Optional["TracebackType"], + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, ) -> None: # Ensure the TTL sweeper task is stopped when exiting the context if hasattr(self, "_ttl_sweeper_task") and self._ttl_sweeper_task is not None: diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/base.py b/libs/checkpoint-postgres/langgraph/store/postgres/base.py index 31c598d55..dcdb47f2c 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import concurrent.futures import json @@ -14,7 +16,6 @@ from typing import ( Generic, Literal, NamedTuple, - Optional, TypeVar, Union, cast, @@ -56,8 +57,8 @@ class Migration(NamedTuple): """A database migration with optional conditions and parameters.""" sql: str - params: Optional[dict[str, Any]] = None - condition: Optional[Callable[["BasePostgresStore"], bool]] = None + params: dict[str, Any] | None = None + condition: Callable[[BasePostgresStore], bool] | None = None MIGRATIONS: Sequence[str] = [ @@ -155,7 +156,7 @@ class PoolConfig(TypedDict, total=False): min_size: int """Minimum number of connections maintained in the pool. Defaults to 1.""" - max_size: Optional[int] + max_size: int | None """Maximum number of connections allowed in the pool. None means unlimited.""" kwargs: dict @@ -230,8 +231,8 @@ 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] + _deserializer: Callable[[bytes | orjson.Fragment], dict[str, Any]] | None + index_config: PostgresIndexConfig | None def _get_batch_GET_ops_queries( self, @@ -293,7 +294,7 @@ class BasePostgresStore(Generic[C]): put_ops: Sequence[tuple[int, PutOp]], ) -> tuple[ list[tuple[str, Sequence]], - Optional[tuple[str, Sequence[tuple[str, str, str, str]]]], + tuple[str, Sequence[tuple[str, str, str, str]]] | None, ]: dedupped_ops: dict[tuple[tuple[str, ...], str], PutOp] = {} for _, op in put_ops: @@ -320,9 +321,7 @@ class BasePostgresStore(Generic[C]): ) params = (_namespace_to_text(namespace), *keys) queries.append((query, params)) - embedding_request: Optional[tuple[str, Sequence[tuple[str, str, str, str]]]] = ( - None - ) + embedding_request: tuple[str, Sequence[tuple[str, str, str, str]]] | None = None if inserts: values = [] insertion_params = [] @@ -403,7 +402,7 @@ class BasePostgresStore(Generic[C]): self, search_ops: Sequence[tuple[int, SearchOp]], ) -> tuple[ - list[tuple[str, list[Union[None, str, list[float]]]]], # queries, params + list[tuple[str, list[None | str | list[float]]]], # queries, params list[tuple[int, str]], # idx, query_text pairs to embed ]: """ @@ -432,7 +431,7 @@ class BasePostgresStore(Generic[C]): filter_params.extend([key, orjson.dumps(value).decode("utf-8")]) ns_condition = "TRUE" - ns_param: Optional[Sequence[Union[str]]] = None + ns_param: Sequence[str] | None = None if op.namespace_prefix: ns_condition = "store.prefix LIKE %s" ns_param = (f"{_namespace_to_text(op.namespace_prefix)}%",) @@ -719,12 +718,10 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): self, conn: _pg_internal.Conn, *, - pipe: Optional[Pipeline] = None, - deserializer: Optional[ - Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]] - ] = None, - index: Optional[PostgresIndexConfig] = None, - ttl: Optional[TTLConfig] = None, + pipe: Pipeline | None = None, + deserializer: Callable[[bytes | orjson.Fragment], dict[str, Any]] | None = None, + index: PostgresIndexConfig | None = None, + ttl: TTLConfig | None = None, ) -> None: super().__init__() self._deserializer = deserializer @@ -738,7 +735,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): else: self.embeddings = None self.ttl_config = ttl - self._ttl_sweeper_thread: Optional[threading.Thread] = None + self._ttl_sweeper_thread: threading.Thread | None = None self._ttl_stop_event = threading.Event() @classmethod @@ -748,10 +745,10 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): conn_string: str, *, pipeline: bool = False, - pool_config: Optional[PoolConfig] = None, - index: Optional[PostgresIndexConfig] = None, - ttl: Optional[TTLConfig] = None, - ) -> Iterator["PostgresStore"]: + pool_config: PoolConfig | None = None, + index: PostgresIndexConfig | None = None, + ttl: TTLConfig | None = None, + ) -> Iterator[PostgresStore]: """Create a new PostgresStore instance from a connection string. Args: @@ -810,7 +807,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): return deleted_count def start_ttl_sweeper( - self, sweep_interval_minutes: Optional[int] = None + self, sweep_interval_minutes: int | None = None ) -> concurrent.futures.Future[None]: """Periodically delete expired store items based on TTL. @@ -867,7 +864,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): ) return future - def stop_ttl_sweeper(self, timeout: Optional[float] = None) -> bool: + def stop_ttl_sweeper(self, timeout: float | None = None) -> bool: """Stop the TTL sweeper thread if it's running. Args: @@ -1196,7 +1193,7 @@ def _row_to_item( namespace: tuple[str, ...], row: Row, *, - loader: Optional[Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]] = None, + loader: Callable[[bytes | orjson.Fragment], dict[str, Any]] | None = None, ) -> Item: """Convert a row from the database into an Item. @@ -1224,7 +1221,7 @@ def _row_to_search_item( namespace: tuple[str, ...], row: Row, *, - loader: Optional[Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]] = None, + loader: Callable[[bytes | orjson.Fragment], dict[str, Any]] | None = None, ) -> SearchItem: """Convert a row from the database into an Item.""" loader = loader or _json_loads @@ -1255,7 +1252,7 @@ def _group_ops(ops: Iterable[Op]) -> tuple[dict[type, list[tuple[int, Op]]], int return grouped_ops, tot -def _json_loads(content: Union[bytes, orjson.Fragment]) -> Any: +def _json_loads(content: bytes | orjson.Fragment) -> Any: if isinstance(content, orjson.Fragment): if hasattr(content, "buf"): content = content.buf @@ -1267,7 +1264,7 @@ def _json_loads(content: Union[bytes, orjson.Fragment]) -> Any: return orjson.loads(cast(bytes, content)) -def _decode_ns_bytes(namespace: Union[str, bytes, list]) -> tuple[str, ...]: +def _decode_ns_bytes(namespace: str | bytes | list) -> tuple[str, ...]: if isinstance(namespace, list): return tuple(namespace) if isinstance(namespace, bytes): @@ -1316,9 +1313,9 @@ def get_distance_operator(store: Any) -> tuple[str, str]: def _ensure_index_config( index_config: PostgresIndexConfig, -) -> tuple[Optional["Embeddings"], PostgresIndexConfig]: +) -> tuple[Embeddings | None, PostgresIndexConfig]: index_config = index_config.copy() - tokenized: list[tuple[str, Union[Literal["$"], list[str]]]] = [] + tokenized: list[tuple[str, Literal["$"] | list[str]]] = [] tot = 0 text_fields = index_config.get("fields") or ["$"] if isinstance(text_fields, str): diff --git a/libs/checkpoint-postgres/pyproject.toml b/libs/checkpoint-postgres/pyproject.toml index e2a3b81e7..0442cfcd5 100644 --- a/libs/checkpoint-postgres/pyproject.toml +++ b/libs/checkpoint-postgres/pyproject.toml @@ -56,7 +56,7 @@ lint.select = [ "B", # flake8-bugbear "I", # isort ] -lint.ignore = ["E501", "B008", "UP007", "UP006"] +lint.ignore = ["E501", "B008"] [tool.mypy] # https://mypy.readthedocs.io/en/stable/config_file.html diff --git a/libs/checkpoint-postgres/tests/checkpoint_utils.py b/libs/checkpoint-postgres/tests/checkpoint_utils.py index 881bef7f4..f38afd740 100644 --- a/libs/checkpoint-postgres/tests/checkpoint_utils.py +++ b/libs/checkpoint-postgres/tests/checkpoint_utils.py @@ -1,13 +1,15 @@ +from __future__ import annotations + from collections.abc import Mapping from datetime import datetime, timezone -from typing import Any, Optional, Protocol +from typing import Any, Protocol from langgraph.checkpoint.base import Checkpoint, EmptyChannelError from langgraph.checkpoint.base.id import uuid6 class ChannelProtocol(Protocol): - def checkpoint(self) -> Optional[Any]: ... + def checkpoint(self) -> Any | None: ... def empty_checkpoint() -> Checkpoint: @@ -23,10 +25,10 @@ def empty_checkpoint() -> Checkpoint: def create_checkpoint( checkpoint: Checkpoint, - channels: Optional[Mapping[str, ChannelProtocol]], + channels: Mapping[str, ChannelProtocol] | None, step: int, *, - id: Optional[str] = None, + id: str | None = None, ) -> Checkpoint: """Create a checkpoint for the given channels.""" ts = datetime.now(timezone.utc).isoformat() diff --git a/libs/checkpoint-postgres/tests/test_async_store.py b/libs/checkpoint-postgres/tests/test_async_store.py index 03a7c12d6..1ea7df1b7 100644 --- a/libs/checkpoint-postgres/tests/test_async_store.py +++ b/libs/checkpoint-postgres/tests/test_async_store.py @@ -1,4 +1,6 @@ # type: ignore +from __future__ import annotations + import asyncio import itertools import sys @@ -6,7 +8,7 @@ import uuid from collections.abc import AsyncIterator from concurrent.futures import ThreadPoolExecutor from contextlib import asynccontextmanager -from typing import Any, Optional +from typing import Any import pytest from langchain_core.embeddings import Embeddings @@ -353,7 +355,7 @@ async def _create_vector_store( vector_type: str, distance_type: str, fake_embeddings: CharacterEmbeddings, - text_fields: Optional[list[str]] = None, + text_fields: list[str] | None = None, ) -> AsyncIterator[AsyncPostgresStore]: """Create a store with vector search enabled.""" if sys.version_info < (3, 10): diff --git a/libs/checkpoint-postgres/tests/test_store.py b/libs/checkpoint-postgres/tests/test_store.py index 397c6054e..f7c7a1b2c 100644 --- a/libs/checkpoint-postgres/tests/test_store.py +++ b/libs/checkpoint-postgres/tests/test_store.py @@ -1,9 +1,10 @@ # type: ignore +from __future__ import annotations import re import time from contextlib import contextmanager -from typing import Any, Optional +from typing import Any from uuid import uuid4 import pytest @@ -379,7 +380,7 @@ def _create_vector_store( vector_type: str, distance_type: str, fake_embeddings: Embeddings, - text_fields: Optional[list[str]] = None, + text_fields: list[str] | None = None, enable_ttl: bool = True, ) -> PostgresStore: """Create a store with vector search enabled.""" diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py index e13cfbe29..caf9cdf3a 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py @@ -1,9 +1,11 @@ +from __future__ import annotations + import random import sqlite3 import threading from collections.abc import AsyncIterator, Iterator, Sequence from contextlib import closing, contextmanager -from typing import Any, Optional, cast +from typing import Any, cast from langchain_core.runnables import RunnableConfig @@ -76,7 +78,7 @@ class SqliteSaver(BaseCheckpointSaver[str]): self, conn: sqlite3.Connection, *, - serde: Optional[SerializerProtocol] = None, + serde: SerializerProtocol | None = None, ) -> None: super().__init__(serde=serde) self.jsonplus_serde = JsonPlusSerializer() @@ -86,7 +88,7 @@ class SqliteSaver(BaseCheckpointSaver[str]): @classmethod @contextmanager - def from_conn_string(cls, conn_string: str) -> Iterator["SqliteSaver"]: + def from_conn_string(cls, conn_string: str) -> Iterator[SqliteSaver]: """Create a new SqliteSaver instance from a connection string. Args: @@ -178,7 +180,7 @@ class SqliteSaver(BaseCheckpointSaver[str]): self.conn.commit() cur.close() - 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 SQLite database based on the @@ -286,11 +288,11 @@ class SqliteSaver(BaseCheckpointSaver[str]): 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. @@ -493,7 +495,7 @@ class SqliteSaver(BaseCheckpointSaver[str]): (str(thread_id),), ) - 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. Note: @@ -504,11 +506,11 @@ class SqliteSaver(BaseCheckpointSaver[str]): 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. @@ -534,7 +536,7 @@ class SqliteSaver(BaseCheckpointSaver[str]): """ raise NotImplementedError(_AIO_ERROR_MSG) - def get_next_version(self, current: Optional[str]) -> str: + def get_next_version(self, current: str | None) -> str: """Generate the next version ID for a channel. This method creates a new version identifier for a channel based on its current version. diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py index 77b499326..dd4a61ab9 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py @@ -1,8 +1,10 @@ +from __future__ import annotations + import asyncio import random from collections.abc import AsyncIterator, Iterator, Sequence from contextlib import asynccontextmanager -from typing import Any, Callable, Optional, TypeVar, cast +from typing import Any, Callable, TypeVar, cast import aiosqlite from langchain_core.runnables import RunnableConfig @@ -108,7 +110,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]): self, conn: aiosqlite.Connection, *, - serde: Optional[SerializerProtocol] = None, + serde: SerializerProtocol | None = None, ): super().__init__(serde=serde) self.jsonplus_serde = JsonPlusSerializer() @@ -121,7 +123,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]): @asynccontextmanager async def from_conn_string( cls, conn_string: str - ) -> AsyncIterator["AsyncSqliteSaver"]: + ) -> AsyncIterator[AsyncSqliteSaver]: """Create a new AsyncSqliteSaver instance from a connection string. Args: @@ -133,7 +135,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]): async with aiosqlite.connect(conn_string) as conn: yield cls(conn) - 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 SQLite database based on the @@ -165,11 +167,11 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]): 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 asynchronously. @@ -310,7 +312,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]): self.is_setup = True - 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 SQLite database based on the @@ -398,11 +400,11 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]): 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. @@ -589,7 +591,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]): ) await self.conn.commit() - def get_next_version(self, current: Optional[str]) -> str: + def get_next_version(self, current: str | None) -> str: """Generate the next version ID for a channel. This method creates a new version identifier for a channel based on its current version. diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/utils.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/utils.py index a6ae576b0..9ab45d289 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/utils.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/utils.py @@ -1,6 +1,8 @@ +from __future__ import annotations + import json from collections.abc import Sequence -from typing import Any, Optional +from typing import Any from langchain_core.runnables import RunnableConfig @@ -52,9 +54,9 @@ def _metadata_predicate( def search_where( - config: Optional[RunnableConfig], - filter: Optional[dict[str, Any]], - before: Optional[RunnableConfig] = None, + config: RunnableConfig | None, + filter: dict[str, Any] | None, + before: RunnableConfig | None = None, ) -> tuple[str, Sequence[Any]]: """Return WHERE clause predicates for (a)search() given metadata filter and `before` config. diff --git a/libs/checkpoint-sqlite/langgraph/store/sqlite/aio.py b/libs/checkpoint-sqlite/langgraph/store/sqlite/aio.py index 75750393c..3b19e512d 100644 --- a/libs/checkpoint-sqlite/langgraph/store/sqlite/aio.py +++ b/libs/checkpoint-sqlite/langgraph/store/sqlite/aio.py @@ -1,10 +1,12 @@ +from __future__ import annotations + import asyncio import logging from collections import defaultdict from collections.abc import AsyncIterator, Iterable, Sequence from contextlib import asynccontextmanager from types import TracebackType -from typing import Any, Callable, Optional, Union, cast +from typing import Any, Callable, cast import aiosqlite import orjson @@ -88,11 +90,10 @@ class AsyncSqliteStore(AsyncBatchedBaseStore, BaseSqliteStore): self, conn: aiosqlite.Connection, *, - deserializer: Optional[ - Callable[[Union[bytes, str, orjson.Fragment]], dict[str, Any]] - ] = None, - index: Optional[SqliteIndexConfig] = None, - ttl: Optional[TTLConfig] = None, + deserializer: Callable[[bytes | str | orjson.Fragment], dict[str, Any]] + | None = None, + index: SqliteIndexConfig | None = None, + ttl: TTLConfig | None = None, ): """Initialize the async SQLite store. @@ -114,7 +115,7 @@ class AsyncSqliteStore(AsyncBatchedBaseStore, BaseSqliteStore): else: self.embeddings = None self.ttl_config = ttl - self._ttl_sweeper_task: Optional[asyncio.Task[None]] = None + self._ttl_sweeper_task: asyncio.Task[None] | None = None self._ttl_stop_event = asyncio.Event() @classmethod @@ -123,9 +124,9 @@ class AsyncSqliteStore(AsyncBatchedBaseStore, BaseSqliteStore): cls, conn_string: str, *, - index: Optional[SqliteIndexConfig] = None, - ttl: Optional[TTLConfig] = None, - ) -> AsyncIterator["AsyncSqliteStore"]: + index: SqliteIndexConfig | None = None, + ttl: TTLConfig | None = None, + ) -> AsyncIterator[AsyncSqliteStore]: """Create a new AsyncSqliteStore instance from a connection string. Args: @@ -253,7 +254,7 @@ class AsyncSqliteStore(AsyncBatchedBaseStore, BaseSqliteStore): return deleted_count async def start_ttl_sweeper( - self, sweep_interval_minutes: Optional[int] = None + self, sweep_interval_minutes: int | None = None ) -> asyncio.Task[None]: """Periodically delete expired store items based on TTL. @@ -298,7 +299,7 @@ class AsyncSqliteStore(AsyncBatchedBaseStore, BaseSqliteStore): self._ttl_sweeper_task = task return task - async def stop_ttl_sweeper(self, timeout: Optional[float] = None) -> bool: + async def stop_ttl_sweeper(self, timeout: float | None = None) -> bool: """Stop the TTL sweeper task if it's running. Args: @@ -333,14 +334,14 @@ class AsyncSqliteStore(AsyncBatchedBaseStore, BaseSqliteStore): return success - async def __aenter__(self) -> "AsyncSqliteStore": + async def __aenter__(self) -> AsyncSqliteStore: return self async def __aexit__( self, - exc_type: Optional[type[BaseException]], - exc_val: Optional[BaseException], - exc_tb: Optional["TracebackType"], + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, ) -> None: # Ensure the TTL sweeper task is stopped when exiting the context if hasattr(self, "_ttl_sweeper_task") and self._ttl_sweeper_task is not None: diff --git a/libs/checkpoint-sqlite/langgraph/store/sqlite/base.py b/libs/checkpoint-sqlite/langgraph/store/sqlite/base.py index 033434578..5ac1ee0e0 100644 --- a/libs/checkpoint-sqlite/langgraph/store/sqlite/base.py +++ b/libs/checkpoint-sqlite/langgraph/store/sqlite/base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import concurrent.futures import datetime import logging @@ -6,7 +8,7 @@ import threading from collections import defaultdict from collections.abc import Iterable, Iterator, Sequence from contextlib import contextmanager -from typing import Any, Callable, Literal, NamedTuple, Optional, Union, cast +from typing import Any, Callable, Literal, NamedTuple, cast import orjson import sqlite_vec # type: ignore[import-untyped] @@ -105,7 +107,7 @@ def _decode_ns_text(namespace: str) -> tuple[str, ...]: return tuple(namespace.split(".")) -def _json_loads(content: Union[bytes, str, orjson.Fragment]) -> Any: +def _json_loads(content: bytes | str | orjson.Fragment) -> Any: if isinstance(content, orjson.Fragment): if hasattr(content, "buf"): content = content.buf @@ -125,9 +127,7 @@ def _row_to_item( namespace: tuple[str, ...], row: dict[str, Any], *, - loader: Optional[ - Callable[[Union[bytes, str, orjson.Fragment]], dict[str, Any]] - ] = None, + loader: Callable[[bytes | str | orjson.Fragment], dict[str, Any]] | None = None, ) -> Item: """Convert a row from the database into an Item.""" val = row["value"] @@ -149,9 +149,7 @@ def _row_to_search_item( namespace: tuple[str, ...], row: dict[str, Any], *, - loader: Optional[ - Callable[[Union[bytes, str, orjson.Fragment]], dict[str, Any]] - ] = None, + loader: Callable[[bytes | str | orjson.Fragment], dict[str, Any]] | None = None, ) -> SearchItem: """Convert a row from the database into a SearchItem.""" loader = loader or _json_loads @@ -196,8 +194,8 @@ class BaseSqliteStore: MIGRATIONS = MIGRATIONS VECTOR_MIGRATIONS = VECTOR_MIGRATIONS supports_ttl = True - index_config: Optional[SqliteIndexConfig] = None - ttl_config: Optional[TTLConfig] = None + index_config: SqliteIndexConfig | None = None + ttl_config: TTLConfig | None = None def _get_batch_GET_ops_queries( self, get_ops: Sequence[tuple[int, GetOp]] @@ -259,7 +257,7 @@ class BaseSqliteStore: self, put_ops: Sequence[tuple[int, PutOp]] ) -> tuple[ list[tuple[str, Sequence]], - Optional[tuple[str, Sequence[tuple[str, str, str, str]]]], + tuple[str, Sequence[tuple[str, str, str, str]]] | None, ]: # Last-write wins dedupped_ops: dict[tuple[tuple[str, ...], str], PutOp] = {} @@ -288,9 +286,7 @@ class BaseSqliteStore: params = (_namespace_to_text(namespace), *keys) queries.append((query, params)) - embedding_request: Optional[tuple[str, Sequence[tuple[str, str, str, str]]]] = ( - None - ) + embedding_request: tuple[str, Sequence[tuple[str, str, str, str]]] | None = None if inserts: values = [] insertion_params = [] @@ -358,7 +354,7 @@ class BaseSqliteStore: def _prepare_batch_search_queries( self, search_ops: Sequence[tuple[int, SearchOp]] ) -> tuple[ - list[tuple[str, list[Union[None, str, list[float]]]]], # queries, params + list[tuple[str, list[None | str | list[float]]]], # queries, params list[tuple[int, str]], # idx, query_text pairs to embed ]: """ @@ -785,11 +781,10 @@ class SqliteStore(BaseSqliteStore, BaseStore): self, conn: sqlite3.Connection, *, - deserializer: Optional[ - Callable[[Union[bytes, str, orjson.Fragment]], dict[str, Any]] - ] = None, - index: Optional[SqliteIndexConfig] = None, - ttl: Optional[TTLConfig] = None, + deserializer: Callable[[bytes | str | orjson.Fragment], dict[str, Any]] + | None = None, + index: SqliteIndexConfig | None = None, + ttl: TTLConfig | None = None, ): super().__init__() self._deserializer = deserializer @@ -802,7 +797,7 @@ class SqliteStore(BaseSqliteStore, BaseStore): else: self.embeddings = None self.ttl_config = ttl - self._ttl_sweeper_thread: Optional[threading.Thread] = None + self._ttl_sweeper_thread: threading.Thread | None = None self._ttl_stop_event = threading.Event() def _get_batch_GET_ops_queries( @@ -956,9 +951,9 @@ class SqliteStore(BaseSqliteStore, BaseStore): cls, conn_string: str, *, - index: Optional[SqliteIndexConfig] = None, - ttl: Optional[TTLConfig] = None, - ) -> Iterator["SqliteStore"]: + index: SqliteIndexConfig | None = None, + ttl: TTLConfig | None = None, + ) -> Iterator[SqliteStore]: """Create a new SqliteStore instance from a connection string. Args: @@ -1087,7 +1082,7 @@ class SqliteStore(BaseSqliteStore, BaseStore): return deleted_count def start_ttl_sweeper( - self, sweep_interval_minutes: Optional[int] = None + self, sweep_interval_minutes: int | None = None ) -> concurrent.futures.Future[None]: """Periodically delete expired store items based on TTL. @@ -1144,7 +1139,7 @@ class SqliteStore(BaseSqliteStore, BaseStore): ) return future - def stop_ttl_sweeper(self, timeout: Optional[float] = None) -> bool: + def stop_ttl_sweeper(self, timeout: float | None = None) -> bool: """Stop the TTL sweeper thread if it's running. Args: @@ -1396,7 +1391,7 @@ def _ensure_index_config( ) -> tuple[Any, SqliteIndexConfig]: """Process and validate index configuration.""" index_config = index_config.copy() - tokenized: list[tuple[str, Union[Literal["$"], list[str]]]] = [] + tokenized: list[tuple[str, Literal["$"] | list[str]]] = [] tot = 0 text_fields = index_config.get("text_fields") or ["$"] if isinstance(text_fields, str): diff --git a/libs/checkpoint-sqlite/pyproject.toml b/libs/checkpoint-sqlite/pyproject.toml index e7de9e46f..7f378d892 100644 --- a/libs/checkpoint-sqlite/pyproject.toml +++ b/libs/checkpoint-sqlite/pyproject.toml @@ -54,7 +54,7 @@ lint.select = [ "B", # flake8-bugbear "I", # isort ] -lint.ignore = ["E501", "B008", "UP007", "UP006"] +lint.ignore = ["E501", "B008"] [tool.pytest-watcher] now = true diff --git a/libs/checkpoint-sqlite/tests/checkpoint_utils.py b/libs/checkpoint-sqlite/tests/checkpoint_utils.py index 881bef7f4..f38afd740 100644 --- a/libs/checkpoint-sqlite/tests/checkpoint_utils.py +++ b/libs/checkpoint-sqlite/tests/checkpoint_utils.py @@ -1,13 +1,15 @@ +from __future__ import annotations + from collections.abc import Mapping from datetime import datetime, timezone -from typing import Any, Optional, Protocol +from typing import Any, Protocol from langgraph.checkpoint.base import Checkpoint, EmptyChannelError from langgraph.checkpoint.base.id import uuid6 class ChannelProtocol(Protocol): - def checkpoint(self) -> Optional[Any]: ... + def checkpoint(self) -> Any | None: ... def empty_checkpoint() -> Checkpoint: @@ -23,10 +25,10 @@ def empty_checkpoint() -> Checkpoint: def create_checkpoint( checkpoint: Checkpoint, - channels: Optional[Mapping[str, ChannelProtocol]], + channels: Mapping[str, ChannelProtocol] | None, step: int, *, - id: Optional[str] = None, + id: str | None = None, ) -> Checkpoint: """Create a checkpoint for the given channels.""" ts = datetime.now(timezone.utc).isoformat() diff --git a/libs/checkpoint/langgraph/checkpoint/base/__init__.py b/libs/checkpoint/langgraph/checkpoint/base/__init__.py index 4c57d033a..8117254a5 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/base/__init__.py @@ -1,11 +1,11 @@ +from __future__ import annotations + from collections.abc import AsyncIterator, Iterator, Sequence from typing import ( # noqa: UP035 Any, Generic, - List, Literal, NamedTuple, - Optional, TypedDict, TypeVar, Union, @@ -98,8 +98,8 @@ class CheckpointTuple(NamedTuple): config: RunnableConfig checkpoint: Checkpoint metadata: CheckpointMetadata - parent_config: Optional[RunnableConfig] = None - pending_writes: Optional[List[PendingWrite]] = None + parent_config: RunnableConfig | None = None + pending_writes: list[PendingWrite] | None = None class BaseCheckpointSaver(Generic[V]): @@ -121,11 +121,11 @@ class BaseCheckpointSaver(Generic[V]): def __init__( self, *, - serde: Optional[SerializerProtocol] = None, + serde: SerializerProtocol | None = None, ) -> None: self.serde = maybe_add_typed_methods(serde or self.serde) - def get(self, config: RunnableConfig) -> Optional[Checkpoint]: + def get(self, config: RunnableConfig) -> Checkpoint | None: """Fetch a checkpoint using the given configuration. Args: @@ -137,7 +137,7 @@ class BaseCheckpointSaver(Generic[V]): if value := self.get_tuple(config): return value.checkpoint - def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: + def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None: """Fetch a checkpoint tuple using the given configuration. Args: @@ -153,11 +153,11 @@ class BaseCheckpointSaver(Generic[V]): 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 that match the given criteria. @@ -229,7 +229,7 @@ class BaseCheckpointSaver(Generic[V]): """ raise NotImplementedError - async def aget(self, config: RunnableConfig) -> Optional[Checkpoint]: + async def aget(self, config: RunnableConfig) -> Checkpoint | None: """Asynchronously fetch a checkpoint using the given configuration. Args: @@ -241,7 +241,7 @@ class BaseCheckpointSaver(Generic[V]): if value := await self.aget_tuple(config): return value.checkpoint - async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: + async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None: """Asynchronously fetch a checkpoint tuple using the given configuration. Args: @@ -257,11 +257,11 @@ class BaseCheckpointSaver(Generic[V]): 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]: """Asynchronously list checkpoints that match the given criteria. @@ -334,7 +334,7 @@ class BaseCheckpointSaver(Generic[V]): """ raise NotImplementedError - def get_next_version(self, current: Optional[V]) -> V: + def get_next_version(self, current: V | None) -> V: """Generate the next version ID for a channel. Default is to use integer versions, incrementing by 1. If you override, you can use str/int/float versions, @@ -361,7 +361,7 @@ class EmptyChannelError(Exception): pass -def get_checkpoint_id(config: RunnableConfig) -> Optional[str]: +def get_checkpoint_id(config: RunnableConfig) -> str | None: """Get checkpoint ID in a backwards-compatible manner (fallback on thread_ts).""" return config["configurable"].get( "checkpoint_id", config["configurable"].get("thread_ts") diff --git a/libs/checkpoint/langgraph/checkpoint/base/id.py b/libs/checkpoint/langgraph/checkpoint/base/id.py index bfa503b87..30abc48a5 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/id.py +++ b/libs/checkpoint/langgraph/checkpoint/base/id.py @@ -3,10 +3,11 @@ https://github.com/oittaa/uuid6-python/blob/main/src/uuid6/__init__.py#L95 Bundled in to avoid install issues with uuid6 package """ +from __future__ import annotations + import random import time import uuid -from typing import Optional _last_v6_timestamp = None @@ -18,12 +19,12 @@ class UUID(uuid.UUID): def __init__( self, - hex: Optional[str] = None, - bytes: Optional[bytes] = None, - bytes_le: Optional[bytes] = None, - fields: Optional[tuple[int, int, int, int, int, int]] = None, - int: Optional[int] = None, - version: Optional[int] = None, + hex: str | None = None, + bytes: bytes | None = None, + bytes_le: bytes | None = None, + fields: tuple[int, int, int, int, int, int] | None = None, + int: int | None = None, + version: int | None = None, *, is_safe: uuid.SafeUUID = uuid.SafeUUID.unknown, ) -> None: @@ -75,7 +76,7 @@ def _subsec_decode(value: int) -> int: return -(-value * 10**6 // 2**20) -def uuid6(node: Optional[int] = None, clock_seq: Optional[int] = None) -> UUID: +def uuid6(node: int | None = None, clock_seq: int | None = None) -> UUID: r"""UUID version 6 is a field-compatible version of UUIDv1, reordered for improved DB locality. It is expected that UUIDv6 will primarily be used in contexts where there are existing v1 UUIDs. Systems that do diff --git a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py index 5caf43ff6..dc6d089e6 100644 --- a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import logging import os import pickle @@ -7,7 +9,7 @@ from collections import defaultdict from collections.abc import AsyncIterator, Iterator, Sequence from contextlib import AbstractAsyncContextManager, AbstractContextManager, ExitStack from types import TracebackType -from typing import Any, Optional, Union +from typing import Any from langchain_core.runnables import RunnableConfig @@ -63,9 +65,7 @@ class InMemorySaver( # thread ID -> checkpoint NS -> checkpoint ID -> checkpoint mapping storage: defaultdict[ str, - dict[ - str, dict[str, tuple[tuple[str, bytes], tuple[str, bytes], Optional[str]]] - ], + dict[str, dict[str, tuple[tuple[str, bytes], tuple[str, bytes], str | None]]], ] # (thread ID, checkpoint NS, checkpoint ID) -> (task ID, write idx) writes: defaultdict[ @@ -74,7 +74,7 @@ class InMemorySaver( ] blobs: dict[ tuple[ - str, str, str, Union[str, int, float] + str, str, str, str | int | float ], # thread id, checkpoint ns, channel, version tuple[str, bytes], ] @@ -82,7 +82,7 @@ class InMemorySaver( def __init__( self, *, - serde: Optional[SerializerProtocol] = None, + serde: SerializerProtocol | None = None, factory: type[defaultdict] = defaultdict, ) -> None: super().__init__(serde=serde) @@ -95,26 +95,26 @@ class InMemorySaver( self.stack.enter_context(self.writes) # type: ignore[arg-type] self.stack.enter_context(self.blobs) # type: ignore[arg-type] - def __enter__(self) -> "InMemorySaver": + def __enter__(self) -> InMemorySaver: return self.stack.__enter__() def __exit__( self, - exc_type: Optional[type[BaseException]], - exc_value: Optional[BaseException], - traceback: Optional[TracebackType], - ) -> Optional[bool]: + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: return self.stack.__exit__(exc_type, exc_value, traceback) - async def __aenter__(self) -> "InMemorySaver": + async def __aenter__(self) -> InMemorySaver: return self.stack.__enter__() async def __aexit__( self, - __exc_type: Optional[type[BaseException]], - __exc_value: Optional[BaseException], - __traceback: Optional[TracebackType], - ) -> Optional[bool]: + __exc_type: type[BaseException] | None, + __exc_value: BaseException | None, + __traceback: TracebackType | None, + ) -> bool | None: return self.stack.__exit__(__exc_type, __exc_value, __traceback) def _load_blobs( @@ -129,7 +129,7 @@ class InMemorySaver( channel_values[k] = self.serde.loads_typed(vv) return channel_values - def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: + def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None: """Get a checkpoint tuple from the in-memory storage. This method retrieves a checkpoint tuple from the in-memory storage based on the @@ -213,11 +213,11 @@ class InMemorySaver( 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 in-memory storage. @@ -422,7 +422,7 @@ class InMemorySaver( if k[0] == thread_id: del self.blobs[k] - async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: + async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None: """Asynchronous version of get_tuple. This method is an asynchronous wrapper around get_tuple that runs the synchronous @@ -438,11 +438,11 @@ class InMemorySaver( 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]: """Asynchronous version of list. @@ -512,7 +512,7 @@ class InMemorySaver( """ return self.delete_thread(thread_id) - def get_next_version(self, current: Optional[str]) -> str: + def get_next_version(self, current: str | None) -> str: if current is None: current_v = 0 elif isinstance(current, int): @@ -571,7 +571,7 @@ class PersistentDict(defaultdict): self.sync() self.clear() - def __enter__(self) -> "PersistentDict": + def __enter__(self) -> PersistentDict: return self def __exit__(self, *exc_info: Any) -> None: diff --git a/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py b/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py index 9699ba30d..db82b8079 100644 --- a/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py +++ b/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import dataclasses import decimal import importlib @@ -18,7 +20,7 @@ from ipaddress import ( IPv6Interface, IPv6Network, ) -from typing import Any, Callable, Optional, Union, cast +from typing import Any, Callable, cast from uuid import UUID from zoneinfo import ZoneInfo @@ -41,7 +43,7 @@ class JsonPlusSerializer(SerializerProtocol): self, *, pickle_fallback: bool = False, - __unpack_ext_hook__: Optional[Callable[[int, bytes], Any]] = None, + __unpack_ext_hook__: Callable[[int, bytes], Any] | None = None, ) -> None: self.pickle_fallback = pickle_fallback self._unpack_ext_hook = ( @@ -52,11 +54,11 @@ class JsonPlusSerializer(SerializerProtocol): def _encode_constructor_args( self, - constructor: Union[Callable, type[Any]], + constructor: Callable | type[Any], *, - method: Union[None, str, Sequence[Union[None, str]]] = None, - args: Optional[Sequence[Any]] = None, - kwargs: Optional[dict[str, Any]] = None, + method: None | str | Sequence[None | str] = None, + args: Sequence[Any] | None = None, + kwargs: dict[str, Any] | None = None, ) -> dict[str, Any]: out = { "lc": 2, @@ -71,7 +73,7 @@ class JsonPlusSerializer(SerializerProtocol): out["kwargs"] = kwargs return out - def _default(self, obj: Any) -> Union[str, dict[str, Any]]: + def _default(self, obj: Any) -> str | dict[str, Any]: if isinstance(obj, Serializable): return cast(dict[str, Any], obj.to_json()) elif hasattr(obj, "model_dump") and callable(obj.model_dump): @@ -251,7 +253,7 @@ EXT_PYDANTIC_V1 = 4 EXT_PYDANTIC_V2 = 5 -def _msgpack_default(obj: Any) -> Union[str, ormsgpack.Ext]: +def _msgpack_default(obj: Any) -> str | ormsgpack.Ext: if hasattr(obj, "model_dump") and callable(obj.model_dump): # pydantic v2 return ormsgpack.Ext( EXT_PYDANTIC_V2, diff --git a/libs/checkpoint/langgraph/store/base/__init__.py b/libs/checkpoint/langgraph/store/base/__init__.py index b20559ee9..b3379bf74 100644 --- a/libs/checkpoint/langgraph/store/base/__init__.py +++ b/libs/checkpoint/langgraph/store/base/__init__.py @@ -9,6 +9,8 @@ Core types: - Op: Get/Put/Search/List operations """ +from __future__ import annotations + from abc import ABC, abstractmethod from collections.abc import Iterable from datetime import datetime @@ -16,7 +18,6 @@ from typing import ( Any, Literal, NamedTuple, - Optional, TypedDict, Union, cast, @@ -127,7 +128,7 @@ class SearchItem(Item): value: dict[str, Any], created_at: datetime, updated_at: datetime, - score: Optional[float] = None, + score: float | None = None, ) -> None: """Initialize a result item. @@ -242,7 +243,7 @@ class SearchOp(NamedTuple): ``` """ - filter: Optional[dict[str, Any]] = None + filter: dict[str, Any] | None = None """Key-value pairs for filtering results based on exact matches or comparison operators. The filter supports both exact matches and operator-based comparisons. @@ -284,7 +285,7 @@ class SearchOp(NamedTuple): offset: int = 0 """Number of matching items to skip for pagination.""" - query: Optional[str] = None + query: str | None = None """Natural language search query for semantic search capabilities. ???+ example "Examples" @@ -379,7 +380,7 @@ class ListNamespacesOp(NamedTuple): """ - match_conditions: Optional[tuple[MatchCondition, ...]] = None + match_conditions: tuple[MatchCondition, ...] | None = None """Optional conditions for filtering namespaces. ???+ example "Examples" @@ -397,7 +398,7 @@ class ListNamespacesOp(NamedTuple): ``` """ - max_depth: Optional[int] = None + max_depth: int | None = None """Maximum depth of namespace hierarchy to return. Note: @@ -452,7 +453,7 @@ class PutOp(NamedTuple): the full path would effectively be "documents/user123/report1" """ - value: Optional[dict[str, Any]] + value: dict[str, Any] | None """The data to store, or None to mark the item for deletion. The value must be a dictionary with string keys and JSON-serializable values. @@ -466,7 +467,7 @@ class PutOp(NamedTuple): } """ - index: Optional[Union[Literal[False], list[str]]] = None # type: ignore[assignment] + index: Literal[False] | list[str] | None = None # type: ignore[assignment] """Controls how the item's fields are indexed for search operations. Indexing configuration determines how the item can be found through search: @@ -501,7 +502,7 @@ class PutOp(NamedTuple): ] ``` """ - ttl: Optional[float] = None + ttl: float | None = None """Controls the TTL (time-to-live) for the item in minutes. If provided, and if the store you are using supports this feature, the item @@ -530,14 +531,14 @@ class TTLConfig(TypedDict, total=False): This can be overridden per-operation by explicitly setting refresh_ttl. Defaults to True if not configured. """ - default_ttl: Optional[float] + default_ttl: float | None """Default TTL (time-to-live) in minutes for new items. If provided, new items will expire after this many minutes after their last access. The expiration timer refreshes on both read and write operations. Defaults to None (no expiration). """ - sweep_interval_minutes: Optional[int] + sweep_interval_minutes: int | None """Interval in minutes between TTL sweep operations. If provided, the store will periodically delete expired items based on TTL. @@ -565,7 +566,7 @@ class IndexConfig(TypedDict, total=False): - cohere:embed-multilingual-light-v3.0: 384 """ - embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, str] + embed: Embeddings | EmbeddingsFunc | AEmbeddingsFunc | str """Optional function to generate embeddings from text. Can be specified in three ways: @@ -633,7 +634,7 @@ class IndexConfig(TypedDict, total=False): ``` """ - fields: Optional[list[str]] + fields: list[str] | None """Fields to extract text from for embedding generation. Controls which parts of stored items are embedded for semantic search. Follows JSON path syntax: @@ -690,7 +691,7 @@ class BaseStore(ABC): """ supports_ttl: bool = False - ttl_config: Optional[TTLConfig] = None + ttl_config: TTLConfig | None = None __slots__ = ("__weakref__",) @@ -723,8 +724,8 @@ class BaseStore(ABC): namespace: tuple[str, ...], key: str, *, - refresh_ttl: Optional[bool] = None, - ) -> Optional[Item]: + refresh_ttl: bool | None = None, + ) -> Item | None: """Retrieve a single item. Args: @@ -746,11 +747,11 @@ class BaseStore(ABC): namespace_prefix: tuple[str, ...], /, *, - query: Optional[str] = None, - filter: Optional[dict[str, Any]] = None, + query: str | None = None, + filter: dict[str, Any] | None = None, limit: int = 10, offset: int = 0, - refresh_ttl: Optional[bool] = None, + refresh_ttl: bool | None = None, ) -> list[SearchItem]: """Search for items within a namespace prefix. @@ -817,9 +818,9 @@ class BaseStore(ABC): namespace: tuple[str, ...], key: str, value: dict[str, Any], - index: Optional[Union[Literal[False], list[str]]] = None, + index: Literal[False] | list[str] | None = None, *, - ttl: Union[Optional[float], "NotProvided"] = NOT_PROVIDED, + ttl: float | None | NotProvided = NOT_PROVIDED, ) -> None: """Store or update an item in the store. @@ -901,9 +902,9 @@ class BaseStore(ABC): def list_namespaces( self, *, - prefix: Optional[NamespacePath] = None, - suffix: Optional[NamespacePath] = None, - max_depth: Optional[int] = None, + prefix: NamespacePath | None = None, + suffix: NamespacePath | None = None, + max_depth: int | None = None, limit: int = 100, offset: int = 0, ) -> list[tuple[str, ...]]: @@ -956,8 +957,8 @@ class BaseStore(ABC): namespace: tuple[str, ...], key: str, *, - refresh_ttl: Optional[bool] = None, - ) -> Optional[Item]: + refresh_ttl: bool | None = None, + ) -> Item | None: """Asynchronously retrieve a single item. Args: @@ -984,11 +985,11 @@ class BaseStore(ABC): namespace_prefix: tuple[str, ...], /, *, - query: Optional[str] = None, - filter: Optional[dict[str, Any]] = None, + query: str | None = None, + filter: dict[str, Any] | None = None, limit: int = 10, offset: int = 0, - refresh_ttl: Optional[bool] = None, + refresh_ttl: bool | None = None, ) -> list[SearchItem]: """Asynchronously search for items within a namespace prefix. @@ -1058,9 +1059,9 @@ class BaseStore(ABC): namespace: tuple[str, ...], key: str, value: dict[str, Any], - index: Optional[Union[Literal[False], list[str]]] = None, + index: Literal[False] | list[str] | None = None, *, - ttl: Union[Optional[float], "NotProvided"] = NOT_PROVIDED, + ttl: float | None | NotProvided = NOT_PROVIDED, ) -> None: """Asynchronously store or update an item in the store. @@ -1150,9 +1151,9 @@ class BaseStore(ABC): async def alist_namespaces( self, *, - prefix: Optional[NamespacePath] = None, - suffix: Optional[NamespacePath] = None, - max_depth: Optional[int] = None, + prefix: NamespacePath | None = None, + suffix: NamespacePath | None = None, + max_depth: int | None = None, limit: int = 100, offset: int = 0, ) -> list[tuple[str, ...]]: @@ -1226,7 +1227,7 @@ def _validate_namespace(namespace: tuple[str, ...]) -> None: def _ensure_refresh( - ttl_config: Optional[TTLConfig], refresh_ttl: Optional[bool] = None + ttl_config: TTLConfig | None, refresh_ttl: bool | None = None ) -> bool: if refresh_ttl is not None: return refresh_ttl @@ -1236,9 +1237,9 @@ def _ensure_refresh( def _ensure_ttl( - ttl_config: Optional[TTLConfig], - ttl: Union[Optional[float], "NotProvided"] = NOT_PROVIDED, -) -> Optional[float]: + ttl_config: TTLConfig | None, + ttl: float | None | NotProvided = NOT_PROVIDED, +) -> float | None: if ttl is NOT_PROVIDED: if ttl_config: return ttl_config.get("default_ttl") diff --git a/libs/checkpoint/langgraph/store/base/batch.py b/libs/checkpoint/langgraph/store/base/batch.py index 63ea02031..7d9daad4a 100644 --- a/libs/checkpoint/langgraph/store/base/batch.py +++ b/libs/checkpoint/langgraph/store/base/batch.py @@ -1,10 +1,12 @@ """Utilities for batching operations in a background task.""" +from __future__ import annotations + import asyncio import functools import weakref from collections.abc import Iterable -from typing import Any, Callable, Literal, Optional, TypeVar, Union +from typing import Any, Callable, Literal, TypeVar from langgraph.store.base import ( NOT_PROVIDED, @@ -30,7 +32,7 @@ F = TypeVar("F", bound=Callable) def _check_loop(func: F) -> F: @functools.wraps(func) - def wrapper(store: "AsyncBatchedBaseStore", *args: Any, **kwargs: Any) -> Any: + def wrapper(store: AsyncBatchedBaseStore, *args: Any, **kwargs: Any) -> Any: method_name: str = func.__name__ try: current_loop = asyncio.get_running_loop() @@ -75,8 +77,8 @@ class AsyncBatchedBaseStore(BaseStore): namespace: tuple[str, ...], key: str, *, - refresh_ttl: Optional[bool] = None, - ) -> Optional[Item]: + refresh_ttl: bool | None = None, + ) -> Item | None: assert not self._task.done() fut = self._loop.create_future() self._aqueue.put_nowait( @@ -96,11 +98,11 @@ class AsyncBatchedBaseStore(BaseStore): namespace_prefix: tuple[str, ...], /, *, - query: Optional[str] = None, - filter: Optional[dict[str, Any]] = None, + query: str | None = None, + filter: dict[str, Any] | None = None, limit: int = 10, offset: int = 0, - refresh_ttl: Optional[bool] = None, + refresh_ttl: bool | None = None, ) -> list[SearchItem]: assert not self._task.done() fut = self._loop.create_future() @@ -124,9 +126,9 @@ class AsyncBatchedBaseStore(BaseStore): namespace: tuple[str, ...], key: str, value: dict[str, Any], - index: Optional[Union[Literal[False], list[str]]] = None, + index: Literal[False] | list[str] | None = None, *, - ttl: Union[Optional[float], "NotProvided"] = NOT_PROVIDED, + ttl: float | None | NotProvided = NOT_PROVIDED, ) -> None: assert not self._task.done() _validate_namespace(namespace) @@ -154,9 +156,9 @@ class AsyncBatchedBaseStore(BaseStore): async def alist_namespaces( self, *, - prefix: Optional[NamespacePath] = None, - suffix: Optional[NamespacePath] = None, - max_depth: Optional[int] = None, + prefix: NamespacePath | None = None, + suffix: NamespacePath | None = None, + max_depth: int | None = None, limit: int = 100, offset: int = 0, ) -> list[tuple[str, ...]]: @@ -187,8 +189,8 @@ class AsyncBatchedBaseStore(BaseStore): namespace: tuple[str, ...], key: str, *, - refresh_ttl: Optional[bool] = None, - ) -> Optional[Item]: + refresh_ttl: bool | None = None, + ) -> Item | None: return asyncio.run_coroutine_threadsafe( self.aget(namespace, key=key, refresh_ttl=refresh_ttl), self._loop ).result() @@ -199,11 +201,11 @@ class AsyncBatchedBaseStore(BaseStore): namespace_prefix: tuple[str, ...], /, *, - query: Optional[str] = None, - filter: Optional[dict[str, Any]] = None, + query: str | None = None, + filter: dict[str, Any] | None = None, limit: int = 10, offset: int = 0, - refresh_ttl: Optional[bool] = None, + refresh_ttl: bool | None = None, ) -> list[SearchItem]: return asyncio.run_coroutine_threadsafe( self.asearch( @@ -223,9 +225,9 @@ class AsyncBatchedBaseStore(BaseStore): namespace: tuple[str, ...], key: str, value: dict[str, Any], - index: Optional[Union[Literal[False], list[str]]] = None, + index: Literal[False] | list[str] | None = None, *, - ttl: Union[Optional[float], "NotProvided"] = NOT_PROVIDED, + ttl: float | None | NotProvided = NOT_PROVIDED, ) -> None: _validate_namespace(namespace) asyncio.run_coroutine_threadsafe( @@ -253,9 +255,9 @@ class AsyncBatchedBaseStore(BaseStore): def list_namespaces( self, *, - prefix: Optional[NamespacePath] = None, - suffix: Optional[NamespacePath] = None, - max_depth: Optional[int] = None, + prefix: NamespacePath | None = None, + suffix: NamespacePath | None = None, + max_depth: int | None = None, limit: int = 100, offset: int = 0, ) -> list[tuple[str, ...]]: @@ -271,7 +273,7 @@ class AsyncBatchedBaseStore(BaseStore): ).result() -def _dedupe_ops(values: list[Op]) -> tuple[Optional[list[int]], list[Op]]: +def _dedupe_ops(values: list[Op]) -> tuple[list[int] | None, list[Op]]: """Dedupe operations while preserving order for results. Args: diff --git a/libs/checkpoint/langgraph/store/base/embed.py b/libs/checkpoint/langgraph/store/base/embed.py index 2347821c0..738a9e53f 100644 --- a/libs/checkpoint/langgraph/store/base/embed.py +++ b/libs/checkpoint/langgraph/store/base/embed.py @@ -6,11 +6,13 @@ with LangChain-compatible tools while maintaining support for both synchronous a asynchronous operations. """ +from __future__ import annotations + import asyncio import functools import json from collections.abc import Awaitable, Sequence -from typing import Any, Callable, Optional, Union +from typing import Any, Callable from langchain_core.embeddings import Embeddings @@ -30,7 +32,7 @@ Similar to EmbeddingsFunc, but returns an awaitable that resolves to the embeddi def ensure_embeddings( - embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, str, None], + embed: Embeddings | EmbeddingsFunc | AEmbeddingsFunc | str | None, ) -> Embeddings: """Ensure that an embedding function conforms to LangChain's Embeddings interface. @@ -141,7 +143,7 @@ class EmbeddingsLambda(Embeddings): def __init__( self, - func: Union[EmbeddingsFunc, AEmbeddingsFunc], + func: EmbeddingsFunc | AEmbeddingsFunc, ) -> None: if func is None: raise ValueError("func must be provided") @@ -221,7 +223,7 @@ class EmbeddingsLambda(Embeddings): return (await afunc([text]))[0] -def get_text_at_path(obj: Any, path: Union[str, list[str]]) -> list[str]: +def get_text_at_path(obj: Any, path: str | list[str]) -> list[str]: """Extract text from an object using a path expression or pre-tokenized path. Args: @@ -279,7 +281,7 @@ def get_text_at_path(obj: Any, path: Union[str, list[str]]) -> list[str]: for field in fields: nested_tokens = tokenize_path(field) if nested_tokens: - current_obj: Optional[dict] = obj + current_obj: dict | None = obj for nested_token in nested_tokens: if ( isinstance(current_obj, dict) @@ -404,7 +406,7 @@ def _is_async_callable( @functools.lru_cache -def _get_init_embeddings() -> Optional[Callable[[str], Embeddings]]: +def _get_init_embeddings() -> Callable[[str], Embeddings] | None: try: from langchain.embeddings import init_embeddings # type: ignore diff --git a/libs/checkpoint/langgraph/store/memory/__init__.py b/libs/checkpoint/langgraph/store/memory/__init__.py index ba6943de4..62fb3085f 100644 --- a/libs/checkpoint/langgraph/store/memory/__init__.py +++ b/libs/checkpoint/langgraph/store/memory/__init__.py @@ -99,6 +99,8 @@ Tip: ``` """ +from __future__ import annotations + import asyncio import concurrent.futures as cf import functools @@ -107,7 +109,7 @@ from collections import defaultdict from collections.abc import Iterable from datetime import datetime, timezone from importlib import util -from typing import Any, Optional +from typing import Any from langchain_core.embeddings import Embeddings @@ -178,7 +180,7 @@ class InMemoryStore(BaseStore): "embeddings", ) - def __init__(self, *, index: Optional[IndexConfig] = None) -> None: + def __init__(self, *, index: IndexConfig | None = None) -> None: # Both _data and _vectors are wrapped in the In-memory API # Do not change their names self._data: dict[tuple[str, ...], dict[str, Item]] = defaultdict(dict) @@ -189,7 +191,7 @@ class InMemoryStore(BaseStore): self.index_config = index if self.index_config: self.index_config = self.index_config.copy() - self.embeddings: Optional[Embeddings] = ensure_embeddings( + self.embeddings: Embeddings | None = ensure_embeddings( self.index_config.get("embed"), ) self.index_config["__tokenized_fields"] = [ @@ -325,7 +327,7 @@ class InMemoryStore(BaseStore): ) # max pooling seen: set[tuple[tuple[str, ...], str]] = set() - kept: list[tuple[Optional[float], Item]] = [] + kept: list[tuple[float | None, Item]] = [] for score, item in sorted_results: key = (item.namespace, item.key) if key in seen: diff --git a/libs/checkpoint/pyproject.toml b/libs/checkpoint/pyproject.toml index bf2fbd2a5..86a8802cd 100644 --- a/libs/checkpoint/pyproject.toml +++ b/libs/checkpoint/pyproject.toml @@ -46,7 +46,7 @@ lint.select = [ "B", # flake8-bugbear "I", # isort ] -lint.ignore = ["E501", "B008", "UP007", "UP006"] +lint.ignore = ["E501", "B008"] [tool.pytest-watcher] now = true diff --git a/libs/checkpoint/tests/checkpoint_utils.py b/libs/checkpoint/tests/checkpoint_utils.py index 881bef7f4..f38afd740 100644 --- a/libs/checkpoint/tests/checkpoint_utils.py +++ b/libs/checkpoint/tests/checkpoint_utils.py @@ -1,13 +1,15 @@ +from __future__ import annotations + from collections.abc import Mapping from datetime import datetime, timezone -from typing import Any, Optional, Protocol +from typing import Any, Protocol from langgraph.checkpoint.base import Checkpoint, EmptyChannelError from langgraph.checkpoint.base.id import uuid6 class ChannelProtocol(Protocol): - def checkpoint(self) -> Optional[Any]: ... + def checkpoint(self) -> Any | None: ... def empty_checkpoint() -> Checkpoint: @@ -23,10 +25,10 @@ def empty_checkpoint() -> Checkpoint: def create_checkpoint( checkpoint: Checkpoint, - channels: Optional[Mapping[str, ChannelProtocol]], + channels: Mapping[str, ChannelProtocol] | None, step: int, *, - id: Optional[str] = None, + id: str | None = None, ) -> Checkpoint: """Create a checkpoint for the given channels.""" ts = datetime.now(timezone.utc).isoformat() diff --git a/libs/langgraph/langgraph/channels/topic.py b/libs/langgraph/langgraph/channels/topic.py index e6cf1b2b6..fa96bcd59 100644 --- a/libs/langgraph/langgraph/channels/topic.py +++ b/libs/langgraph/langgraph/channels/topic.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from collections.abc import Iterator, Sequence from typing import Any, Generic, Union @@ -8,7 +10,7 @@ from langgraph.constants import MISSING from langgraph.errors import EmptyChannelError -def flatten(values: Sequence[Union[Value, list[Value]]]) -> Iterator[Value]: +def flatten(values: Sequence[Value | list[Value]]) -> Iterator[Value]: for value in values: if isinstance(value, list): yield from value @@ -70,7 +72,7 @@ class Topic( empty.values = checkpoint return empty - def update(self, values: Sequence[Union[Value, list[Value]]]) -> bool: + def update(self, values: Sequence[Value | list[Value]]) -> bool: updated = False if not self.accumulate: updated = bool(self.values) diff --git a/libs/langgraph/langgraph/func/__init__.py b/libs/langgraph/langgraph/func/__init__.py index 825441b07..e85ccecf8 100644 --- a/libs/langgraph/langgraph/func/__init__.py +++ b/libs/langgraph/langgraph/func/__init__.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import concurrent.futures import functools @@ -9,9 +11,7 @@ from typing import ( Any, Callable, Generic, - Optional, TypeVar, - Union, get_args, get_origin, overload, @@ -47,8 +47,8 @@ class TaskFunction(Generic[P, T]): func: Callable[P, T], *, retry_policy: Sequence[RetryPolicy], - cache_policy: Optional[CachePolicy[Callable[P, Union[str, bytes]]]] = None, - name: Optional[str] = None, + cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None, + name: str | None = None, ) -> None: if name is not None: if hasattr(func, "__func__"): @@ -91,36 +91,33 @@ class TaskFunction(Generic[P, T]): @overload def task( *, - name: Optional[str] = None, - retry_policy: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, - cache_policy: Optional[CachePolicy[Callable[P, Union[str, bytes]]]] = None, + name: str | None = None, + retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Callable[ - [Union[Callable[P, Awaitable[T]], Callable[P, T]]], + [Callable[P, Awaitable[T]] | Callable[P, T]], TaskFunction[P, T], ]: ... @overload def task( - __func_or_none__: Union[Callable[P, Awaitable[T]], Callable[P, T]], + __func_or_none__: Callable[P, Awaitable[T]] | Callable[P, T], ) -> TaskFunction[P, T]: ... def task( - __func_or_none__: Optional[Union[Callable[P, Awaitable[T]], Callable[P, T]]] = None, + __func_or_none__: Callable[P, Awaitable[T]] | Callable[P, T] | None = None, *, - name: Optional[str] = None, - retry_policy: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, - cache_policy: Optional[CachePolicy[Callable[P, Union[str, bytes]]]] = None, + name: str | None = None, + retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None, **kwargs: Unpack[DeprecatedKwargs], -) -> Union[ - Callable[ - [Union[Callable[P, Awaitable[T]], Callable[P, T]]], - TaskFunction[P, T], - ], - TaskFunction[P, T], -]: +) -> ( + Callable[[Callable[P, Awaitable[T]] | Callable[P, T]], TaskFunction[P, T]] + | TaskFunction[P, T] +): """Define a LangGraph task using the `task` decorator. !!! important "Requires python 3.11 or higher for async functions" @@ -196,10 +193,8 @@ def task( ) def decorator( - func: Union[Callable[P, Awaitable[T]], Callable[P, T]], - ) -> Union[ - Callable[P, concurrent.futures.Future[T]], Callable[P, asyncio.Future[T]] - ]: + func: Callable[P, Awaitable[T]] | Callable[P, T], + ) -> Callable[P, concurrent.futures.Future[T]] | Callable[P, asyncio.Future[T]]: return TaskFunction( func, retry_policy=retry_policies, cache_policy=cache_policy, name=name ) @@ -376,12 +371,12 @@ class entrypoint: def __init__( self, - checkpointer: Optional[BaseCheckpointSaver] = None, - store: Optional[BaseStore] = None, - cache: Optional[BaseCache] = None, - config_schema: Optional[type[Any]] = None, - cache_policy: Optional[CachePolicy] = None, - retry_policy: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, + checkpointer: BaseCheckpointSaver | None = None, + store: BaseStore | None = None, + cache: BaseCache | None = None, + config_schema: type[Any] | None = None, + cache_policy: CachePolicy | None = None, + retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> None: """Initialize the entrypoint decorator.""" diff --git a/libs/langgraph/langgraph/graph/branch.py b/libs/langgraph/langgraph/graph/branch.py index 6df2ab1b5..f120167d6 100644 --- a/libs/langgraph/langgraph/graph/branch.py +++ b/libs/langgraph/langgraph/graph/branch.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from collections.abc import Awaitable, Hashable, Sequence from inspect import ( isfunction, @@ -11,7 +13,6 @@ from typing import ( Callable, Literal, NamedTuple, - Optional, Union, cast, get_args, @@ -40,21 +41,18 @@ Writer = Callable[ def _get_branch_path_input_schema( - path: Union[ - Callable[..., Union[Hashable, list[Hashable]]], - Callable[..., Awaitable[Union[Hashable, list[Hashable]]]], - Runnable[Any, Union[Hashable, list[Hashable]]], - ], -) -> Optional[type[Any]]: + path: Callable[..., Hashable | list[Hashable]] + | Callable[..., Awaitable[Hashable | list[Hashable]]] + | Runnable[Any, Hashable | list[Hashable]], +) -> type[Any] | None: input = None # detect input schema annotation in the branch callable try: - callable_: Optional[ - Union[ - Callable[..., Union[Hashable, list[Hashable]]], - Callable[..., Awaitable[Union[Hashable, list[Hashable]]]], - ] - ] = None + callable_: ( + Callable[..., Hashable | list[Hashable]] + | Callable[..., Awaitable[Hashable | list[Hashable]]] + | None + ) = None if isinstance(path, (RunnableCallable, RunnableLambda)): if isfunction(path.func) or ismethod(path.func): callable_ = path.func @@ -85,19 +83,19 @@ def _get_branch_path_input_schema( class Branch(NamedTuple): - path: Runnable[Any, Union[Hashable, list[Hashable]]] - ends: Optional[dict[Hashable, str]] - input_schema: Optional[type[Any]] = None + path: Runnable[Any, Hashable | list[Hashable]] + ends: dict[Hashable, str] | None + input_schema: type[Any] | None = None @classmethod def from_path( cls, - path: Runnable[Any, Union[Hashable, list[Hashable]]], - path_map: Optional[Union[dict[Hashable, str], list[str]]], + path: Runnable[Any, Hashable | list[Hashable]], + path_map: dict[Hashable, str] | list[str] | None, infer_schema: bool = False, - ) -> "Branch": + ) -> Branch: # coerce path_map to a dictionary - path_map_: Optional[dict[Hashable, str]] = None + path_map_: dict[Hashable, str] | None = None try: if isinstance(path_map, dict): path_map_ = path_map.copy() @@ -105,7 +103,7 @@ class Branch(NamedTuple): path_map_ = {name: name for name in path_map} else: # find func - func: Optional[Callable] = None + func: Callable | None = None if isinstance(path, (RunnableCallable, RunnableLambda)): func = path.func or path.afunc if func is not None: @@ -126,7 +124,7 @@ class Branch(NamedTuple): def run( self, writer: Writer, - reader: Optional[Callable[[RunnableConfig], Any]] = None, + reader: Callable[[RunnableConfig], Any] | None = None, ) -> RunnableCallable: return ChannelWrite.register_writer( RunnableCallable( @@ -153,7 +151,7 @@ class Branch(NamedTuple): input: Any, config: RunnableConfig, *, - reader: Optional[Callable[[RunnableConfig], Any]], + reader: Callable[[RunnableConfig], Any] | None, writer: Writer, ) -> Runnable: if reader: @@ -176,7 +174,7 @@ class Branch(NamedTuple): input: Any, config: RunnableConfig, *, - reader: Optional[Callable[[RunnableConfig], Any]], + reader: Callable[[RunnableConfig], Any] | None, writer: Writer, ) -> Runnable: if reader: @@ -200,11 +198,11 @@ class Branch(NamedTuple): input: Any, result: Any, config: RunnableConfig, - ) -> Union[Runnable, Any]: + ) -> Runnable | Any: if not isinstance(result, (list, tuple)): result = [result] if self.ends: - destinations: Sequence[Union[Send, str]] = [ + destinations: Sequence[Send | str] = [ r if isinstance(r, Send) else self.ends[r] for r in result ] else: diff --git a/libs/langgraph/langgraph/graph/message.py b/libs/langgraph/langgraph/graph/message.py index 2a752e0c5..059102027 100644 --- a/libs/langgraph/langgraph/graph/message.py +++ b/libs/langgraph/langgraph/graph/message.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import uuid import warnings from collections.abc import Sequence @@ -7,7 +9,6 @@ from typing import ( Any, Callable, Literal, - Optional, Union, cast, ) @@ -32,8 +33,8 @@ REMOVE_ALL_MESSAGES = "__remove_all__" def _add_messages_wrapper(func: Callable) -> Callable[[Messages, Messages], Messages]: def _add_messages( - left: Optional[Messages] = None, right: Optional[Messages] = None, **kwargs: Any - ) -> Union[Messages, Callable[[Messages, Messages], Messages]]: + left: Messages | None = None, right: Messages | None = None, **kwargs: Any + ) -> Messages | Callable[[Messages, Messages], Messages]: if left is not None and right is not None: return func(left, right, **kwargs) elif left is not None or right is not None: @@ -54,7 +55,7 @@ def add_messages( left: Messages, right: Messages, *, - format: Optional[Literal["langchain-openai"]] = None, + format: Literal["langchain-openai"] | None = None, ) -> Messages: """Merges two lists of messages, updating existing messages by ID. @@ -246,9 +247,9 @@ def _format_messages(messages: Sequence[BaseMessage]) -> list[BaseMessage]: def push_message( - message: Union[MessageLikeRepresentation, BaseMessageChunk], + message: MessageLikeRepresentation | BaseMessageChunk, *, - state_key: Optional[str] = "messages", + state_key: str | None = "messages", ) -> AnyMessage: """Write a message manually to the `messages` / `messages-tuple` stream mode. diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 72f0f9f66..b6b73310e 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -15,7 +15,6 @@ from typing import ( Generic, Literal, NamedTuple, - Optional, Protocol, Union, cast, @@ -91,7 +90,7 @@ from langgraph.warnings import LangGraphDeprecatedSinceV10 logger = logging.getLogger(__name__) -def _warn_invalid_state_schema(schema: Union[type[Any], Any]) -> None: +def _warn_invalid_state_schema(schema: type[Any] | Any) -> None: if isinstance(schema, type): return if typing.get_args(schema): @@ -174,11 +173,11 @@ class StateNodeSpec(NamedTuple): # TODO: rename this callable, also move away from NamedTuple so that we can use # a generic StateNode, so maybe a dataclass runnable: StateNode - metadata: Optional[dict[str, Any]] + metadata: dict[str, Any] | None input: type[Any] - retry_policy: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] - cache_policy: Optional[CachePolicy] - ends: Optional[Union[tuple[str, ...], dict[str, str]]] = EMPTY_SEQ + retry_policy: RetryPolicy | Sequence[RetryPolicy] | None + cache_policy: CachePolicy | None + ends: tuple[str, ...] | dict[str, str] | None = EMPTY_SEQ defer: bool = False @@ -239,7 +238,7 @@ class StateGraph(Generic[StateT, InputT]): branches: defaultdict[str, dict[str, Branch]] channels: dict[str, BaseChannel] managed: dict[str, ManagedValueSpec] - schemas: dict[type[Any], dict[str, Union[BaseChannel, ManagedValueSpec]]] + schemas: dict[type[Any], dict[str, BaseChannel | ManagedValueSpec]] def __init__( self, @@ -313,11 +312,11 @@ class StateGraph(Generic[StateT, InputT]): node: StateNode[StateT], *, defer: bool = False, - metadata: Optional[dict[str, Any]] = None, - input: Optional[type[Any]] = None, - retry_policy: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, - cache_policy: Optional[CachePolicy] = None, - destinations: Optional[Union[dict[str, str], tuple[str, ...]]] = None, + metadata: dict[str, Any] | None = None, + input: type[Any] | None = None, + retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy | None = None, + destinations: dict[str, str] | tuple[str, ...] | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Self: """Add a new node to the state graph. @@ -332,11 +331,11 @@ class StateGraph(Generic[StateT, InputT]): action: StateNode[StateT], *, defer: bool = False, - metadata: Optional[dict[str, Any]] = None, - input: Optional[type[Any]] = None, - retry_policy: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, - cache_policy: Optional[CachePolicy] = None, - destinations: Optional[Union[dict[str, str], tuple[str, ...]]] = None, + metadata: dict[str, Any] | None = None, + input: type[Any] | None = None, + retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy | None = None, + destinations: dict[str, str] | tuple[str, ...] | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Self: """Add a new node to the state graph.""" @@ -344,15 +343,15 @@ class StateGraph(Generic[StateT, InputT]): def add_node( self, - node: Union[str, StateNode[StateT]], - action: Optional[StateNode[StateT]] = None, + node: str | StateNode[StateT], + action: StateNode[StateT] | None = None, *, defer: bool = False, - metadata: Optional[dict[str, Any]] = None, - input: Optional[type[Any]] = None, - retry_policy: Optional[Union[RetryPolicy, Sequence[RetryPolicy]]] = None, - cache_policy: Optional[CachePolicy] = None, - destinations: Optional[Union[dict[str, str], tuple[str, ...]]] = None, + metadata: dict[str, Any] | None = None, + input: type[Any] | None = None, + retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy | None = None, + destinations: dict[str, str] | tuple[str, ...] | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Self: """Add a new node to the state graph. @@ -451,7 +450,7 @@ class StateGraph(Generic[StateT, InputT]): f"'{character}' is a reserved character and is not allowed in the node names." ) - ends: Union[tuple[str, ...], dict[str, str]] = EMPTY_SEQ + ends: tuple[str, ...] | dict[str, str] = EMPTY_SEQ try: if ( isfunction(action) @@ -512,7 +511,7 @@ class StateGraph(Generic[StateT, InputT]): ) return self - def add_edge(self, start_key: Union[str, list[str]], end_key: str) -> Self: + def add_edge(self, start_key: str | list[str], end_key: str) -> Self: """Add a directed edge from the start node (or list of start nodes) to the end node. When a single start node is provided, the graph will wait for that node to complete @@ -569,12 +568,10 @@ class StateGraph(Generic[StateT, InputT]): def add_conditional_edges( self, source: str, - path: Union[ - Callable[..., Union[Hashable, list[Hashable]]], - Callable[..., Awaitable[Union[Hashable, list[Hashable]]]], - Runnable[Any, Union[Hashable, list[Hashable]]], - ], - path_map: Optional[Union[dict[Hashable, str], list[str]]] = None, + path: Callable[..., Hashable | list[Hashable]] + | Callable[..., Awaitable[Hashable | list[Hashable]]] + | Runnable[Any, Hashable | list[Hashable]], + path_map: dict[Hashable, str] | list[str] | None = None, ) -> Self: """Add a conditional edge from the starting node to any number of destination nodes. @@ -616,7 +613,7 @@ class StateGraph(Generic[StateT, InputT]): def add_sequence( self, - nodes: Sequence[Union[StateNode[StateT], tuple[str, StateNode[StateT]]]], + nodes: Sequence[StateNode[StateT] | tuple[str, StateNode[StateT]]], ) -> Self: """Add a sequence of nodes that will be executed in the provided order. @@ -635,7 +632,7 @@ class StateGraph(Generic[StateT, InputT]): if len(nodes) < 1: raise ValueError("Sequence requires at least one node.") - previous_name: Optional[str] = None + previous_name: str | None = None for node in nodes: if isinstance(node, tuple) and len(node) == 2: name, node = node @@ -671,12 +668,10 @@ class StateGraph(Generic[StateT, InputT]): def set_conditional_entry_point( self, - path: Union[ - Callable[..., Union[Hashable, list[Hashable]]], - Callable[..., Awaitable[Union[Hashable, list[Hashable]]]], - Runnable[Any, Union[Hashable, list[Hashable]]], - ], - path_map: Optional[Union[dict[Hashable, str], list[str]]] = None, + path: Callable[..., Hashable | list[Hashable]] + | Callable[..., Awaitable[Hashable | list[Hashable]]] + | Runnable[Any, Hashable | list[Hashable]], + path_map: dict[Hashable, str] | list[str] | None = None, ) -> Self: """Sets a conditional entry point in the graph. @@ -705,7 +700,7 @@ class StateGraph(Generic[StateT, InputT]): """ return self.add_edge(key, END) - def validate(self, interrupt: Optional[Sequence[str]] = None) -> Self: + def validate(self, interrupt: Sequence[str] | None = None) -> Self: # assemble sources all_sources = {src for src, _ in self._all_edges} for start, branches in self.branches.items(): @@ -759,12 +754,12 @@ class StateGraph(Generic[StateT, InputT]): self: StateGraph[StateT, Unset], checkpointer: Checkpointer = None, *, - cache: Optional[BaseCache] = None, - store: Optional[BaseStore] = None, - interrupt_before: Optional[Union[All, list[str]]] = None, - interrupt_after: Optional[Union[All, list[str]]] = None, + cache: BaseCache | None = None, + store: BaseStore | None = None, + interrupt_before: All | list[str] | None = None, + interrupt_after: All | list[str] | None = None, debug: bool = False, - name: Optional[str] = None, + name: str | None = None, ) -> CompiledStateGraph[StateT, StateT]: ... @overload @@ -772,25 +767,25 @@ class StateGraph(Generic[StateT, InputT]): self: StateGraph[StateT, InputT], checkpointer: Checkpointer = None, *, - cache: Optional[BaseCache] = None, - store: Optional[BaseStore] = None, - interrupt_before: Optional[Union[All, list[str]]] = None, - interrupt_after: Optional[Union[All, list[str]]] = None, + cache: BaseCache | None = None, + store: BaseStore | None = None, + interrupt_before: All | list[str] | None = None, + interrupt_after: All | list[str] | None = None, debug: bool = False, - name: Optional[str] = None, + name: str | None = None, ) -> CompiledStateGraph[StateT, InputT]: ... def compile( self, checkpointer: Checkpointer = None, *, - cache: Optional[BaseCache] = None, - store: Optional[BaseStore] = None, - interrupt_before: Optional[Union[All, list[str]]] = None, - interrupt_after: Optional[Union[All, list[str]]] = None, + cache: BaseCache | None = None, + store: BaseStore | None = None, + interrupt_before: All | list[str] | None = None, + interrupt_after: All | list[str] | None = None, debug: bool = False, - name: Optional[str] = None, - ) -> Union[CompiledStateGraph[StateT, StateT], CompiledStateGraph[StateT, InputT]]: + name: str | None = None, + ) -> CompiledStateGraph[StateT, StateT] | CompiledStateGraph[StateT, InputT]: """Compiles the state graph into a `CompiledStateGraph` object. The compiled graph implements the `Runnable` interface and can be invoked, @@ -842,7 +837,7 @@ class StateGraph(Generic[StateT, InputT]): ] ) - ResolvedInputT: Union[type[InputT], type[StateT]] = self.input or self.schema + ResolvedInputT: type[InputT] | type[StateT] = self.input or self.schema compiled = CompiledStateGraph[StateT, ResolvedInputT]( # type: ignore[valid-type] builder=self, schema_to_mapper={}, @@ -893,22 +888,20 @@ class StateGraph(Generic[StateT, InputT]): class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): builder: StateGraph[StateT, InputT] - schema_to_mapper: dict[type[Any], Optional[Callable[[Any], Any]]] + schema_to_mapper: dict[type[Any], Callable[[Any], Any] | None] def __init__( self, *, builder: StateGraph[StateT, InputT], - schema_to_mapper: dict[type[Any], Optional[Callable[[Any], Any]]], + schema_to_mapper: dict[type[Any], Callable[[Any], Any] | None], **kwargs: Any, ) -> None: super().__init__(**kwargs) self.builder = builder self.schema_to_mapper = schema_to_mapper - def get_input_schema( - self, config: Optional[RunnableConfig] = None - ) -> type[BaseModel]: + def get_input_schema(self, config: RunnableConfig | None = None) -> type[BaseModel]: return _get_schema( typ=self.builder.input, schemas=self.builder.schemas, @@ -917,7 +910,7 @@ class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): ) def get_output_schema( - self, config: Optional[RunnableConfig] = None + self, config: RunnableConfig | None = None ) -> type[BaseModel]: return _get_schema( typ=self.builder.output, @@ -926,7 +919,7 @@ class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): name=self.get_name("Output"), ) - def attach_node(self, key: str, node: Optional[StateNodeSpec]) -> None: + def attach_node(self, key: str, node: StateNodeSpec | None) -> None: if key == START: output_keys = [ k @@ -939,8 +932,8 @@ class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): ] def _get_updates( - input: Union[None, dict, Any], - ) -> Optional[Sequence[tuple[str, Any]]]: + input: None | dict | Any, + ) -> Sequence[tuple[str, Any]] | None: if input is None: return None elif isinstance(input, dict): @@ -977,7 +970,7 @@ class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): raise InvalidUpdateError(msg) # state updaters - write_entries: tuple[Union[ChannelWriteEntry, ChannelWriteTupleEntry], ...] = ( + write_entries: tuple[ChannelWriteEntry | ChannelWriteTupleEntry, ...] = ( ChannelWriteTupleEntry( mapper=_get_root if output_keys == ["__root__"] else _get_updates ), @@ -1032,7 +1025,7 @@ class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): else: raise RuntimeError - def attach_edge(self, starts: Union[str, Sequence[str]], end: str) -> None: + def attach_edge(self, starts: str | Sequence[str], end: str) -> None: if isinstance(starts, str): # subscribe to start channel if end != END: @@ -1062,8 +1055,8 @@ class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): self, start: str, name: str, branch: Branch, *, with_reader: bool = True ) -> None: def get_writes( - packets: Sequence[Union[str, Send]], static: bool = False - ) -> Sequence[Union[ChannelWriteEntry, Send]]: + packets: Sequence[str | Send], static: bool = False + ) -> Sequence[ChannelWriteEntry | Send]: writes = [ ( ChannelWriteEntry( @@ -1094,7 +1087,7 @@ class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): mapper = _pick_mapper(channels, schema) self.schema_to_mapper[schema] = mapper # create reader - reader: Optional[Callable[[RunnableConfig], Any]] = partial( + reader: Callable[[RunnableConfig], Any] | None = partial( ChannelRead.do_read, select=channels[0] if channels == ["__root__"] else channels, fresh=True, @@ -1214,7 +1207,7 @@ class CompiledStateGraph(Pregel[InputT], Generic[StateT, InputT]): def _pick_mapper( state_keys: Sequence[str], schema: type[Any] -) -> Optional[Callable[[Any], Any]]: +) -> Callable[[Any], Any] | None: if state_keys == ["__root__"]: return None if isclass(schema) and issubclass(schema, dict): @@ -1255,8 +1248,8 @@ def _control_branch(value: Any) -> Sequence[tuple[str, Any]]: def _control_static( - ends: Union[tuple[str, ...], dict[str, str]], -) -> Sequence[tuple[str, Any, Optional[str]]]: + ends: tuple[str, ...] | dict[str, str], +) -> Sequence[tuple[str, Any, str | None]]: if isinstance(ends, dict): return [ (k if k == END else CHANNEL_BRANCH_TO.format(k), None, label) @@ -1268,7 +1261,7 @@ def _control_static( ] -def _get_root(input: Any) -> Optional[Sequence[tuple[str, Any]]]: +def _get_root(input: Any) -> Sequence[tuple[str, Any]] | None: if isinstance(input, Command): if input.graph == Command.PARENT: return () @@ -1323,12 +1316,12 @@ def _get_channel( @overload def _get_channel( name: str, annotation: Any, *, allow_managed: Literal[True] = True -) -> Union[BaseChannel, ManagedValueSpec]: ... +) -> BaseChannel | ManagedValueSpec: ... def _get_channel( name: str, annotation: Any, *, allow_managed: bool = True -) -> Union[BaseChannel, ManagedValueSpec]: +) -> BaseChannel | ManagedValueSpec: if manager := _is_field_managed_value(name, annotation): if allow_managed: return manager @@ -1346,7 +1339,7 @@ def _get_channel( return fallback -def _is_field_channel(typ: type[Any]) -> Optional[BaseChannel]: +def _is_field_channel(typ: type[Any]) -> BaseChannel | None: if hasattr(typ, "__metadata__"): meta = typ.__metadata__ if len(meta) >= 1 and isinstance(meta[-1], BaseChannel): @@ -1356,7 +1349,7 @@ def _is_field_channel(typ: type[Any]) -> Optional[BaseChannel]: return None -def _is_field_binop(typ: type[Any]) -> Optional[BinaryOperatorAggregate]: +def _is_field_binop(typ: type[Any]) -> BinaryOperatorAggregate | None: if hasattr(typ, "__metadata__"): meta = typ.__metadata__ if len(meta) >= 1 and callable(meta[-1]): @@ -1377,7 +1370,7 @@ def _is_field_binop(typ: type[Any]) -> Optional[BinaryOperatorAggregate]: return None -def _is_field_managed_value(name: str, typ: type[Any]) -> Optional[ManagedValueSpec]: +def _is_field_managed_value(name: str, typ: type[Any]) -> ManagedValueSpec | None: if hasattr(typ, "__metadata__"): meta = typ.__metadata__ if len(meta) >= 1: diff --git a/libs/langgraph/langgraph/graph/ui.py b/libs/langgraph/langgraph/graph/ui.py index a2e66f91a..181587740 100644 --- a/libs/langgraph/langgraph/graph/ui.py +++ b/libs/langgraph/langgraph/graph/ui.py @@ -1,4 +1,6 @@ -from typing import Any, Literal, Optional, Union, cast +from __future__ import annotations + +from typing import Any, Literal, Union, cast from uuid import uuid4 from langchain_core.messages import AnyMessage @@ -51,10 +53,10 @@ def push_ui_message( name: str, props: dict[str, Any], *, - id: Optional[str] = None, - metadata: Optional[dict[str, Any]] = None, - message: Optional[AnyMessage] = None, - state_key: Optional[str] = "ui", + id: str | None = None, + metadata: dict[str, Any] | None = None, + message: AnyMessage | None = None, + state_key: str | None = "ui", merge: bool = False, ) -> UIMessage: """Push a new UI message to update the UI state. @@ -149,8 +151,8 @@ def delete_ui_message(id: str, *, state_key: str = "ui") -> RemoveUIMessage: def ui_message_reducer( - left: Union[list[AnyUIMessage], AnyUIMessage], - right: Union[list[AnyUIMessage], AnyUIMessage], + left: list[AnyUIMessage] | AnyUIMessage, + right: list[AnyUIMessage] | AnyUIMessage, ) -> list[AnyUIMessage]: """Merge two lists of UI messages, supporting removing UI messages. diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index a4ea464ab..8fc14a0d1 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -2766,8 +2766,8 @@ class Pregel(PregelProtocol[InputT], Generic[InputT]): """ output_keys = output_keys if output_keys is not None else self.output_channels - latest: Union[dict[str, Any], Any] = None - chunks: list[Union[dict[str, Any], Any]] = [] + latest: dict[str, Any] | Any = None + chunks: list[dict[str, Any] | Any] = [] interrupts: list[Interrupt] = [] for chunk in self.stream( @@ -2833,8 +2833,8 @@ class Pregel(PregelProtocol[InputT], Generic[InputT]): output_keys = output_keys if output_keys is not None else self.output_channels - latest: Union[dict[str, Any], Any] = None - chunks: list[Union[dict[str, Any], Any]] = [] + latest: dict[str, Any] | Any = None + chunks: list[dict[str, Any] | Any] = [] interrupts: list[Interrupt] = [] async for chunk in self.astream( diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 1bd29e016..eb83de225 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import binascii import itertools import sys @@ -14,7 +16,6 @@ from typing import ( NamedTuple, Optional, Protocol, - Union, cast, overload, ) @@ -91,7 +92,7 @@ class WritesProtocol(Protocol): Implemented by PregelTaskWrites and PregelExecutableTask.""" @property - def path(self) -> tuple[Union[str, int, tuple], ...]: ... + def path(self) -> tuple[str | int | tuple, ...]: ... @property def name(self) -> str: ... @@ -107,7 +108,7 @@ class PregelTaskWrites(NamedTuple): """Simplest implementation of WritesProtocol, for usage with writes that don't originate from a runnable task, eg. graph input, update_state, etc.""" - path: tuple[Union[str, int, tuple], ...] + path: tuple[str | int | tuple, ...] name: str writes: Sequence[tuple[str, Any]] triggers: Sequence[str] @@ -118,8 +119,8 @@ class Call: func: Callable input: tuple[tuple[Any, ...], dict[str, Any]] - retry_policy: Optional[Sequence[RetryPolicy]] - cache_policy: Optional[CachePolicy] + retry_policy: Sequence[RetryPolicy] | None + cache_policy: CachePolicy | None callbacks: Callbacks def __init__( @@ -127,8 +128,8 @@ class Call: func: Callable, input: tuple[tuple[Any, ...], dict[str, Any]], *, - retry_policy: Optional[Sequence[RetryPolicy]], - cache_policy: Optional[CachePolicy], + retry_policy: Sequence[RetryPolicy] | None, + cache_policy: CachePolicy | None, callbacks: Callbacks, ) -> None: self.func = func @@ -140,7 +141,7 @@ class Call: def should_interrupt( checkpoint: Checkpoint, - interrupt_nodes: Union[All, Sequence[str]], + interrupt_nodes: All | Sequence[str], tasks: Iterable[PregelExecutableTask], ) -> list[PregelExecutableTask]: """Check if the graph should be interrupted based on current state.""" @@ -176,9 +177,9 @@ def local_read( channels: Mapping[str, BaseChannel], managed: ManagedValueMapping, task: WritesProtocol, - select: Union[list[str], str], + select: list[str] | str, fresh: bool = False, -) -> Union[dict[str, Any], Any]: +) -> dict[str, Any] | Any: """Function injected under CONFIG_KEY_READ in task config, to read current state. Used by conditional edges to read a copy of the state with reflecting the writes from that node only.""" @@ -213,7 +214,7 @@ def local_read( return values -def increment(current: Optional[int]) -> int: +def increment(current: int | None) -> int: """Default channel versioning function, increments the current int version.""" return current + 1 if current is not None else 1 @@ -222,7 +223,7 @@ def apply_writes( checkpoint: Checkpoint, channels: Mapping[str, BaseChannel], tasks: Iterable[WritesProtocol], - get_next_version: Optional[GetNextVersion], + get_next_version: GetNextVersion | None, trigger_to_nodes: Mapping[str, Sequence[str]], ) -> set[str]: """Apply writes from a set of tasks (usually the tasks from a Pregel step) @@ -338,8 +339,8 @@ def prepare_next_tasks( store: Literal[None] = None, checkpointer: Literal[None] = None, manager: Literal[None] = None, - trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None, - updated_channels: Optional[set[str]] = None, + trigger_to_nodes: Mapping[str, Sequence[str]] | None = None, + updated_channels: set[str] | None = None, retry_policy: Sequence[RetryPolicy] = (), cache_policy: Literal[None] = None, ) -> dict[str, PregelTask]: ... @@ -357,13 +358,13 @@ def prepare_next_tasks( stop: int, *, for_execution: Literal[True], - store: Optional[BaseStore], - checkpointer: Optional[BaseCheckpointSaver], - manager: Union[None, ParentRunManager, AsyncParentRunManager], - trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None, - updated_channels: Optional[set[str]] = None, + store: BaseStore | None, + checkpointer: BaseCheckpointSaver | None, + manager: None | ParentRunManager | AsyncParentRunManager, + trigger_to_nodes: Mapping[str, Sequence[str]] | None = None, + updated_channels: set[str] | None = None, retry_policy: Sequence[RetryPolicy] = (), - cache_policy: Optional[CachePolicy] = None, + cache_policy: CachePolicy | None = None, ) -> dict[str, PregelExecutableTask]: ... @@ -378,14 +379,14 @@ def prepare_next_tasks( stop: int, *, for_execution: bool, - store: Optional[BaseStore] = None, - checkpointer: Optional[BaseCheckpointSaver] = None, - manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, - trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None, - updated_channels: Optional[set[str]] = None, + store: BaseStore | None = None, + checkpointer: BaseCheckpointSaver | None = None, + manager: None | ParentRunManager | AsyncParentRunManager = None, + trigger_to_nodes: Mapping[str, Sequence[str]] | None = None, + updated_channels: set[str] | None = None, retry_policy: Sequence[RetryPolicy] = (), - cache_policy: Optional[CachePolicy] = None, -) -> Union[dict[str, PregelTask], dict[str, PregelExecutableTask]]: + cache_policy: CachePolicy | None = None, +) -> dict[str, PregelTask] | dict[str, PregelExecutableTask]: """Prepare the set of tasks that will make up the next Pregel step. Args: @@ -415,7 +416,7 @@ def prepare_next_tasks( input_cache: dict[INPUT_CACHE_KEY_TYPE, Any] = {} checkpoint_id_bytes = binascii.unhexlify(checkpoint["id"].replace("-", "")) null_version = checkpoint_null_version(checkpoint) - tasks: list[Union[PregelTask, PregelExecutableTask]] = [] + tasks: list[PregelTask | PregelExecutableTask] = [] # Consume pending tasks tasks_channel = cast(Optional[Topic[Send]], channels.get(TASKS)) if tasks_channel and tasks_channel.is_available(): @@ -496,11 +497,11 @@ PUSH_TRIGGER = (PUSH,) def prepare_single_task( task_path: tuple[Any, ...], - task_id_checksum: Optional[str], + task_id_checksum: str | None, *, checkpoint: Checkpoint, checkpoint_id_bytes: bytes, - checkpoint_null_version: Optional[V], + checkpoint_null_version: V | None, pending_writes: list[PendingWrite], processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], @@ -509,13 +510,13 @@ def prepare_single_task( step: int, stop: int, for_execution: bool, - store: Optional[BaseStore] = None, - checkpointer: Optional[BaseCheckpointSaver] = None, - manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, - input_cache: Optional[dict[INPUT_CACHE_KEY_TYPE, Any]] = None, - cache_policy: Optional[CachePolicy] = None, + store: BaseStore | None = None, + checkpointer: BaseCheckpointSaver | None = None, + manager: None | ParentRunManager | AsyncParentRunManager = None, + input_cache: dict[INPUT_CACHE_KEY_TYPE, Any] | None = None, + cache_policy: CachePolicy | None = None, retry_policy: Sequence[RetryPolicy] = (), -) -> Union[None, PregelTask, PregelExecutableTask]: +) -> None | PregelTask | PregelExecutableTask: """Prepares a single task for the next Pregel step, given a task path, which uniquely identifies a PUSH or PULL task within the graph.""" configurable = config.get(CONF, {}) @@ -560,7 +561,7 @@ def prepare_single_task( cache_policy = call.cache_policy or cache_policy if cache_policy: args_key = cache_policy.key_func(*call.input[0], **call.input[1]) - cache_key: Optional[CacheKey] = CacheKey( + cache_key: CacheKey | None = CacheKey( ( CACHE_NS_WRITES, (identifier(call.func) or "__dynamic__"), @@ -908,7 +909,7 @@ def prepare_single_task( def checkpoint_null_version( checkpoint: Checkpoint, -) -> Optional[V]: +) -> V | None: """Get the null version for the checkpoint, if available.""" for version in checkpoint["channel_versions"].values(): return type(version)() @@ -918,7 +919,7 @@ def checkpoint_null_version( def _triggers( channels: Mapping[str, BaseChannel], versions: ChannelVersions, - seen: Optional[ChannelVersions], + seen: ChannelVersions | None, null_version: V, proc: PregelNode, ) -> Sequence[str]: @@ -936,11 +937,11 @@ def _triggers( def _scratchpad( - parent_scratchpad: Optional[PregelScratchpad], + parent_scratchpad: PregelScratchpad | None, pending_writes: list[PendingWrite], task_id: str, namespace_hash: str, - resume_map: Optional[dict[str, Any]], + resume_map: dict[str, Any] | None, step: int, stop: int, ) -> PregelScratchpad: @@ -1010,7 +1011,7 @@ def _proc_input( *, for_execution: bool, scratchpad: PregelScratchpad, - input_cache: Optional[dict[INPUT_CACHE_KEY_TYPE, Any]], + input_cache: dict[INPUT_CACHE_KEY_TYPE, Any] | None, ) -> Any: """Prepare input for a PULL task, based on the process's channels and triggers.""" # if in cache return shallow copy @@ -1053,7 +1054,7 @@ def _proc_input( return val -def _uuid5_str(namespace: bytes, *parts: Union[str, bytes]) -> str: +def _uuid5_str(namespace: bytes, *parts: str | bytes) -> str: """Generate a UUID from the SHA-1 hash of a namespace and str parts.""" sha = sha1(namespace, usedforsecurity=False) @@ -1062,7 +1063,7 @@ def _uuid5_str(namespace: bytes, *parts: Union[str, bytes]) -> str: return f"{hex[:8]}-{hex[8:12]}-{hex[12:16]}-{hex[16:20]}-{hex[20:32]}" -def _xxhash_str(namespace: bytes, *parts: Union[str, bytes]) -> str: +def _xxhash_str(namespace: bytes, *parts: str | bytes) -> str: """Generate a UUID from the XXH3 hash of a namespace and str parts.""" hex = xxh3_128_hexdigest( namespace + b"".join(p.encode() if isinstance(p, str) else p for p in parts) @@ -1070,7 +1071,7 @@ def _xxhash_str(namespace: bytes, *parts: Union[str, bytes]) -> str: return f"{hex[:8]}-{hex[8:12]}-{hex[12:16]}-{hex[16:20]}-{hex[20:32]}" -def task_path_str(tup: Union[str, int, tuple]) -> str: +def task_path_str(tup: str | int | tuple) -> str: """Generate a string representation of the task path.""" return ( f"~{', '.join(task_path_str(x) for x in tup)}" @@ -1087,7 +1088,7 @@ LAZY_ATOMIC_COUNTER_LOCK = threading.Lock() class LazyAtomicCounter: __slots__ = ("_counter",) - _counter: Optional[Callable[[], int]] + _counter: Callable[[], int] | None def __init__(self) -> None: self._counter = None diff --git a/libs/langgraph/langgraph/pregel/call.py b/libs/langgraph/langgraph/pregel/call.py index 3cb39586b..df1e158cb 100644 --- a/libs/langgraph/langgraph/pregel/call.py +++ b/libs/langgraph/langgraph/pregel/call.py @@ -1,12 +1,14 @@ """Utility to convert a user provided function into a Runnable with a ChannelWrite.""" +from __future__ import annotations + import concurrent.futures import functools import inspect import sys import types from collections.abc import Generator, Sequence -from typing import Any, Callable, Generic, Optional, TypeVar, cast +from typing import Any, Callable, Generic, TypeVar, cast from langchain_core.runnables import Runnable from typing_extensions import ParamSpec @@ -40,7 +42,7 @@ def _getattribute(obj: Any, name: str) -> Any: return obj, parent -def _whichmodule(obj: Any, name: str) -> Optional[str]: +def _whichmodule(obj: Any, name: str) -> str | None: """Find the module an object belongs to. This function differs from ``pickle.whichmodule`` in two ways: @@ -74,7 +76,7 @@ def _whichmodule(obj: Any, name: str) -> Optional[str]: return None -def identifier(obj: Any, name: Optional[str] = None) -> Optional[str]: +def identifier(obj: Any, name: str | None = None) -> str | None: """Return the module and name of an object.""" from langgraph.pregel.read import PregelNode from langgraph.utils.runnable import RunnableCallable, RunnableSeq @@ -104,8 +106,8 @@ def identifier(obj: Any, name: Optional[str] = None) -> Optional[str]: def _lookup_module_and_qualname( - obj: Any, name: Optional[str] = None -) -> Optional[tuple[types.ModuleType, str]]: + obj: Any, name: str | None = None +) -> tuple[types.ModuleType, str] | None: if name is None: name = getattr(obj, "__qualname__", None) if name is None: # pragma: no cover @@ -251,8 +253,8 @@ class SyncAsyncFuture(Generic[T], concurrent.futures.Future[T]): def call( func: Callable[P, T], *args: Any, - retry_policy: Optional[Sequence[RetryPolicy]] = None, - cache_policy: Optional[CachePolicy] = None, + retry_policy: Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy | None = None, **kwargs: Any, ) -> SyncAsyncFuture[T]: config = get_config() diff --git a/libs/langgraph/langgraph/pregel/checkpoint.py b/libs/langgraph/langgraph/pregel/checkpoint.py index 0a0f95078..b8ca90db4 100644 --- a/libs/langgraph/langgraph/pregel/checkpoint.py +++ b/libs/langgraph/langgraph/pregel/checkpoint.py @@ -1,6 +1,7 @@ +from __future__ import annotations + from collections.abc import Mapping from datetime import datetime, timezone -from typing import Optional, Union from langgraph.channels.base import BaseChannel from langgraph.checkpoint.base import Checkpoint @@ -24,10 +25,10 @@ def empty_checkpoint() -> Checkpoint: def create_checkpoint( checkpoint: Checkpoint, - channels: Optional[Mapping[str, BaseChannel]], + channels: Mapping[str, BaseChannel] | None, step: int, *, - id: Optional[str] = None, + id: str | None = None, ) -> Checkpoint: """Create a checkpoint for the given channels.""" ts = datetime.now(timezone.utc).isoformat() @@ -52,7 +53,7 @@ def create_checkpoint( def channels_from_checkpoint( - specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], + specs: Mapping[str, BaseChannel | ManagedValueSpec], checkpoint: Checkpoint, ) -> tuple[Mapping[str, BaseChannel], ManagedValueMapping]: """Get channels from a checkpoint.""" diff --git a/libs/langgraph/langgraph/pregel/debug.py b/libs/langgraph/langgraph/pregel/debug.py index 6429cf0c4..1733d1ff3 100644 --- a/libs/langgraph/langgraph/pregel/debug.py +++ b/libs/langgraph/langgraph/pregel/debug.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from collections import defaultdict from collections.abc import Iterable, Iterator, Mapping, Sequence from dataclasses import asdict @@ -6,7 +8,6 @@ from pprint import pformat from typing import ( Any, Literal, - Optional, Union, ) from uuid import UUID @@ -43,7 +44,7 @@ class TaskPayload(TypedDict): class TaskResultPayload(TypedDict): id: str name: str - error: Optional[str] + error: str | None interrupts: list[dict] result: list[tuple[str, Any]] @@ -51,17 +52,17 @@ class TaskResultPayload(TypedDict): class CheckpointTask(TypedDict): id: str name: str - error: Optional[str] + error: str | None interrupts: list[dict] - state: Optional[RunnableConfig] + state: RunnableConfig | None class CheckpointPayload(TypedDict): - config: Optional[RunnableConfig] + config: RunnableConfig | None metadata: CheckpointMetadata values: dict[str, Any] next: list[str] - parent_config: Optional[RunnableConfig] + parent_config: RunnableConfig | None tasks: list[CheckpointTask] @@ -116,7 +117,7 @@ def map_debug_tasks( def map_debug_task_results( step: int, task_tup: tuple[PregelExecutableTask, Sequence[tuple[str, Any]]], - stream_keys: Union[str, Sequence[str]], + stream_keys: str | Sequence[str], ) -> Iterator[DebugOutputTaskResult]: """Produce "task_result" events for stream_mode=debug.""" stream_channels_list = ( @@ -144,7 +145,7 @@ def map_debug_task_results( } -def rm_pregel_keys(config: Optional[RunnableConfig]) -> Optional[RunnableConfig]: +def rm_pregel_keys(config: RunnableConfig | None) -> RunnableConfig | None: """Remove pregel-specific keys from the config.""" if config is None: return config @@ -161,18 +162,18 @@ def map_debug_checkpoint( step: int, config: RunnableConfig, channels: Mapping[str, BaseChannel], - stream_channels: Union[str, Sequence[str]], + stream_channels: str | Sequence[str], metadata: CheckpointMetadata, checkpoint: Checkpoint, tasks: Iterable[PregelExecutableTask], pending_writes: list[PendingWrite], - parent_config: Optional[RunnableConfig], - output_keys: Union[str, Sequence[str]], + parent_config: RunnableConfig | None, + output_keys: str | Sequence[str], ) -> Iterator[DebugOutputCheckpoint]: """Produce "checkpoint" events for stream_mode=debug.""" parent_ns = config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") - task_states: dict[str, Union[RunnableConfig, StateSnapshot]] = {} + task_states: dict[str, RunnableConfig | StateSnapshot] = {} for task in tasks: if not task.subgraphs: @@ -278,10 +279,10 @@ def print_step_checkpoint( def tasks_w_writes( - tasks: Iterable[Union[PregelTask, PregelExecutableTask]], - pending_writes: Optional[list[PendingWrite]], - states: Optional[dict[str, Union[RunnableConfig, StateSnapshot]]], - output_keys: Union[str, Sequence[str]], + tasks: Iterable[PregelTask | PregelExecutableTask], + pending_writes: list[PendingWrite] | None, + states: dict[str, RunnableConfig | StateSnapshot] | None, + output_keys: str | Sequence[str], ) -> tuple[PregelTask, ...]: """Apply writes / subgraph states to tasks to be returned in a StateSnapshot.""" pending_writes = pending_writes or [] diff --git a/libs/langgraph/langgraph/pregel/draw.py b/libs/langgraph/langgraph/pregel/draw.py index 2e7783f21..091e92be3 100644 --- a/libs/langgraph/langgraph/pregel/draw.py +++ b/libs/langgraph/langgraph/pregel/draw.py @@ -1,6 +1,8 @@ +from __future__ import annotations + from collections import defaultdict from collections.abc import Mapping, Sequence -from typing import Any, Optional, Union, cast +from typing import Any, cast from langchain_core.runnables.config import RunnableConfig from langchain_core.runnables.graph import Graph, Node @@ -26,10 +28,10 @@ def draw_graph( config: RunnableConfig, *, nodes: dict[str, PregelNode], - specs: dict[str, Union[BaseChannel, ManagedValueSpec]], - input_channels: Union[str, Sequence[str]], - interrupt_after_nodes: Union[All, Sequence[str]], - interrupt_before_nodes: Union[All, Sequence[str]], + specs: dict[str, BaseChannel | ManagedValueSpec], + input_channels: str | Sequence[str], + interrupt_after_nodes: All | Sequence[str], + interrupt_before_nodes: All | Sequence[str], trigger_to_nodes: Mapping[str, Sequence[str]], checkpointer: Checkpointer, subgraphs: dict[str, Graph], @@ -46,7 +48,7 @@ def draw_graph( The graph for this Pregel instance. """ # (src, dest, is_conditional, label) - edges: set[tuple[str, str, bool, Optional[str]]] = set() + edges: set[tuple[str, str, bool, str | None]] = set() step = -1 checkpoint = empty_checkpoint() @@ -60,8 +62,8 @@ def draw_graph( checkpoint, ) static_seen: set[Any] = set() - sources: dict[str, set[tuple[str, bool, Optional[str]]]] = {} - step_sources: dict[str, set[tuple[str, bool, Optional[str]]]] = {} + sources: dict[str, set[tuple[str, bool, str | None]]] = {} + step_sources: dict[str, set[tuple[str, bool, str | None]]] = {} # remove node mappers nodes = { k: v.copy(update={"mapper": None}) if v.mapper is not None else v @@ -100,7 +102,7 @@ def draw_graph( for step in range(step, limit): if not tasks: break - conditionals: dict[tuple[str, str, Any], Optional[str]] = {} + conditionals: dict[tuple[str, str, Any], str | None] = {} # run task writers for task in tasks.values(): for w in task.writers: @@ -140,8 +142,8 @@ def draw_graph( } sources.update(step_sources) # invert triggers - trigger_to_sources: dict[str, set[tuple[str, bool, Optional[str]]]] = ( - defaultdict(set) + trigger_to_sources: dict[str, set[tuple[str, bool, str | None]]] = defaultdict( + set ) for src, triggers in sources.items(): for trigger, cond, label in triggers: @@ -246,7 +248,7 @@ def add_edge( source: str, target: str, *, - data: Optional[Any] = None, + data: Any | None = None, conditional: bool = False, ) -> None: """Add an edge to the graph.""" diff --git a/libs/langgraph/langgraph/pregel/executor.py b/libs/langgraph/langgraph/pregel/executor.py index 488f21dcf..62df4b19f 100644 --- a/libs/langgraph/langgraph/pregel/executor.py +++ b/libs/langgraph/langgraph/pregel/executor.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import concurrent.futures import time @@ -7,7 +9,6 @@ from contextvars import copy_context from types import TracebackType from typing import ( Callable, - Optional, Protocol, TypeVar, cast, @@ -29,7 +30,7 @@ class Submit(Protocol[P, T]): self, fn: Callable[P, T], *args: P.args, - __name__: Optional[str] = None, + __name__: str | None = None, __cancel_on_exit__: bool = False, __reraise_on_exit__: bool = True, __next_tick__: bool = False, @@ -55,7 +56,7 @@ class BackgroundExecutor(AbstractContextManager): self, fn: Callable[P, T], *args: P.args, - __name__: Optional[str] = None, # currently not used in sync version + __name__: str | None = None, # currently not used in sync version __cancel_on_exit__: bool = False, # for sync, can cancel only if not started __reraise_on_exit__: bool = True, __next_tick__: bool = False, @@ -92,10 +93,10 @@ class BackgroundExecutor(AbstractContextManager): def __exit__( self, - exc_type: Optional[type[BaseException]], - exc_value: Optional[BaseException], - traceback: Optional[TracebackType], - ) -> Optional[bool]: + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: # copy the tasks as done() callback may modify the dict tasks = self.tasks.copy() # cancel all tasks that should be cancelled @@ -133,7 +134,7 @@ class AsyncBackgroundExecutor(AbstractAsyncContextManager): self.sentinel = object() self.loop = asyncio.get_running_loop() if max_concurrency := config.get("max_concurrency"): - self.semaphore: Optional[asyncio.Semaphore] = asyncio.Semaphore( + self.semaphore: asyncio.Semaphore | None = asyncio.Semaphore( max_concurrency ) else: @@ -143,7 +144,7 @@ class AsyncBackgroundExecutor(AbstractAsyncContextManager): self, fn: Callable[P, Awaitable[T]], *args: P.args, - __name__: Optional[str] = None, + __name__: str | None = None, __cancel_on_exit__: bool = False, __reraise_on_exit__: bool = True, __next_tick__: bool = False, # noop in async (always True) @@ -185,9 +186,9 @@ class AsyncBackgroundExecutor(AbstractAsyncContextManager): async def __aexit__( self, - exc_type: Optional[type[BaseException]], - exc_value: Optional[BaseException], - traceback: Optional[TracebackType], + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, ) -> None: # copy the tasks as done() callback may modify the dict tasks = self.tasks.copy() diff --git a/libs/langgraph/langgraph/pregel/io.py b/libs/langgraph/langgraph/pregel/io.py index 157be21f2..48268af76 100644 --- a/libs/langgraph/langgraph/pregel/io.py +++ b/libs/langgraph/langgraph/pregel/io.py @@ -1,6 +1,8 @@ +from __future__ import annotations + from collections import Counter from collections.abc import Iterator, Mapping, Sequence -from typing import Any, Literal, Optional, Union +from typing import Any, Literal from langgraph.channels.base import BaseChannel, EmptyChannelError from langgraph.constants import ( @@ -37,10 +39,10 @@ def read_channel( def read_channels( channels: Mapping[str, BaseChannel], - select: Union[Sequence[str], str], + select: Sequence[str] | str, *, skip_empty: bool = True, -) -> Union[dict[str, Any], Any]: +) -> dict[str, Any] | Any: if isinstance(select, str): return read_channel(channels, select) else: @@ -79,8 +81,8 @@ def map_command(cmd: Command) -> Iterator[tuple[str, str, Any]]: def map_input( - input_channels: Union[str, Sequence[str]], - chunk: Optional[Union[dict[str, Any], Any]], + input_channels: str | Sequence[str], + chunk: dict[str, Any] | Any | None, ) -> Iterator[tuple[str, Any]]: """Map input chunk to a sequence of pending writes in the form (channel, value).""" if chunk is None: @@ -98,10 +100,10 @@ def map_input( def map_output_values( - output_channels: Union[str, Sequence[str]], - pending_writes: Union[Literal[True], Sequence[tuple[str, Any]]], + output_channels: str | Sequence[str], + pending_writes: Literal[True] | Sequence[tuple[str, Any]], channels: Mapping[str, BaseChannel], -) -> Iterator[Union[dict[str, Any], Any]]: +) -> Iterator[dict[str, Any] | Any]: """Map pending writes (a sequence of tuples (channel, value)) to output chunk.""" if isinstance(output_channels, str): if pending_writes is True or any( @@ -116,10 +118,10 @@ def map_output_values( def map_output_updates( - output_channels: Union[str, Sequence[str]], + output_channels: str | Sequence[str], tasks: list[tuple[PregelExecutableTask, Sequence[tuple[str, Any]]]], cached: bool = False, -) -> Iterator[dict[str, Union[Any, dict[str, Any]]]]: +) -> Iterator[dict[str, Any | dict[str, Any]]]: """Map pending writes (a sequence of tuples (channel, value)) to output chunk.""" output_tasks = [ (t, ww) diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 80aaab9bc..3031ec525 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import binascii import concurrent.futures @@ -18,7 +20,6 @@ from typing import ( Literal, Optional, TypeVar, - Union, cast, ) @@ -148,36 +149,36 @@ def DuplexStream(*streams: StreamProtocol) -> StreamProtocol: class PregelLoop: config: RunnableConfig - store: Optional["BaseStore"] - stream: Optional[StreamProtocol] + store: BaseStore | None + stream: StreamProtocol | None step: int stop: int - input: Optional[Any] - input_model: Optional[type[BaseModel]] - cache: Optional[BaseCache[WritesT]] - checkpointer: Optional[BaseCheckpointSaver] + input: Any | None + input_model: type[BaseModel] | None + cache: BaseCache[WritesT] | None + checkpointer: BaseCheckpointSaver | None nodes: Mapping[str, PregelNode] - specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]] - output_keys: Union[str, Sequence[str]] - stream_keys: Union[str, Sequence[str]] + specs: Mapping[str, BaseChannel | ManagedValueSpec] + output_keys: str | Sequence[str] + stream_keys: str | Sequence[str] skip_done_tasks: bool is_nested: bool - manager: Union[None, AsyncParentRunManager, ParentRunManager] - interrupt_after: Union[All, Sequence[str]] - interrupt_before: Union[All, Sequence[str]] + manager: None | AsyncParentRunManager | ParentRunManager + interrupt_after: All | Sequence[str] + interrupt_before: All | Sequence[str] checkpoint_during: bool debug: bool retry_policy: Sequence[RetryPolicy] - cache_policy: Optional[CachePolicy] + cache_policy: CachePolicy | None checkpointer_get_next_version: GetNextVersion - checkpointer_put_writes: Optional[Callable[[RunnableConfig, WritesT, str], Any]] + checkpointer_put_writes: Callable[[RunnableConfig, WritesT, str], Any] | None checkpointer_put_writes_accepts_task_path: bool - _checkpointer_put_after_previous: Optional[ + _checkpointer_put_after_previous: ( Callable[ [ - Optional[concurrent.futures.Future], + concurrent.futures.Future | None, RunnableConfig, Checkpoint, str, @@ -185,8 +186,9 @@ class PregelLoop: ], Any, ] - ] - _migrate_checkpoint: Optional[Callable[[Checkpoint], None]] + | None + ) + _migrate_checkpoint: Callable[[Checkpoint], None] | None submit: Submit channels: Mapping[str, BaseChannel] managed: ManagedValueMapping @@ -196,40 +198,40 @@ class PregelLoop: checkpoint_config: RunnableConfig checkpoint_metadata: CheckpointMetadata checkpoint_pending_writes: list[PendingWrite] - checkpoint_previous_versions: dict[str, Union[str, float, int]] - prev_checkpoint_config: Optional[RunnableConfig] + checkpoint_previous_versions: dict[str, str | float | int] + prev_checkpoint_config: RunnableConfig | None status: Literal[ "pending", "done", "interrupt_before", "interrupt_after", "out_of_steps" ] tasks: dict[str, PregelExecutableTask] to_interrupt: list[PregelExecutableTask] - output: Union[None, dict[str, Any], Any] = None + output: None | dict[str, Any] | Any = None # public def __init__( self, - input: Optional[Any], + input: Any | None, *, - stream: Optional[StreamProtocol], + stream: StreamProtocol | None, config: RunnableConfig, - store: Optional[BaseStore], - cache: Optional[BaseCache], - checkpointer: Optional[BaseCheckpointSaver], + store: BaseStore | None, + cache: BaseCache | None, + checkpointer: BaseCheckpointSaver | None, nodes: Mapping[str, PregelNode], - specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], - output_keys: Union[str, Sequence[str]], - stream_keys: Union[str, Sequence[str]], + specs: Mapping[str, BaseChannel | ManagedValueSpec], + output_keys: str | Sequence[str], + stream_keys: str | Sequence[str], trigger_to_nodes: Mapping[str, Sequence[str]], - interrupt_after: Union[All, Sequence[str]] = EMPTY_SEQ, - interrupt_before: Union[All, Sequence[str]] = EMPTY_SEQ, - manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, - input_model: Optional[type[BaseModel]] = None, + interrupt_after: All | Sequence[str] = EMPTY_SEQ, + interrupt_before: All | Sequence[str] = EMPTY_SEQ, + manager: None | AsyncParentRunManager | ParentRunManager = None, + input_model: type[BaseModel] | None = None, debug: bool = False, - migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None, + migrate_checkpoint: Callable[[Checkpoint], None] | None = None, retry_policy: Sequence[RetryPolicy] = (), - cache_policy: Optional[CachePolicy] = None, + cache_policy: CachePolicy | None = None, checkpoint_during: bool = True, ) -> None: self.stream = stream @@ -261,7 +263,7 @@ class PregelLoop: self.debug = debug if self.stream is not None and CONFIG_KEY_STREAM in config[CONF]: self.stream = DuplexStream(self.stream, config[CONF][CONFIG_KEY_STREAM]) - scratchpad: Optional[PregelScratchpad] = config[CONF].get(CONFIG_KEY_SCRATCHPAD) + scratchpad: PregelScratchpad | None = config[CONF].get(CONFIG_KEY_SCRATCHPAD) if not self.config[CONF].get(CONFIG_KEY_DELEGATE) and isinstance( scratchpad, PregelScratchpad ): @@ -399,8 +401,8 @@ class PregelLoop: ) def accept_push( - self, task: PregelExecutableTask, write_idx: int, call: Optional[Call] = None - ) -> Optional[PregelExecutableTask]: + self, task: PregelExecutableTask, write_idx: int, call: Call | None = None + ) -> PregelExecutableTask | None: """Accept a PUSH from a task, potentially returning a new task to start.""" # don't start if we should interrupt *after* the original task if self.interrupt_after and should_interrupt( @@ -455,7 +457,7 @@ class PregelLoop: def tick( self, *, - input_keys: Union[str, Sequence[str]], + input_keys: str | Sequence[str], ) -> bool: """Execute a single iteration of the Pregel loop. @@ -649,7 +651,7 @@ class PregelLoop: else: task.writes.append((k, v)) - def _first(self, *, input_keys: Union[str, Sequence[str]]) -> Optional[set[str]]: + def _first(self, *, input_keys: str | Sequence[str]) -> set[str] | None: # resuming from previous checkpoint requires # - finding a previous checkpoint # - receiving None input (outer graph) or RESUMING flag (subgraph) @@ -667,7 +669,7 @@ class PregelLoop: ) ) # this can be set only when there are input_writes - updated_channels: Optional[set[str]] = None + updated_channels: set[str] | None = None # map command to writes if isinstance(self.input, Command): @@ -861,10 +863,10 @@ class PregelLoop: def _suppress_interrupt( self, - exc_type: Optional[type[BaseException]], - exc_value: Optional[BaseException], - traceback: Optional[TracebackType], - ) -> Optional[bool]: + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: # persist current checkpoint and writes if not self.checkpoint_during: self._put_checkpoint(self.checkpoint_metadata) @@ -977,26 +979,26 @@ class PregelLoop: class SyncPregelLoop(PregelLoop, AbstractContextManager): def __init__( self, - input: Optional[Any], + input: Any | None, *, - stream: Optional[StreamProtocol], + stream: StreamProtocol | None, config: RunnableConfig, - store: Optional[BaseStore], - cache: Optional[BaseCache], - checkpointer: Optional[BaseCheckpointSaver], + store: BaseStore | None, + cache: BaseCache | None, + checkpointer: BaseCheckpointSaver | None, nodes: Mapping[str, PregelNode], - specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], + specs: Mapping[str, BaseChannel | ManagedValueSpec], trigger_to_nodes: Mapping[str, Sequence[str]], - manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, - interrupt_after: Union[All, Sequence[str]] = EMPTY_SEQ, - interrupt_before: Union[All, Sequence[str]] = EMPTY_SEQ, - output_keys: Union[str, Sequence[str]] = EMPTY_SEQ, - stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ, - input_model: Optional[type[BaseModel]] = None, + manager: None | AsyncParentRunManager | ParentRunManager = None, + interrupt_after: All | Sequence[str] = EMPTY_SEQ, + interrupt_before: All | Sequence[str] = EMPTY_SEQ, + output_keys: str | Sequence[str] = EMPTY_SEQ, + stream_keys: str | Sequence[str] = EMPTY_SEQ, + input_model: type[BaseModel] | None = None, debug: bool = False, - migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None, + migrate_checkpoint: Callable[[Checkpoint], None] | None = None, retry_policy: Sequence[RetryPolicy] = (), - cache_policy: Optional[CachePolicy] = None, + cache_policy: CachePolicy | None = None, checkpoint_during: bool = True, ) -> None: super().__init__( @@ -1037,7 +1039,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): def _checkpointer_put_after_previous( self, - prev: Optional[concurrent.futures.Future], + prev: concurrent.futures.Future | None, config: RunnableConfig, checkpoint: Checkpoint, metadata: CheckpointMetadata, @@ -1067,8 +1069,8 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): return matched def accept_push( - self, task: PregelExecutableTask, write_idx: int, call: Optional[Call] = None - ) -> Optional[PregelExecutableTask]: + self, task: PregelExecutableTask, write_idx: int, call: Call | None = None + ) -> PregelExecutableTask | None: if pushed := super().accept_push(task, write_idx, call): for task in self.match_cached_writes(): self.output_writes(task.id, task.writes, cached=True) @@ -1156,10 +1158,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): def __exit__( self, - exc_type: Optional[type[BaseException]], - exc_value: Optional[BaseException], - traceback: Optional[TracebackType], - ) -> Optional[bool]: + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: # unwind stack return self.stack.__exit__(exc_type, exc_value, traceback) @@ -1167,26 +1169,26 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): def __init__( self, - input: Optional[Any], + input: Any | None, *, - stream: Optional[StreamProtocol], + stream: StreamProtocol | None, config: RunnableConfig, - store: Optional[BaseStore], - cache: Optional[BaseCache], - checkpointer: Optional[BaseCheckpointSaver], + store: BaseStore | None, + cache: BaseCache | None, + checkpointer: BaseCheckpointSaver | None, nodes: Mapping[str, PregelNode], - specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], + specs: Mapping[str, BaseChannel | ManagedValueSpec], trigger_to_nodes: Mapping[str, Sequence[str]], - interrupt_after: Union[All, Sequence[str]] = EMPTY_SEQ, - interrupt_before: Union[All, Sequence[str]] = EMPTY_SEQ, - manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, - output_keys: Union[str, Sequence[str]] = EMPTY_SEQ, - stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ, - input_model: Optional[type[BaseModel]] = None, + interrupt_after: All | Sequence[str] = EMPTY_SEQ, + interrupt_before: All | Sequence[str] = EMPTY_SEQ, + manager: None | AsyncParentRunManager | ParentRunManager = None, + output_keys: str | Sequence[str] = EMPTY_SEQ, + stream_keys: str | Sequence[str] = EMPTY_SEQ, + input_model: type[BaseModel] | None = None, debug: bool = False, - migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None, + migrate_checkpoint: Callable[[Checkpoint], None] | None = None, retry_policy: Sequence[RetryPolicy] = (), - cache_policy: Optional[CachePolicy] = None, + cache_policy: CachePolicy | None = None, checkpoint_during: bool = True, ) -> None: super().__init__( @@ -1227,7 +1229,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): async def _checkpointer_put_after_previous( self, - prev: Optional[asyncio.Task], + prev: asyncio.Task | None, config: RunnableConfig, checkpoint: Checkpoint, metadata: CheckpointMetadata, @@ -1257,8 +1259,8 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): return matched async def aaccept_push( - self, task: PregelExecutableTask, write_idx: int, call: Optional[Call] = None - ) -> Optional[PregelExecutableTask]: + self, task: PregelExecutableTask, write_idx: int, call: Call | None = None + ) -> PregelExecutableTask | None: if pushed := super().accept_push(task, write_idx, call): for task in await self.amatch_cached_writes(): self.output_writes(task.id, task.writes, cached=True) @@ -1352,10 +1354,10 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): async def __aexit__( self, - exc_type: Optional[type[BaseException]], - exc_value: Optional[BaseException], - traceback: Optional[TracebackType], - ) -> Optional[bool]: + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: # unwind stack exit_task = asyncio.create_task( self.stack.__aexit__(exc_type, exc_value, traceback) diff --git a/libs/langgraph/langgraph/pregel/messages.py b/libs/langgraph/langgraph/pregel/messages.py index 5014d1cbb..9a9210aa6 100644 --- a/libs/langgraph/langgraph/pregel/messages.py +++ b/libs/langgraph/langgraph/pregel/messages.py @@ -1,10 +1,10 @@ +from __future__ import annotations + from collections.abc import AsyncIterator, Iterator, Sequence from typing import ( Any, Callable, - Optional, TypeVar, - Union, cast, ) from uuid import UUID, uuid4 @@ -36,7 +36,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): self.stream = stream self.subgraphs = subgraphs self.metadata: dict[UUID, Meta] = {} - self.seen: set[Union[int, str]] = set() + self.seen: set[int | str] = set() def _emit(self, meta: Meta, message: BaseMessage, *, dedupe: bool = False) -> None: if dedupe and message.id in self.seen: @@ -89,9 +89,9 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): messages: list[list[BaseMessage]], *, run_id: UUID, - parent_run_id: Optional[UUID] = None, - tags: Optional[list[str]] = None, - metadata: Optional[dict[str, Any]] = None, + parent_run_id: UUID | None = None, + tags: list[str] | None = None, + metadata: dict[str, Any] | None = None, **kwargs: Any, ) -> Any: if metadata and ( @@ -111,10 +111,10 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): self, token: str, *, - chunk: Optional[ChatGenerationChunk] = None, + chunk: ChatGenerationChunk | None = None, run_id: UUID, - parent_run_id: Optional[UUID] = None, - tags: Optional[list[str]] = None, + parent_run_id: UUID | None = None, + tags: list[str] | None = None, **kwargs: Any, ) -> Any: if not isinstance(chunk, ChatGenerationChunk): @@ -127,7 +127,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): response: LLMResult, *, run_id: UUID, - parent_run_id: Optional[UUID] = None, + parent_run_id: UUID | None = None, **kwargs: Any, ) -> Any: if meta := self.metadata.get(run_id): @@ -142,7 +142,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): error: BaseException, *, run_id: UUID, - parent_run_id: Optional[UUID] = None, + parent_run_id: UUID | None = None, **kwargs: Any, ) -> Any: self.metadata.pop(run_id, None) @@ -153,9 +153,9 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): inputs: dict[str, Any], *, run_id: UUID, - parent_run_id: Optional[UUID] = None, - tags: Optional[list[str]] = None, - metadata: Optional[dict[str, Any]] = None, + parent_run_id: UUID | None = None, + tags: list[str] | None = None, + metadata: dict[str, Any] | None = None, **kwargs: Any, ) -> Any: if ( @@ -185,7 +185,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): response: Any, *, run_id: UUID, - parent_run_id: Optional[UUID] = None, + parent_run_id: UUID | None = None, **kwargs: Any, ) -> Any: if meta := self.metadata.pop(run_id, None): @@ -210,7 +210,7 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): error: BaseException, *, run_id: UUID, - parent_run_id: Optional[UUID] = None, + parent_run_id: UUID | None = None, **kwargs: Any, ) -> Any: self.metadata.pop(run_id, None) diff --git a/libs/langgraph/langgraph/pregel/protocol.py b/libs/langgraph/langgraph/pregel/protocol.py index 391bf2ba5..9441f4514 100644 --- a/libs/langgraph/langgraph/pregel/protocol.py +++ b/libs/langgraph/langgraph/pregel/protocol.py @@ -2,7 +2,7 @@ from __future__ import annotations from abc import ABC, abstractmethod from collections.abc import AsyncIterator, Iterator, Sequence -from typing import Any, Generic, Optional, Union +from typing import Any, Generic from langchain_core.runnables import Runnable, RunnableConfig from langchain_core.runnables.graph import Graph as DrawableGraph @@ -16,23 +16,23 @@ from langgraph.typing import InputT class PregelProtocol(Runnable[InputT, Any], Generic[InputT], ABC): @abstractmethod def with_config( - self, config: Optional[RunnableConfig] = None, **kwargs: Any + self, config: RunnableConfig | None = None, **kwargs: Any ) -> Self: ... @abstractmethod def get_graph( self, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, - xray: Union[int, bool] = False, + xray: int | bool = False, ) -> DrawableGraph: ... @abstractmethod async def aget_graph( self, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, - xray: Union[int, bool] = False, + xray: int | bool = False, ) -> DrawableGraph: ... @abstractmethod @@ -50,9 +50,9 @@ class PregelProtocol(Runnable[InputT, Any], Generic[InputT], ABC): self, config: RunnableConfig, *, - 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[StateSnapshot]: ... @abstractmethod @@ -60,9 +60,9 @@ class PregelProtocol(Runnable[InputT, Any], Generic[InputT], ABC): self, config: RunnableConfig, *, - 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[StateSnapshot]: ... @abstractmethod @@ -83,58 +83,58 @@ class PregelProtocol(Runnable[InputT, Any], Generic[InputT], ABC): def update_state( self, config: RunnableConfig, - values: Optional[Union[dict[str, Any], Any]], - as_node: Optional[str] = None, + values: dict[str, Any] | Any | None, + as_node: str | None = None, ) -> RunnableConfig: ... @abstractmethod async def aupdate_state( self, config: RunnableConfig, - values: Optional[Union[dict[str, Any], Any]], - as_node: Optional[str] = None, + values: dict[str, Any] | Any | None, + as_node: str | None = None, ) -> RunnableConfig: ... @abstractmethod def stream( self, input: InputT, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, - stream_mode: Optional[Union[StreamMode, list[StreamMode]]] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, + stream_mode: StreamMode | list[StreamMode] | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, subgraphs: bool = False, - ) -> Iterator[Union[dict[str, Any], Any]]: ... + ) -> Iterator[dict[str, Any] | Any]: ... @abstractmethod def astream( self, input: InputT, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, - stream_mode: Optional[Union[StreamMode, list[StreamMode]]] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, + stream_mode: StreamMode | list[StreamMode] | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, subgraphs: bool = False, - ) -> AsyncIterator[Union[dict[str, Any], Any]]: ... + ) -> AsyncIterator[dict[str, Any] | Any]: ... @abstractmethod def invoke( self, input: InputT, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - ) -> Union[dict[str, Any], Any]: ... + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + ) -> dict[str, Any] | Any: ... @abstractmethod async def ainvoke( self, input: InputT, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - ) -> Union[dict[str, Any], Any]: ... + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + ) -> dict[str, Any] | Any: ... diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index 1f20f96ab..38b334e46 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -1,10 +1,10 @@ +from __future__ import annotations + from collections.abc import AsyncIterator, Iterator, Sequence from dataclasses import asdict from typing import ( Any, Literal, - Optional, - Union, cast, ) @@ -96,20 +96,20 @@ class RemoteGraph(PregelProtocol): """ assistant_id: str - name: Optional[str] + name: str | None def __init__( self, assistant_id: str, # graph_id /, *, - url: Optional[str] = None, - api_key: Optional[str] = None, - headers: Optional[dict[str, str]] = None, - client: Optional[LangGraphClient] = None, - sync_client: Optional[SyncLangGraphClient] = None, - config: Optional[RunnableConfig] = None, - name: Optional[str] = None, + url: str | None = None, + api_key: str | None = None, + headers: dict[str, str] | None = None, + client: LangGraphClient | None = None, + sync_client: SyncLangGraphClient | None = None, + config: RunnableConfig | None = None, + name: str | None = None, ): """Specify `url`, `api_key`, and/or `headers` to create default sync and async clients. @@ -162,9 +162,7 @@ class RemoteGraph(PregelProtocol): attrs = {**self.__dict__, **update} return self.__class__(attrs.pop("assistant_id"), **attrs) - def with_config( - self, config: Optional[RunnableConfig] = None, **kwargs: Any - ) -> Self: + def with_config(self, config: RunnableConfig | None = None, **kwargs: Any) -> Self: return self.copy( {"config": merge_configs(self.config, config, cast(RunnableConfig, kwargs))} ) @@ -195,9 +193,9 @@ class RemoteGraph(PregelProtocol): def get_graph( self, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, - xray: Union[int, bool] = False, + xray: int | bool = False, ) -> DrawableGraph: """Get graph by graph name. @@ -224,9 +222,9 @@ class RemoteGraph(PregelProtocol): async def aget_graph( self, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, - xray: Union[int, bool] = False, + xray: int | bool = False, ) -> DrawableGraph: """Get graph by graph name. @@ -309,7 +307,7 @@ class RemoteGraph(PregelProtocol): interrupts=tuple([i for task in tasks for i in task.interrupts]), ) - def _get_checkpoint(self, config: Optional[RunnableConfig]) -> Optional[Checkpoint]: + def _get_checkpoint(self, config: RunnableConfig | None) -> Checkpoint | None: if config is None: return None @@ -423,9 +421,9 @@ class RemoteGraph(PregelProtocol): self, config: RunnableConfig, *, - 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[StateSnapshot]: """Get the state history of a thread. @@ -458,9 +456,9 @@ class RemoteGraph(PregelProtocol): self, config: RunnableConfig, *, - 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[StateSnapshot]: """Get the state history of a thread. @@ -492,22 +490,22 @@ class RemoteGraph(PregelProtocol): def bulk_update_state( self, config: RunnableConfig, - updates: list[tuple[Optional[dict[str, Any]], Optional[str]]], + updates: list[tuple[dict[str, Any] | None, str | None]], ) -> RunnableConfig: raise NotImplementedError async def abulk_update_state( self, config: RunnableConfig, - updates: list[tuple[Optional[dict[str, Any]], Optional[str]]], + updates: list[tuple[dict[str, Any] | None, str | None]], ) -> RunnableConfig: raise NotImplementedError def update_state( self, config: RunnableConfig, - values: Optional[Union[dict[str, Any], Any]], - as_node: Optional[str] = None, + values: dict[str, Any] | Any | None, + as_node: str | None = None, ) -> RunnableConfig: """Update the state of a thread. @@ -536,8 +534,8 @@ class RemoteGraph(PregelProtocol): async def aupdate_state( self, config: RunnableConfig, - values: Optional[Union[dict[str, Any], Any]], - as_node: Optional[str] = None, + values: dict[str, Any] | Any | None, + as_node: str | None = None, ) -> RunnableConfig: """Update the state of a thread. @@ -565,12 +563,10 @@ class RemoteGraph(PregelProtocol): def _get_stream_modes( self, - stream_mode: Optional[Union[StreamMode, list[StreamMode]]], - config: Optional[RunnableConfig], + stream_mode: StreamMode | list[StreamMode] | None, + config: RunnableConfig | None, default: StreamMode = "updates", - ) -> tuple[ - list[StreamModeSDK], list[StreamModeSDK], bool, Optional[StreamProtocol] - ]: + ) -> tuple[list[StreamModeSDK], list[StreamModeSDK], bool, StreamProtocol | None]: """Return a tuple of the final list of stream modes sent to the remote graph and a boolean flag indicating if stream mode 'updates' was present in the original list of stream modes. @@ -591,7 +587,7 @@ class RemoteGraph(PregelProtocol): updated_stream_modes.append(default) requested_stream_modes = updated_stream_modes.copy() # add any from parent graph - stream: Optional[StreamProtocol] = ( + stream: StreamProtocol | None = ( (config or {}).get(CONF, {}).get(CONFIG_KEY_STREAM) ) if stream: @@ -618,15 +614,15 @@ class RemoteGraph(PregelProtocol): def stream( self, - input: Union[dict[str, Any], Any], - config: Optional[RunnableConfig] = None, + input: dict[str, Any] | Any, + config: RunnableConfig | None = None, *, - stream_mode: Optional[Union[StreamMode, list[StreamMode]]] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, + stream_mode: StreamMode | list[StreamMode] | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, subgraphs: bool = False, **kwargs: Any, - ) -> Iterator[Union[dict[str, Any], Any]]: + ) -> Iterator[dict[str, Any] | Any]: """Create a run and stream the results. This method calls `POST /threads/{thread_id}/runs/stream` if a `thread_id` @@ -652,7 +648,7 @@ class RemoteGraph(PregelProtocol): stream_mode, config ) if isinstance(input, Command): - command: Optional[CommandSDK] = cast(CommandSDK, asdict(input)) + command: CommandSDK | None = cast(CommandSDK, asdict(input)) input = None else: command = None @@ -717,15 +713,15 @@ class RemoteGraph(PregelProtocol): async def astream( self, - input: Union[dict[str, Any], Any], - config: Optional[RunnableConfig] = None, + input: dict[str, Any] | Any, + config: RunnableConfig | None = None, *, - stream_mode: Optional[Union[StreamMode, list[StreamMode]]] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, + stream_mode: StreamMode | list[StreamMode] | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, subgraphs: bool = False, **kwargs: Any, - ) -> AsyncIterator[Union[dict[str, Any], Any]]: + ) -> AsyncIterator[dict[str, Any] | Any]: """Create a run and stream the results. This method calls `POST /threads/{thread_id}/runs/stream` if a `thread_id` @@ -751,7 +747,7 @@ class RemoteGraph(PregelProtocol): stream_mode, config ) if isinstance(input, Command): - command: Optional[CommandSDK] = cast(CommandSDK, asdict(input)) + command: CommandSDK | None = cast(CommandSDK, asdict(input)) input = None else: command = None @@ -817,28 +813,28 @@ class RemoteGraph(PregelProtocol): async def astream_events( self, input: Any, - config: Optional[RunnableConfig] = None, + config: RunnableConfig | None = None, *, version: Literal["v1", "v2"], - include_names: Optional[Sequence[All]] = None, - include_types: Optional[Sequence[All]] = None, - include_tags: Optional[Sequence[All]] = None, - exclude_names: Optional[Sequence[All]] = None, - exclude_types: Optional[Sequence[All]] = None, - exclude_tags: Optional[Sequence[All]] = None, + include_names: Sequence[All] | None = None, + include_types: Sequence[All] | None = None, + include_tags: Sequence[All] | None = None, + exclude_names: Sequence[All] | None = None, + exclude_types: Sequence[All] | None = None, + exclude_tags: Sequence[All] | None = None, **kwargs: Any, ) -> AsyncIterator[dict[str, Any]]: raise NotImplementedError def invoke( self, - input: Union[dict[str, Any], Any], - config: Optional[RunnableConfig] = None, + input: dict[str, Any] | Any, + config: RunnableConfig | None = None, *, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, **kwargs: Any, - ) -> Union[dict[str, Any], Any]: + ) -> dict[str, Any] | Any: """Create a run, wait until it finishes and return the final state. Args: @@ -867,13 +863,13 @@ class RemoteGraph(PregelProtocol): async def ainvoke( self, - input: Union[dict[str, Any], Any], - config: Optional[RunnableConfig] = None, + input: dict[str, Any] | Any, + config: RunnableConfig | None = None, *, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, **kwargs: Any, - ) -> Union[dict[str, Any], Any]: + ) -> dict[str, Any] | Any: """Create a run, wait until it finishes and return the final state. Args: diff --git a/libs/langgraph/langgraph/pregel/retry.py b/libs/langgraph/langgraph/pregel/retry.py index dd4987fd4..a91edb9c4 100644 --- a/libs/langgraph/langgraph/pregel/retry.py +++ b/libs/langgraph/langgraph/pregel/retry.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import logging import random @@ -5,7 +7,7 @@ import sys import time from collections.abc import Awaitable, Sequence from dataclasses import replace -from typing import Any, Callable, Optional +from typing import Any, Callable from langgraph.constants import ( CONF, @@ -23,8 +25,8 @@ SUPPORTS_EXC_NOTES = sys.version_info >= (3, 11) def run_with_retry( task: PregelExecutableTask, - retry_policy: Optional[Sequence[RetryPolicy]], - configurable: Optional[dict[str, Any]] = None, + retry_policy: Sequence[RetryPolicy] | None, + configurable: dict[str, Any] | None = None, ) -> None: """Run a task with retries.""" retry_policy = task.retry_policy or retry_policy @@ -104,12 +106,11 @@ def run_with_retry( async def arun_with_retry( task: PregelExecutableTask, - retry_policy: Optional[Sequence[RetryPolicy]], + retry_policy: Sequence[RetryPolicy] | None, stream: bool = False, - match_cached_writes: Optional[ - Callable[[], Awaitable[Sequence[PregelExecutableTask]]] - ] = None, - configurable: Optional[dict[str, Any]] = None, + match_cached_writes: Callable[[], Awaitable[Sequence[PregelExecutableTask]]] + | None = None, + configurable: dict[str, Any] | None = None, ) -> None: """Run a task asynchronously with retries.""" retry_policy = task.retry_policy or retry_policy diff --git a/libs/langgraph/langgraph/pregel/runner.py b/libs/langgraph/langgraph/pregel/runner.py index d17463b3e..47afc199d 100644 --- a/libs/langgraph/langgraph/pregel/runner.py +++ b/libs/langgraph/langgraph/pregel/runner.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import concurrent.futures import threading @@ -56,16 +58,14 @@ EXCLUDED_FRAME_FNAMES = ( "concurrent/futures/_base.py", ) -SKIP_RERAISE_SET: weakref.WeakSet[Union[concurrent.futures.Future, asyncio.Future]] = ( +SKIP_RERAISE_SET: weakref.WeakSet[concurrent.futures.Future | asyncio.Future] = ( weakref.WeakSet() ) class FuturesDict(Generic[F, E], dict[F, Optional[PregelExecutableTask]]): event: E - callback: weakref.ref[ - Callable[[PregelExecutableTask, Optional[BaseException]], None] - ] + callback: weakref.ref[Callable[[PregelExecutableTask, BaseException | None], None]] counter: int done: set[F] lock: threading.Lock @@ -74,7 +74,7 @@ class FuturesDict(Generic[F, E], dict[F, Optional[PregelExecutableTask]]): self, event: E, callback: weakref.ref[ - Callable[[PregelExecutableTask, Optional[BaseException]], None] + Callable[[PregelExecutableTask, BaseException | None], None] ], future_type: type[F], # used for generic typing, newer py supports FutureDict[...](...) @@ -89,7 +89,7 @@ class FuturesDict(Generic[F, E], dict[F, Optional[PregelExecutableTask]]): def __setitem__( self, key: F, - value: Optional[PregelExecutableTask], + value: PregelExecutableTask | None, ) -> None: super().__setitem__(key, value) # type: ignore[index] if value is not None: @@ -124,7 +124,7 @@ class PregelRunner: submit: weakref.ref[Submit], put_writes: weakref.ref[Callable[[str, Sequence[tuple[str, Any]]], None]], use_astream: bool = False, - node_finished: Optional[Callable[[str], None]] = None, + node_finished: Callable[[str], None] | None = None, ) -> None: self.submit = submit self.put_writes = put_writes @@ -136,12 +136,12 @@ class PregelRunner: tasks: Iterable[PregelExecutableTask], *, reraise: bool = True, - timeout: Optional[float] = None, - retry_policy: Optional[Sequence[RetryPolicy]] = None, - get_waiter: Optional[Callable[[], concurrent.futures.Future[None]]] = None, + timeout: float | None = None, + retry_policy: Sequence[RetryPolicy] | None = None, + get_waiter: Callable[[], concurrent.futures.Future[None]] | None = None, schedule_task: Callable[ - [PregelExecutableTask, int, Optional[Call]], - Optional[PregelExecutableTask], + [PregelExecutableTask, int, Call | None], + PregelExecutableTask | None, ], ) -> Iterator[None]: tasks = tuple(tasks) @@ -268,12 +268,12 @@ class PregelRunner: tasks: Iterable[PregelExecutableTask], *, reraise: bool = True, - timeout: Optional[float] = None, - retry_policy: Optional[Sequence[RetryPolicy]] = None, - get_waiter: Optional[Callable[[], asyncio.Future[None]]] = None, + timeout: float | None = None, + retry_policy: Sequence[RetryPolicy] | None = None, + get_waiter: Callable[[], asyncio.Future[None]] | None = None, schedule_task: Callable[ - [PregelExecutableTask, int, Optional[Call]], - Awaitable[Optional[PregelExecutableTask]], + [PregelExecutableTask, int, Call | None], + Awaitable[PregelExecutableTask | None], ], ) -> AsyncIterator[None]: loop = asyncio.get_event_loop() @@ -415,7 +415,7 @@ class PregelRunner: def commit( self, task: PregelExecutableTask, - exception: Optional[BaseException], + exception: BaseException | None, ) -> None: if isinstance(exception, asyncio.CancelledError): # for cancelled tasks, also save error in task, @@ -465,8 +465,8 @@ def _should_stop_others( def _exception( - fut: Union[concurrent.futures.Future[Any], asyncio.Future[Any]], -) -> Optional[BaseException]: + fut: concurrent.futures.Future[Any] | asyncio.Future[Any], +) -> BaseException | None: """Return the exception from a future, without raising CancelledError.""" if fut.cancelled(): if isinstance(fut, asyncio.Future): @@ -478,14 +478,14 @@ def _exception( def _panic_or_proceed( - futs: Union[set[concurrent.futures.Future], set[asyncio.Future]], + futs: set[concurrent.futures.Future] | set[asyncio.Future], *, timeout_exc_cls: type[Exception] = TimeoutError, panic: bool = True, ) -> None: """Cancel remaining tasks if any failed, re-raise exception if panic is True.""" - done: set[Union[concurrent.futures.Future[Any], asyncio.Future[Any]]] = set() - inflight: set[Union[concurrent.futures.Future[Any], asyncio.Future[Any]]] = set() + done: set[concurrent.futures.Future[Any] | asyncio.Future[Any]] = set() + inflight: set[concurrent.futures.Future[Any] | asyncio.Future[Any]] = set() for fut in futs: if fut.cancelled(): continue @@ -522,22 +522,22 @@ def _panic_or_proceed( def _call( task: weakref.ref[PregelExecutableTask], - func: Callable[[Any], Union[Awaitable[Any], Any]], + func: Callable[[Any], Awaitable[Any] | Any], input: Any, *, - retry_policy: Optional[Sequence[RetryPolicy]] = None, - cache_policy: Optional[CachePolicy] = None, + retry_policy: Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy | None = None, callbacks: Callbacks = None, futures: weakref.ref[FuturesDict], schedule_task: Callable[ - [PregelExecutableTask, int, Optional[Call]], Optional[PregelExecutableTask] + [PregelExecutableTask, int, Call | None], PregelExecutableTask | None ], submit: weakref.ref[Submit], ) -> concurrent.futures.Future[Any]: if asyncio.iscoroutinefunction(func): raise RuntimeError("In an sync context async tasks cannot be called") - fut: Optional[concurrent.futures.Future] = None + fut: concurrent.futures.Future | None = None # schedule PUSH tasks, collect futures scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr] # schedule the next task, if the callback returns one @@ -609,22 +609,22 @@ def _call( def _acall( task: weakref.ref[PregelExecutableTask], - func: Callable[[Any], Union[Awaitable[Any], Any]], + func: Callable[[Any], Awaitable[Any] | Any], input: Any, *, - retry_policy: Optional[Sequence[RetryPolicy]] = None, - cache_policy: Optional[CachePolicy] = None, + retry_policy: Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy | None = None, callbacks: Callbacks = None, # injected dependencies futures: weakref.ref[FuturesDict], schedule_task: Callable[ - [PregelExecutableTask, int, Optional[Call]], - Awaitable[Optional[PregelExecutableTask]], + [PregelExecutableTask, int, Call | None], + Awaitable[PregelExecutableTask | None], ], submit: weakref.ref[Submit], loop: asyncio.AbstractEventLoop, stream: bool = False, -) -> Union[asyncio.Future[Any], concurrent.futures.Future[Any]]: +) -> asyncio.Future[Any] | concurrent.futures.Future[Any]: # return a chained future to ensure commit() callback is called # before the returned future is resolved, to ensure stream order etc try: @@ -633,8 +633,8 @@ def _acall( in_async = False # if in async context return an async future, otherwise return a sync future if in_async: - fut: Union[asyncio.Future[Any], concurrent.futures.Future[Any]] = ( - asyncio.Future(loop=loop) + fut: asyncio.Future[Any] | concurrent.futures.Future[Any] = asyncio.Future( + loop=loop ) else: fut = concurrent.futures.Future() @@ -661,26 +661,26 @@ def _acall( async def _acall_impl( - destination: Union[asyncio.Future[Any], concurrent.futures.Future[Any]], + destination: asyncio.Future[Any] | concurrent.futures.Future[Any], task: weakref.ref[PregelExecutableTask], - func: Callable[[Any], Union[Awaitable[Any], Any]], + func: Callable[[Any], Awaitable[Any] | Any], input: Any, *, - retry_policy: Optional[Sequence[RetryPolicy]] = None, - cache_policy: Optional[CachePolicy] = None, + retry_policy: Sequence[RetryPolicy] | None = None, + cache_policy: CachePolicy | None = None, callbacks: Callbacks = None, # injected dependencies futures: weakref.ref[FuturesDict[asyncio.Future, asyncio.Event]], schedule_task: Callable[ - [PregelExecutableTask, int, Optional[Call]], - Awaitable[Optional[PregelExecutableTask]], + [PregelExecutableTask, int, Call | None], + Awaitable[PregelExecutableTask | None], ], submit: weakref.ref[Submit], loop: asyncio.AbstractEventLoop, stream: bool = False, ) -> None: try: - fut: Optional[asyncio.Future] = None + fut: asyncio.Future | None = None # schedule PUSH tasks, collect futures scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr] # schedule the next task, if the callback returns one diff --git a/libs/langgraph/langgraph/pregel/utils.py b/libs/langgraph/langgraph/pregel/utils.py index da0076f6b..a37228c05 100644 --- a/libs/langgraph/langgraph/pregel/utils.py +++ b/libs/langgraph/langgraph/pregel/utils.py @@ -1,8 +1,10 @@ +from __future__ import annotations + import ast import inspect import re import textwrap -from typing import Any, Callable, Optional +from typing import Any, Callable from langchain_core.runnables import RunnableLambda, RunnableSequence from typing_extensions import override @@ -30,7 +32,7 @@ def get_new_channel_versions( return new_versions -def find_subgraph_pregel(candidate: Runnable) -> Optional[PregelProtocol]: +def find_subgraph_pregel(candidate: Runnable) -> PregelProtocol | None: from langgraph.pregel import Pregel candidates: list[Runnable] = [candidate] diff --git a/libs/langgraph/langgraph/pregel/validate.py b/libs/langgraph/langgraph/pregel/validate.py index 169268169..acf174dba 100644 --- a/libs/langgraph/langgraph/pregel/validate.py +++ b/libs/langgraph/langgraph/pregel/validate.py @@ -1,5 +1,7 @@ +from __future__ import annotations + from collections.abc import Mapping, Sequence -from typing import Any, Optional, Union +from typing import Any from langgraph.channels.base import BaseChannel from langgraph.constants import RESERVED @@ -10,11 +12,11 @@ from langgraph.types import All def validate_graph( nodes: Mapping[str, PregelNode], channels: dict[str, BaseChannel], - input_channels: Union[str, Sequence[str]], - output_channels: Union[str, Sequence[str]], - stream_channels: Optional[Union[str, Sequence[str]]], - interrupt_after_nodes: Union[All, Sequence[str]], - interrupt_before_nodes: Union[All, Sequence[str]], + input_channels: str | Sequence[str], + output_channels: str | Sequence[str], + stream_channels: str | Sequence[str] | None, + interrupt_after_nodes: All | Sequence[str], + interrupt_before_nodes: All | Sequence[str], ) -> None: for chan in channels: if chan in RESERVED: @@ -88,7 +90,7 @@ def validate_graph( def validate_keys( - keys: Optional[Union[str, Sequence[str]]], + keys: str | Sequence[str] | None, channels: Mapping[str, Any], ) -> None: if isinstance(keys, str): diff --git a/libs/langgraph/langgraph/pregel/write.py b/libs/langgraph/langgraph/pregel/write.py index ee386ad51..7f1fdc73f 100644 --- a/libs/langgraph/langgraph/pregel/write.py +++ b/libs/langgraph/langgraph/pregel/write.py @@ -40,7 +40,7 @@ class ChannelWriteTupleEntry(NamedTuple): """Function to extract tuples from value.""" value: Any = PASSTHROUGH """Value to write, or PASSTHROUGH to use the input.""" - static: Optional[Sequence[tuple[str, Any, Optional[str]]]] = None + static: Sequence[tuple[str, Any, str | None]] | None = None """Optional, declared writes for static analysis.""" @@ -138,7 +138,7 @@ class ChannelWrite(RunnableCallable): @staticmethod def get_static_writes( runnable: Runnable, - ) -> Optional[Sequence[tuple[str, Any, Optional[str]]]]: + ) -> Sequence[tuple[str, Any, str | None]] | None: """Used to get conditional writes a writer declares for static analysis.""" if isinstance(runnable, ChannelWrite): return [ @@ -160,9 +160,7 @@ class ChannelWrite(RunnableCallable): @staticmethod def register_writer( runnable: R, - static: Optional[ - Sequence[tuple[Union[ChannelWriteEntry, Send], Optional[str]]] - ] = None, + static: Sequence[tuple[ChannelWriteEntry | Send, str | None]] | None = None, ) -> R: """Used to mark a runnable as a writer, so that it can be detected by is_writer. Instances of ChannelWrite are automatically marked as writers. @@ -174,7 +172,7 @@ class ChannelWrite(RunnableCallable): def _assemble_writes( - writes: Sequence[Union[ChannelWriteEntry, ChannelWriteTupleEntry, Send]], + writes: Sequence[ChannelWriteEntry | ChannelWriteTupleEntry | Send], ) -> list[tuple[str, Any]]: """Assembles the writes into a list of tuples.""" tuples: list[tuple[str, Any]] = [] diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index d6db21db9..ddcca374f 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import dataclasses import sys from collections import deque @@ -10,7 +12,6 @@ from typing import ( Generic, Literal, NamedTuple, - Optional, TypeVar, Union, cast, @@ -115,9 +116,9 @@ class RetryPolicy(NamedTuple): """Maximum number of attempts to make before giving up, including the first.""" jitter: bool = True """Whether to add random jitter to the interval between retries.""" - retry_on: Union[ - type[Exception], Sequence[type[Exception]], Callable[[Exception], bool] - ] = default_retry_on + retry_on: ( + type[Exception] | Sequence[type[Exception]] | Callable[[Exception], bool] + ) = default_retry_on """List of exception classes that should trigger a retry, or a callable that returns True for exceptions that should trigger a retry.""" @@ -132,7 +133,7 @@ class CachePolicy(Generic[KeyFuncT]): """Function to generate a cache key from the node's input. Defaults to hashing the input with pickle.""" - ttl: Optional[int] = None + ttl: int | None = None """Time to live for the cache entry in seconds. If None, the entry never expires.""" @@ -145,7 +146,7 @@ class Interrupt: value: Any resumable: bool = False - ns: Optional[Sequence[str]] = None + ns: Sequence[str] | None = None when: Literal["during"] = dataclasses.field(default="during", repr=False) @property @@ -157,8 +158,8 @@ class Interrupt: class StateUpdate(NamedTuple): - values: Optional[dict[str, Any]] - as_node: Optional[str] = None + values: dict[str, Any] | None + as_node: str | None = None class PregelTask(NamedTuple): @@ -166,11 +167,11 @@ class PregelTask(NamedTuple): id: str name: str - path: tuple[Union[str, int, tuple], ...] - error: Optional[Exception] = None + path: tuple[str | int | tuple, ...] + error: Exception | None = None interrupts: tuple[Interrupt, ...] = () - state: Union[None, RunnableConfig, "StateSnapshot"] = None - result: Optional[Any] = None + state: None | RunnableConfig | StateSnapshot = None + result: Any | None = None if sys.version_info > (3, 11): @@ -186,7 +187,7 @@ class CacheKey(NamedTuple): """Namespace for the cache entry.""" key: str """Key for the cache entry.""" - ttl: Optional[int] + ttl: int | None """Time to live for the cache entry in seconds.""" @@ -199,28 +200,28 @@ class PregelExecutableTask: config: RunnableConfig triggers: Sequence[str] retry_policy: Sequence[RetryPolicy] - cache_key: Optional[CacheKey] + cache_key: CacheKey | None id: str - path: tuple[Union[str, int, tuple], ...] + path: tuple[str | int | tuple, ...] scheduled: bool = False writers: Sequence[Runnable] = () - subgraphs: Sequence["PregelProtocol"] = () + subgraphs: Sequence[PregelProtocol] = () class StateSnapshot(NamedTuple): """Snapshot of the state of the graph at the beginning of a step.""" - values: Union[dict[str, Any], Any] + values: dict[str, Any] | Any """Current values of channels.""" next: tuple[str, ...] """The name of the node to execute in each task for this step.""" config: RunnableConfig """Config used to fetch this snapshot.""" - metadata: Optional[CheckpointMetadata] + metadata: CheckpointMetadata | None """Metadata associated with this snapshot.""" - created_at: Optional[str] + created_at: str | None """Timestamp of snapshot creation.""" - parent_config: Optional[RunnableConfig] + parent_config: RunnableConfig | None """Config used to fetch the parent snapshot, if any.""" tasks: tuple[PregelTask, ...] """Tasks to execute in this step. If already attempted, may contain an error.""" @@ -327,10 +328,10 @@ class Command(Generic[N], ToolOutputMixin): - sequence of `Send` objects """ - graph: Optional[str] = None - update: Optional[Any] = None - resume: Optional[Union[dict[str, Any], Any]] = None - goto: Union[Send, Sequence[Union[Send, N]], N] = () + graph: str | None = None + update: Any | None = None + resume: dict[str, Any] | Any | None = None + goto: Send | Sequence[Send | N] | N = () def __repr__(self) -> str: # get all non-None values diff --git a/libs/langgraph/langgraph/utils/config.py b/libs/langgraph/langgraph/utils/config.py index 51259c353..ace98cd91 100644 --- a/libs/langgraph/langgraph/utils/config.py +++ b/libs/langgraph/langgraph/utils/config.py @@ -1,7 +1,9 @@ +from __future__ import annotations + from collections import ChainMap from collections.abc import Sequence from os import getenv -from typing import Any, Optional, cast +from typing import Any, cast from langchain_core.callbacks import ( AsyncCallbackManager, @@ -45,7 +47,7 @@ def recast_checkpoint_ns(ns: str) -> str: def patch_configurable( - config: Optional[RunnableConfig], patch: dict[str, Any] + config: RunnableConfig | None, patch: dict[str, Any] ) -> RunnableConfig: if config is None: return {CONF: patch} @@ -56,7 +58,7 @@ def patch_configurable( def patch_checkpoint_map( - config: Optional[RunnableConfig], metadata: Optional[CheckpointMetadata] + config: RunnableConfig | None, metadata: CheckpointMetadata | None ) -> RunnableConfig: if config is None: return config @@ -75,7 +77,7 @@ def patch_checkpoint_map( return config -def merge_configs(*configs: Optional[RunnableConfig]) -> RunnableConfig: +def merge_configs(*configs: RunnableConfig | None) -> RunnableConfig: """Merge multiple configs into one. Args: @@ -148,13 +150,13 @@ def merge_configs(*configs: Optional[RunnableConfig]) -> RunnableConfig: def patch_config( - config: Optional[RunnableConfig], + config: RunnableConfig | None, *, callbacks: Callbacks = None, - recursion_limit: Optional[int] = None, - max_concurrency: Optional[int] = None, - run_name: Optional[str] = None, - configurable: Optional[dict[str, Any]] = None, + recursion_limit: int | None = None, + max_concurrency: int | None = None, + run_name: str | None = None, + configurable: dict[str, Any] | None = None, ) -> RunnableConfig: """Patch a config with new values. @@ -194,7 +196,7 @@ def patch_config( def get_callback_manager_for_config( - config: RunnableConfig, tags: Optional[Sequence[str]] = None + config: RunnableConfig, tags: Sequence[str] | None = None ) -> CallbackManager: """Get a callback manager for a config. @@ -232,7 +234,7 @@ def get_callback_manager_for_config( def get_async_callback_manager_for_config( config: RunnableConfig, - tags: Optional[Sequence[str]] = None, + tags: Sequence[str] | None = None, ) -> AsyncCallbackManager: """Get an async callback manager for a config. @@ -275,7 +277,7 @@ def _is_not_empty(value: Any) -> bool: return value is not None -def ensure_config(*configs: Optional[RunnableConfig]) -> RunnableConfig: +def ensure_config(*configs: RunnableConfig | None) -> RunnableConfig: """Return a config with all keys, merging any provided configs. Args: diff --git a/libs/langgraph/langgraph/utils/fields.py b/libs/langgraph/langgraph/utils/fields.py index 4f8cd78cc..33a8f2951 100644 --- a/libs/langgraph/langgraph/utils/fields.py +++ b/libs/langgraph/langgraph/utils/fields.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import dataclasses import types import weakref @@ -31,7 +33,7 @@ def _is_optional_type(type_: Any) -> bool: return type_ is None -def _is_required_type(type_: Any) -> Optional[bool]: +def _is_required_type(type_: Any) -> bool | None: """Check if an annotation is marked as Required/NotRequired. Returns: @@ -118,7 +120,7 @@ def get_field_default(name: str, type_: Any, schema: type[Any]) -> Any: def get_enhanced_type_hints( type: type[Any], -) -> Generator[tuple[str, Any, Any, Optional[str]], None, None]: +) -> Generator[tuple[str, Any, Any, str | None], None, None]: """Attempt to extract default values and descriptions from provided type, used for config schema.""" for name, typ in get_type_hints(type).items(): default = None diff --git a/libs/langgraph/langgraph/utils/future.py b/libs/langgraph/langgraph/utils/future.py index c69ff688e..38face045 100644 --- a/libs/langgraph/langgraph/utils/future.py +++ b/libs/langgraph/langgraph/utils/future.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import concurrent.futures import contextvars @@ -5,7 +7,7 @@ import inspect import sys import types from collections.abc import Awaitable, Coroutine, Generator -from typing import Optional, TypeVar, Union, cast +from typing import TypeVar, Union, cast T = TypeVar("T") AnyFuture = Union[asyncio.Future, concurrent.futures.Future] @@ -139,11 +141,11 @@ def chain_future(source: AnyFuture, destination: AnyFuture) -> AnyFuture: def _ensure_future( - coro_or_future: Union[Coroutine[None, None, T], Awaitable[T]], + coro_or_future: Coroutine[None, None, T] | Awaitable[T], *, loop: asyncio.AbstractEventLoop, - name: Optional[str] = None, - context: Optional[contextvars.Context] = None, + name: str | None = None, + context: contextvars.Context | None = None, lazy: bool = True, ) -> asyncio.Task[T]: called_wrap_awaitable = False @@ -189,8 +191,8 @@ def run_coroutine_threadsafe( loop: asyncio.AbstractEventLoop, *, lazy: bool, - name: Optional[str] = None, - context: Optional[contextvars.Context] = None, + name: str | None = None, + context: contextvars.Context | None = None, ) -> asyncio.Future[T]: """Submit a coroutine object to a given event loop. diff --git a/libs/langgraph/langgraph/utils/pydantic.py b/libs/langgraph/langgraph/utils/pydantic.py index a935aa16b..0d93f085e 100644 --- a/libs/langgraph/langgraph/utils/pydantic.py +++ b/libs/langgraph/langgraph/utils/pydantic.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import sys import typing import warnings @@ -6,8 +8,6 @@ from dataclasses import is_dataclass from functools import lru_cache from typing import ( Any, - Optional, - Union, cast, overload, ) @@ -39,7 +39,7 @@ def get_fields(model: BaseModel) -> dict[str, FieldInfo]: ... def get_fields( - model: Union[type[BaseModel], BaseModel], + model: type[BaseModel] | BaseModel, ) -> dict[str, FieldInfo]: """Get the field names of a Pydantic model.""" if hasattr(model, "model_fields"): @@ -61,7 +61,7 @@ NO_DEFAULT = object() def _create_root_model( name: str, type_: Any, - module_name: Optional[str] = None, + module_name: str | None = None, default_: object = NO_DEFAULT, ) -> type[BaseModel]: """Create a base class.""" @@ -115,7 +115,7 @@ def _create_root_model_cached( model_name: str, type_: Any, *, - module_name: Optional[str] = None, + module_name: str | None = None, default_: object = NO_DEFAULT, ) -> type[BaseModel]: return _create_root_model( @@ -181,8 +181,8 @@ def _remap_field_definitions(field_definitions: dict[str, Any]) -> dict[str, Any def create_model( model_name: str, *, - field_definitions: Optional[dict[str, Any]] = None, - root: Optional[Any] = None, + field_definitions: dict[str, Any] | None = None, + root: Any | None = None, ) -> type[BaseModel]: """Create a pydantic model with the given field definitions. diff --git a/libs/langgraph/langgraph/utils/queue.py b/libs/langgraph/langgraph/utils/queue.py index b68e15322..14ff7875b 100644 --- a/libs/langgraph/langgraph/utils/queue.py +++ b/libs/langgraph/langgraph/utils/queue.py @@ -1,4 +1,5 @@ # type: ignore +from __future__ import annotations import asyncio import queue @@ -7,7 +8,6 @@ import threading import types from collections import deque from time import monotonic -from typing import Optional PY_310 = sys.version_info >= (3, 10) @@ -50,7 +50,7 @@ class AsyncQueue(asyncio.Queue): class Semaphore(threading.Semaphore): """Semaphore subclass with a wait() method.""" - def wait(self, blocking: bool = True, timeout: Optional[float] = None): + def wait(self, blocking: bool = True, timeout: float | None = None): """Block until the semaphore can be acquired, but don't acquire it.""" if not blocking and timeout is not None: raise ValueError("can't specify timeout for non-blocking acquire") diff --git a/libs/langgraph/langgraph/utils/runnable.py b/libs/langgraph/langgraph/utils/runnable.py index 76699c947..3ca48aec3 100644 --- a/libs/langgraph/langgraph/utils/runnable.py +++ b/libs/langgraph/langgraph/utils/runnable.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import enum import inspect @@ -63,7 +65,7 @@ except ImportError: def _set_config_context( config: RunnableConfig, run: Any = None -) -> Token[Optional[RunnableConfig]]: +) -> Token[RunnableConfig | None]: """Set the child Runnable config + tracing context. Args: @@ -77,9 +79,7 @@ def _set_config_context( return config_token -def _unset_config_context( - token: Token[Optional[RunnableConfig]], run: Any = None -) -> None: +def _unset_config_context(token: Token[RunnableConfig | None], run: Any = None) -> None: """Set the child Runnable config + tracing context. Args: @@ -242,15 +242,15 @@ class RunnableCallable(Runnable): def __init__( self, - func: Optional[Callable[..., Union[Any, Runnable]]], - afunc: Optional[Callable[..., Awaitable[Union[Any, Runnable]]]] = None, + func: Callable[..., Any | Runnable] | None, + afunc: Callable[..., Awaitable[Any | Runnable]] | None = None, *, - name: Optional[str] = None, - tags: Optional[Sequence[str]] = None, + name: str | None = None, + tags: Sequence[str] | None = None, trace: bool = True, recurse: bool = True, explode_args: bool = False, - func_accepts_config: Optional[bool] = None, + func_accepts_config: bool | None = None, **kwargs: Any, ) -> None: self.name = name @@ -312,7 +312,7 @@ class RunnableCallable(Runnable): return f"{self.get_name()}({', '.join(f'{k}={v!r}' for k, v in repr_args.items())})" def invoke( - self, input: Any, config: Optional[RunnableConfig] = None, **kwargs: Any + self, input: Any, config: RunnableConfig | None = None, **kwargs: Any ) -> Any: if self.func is None: raise TypeError( @@ -380,7 +380,7 @@ class RunnableCallable(Runnable): return ret async def ainvoke( - self, input: Any, config: Optional[RunnableConfig] = None, **kwargs: Any + self, input: Any, config: RunnableConfig | None = None, **kwargs: Any ) -> Any: if not self.afunc: return self.invoke(input, config) @@ -466,7 +466,7 @@ def is_async_generator( def coerce_to_runnable( - thing: RunnableLike, *, name: Optional[str], trace: bool + thing: RunnableLike, *, name: str | None, trace: bool ) -> Runnable: """Coerce a runnable-like object into a Runnable. @@ -509,8 +509,8 @@ class RunnableSeq(Runnable): def __init__( self, *steps: RunnableLike, - name: Optional[str] = None, - trace_inputs: Optional[Callable[[Any], Any]] = None, + name: str | None = None, + trace_inputs: Callable[[Any], Any] | None = None, ) -> None: """Create a new RunnableSeq. @@ -588,7 +588,7 @@ class RunnableSeq(Runnable): ) def invoke( - self, input: Input, config: Optional[RunnableConfig] = None, **kwargs: Any + self, input: Input, config: RunnableConfig | None = None, **kwargs: Any ) -> Any: if config is None: config = ensure_config() @@ -634,8 +634,8 @@ class RunnableSeq(Runnable): async def ainvoke( self, input: Input, - config: Optional[RunnableConfig] = None, - **kwargs: Optional[Any], + config: RunnableConfig | None = None, + **kwargs: Any | None, ) -> Any: if config is None: config = ensure_config() @@ -687,8 +687,8 @@ class RunnableSeq(Runnable): def stream( self, input: Input, - config: Optional[RunnableConfig] = None, - **kwargs: Optional[Any], + config: RunnableConfig | None = None, + **kwargs: Any | None, ) -> Iterator[Any]: if config is None: config = ensure_config() @@ -747,8 +747,8 @@ class RunnableSeq(Runnable): async def astream( self, input: Input, - config: Optional[RunnableConfig] = None, - **kwargs: Optional[Any], + config: RunnableConfig | None = None, + **kwargs: Any | None, ) -> AsyncIterator[Any]: if config is None: config = ensure_config() diff --git a/libs/langgraph/pyproject.toml b/libs/langgraph/pyproject.toml index 34fe40a48..04037001e 100644 --- a/libs/langgraph/pyproject.toml +++ b/libs/langgraph/pyproject.toml @@ -63,12 +63,15 @@ langgraph-sdk = { path = "../sdk-py", editable = true } [tool.ruff] lint.select = [ "E", "F", "I", "TID251", "UP" ] -lint.ignore = [ "E501", "UP007" ] +lint.ignore = [ "E501" ] line-length = 88 indent-width = 4 extend-include = ["*.ipynb"] target-version = "py39" +[tool.ruff.lint.per-file-ignores] +"tests/bench/*" = ["UP006", "UP007"] + [tool.ruff.format] quote-style = "double" indent-style = "space" diff --git a/libs/langgraph/tests/test_runnable.py b/libs/langgraph/tests/test_runnable.py index 0a81be368..1dd8e35cc 100644 --- a/libs/langgraph/tests/test_runnable.py +++ b/libs/langgraph/tests/test_runnable.py @@ -85,7 +85,7 @@ def test_runnable_callable_injectable_arguments() -> None: """ # Test Optional[BaseStore] annotation. - def func_optional_store(inputs: Any, store: Optional[BaseStore]) -> str: + def func_optional_store(inputs: Any, store: Optional[BaseStore]) -> str: # noqa: UP007 """Test function that accepts an optional store parameter.""" assert store is None return "success" @@ -159,12 +159,12 @@ async def test_runnable_callable_injectable_arguments_async() -> None: """ # Test Optional[BaseStore] annotation. - def func_optional_store(inputs: Any, store: Optional[BaseStore]) -> str: + def func_optional_store(inputs: Any, store: Optional[BaseStore]) -> str: # noqa: UP007 """Test function that accepts an optional store parameter.""" assert store is None return "success" - async def afunc_optional_store(inputs: Any, store: Optional[BaseStore]) -> str: + async def afunc_optional_store(inputs: Any, store: BaseStore | None) -> str: """Async version of func_optional_store.""" assert store is None return "success" diff --git a/libs/sdk-py/langgraph_sdk/auth/__init__.py b/libs/sdk-py/langgraph_sdk/auth/__init__.py index 01dfd9143..7ebab9719 100644 --- a/libs/sdk-py/langgraph_sdk/auth/__init__.py +++ b/libs/sdk-py/langgraph_sdk/auth/__init__.py @@ -175,7 +175,7 @@ class Auth: # will be considered a breaking change. self._handlers: dict[tuple[str, str], list[types.Handler]] = {} self._global_handlers: list[types.Handler] = [] - self._authenticate_handler: typing.Optional[types.Authenticator] = None + self._authenticate_handler: types.Authenticator | None = None self._handler_cache: dict[tuple[str, str], types.Handler] = {} def authenticate(self, fn: AH) -> AH: @@ -301,7 +301,7 @@ class _ResourceOn(typing.Generic[VCreate, VRead, VUpdate, VDelete, VSearch]): Generic base class for resource-specific handlers. """ - value: type[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]] + value: type[VCreate | VUpdate | VRead | VDelete | VSearch] Create: type[VCreate] Read: type[VRead] @@ -335,40 +335,36 @@ class _ResourceOn(typing.Generic[VCreate, VRead, VUpdate, VDelete, VSearch]): @typing.overload def __call__( self, - fn: typing.Union[ - _ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]], - _ActionHandler[dict[str, typing.Any]], - ], - ) -> _ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]]: ... + fn: _ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch] + | _ActionHandler[dict[str, typing.Any]], + ) -> _ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch]: ... @typing.overload def __call__( self, *, - resources: typing.Union[str, Sequence[str]], - actions: typing.Optional[typing.Union[str, Sequence[str]]] = None, + resources: str | Sequence[str], + actions: str | Sequence[str] | None = None, ) -> Callable[ - [_ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]]], - _ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]], + [_ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch]], + _ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch], ]: ... def __call__( self, - fn: typing.Union[ - _ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]], - _ActionHandler[dict[str, typing.Any]], - None, - ] = None, + fn: _ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch] + | _ActionHandler[dict[str, typing.Any]] + | None = None, *, - resources: typing.Union[str, Sequence[str], None] = None, - actions: typing.Optional[typing.Union[str, Sequence[str]]] = None, - ) -> typing.Union[ - _ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]], - Callable[ - [_ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]]], - _ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]], - ], - ]: + resources: str | Sequence[str] | None = None, + actions: str | Sequence[str] | None = None, + ) -> ( + _ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch] + | Callable[ + [_ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch]], + _ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch], + ] + ): if fn is not None: _validate_handler(fn) return typing.cast( @@ -377,10 +373,8 @@ class _ResourceOn(typing.Generic[VCreate, VRead, VUpdate, VDelete, VSearch]): ) def decorator( - handler: _ActionHandler[ - typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch] - ], - ) -> _ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]]: + handler: _ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch], + ) -> _ActionHandler[VCreate | VUpdate | VRead | VDelete | VSearch]: _validate_handler(handler) return typing.cast( _ActionHandler[typing.Union[VCreate, VUpdate, VRead, VDelete, VSearch]], @@ -482,14 +476,9 @@ class _StoreOn: def __call__( self, *, - actions: typing.Optional[ - typing.Union[ - typing.Literal["put", "get", "search", "list_namespaces", "delete"], - Sequence[ - typing.Literal["put", "get", "search", "list_namespaces", "delete"] - ], - ] - ] = None, + actions: typing.Literal["put", "get", "search", "list_namespaces", "delete"] + | Sequence[typing.Literal["put", "get", "search", "list_namespaces", "delete"]] + | None = None, ) -> Callable[[AHO], AHO]: ... @typing.overload @@ -497,17 +486,12 @@ class _StoreOn: def __call__( self, - fn: typing.Optional[AHO] = None, + fn: AHO | None = None, *, - actions: typing.Optional[ - typing.Union[ - typing.Literal["put", "get", "search", "list_namespaces", "delete"], - Sequence[ - typing.Literal["put", "get", "search", "list_namespaces", "delete"] - ], - ] - ] = None, - ) -> typing.Union[AHO, Callable[[AHO], AHO]]: + actions: typing.Literal["put", "get", "search", "list_namespaces", "delete"] + | Sequence[typing.Literal["put", "get", "search", "list_namespaces", "delete"]] + | None = None, + ) -> AHO | Callable[[AHO], AHO]: """Register a handler for specific resources and actions. Can be used as a decorator or with explicit resource/action parameters: @@ -620,8 +604,8 @@ class _On: def __call__( self, *, - resources: typing.Union[str, Sequence[str]], - actions: typing.Optional[typing.Union[str, Sequence[str]]] = None, + resources: str | Sequence[str], + actions: str | Sequence[str] | None = None, ) -> Callable[[AHO], AHO]: ... @typing.overload @@ -629,11 +613,11 @@ class _On: def __call__( self, - fn: typing.Optional[AHO] = None, + fn: AHO | None = None, *, - resources: typing.Union[str, Sequence[str], None] = None, - actions: typing.Optional[typing.Union[str, Sequence[str]]] = None, - ) -> typing.Union[AHO, Callable[[AHO], AHO]]: + resources: str | Sequence[str] | None = None, + actions: str | Sequence[str] | None = None, + ) -> AHO | Callable[[AHO], AHO]: """Register a handler for specific resources and actions. Can be used as a decorator or with explicit resource/action parameters: @@ -675,8 +659,8 @@ class _On: def _register_handler( auth: Auth, - resource: typing.Optional[str], - action: typing.Optional[str], + resource: str | None, + action: str | None, fn: types.Handler, ) -> types.Handler: _validate_handler(fn) diff --git a/libs/sdk-py/langgraph_sdk/auth/exceptions.py b/libs/sdk-py/langgraph_sdk/auth/exceptions.py index 8dc0311a9..c58e67825 100644 --- a/libs/sdk-py/langgraph_sdk/auth/exceptions.py +++ b/libs/sdk-py/langgraph_sdk/auth/exceptions.py @@ -1,7 +1,9 @@ """Exceptions used in the auth system.""" +from __future__ import annotations + import http -import typing +from collections.abc import Mapping class HTTPException(Exception): @@ -37,8 +39,8 @@ class HTTPException(Exception): def __init__( self, status_code: int = 401, - detail: typing.Optional[str] = None, - headers: typing.Optional[typing.Mapping[str, str]] = None, + detail: str | None = None, + headers: Mapping[str, str] | None = None, ) -> None: if detail is None: detail = http.HTTPStatus(status_code).phrase diff --git a/libs/sdk-py/langgraph_sdk/auth/types.py b/libs/sdk-py/langgraph_sdk/auth/types.py index b8d7e408b..c9a5bd523 100644 --- a/libs/sdk-py/langgraph_sdk/auth/types.py +++ b/libs/sdk-py/langgraph_sdk/auth/types.py @@ -8,6 +8,8 @@ Note: All typing.TypedDict classes use total=False to make all fields typing.Optional by default. """ +from __future__ import annotations + import functools import sys import typing @@ -56,10 +58,8 @@ Values: """ FilterType = typing.Union[ - typing.Dict[ - str, typing.Union[str, typing.Dict[typing.Literal["$eq", "$contains"], str]] - ], - typing.Dict[str, str], + dict[str, typing.Union[str, dict[typing.Literal["$eq", "$contains"], str]]], + dict[str, str], ] """Response type for authorization handlers. @@ -100,7 +100,7 @@ Values: - error: Thread encountered an error """ -MetadataInput = typing.Dict[str, typing.Any] +MetadataInput = dict[str, typing.Any] """Type for arbitrary metadata attached to entities. Allows storing custom key-value pairs with any entity. @@ -434,7 +434,7 @@ class ThreadsRead(typing.TypedDict, total=False): thread_id: UUID """Unique identifier for the thread.""" - run_id: typing.Optional[UUID] + run_id: UUID | None """Run ID to filter by. Only used when reading run information within a thread.""" @@ -451,7 +451,7 @@ class ThreadsUpdate(typing.TypedDict, total=False): metadata: MetadataInput """typing.Optional metadata to update.""" - action: typing.Optional[typing.Literal["interrupt", "rollback"]] + action: typing.Literal["interrupt", "rollback"] | None """typing.Optional action to perform on the thread.""" @@ -464,7 +464,7 @@ class ThreadsDelete(typing.TypedDict, total=False): thread_id: UUID """Unique identifier for the thread.""" - run_id: typing.Optional[UUID] + run_id: UUID | None """typing.Optional run ID to filter by.""" @@ -480,7 +480,7 @@ class ThreadsSearch(typing.TypedDict, total=False): values: MetadataInput """typing.Optional values to filter by.""" - status: typing.Optional[ThreadStatus] + status: ThreadStatus | None """typing.Optional status to filter by.""" limit: int @@ -489,7 +489,7 @@ class ThreadsSearch(typing.TypedDict, total=False): offset: int """Offset for pagination.""" - thread_id: typing.Optional[UUID] + thread_id: UUID | None """typing.Optional thread ID to filter by.""" @@ -514,16 +514,16 @@ class RunsCreate(typing.TypedDict, total=False): ``` """ - assistant_id: typing.Optional[UUID] + assistant_id: UUID | None """typing.Optional assistant ID to use for this run.""" - thread_id: typing.Optional[UUID] + thread_id: UUID | None """typing.Optional thread ID to use for this run.""" - run_id: typing.Optional[UUID] + run_id: UUID | None """typing.Optional run ID to use for this run.""" - status: typing.Optional[RunStatus] + status: RunStatus | None """typing.Optional status for this run.""" metadata: MetadataInput @@ -541,10 +541,10 @@ class RunsCreate(typing.TypedDict, total=False): after_seconds: int """Number of seconds to wait before creating the run.""" - kwargs: typing.Dict[str, typing.Any] + kwargs: dict[str, typing.Any] """Keyword arguments to pass to the run.""" - action: typing.Optional[typing.Literal["interrupt", "rollback"]] + action: typing.Literal["interrupt", "rollback"] | None """Action to take if updating an existing run.""" @@ -570,7 +570,7 @@ class AssistantsCreate(typing.TypedDict, total=False): graph_id: str """Graph ID to use for this assistant.""" - config: typing.Optional[typing.Union[typing.Dict[str, typing.Any], typing.Any]] + config: dict[str, typing.Any] | typing.Any | None """typing.Optional configuration for the assistant.""" metadata: MetadataInput @@ -621,19 +621,19 @@ class AssistantsUpdate(typing.TypedDict, total=False): assistant_id: UUID """Unique identifier for the assistant.""" - graph_id: typing.Optional[str] + graph_id: str | None """typing.Optional graph ID to update.""" - config: typing.Optional[typing.Union[typing.Dict[str, typing.Any], typing.Any]] + config: dict[str, typing.Any] | typing.Any | None """typing.Optional configuration to update.""" metadata: MetadataInput """typing.Optional metadata to update.""" - name: typing.Optional[str] + name: str | None """typing.Optional name to update.""" - version: typing.Optional[int] + version: int | None """typing.Optional version to update.""" @@ -666,7 +666,7 @@ class AssistantsSearch(typing.TypedDict): ``` """ - graph_id: typing.Optional[str] + graph_id: str | None """typing.Optional graph ID to filter by.""" metadata: MetadataInput @@ -695,22 +695,22 @@ class CronsCreate(typing.TypedDict, total=False): ``` """ - payload: typing.Dict[str, typing.Any] + payload: dict[str, typing.Any] """Payload for the cron job.""" schedule: str """Schedule for the cron job.""" - cron_id: typing.Optional[UUID] + cron_id: UUID | None """typing.Optional unique identifier for the cron job.""" - thread_id: typing.Optional[UUID] + thread_id: UUID | None """typing.Optional thread ID to use for this cron job.""" - user_id: typing.Optional[str] + user_id: str | None """typing.Optional user ID to use for this cron job.""" - end_time: typing.Optional[datetime] + end_time: datetime | None """typing.Optional end time for the cron job.""" @@ -760,10 +760,10 @@ class CronsUpdate(typing.TypedDict, total=False): cron_id: UUID """Unique identifier for the cron job.""" - payload: typing.Optional[typing.Dict[str, typing.Any]] + payload: dict[str, typing.Any] | None """typing.Optional payload to update.""" - schedule: typing.Optional[str] + schedule: str | None """typing.Optional schedule to update.""" @@ -781,10 +781,10 @@ class CronsSearch(typing.TypedDict, total=False): ``` """ - assistant_id: typing.Optional[UUID] + assistant_id: UUID | None """typing.Optional assistant ID to filter by.""" - thread_id: typing.Optional[UUID] + thread_id: UUID | None """typing.Optional thread ID to filter by.""" limit: int @@ -810,7 +810,7 @@ class StoreSearch(typing.TypedDict): namespace: tuple[str, ...] """Prefix filter for defining the search scope.""" - filter: typing.Optional[dict[str, typing.Any]] + filter: dict[str, typing.Any] | None """Key-value pairs for filtering results based on exact matches or comparison operators.""" limit: int @@ -819,20 +819,20 @@ class StoreSearch(typing.TypedDict): offset: int """Number of matching items to skip for pagination.""" - query: typing.Optional[str] + query: str | None """Naturalj language search query for semantic search capabilities.""" class StoreListNamespaces(typing.TypedDict): """Operation to list and filter namespaces in the store.""" - namespace: typing.Optional[tuple[str, ...]] + namespace: tuple[str, ...] | None """Prefix filter namespaces.""" - suffix: typing.Optional[tuple[str, ...]] + suffix: tuple[str, ...] | None """Optional conditions for filtering namespaces.""" - max_depth: typing.Optional[int] + max_depth: int | None """Maximum depth of namespace hierarchy to return. Note: @@ -855,10 +855,10 @@ class StorePut(typing.TypedDict): key: str """Unique identifier for the item within its namespace.""" - value: typing.Optional[dict[str, typing.Any]] + value: dict[str, typing.Any] | None """The data to store, or None to mark the item for deletion.""" - index: typing.Optional[typing.Union[typing.Literal[False], list[str]]] + index: typing.Literal[False] | list[str] | None """Optional index configuration for full-text search.""" @@ -900,7 +900,7 @@ class on: ``` """ - value = typing.Dict[str, typing.Any] + value = dict[str, typing.Any] class threads: """Types for thread-related operations.""" diff --git a/libs/sdk-py/langgraph_sdk/client.py b/libs/sdk-py/langgraph_sdk/client.py index 4278f4dfd..52863229f 100644 --- a/libs/sdk-py/langgraph_sdk/client.py +++ b/libs/sdk-py/langgraph_sdk/client.py @@ -72,7 +72,7 @@ logger = logging.getLogger(__name__) RESERVED_HEADERS = ("x-api-key",) -def _get_api_key(api_key: Optional[str] = None) -> Optional[str]: +def _get_api_key(api_key: str | None = None) -> str | None: """Get the API key from the environment. Precedence: 1. explicit argument @@ -89,7 +89,7 @@ def _get_api_key(api_key: Optional[str] = None) -> Optional[str]: def _get_headers( - api_key: Optional[str], custom_headers: Optional[dict[str, str]] + api_key: str | None, custom_headers: dict[str, str] | None ) -> dict[str, str]: """Combine api_key and custom user-provided headers.""" custom_headers = custom_headers or {} @@ -127,7 +127,7 @@ _RUN_METADATA_PATTERN = re.compile( def _get_run_metadata_from_response( response: httpx.Response, -) -> Optional[RunCreateMetadata]: +) -> RunCreateMetadata | None: """Extract run metadata from the response headers.""" if (content_location := response.headers.get("Content-Location")) and ( match := _RUN_METADATA_PATTERN.search(content_location) @@ -142,10 +142,10 @@ def _get_run_metadata_from_response( def get_client( *, - url: Optional[str] = None, - api_key: Optional[str] = None, - headers: Optional[dict[str, str]] = None, - timeout: Optional[TimeoutTypes] = None, + url: str | None = None, + api_key: str | None = None, + headers: dict[str, str] | None = None, + timeout: TimeoutTypes | None = None, ) -> LangGraphClient: """Get a LangGraphClient instance. @@ -180,7 +180,7 @@ def get_client( ``` """ - transport: Optional[httpx.AsyncBaseTransport] = None + transport: httpx.AsyncBaseTransport | None = None if url is None: if os.environ.get("__LANGGRAPH_DEFER_LOOPBACK_TRANSPORT") == "true": transport = get_asgi_transport()(app=None, root_path="/noauth") @@ -248,9 +248,9 @@ class HttpClient: self, path: str, *, - params: Optional[QueryParamTypes] = None, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + params: QueryParamTypes | None = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Any: """Send a GET request.""" r = await self.client.get(path, params=params, headers=headers) @@ -271,9 +271,9 @@ class HttpClient: self, path: str, *, - json: Optional[dict], - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + json: dict | None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Any: """Send a POST request.""" if json is not None: @@ -302,8 +302,8 @@ class HttpClient: path: str, *, json: dict, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Any: """Send a PUT request.""" request_headers, content = await _aencode_json(json) @@ -328,8 +328,8 @@ class HttpClient: path: str, *, json: dict, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Any: """Send a PATCH request.""" request_headers, content = await _aencode_json(json) @@ -353,9 +353,9 @@ class HttpClient: self, path: str, *, - json: Optional[Any] = None, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + json: Any | None = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> None: """Send a DELETE request.""" r = await self.client.request("DELETE", path, json=json, headers=headers) @@ -376,10 +376,10 @@ class HttpClient: path: str, method: str, *, - json: Optional[dict] = None, - params: Optional[QueryParamTypes] = None, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + json: dict | None = None, + params: QueryParamTypes | None = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> AsyncIterator[StreamPart]: """Stream results using SSE.""" request_headers, content = await _aencode_json(json) @@ -462,7 +462,7 @@ class AssistantsClient: self.http = http async def get( - self, assistant_id: str, *, headers: Optional[dict[str, str]] = None + self, assistant_id: str, *, headers: dict[str, str] | None = None ) -> Assistant: """Get an assistant by ID. @@ -503,8 +503,8 @@ class AssistantsClient: self, assistant_id: str, *, - xray: Union[int, bool] = False, - headers: Optional[dict[str, str]] = None, + xray: int | bool = False, + headers: dict[str, str] | None = None, ) -> dict[str, list[dict[str, Any]]]: """Get the graph of an assistant by ID. @@ -552,7 +552,7 @@ class AssistantsClient: ) async def get_schemas( - self, assistant_id: str, *, headers: Optional[dict[str, str]] = None + self, assistant_id: str, *, headers: dict[str, str] | None = None ) -> GraphSchema: """Get the schemas of an assistant by ID. @@ -670,10 +670,10 @@ class AssistantsClient: async def get_subgraphs( self, assistant_id: str, - namespace: Optional[str] = None, + namespace: str | None = None, recurse: bool = False, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Subgraphs: """Get the schemas of an assistant by ID. @@ -702,15 +702,15 @@ class AssistantsClient: async def create( self, - graph_id: Optional[str], - config: Optional[Config] = None, + graph_id: str | None, + config: Config | None = None, *, metadata: Json = None, - assistant_id: Optional[str] = None, - if_exists: Optional[OnConflictBehavior] = None, - name: Optional[str] = None, - headers: Optional[dict[str, str]] = None, - description: Optional[str] = None, + assistant_id: str | None = None, + if_exists: OnConflictBehavior | None = None, + name: str | None = None, + headers: dict[str, str] | None = None, + description: str | None = None, ) -> Assistant: """Create a new assistant. @@ -766,12 +766,12 @@ class AssistantsClient: self, assistant_id: str, *, - graph_id: Optional[str] = None, - config: Optional[Config] = None, + graph_id: str | None = None, + config: Config | None = None, metadata: Json = None, - name: Optional[str] = None, - headers: Optional[dict[str, str]] = None, - description: Optional[str] = None, + name: str | None = None, + headers: dict[str, str] | None = None, + description: str | None = None, ) -> Assistant: """Update an assistant. @@ -825,7 +825,7 @@ class AssistantsClient: self, assistant_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Delete an assistant. @@ -852,12 +852,12 @@ class AssistantsClient: self, *, metadata: Json = None, - graph_id: Optional[str] = None, + graph_id: str | None = None, limit: int = 10, offset: int = 0, - sort_by: Optional[AssistantSortBy] = None, - sort_order: Optional[SortOrder] = None, - headers: Optional[dict[str, str]] = None, + sort_by: AssistantSortBy | None = None, + sort_order: SortOrder | None = None, + headers: dict[str, str] | None = None, ) -> list[Assistant]: """Search for assistants. @@ -911,7 +911,7 @@ class AssistantsClient: limit: int = 10, offset: int = 0, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> list[AssistantVersion]: """List all versions of an assistant. @@ -950,7 +950,7 @@ class AssistantsClient: assistant_id: str, version: int, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Assistant: """Change the version of an assistant. @@ -1000,7 +1000,7 @@ class ThreadsClient: self.http = http async def get( - self, thread_id: str, *, headers: Optional[dict[str, str]] = None + self, thread_id: str, *, headers: dict[str, str] | None = None ) -> Thread: """Get a thread by ID. @@ -1040,11 +1040,11 @@ class ThreadsClient: self, *, metadata: Json = None, - thread_id: Optional[str] = None, - if_exists: Optional[OnConflictBehavior] = None, - supersteps: Optional[Sequence[dict[str, Sequence[dict[str, Any]]]]] = None, - graph_id: Optional[str] = None, - headers: Optional[dict[str, str]] = None, + thread_id: str | None = None, + if_exists: OnConflictBehavior | None = None, + supersteps: Sequence[dict[str, Sequence[dict[str, Any]]]] | None = None, + graph_id: str | None = None, + headers: dict[str, str] | None = None, ) -> Thread: """Create a new thread. @@ -1105,7 +1105,7 @@ class ThreadsClient: thread_id: str, *, metadata: dict[str, Any], - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Thread: """Update a thread. @@ -1132,7 +1132,7 @@ class ThreadsClient: ) async def delete( - self, thread_id: str, *, headers: Optional[dict[str, str]] = None + self, thread_id: str, *, headers: dict[str, str] | None = None ) -> None: """Delete a thread. @@ -1160,12 +1160,12 @@ class ThreadsClient: *, metadata: Json = None, values: Json = None, - status: Optional[ThreadStatus] = None, + status: ThreadStatus | None = None, limit: int = 10, offset: int = 0, - sort_by: Optional[ThreadSortBy] = None, - sort_order: Optional[SortOrder] = None, - headers: Optional[dict[str, str]] = None, + sort_by: ThreadSortBy | None = None, + sort_order: SortOrder | None = None, + headers: dict[str, str] | None = None, ) -> list[Thread]: """Search for threads. @@ -1217,7 +1217,7 @@ class ThreadsClient: ) async def copy( - self, thread_id: str, *, headers: Optional[dict[str, str]] = None + self, thread_id: str, *, headers: dict[str, str] | None = None ) -> None: """Copy a thread. @@ -1245,11 +1245,11 @@ class ThreadsClient: async def get_state( self, thread_id: str, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, # deprecated + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, # deprecated *, subgraphs: bool = False, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> ThreadState: """Get the state of a thread. @@ -1372,12 +1372,12 @@ class ThreadsClient: async def update_state( self, thread_id: str, - values: Optional[Union[dict, Sequence[dict]]], + values: dict | Sequence[dict] | None, *, - as_node: Optional[str] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, # deprecated - headers: Optional[dict[str, str]] = None, + as_node: str | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, # deprecated + headers: dict[str, str] | None = None, ) -> ThreadUpdateStateResponse: """Update the state of a thread. @@ -1435,10 +1435,10 @@ class ThreadsClient: thread_id: str, *, limit: int = 10, - before: Optional[str | Checkpoint] = None, - metadata: Optional[dict] = None, - checkpoint: Optional[Checkpoint] = None, - headers: Optional[dict[str, str]] = None, + before: str | Checkpoint | None = None, + metadata: dict | None = None, + checkpoint: Checkpoint | None = None, + headers: dict[str, str] | None = None, ) -> list[ThreadState]: """Get the state history of a thread. @@ -1501,26 +1501,26 @@ class RunsClient: thread_id: str, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - feedback_keys: Optional[Sequence[str]] = None, - on_disconnect: Optional[DisconnectMode] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + feedback_keys: Sequence[str] | None = None, + on_disconnect: DisconnectMode | None = None, + webhook: str | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> AsyncIterator[StreamPart]: ... @overload @@ -1529,52 +1529,52 @@ class RunsClient: thread_id: None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - feedback_keys: Optional[Sequence[str]] = None, - on_disconnect: Optional[DisconnectMode] = None, - on_completion: Optional[OnCompletionBehavior] = None, - if_not_exists: Optional[IfNotExists] = None, - webhook: Optional[str] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + feedback_keys: Sequence[str] | None = None, + on_disconnect: DisconnectMode | None = None, + on_completion: OnCompletionBehavior | None = None, + if_not_exists: IfNotExists | None = None, + webhook: str | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> AsyncIterator[StreamPart]: ... def stream( self, - thread_id: Optional[str], + thread_id: str | None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - feedback_keys: Optional[Sequence[str]] = None, - on_disconnect: Optional[DisconnectMode] = None, - on_completion: Optional[OnCompletionBehavior] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + feedback_keys: Sequence[str] | None = None, + on_disconnect: DisconnectMode | None = None, + on_completion: OnCompletionBehavior | None = None, + webhook: str | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> AsyncIterator[StreamPart]: """Create a run and stream the results. @@ -1692,22 +1692,22 @@ class RunsClient: thread_id: None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - checkpoint_during: Optional[bool] = None, - config: Optional[Config] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - on_completion: Optional[OnCompletionBehavior] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + checkpoint_during: bool | None = None, + config: Config | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + on_completion: OnCompletionBehavior | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Run: ... @overload @@ -1716,50 +1716,50 @@ class RunsClient: thread_id: str, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Run: ... async def create( self, - thread_id: Optional[str], + thread_id: str | None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - on_completion: Optional[OnCompletionBehavior] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + on_completion: OnCompletionBehavior | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Run: """Create a background run. @@ -1914,24 +1914,24 @@ class RunsClient: thread_id: str, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - on_disconnect: Optional[DisconnectMode] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, + input: dict | None = None, + command: Command | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + on_disconnect: DisconnectMode | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, raise_error: bool = True, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, - ) -> Union[list[dict], dict[str, Any]]: ... + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, + ) -> list[dict] | dict[str, Any]: ... @overload async def wait( @@ -1939,47 +1939,47 @@ class RunsClient: thread_id: None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - on_disconnect: Optional[DisconnectMode] = None, - on_completion: Optional[OnCompletionBehavior] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, + input: dict | None = None, + command: Command | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + on_disconnect: DisconnectMode | None = None, + on_completion: OnCompletionBehavior | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, raise_error: bool = True, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, - ) -> Union[list[dict], dict[str, Any]]: ... + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, + ) -> list[dict] | dict[str, Any]: ... async def wait( self, - thread_id: Optional[str], + thread_id: str | None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - on_disconnect: Optional[DisconnectMode] = None, - on_completion: Optional[OnCompletionBehavior] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, + input: dict | None = None, + command: Command | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + on_disconnect: DisconnectMode | None = None, + on_completion: OnCompletionBehavior | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, raise_error: bool = True, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, - ) -> Union[list[dict], dict[str, Any]]: + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, + ) -> list[dict] | dict[str, Any]: """Create a run, wait until it finishes and return the final state. Args: @@ -2113,8 +2113,8 @@ class RunsClient: *, limit: int = 10, offset: int = 0, - status: Optional[RunStatus] = None, - headers: Optional[dict[str, str]] = None, + status: RunStatus | None = None, + headers: dict[str, str] | None = None, ) -> list[Run]: """List runs. @@ -2151,7 +2151,7 @@ class RunsClient: ) async def get( - self, thread_id: str, run_id: str, *, headers: Optional[dict[str, str]] = None + self, thread_id: str, run_id: str, *, headers: dict[str, str] | None = None ) -> Run: """Get a run. @@ -2186,7 +2186,7 @@ class RunsClient: *, wait: bool = False, action: CancelAction = "interrupt", - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Get a run. @@ -2221,7 +2221,7 @@ class RunsClient: ) async def join( - self, thread_id: str, run_id: str, *, headers: Optional[dict[str, str]] = None + self, thread_id: str, run_id: str, *, headers: dict[str, str] | None = None ) -> dict: """Block until a run is done. Returns the final state of the thread. @@ -2254,9 +2254,9 @@ class RunsClient: run_id: str, *, cancel_on_disconnect: bool = False, - stream_mode: Optional[Union[StreamMode, Sequence[StreamMode]]] = None, - headers: Optional[dict[str, str]] = None, - last_event_id: Optional[str] = None, + stream_mode: StreamMode | Sequence[StreamMode] | None = None, + headers: dict[str, str] | None = None, + last_event_id: str | None = None, ) -> 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 @@ -2302,7 +2302,7 @@ class RunsClient: ) async def delete( - self, thread_id: str, run_id: str, *, headers: Optional[dict[str, str]] = None + self, thread_id: str, run_id: str, *, headers: dict[str, str] | None = None ) -> None: """Delete a run. @@ -2361,15 +2361,15 @@ class CronClient: assistant_id: str, *, schedule: str, - input: Optional[dict] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, list[str]]] = None, - interrupt_after: Optional[Union[All, list[str]]] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[str] = None, - headers: Optional[dict[str, str]] = None, + input: dict | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | list[str] | None = None, + interrupt_after: All | list[str] | None = None, + webhook: str | None = None, + multitask_strategy: str | None = None, + headers: dict[str, str] | None = None, ) -> Run: """Create a cron job for a thread. @@ -2435,15 +2435,15 @@ class CronClient: assistant_id: str, *, schedule: str, - input: Optional[dict] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, list[str]]] = None, - interrupt_after: Optional[Union[All, list[str]]] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[str] = None, - headers: Optional[dict[str, str]] = None, + input: dict | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | list[str] | None = None, + interrupt_after: All | list[str] | None = None, + webhook: str | None = None, + multitask_strategy: str | None = None, + headers: dict[str, str] | None = None, ) -> Run: """Create a cron run. @@ -2502,7 +2502,7 @@ class CronClient: async def delete( self, cron_id: str, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Delete a cron. @@ -2528,11 +2528,11 @@ class CronClient: async def search( self, *, - assistant_id: Optional[str] = None, - thread_id: Optional[str] = None, + assistant_id: str | None = None, + thread_id: str | None = None, limit: int = 10, offset: int = 0, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> list[Cron]: """Get a list of cron jobs. @@ -2617,9 +2617,9 @@ class StoreClient: /, key: str, value: dict[str, Any], - index: Optional[Union[Literal[False], list[str]]] = None, - ttl: Optional[int] = None, - headers: Optional[dict[str, str]] = None, + index: Literal[False] | list[str] | None = None, + ttl: int | None = None, + headers: dict[str, str] | None = None, ) -> None: """Store or update an item. @@ -2667,8 +2667,8 @@ class StoreClient: /, key: str, *, - refresh_ttl: Optional[bool] = None, - headers: Optional[dict[str, str]] = None, + refresh_ttl: bool | None = None, + headers: dict[str, str] | None = None, ) -> Item: """Retrieve a single item. @@ -2719,7 +2719,7 @@ class StoreClient: namespace: Sequence[str], /, key: str, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Delete an item. @@ -2751,12 +2751,12 @@ class StoreClient: self, namespace_prefix: Sequence[str], /, - filter: Optional[dict[str, Any]] = None, + filter: dict[str, Any] | None = None, limit: int = 10, offset: int = 0, - query: Optional[str] = None, - refresh_ttl: Optional[bool] = None, - headers: Optional[dict[str, str]] = None, + query: str | None = None, + refresh_ttl: bool | None = None, + headers: dict[str, str] | None = None, ) -> SearchItemsResponse: """Search for items within a namespace prefix. @@ -2822,12 +2822,12 @@ class StoreClient: async def list_namespaces( self, - prefix: Optional[list[str]] = None, - suffix: Optional[list[str]] = None, - max_depth: Optional[int] = None, + prefix: list[str] | None = None, + suffix: list[str] | None = None, + max_depth: int | None = None, limit: int = 100, offset: int = 0, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> ListNamespaceResponse: """List namespaces with optional match conditions. @@ -2879,10 +2879,10 @@ class StoreClient: def get_sync_client( *, - url: Optional[str] = None, - api_key: Optional[str] = None, - headers: Optional[dict[str, str]] = None, - timeout: Optional[TimeoutTypes] = None, + url: str | None = None, + api_key: str | None = None, + headers: dict[str, str] | None = None, + timeout: TimeoutTypes | None = None, ) -> SyncLangGraphClient: """Get a synchronous LangGraphClient instance. @@ -2974,9 +2974,9 @@ class SyncHttpClient: self, path: str, *, - params: Optional[QueryParamTypes] = None, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + params: QueryParamTypes | None = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Any: """Send a GET request.""" r = self.client.get(path, params=params, headers=headers) @@ -2997,9 +2997,9 @@ class SyncHttpClient: self, path: str, *, - json: Optional[dict], - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + json: dict | None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Any: """Send a POST request.""" if json is not None: @@ -3027,8 +3027,8 @@ class SyncHttpClient: path: str, *, json: dict, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Any: """Send a PUT request.""" request_headers, content = _encode_json(json) @@ -3054,8 +3054,8 @@ class SyncHttpClient: path: str, *, json: dict, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Any: """Send a PATCH request.""" request_headers, content = _encode_json(json) @@ -3079,9 +3079,9 @@ class SyncHttpClient: self, path: str, *, - json: Optional[Any] = None, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + json: Any | None = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> None: """Send a DELETE request.""" r = self.client.request("DELETE", path, json=json, headers=headers) @@ -3102,10 +3102,10 @@ class SyncHttpClient: path: str, method: str, *, - json: Optional[dict] = None, - params: Optional[QueryParamTypes] = None, - headers: Optional[dict[str, str]] = None, - on_response: Optional[Callable[[httpx.Response], None]] = None, + json: dict | None = None, + params: QueryParamTypes | None = None, + headers: dict[str, str] | None = None, + on_response: Callable[[httpx.Response], None] | None = None, ) -> Iterator[StreamPart]: """Stream the results of a request using SSE.""" request_headers, content = _encode_json(json) @@ -3180,7 +3180,7 @@ class SyncAssistantsClient: self, assistant_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Assistant: """Get an assistant by ID. @@ -3220,8 +3220,8 @@ class SyncAssistantsClient: self, assistant_id: str, *, - xray: Union[int, bool] = False, - headers: Optional[dict[str, str]] = None, + xray: int | bool = False, + headers: dict[str, str] | None = None, ) -> dict[str, list[dict[str, Any]]]: """Get the graph of an assistant by ID. @@ -3268,7 +3268,7 @@ class SyncAssistantsClient: self, assistant_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> GraphSchema: """Get the schemas of an assistant by ID. @@ -3382,10 +3382,10 @@ class SyncAssistantsClient: def get_subgraphs( self, assistant_id: str, - namespace: Optional[str] = None, + namespace: str | None = None, recurse: bool = False, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Subgraphs: """Get the schemas of an assistant by ID. @@ -3412,15 +3412,15 @@ class SyncAssistantsClient: def create( self, - graph_id: Optional[str], - config: Optional[Config] = None, + graph_id: str | None, + config: Config | None = None, *, metadata: Json = None, - assistant_id: Optional[str] = None, - if_exists: Optional[OnConflictBehavior] = None, - name: Optional[str] = None, - headers: Optional[dict[str, str]] = None, - description: Optional[str] = None, + assistant_id: str | None = None, + if_exists: OnConflictBehavior | None = None, + name: str | None = None, + headers: dict[str, str] | None = None, + description: str | None = None, ) -> Assistant: """Create a new assistant. @@ -3476,12 +3476,12 @@ class SyncAssistantsClient: self, assistant_id: str, *, - graph_id: Optional[str] = None, - config: Optional[Config] = None, + graph_id: str | None = None, + config: Config | None = None, metadata: Json = None, - name: Optional[str] = None, - headers: Optional[dict[str, str]] = None, - description: Optional[str] = None, + name: str | None = None, + headers: dict[str, str] | None = None, + description: str | None = None, ) -> Assistant: """Update an assistant. @@ -3534,7 +3534,7 @@ class SyncAssistantsClient: self, assistant_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Delete an assistant. @@ -3561,10 +3561,10 @@ class SyncAssistantsClient: self, *, metadata: Json = None, - graph_id: Optional[str] = None, + graph_id: str | None = None, limit: int = 10, offset: int = 0, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> list[Assistant]: """Search for assistants. @@ -3612,7 +3612,7 @@ class SyncAssistantsClient: limit: int = 10, offset: int = 0, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> list[AssistantVersion]: """List all versions of an assistant. @@ -3652,7 +3652,7 @@ class SyncAssistantsClient: assistant_id: str, version: int, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Assistant: """Change the version of an assistant. @@ -3704,7 +3704,7 @@ class SyncThreadsClient: self, thread_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Thread: """Get a thread by ID. @@ -3743,11 +3743,11 @@ class SyncThreadsClient: self, *, metadata: Json = None, - thread_id: Optional[str] = None, - if_exists: Optional[OnConflictBehavior] = None, - supersteps: Optional[Sequence[dict[str, Sequence[dict[str, Any]]]]] = None, - graph_id: Optional[str] = None, - headers: Optional[dict[str, str]] = None, + thread_id: str | None = None, + if_exists: OnConflictBehavior | None = None, + supersteps: Sequence[dict[str, Sequence[dict[str, Any]]]] | None = None, + graph_id: str | None = None, + headers: dict[str, str] | None = None, ) -> Thread: """Create a new thread. @@ -3809,7 +3809,7 @@ class SyncThreadsClient: thread_id: str, *, metadata: dict[str, Any], - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Thread: """Update a thread. @@ -3839,7 +3839,7 @@ class SyncThreadsClient: self, thread_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Delete a thread. @@ -3866,10 +3866,10 @@ class SyncThreadsClient: *, metadata: Json = None, values: Json = None, - status: Optional[ThreadStatus] = None, + status: ThreadStatus | None = None, limit: int = 10, offset: int = 0, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> list[Thread]: """Search for threads. @@ -3913,7 +3913,7 @@ class SyncThreadsClient: self, thread_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Copy a thread. @@ -3939,11 +3939,11 @@ class SyncThreadsClient: def get_state( self, thread_id: str, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, # deprecated + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, # deprecated *, subgraphs: bool = False, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> ThreadState: """Get the state of a thread. @@ -4066,12 +4066,12 @@ class SyncThreadsClient: def update_state( self, thread_id: str, - values: Optional[Union[dict, Sequence[dict]]], + values: dict | Sequence[dict] | None, *, - as_node: Optional[str] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, # deprecated - headers: Optional[dict[str, str]] = None, + as_node: str | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, # deprecated + headers: dict[str, str] | None = None, ) -> ThreadUpdateStateResponse: """Update the state of a thread. @@ -4127,10 +4127,10 @@ class SyncThreadsClient: thread_id: str, *, limit: int = 10, - before: Optional[str | Checkpoint] = None, - metadata: Optional[dict] = None, - checkpoint: Optional[Checkpoint] = None, - headers: Optional[dict[str, str]] = None, + before: str | Checkpoint | None = None, + metadata: dict | None = None, + checkpoint: Checkpoint | None = None, + headers: dict[str, str] | None = None, ) -> list[ThreadState]: """Get the state history of a thread. @@ -4195,25 +4195,25 @@ class SyncRunsClient: thread_id: str, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - feedback_keys: Optional[Sequence[str]] = None, - on_disconnect: Optional[DisconnectMode] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + feedback_keys: Sequence[str] | None = None, + on_disconnect: DisconnectMode | None = None, + webhook: str | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Iterator[StreamPart]: ... @overload @@ -4222,52 +4222,52 @@ class SyncRunsClient: thread_id: None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - feedback_keys: Optional[Sequence[str]] = None, - on_disconnect: Optional[DisconnectMode] = None, - on_completion: Optional[OnCompletionBehavior] = None, - if_not_exists: Optional[IfNotExists] = None, - webhook: Optional[str] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + feedback_keys: Sequence[str] | None = None, + on_disconnect: DisconnectMode | None = None, + on_completion: OnCompletionBehavior | None = None, + if_not_exists: IfNotExists | None = None, + webhook: str | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Iterator[StreamPart]: ... def stream( self, - thread_id: Optional[str], + thread_id: str | None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - feedback_keys: Optional[Sequence[str]] = None, - on_disconnect: Optional[DisconnectMode] = None, - on_completion: Optional[OnCompletionBehavior] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + feedback_keys: Sequence[str] | None = None, + on_disconnect: DisconnectMode | None = None, + on_completion: OnCompletionBehavior | None = None, + webhook: str | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Iterator[StreamPart]: """Create a run and stream the results. @@ -4383,22 +4383,22 @@ class SyncRunsClient: thread_id: None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - on_completion: Optional[OnCompletionBehavior] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + on_completion: OnCompletionBehavior | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Run: ... @overload @@ -4407,50 +4407,50 @@ class SyncRunsClient: thread_id: str, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Run: ... def create( self, - thread_id: Optional[str], + thread_id: str | None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - stream_mode: Union[StreamMode, Sequence[StreamMode]] = "values", + input: dict | None = None, + command: Command | None = None, + stream_mode: StreamMode | Sequence[StreamMode] = "values", stream_subgraphs: bool = False, stream_resumable: bool = False, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - on_completion: Optional[OnCompletionBehavior] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + multitask_strategy: MultitaskStrategy | None = None, + on_completion: OnCompletionBehavior | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, ) -> Run: """Create a background run. @@ -4591,7 +4591,7 @@ class SyncRunsClient: ) def create_batch( - self, payloads: list[RunCreate], *, headers: Optional[dict[str, str]] = None + self, payloads: list[RunCreate], *, headers: dict[str, str] | None = None ) -> list[Run]: """Create a batch of stateless background runs.""" @@ -4607,23 +4607,23 @@ class SyncRunsClient: thread_id: str, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - on_disconnect: Optional[DisconnectMode] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, - ) -> Union[list[dict], dict[str, Any]]: ... + input: dict | None = None, + command: Command | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + on_disconnect: DisconnectMode | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, + ) -> list[dict] | dict[str, Any]: ... @overload def wait( @@ -4631,45 +4631,45 @@ class SyncRunsClient: thread_id: None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - on_disconnect: Optional[DisconnectMode] = None, - on_completion: Optional[OnCompletionBehavior] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, - ) -> Union[list[dict], dict[str, Any]]: ... + input: dict | None = None, + command: Command | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + on_disconnect: DisconnectMode | None = None, + on_completion: OnCompletionBehavior | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, + ) -> list[dict] | dict[str, Any]: ... def wait( self, - thread_id: Optional[str], + thread_id: str | None, assistant_id: str, *, - input: Optional[dict] = None, - command: Optional[Command] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - checkpoint: Optional[Checkpoint] = None, - checkpoint_id: Optional[str] = None, - interrupt_before: Optional[Union[All, Sequence[str]]] = None, - interrupt_after: Optional[Union[All, Sequence[str]]] = None, - webhook: Optional[str] = None, - on_disconnect: Optional[DisconnectMode] = None, - on_completion: Optional[OnCompletionBehavior] = None, - multitask_strategy: Optional[MultitaskStrategy] = None, - if_not_exists: Optional[IfNotExists] = None, - after_seconds: Optional[int] = None, - headers: Optional[dict[str, str]] = None, - on_run_created: Optional[Callable[[RunCreateMetadata], None]] = None, - ) -> Union[list[dict], dict[str, Any]]: + input: dict | None = None, + command: Command | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + checkpoint: Checkpoint | None = None, + checkpoint_id: str | None = None, + interrupt_before: All | Sequence[str] | None = None, + interrupt_after: All | Sequence[str] | None = None, + webhook: str | None = None, + on_disconnect: DisconnectMode | None = None, + on_completion: OnCompletionBehavior | None = None, + multitask_strategy: MultitaskStrategy | None = None, + if_not_exists: IfNotExists | None = None, + after_seconds: int | None = None, + headers: dict[str, str] | None = None, + on_run_created: Callable[[RunCreateMetadata], None] | None = None, + ) -> list[dict] | dict[str, Any]: """Create a run, wait until it finishes and return the final state. Args: @@ -4794,7 +4794,7 @@ class SyncRunsClient: *, limit: int = 10, offset: int = 0, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> list[Run]: """List runs. @@ -4828,7 +4828,7 @@ class SyncRunsClient: thread_id: str, run_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> Run: """Get a run. @@ -4860,7 +4860,7 @@ class SyncRunsClient: *, wait: bool = False, action: CancelAction = "interrupt", - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Get a run. @@ -4899,7 +4899,7 @@ class SyncRunsClient: thread_id: str, run_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> dict: """Block until a run is done. Returns the final state of the thread. @@ -4931,10 +4931,10 @@ class SyncRunsClient: thread_id: str, run_id: str, *, - stream_mode: Optional[Union[StreamMode, Sequence[StreamMode]]] = None, + stream_mode: StreamMode | Sequence[StreamMode] | None = None, cancel_on_disconnect: bool = False, - headers: Optional[dict[str, str]] = None, - last_event_id: Optional[str] = None, + headers: dict[str, str] | None = None, + last_event_id: str | None = None, ) -> Iterator[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 @@ -4983,7 +4983,7 @@ class SyncRunsClient: thread_id: str, run_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Delete a run. @@ -5036,15 +5036,15 @@ class SyncCronClient: assistant_id: str, *, schedule: str, - input: Optional[dict] = None, - metadata: Optional[dict] = None, - checkpoint_during: Optional[bool] = None, - config: Optional[Config] = None, - interrupt_before: Optional[Union[All, list[str]]] = None, - interrupt_after: Optional[Union[All, list[str]]] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[str] = None, - headers: Optional[dict[str, str]] = None, + input: dict | None = None, + metadata: dict | None = None, + checkpoint_during: bool | None = None, + config: Config | None = None, + interrupt_before: All | list[str] | None = None, + interrupt_after: All | list[str] | None = None, + webhook: str | None = None, + multitask_strategy: str | None = None, + headers: dict[str, str] | None = None, ) -> Run: """Create a cron job for a thread. @@ -5107,15 +5107,15 @@ class SyncCronClient: assistant_id: str, *, schedule: str, - input: Optional[dict] = None, - metadata: Optional[dict] = None, - config: Optional[Config] = None, - checkpoint_during: Optional[bool] = None, - interrupt_before: Optional[Union[All, list[str]]] = None, - interrupt_after: Optional[Union[All, list[str]]] = None, - webhook: Optional[str] = None, - multitask_strategy: Optional[str] = None, - headers: Optional[dict[str, str]] = None, + input: dict | None = None, + metadata: dict | None = None, + config: Config | None = None, + checkpoint_during: bool | None = None, + interrupt_before: All | list[str] | None = None, + interrupt_after: All | list[str] | None = None, + webhook: str | None = None, + multitask_strategy: str | None = None, + headers: dict[str, str] | None = None, ) -> Run: """Create a cron run. @@ -5175,7 +5175,7 @@ class SyncCronClient: self, cron_id: str, *, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Delete a cron. @@ -5201,11 +5201,11 @@ class SyncCronClient: def search( self, *, - assistant_id: Optional[str] = None, - thread_id: Optional[str] = None, + assistant_id: str | None = None, + thread_id: str | None = None, limit: int = 10, offset: int = 0, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> list[Cron]: """Get a list of cron jobs. @@ -5289,9 +5289,9 @@ class SyncStoreClient: /, key: str, value: dict[str, Any], - index: Optional[Union[Literal[False], list[str]]] = None, - ttl: Optional[int] = None, - headers: Optional[dict[str, str]] = None, + index: Literal[False] | list[str] | None = None, + ttl: int | None = None, + headers: dict[str, str] | None = None, ) -> None: """Store or update an item. @@ -5337,8 +5337,8 @@ class SyncStoreClient: /, key: str, *, - refresh_ttl: Optional[bool] = None, - headers: Optional[dict[str, str]] = None, + refresh_ttl: bool | None = None, + headers: dict[str, str] | None = None, ) -> Item: """Retrieve a single item. @@ -5390,7 +5390,7 @@ class SyncStoreClient: namespace: Sequence[str], /, key: str, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> None: """Delete an item. @@ -5420,12 +5420,12 @@ class SyncStoreClient: self, namespace_prefix: Sequence[str], /, - filter: Optional[dict[str, Any]] = None, + filter: dict[str, Any] | None = None, limit: int = 10, offset: int = 0, - query: Optional[str] = None, - refresh_ttl: Optional[bool] = None, - headers: Optional[dict[str, str]] = None, + query: str | None = None, + refresh_ttl: bool | None = None, + headers: dict[str, str] | None = None, ) -> SearchItemsResponse: """Search for items within a namespace prefix. @@ -5487,12 +5487,12 @@ class SyncStoreClient: def list_namespaces( self, - prefix: Optional[list[str]] = None, - suffix: Optional[list[str]] = None, - max_depth: Optional[int] = None, + prefix: list[str] | None = None, + suffix: list[str] | None = None, + max_depth: int | None = None, limit: int = 100, offset: int = 0, - headers: Optional[dict[str, str]] = None, + headers: dict[str, str] | None = None, ) -> ListNamespaceResponse: """List namespaces with optional match conditions. diff --git a/libs/sdk-py/langgraph_sdk/schema.py b/libs/sdk-py/langgraph_sdk/schema.py index 63b9855d9..c3e135603 100644 --- a/libs/sdk-py/langgraph_sdk/schema.py +++ b/libs/sdk-py/langgraph_sdk/schema.py @@ -1,5 +1,7 @@ """Data models for interacting with the LangGraph API.""" +from __future__ import annotations + from collections.abc import Sequence from datetime import datetime from typing import ( @@ -8,7 +10,6 @@ from typing import ( NamedTuple, Optional, TypedDict, - Union, ) Json = Optional[dict[str, Any]] @@ -142,9 +143,9 @@ class Checkpoint(TypedDict): """Unique identifier for the thread associated with this checkpoint.""" checkpoint_ns: str """Namespace for the checkpoint; used internally to manage subgraph state.""" - checkpoint_id: Optional[str] + checkpoint_id: str | None """Optional unique identifier for the checkpoint itself.""" - checkpoint_map: Optional[dict[str, Any]] + checkpoint_map: dict[str, Any] | None """Optional dictionary containing checkpoint-specific data.""" @@ -153,16 +154,16 @@ class GraphSchema(TypedDict): graph_id: str """The ID of the graph.""" - input_schema: Optional[dict] + input_schema: dict | None """The schema for the graph input. Missing if unable to generate JSON schema from graph.""" - output_schema: Optional[dict] + output_schema: dict | None """The schema for the graph output. Missing if unable to generate JSON schema from graph.""" - state_schema: Optional[dict] + state_schema: dict | None """The schema for the graph state. Missing if unable to generate JSON schema from graph.""" - config_schema: Optional[dict] + config_schema: dict | None """The schema for the graph config. Missing if unable to generate JSON schema from graph.""" @@ -187,7 +188,7 @@ class AssistantBase(TypedDict): """The version of the assistant""" name: str """The name of the assistant""" - description: Optional[str] + description: str | None """The description of the assistant""" @@ -213,7 +214,7 @@ class Interrupt(TypedDict, total=False): """When the interrupt occurred.""" resumable: bool """Whether the interrupt can be resumed.""" - ns: Optional[list[str]] + ns: list[str] | None """Optional namespace for the interrupt.""" @@ -241,17 +242,17 @@ class ThreadTask(TypedDict): id: str name: str - error: Optional[str] + error: str | None interrupts: list[Interrupt] - checkpoint: Optional[Checkpoint] - state: Optional["ThreadState"] - result: Optional[dict[str, Any]] + checkpoint: Checkpoint | None + state: ThreadState | None + result: dict[str, Any] | None class ThreadState(TypedDict): """Represents the state of a thread.""" - values: Union[list[dict], dict[str, Any]] + values: list[dict] | dict[str, Any] """The state values.""" next: Sequence[str] """The next nodes to execute. If empty, the thread is done until new input is @@ -260,9 +261,9 @@ class ThreadState(TypedDict): """The ID of the checkpoint.""" metadata: Json """Metadata for this state""" - created_at: Optional[str] + created_at: str | None """Timestamp of state creation""" - parent_checkpoint: Optional[Checkpoint] + parent_checkpoint: Checkpoint | None """The ID of the parent checkpoint. If missing, this is the root checkpoint.""" tasks: Sequence[ThreadTask] """Tasks to execute in this step. If already attempted, may contain an error.""" @@ -301,9 +302,9 @@ class Cron(TypedDict): cron_id: str """The ID of the cron.""" - thread_id: Optional[str] + thread_id: str | None """The ID of the thread.""" - end_time: Optional[datetime] + end_time: datetime | None """The end date to stop running the cron.""" schedule: str """The schedule to run, cron format.""" @@ -318,25 +319,25 @@ class Cron(TypedDict): class RunCreate(TypedDict): """Defines the parameters for initiating a background run.""" - thread_id: Optional[str] + thread_id: str | None """The identifier of the thread to run. If not provided, the run is stateless.""" assistant_id: str """The identifier of the assistant to use for this run.""" - input: Optional[dict] + input: dict | None """Initial input data for the run.""" - metadata: Optional[dict] + metadata: dict | None """Additional metadata to associate with the run.""" - config: Optional[Config] + config: Config | None """Configuration options for the run.""" - checkpoint_id: Optional[str] + checkpoint_id: str | None """The identifier of a checkpoint to resume from.""" - interrupt_before: Optional[list[str]] + interrupt_before: list[str] | None """List of node names to interrupt execution before.""" - interrupt_after: Optional[list[str]] + interrupt_after: list[str] | None """List of node names to interrupt execution after.""" - webhook: Optional[str] + webhook: str | None """URL to send webhook notifications about the run's progress.""" - multitask_strategy: Optional[MultitaskStrategy] + multitask_strategy: MultitaskStrategy | None """Strategy for handling concurrent runs on the same thread.""" @@ -376,7 +377,7 @@ class SearchItem(Item, total=False): searching a compatible store with a natural language query. """ - score: Optional[float] + score: float | None class SearchItemsResponse(TypedDict): @@ -404,7 +405,7 @@ class Send(TypedDict): node: str """The name of the target node to send the message to.""" - input: Optional[dict[str, Any]] + input: dict[str, Any] | None """Optional dictionary containing the input data to be passed to the node. If None, the node will be called with no input.""" @@ -418,14 +419,14 @@ class Command(TypedDict, total=False): and resume from interruptions. """ - goto: Union[Send, str, Sequence[Union[Send, str]]] + goto: Send | str | Sequence[Send | str] """Specifies where execution should continue. Can be: - A string node name to navigate to - A Send object to execute a node with specific input - A sequence of node names or Send objects to execute in order """ - update: Union[dict[str, Any], Sequence[tuple[str, Any]]] + update: dict[str, Any] | Sequence[tuple[str, Any]] """Updates to apply to the graph's state. Can be: - A dictionary of state updates to merge @@ -443,5 +444,5 @@ class RunCreateMetadata(TypedDict): run_id: str """The ID of the run.""" - thread_id: Optional[str] + thread_id: str | None """The ID of the thread.""" diff --git a/libs/sdk-py/langgraph_sdk/sse.py b/libs/sdk-py/langgraph_sdk/sse.py index a3c9be790..678f6977e 100644 --- a/libs/sdk-py/langgraph_sdk/sse.py +++ b/libs/sdk-py/langgraph_sdk/sse.py @@ -1,7 +1,9 @@ """Adapted from httpx_sse to split lines on \n, \r, \r\n per the SSE spec.""" +from __future__ import annotations + from collections.abc import AsyncIterator, Iterator -from typing import Optional, Union +from typing import Union import httpx import orjson @@ -77,9 +79,9 @@ class SSEDecoder: self._event = "" self._data = bytearray() self._last_event_id = "" - self._retry: Optional[int] = None + self._retry: int | None = None - def decode(self, line: bytes) -> Optional[StreamPart]: + def decode(self, line: bytes) -> StreamPart | None: # See: https://html.spec.whatwg.org/multipage/server-sent-events.html#event-stream-interpretation # noqa: E501 if not line: diff --git a/libs/sdk-py/pyproject.toml b/libs/sdk-py/pyproject.toml index 4c4b050aa..8f1b6d92a 100644 --- a/libs/sdk-py/pyproject.toml +++ b/libs/sdk-py/pyproject.toml @@ -48,4 +48,4 @@ lint.select = [ "B", # flake8-bugbear "I", # isort ] -lint.ignore = ["E501", "B008", "UP007", "UP006"] +lint.ignore = ["E501", "B008"]