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
62 changes: 62 additions & 0 deletions campus_python/auth/v1/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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,
Expand All @@ -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,
)
183 changes: 156 additions & 27 deletions campus_python/auth/v1/oauth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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,
Expand Down
35 changes: 35 additions & 0 deletions campus_python/errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,33 @@ class FieldError:
message: str


# Campus envelope error codes that do not follow the AUTH_<OAUTH_ERROR>
# 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_<OAUTH_ERROR> (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.

Expand Down Expand Up @@ -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):
Expand Down
Loading
Loading