Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
10 changes: 10 additions & 0 deletions features/contract/auth.feature
Original file line number Diff line number Diff line change
Expand Up @@ -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
3 changes: 2 additions & 1 deletion features/contract_support.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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] = []
Expand Down
18 changes: 18 additions & 0 deletions features/steps/sdk_contract_steps.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 2 additions & 0 deletions src/volcano_sdk/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
NotFoundError,
RateLimitedError,
ServerError,
SessionChangedError,
TransportError,
ValidationError,
VolcanoError,
Expand All @@ -21,6 +22,7 @@
"RateLimitedError",
"ServerError",
"Session",
"SessionChangedError",
"TransportError",
"ValidationError",
"VolcanoClient",
Expand Down
25 changes: 24 additions & 1 deletion src/volcano_sdk/_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,14 +11,15 @@

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 (
download_storage_object,
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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
*,
Expand Down
60 changes: 53 additions & 7 deletions src/volcano_sdk/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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."""

Expand All @@ -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."""
Expand Down Expand Up @@ -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
34 changes: 30 additions & 4 deletions src/volcano_sdk/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

from __future__ import annotations

import threading
from typing import TYPE_CHECKING

from ._transport import GeneratedTransport, Transport
Expand Down Expand Up @@ -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
Expand All @@ -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."""
Expand All @@ -66,14 +69,37 @@ 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:
raise RuntimeError(_NO_SERVICE_KEY)
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
12 changes: 12 additions & 0 deletions src/volcano_sdk/errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
4 changes: 3 additions & 1 deletion tests/unit/test_contract_bindings.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@

ROOT = Path(__file__).parents[2]
FEATURE_SHA256 = {
"auth.feature": "4145af3f1120331d100d8548b2abfcb98c0aeef552fd2b3304dc7e074e526ce2",
"auth.feature": "ca01037d4f7c40a1d17be3e79ff24d2433c2085e991332f594ca698466c33f8d",
"database.feature": (
"4685b29357a621068b25984ff0de29cd4c504eebe5cfb597f0b999e29878a668"
),
Expand Down Expand Up @@ -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",
Expand Down
Loading
Loading