feat(image): apply generation settings live

This commit is contained in:
Xubin Ren
2026-07-23 12:42:24 +08:00
parent c7393c785e
commit 1616fa9f14
11 changed files with 481 additions and 18 deletions
+8 -1
View File
@@ -8,6 +8,7 @@ from typing import Any, Mapping, Sequence
from nanobot.agent.memory import MemoryStore from nanobot.agent.memory import MemoryStore
from nanobot.agent.skills import SkillsLoader 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 import mcp as mcp_tools
from nanobot.agent.tools.registry import ToolRegistry from nanobot.agent.tools.registry import ToolRegistry
from nanobot.apps.cli import utils as cli_app_utils 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: 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: class ContextBuilder:
+121
View File
@@ -2,24 +2,34 @@
from __future__ import annotations from __future__ import annotations
import asyncio
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from loguru import logger
from pydantic import Field from pydantic import Field
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.agent.tools.schema import ( from nanobot.agent.tools.schema import (
ArraySchema, ArraySchema,
IntegerSchema, IntegerSchema,
StringSchema, StringSchema,
tool_parameters_schema, 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.paths import get_media_dir
from nanobot.config_base import Base from nanobot.config_base import Base
from nanobot.providers.image_generation import ( from nanobot.providers.image_generation import (
ImageGenerationError, ImageGenerationError,
ImageGenerationProvider, ImageGenerationProvider,
get_image_gen_provider, get_image_gen_provider,
image_gen_provider_configs,
) )
from nanobot.security.workspace_access import current_tool_workspace from nanobot.security.workspace_access import current_tool_workspace
from nanobot.security.workspace_policy import WorkspaceBoundaryError, resolve_allowed_path from nanobot.security.workspace_policy import WorkspaceBoundaryError, resolve_allowed_path
@@ -208,3 +218,114 @@ class ImageGenerationTool(Tool):
return generated_image_tool_result(artifacts) return generated_image_tool_result(artifacts)
except (ArtifactError, ImageGenerationError, OSError) as exc: except (ArtifactError, ImageGenerationError, OSError) as exc:
return ToolResult.error(f"Error: {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
+1
View File
@@ -17,6 +17,7 @@ OUTBOUND_META_AGENT_UI = "_agent_ui"
INBOUND_META_RUNTIME_CONTROL = "_runtime_control" INBOUND_META_RUNTIME_CONTROL = "_runtime_control"
RUNTIME_CONTROL_ACK = "_ack" RUNTIME_CONTROL_ACK = "_ack"
RUNTIME_CONTROL_MCP_RELOAD = "mcp_reload" RUNTIME_CONTROL_MCP_RELOAD = "mcp_reload"
RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD = "image_generation_reload"
@dataclass @dataclass
@@ -1969,6 +1969,17 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
"login_supported": True, "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 = _ch(bus, port=port)
channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300 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"]["label"] == "OpenRouter"
assert image_providers["openrouter"]["configured"] is False 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"]["auth_type"] == "oauth"
assert image_providers["openai_codex"]["configured"] is False 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 image_providers["gemini"]["label"] == "Gemini"
assert body["runtime"]["config_path"] == str(config_path) assert body["runtime"]["config_path"] == str(config_path)
workspace_path = body["runtime"]["workspace_path"].replace("\\", "/") 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 assert image_updated.status_code == 200
image_body = image_updated.json() image_body = image_updated.json()
assert image_body["requires_restart"] is True 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"]["enabled"] is True
assert image_body["image_generation"]["model"] == "openai/gpt-image-1" assert image_body["image_generation"]["model"] == "openai/gpt-image-1"
assert image_body["image_generation"]["default_aspect_ratio"] == "16:9" 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.status_code == 200
assert image_provider_updated.json()["requires_restart"] is True assert image_provider_updated.json()["requires_restart"] is True
assert image_provider_updated.json()["restart_required_sections"] == [ assert image_provider_updated.json()["restart_required_sections"] == ["browser", "runtime"]
"browser",
"image",
"runtime",
]
assert "sk-or-next" not in image_provider_updated.text assert "sk-or-next" not in image_provider_updated.text
assert image_reload.await_count == 2
bad_web = await _http_get( bad_web = await _http_get(
"http://127.0.0.1:" "http://127.0.0.1:"
@@ -2255,6 +2269,92 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
await server_task 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 @pytest.mark.asyncio
async def test_commands_api_returns_slash_command_metadata(bus: MagicMock) -> None: async def test_commands_api_returns_slash_command_metadata(bus: MagicMock) -> None:
port = 29892 port = 29892
+10
View File
@@ -177,6 +177,7 @@ class ImageGenerationProvider(ABC):
"""Base class for image generation provider clients.""" """Base class for image generation provider clients."""
provider_name: str = "" provider_name: str = ""
model_options: tuple[str, ...] = ()
missing_key_message: str = "" missing_key_message: str = ""
default_timeout: float = _DEFAULT_TIMEOUT_S default_timeout: float = _DEFAULT_TIMEOUT_S
@@ -254,6 +255,7 @@ class OpenRouterImageGenerationClient(ImageGenerationProvider):
"""Small async client for OpenRouter Chat Completions image generation.""" """Small async client for OpenRouter Chat Completions image generation."""
provider_name = "openrouter" provider_name = "openrouter"
model_options = ("openai/gpt-5.4-image-2",)
missing_key_message = ( missing_key_message = (
"OpenRouter API key is not configured. Set providers.openrouter.apiKey." "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.""" """Small async client for AIHubMix unified image generation."""
provider_name = "aihubmix" provider_name = "aihubmix"
model_options = ("gpt-image-2-free",)
missing_key_message = ( missing_key_message = (
"AIHubMix API key is not configured. Set providers.aihubmix.apiKey." "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.""" """Async client for Ollama native image generation models."""
provider_name = "ollama" provider_name = "ollama"
model_options = ("x/z-image-turbo",)
default_timeout = 300.0 default_timeout = 300.0
def _default_base_url(self) -> str: 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.""" """Async client for Gemini/Imagen image generation via the Generative Language API."""
provider_name = "gemini" provider_name = "gemini"
model_options = ("gemini-2.5-flash-image", "imagen-4.0-generate-001")
missing_key_message = ( missing_key_message = (
"Gemini API key is not configured. Set providers.gemini.apiKey." "Gemini API key is not configured. Set providers.gemini.apiKey."
) )
@@ -815,6 +820,7 @@ class MiniMaxImageGenerationClient(ImageGenerationProvider):
"""Async client for MiniMax image generation API.""" """Async client for MiniMax image generation API."""
provider_name = "minimax" provider_name = "minimax"
model_options = ("image-01",)
missing_key_message = ( missing_key_message = (
"MiniMax API key is not configured. Set providers.minimax.apiKey." "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``).""" """OpenAI Images API using an API key (``providers.openai.apiKey``)."""
provider_name = "openai" provider_name = "openai"
model_options = ("gpt-image-2", "gpt-image-1", "dall-e-3", "dall-e-2")
missing_key_message = ( missing_key_message = (
"OpenAI API key is not configured. Set providers.openai.apiKey." "OpenAI API key is not configured. Set providers.openai.apiKey."
) )
@@ -1210,6 +1217,7 @@ class CodexImageGenerationClient(ImageGenerationProvider):
""" """
provider_name = "openai_codex" provider_name = "openai_codex"
model_options = ("gpt-5.4",)
missing_key_message = ( missing_key_message = (
"Codex OAuth token is unavailable. " "Codex OAuth token is unavailable. "
"Log in with Codex subscription first." "Log in with Codex subscription first."
@@ -1508,6 +1516,7 @@ class StepFunImageGenerationClient(ImageGenerationProvider):
""" """
provider_name = "stepfun" provider_name = "stepfun"
model_options = ("step-image-edit-2", "step-1x-medium")
missing_key_message = ( missing_key_message = (
"StepFun API key is not configured. Set providers.stepfun.apiKey." "StepFun API key is not configured. Set providers.stepfun.apiKey."
) )
@@ -1634,6 +1643,7 @@ class ZhipuImageGenerationClient(ImageGenerationProvider):
""" """
provider_name = "zhipu" 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." missing_key_message = "Zhipu API key is not configured. Set providers.zhipu.apiKey."
default_timeout = _ZHIPU_TIMEOUT_S default_timeout = _ZHIPU_TIMEOUT_S
+7
View File
@@ -703,6 +703,7 @@ def _validate_configured_provider(config: Any, provider: str) -> None:
def _image_generation_provider_rows(config: Any) -> list[dict[str, Any]]: def _image_generation_provider_rows(config: Any) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = [] rows: list[dict[str, Any]] = []
for name in image_gen_provider_names(): for name in image_gen_provider_names():
image_provider = get_image_gen_provider(name)
spec = find_by_name(name) spec = find_by_name(name)
provider_config = getattr(config.providers, name, None) provider_config = getattr(config.providers, name, None)
configured = ( configured = (
@@ -723,6 +724,12 @@ def _image_generation_provider_rows(config: Any) -> list[dict[str, Any]]:
"default_api_base": ( "default_api_base": (
spec.default_api_base if spec and spec.default_api_base else None 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 return rows
+32 -4
View File
@@ -17,6 +17,7 @@ from typing import Any
from websockets.http11 import Request as WsRequest from websockets.http11 import Request as WsRequest
from websockets.http11 import Response 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.agent.tools.mcp import request_mcp_reload
from nanobot.api.runtime import ApiRuntime, ApiStartOptions, api_runtime_paths from nanobot.api.runtime import ApiRuntime, ApiStartOptions, api_runtime_paths
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
@@ -145,7 +146,7 @@ class WebUISettingsRouter:
if path == "/api/settings/model-configurations/update": if path == "/api/settings/model-configurations/update":
return self._handle_settings_model_configuration_update(request) return self._handle_settings_model_configuration_update(request)
if path == "/api/settings/provider/update": 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": if path == "/api/settings/provider-models":
return await self._handle_settings_provider_models(request) return await self._handle_settings_provider_models(request)
if path == "/api/settings/provider/oauth-login": if path == "/api/settings/provider/oauth-login":
@@ -163,7 +164,7 @@ class WebUISettingsRouter:
if path == "/api/settings/api-service/stop": if path == "/api/settings/api-service/stop":
return await self._handle_settings_api_service_stop(request) return await self._handle_settings_api_service_stop(request)
if path == "/api/settings/image-generation/update": 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": if path == "/api/settings/transcription/update":
return self._handle_settings_transcription_update(request) return self._handle_settings_transcription_update(request)
if path == "/api/settings/network-safety/update": if path == "/api/settings/network-safety/update":
@@ -352,13 +353,14 @@ class WebUISettingsRouter:
return self._error_response(e.status, e.message) return self._error_response(e.status, e.message)
return self._json_response(self._with_restart_state(payload)) 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): if not self._authorized(request):
return self._unauthorized() return self._unauthorized()
try: try:
payload = update_provider_settings(self._query(request)) payload = update_provider_settings(self._query(request))
except WebUISettingsError as e: except WebUISettingsError as e:
return self._error_response(e.status, e.message) 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")) return self._json_response(self._with_restart_state(payload, section="image"))
async def _handle_settings_provider_models(self, request: WsRequest) -> Response: 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 f"API server {message.removeprefix('api_').replace('_', ' ')}"
return message.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): if not self._authorized(request):
return self._unauthorized() return self._unauthorized()
try: try:
payload = update_image_generation_settings(self._query(request)) payload = update_image_generation_settings(self._query(request))
except WebUISettingsError as e: except WebUISettingsError as e:
return self._error_response(e.status, e.message) 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")) 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: def _handle_settings_transcription_update(self, request: WsRequest) -> Response:
if not self._authorized(request): if not self._authorized(request):
return self._unauthorized() return self._unauthorized()
@@ -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")
+40 -7
View File
@@ -3574,6 +3574,20 @@ function ImageGenerationSettings({
IMAGE_SIZE_OPTIONS.map((value) => ({ name: value, label: value })), IMAGE_SIZE_OPTIONS.map((value) => ({ name: value, label: value })),
form.defaultImageSize, 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 ( return (
<div className="space-y-7"> <div className="space-y-7">
@@ -3600,7 +3614,7 @@ function ImageGenerationSettings({
value={form.provider} value={form.provider}
emptyLabel={tx("settings.image.selectProvider", "Select provider")} emptyLabel={tx("settings.image.selectProvider", "Select provider")}
showProviderLogos={showBrandLogos} showProviderLogos={showBrandLogos}
onChange={(provider) => onChangeForm((prev) => ({ ...prev, provider }))} onChange={selectProvider}
/> />
</SettingsRow> </SettingsRow>
<SettingsRow <SettingsRow
@@ -3635,11 +3649,22 @@ function ImageGenerationSettings({
title={tx("settings.rows.imageModel", "Image model")} title={tx("settings.rows.imageModel", "Image model")}
description={tx("settings.help.imageModel", "Model name sent to the selected image provider.")} description={tx("settings.help.imageModel", "Model name sent to the selected image provider.")}
> >
<Input {hasModelCatalog ? (
value={form.model} <ProviderPicker
onChange={(event) => onChangeForm((prev) => ({ ...prev, model: event.target.value }))} providers={modelOptions}
className="h-8 w-[min(300px,70vw)] rounded-full text-[13px]" value={form.model}
/> emptyLabel={tx("settings.image.selectModel", "Select image model")}
onChange={(model) => onChangeForm((prev) => ({ ...prev, model }))}
triggerClassName="w-[min(300px,70vw)]"
contentClassName="w-[320px]"
/>
) : (
<Input
value={form.model}
onChange={(event) => onChangeForm((prev) => ({ ...prev, model: event.target.value }))}
className="h-8 w-[min(300px,70vw)] rounded-full text-[13px]"
/>
)}
</SettingsRow> </SettingsRow>
<SettingsRow <SettingsRow
title={tx("settings.rows.defaultAspectRatio", "Default aspect")} title={tx("settings.rows.defaultAspectRatio", "Default aspect")}
@@ -7632,12 +7657,16 @@ function ProviderPicker({
value, value,
emptyLabel, emptyLabel,
showProviderLogos = false, showProviderLogos = false,
triggerClassName,
contentClassName,
onChange, onChange,
}: { }: {
providers: Array<{ name: string; label: string }>; providers: Array<{ name: string; label: string }>;
value: string; value: string;
emptyLabel: string; emptyLabel: string;
showProviderLogos?: boolean; showProviderLogos?: boolean;
triggerClassName?: string;
contentClassName?: string;
onChange: (provider: string) => void; onChange: (provider: string) => void;
}) { }) {
const selectedProvider = providers.find((provider) => provider.name === value) ?? null; 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", "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", "hover:bg-accent/55 focus-visible:ring-2 focus-visible:ring-ring",
disabled && "text-muted-foreground", disabled && "text-muted-foreground",
triggerClassName,
)} )}
> >
<span className="flex min-w-0 items-center gap-2"> <span className="flex min-w-0 items-center gap-2">
@@ -7670,7 +7700,10 @@ function ProviderPicker({
</DropdownMenuTrigger> </DropdownMenuTrigger>
<DropdownMenuContent <DropdownMenuContent
align="end" align="end"
className="max-h-[18rem] w-[240px] overflow-y-auto scrollbar-thin scrollbar-track-transparent" className={cn(
"max-h-[18rem] w-[240px] overflow-y-auto scrollbar-thin scrollbar-track-transparent",
contentClassName,
)}
> >
{providers.map((provider) => { {providers.map((provider) => {
const selected = provider.name === value; const selected = provider.name === value;
+2
View File
@@ -501,6 +501,8 @@ export interface SettingsPayload {
api_key_hint?: string | null; api_key_hint?: string | null;
api_base?: string | null; api_base?: string | null;
default_api_base?: string | null; default_api_base?: string | null;
models?: string[];
default_model?: string | null;
}>; }>;
}; };
transcription?: { transcription?: {
+51
View File
@@ -303,6 +303,7 @@ function renderSettingsView(
| "automations" | "automations"
| "advanced" | "advanced"
| "models" | "models"
| "image"
| "browser" | "browser"
| "runtime"; | "runtime";
initialSettings?: SettingsPayload; 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 () => { it("keeps the default model distinct from the active named configuration", async () => {
const base = settingsPayload(); const base = settingsPayload();
const payload: SettingsPayload = { const payload: SettingsPayload = {