From 656f89e16ab7074c972d1920d774bf0852346996 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 16 Aug 2024 16:38:46 -0700 Subject: [PATCH] Lint --- libs/langgraph/langgraph/graph/state.py | 3 +++ libs/langgraph/langgraph/managed/shared_value.py | 9 ++++++++- 2 files changed, 11 insertions(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 1c9242ea1..261fbf1d5 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -43,6 +43,7 @@ from langgraph.graph.graph import ( from langgraph.kv.base import BaseKV from langgraph.managed.base import ( ChannelKeyPlaceholder, + ChannelTypePlaceholder, ConfiguredManagedValue, ManagedValue, is_managed_value, @@ -760,6 +761,8 @@ def _is_field_managed_value(name: str, typ: Type[Any]) -> Optional[Type[ManagedV for k, v in decoration.kwargs.items(): if v is ChannelKeyPlaceholder: decoration.kwargs[k] = name + if v is ChannelTypePlaceholder: + decoration.kwargs[k] = typ.__origin__ return decoration return None diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index f1b8d9516..7c55fc858 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -17,6 +17,7 @@ from langgraph.errors import InvalidUpdateError from langgraph.kv.base import BaseKV from langgraph.managed.base import ( ChannelKeyPlaceholder, + ChannelTypePlaceholder, ConfiguredManagedValue, WritableManagedValue, ) @@ -43,7 +44,12 @@ class SharedValue(WritableManagedValue[Value, Update]): @staticmethod def on(scope: str) -> ConfiguredManagedValue: return ConfiguredManagedValue( - SharedValue, {"scope": scope, "key": ChannelKeyPlaceholder} + SharedValue, + { + "scope": scope, + "key": ChannelKeyPlaceholder, + "typ": ChannelTypePlaceholder, + }, ) @classmethod @@ -68,6 +74,7 @@ class SharedValue(WritableManagedValue[Value, Update]): self, config: RunnableConfig, *, typ: Type[Any], scope: str, key: str ) -> None: if typ := _strip_extras(typ): + print(typ) if typ not in ( dict, collections.abc.Mapping,