Addressing the comments left in the PR review.

This commit is contained in:
filipi87
2026-05-12 17:12:19 -03:00
parent c2bdc1aada
commit 0146947b68
2 changed files with 18 additions and 25 deletions

View File

@@ -46,15 +46,15 @@ from pipecat.utils.tracing.service_decorators import traced_stt
@dataclass @dataclass
class NvidiaSageMakerWSSTTSettings(STTSettings): class NvidiaSageMakerSTTSettings(STTSettings):
"""Settings for NvidiaSageMakerWebsocketSTTService. """Settings for NvidiaSageMakerSTTService.
Parameters: Parameters:
language: ISO-639-1 language code passed to NIM (e.g. ``en-US``). language: ISO-639-1 language code passed to NIM (e.g. ``en-US``).
""" """
class NvidiaSageMakerWebsocketSTTService(STTService): class NvidiaSageMakerSTTService(STTService):
"""NVIDIA Nemotron ASR STT service using SageMaker bidirectional streaming. """NVIDIA Nemotron ASR STT service using SageMaker bidirectional streaming.
Maintains a persistent HTTP/2 bidi-stream session to the SageMaker endpoint Maintains a persistent HTTP/2 bidi-stream session to the SageMaker endpoint
@@ -65,16 +65,16 @@ class NvidiaSageMakerWebsocketSTTService(STTService):
Example:: Example::
stt = NvidiaSageMakerWebsocketSTTService( stt = NvidiaSageMakerSTTService(
endpoint_name=os.getenv("SAGEMAKER_ASR_ENDPOINT_NAME"), endpoint_name=os.getenv("SAGEMAKER_ASR_ENDPOINT_NAME"),
region=os.getenv("AWS_REGION", "us-west-2"), region=os.getenv("AWS_REGION", "us-west-2"),
settings=NvidiaSageMakerWebsocketSTTService.Settings( settings=NvidiaSageMakerSTTService.Settings(
language="en-US", language="en-US",
), ),
) )
""" """
Settings = NvidiaSageMakerWSSTTSettings Settings = NvidiaSageMakerSTTSettings
def __init__( def __init__(
self, self,
@@ -82,7 +82,7 @@ class NvidiaSageMakerWebsocketSTTService(STTService):
endpoint_name: str, endpoint_name: str,
region: str = "us-west-2", region: str = "us-west-2",
sample_rate: int | None = None, sample_rate: int | None = None,
settings: NvidiaSageMakerWSSTTSettings | None = None, settings: NvidiaSageMakerSTTSettings | None = None,
ttfs_p99_latency: float | None = 1.5, ttfs_p99_latency: float | None = 1.5,
**kwargs, **kwargs,
): ):
@@ -301,6 +301,8 @@ class NvidiaSageMakerWebsocketSTTService(STTService):
delta, delta,
self._user_id, self._user_id,
time_now_iso8601(), time_now_iso8601(),
language=self._settings.language,
result=msg,
) )
) )
@@ -313,7 +315,9 @@ class NvidiaSageMakerWebsocketSTTService(STTService):
transcript, transcript,
self._user_id, self._user_id,
time_now_iso8601(), time_now_iso8601(),
language=self._settings.language,
result=msg, result=msg,
finalized=True,
) )
) )
await self._handle_transcription(transcript, True) await self._handle_transcription(transcript, True)

View File

@@ -35,7 +35,7 @@ from pipecat.utils.tracing.service_decorators import traced_tts
@dataclass @dataclass
class NvidiaSageMakerTTSSettings(TTSSettings): class NvidiaSageMakerTTSSettings(TTSSettings):
"""Settings for NvidiaSageMakerHTTPTTSService. """Settings for NVIDIA SageMaker TTS services.
Parameters: Parameters:
voice: NIM voice name (e.g. ``Magpie-Multilingual.EN-US.Aria``). voice: NIM voice name (e.g. ``Magpie-Multilingual.EN-US.Aria``).
@@ -79,7 +79,6 @@ class NvidiaSageMakerHTTPTTSService(TTSService):
endpoint_name: Name of the deployed SageMaker endpoint. endpoint_name: Name of the deployed SageMaker endpoint.
region: AWS region where the endpoint lives. region: AWS region where the endpoint lives.
sample_rate: Output sample rate in Hz. Defaults to bot's pipeline rate. sample_rate: Output sample rate in Hz. Defaults to bot's pipeline rate.
params: Deprecated — use ``settings`` instead.
settings: Runtime-updatable settings (voice, language). settings: Runtime-updatable settings (voice, language).
**kwargs: Forwarded to :class:`TTSService`. **kwargs: Forwarded to :class:`TTSService`.
""" """
@@ -218,17 +217,7 @@ class NvidiaSageMakerHTTPTTSService(TTSService):
await self.start_tts_usage_metrics(text) await self.start_tts_usage_metrics(text)
@dataclass class NvidiaSageMakerTTSService(InterruptibleTTSService):
class NvidiaSageMakerWSTTSSettings(TTSSettings):
"""Settings for NvidiaSageMakerWebsocketTTSService.
Parameters:
voice: NIM voice name (e.g. ``Magpie-Multilingual.EN-US.Aria``).
language: BCP-47 language code passed to NIM (e.g. ``en-US``).
"""
class NvidiaSageMakerWebsocketTTSService(InterruptibleTTSService):
"""NVIDIA Magpie TTS service using SageMaker bidirectional streaming. """NVIDIA Magpie TTS service using SageMaker bidirectional streaming.
Maintains a persistent HTTP/2 bidi-stream session to the SageMaker endpoint Maintains a persistent HTTP/2 bidi-stream session to the SageMaker endpoint
@@ -238,17 +227,17 @@ class NvidiaSageMakerWebsocketTTSService(InterruptibleTTSService):
Example:: Example::
tts = NvidiaSageMakerWebsocketTTSService( tts = NvidiaSageMakerTTSService(
endpoint_name=os.getenv("SAGEMAKER_MAGPIE_ENDPOINT_NAME"), endpoint_name=os.getenv("SAGEMAKER_MAGPIE_ENDPOINT_NAME"),
region=os.getenv("AWS_REGION", "us-west-2"), region=os.getenv("AWS_REGION", "us-west-2"),
settings=NvidiaSageMakerWebsocketTTSService.Settings( settings=NvidiaSageMakerTTSService.Settings(
voice="Magpie-Multilingual.EN-US.Aria", voice="Magpie-Multilingual.EN-US.Aria",
language="en-US", language="en-US",
), ),
) )
""" """
Settings = NvidiaSageMakerWSTTSSettings Settings = NvidiaSageMakerTTSSettings
def __init__( def __init__(
self, self,
@@ -256,7 +245,7 @@ class NvidiaSageMakerWebsocketTTSService(InterruptibleTTSService):
endpoint_name: str, endpoint_name: str,
region: str = "us-west-2", region: str = "us-west-2",
sample_rate: int | None = None, sample_rate: int | None = None,
settings: NvidiaSageMakerWSTTSSettings | None = None, settings: NvidiaSageMakerTTSSettings | None = None,
**kwargs, **kwargs,
): ):
"""Initialize the SageMaker WebSocket TTS service. """Initialize the SageMaker WebSocket TTS service.
@@ -507,4 +496,4 @@ class NvidiaSageMakerWebsocketTTSService(InterruptibleTTSService):
yield None yield None
except Exception as e: except Exception as e:
logger.error(f"{self}: TTS error: {e}") logger.error(f"{self}: TTS error: {e}")
yield ErrorFrame(error=f"NvidiaSageMakerWebsocketTTSService error: {e}") yield ErrorFrame(error=f"NvidiaSageMakerTTSService error: {e}")