diff --git a/README.md b/README.md index dd86d02..2736630 100644 --- a/README.md +++ b/README.md @@ -150,6 +150,26 @@ client.auth.client.set_bearer_authorization(access_token) # Now you can make authenticated requests ``` +### 401 Auto-Refresh + +`with_user_session()` installs a one-shot 401 auto-refresh hook on the auth and api clients by default (issue #89): when a request comes back 401 because the bearer expired mid-session, the session token is force-refreshed once and the request retried with the new Authorization header. If the refresh fails, the original 401 surfaces to the caller as before. Pass `refresh_on_401=False` for the old behaviour. + +```python +with campus.with_user_session() as client: # hook on (default) + ... + +with campus.with_user_session(refresh_on_401=False) as client: # hook off + ... +``` + +The retry cannot double-execute work: Campus services authenticate requests in a `before_request` hook before any handler runs, so a 401 response means no handler executed. + +Raw `CampusRequest` users can install their own hook (e.g. public clients driving `auth.refresh(stored)`): + +```python +client.set_unauthorized_hook(lambda: campus.auth.refresh(stored, client_id="campus-cli").access_token) +``` + ## Service Base URLs Each service client (`campus.auth`, `campus.api`, `campus.audit`) resolves its base URL in this order: diff --git a/campus_python/__init__.py b/campus_python/__init__.py index 96b4424..65168cb 100644 --- a/campus_python/__init__.py +++ b/campus_python/__init__.py @@ -243,21 +243,78 @@ def with_app_session(self) -> Iterator["Campus"]: self.revoke_session() @contextmanager - def with_user_session(self) -> Iterator["Campus"]: + def with_user_session( + self, + *, + refresh_on_401: bool = True, + ) -> Iterator["Campus"]: """Context manager yielding CampusRequest with user credentials. Usage: with campus.with_user_session() as client: # use client for requests + By default a 401 auto-refresh hook is installed on the auth and + api clients for the duration of the session (issue #89): when a + request comes back 401 because the bearer expired after the + proactive refresh below, the token is force-refreshed once and + the request retried with the new Authorization header. The + retry cannot double-execute work — Campus services reject + unauthenticated requests in a before_request authenticator + before any handler runs. If the refresh itself fails, the + original 401 response surfaces to the caller as before. Pass + refresh_on_401=False for the old behaviour. + Yields: - CampusRequest: JSON client with user credentials set. + Campus: JSON client with user credentials set. """ + token = self._get_token_from_session() + self.use_token(token) + + hooked_clients = [] + if refresh_on_401: + refreshing = False + + def refresh_bearer() -> str | None: + """Force-refresh the session token once. + + Returns the new bearer token for the client to retry + with, or None on failure (the original 401 then + surfaces). Re-entrant calls (the refresh round-trip + itself hitting a 401) bail out immediately. + + The refresh round-trip authenticates as the client + (Basic), the same mode _get_token_from_session runs in + at session establishment — auth routes accept both, + but the session's stale bearer must not be presented + mid-rotation. + """ + nonlocal refreshing + if refreshing: + return None + refreshing = True + try: + self.revoke_session() + refreshed = self._get_token_from_session( + force_refresh=True + ) + except errors.APIError: + self.use_token(token) + return None + finally: + refreshing = False + self.use_token(refreshed) + return refreshed.access_token + + for client in (self.api.client, self.auth.client): + client.set_unauthorized_hook(refresh_bearer) + hooked_clients.append(client) + try: - token = self._get_token_from_session() - self.use_token(token) yield self except Exception: raise finally: + for client in hooked_clients: + client.set_unauthorized_hook(None) self.revoke_session() diff --git a/campus_python/json_client/__init__.py b/campus_python/json_client/__init__.py index c90c396..5ca1e40 100644 --- a/campus_python/json_client/__init__.py +++ b/campus_python/json_client/__init__.py @@ -10,6 +10,7 @@ ] import base64 +from collections.abc import Callable from typing import Any, Mapping, Self, cast import campus.model @@ -87,6 +88,10 @@ def __init__( self._headers = dict(headers or {}) # allow optional default timeout via kwargs self._timeout = kwargs.get("timeout", 10) + # Optional 401 auto-refresh hook (issue #89); see + # set_unauthorized_hook(). Off by default. + self._unauthorized_hook: Callable[[], str | None] | None = None + self._in_unauthorized_hook = False # Session to persist headers and connection pooling self._session = requests.Session() self._session.headers.update(self._headers) @@ -140,6 +145,62 @@ def set_bearer_authorization(self, token: str) -> None: """ self._session.headers["Authorization"] = "Bearer " + token + def set_unauthorized_hook( + self, + hook: Callable[[], str | None] | None, + ) -> None: + """Install a 401 auto-refresh hook (issue #89). + + When a request comes back 401, the hook is invoked once; if it + returns a bearer token, the Authorization header is refreshed + and the request is retried once with it. A second 401 on the + retry is returned to the caller as-is, and a hook returning + None (refresh failed) leaves the original 401 response in + place — failures surface exactly as they would without the + hook. + + The hook runs at most once per request, and never re-enters + itself: requests made from inside the hook skip the hook, so a + hook that calls back into this client cannot recurse. + + Args: + hook: Callable returning the new bearer token, or None to + signal refresh failure. None clears the hook. + """ + self._unauthorized_hook = hook + + def _send(self, method: str, url: str, **kwargs: Any) -> JsonResponse: + """Send a request, with one 401-triggered retry (issue #89). + + The retry is safe even for non-idempotent requests: Campus + services validate the bearer in a before_request authenticator + before dispatching, so a 401 response is produced by the auth + layer before any handler runs — there is no partial work to + replay. + """ + try: + resp = self._session.request( + method, url, timeout=self._timeout, **kwargs + ) + if ( + resp.status_code == 401 + and self._unauthorized_hook is not None + and not self._in_unauthorized_hook + ): + self._in_unauthorized_hook = True + try: + token = self._unauthorized_hook() + finally: + self._in_unauthorized_hook = False + if token: + self.set_bearer_authorization(token) + resp = self._session.request( + method, url, timeout=self._timeout, **kwargs + ) + except requests.RequestException as exc: + raise errors.ServerError(error_description=str(exc)) from None + return CampusResponse(resp) + def get( self: Self, path: str, @@ -147,27 +208,14 @@ def get( ) -> JsonResponse: """Sends a GET request.""" url = self._build_url(path) - try: - if query: - resp = self._session.get( - url, - params=query, - timeout=self._timeout - ) - else: - resp = self._session.get(url, timeout=self._timeout) - except requests.RequestException as exc: - raise errors.ServerError(error_description=str(exc)) from None - return CampusResponse(resp) + if query: + return self._send("GET", url, params=query) + return self._send("GET", url) def post(self: Self, path: str, json: JsonDict | None = None) -> JsonResponse: """Sends a POST request.""" url = self._build_url(path) - try: - resp = self._session.post(url, json=json, timeout=self._timeout) - except requests.RequestException as exc: - raise errors.ServerError(error_description=str(exc)) from None - return CampusResponse(resp) + return self._send("POST", url, json=json) def put( self: Self, @@ -177,13 +225,7 @@ def put( ) -> JsonResponse: """Sends a PUT request.""" url = self._build_url(path) - try: - resp = self._session.put( - url, json=json, params=query, timeout=self._timeout - ) - except requests.RequestException as exc: - raise errors.ServerError(error_description=str(exc)) from None - return CampusResponse(resp) + return self._send("PUT", url, json=json, params=query) def delete( self: Self, @@ -193,13 +235,7 @@ def delete( ) -> JsonResponse: """Sends a DELETE request.""" url = self._build_url(path) - try: - resp = self._session.delete( - url, json=json, params=query, timeout=self._timeout - ) - except requests.RequestException as exc: - raise errors.ServerError(error_description=str(exc)) from None - return CampusResponse(resp) + return self._send("DELETE", url, json=json, params=query) def patch( self: Self, @@ -209,10 +245,4 @@ def patch( ) -> JsonResponse: """Sends a PATCH request.""" url = self._build_url(path) - try: - resp = self._session.patch( - url, json=json, params=query, timeout=self._timeout - ) - except requests.RequestException as exc: - raise errors.ServerError(error_description=str(exc)) from None - return CampusResponse(resp) + return self._send("PATCH", url, json=json, params=query) diff --git a/campus_python/json_client/interface.py b/campus_python/json_client/interface.py index b08f09e..2d9e886 100644 --- a/campus_python/json_client/interface.py +++ b/campus_python/json_client/interface.py @@ -107,6 +107,12 @@ def raise_for_status(self) -> None: class JsonClient(ABC): """This class describes the public interface required from Client classes, which are used to send JSON requests. + + Optional capability, not part of this abstract interface: concrete + clients may support a 401 auto-refresh hook + (CampusRequest.set_unauthorized_hook, issue #89) that refreshes the + bearer token and retries a request once when it comes back 401. + Callers must feature-detect (hasattr) rather than assume it. """ base_url: str # pylint: disable=unnecessary-ellipsis diff --git a/tests/unit/test_json_client_query.py b/tests/unit/test_json_client_query.py index d4f727b..fdb1438 100644 --- a/tests/unit/test_json_client_query.py +++ b/tests/unit/test_json_client_query.py @@ -23,33 +23,34 @@ def setUp(self): base_url="https://auth.example.test", mode="device" ) - def _assert_params_passed(self, verb: str, session_method: mock.Mock): + def _assert_params_passed(self, verb: str, expected_method: str): self.client._timeout = 5 with mock.patch.object( - self.client._session, session_method, - return_value=mock.Mock()) as method: + self.client._session, "request", + return_value=mock.Mock()) as request: getattr(self.client, verb)( "/some/path", json={"k": "v"}, query={"user_id": "u"}) - _, kwargs = method.call_args + args, kwargs = request.call_args + self.assertEqual(args[0], expected_method.upper()) self.assertEqual(kwargs["json"], {"k": "v"}) self.assertEqual(kwargs["params"], {"user_id": "u"}) def test_put_passes_query_as_params(self): - self._assert_params_passed("put", "put") + self._assert_params_passed("put", "PUT") def test_delete_passes_query_as_params(self): - self._assert_params_passed("delete", "delete") + self._assert_params_passed("delete", "DELETE") def test_patch_passes_query_as_params(self): - self._assert_params_passed("patch", "patch") + self._assert_params_passed("patch", "PATCH") def test_query_defaults_to_none(self): """Omitting query= sends params=None, preserving old behavior.""" with mock.patch.object( - self.client._session, "delete", - return_value=mock.Mock()) as delete: + self.client._session, "request", + return_value=mock.Mock()) as request: self.client.delete("/some/path") - self.assertIsNone(delete.call_args.kwargs["params"]) + self.assertIsNone(request.call_args.kwargs["params"]) if __name__ == "__main__": diff --git a/tests/unit/test_unauthorized_refresh.py b/tests/unit/test_unauthorized_refresh.py new file mode 100644 index 0000000..51895ae --- /dev/null +++ b/tests/unit/test_unauthorized_refresh.py @@ -0,0 +1,244 @@ +"""Tests for the 401 auto-refresh hook (issue #89). + +CampusRequest.set_unauthorized_hook() installs a one-shot refresh hook: +on a 401 response the hook runs once, and if it returns a bearer token +the Authorization header is refreshed and the request retried once. +Failures (hook returns None, or the retry 401s again) surface the +original 401 to the caller — the hook is strictly a safety net. + +Campus.with_user_session(refresh_on_401=True) wires the hook so the +refresh round-trip force-refreshes the session token, authenticating as +the client (Basic) like session establishment does. +""" + +import json +import os +import unittest +from unittest import mock +from unittest.mock import MagicMock, Mock + +import requests + +os.environ.setdefault("CLIENT_ID", "test-client-id") +os.environ.setdefault("CLIENT_SECRET", "test-client-secret") + +from campus_python import Campus, errors +from campus_python.json_client import CampusRequest + + +def raw_response(status_code: int, payload: dict | None = None) -> requests.Response: + """Build a real requests.Response so CampusResponse parses it.""" + raw = requests.Response() + raw.status_code = status_code + raw._content = json.dumps(payload or {}).encode("utf-8") + raw.headers["Content-Type"] = "application/json" + return raw + + +def make_client() -> CampusRequest: + client = CampusRequest(base_url="https://api.example.test", mode="device") + client._session = MagicMock() + # A real dict: set_bearer_authorization writes headers["..."] and + # the tests read it back + client._session.headers = {} + return client + + +class TestUnauthorizedHook(unittest.TestCase): + """set_unauthorized_hook(): retry-once semantics on 401.""" + + def setUp(self): + self.client = make_client() + + def set_responses(self, *responses: requests.Response) -> MagicMock: + self.client._session.request.side_effect = list(responses) + return self.client._session.request + + def test_401_triggers_hook_and_retries(self): + request_mock = self.set_responses( + raw_response(401, {"error": {"code": "AUTH_TOKEN_INVALID"}}), + raw_response(200, {"ok": True}), + ) + self.client.set_unauthorized_hook(lambda: "fresh-token") + + resp = self.client.get("/things") + + self.assertEqual(resp.status_code, 200) + self.assertEqual(request_mock.call_count, 2) + self.assertEqual( + self.client._session.headers["Authorization"], "Bearer fresh-token" + ) + + def test_hook_returning_none_keeps_original_401(self): + request_mock = self.set_responses(raw_response(401, {})) + self.client.set_unauthorized_hook(lambda: None) + + resp = self.client.get("/things") + + self.assertEqual(resp.status_code, 401) + self.assertEqual(request_mock.call_count, 1) + self.assertNotIn("Authorization", self.client._session.headers) + + def test_no_hook_401_passes_through(self): + request_mock = self.set_responses(raw_response(401, {})) + + resp = self.client.get("/things") + + self.assertEqual(resp.status_code, 401) + self.assertEqual(request_mock.call_count, 1) + + def test_second_401_not_retried_again(self): + request_mock = self.set_responses( + raw_response(401, {}), raw_response(401, {}) + ) + calls: list[int] = [] + self.client.set_unauthorized_hook(lambda: calls.append(1) or "tok") + + resp = self.client.get("/things") + + self.assertEqual(resp.status_code, 401) + self.assertEqual(len(calls), 1) + self.assertEqual(request_mock.call_count, 2) + + def test_reentrant_hook_call_skips_hook(self): + """Requests made from inside the hook skip the hook — a hook + that calls back into this client cannot recurse.""" + request_mock = self.set_responses( + raw_response(401, {}), # outer request + raw_response(401, {}), # hook's own inner request + raw_response(401, {}), # outer retry + ) + + def hook() -> str: + inner = self.client.get("/inner") + self.assertEqual(inner.status_code, 401) + return "tok" + + self.client.set_unauthorized_hook(hook) + resp = self.client.get("/outer") + + self.assertEqual(resp.status_code, 401) + self.assertEqual(request_mock.call_count, 3) + + def test_clearing_hook_disables_retry(self): + request_mock = self.set_responses(raw_response(401, {})) + self.client.set_unauthorized_hook(lambda: "tok") + self.client.set_unauthorized_hook(None) + + resp = self.client.get("/things") + + self.assertEqual(resp.status_code, 401) + self.assertEqual(request_mock.call_count, 1) + + def test_post_retries_like_get(self): + request_mock = self.set_responses( + raw_response(401, {}), raw_response(200, {"id": "x1"}) + ) + self.client.set_unauthorized_hook(lambda: "fresh-token") + + resp = self.client.post("/things", json={"k": "v"}) + + self.assertEqual(resp.status_code, 200) + self.assertEqual(request_mock.call_count, 2) + # The retry replays the same body + self.assertEqual(request_mock.call_args.kwargs["json"], {"k": "v"}) + + +class TestWithUserSessionWiring(unittest.TestCase): + """with_user_session() installs a force-refresh hook on both + service clients and removes it when the session closes.""" + + def setUp(self): + self.campus = Campus.__new__(Campus) + self.campus._mode = "server" + + self.api_client = make_client() + self.auth_client = make_client() + self.api_root = Mock(client=self.api_client) + self.auth_root = Mock(client=self.auth_client) + + patchers = [ + mock.patch.object( + Campus, "api", new_callable=mock.PropertyMock, + return_value=self.api_root, + ), + mock.patch.object( + Campus, "auth", new_callable=mock.PropertyMock, + return_value=self.auth_root, + ), + ] + for patcher in patchers: + patcher.start() + self.addCleanup(patcher.stop) + + self.get_token = mock.patch.object( + Campus, "_get_token_from_session" + ).start() + self.addCleanup(mock.patch.stopall) + + token1 = Mock(access_token="tok1") + self.get_token.return_value = token1 + + def test_401_mid_session_force_refreshes_and_retries(self): + self.api_client._session.request.side_effect = [ + raw_response(401, {}), raw_response(200, {"ok": True}), + ] + token2 = Mock(access_token="tok2") + self.get_token.side_effect = [Mock(access_token="tok1"), token2] + + with self.campus.with_user_session(): + resp = self.campus.api.client.get("/things") + # The new bearer is live on BOTH service clients mid-session + self.assertEqual( + self.api_client._session.headers["Authorization"], + "Bearer tok2", + ) + self.assertEqual( + self.auth_client._session.headers["Authorization"], + "Bearer tok2", + ) + + self.assertEqual(resp.status_code, 200) + # Second call is the hook's force-refresh + self.assertEqual(self.get_token.call_count, 2) + self.assertTrue(self.get_token.call_args.kwargs["force_refresh"]) + + def test_refresh_failure_surfaces_original_401(self): + self.api_client._session.request.side_effect = [raw_response(401, {})] + self.get_token.side_effect = [ + Mock(access_token="tok1"), + errors.AuthenticationError(error_description="session gone"), + ] + + with self.campus.with_user_session(): + resp = self.campus.api.client.get("/things") + # Pre-hook bearer state restored after the failed refresh + self.assertEqual( + self.api_client._session.headers["Authorization"], + "Bearer tok1", + ) + + self.assertEqual(resp.status_code, 401) + self.assertEqual(self.api_client._session.request.call_count, 1) + + def test_hooks_cleared_after_session(self): + with self.campus.with_user_session(): + self.assertIsNotNone(self.api_client._unauthorized_hook) + self.assertIsNotNone(self.auth_client._unauthorized_hook) + + self.assertIsNone(self.api_client._unauthorized_hook) + self.assertIsNone(self.auth_client._unauthorized_hook) + + def test_refresh_on_401_false_keeps_old_behaviour(self): + self.api_client._session.request.side_effect = [raw_response(401, {})] + + with self.campus.with_user_session(refresh_on_401=False): + self.assertIsNone(self.api_client._unauthorized_hook) + resp = self.campus.api.client.get("/things") + + self.assertEqual(resp.status_code, 401) + self.get_token.assert_called_once() # only the proactive refresh + + +if __name__ == "__main__": + unittest.main()