From 58076228ec7cd712fce193108915147676fa00fc Mon Sep 17 00:00:00 2001 From: Paulus Schoutsen Date: Sun, 30 Aug 2026 07:31:25 +0200 Subject: [PATCH] Add options flow to MCP Server integration (#180629) Co-authored-by: Claude --- .../components/mcp_server/config_flow.py | 119 ++++++++++++++---- .../components/mcp_server/strings.json | 16 +++ tests/components/mcp_server/conftest.py | 17 +++ .../components/mcp_server/test_config_flow.py | 100 +++++++++++++++ tests/components/mcp_server/test_http.py | 52 +++++--- 5 files changed, 265 insertions(+), 39 deletions(-) diff --git a/homeassistant/components/mcp_server/config_flow.py b/homeassistant/components/mcp_server/config_flow.py index 8dfa10a7dfb2..0609a7aaebd8 100644 --- a/homeassistant/components/mcp_server/config_flow.py +++ b/homeassistant/components/mcp_server/config_flow.py @@ -5,8 +5,14 @@ from typing import Any, override import voluptuous as vol -from homeassistant.config_entries import ConfigFlow, ConfigFlowResult +from homeassistant.config_entries import ( + ConfigEntry, + ConfigFlow, + ConfigFlowResult, + OptionsFlow, +) from homeassistant.const import CONF_LLM_HASS_API +from homeassistant.core import HomeAssistant, callback from homeassistant.helpers import llm from homeassistant.helpers.selector import ( SelectOptionDict, @@ -21,50 +27,117 @@ _LOGGER = logging.getLogger(__name__) MORE_INFO_URL = "https://www.home-assistant.io/integrations/mcp_server/#configuration" +def _llm_api_names(hass: HomeAssistant) -> dict[str, str]: + """Return the registered LLM API names keyed by API id.""" + return {api.id: api.name for api in llm.async_get_apis(hass)} + + +def _llm_api_title(llm_apis: dict[str, str], api_ids: list[str]) -> str: + """Return the entry title generated for the selected LLM APIs.""" + return ", ".join(llm_apis[api_id] for api_id in api_ids if api_id in llm_apis) + + +def _selected_llm_apis(entry: ConfigEntry, llm_apis: dict[str, str]) -> list[str]: + """Return the still registered LLM APIs selected by the config entry.""" + api_ids = entry.data.get(CONF_LLM_HASS_API) or [] + if isinstance(api_ids, str): # Old config entries stored a single API + api_ids = [api_ids] + return [api_id for api_id in api_ids if api_id in llm_apis] + + +def _llm_api_schema(llm_apis: dict[str, str], default: list[str]) -> vol.Schema: + """Return the schema for selecting LLM APIs.""" + return vol.Schema( + { + vol.Optional( + CONF_LLM_HASS_API, + default=default, + ): SelectSelector( + SelectSelectorConfig( + options=[ + SelectOptionDict( + label=name, + value=llm_api_id, + ) + for llm_api_id, name in llm_apis.items() + ], + multiple=True, + ) + ), + } + ) + + class ModelContextServerProtocolConfigFlow(ConfigFlow, domain=DOMAIN): """Handle a config flow for Model Context Protocol Server.""" VERSION = 1 + @staticmethod + @callback + @override + def async_get_options_flow( + config_entry: ConfigEntry, + ) -> ModelContextServerProtocolOptionsFlow: + """Create the options flow.""" + return ModelContextServerProtocolOptionsFlow() + @override async def async_step_user( self, user_input: dict[str, Any] | None = None ) -> ConfigFlowResult: """Handle the initial step.""" errors: dict[str, str] = {} - llm_apis = {api.id: api.name for api in llm.async_get_apis(self.hass)} + llm_apis = _llm_api_names(self.hass) if user_input is not None: if not user_input[CONF_LLM_HASS_API]: errors[CONF_LLM_HASS_API] = "llm_api_required" else: return self.async_create_entry( - title=", ".join( - llm_apis[api_id] for api_id in user_input[CONF_LLM_HASS_API] - ), + title=_llm_api_title(llm_apis, user_input[CONF_LLM_HASS_API]), data=user_input, ) return self.async_show_form( step_id="user", - data_schema=vol.Schema( - { - vol.Optional( - CONF_LLM_HASS_API, - default=[llm.LLM_API_ASSIST], - ): SelectSelector( - SelectSelectorConfig( - options=[ - SelectOptionDict( - label=name, - value=llm_api_id, - ) - for llm_api_id, name in llm_apis.items() - ], - multiple=True, - ) - ), + data_schema=_llm_api_schema(llm_apis, [llm.LLM_API_ASSIST]), + description_placeholders={"more_info_url": MORE_INFO_URL}, + errors=errors, + ) + + +class ModelContextServerProtocolOptionsFlow(OptionsFlow): + """Handle an options flow to change the exposed LLM APIs.""" + + async def async_step_init( + self, user_input: dict[str, Any] | None = None + ) -> ConfigFlowResult: + """Handle the options step.""" + errors: dict[str, str] = {} + llm_apis = _llm_api_names(self.hass) + current = _selected_llm_apis(self.config_entry, llm_apis) + if user_input is not None: + if not user_input[CONF_LLM_HASS_API]: + errors[CONF_LLM_HASS_API] = "llm_api_required" + else: + updates: dict[str, Any] = { + "data": {**self.config_entry.data, **user_input} } - ), + # Keep a title the user renamed, only refresh a generated one. + if self.config_entry.title == _llm_api_title(llm_apis, current): + updates["title"] = _llm_api_title( + llm_apis, user_input[CONF_LLM_HASS_API] + ) + self.hass.config_entries.async_update_entry( + self.config_entry, **updates + ) + # An open SSE session keeps serving the APIs it started with. + await self.hass.config_entries.async_reload(self.config_entry.entry_id) + return self.async_create_entry(data={}) + + return self.async_show_form( + step_id="init", + data_schema=_llm_api_schema(llm_apis, current), description_placeholders={"more_info_url": MORE_INFO_URL}, errors=errors, ) diff --git a/homeassistant/components/mcp_server/strings.json b/homeassistant/components/mcp_server/strings.json index 44ef16f414d9..5faaac97026d 100644 --- a/homeassistant/components/mcp_server/strings.json +++ b/homeassistant/components/mcp_server/strings.json @@ -17,5 +17,21 @@ "description": "See the [integration documentation]({more_info_url}) for setup instructions." } } + }, + "options": { + "error": { + "llm_api_required": "[%key:component::mcp_server::config::error::llm_api_required%]" + }, + "step": { + "init": { + "data": { + "llm_hass_api": "[%key:common::config_flow::data::llm_hass_api%]" + }, + "data_description": { + "llm_hass_api": "[%key:component::mcp_server::config::step::user::data_description::llm_hass_api%]" + }, + "description": "[%key:component::mcp_server::config::step::user::description%]" + } + } } } diff --git a/tests/components/mcp_server/conftest.py b/tests/components/mcp_server/conftest.py index a2e98c1121ac..40b82c2a2301 100644 --- a/tests/components/mcp_server/conftest.py +++ b/tests/components/mcp_server/conftest.py @@ -13,6 +13,23 @@ from homeassistant.setup import async_setup_component from tests.common import MockConfigEntry +TEST_LLM_API_ID = "test-api" + + +class MockLLMAPI(llm.API): + """Test LLM API that does not expose any tools.""" + + async def async_get_api_instance( + self, llm_context: llm.LLMContext + ) -> llm.APIInstance: + """Return a test API instance.""" + return llm.APIInstance( + api=self, + api_prompt="Test prompt", + llm_context=llm_context, + tools=[], + ) + @pytest.fixture(autouse=True) async def ensure_homeassistant_loaded(hass: HomeAssistant) -> None: diff --git a/tests/components/mcp_server/test_config_flow.py b/tests/components/mcp_server/test_config_flow.py index d084b6064d3c..68aa099d2284 100644 --- a/tests/components/mcp_server/test_config_flow.py +++ b/tests/components/mcp_server/test_config_flow.py @@ -10,6 +10,11 @@ from homeassistant.components.mcp_server.const import DOMAIN from homeassistant.const import CONF_LLM_HASS_API from homeassistant.core import HomeAssistant from homeassistant.data_entry_flow import FlowResultType +from homeassistant.helpers import llm + +from .conftest import TEST_LLM_API_ID, MockLLMAPI + +from tests.common import MockConfigEntry @pytest.mark.parametrize( @@ -66,3 +71,98 @@ async def test_form_errors( assert result["type"] is FlowResultType.FORM assert result["errors"] == errors + + +async def test_options_flow(hass: HomeAssistant, config_entry: MockConfigEntry) -> None: + """Test changing the LLM APIs in the options flow.""" + llm.async_register_api(hass, MockLLMAPI(hass=hass, id=TEST_LLM_API_ID, name="Test")) + # The title generated for the APIs the entry was created with + hass.config_entries.async_update_entry(config_entry, title="Assist") + assert await hass.config_entries.async_setup(config_entry.entry_id) + + result = await hass.config_entries.options.async_init(config_entry.entry_id) + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "init" + assert not result["errors"] + assert result["data_schema"]({}) == {CONF_LLM_HASS_API: [llm.LLM_API_ASSIST]} + + result = await hass.config_entries.options.async_configure( + result["flow_id"], + {CONF_LLM_HASS_API: [llm.LLM_API_ASSIST, TEST_LLM_API_ID]}, + ) + await hass.async_block_till_done() + + assert result["type"] is FlowResultType.CREATE_ENTRY + assert config_entry.data == { + CONF_LLM_HASS_API: [llm.LLM_API_ASSIST, TEST_LLM_API_ID] + } + assert config_entry.title == "Assist, Test" + + +@pytest.mark.parametrize("llm_hass_api", [llm.LLM_API_ASSIST]) +async def test_options_flow_legacy_single_api( + hass: HomeAssistant, config_entry: MockConfigEntry +) -> None: + """Test the form defaults for an entry that stored a single API as a string.""" + assert await hass.config_entries.async_setup(config_entry.entry_id) + + result = await hass.config_entries.options.async_init(config_entry.entry_id) + + assert result["type"] is FlowResultType.FORM + assert result["data_schema"]({}) == {CONF_LLM_HASS_API: [llm.LLM_API_ASSIST]} + + result = await hass.config_entries.options.async_configure( + result["flow_id"], + {CONF_LLM_HASS_API: [llm.LLM_API_ASSIST]}, + ) + await hass.async_block_till_done() + + assert result["type"] is FlowResultType.CREATE_ENTRY + assert config_entry.data == {CONF_LLM_HASS_API: [llm.LLM_API_ASSIST]} + + +async def test_options_flow_keeps_custom_title( + hass: HomeAssistant, config_entry: MockConfigEntry +) -> None: + """Test the options flow does not overwrite a title the user changed.""" + llm.async_register_api(hass, MockLLMAPI(hass=hass, id=TEST_LLM_API_ID, name="Test")) + hass.config_entries.async_update_entry(config_entry, title="My MCP server") + assert await hass.config_entries.async_setup(config_entry.entry_id) + + result = await hass.config_entries.options.async_init(config_entry.entry_id) + result = await hass.config_entries.options.async_configure( + result["flow_id"], + {CONF_LLM_HASS_API: [TEST_LLM_API_ID]}, + ) + await hass.async_block_till_done() + + assert result["type"] is FlowResultType.CREATE_ENTRY + assert config_entry.data == {CONF_LLM_HASS_API: [TEST_LLM_API_ID]} + assert config_entry.title == "My MCP server" + + +async def test_options_flow_errors( + hass: HomeAssistant, config_entry: MockConfigEntry +) -> None: + """Test the options flow requires at least one LLM API.""" + assert await hass.config_entries.async_setup(config_entry.entry_id) + + result = await hass.config_entries.options.async_init(config_entry.entry_id) + assert result["type"] is FlowResultType.FORM + + result = await hass.config_entries.options.async_configure( + result["flow_id"], + {CONF_LLM_HASS_API: []}, + ) + + assert result["type"] is FlowResultType.FORM + assert result["errors"] == {CONF_LLM_HASS_API: "llm_api_required"} + + result = await hass.config_entries.options.async_configure( + result["flow_id"], + {CONF_LLM_HASS_API: [llm.LLM_API_ASSIST]}, + ) + await hass.async_block_till_done() + + assert result["type"] is FlowResultType.CREATE_ENTRY + assert config_entry.data == {CONF_LLM_HASS_API: [llm.LLM_API_ASSIST]} diff --git a/tests/components/mcp_server/test_http.py b/tests/components/mcp_server/test_http.py index 4a85ae5adf70..d908cc5177b4 100644 --- a/tests/components/mcp_server/test_http.py +++ b/tests/components/mcp_server/test_http.py @@ -42,6 +42,8 @@ from homeassistant.helpers import ( from homeassistant.helpers.httpx_client import create_async_httpx_client from homeassistant.setup import async_setup_component +from .conftest import TEST_LLM_API_ID, MockLLMAPI + from tests.common import MockConfigEntry, setup_test_component_platform from tests.components.light.common import MockLight from tests.typing import ClientSessionGenerator @@ -50,7 +52,6 @@ _LOGGER = logging.getLogger(__name__) TEST_ENTITY = "light.kitchen" SNAPSHOT_RESOURCE_URI = "homeassistant://assist/context-snapshot" -TEST_LLM_API_ID = "test-api" INITIALIZE_MESSAGE = { "jsonrpc": "2.0", "id": "request-id-1", @@ -73,21 +74,6 @@ EXPECTED_PROMPT_ENTITY_DEFINITION = """ """ -class MockLLMAPI(llm.API): - """Test LLM API that does not expose any tools.""" - - async def async_get_api_instance( - self, llm_context: llm.LLMContext - ) -> llm.APIInstance: - """Return a test API instance.""" - return llm.APIInstance( - api=self, - api_prompt="Test prompt", - llm_context=llm_context, - tools=[], - ) - - @pytest.fixture async def setup_integration(hass: HomeAssistant, config_entry: MockConfigEntry) -> None: """Set up the config entry.""" @@ -289,6 +275,40 @@ async def test_http_messages_no_config_entry( assert "Could not find session ID" in response_data +async def test_options_flow_closes_sessions( + hass: HomeAssistant, + setup_integration: None, + config_entry: MockConfigEntry, + hass_client: ClientSessionGenerator, +) -> None: + """Test an open SSE session is closed when the selected APIs change.""" + llm.async_register_api( + hass, MockLLMAPI(hass=hass, id=TEST_LLM_API_ID, name="Test API") + ) + client = await hass_client() + + # Start an SSE session + response = await client.get(SSE_API) + assert response.status == HTTPStatus.OK + reader = sse_response_reader(response) + event, endpoint_url = await anext(reader) + assert event == "endpoint" + + # Change the exposed LLM APIs + result = await hass.config_entries.options.async_init(config_entry.entry_id) + await hass.config_entries.options.async_configure( + result["flow_id"], + {CONF_LLM_HASS_API: [TEST_LLM_API_ID]}, + ) + await hass.async_block_till_done() + + # The session serving the previous APIs is gone + response = await client.post(endpoint_url, json=INITIALIZE_MESSAGE) + assert response.status == HTTPStatus.NOT_FOUND + response_data = await response.text() + assert "Could not find session ID" in response_data + + async def test_http_requires_authentication( hass: HomeAssistant, setup_integration: None,