Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 14 additions & 5 deletions mellea/backends/openai.py

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Both of these sit outside the diff, hence the file-level comment.

Since mellea now resolves OPENAI_BASE_URL itself, _base_url gets filled in where it used to be None. Two knock-ons:

  • The docstring for base_url (line 79) doesn't mention the env fallback, though the api_key entry at line 92 documents its own. Worth a matching line.
  • __repr__ (line 256) prints base_url as-is, and it's the one method that goes out of its way to mask the api key. If someone's OPENAI_BASE_URL has credentials in it (https://user:token@host/v1), that used to print as None and now prints in full. Nothing in mellea logs a backend, so it only surfaces if you print one yourself — your call whether that's worth masking.

Original file line number Diff line number Diff line change
Expand Up @@ -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")
Comment thread
planetf1 marked this conversation as resolved.

# Validate that we have the required configuration
if self._api_key is None and os.getenv("OPENAI_API_KEY") is None:
Expand All @@ -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`"
Expand All @@ -207,6 +209,16 @@ def __init__(
if self._base_url is not None
else _ServerType.OPENAI
) # type: ignore
if self._server_type != _ServerType.OPENAI:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Flagging in case another reviewer raises this: moving the log to fire on server type alone, regardless of whether format= is ever used, isn't scope creep — issue #1502's own "Proposed fix, Option 1" asks for exactly this ("tie the message to backend setup rather than the generate loop"). This matches the issue as written.

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)

Expand Down Expand Up @@ -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": {
Expand Down
101 changes: 101 additions & 0 deletions test/backends/test_openai_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"
)
Comment thread
planetf1 marked this conversation as resolved.

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),
):
Comment on lines +462 to +467

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
with (
patch(
"mellea.backends.openai.MelleaLogger.get_logger", return_value=mock_logger
),
patch.dict(os.environ),
):
# These backends point at api.openai.com, so mock the vLLM version probe:
# __init__ calls is_vllm_server_with_structured_output unconditionally,
# which would otherwise make a real GET to api.openai.com/version.
with (
patch(
"mellea.backends.openai.MelleaLogger.get_logger", return_value=mock_logger
),
patch(
"mellea.backends.openai.is_vllm_server_with_structured_output",
return_value=False,
),
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"
)
Comment thread
planetf1 marked this conversation as resolved.
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():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Between them these cover LOCALHOST and OPENAI, but nothing lands on _ServerType.UNKNOWN — any hosted endpoint, which is probably the commonest non-OpenAI case in the wild. bedrock.py:89 builds exactly that (https://bedrock-mantle.{region}.api.aws/v1) and hands it to OpenAIBackend. Same branch by inspection, so I'm not expecting a surprise; it's just the one classification outcome nothing pins. Fine to leave for later.

"""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
),
Comment thread
planetf1 marked this conversation as resolved.
):
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
Comment thread
planetf1 marked this conversation as resolved.


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"])
Loading