diff --git a/src/agents/models/openai_provider.py b/src/agents/models/openai_provider.py index f46fdc55ca..ee1c6e9367 100644 --- a/src/agents/models/openai_provider.py +++ b/src/agents/models/openai_provider.py @@ -259,8 +259,8 @@ def get_model(self, model_name: str | None) -> Model: if running_loop is not None else None ) - if loop_cache is not None and (cached_model := loop_cache.get(cache_key)): - return cached_model + if loop_cache is not None and cache_key in loop_cache: + return loop_cache[cache_key] client = self._get_client() model: Model diff --git a/tests/test_config.py b/tests/test_config.py index 845bf213b5..9eaccab587 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -298,6 +298,38 @@ async def get_model(): assert model2 is not model1 +def test_openai_provider_reuses_falsy_websocket_model(monkeypatch): + class DummyAsyncOpenAI: + pass + + class FalsyWebsocketModel: + def __init__(self, **kwargs): + self.kwargs = kwargs + + def __bool__(self) -> bool: + return False + + monkeypatch.setattr("agents.models.openai_provider.OpenAIResponsesWSModel", FalsyWebsocketModel) + provider = OpenAIProvider( + use_responses=True, + use_responses_websocket=True, + openai_client=DummyAsyncOpenAI(), # type: ignore[arg-type] + ) + + async def get_model(): + return provider.get_model("gpt-4") + + loop = asyncio.new_event_loop() + try: + model = loop.run_until_complete(get_model()) + model_again = loop.run_until_complete(get_model()) + finally: + loop.close() + asyncio.set_event_loop(None) + + assert model is model_again + + def test_openai_provider_websocket_loop_cache_does_not_keep_closed_loop_alive(monkeypatch): class DummyAsyncOpenAI: pass