mirror of
https://github.com/home-assistant/core.git
synced 2026-10-06 22:38:02 -04:00
Return HTTP 403 for MCP permission failures (#184253)
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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])
|
||||
|
||||
Reference in New Issue
Block a user