fix(image_gen/meta-ai): honor dispatcher model kwarg, standard paid badge

- generate() now passes kwargs.get("model") into _resolve_model(), so the
  user's hermes tools pick (forwarded by the dispatcher as top-level
  image_gen.model) is honored instead of silently dropped (#55893 class;
  matches xai/krea/openrouter).
- Setup schema badge "internal" -> "paid" to match every other paid
  image backend in the hermes tools picker.
- Tests: caller-model precedence, unknown caller model falls through,
  model kwarg reaches the API payload, badge contract.
This commit is contained in:
Teknium
2026-09-01 22:09:48 -07:00
parent d7e92ab7e3
commit eb4d77c2a8
2 changed files with 87 additions and 23 deletions

View File

@@ -8,10 +8,11 @@ the OpenAI Python SDK pointed at Meta's base URL and authenticate with
Output is base64 JSON (WebP) -> saved under ``$HERMES_HOME/cache/images/``.
Selection precedence (first hit wins):
1. ``META_IMAGE_MODEL`` env var (escape hatch for scripts / tests)
2. ``image_gen.meta-ai.model`` in ``config.yaml``
3. ``image_gen.model`` in ``config.yaml`` (when it's one of our IDs)
4. :data:`DEFAULT_MODEL`
1. ``model`` kwarg forwarded by the dispatcher (the ``hermes tools`` pick)
2. ``META_IMAGE_MODEL`` env var (escape hatch for scripts / tests)
3. ``image_gen.meta-ai.model`` in ``config.yaml``
4. ``image_gen.model`` in ``config.yaml`` (when it's one of our IDs)
5. :data:`DEFAULT_MODEL`
"""
from __future__ import annotations
@@ -81,8 +82,17 @@ _SIZES: Dict[str, str] = {
}
def _resolve_model() -> Tuple[str, Dict[str, Any]]:
"""Return (model_id, metadata) using the documented precedence chain."""
def _resolve_model(caller_model: Optional[str] = None) -> Tuple[str, Dict[str, Any]]:
"""Return (model_id, metadata) using the documented precedence chain.
``caller_model`` is the ``model`` kwarg the dispatcher forwards from the
top-level ``image_gen.model`` config key (what ``hermes tools`` writes).
It wins when it names one of our models, mirroring the xai/krea/openrouter
providers, so a user's picker choice is never silently dropped.
"""
if caller_model and caller_model in _MODELS:
return caller_model, _MODELS[caller_model]
env_model = os.environ.get("META_IMAGE_MODEL")
if env_model and env_model in _MODELS:
return env_model, _MODELS[env_model]
@@ -142,7 +152,7 @@ class MetaImageGenProvider(ImageGenProvider):
def get_setup_schema(self) -> Dict[str, Any]:
return {
"name": "Meta Model API",
"badge": "internal",
"badge": "paid",
"tag": "Muse Image via Meta Model API (api.meta.ai)",
"env_vars": [
{
@@ -200,7 +210,7 @@ class MetaImageGenProvider(ImageGenProvider):
aspect_ratio=aspect,
)
model_id, _meta = _resolve_model()
model_id, _meta = _resolve_model(kwargs.get("model"))
size = _SIZES.get(aspect, _SIZES["square"])
client = openai.OpenAI(api_key=api_key, base_url=_resolve_base_url())

View File

@@ -25,6 +25,7 @@ _PNG_HEX = (
def _b64_png() -> str:
import base64
return base64.b64encode(bytes.fromhex(_PNG_HEX)).decode()
@@ -37,8 +38,13 @@ def _fake_response(*, b64=None, url=None, revised_prompt=None):
def _tmp_hermes_home(tmp_path, monkeypatch):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
# Clear every auth + override env var so tests start from a clean slate.
for env in ("MODEL_API_KEY", "META_API_KEY", "META_MODEL_API_KEY",
"META_BASE_URL", "META_IMAGE_MODEL"):
for env in (
"MODEL_API_KEY",
"META_API_KEY",
"META_MODEL_API_KEY",
"META_BASE_URL",
"META_IMAGE_MODEL",
):
monkeypatch.delenv(env, raising=False)
yield tmp_path
@@ -133,11 +139,45 @@ class TestModelResolution:
# Unknown id is ignored; falls through to the default.
assert model_id == "muse-image-1.0"
def test_caller_model_kwarg_wins(self, monkeypatch):
# The dispatcher forwards top-level image_gen.model as the `model`
# kwarg; it must beat the env override (#55893 bug class).
monkeypatch.setitem(
meta_plugin._MODELS,
"muse-image-test",
dict(meta_plugin._MODELS["muse-image-1.0"]),
)
monkeypatch.setenv("META_IMAGE_MODEL", "muse-image-1.0")
model_id, _meta = meta_plugin._resolve_model("muse-image-test")
assert model_id == "muse-image-test"
def test_caller_model_unknown_falls_through(self):
model_id, _meta = meta_plugin._resolve_model("not-a-real-model")
assert model_id == "muse-image-1.0"
# ── Generate ──────────────────────────────────────────────────────────────────
class TestGenerate:
def test_model_kwarg_reaches_payload(self, provider, monkeypatch):
monkeypatch.setitem(
meta_plugin._MODELS,
"muse-image-test",
dict(meta_plugin._MODELS["muse-image-1.0"]),
)
fake_client = MagicMock()
fake_client.images.generate.return_value = _fake_response(b64=_b64_png())
with _patched_openai(fake_client):
result = provider.generate("a cat", model="muse-image-test")
assert result["success"] is True
assert (
fake_client.images.generate.call_args.kwargs["model"] == "muse-image-test"
)
def test_badge_is_standard_paid(self, provider):
assert provider.get_setup_schema()["badge"] == "paid"
def test_empty_prompt_rejected(self, provider):
result = provider.generate("", aspect_ratio="square")
assert result["success"] is False
@@ -182,7 +222,9 @@ class TestGenerate:
with patch.dict("sys.modules", {"openai": fake_openai}):
provider.generate("a cat")
assert fake_openai.OpenAI.call_args.kwargs["base_url"] == "https://api.meta.ai/v1"
assert (
fake_openai.OpenAI.call_args.kwargs["base_url"] == "https://api.meta.ai/v1"
)
def test_base_url_override_reaches_client(self, provider, monkeypatch):
monkeypatch.setenv("META_BASE_URL", "https://proxy.internal/v1")
@@ -194,13 +236,19 @@ class TestGenerate:
with patch.dict("sys.modules", {"openai": fake_openai}):
provider.generate("a cat")
assert fake_openai.OpenAI.call_args.kwargs["base_url"] == "https://proxy.internal/v1"
assert (
fake_openai.OpenAI.call_args.kwargs["base_url"]
== "https://proxy.internal/v1"
)
@pytest.mark.parametrize("aspect,expected_size", [
("landscape", "1536x1024"),
("square", "1024x1024"),
("portrait", "1024x1536"),
])
@pytest.mark.parametrize(
"aspect,expected_size",
[
("landscape", "1536x1024"),
("square", "1024x1024"),
("portrait", "1024x1536"),
],
)
def test_aspect_ratio_mapping(self, provider, aspect, expected_size):
fake_client = MagicMock()
fake_client.images.generate.return_value = _fake_response(b64=_b64_png())
@@ -213,7 +261,8 @@ class TestGenerate:
def test_revised_prompt_passed_through(self, provider):
fake_client = MagicMock()
fake_client.images.generate.return_value = _fake_response(
b64=_b64_png(), revised_prompt="A photo of a cat",
b64=_b64_png(),
revised_prompt="A photo of a cat",
)
with _patched_openai(fake_client):
@@ -226,13 +275,18 @@ class TestGenerate:
providers) so ephemeral signed URLs can't expire mid-flight."""
fake_client = MagicMock()
fake_client.images.generate.return_value = _fake_response(
b64=None, url="https://example.com/img.webp",
b64=None,
url="https://example.com/img.webp",
)
with _patched_openai(fake_client), patch.object(
meta_plugin, "save_url_image",
return_value=Path("/tmp/meta_20260524_000000_deadbeef.webp"),
) as mock_save_url:
with (
_patched_openai(fake_client),
patch.object(
meta_plugin,
"save_url_image",
return_value=Path("/tmp/meta_20260524_000000_deadbeef.webp"),
) as mock_save_url,
):
result = provider.generate("a cat")
assert result["success"] is True