From 0c5e19109dd414b39a4d7bd016992eb07361eb5a Mon Sep 17 00:00:00 2001 From: JS Ng Date: Sun, 4 Oct 2026 13:55:54 +0800 Subject: [PATCH] feat(auth): public-client refresh grant + device-flow polling parity (#87) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - auth.token(refresh_token grant) now sends client_id — required by the token endpoint for every grant (campus auth/routes/oauth.py); explicit arg for public clients, CLIENT_ID env fallback for server mode. The grant previously shipped without client_id and could only fail server-side. - auth.refresh(stored) — public-client helper: presents the stored (single-use) refresh token and returns the rotated pair. - auth.oauth.wait_for_token() — RFC 8628 §3.5 polling loop: persisted slow_down (+5s), authorization_pending polling with on_pending hook, network retry on non-final attempts, timeout error. Parity table in the oauth module docstring lets campus-cli retire login.py's copies. - poll_for_token() parses all three server error shapes: Campus envelope (dev), envelope with details stripped (production — OAuth error recovered from the AUTH_* code via errors.oauth_error_from_code), and flat RFC 6749. - 218 tests (was 197). --- campus_python/auth/v1/__init__.py | 62 ++++++ campus_python/auth/v1/oauth.py | 183 +++++++++++++++--- campus_python/errors.py | 35 ++++ tests/unit/test_oauth_device_flow.py | 245 +++++++++++++++++++++--- tests/unit/test_oauth_token_contract.py | 31 ++- tests/unit/test_token_refresh.py | 112 ++++++++++- 6 files changed, 613 insertions(+), 55 deletions(-) diff --git a/campus_python/auth/v1/__init__.py b/campus_python/auth/v1/__init__.py index 1d059f4..7271908 100644 --- a/campus_python/auth/v1/__init__.py +++ b/campus_python/auth/v1/__init__.py @@ -341,6 +341,7 @@ def token( ], *, refresh_token: str | None = None, + client_id: str | None = None, ) -> campus.model.OAuthToken: """Get OAuth token from the token endpoint. @@ -356,6 +357,17 @@ def token( client's base_url itself, so an absolute URL here would be double-prefixed into `https://host/https://host/...` and 404 against real deployments. + + Args: + grant_type: "client_credentials" or "refresh_token". + refresh_token: The refresh token to present (refresh_token + grant). + client_id: The OAuth client the grant is made as. The token + endpoint requires client_id for every grant (campus + auth/routes/oauth.py token()), including refresh_token; + defaults to CLIENT_ID from the environment (server + mode). Public clients without CLIENT_ID configured must + pass it explicitly (issue #87). """ json_body: dict[str, str] = { "grant_type": grant_type, @@ -370,9 +382,59 @@ def token( error_description="Refresh token required for " "refresh_token grant type." ) + resolved_client_id = client_id or env.get("CLIENT_ID") + if not resolved_client_id: + raise errors.AuthenticationError( + error_description="client_id is required for the " + "refresh_token grant; pass it " + "explicitly or set CLIENT_ID." + ) + json_body["client_id"] = resolved_client_id json_body["refresh_token"] = refresh_token token_path = self.url_prefix + "/oauth/token" resp = self.client.post(token_path, json=json_body) resp.raise_for_status() return campus.model.OAuthToken.from_resource(resp.json()) + + def refresh( + self, + stored: campus.model.OAuthToken, + *, + client_id: str | None = None, + ) -> campus.model.OAuthToken: + """Refresh an OAuth token pair (RFC 6749 section 6). + + Takes a stored token and returns the rotated pair issued by the + server: a new access token with a new refresh token. The + presented refresh token is single-use — refresh-token grants + rotate server-side (campus/auth/routes/oauth.py + _handle_refresh_token_grant) — so the returned token must be + persisted before the stale one is presented again. + + This is the public-client entry point (issue #87): it needs no + client secret, only the client_id the stored token was issued + to. Error responses (invalid_grant, invalid_client, ...) raise + APIError subclasses carrying the OAuth error code in details, + readable via APIError.oauth_error. + + Args: + stored: The stored OAuthToken whose refresh token is + presented to the server. + client_id: The OAuth client the token was issued to; + defaults to CLIENT_ID from the environment (server + mode). + + Returns: + The refreshed OAuthToken (new access + refresh tokens). + """ + if not stored.refresh_token: + raise errors.AuthenticationError( + error_description="Stored token has no refresh token; " + "re-authentication is required." + ) + return self.token( + grant_type="refresh_token", + refresh_token=stored.refresh_token, + client_id=client_id, + ) diff --git a/campus_python/auth/v1/oauth.py b/campus_python/auth/v1/oauth.py index 44496cc..e8be71c 100644 --- a/campus_python/auth/v1/oauth.py +++ b/campus_python/auth/v1/oauth.py @@ -4,13 +4,75 @@ This module provides methods for the device authorization flow, which is used by CLI and other device applications. + +Device-flow parity (issue #87) +------------------------------ + +The error mapping and polling semantics implemented here pin the +behaviour campus-cli's device flow (campus_cli/auth/login.py) relies +on, so the CLI can retire its raw `requests` copies and drive this +resource instead. Parity table: + +| Server response (RFC 8628 §3.5) | Library behaviour | +|--------------------------------------|--------------------------------------| +| authorization_pending | wait_for_token() invokes on_pending() and keeps polling at the current interval | +| slow_down | wait_for_token() raises the poll interval by 5s and the raised interval persists for the remainder of the flow (not just the next attempt) | +| expired_token | fatal: AuthenticationError with oauth_error="expired_token" | +| access_denied | fatal: AuthenticationError with oauth_error="access_denied" | +| any other 400 | fatal: AuthenticationError carrying the server's OAuth error code | +| network failure / 5xx | retried at the current interval on non-final attempts, then ServerError propagates | +| max_attempts exhausted | fatal: AuthenticationError (timeout) | + +The auth server emits token-endpoint errors as the Campus envelope +{"error": {code, message, details.oauth_error}} (dev/staging), the same +envelope with details stripped in production, or the flat RFC 6749 form +{"error": ..., "error_description": ...}; poll_for_token() accepts all +three and normalizes them to AuthenticationError with the OAuth error +code in details (APIError.oauth_error). """ +import time +from collections.abc import Callable from typing import Literal from ... import errors from ...interface import ResourceRoot +# Fallback poll interval (seconds) when the server omits interval from +# the device_authorize response (RFC 8628 §3.2 default is 5). +DEFAULT_POLL_INTERVAL = 5 + +# RFC 8628 §3.5: the interval adjustment requested by slow_down. +SLOW_DOWN_ADJUSTMENT = 5 + + +def _parse_oauth_error(payload: dict) -> "tuple[str | None, str]": + """Extract the OAuth error code and message from a token-endpoint + error payload. + + Handles the three shapes the auth server emits: + + - flat RFC 6749: {"error": "authorization_pending", ...} + - Campus envelope (dev/staging): {"error": {"code": "AUTH_...", + "message": ..., "details": {"oauth_error": ...}}} + - Campus envelope with details stripped (production): the OAuth + error is recovered from the AUTH_* code + + Returns: + (oauth_error, message); oauth_error is None when the payload + carries neither a details.oauth_error key nor a recoverable + AUTH_* code. + """ + error = payload.get("error", "") + if isinstance(error, str): + # Flat RFC 6749 format + return (error or None, payload.get("error_description", "")) + + oauth_error = (error.get("details") or {}).get("oauth_error") + if not oauth_error: + oauth_error = errors.oauth_error_from_code(error.get("code", "")) + return (oauth_error, error.get("message", "")) + class OAuth(ResourceRoot): """OAuth 2.0 Device Authorization Flow resource. @@ -93,41 +155,108 @@ def poll_for_token( # Handle OAuth error responses if resp.status_code == 400: - error_data = resp.json() - error = error_data.get("error", "") + oauth_error, message = _parse_oauth_error(resp.json()) # Map RFC 8628 errors to AuthenticationError; the OAuth error # code travels in details so callers can read it back via the # APIError.oauth_error property. - if error == "authorization_pending": - raise errors.AuthenticationError( - error_description="Authorization pending", - details={"oauth_error": "authorization_pending"} - ) - elif error == "slow_down": - raise errors.AuthenticationError( - error_description="Slow down", - details={"oauth_error": "slow_down"} - ) - elif error == "expired_token": - raise errors.AuthenticationError( - error_description="Device code has expired", - details={"oauth_error": "expired_token"} - ) - elif error == "access_denied": - raise errors.AuthenticationError( - error_description="Access denied by user", - details={"oauth_error": "access_denied"} - ) - else: - raise errors.AuthenticationError( - error_description=error_data.get("error_description", "Unknown error"), - details={"oauth_error": error} - ) + descriptions = { + "authorization_pending": "Authorization pending", + "slow_down": "Slow down", + "expired_token": "Device code has expired", + "access_denied": "Access denied by user", + } + raise errors.AuthenticationError( + status_code=400, + error_description=( + descriptions.get(oauth_error) + or message + or "Unknown error" + ), + details={"oauth_error": oauth_error}, + ) resp.raise_for_status() return resp.json() + def wait_for_token( + self, + client_id: str, + device_code: str, + *, + interval: int | None = None, + max_attempts: int = 60, + on_pending: Callable[[], None] | None = None, + sleep: Callable[[float], None] = time.sleep, + ) -> dict: + """Poll the token endpoint until the device is authorized. + + Implements the polling loop RFC 8628 §3.5 clients must run on + top of poll_for_token(): authorization_pending keeps polling, + slow_down raises the interval by SLOW_DOWN_ADJUSTMENT seconds + with the raised interval persisting for the remainder of the + flow, network failures are retried on non-final attempts, and + every other error (expired_token, access_denied, unknown) is + fatal. See the parity table in this module's docstring. + + Args: + client_id: The OAuth client ID (e.g., "campus-cli") + device_code: The device code from request_device_code() + interval: Minimum seconds between poll attempts, from the + request_device_code() response. Defaults to + DEFAULT_POLL_INTERVAL when absent or zero. + max_attempts: Maximum number of poll attempts; callers + typically derive this from the device code's expires_in. + on_pending: Invoked after each authorization_pending response, + for progress reporting (e.g. printing a dot per poll). + sleep: The sleep function; injectable for tests. + + Returns: + The token response dict (access_token, refresh_token, ...), + as returned by poll_for_token(). + + Raises: + AuthenticationError: On fatal OAuth errors or when + max_attempts is exhausted without authorization. + errors.ServerError: When the final attempt fails at the + network/5xx level. + """ + poll_interval = ( + max(1, int(interval)) if interval else DEFAULT_POLL_INTERVAL + ) + last_attempt = max_attempts - 1 + + for attempt in range(max_attempts): + try: + return self.poll_for_token( + client_id=client_id, device_code=device_code + ) + except errors.AuthenticationError as err: + if err.oauth_error == "authorization_pending": + if on_pending is not None: + on_pending() + sleep(poll_interval) + elif err.oauth_error == "slow_down": + # RFC 8628 §3.5: the raised interval persists for + # the remainder of the flow, not just the next + # attempt. + poll_interval += SLOW_DOWN_ADJUSTMENT + sleep(poll_interval) + else: + raise + except errors.ServerError: + # Network failure: retry on non-final attempts only + if attempt == last_attempt: + raise + sleep(poll_interval) + + raise errors.AuthenticationError( + error_description=( + "Device authorization timed out after " + f"{max_attempts} poll attempts; restart the flow." + ) + ) + def authorize_device( self, user_code: str, diff --git a/campus_python/errors.py b/campus_python/errors.py index 9b24e8f..97bff07 100644 --- a/campus_python/errors.py +++ b/campus_python/errors.py @@ -25,6 +25,33 @@ class FieldError: message: str +# Campus envelope error codes that do not follow the AUTH_ +# pattern (campus/common/errors/base.py _OAUTH_TO_CAMPUS_ERROR_CODES). +_CAMPUS_CODE_TO_OAUTH_ERROR = { + "AUTH_UNSUPPORTED_GRANT": "unsupported_grant_type", +} + + +def oauth_error_from_code(code: str) -> str | None: + """Recover the OAuth error code from a Campus envelope AUTH_* code. + + The auth server strips the error envelope's details in production + (campus/common/errors/handlers.py), taking details.oauth_error with + it; the code remains and encodes the OAuth error as + AUTH_ (upper snake case), with the exceptions mapped + in _CAMPUS_CODE_TO_OAUTH_ERROR. + + Returns None for codes that carry no OAuth error (plain API codes). + """ + if not code: + return None + if code in _CAMPUS_CODE_TO_OAUTH_ERROR: + return _CAMPUS_CODE_TO_OAUTH_ERROR[code] + if code.startswith("AUTH_"): + return code[len("AUTH_"):].lower() + return None + + class APIError(Exception): """Base exception for all campus client errors. @@ -140,6 +167,14 @@ def with_status_code( error_description = error_description or error_obj.get("message") request_id = request_id or error_obj.get("request_id") details = details or error_obj.get("details") + if not details: + # Production strips envelope details, taking the + # OAuth error code with it; recover it from the + # AUTH_* code so auth callers can still branch on + # APIError.oauth_error (#87). + derived = oauth_error_from_code(error_obj.get("code", "")) + if derived: + details = {"oauth_error": derived} # Parse field-level errors for validation errors if "errors" in error_obj and isinstance(error_obj["errors"], list): diff --git a/tests/unit/test_oauth_device_flow.py b/tests/unit/test_oauth_device_flow.py index 7deed64..358ddc6 100644 --- a/tests/unit/test_oauth_device_flow.py +++ b/tests/unit/test_oauth_device_flow.py @@ -67,55 +67,248 @@ class TestPollForTokenErrors(unittest.TestCase): def setUp(self): self.auth, self.client = make_auth() - def make_error_response(self, error: str) -> Mock: + def make_error_response(self, error) -> Mock: response = Mock() response.status_code = 400 - response.json.return_value = { - "error": error, - "error_description": f"desc: {error}", - } + response.json.return_value = error self.client.post.return_value = response return response - def test_authorization_pending(self): - self.make_error_response("authorization_pending") + def assert_oauth_error(self, oauth_error: str): with self.assertRaises(errors.AuthenticationError) as ctx: self.auth.oauth.poll_for_token( client_id="campus-cli", device_code="dev123" ) - self.assertEqual(ctx.exception.oauth_error, "authorization_pending") + self.assertEqual(ctx.exception.oauth_error, oauth_error) + + def test_authorization_pending(self): + self.make_error_response({ + "error": "authorization_pending", + "error_description": "desc: authorization_pending", + }) + self.assert_oauth_error("authorization_pending") def test_slow_down(self): - self.make_error_response("slow_down") - with self.assertRaises(errors.AuthenticationError) as ctx: - self.auth.oauth.poll_for_token( - client_id="campus-cli", device_code="dev123" - ) - self.assertEqual(ctx.exception.oauth_error, "slow_down") + self.make_error_response({ + "error": "slow_down", + "error_description": "desc: slow_down", + }) + self.assert_oauth_error("slow_down") def test_expired_token(self): - self.make_error_response("expired_token") + self.make_error_response({ + "error": "expired_token", + "error_description": "desc: expired_token", + }) + self.assert_oauth_error("expired_token") + + def test_access_denied(self): + self.make_error_response({ + "error": "access_denied", + "error_description": "desc: access_denied", + }) + self.assert_oauth_error("access_denied") + + def test_unknown_error(self): + self.make_error_response({ + "error": "something_else", + "error_description": "desc: something_else", + }) + self.assert_oauth_error("something_else") + + +class TestPollForTokenErrorEnvelopes(unittest.TestCase): + """The token endpoint emits three error shapes; all must resolve to + the same oauth_error (#87): + + - Campus envelope (dev/staging): oauth_error lives in details + - Campus envelope with details stripped (production): recovered + from the AUTH_* code + - flat RFC 6749: {"error": } + """ + + def setUp(self): + self.auth, self.client = make_auth() + + def poll_error(self, payload: dict) -> errors.AuthenticationError: + response = Mock() + response.status_code = 400 + response.json.return_value = payload + self.client.post.return_value = response with self.assertRaises(errors.AuthenticationError) as ctx: self.auth.oauth.poll_for_token( client_id="campus-cli", device_code="dev123" ) + return ctx.exception + + def test_campus_envelope_with_details(self): + err = self.poll_error({ + "error": { + "code": "AUTH_SLOW_DOWN", + "message": "Slow down", + "details": {"oauth_error": "slow_down"}, + "request_id": None, + } + }) + self.assertEqual(err.oauth_error, "slow_down") + self.assertEqual(err.error_description, "Slow down") + + def test_campus_envelope_stripped_details(self): + """Production strips details; the AUTH_* code identifies the + OAuth error.""" + err = self.poll_error({ + "error": { + "code": "AUTH_AUTHORIZATION_PENDING", + "message": "Authorization pending", + "request_id": None, + } + }) + self.assertEqual(err.oauth_error, "authorization_pending") + + def test_campus_envelope_nonstandard_code_mapping(self): + """Codes that don't follow AUTH_ map explicitly.""" + err = self.poll_error({ + "error": { + "code": "AUTH_UNSUPPORTED_GRANT", + "message": "Unsupported grant type", + "request_id": None, + } + }) + self.assertEqual(err.oauth_error, "unsupported_grant_type") + + +class TestWaitForToken(unittest.TestCase): + """wait_for_token() implements the RFC 8628 §3.5 polling loop the + CLI's poll_for_token() owned (issue #87).""" + + def setUp(self): + self.auth, self.client = make_auth() + self.sleeps: list[float] = [] + + def sleep(self, seconds: float) -> None: + self.sleeps.append(seconds) + + def token_response(self) -> Mock: + response = Mock() + response.status_code = 200 + response.json.return_value = {"access_token": "tok123"} + return response + + def error_response(self, payload: dict) -> Mock: + response = Mock() + response.status_code = 400 + response.json.return_value = payload + return response + + def pending(self) -> Mock: + return self.error_response({ + "error": { + "code": "AUTH_AUTHORIZATION_PENDING", + "message": "Authorization pending", + "request_id": None, + } + }) + + def with_error(self, oauth_error: str) -> Mock: + return self.error_response({ + "error": { + "code": f"AUTH_{oauth_error.upper()}", + "message": oauth_error, + "request_id": None, + } + }) + + def wait(self, max_attempts: int = 5, **kwargs) -> dict: + return self.auth.oauth.wait_for_token( + client_id="campus-cli", + device_code="dev123", + interval=5, + max_attempts=max_attempts, + sleep=self.sleep, + **kwargs, + ) + + def test_returns_token_after_pending_polls(self): + self.client.post.side_effect = [ + self.pending(), self.pending(), self.token_response(), + ] + result = self.wait() + self.assertEqual(result, {"access_token": "tok123"}) + self.assertEqual(self.sleeps, [5, 5]) + + def test_on_pending_fires_per_pending_poll(self): + self.client.post.side_effect = [ + self.pending(), self.token_response(), + ] + pends: list[int] = [] + self.wait(on_pending=lambda: pends.append(1)) + self.assertEqual(len(pends), 1) + + def test_slow_down_raises_interval_persistently(self): + """RFC 8628 §3.5: the +5s adjustment persists for the remainder + of the flow, not just the next attempt.""" + self.client.post.side_effect = [ + self.with_error("slow_down"), self.pending(), self.token_response(), + ] + self.wait() + self.assertEqual(self.sleeps, [10, 10]) + + def test_expired_token_is_fatal(self): + self.client.post.side_effect = [self.with_error("expired_token")] + with self.assertRaises(errors.AuthenticationError) as ctx: + self.wait() self.assertEqual(ctx.exception.oauth_error, "expired_token") + self.assertEqual(self.sleeps, []) - def test_access_denied(self): - self.make_error_response("access_denied") + def test_access_denied_is_fatal(self): + self.client.post.side_effect = [self.with_error("access_denied")] with self.assertRaises(errors.AuthenticationError) as ctx: - self.auth.oauth.poll_for_token( - client_id="campus-cli", device_code="dev123" - ) + self.wait() self.assertEqual(ctx.exception.oauth_error, "access_denied") - def test_unknown_error(self): - self.make_error_response("something_else") + def test_unknown_error_is_fatal(self): + self.client.post.side_effect = [ + self.with_error("unsupported_grant_type") + ] + with self.assertRaises(errors.AuthenticationError): + self.wait() + + def test_network_error_retries_on_nonfinal_attempt(self): + self.client.post.side_effect = [ + errors.ServerError(error_description="connection reset"), + self.token_response(), + ] + result = self.wait() + self.assertEqual(result, {"access_token": "tok123"}) + self.assertEqual(self.sleeps, [5]) + + def test_network_error_on_final_attempt_propagates(self): + self.client.post.side_effect = [ + errors.ServerError(error_description="connection reset") + ] * 3 + with self.assertRaises(errors.ServerError): + self.wait(max_attempts=3) + # Two retries after the first failure, no third + self.assertEqual(self.client.post.call_count, 3) + self.assertEqual(self.sleeps, [5, 5]) + + def test_timeout_raises_authentication_error(self): + self.client.post.side_effect = [self.pending()] * 5 with self.assertRaises(errors.AuthenticationError) as ctx: - self.auth.oauth.poll_for_token( - client_id="campus-cli", device_code="dev123" - ) - self.assertEqual(ctx.exception.oauth_error, "something_else") + self.wait(max_attempts=5) + self.assertIsNone(ctx.exception.oauth_error) + self.assertIn("timed out", ctx.exception.error_description) + + def test_default_interval_when_server_omits_it(self): + self.client.post.side_effect = [self.pending(), self.token_response()] + self.auth.oauth.wait_for_token( + client_id="campus-cli", + device_code="dev123", + interval=None, + max_attempts=5, + sleep=self.sleep, + ) + self.assertEqual(self.sleeps, [5]) if __name__ == "__main__": diff --git a/tests/unit/test_oauth_token_contract.py b/tests/unit/test_oauth_token_contract.py index a4b7fa7..494baa6 100644 --- a/tests/unit/test_oauth_token_contract.py +++ b/tests/unit/test_oauth_token_contract.py @@ -28,6 +28,7 @@ import campus.model +from campus_python import errors from campus_python.auth.v1 import AuthRoot # Exact token-endpoint response shape emitted by campus weekly @@ -205,13 +206,41 @@ def test_client_credentials_targets_oauth_token_endpoint(self): self.assertEqual(kwargs["json"]["client_secret"], "sec123") def test_refresh_token_targets_oauth_token_endpoint(self): - self.auth.token(grant_type="refresh_token", refresh_token="rt123") + """The refresh grant must carry client_id: the token endpoint + requires it for every grant (campus auth/routes/oauth.py + token()), including refresh_token (#87).""" + with patch.dict(os.environ, {"CLIENT_ID": "cid123"}): + self.auth.token(grant_type="refresh_token", refresh_token="rt123") args, kwargs = self.client.post.call_args self.assertEqual(args[0], "/auth/v1/oauth/token") self.assertEqual(kwargs["json"]["grant_type"], "refresh_token") + self.assertEqual(kwargs["json"]["client_id"], "cid123") self.assertEqual(kwargs["json"]["refresh_token"], "rt123") + def test_refresh_token_explicit_client_id_for_public_clients(self): + """A public client passes client_id explicitly (#87): device + mode has no CLIENT_ID env to fall back to.""" + with patch.dict(os.environ, {}, clear=True): + self.auth.token( + grant_type="refresh_token", + refresh_token="rt123", + client_id="campus-cli", + ) + + kwargs = self.client.post.call_args.kwargs + self.assertEqual(kwargs["json"]["client_id"], "campus-cli") + + def test_refresh_token_without_client_id_is_rejected(self): + """No client_id from arg or env — refuse to send a request the + server would reject anyway (missing required parameter).""" + with patch.dict(os.environ, {}, clear=True): + with self.assertRaises(errors.AuthenticationError): + self.auth.token( + grant_type="refresh_token", refresh_token="rt123" + ) + self.client.post.assert_not_called() + def test_token_path_is_relative_not_double_prefixed(self): """Regression (client issue #62): with a non-empty client base_url, the path handed to the client must stay relative — diff --git a/tests/unit/test_token_refresh.py b/tests/unit/test_token_refresh.py index 778d3cb..996d71f 100644 --- a/tests/unit/test_token_refresh.py +++ b/tests/unit/test_token_refresh.py @@ -3,13 +3,47 @@ After exchanging the refresh token, the refreshed token must be the one returned — the credentials resource still holds the pre-refresh access token, and refresh-token grants rotate (single-use) on the server. + +Also covers the public-client refresh helper auth.refresh(stored) +(client issue #87): it presents the stored refresh token with the +issuing client_id and returns the rotated pair. """ +import json import os import unittest from unittest.mock import MagicMock, Mock, patch -from campus_python import Campus +import campus.model +import requests + +from campus_python import Campus, errors +from campus_python.auth.v1 import AuthRoot +from campus_python.json_client import CampusResponse + +RFC_REFRESH_PAYLOAD = { + "access_token": "tok-new", + "token_type": "Bearer", + "expires_in": 86400, + "refresh_token": "rt-new", + "scope": "campus.profile", +} + + +def make_auth() -> tuple[AuthRoot, Mock]: + """Create an AuthRoot backed by a mock JSON client.""" + client = Mock() + return AuthRoot(json_client=client), client + + +def make_response(status_code: int, payload: dict) -> CampusResponse: + """Build a CampusResponse over a real requests.Response so + raise_for_status() runs the actual error-envelope mapping.""" + raw = requests.Response() + raw.status_code = status_code + raw._content = json.dumps(payload).encode("utf-8") + raw.headers["Content-Type"] = "application/json" + return CampusResponse(raw) class TestGetTokenFromSession(unittest.TestCase): @@ -48,5 +82,81 @@ def test_returns_refreshed_token_after_refresh(self): creds_resource.update.assert_called_once_with(token=new_token) +class TestPublicClientRefresh(unittest.TestCase): + """auth.refresh(stored) presents the stored refresh token with the + issuing client_id and returns the rotated pair (#87).""" + + def setUp(self): + self.auth, self.client = make_auth() + self.client.post.return_value = make_response( + 200, dict(RFC_REFRESH_PAYLOAD) + ) + self.stored = campus.model.OAuthToken( + id="tok-old", + expires_in=3600, + refresh_token="rt-old", + scopes=["campus.profile"], + ) + + def test_returns_rotated_pair(self): + refreshed = self.auth.refresh(self.stored, client_id="campus-cli") + + self.assertEqual(refreshed.access_token, "tok-new") + self.assertEqual(refreshed.refresh_token, "rt-new") + # The stored token is unchanged; rotation is the caller's to persist + self.assertEqual(self.stored.refresh_token, "rt-old") + + def test_sends_stored_refresh_token_with_client_id(self): + self.auth.refresh(self.stored, client_id="campus-cli") + + args, kwargs = self.client.post.call_args + self.assertEqual(args[0], "/auth/v1/oauth/token") + self.assertEqual(kwargs["json"]["grant_type"], "refresh_token") + self.assertEqual(kwargs["json"]["client_id"], "campus-cli") + self.assertEqual(kwargs["json"]["refresh_token"], "rt-old") + + def test_client_id_falls_back_to_env(self): + with patch.dict(os.environ, {"CLIENT_ID": "cid123"}): + self.auth.refresh(self.stored) + + kwargs = self.client.post.call_args.kwargs + self.assertEqual(kwargs["json"]["client_id"], "cid123") + + def test_stored_token_without_refresh_token_is_rejected(self): + bare = campus.model.OAuthToken(id="tok-bare", expires_in=3600) + with self.assertRaises(errors.AuthenticationError): + self.auth.refresh(bare, client_id="campus-cli") + self.client.post.assert_not_called() + + def test_invalid_grant_maps_to_apierror_with_oauth_error(self): + """A rejected refresh token raises with the OAuth error code in + details — the caller's cue to fall back to re-login.""" + self.client.post.return_value = make_response(400, { + "error": { + "code": "AUTH_INVALID_GRANT", + "message": "Invalid or expired refresh token", + "details": {"oauth_error": "invalid_grant"}, + "request_id": None, + } + }) + with self.assertRaises(errors.APIError) as ctx: + self.auth.refresh(self.stored, client_id="campus-cli") + self.assertEqual(ctx.exception.oauth_error, "invalid_grant") + + def test_invalid_grant_stripped_details_still_identifies_error(self): + """Production strips error details; the AUTH_* code still + identifies the OAuth error (#87).""" + self.client.post.return_value = make_response(400, { + "error": { + "code": "AUTH_INVALID_GRANT", + "message": "Invalid or expired refresh token", + "request_id": None, + } + }) + with self.assertRaises(errors.APIError) as ctx: + self.auth.refresh(self.stored, client_id="campus-cli") + self.assertEqual(ctx.exception.oauth_error, "invalid_grant") + + if __name__ == "__main__": unittest.main()