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:
@@ -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__":
|
||||
|
||||
@@ -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
@@ -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 = []
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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__":
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user