diff --git a/homeassistant/components/mcp_server/http.py b/homeassistant/components/mcp_server/http.py index 903a42d0ed6f..7d2488586ef6 100644 --- a/homeassistant/components/mcp_server/http.py +++ b/homeassistant/components/mcp_server/http.py @@ -36,7 +36,7 @@ import logging from typing import get_args from aiohttp import web -from aiohttp.web_exceptions import HTTPBadRequest, HTTPNotFound +from aiohttp.web_exceptions import HTTPBadRequest, HTTPForbidden, HTTPNotFound from aiohttp_sse import sse_response import anyio from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream @@ -48,7 +48,6 @@ from homeassistant.components import conversation from homeassistant.components.http import KEY_HASS, HomeAssistantView from homeassistant.const import CONF_LLM_HASS_API, CONTENT_TYPE_JSON from homeassistant.core import Context, HomeAssistant, callback -from homeassistant.exceptions import Unauthorized from homeassistant.helpers import llm from .const import CONF_ALL_LLM_APIS, CONF_REQUIRE_ADMIN, DOMAIN @@ -113,7 +112,7 @@ def _entry_llm_api_ids( def _validate_admin(request: web.Request, entry: MCPServerConfigEntry) -> None: """Verify the user may use the endpoints serving the configured LLM APIs.""" if entry.data[CONF_REQUIRE_ADMIN] and not request["hass_user"].is_admin: - raise Unauthorized + raise HTTPForbidden @dataclass @@ -367,7 +366,7 @@ class ModelContextProtocolStreamableApiView(HomeAssistantView): """Process JSON-RPC messages for the LLM API identified by api_id.""" hass = request.app[KEY_HASS] if api_id != llm.LLM_API_ASSIST and not request["hass_user"].is_admin: - raise Unauthorized + raise HTTPForbidden if api_id not in {api.id for api in llm.async_get_apis(hass)}: raise HTTPNotFound(text=f"Unknown LLM API '{api_id}'") return await _async_handle_streamable_message( diff --git a/tests/components/mcp_server/test_http.py b/tests/components/mcp_server/test_http.py index 3fa02a2778ea..3b21dc62e582 100644 --- a/tests/components/mcp_server/test_http.py +++ b/tests/components/mcp_server/test_http.py @@ -337,19 +337,31 @@ async def test_options_flow_closes_sessions( assert "Could not find session ID" in response_data +@pytest.mark.parametrize( + ("method", "path"), + [ + pytest.param("GET", SSE_API, id="sse"), + pytest.param( + "POST", MESSAGES_API.format(session_id="session-id"), id="messages" + ), + pytest.param("POST", STREAMABLE_API, id="streamable"), + pytest.param( + "POST", f"{STREAMABLE_API}/{TEST_LLM_API_ID}", id="streamable-api" + ), + ], +) async def test_http_requires_authentication( hass: HomeAssistant, setup_integration: None, hass_client_no_auth: ClientSessionGenerator, + method: str, + path: str, ) -> None: - """Test the SSE endpoint requires authentication.""" + """Test every MCP endpoint requires authentication.""" client = await hass_client_no_auth() - response = await client.get(SSE_API) - assert response.status == HTTPStatus.UNAUTHORIZED - - response = await client.post(MESSAGES_API.format(session_id="session-id")) + response = await client.request(method, path) assert response.status == HTTPStatus.UNAUTHORIZED @@ -1125,7 +1137,7 @@ async def test_streamable_api_id_requires_admin( json=INITIALIZE_MESSAGE, headers={"accept": CONTENT_TYPE_JSON}, ) - assert response.status == HTTPStatus.UNAUTHORIZED + assert response.status == HTTPStatus.FORBIDDEN async def test_streamable_api_id_assist_allows_non_admin( @@ -1164,7 +1176,7 @@ async def test_streamable_api_id_unknown( ("require_admin", "expected_status"), [ pytest.param(False, HTTPStatus.OK, id="not_required"), - pytest.param(True, HTTPStatus.UNAUTHORIZED, id="required"), + pytest.param(True, HTTPStatus.FORBIDDEN, id="required"), ], ) async def test_require_admin_option( @@ -1196,10 +1208,10 @@ async def test_require_admin_blocks_sse_endpoints( client = await hass_client(hass_read_only_access_token) response = await client.get(SSE_API) - assert response.status == HTTPStatus.UNAUTHORIZED + assert response.status == HTTPStatus.FORBIDDEN response = await client.post(MESSAGES_API.format(session_id="session-id")) - assert response.status == HTTPStatus.UNAUTHORIZED + assert response.status == HTTPStatus.FORBIDDEN @pytest.mark.parametrize("require_admin", [True])