diff --git a/libs/checkpoint/Makefile b/libs/checkpoint/Makefile index 94b1963d6..ddf087ef5 100644 --- a/libs/checkpoint/Makefile +++ b/libs/checkpoint/Makefile @@ -27,7 +27,8 @@ lint lint_diff lint_package lint_tests: poetry run ruff check . [ "$(PYTHON_FILES)" = "" ] || poetry run ruff format $(PYTHON_FILES) --diff [ "$(PYTHON_FILES)" = "" ] || poetry run ruff check --select I $(PYTHON_FILES) - [ "$(PYTHON_FILES)" = "" ] || mkdir -p $(MYPY_CACHE) || poetry run mypy $(PYTHON_FILES) --cache-dir $(MYPY_CACHE) + [ "$(PYTHON_FILES)" = "" ] || mkdir -p $(MYPY_CACHE) + [ "$(PYTHON_FILES)" = "" ] || poetry run mypy $(PYTHON_FILES) --cache-dir $(MYPY_CACHE) format format_diff: poetry run ruff format $(PYTHON_FILES) diff --git a/libs/checkpoint/langgraph/checkpoint/base/__init__.py b/libs/checkpoint/langgraph/checkpoint/base/__init__.py index 17caf46d8..436c7e742 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/base/__init__.py @@ -3,6 +3,7 @@ from typing import ( Any, AsyncIterator, Dict, + Generic, Iterator, List, Literal, @@ -135,7 +136,7 @@ def create_checkpoint( if channels is None: values = checkpoint["channel_values"] else: - values: dict[str, Any] = {} + values = {} for k, v in channels.items(): if k not in checkpoint["channel_versions"]: continue @@ -192,7 +193,7 @@ CheckpointId = ConfigurableFieldSpec( ) -class BaseCheckpointSaver: +class BaseCheckpointSaver(Generic[V]): """Base class for creating a graph checkpointer. Checkpointers allow LangGraph agents to persist their state @@ -420,7 +421,12 @@ class BaseCheckpointSaver: Returns: V: The next version identifier, which must be increasing. """ - return current + 1 if current is not None else 1 + if isinstance(current, str): + raise NotImplementedError + elif current is None: + return 1 # type: ignore[return-value] + else: + return current + 1 class EmptyChannelError(Exception): diff --git a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py index a0c3237ac..20d53ea06 100644 --- a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py @@ -22,7 +22,7 @@ from langgraph.checkpoint.serde.types import TASKS, ChannelProtocol class MemorySaver( - BaseCheckpointSaver, AbstractContextManager, AbstractAsyncContextManager + BaseCheckpointSaver[str], AbstractContextManager, AbstractAsyncContextManager ): """An in-memory checkpoint saver. @@ -54,9 +54,14 @@ class MemorySaver( """ # thread ID -> checkpoint NS -> checkpoint ID -> checkpoint mapping - storage: defaultdict[str, dict[str, dict[str, tuple[bytes, bytes, Optional[str]]]]] + storage: defaultdict[ + str, + dict[ + str, dict[str, tuple[tuple[str, bytes], tuple[str, bytes], Optional[str]]] + ], + ] writes: defaultdict[ - tuple[str, str, str], dict[tuple[str, int], tuple[str, str, bytes]] + tuple[str, str, str], dict[tuple[str, int], tuple[str, str, tuple[str, bytes]]] ] def __init__( @@ -316,7 +321,7 @@ class MemorySaver( RunnableConfig: The updated config containing the saved checkpoint's timestamp. """ c = checkpoint.copy() - c.pop("pending_sends") + c.pop("pending_sends") # type: ignore[misc] thread_id = config["configurable"]["thread_id"] checkpoint_ns = config["configurable"]["checkpoint_ns"] self.storage[thread_id][checkpoint_ns].update( @@ -341,7 +346,7 @@ class MemorySaver( config: RunnableConfig, writes: List[Tuple[str, Any]], task_id: str, - ) -> RunnableConfig: + ) -> None: """Save a list of writes to the in-memory storage. This method saves a list of writes to the in-memory storage. The writes are associated @@ -444,7 +449,7 @@ class MemorySaver( config: RunnableConfig, writes: List[Tuple[str, Any]], task_id: str, - ) -> RunnableConfig: + ) -> None: """Asynchronous version of put_writes. This method is an asynchronous wrapper around put_writes that runs the synchronous diff --git a/libs/checkpoint/langgraph/checkpoint/serde/base.py b/libs/checkpoint/langgraph/checkpoint/serde/base.py index 58be7e08c..229837735 100644 --- a/libs/checkpoint/langgraph/checkpoint/serde/base.py +++ b/libs/checkpoint/langgraph/checkpoint/serde/base.py @@ -25,6 +25,12 @@ class SerializerCompat(SerializerProtocol): def __init__(self, serde: SerializerProtocol) -> None: self.serde = serde + def dumps(self, obj: Any) -> bytes: + return self.serde.dumps(obj) + + def loads(self, data: bytes) -> Any: + return self.serde.loads(data) + def dumps_typed(self, obj: Any) -> tuple[str, bytes]: return type(obj).__name__, self.serde.dumps(obj) diff --git a/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py b/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py index 139a26b67..707442d34 100644 --- a/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py +++ b/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py @@ -16,10 +16,10 @@ from ipaddress import ( IPv6Interface, IPv6Network, ) -from typing import Any, Optional, Sequence +from typing import Any, Callable, Optional, Sequence, Union, cast from uuid import UUID -import msgpack +import msgpack # type: ignore[import-untyped] from langchain_core.load.load import Reviver from langchain_core.load.serializable import Serializable from zoneinfo import ZoneInfo @@ -33,12 +33,12 @@ LC_REVIVER = Reviver() class JsonPlusSerializer(SerializerProtocol): def _encode_constructor_args( self, - constructor: type[Any], + constructor: Union[Callable, type[Any]], *, - method: Optional[str] = None, + method: Union[None, str, Sequence[Union[None, str]]] = None, args: Optional[Sequence[Any]] = None, kwargs: Optional[dict[str, Any]] = None, - ): + ) -> dict[str, Any]: out = { "lc": 2, "type": "constructor", @@ -52,9 +52,9 @@ class JsonPlusSerializer(SerializerProtocol): out["kwargs"] = kwargs return out - def _default(self, obj): + def _default(self, obj: Any) -> Union[str, dict[str, Any]]: if isinstance(obj, Serializable): - return obj.to_json() + return cast(dict[str, Any], obj.to_json()) elif hasattr(obj, "model_dump") and callable(obj.model_dump): return self._encode_constructor_args( obj.__class__, method=(None, "model_construct"), kwargs=obj.model_dump() @@ -87,7 +87,10 @@ class JsonPlusSerializer(SerializerProtocol): datetime, method="fromisoformat", args=(obj.isoformat(),) ) elif isinstance(obj, timezone): - return self._encode_constructor_args(timezone, args=obj.__getinitargs__()) + return self._encode_constructor_args( + timezone, + args=obj.__getinitargs__(), # type: ignore[attr-defined] + ) elif isinstance(obj, ZoneInfo): return self._encode_constructor_args(ZoneInfo, args=(obj.key,)) elif isinstance(obj, timedelta): @@ -217,7 +220,7 @@ EXT_PYDANTIC_V1 = 4 EXT_PYDANTIC_V2 = 5 -def _msgpack_default(obj): +def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]: if hasattr(obj, "model_dump") and callable(obj.model_dump): # pydantic v2 return msgpack.ExtType( EXT_PYDANTIC_V2, @@ -360,7 +363,7 @@ def _msgpack_default(obj): ( obj.__class__.__module__, obj.__class__.__name__, - obj.__getinitargs__(), + obj.__getinitargs__(), # type: ignore[attr-defined] ), ), ) @@ -406,7 +409,7 @@ def _msgpack_default(obj): raise TypeError(f"Object of type {obj.__class__.__name__} is not serializable") -def _msgpack_ext_hook(code: int, data: bytes): +def _msgpack_ext_hook(code: int, data: bytes) -> Any: if code == EXT_CONSTRUCTOR_SINGLE_ARG: try: tup = msgpack.unpackb(data, ext_hook=_msgpack_ext_hook) @@ -461,7 +464,7 @@ def _msgpack_ext_hook(code: int, data: bytes): return -ENC_POOL = deque(maxlen=32) +ENC_POOL: deque[msgpack.Packer] = deque(maxlen=32) def _msgpack_enc(data: Any) -> bytes: diff --git a/libs/checkpoint/langgraph/checkpoint/serde/types.py b/libs/checkpoint/langgraph/checkpoint/serde/types.py index 3fc82b68d..f86c2e558 100644 --- a/libs/checkpoint/langgraph/checkpoint/serde/types.py +++ b/libs/checkpoint/langgraph/checkpoint/serde/types.py @@ -16,8 +16,8 @@ ERROR = "__error__" SCHEDULED = "__scheduled__" TASKS = "__pregel_tasks" -Value = TypeVar("Value") -Update = TypeVar("Update") +Value = TypeVar("Value", covariant=True) +Update = TypeVar("Update", contravariant=True) C = TypeVar("C") diff --git a/libs/checkpoint/pyproject.toml b/libs/checkpoint/pyproject.toml index 5b77649fe..0d2d3044e 100644 --- a/libs/checkpoint/pyproject.toml +++ b/libs/checkpoint/pyproject.toml @@ -53,3 +53,13 @@ now = true delay = 0.1 runner_args = ["--ff", "-v", "--tb", "short"] patterns = ["*.py"] + +[tool.mypy] +# https://mypy.readthedocs.io/en/stable/config_file.html +disallow_untyped_defs = "True" +explicit_package_bases = "True" +warn_no_return = "False" +warn_unused_ignores = "True" +warn_redundant_casts = "True" +allow_redefinition = "True" +disable_error_code = "typeddict-item, return-value"