Remove pydantic base class from Pregel/PregelNode

This commit is contained in:
Nuno Campos
2024-09-03 09:20:17 -07:00
committed by Nuno Campos
parent 43d5dbf174
commit 3c5913aa29
6 changed files with 155 additions and 139 deletions
+4
View File
@@ -424,6 +424,10 @@ class Graph:
class CompiledGraph(Pregel):
builder: Graph
def __init__(self, *, builder: Graph, **kwargs):
super().__init__(**kwargs)
self.builder = builder
def attach_node(self, key: str, node: NodeSpec) -> None:
self.channels[key] = EphemeralValue(Any)
self.nodes[key] = (
+1 -1
View File
@@ -17,10 +17,10 @@ from typing import (
overload,
)
from langchain_core.pydantic_v1 import BaseModel
from langchain_core.runnables import Runnable, RunnableConfig
from langchain_core.runnables.base import RunnableLike
from langchain_core.runnables.utils import create_model
from pydantic import BaseModel
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import BinaryOperatorAggregate
+48 -41
View File
@@ -25,12 +25,10 @@ from uuid import UUID, uuid5
from langchain_core.globals import get_debug
from langchain_core.load.dump import dumpd
from langchain_core.pydantic_v1 import BaseModel, Field, root_validator
from langchain_core.runnables import (
Runnable,
RunnableLambda,
RunnableSequence,
RunnableSerializable,
)
from langchain_core.runnables.base import Input, Output, coerce_to_runnable
from langchain_core.runnables.config import (
@@ -48,6 +46,7 @@ from langchain_core.runnables.utils import (
get_unique_config_specs,
)
from langchain_core.tracers._streaming import _StreamingCallbackHandler
from pydantic import BaseModel
from typing_extensions import Self
from langgraph.channels.base import (
@@ -186,16 +185,10 @@ class Channel:
)
class Pregel(
RunnableSerializable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]
):
class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
nodes: Mapping[str, PregelNode]
channels: Mapping[str, Union[BaseChannel, ManagedValueSpec]] = Field(
default_factory=dict
)
auto_validate: bool = True
channels: Mapping[str, Union[BaseChannel, ManagedValueSpec]]
stream_mode: StreamMode = "values"
"""Mode to stream output, defaults to 'values'."""
@@ -205,16 +198,16 @@ class Pregel(
stream_channels: Optional[Union[str, Sequence[str]]] = None
"""Channels to stream, defaults to all channels not in reserved channels"""
interrupt_after_nodes: Union[All, Sequence[str]] = Field(default_factory=list)
interrupt_after_nodes: Union[All, Sequence[str]]
interrupt_before_nodes: Union[All, Sequence[str]] = Field(default_factory=list)
interrupt_before_nodes: Union[All, Sequence[str]]
input_channels: Union[str, Sequence[str]]
step_timeout: Optional[float] = None
"""Maximum time to wait for a step to complete, in seconds. Defaults to None."""
debug: bool = Field(default_factory=get_debug)
debug: bool
"""Whether to print debug information during execution. Defaults to False."""
checkpointer: Optional[BaseCheckpointSaver] = None
@@ -232,36 +225,50 @@ class Pregel(
name: str = "LangGraph"
class Config:
arbitrary_types_allowed = True
def __init__(
self,
*,
nodes: Mapping[str, PregelNode],
channels: Mapping[str, Union[BaseChannel, ManagedValueSpec]] = None,
auto_validate: bool = True,
stream_mode: StreamMode = "values",
output_channels: Union[str, Sequence[str]],
stream_channels: Optional[Union[str, Sequence[str]]] = None,
interrupt_after_nodes: Union[All, Sequence[str]] = (),
interrupt_before_nodes: Union[All, Sequence[str]] = (),
input_channels: Union[str, Sequence[str]],
step_timeout: Optional[float] = None,
debug: bool = False,
checkpointer: Optional[BaseCheckpointSaver] = None,
store: Optional[BaseStore] = None,
retry_policy: Optional[RetryPolicy] = None,
config_type: Optional[Type[Any]] = None,
config: Optional[RunnableConfig] = None,
name: str = "LangGraph",
) -> None:
self.nodes = nodes
self.channels = channels or {}
self.stream_mode = stream_mode
self.output_channels = output_channels
self.stream_channels = stream_channels
self.interrupt_after_nodes = interrupt_after_nodes
self.interrupt_before_nodes = interrupt_before_nodes
self.input_channels = input_channels
self.step_timeout = step_timeout
self.debug = debug if debug is not None else get_debug()
self.checkpointer = checkpointer
self.store = store
self.retry_policy = retry_policy
self.config_type = config_type
self.config = config
self.name = name
if auto_validate:
self.validate()
def with_config(self, config: RunnableConfig | None = None, **kwargs: Any) -> Self:
return self.copy(
update={"config": cast(RunnableConfig, {**(config or {}), **kwargs})}
)
@classmethod
def is_lc_serializable(cls) -> bool:
"""Return whether the graph can be serialized by Langchain."""
return True
@root_validator(skip_on_failure=True)
def validate_on_init(cls, values: dict[str, Any]) -> dict[str, Any]:
if not values["auto_validate"]:
return values
validate_graph(
values["nodes"],
values["channels"],
values["input_channels"],
values["output_channels"],
values["stream_channels"],
values["interrupt_after_nodes"],
values["interrupt_before_nodes"],
)
if values["interrupt_after_nodes"] or values["interrupt_before_nodes"]:
if not values["checkpointer"]:
raise ValueError("Interrupts require a checkpointer")
return values
attrs = {**self.__dict__}
attrs["config"] = merge_configs(self.config, config, kwargs)
return self.__class__(**attrs)
def validate(self) -> Self:
validate_graph(
+84 -42
View File
@@ -1,8 +1,16 @@
from __future__ import annotations
from typing import Any, Callable, Mapping, Optional, Sequence, Union
from typing import (
Any,
AsyncIterator,
Callable,
Iterator,
Mapping,
Optional,
Sequence,
Union,
)
from langchain_core.pydantic_v1 import Field
from langchain_core.runnables import (
Runnable,
RunnableConfig,
@@ -10,7 +18,7 @@ from langchain_core.runnables import (
RunnableSequence,
RunnableSerializable,
)
from langchain_core.runnables.base import Other, RunnableBindingBase, coerce_to_runnable
from langchain_core.runnables.base import Input, Other, Output, coerce_to_runnable
from langchain_core.runnables.config import merge_configs
from langchain_core.runnables.utils import ConfigurableFieldSpec
@@ -99,20 +107,47 @@ class ChannelRead(RunnableCallable):
DEFAULT_BOUND: RunnablePassthrough = RunnablePassthrough()
class PregelNode(RunnableBindingBase):
class PregelNode(Runnable):
channels: Union[list[str], Mapping[str, str]]
triggers: list[str] = Field(default_factory=list)
triggers: list[str]
mapper: Optional[Callable[[Any], Any]] = None
mapper: Optional[Callable[[Any], Any]]
writers: list[Runnable] = Field(default_factory=list)
writers: list[Runnable]
bound: Runnable[Any, Any] = Field(default=DEFAULT_BOUND)
bound: Runnable[Any, Any]
kwargs: Mapping[str, Any] = Field(default_factory=dict)
retry_policy: Optional[RetryPolicy]
retry_policy: Optional[RetryPolicy] = None
config: RunnableConfig
def __init__(
self,
*,
channels: Union[list[str], Mapping[str, str]],
triggers: Sequence[str],
mapper: Optional[Callable[[Any], Any]] = None,
writers: Optional[list[Runnable]] = None,
tags: Optional[list[str]] = None,
metadata: Optional[Mapping[str, Any]] = None,
bound: Optional[Runnable[Any, Any]] = None,
retry_policy: Optional[RetryPolicy] = None,
config: Optional[RunnableConfig] = None,
) -> None:
self.channels = channels
self.triggers = list(triggers)
self.mapper = mapper
self.writers = writers or []
self.bound = bound if bound is not None else DEFAULT_BOUND
self.retry_policy = retry_policy
self.config = merge_configs(
config, {"tags": tags or [], "metadata": metadata or {}}
)
def copy(self, update: dict[str, Any]) -> PregelNode:
attrs = {**self.__dict__, **update}
return PregelNode(**attrs)
def get_writers(self) -> list[Runnable]:
"""Get writers with optimizations applied."""
@@ -145,38 +180,6 @@ class PregelNode(RunnableBindingBase):
else:
return self.bound
def __init__(
self,
*,
channels: Union[list[str], Mapping[str, str]],
triggers: Sequence[str],
mapper: Optional[Callable[[Any], Any]] = None,
writers: Optional[list[Runnable]] = None,
tags: Optional[list[str]] = None,
metadata: Optional[Mapping[str, Any]] = None,
bound: Optional[Runnable[Any, Any]] = None,
kwargs: Optional[Mapping[str, Any]] = None,
config: Optional[RunnableConfig] = None,
retry_policy: Optional[RetryPolicy] = None,
**other_kwargs: Any,
) -> None:
super().__init__(
channels=channels,
triggers=triggers,
mapper=mapper,
writers=writers or [],
bound=bound or DEFAULT_BOUND,
kwargs=kwargs or {},
retry_policy=retry_policy,
config=merge_configs(
config, {"tags": tags or [], "metadata": metadata or {}}
),
**other_kwargs,
)
def __repr_args__(self) -> Any:
return [(k, v) for k, v in super().__repr_args__() if k != "bound"]
def join(self, channels: Sequence[str]) -> PregelNode:
assert isinstance(channels, list) or isinstance(
channels, tuple
@@ -226,3 +229,42 @@ class PregelNode(RunnableBindingBase):
],
) -> RunnableSerializable:
raise NotImplementedError()
def invoke(
self,
input: Input,
config: Optional[RunnableConfig] = None,
**kwargs: Optional[Any],
) -> Output:
return self.bound.invoke(input, merge_configs(self.config, config), **kwargs)
async def ainvoke(
self,
input: Input,
config: Optional[RunnableConfig] = None,
**kwargs: Optional[Any],
) -> Output:
return await self.bound.ainvoke(
input, merge_configs(self.config, config), **kwargs
)
def stream(
self,
input: Input,
config: Optional[RunnableConfig] = None,
**kwargs: Optional[Any],
) -> Iterator[Output]:
yield from self.bound.stream(
input, merge_configs(self.config, config), **kwargs
)
async def astream(
self,
input: Input,
config: Optional[RunnableConfig] = None,
**kwargs: Optional[Any],
) -> AsyncIterator[Output]:
async for item in self.bound.astream(
input, merge_configs(self.config, config), **kwargs
):
yield item
@@ -3300,11 +3300,7 @@
'type': 'string',
}),
}),
'required': list([
'query',
'inner',
]),
'title': 'Input',
'title': 'LangGraphInput',
'type': 'object',
})
# ---
@@ -3323,11 +3319,7 @@
'type': 'array',
}),
}),
'required': list([
'answer',
'docs',
]),
'title': 'Output',
'title': 'LangGraphOutput',
'type': 'object',
})
# ---
@@ -3370,11 +3362,7 @@
'type': 'string',
}),
}),
'required': list([
'query',
'inner',
]),
'title': 'Input',
'title': 'LangGraphInput',
'type': 'object',
})
# ---
@@ -3393,11 +3381,7 @@
'type': 'array',
}),
}),
'required': list([
'answer',
'docs',
]),
'title': 'Output',
'title': 'LangGraphOutput',
'type': 'object',
})
# ---
@@ -3440,11 +3424,7 @@
'type': 'string',
}),
}),
'required': list([
'query',
'inner',
]),
'title': 'Input',
'title': 'LangGraphInput',
'type': 'object',
})
# ---
@@ -3463,11 +3443,7 @@
'type': 'array',
}),
}),
'required': list([
'answer',
'docs',
]),
'title': 'Output',
'title': 'LangGraphOutput',
'type': 'object',
})
# ---
@@ -3510,11 +3486,7 @@
'type': 'string',
}),
}),
'required': list([
'query',
'inner',
]),
'title': 'Input',
'title': 'LangGraphInput',
'type': 'object',
})
# ---
@@ -3533,11 +3505,7 @@
'type': 'array',
}),
}),
'required': list([
'answer',
'docs',
]),
'title': 'Output',
'title': 'LangGraphOutput',
'type': 'object',
})
# ---
@@ -3580,11 +3548,7 @@
'type': 'string',
}),
}),
'required': list([
'query',
'inner',
]),
'title': 'Input',
'title': 'LangGraphInput',
'type': 'object',
})
# ---
@@ -3603,11 +3567,7 @@
'type': 'array',
}),
}),
'required': list([
'answer',
'docs',
]),
'title': 'Output',
'title': 'LangGraphOutput',
'type': 'object',
})
# ---
+8 -5
View File
@@ -1903,7 +1903,7 @@ def test_invoke_two_processes_no_in(mocker: MockerFixture) -> None:
one = Channel.subscribe_to("between") | add_one | Channel.write_to("output")
two = Channel.subscribe_to("between") | add_one
with pytest.raises(ValueError):
with pytest.raises(TypeError):
Pregel(nodes={"one": one, "two": two})
@@ -9784,6 +9784,7 @@ def test_doubly_nested_graph_state(
]
@pytest.mark.repeat(10)
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_send_to_nested_graphs(
request: pytest.FixtureRequest, checkpointer_name: str
@@ -9943,10 +9944,12 @@ def test_send_to_nested_graphs(
graph.update_state(outer_state.tasks[1].state, {"subject": "turtles - hohoho"})
# continue past interrupt
assert graph.invoke(None, config=config) == {
"subjects": ["cats", "dogs"],
"jokes": ["Joke about cats - hohoho", "Joke about turtles - hohoho"],
}
assert sorted(
graph.stream(None, config=config), key=lambda d: d["generate_joke"]["jokes"][0]
) == [
{"generate_joke": {"jokes": ["Joke about cats - hohoho"]}},
{"generate_joke": {"jokes": ["Joke about turtles - hohoho"]}},
]
actual_snapshot = graph.get_state(config)
expected_snapshot = StateSnapshot(