mirror of
https://github.com/home-assistant/core.git
synced 2026-09-30 00:03:25 +01:00
Use HassKey for webhook data (#169360)
Co-authored-by: Copilot <copilot@github.com>
This commit is contained in:
@@ -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()
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user