mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
ci: Enable mypy checks for checkpoint lib
This commit is contained in:
@@ -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")
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user