diff --git a/mellea/backends/openai.py b/mellea/backends/openai.py index 01ef2535e..7928cfe33 100644 --- a/mellea/backends/openai.py +++ b/mellea/backends/openai.py @@ -186,7 +186,9 @@ def __init__( # Use provided parameters or fall back to environment variables self._api_key = api_key - self._base_url = base_url + # Resolve env here (not only in the SDK) so _server_type / init logging + # see the same host the client will actually call. + self._base_url = base_url or os.getenv("OPENAI_BASE_URL") # Validate that we have the required configuration if self._api_key is None and os.getenv("OPENAI_API_KEY") is None: @@ -196,7 +198,7 @@ def __init__( " 2. Pass it as a parameter: OpenAIBackend(api_key='your-key-here')" ) - if self._base_url is None and os.getenv("OPENAI_BASE_URL") is None: + if self._base_url is None: MelleaLogger.get_logger().warning( "OPENAI_BASE_URL or base_url is not set.\n" "The openai SDK is going to assume that the base_url is `https://api.openai.com/v1`" @@ -207,6 +209,16 @@ def __init__( if self._base_url is not None else _ServerType.OPENAI ) # type: ignore + if self._server_type != _ServerType.OPENAI: + MelleaLogger.get_logger().info( + "Mellea assumes you are NOT using the OpenAI platform, and that " + "other model providers have less strict requirements on supporting " + "JSON schemas passed into `format=`. If you encounter a server-side " + "error when using format=, then you found an exception to this " + "assumption. Please open an issue at " + "github.com/generative-computing/mellea with the stack trace and " + "your inference engine / model provider." + ) self._openai_client_kwargs = self.filter_openai_client_kwargs(**kwargs) @@ -934,9 +946,6 @@ async def _generate_from_chat_context_standard( }, } else: - MelleaLogger.get_logger().info( - "Mellea assumes you are NOT using the OpenAI platform, and that other model providers have less strict requirements on supporting JSON schemas passed into `format=`. If you encounter a server-side error following this message, then you found an exception to this assumption. Please open an issue at github.com/generative_computing/mellea with this stack trace and your inference engine / model provider." - ) extra_params["response_format"] = { "type": "json_schema", "json_schema": { diff --git a/test/backends/test_openai_unit.py b/test/backends/test_openai_unit.py index a40db14c5..950ba5db1 100644 --- a/test/backends/test_openai_unit.py +++ b/test/backends/test_openai_unit.py @@ -7,6 +7,7 @@ _simplify_and_merge, and _make_backend_specific_and_remove. """ +import os from unittest.mock import AsyncMock, MagicMock, PropertyMock, patch import pytest @@ -429,5 +430,105 @@ class Answer(pydantic.BaseModel): assert "guided_json" in extra_body or "structured_outputs" in extra_body +# --- #1502: non-OpenAI format= warning only at init --- + +_FORMAT_ASSUMPTION = "NOT using the OpenAI platform" + + +def _info_msgs(mock_logger) -> list[str]: + return [str(c.args[0]) for c in mock_logger.info.call_args_list if c.args] + + +def test_non_openai_format_assumption_logged_once_at_init(): + mock_logger = MagicMock() + with patch( + "mellea.backends.openai.MelleaLogger.get_logger", return_value=mock_logger + ): + OpenAIBackend( + model_id="gpt-4o", api_key="fake-key", base_url="http://localhost:9999/v1" + ) + OpenAIBackend( + model_id="gpt-4o", api_key="fake-key", base_url="http://localhost:9999/v1" + ) + + matches = [m for m in _info_msgs(mock_logger) if _FORMAT_ASSUMPTION in m] + assert len(matches) == 2 # once per backend instance + + +def test_openai_platform_skips_format_assumption_log(): + mock_logger = MagicMock() + # Unset OPENAI_BASE_URL for this test: after resolving env into _base_url, + # a leftover non-OpenAI env would make the no-base_url construction log. + with ( + patch( + "mellea.backends.openai.MelleaLogger.get_logger", return_value=mock_logger + ), + patch.dict(os.environ), + ): + os.environ.pop("OPENAI_BASE_URL", None) + OpenAIBackend( + model_id="gpt-4o", api_key="fake-key", base_url="https://api.openai.com/v1" + ) + OpenAIBackend(model_id="gpt-4o", api_key="fake-key") + + matches = [m for m in _info_msgs(mock_logger) if _FORMAT_ASSUMPTION in m] + assert matches == [] + + +def test_format_assumption_log_honors_openai_base_url_env(): + """Env-only non-OpenAI base_url must still classify as non-OpenAI at init.""" + mock_logger = MagicMock() + with ( + patch( + "mellea.backends.openai.MelleaLogger.get_logger", return_value=mock_logger + ), + patch.dict( + os.environ, {"OPENAI_BASE_URL": "http://localhost:9999/v1"}, clear=False + ), + ): + OpenAIBackend(model_id="gpt-4o", api_key="fake-key") + + matches = [m for m in _info_msgs(mock_logger) if _FORMAT_ASSUMPTION in m] + assert len(matches) == 1 + + +async def test_format_assumption_not_relogged_per_generation(): + """#1502: the notice must not repeat on every format= generation.""" + import pydantic + + from mellea.core.base import CBlock + from mellea.stdlib.context import ChatContext + + class Answer(pydantic.BaseModel): + value: int + + backend = OpenAIBackend( + model_id="gpt-4o", api_key="fake-key", base_url="http://localhost:9999/v1" + ) + ctx = ChatContext().add(CBlock(value="q")) + resp = MagicMock() + resp.choices = [MagicMock()] + resp.choices[0].message.content = "{}" + resp.choices[0].message.role = "assistant" + + mock_logger = MagicMock() + with ( + patch( + "mellea.backends.openai.MelleaLogger.get_logger", return_value=mock_logger + ), + patch.object( + backend._async_client.chat.completions, "create", new_callable=AsyncMock + ) as create, + ): + create.return_value = resp + for _ in range(3): + await backend.generate_from_chat_context( + CBlock(value="q"), ctx, _format=Answer, model_options={} + ) + + msgs = [str(c.args[0]) for c in mock_logger.info.call_args_list if c.args] + assert [m for m in msgs if _FORMAT_ASSUMPTION in m] == [] + + if __name__ == "__main__": pytest.main([__file__, "-v"])