mirror of
https://github.com/imartinez/privateGPT.git
synced 2026-08-08 23:35:29 +00:00
fix: response output
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user