diff --git a/src/agents/models/openai_provider.py b/src/agents/models/openai_provider.py index 59a04071..dd4b888c 100644 --- a/src/agents/models/openai_provider.py +++ b/src/agents/models/openai_provider.py @@ -7,6 +7,7 @@ import weakref import httpx from openai import AsyncOpenAI, DefaultAsyncHttpxClient +from ..exceptions import UserError from . import _openai_shared from .default_models import get_default_model from .interface import Model, ModelProvider @@ -86,10 +87,11 @@ class OpenAIProvider(ModelProvider): chunk semantics are not reliable enough for incremental processing. """ if openai_client is not None: - assert api_key is None and base_url is None and websocket_base_url is None, ( - "Don't provide api_key, base_url, or websocket_base_url if you provide " - "openai_client" - ) + if api_key is not None or base_url is not None or websocket_base_url is not None: + raise UserError( + "Don't provide api_key, base_url, or websocket_base_url if you provide " + "openai_client" + ) self._client: AsyncOpenAI | None = openai_client else: self._client = None diff --git a/src/agents/voice/models/openai_model_provider.py b/src/agents/voice/models/openai_model_provider.py index 31482570..b992f9b4 100644 --- a/src/agents/voice/models/openai_model_provider.py +++ b/src/agents/voice/models/openai_model_provider.py @@ -3,6 +3,7 @@ from __future__ import annotations import httpx from openai import AsyncOpenAI, DefaultAsyncHttpxClient +from ...exceptions import UserError from ...models import _openai_shared from ...models.openai_agent_registration import ( OpenAIAgentRegistrationConfig, @@ -56,9 +57,8 @@ class OpenAIVoiceModelProvider(VoiceModelProvider): agent_registration: Optional agent registration configuration. """ if openai_client is not None: - assert api_key is None and base_url is None, ( - "Don't provide api_key or base_url if you provide openai_client" - ) + if api_key is not None or base_url is not None: + raise UserError("Don't provide api_key or base_url if you provide openai_client") self._client: AsyncOpenAI | None = openai_client else: self._client = None diff --git a/tests/test_config.py b/tests/test_config.py index debb6c16..0eefc367 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -7,6 +7,7 @@ import openai import pytest from agents import ( + UserError, responses_websocket_session, set_default_openai_api, set_default_openai_client, @@ -97,6 +98,27 @@ def test_set_default_openai_responses_transport_rejects_invalid_value(): set_default_openai_responses_transport("ws") # type: ignore[arg-type] +@pytest.mark.parametrize( + "conflicting_kwargs", + [ + {"api_key": "other_key"}, + {"base_url": "https://example.com"}, + {"websocket_base_url": "wss://example.com"}, + { + "api_key": "other_key", + "base_url": "https://example.com", + "websocket_base_url": "wss://example.com", + }, + ], +) +def test_openai_provider_rejects_client_with_conflicting_args(conflicting_kwargs): + # Regression test for #3808: this validation used a bare `assert`, which is + # stripped under `python -O`, silently ignoring the conflicting arguments. + client = openai.AsyncOpenAI(api_key="test_key") + with pytest.raises(UserError, match="Don't provide"): + OpenAIProvider(openai_client=client, **conflicting_kwargs) + + def test_openai_provider_transport_override_beats_default(): set_default_openai_api("responses") set_default_openai_responses_transport("websocket") diff --git a/tests/voice/test_openai_model_provider.py b/tests/voice/test_openai_model_provider.py new file mode 100644 index 00000000..2d9de3ae --- /dev/null +++ b/tests/voice/test_openai_model_provider.py @@ -0,0 +1,29 @@ +# Tests for the OpenAI voice model provider (OpenAIVoiceModelProvider). + +import openai +import pytest + +from agents.exceptions import UserError +from agents.voice.models.openai_model_provider import OpenAIVoiceModelProvider + + +@pytest.mark.parametrize( + "conflicting_kwargs", + [ + {"api_key": "other_key"}, + {"base_url": "https://example.com"}, + {"api_key": "other_key", "base_url": "https://example.com"}, + ], +) +def test_voice_provider_rejects_client_with_conflicting_args(conflicting_kwargs): + # Regression test for #3808: this validation used a bare `assert`, which is + # stripped under `python -O`, silently ignoring the conflicting arguments. + client = openai.AsyncOpenAI(api_key="test_key") + with pytest.raises(UserError, match="Don't provide"): + OpenAIVoiceModelProvider(openai_client=client, **conflicting_kwargs) + + +def test_voice_provider_accepts_client_without_conflicting_args(): + client = openai.AsyncOpenAI(api_key="test_key") + provider = OpenAIVoiceModelProvider(openai_client=client) + assert provider._get_client() is client