Support returning non-dict objects as state updates from nodes

This commit is contained in:
Nuno Campos
2024-05-30 08:38:40 -07:00
parent 978d7aa539
commit e4c32248aa
3 changed files with 19 additions and 5 deletions
+7 -3
View File
@@ -1,3 +1,4 @@
from dataclasses import dataclass
import logging
import typing
import warnings
@@ -291,10 +292,13 @@ class CompiledStateGraph(CompiledGraph):
def _get_state_key(input: dict, config: RunnableConfig, *, key: str) -> Any:
if input is None:
return SKIP_WRITE
elif not isinstance(input, dict):
raise InvalidUpdateError(f"Expected dict, got {input}")
else:
elif isinstance(input, dict):
return input.get(key, SKIP_WRITE)
elif get_type_hints(type(input)).get(key):
value = getattr(input, key, SKIP_WRITE)
return value if value is not None else SKIP_WRITE
else:
raise InvalidUpdateError(f"Expected dict, got {input}")
# state updaters
state_write_entries = (
+6 -1
View File
@@ -6029,11 +6029,16 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class(
answer: Optional[str] = None
docs: Annotated[list[str], sorted_add]
class StateUpdate(BaseModel):
query: Optional[str] = None
answer: Optional[str] = None
docs: Optional[list[str]] = None
def rewrite_query(data: State) -> State:
return {"query": f"query: {data.query}"}
def analyzer_one(data: State) -> State:
return {"query": f"analyzed: {data.query}"}
return StateUpdate(query=f"analyzed: {data.query}")
def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]}
+6 -1
View File
@@ -4524,11 +4524,16 @@ async def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class(
answer: Optional[str] = None
docs: Annotated[list[str], sorted_add]
class StateUpdate(BaseModel):
query: Optional[str] = None
answer: Optional[str] = None
docs: Optional[list[str]] = None
async def rewrite_query(data: State) -> State:
return {"query": f"query: {data.query}"}
async def analyzer_one(data: State) -> State:
return {"query": f"analyzed: {data.query}"}
return StateUpdate(query=f"analyzed: {data.query}")
async def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]}