diff --git a/nanobot/providers/openai_compat_provider.py b/nanobot/providers/openai_compat_provider.py index 9a7da7522..83db7e8f8 100644 --- a/nanobot/providers/openai_compat_provider.py +++ b/nanobot/providers/openai_compat_provider.py @@ -157,6 +157,16 @@ def _is_direct_openai_base(api_base: str | None) -> bool: return "api.openai.com" in normalized and "openrouter" not in normalized +def _responses_circuit_key( + model: str | None, + default_model: str, + reasoning_effort: str | None, +) -> str: + model_name = (model or default_model).lower() + effort = reasoning_effort.lower() if isinstance(reasoning_effort, str) else "" + return f"{model_name}:{effort}" + + class OpenAICompatProvider(LLMProvider): """Unified provider for all OpenAI-compatible APIs. @@ -434,7 +444,7 @@ class OpenAICompatProvider(LLMProvider): return False # Circuit breaker: skip after repeated failures, probe periodically. - key = f"{model_name}:{reasoning_effort or ''}" + key = _responses_circuit_key(model, self.default_model, reasoning_effort) failures = self._responses_failures.get(key, 0) if failures >= _RESPONSES_FAILURE_THRESHOLD: tripped = self._responses_tripped_at.get(key, 0.0) @@ -444,7 +454,7 @@ class OpenAICompatProvider(LLMProvider): return True def _record_responses_failure(self, model: str | None, reasoning_effort: str | None) -> None: - key = f"{(model or self.default_model).lower()}:{reasoning_effort or ''}" + key = _responses_circuit_key(model, self.default_model, reasoning_effort) count = self._responses_failures.get(key, 0) + 1 self._responses_failures[key] = count if count >= _RESPONSES_FAILURE_THRESHOLD: @@ -455,7 +465,7 @@ class OpenAICompatProvider(LLMProvider): ) def _record_responses_success(self, model: str | None, reasoning_effort: str | None) -> None: - key = f"{(model or self.default_model).lower()}:{reasoning_effort or ''}" + key = _responses_circuit_key(model, self.default_model, reasoning_effort) self._responses_failures.pop(key, None) self._responses_tripped_at.pop(key, None) diff --git a/tests/providers/test_litellm_kwargs.py b/tests/providers/test_litellm_kwargs.py index 8304aae8f..47db20398 100644 --- a/tests/providers/test_litellm_kwargs.py +++ b/tests/providers/test_litellm_kwargs.py @@ -441,6 +441,35 @@ async def test_direct_openai_responses_404_falls_back_to_chat_completions() -> N mock_chat.assert_awaited_once() +@pytest.mark.asyncio +async def test_direct_openai_open_circuit_skips_responses_api() -> None: + mock_chat = AsyncMock(return_value=_fake_chat_response("from chat")) + mock_responses = AsyncMock(return_value=_fake_responses_response("from responses")) + spec = find_by_name("openai") + + with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as MockClient: + client_instance = MockClient.return_value + client_instance.chat.completions.create = mock_chat + client_instance.responses.create = mock_responses + + provider = OpenAICompatProvider( + api_key="sk-test-key", + default_model="gpt-5-chat", + spec=spec, + ) + for _ in range(3): + provider._record_responses_failure("gpt-5-chat", None) + + result = await provider.chat( + messages=[{"role": "user", "content": "hello"}], + model="gpt-5-chat", + ) + + assert result.content == "from chat" + mock_responses.assert_not_awaited() + mock_chat.assert_awaited_once() + + @pytest.mark.asyncio async def test_direct_openai_stream_responses_unsupported_param_falls_back() -> None: mock_chat = AsyncMock(return_value=_fake_chat_stream("fallback stream")) diff --git a/tests/providers/test_responses_circuit_breaker.py b/tests/providers/test_responses_circuit_breaker.py index 4787459c7..409aea1d5 100644 --- a/tests/providers/test_responses_circuit_breaker.py +++ b/tests/providers/test_responses_circuit_breaker.py @@ -69,3 +69,9 @@ def test_reasoning_effort_keyed_separately(provider): provider._record_responses_failure("o3", "high") assert provider._should_use_responses_api("o3", "high") is False assert provider._should_use_responses_api("o3", "low") is True + + +def test_reasoning_effort_key_is_case_insensitive(provider): + for _ in range(_RESPONSES_FAILURE_THRESHOLD): + provider._record_responses_failure("o3", "High") + assert provider._should_use_responses_api("o3", "high") is False