Add user_id to TranscriptionFrame and InterimTranscriptionFrame pushed by STTServices
This commit is contained in:
@@ -9,6 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|
||||||
|
<<<<<<< HEAD
|
||||||
- Added support for providing "direct" functions, which don't need an
|
- Added support for providing "direct" functions, which don't need an
|
||||||
accompanying `FunctionSchema` or function definition dict. Instead, metadata
|
accompanying `FunctionSchema` or function definition dict. Instead, metadata
|
||||||
(i.e. `name`, `description`, `properties`, and `required`) are automatically
|
(i.e. `name`, `description`, `properties`, and `required`) are automatically
|
||||||
@@ -39,6 +40,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
tools = ToolsSchema(standard_tools=[do_something])
|
tools = ToolsSchema(standard_tools=[do_something])
|
||||||
```
|
```
|
||||||
|
=======
|
||||||
|
- `user_id` is now populated in the `TranscriptionFrame` and
|
||||||
|
`InterimTranscriptionFrame` when using a service that provides a `user_id`,
|
||||||
|
like `DailyTransport` or `LiveKitTransport`.
|
||||||
|
>>>>>>> 5f958226 (Add user_id to TranscriptionFrame and InterimTranscriptionFrame pushed by STTServices)
|
||||||
|
|
||||||
- Added `watchdog_coroutine()`. This is a watchdog helper for couroutines. So,
|
- Added `watchdog_coroutine()`. This is a watchdog helper for couroutines. So,
|
||||||
if you have a coroutine that is waiting for a result and that takes a long
|
if you have a coroutine that is waiting for a result and that takes a long
|
||||||
|
|||||||
@@ -311,7 +311,7 @@ class AssemblyAISTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(
|
TranscriptionFrame(
|
||||||
message.transcript,
|
message.transcript,
|
||||||
"", # participant
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
self._language,
|
self._language,
|
||||||
message,
|
message,
|
||||||
@@ -323,7 +323,7 @@ class AssemblyAISTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
InterimTranscriptionFrame(
|
InterimTranscriptionFrame(
|
||||||
message.transcript,
|
message.transcript,
|
||||||
"", # participant
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
self._language,
|
self._language,
|
||||||
message,
|
message,
|
||||||
|
|||||||
@@ -366,7 +366,7 @@ class AWSTranscribeSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(
|
TranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
self._settings["language"],
|
self._settings["language"],
|
||||||
result=result,
|
result=result,
|
||||||
@@ -382,7 +382,7 @@ class AWSTranscribeSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
InterimTranscriptionFrame(
|
InterimTranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
self._settings["language"],
|
self._settings["language"],
|
||||||
result=result,
|
result=result,
|
||||||
|
|||||||
@@ -183,7 +183,7 @@ class AzureSTTService(STTService):
|
|||||||
language = getattr(event.result, "language", None) or self._settings.get("language")
|
language = getattr(event.result, "language", None) or self._settings.get("language")
|
||||||
frame = TranscriptionFrame(
|
frame = TranscriptionFrame(
|
||||||
event.result.text,
|
event.result.text,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
language,
|
language,
|
||||||
result=event,
|
result=event,
|
||||||
|
|||||||
@@ -289,14 +289,24 @@ class CartesiaSTTService(STTService):
|
|||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
if is_final:
|
if is_final:
|
||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(transcript, "", time_now_iso8601(), language)
|
TranscriptionFrame(
|
||||||
|
transcript,
|
||||||
|
self._user_id,
|
||||||
|
time_now_iso8601(),
|
||||||
|
language,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
await self._handle_transcription(transcript, is_final, language)
|
await self._handle_transcription(transcript, is_final, language)
|
||||||
await self.stop_processing_metrics()
|
await self.stop_processing_metrics()
|
||||||
else:
|
else:
|
||||||
# For interim transcriptions, just push the frame without tracing
|
# For interim transcriptions, just push the frame without tracing
|
||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
InterimTranscriptionFrame(transcript, "", time_now_iso8601(), language)
|
InterimTranscriptionFrame(
|
||||||
|
transcript,
|
||||||
|
self._user_id,
|
||||||
|
time_now_iso8601(),
|
||||||
|
language,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
|
|||||||
@@ -278,7 +278,7 @@ class DeepgramSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(
|
TranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
language,
|
language,
|
||||||
result=result,
|
result=result,
|
||||||
@@ -291,7 +291,7 @@ class DeepgramSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
InterimTranscriptionFrame(
|
InterimTranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
language,
|
language,
|
||||||
result=result,
|
result=result,
|
||||||
|
|||||||
@@ -291,7 +291,7 @@ class FalSTTService(SegmentedSTTService):
|
|||||||
logger.debug(f"Transcription: [{text}]")
|
logger.debug(f"Transcription: [{text}]")
|
||||||
yield TranscriptionFrame(
|
yield TranscriptionFrame(
|
||||||
text,
|
text,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
Language(self._settings["language"]),
|
Language(self._settings["language"]),
|
||||||
result=response,
|
result=response,
|
||||||
|
|||||||
@@ -567,7 +567,7 @@ class GladiaSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(
|
TranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
language,
|
language,
|
||||||
result=content,
|
result=content,
|
||||||
@@ -582,7 +582,7 @@ class GladiaSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
InterimTranscriptionFrame(
|
InterimTranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
language,
|
language,
|
||||||
result=content,
|
result=content,
|
||||||
|
|||||||
@@ -862,7 +862,7 @@ class GoogleSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(
|
TranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
primary_language,
|
primary_language,
|
||||||
result=result,
|
result=result,
|
||||||
@@ -880,7 +880,7 @@ class GoogleSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
InterimTranscriptionFrame(
|
InterimTranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
primary_language,
|
primary_language,
|
||||||
result=result,
|
result=result,
|
||||||
|
|||||||
@@ -314,7 +314,7 @@ class RivaSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(
|
TranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
self._language_code,
|
self._language_code,
|
||||||
result=result,
|
result=result,
|
||||||
@@ -329,7 +329,7 @@ class RivaSTTService(STTService):
|
|||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
InterimTranscriptionFrame(
|
InterimTranscriptionFrame(
|
||||||
transcript,
|
transcript,
|
||||||
"",
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
self._language_code,
|
self._language_code,
|
||||||
result=result,
|
result=result,
|
||||||
@@ -636,7 +636,10 @@ class RivaSegmentedSTTService(SegmentedSTTService):
|
|||||||
if text:
|
if text:
|
||||||
logger.debug(f"Transcription: [{text}]")
|
logger.debug(f"Transcription: [{text}]")
|
||||||
yield TranscriptionFrame(
|
yield TranscriptionFrame(
|
||||||
text, "", time_now_iso8601(), self._language_enum
|
text,
|
||||||
|
self._user_id,
|
||||||
|
time_now_iso8601(),
|
||||||
|
self._language_enum,
|
||||||
)
|
)
|
||||||
transcription_found = True
|
transcription_found = True
|
||||||
|
|
||||||
|
|||||||
@@ -57,6 +57,7 @@ class STTService(AIService):
|
|||||||
self._sample_rate = 0
|
self._sample_rate = 0
|
||||||
self._settings: Dict[str, Any] = {}
|
self._settings: Dict[str, Any] = {}
|
||||||
self._muted: bool = False
|
self._muted: bool = False
|
||||||
|
self._user_id: str = ""
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_muted(self) -> bool:
|
def is_muted(self) -> bool:
|
||||||
@@ -132,6 +133,11 @@ class STTService(AIService):
|
|||||||
async def process_audio_frame(self, frame: AudioRawFrame, direction: FrameDirection):
|
async def process_audio_frame(self, frame: AudioRawFrame, direction: FrameDirection):
|
||||||
"""Process an audio frame for speech recognition.
|
"""Process an audio frame for speech recognition.
|
||||||
|
|
||||||
|
If the service is muted, this method does nothing. Otherwise, it
|
||||||
|
processes the audio frame and runs speech-to-text on it, yielding
|
||||||
|
transcription results. If the frame has a user_id, it is stored
|
||||||
|
for later use in transcription.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
frame: The audio frame to process.
|
frame: The audio frame to process.
|
||||||
direction: The direction of frame processing.
|
direction: The direction of frame processing.
|
||||||
@@ -139,6 +145,13 @@ class STTService(AIService):
|
|||||||
if self._muted:
|
if self._muted:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# UserAudioRawFrame contains a user_id (e.g. Daily, Livekit)
|
||||||
|
if hasattr(frame, "user_id"):
|
||||||
|
self._user_id = frame.user_id
|
||||||
|
# AudioRawFrame does not have a user_id (e.g. SmallWebRTCTransport, websockets)
|
||||||
|
else:
|
||||||
|
self._user_id = ""
|
||||||
|
|
||||||
await self.process_generator(self.run_stt(frame.audio))
|
await self.process_generator(self.run_stt(frame.audio))
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
@@ -241,10 +254,19 @@ class SegmentedSTTService(STTService):
|
|||||||
Continuously buffers audio, growing the buffer while user is speaking and
|
Continuously buffers audio, growing the buffer while user is speaking and
|
||||||
maintaining a small buffer when not speaking to account for VAD delay.
|
maintaining a small buffer when not speaking to account for VAD delay.
|
||||||
|
|
||||||
|
If the frame has a user_id, it is stored for later use in transcription.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
frame: The audio frame to process.
|
frame: The audio frame to process.
|
||||||
direction: The direction of frame processing.
|
direction: The direction of frame processing.
|
||||||
"""
|
"""
|
||||||
|
# UserAudioRawFrame contains a user_id (e.g. Daily, Livekit)
|
||||||
|
if hasattr(frame, "user_id"):
|
||||||
|
self._user_id = frame.user_id
|
||||||
|
# AudioRawFrame does not have a user_id (e.g. SmallWebRTCTransport, websockets)
|
||||||
|
else:
|
||||||
|
self._user_id = ""
|
||||||
|
|
||||||
# If the user is speaking the audio buffer will keep growing.
|
# If the user is speaking the audio buffer will keep growing.
|
||||||
self._audio_buffer += frame.audio
|
self._audio_buffer += frame.audio
|
||||||
|
|
||||||
|
|||||||
@@ -219,7 +219,11 @@ class BaseWhisperSTTService(SegmentedSTTService):
|
|||||||
if text:
|
if text:
|
||||||
await self._handle_transcription(text, True, self._language)
|
await self._handle_transcription(text, True, self._language)
|
||||||
logger.debug(f"Transcription: [{text}]")
|
logger.debug(f"Transcription: [{text}]")
|
||||||
yield TranscriptionFrame(text, "", time_now_iso8601())
|
yield TranscriptionFrame(
|
||||||
|
text,
|
||||||
|
self._user_id,
|
||||||
|
time_now_iso8601(),
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
logger.warning("Received empty transcription from API")
|
logger.warning("Received empty transcription from API")
|
||||||
|
|
||||||
|
|||||||
@@ -395,7 +395,12 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
if text:
|
if text:
|
||||||
await self._handle_transcription(text, True, self._settings["language"])
|
await self._handle_transcription(text, True, self._settings["language"])
|
||||||
logger.debug(f"Transcription: [{text}]")
|
logger.debug(f"Transcription: [{text}]")
|
||||||
yield TranscriptionFrame(text, "", time_now_iso8601(), self._settings["language"])
|
yield TranscriptionFrame(
|
||||||
|
text,
|
||||||
|
self._user_id,
|
||||||
|
time_now_iso8601(),
|
||||||
|
self._settings["language"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class WhisperSTTServiceMLX(WhisperSTTService):
|
class WhisperSTTServiceMLX(WhisperSTTService):
|
||||||
@@ -500,7 +505,12 @@ class WhisperSTTServiceMLX(WhisperSTTService):
|
|||||||
if text:
|
if text:
|
||||||
await self._handle_transcription(text, True, self._settings["language"])
|
await self._handle_transcription(text, True, self._settings["language"])
|
||||||
logger.debug(f"Transcription: [{text}]")
|
logger.debug(f"Transcription: [{text}]")
|
||||||
yield TranscriptionFrame(text, "", time_now_iso8601(), self._settings["language"])
|
yield TranscriptionFrame(
|
||||||
|
text,
|
||||||
|
self._user_id,
|
||||||
|
time_now_iso8601(),
|
||||||
|
self._settings["language"],
|
||||||
|
)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"MLX Whisper transcription error: {e}")
|
logger.exception(f"MLX Whisper transcription error: {e}")
|
||||||
|
|||||||
Reference in New Issue
Block a user