mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
Merge pull request #21986 from dannon/fix/plugin-chat-inference-services
[26.0] Use inference_services config for plugin chat proxy
This commit is contained in:
@@ -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.
|
||||
Enable beta tool formats (yaml, cwl, ...) which is a prerequisite for user defined tools.
|
||||
|
||||
@@ -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.<plugin_name>.<key>
|
||||
2. Default inference: inference_services.default.<key>
|
||||
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)
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user