Skip to content
Closed
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
20 changes: 17 additions & 3 deletions plugins/image_gen/xai/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,8 +100,22 @@ def _load_xai_config() -> Dict[str, Any]:
return {}


def _resolve_model() -> Tuple[str, Dict[str, Any]]:
"""Decide which model to use and return ``(model_id, meta)``."""
def _resolve_model(caller_model: Optional[str] = None) -> Tuple[str, Dict[str, Any]]:
"""Decide which model to use and return ``(model_id, meta)``.

Priority:
1. Caller-supplied ``caller_model`` (from ``generate(**kwargs)`` via
``image_gen.model`` config key) — mirrors how the openrouter provider
threads ``kwargs.get("model")`` into its own resolver.
2. ``XAI_IMAGE_MODEL`` env override.
3. ``image_gen.model`` in config.yaml (already read by the caller into
``caller_model``; this branch handles the legacy path where config is
read directly here rather than forwarded by the caller).
4. Hard-coded default.
"""
if caller_model and caller_model in _MODELS:
return caller_model, _MODELS[caller_model]

env_override = os.environ.get("XAI_IMAGE_MODEL")
if env_override and env_override in _MODELS:
return env_override, _MODELS[env_override]
Expand Down Expand Up @@ -234,7 +248,7 @@ def generate(
aspect_ratio=aspect_ratio,
)

model_id, meta = _resolve_model()
model_id, meta = _resolve_model(kwargs.get("model"))
aspect = resolve_aspect_ratio(aspect_ratio)
xai_ar = _XAI_ASPECT_RATIOS.get(aspect, "1:1")
resolution = _resolve_resolution()
Expand Down
37 changes: 37 additions & 0 deletions tests/plugins/image_gen/test_xai_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,43 @@ def test_custom_model(self, monkeypatch):
model_id, _ = _resolve_model()
assert model_id == "grok-imagine-image"

def test_caller_model_overrides_env(self, monkeypatch):
"""caller_model (from image_gen.model config key) must take priority
over XAI_IMAGE_MODEL env — mirrors the fix applied to the openrouter
provider in #55672."""
monkeypatch.setenv("XAI_IMAGE_MODEL", "grok-imagine-image")
from plugins.image_gen.xai import _resolve_model

model_id, _ = _resolve_model("grok-imagine-image-quality")
assert model_id == "grok-imagine-image-quality"

def test_unknown_caller_model_falls_back_to_env(self, monkeypatch):
"""An unrecognised caller_model must not crash — fall through to env."""
monkeypatch.setenv("XAI_IMAGE_MODEL", "grok-imagine-image")
from plugins.image_gen.xai import _resolve_model

model_id, _ = _resolve_model("not-a-real-model")
assert model_id == "grok-imagine-image"

def test_model_kwarg_forwarded_to_generate(self):
"""generate(model=...) must use the supplied model, not the default."""
from plugins.image_gen.xai import XAIImageGenProvider

mock_resp = MagicMock()
mock_resp.status_code = 200
mock_resp.raise_for_status = MagicMock()
mock_resp.json.return_value = {"data": [{"b64_json": "dGVzdA=="}]}

with patch("plugins.image_gen.xai.requests.post", return_value=mock_resp) as mock_post:
with patch("plugins.image_gen.xai.save_b64_image", return_value="/tmp/out.png"):
provider = XAIImageGenProvider()
result = provider.generate(prompt="test", model="grok-imagine-image-quality")

assert result["success"] is True
assert result["model"] == "grok-imagine-image-quality"
payload = mock_post.call_args.kwargs.get("json") or mock_post.call_args[1].get("json", {})
assert payload.get("model") == "grok-imagine-image-quality"


# ---------------------------------------------------------------------------
# Generate tests
Expand Down
Loading