Fixing sagemaker merge conflicts.

This commit is contained in:
filipi87
2026-03-03 11:03:51 -03:00
parent fc905a7ef5
commit 0fdf7dc16a

View File

@@ -15,7 +15,7 @@ languages, and various Deepgram features.
import asyncio import asyncio
import json import json
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, AsyncGenerator, Dict, Optional from typing import Any, AsyncGenerator, Optional
from loguru import logger from loguru import logger
@@ -32,7 +32,7 @@ from pipecat.frames.frames import (
) )
from pipecat.processors.frame_processor import FrameDirection from pipecat.processors.frame_processor import FrameDirection
from pipecat.services.aws.sagemaker.bidi_client import SageMakerBidiClient from pipecat.services.aws.sagemaker.bidi_client import SageMakerBidiClient
from pipecat.services.deepgram.stt import _DeepgramSTTSettingsBase from pipecat.services.deepgram.stt import DeepgramSTTSettings
from pipecat.services.settings import STTSettings from pipecat.services.settings import STTSettings
from pipecat.services.stt_latency import DEEPGRAM_SAGEMAKER_TTFS_P99 from pipecat.services.stt_latency import DEEPGRAM_SAGEMAKER_TTFS_P99
from pipecat.services.stt_service import STTService from pipecat.services.stt_service import STTService
@@ -41,7 +41,7 @@ from pipecat.utils.time import time_now_iso8601
from pipecat.utils.tracing.service_decorators import traced_stt from pipecat.utils.tracing.service_decorators import traced_stt
try: try:
from deepgram import LiveOptions from pipecat.services.deepgram.stt import LiveOptions
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
logger.error(f"Exception: {e}") logger.error(f"Exception: {e}")
logger.error( logger.error(
@@ -51,10 +51,10 @@ except ModuleNotFoundError as e:
@dataclass @dataclass
class DeepgramSageMakerSTTSettings(_DeepgramSTTSettingsBase): class DeepgramSageMakerSTTSettings(DeepgramSTTSettings):
"""Settings for the Deepgram SageMaker STT service. """Settings for the Deepgram SageMaker STT service.
See ``_DeepgramSTTSettingsBase`` for full documentation. See ``DeepgramSTTSettings`` for full documentation.
""" """
pass pass
@@ -117,22 +117,26 @@ class DeepgramSageMakerSTTService(STTService):
""" """
sample_rate = sample_rate or (live_options.sample_rate if live_options else None) sample_rate = sample_rate or (live_options.sample_rate if live_options else None)
default_options = LiveOptions( settings = DeepgramSageMakerSTTSettings(
encoding="linear16",
language=Language.EN,
model="nova-3", model="nova-3",
language=Language.EN,
encoding="linear16",
channels=1, channels=1,
interim_results=True, interim_results=True,
smart_format=False,
punctuate=True, punctuate=True,
profanity_filter=True,
vad_events=False,
diarize=False,
endpointing=None,
) )
settings = DeepgramSageMakerSTTSettings(
model=default_options.model,
language=default_options.language,
live_options=default_options,
)
if live_options: if live_options:
settings._merge_live_options_delta(live_options) lo_dict = live_options.to_dict()
delta = DeepgramSageMakerSTTSettings.from_mapping(
{k: v for k, v in lo_dict.items() if k != "sample_rate"}
)
settings.apply_update(delta)
super().__init__( super().__init__(
sample_rate=sample_rate, sample_rate=sample_rate,
@@ -224,9 +228,8 @@ class DeepgramSageMakerSTTService(STTService):
""" """
logger.debug("Connecting to Deepgram on SageMaker...") logger.debug("Connecting to Deepgram on SageMaker...")
live_options = LiveOptions( # Reconstruct a LiveOptions from the flat settings to build the query string.
**{**self._settings.live_options.to_dict(), "sample_rate": self.sample_rate} live_options = LiveOptions(**self._settings.given_fields())
)
# Build query string from live_options, converting booleans to strings # Build query string from live_options, converting booleans to strings
query_params = {} query_params = {}
@@ -237,6 +240,7 @@ class DeepgramSageMakerSTTService(STTService):
query_params[key] = str(value).lower() query_params[key] = str(value).lower()
else: else:
query_params[key] = str(value) query_params[key] = str(value)
query_params["sample_rate"] = str(self.sample_rate)
query_string = "&".join(f"{k}={v}" for k, v in query_params.items()) query_string = "&".join(f"{k}={v}" for k, v in query_params.items())