feat(qwen-audio): enhance user transcript handling in QwenAudioRealtimeService
- Added support for managing user transcripts, including pending states and timestamps. - Implemented methods to handle user transcript completion and deferred assistant messages. - Updated event handling to ensure user transcripts are emitted before assistant responses. - Enhanced tests to verify the correct order of transcript and assistant message emissions during user interactions.
This commit is contained in:
@@ -111,6 +111,10 @@ class QwenAudioRealtimeService(AIService):
|
||||
self._assistant_text = ""
|
||||
self._assistant_timestamp = ""
|
||||
self._greeting_request_item_id: str | None = None
|
||||
self._user_transcript_pending = False
|
||||
self._user_transcript_item_id = ""
|
||||
self._user_transcript_timestamp = ""
|
||||
self._deferred_assistant_messages: list[dict[str, Any]] = []
|
||||
|
||||
async def start(self, frame: StartFrame) -> None:
|
||||
await super().start(frame)
|
||||
@@ -306,6 +310,10 @@ class QwenAudioRealtimeService(AIService):
|
||||
self._session_ready.clear()
|
||||
self._pending_events.clear()
|
||||
self._response_active = False
|
||||
self._user_transcript_pending = False
|
||||
self._user_transcript_item_id = ""
|
||||
self._user_transcript_timestamp = ""
|
||||
self._deferred_assistant_messages.clear()
|
||||
if websocket and websocket.state is State.OPEN:
|
||||
try:
|
||||
await websocket.close()
|
||||
@@ -370,10 +378,19 @@ class QwenAudioRealtimeService(AIService):
|
||||
else:
|
||||
await self._append_assistant_text(transcript)
|
||||
elif event_type == "conversation.item.input_audio_transcription.completed":
|
||||
await self._send_transcript("user", str(event.get("transcript") or ""))
|
||||
await self._handle_user_transcript_completed(event)
|
||||
elif event_type == "conversation.item.input_audio_transcription.failed":
|
||||
await self._release_deferred_assistant_messages()
|
||||
elif event_type == "input_audio_buffer.speech_started":
|
||||
user_turn_timestamp = time_now_iso8601()
|
||||
await self._cancel_active_response()
|
||||
await self.broadcast_interruption()
|
||||
await self._start_user_transcript_turn(event, user_turn_timestamp)
|
||||
elif (
|
||||
event_type == "input_audio_buffer.speech_stopped"
|
||||
and event.get("reason") == "turn_invalid"
|
||||
):
|
||||
await self._release_deferred_assistant_messages()
|
||||
elif event_type == "response.done":
|
||||
response = event.get("response")
|
||||
status = response.get("status") if isinstance(response, dict) else None
|
||||
@@ -413,6 +430,67 @@ class QwenAudioRealtimeService(AIService):
|
||||
wait_until_ready=False,
|
||||
)
|
||||
|
||||
async def _start_user_transcript_turn(
|
||||
self,
|
||||
event: dict[str, Any],
|
||||
timestamp: str,
|
||||
) -> None:
|
||||
"""Hold the next assistant transcript until this user turn is visible.
|
||||
|
||||
Qwen produces the assistant response and the final ASR transcript on
|
||||
independent streams. The assistant delta can therefore arrive first.
|
||||
Keep this provider-specific ordering rule inside the adapter so the
|
||||
shared cascade pipeline and Debug Drawer protocol remain unchanged.
|
||||
"""
|
||||
if self._user_transcript_pending:
|
||||
await self._release_deferred_assistant_messages()
|
||||
self._user_transcript_pending = True
|
||||
self._user_transcript_item_id = str(event.get("item_id") or "")
|
||||
self._user_transcript_timestamp = timestamp
|
||||
|
||||
async def _handle_user_transcript_completed(
|
||||
self,
|
||||
event: dict[str, Any],
|
||||
) -> None:
|
||||
item_id = str(event.get("item_id") or "")
|
||||
timestamp = (
|
||||
self._user_transcript_timestamp
|
||||
if self._user_transcript_pending
|
||||
and (
|
||||
not self._user_transcript_item_id
|
||||
or not item_id
|
||||
or item_id == self._user_transcript_item_id
|
||||
)
|
||||
else ""
|
||||
)
|
||||
await self._send_transcript(
|
||||
"user",
|
||||
str(event.get("transcript") or ""),
|
||||
timestamp=timestamp or None,
|
||||
)
|
||||
if timestamp:
|
||||
await self._release_deferred_assistant_messages()
|
||||
|
||||
async def _send_assistant_transport_message(
|
||||
self,
|
||||
message: dict[str, Any],
|
||||
) -> None:
|
||||
if self._user_transcript_pending:
|
||||
self._deferred_assistant_messages.append(message)
|
||||
return
|
||||
await self._send_transport_message(message)
|
||||
|
||||
async def _release_deferred_assistant_messages(self) -> None:
|
||||
self._user_transcript_pending = False
|
||||
self._user_transcript_item_id = ""
|
||||
self._user_transcript_timestamp = ""
|
||||
pending, self._deferred_assistant_messages = (
|
||||
self._deferred_assistant_messages,
|
||||
[],
|
||||
)
|
||||
for message in pending:
|
||||
await self._send_transport_message(message)
|
||||
|
||||
async def _send_event(
|
||||
self, payload: dict[str, Any], *, wait_until_ready: bool = True
|
||||
) -> None:
|
||||
@@ -430,7 +508,7 @@ class QwenAudioRealtimeService(AIService):
|
||||
if not self._assistant_turn_id:
|
||||
self._assistant_turn_id = uuid4().hex
|
||||
self._assistant_timestamp = time_now_iso8601()
|
||||
await self._send_transport_message(
|
||||
await self._send_assistant_transport_message(
|
||||
{
|
||||
"type": "assistant-text-start",
|
||||
"turn_id": self._assistant_turn_id,
|
||||
@@ -438,7 +516,7 @@ class QwenAudioRealtimeService(AIService):
|
||||
}
|
||||
)
|
||||
self._assistant_text += delta
|
||||
await self._send_transport_message(
|
||||
await self._send_assistant_transport_message(
|
||||
{
|
||||
"type": "assistant-text-delta",
|
||||
"turn_id": self._assistant_turn_id,
|
||||
@@ -449,7 +527,7 @@ class QwenAudioRealtimeService(AIService):
|
||||
async def _finish_assistant_text(self, *, interrupted: bool) -> None:
|
||||
if not self._assistant_turn_id:
|
||||
return
|
||||
await self._send_transport_message(
|
||||
await self._send_assistant_transport_message(
|
||||
{
|
||||
"type": "assistant-text-end",
|
||||
"turn_id": self._assistant_turn_id,
|
||||
@@ -461,14 +539,20 @@ class QwenAudioRealtimeService(AIService):
|
||||
self._assistant_text = ""
|
||||
self._assistant_timestamp = ""
|
||||
|
||||
async def _send_transcript(self, role: str, content: str) -> None:
|
||||
async def _send_transcript(
|
||||
self,
|
||||
role: str,
|
||||
content: str,
|
||||
*,
|
||||
timestamp: str | None = None,
|
||||
) -> None:
|
||||
if content:
|
||||
await self._send_transport_message(
|
||||
{
|
||||
"type": "transcript",
|
||||
"role": role,
|
||||
"content": content,
|
||||
"timestamp": time_now_iso8601(),
|
||||
"timestamp": timestamp or time_now_iso8601(),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -150,6 +150,88 @@ class QwenAudioRealtimeServiceTest(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertEqual(end_messages[0]["content"], "Hi")
|
||||
self.assertTrue(end_messages[0]["interrupted"])
|
||||
|
||||
async def test_user_transcript_is_emitted_before_early_assistant_delta(self):
|
||||
service = _service()
|
||||
service.push_frame = AsyncMock()
|
||||
service.broadcast_interruption = AsyncMock()
|
||||
|
||||
with patch(
|
||||
"services.pipecat.qwen_audio_realtime.time_now_iso8601",
|
||||
side_effect=[
|
||||
"2026-07-29T10:00:00+00:00",
|
||||
"2026-07-29T10:00:01+00:00",
|
||||
],
|
||||
):
|
||||
await service._handle_server_event(
|
||||
{
|
||||
"type": "input_audio_buffer.speech_started",
|
||||
"item_id": "item_user",
|
||||
}
|
||||
)
|
||||
await service._handle_server_event({"type": "response.created"})
|
||||
await service._handle_server_event(
|
||||
{
|
||||
"type": "response.audio_transcript.delta",
|
||||
"delta": "您好",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(service.push_frame.await_count, 0)
|
||||
|
||||
await service._handle_server_event(
|
||||
{
|
||||
"type": "conversation.item.input_audio_transcription.completed",
|
||||
"item_id": "item_user",
|
||||
"transcript": "你好",
|
||||
}
|
||||
)
|
||||
|
||||
messages = [
|
||||
call.args[0].message
|
||||
for call in service.push_frame.await_args_list
|
||||
if isinstance(call.args[0], OutputTransportMessageUrgentFrame)
|
||||
]
|
||||
self.assertEqual(
|
||||
[message["type"] for message in messages],
|
||||
["transcript", "assistant-text-start", "assistant-text-delta"],
|
||||
)
|
||||
self.assertEqual(messages[0]["timestamp"], "2026-07-29T10:00:00+00:00")
|
||||
self.assertEqual(messages[1]["timestamp"], "2026-07-29T10:00:01+00:00")
|
||||
|
||||
async def test_transcription_failure_releases_assistant_delta(self):
|
||||
service = _service()
|
||||
service.push_frame = AsyncMock()
|
||||
service.broadcast_interruption = AsyncMock()
|
||||
|
||||
await service._handle_server_event(
|
||||
{
|
||||
"type": "input_audio_buffer.speech_started",
|
||||
"item_id": "item_user",
|
||||
}
|
||||
)
|
||||
await service._handle_server_event({"type": "response.created"})
|
||||
await service._handle_server_event(
|
||||
{"type": "response.audio_transcript.delta", "delta": "您好"}
|
||||
)
|
||||
self.assertEqual(service.push_frame.await_count, 0)
|
||||
|
||||
await service._handle_server_event(
|
||||
{
|
||||
"type": "conversation.item.input_audio_transcription.failed",
|
||||
"item_id": "item_user",
|
||||
}
|
||||
)
|
||||
|
||||
messages = [
|
||||
call.args[0].message
|
||||
for call in service.push_frame.await_args_list
|
||||
if isinstance(call.args[0], OutputTransportMessageUrgentFrame)
|
||||
]
|
||||
self.assertEqual(
|
||||
[message["type"] for message in messages],
|
||||
["assistant-text-start", "assistant-text-delta"],
|
||||
)
|
||||
|
||||
async def test_greeting_request_is_removed_after_response(self):
|
||||
service = _service()
|
||||
websocket = _OpenWebSocket()
|
||||
|
||||
Reference in New Issue
Block a user