mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-12 20:57:52 +02:00
Remove dict subclasses used for values/updates stream chunks (#4816)
This commit is contained in:
@@ -2,8 +2,6 @@ from collections import Counter
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from typing import Any, Literal, Optional, TypeVar, Union
|
||||
|
||||
from langchain_core.runnables.utils import AddableDict
|
||||
|
||||
from langgraph.channels.base import BaseChannel, EmptyChannelError
|
||||
from langgraph.constants import (
|
||||
EMPTY_SEQ,
|
||||
@@ -99,14 +97,6 @@ def map_input(
|
||||
logger.warning(f"Input channel {k} not found in {input_channels}")
|
||||
|
||||
|
||||
class AddableValuesDict(AddableDict):
|
||||
def __add__(self, other: dict[str, Any]) -> "AddableValuesDict":
|
||||
return self | other
|
||||
|
||||
def __radd__(self, other: dict[str, Any]) -> "AddableValuesDict":
|
||||
return other | self
|
||||
|
||||
|
||||
def map_output_values(
|
||||
output_channels: Union[str, Sequence[str]],
|
||||
pending_writes: Union[Literal[True], Sequence[tuple[str, Any]]],
|
||||
@@ -122,15 +112,7 @@ def map_output_values(
|
||||
if pending_writes is True or {
|
||||
c for c, _ in pending_writes if c in output_channels
|
||||
}:
|
||||
yield AddableValuesDict(read_channels(channels, output_channels))
|
||||
|
||||
|
||||
class AddableUpdatesDict(AddableDict):
|
||||
def __add__(self, other: dict[str, Any]) -> "AddableUpdatesDict":
|
||||
return [self, other]
|
||||
|
||||
def __radd__(self, other: dict[str, Any]) -> "AddableUpdatesDict":
|
||||
raise TypeError("AddableUpdatesDict does not support right-side addition")
|
||||
yield read_channels(channels, output_channels)
|
||||
|
||||
|
||||
def map_output_updates(
|
||||
@@ -179,17 +161,17 @@ def map_output_updates(
|
||||
},
|
||||
)
|
||||
)
|
||||
grouped: dict[str, list[Any]] = {t.name: [] for t, _ in output_tasks}
|
||||
grouped: dict[str, Any] = {t.name: [] for t, _ in output_tasks}
|
||||
for node, value in updated:
|
||||
grouped[node].append(value)
|
||||
for node, value in grouped.items():
|
||||
if len(value) == 0:
|
||||
grouped[node] = None # type: ignore[assignment]
|
||||
grouped[node] = None
|
||||
if len(value) == 1:
|
||||
grouped[node] = value[0]
|
||||
if cached:
|
||||
grouped["__metadata__"] = {"cached": cached} # type: ignore[assignment]
|
||||
yield AddableUpdatesDict(grouped)
|
||||
grouped["__metadata__"] = {"cached": cached}
|
||||
yield grouped
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
@@ -5456,10 +5456,7 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None:
|
||||
if event["event"] == "on_chain_end" and event["run_id"] == str(UUID(int=0)):
|
||||
times_called += 1
|
||||
assert event["data"] == {
|
||||
"output": [
|
||||
{"inner": {"my_key": "my value there"}},
|
||||
{"side": {"my_key": "my value there and back again"}},
|
||||
]
|
||||
"output": {"side": {"my_key": "my value there and back again"}}
|
||||
}
|
||||
assert times_called == 1
|
||||
|
||||
|
||||
Reference in New Issue
Block a user