429 files

This commit is contained in:
Home Assistant Version Control
2026-08-30 02:43:31 +00:00
parent 472c19bceb
commit 963cce81a7
566 changed files with 50311 additions and 47616 deletions
+328 -216
View File
@@ -40,11 +40,12 @@ import secrets
import time
from collections.abc import Callable
from html import escape
from typing import TYPE_CHECKING, TypedDict
from typing import TYPE_CHECKING, Any, TypedDict
from urllib.parse import unquote_plus, urlparse
from aiohttp import web
from homeassistant.components.http import HomeAssistantView
from multidict import MultiDictProxy
from .const import WEBHOOK_AUTH_LEGACY
@@ -86,8 +87,8 @@ PKCE_VERIFIER_MIN = 43
PKCE_VERIFIER_MAX = 128
# SHA-256 → 32 bytes → 43 base64url chars (no padding).
PKCE_S256_CHALLENGE_LEN = 43
_PKCE_VERIFIER_RE = re.compile(r"^[A-Za-z0-9._~-]+$")
_PKCE_CHALLENGE_RE = re.compile(r"^[A-Za-z0-9_-]{43}$")
_PKCE_VERIFIER_RE = re.compile(r"[A-Za-z0-9._~-]+")
_PKCE_CHALLENGE_RE = re.compile(r"[A-Za-z0-9_-]{43}")
# Pending-code dict cap. An attacker spamming /authorize with valid params
# could grow the dict between the prune passes that run on each issuance.
@@ -138,6 +139,13 @@ _LEGACY_PENDING_RESTART_KEY = "ha_mcp_tools_oauth_legacy_pending_restart"
_DOMAIN = "ha_mcp_tools"
# TOP-LEVEL hass.data key holding the fingerprint of the credentials the
# SCOPED-ONLY legacy provider serves (root routes owned by another integration
# this session — see build_unbound_legacy_provider). Separate from
# OAUTH_ROUTE_KEY_FINGERPRINT, which records what the bound ROOT views serve.
_LEGACY_SCOPED_FP_KEY = "ha_mcp_tools_oauth_legacy_scoped_fingerprint"
class LegacyOAuthRouteConflict(RuntimeError):
"""Raised when another integration already owns the root OAuth routes."""
@@ -193,14 +201,6 @@ def bind_legacy_views(
owner = hass.data.get(OAUTH_ROUTE_OWNER_KEY)
if owner is not None and owner != _DOMAIN:
_LOGGER.error(
"HA-MCP: cannot enable legacy OAuth mode -- the Webhook Proxy "
"add-on ('%s') already owns the root /authorize and /token routes "
"in this Home Assistant instance, and Home Assistant cannot "
"release them until it restarts. Stop that add-on and restart "
"Home Assistant, then enable legacy mode again.",
owner,
)
raise LegacyOAuthRouteConflict(owner)
bound_provider = hass.data.get(_LEGACY_PROVIDER_KEY)
@@ -242,30 +242,67 @@ def bind_legacy_views(
return provider, pending_restart
def build_unbound_legacy_provider(
hass: HomeAssistant,
client_id: str,
client_secret: str,
signing_key: bytes | str,
) -> LegacyOAuthProvider:
"""Build a legacy provider for the scoped endpoints only.
Root routes are owned by another integration this session. Nothing is
bound, so no restart bookkeeping applies.
"""
key_bytes = _normalize_signing_key(signing_key)
# The scoped views serve THIS provider from cfg immediately (no restart
# semantics apply) — record its identity so legacy_credentials_active
# reports these working credentials as live (#2213 review).
hass.data[_LEGACY_SCOPED_FP_KEY] = _oauth_route_fingerprint(
client_id, client_secret, key_bytes
)
return LegacyOAuthProvider(
client_id=client_id,
client_secret=client_secret,
signing_key=key_bytes,
active_mode_getter=lambda: _live_auth_mode(hass),
)
def clear_scoped_legacy_credentials(hass: HomeAssistant) -> None:
"""Forget credentials when no scoped-only legacy provider is live."""
hass.data.pop(_LEGACY_SCOPED_FP_KEY, None)
def legacy_credentials_active(
hass: HomeAssistant,
client_id: str,
client_secret: str,
signing_key: bytes | str,
) -> bool:
"""Whether the bound root views currently serve exactly these credentials.
"""Whether a live legacy surface currently serves exactly these credentials.
False while a credential rotation is pending a restart (the bound provider
keeps the previous identity until then — see :func:`bind_legacy_views`),
when another integration owns the routes, or when legacy OAuth was never
bound this session. Callers use this to withhold rotated credentials from
surfaces a still-valid old-identity token can read — the admin startup log
in particular, which is reachable through the server's own log tools
(review finding on #1880).
keeps the previous identity until then — see :func:`bind_legacy_views`) or
when legacy OAuth was never bound this session. When another integration
owns the root routes, the scoped endpoints may still serve via the unbound
provider — True then iff the scoped provider's fingerprint matches these
credentials (#2213 review). Callers use this to withhold rotated
credentials from surfaces a still-valid old-identity token can read — the
admin startup log in particular, which is reachable through the server's
own log tools (review finding on #1880).
"""
if hass.data.get(OAUTH_ROUTE_OWNER_KEY) != _DOMAIN:
return False
bound = hass.data.get(OAUTH_ROUTE_KEY_FINGERPRINT)
if not isinstance(bound, str):
return False
current = _oauth_route_fingerprint(
client_id, client_secret, _normalize_signing_key(signing_key)
)
if hass.data.get(OAUTH_ROUTE_OWNER_KEY) != _DOMAIN:
# Root routes owned elsewhere (webhook-proxy add-on). The scoped
# endpoints may still serve these credentials via the unbound provider
# (#2213 review) — surface them when its fingerprint matches.
scoped = hass.data.get(_LEGACY_SCOPED_FP_KEY)
return isinstance(scoped, str) and hmac.compare_digest(scoped, current)
bound = hass.data.get(OAUTH_ROUTE_KEY_FINGERPRINT)
if not isinstance(bound, str):
return False
return hmac.compare_digest(bound, current)
@@ -330,6 +367,12 @@ def _is_loopback_host(hostname: str) -> bool:
return False
# RFC 3986 §3.2 authority charset (unreserved / pct-encoded / sub-delims /
# ':' '@' and IPv6 brackets). Anything outside it (a backslash, a space, raw
# unicode) is an illegal authority that downstream URL builders reject.
_AUTHORITY_CHARS_RE = re.compile(r"[A-Za-z0-9._~%!$&'()*+,;=:@\[\]-]*")
def _is_valid_redirect_uri(redirect_uri: str) -> bool:
"""Spec-floor validation for OAuth redirect_uri: an https:// URL — or an
http:// loopback URL (RFC 8252 §7.3, for native/CLI clients) — with a
@@ -349,6 +392,13 @@ def _is_valid_redirect_uri(redirect_uri: str) -> bool:
return False
if not parsed.hostname:
return False
if not _AUTHORITY_CHARS_RE.fullmatch(parsed.netloc):
# Same contract as the .port access above: urlparse and yarl split
# authorities differently (a backslash before '@', a zero-width
# character in the host), so an RFC 3986-illegal authority must fail
# HERE rather than escape as a 500 out of _redirect_with
# (#2219 codex review).
return False
if parsed.scheme == "http":
# Plain http only for loopback callbacks (native-client flow).
if not _is_loopback_host(parsed.hostname):
@@ -460,7 +510,7 @@ class PKCECodeStore:
# rejected explicitly rather than silently hashing junk.
if not (PKCE_VERIFIER_MIN <= len(code_verifier) <= PKCE_VERIFIER_MAX):
return False
if not _PKCE_VERIFIER_RE.match(code_verifier):
if not _PKCE_VERIFIER_RE.fullmatch(code_verifier):
return False
entry = self._codes.pop(code, None)
if entry is None:
@@ -606,6 +656,37 @@ class LegacyOAuthProvider:
"""Verify PKCE S256 + one-shot consume a code (see PKCECodeStore)."""
return self._code_store.consume_code(code, redirect_uri, code_verifier)
def validate_authorize_params(
self,
*,
response_type: str,
client_id: str,
redirect_uri: str,
code_challenge: str,
code_challenge_method: str,
) -> web.Response | None:
"""Return a 400 web.Response if any /authorize param is invalid, or
None if all checks pass. Centralized so GET and POST share identical
validation — the POST path explicitly re-validates the hidden form
fields rather than trusting them."""
if response_type != "code":
return _text_error(400, "unsupported_response_type")
if code_challenge_method != "S256":
return _text_error(400, "invalid code_challenge_method (S256 required)")
if not _PKCE_CHALLENGE_RE.fullmatch(code_challenge):
return _text_error(
400, "invalid code_challenge (must be 43-char base64url)"
)
if client_id != self.client_id:
return _text_error(400, "invalid client_id", restart_hint=True)
if not _is_valid_redirect_uri(redirect_uri):
return _text_error(
400,
"redirect_uri must be an https:// URL (or an http:// loopback "
"URL) with a valid host and port",
)
return None
# -----------------------------------------------------------------
# Client authentication
# -----------------------------------------------------------------
@@ -621,58 +702,49 @@ class LegacyOAuthProvider:
# ---------------------------------------------------------------------------
# Views (root /authorize + /token)
# Shared handlers + views (root /authorize + /token)
# ---------------------------------------------------------------------------
class AuthorizeView(HomeAssistantView):
"""OAuth /authorize endpoint with a minimal consent page."""
def _redirect_with_params(redirect_uri: str, **params: str) -> web.Response:
# yarl ships with aiohttp and handles existing-query-string merging
# plus parameter encoding correctly — safer than hand-rolling.
import yarl
requires_auth = False
url = AUTHORIZE_PATH
name = "ha_mcp_tools:oauth:authorize"
url = yarl.URL(redirect_uri).update_query(params)
return web.Response(status=302, headers={"Location": str(url)})
def __init__(self, provider: LegacyOAuthProvider) -> None:
self._provider = provider
@staticmethod
def _redirect_with(redirect_uri: str, **params: str) -> web.Response:
# yarl ships with aiohttp and handles existing-query-string merging
# plus parameter encoding correctly — safer than hand-rolling.
import yarl
async def handle_legacy_authorize_get(
provider: LegacyOAuthProvider, request: web.Request
) -> web.Response:
"""Serve the legacy consent page (shared by the root and scoped routes).
url = yarl.URL(redirect_uri).update_query(params)
return web.Response(status=302, headers={"Location": str(url)})
The form posts back to ``request.path`` so the scoped
``{OAUTH_BASE}/authorize`` route and the legacy root ``/authorize`` route
each round-trip through themselves.
"""
params = request.query
client_id = params.get("client_id", "")
redirect_uri = params.get("redirect_uri", "")
state = params.get("state", "")
code_challenge = params.get("code_challenge", "")
code_challenge_method = params.get("code_challenge_method", "")
response_type = params.get("response_type", "")
async def get(self, request: web.Request) -> web.Response:
if not self._provider.is_active():
# Serve ONLY while legacy is the live mode. Both ha_auth/none (HA
# core or the secret URL is the authority) and a not-yet-live
# entry mean this root view must not serve — HA can't rebind or
# drop it without a restart, so a mode switch away leaves it
# bound. Refuse it.
return _text_error(404, "not found")
params = request.query
client_id = params.get("client_id", "")
redirect_uri = params.get("redirect_uri", "")
state = params.get("state", "")
code_challenge = params.get("code_challenge", "")
code_challenge_method = params.get("code_challenge_method", "")
response_type = params.get("response_type", "")
err = provider.validate_authorize_params(
response_type=response_type,
client_id=client_id,
redirect_uri=redirect_uri,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
)
if err is not None:
return err
err = self._validate_authorize_params(
response_type=response_type,
client_id=client_id,
redirect_uri=redirect_uri,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
)
if err is not None:
return err
# Render minimal consent page. Showing the redirect_uri lets the user
# verify the flow goes back to a domain they recognize.
html = f"""<!DOCTYPE html>
# Render minimal consent page. Showing the redirect_uri lets the user
# verify the flow goes back to a domain they recognize.
html = f"""<!DOCTYPE html>
<html>
<head>
<meta charset="utf-8">
@@ -690,7 +762,7 @@ class AuthorizeView(HomeAssistantView):
<p>An MCP client is requesting access to your Home Assistant MCP server.</p>
<p>It will redirect to:<br><code>{escape(redirect_uri)}</code></p>
<p>Only allow this if you started this connection yourself.</p>
<form method="POST" action="{AUTHORIZE_PATH}">
<form method="POST" action="{escape(request.path)}">
<input type="hidden" name="client_id" value="{escape(client_id)}">
<input type="hidden" name="redirect_uri" value="{escape(redirect_uri)}">
<input type="hidden" name="state" value="{escape(state)}">
@@ -700,83 +772,197 @@ class AuthorizeView(HomeAssistantView):
</form>
</body>
</html>"""
return web.Response(text=html, content_type="text/html")
return web.Response(text=html, content_type="text/html")
async def handle_legacy_authorize_post(
provider: LegacyOAuthProvider, request: web.Request
) -> web.Response:
"""Consume the consent form and redirect with a code (shared by both routes)."""
data = await read_form(request)
if data is None:
return _text_error(400, "Invalid form submission")
action = str(data.get("action", ""))
client_id = str(data.get("client_id", ""))
redirect_uri = str(data.get("redirect_uri", ""))
state = str(data.get("state", ""))
code_challenge = str(data.get("code_challenge", ""))
# Re-validate everything from the form — never trust hidden fields.
# response_type/method aren't carried on the POST so we hard-code the
# spec values here; the validator still applies all the same rules to
# the user-influenceable fields.
err = provider.validate_authorize_params(
response_type="code",
client_id=client_id,
redirect_uri=redirect_uri,
code_challenge=code_challenge,
code_challenge_method="S256",
)
if err is not None:
return err
# RFC 9207: every authorization response — success or error — names the
# issuer that produced it, so a client registered with several
# authorization servers cannot be fed a response minted by another one.
iss = _issuer_for(request)
if action == "deny":
return _redirect_with_params(
redirect_uri, error="access_denied", state=state, iss=iss
)
if action != "approve":
return _text_error(400, "invalid action")
code = provider.issue_code(redirect_uri, code_challenge)
if code is None:
# Pending-code store at cap → signal back per RFC 6749 §4.1.2.1
# instead of silently failing.
return _redirect_with_params(
redirect_uri, error="temporarily_unavailable", state=state, iss=iss
)
return _redirect_with_params(redirect_uri, code=code, state=state, iss=iss)
async def read_form(request: web.Request) -> MultiDictProxy[Any] | None:
"""Parse a form body, or None when the request body is undecodable.
aiohttp raises LookupError — NOT ValueError — when the Content-Type names
an unknown charset (``application/x-www-form-urlencoded; charset=nope``),
and every caller here is an ANONYMOUS view, so an unguarded parse turns
one header into a 500 with a traceback (#2219 review round 3).
"""
try:
return await request.post()
except (ValueError, LookupError):
return None
def _extract_client_creds(
request: web.Request, form: dict
) -> list[tuple[str | None, str | None]]:
"""Candidate client_id/secret pairs from Basic auth OR the form body.
RFC 6749 §2.3.1: client_secret_basic values are form-urlencoded before
base64, so the decoded candidate applies unquote_PLUS ("+" means space in
that encoding). Many clients skip the encoding step, though, and custom
credential overrides may contain a literal "+" or "%XX" that decoding
would mangle (#2218 review, mirrored from the proxy) — so the raw split
is offered as a second candidate and either may authenticate. Both are
no-ops for the generated credentials (URL-safe alphabets).
"""
header = request.headers.get("Authorization", "")
if header.lower().startswith("basic "):
try:
decoded = base64.b64decode(header[6:].strip(), validate=True).decode(
"utf-8"
)
except (ValueError, UnicodeDecodeError, binascii.Error):
return []
if ":" not in decoded:
return []
cid, _, sec = decoded.partition(":")
candidates: list[tuple[str | None, str | None]] = [
(unquote_plus(cid), unquote_plus(sec))
]
if (cid, sec) != candidates[0]:
candidates.append((cid, sec))
return candidates
# str()-wrapped like every sibling site: request.post() yields
# str | bytes | FileField, and the non-str shapes are truthy, so they
# would slip past authenticate_client's emptiness guard and raise
# AttributeError on .encode() (#2219 review round 3).
return [(str(form.get("client_id", "")), str(form.get("client_secret", "")))]
async def _handle_authorization_code(
provider: LegacyOAuthProvider, form: dict
) -> web.Response:
code = str(form.get("code", ""))
redirect_uri = str(form.get("redirect_uri", ""))
code_verifier = str(form.get("code_verifier", ""))
if not (code and redirect_uri and code_verifier):
return _json_error("invalid_request", 400)
if not provider.consume_code(code, redirect_uri, code_verifier):
return _json_error("invalid_grant", 400)
return web.json_response(
{
"access_token": provider.issue_access_token(),
"token_type": "Bearer",
"expires_in": ACCESS_TOKEN_TTL,
"refresh_token": provider.issue_refresh_token(),
},
headers=_TOKEN_RESPONSE_HEADERS,
)
async def _handle_refresh(provider: LegacyOAuthProvider, form: dict) -> web.Response:
refresh = str(form.get("refresh_token", ""))
if not refresh or not provider.validate_refresh_token(refresh):
return _json_error("invalid_grant", 400)
return web.json_response(
{
"access_token": provider.issue_access_token(),
"token_type": "Bearer",
"expires_in": ACCESS_TOKEN_TTL,
"refresh_token": provider.issue_refresh_token(),
},
headers=_TOKEN_RESPONSE_HEADERS,
)
async def handle_legacy_token_post(
provider: LegacyOAuthProvider, request: web.Request
) -> web.Response:
"""Token endpoint body (shared by the root and scoped routes)."""
raw_form = await read_form(request)
if raw_form is None:
return _json_error("invalid_request", 400)
form = dict(raw_form)
candidates = _extract_client_creds(request, form)
if not any(
provider.authenticate_client(client_id, client_secret)
for client_id, client_secret in candidates
):
return _json_error(
"invalid_client",
401,
headers={"WWW-Authenticate": 'Basic realm="HA-MCP OAuth"'},
restart_hint=True,
)
grant_type = form.get("grant_type", "")
if grant_type == "authorization_code":
return await _handle_authorization_code(provider, form)
if grant_type == "refresh_token":
return await _handle_refresh(provider, form)
return _json_error("unsupported_grant_type", 400)
class AuthorizeView(HomeAssistantView):
"""OAuth /authorize endpoint with a minimal consent page."""
requires_auth = False
url = AUTHORIZE_PATH
name = "ha_mcp_tools:oauth:authorize"
def __init__(self, provider: LegacyOAuthProvider) -> None:
self._provider = provider
async def get(self, request: web.Request) -> web.Response:
if not self._provider.is_active():
# Serve ONLY while legacy is the live mode. Both ha_auth/none (HA
# core or the secret URL is the authority) and a not-yet-live
# entry mean this root view must not serve — HA can't rebind or
# drop it without a restart, so a mode switch away leaves it
# bound. Refuse it.
return _text_error(404, "not found")
return await handle_legacy_authorize_get(self._provider, request)
async def post(self, request: web.Request) -> web.Response:
if not self._provider.is_active():
return _text_error(404, "not found")
data = await request.post()
action = str(data.get("action", ""))
client_id = str(data.get("client_id", ""))
redirect_uri = str(data.get("redirect_uri", ""))
state = str(data.get("state", ""))
code_challenge = str(data.get("code_challenge", ""))
# Re-validate everything from the form — never trust hidden fields.
# response_type/method aren't carried on the POST so we hard-code the
# spec values here; the validator still applies all the same rules to
# the user-influenceable fields.
err = self._validate_authorize_params(
response_type="code",
client_id=client_id,
redirect_uri=redirect_uri,
code_challenge=code_challenge,
code_challenge_method="S256",
)
if err is not None:
return err
# RFC 9207: every authorization response — success or error — names the
# issuer that produced it, so a client registered with several
# authorization servers cannot be fed a response minted by another one.
iss = _issuer_for(request)
if action == "deny":
return self._redirect_with(
redirect_uri, error="access_denied", state=state, iss=iss
)
if action != "approve":
return _text_error(400, "invalid action")
code = self._provider.issue_code(redirect_uri, code_challenge)
if code is None:
# Pending-code store at cap → signal back per RFC 6749 §4.1.2.1
# instead of silently failing.
return self._redirect_with(
redirect_uri, error="temporarily_unavailable", state=state, iss=iss
)
return self._redirect_with(redirect_uri, code=code, state=state, iss=iss)
def _validate_authorize_params(
self,
*,
response_type: str,
client_id: str,
redirect_uri: str,
code_challenge: str,
code_challenge_method: str,
) -> web.Response | None:
"""Return a 400 web.Response if any /authorize param is invalid, or
None if all checks pass. Centralized so GET and POST share identical
validation — the POST path explicitly re-validates the hidden form
fields rather than trusting them."""
if response_type != "code":
return _text_error(400, "unsupported_response_type")
if code_challenge_method != "S256":
return _text_error(400, "invalid code_challenge_method (S256 required)")
if not _PKCE_CHALLENGE_RE.match(code_challenge):
return _text_error(
400, "invalid code_challenge (must be 43-char base64url)"
)
if client_id != self._provider.client_id:
return _text_error(400, "invalid client_id", restart_hint=True)
if not _is_valid_redirect_uri(redirect_uri):
return _text_error(
400,
"redirect_uri must be an https:// URL (or an http:// loopback "
"URL) with a valid host and port",
)
return None
return await handle_legacy_authorize_post(self._provider, request)
class TokenView(HomeAssistantView):
@@ -790,81 +976,7 @@ class TokenView(HomeAssistantView):
def __init__(self, provider: LegacyOAuthProvider) -> None:
self._provider = provider
@staticmethod
def _extract_client_creds(
request: web.Request, form: dict
) -> tuple[str | None, str | None]:
"""Pull client_id/secret from Basic auth header OR form body."""
header = request.headers.get("Authorization", "")
if header.lower().startswith("basic "):
try:
decoded = base64.b64decode(header[6:].strip(), validate=True).decode(
"utf-8"
)
except (ValueError, UnicodeDecodeError, binascii.Error):
return None, None
if ":" in decoded:
cid, _, sec = decoded.partition(":")
# RFC 6749 §2.3.1: client_secret_basic values are
# application/x-www-form-urlencoded before base64, so decode
# them back with unquote_PLUS ("+" means space in that
# encoding — plain unquote would leave it literal). A no-op for
# the generated credentials (URL-safe alphabets, nothing to
# decode) but required for custom overrides containing reserved
# characters — matches the form-body path below, which aiohttp
# also form-decodes ("+" → space).
return unquote_plus(cid), unquote_plus(sec)
return None, None
return form.get("client_id"), form.get("client_secret")
async def post(self, request: web.Request) -> web.Response:
if not self._provider.is_active():
return _json_not_found()
form = dict(await request.post())
client_id, client_secret = self._extract_client_creds(request, form)
if not self._provider.authenticate_client(client_id, client_secret):
return _json_error(
"invalid_client",
401,
headers={"WWW-Authenticate": 'Basic realm="HA-MCP OAuth"'},
restart_hint=True,
)
grant_type = form.get("grant_type", "")
if grant_type == "authorization_code":
return await self._handle_authorization_code(form)
if grant_type == "refresh_token":
return await self._handle_refresh(form)
return _json_error("unsupported_grant_type", 400)
async def _handle_authorization_code(self, form: dict) -> web.Response:
code = str(form.get("code", ""))
redirect_uri = str(form.get("redirect_uri", ""))
code_verifier = str(form.get("code_verifier", ""))
if not (code and redirect_uri and code_verifier):
return _json_error("invalid_request", 400)
if not self._provider.consume_code(code, redirect_uri, code_verifier):
return _json_error("invalid_grant", 400)
return web.json_response(
{
"access_token": self._provider.issue_access_token(),
"token_type": "Bearer",
"expires_in": ACCESS_TOKEN_TTL,
"refresh_token": self._provider.issue_refresh_token(),
},
headers=_TOKEN_RESPONSE_HEADERS,
)
async def _handle_refresh(self, form: dict) -> web.Response:
refresh = str(form.get("refresh_token", ""))
if not refresh or not self._provider.validate_refresh_token(refresh):
return _json_error("invalid_grant", 400)
return web.json_response(
{
"access_token": self._provider.issue_access_token(),
"token_type": "Bearer",
"expires_in": ACCESS_TOKEN_TTL,
"refresh_token": self._provider.issue_refresh_token(),
},
headers=_TOKEN_RESPONSE_HEADERS,
)
return await handle_legacy_token_post(self._provider, request)