Make PipelineTask.add_observer() synchronous. This allows callers to call it before run()ning the PipelineTask first. Without this change, if they tried to do that, they would get an error because the TaskManager's event loop hadn't been set yet.
This commit is contained in:
@@ -310,8 +310,8 @@ class PipelineTask(BaseTask):
|
|||||||
"""Return the turn trace observer if enabled."""
|
"""Return the turn trace observer if enabled."""
|
||||||
return self._turn_trace_observer
|
return self._turn_trace_observer
|
||||||
|
|
||||||
async def add_observer(self, observer: BaseObserver):
|
def add_observer(self, observer: BaseObserver):
|
||||||
await self._observer.add_observer(observer)
|
self._observer.add_observer(observer)
|
||||||
|
|
||||||
async def remove_observer(self, observer: BaseObserver):
|
async def remove_observer(self, observer: BaseObserver):
|
||||||
await self._observer.remove_observer(observer)
|
await self._observer.remove_observer(observer)
|
||||||
|
|||||||
@@ -49,21 +49,31 @@ class TaskObserver(BaseObserver):
|
|||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._observers = observers or []
|
self._observers = observers or []
|
||||||
self._task_manager = task_manager
|
self._task_manager = task_manager
|
||||||
self._proxies: Dict[BaseObserver, Proxy] = {}
|
self._proxies: Optional[Dict[BaseObserver, Proxy]] = (
|
||||||
|
None # Becomes a dict after start() is called
|
||||||
|
)
|
||||||
|
|
||||||
async def add_observer(self, observer: BaseObserver):
|
def add_observer(self, observer: BaseObserver):
|
||||||
proxy = self._create_proxy(observer)
|
# Add the observer to the list.
|
||||||
self._proxies[observer] = proxy
|
|
||||||
self._observers.append(observer)
|
self._observers.append(observer)
|
||||||
|
|
||||||
|
# If we already started, create a new proxy for the observer.
|
||||||
|
# Otherwise, it will be created in start().
|
||||||
|
if self._started():
|
||||||
|
proxy = self._create_proxy(observer)
|
||||||
|
self._proxies[observer] = proxy
|
||||||
|
|
||||||
async def remove_observer(self, observer: BaseObserver):
|
async def remove_observer(self, observer: BaseObserver):
|
||||||
|
# If the observer has a proxy, remove it.
|
||||||
if observer in self._proxies:
|
if observer in self._proxies:
|
||||||
proxy = self._proxies[observer]
|
proxy = self._proxies[observer]
|
||||||
# Remove the proxy so it doesn't get called anymore.
|
# Remove the proxy so it doesn't get called anymore.
|
||||||
del self._proxies[observer]
|
del self._proxies[observer]
|
||||||
# Cancel the proxy task right away.
|
# Cancel the proxy task right away.
|
||||||
await self._task_manager.cancel_task(proxy.task)
|
await self._task_manager.cancel_task(proxy.task)
|
||||||
# Remove the observer.
|
|
||||||
|
# Remove the observer from the list.
|
||||||
|
if observer in self._observers:
|
||||||
self._observers.remove(observer)
|
self._observers.remove(observer)
|
||||||
|
|
||||||
async def start(self):
|
async def start(self):
|
||||||
@@ -79,6 +89,9 @@ class TaskObserver(BaseObserver):
|
|||||||
for proxy in self._proxies.values():
|
for proxy in self._proxies.values():
|
||||||
await proxy.queue.put(data)
|
await proxy.queue.put(data)
|
||||||
|
|
||||||
|
def _started(self) -> bool:
|
||||||
|
return self._proxies is not None
|
||||||
|
|
||||||
def _create_proxy(self, observer: BaseObserver) -> Proxy:
|
def _create_proxy(self, observer: BaseObserver) -> Proxy:
|
||||||
queue = asyncio.Queue()
|
queue = asyncio.Queue()
|
||||||
task = self._task_manager.create_task(
|
task = self._task_manager.create_task(
|
||||||
|
|||||||
@@ -142,7 +142,7 @@ class TestPipelineTask(unittest.IsolatedAsyncioTestCase):
|
|||||||
observer = CustomAddObserver()
|
observer = CustomAddObserver()
|
||||||
# Wait after the pipeline is started and add an observer.
|
# Wait after the pipeline is started and add an observer.
|
||||||
await asyncio.sleep(0.1)
|
await asyncio.sleep(0.1)
|
||||||
await task.add_observer(observer)
|
task.add_observer(observer)
|
||||||
# Push a TextFrame and wait for the observer to pick it up.
|
# Push a TextFrame and wait for the observer to pick it up.
|
||||||
await task.queue_frame(TextFrame(text="Hello Downstream!"))
|
await task.queue_frame(TextFrame(text="Hello Downstream!"))
|
||||||
await asyncio.sleep(0.1)
|
await asyncio.sleep(0.1)
|
||||||
|
|||||||
Reference in New Issue
Block a user