ci: Enable mypy checks for checkpoint lib

This commit is contained in:
Nuno Campos
2024-09-18 18:12:11 -07:00
parent d4ba315ae2
commit 4d26c3599f
7 changed files with 55 additions and 24 deletions
+2 -1
View File
@@ -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)
@@ -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):
@@ -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
@@ -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)
@@ -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:
@@ -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")
+10
View File
@@ -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"