Rename BaseTask → BaseWorker and reserve "task" for asyncio

Replaces every "task" identifier that referred to the BaseTask
abstraction with "worker". Asyncio task plumbing (asyncio.Task,
BaseTaskManager, TaskManager, create_task, cancel_task, etc.) stays
untouched. Highlights:

- Classes: BaseTask → BaseWorker, PipelineTask → PipelineWorker,
  LLMTask → LLMWorker, LLMContextTask → LLMContextWorker, TaskBus →
  WorkerBus, TaskRegistry → WorkerRegistry, TaskActivationArgs →
  WorkerActivationArgs, TaskReadyData → WorkerReadyData,
  TaskRegistryEntry → WorkerRegistryEntry, TaskObserver →
  WorkerObserver, all Bus*TaskMessage → Bus*WorkerMessage,
  BusAddTaskMessage.task field → worker, BusWorkerRegistryMessage.tasks
  field → workers.
- Methods/decorators: activate_task → activate_worker, deactivate_task
  → deactivate_worker, add_task → add_worker, watch_task →
  watch_worker, @task_ready → @worker_ready, setup_pipeline_task hook
  → setup_pipeline_worker.
- Params/fields: FrameProcessorSetup.pipeline_task and
  FunctionCallParams.pipeline_task → pipeline_worker. Parameter names
  like task_name → worker_name; spawn/run accept worker:.
- Files: pipeline/base_task.py → base_worker.py, pipeline/task.py →
  worker.py (plus a re-export shim at pipeline/task.py),
  task_observer.py → worker_observer.py, task_ready_decorator.py →
  worker_ready_decorator.py, pipecat.tasks → pipecat.workers,
  llm_task.py → llm_worker.py, llm_context_task.py →
  llm_context_worker.py, examples/multi-task → examples/multi-worker.

Back-compat:
- PipelineTask kept as a deprecated subclass of PipelineWorker that
  warns on construction.
- pipecat.pipeline.task re-exports PipelineWorker/PipelineTask/etc. so
  existing user imports keep working.
- FrameProcessor.pipeline_task kept as a deprecated property that
  forwards to pipeline_worker.

Local variables in examples that hold a worker (task = PipelineTask(...))
are renamed to worker = PipelineWorker(...). Asyncio-task locals
(runner_task, etc.) are preserved.
This commit is contained in:
Aleix Conchillo Flaqué
2026-05-20 16:39:45 -07:00
parent b9aed0d673
commit b03247f360
394 changed files with 4602 additions and 4487 deletions

View File

@@ -15,7 +15,7 @@ from pipecat.adapters.schemas.direct_function import DirectFunctionWrapper
from pipecat.clocks.system_clock import SystemClock
from pipecat.frames.frames import EndFrame, Frame, StartFrame
from pipecat.pipeline.pipeline import Pipeline
from pipecat.pipeline.task import PipelineTask, PipelineTaskParams
from pipecat.pipeline.worker import PipelineWorker, PipelineWorkerParams
from pipecat.processors.aggregators.llm_context import LLMContext
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor, FrameProcessorSetup
from pipecat.services.llm_service import (
@@ -64,7 +64,7 @@ class TestFunctionCallParamsAppResources(unittest.TestCase):
tool_call_id="1",
arguments={},
llm=None, # type: ignore[arg-type]
pipeline_task=None, # type: ignore[arg-type]
pipeline_worker=None, # type: ignore[arg-type]
context=LLMContext(),
result_callback=AsyncMock(),
)
@@ -77,7 +77,7 @@ class TestFunctionCallParamsAppResources(unittest.TestCase):
tool_call_id="1",
arguments={},
llm=None, # type: ignore[arg-type]
pipeline_task=None, # type: ignore[arg-type]
pipeline_worker=None, # type: ignore[arg-type]
context=LLMContext(),
result_callback=AsyncMock(),
app_resources=resources,
@@ -91,7 +91,7 @@ class TestFunctionCallParamsAppResources(unittest.TestCase):
tool_call_id="1",
arguments={},
llm=None, # type: ignore[arg-type]
pipeline_task=None, # type: ignore[arg-type]
pipeline_worker=None, # type: ignore[arg-type]
context=LLMContext(),
result_callback=AsyncMock(),
app_resources=resources,
@@ -105,8 +105,8 @@ class TestLLMServiceFunctionCallReadsAppResources(unittest.IsolatedAsyncioTestCa
async def test_function_call_params_receives_app_resources(self):
service = _MockLLMService()
resources = _Resources(user_name="John")
# Stub the pipeline task with just the bit LLMService reads.
service._pipeline_task = SimpleNamespace(app_resources=resources) # type: ignore[assignment]
# Stub the pipeline worker with just the bit LLMService reads.
service._pipeline_worker = SimpleNamespace(app_resources=resources) # type: ignore[assignment]
captured: dict[str, Any] = {}
@@ -137,7 +137,7 @@ class TestLLMServiceFunctionCallReadsAppResources(unittest.IsolatedAsyncioTestCa
async def test_direct_function_params_receives_app_resources(self):
service = _MockLLMService()
resources = _Resources(user_name="John")
service._pipeline_task = SimpleNamespace(app_resources=resources) # type: ignore[assignment]
service._pipeline_worker = SimpleNamespace(app_resources=resources) # type: ignore[assignment]
captured: dict[str, Any] = {}
async def lookup(params: FunctionCallParams):
@@ -174,7 +174,7 @@ class TestLLMServiceFunctionCallReadsAppResources(unittest.IsolatedAsyncioTestCa
setup = FrameProcessorSetup(
clock=SystemClock(),
task_manager=task_manager,
pipeline_task=SimpleNamespace(app_resources=None), # type: ignore[arg-type]
pipeline_worker=SimpleNamespace(app_resources=None), # type: ignore[arg-type]
tool_resources=resources,
)
@@ -186,29 +186,29 @@ class TestLLMServiceFunctionCallReadsAppResources(unittest.IsolatedAsyncioTestCa
class TestPipelineTaskAppResources(unittest.TestCase):
def test_getter_returns_constructor_value(self):
resources = _Resources(user_name="John")
task = PipelineTask(Pipeline([]), app_resources=resources)
self.assertIs(task.app_resources, resources)
worker = PipelineWorker(Pipeline([]), app_resources=resources)
self.assertIs(worker.app_resources, resources)
def test_default_app_resources_is_none(self):
task = PipelineTask(Pipeline([]))
self.assertIsNone(task.app_resources)
worker = PipelineWorker(Pipeline([]))
self.assertIsNone(worker.app_resources)
def test_tool_resources_kwarg_warns_and_aliases_app_resources(self):
resources = _Resources(user_name="John")
with self.assertWarns(DeprecationWarning):
task = PipelineTask(Pipeline([]), tool_resources=resources)
self.assertIs(task.app_resources, resources)
worker = PipelineWorker(Pipeline([]), tool_resources=resources)
self.assertIs(worker.app_resources, resources)
def test_app_resources_takes_precedence_over_tool_resources(self):
new = _Resources(user_name="new")
old = _Resources(user_name="old")
with self.assertWarns(DeprecationWarning):
task = PipelineTask(Pipeline([]), app_resources=new, tool_resources=old)
self.assertIs(task.app_resources, new)
worker = PipelineWorker(Pipeline([]), app_resources=new, tool_resources=old)
self.assertIs(worker.app_resources, new)
class _RecordingProcessor(FrameProcessor):
"""Records the pipeline_task it sees once StartFrame reaches it."""
"""Records the pipeline_worker it sees once StartFrame reaches it."""
def __init__(self):
super().__init__()
@@ -218,10 +218,10 @@ class _RecordingProcessor(FrameProcessor):
async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction)
if isinstance(frame, StartFrame):
# setup() runs before any frame reaches us, so pipeline_task is wired up.
assert self.pipeline_task is not None
self.observed_task = self.pipeline_task
self.observed_app_resources = self.pipeline_task.app_resources
# setup() runs before any frame reaches us, so pipeline_worker is wired up.
assert self.pipeline_worker is not None
self.observed_task = self.pipeline_worker
self.observed_app_resources = self.pipeline_worker.app_resources
await self.push_frame(frame, direction)
@@ -230,7 +230,7 @@ class _LegacyToolResourcesReader(FrameProcessor):
Models a previously-written user FrameProcessor whose ``setup()``
override hasn't been migrated yet. The field is populated by
``PipelineTask`` for backwards compatibility; reading it emits a
``PipelineWorker`` for backwards compatibility; reading it emits a
DeprecationWarning.
"""
@@ -244,7 +244,7 @@ class _LegacyToolResourcesReader(FrameProcessor):
async def process_frame(self, frame: Frame, direction: FrameDirection):
# Forward all frames so the EndFrame reaches the pipeline sink and
# ``task.run()`` can return cleanly.
# ``worker.run()`` can return cleanly.
await super().process_frame(frame, direction)
await self.push_frame(frame, direction)
@@ -254,29 +254,29 @@ class TestFrameProcessorSetupToolResourcesBackwardsCompat(unittest.IsolatedAsync
resources = _Resources(user_name="John")
legacy = _LegacyToolResourcesReader()
pipeline = Pipeline([legacy])
task = PipelineTask(pipeline, app_resources=resources)
worker = PipelineWorker(pipeline, app_resources=resources)
await task.queue_frame(EndFrame())
await worker.queue_frame(EndFrame())
with self.assertWarns(DeprecationWarning):
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
self.assertIs(legacy.captured_tool_resources, resources)
async def test_legacy_processor_receives_value_via_deprecated_tool_resources_kwarg(
self,
):
# If the user is still constructing PipelineTask with the deprecated
# If the user is still constructing PipelineWorker with the deprecated
# ``tool_resources`` kwarg (and hasn't migrated to ``app_resources``),
# legacy processors must still see the value too.
resources = _Resources(user_name="John")
legacy = _LegacyToolResourcesReader()
pipeline = Pipeline([legacy])
with self.assertWarns(DeprecationWarning):
task = PipelineTask(pipeline, tool_resources=resources)
worker = PipelineWorker(pipeline, tool_resources=resources)
await task.queue_frame(EndFrame())
await worker.queue_frame(EndFrame())
with self.assertWarns(DeprecationWarning):
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
self.assertIs(legacy.captured_tool_resources, resources)
@@ -286,18 +286,18 @@ class TestFrameProcessorPipelineTaskAccess(unittest.IsolatedAsyncioTestCase):
resources = _Resources(user_name="John")
recorder = _RecordingProcessor()
pipeline = Pipeline([recorder])
task = PipelineTask(pipeline, app_resources=resources)
worker = PipelineWorker(pipeline, app_resources=resources)
await task.queue_frame(EndFrame())
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.queue_frame(EndFrame())
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
self.assertIs(recorder.observed_task, task)
self.assertIs(recorder.observed_task, worker)
self.assertIs(recorder.observed_app_resources, resources)
def test_pipeline_task_raises_when_not_set_up(self):
recorder = _RecordingProcessor()
with self.assertRaises(Exception):
_ = recorder.pipeline_task
_ = recorder.pipeline_worker
if __name__ == "__main__":

View File

@@ -49,7 +49,7 @@ async def _make_processor(*, buffer_size: int = 0) -> AudioBufferProcessor:
FrameProcessorSetup(
clock=SystemClock(),
task_manager=task_manager,
pipeline_task=SimpleNamespace(app_resources=None), # type: ignore[arg-type]
pipeline_worker=SimpleNamespace(app_resources=None), # type: ignore[arg-type]
)
)

File diff suppressed because it is too large Load Diff

View File

@@ -29,7 +29,7 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
processor = BusBridgeProcessor(
bus=bus,
task_name="test_task",
worker_name="test_task",
)
pipeline = Pipeline([processor])
@@ -66,7 +66,7 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
processor = BusBridgeProcessor(
bus=bus,
task_name="test_task",
worker_name="test_task",
)
pipeline = Pipeline([processor])
@@ -98,7 +98,7 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
processor = BusBridgeProcessor(
bus=bus,
task_name="test_task",
worker_name="test_task",
exclude_frames=(TextFrame,),
)
pipeline = Pipeline([processor])
@@ -125,7 +125,7 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
travel downstream alongside frames from later pipeline processors."""
from pipecat.frames.frames import EndFrame
from pipecat.pipeline.runner import PipelineRunner
from pipecat.pipeline.task import PipelineTask
from pipecat.pipeline.worker import PipelineWorker
from pipecat.processors.frame_processor import FrameProcessor
class AppendFrameProcessor(FrameProcessor):
@@ -140,16 +140,16 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
bus = AsyncQueueBus()
bridge = BusBridgeProcessor(
bus=bus,
task_name="main_task",
worker_name="main_task",
)
pipeline = Pipeline([bridge, AppendFrameProcessor()])
task = PipelineTask(pipeline, cancel_on_idle_timeout=False)
worker = PipelineWorker(pipeline, cancel_on_idle_timeout=False)
received = []
task.set_reached_downstream_filter((TextFrame,))
worker.set_reached_downstream_filter((TextFrame,))
@task.event_handler("on_frame_reached_downstream")
async def on_frame(task, frame):
@worker.event_handler("on_frame_reached_downstream")
async def on_frame(worker, frame):
received.append(frame)
msg = BusFrameMessage(
@@ -162,21 +162,21 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
await asyncio.sleep(0.02)
await bridge.on_bus_message(msg)
await asyncio.sleep(0.02)
await task.queue_frame(EndFrame())
await worker.queue_frame(EndFrame())
runner = PipelineRunner()
await asyncio.gather(runner.run(task), inject_and_end())
await asyncio.gather(runner.run(worker), inject_and_end())
texts = [f.text for f in received if isinstance(f, TextFrame)]
self.assertIn("from_child", texts)
self.assertIn("after_bridge", texts)
async def test_skips_own_frames(self):
"""Bridge ignores bus frames from its own task."""
"""Bridge ignores bus frames from its own worker."""
bus = AsyncQueueBus()
processor = BusBridgeProcessor(
bus=bus,
task_name="test_task",
worker_name="test_task",
)
injected = []
@@ -200,11 +200,11 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
self.assertEqual(len(injected), 0)
async def test_target_task_filtering(self):
"""Bridge with target_task only accepts frames from that task."""
"""Bridge with target_task only accepts frames from that worker."""
bus = AsyncQueueBus()
processor = BusBridgeProcessor(
bus=bus,
task_name="main_task",
worker_name="main_task",
target_task="specific_child",
)
@@ -217,7 +217,7 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
processor.push_frame = capture_push
# Frame from wrong task — should be ignored
# Frame from wrong worker — should be ignored
wrong_msg = BusFrameMessage(
source="other_child",
frame=TextFrame(text="wrong"),
@@ -226,7 +226,7 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
await processor.on_bus_message(wrong_msg)
self.assertEqual(len(injected), 0)
# Frame from correct task — should be injected
# Frame from correct worker — should be injected
right_msg = BusFrameMessage(
source="specific_child",
frame=TextFrame(text="right"),
@@ -237,11 +237,11 @@ class TestBusBridgeProcessor(unittest.IsolatedAsyncioTestCase):
self.assertEqual(injected[0].text, "right")
async def test_targeted_message_for_other_task_skipped(self):
"""Bridge skips bus messages targeted at a different task."""
"""Bridge skips bus messages targeted at a different worker."""
bus = AsyncQueueBus()
processor = BusBridgeProcessor(
bus=bus,
task_name="main_task",
worker_name="main_task",
)
injected = []

View File

@@ -11,7 +11,7 @@ import unittest
from pipecat.bus import (
AsyncQueueBus,
BusCancelMessage,
BusCancelTaskMessage,
BusCancelWorkerMessage,
BusDataMessage,
BusJobCancelMessage,
BusSubscriber,
@@ -245,7 +245,7 @@ class TestBusMessagePriority(unittest.IsolatedAsyncioTestCase):
await bus.send(BusDataMessage(source="data_2"))
await bus.send(BusCancelMessage(source="cancel_1"))
await bus.send(BusDataMessage(source="data_3"))
await bus.send(BusCancelTaskMessage(source="cancel_2", target="task"))
await bus.send(BusCancelWorkerMessage(source="cancel_2", target="task"))
await bus.start()
await asyncio.sleep(0.1)

View File

@@ -13,7 +13,7 @@ from pipecat.bus import (
BusJobRequestMessage,
BusJobResponseMessage,
)
from pipecat.pipeline.base_task import BaseTask
from pipecat.pipeline.base_worker import BaseWorker
from pipecat.pipeline.job_context import (
JobError,
JobEvent,
@@ -21,16 +21,16 @@ from pipecat.pipeline.job_context import (
JobGroupEvent,
JobStatus,
)
from pipecat.registry import TaskRegistry
from pipecat.registry.types import TaskReadyData
from pipecat.registry import WorkerRegistry
from pipecat.registry.types import WorkerReadyData
from pipecat.utils.asyncio.task_manager import TaskManager, TaskManagerParams
class StubTask(BaseTask):
class StubTask(BaseWorker):
pass
class JobWorkerTask(BaseTask):
class JobWorkerTask(BaseWorker):
"""Worker that automatically responds to job requests via the bus."""
def __init__(self, name, *, response=None, status=JobStatus.COMPLETED):
@@ -43,7 +43,7 @@ class JobWorkerTask(BaseTask):
await self.send_job_response(message.job_id, self._auto_response, status=self._auto_status)
class UrgentJobWorkerTask(BaseTask):
class UrgentJobWorkerTask(BaseWorker):
"""Worker that responds urgently to job requests."""
def __init__(self, name, *, response=None, status=JobStatus.COMPLETED):
@@ -58,7 +58,7 @@ class UrgentJobWorkerTask(BaseTask):
)
class UpdatingWorkerTask(BaseTask):
class UpdatingWorkerTask(BaseWorker):
"""Worker that sends updates before responding."""
def __init__(self, name, *, updates, response=None):
@@ -73,7 +73,7 @@ class UpdatingWorkerTask(BaseTask):
await self.send_job_response(message.job_id, self._auto_response)
class StreamingWorkerTask(BaseTask):
class StreamingWorkerTask(BaseWorker):
"""Worker that streams data before responding."""
def __init__(self, name, *, chunks, response=None):
@@ -89,7 +89,7 @@ class StreamingWorkerTask(BaseTask):
await self.send_job_stream_end(message.job_id, self._auto_response)
class SlowWorkerTask(BaseTask):
class SlowWorkerTask(BaseWorker):
"""Worker that blocks during job execution until cancelled."""
def __init__(self, name):
@@ -113,7 +113,7 @@ async def create_test_env():
tm.setup(TaskManagerParams(loop=asyncio.get_running_loop()))
await bus.setup(tm)
await bus.start()
registry = TaskRegistry(runner_name="test-runner")
registry = WorkerRegistry(runner_name="test-runner")
return bus, tm, registry
@@ -122,7 +122,7 @@ async def setup_task(bus, registry, task):
task.attach(registry=registry, bus=bus)
await task.setup(bus.task_manager)
await bus.subscribe(task)
await registry.register(TaskReadyData(task_name=task.name, runner="test-runner"))
await registry.register(WorkerReadyData(worker_name=task.name, runner="test-runner"))
def capture_bus(bus):
@@ -351,7 +351,7 @@ class TestJobGroupContext(unittest.IsolatedAsyncioTestCase):
self.assertEqual(len(events), 2)
self.assertEqual(events[0].type, JobGroupEvent.UPDATE)
self.assertEqual(events[0].task_name, "worker")
self.assertEqual(events[0].worker_name, "worker")
self.assertEqual(events[0].data, {"progress": 25})
self.assertEqual(events[1].data, {"progress": 75})
self.assertEqual(tg.responses, {"worker": {"result": "done"}})
@@ -405,7 +405,7 @@ class TestJobGroupContext(unittest.IsolatedAsyncioTestCase):
self.assertEqual(len(events), 1)
self.assertEqual(events[0].type, JobGroupEvent.UPDATE)
self.assertEqual(events[0].task_name, "w1")
self.assertEqual(events[0].worker_name, "w1")
self.assertEqual(tg.responses, {"w1": {"a": 1}, "w2": {"b": 2}})
async def test_job_group_no_iteration_still_works(self):

View File

@@ -47,7 +47,7 @@ class MockLLMService(LLMService):
)
super().__init__(settings=settings, **kwargs)
# Stub the pipeline task so FunctionCallParams can be constructed.
self._pipeline_task = SimpleNamespace(app_resources=None)
self._pipeline_worker = SimpleNamespace(app_resources=None)
class TestUnparameterizedSubclass(unittest.TestCase):

View File

@@ -9,16 +9,16 @@ import unittest
from unittest.mock import MagicMock
from pipecat.frames.frames import LLMMessagesAppendFrame
from pipecat.pipeline.task import PipelineTask
from pipecat.pipeline.worker import PipelineWorker
from pipecat.processors.frame_processor import FrameDirection
from pipecat.tasks.llm import LLMTask, tool
from pipecat.tasks.llm.llm_task import PipelineFlushFrame
from pipecat.workers.llm import LLMWorker, tool
from pipecat.workers.llm.llm_worker import PipelineFlushFrame
def _create_task():
"""Create a StubLLMTask with mocked parent queue_frame for testing."""
class StubLLMTask(LLMTask):
class StubLLMTask(LLMWorker):
@tool
async def fast_tool(self, params):
"""A quick tool."""
@@ -33,11 +33,11 @@ def _create_task():
llm.register_direct_function = MagicMock()
task = StubLLMTask("test_task", llm=llm, bridged=())
# Capture frames passed to PipelineTask.queue_frame (i.e. super().queue_frame).
# Capture frames passed to PipelineWorker.queue_frame (i.e. super().queue_frame).
# Auto-set _flush_done when PipelineFlushFrame is queued, simulating the flush
# round-trip through the pipeline.
delivered: list[tuple] = []
original_pt_queue_frame = PipelineTask.queue_frame
original_pt_queue_frame = PipelineWorker.queue_frame
async def class_replacement(self, frame, direction=FrameDirection.DOWNSTREAM):
# Only intercept for this specific instance; otherwise fall through.
@@ -48,9 +48,9 @@ def _create_task():
return
await original_pt_queue_frame(self, frame, direction)
PipelineTask.queue_frame = class_replacement
PipelineWorker.queue_frame = class_replacement
task._restore_pt_queue_frame = lambda: setattr(
PipelineTask, "queue_frame", original_pt_queue_frame
PipelineWorker, "queue_frame", original_pt_queue_frame
)
task._delivered_frames = delivered

View File

@@ -12,7 +12,7 @@ from dataclasses import dataclass, field
from datetime import datetime
from pipecat.bus import (
BusAddTaskMessage,
BusAddWorkerMessage,
BusDataMessage,
BusEndMessage,
BusFrameMessage,
@@ -21,7 +21,7 @@ from pipecat.bus import (
)
from pipecat.bus.serializers import JSONMessageSerializer
from pipecat.frames.frames import TextFrame
from pipecat.pipeline.base_task import BaseTask
from pipecat.pipeline.base_worker import BaseWorker
from pipecat.processors.frame_processor import FrameDirection
from pipecat.utils.asyncio.task_manager import TaskManager, TaskManagerParams
@@ -171,8 +171,8 @@ class TestPgmqBus(unittest.IsolatedAsyncioTestCase):
received = []
await self.bus.subscribe(_make_sub(received))
task = BaseTask("test")
msg = BusAddTaskMessage(source="parent", task=task)
worker = BaseWorker("test")
msg = BusAddWorkerMessage(source="parent", worker=worker)
await self.bus.send(msg)
await asyncio.sleep(0.05)
@@ -181,8 +181,8 @@ class TestPgmqBus(unittest.IsolatedAsyncioTestCase):
self.assertEqual(len(self.pgmq.sent), 0)
# But delivered locally
self.assertEqual(len(received), 1)
self.assertIsInstance(received[0], BusAddTaskMessage)
self.assertIs(received[0].task, task)
self.assertIsInstance(received[0], BusAddWorkerMessage)
self.assertIs(received[0].worker, worker)
async def test_round_trip_via_subscriber(self):
"""Messages published are received by subscribers."""
@@ -367,7 +367,7 @@ class TestPgmqBusEdgeCases(unittest.IsolatedAsyncioTestCase):
await bus.stop()
async def test_stop_cleans_up(self):
"""stop() cancels the reader task and drops the queue."""
"""stop() cancels the reader worker and drops the queue."""
pgmq = FakePgmq()
bus = PgmqBus(pgmq=pgmq, channel="cleanup_test", poll_interval_ms=10, max_poll_seconds=1)
tm = TaskManager()

View File

@@ -25,7 +25,7 @@ from pipecat.frames.frames import (
from pipecat.observers.base_observer import BaseObserver, FramePushed
from pipecat.pipeline.parallel_pipeline import ParallelPipeline
from pipecat.pipeline.pipeline import Pipeline
from pipecat.pipeline.task import PipelineParams, PipelineTask, PipelineTaskParams
from pipecat.pipeline.worker import PipelineParams, PipelineWorker, PipelineWorkerParams
from pipecat.processors.filters.frame_filter import FrameFilter
from pipecat.processors.filters.identity_filter import IdentityFilter
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
@@ -130,12 +130,12 @@ class TestParallelPipeline(unittest.IsolatedAsyncioTestCase):
class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
async def test_task_single(self):
pipeline = Pipeline([IdentityFilter()])
task = PipelineTask(pipeline)
worker = PipelineWorker(pipeline)
await task.queue_frame(TextFrame(text="Hello!"))
await task.queue_frames([TextFrame(text="Bye!"), EndFrame()])
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
assert task.has_finished()
await worker.queue_frame(TextFrame(text="Hello!"))
await worker.queue_frames([TextFrame(text="Bye!"), EndFrame()])
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
assert worker.has_finished()
async def test_task_observers(self):
frame_received = False
@@ -149,10 +149,10 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline, observers=[CustomObserver()])
worker = PipelineWorker(pipeline, observers=[CustomObserver()])
await task.queue_frames([TextFrame(text="Hello Downstream!"), EndFrame()])
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.queue_frames([TextFrame(text="Hello Downstream!"), EndFrame()])
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
assert frame_received
async def test_task_add_observer(self):
@@ -183,32 +183,32 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline, observers=[CustomObserver()])
worker = PipelineWorker(pipeline, observers=[CustomObserver()])
# Add a new observer right away, before doing anything else with the task.
# Add a new observer right away, before doing anything else with the worker.
observer1 = CustomAddObserver1()
task.add_observer(observer1)
worker.add_observer(observer1)
async def delayed_add_observer():
observer2 = CustomAddObserver2()
# Wait after the pipeline is started and add another observer.
await asyncio.sleep(0.1)
task.add_observer(observer2)
worker.add_observer(observer2)
# Push a TextFrame and wait for the observer to pick it up.
await task.queue_frame(TextFrame(text="Hello Downstream!"))
await worker.queue_frame(TextFrame(text="Hello Downstream!"))
await asyncio.sleep(0.1)
# Remove both observers.
await task.remove_observer(observer1)
await task.remove_observer(observer2)
await worker.remove_observer(observer1)
await worker.remove_observer(observer2)
# Push another TextFrame. This time the counter should not
# increments since we have removed the observer.
await task.queue_frame(TextFrame(text="Hello Downstream!"))
await worker.queue_frame(TextFrame(text="Hello Downstream!"))
await asyncio.sleep(0.1)
# Finally end the pipeline.
await task.queue_frame(EndFrame())
await worker.queue_frame(EndFrame())
await asyncio.gather(
task.run(PipelineTaskParams(loop=asyncio.get_event_loop())), delayed_add_observer()
worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop())), delayed_add_observer()
)
assert frame_received
@@ -221,20 +221,20 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline)
worker = PipelineWorker(pipeline)
@task.event_handler("on_pipeline_started")
async def on_pipeline_started(task, frame: StartFrame):
@worker.event_handler("on_pipeline_started")
async def on_pipeline_started(worker, frame: StartFrame):
nonlocal start_received
start_received = True
@task.event_handler("on_pipeline_finished")
async def on_pipeline_finished(task, frame: Frame):
@worker.event_handler("on_pipeline_finished")
async def on_pipeline_finished(worker, frame: Frame):
nonlocal end_received
end_received = isinstance(frame, EndFrame)
await task.queue_frame(EndFrame())
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.queue_frame(EndFrame())
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
assert start_received
assert end_received
@@ -244,15 +244,15 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline)
worker = PipelineWorker(pipeline)
@task.event_handler("on_pipeline_finished")
async def on_pipeline_finished(task, frame: Frame):
@worker.event_handler("on_pipeline_finished")
async def on_pipeline_finished(worker, frame: Frame):
nonlocal stop_received
stop_received = isinstance(frame, StopFrame)
await task.queue_frame(StopFrame())
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.queue_frame(StopFrame())
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
assert stop_received
@@ -262,18 +262,18 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline, cancel_on_idle_timeout=False)
task.set_reached_upstream_filter((TextFrame,))
task.set_reached_downstream_filter((TextFrame,))
worker = PipelineWorker(pipeline, cancel_on_idle_timeout=False)
worker.set_reached_upstream_filter((TextFrame,))
worker.set_reached_downstream_filter((TextFrame,))
@task.event_handler("on_frame_reached_upstream")
async def on_frame_reached_upstream(task, frame):
@worker.event_handler("on_frame_reached_upstream")
async def on_frame_reached_upstream(worker, frame):
nonlocal upstream_received
if isinstance(frame, TextFrame) and frame.text == "Hello Upstream!":
upstream_received = True
@task.event_handler("on_frame_reached_downstream")
async def on_frame_reached_downstream(task, frame):
@worker.event_handler("on_frame_reached_downstream")
async def on_frame_reached_downstream(worker, frame):
nonlocal downstream_received
if isinstance(frame, TextFrame) and frame.text == "Hello Downstream!":
downstream_received = True
@@ -281,11 +281,11 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
TextFrame(text="Hello Upstream!"), FrameDirection.UPSTREAM
)
await task.queue_frame(TextFrame(text="Hello Downstream!"))
await worker.queue_frame(TextFrame(text="Hello Downstream!"))
try:
await asyncio.wait_for(
task.run(PipelineTaskParams(loop=asyncio.get_event_loop())),
worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop())),
timeout=1.0,
)
except TimeoutError:
@@ -298,22 +298,22 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
upstream_received = False
pipeline = Pipeline([IdentityFilter()])
task = PipelineTask(pipeline, cancel_on_idle_timeout=False)
task.set_reached_upstream_filter((TextFrame,))
worker = PipelineWorker(pipeline, cancel_on_idle_timeout=False)
worker.set_reached_upstream_filter((TextFrame,))
@task.event_handler("on_frame_reached_upstream")
async def on_frame_reached_upstream(task, frame):
@worker.event_handler("on_frame_reached_upstream")
async def on_frame_reached_upstream(worker, frame):
nonlocal upstream_received
if isinstance(frame, TextFrame) and frame.text == "Hello Upstream!":
upstream_received = True
@task.event_handler("on_pipeline_started")
async def on_pipeline_started(task, frame):
await task.queue_frame(TextFrame(text="Hello Upstream!"), FrameDirection.UPSTREAM)
@worker.event_handler("on_pipeline_started")
async def on_pipeline_started(worker, frame):
await worker.queue_frame(TextFrame(text="Hello Upstream!"), FrameDirection.UPSTREAM)
try:
await asyncio.wait_for(
task.run(PipelineTaskParams(loop=asyncio.get_event_loop())),
worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop())),
timeout=1.0,
)
except TimeoutError:
@@ -325,24 +325,24 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
upstream_texts = []
pipeline = Pipeline([IdentityFilter()])
task = PipelineTask(pipeline, cancel_on_idle_timeout=False)
task.set_reached_upstream_filter((TextFrame,))
worker = PipelineWorker(pipeline, cancel_on_idle_timeout=False)
worker.set_reached_upstream_filter((TextFrame,))
@task.event_handler("on_frame_reached_upstream")
async def on_frame_reached_upstream(task, frame):
@worker.event_handler("on_frame_reached_upstream")
async def on_frame_reached_upstream(worker, frame):
if isinstance(frame, TextFrame):
upstream_texts.append(frame.text)
@task.event_handler("on_pipeline_started")
async def on_pipeline_started(task, frame):
await task.queue_frames(
@worker.event_handler("on_pipeline_started")
async def on_pipeline_started(worker, frame):
await worker.queue_frames(
[TextFrame(text="First"), TextFrame(text="Second")],
FrameDirection.UPSTREAM,
)
try:
await asyncio.wait_for(
task.run(PipelineTaskParams(loop=asyncio.get_event_loop())),
worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop())),
timeout=1.0,
)
except TimeoutError:
@@ -363,7 +363,7 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
heartbeats_observer = HeartbeatsObserver(
target=identity, heartbeat_callback=heartbeat_received
)
task = PipelineTask(
worker = PipelineWorker(
pipeline,
params=PipelineParams(
enable_heartbeats=True,
@@ -375,10 +375,10 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
expected_heartbeats = 1.0 / 0.2
await task.queue_frame(TextFrame(text="Hello!"))
await worker.queue_frame(TextFrame(text="Hello!"))
try:
await asyncio.wait_for(
task.run(PipelineTaskParams(loop=asyncio.get_event_loop())),
worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop())),
timeout=1.0,
)
except TimeoutError:
@@ -401,7 +401,7 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
try:
pipeline = Pipeline([HeartbeatBlocker()])
task = PipelineTask(
worker = PipelineWorker(
pipeline,
params=PipelineParams(
enable_heartbeats=True,
@@ -413,7 +413,7 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
try:
await asyncio.wait_for(
task.run(PipelineTaskParams(loop=asyncio.get_event_loop())),
worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop())),
timeout=0.6,
)
except TimeoutError:
@@ -427,17 +427,17 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
async def test_idle_task(self):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline, idle_timeout_secs=0.2)
worker = PipelineWorker(pipeline, idle_timeout_secs=0.2)
# This shouldn't freeze, so nothing to check really.
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
async def test_no_idle_task(self):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline, idle_timeout_secs=0.2, cancel_on_idle_timeout=False)
worker = PipelineWorker(pipeline, idle_timeout_secs=0.2, cancel_on_idle_timeout=False)
try:
await asyncio.wait_for(
task.run(PipelineTaskParams(loop=asyncio.get_event_loop())),
worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop())),
timeout=0.3,
)
except TimeoutError:
@@ -448,7 +448,7 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
async def test_idle_task_heartbeats(self):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(
worker = PipelineWorker(
pipeline,
params=PipelineParams(
enable_heartbeats=True,
@@ -456,39 +456,39 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
),
idle_timeout_secs=0.3,
)
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
async def test_idle_task_event_handler_no_frames(self):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline, idle_timeout_secs=0.2, cancel_on_idle_timeout=False)
worker = PipelineWorker(pipeline, idle_timeout_secs=0.2, cancel_on_idle_timeout=False)
idle_timeout = False
@task.event_handler("on_idle_timeout")
async def on_idle_timeout(task: PipelineTask):
@worker.event_handler("on_idle_timeout")
async def on_idle_timeout(worker: PipelineWorker):
nonlocal idle_timeout
idle_timeout = True
await task.cancel()
await worker.cancel()
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
assert idle_timeout
async def test_idle_task_event_handler_quiet_user(self):
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(pipeline, idle_timeout_secs=0.2, cancel_on_idle_timeout=False)
worker = PipelineWorker(pipeline, idle_timeout_secs=0.2, cancel_on_idle_timeout=False)
idle_timeout = 0
@task.event_handler("on_idle_timeout")
async def on_idle_timeout(task: PipelineTask):
@worker.event_handler("on_idle_timeout")
async def on_idle_timeout(worker: PipelineWorker):
nonlocal idle_timeout
idle_timeout += 1
# Stay a bit longer here while user audio frames are still being
# pushed. We do this to make sure this function is only called once.
await asyncio.sleep(0.1)
await task.queue_frame(EndFrame())
await worker.queue_frame(EndFrame())
async def send_audio():
# We send audio during and after the 0.2 seconds of idle
@@ -496,13 +496,13 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
# simulating the pipeline finishing (e.g. goodbye message from bot
# flushing).
for i in range(30):
await task.queue_frame(
await worker.queue_frame(
InputAudioRawFrame(audio=b"\x00", sample_rate=16000, num_channels=1)
)
await asyncio.sleep(0.01)
await asyncio.gather(
send_audio(), task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
send_audio(), worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
)
assert idle_timeout == 1
@@ -513,7 +513,7 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
# Use the identify filter so the frames just reach the end of the pipeline.
identity = IdentityFilter()
pipeline = Pipeline([identity])
task = PipelineTask(
worker = PipelineWorker(
pipeline,
idle_timeout_secs=idle_timeout_secs,
idle_timeout_frames=(TextFrame,),
@@ -523,20 +523,20 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
"""Sending multiple text frames.
The total amount of elapsed time in this function should be greater
than the task idle timeout. If an idle timeout event is triggered it
than the worker idle timeout. If an idle timeout event is triggered it
means we haven't detected that the TextFrames have been pushed.
"""
await asyncio.sleep(sleep_time_secs)
await task.queue_frame(TextFrame("Hello Pipecat!"))
await worker.queue_frame(TextFrame("Hello Pipecat!"))
await asyncio.sleep(sleep_time_secs)
await task.queue_frame(TextFrame("Hello Pipecat!"))
await worker.queue_frame(TextFrame("Hello Pipecat!"))
await asyncio.sleep(sleep_time_secs)
await task.queue_frame(TextFrame("Hello Pipecat!"))
await worker.queue_frame(TextFrame("Hello Pipecat!"))
start_time = time.time()
tasks = [
asyncio.create_task(task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))),
asyncio.create_task(worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))),
asyncio.create_task(delayed_frames()),
]
@@ -558,7 +558,7 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
# reach the end of the pipeline).
filter = FrameFilter(types=())
pipeline = Pipeline([filter])
task = PipelineTask(
worker = PipelineWorker(
pipeline,
idle_timeout_secs=idle_timeout_secs,
idle_timeout_frames=(TextFrame,),
@@ -570,18 +570,18 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
"""Sending multiple text frames.
The total amount of elapsed time in this function should be greater
than the task idle timeout. If an idle timeout event is triggered it
than the worker idle timeout. If an idle timeout event is triggered it
means we haven't detected that the TextFrames have been pushed.
"""
await asyncio.sleep(sleep_time_secs)
await task.queue_frame(TextFrame("Hello Pipecat!"))
await worker.queue_frame(TextFrame("Hello Pipecat!"))
await asyncio.sleep(sleep_time_secs)
await task.queue_frame(TextFrame("Hello Pipecat!"))
await worker.queue_frame(TextFrame("Hello Pipecat!"))
await asyncio.sleep(sleep_time_secs)
await task.queue_frame(TextFrame("Hello Pipecat!"))
await worker.queue_frame(TextFrame("Hello Pipecat!"))
tasks = [
asyncio.create_task(task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))),
asyncio.create_task(worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))),
asyncio.create_task(delayed_frames()),
]
@@ -606,21 +606,21 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
await self.push_frame(frame, direction)
pipeline = Pipeline([CancelFilter()])
task = PipelineTask(pipeline, cancel_timeout_secs=0.2)
worker = PipelineWorker(pipeline, cancel_timeout_secs=0.2)
cancelled = False
@task.event_handler("on_pipeline_started")
async def on_pipeline_started(task: PipelineTask, frame: StartFrame):
await task.cancel()
@worker.event_handler("on_pipeline_started")
async def on_pipeline_started(worker: PipelineWorker, frame: StartFrame):
await worker.cancel()
@task.event_handler("on_pipeline_finished")
async def on_pipeline_finished(task: PipelineTask, frame: Frame):
@worker.event_handler("on_pipeline_finished")
async def on_pipeline_finished(worker: PipelineWorker, frame: Frame):
nonlocal cancelled
cancelled = isinstance(frame, CancelFrame)
try:
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
except asyncio.CancelledError:
assert cancelled
@@ -640,18 +640,18 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
error_received = False
pipeline = Pipeline([ErrorProcessor()])
task = PipelineTask(pipeline)
worker = PipelineWorker(pipeline)
@task.event_handler("on_pipeline_error")
async def on_pipeline_error(task: PipelineTask, frame: ErrorFrame):
@worker.event_handler("on_pipeline_error")
async def on_pipeline_error(worker: PipelineWorker, frame: ErrorFrame):
nonlocal error_received
error_received = True
await task.cancel()
await worker.cancel()
await task.queue_frame(TextFrame(text="Hello from Pipecat!"))
await worker.queue_frame(TextFrame(text="Hello from Pipecat!"))
try:
await task.run(PipelineTaskParams(loop=asyncio.get_event_loop()))
await worker.run(PipelineWorkerParams(loop=asyncio.get_event_loop()))
except asyncio.CancelledError:
assert error_received

View File

@@ -9,7 +9,7 @@ import itertools
import unittest
from pipecat.bus import (
BusAddTaskMessage,
BusAddWorkerMessage,
BusDataMessage,
BusEndMessage,
BusFrameMessage,
@@ -18,7 +18,7 @@ from pipecat.bus import (
)
from pipecat.bus.serializers import JSONMessageSerializer
from pipecat.frames.frames import TextFrame
from pipecat.pipeline.base_task import BaseTask
from pipecat.pipeline.base_worker import BaseWorker
from pipecat.processors.frame_processor import FrameDirection
from pipecat.utils.asyncio.task_manager import TaskManager, TaskManagerParams
@@ -133,8 +133,8 @@ class TestRedisBus(unittest.IsolatedAsyncioTestCase):
await self.bus.subscribe(_make_sub(received))
await self.bus.start()
task = BaseTask("test")
msg = BusAddTaskMessage(source="parent", task=task)
worker = BaseWorker("test")
msg = BusAddWorkerMessage(source="parent", worker=worker)
await self.bus.send(msg)
await asyncio.sleep(0.05)
@@ -144,8 +144,8 @@ class TestRedisBus(unittest.IsolatedAsyncioTestCase):
self.assertEqual(len(self.redis._published), 0)
# But delivered locally
self.assertEqual(len(received), 1)
self.assertIsInstance(received[0], BusAddTaskMessage)
self.assertIs(received[0].task, task)
self.assertIsInstance(received[0], BusAddWorkerMessage)
self.assertIs(received[0].worker, worker)
async def test_round_trip_via_subscriber(self):
"""Messages published are received by subscribers."""
@@ -156,7 +156,7 @@ class TestRedisBus(unittest.IsolatedAsyncioTestCase):
msg = BusEndMessage(source="task_a", reason="done")
await self.bus.send(msg)
# Give the reader task time to process
# Give the reader worker time to process
await asyncio.sleep(0.1)
await self.bus.stop()
@@ -239,7 +239,7 @@ class TestRedisBus(unittest.IsolatedAsyncioTestCase):
self.assertEqual(self.redis._published[0][0], "custom:channel")
async def test_stop_cleans_up(self):
"""stop() cancels the reader task and unsubscribes from Redis."""
"""stop() cancels the reader worker and unsubscribes from Redis."""
await self.bus.start()
self.assertIsNotNone(self.bus._pubsub)
self.assertIsNotNone(self.bus._reader_task)

View File

@@ -6,45 +6,45 @@
import unittest
from pipecat.registry import TaskRegistry
from pipecat.registry.types import TaskReadyData
from pipecat.registry import WorkerRegistry
from pipecat.registry.types import WorkerReadyData
class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
def setUp(self):
self.registry = TaskRegistry(runner_name="runner_a")
self.registry = WorkerRegistry(runner_name="runner_a")
async def test_register_local_task(self):
"""Local task is registered and appears in local_tasks."""
data = TaskReadyData(task_name="greeter", runner="runner_a")
"""Local task is registered and appears in local_workers."""
data = WorkerReadyData(worker_name="greeter", runner="runner_a")
result = await self.registry.register(data)
self.assertTrue(result)
self.assertIn("greeter", self.registry.local_tasks)
self.assertNotIn("greeter", self.registry.remote_tasks)
self.assertIn("greeter", self.registry.local_workers)
self.assertNotIn("greeter", self.registry.remote_workers)
async def test_register_remote_task(self):
"""Remote task is registered and appears in remote_tasks."""
data = TaskReadyData(task_name="support", runner="runner_b")
"""Remote task is registered and appears in remote_workers."""
data = WorkerReadyData(worker_name="support", runner="runner_b")
result = await self.registry.register(data)
self.assertTrue(result)
self.assertIn("support", self.registry.remote_tasks)
self.assertNotIn("support", self.registry.local_tasks)
self.assertIn("support", self.registry.remote_workers)
self.assertNotIn("support", self.registry.local_workers)
async def test_duplicate_registration_returns_false(self):
"""Registering the same task twice returns False."""
data = TaskReadyData(task_name="greeter", runner="runner_a")
data = WorkerReadyData(worker_name="greeter", runner="runner_a")
first = await self.registry.register(data)
second = await self.registry.register(data)
self.assertTrue(first)
self.assertFalse(second)
self.assertEqual(self.registry.local_tasks.count("greeter"), 1)
self.assertEqual(self.registry.local_workers.count("greeter"), 1)
async def test_get_local_task(self):
"""get() returns data for a local task."""
data = TaskReadyData(task_name="greeter", runner="runner_a")
data = WorkerReadyData(worker_name="greeter", runner="runner_a")
await self.registry.register(data)
result = self.registry.get("greeter")
@@ -52,7 +52,7 @@ class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
async def test_get_remote_task(self):
"""get() returns data for a remote task."""
data = TaskReadyData(task_name="support", runner="runner_b")
data = WorkerReadyData(worker_name="support", runner="runner_b")
await self.registry.register(data)
result = self.registry.get("support")
@@ -64,7 +64,7 @@ class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
async def test_contains(self):
"""__contains__ works for registered and unregistered tasks."""
data = TaskReadyData(task_name="greeter", runner="runner_a")
data = WorkerReadyData(worker_name="greeter", runner="runner_a")
await self.registry.register(data)
self.assertIn("greeter", self.registry)
@@ -79,7 +79,7 @@ class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
await self.registry.watch("greeter", handler)
data = TaskReadyData(task_name="greeter", runner="runner_a")
data = WorkerReadyData(worker_name="greeter", runner="runner_a")
await self.registry.register(data)
self.assertEqual(len(received), 1)
@@ -94,7 +94,7 @@ class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
await self.registry.watch("greeter", handler)
data = TaskReadyData(task_name="support", runner="runner_a")
data = WorkerReadyData(worker_name="support", runner="runner_a")
await self.registry.register(data)
self.assertEqual(len(received), 0)
@@ -108,7 +108,7 @@ class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
await self.registry.watch("greeter", handler)
data = TaskReadyData(task_name="greeter", runner="runner_a")
data = WorkerReadyData(worker_name="greeter", runner="runner_a")
await self.registry.register(data)
await self.registry.register(data)
@@ -128,7 +128,7 @@ class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
await self.registry.watch("greeter", handler_a)
await self.registry.watch("greeter", handler_b)
data = TaskReadyData(task_name="greeter", runner="runner_a")
data = WorkerReadyData(worker_name="greeter", runner="runner_a")
await self.registry.register(data)
self.assertEqual(len(received_a), 1)
@@ -136,7 +136,7 @@ class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
async def test_watch_fires_immediately_if_already_registered(self):
"""Watch handler fires immediately when the task is already registered."""
data = TaskReadyData(task_name="greeter", runner="runner_a")
data = WorkerReadyData(worker_name="greeter", runner="runner_a")
await self.registry.register(data)
received = []
@@ -155,12 +155,12 @@ class TestTaskRegistry(unittest.IsolatedAsyncioTestCase):
async def test_multiple_remote_runners(self):
"""Tasks from multiple remote runners are tracked separately."""
data_b = TaskReadyData(task_name="task_b", runner="runner_b")
data_c = TaskReadyData(task_name="task_c", runner="runner_c")
data_b = WorkerReadyData(worker_name="task_b", runner="runner_b")
data_c = WorkerReadyData(worker_name="task_c", runner="runner_c")
await self.registry.register(data_b)
await self.registry.register(data_c)
remote = self.registry.remote_tasks
remote = self.registry.remote_workers
self.assertIn("task_b", remote)
self.assertIn("task_c", remote)
self.assertEqual(len(remote), 2)

View File

@@ -8,44 +8,44 @@ import asyncio
import unittest
from pipecat.bus import (
BusAddTaskMessage,
BusAddWorkerMessage,
BusCancelMessage,
BusCancelTaskMessage,
BusCancelWorkerMessage,
BusEndMessage,
BusEndTaskMessage,
BusEndWorkerMessage,
)
from pipecat.pipeline.base_task import BaseTask
from pipecat.pipeline.base_worker import BaseWorker
from pipecat.pipeline.runner import PipelineRunner
class StubTask(BaseTask):
"""BaseTask subclass that stops on end/cancel so the runner can exit."""
class StubTask(BaseWorker):
"""BaseWorker subclass that stops on end/cancel so the runner can exit."""
async def _handle_task_end(self, message):
await super()._handle_task_end(message)
async def _handle_worker_end(self, message):
await super()._handle_worker_end(message)
self._finished_event.set()
async def _handle_task_cancel(self, message):
await super()._handle_task_cancel(message)
async def _handle_worker_cancel(self, message):
await super()._handle_worker_cancel(message)
self._finished_event.set()
class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
async def test_spawn_registers_task(self):
"""spawn() registers the task by name (duplicate is silently skipped)."""
"""add_worker() registers the task by name (duplicate is silently skipped)."""
runner = PipelineRunner(handle_sigint=False)
task = StubTask("task_a")
await runner.spawn(task)
await runner.add_worker(task)
# Duplicate is silently skipped (logs error)
await runner.spawn(StubTask("task_a"))
await runner.add_worker(StubTask("task_a"))
async def test_run_starts_bus_and_tasks(self):
"""run() starts bus, starts all tasks, fires on_ready."""
runner = PipelineRunner(handle_sigint=False)
task = StubTask("task_a")
await runner.spawn(task)
await runner.add_worker(task)
runner_started = asyncio.Event()
@@ -63,7 +63,7 @@ class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
"""end() is idempotent — subsequent calls are no-ops."""
runner = PipelineRunner(handle_sigint=False)
task = StubTask("task_a")
await runner.spawn(task)
await runner.add_worker(task)
@runner.event_handler("on_ready")
async def on_ready(runner):
@@ -77,7 +77,7 @@ class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
"""cancel() is idempotent — subsequent calls are no-ops."""
runner = PipelineRunner(handle_sigint=False)
task = StubTask("task_a")
await runner.spawn(task)
await runner.add_worker(task)
@runner.event_handler("on_ready")
async def on_ready(runner):
@@ -90,14 +90,14 @@ class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
pass
async def test_end_sends_end_task_message_to_root_tasks_only(self):
"""end() sends BusEndTaskMessage only to root tasks (no parent)."""
"""end() sends BusEndWorkerMessage only to root tasks (no parent)."""
runner = PipelineRunner(handle_sigint=False)
root = StubTask("root")
child = StubTask("child")
# Manually mark child as having root as parent
child._parent = root.name
await runner.spawn(root)
await runner.spawn(child)
await runner.add_worker(root)
await runner.add_worker(child)
sent = []
bus = runner.bus
@@ -112,19 +112,19 @@ class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
# Call end() directly — no need to run the full pipeline lifecycle
await runner.end()
end_msgs = [m for m in sent if isinstance(m, BusEndTaskMessage)]
end_msgs = [m for m in sent if isinstance(m, BusEndWorkerMessage)]
targets = {m.target for m in end_msgs}
self.assertIn("root", targets)
self.assertNotIn("child", targets)
async def test_cancel_sends_cancel_task_message_to_root_tasks_only(self):
"""cancel() sends BusCancelTaskMessage only to root tasks (no parent)."""
"""cancel() sends BusCancelWorkerMessage only to root tasks (no parent)."""
runner = PipelineRunner(handle_sigint=False)
root = StubTask("root")
child = StubTask("child")
child._parent = root.name
await runner.spawn(root)
await runner.spawn(child)
await runner.add_worker(root)
await runner.add_worker(child)
sent = []
bus = runner.bus
@@ -139,7 +139,7 @@ class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
# Call cancel() directly — no need to run the full pipeline lifecycle
await runner.cancel()
cancel_msgs = [m for m in sent if isinstance(m, BusCancelTaskMessage)]
cancel_msgs = [m for m in sent if isinstance(m, BusCancelWorkerMessage)]
targets = {m.target for m in cancel_msgs}
self.assertIn("root", targets)
self.assertNotIn("child", targets)
@@ -148,7 +148,7 @@ class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
"""BusEndMessage on bus triggers runner.end()."""
runner = PipelineRunner(handle_sigint=False)
task = StubTask("task_a")
await runner.spawn(task)
await runner.add_worker(task)
bus = runner.bus
@@ -164,7 +164,7 @@ class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
"""BusCancelMessage on bus triggers runner.cancel()."""
runner = PipelineRunner(handle_sigint=False)
task = StubTask("task_a")
await runner.spawn(task)
await runner.add_worker(task)
bus = runner.bus
@@ -177,25 +177,25 @@ class TestPipelineRunner(unittest.IsolatedAsyncioTestCase):
except asyncio.CancelledError:
pass
async def test_bus_add_task_message_triggers_spawn(self):
"""BusAddTaskMessage on bus triggers spawn()."""
async def test_bus_add_task_message_triggers_add(self):
"""BusAddWorkerMessage on bus triggers add_worker()."""
runner = PipelineRunner(handle_sigint=False)
task_a = StubTask("task_a")
await runner.spawn(task_a)
await runner.add_worker(task_a)
task_b = StubTask("task_b")
bus = runner.bus
@runner.event_handler("on_ready")
async def on_ready(runner):
await bus.send(BusAddTaskMessage(source="task_a", task=task_b))
await bus.send(BusAddWorkerMessage(source="task_a", task=task_b))
await asyncio.sleep(0.1)
await runner.end()
await asyncio.wait_for(runner.run(), timeout=5.0)
# Verify task_b was added (duplicate is silently skipped)
await runner.spawn(StubTask("task_b"))
await runner.add_worker(StubTask("task_b"))
if __name__ == "__main__":

View File

@@ -9,7 +9,7 @@ import unittest
from pydantic import BaseModel
from pipecat.bus.messages import (
BusActivateTaskMessage,
BusActivateWorkerMessage,
BusCancelMessage,
BusDataMessage,
BusEndMessage,
@@ -60,8 +60,8 @@ class TestJSONMessageSerializer(unittest.TestCase):
self.assertIsNone(restored.target)
def test_round_trip_activate_message(self):
"""BusActivateTaskMessage with args round-trips."""
msg = BusActivateTaskMessage(
"""BusActivateWorkerMessage with args round-trips."""
msg = BusActivateWorkerMessage(
source="parent",
target="child",
args={"messages": [{"role": "user", "content": "hello"}]},
@@ -69,7 +69,7 @@ class TestJSONMessageSerializer(unittest.TestCase):
data = self.serializer.serialize(msg)
restored = self.serializer.deserialize(data)
self.assertIsInstance(restored, BusActivateTaskMessage)
self.assertIsInstance(restored, BusActivateWorkerMessage)
self.assertEqual(restored.source, "parent")
self.assertEqual(restored.target, "child")
self.assertEqual(restored.args["messages"][0]["content"], "hello")

View File

@@ -181,7 +181,7 @@ class TestStartupTimingObserver(unittest.IsolatedAsyncioTestCase):
report = reports[0]
# No internal processors (PipelineSource, PipelineSink, Pipeline) in the report.
internal_names = ("Pipeline#", "PipelineTask#")
internal_names = ("Pipeline#", "PipelineWorker#")
for t in report.processor_timings:
for prefix in internal_names:
self.assertNotIn(

View File

@@ -133,7 +133,7 @@ class TestTurnTraceObserver(unittest.IsolatedAsyncioTestCase):
observers=self._all_observers(trace_observer),
)
# End conversation to flush the conversation span (normally done by PipelineTask._cleanup)
# End conversation to flush the conversation span (normally done by PipelineWorker._cleanup)
trace_observer.end_conversation_tracing()
conv_spans = self._get_spans_by_name("conversation")
@@ -320,7 +320,7 @@ class TestTurnTraceObserver(unittest.IsolatedAsyncioTestCase):
observers=self._all_observers(trace_observer),
)
# Manually end conversation tracing (as PipelineTask._cleanup does)
# Manually end conversation tracing (as PipelineWorker._cleanup does)
trace_observer.end_conversation_tracing()
self.assertIsNone(tracing_ctx.get_conversation_context())

View File

@@ -12,6 +12,7 @@ from collections.abc import Awaitable, Callable
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import Any, Optional
from unittest.mock import ANY, AsyncMock, MagicMock, Mock, call, patch
@@ -247,7 +248,11 @@ class TestVonageVideoConnectorTransport:
clock: SystemClock = SystemClock() # type: ignore[no-untyped-call]
task_manager = TaskManager()
task_manager.setup(TaskManagerParams(loop=asyncio.get_running_loop()))
self._frame_processor_setup = FrameProcessorSetup(clock=clock, task_manager=task_manager)
self._frame_processor_setup = FrameProcessorSetup(
clock=clock,
task_manager=task_manager,
pipeline_worker=SimpleNamespace(app_resources=None), # type: ignore[arg-type]
)
return self._frame_processor_setup
async def _wait_for_condition(

View File

@@ -10,12 +10,12 @@ from unittest.mock import MagicMock
from pipecat.bus import (
AsyncQueueBus,
BusAddTaskMessage,
BusAddWorkerMessage,
BusDataMessage,
)
from pipecat.bus.serializers import JSONMessageSerializer
from pipecat.pipeline.base_task import BaseTask
from pipecat.registry import TaskRegistry
from pipecat.pipeline.base_worker import BaseWorker
from pipecat.registry import WorkerRegistry
from pipecat.utils.asyncio.task_manager import TaskManager, TaskManagerParams
@@ -85,17 +85,17 @@ class FakeStarletteWebSocket:
class TestWebSocketProxyClientTask(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
self.bus, self.tm = await create_test_bus()
self.registry = TaskRegistry(runner_name="test-runner")
self.registry = WorkerRegistry(runner_name="test-runner")
self.serializer = JSONMessageSerializer()
async def _create_client(self, fake_ws):
from pipecat.tasks.proxy.websocket.client import WebSocketProxyClientTask
from pipecat.workers.proxy.websocket.client import WebSocketProxyClientTask
task = WebSocketProxyClientTask(
"proxy",
url="ws://fake",
remote_task_name="worker",
local_task_name="voice",
remote_worker_name="worker",
local_worker_name="voice",
serializer=self.serializer,
)
task.attach(registry=self.registry, bus=self.bus)
@@ -141,9 +141,9 @@ class TestWebSocketProxyClientTask(unittest.IsolatedAsyncioTestCase):
fake_ws = FakeWebSocket()
task = await self._create_client(fake_ws)
stub = MagicMock(spec=BaseTask)
stub = MagicMock(spec=BaseWorker)
stub.name = "child"
msg = BusAddTaskMessage(source="parent", target="worker", task=stub)
msg = BusAddWorkerMessage(source="parent", target="worker", worker=stub)
await task.on_bus_message(msg)
self.assertEqual(len(fake_ws._sent), 0)
@@ -208,17 +208,17 @@ class TestWebSocketProxyClientTask(unittest.IsolatedAsyncioTestCase):
class TestWebSocketProxyServerTask(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
self.bus, self.tm = await create_test_bus()
self.registry = TaskRegistry(runner_name="test-runner")
self.registry = WorkerRegistry(runner_name="test-runner")
self.serializer = JSONMessageSerializer()
async def _create_server(self, fake_ws):
from pipecat.tasks.proxy.websocket.server import WebSocketProxyServerTask
from pipecat.workers.proxy.websocket.server import WebSocketProxyServerTask
task = WebSocketProxyServerTask(
"gateway",
websocket=fake_ws,
task_name="worker",
remote_task_name="voice",
worker_name="worker",
remote_worker_name="voice",
serializer=self.serializer,
)
task.attach(registry=self.registry, bus=self.bus)