diff --git a/permchain/connection.py b/permchain/connection.py index bfaa71efc..9c8528816 100644 --- a/permchain/connection.py +++ b/permchain/connection.py @@ -2,61 +2,64 @@ import asyncio from abc import ABC, abstractmethod from typing import Any, Callable, Iterator, TypedDict -PubSubListener = Callable[[Any], None] - -class LogMessage(TypedDict): +class PubSubMessage(TypedDict): topic: str value: Any published_at: str correlation_id: str +PubSubListener = Callable[[PubSubMessage], None] + + class PubSubConnection(ABC): - def full_topic_name(self, prefix: str, *parts: str) -> str: + def full_name(self, prefix: str, *parts: str) -> str: """Return the full topic name for a given prefix and topic name.""" return ":".join(map(str, [prefix, *parts])) @abstractmethod - def iterate(self, prefix: str, topic_name: str, wait: bool) -> Iterator[Any]: - """Iterate over all currently queued messages for a topic, consuming them.""" + def observe(self, prefix: str) -> Iterator[PubSubMessage]: + """Iterate over messages for all topics under this prefix, + without affecting listeners/iterators on each topic. + This method waits for new messages to arrive.""" + ... + + @abstractmethod + def iterate( + self, prefix: str, topic: str, *, wait: bool + ) -> Iterator[PubSubMessage]: + """Iterate over all currently queued messages for a topic, consuming them. + Optionally wait for new messages to arrive.""" ... # TODO add aiterate() method @abstractmethod - def listen( - self, prefix: str, topic_name: str, listeners: list[PubSubListener] - ) -> None: + def listen(self, prefix: str, topic: str, listeners: list[PubSubListener]) -> None: ... async def alisten( - self, prefix: str, topic_name: str, listeners: list[PubSubListener] + self, prefix: str, topic: str, listeners: list[PubSubListener] ) -> None: return await asyncio.get_event_loop().run_in_executor( - None, self.listen, prefix, topic_name, listeners + None, self.listen, prefix, topic, listeners ) @abstractmethod - def send(self, prefix: str, topic_name: str, message: Any) -> None: + def send(self, prefix: str, topic: str, value: Any) -> None: ... - async def asend(self, prefix: str, topic_name: str, message: Any) -> None: + async def asend(self, prefix: str, topic: str, value: Any) -> None: return await asyncio.get_event_loop().run_in_executor( - None, self.send, prefix, topic_name, message + None, self.send, prefix, topic, value ) @abstractmethod - def disconnect(self, prefix: str) -> None: + def disconnect(self, name: str) -> None: ... - async def adisconnect(self, prefix: str) -> None: + async def adisconnect(self, name: str) -> None: return await asyncio.get_event_loop().run_in_executor( - None, self.disconnect, prefix + None, self.disconnect, name ) - - @abstractmethod - def peek(self, prefix: str) -> Iterator[LogMessage]: - """Iterate over all previously published messages for all topics, - without consuming them.""" - ... diff --git a/permchain/connection_inmemory.py b/permchain/connection_inmemory.py index 82c98fb33..53b288ddf 100644 --- a/permchain/connection_inmemory.py +++ b/permchain/connection_inmemory.py @@ -2,18 +2,25 @@ import queue import threading from collections import defaultdict from datetime import datetime -from typing import Any, Iterator +from typing import Any, Iterator, cast -from permchain.connection import LogMessage, PubSubConnection, PubSubListener +from permchain.connection import PubSubConnection, PubSubListener, PubSubMessage class IterableQueue(queue.SimpleQueue): done_sentinel = object() - def get(self, block: bool = True, timeout: float | None = None) -> Any: + def put( + self, item: PubSubMessage, block: bool = True, timeout: float | None = None + ) -> None: + return super().put(item, block, timeout) + + def get( + self, block: bool = True, timeout: float | None = None + ) -> PubSubMessage | object: return super().get(block=block, timeout=timeout) - def __iter__(self) -> Iterator[Any]: + def __iter__(self) -> Iterator[PubSubMessage]: return iter(self.get, self.done_sentinel) def close(self) -> None: @@ -22,23 +29,25 @@ class IterableQueue(queue.SimpleQueue): class InMemoryPubSubConnection(PubSubConnection): clear_on_disconnect: bool - logs: defaultdict[str, list[Any]] + logs: defaultdict[str, IterableQueue] topics: defaultdict[str, IterableQueue] listeners: defaultdict[str, list[PubSubListener]] lock: threading.RLock def __init__(self, clear_on_disconnect: bool = True) -> None: self.clear_on_disconnect = clear_on_disconnect - self.logs = defaultdict(list) + self.logs = defaultdict(IterableQueue) self.topics = defaultdict(IterableQueue) self.listeners = defaultdict(list) self.lock = threading.RLock() - def peek(self, prefix: str) -> Iterator[LogMessage]: + def observe(self, prefix: str) -> Iterator[PubSubMessage]: return iter(self.logs[str(prefix)]) - def iterate(self, prefix: str, topic_name: str, wait: bool) -> Iterator[Any]: - topic = self.full_topic_name(prefix, topic_name) + def iterate( + self, prefix: str, topic: str, *, wait: bool + ) -> Iterator[PubSubMessage]: + topic = self.full_name(prefix, topic) # This connection doesn't support iterating over topics with listeners connected with self.lock: @@ -53,56 +62,54 @@ class InMemoryPubSubConnection(PubSubConnection): return iter(self.topics[topic]) - def listen( - self, prefix: str, topic_name: str, listeners: list[PubSubListener] - ) -> None: - topic = self.full_topic_name(prefix, topic_name) - self.disconnect(topic) + def listen(self, prefix: str, topic: str, listeners: list[PubSubListener]) -> None: + full_name = self.full_name(prefix, topic) + self.disconnect(full_name) with self.lock: # Add the listeners for future messages - self.listeners[topic].extend(listeners) + self.listeners[full_name].extend(listeners) # Send any pending messages to the listeners - topic_queue = self.topics[topic] + topic_queue = self.topics[full_name] while not topic_queue.empty(): message = topic_queue.get() if message is not topic_queue.done_sentinel: - for listener in self.listeners[topic]: - listener(message) + for listener in self.listeners[full_name]: + listener(cast(PubSubMessage, message)) - def send(self, prefix: str, topic_name: str, message: Any) -> None: - topic = self.full_topic_name(prefix, topic_name) + def send(self, prefix: str, topic: str, value: Any) -> None: + full_name = self.full_name(prefix, topic) + message = PubSubMessage( + value=value, + topic=topic, + correlation_id=str(prefix), + published_at=datetime.now().isoformat(), + ) + # Add the message to the log + self.logs[str(prefix)].put(message) with self.lock: - # Add the message to the log - self.logs[str(prefix)].append( - LogMessage( - value=message, - topic=topic_name, - correlation_id=str(prefix), - published_at=datetime.now().isoformat(), - ) - ) - listeners = self.listeners[topic] + listeners = self.listeners[full_name] if listeners: # Send the message to listeners if any are connected for listener in listeners: listener(message) else: # Otherwise add the message to the topic queue for later - self.topics[topic].put(message) + self.topics[full_name].put(message) - def disconnect(self, prefix_or_topic: str) -> None: + def disconnect(self, name: str) -> None: with self.lock: - if self.clear_on_disconnect: - if prefix_or_topic in self.logs: - del self.logs[prefix_or_topic] + if name in self.logs: + self.logs[name].close() + if self.clear_on_disconnect: + del self.logs[name] to_delete = [] - for topic, queue in self.topics.items(): - if topic.startswith(prefix_or_topic): - queue.close() + for topic, q in self.topics.items(): + if topic.startswith(name): + q.close() if self.clear_on_disconnect: to_delete.append(topic) # can't delete while iterating @@ -111,7 +118,7 @@ class InMemoryPubSubConnection(PubSubConnection): to_delete = [] for topic in self.listeners: - if topic.startswith(prefix_or_topic): + if topic.startswith(name): to_delete.append(topic) # can't delete while iterating for topic in to_delete: diff --git a/permchain/pubsub.py b/permchain/pubsub.py index 087a80e8d..3769bafd5 100644 --- a/permchain/pubsub.py +++ b/permchain/pubsub.py @@ -8,10 +8,9 @@ from typing import Any, Iterator, List, Optional, Sequence, Set, TypeVar from langchain.callbacks.manager import CallbackManagerForChainRun from langchain.schema.runnable import Runnable, RunnableConfig, patch_config -from langchain.schema.runnable.base import Runnable from langchain.schema.runnable.config import get_executor_for_config -from permchain.connection import PubSubConnection +from permchain.connection import PubSubConnection, PubSubMessage from permchain.constants import CONFIG_GET_KEY, CONFIG_SEND_KEY from permchain.topic import ( INPUT_TOPIC, @@ -84,7 +83,7 @@ class PubSub(Runnable[Any, Any], ABC): # Track inflight futures inflight: Set[Future] = set() # Track exceptions - exceptions: List[Exception] = [] + exceptions: List[BaseException] = [] def on_idle() -> None: """Called when all subscribed topics are empty. @@ -127,9 +126,15 @@ class PubSub(Runnable[Any, Any], ABC): on_idle() def run_once( - process: RunnableSubscriber[Any] | RunnableReducer[Any], value: Any + process: RunnableSubscriber[Any] | RunnableReducer[Any], + messages: PubSubMessage | list[PubSubMessage], ) -> None: """Run a process once.""" + value = ( + [m["value"] for m in messages] + if isinstance(messages, list) + else messages["value"] + ) def get(topic_name: str) -> Any: if topic_name == INPUT_TOPIC: @@ -151,10 +156,10 @@ class PubSub(Runnable[Any, Any], ABC): callbacks=run_manager.get_child(), run_name=f"Topic: {process.topic.name}", ), - "correlation_id": self.connection.full_topic_name( + "correlation_id": self.connection.full_name( topic_prefix, process.topic.name, - self.processes.index(process), + str(self.processes.index(process)), ), CONFIG_SEND_KEY: partial(self.connection.send, topic_prefix), CONFIG_GET_KEY: get, @@ -181,10 +186,11 @@ class PubSub(Runnable[Any, Any], ABC): # Yield output until all processes are done # This blocks the current thread, all other work needs to go # through the executor - for chunk in self.connection.iterate( - topic_prefix, OUTPUT_TOPIC, wait=True - ): + for chunk in self.connection.observe(topic_prefix): yield chunk + if chunk["topic"] == OUTPUT_TOPIC: + self.connection.disconnect(topic_prefix) + break else: on_idle() finally: @@ -201,13 +207,12 @@ class PubSub(Runnable[Any, Any], ABC): input: Any, config: Optional[RunnableConfig] = None, **kwargs: Optional[Any], - ) -> Iterator[Any]: + ) -> Iterator[PubSubMessage]: yield from self._transform_stream_with_config( iter([input]), self._transform, config, **kwargs ) def invoke(self, input: Any, config: Optional[RunnableConfig] = None) -> Any: - collected = [] for chunk in self.stream(input, config): - collected.append(chunk) - return collected + if chunk["topic"] == OUTPUT_TOPIC: + return chunk["value"] diff --git a/tests/test_invoke.py b/tests/test_invoke.py index 8b2d5f5ef..914fdeb1d 100644 --- a/tests/test_invoke.py +++ b/tests/test_invoke.py @@ -1,13 +1,24 @@ +from typing import Iterator from uuid import uuid4 import pytest from pytest_mock import MockerFixture +from permchain.connection import PubSubMessage from permchain.connection_inmemory import InMemoryPubSubConnection from permchain.pubsub import PubSub from permchain.topic import RunnableSubscriber, Topic +def clean_log( + logs: Iterator[PubSubMessage], correlation_id: bool | None = None +) -> list[PubSubMessage]: + if correlation_id is False: + return [{**m, "published_at": None, "correlation_id": None} for m in logs] + else: + return [{**m, "published_at": None} for m in logs] + + def test_invoke_single_process_in_out(mocker: MockerFixture): add_one = mocker.Mock(side_effect=lambda x: x + 1) chain = Topic.IN.subscribe() | add_one | Topic.OUT.publish() @@ -22,7 +33,7 @@ def test_invoke_single_process_in_out(mocker: MockerFixture): # If we start with 0 listeners assert conn.listeners == {} # Then invoke pubsub - assert pubsub.invoke(2) == [3] + assert pubsub.invoke(2) == 3 # After invoke returns the listeners were cleaned up assert conn.listeners == {} @@ -44,7 +55,7 @@ def test_invoke_two_processes_in_out(mocker: MockerFixture): # If we start with 0 listeners assert conn.listeners == {} # Then invoke pubsub - assert pubsub.invoke(2) == [4] + assert pubsub.invoke(2) == 4 # After invoke returns the listeners were cleaned up assert conn.listeners == {} @@ -71,12 +82,7 @@ def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture): correlation_id = uuid4() # invoke() step 1 - assert pubsub_one.invoke(2, {"correlation_id": correlation_id}) == [] - - # listeners are still cleared, even though state is preserved - assert conn.listeners == {} - # The log contains all messages published to all topics, in order - assert [{**m, "published_at": None} for m in conn.peek(correlation_id)] == [ + assert clean_log(pubsub_one.stream(2, {"correlation_id": correlation_id})) == [ { "value": 2, "topic": "__in__", @@ -90,33 +96,17 @@ def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture): "published_at": None, }, ] - # IN, OUT, one - assert len(conn.topics) == 3 - topic_one_full_name = conn.full_topic_name(correlation_id, topic_one.name) + + # IN, one + assert len(conn.topics) == 2 + topic_one_full_name = conn.full_name(correlation_id, topic_one.name) # the actual message publishd by chain_one, and a sentinel "end" value assert conn.topics[topic_one_full_name].qsize() == 2 # invoke() step 2 # this picks up where the first left off, and produces same result as # `test_invoke_two_processes_in_out` - assert pubsub_two.invoke(None, {"correlation_id": correlation_id}) == [4] - - # listeners are still cleared, even though state is preserved - assert conn.listeners == {} - # The log contains all messages published to all topics, in order - assert [{**m, "published_at": None} for m in conn.peek(correlation_id)] == [ - { - "value": 2, - "topic": "__in__", - "correlation_id": str(correlation_id), - "published_at": None, - }, - { - "value": 3, - "topic": "one", - "correlation_id": str(correlation_id), - "published_at": None, - }, + assert clean_log(pubsub_two.stream(None, {"correlation_id": correlation_id})) == [ { "value": None, "topic": "__in__", @@ -130,6 +120,8 @@ def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture): "published_at": None, }, ] + # listeners are still cleared, even though state is preserved + assert conn.listeners == {} # IN, OUT, one assert len(conn.topics) == 3 @@ -142,8 +134,8 @@ def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture): # Empty because this was consumed by invoke() assert queue.qsize() == 0 if topic_name.endswith("one"): - # Contains sentinel "end" value - assert queue.qsize() == 1 + # Contains 2 sentinel "end" values + assert queue.qsize() == 2 def test_invoke_many_processes_in_out(mocker: MockerFixture): @@ -170,7 +162,7 @@ def test_invoke_many_processes_in_out(mocker: MockerFixture): # If we start with 0 listeners assert conn.listeners == {} # Then invoke pubsub - assert pubsub.invoke(2) == [2 + test_size] + assert pubsub.invoke(2) == 2 + test_size # After invoke returns the listeners were cleaned up assert conn.listeners == {} @@ -192,8 +184,9 @@ def test_invoke_two_processes_two_in_two_out(mocker: MockerFixture): assert conn.listeners == {} # Then invoke pubsub - # We get two equal results as the two chains do the same thing - assert pubsub.invoke(2) == [3, 3] + # We get only one of the two return values, as computation is closed + # as soon as we publish to OUT for the first time + assert pubsub.invoke(2) == 3 # After invoke returns the listeners were cleaned up assert conn.listeners == {} @@ -223,7 +216,7 @@ def test_invoke_two_processes_two_in_join_two_out(mocker: MockerFixture): # Then invoke pubsub # We get a single array result as chain_four waits for all publishers to finish # before operating on all elements published to topic_two as an array - assert pubsub.invoke(2) == [[13, 14]] + assert pubsub.invoke(2) == [13, 14] # After invoke returns the listeners were cleaned up assert conn.listeners == {} @@ -255,8 +248,7 @@ def test_invoke_join_then_subscribe(mocker: MockerFixture): # Then invoke pubsub # We get a single array result as chain_four waits for all publishers to finish # before operating on all elements published to topic_two as an array - assert pubsub.invoke([2, 3], {"correlation_id": correlation_id}) == [26] - assert [{**m, "published_at": None} for m in conn.peek(correlation_id)] == [ + assert clean_log(pubsub.stream([2, 3], {"correlation_id": correlation_id})) == [ { "value": [2, 3], "topic": "__in__", @@ -293,7 +285,6 @@ def test_invoke_join_then_subscribe(mocker: MockerFixture): assert conn.listeners == {} -@pytest.mark.skip("TODO") def test_invoke_join_then_call_other_pubsub(mocker: MockerFixture): conn = InMemoryPubSubConnection(clear_on_disconnect=False) add_one = mocker.Mock(side_effect=lambda x: x + 1) @@ -308,8 +299,8 @@ def test_invoke_join_then_call_other_pubsub(mocker: MockerFixture): topic_two = Topic("two") chain_one = Topic.IN.subscribe() | add_10_each | topic_one.publish_each() - chain_two = topic_one.join() | inner_pubsub.map() | topic_two.publish() - chain_three = topic_two.subscribe() | Topic.OUT.publish() + chain_two = topic_one.join() | inner_pubsub.map() | sorted | topic_two.publish() + chain_three = topic_two.subscribe() | sum | Topic.OUT.publish() correlation_id = uuid4() pubsub = PubSub((chain_one, chain_two, chain_three), connection=conn) @@ -319,21 +310,37 @@ def test_invoke_join_then_call_other_pubsub(mocker: MockerFixture): assert conn.listeners == {} # Then invoke pubsub - # We get a single array result as chain_four waits for all publishers to finish - # before operating on all elements published to topic_two as an array - # assert pubsub.invoke([2, 3], {"correlation_id": correlation_id}) == [[13, 14]] - pubsub.invoke([2, 3], {"correlation_id": correlation_id}) - assert [{**m, "started_at": None} for m in conn.peek(correlation_id)] == [ - {"message": [2, 3], "topic_name": "__in__", "started_at": None}, - {"message": 12, "topic_name": "one", "started_at": None}, - {"message": 13, "topic_name": "one", "started_at": None}, - {"message": 25, "topic_name": "two", "started_at": None}, - {"message": 26, "topic_name": "__out__", "started_at": None}, - {"message": [2, 3], "started_at": None, "topic_name": "__in__"}, - {"message": 12, "started_at": None, "topic_name": "one"}, - {"message": 13, "started_at": None, "topic_name": "one"}, - {"message": 25, "started_at": None, "topic_name": "two"}, - {"message": 26, "started_at": None, "topic_name": "__out__"}, + assert clean_log(pubsub.stream([2, 3], {"correlation_id": correlation_id})) == [ + { + "value": [2, 3], + "topic": "__in__", + "correlation_id": str(correlation_id), + "published_at": None, + }, + { + "value": 12, + "topic": "one", + "correlation_id": str(correlation_id), + "published_at": None, + }, + { + "value": 13, + "topic": "one", + "correlation_id": str(correlation_id), + "published_at": None, + }, + { + "value": [13, 14], + "topic": "two", + "correlation_id": str(correlation_id), + "published_at": None, + }, + { + "value": 27, + "topic": "__out__", + "correlation_id": str(correlation_id), + "published_at": None, + }, ] # After invoke returns the listeners were cleaned up @@ -361,9 +368,21 @@ def test_invoke_two_processes_one_in_two_out(mocker: MockerFixture): assert conn.listeners == {} # Then invoke pubsub - # pubsub didn't stop executing after getting the first return value - # the values arrive in the order they are produced - assert pubsub.invoke(2) == [3, 4] + # pubsub stopped executing after publishing to OUT, so only one value is returned + assert clean_log(pubsub.stream(2), correlation_id=False) == [ + { + "value": 2, + "topic": "__in__", + "correlation_id": None, + "published_at": None, + }, + { + "value": 3, + "topic": "__out__", + "correlation_id": None, + "published_at": None, + }, + ] # After invoke returns the listeners were cleaned up assert conn.listeners == {} @@ -389,7 +408,7 @@ def test_invoke_two_processes_no_out(mocker: MockerFixture): # Then invoke pubsub # It finishes executing (once no more messages being published) # but returns nothing, as nothing was published to OUT topic - assert pubsub.invoke(2) == [] + assert pubsub.invoke(2) is None # After invoke returns the listeners were cleaned up assert conn.listeners == {} @@ -414,14 +433,14 @@ def test_invoke_two_processes_no_in(mocker: MockerFixture): # Then invoke pubsub # It returns without any output as there is nothing to run - assert pubsub.invoke(2) == [] + assert pubsub.invoke(2) is None # After invoke returns the listeners were cleaned up assert conn.listeners == {} @pytest.mark.skip("TODO") -def test_invoke_two_processes_simple_cycle(mocker: MockerFixture): +def test_invoke_two_processes_simple_cycle(mocker: MockerFixture) -> None: add_one = mocker.Mock(side_effect=lambda x: x + 1) topic_one = Topic("one") chain_one = Topic.IN.subscribe() | add_one | topic_one.publish()