From 04d3c9d30fea7fdc7a65f67d96980c9895cbfa58 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 11 Apr 2025 08:38:14 -0700 Subject: [PATCH] Use cache in attach_branch too --- libs/langgraph/langgraph/graph/state.py | 50 ++++++++++++------------- 1 file changed, 25 insertions(+), 25 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 9555be1be..fde945709 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -916,19 +916,33 @@ class CompiledStateGraph(CompiledGraph): config, cast(Sequence[Union[Send, ChannelWriteEntry]], writes) ) - schema = branch.input_schema or ( - self.builder.nodes[start].input - if start in self.builder.nodes - else self.builder.schema - ) + if with_reader: + # get schema + schema = branch.input_schema or ( + self.builder.nodes[start].input + if start in self.builder.nodes + else self.builder.schema + ) + channels = list(self.builder.schemas[schema]) + # get mapper + if schema in self.schema_to_mapper: + mapper = self.schema_to_mapper[schema] + else: + mapper = _pick_mapper(channels, schema, self.builder.type_hints[schema]) + self.schema_to_mapper[schema] = mapper + # create reader + reader: Optional[Callable[[RunnableConfig], Any]] = partial( + ChannelRead.do_read, + select=channels[0] if channels == ["__root__"] else channels, + fresh=True, + # coerce state dict to schema class (eg. pydantic model) + mapper=mapper, + ) + else: + reader = None # attach branch publisher - self.nodes[start].writers.append( - branch.run( - branch_writer, - _get_state_reader(self.builder, schema) if with_reader else None, - ) - ) + self.nodes[start].writers.append(branch.run(branch_writer, reader)) # attach then subscriber if branch.then and branch.then != END: @@ -1053,20 +1067,6 @@ class CompiledStateGraph(CompiledGraph): seen[INTERRUPT].pop(k, MISSING) -def _get_state_reader( - builder: StateGraph, schema: Type[Any] -) -> Callable[[RunnableConfig], Any]: - state_keys = list(builder.channels) - select = list(builder.schemas[schema]) - return partial( - ChannelRead.do_read, - select=select[0] if select == ["__root__"] else select, - fresh=True, - # coerce state dict to schema class (eg. pydantic model) - mapper=_pick_mapper(state_keys, schema, builder.type_hints[schema]), - ) - - def _pick_mapper( state_keys: Sequence[str], schema: Type[Any], type_hints: Optional[dict[str, Any]] ) -> Optional[Callable[[Any], Any]]: