diff --git a/enlace_auth/__main__.py b/enlace_auth/__main__.py index e1e6cf6..c216e49 100644 --- a/enlace_auth/__main__.py +++ b/enlace_auth/__main__.py @@ -253,6 +253,8 @@ def set_password(email: str, *, toml: str = "platform.toml"): updated = dict(record) updated["password_hash"] = _hash(pw) + # Account recovery: unlink external sign-ins too (see the HTTP reset paths). + updated.pop("oauth_links", None) store[key] = updated # Same rule as the HTTP reset paths: a new password ends the old sessions. revoked = _load_session_store(Path(toml)).revoke_user(key) diff --git a/enlace_auth/admin/routes.py b/enlace_auth/admin/routes.py index 5479a58..700eed5 100644 --- a/enlace_auth/admin/routes.py +++ b/enlace_auth/admin/routes.py @@ -242,6 +242,9 @@ async def admin_reset_password( raise HTTPException(status_code=500, detail="Corrupt user record") record = dict(record) record["password_hash"] = hash_password(body.password) + # A reset is account recovery: external sign-ins linked to the account + # are unlinked too, or whoever linked one walks straight back in. + record.pop("oauth_links", None) user_store[target] = record # An admin reset is how a compromised account is recovered: whoever # holds a session opened with the old password must lose it. diff --git a/enlace_auth/auth/oauth.py b/enlace_auth/auth/oauth.py index 6ce4ff4..0592971 100644 --- a/enlace_auth/auth/oauth.py +++ b/enlace_auth/auth/oauth.py @@ -6,12 +6,35 @@ in TOML). On callback we create a local session — the upstream tokens are discarded because we use OAuth for identity only, not API access. +Two rules keep an OAuth login from being weaker than the account it opens: + +- **The anti-CSRF state lives in a signed cookie** scoped to ``/auth`` + (:func:`_oauth_state_session`). Authlib keeps the ``state``, nonce and PKCE + verifier in ``request.session``; the plugin installs no Starlette + ``SessionMiddleware``, so this module supplies that session itself. A callback + whose ``state`` was not issued to *this* browser is refused. +- **An identity is bound to the provider's stable subject** (``sub``, or + ``tid``/``oid`` for Microsoft, GitHub's numeric ``id``), recorded as + ``oauth_links[provider]`` on the account. A later login must present the same + subject. An existing *password* account, or one linked to another provider, is + never taken over by an email match alone. + +Residual limits, by design: the state cookie is signed but not bound to the +browser, so a script that can set cookies on the platform origin (any +co-hosted app) could plant its own state for a browser that has none and so +log that browser in as the attacker (two cookies of the name are refused). An +account created by a provider before links existed is bound to the first +subject that signs in to it after the upgrade. A password reset (admin, emailed +link, CLI) unlinks every external sign-in. The cookie path assumes the router +is mounted at ``/auth`` with no root path. + Built-in provider presets for Google and GitHub auto-fill the well-known endpoints; other providers need explicit URLs in the config. """ from __future__ import annotations +import json import logging import os import time @@ -20,7 +43,7 @@ from fastapi import APIRouter, HTTPException, Request, Response from fastapi.responses import JSONResponse -from enlace_auth.auth.cookies import sign_cookie +from enlace_auth.auth.cookies import sign_cookie, verify_cookie from enlace_auth.auth.sessions import SessionStore from enlace_auth.config import OAuthProviderConfig @@ -76,6 +99,82 @@ def _email_trusted(provider: str, cfg, claims: dict) -> bool: ) +def _is_entra_issuer(iss: str) -> bool: + """True for Microsoft Entra ID token issuers.""" + from urllib.parse import urlparse + + host = (urlparse(iss).hostname or "").lower() + return host in {"login.microsoftonline.com", "sts.windows.net"} + + +def _cookie_count(request: Request, name: str) -> int: + """How many cookies called *name* the request carries (Starlette keeps one).""" + raw = request.headers.get("cookie", "") + return sum(1 for part in raw.split(";") if part.strip().split("=", 1)[0] == name) + + +def _stable_subject(provider: str, claims: dict) -> Optional[str]: + """The provider's stable, non-reassignable id for the signed-in identity. + + >>> _stable_subject("google", {"sub": "123", "email": "a@x.io"}) + '123' + >>> _stable_subject("microsoft", {"sub": "s", "tid": "T", "oid": "O", + ... "iss": "https://login.microsoftonline.com/T/v2.0"}) + 'T/O' + >>> _stable_subject("custom", {"sub": "s", "tid": "T", "oid": "O"}) + 's' + >>> _stable_subject("github", {"id": 42}) + '42' + >>> _stable_subject("x", {"email": "a@x.io"}) is None + True + """ + tid, oid = claims.get("tid"), claims.get("oid") + iss = str(claims.get("iss") or "") + if tid and oid and _is_entra_issuer(iss): + # Entra ID: `sub` is pairwise per app; tid/oid is the user. Only for + # Entra's own issuers -- elsewhere these are ordinary, maybe + # user-influenced, claims. + return f"{tid}/{oid}" + for key in ("sub", "id"): # OIDC, then GitHub's /user + value = claims.get(key) + if value not in (None, ""): + return str(value) + return None + + +def _login_refusal(record: Any, provider: str, subject: Optional[str]) -> Optional[str]: + """Why *record* must not be opened by this OAuth identity, or None if it may. + + >>> _login_refusal({"password_hash": "h"}, "google", "1") is not None + True + >>> _login_refusal({"oauth_links": {"google": "1"}}, "google", "1") is None + True + >>> _login_refusal({"oauth_links": {"google": "1"}}, "google", "2") is not None + True + >>> _login_refusal({"password_hash": None, "oauth_provider": "google"}, + ... "google", "1") is None + True + >>> _login_refusal({"password_hash": None, "oauth_provider": "github"}, + ... "google", "1") is not None + True + """ + if not isinstance(record, dict): + return "This account cannot be opened with an external sign-in." + links = record.get("oauth_links") or {} + if provider in links: + if subject is not None and links[provider] == subject: + return None + return "This sign-in does not match the identity linked to this account." + if record.get("password_hash"): + return ( + "This email belongs to a password account. Sign in with your " + "password; an external sign-in is not linked to it." + ) + if record.get("oauth_provider") == provider: + return None # created by this provider before links were recorded + return "This account is linked to a different sign-in method." + + def _build_oauth_registry(providers: dict[str, OAuthProviderConfig]): OAuth = _import_authlib() oauth = OAuth() @@ -122,8 +221,16 @@ def make_oauth_router( session_max_age: int = 86400, secure_cookies: bool = True, can_register: Callable[[str], bool] = lambda _: False, + state_cookie_name: str = "enlace_oauth_state", + state_max_age: int = 600, ) -> Optional[APIRouter]: - """Build an OAuth router or return None if no providers are configured.""" + """Build an OAuth router or return None if no providers are configured. + + *state_cookie_name* / *state_max_age* name and bound the signed cookie that + carries Authlib's per-login state between ``/auth/login/{provider}`` and the + callback (see the module docstring). It is only used when no Starlette + ``SessionMiddleware`` already provides ``request.session``. + """ if not providers: return None @@ -143,6 +250,61 @@ def _set_session_cookie(response: Response, session_id: str): attrs.append("Secure") response.headers.append("set-cookie", "; ".join(attrs)) + _state_salt = "oauth-state" + + def _oauth_state_session(request: Request) -> bool: + """Give Authlib a ``request.session`` backed by the signed state cookie. + + Returns True when this module owns the session (and so must write it + back), False when a real ``SessionMiddleware`` already provides one. + """ + if "session" in request.scope: + return False + data: dict = {} + token = request.cookies.get(state_cookie_name) + if _cookie_count(request, state_cookie_name) > 1: + # Two cookies of this name means one was planted at another path + # (cookie tossing) to smuggle in someone else's login state. + token = None + raw = ( + verify_cookie(token, signing_key, max_age=state_max_age, salt=_state_salt) + if token + else None + ) + if raw: + try: + loaded = json.loads(raw) + if isinstance(loaded, dict): + data = loaded + except ValueError: + pass + request.scope["session"] = data + return True + + def _state_cookie_header(session: dict) -> str: + if session: + value = sign_cookie(json.dumps(session), signing_key, salt=_state_salt) + attrs = [ + f"{state_cookie_name}={value}", + "Path=/auth", + "HttpOnly", + f"Max-Age={state_max_age}", + "SameSite=Lax", + ] + else: + attrs = [f"{state_cookie_name}=", "Path=/auth", "HttpOnly", "Max-Age=0"] + if secure_cookies: + attrs.append("Secure") + return "; ".join(attrs) + + def _write_state_cookie(response: Response, session: dict) -> None: + response.headers.append("set-cookie", _state_cookie_header(session)) + + def _drop_provider_states(session: dict, provider: str) -> None: + """A callback spends every pending state of its provider, win or lose.""" + for key in [k for k in session if k.startswith(f"_state_{provider}_")]: + session.pop(key, None) + @router.get("/login/{provider}") async def login(provider: str, request: Request): client = getattr(oauth, provider, None) @@ -150,8 +312,12 @@ async def login(provider: str, request: Request): raise HTTPException( status_code=404, detail=f"Unknown provider '{provider}'" ) + owns_session = _oauth_state_session(request) redirect_uri = str(request.url_for("oauth_callback", provider=provider)) - return await client.authorize_redirect(request, redirect_uri) + resp = await client.authorize_redirect(request, redirect_uri) + if owns_session: + _write_state_cookie(resp, request.session) + return resp @router.get("/callback/{provider}", name="oauth_callback") async def callback(provider: str, request: Request): @@ -160,6 +326,23 @@ async def callback(provider: str, request: Request): raise HTTPException( status_code=404, detail=f"Unknown provider '{provider}'" ) + owns_session = _oauth_state_session(request) + try: + resp = await _complete_login(provider, client, request) + except HTTPException as e: + if owns_session: + _drop_provider_states(request.session, provider) + e.headers = { + **(e.headers or {}), + "set-cookie": _state_cookie_header(request.session), + } + raise + if owns_session: + _drop_provider_states(request.session, provider) + _write_state_cookie(resp, request.session) + return resp + + async def _complete_login(provider: str, client, request: Request) -> Response: try: token = await client.authorize_access_token(request) except Exception as e: @@ -197,7 +380,12 @@ async def callback(provider: str, request: Request): ) email = email.lower() - if email not in user_store: + subject = _stable_subject(provider, claims) + try: + record = user_store[email] + except KeyError: + record = None + if record is None: if not can_register(email): raise HTTPException( status_code=403, @@ -210,7 +398,29 @@ async def callback(provider: str, request: Request): "password_hash": None, "created_at": time.time(), "oauth_provider": provider, + "oauth_links": {provider: subject} if subject else {}, } + else: + refusal = _login_refusal(record, provider, subject) + if refusal is not None: + _logger.warning( + "OAuth sign-in via %s refused for %s: %s", provider, email, refusal + ) + raise HTTPException(status_code=403, detail=refusal) + links = record.get("oauth_links") or {} + if subject and provider not in links: + # A legacy account this provider created: bind it to the + # subject now, so the email alone never opens it again. + # Re-read so a password change racing this login is not + # overwritten with the copy read above. + current = user_store.get(email, record) + user_store[email] = { + **current, + "oauth_links": { + **(current.get("oauth_links") or {}), + provider: subject, + }, + } session_id = session_store.create(user_id=email, email=email) resp = JSONResponse({"ok": True, "email": email}) _set_session_cookie(resp, session_id) diff --git a/enlace_auth/auth/routes.py b/enlace_auth/auth/routes.py index e4009d4..01e9f9f 100644 --- a/enlace_auth/auth/routes.py +++ b/enlace_auth/auth/routes.py @@ -449,6 +449,9 @@ async def password_reset_confirm( ) record = dict(record) record["password_hash"] = hash_password(body.new_password) + # A reset is account recovery: external sign-ins linked to the account + # are unlinked too, or whoever linked one walks straight back in. + record.pop("oauth_links", None) user_store[email] = record # A reset is the recovery path for a compromised account, so every # session opened with the old password ends here. diff --git a/pyproject.toml b/pyproject.toml index 571735f..42d473c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -31,12 +31,12 @@ Homepage = "https://github.com/i2mint/enlace_auth" enlace-auth = "enlace_auth.__main__:main" [project.optional-dependencies] -oauth = ["authlib>=1.3", "httpx>=0.24.0", "cryptography>=42", "python-multipart>=0.0.9"] +oauth = ["authlib>=1.4", "httpx>=0.24.0", "cryptography>=42", "python-multipart>=0.0.9"] dev = [ "pytest", "httpx", "pytest-asyncio", - "authlib>=1.3", + "authlib>=1.4", "cryptography>=42", "python-multipart>=0.0.9", ] diff --git a/tests/test_oauth_state_and_links.py b/tests/test_oauth_state_and_links.py new file mode 100644 index 0000000..5e72262 --- /dev/null +++ b/tests/test_oauth_state_and_links.py @@ -0,0 +1,265 @@ +# ruff: noqa: F811 -- admin_client is a fixture imported from test_admin. +"""OAuth login: real Authlib state round trip, and subject-bound account links. + +i2mint/enlace_auth#28. The older ``test_oauth.py`` mocks both +``authorize_redirect`` and ``authorize_access_token``, so it never exercised the +anti-CSRF ``state`` check, which needs a ``request.session`` that nothing in +enlace_auth used to provide. These tests keep Authlib's real redirect/state code +and stub only the two network calls (token exchange, userinfo). +""" + +from __future__ import annotations + +import os +from unittest.mock import AsyncMock, patch +from urllib.parse import parse_qs, urlparse + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from enlace_auth.auth import SessionStore +from enlace_auth.config import OAuthProviderConfig +from tests.test_admin import admin_client # noqa: F401 - fixture + +SIGNING_KEY = "oauth-signing-key-32bytes-minlen" + + +def _make_app(user_store, *, claims, can_register=lambda _e: True): + """Router with a real Authlib registry for a stub, non-OIDC provider.""" + os.environ["STUB_ID"] = "stub-client-id" + os.environ["STUB_SECRET"] = "stub-client-secret" + from enlace_auth.auth import oauth as oauth_mod + + providers = { + "stub": OAuthProviderConfig( + client_id_env="STUB_ID", + client_secret_env="STUB_SECRET", + authorize_url="https://idp.example/authorize", + token_url="https://idp.example/token", + userinfo_url="https://idp.example/userinfo", + scopes=["profile", "email"], + trust_unverified_email=True, + ) + } + real_build = oauth_mod._build_oauth_registry + captured = {} + + def _build(p): + registry = real_build(p) + client = registry.stub + # Stub only the network: the state check stays Authlib's own. + client.fetch_access_token = AsyncMock( + return_value={"access_token": "at", "token_type": "Bearer"} + ) + client.userinfo = AsyncMock(return_value=claims) + captured["client"] = client + return registry + + sessions = SessionStore({}) + with patch.object(oauth_mod, "_build_oauth_registry", _build): + router = oauth_mod.make_oauth_router( + providers=providers, + session_store=sessions, + user_store=user_store, + signing_key=SIGNING_KEY, + secure_cookies=False, + can_register=can_register, + ) + app = FastAPI() + app.include_router(router) + return TestClient(app), sessions + + +def _start_login(client: TestClient) -> str: + r = client.get("/auth/login/stub", follow_redirects=False) + assert r.status_code in (302, 303, 307), r.text + assert "enlace_oauth_state=" in r.headers.get("set-cookie", "") + assert "Path=/auth" in r.headers["set-cookie"] + state = parse_qs(urlparse(r.headers["location"]).query)["state"][0] + return state + + +CLAIMS = {"sub": "subject-1", "email": "alice@example.com"} + + +def test_real_state_round_trip_signs_in_and_clears_the_state_cookie(): + users: dict = {} + client, sessions = _make_app(users, claims=CLAIMS) + state = _start_login(client) + r = client.get(f"/auth/callback/stub?code=c&state={state}") + assert r.status_code == 200, r.text + assert users["alice@example.com"]["oauth_links"] == {"stub": "subject-1"} + assert len(sessions.list_all()) == 1 + assert "enlace_oauth_state=;" in r.headers.get("set-cookie", "") + + +def test_callback_without_the_browsers_state_cookie_is_refused(): + """Login CSRF: a callback URL minted in another browser must not sign in.""" + attacker, _ = _make_app({}, claims=CLAIMS) + state = _start_login(attacker) + users: dict = {} + victim, sessions = _make_app(users, claims=CLAIMS) + r = victim.get(f"/auth/callback/stub?code=c&state={state}") + assert r.status_code == 401 + assert users == {} and sessions.list_all() == [] + + +def test_callback_with_a_different_state_is_refused(): + users: dict = {} + client, sessions = _make_app(users, claims=CLAIMS) + _start_login(client) + r = client.get("/auth/callback/stub?code=c&state=forged") + assert r.status_code == 401 + assert sessions.list_all() == [] + + +def test_state_is_single_use(): + client, sessions = _make_app({}, claims=CLAIMS) + state = _start_login(client) + assert client.get(f"/auth/callback/stub?code=c&state={state}").status_code == 200 + r = client.get(f"/auth/callback/stub?code=c&state={state}") + assert r.status_code == 401 + + +def test_a_forged_state_cookie_is_ignored(): + client, sessions = _make_app({}, claims=CLAIMS) + client.cookies.set( + "enlace_oauth_state", + '{"_state_stub_x": {"data": {}, "exp": 9999999999}}', + path="/auth", + ) + r = client.get("/auth/callback/stub?code=c&state=x") + assert r.status_code == 401 + + +# ---- account links ---------------------------------------------------------- + + +def _signin(users, claims): + client, sessions = _make_app(users, claims=claims) + state = _start_login(client) + return client.get(f"/auth/callback/stub?code=c&state={state}"), sessions + + +def test_existing_password_account_is_not_taken_over_by_email(): + users = {"alice@example.com": {"password_hash": "h", "created_at": 0}} + r, sessions = _signin(users, CLAIMS) + assert r.status_code == 403 + assert sessions.list_all() == [] + assert "oauth_links" not in users["alice@example.com"] + + +def test_linked_account_requires_the_same_subject(): + users = { + "alice@example.com": { + "password_hash": None, + "oauth_provider": "stub", + "oauth_links": {"stub": "subject-1"}, + } + } + r, _ = _signin(users, {**CLAIMS, "sub": "someone-else"}) + assert r.status_code == 403 + r, _ = _signin(users, CLAIMS) + assert r.status_code == 200 + + +def test_legacy_account_of_this_provider_is_linked_on_next_login(): + users = {"alice@example.com": {"password_hash": None, "oauth_provider": "stub"}} + r, _ = _signin(users, CLAIMS) + assert r.status_code == 200 + assert users["alice@example.com"]["oauth_links"] == {"stub": "subject-1"} + r, _ = _signin(users, {**CLAIMS, "sub": "someone-else"}) + assert r.status_code == 403 + + +def test_account_of_another_provider_is_not_opened(): + users = {"alice@example.com": {"password_hash": None, "oauth_provider": "google"}} + r, _ = _signin(users, CLAIMS) + assert r.status_code == 403 + + +@pytest.mark.parametrize("claims", [{"email": "alice@example.com"}]) +def test_linked_account_refuses_an_identity_without_subject(claims): + users = {"alice@example.com": {"password_hash": None, "oauth_links": {"stub": "s"}}} + r, _ = _signin(users, claims) + assert r.status_code == 403 + + +# ---- review of #30 ----------------------------------------------------------- + + +def test_a_failed_callback_spends_the_state(): + """An error callback must not leave the state usable for a later code.""" + client, sessions = _make_app({}, claims=CLAIMS) + state = _start_login(client) + r = client.get(f"/auth/callback/stub?error=access_denied&state={state}") + assert r.status_code == 401 + r = client.get(f"/auth/callback/stub?code=c&state={state}") + assert r.status_code == 401 + assert sessions.list_all() == [] + + +def test_a_refused_login_spends_the_state(): + users = {"alice@example.com": {"password_hash": "h", "created_at": 0}} + client, _ = _make_app(users, claims=CLAIMS) + state = _start_login(client) + assert client.get(f"/auth/callback/stub?code=c&state={state}").status_code == 403 + assert client.get(f"/auth/callback/stub?code=c&state={state}").status_code == 401 + + +def test_duplicate_state_cookies_are_refused(): + """Cookie tossing: a second state cookie planted at another path.""" + attacker, _ = _make_app({}, claims=CLAIMS) + state = _start_login(attacker) + planted = attacker.cookies.get("enlace_oauth_state") + victim, sessions = _make_app({}, claims=CLAIMS) + _start_login(victim) + genuine = victim.cookies.get("enlace_oauth_state") + r = victim.get( + f"/auth/callback/stub?code=c&state={state}", + headers={ + "cookie": f"enlace_oauth_state={genuine}; enlace_oauth_state={planted}" + }, + ) + assert r.status_code == 401 + assert sessions.list_all() == [] + + +def test_admin_password_reset_unlinks_external_sign_ins(admin_client, tmp_path): + from enlace_auth.stores import make_file_store_factory + from tests.test_admin import _csrf, _register + + csrf = _csrf(admin_client) + _register(admin_client, "boss@example.com", "bosspw1!", csrf) + r = admin_client.post( + "/_admin/api/users", + json={"email": "vic@example.com", "password": "victim-pw1"}, + headers=csrf, + ) + assert r.status_code == 200, r.text + users = make_file_store_factory(str(tmp_path / "platform"))("users") + users["vic@example.com"] = { + **users["vic@example.com"], + "oauth_links": {"google": "attacker-subject"}, + } + r = admin_client.post( + "/_admin/api/users/vic@example.com/password", + json={"password": "brand-new-pw1"}, + headers=csrf, + ) + assert r.status_code == 200, r.text + assert "oauth_links" not in users["vic@example.com"] + + +def test_stable_subject_ignores_tid_oid_outside_entra(): + from enlace_auth.auth.oauth import _stable_subject + + assert _stable_subject("x", {"sub": "s", "tid": "t", "oid": "o"}) == "s" + assert ( + _stable_subject( + "x", + {"sub": "s", "tid": "t", "oid": "o", "iss": "https://sts.windows.net/t/"}, + ) + == "t/o" + )