diff --git a/gnexus_creds/oauth.py b/gnexus_creds/oauth.py index 09367f7..f47b0cb 100644 --- a/gnexus_creds/oauth.py +++ b/gnexus_creds/oauth.py @@ -7,13 +7,14 @@ from gnexus_gauth.config import GAuthConfig from gnexus_gauth.oauth import AuthorizationUrlBuilder, HttpTokenEndpoint, PkceGenerator from gnexus_gauth.runtime import HttpRuntimeUserProvider +from gnexus_gauth.exceptions import WebhookPayloadException, WebhookVerificationException from gnexus_gauth.webhook import HmacWebhookVerifier, JsonWebhookParser from sqlalchemy.orm import Session from gnexus_creds.config import get_settings from gnexus_creds.db import get_db from gnexus_creds.errors import AppError -from gnexus_creds.models import OAuthState, User, utcnow +from gnexus_creds.models import OAuthState, SessionRecord, User, utcnow from gnexus_creds.services import ( check_rate_limit, create_session, @@ -134,14 +135,26 @@ return response +# The auth platform's webhook target may be stored with a trailing slash. +# A mismatched slash would not redirect: the SPA catch-all is a GET route +# that matches the path and answers 405 (with allow: GET) — seen in the +# platform's queue as "HTTP 405" retries for auth.global_logout. Accept +# both spellings explicitly. @webhook_router.post("/gnexus-auth") +@webhook_router.post("/gnexus-auth/") async def gnexus_auth_webhook(request: Request, db: Session = Depends(get_db)) -> dict[str, str]: settings = get_settings() raw = (await request.body()).decode() headers = dict(request.headers) config = _config() - HmacWebhookVerifier(config).verify(raw, headers, settings.auth_webhook_secret) - event = JsonWebhookParser().parse(raw) + try: + HmacWebhookVerifier(config).verify(raw, headers, settings.auth_webhook_secret) + except WebhookVerificationException as exc: + raise AppError("webhook_unauthorized", str(exc), status_code=401) + try: + event = JsonWebhookParser().parse(raw) + except WebhookPayloadException as exc: + raise AppError("webhook_bad_payload", str(exc), status_code=400) subject = ( event.target_identifiers.get("sub") or event.target_identifiers.get("user_id") @@ -163,5 +176,9 @@ user.status = "disabled" elif status == "enabled": user.status = "enabled" + if event.event_type == "auth.global_logout": + # logging out of the auth system ends every web session of + # this user here too — the SPA cookie must stop working + db.query(SessionRecord).filter(SessionRecord.user_id == user.id).delete() db.commit() return {"status": "ok"} diff --git a/tests/test_auth.py b/tests/test_auth.py index 1ff8503..07dabf0 100644 --- a/tests/test_auth.py +++ b/tests/test_auth.py @@ -118,6 +118,7 @@ class FakeParser: def parse(self, raw): return SimpleNamespace( + event_type="auth.updated", target_identifiers={"sub": "auth-user-1"}, metadata={ "status": "disabled", @@ -145,3 +146,123 @@ # per-service locale override assert user.locale == "en" assert user.profile == {"display_name": "Disabled User", "locale": "uk"} + + +@pytest.mark.anyio +async def test_webhook_accepts_trailing_slash(auth_app, db_session, user, monkeypatch): + # the auth platform's target URL may carry a trailing slash; the SPA + # catch-all (a GET route) must not answer the POST with 405 + class FakeVerifier: + def __init__(self, config): + self.config = config + + def verify(self, raw, headers, secret): + pass + + class FakeParser: + def parse(self, raw): + return SimpleNamespace( + target_identifiers={"sub": "auth-user-1"}, metadata={}, event_type="auth.updated" + ) + + monkeypatch.setattr("gnexus_creds.oauth.HmacWebhookVerifier", FakeVerifier) + monkeypatch.setattr("gnexus_creds.oauth.JsonWebhookParser", FakeParser) + + async with AsyncClient( + transport=ASGITransport(app=auth_app), base_url="http://test" + ) as client: + response = await client.post("/webhooks/gnexus-auth/", json={"x": 1}) + + assert response.status_code == 200 + + +@pytest.mark.anyio +async def test_webhook_rejects_bad_signature_with_401(auth_app, db_session, monkeypatch): + from gnexus_gauth.exceptions import WebhookVerificationException + + class Explode: + def __init__(self, config): + self.config = config + + def verify(self, raw, headers, secret): + raise WebhookVerificationException("signature mismatch") + + monkeypatch.setattr("gnexus_creds.oauth.HmacWebhookVerifier", Explode) + + async with AsyncClient( + transport=ASGITransport(app=auth_app), base_url="http://test" + ) as client: + response = await client.post("/webhooks/gnexus-auth", content=b"{}", headers={}) + + assert response.status_code == 401 + + +@pytest.mark.anyio +async def test_webhook_global_logout_ends_sessions(auth_app, db_session, user, monkeypatch): + db_session.add_all( + [ + SessionRecord( + id="logged-out-session", + user_id=user.id, + data={}, + expires_at=utcnow() + timedelta(days=1), + ), + SessionRecord( + id="foreign-session", + user_id=user.id, + data={}, + expires_at=utcnow() + timedelta(days=1), + ), + ] + ) + db_session.commit() + other_user = User( + auth_subject="other-sub", + email="other@example.test", + display_name="Other", + status="enabled", + profile={}, + locale="en", + ) + db_session.add(other_user) + db_session.commit() + db_session.add( + SessionRecord( + id="kept-session", + user_id=other_user.id, + data={}, + expires_at=utcnow() + timedelta(days=1), + ) + ) + db_session.commit() + + class FakeVerifier: + def __init__(self, config): + self.config = config + + def verify(self, raw, headers, secret): + pass + + class FakeParser: + def parse(self, raw): + return SimpleNamespace( + event_type="auth.global_logout", + target_identifiers={"sub": "auth-user-1"}, + metadata={}, + ) + + monkeypatch.setattr("gnexus_creds.oauth.HmacWebhookVerifier", FakeVerifier) + monkeypatch.setattr("gnexus_creds.oauth.JsonWebhookParser", FakeParser) + + async with AsyncClient( + transport=ASGITransport(app=auth_app), base_url="http://test" + ) as client: + response = await client.post("/webhooks/gnexus-auth", json={"x": 1}) + + assert response.status_code == 200 + db_session.expire_all() + assert db_session.get(SessionRecord, "logged-out-session") is None + # the second "own" row was part of the delete-all; only other users' + # sessions survive a global logout + assert db_session.get(SessionRecord, "foreign-session") is None + assert db_session.get(SessionRecord, "kept-session") is not None