From 1616fa9f14e1a101e93129a3f3cff7875bfb0fcc Mon Sep 17 00:00:00 2001 From: Xubin Ren <52506698+Re-bin@users.noreply.github.com> Date: Fri, 17 Jul 2026 13:02:49 +0800 Subject: [PATCH] feat(image): apply generation settings live --- nanobot/agent/context.py | 9 +- nanobot/agent/tools/image_generation.py | 121 ++++++++++++++++++ nanobot/bus/events.py | 1 + .../websocket/tests/test_websocket_channel.py | 112 +++++++++++++++- nanobot/providers/image_generation.py | 10 ++ nanobot/webui/settings_api.py | 7 + nanobot/webui/settings_routes.py | 36 +++++- .../agent/test_image_generation_hot_reload.py | 103 +++++++++++++++ .../src/components/settings/SettingsView.tsx | 47 ++++++- webui/src/lib/types.ts | 2 + webui/src/tests/settings-view.test.tsx | 51 ++++++++ 11 files changed, 481 insertions(+), 18 deletions(-) create mode 100644 tests/agent/test_image_generation_hot_reload.py diff --git a/nanobot/agent/context.py b/nanobot/agent/context.py index 85c90556..5489c4f7 100644 --- a/nanobot/agent/context.py +++ b/nanobot/agent/context.py @@ -8,6 +8,7 @@ from typing import Any, Mapping, Sequence from nanobot.agent.memory import MemoryStore from nanobot.agent.skills import SkillsLoader +from nanobot.agent.tools import image_generation as image_generation_tools from nanobot.agent.tools import mcp as mcp_tools from nanobot.agent.tools.registry import ToolRegistry from nanobot.apps.cli import utils as cli_app_utils @@ -41,7 +42,13 @@ async def close_mcp(state: Any) -> None: async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool: - return await mcp_tools.handle_runtime_control(state, msg, tools) + for handler in ( + image_generation_tools.handle_runtime_control, + mcp_tools.handle_runtime_control, + ): + if await handler(state, msg, tools): + return True + return False class ContextBuilder: diff --git a/nanobot/agent/tools/image_generation.py b/nanobot/agent/tools/image_generation.py index f70e9429..b1e448e7 100644 --- a/nanobot/agent/tools/image_generation.py +++ b/nanobot/agent/tools/image_generation.py @@ -2,24 +2,34 @@ from __future__ import annotations +import asyncio from pathlib import Path from typing import TYPE_CHECKING, Any +from loguru import logger from pydantic import Field from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters +from nanobot.agent.tools.registry import ToolRegistry from nanobot.agent.tools.schema import ( ArraySchema, IntegerSchema, StringSchema, tool_parameters_schema, ) +from nanobot.bus.events import ( + INBOUND_META_RUNTIME_CONTROL, + RUNTIME_CONTROL_ACK, + RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD, + InboundMessage, +) from nanobot.config.paths import get_media_dir from nanobot.config_base import Base from nanobot.providers.image_generation import ( ImageGenerationError, ImageGenerationProvider, get_image_gen_provider, + image_gen_provider_configs, ) from nanobot.security.workspace_access import current_tool_workspace from nanobot.security.workspace_policy import WorkspaceBoundaryError, resolve_allowed_path @@ -208,3 +218,114 @@ class ImageGenerationTool(Tool): return generated_image_tool_result(artifacts) except (ArtifactError, ImageGenerationError, OSError) as exc: return ToolResult.error(f"Error: {exc}") + + +async def reload_image_generation_tool(state: Any, registry: ToolRegistry) -> dict[str, Any]: + """Apply the persisted image configuration to the running agent.""" + try: + from nanobot.config.loader import load_config, resolve_config_env_vars + + config = resolve_config_env_vars(load_config()) + tool_config = config.tools.image_generation + provider_configs = image_gen_provider_configs(config) + except Exception as exc: + logger.warning("Image generation hot reload could not read config: {}", exc) + return { + "ok": False, + "message": "Could not reload image generation config.", + "requires_restart": True, + "error": str(exc), + } + + next_tool = ( + ImageGenerationTool( + workspace=state.workspace, + config=tool_config, + provider_configs=provider_configs, + ) + if tool_config.enabled + else None + ) + + state.tools_config.image_generation = tool_config + state._image_generation_provider_configs = provider_configs + if next_tool is not None: + registry.register(next_tool) + else: + registry.unregister("generate_image") + + logger.info( + "Image generation config reloaded: enabled={} provider={} model={}", + tool_config.enabled, + tool_config.provider, + tool_config.model, + ) + return { + "ok": True, + "message": "Image generation settings applied without restarting nanobot.", + "enabled": tool_config.enabled, + "provider": tool_config.provider, + "model": tool_config.model, + "requires_restart": False, + } + + +async def request_image_generation_reload( + bus: Any, + *, + timeout: float = 5.0, +) -> dict[str, Any]: + """Ask the running agent loop to refresh its image generation tool.""" + loop = asyncio.get_running_loop() + ack: asyncio.Future[dict[str, Any]] = loop.create_future() + await bus.publish_inbound( + InboundMessage( + channel="system", + sender_id="webui-settings", + chat_id="runtime", + content=RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD, + metadata={ + INBOUND_META_RUNTIME_CONTROL: RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD, + RUNTIME_CONTROL_ACK: ack, + }, + ) + ) + try: + result = await asyncio.wait_for(ack, timeout=timeout) + except asyncio.TimeoutError: + return { + "ok": False, + "message": "Image generation hot reload timed out.", + "requires_restart": True, + } + return result if isinstance(result, dict) else { + "ok": False, + "message": "Image generation hot reload returned an unexpected response.", + "requires_restart": True, + } + + +async def handle_runtime_control( + state: Any, + msg: InboundMessage, + registry: ToolRegistry, +) -> bool: + """Handle an in-process image generation reload request.""" + metadata = msg.metadata if isinstance(msg.metadata, dict) else {} + if metadata.get(INBOUND_META_RUNTIME_CONTROL) != RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD: + return False + + ack = metadata.get(RUNTIME_CONTROL_ACK) + try: + result = await reload_image_generation_tool(state, registry) + except Exception as exc: + logger.exception("Image generation hot reload failed") + result = { + "ok": False, + "message": "Image generation hot reload failed.", + "requires_restart": True, + "error": str(exc), + } + if isinstance(ack, asyncio.Future) and not ack.done(): + ack.set_result(result) + return True diff --git a/nanobot/bus/events.py b/nanobot/bus/events.py index 5bfdd6db..def7703f 100644 --- a/nanobot/bus/events.py +++ b/nanobot/bus/events.py @@ -17,6 +17,7 @@ OUTBOUND_META_AGENT_UI = "_agent_ui" INBOUND_META_RUNTIME_CONTROL = "_runtime_control" RUNTIME_CONTROL_ACK = "_ack" RUNTIME_CONTROL_MCP_RELOAD = "mcp_reload" +RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD = "image_generation_reload" @dataclass diff --git a/nanobot/channels/websocket/tests/test_websocket_channel.py b/nanobot/channels/websocket/tests/test_websocket_channel.py index 7898e581..935203ba 100644 --- a/nanobot/channels/websocket/tests/test_websocket_channel.py +++ b/nanobot/channels/websocket/tests/test_websocket_channel.py @@ -1969,6 +1969,17 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist( "login_supported": True, }, ) + image_reload = AsyncMock( + return_value={ + "ok": True, + "message": "Image generation settings applied.", + "requires_restart": False, + } + ) + monkeypatch.setattr( + "nanobot.webui.settings_routes.request_image_generation_reload", + image_reload, + ) channel = _ch(bus, port=port) channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300 @@ -2029,8 +2040,14 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist( } assert image_providers["openrouter"]["label"] == "OpenRouter" assert image_providers["openrouter"]["configured"] is False + assert image_providers["openrouter"]["default_model"] == "openai/gpt-5.4-image-2" + assert image_providers["openrouter"]["models"] == ["openai/gpt-5.4-image-2"] assert image_providers["openai_codex"]["auth_type"] == "oauth" assert image_providers["openai_codex"]["configured"] is False + assert image_providers["gemini"]["models"] == [ + "gemini-2.5-flash-image", + "imagen-4.0-generate-001", + ] assert image_providers["gemini"]["label"] == "Gemini" assert body["runtime"]["config_path"] == str(config_path) workspace_path = body["runtime"]["workspace_path"].replace("\\", "/") @@ -2187,7 +2204,7 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist( assert image_updated.status_code == 200 image_body = image_updated.json() assert image_body["requires_restart"] is True - assert image_body["restart_required_sections"] == ["browser", "image", "runtime"] + assert image_body["restart_required_sections"] == ["browser", "runtime"] assert image_body["image_generation"]["enabled"] is True assert image_body["image_generation"]["model"] == "openai/gpt-image-1" assert image_body["image_generation"]["default_aspect_ratio"] == "16:9" @@ -2202,12 +2219,9 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist( ) assert image_provider_updated.status_code == 200 assert image_provider_updated.json()["requires_restart"] is True - assert image_provider_updated.json()["restart_required_sections"] == [ - "browser", - "image", - "runtime", - ] + assert image_provider_updated.json()["restart_required_sections"] == ["browser", "runtime"] assert "sk-or-next" not in image_provider_updated.text + assert image_reload.await_count == 2 bad_web = await _http_get( "http://127.0.0.1:" @@ -2255,6 +2269,92 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist( await server_task +@pytest.mark.asyncio +async def test_image_settings_hot_reload_without_restart( + bus: MagicMock, + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + port = 29935 + config_path = tmp_path / "config.json" + config = Config() + config.providers.openrouter.api_key = "image-key" + save_config(config, config_path) + monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path) + image_reload = AsyncMock( + return_value={ + "ok": True, + "message": "Image generation settings applied.", + "requires_restart": False, + } + ) + monkeypatch.setattr( + "nanobot.webui.settings_routes.request_image_generation_reload", + image_reload, + ) + + channel = _ch(bus, port=port) + channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300 + server_task = asyncio.create_task(channel.start()) + await asyncio.sleep(0.3) + try: + response = await _http_get( + f"http://127.0.0.1:{port}/api/settings/image-generation/update" + "?enabled=true&provider=openrouter&model=openai%2Fgpt-image-1", + headers={"Authorization": "Bearer tok"}, + ) + + assert response.status_code == 200 + assert response.json()["requires_restart"] is False + assert response.json()["restart_required_sections"] == [] + image_reload.assert_awaited_once_with(bus) + finally: + await channel.stop() + await server_task + + +@pytest.mark.asyncio +async def test_image_settings_fall_back_to_restart_when_hot_reload_fails( + bus: MagicMock, + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + port = 29936 + config_path = tmp_path / "config.json" + config = Config() + config.providers.openrouter.api_key = "image-key" + save_config(config, config_path) + monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path) + monkeypatch.setattr( + "nanobot.webui.settings_routes.request_image_generation_reload", + AsyncMock( + return_value={ + "ok": False, + "message": "Image generation hot reload timed out.", + "requires_restart": True, + } + ), + ) + + channel = _ch(bus, port=port) + channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300 + server_task = asyncio.create_task(channel.start()) + await asyncio.sleep(0.3) + try: + response = await _http_get( + f"http://127.0.0.1:{port}/api/settings/image-generation/update" + "?enabled=true&provider=openrouter&model=openai%2Fgpt-image-1", + headers={"Authorization": "Bearer tok"}, + ) + + assert response.status_code == 200 + assert response.json()["requires_restart"] is True + assert response.json()["restart_required_sections"] == ["image"] + finally: + await channel.stop() + await server_task + + @pytest.mark.asyncio async def test_commands_api_returns_slash_command_metadata(bus: MagicMock) -> None: port = 29892 diff --git a/nanobot/providers/image_generation.py b/nanobot/providers/image_generation.py index 04ee4631..58951c34 100644 --- a/nanobot/providers/image_generation.py +++ b/nanobot/providers/image_generation.py @@ -177,6 +177,7 @@ class ImageGenerationProvider(ABC): """Base class for image generation provider clients.""" provider_name: str = "" + model_options: tuple[str, ...] = () missing_key_message: str = "" default_timeout: float = _DEFAULT_TIMEOUT_S @@ -254,6 +255,7 @@ class OpenRouterImageGenerationClient(ImageGenerationProvider): """Small async client for OpenRouter Chat Completions image generation.""" provider_name = "openrouter" + model_options = ("openai/gpt-5.4-image-2",) missing_key_message = ( "OpenRouter API key is not configured. Set providers.openrouter.apiKey." ) @@ -345,6 +347,7 @@ class AIHubMixImageGenerationClient(ImageGenerationProvider): """Small async client for AIHubMix unified image generation.""" provider_name = "aihubmix" + model_options = ("gpt-image-2-free",) missing_key_message = ( "AIHubMix API key is not configured. Set providers.aihubmix.apiKey." ) @@ -515,6 +518,7 @@ class OllamaImageGenerationClient(ImageGenerationProvider): """Async client for Ollama native image generation models.""" provider_name = "ollama" + model_options = ("x/z-image-turbo",) default_timeout = 300.0 def _default_base_url(self) -> str: @@ -591,6 +595,7 @@ class GeminiImageGenerationClient(ImageGenerationProvider): """Async client for Gemini/Imagen image generation via the Generative Language API.""" provider_name = "gemini" + model_options = ("gemini-2.5-flash-image", "imagen-4.0-generate-001") missing_key_message = ( "Gemini API key is not configured. Set providers.gemini.apiKey." ) @@ -815,6 +820,7 @@ class MiniMaxImageGenerationClient(ImageGenerationProvider): """Async client for MiniMax image generation API.""" provider_name = "minimax" + model_options = ("image-01",) missing_key_message = ( "MiniMax API key is not configured. Set providers.minimax.apiKey." ) @@ -947,6 +953,7 @@ class OpenAIImageGenerationClient(ImageGenerationProvider): """OpenAI Images API using an API key (``providers.openai.apiKey``).""" provider_name = "openai" + model_options = ("gpt-image-2", "gpt-image-1", "dall-e-3", "dall-e-2") missing_key_message = ( "OpenAI API key is not configured. Set providers.openai.apiKey." ) @@ -1210,6 +1217,7 @@ class CodexImageGenerationClient(ImageGenerationProvider): """ provider_name = "openai_codex" + model_options = ("gpt-5.4",) missing_key_message = ( "Codex OAuth token is unavailable. " "Log in with Codex subscription first." @@ -1508,6 +1516,7 @@ class StepFunImageGenerationClient(ImageGenerationProvider): """ provider_name = "stepfun" + model_options = ("step-image-edit-2", "step-1x-medium") missing_key_message = ( "StepFun API key is not configured. Set providers.stepfun.apiKey." ) @@ -1634,6 +1643,7 @@ class ZhipuImageGenerationClient(ImageGenerationProvider): """ provider_name = "zhipu" + model_options = ("glm-image", "cogview-4", "cogview-4-250304", "cogview-3-flash") missing_key_message = "Zhipu API key is not configured. Set providers.zhipu.apiKey." default_timeout = _ZHIPU_TIMEOUT_S diff --git a/nanobot/webui/settings_api.py b/nanobot/webui/settings_api.py index 152e5a0c..0ba46362 100644 --- a/nanobot/webui/settings_api.py +++ b/nanobot/webui/settings_api.py @@ -703,6 +703,7 @@ def _validate_configured_provider(config: Any, provider: str) -> None: def _image_generation_provider_rows(config: Any) -> list[dict[str, Any]]: rows: list[dict[str, Any]] = [] for name in image_gen_provider_names(): + image_provider = get_image_gen_provider(name) spec = find_by_name(name) provider_config = getattr(config.providers, name, None) configured = ( @@ -723,6 +724,12 @@ def _image_generation_provider_rows(config: Any) -> list[dict[str, Any]]: "default_api_base": ( spec.default_api_base if spec and spec.default_api_base else None ), + "models": list(image_provider.model_options) if image_provider else [], + "default_model": ( + image_provider.model_options[0] + if image_provider and image_provider.model_options + else None + ), } ) return rows diff --git a/nanobot/webui/settings_routes.py b/nanobot/webui/settings_routes.py index 1060f0f2..19a4fcff 100644 --- a/nanobot/webui/settings_routes.py +++ b/nanobot/webui/settings_routes.py @@ -17,6 +17,7 @@ from typing import Any from websockets.http11 import Request as WsRequest from websockets.http11 import Response +from nanobot.agent.tools.image_generation import request_image_generation_reload from nanobot.agent.tools.mcp import request_mcp_reload from nanobot.api.runtime import ApiRuntime, ApiStartOptions, api_runtime_paths from nanobot.bus.queue import MessageBus @@ -145,7 +146,7 @@ class WebUISettingsRouter: if path == "/api/settings/model-configurations/update": return self._handle_settings_model_configuration_update(request) if path == "/api/settings/provider/update": - return self._handle_settings_provider_update(request) + return await self._handle_settings_provider_update(request) if path == "/api/settings/provider-models": return await self._handle_settings_provider_models(request) if path == "/api/settings/provider/oauth-login": @@ -163,7 +164,7 @@ class WebUISettingsRouter: if path == "/api/settings/api-service/stop": return await self._handle_settings_api_service_stop(request) if path == "/api/settings/image-generation/update": - return self._handle_settings_image_generation_update(request) + return await self._handle_settings_image_generation_update(request) if path == "/api/settings/transcription/update": return self._handle_settings_transcription_update(request) if path == "/api/settings/network-safety/update": @@ -352,13 +353,14 @@ class WebUISettingsRouter: return self._error_response(e.status, e.message) return self._json_response(self._with_restart_state(payload)) - def _handle_settings_provider_update(self, request: WsRequest) -> Response: + async def _handle_settings_provider_update(self, request: WsRequest) -> Response: if not self._authorized(request): return self._unauthorized() try: payload = update_provider_settings(self._query(request)) except WebUISettingsError as e: return self._error_response(e.status, e.message) + payload = await self._apply_image_generation_runtime_change(payload) return self._json_response(self._with_restart_state(payload, section="image")) async def _handle_settings_provider_models(self, request: WsRequest) -> Response: @@ -541,15 +543,41 @@ class WebUISettingsRouter: return f"API server {message.removeprefix('api_').replace('_', ' ')}" return message.replace("_", " ") - def _handle_settings_image_generation_update(self, request: WsRequest) -> Response: + async def _handle_settings_image_generation_update(self, request: WsRequest) -> Response: if not self._authorized(request): return self._unauthorized() try: payload = update_image_generation_settings(self._query(request)) except WebUISettingsError as e: return self._error_response(e.status, e.message) + payload = await self._apply_image_generation_runtime_change(payload) return self._json_response(self._with_restart_state(payload, section="image")) + async def _apply_image_generation_runtime_change( + self, + payload: dict[str, Any], + ) -> dict[str, Any]: + """Hot-apply image settings, preserving restart fallback on failure.""" + if not payload.get("requires_restart"): + return payload + try: + result = await request_image_generation_reload(self.bus) + except Exception: + self.logger.exception("failed to hot-reload image generation settings") + return payload + + applied = bool(result.get("ok")) and not result.get("requires_restart") + payload = dict(payload) + payload["requires_restart"] = not applied + if applied: + self._restart_sections.discard("image") + else: + self.logger.warning( + "image generation settings were saved but require restart: {}", + result.get("message") or "hot reload failed", + ) + return payload + def _handle_settings_transcription_update(self, request: WsRequest) -> Response: if not self._authorized(request): return self._unauthorized() diff --git a/tests/agent/test_image_generation_hot_reload.py b/tests/agent/test_image_generation_hot_reload.py new file mode 100644 index 00000000..48ee1a24 --- /dev/null +++ b/tests/agent/test_image_generation_hot_reload.py @@ -0,0 +1,103 @@ +"""Image generation runtime reload tests.""" + +from __future__ import annotations + +import asyncio +from types import SimpleNamespace + +import pytest + +from nanobot.agent import context as agent_context +from nanobot.agent.tools.image_generation import ( + ImageGenerationTool, + reload_image_generation_tool, + request_image_generation_reload, +) +from nanobot.agent.tools.registry import ToolRegistry +from nanobot.bus.queue import MessageBus +from nanobot.config.loader import load_config, save_config +from nanobot.config.schema import ToolsConfig + + +def _runtime_state(tmp_path): + return SimpleNamespace( + workspace=tmp_path, + tools_config=ToolsConfig(), + _image_generation_provider_configs={}, + ) + + +@pytest.mark.asyncio +async def test_image_generation_reload_replaces_and_removes_live_tool( + tmp_path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + config_path = tmp_path / "config.json" + monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path) + config = load_config() + config.providers.openrouter.api_key = "first-key" + config.tools.image_generation.enabled = True + config.tools.image_generation.provider = "openrouter" + config.tools.image_generation.model = "openai/first-image-model" + save_config(config) + + state = _runtime_state(tmp_path) + registry = ToolRegistry() + result = await reload_image_generation_tool(state, registry) + + first_tool = registry.get("generate_image") + assert result["requires_restart"] is False + assert isinstance(first_tool, ImageGenerationTool) + assert first_tool.config.model == "openai/first-image-model" + assert first_tool.provider_configs["openrouter"].api_key == "first-key" + + config = load_config() + config.providers.openrouter.api_key = "second-key" + config.tools.image_generation.model = "openai/second-image-model" + save_config(config) + + result = await reload_image_generation_tool(state, registry) + second_tool = registry.get("generate_image") + assert result["requires_restart"] is False + assert isinstance(second_tool, ImageGenerationTool) + assert second_tool is not first_tool + assert second_tool.config.model == "openai/second-image-model" + assert second_tool.provider_configs["openrouter"].api_key == "second-key" + assert state.tools_config.image_generation.model == "openai/second-image-model" + + config = load_config() + config.tools.image_generation.enabled = False + save_config(config) + + result = await reload_image_generation_tool(state, registry) + assert result["requires_restart"] is False + assert not registry.has("generate_image") + + +@pytest.mark.asyncio +async def test_image_generation_reload_reaches_agent_runtime_control( + tmp_path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + config_path = tmp_path / "config.json" + monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path) + config = load_config() + config.providers.openrouter.api_key = "image-key" + config.tools.image_generation.enabled = True + save_config(config) + + bus = MessageBus() + state = _runtime_state(tmp_path) + registry = ToolRegistry() + + async def consume_control() -> None: + message = await bus.consume_inbound() + assert await agent_context.handle_runtime_control(state, message, registry) is True + + consumer = asyncio.create_task(consume_control()) + result = await request_image_generation_reload(bus, timeout=2.0) + await consumer + + assert result["ok"] is True + assert result["requires_restart"] is False + assert registry.has("generate_image") diff --git a/webui/src/components/settings/SettingsView.tsx b/webui/src/components/settings/SettingsView.tsx index 2a4a0f7e..ebfdc527 100644 --- a/webui/src/components/settings/SettingsView.tsx +++ b/webui/src/components/settings/SettingsView.tsx @@ -3574,6 +3574,20 @@ function ImageGenerationSettings({ IMAGE_SIZE_OPTIONS.map((value) => ({ name: value, label: value })), form.defaultImageSize, ); + const modelOptions = optionRowsWithCurrent( + (selectedProvider?.models ?? []).map((model) => ({ name: model, label: model })), + form.model, + ); + const hasModelCatalog = Boolean(selectedProvider?.models?.length); + + const selectProvider = (provider: string) => { + const nextProvider = settings.image_generation.providers.find((row) => row.name === provider); + onChangeForm((prev) => ({ + ...prev, + provider, + model: nextProvider?.default_model || nextProvider?.models?.[0] || prev.model, + })); + }; return (
@@ -3600,7 +3614,7 @@ function ImageGenerationSettings({ value={form.provider} emptyLabel={tx("settings.image.selectProvider", "Select provider")} showProviderLogos={showBrandLogos} - onChange={(provider) => onChangeForm((prev) => ({ ...prev, provider }))} + onChange={selectProvider} /> - onChangeForm((prev) => ({ ...prev, model: event.target.value }))} - className="h-8 w-[min(300px,70vw)] rounded-full text-[13px]" - /> + {hasModelCatalog ? ( + onChangeForm((prev) => ({ ...prev, model }))} + triggerClassName="w-[min(300px,70vw)]" + contentClassName="w-[320px]" + /> + ) : ( + onChangeForm((prev) => ({ ...prev, model: event.target.value }))} + className="h-8 w-[min(300px,70vw)] rounded-full text-[13px]" + /> + )} ; value: string; emptyLabel: string; showProviderLogos?: boolean; + triggerClassName?: string; + contentClassName?: string; onChange: (provider: string) => void; }) { const selectedProvider = providers.find((provider) => provider.name === value) ?? null; @@ -7654,6 +7683,7 @@ function ProviderPicker({ "h-8 w-[210px] justify-between rounded-full border-input bg-background px-3 text-[13px] font-normal shadow-none", "hover:bg-accent/55 focus-visible:ring-2 focus-visible:ring-ring", disabled && "text-muted-foreground", + triggerClassName, )} > @@ -7670,7 +7700,10 @@ function ProviderPicker({ {providers.map((provider) => { const selected = provider.name === value; diff --git a/webui/src/lib/types.ts b/webui/src/lib/types.ts index 10430e99..ad873c26 100644 --- a/webui/src/lib/types.ts +++ b/webui/src/lib/types.ts @@ -501,6 +501,8 @@ export interface SettingsPayload { api_key_hint?: string | null; api_base?: string | null; default_api_base?: string | null; + models?: string[]; + default_model?: string | null; }>; }; transcription?: { diff --git a/webui/src/tests/settings-view.test.tsx b/webui/src/tests/settings-view.test.tsx index 043c386d..70b59eb9 100644 --- a/webui/src/tests/settings-view.test.tsx +++ b/webui/src/tests/settings-view.test.tsx @@ -303,6 +303,7 @@ function renderSettingsView( | "automations" | "advanced" | "models" + | "image" | "browser" | "runtime"; initialSettings?: SettingsPayload; @@ -2277,6 +2278,56 @@ describe("SettingsView Apps catalog", () => { ); }); + it("selects image models from provider-specific options", async () => { + const base = settingsPayload(); + const payload: SettingsPayload = { + ...base, + image_generation: { + ...base.image_generation, + providers: [ + { + name: "openrouter", + label: "OpenRouter", + configured: true, + models: ["openai/gpt-5.4-image-2"], + default_model: "openai/gpt-5.4-image-2", + }, + { + name: "gemini", + label: "Gemini", + configured: true, + models: ["gemini-2.5-flash-image", "imagen-4.0-generate-001"], + default_model: "gemini-2.5-flash-image", + }, + { + name: "custom", + label: "Custom", + configured: true, + models: [], + default_model: null, + }, + ], + }, + }; + + renderSettingsView({ initialSection: "image", initialSettings: payload }); + + expect(screen.queryByDisplayValue("openai/gpt-5.4-image-2")).not.toBeInTheDocument(); + fireEvent.pointerDown(screen.getByRole("button", { name: "OpenRouter" })); + fireEvent.click(await screen.findByRole("menuitem", { name: "Gemini" })); + + expect(await screen.findByRole("button", { name: "gemini-2.5-flash-image" })).toBeInTheDocument(); + fireEvent.pointerDown(screen.getByRole("button", { name: "gemini-2.5-flash-image" })); + fireEvent.click(await screen.findByRole("menuitem", { name: "imagen-4.0-generate-001" })); + await waitFor(() => + expect(screen.getByRole("button", { name: "imagen-4.0-generate-001" })).toBeInTheDocument(), + ); + + fireEvent.pointerDown(screen.getByRole("button", { name: "Gemini" })); + fireEvent.click(await screen.findByRole("menuitem", { name: "Custom" })); + expect(screen.getByDisplayValue("imagen-4.0-generate-001")).toBeInTheDocument(); + }); + it("keeps the default model distinct from the active named configuration", async () => { const base = settingsPayload(); const payload: SettingsPayload = {