This commit is contained in:
Nuno Campos
2024-11-11 13:59:18 -08:00
parent 0cc45f7a35
commit 75c502cd93
3 changed files with 14 additions and 21 deletions
+5
View File
@@ -261,6 +261,11 @@ class Command:
getattr(self, key) == getattr(value, key) for key in self.__all_slots__
)
def copy(self, **kwargs: Any) -> Self:
for slot in self.__all_slots__:
kwargs.setdefault(slot, getattr(self, slot))
return self.__class__(**kwargs)
StreamChunk = tuple[tuple[str, ...], str, Any]
+6 -1
View File
@@ -21,7 +21,8 @@ from langgraph.store.base import BaseStore
from langgraph.store.duckdb import AsyncDuckDBStore, DuckDBStore
from langgraph.store.memory import InMemoryStore
from langgraph.store.postgres import AsyncPostgresStore, PostgresStore
from tests.memory_assert import MemorySaverAssertImmutable
pytest.register_assert_rewrite("tests.memory_assert")
DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5442/"
# TODO: fix this once core is released
@@ -49,6 +50,8 @@ def deterministic_uuids(mocker: MockerFixture) -> MockerFixture:
@pytest.fixture(scope="function")
def checkpointer_memory():
from tests.memory_assert import MemorySaverAssertImmutable
yield MemorySaverAssertImmutable()
@@ -225,6 +228,8 @@ async def awith_checkpointer(
if checkpointer_name is None:
yield None
elif checkpointer_name == "memory":
from tests.memory_assert import MemorySaverAssertImmutable
yield MemorySaverAssertImmutable()
elif checkpointer_name == "sqlite_aio":
async with _checkpointer_sqlite_aio() as checkpointer:
+3 -20
View File
@@ -2034,8 +2034,7 @@ async def test_concurrent_emit_sends() -> None:
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_send_sequences(checkpointer_name: str) -> None:
async def test_send_sequences() -> None:
class Node:
def __init__(self, name: str):
self.name = name
@@ -2048,8 +2047,7 @@ async def test_send_sequences(checkpointer_name: str) -> None:
else ["|".join((self.name, str(state)))]
)
if isinstance(state, GraphCommand):
state.update = update
return state
return state.copy(update=update)
else:
return update
@@ -2084,21 +2082,6 @@ async def test_send_sequences(checkpointer_name: str) -> None:
"3",
]
async with awith_checkpointer(checkpointer_name) as checkpointer:
graph = builder.compile(checkpointer=checkpointer)
thread1 = {"configurable": {"thread_id": "1"}}
assert await graph.ainvoke(["0"], thread1) == [
"0",
"1",
"3.1",
"2|Command(send=Send(node='2', arg=3))",
"2|Command(send=Send(node='2', arg=4))",
"3",
"2|3",
"2|4",
"3",
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_send_react_interrupt(checkpointer_name: str) -> None:
@@ -8695,7 +8678,7 @@ async def test_stream_subgraphs_during_execution(checkpointer_name: str) -> None
{"inner_1": {"my_key": "got here", "my_other_key": ""}},
),
),
(FloatBetween(0.2, 0.3), ((), {"outer_1": {"my_key": " and parallel"}})),
(FloatBetween(0.2, 0.4), ((), {"outer_1": {"my_key": " and parallel"}})),
(
FloatBetween(0.5, 0.7),
(