diff --git a/private_gpt/components/engines/chat/async_chat_engine.py b/private_gpt/components/engines/chat/async_chat_engine.py index 9d77986f..8c8fbbb6 100644 --- a/private_gpt/components/engines/chat/async_chat_engine.py +++ b/private_gpt/components/engines/chat/async_chat_engine.py @@ -753,23 +753,34 @@ class AsyncChatEngine: if tool.name is not None } - for response in payload.tool_responses: - tool_spec = tool_specs_by_name.get(response.tool_name) - if tool_spec is None: - raise RuntimeError( - f"Missing tool spec for resumed tool {response.tool_name!r}" - ) - result_start = RawContentBlockStartEvent( - index=run.block_count, - block_id=f"block_{uuid4().hex}", - content_block=tool_spec.resolve_event_adapter().build_tool_result( - tool_use_id=response.tool_id, - outcome=response.outcome, - ), - ) - run.block_count += 1 - channel.emit(result_start) - channel.emit(RawContentBlockStopEvent.from_start(result_start)) + async def _emit_tool_results(handler: _EventHandler) -> None: + try: + for response in payload.tool_responses: + tool_spec = tool_specs_by_name.get(response.tool_name) + if tool_spec is None: + raise RuntimeError( + f"Missing tool spec for resumed tool {response.tool_name!r}" + ) + result_start = RawContentBlockStartEvent( + index=run.block_count, + block_id=f"block_{uuid4().hex}", + content_block=tool_spec.resolve_event_adapter().build_tool_result( + tool_use_id=response.tool_id, + outcome=response.outcome, + ), + ) + run.block_count += 1 + handler.emit(result_start) + handler.emit(RawContentBlockStopEvent.from_start(result_start)) + finally: + await handler.close() + + await self._pipe_events_through_interceptors( + run=run, + channel=channel, + producer=_emit_tool_results, + call_iteration_hooks=True, + ) if pending_external: run.state = run.state.model_copy(deep=True) @@ -862,11 +873,22 @@ class AsyncChatEngine: ) -> _Run: return self._initialize_run(request, context_stack=context_stack, hooks=hooks) - async def _run_intercepted_iteration( + async def _pipe_events_through_interceptors( self, run: _Run, channel: EventChannel, + producer: Callable[["_EventHandler"], Awaitable[None]], + call_iteration_hooks: bool = False, ) -> None: + """Emit events from *producer* through all response interceptors into *channel*. + + *producer* receives the handler it should emit events to and must close + it (via ``await handler.close()``) when done so the consumer loop exits. + + When *call_iteration_hooks* is True, ``on_iteration_start`` and + ``on_iteration_end`` are called on each interceptor around the producer, + which is required for stateful interceptors (e.g. index tracking). + """ iter_handler = _EventHandler(queue=asyncio.Queue()) def current_context() -> ChatInterceptorContext: @@ -877,18 +899,21 @@ class AsyncChatEngine: emit_fn=iter_handler.emit, ) - async def _run_and_close() -> None: - try: - for interceptor in self._response_interceptors: - await interceptor.on_iteration_start(current_context()) - await self._run_iteration(run, iter_handler) - finally: - for interceptor in self._response_interceptors: - with suppress(Exception): - await interceptor.on_iteration_end(current_context()) - await iter_handler.close() + ctx = current_context() + if call_iteration_hooks: + for interceptor in self._response_interceptors: + await interceptor.on_iteration_start(ctx) - iter_task = asyncio.create_task(_run_and_close()) + async def _wrapped_producer(handler: "_EventHandler") -> None: + try: + await producer(handler=handler) + finally: + if call_iteration_hooks: + for interceptor in self._response_interceptors: + with suppress(Exception): + await interceptor.on_iteration_end(current_context()) + + iter_task = asyncio.create_task(_wrapped_producer(handler=iter_handler)) async for raw_event in iter_handler.stream(iter_task): processed: Event | None = raw_event @@ -903,6 +928,24 @@ class AsyncChatEngine: if processed is not None: channel.emit(processed) + async def _run_intercepted_iteration( + self, + run: _Run, + channel: EventChannel, + ) -> None: + async def _run_and_close(handler: "_EventHandler") -> None: + try: + await self._run_iteration(run=run, handler=handler) + finally: + await handler.close() + + await self._pipe_events_through_interceptors( + run=run, + channel=channel, + producer=_run_and_close, + call_iteration_hooks=True, + ) + async def _run_iteration(self, run: _Run, handler: _EventHandler) -> None: """One LLM call. Sets status; stop events are emitted by run_close.""" run.state = run.state.model_copy(deep=True)