Support PKCE S256 in OAuth server (#181957)

Co-authored-by: Paulus Schoutsen <balloob@gmail.com>
Co-authored-by: Simon Lamon <32477463+silamon@users.noreply.github.com>
Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
Allen Porter
2026-09-26 10:21:45 -04:00
committed by GitHub
co-authored by Paulus Schoutsen Simon Lamon Claude Opus 5.5
parent f45151cfd0
commit 7398141b4c
5 changed files with 366 additions and 32 deletions
+2
View File
@@ -28,6 +28,8 @@ class AuthFlowContext(FlowContext, total=False):
ip_address: IPv4Address | IPv6Address
redirect_uri: str
code_challenge: str
code_challenge_method: str
class AuthFlowResult(FlowResult[AuthFlowContext, tuple[str, str]], total=False):
+97 -19
View File
@@ -14,7 +14,8 @@ Exchange the authorization code retrieved from the login flow for tokens.
{
"client_id": "https://hassbian.local:8123/",
"grant_type": "authorization_code",
"code": "411ee2f916e648d691e937ae9344681e"
"code": "411ee2f916e648d691e937ae9344681e",
"code_verifier": "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"
}
Return value will be the access and refresh tokens. The access token will have
@@ -124,11 +125,16 @@ as part of a config flow.
"""
import asyncio
import base64
from collections.abc import Callable
from dataclasses import dataclass
from datetime import datetime, timedelta
import hashlib
import hmac
from http import HTTPStatus
from logging import getLogger
from typing import Any, cast
import re
from typing import Any, Protocol, cast
import uuid
from aiohttp import web
@@ -162,8 +168,31 @@ from . import indieauth, login_flow, mfa_setup_flow
DOMAIN = "auth"
type StoreResultType = Callable[[str, Credentials], str]
type RetrieveResultType = Callable[[str, str], Credentials | None]
@dataclass(slots=True)
class AuthCodeEntry:
"""Entry stored in the auth code store."""
created: datetime
credentials: Credentials
code_challenge: str | None = None
code_challenge_method: str | None = None
class StoreResultType(Protocol):
"""Protocol for storing auth flow results."""
def __call__(
self,
client_id: str,
result: Credentials,
code_challenge: str | None = None,
code_challenge_method: str | None = None,
) -> str:
"""Store flow result and return a code to retrieve it."""
type RetrieveResultType = Callable[[str, str], AuthCodeEntry | None]
DATA_STORE: HassKey[StoreResultType] = HassKey(DOMAIN)
CONFIG_SCHEMA = cv.empty_config_schema(DOMAIN)
@@ -231,6 +260,19 @@ class RevokeTokenView(HomeAssistantView):
return web.Response(status=HTTPStatus.OK)
# RFC 7636 4.1: code_verifier is 43-128 unreserved characters.
_CODE_VERIFIER_RE = re.compile(r"^[A-Za-z0-9._~-]{43,128}\Z")
def _verify_code_verifier(code_verifier: str, code_challenge: str) -> bool:
"""Verify code_verifier against code_challenge per RFC 7636 (S256)."""
if not _CODE_VERIFIER_RE.match(code_verifier):
return False
hashed = hashlib.sha256(code_verifier.encode("ascii")).digest()
computed_challenge = base64.urlsafe_b64encode(hashed).decode("ascii").rstrip("=")
return hmac.compare_digest(computed_challenge, code_challenge)
class TokenView(HomeAssistantView):
"""View to issue tokens."""
@@ -290,14 +332,43 @@ class TokenView(HomeAssistantView):
status_code=HTTPStatus.BAD_REQUEST,
)
credential = self._retrieve_auth(client_id, code)
entry = self._retrieve_auth(client_id, code)
if credential is None or not isinstance(credential, Credentials):
if entry is None:
return self.json(
{"error": "invalid_request", "error_description": "Invalid code"},
status_code=HTTPStatus.BAD_REQUEST,
)
if entry.code_challenge is not None:
if not (code_verifier := data.get("code_verifier")):
return self.json(
{
"error": "invalid_request",
"error_description": "Code verifier required",
},
status_code=HTTPStatus.BAD_REQUEST,
)
if not _verify_code_verifier(code_verifier, entry.code_challenge):
return self.json(
{
"error": "invalid_grant",
"error_description": "Invalid code verifier",
},
status_code=HTTPStatus.BAD_REQUEST,
)
elif "code_verifier" in data:
return self.json(
{
"error": "invalid_request",
"error_description": (
"Code verifier provided but no code challenge was present"
),
},
status_code=HTTPStatus.BAD_REQUEST,
)
credential = entry.credentials
user = await hass.auth.async_get_or_create_user(credential)
if user_access_error := async_user_not_allowed_do_auth(hass, user):
@@ -421,12 +492,12 @@ class LinkUserView(HomeAssistantView):
hass = request.app[KEY_HASS]
user: User = request["hass_user"]
credentials = self._retrieve_credentials(data["client_id"], data["code"])
entry = self._retrieve_credentials(data["client_id"], data["code"])
if credentials is None:
if entry is None:
return self.json_message("Invalid code", status_code=HTTPStatus.BAD_REQUEST)
linked_user = await hass.auth.async_get_user_by_credentials(credentials)
linked_user = await hass.auth.async_get_user_by_credentials(entry.credentials)
if linked_user != user and linked_user is not None:
return self.json_message(
"Credential already linked", status_code=HTTPStatus.BAD_REQUEST
@@ -434,44 +505,51 @@ class LinkUserView(HomeAssistantView):
# No-op if credential is already linked to the user it will be linked to
if linked_user != user:
await hass.auth.async_link_user(user, credentials)
await hass.auth.async_link_user(user, entry.credentials)
return self.json_message("User linked")
@callback
def _create_auth_code_store() -> tuple[StoreResultType, RetrieveResultType]:
"""Create an in memory store."""
temp_results: dict[tuple[str, str], tuple[datetime, Credentials]] = {}
temp_results: dict[tuple[str, str], AuthCodeEntry] = {}
@callback
def store_result(client_id: str, result: Credentials) -> str:
def store_result(
client_id: str,
result: Credentials,
code_challenge: str | None = None,
code_challenge_method: str | None = None,
) -> str:
"""Store flow result and return a code to retrieve it."""
if not isinstance(result, Credentials):
raise TypeError("result has to be a Credentials instance")
code = uuid.uuid4().hex
temp_results[(client_id, code)] = (
dt_util.utcnow(),
result,
temp_results[(client_id, code)] = AuthCodeEntry(
created=dt_util.utcnow(),
credentials=result,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
)
return code
@callback
def retrieve_result(client_id: str, code: str) -> Credentials | None:
def retrieve_result(client_id: str, code: str) -> AuthCodeEntry | None:
"""Retrieve flow result."""
key = (client_id, code)
if key not in temp_results:
return None
created, result = temp_results.pop(key)
entry = temp_results.pop(key)
# OAuth 4.2.1
# The authorization code MUST expire shortly after it is issued to
# mitigate the risk of leaks. A maximum authorization code lifetime of
# 10 minutes is RECOMMENDED.
if dt_util.utcnow() - created < timedelta(minutes=10):
return result
if dt_util.utcnow() - entry.created < timedelta(minutes=10):
return entry
return None
+43 -12
View File
@@ -22,13 +22,18 @@ Pass in parameter 'client_id' and 'redirect_url' validate by indieauth.
Pass in parameter 'handler' to specify the auth provider to use. Auth providers
are identified by type and id.
Pass in optional parameters 'code_challenge' and 'code_challenge_method' for
PKCE (RFC 7636). The only supported method is 'S256'.
The default 'type' is 'authorize'.
{
"client_id": "https://hassbian.local:8123/",
"handler": ["local_provider", null],
"redirect_url": "https://hassbian.local:8123/",
"type': "authorize"
"type': "authorize",
"code_challenge": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM",
"code_challenge_method": "S256"
}
Return value will be a step in a data entry flow. See the docs for data entry
@@ -64,7 +69,6 @@ an authorization code.
}
"""
from collections.abc import Callable
from http import HTTPStatus
from ipaddress import ip_address
from typing import TYPE_CHECKING, Any, cast
@@ -74,7 +78,7 @@ import probatio
from homeassistant import data_entry_flow
from homeassistant.auth import AuthManagerFlowManager, InvalidAuthError
from homeassistant.auth.models import AuthFlowContext, AuthFlowResult, Credentials
from homeassistant.auth.models import AuthFlowContext, AuthFlowResult
from homeassistant.components import onboarding
from homeassistant.components.http import KEY_HASS
from homeassistant.components.http.auth import async_user_not_allowed_do_auth
@@ -104,9 +108,7 @@ if TYPE_CHECKING:
@callback
def async_setup(
hass: HomeAssistant, store_result: Callable[[str, Credentials], str]
) -> None:
def async_setup(hass: HomeAssistant, store_result: StoreResultType) -> None:
"""Component to allow users to login."""
hass.http.register_view(WellKnownOAuthInfoView)
hass.http.register_view(WellKnownProtectedResourceView)
@@ -142,6 +144,7 @@ class WellKnownOAuthInfoView(HomeAssistantView):
# This flag advertises that support
# (draft-ietf-oauth-client-id-metadata-document).
"client_id_metadata_document_supported": True,
"code_challenge_methods_supported": ["S256"],
"response_types_supported": ["code"],
"service_documentation": (
"https://developers.home-assistant.io/docs/auth_api"
@@ -306,7 +309,7 @@ class LoginFlowBaseView(HomeAssistantView):
return self.json_message("Invalid redirect URI", HTTPStatus.FORBIDDEN)
result.pop("data")
result.pop("context")
context = result.pop("context")
result_obj = result.pop("result")
@@ -322,7 +325,12 @@ class LoginFlowBaseView(HomeAssistantView):
process_success_login(request)
# We overwrite the Credentials object with the string code to retrieve it.
result["result"] = self._store_result(client_id, result_obj) # type: ignore[typeddict-item]
result["result"] = self._store_result(
client_id,
result_obj,
code_challenge=context.get("code_challenge"),
code_challenge_method=context.get("code_challenge_method"),
) # type: ignore[typeddict-item]
return self.json(result)
@@ -347,6 +355,11 @@ class LoginFlowIndexView(LoginFlowBaseView):
probatio.Coerce(tuple),
),
probatio.Required("redirect_uri"): str,
# S256 challenges are always 43 unpadded base64url characters.
probatio.Optional("code_challenge"): probatio.Match(
r"^[A-Za-z0-9_-]{43}\Z"
),
probatio.Optional("code_challenge_method"): str,
probatio.Optional(
"type", default="authorize"
): str, # not used, kept for backwards compatibility
@@ -362,15 +375,33 @@ class LoginFlowIndexView(LoginFlowBaseView):
if not indieauth.verify_client_id(client_id):
return self.json_message("Invalid client id", HTTPStatus.BAD_REQUEST)
code_challenge = data.get("code_challenge")
code_challenge_method = data.get("code_challenge_method")
if code_challenge_method is not None and not code_challenge:
return self.json_message(
"code_challenge required when code_challenge_method is provided",
HTTPStatus.BAD_REQUEST,
)
# RFC 7636 4.3: the method defaults to "plain", which is not supported.
if code_challenge is not None and code_challenge_method != "S256":
return self.json_message(
"Transform algorithm not supported", HTTPStatus.BAD_REQUEST
)
handler: tuple[str, str] = tuple(data["handler"])
flow_context = AuthFlowContext(
ip_address=ip_address(request.remote), # type: ignore[arg-type]
redirect_uri=redirect_uri,
)
if code_challenge and code_challenge_method:
flow_context["code_challenge"] = code_challenge
flow_context["code_challenge_method"] = code_challenge_method
try:
result = await self._flow_mgr.async_init(
handler,
context=AuthFlowContext(
ip_address=ip_address(request.remote), # type: ignore[arg-type]
redirect_uri=redirect_uri,
),
context=flow_context,
)
except data_entry_flow.UnknownHandler:
return self.json_message("Invalid handler specified", HTTPStatus.NOT_FOUND)
+156 -1
View File
@@ -3,8 +3,10 @@
from datetime import timedelta
from http import HTTPStatus
import logging
from typing import Any
from unittest.mock import patch
from aiohttp.test_utils import TestClient
from freezegun.api import FrozenDateTimeFactory
import pytest
@@ -189,7 +191,9 @@ def test_auth_code_store_expiration(
code = store(client_id, mock_credential)
freezer.move_to(now + timedelta(minutes=9, seconds=59))
assert retrieve(client_id, code) == mock_credential
entry = retrieve(client_id, code)
assert entry is not None
assert entry.credentials == mock_credential
def test_auth_code_store_requires_credentials(mock_credential) -> None:
@@ -761,3 +765,154 @@ async def test_ws_refresh_token_set_expiry_error(
"code": "invalid_token_id",
"message": "Received invalid token",
}
RFC7636_VERIFIER = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"
RFC7636_CHALLENGE = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
async def _async_login_for_code(
client: TestClient, code_challenge: str | None = None
) -> str:
"""Run the login flow and return the authorization code."""
payload: dict[str, Any] = {
"client_id": CLIENT_ID,
"handler": ["insecure_example", None],
"redirect_uri": CLIENT_REDIRECT_URI,
}
if code_challenge is not None:
payload["code_challenge"] = code_challenge
payload["code_challenge_method"] = "S256"
resp = await client.post("/auth/login_flow", json=payload)
assert resp.status == HTTPStatus.OK
step = await resp.json()
resp = await client.post(
f"/auth/login_flow/{step['flow_id']}",
json={
"client_id": CLIENT_ID,
"username": "test-user",
"password": "test-pass",
},
)
assert resp.status == HTTPStatus.OK
step = await resp.json()
return step["result"]
async def test_auth_code_pkce_success(
hass: HomeAssistant, aiohttp_client: ClientSessionGenerator
) -> None:
"""Test login flow and token exchange with PKCE S256."""
client = await async_setup_auth(hass, aiohttp_client, setup_api=True)
code = await _async_login_for_code(client, RFC7636_CHALLENGE)
resp = await client.post(
"/auth/token",
data={
"client_id": CLIENT_ID,
"grant_type": "authorization_code",
"code": code,
"code_verifier": RFC7636_VERIFIER,
},
)
assert resp.status == HTTPStatus.OK
tokens = await resp.json()
assert hass.auth.async_validate_access_token(tokens["access_token"]) is not None
async def test_auth_code_pkce_missing_code_verifier(
hass: HomeAssistant, aiohttp_client: ClientSessionGenerator
) -> None:
"""Test token exchange fails when code_verifier is missing for PKCE code."""
client = await async_setup_auth(hass, aiohttp_client, setup_api=True)
code = await _async_login_for_code(client, RFC7636_CHALLENGE)
resp = await client.post(
"/auth/token",
data={
"client_id": CLIENT_ID,
"grant_type": "authorization_code",
"code": code,
},
)
assert resp.status == HTTPStatus.BAD_REQUEST
result = await resp.json()
assert result["error"] == "invalid_request"
assert result["error_description"] == "Code verifier required"
@pytest.mark.parametrize(
"invalid_verifier",
[
"wrong_verifier_123456789012345678901234567890123", # valid length, wrong content
"short", # < 43 chars
"a" * 129, # > 128 chars
"non_ascii_verifier_with_unicode_characters_✓_123456", # non-ascii
],
ids=["wrong", "too_short", "too_long", "non_ascii"],
)
async def test_auth_code_pkce_invalid_code_verifier(
hass: HomeAssistant,
aiohttp_client: ClientSessionGenerator,
invalid_verifier: str,
) -> None:
"""Test token exchange fails when code_verifier is invalid."""
client = await async_setup_auth(hass, aiohttp_client, setup_api=True)
code = await _async_login_for_code(client, RFC7636_CHALLENGE)
resp = await client.post(
"/auth/token",
data={
"client_id": CLIENT_ID,
"grant_type": "authorization_code",
"code": code,
"code_verifier": invalid_verifier,
},
)
assert resp.status == HTTPStatus.BAD_REQUEST
result = await resp.json()
assert result["error"] == "invalid_grant"
assert result["error_description"] == "Invalid code verifier"
async def test_auth_code_without_challenge_succeeds(
hass: HomeAssistant, aiohttp_client: ClientSessionGenerator
) -> None:
"""Test token exchange succeeds when flow was started without code_challenge and no verifier sent."""
client = await async_setup_auth(hass, aiohttp_client, setup_api=True)
code = await _async_login_for_code(client)
resp = await client.post(
"/auth/token",
data={
"client_id": CLIENT_ID,
"grant_type": "authorization_code",
"code": code,
},
)
assert resp.status == HTTPStatus.OK
tokens = await resp.json()
assert hass.auth.async_validate_access_token(tokens["access_token"]) is not None
async def test_auth_code_unexpected_verifier_rejected(
hass: HomeAssistant, aiohttp_client: ClientSessionGenerator
) -> None:
"""Test token exchange fails when client sends code_verifier but no code_challenge was registered."""
client = await async_setup_auth(hass, aiohttp_client, setup_api=True)
code = await _async_login_for_code(client)
resp = await client.post(
"/auth/token",
data={
"client_id": CLIENT_ID,
"grant_type": "authorization_code",
"code": code,
"code_verifier": RFC7636_VERIFIER,
},
)
assert resp.status == HTTPStatus.BAD_REQUEST
result = await resp.json()
assert result["error"] == "invalid_request"
assert "no code challenge was present" in result["error_description"]
+68
View File
@@ -426,6 +426,7 @@ async def test_well_known_auth_info(
"token_endpoint": f"{expected_url_prefix}/auth/token",
"revocation_endpoint": f"{expected_url_prefix}/auth/revoke",
"client_id_metadata_document_supported": True,
"code_challenge_methods_supported": ["S256"],
"response_types_supported": ["code"],
"service_documentation": "https://developers.home-assistant.io/docs/auth_api",
}
@@ -496,3 +497,70 @@ async def test_well_known_protected_resource_no_url(
"/.well-known/oauth-protected-resource",
)
assert resp.status == 404
@pytest.mark.parametrize(
("payload", "expected_message"),
[
(
{
"code_challenge_method": "S256",
},
"code_challenge required when code_challenge_method is provided",
),
(
{
"code_challenge": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM",
},
"Transform algorithm not supported",
),
(
{
"code_challenge": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM",
"code_challenge_method": "plain",
},
"Transform algorithm not supported",
),
(
{
"code_challenge": "short",
"code_challenge_method": "S256",
},
"Message format incorrect",
),
(
{
"code_challenge": "a" * 43 + "=",
"code_challenge_method": "S256",
},
"Message format incorrect",
),
],
ids=[
"method_without_challenge",
"challenge_without_method",
"unsupported_plain_method",
"challenge_too_short",
"challenge_padded",
],
)
async def test_login_flow_pkce_validation(
hass: HomeAssistant,
aiohttp_client: ClientSessionGenerator,
payload: dict[str, str],
expected_message: str,
) -> None:
"""Test PKCE parameter validation in login_flow."""
client = await async_setup_auth(hass, aiohttp_client, setup_api=True)
resp = await client.post(
"/auth/login_flow",
json={
"client_id": CLIENT_ID,
"handler": ["insecure_example", None],
"redirect_uri": CLIENT_REDIRECT_URI,
**payload,
},
)
assert resp.status == HTTPStatus.BAD_REQUEST
result = await resp.json()
assert expected_message in result["message"]