mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-13 13:17:52 +02:00
Add more comments
This commit is contained in:
+34
-40
@@ -1,9 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC
|
||||
from collections import defaultdict
|
||||
from concurrent.futures import CancelledError, Future
|
||||
from functools import partial
|
||||
from itertools import groupby
|
||||
from typing import Any, Iterator, List, Optional, Sequence, Set, TypeVar
|
||||
|
||||
from langchain.callbacks.manager import CallbackManagerForChainRun
|
||||
@@ -39,10 +39,6 @@ class PubSub(Runnable[Any, Any], ABC):
|
||||
self.processes = processes
|
||||
self.connection = connection
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
revalidate_instances = False
|
||||
|
||||
def with_retry(self, **kwargs: Any) -> Runnable[Any, Any]:
|
||||
return self.__class__(
|
||||
processes=[p.with_retry(**kwargs) for p in self.processes],
|
||||
@@ -55,13 +51,13 @@ class PubSub(Runnable[Any, Any], ABC):
|
||||
run_manager: CallbackManagerForChainRun,
|
||||
config: RunnableConfig,
|
||||
) -> Iterator[Any]:
|
||||
subscribers: list[RunnableSubscriber[Any]] = []
|
||||
reducers: list[RunnableReducer[Any]] = []
|
||||
subscribers: defaultdict[str, list[RunnableSubscriber[Any]]] = defaultdict(list)
|
||||
reducers: defaultdict[str, list[RunnableReducer[Any]]] = defaultdict(list)
|
||||
for process in self.processes:
|
||||
if isinstance(process, RunnableReducer):
|
||||
reducers.append(process)
|
||||
reducers[process.topic.name].append(process)
|
||||
elif isinstance(process, RunnableSubscriber):
|
||||
subscribers.append(process)
|
||||
subscribers[process.topic.name].append(process)
|
||||
else:
|
||||
raise ValueError(f"Unknown process type: {process}")
|
||||
|
||||
@@ -74,7 +70,7 @@ class PubSub(Runnable[Any, Any], ABC):
|
||||
input_value += chunk
|
||||
|
||||
with get_executor_for_config(config) as executor:
|
||||
# Namespace topics for each run, default to run_id, ie. isolate runs
|
||||
# Namespace topics for each run, default to run_id, ie. isolated
|
||||
topic_prefix = str(config.get("state_id") or run_manager.run_id)
|
||||
# Track inflight futures
|
||||
inflight: Set[Future] = set()
|
||||
@@ -82,17 +78,19 @@ class PubSub(Runnable[Any, Any], ABC):
|
||||
exceptions: List[Exception] = []
|
||||
|
||||
def on_idle():
|
||||
"""Called when all subscribed topics are empty.
|
||||
It first runs any topic reducers. Then, if all subscribed topics
|
||||
still empty, it closes the computation.
|
||||
"""
|
||||
if reducers:
|
||||
reducers_by_topic = groupby(
|
||||
sorted(reducers, key=lambda p: p.topic.name),
|
||||
lambda p: p.topic.name,
|
||||
)
|
||||
for topic_name, processes in reducers_by_topic:
|
||||
for topic_name, processes in reducers.items():
|
||||
# Collect all pending messages for each topic
|
||||
messages = list(
|
||||
self.connection.iterate(
|
||||
topic_prefix, topic_name, wait=False
|
||||
)
|
||||
)
|
||||
# Run each reducer once with the collected messages
|
||||
if messages:
|
||||
for process in processes:
|
||||
run_once(process, messages)
|
||||
@@ -100,28 +98,28 @@ class PubSub(Runnable[Any, Any], ABC):
|
||||
if not inflight:
|
||||
self.connection.disconnect(topic_prefix)
|
||||
|
||||
def check_if_idle(fut: Future) -> None:
|
||||
"""Cleanup after a process runs."""
|
||||
inflight.discard(fut)
|
||||
|
||||
try:
|
||||
exc = fut.exception()
|
||||
except CancelledError:
|
||||
exc = None
|
||||
except Exception as e:
|
||||
exc = e
|
||||
if exc is not None:
|
||||
exceptions.append(exc)
|
||||
|
||||
# Close output iterator if
|
||||
# - all processes are done, or
|
||||
# - an exception occurred
|
||||
if not inflight or exc is not None:
|
||||
on_idle()
|
||||
|
||||
def run_once(process: RunnableSubscriber[Any], value: Any) -> None:
|
||||
"""Run a process once."""
|
||||
|
||||
def cleanup_run(fut: Future) -> None:
|
||||
"""Cleanup after a process runs."""
|
||||
inflight.discard(fut)
|
||||
|
||||
try:
|
||||
exc = fut.exception()
|
||||
except CancelledError:
|
||||
exc = None
|
||||
except Exception as e:
|
||||
exc = e
|
||||
if exc is not None:
|
||||
exceptions.append(exc)
|
||||
|
||||
# Close output iterator if
|
||||
# - all processes are done, or
|
||||
# - an exception occurred
|
||||
if not inflight or exc is not None:
|
||||
on_idle()
|
||||
|
||||
def get(topic_name: str) -> Any:
|
||||
if topic_name == INPUT_TOPIC:
|
||||
return input_value
|
||||
@@ -149,14 +147,10 @@ class PubSub(Runnable[Any, Any], ABC):
|
||||
|
||||
# Add callback to cleanup
|
||||
inflight.add(fut)
|
||||
fut.add_done_callback(cleanup_run)
|
||||
fut.add_done_callback(check_if_idle)
|
||||
|
||||
# Listen on all subscribed topics
|
||||
processes_by_topic = groupby(
|
||||
sorted(subscribers, key=lambda p: p.topic.name),
|
||||
lambda p: p.topic.name,
|
||||
)
|
||||
for topic_name, processes in processes_by_topic:
|
||||
for topic_name, processes in subscribers.items():
|
||||
self.connection.listen(
|
||||
topic_prefix,
|
||||
topic_name,
|
||||
|
||||
@@ -171,7 +171,7 @@ def test_invoke_two_processes_two_in_two_out(mocker: MockerFixture):
|
||||
|
||||
def test_invoke_two_processes_two_in_join_two_out(mocker: MockerFixture):
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
add_10_each = mocker.Mock(side_effect=lambda x: [y + 10 for y in x])
|
||||
add_10_each = mocker.Mock(side_effect=lambda x: sorted(y + 10 for y in x))
|
||||
topic_one = Topic("one")
|
||||
topic_two = Topic("two")
|
||||
chain_one = Topic.IN.subscribe() | add_one | topic_one.publish()
|
||||
|
||||
Reference in New Issue
Block a user