Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 49 additions & 22 deletions livekit-agents/livekit/agents/voice/agent_activity.py
Original file line number Diff line number Diff line change
Expand Up @@ -3814,6 +3814,9 @@ async def _realtime_reply_task(
# cancel the pending generation; the plugin emits response.cancel
if not generate_reply_fut.done():
generate_reply_fut.cancel()
else:
# the response already landed, cancelling the future no longer reaches it
self._rt_session.interrupt()
return

try:
Expand Down Expand Up @@ -3909,6 +3912,41 @@ async def _realtime_generation_task_impl(
)
tool_ctx = llm.ToolContext(self.tools)

tasks: list[asyncio.Task[Any]] = []
tees: list[utils.aio.itertools.Tee[Any]] = []

msg_tee = utils.aio.itertools.tee(generation_ev.message_stream, 2)
msg_stream, msg_stream_to_drain = msg_tee
tees.append(msg_tee)

fnc_tee = utils.aio.itertools.tee(generation_ev.function_stream, 2)
fnc_stream, fnc_stream_for_tracing = fnc_tee
tees.append(fnc_tee)

function_calls: list[llm.FunctionCall] = []
generation_ended = False

async def _read_fnc_stream() -> None:
async for fnc in fnc_stream_for_tracing:
function_calls.append(fnc)

async def _drain_msg_stream() -> None:
async for _ in msg_stream_to_drain:
pass

async def _watch_generation_end() -> None:
# the provider closes both streams when it ends the response
nonlocal generation_ended
await asyncio.gather(_drain_msg_stream(), _read_fnc_stream())
generation_ended = True

tasks.append(
asyncio.create_task(
_watch_generation_end(),
name="AgentActivity.realtime_generation.watch_end",
)
)

authorization_tasks: list[asyncio.Future[Any]] = [
asyncio.ensure_future(speech_handle._wait_for_authorization()),
asyncio.ensure_future(self._authorization_allowed.wait()),
Expand All @@ -3919,7 +3957,12 @@ async def _realtime_generation_task_impl(
speech_handle._clear_authorization()

if speech_handle.interrupted:
await utils.aio.cancel_and_wait(*authorization_tasks)
# nothing was played, but the response may still be generating server-side
if not generation_ended:
self._rt_session.interrupt()
await utils.aio.cancel_and_wait(*authorization_tasks, *tasks)
for tee in tees:
await tee.aclose()
current_span.set_attribute(trace_types.ATTR_SPEECH_INTERRUPTED, True)
return # TODO(theomonnom): remove the message from the serverside history

Expand Down Expand Up @@ -3959,9 +4002,6 @@ def _on_first_frame(
if self.interruption_enabled:
self._disable_vad_interruption_soon()

tasks: list[asyncio.Task[Any]] = []
tees: list[utils.aio.itertools.Tee[Any]] = []

read_transcript_from_tts = False

# multiple message items may be produced for a single realtime response
Expand Down Expand Up @@ -4051,7 +4091,7 @@ async def _process_one_message(msg: MessageGeneration) -> _MsgOutput:

@utils.log_exceptions(logger=logger)
async def _process_messages() -> None:
async for msg in generation_ev.message_stream:
async for msg in msg_stream:
if speech_handle.interrupted:
# remaining messages are left out of message_outputs so
# update_chat_ctx below removes them server-side.
Expand All @@ -4066,23 +4106,6 @@ async def _process_messages() -> None:
)
tasks.append(process_msg_task)

# read function calls
fnc_tee = utils.aio.itertools.tee(generation_ev.function_stream, 2)
fnc_stream, fnc_stream_for_tracing = fnc_tee
tees.append(fnc_tee)
function_calls: list[llm.FunctionCall] = []

async def _read_fnc_stream() -> None:
async for fnc in fnc_stream_for_tracing:
function_calls.append(fnc)

tasks.append(
asyncio.create_task(
_read_fnc_stream(),
name="AgentActivity.realtime_generation.read_fnc_stream",
)
)

# messages in RunResult are ordered by the `created_at` field
def _tool_execution_started_cb(fnc_call: llm.FunctionCall) -> None:
# function call is created during the realtime generation, before the assistant
Expand All @@ -4108,6 +4131,10 @@ def _tool_execution_completed_cb(out: ToolExecutionOutput) -> None:

await speech_handle.wait_if_not_interrupted([*tasks])

# the cancel is session-wide, so send it only while this response is still generating
if speech_handle.interrupted and not generation_ended:
self._rt_session.interrupt()

current_span.set_attribute(trace_types.ATTR_SPEECH_INTERRUPTED, speech_handle.interrupted)
current_span.set_attribute(
trace_types.ATTR_RESPONSE_FUNCTION_CALLS,
Expand Down
Loading