fix(models): raise UserError for conflicting provider args instead of assert (#3809)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user