Use HassKey for webhook data (#169360)

Co-authored-by: Copilot <copilot@github.com>
This commit is contained in:
Robert Resch
2026-04-30 12:44:54 +02:00
committed by GitHub
co-authored by Copilot
parent 13d285298c
commit b0e18e432e
3 changed files with 43 additions and 29 deletions
+43 -21
View File
@@ -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()
]
-4
View File
@@ -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)
-4
View File
@@ -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)