From 14c22418539d3ffaf65b0516b4a747d7150d171f Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 11 Mar 2025 17:44:24 -0700 Subject: [PATCH] Lint --- libs/langgraph/langgraph/graph/state.py | 2 +- libs/langgraph/langgraph/pregel/loop.py | 12 +++++++----- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 285542bdc..45564334f 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -936,7 +936,7 @@ def _get_state_reader( def _pick_mapper( state_keys: Sequence[str], schema: Type[Any] -) -> Optional[Callable[[Type[Any], dict[str, Any]], dict[str, Any]]]: +) -> Optional[Callable[[Any], Any]]: if state_keys == ["__root__"]: return None if issubclass(schema, dict): diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 1ba444dad..169943718 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -141,7 +141,7 @@ def DuplexStream(*streams: StreamProtocol) -> StreamProtocol: class PregelLoop(LoopProtocol): input: Optional[Any] - input_model: Optional[BaseModel] + input_model: Optional[Type[BaseModel]] checkpointer: Optional[BaseCheckpointSaver] nodes: Mapping[str, PregelNode] specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]] @@ -205,7 +205,7 @@ class PregelLoop(LoopProtocol): 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[BaseModel] = None, + input_model: Optional[Type[BaseModel]] = None, debug: bool = False, ) -> None: super().__init__( @@ -434,7 +434,9 @@ class PregelLoop(LoopProtocol): if self.input is INPUT_SHOULD_VALIDATE: self.input = INPUT_DONE # validate - self.input_model(**read_channels(self.channels, self.stream_keys)) + cast(Type[BaseModel], self.input_model)( + **read_channels(self.channels, self.stream_keys) + ) # produce values output self._emit( "values", map_output_values, self.output_keys, writes, self.channels @@ -859,7 +861,7 @@ class SyncPregelLoop(PregelLoop, ContextManager): 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[BaseModel] = None, + input_model: Optional[Type[BaseModel]] = None, debug: bool = False, ) -> None: super().__init__( @@ -1000,7 +1002,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): 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[BaseModel] = None, + input_model: Optional[Type[BaseModel]] = None, debug: bool = False, ) -> None: super().__init__(