From 2c4a465842f5b5bc874a5d49f9a2be0f9fffb4c8 Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Sun, 30 Aug 2026 11:24:19 -0400 Subject: [PATCH] feat(auth): refresh current sessions --- README.md | 12 ++++ features/contract/auth.feature | 10 +++ features/contract_support.py | 3 +- features/steps/sdk_contract_steps.py | 18 +++++ src/volcano_sdk/__init__.py | 2 + src/volcano_sdk/_transport.py | 25 ++++++- src/volcano_sdk/auth.py | 60 ++++++++++++++-- src/volcano_sdk/client.py | 34 +++++++-- src/volcano_sdk/errors.py | 12 ++++ tests/unit/test_contract_bindings.py | 4 +- tests/unit/test_generated_transport.py | 31 ++++++--- tests/unit/test_state.py | 95 +++++++++++++++++++++++++- 12 files changed, 282 insertions(+), 24 deletions(-) diff --git a/README.md b/README.md index 12a73490..bc625012 100644 --- a/README.md +++ b/README.md @@ -45,6 +45,18 @@ if session is not None: `set_session()` copies the session without making a request or persisting credentials. It raises `ValueError` when the session type or any credential field is incomplete. +Refresh the session with its current refresh token: + +```python +refreshed = client.auth.refresh_session() +assert client.auth.get_session() is refreshed +``` + +On success, `refresh_session()` replaces the in-memory session and returns the immutable new +snapshot. An authentication failure clears the session that initiated the request. Server and +transport failures preserve it, and a late response never replaces a newer session. The SDK does +not persist sessions. + Realtime is async. Channels wrap `centrifuge-python`; the underlying client and subscription objects are not part of the public API. diff --git a/features/contract/auth.feature b/features/contract/auth.feature index f280a45c..682cbde0 100644 --- a/features/contract/auth.feature +++ b/features/contract/auth.feature @@ -25,3 +25,13 @@ Feature: SDK authentication contract Then the SDK operation succeeds And the current session belongs to the contract user And the current session exposes access and refresh tokens + + @auth @SDK-AUTH-004 + Scenario: A client refreshes its current session + Given the confirmed contract user + When the client signs in with the contract user's credentials + And the client refreshes the current session + Then the SDK operation succeeds + And the refreshed session replaces the previous credentials + And the current session belongs to the contract user + And the current session exposes access and refresh tokens diff --git a/features/contract_support.py b/features/contract_support.py index fc626414..576debcf 100644 --- a/features/contract_support.py +++ b/features/contract_support.py @@ -23,7 +23,7 @@ if TYPE_CHECKING: from collections.abc import Awaitable, Callable - from volcano_sdk.models import LockLease + from volcano_sdk.models import LockLease, Session from volcano_sdk.realtime import Channel HTTP_NOT_FOUND = 404 @@ -114,6 +114,7 @@ def __init__(self, fixture: dict[str, Any]) -> None: "value": f"volcano-sdk-contract-{suffix}", } self.last_outcome: Outcome | None = None + self.previous_session: Session | None = None self.subscriber: Channel | None = None self.publisher: Channel | None = None self.realtime_clients: list[VolcanoClient] = [] diff --git a/features/steps/sdk_contract_steps.py b/features/steps/sdk_contract_steps.py index 10c01b49..302d36d4 100644 --- a/features/steps/sdk_contract_steps.py +++ b/features/steps/sdk_contract_steps.py @@ -53,6 +53,24 @@ def adopt_current_session(context: Any) -> None: world.client = target +@when("the client refreshes the current session") +def refresh_current_session(context: Any) -> None: + world = _world(context) + world.previous_session = world.client.auth.get_session() + assert world.previous_session is not None + world.record(world.client.auth.refresh_session) + + +@then("the refreshed session replaces the previous credentials") +def refreshed_session_replaces_credentials(context: Any) -> None: + world = _world(context) + assert world.last_outcome is not None + assert world.previous_session is not None + assert ( + world.last_outcome.value.refresh_token != world.previous_session.refresh_token + ) + + @then("the SDK operation succeeds") def operation_succeeds(context: Any) -> None: outcome = _world(context).last_outcome diff --git a/src/volcano_sdk/__init__.py b/src/volcano_sdk/__init__.py index ab9a77cb..d6e75f26 100644 --- a/src/volcano_sdk/__init__.py +++ b/src/volcano_sdk/__init__.py @@ -7,6 +7,7 @@ NotFoundError, RateLimitedError, ServerError, + SessionChangedError, TransportError, ValidationError, VolcanoError, @@ -21,6 +22,7 @@ "RateLimitedError", "ServerError", "Session", + "SessionChangedError", "TransportError", "ValidationError", "VolcanoClient", diff --git a/src/volcano_sdk/_transport.py b/src/volcano_sdk/_transport.py index d46bfb50..d692777e 100644 --- a/src/volcano_sdk/_transport.py +++ b/src/volcano_sdk/_transport.py @@ -11,7 +11,7 @@ import httpx -from ._generated.api.authentication import auth_signin +from ._generated.api.authentication import auth_refresh, auth_signin from ._generated.api.database_queries import query_database_select from ._generated.api.locks import acquire_project_lock, release_project_lock from ._generated.api.storage_objects import ( @@ -19,6 +19,7 @@ upload_storage_object, ) from ._generated.client import AuthenticatedClient +from ._generated.models.auth_refresh_body import AuthRefreshBody from ._generated.models.auth_signin_body import AuthSigninBody from ._generated.models.database_select_request import DatabaseSelectRequest from ._generated.models.project_lock_lease_request import ProjectLockLeaseRequest @@ -78,6 +79,15 @@ class _GeneratedTransportResponse: headers: Mapping[str, str] +class AuthRefreshTransport(Protocol): + def auth_refresh( + self, + *, + authorization: str, + refresh_token: str, + ) -> TransportResponse: ... + + class Transport(Protocol): def auth_signin( self, @@ -243,6 +253,19 @@ def auth_signin( ) return self._response(response) + def auth_refresh( + self, + *, + authorization: str, + refresh_token: str, + ) -> TransportResponse: + with self._client(authorization) as client: + response = auth_refresh.sync_detailed( + client=client, + body=AuthRefreshBody(refresh_token=refresh_token), + ) + return self._response(response) + def query_database_select( self, *, diff --git a/src/volcano_sdk/auth.py b/src/volcano_sdk/auth.py index d98b6326..273fb105 100644 --- a/src/volcano_sdk/auth.py +++ b/src/volcano_sdk/auth.py @@ -2,12 +2,15 @@ from __future__ import annotations -from typing import Protocol +from collections.abc import Mapping +from typing import Protocol, cast -from ._transport import Transport, invoke, response_payload +from ._transport import AuthRefreshTransport, Transport, invoke, response_payload +from .errors import AuthenticationError, SessionChangedError from .models import Session _INCOMPLETE_SESSION = "Expected a complete Session" +_NO_ACTIVE_SESSION = "No active session" def _is_non_empty_string(value: object) -> bool: @@ -35,6 +38,25 @@ def _copy_complete_session(session: object) -> Session: ) +def _session_from_payload(payload: object) -> Session: + values: Mapping[object, object] = ( + cast("Mapping[object, object]", payload) if isinstance(payload, Mapping) else {} + ) + raw_user = values.get("user") + user: Mapping[object, object] = ( + cast("Mapping[object, object]", raw_user) + if isinstance(raw_user, Mapping) + else {} + ) + return _copy_complete_session( + Session( + access_token=cast("str", values.get("access_token")), + refresh_token=cast("str", values.get("refresh_token")), + user_id=cast("str", user.get("id")), + ) + ) + + class AuthContext(Protocol): """Client capabilities required by the authentication facade.""" @@ -49,6 +71,12 @@ def _anon_token(self) -> str: ... def _set_session(self, session: Session) -> None: ... + def _capture_session(self) -> tuple[int, Session | None]: ... + + def _set_session_if_current(self, session: Session, generation: int) -> bool: ... + + def _clear_session_if_current(self, generation: int) -> bool: ... + class Auth: """Authenticate users and update the client session.""" @@ -76,10 +104,28 @@ def sign_in(self, *, email: str, password: str) -> Session: password=password, ) payload = response_payload(response, 200) - session = Session( - access_token=payload["access_token"], - refresh_token=payload["refresh_token"], - user_id=payload["user"]["id"], - ) + session = _session_from_payload(payload) self._client._set_session(session) return session + + def refresh_session(self) -> Session: + """Refresh and replace the current session.""" + generation, current = self._client._capture_session() + if current is None: + raise AuthenticationError(_NO_ACTIVE_SESSION) + + transport = cast("AuthRefreshTransport", self._client._transport) + response = invoke( + transport.auth_refresh, + authorization=self._client._anon_token(), + refresh_token=current.refresh_token, + ) + try: + payload = response_payload(response, 200) + except AuthenticationError: + self._client._clear_session_if_current(generation) + raise + refreshed = _session_from_payload(payload) + if not self._client._set_session_if_current(refreshed, generation): + raise SessionChangedError + return refreshed diff --git a/src/volcano_sdk/client.py b/src/volcano_sdk/client.py index 33c18ad9..72464bcb 100644 --- a/src/volcano_sdk/client.py +++ b/src/volcano_sdk/client.py @@ -2,6 +2,7 @@ from __future__ import annotations +import threading from typing import TYPE_CHECKING from ._transport import GeneratedTransport, Transport @@ -35,6 +36,8 @@ def __init__( self._api_url = api_url.rstrip("/") self._anon_key = anon_key self._service_key = service_key + self._session_lock = threading.Lock() + self._session_generation = 0 self._current_session: Session | None = None self._transport: Transport = ( _transport @@ -56,7 +59,7 @@ def __init__( @property def current_session(self) -> Session | None: """Return the authenticated session, if one exists.""" - return self._current_session + return self._capture_session()[1] def database(self, name: str) -> Database: """Create a query facade for a project database.""" @@ -66,9 +69,10 @@ def _anon_token(self) -> str: return self._anon_key def _session_token(self) -> str: - if self._current_session is None: + session = self._capture_session()[1] + if session is None: raise RuntimeError(_NO_ACTIVE_SESSION) - return self._current_session.access_token + return session.access_token def _service_token(self) -> str: if self._service_key is None: @@ -76,4 +80,26 @@ def _service_token(self) -> str: return self._service_key def _set_session(self, session: Session) -> None: - self._current_session = session + with self._session_lock: + self._current_session = session + self._session_generation += 1 + + def _capture_session(self) -> tuple[int, Session | None]: + with self._session_lock: + return self._session_generation, self._current_session + + def _set_session_if_current(self, session: Session, generation: int) -> bool: + with self._session_lock: + if generation != self._session_generation: + return False + self._current_session = session + self._session_generation += 1 + return True + + def _clear_session_if_current(self, generation: int) -> bool: + with self._session_lock: + if generation != self._session_generation: + return False + self._current_session = None + self._session_generation += 1 + return True diff --git a/src/volcano_sdk/errors.py b/src/volcano_sdk/errors.py index ae97c6ed..577855cc 100644 --- a/src/volcano_sdk/errors.py +++ b/src/volcano_sdk/errors.py @@ -37,6 +37,18 @@ class ConflictError(VolcanoError): """The request conflicts with the current resource state.""" +class SessionChangedError(ConflictError): + """An auth operation completed after the client session changed.""" + + def __init__(self) -> None: + """Create a deterministic stale-auth-operation error.""" + super().__init__( + "Session changed during authentication operation", + status=409, + code="auth_session_changed", + ) + + class RateLimitedError(VolcanoError): """The API rejected the request because of a rate limit.""" diff --git a/tests/unit/test_contract_bindings.py b/tests/unit/test_contract_bindings.py index 959a0a7e..1bac416f 100644 --- a/tests/unit/test_contract_bindings.py +++ b/tests/unit/test_contract_bindings.py @@ -15,7 +15,7 @@ ROOT = Path(__file__).parents[2] FEATURE_SHA256 = { - "auth.feature": "4145af3f1120331d100d8548b2abfcb98c0aeef552fd2b3304dc7e074e526ce2", + "auth.feature": "ca01037d4f7c40a1d17be3e79ff24d2433c2085e991332f594ca698466c33f8d", "database.feature": ( "4685b29357a621068b25984ff0de29cd4c504eebe5cfb597f0b999e29878a668" ), @@ -68,12 +68,14 @@ def test_every_contract_phrase_is_bound_verbatim() -> None: "the client acquires and releases the contract lock", 'the client selects the contract table where "slug" equals the fixture slug', "the client reads the current session", + "the client refreshes the current session", "the client signs in with the contract user's credentials", "the client uploads and downloads the contract object", "the confirmed contract user", "the current session belongs to the contract user", "the current session exposes access and refresh tokens", "the downloaded bytes equal the uploaded bytes", + "the refreshed session replaces the previous credentials", "the released lease is no longer held", "the stored object path equals the contract path", "the subscriber receives the contract message within 10 seconds", diff --git a/tests/unit/test_generated_transport.py b/tests/unit/test_generated_transport.py index 30e5e4b3..82508da4 100644 --- a/tests/unit/test_generated_transport.py +++ b/tests/unit/test_generated_transport.py @@ -7,18 +7,23 @@ from volcano_sdk._transport import GeneratedTransport -def test_generated_transport_calls_the_six_openapi_operations() -> None: +def test_generated_transport_calls_the_seven_openapi_operations() -> None: requests: list[httpx.Request] = [] def handle(request: httpx.Request) -> httpx.Response: requests.append(request) path = request.url.path - if path == "/auth/signin": + if path in {"/auth/signin", "/auth/refresh"}: + refreshing = path == "/auth/refresh" return httpx.Response( 200, json={ - "access_token": "access-token", - "refresh_token": "refresh-token", + "access_token": ( + "refreshed-access-token" if refreshing else "access-token" + ), + "refresh_token": ( + "refreshed-refresh-token" if refreshing else "refresh-token" + ), "user": { "id": "00000000-0000-4000-8000-000000000010", "email": "user@example.com", @@ -75,6 +80,10 @@ def handle(request: httpx.Request) -> httpx.Response: email="user@example.com", password="secret", ) + refresh = transport.auth_refresh( + authorization="anon-key", + refresh_token="refresh-token", + ) query = transport.query_database_select( authorization="access-token", database_name="main", @@ -108,6 +117,7 @@ def handle(request: httpx.Request) -> httpx.Response: ) assert auth.payload["user"]["id"] == "00000000-0000-4000-8000-000000000010" + assert refresh.payload["access_token"] == "refreshed-access-token" assert query.payload == {"data": [{"slug": "a"}], "count": 1} assert upload.payload["name"] == "a.txt" assert download.content == b"hello" @@ -117,11 +127,13 @@ def handle(request: httpx.Request) -> httpx.Response: "POST", "POST", "POST", + "POST", "GET", "POST", "DELETE", ] assert [request.headers["authorization"] for request in requests] == [ + "Bearer anon-key", "Bearer anon-key", "Bearer access-token", "Bearer access-token", @@ -134,15 +146,18 @@ def handle(request: httpx.Request) -> httpx.Response: "password": "secret", } assert json.loads(requests[1].content) == { + "refresh_token": "refresh-token", + } + assert json.loads(requests[2].content) == { "table": "items", "select": ["*"], "filters": [{"column": "slug", "operator": "eq", "value": "a"}], } - assert b"hello" in requests[2].content - assert json.loads(requests[4].content) == {"ttl_seconds": 30} - assert requests[4].headers["x-volcano-lock-token"] == ( + assert b"hello" in requests[3].content + assert json.loads(requests[5].content) == {"ttl_seconds": 30} + assert requests[5].headers["x-volcano-lock-token"] == ( "00000000-0000-4000-8000-000000000001" ) - assert requests[5].headers["x-volcano-lock-token"] == ( + assert requests[6].headers["x-volcano-lock-token"] == ( "00000000-0000-4000-8000-000000000001" ) diff --git a/tests/unit/test_state.py b/tests/unit/test_state.py index f9f7370c..f334dfaa 100644 --- a/tests/unit/test_state.py +++ b/tests/unit/test_state.py @@ -1,11 +1,20 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any +from typing import TYPE_CHECKING, Any import pytest -from volcano_sdk import Session, VolcanoClient +from volcano_sdk import ( + AuthenticationError, + ServerError, + Session, + SessionChangedError, + VolcanoClient, +) + +if TYPE_CHECKING: + from collections.abc import Callable @dataclass(frozen=True) @@ -19,6 +28,15 @@ class Response: class StateTransport: def __init__(self) -> None: self.next_access_token = "access-1" + self.refresh_response = Response( + 200, + { + "access_token": "access-2", + "refresh_token": "refresh-2", + "user": {"id": "user-123"}, + }, + ) + self.on_refresh: Callable[[], None] | None = None self.query_calls: list[dict[str, Any]] = [] self.authorizations: list[tuple[str, str]] = [] @@ -33,6 +51,12 @@ def auth_signin(self, **kwargs: Any) -> Response: }, ) + def auth_refresh(self, **kwargs: Any) -> Response: + self.authorizations.append(("refresh", kwargs["authorization"])) + if self.on_refresh is not None: + self.on_refresh() + return self.refresh_response + def query_database_select(self, **kwargs: Any) -> Response: self.authorizations.append(("query", kwargs["authorization"])) self.query_calls.append(kwargs["body"]) @@ -191,3 +215,70 @@ def test_auth_facade_rejects_incomplete_adoption_without_mutation(invalid: Any) assert client.auth.get_session() is previous assert transport.authorizations == calls_after_sign_in + + +def test_refresh_replaces_the_captured_session() -> None: + transport = StateTransport() + client = VolcanoClient(anon_key="anon", _transport=transport) + established = client.auth.sign_in(email="user@example.com", password="secret") + + refreshed = client.auth.refresh_session() + + assert refreshed == Session("access-2", "refresh-2", established.user_id) + assert client.auth.get_session() is refreshed + + +def test_refresh_without_a_session_fails_without_transport() -> None: + transport = StateTransport() + client = VolcanoClient(anon_key="anon", _transport=transport) + + with pytest.raises(AuthenticationError, match="No active session"): + client.auth.refresh_session() + + assert transport.authorizations == [] + + +def test_refresh_authentication_failure_clears_only_the_captured_session() -> None: + transport = StateTransport() + client = VolcanoClient(anon_key="anon", _transport=transport) + established = client.auth.sign_in(email="user@example.com", password="secret") + transport.refresh_response = Response(401, {"error": "expired"}) + + with pytest.raises(AuthenticationError, match="expired"): + client.auth.refresh_session() + + assert client.auth.get_session() is None + assert established.refresh_token == "refresh-access-1" + + +def test_refresh_server_failure_preserves_the_captured_session() -> None: + transport = StateTransport() + client = VolcanoClient(anon_key="anon", _transport=transport) + established = client.auth.sign_in(email="user@example.com", password="secret") + transport.refresh_response = Response(503, {"error": "unavailable"}) + + with pytest.raises(ServerError, match="unavailable"): + client.auth.refresh_session() + + assert client.auth.get_session() is established + + +def test_refresh_does_not_replace_a_session_established_during_the_request() -> None: + transport = StateTransport() + client = VolcanoClient(anon_key="anon", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + replacement = Session( + "replacement-access", + "replacement-refresh", + "replacement-user", + ) + + def replace_session() -> None: + client.auth.set_session(replacement) + + transport.on_refresh = replace_session + + with pytest.raises(SessionChangedError): + client.auth.refresh_session() + + assert client.auth.get_session() == replacement