diff --git a/lib/galaxy/config/schemas/config_schema.yml b/lib/galaxy/config/schemas/config_schema.yml index 5c4788cfc8f..bd128eb768b 100644 --- a/lib/galaxy/config/schemas/config_schema.yml +++ b/lib/galaxy/config/schemas/config_schema.yml @@ -4094,9 +4094,11 @@ mapping: type: any required: false desc: | - Configuration for AI inference services used by agents. Supports per-agent model, temperature, and token settings. - Agents inherit from 'default' configuration, which itself falls back to global ai_model/ai_api_key settings. - Example: inference_services: { default: { model: gpt-4o-mini, temperature: 0.7 } } + Configuration for AI inference services used by agents and visualization plugins. + Supports per-agent or per-plugin model, temperature, and token settings. + Valid keys include agent types (e.g. router, error_analysis) and plugin names (e.g. jupyterlite). + Agents and plugins inherit from 'default' configuration, which itself falls back to global ai_model/ai_api_key settings. + Example: inference_services: { default: { model: gpt-4o-mini }, jupyterlite: { model: gpt-4o } } enable_tool_recommendations: type: bool @@ -4374,10 +4376,10 @@ mapping: default: 86400 desc: | The interval in seconds between attempts to delete all failed Galaxy job's working directories from the filesystem (every 24 hours by default) if enable_failed_jobs_working_directory_cleanup is ``true``. Runs in a Celery task. - + enable_beta_tool_formats: type: bool default: false required: false desc: | - Enable beta tool formats (yaml, cwl, ...) which is a prerequisite for user defined tools. \ No newline at end of file + Enable beta tool formats (yaml, cwl, ...) which is a prerequisite for user defined tools. diff --git a/lib/galaxy/webapps/galaxy/api/plugins.py b/lib/galaxy/webapps/galaxy/api/plugins.py index fe7cffcfc40..4e7e84f8f0b 100644 --- a/lib/galaxy/webapps/galaxy/api/plugins.py +++ b/lib/galaxy/webapps/galaxy/api/plugins.py @@ -150,23 +150,49 @@ class FastAPIPlugins: plugin_specs = plugin and plugin.config.get("specs") plugin_ai_prompt = plugin_specs and plugin_specs.get("ai_prompt") if plugin_ai_prompt: - return await self._open_ai_adapter(payload, plugin_ai_prompt) + return await self._open_ai_adapter(payload, plugin_ai_prompt, plugin_name) else: return self._create_error("Selected plugin has no AI prompt.") else: return self._create_error("Visualization registry is not available.") + def _get_plugin_config(self, plugin_name: str, key: str) -> Optional[str]: + """Get config for a plugin with fallback through inference_services. + + Precedence: + 1. Plugin-specific: inference_services.. + 2. Default inference: inference_services.default. + 3. Global config: ai_model / ai_api_key / ai_api_base_url + """ + inference_config = getattr(self.config, "inference_services", None) + if isinstance(inference_config, dict): + plugin_specific = inference_config.get(plugin_name) + if isinstance(plugin_specific, dict) and key in plugin_specific: + return plugin_specific[key] + default_config = inference_config.get("default") + if isinstance(default_config, dict) and key in default_config: + return default_config[key] + + if key == "model": + return self.config.ai_model + elif key == "api_key": + return self.config.ai_api_key + elif key == "api_base_url": + return self.config.ai_api_base_url + return None + async def _open_ai_adapter( self, payload: ChatCompletionRequest, prompt: str, + plugin_name: str, ): """Galaxy managed chat completion adapter with prompt injection""" - # Collect configuration - ai_api_key = self.config.ai_api_key - ai_api_base_url = self.config.ai_api_base_url - ai_model = self.config.ai_model + # Collect configuration via inference_services fallback chain + ai_api_key = self._get_plugin_config(plugin_name, "api_key") + ai_api_base_url = self._get_plugin_config(plugin_name, "api_base_url") + ai_model = self._get_plugin_config(plugin_name, "model") if ai_api_key is None: return self._create_error("AI service not configured: API key is required.") if ai_model is None: @@ -220,11 +246,12 @@ class FastAPIPlugins: return self._create_error("Number of tools exceeded or invalid tools list.") # Build openai client with timeout - client_kwargs = dict(api_key=ai_api_key, timeout=TIMEOUT) - if ai_api_base_url: - client_kwargs["base_url"] = ai_api_base_url try: - client = AsyncOpenAI(**client_kwargs) + client = AsyncOpenAI( + api_key=ai_api_key, + timeout=TIMEOUT, + base_url=ai_api_base_url or None, + ) except Exception as e: log.debug("Failed to initialize OpenAI client.", exc_info=e) return self._create_error("Failed to initialize OpenAI client.", 500) diff --git a/test/integration/test_plugins.py b/test/integration/test_plugins.py index fe651773bb4..d620f3bbc8e 100644 --- a/test/integration/test_plugins.py +++ b/test/integration/test_plugins.py @@ -19,6 +19,16 @@ openai = pytest.importorskip("openai") TEST_VISUALIZATION_PLUGINS_DIR = os.path.join(os.path.dirname(__file__), "test_visualization_plugins") +def _create_chat_payload(extra=None): + payload = { + "messages": [{"role": "user", "content": "hi"}], + "tools": [], + } + if extra: + payload.update(extra) + return payload + + class TestVisualizationPluginsApi(IntegrationTestCase): """Tests for the visualization plugins API endpoints.""" @@ -81,15 +91,6 @@ class TestVisualizationPluginsApi(IntegrationTestCase): assert "tests" in plugin assert len(plugin["tests"]) >= 1 - def _create_payload(self, extra=None): - payload = { - "messages": [{"role": "user", "content": "hi"}], - "tools": [], - } - if extra: - payload.update(extra) - return payload - def _post_payload(self, payload=None, anon=False): return self._post("plugins/jupyterlite/chat/completions", payload, json=True, anon=anon) @@ -100,7 +101,7 @@ class TestVisualizationPluginsApi(IntegrationTestCase): mock_instance = MagicMock() mock_instance.chat.completions.create = AsyncMock(return_value=mock_response) mock_client.return_value = mock_instance - payload = self._create_payload() + payload = _create_chat_payload() response = self._post_payload(payload, anon=False) self._assert_status_code_is(response, 200) assert response.json()["id"] == "test" @@ -119,7 +120,7 @@ class TestVisualizationPluginsApi(IntegrationTestCase): mock_instance.chat.completions.create = AsyncMock(return_value=stream_gen()) mock_instance.close = AsyncMock() mock_client.return_value = mock_instance - payload = self._create_payload({"stream": True}) + payload = _create_chat_payload({"stream": True}) response = self._post_payload(payload, anon=False) self._assert_status_code_is(response, 200) body = response.text @@ -131,7 +132,7 @@ class TestVisualizationPluginsApi(IntegrationTestCase): assert mock_instance.close.called def test_tools_exceed_max(self): - payload = self._create_payload( + payload = _create_chat_payload( {"tools": [{"type": "function", "function": {"name": "f", "parameters": {}}}] * 129} ) response = self._post_payload(payload) @@ -139,7 +140,7 @@ class TestVisualizationPluginsApi(IntegrationTestCase): def test_tool_schema_too_large(self): big_params = {"x": "a" * 20000} - payload = self._create_payload( + payload = _create_chat_payload( {"tools": [{"type": "function", "function": {"name": "f", "parameters": big_params}}]} ) response = self._post_payload(payload) @@ -147,7 +148,7 @@ class TestVisualizationPluginsApi(IntegrationTestCase): def test_exceed_max_messages(self): msgs = {"messages": [{"role": "user", "content": "x"}] * (1024 + 1)} - payload = self._create_payload(msgs) + payload = _create_chat_payload(msgs) response = self._post_payload(payload) assert "You have exceeded the number of maximum messages" in response.json()["error"]["message"] @@ -195,7 +196,7 @@ class TestVisualizationPluginsApi(IntegrationTestCase): mock_instance = MagicMock() mock_instance.chat.completions.create = AsyncMock(return_value=mock_response) mock_client.return_value = mock_instance - payload = self._create_payload( + payload = _create_chat_payload( { "tools": [ { @@ -234,8 +235,91 @@ class TestVisualizationPluginsApi(IntegrationTestCase): mock_instance = MagicMock() mock_instance.chat.completions.create = AsyncMock(side_effect=MockOpenAIError()) mock_client.return_value = mock_instance - response = self._post_payload(self._create_payload()) + response = self._post_payload(_create_chat_payload()) self._assert_status_code_is(response, 404) body = response.json() assert body["error"]["message"] == "original error message" assert body["error"]["type"] == "api_error" + + +class TestPluginsInferenceServicesConfig(IntegrationTestCase): + """Tests for inference_services config resolution in plugins.""" + + @classmethod + def handle_galaxy_config_kwds(cls, config) -> None: + config["ai_api_key"] = "global_key" + config["ai_api_base_url"] = "http://global-url" + config["ai_model"] = "global_model" + config["visualization_plugins_directory"] = TEST_VISUALIZATION_PLUGINS_DIR + config["inference_services"] = { + "default": { + "model": "default_model", + "api_key": "default_key", + "api_base_url": "http://default-url", + }, + "jupyterlite": { + "model": "jupyterlite_model", + "api_key": "jupyterlite_key", + "api_base_url": "http://jupyterlite-url", + }, + } + + def _post_payload(self, payload=None, anon=False): + return self._post("plugins/jupyterlite/chat/completions", payload, json=True, anon=anon) + + @patch("galaxy.webapps.galaxy.api.plugins.AsyncOpenAI") + def test_plugin_specific_config_used(self, mock_client): + """Plugin-specific inference_services config overrides default and global.""" + mock_response = MagicMock() + mock_response.model_dump.return_value = {"id": "test", "choices": []} + mock_instance = MagicMock() + mock_instance.chat.completions.create = AsyncMock(return_value=mock_response) + mock_client.return_value = mock_instance + + response = self._post_payload(payload=_create_chat_payload()) + self._assert_status_code_is(response, 200) + + call_kwargs = mock_instance.chat.completions.create.call_args.kwargs + assert call_kwargs["model"] == "jupyterlite_model" + client_kwargs = mock_client.call_args.kwargs + assert client_kwargs["api_key"] == "jupyterlite_key" + assert client_kwargs["base_url"] == "http://jupyterlite-url" + + +class TestPluginsInferenceServicesDefault(IntegrationTestCase): + """Tests that inference_services.default is used when no plugin-specific config exists.""" + + @classmethod + def handle_galaxy_config_kwds(cls, config) -> None: + config["ai_api_key"] = "global_key" + config["ai_api_base_url"] = "http://global-url" + config["ai_model"] = "global_model" + config["visualization_plugins_directory"] = TEST_VISUALIZATION_PLUGINS_DIR + config["inference_services"] = { + "default": { + "model": "default_model", + "api_key": "default_key", + "api_base_url": "http://default-url", + }, + } + + def _post_payload(self, payload=None, anon=False): + return self._post("plugins/jupyterlite/chat/completions", payload, json=True, anon=anon) + + @patch("galaxy.webapps.galaxy.api.plugins.AsyncOpenAI") + def test_default_config_fallback(self, mock_client): + """inference_services.default is used when no plugin-specific entry exists.""" + mock_response = MagicMock() + mock_response.model_dump.return_value = {"id": "test", "choices": []} + mock_instance = MagicMock() + mock_instance.chat.completions.create = AsyncMock(return_value=mock_response) + mock_client.return_value = mock_instance + + response = self._post_payload(payload=_create_chat_payload()) + self._assert_status_code_is(response, 200) + + call_kwargs = mock_instance.chat.completions.create.call_args.kwargs + assert call_kwargs["model"] == "default_model" + client_kwargs = mock_client.call_args.kwargs + assert client_kwargs["api_key"] == "default_key" + assert client_kwargs["base_url"] == "http://default-url"