mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-30 19:59:40 +02:00
Speed up task triggers check
- Using a sentinel value is faster than raising-catching an exception
This commit is contained in:
@@ -3,6 +3,7 @@ from typing import Any, Generic, Optional, Sequence, TypeVar
|
||||
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.constants import MISSING
|
||||
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
||||
|
||||
Value = TypeVar("Value")
|
||||
@@ -64,6 +65,14 @@ class BaseChannel(Generic[Value, Update, C], ABC):
|
||||
"""
|
||||
return False
|
||||
|
||||
def get_catch(self) -> Value:
|
||||
"""Return the current value of the channel, or MISSING if the channel
|
||||
is empty. Subclasses can override to skip the EmptyChannelError check."""
|
||||
try:
|
||||
return self.get()
|
||||
except EmptyChannelError:
|
||||
return MISSING
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BaseChannel",
|
||||
|
||||
@@ -3,6 +3,7 @@ from typing import Any, Generic, Optional, Sequence, Type
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.channels.base import BaseChannel, Value
|
||||
from langgraph.constants import MISSING
|
||||
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
||||
|
||||
|
||||
@@ -14,6 +15,7 @@ class EphemeralValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
||||
def __init__(self, typ: Any, guard: bool = True) -> None:
|
||||
super().__init__(typ)
|
||||
self.guard = guard
|
||||
self.value = MISSING
|
||||
|
||||
def __eq__(self, value: object) -> bool:
|
||||
return isinstance(value, EphemeralValue) and value.guard == self.guard
|
||||
@@ -37,10 +39,10 @@ class EphemeralValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
||||
|
||||
def update(self, values: Sequence[Value]) -> bool:
|
||||
if len(values) == 0:
|
||||
try:
|
||||
del self.value
|
||||
if self.value is not MISSING:
|
||||
self.value = MISSING
|
||||
return True
|
||||
except AttributeError:
|
||||
else:
|
||||
return False
|
||||
if len(values) != 1 and self.guard:
|
||||
raise InvalidUpdateError(
|
||||
@@ -51,7 +53,9 @@ class EphemeralValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
||||
return True
|
||||
|
||||
def get(self) -> Value:
|
||||
try:
|
||||
return self.value
|
||||
except AttributeError:
|
||||
if self.value is MISSING:
|
||||
raise EmptyChannelError()
|
||||
return self.value
|
||||
|
||||
def get_catch(self) -> Value:
|
||||
return self.value
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from typing import Generic, Optional, Sequence, Type
|
||||
from typing import Any, Generic, Optional, Sequence, Type
|
||||
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.channels.base import BaseChannel, Value
|
||||
from langgraph.constants import MISSING
|
||||
from langgraph.errors import (
|
||||
EmptyChannelError,
|
||||
ErrorCode,
|
||||
@@ -16,6 +17,10 @@ class LastValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
||||
|
||||
__slots__ = ("value",)
|
||||
|
||||
def __init__(self, typ: Any, key: str = "") -> None:
|
||||
super().__init__(typ, key)
|
||||
self.value = MISSING
|
||||
|
||||
def __eq__(self, value: object) -> bool:
|
||||
return isinstance(value, LastValue)
|
||||
|
||||
@@ -50,7 +55,9 @@ class LastValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
||||
return True
|
||||
|
||||
def get(self) -> Value:
|
||||
try:
|
||||
return self.value
|
||||
except AttributeError:
|
||||
if self.value is MISSING:
|
||||
raise EmptyChannelError()
|
||||
return self.value
|
||||
|
||||
def get_catch(self) -> Value:
|
||||
return self.value
|
||||
|
||||
@@ -3,6 +3,7 @@ from typing import Generic, Optional, Sequence, Type
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.channels.base import BaseChannel, Value
|
||||
from langgraph.constants import MISSING
|
||||
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
||||
|
||||
|
||||
@@ -14,6 +15,7 @@ class UntrackedValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
||||
def __init__(self, typ: Type[Value], guard: bool = True) -> None:
|
||||
super().__init__(typ)
|
||||
self.guard = guard
|
||||
self.value = MISSING
|
||||
|
||||
def __eq__(self, value: object) -> bool:
|
||||
return isinstance(value, UntrackedValue) and value.guard == self.guard
|
||||
@@ -48,7 +50,9 @@ class UntrackedValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
||||
return True
|
||||
|
||||
def get(self) -> Value:
|
||||
try:
|
||||
return self.value
|
||||
except AttributeError:
|
||||
if self.value is MISSING:
|
||||
raise EmptyChannelError()
|
||||
return self.value
|
||||
|
||||
def get_catch(self) -> Value:
|
||||
return self.value
|
||||
|
||||
@@ -48,6 +48,7 @@ from langgraph.constants import (
|
||||
EMPTY_SEQ,
|
||||
ERROR,
|
||||
INTERRUPT,
|
||||
MISSING,
|
||||
NO_WRITES,
|
||||
NS_END,
|
||||
NS_SEP,
|
||||
@@ -649,9 +650,7 @@ def prepare_single_task(
|
||||
if triggers := sorted(
|
||||
chan
|
||||
for chan in proc.triggers
|
||||
if not isinstance(
|
||||
read_channel(channels, chan, return_exception=True), EmptyChannelError
|
||||
)
|
||||
if channels[chan].get_catch() is not MISSING
|
||||
and checkpoint["channel_versions"].get(chan, null_version) # type: ignore[operator]
|
||||
> seen.get(chan, null_version)
|
||||
):
|
||||
|
||||
@@ -38,14 +38,11 @@ def read_channel(
|
||||
chan: str,
|
||||
*,
|
||||
catch: bool = True,
|
||||
return_exception: bool = False,
|
||||
) -> Any:
|
||||
try:
|
||||
return channels[chan].get()
|
||||
except EmptyChannelError as exc:
|
||||
if return_exception:
|
||||
return exc
|
||||
elif catch:
|
||||
except EmptyChannelError:
|
||||
if catch:
|
||||
return None
|
||||
else:
|
||||
raise
|
||||
|
||||
Reference in New Issue
Block a user