From b0e18e432e3f5544cd0d8a261afe26018dfa43cf Mon Sep 17 00:00:00 2001 From: Robert Resch Date: Thu, 30 Apr 2026 12:44:54 +0200 Subject: [PATCH] Use HassKey for webhook data (#169360) Co-authored-by: Copilot --- homeassistant/components/webhook/__init__.py | 64 +++++++++++++------- tests/components/vegehub/test_sensor.py | 4 -- tests/components/vegehub/test_switch.py | 4 -- 3 files changed, 43 insertions(+), 29 deletions(-) diff --git a/homeassistant/components/webhook/__init__.py b/homeassistant/components/webhook/__init__.py index 0b636367fac2..2fce99bd6cda 100644 --- a/homeassistant/components/webhook/__init__.py +++ b/homeassistant/components/webhook/__init__.py @@ -3,11 +3,12 @@ from __future__ import annotations from collections.abc import Awaitable, Callable, Iterable +from dataclasses import dataclass from http import HTTPStatus from ipaddress import ip_address import logging import secrets -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, cast from aiohttp import StreamReader from aiohttp.hdrs import METH_GET, METH_HEAD, METH_POST, METH_PUT @@ -22,6 +23,7 @@ from homeassistant.helpers.network import get_url, is_cloud_connection from homeassistant.helpers.typing import ConfigType from homeassistant.util import network as network_util from homeassistant.util.aiohttp import MockRequest, MockStreamReader, serialize_response +from homeassistant.util.hass_dict import HassKey _LOGGER = logging.getLogger(__name__) @@ -33,6 +35,22 @@ URL_WEBHOOK_PATH = "/api/webhook/{webhook_id}" CONFIG_SCHEMA = cv.empty_config_schema(DOMAIN) +type HandlerType = Callable[[HomeAssistant, str, Request], Awaitable[Response | None]] + + +@dataclass(frozen=True, slots=True) +class WebhookData: + """Data for a registered webhook.""" + + domain: str + name: str + handler: HandlerType + local_only: bool + allowed_methods: frozenset[str] + + +_HANDLERS: HassKey[dict[str, WebhookData]] = HassKey(DOMAIN) + @callback def async_register( @@ -40,13 +58,13 @@ def async_register( domain: str, name: str, webhook_id: str, - handler: Callable[[HomeAssistant, str, Request], Awaitable[Response | None]], + handler: HandlerType, *, local_only: bool = False, allowed_methods: Iterable[str] | None = None, ) -> None: """Register a webhook.""" - handlers = hass.data.setdefault(DOMAIN, {}) + handlers = hass.data.setdefault(_HANDLERS, {}) if webhook_id in handlers: raise ValueError("Handler is already defined!") @@ -67,19 +85,19 @@ def async_register( # deprecation period has ended and the message was removed. raise TypeError("local_only must be a boolean") - handlers[webhook_id] = { - "domain": domain, - "name": name, - "handler": handler, - "local_only": local_only, - "allowed_methods": allowed_methods, - } + handlers[webhook_id] = WebhookData( + domain=domain, + name=name, + handler=handler, + local_only=local_only, + allowed_methods=allowed_methods, + ) @callback def async_unregister(hass: HomeAssistant, webhook_id: str) -> None: """Remove a webhook.""" - handlers = hass.data.setdefault(DOMAIN, {}) + handlers = hass.data.setdefault(_HANDLERS, {}) handlers.pop(webhook_id, None) @@ -124,7 +142,7 @@ async def async_handle_webhook( hass: HomeAssistant, webhook_id: str, request: Request | MockRequest ) -> Response: """Handle a webhook.""" - handlers: dict[str, dict[str, Any]] = hass.data.setdefault(DOMAIN, {}) + handlers = hass.data.setdefault(_HANDLERS, {}) content_stream: StreamReader | MockStreamReader received_from: str | None @@ -134,6 +152,10 @@ async def async_handle_webhook( received_from += f" ({request.remote})" content_stream = request.content method_name = request.method + if TYPE_CHECKING: + # MockRequest mimics the aiohttp Request interface and is used for + # cloudhooks and webhooks triggered via the WebSocket API. + request = cast(Request, request) else: received_from = request.remote content_stream = request.content @@ -152,7 +174,7 @@ async def async_handle_webhook( _LOGGER.debug("%s", content) return Response(status=HTTPStatus.OK) - if method_name not in webhook["allowed_methods"]: + if method_name not in webhook.allowed_methods: if method_name == METH_HEAD: # Allow websites to verify that the URL exists. return Response(status=HTTPStatus.OK) @@ -160,13 +182,13 @@ async def async_handle_webhook( _LOGGER.warning( "Webhook %s only supports %s methods but %s was received from %s", webhook_id, - ",".join(webhook["allowed_methods"]), + ",".join(webhook.allowed_methods), method_name, received_from, ) return Response(status=HTTPStatus.METHOD_NOT_ALLOWED) - if webhook["local_only"]: + if webhook.local_only: is_local = not (is_cloud_connection(hass) or request.remote is None) if is_local: @@ -186,7 +208,7 @@ async def async_handle_webhook( return Response(status=HTTPStatus.OK) try: - response: Response | None = await webhook["handler"](hass, webhook_id, request) + response = await webhook.handler(hass, webhook_id, request) if response is None: response = Response(status=HTTPStatus.OK) except Exception: @@ -235,14 +257,14 @@ def websocket_list( msg: dict[str, Any], ) -> None: """Return a list of webhooks.""" - handlers = hass.data.setdefault(DOMAIN, {}) + handlers = hass.data.setdefault(_HANDLERS, {}) result = [ { "webhook_id": webhook_id, - "domain": info["domain"], - "name": info["name"], - "local_only": info["local_only"], - "allowed_methods": sorted(info["allowed_methods"]), + "domain": info.domain, + "name": info.name, + "local_only": info.local_only, + "allowed_methods": sorted(info.allowed_methods), } for webhook_id, info in handlers.items() ] diff --git a/tests/components/vegehub/test_sensor.py b/tests/components/vegehub/test_sensor.py index b6b4533c3b95..dfe7eac2edc3 100644 --- a/tests/components/vegehub/test_sensor.py +++ b/tests/components/vegehub/test_sensor.py @@ -44,10 +44,6 @@ async def test_sensor_entities( assert TEST_WEBHOOK_ID in hass.data["webhook"], "Webhook was not registered" - # Verify the webhook handler - webhook_info = hass.data["webhook"][TEST_WEBHOOK_ID] - assert webhook_info["handler"], "Webhook handler is not set" - client = await hass_client_no_auth() resp = await client.post(f"/api/webhook/{TEST_WEBHOOK_ID}", json=UPDATE_DATA) diff --git a/tests/components/vegehub/test_switch.py b/tests/components/vegehub/test_switch.py index ab9768b81490..27bc1d26a32f 100644 --- a/tests/components/vegehub/test_switch.py +++ b/tests/components/vegehub/test_switch.py @@ -46,10 +46,6 @@ async def test_switch_entities( assert TEST_WEBHOOK_ID in hass.data["webhook"], "Webhook was not registered" - # Verify the webhook handler - webhook_info = hass.data["webhook"][TEST_WEBHOOK_ID] - assert webhook_info["handler"], "Webhook handler is not set" - client = await hass_client_no_auth() resp = await client.post(f"/api/webhook/{TEST_WEBHOOK_ID}", json=UPDATE_DATA)