diff --git a/src/agents/voice/models/openai_stt.py b/src/agents/voice/models/openai_stt.py index cf504d8892..5d098507f6 100644 --- a/src/agents/voice/models/openai_stt.py +++ b/src/agents/voice/models/openai_stt.py @@ -175,6 +175,15 @@ async def _event_listener(self) -> None: async def _configure_session(self) -> None: assert self._websocket is not None, "Websocket not initialized" + transcription_config: dict[str, Any] = {"model": self._model} + if self._settings.language is not None: + if self._model == "gpt-live-transcribe": + transcription_config["languages"] = [self._settings.language] + else: + transcription_config["language"] = self._settings.language + if self._settings.prompt is not None: + transcription_config["prompt"] = self._settings.prompt + await self._websocket.send( json.dumps( { @@ -184,7 +193,7 @@ async def _configure_session(self) -> None: "audio": { "input": { "format": {"type": "audio/pcm", "rate": 24000}, - "transcription": {"model": self._model}, + "transcription": transcription_config, "turn_detection": self._turn_detection, } }, diff --git a/tests/voice/test_openai_stt_session_config.py b/tests/voice/test_openai_stt_session_config.py new file mode 100644 index 0000000000..2de1114f0a --- /dev/null +++ b/tests/voice/test_openai_stt_session_config.py @@ -0,0 +1,97 @@ +import json +from unittest.mock import AsyncMock + +import pytest + +from agents.voice import StreamedAudioInput, STTModelSettings +from agents.voice.models.openai_stt import OpenAISTTTranscriptionSession + + +@pytest.mark.asyncio +async def test_streaming_stt_sends_language_and_prompt() -> None: + session = OpenAISTTTranscriptionSession( + input=StreamedAudioInput(), + client=AsyncMock(api_key="FAKE_KEY"), + model="gpt-4o-transcribe", + settings=STTModelSettings(language="fr", prompt="domain vocabulary"), + trace_include_sensitive_data=False, + trace_include_sensitive_audio_data=False, + ) + websocket = AsyncMock() + session._websocket = websocket + + await session._configure_session() + + payload = json.loads(websocket.send.await_args.args[0]) + assert payload["session"]["audio"]["input"]["transcription"] == { + "model": "gpt-4o-transcribe", + "language": "fr", + "prompt": "domain vocabulary", + } + + +@pytest.mark.asyncio +async def test_streaming_stt_sends_singular_language_for_gpt_transcribe() -> None: + session = OpenAISTTTranscriptionSession( + input=StreamedAudioInput(), + client=AsyncMock(api_key="FAKE_KEY"), + model="gpt-transcribe", + settings=STTModelSettings(language="fr", prompt="domain vocabulary"), + trace_include_sensitive_data=False, + trace_include_sensitive_audio_data=False, + ) + websocket = AsyncMock() + session._websocket = websocket + + await session._configure_session() + + payload = json.loads(websocket.send.await_args.args[0]) + assert payload["session"]["audio"]["input"]["transcription"] == { + "model": "gpt-transcribe", + "language": "fr", + "prompt": "domain vocabulary", + } + + +@pytest.mark.asyncio +async def test_streaming_stt_sends_plural_languages_for_gpt_live_transcribe() -> None: + session = OpenAISTTTranscriptionSession( + input=StreamedAudioInput(), + client=AsyncMock(api_key="FAKE_KEY"), + model="gpt-live-transcribe", + settings=STTModelSettings(language="fr", prompt="domain vocabulary"), + trace_include_sensitive_data=False, + trace_include_sensitive_audio_data=False, + ) + websocket = AsyncMock() + session._websocket = websocket + + await session._configure_session() + + payload = json.loads(websocket.send.await_args.args[0]) + assert payload["session"]["audio"]["input"]["transcription"] == { + "model": "gpt-live-transcribe", + "languages": ["fr"], + "prompt": "domain vocabulary", + } + + +@pytest.mark.asyncio +async def test_streaming_stt_omits_unset_language_and_prompt() -> None: + session = OpenAISTTTranscriptionSession( + input=StreamedAudioInput(), + client=AsyncMock(api_key="FAKE_KEY"), + model="gpt-4o-transcribe", + settings=STTModelSettings(), + trace_include_sensitive_data=False, + trace_include_sensitive_audio_data=False, + ) + websocket = AsyncMock() + session._websocket = websocket + + await session._configure_session() + + payload = json.loads(websocket.send.await_args.args[0]) + assert payload["session"]["audio"]["input"]["transcription"] == { + "model": "gpt-4o-transcribe" + }