from __future__ import annotations
import base64
import inspect
import time
import warnings
from contextlib import contextmanager
from contextvars import ContextVar
from typing import TYPE_CHECKING, Any, Literal, cast
import requests
from social_core.exceptions import (
AuthConfigurationError,
AuthInputError,
AuthPolicyError,
AuthProviderError,
AuthResponseError,
AuthSessionError,
ErrorStage,
SocialAuthBaseException,
)
from social_core.registry import REGISTRY
from social_core.utils import (
constant_time_compare,
http_error,
module_member,
normalize_user_names,
parse_qs,
social_logger,
user_agent,
user_is_authenticated,
)
if TYPE_CHECKING:
from collections.abc import Iterator, Mapping
from requests import Response
from requests.auth import AuthBase
from social_core.storage import PartialMixin, PipelineUserProtocol, UserProtocol
from social_core.strategy import BaseStrategy, HttpResponseProtocol
[docs]
class BaseAuth:
"""A authentication backend that authenticates the user based on
the provider response.
Set ``ASSOCIATION_ONLY`` to connect provider access to an authenticated local
user without updating their profile. Browser flows must start through
``do_auth(user=...)`` and validate callbacks with
``validate_association_state()`` before processing provider credentials.
Authentication and disconnect partials remain bound to that local user.
"""
name = "" # provider name, it's stored in database
title: str | None = None # human-readable sign-in label
icon: str | None = None # filename in static/social_auth/icons
supports_inactive_user = False # Django auth
ID_KEY: str = ""
LEGACY_ID_KEYS: tuple[str, ...] = ()
MUTABLE_ID_KEYS: tuple[str, ...] = ()
EXTRA_DATA: list[str | tuple[str, str] | tuple[str, str, bool]] | None = None
GET_ALL_EXTRA_DATA = False
REQUIRES_EMAIL_VALIDATION = False
REQUIRES_USER_ID: bool = False
SEND_USER_AGENT = True
ASSOCIATION_ONLY = False
def __init__(
self, strategy: BaseStrategy | None = None, redirect_uri: str | None = None
) -> None:
self.strategy: BaseStrategy = (
strategy if strategy is not None else REGISTRY.default_strategy
)
self.redirect_uri = redirect_uri
self._pipeline_type: ContextVar[str] = ContextVar(
"pipeline_type", default="authentication"
)
self._mutable_id_key_warned = False
self.data = self.strategy.request_data()
self.redirect_uri = self.strategy.absolute_uri(self.redirect_uri)
[docs]
def log_debug(self, message, *args) -> None:
social_logger.debug(f"{self.name}: {message}", *args)
[docs]
def log_warning(self, message, *args) -> None:
social_logger.warning(f"{self.name}: {message}", *args)
[docs]
def setting(self, name: str, default=None):
"""Return setting value from strategy"""
return self.strategy.setting(name, default=default, backend=self)
[docs]
def prepare_auth(self, user: UserProtocol | None = None) -> None:
"""Bind association-only authorization to an authenticated local user."""
if self.ASSOCIATION_ONLY:
user = self.require_association_user(user, stage="begin")
self.strategy.session_set(
f"{self.name}_state",
{"state": self.strategy.random_string(32), "user_id": str(user.id)},
)
[docs]
def require_association_user(
self, user: UserProtocol | None, *, stage: ErrorStage = "user_info"
) -> UserProtocol:
"""Require an authenticated local user for an association-only flow."""
if user is None or not user_is_authenticated(user):
raise AuthSessionError(
self,
"Association requires an authenticated user",
code="session_context_missing",
stage=stage,
)
return user
[docs]
def get_association_state(self) -> str:
"""Return the token from a prepared, user-bound authorization context."""
context = self.strategy.session_get(f"{self.name}_state")
if not isinstance(context, dict):
raise AuthSessionError(
self, "state", code="session_context_missing", stage="callback"
)
state = context.get("state")
if not isinstance(state, str) or not state:
raise AuthSessionError(
self, "state", code="session_context_missing", stage="callback"
)
return state
[docs]
def validate_association_state(
self, request_state: Any, user: UserProtocol | None = None
) -> str:
"""Validate and consume authorization state bound to the current user."""
if not request_state:
raise AuthInputError(
self, parameter="state", code="missing_parameter", stage="callback"
)
state = self.get_association_state()
if not isinstance(request_state, (str, bytes)) or not constant_time_compare(
request_state, state
):
raise AuthSessionError(self, code="state_mismatch", stage="callback")
user = self.require_association_user(user, stage="callback")
context = self.strategy.session_get(f"{self.name}_state")
if context.get("user_id") != str(user.id):
raise AuthSessionError(
self,
"Association user mismatch",
code="user_mismatch",
stage="callback",
)
self.strategy.session_pop(f"{self.name}_state")
return state
[docs]
def association_user_id_key(self, pipeline_type: str = "authentication") -> str:
"""Name of the initiator binding saved in pipeline arguments."""
operation = "disconnect" if pipeline_type == "disconnect" else "association"
return f"{self.name}_{operation}_user_id"
def _bind_association_user(
self, kwargs: dict[str, Any], pipeline_type: str = "authentication"
) -> None:
stage: ErrorStage = (
"disconnect" if pipeline_type == "disconnect" else "user_info"
)
user = self.require_association_user(kwargs.get("user"), stage=stage)
key = self.association_user_id_key(pipeline_type)
if key in kwargs and kwargs[key] != str(user.id):
raise AuthSessionError(
self,
"Association user mismatch",
code="user_mismatch",
stage=stage,
)
kwargs[key] = str(user.id)
[docs]
def start(self) -> HttpResponseProtocol:
if self.ASSOCIATION_ONLY:
self.get_association_state()
if self.uses_redirect():
return self.strategy.redirect(self.auth_url())
return self.strategy.html(self.auth_html())
[docs]
def complete(self, *args, **kwargs) -> HttpResponseProtocol | UserProtocol | None:
return self.auth_complete(*args, **kwargs)
[docs]
def auth_url(self) -> str:
"""Must return redirect URL to auth provider"""
raise NotImplementedError("Implement in subclass")
[docs]
def auth_html(self) -> str:
"""Must return login HTML content returned by provider"""
return "Implement in subclass"
[docs]
def auth_complete(
self, *args, **kwargs
) -> HttpResponseProtocol | UserProtocol | None:
"""Completes login process, must return user instance"""
raise NotImplementedError("Implement in subclass")
[docs]
def process_error(self, data, *, stage: ErrorStage = "callback") -> None:
"""Hook to process provider response errors.
Default implementation is a no-op. Backends that can detect
provider-specific error payloads should override this method and
raise an appropriate exception when needed.
"""
def _process_error(self, data, *, stage: ErrorStage) -> None:
"""Call the error hook, including overrides predating the stage keyword."""
kwargs: dict[str, ErrorStage]
try:
inspect.signature(self.process_error).bind(data, stage=stage)
except TypeError:
# Check the signature before calling so a TypeError inside the hook
# is propagated without invoking it twice.
kwargs = {}
else:
kwargs = {"stage": stage}
try:
self.process_error(data, **kwargs)
except SocialAuthBaseException as error:
error.stage = stage
raise
[docs]
def authenticate(
self, *args, **kwargs
) -> UserProtocol | HttpResponseProtocol | None:
"""Authenticate user using social credentials
Authentication is made if this is the correct backend, backend
verification is made by kwargs inspection for current backend
name presence.
"""
# Validate backend and arguments. Require that the Social Auth
# response be passed in as a keyword argument, to make sure we
# don't match the username/password calling conventions of
# authenticate.
if (
"backend" not in kwargs
or kwargs["backend"].name != self.name
or "strategy" not in kwargs
or "response" not in kwargs
):
return None
self.strategy = kwargs.get("strategy") or self.strategy
self.redirect_uri = kwargs.get("redirect_uri") or self.redirect_uri
self.data = self.strategy.request_data()
if self.ASSOCIATION_ONLY:
self._bind_association_user(kwargs)
kwargs.setdefault("is_new", False)
pipeline = self.strategy.get_pipeline(self)
args, kwargs = self.strategy.clean_authenticate_args(*args, **kwargs)
return self.pipeline(pipeline, *args, **kwargs)
[docs]
def pipeline(
self, pipeline, pipeline_index: int = 0, *args, **kwargs
) -> UserProtocol | HttpResponseProtocol | None:
token = self._pipeline_type.set("authentication")
try:
out = self.run_pipeline(pipeline, pipeline_index, *args, **kwargs)
finally:
self._pipeline_type.reset(token)
if not isinstance(out, dict):
return cast("HttpResponseProtocol", out)
user = cast("UserProtocol | None", out.get("user"))
if user:
pipeline_user = cast("PipelineUserProtocol", user)
pipeline_user.social_user = cast("Any", out.get("social"))
pipeline_user.is_new = bool(out.get("is_new"))
return user
[docs]
def disconnect(self, *args, **kwargs) -> dict:
if self.ASSOCIATION_ONLY:
self._bind_association_user(kwargs, "disconnect")
pipeline = self.strategy.get_disconnect_pipeline(self)
kwargs["name"] = self.name
kwargs["user_storage"] = self.strategy.storage.user
token = self._pipeline_type.set("disconnect")
try:
return self.run_pipeline(pipeline, *args, **kwargs)
finally:
self._pipeline_type.reset(token)
@property
def pipeline_type(self) -> str:
"""Type of the currently executing pipeline, saved with its partials."""
return self._pipeline_type.get()
[docs]
def run_pipeline(
self, pipeline: list[str], pipeline_index=0, *args, **kwargs
) -> dict:
out = kwargs.copy()
out.setdefault("strategy", self.strategy)
out.setdefault("backend", out.pop(self.name, None) or self)
out.pop("request", None)
out.setdefault("details", {})
if (
not isinstance(pipeline_index, int)
or pipeline_index < 0
or pipeline_index >= len(pipeline)
):
pipeline_index = 0
for idx, name in enumerate(pipeline[pipeline_index:]):
out["pipeline_index"] = pipeline_index + idx
func = module_member(name)
result = func(*args, **out) or {}
if not isinstance(result, dict):
return result
out.update(result)
return out
[docs]
def auth_allowed(self, response, details):
"""Return True if the user should be allowed to authenticate, by
default check if email is whitelisted (if there's a whitelist)"""
emails = [
email.lower()
for email in cast("list[str]", self.setting("WHITELISTED_EMAILS", []))
]
domains = [
domain.lower()
for domain in cast("list[str]", self.setting("WHITELISTED_DOMAINS", []))
]
email = details.get("email")
allowed = True
if email and (emails or domains):
email = email.lower()
parts = email.split("@", 1)
if len(parts) != 2:
allowed = False
else:
domain = parts[1]
allowed = email in emails or domain in domains
allow_groups = self.get_group_setting("ALLOW_GROUPS", response, [])
if not isinstance(allow_groups, (list, tuple, set)) or any(
not isinstance(group, str) or not group for group in allow_groups
):
raise AuthConfigurationError(
self, code="invalid_setting", parameter="ALLOW_GROUPS", stage="pipeline"
)
if allow_groups:
# Backends override the default no-extraction implementation.
# pylint: disable-next=assignment-from-none
groups = self.get_user_groups(response)
if groups is None:
raise AuthConfigurationError(
self,
code="missing_setting",
parameter="group extraction",
stage="pipeline",
)
allowed = allowed and bool(set(groups).intersection(allow_groups))
return allowed
[docs]
def get_user_groups(self, response) -> list[str] | None:
"""Return normalized external memberships, or None when disabled."""
return None
[docs]
def get_group_setting(self, name: str, response, default=None):
"""Resolve group configuration for this provider (or a SAML IdP)."""
return self.setting(name, default)
[docs]
def get_group_mappings(self):
"""Enumerate independently managed membership sources."""
return [("", self.setting("GROUPS_MAP", {}))]
[docs]
def get_group_source(self, response) -> str:
"""Identify the provider's active membership source."""
return ""
[docs]
def id_key(self) -> str:
"""Return the ID_KEY to use for this backend, checking settings first."""
configured = self.setting("ID_KEY")
id_key = configured or self.ID_KEY
if (
configured
and id_key in self.MUTABLE_ID_KEYS
and not self._mutable_id_key_warned
):
self.log_warning(
"configured ID_KEY %r is mutable and is unsafe as an account identifier",
id_key,
)
self._mutable_id_key_warned = True
return id_key
[docs]
def get_user_id_for_key(self, details, response, id_key: str):
"""Return a user identifier selected by an explicit response key."""
return self.get_user_id_from_sources(details, response, id_key=id_key)
[docs]
def get_legacy_user_ids(self, details, response) -> list[str]:
"""Return current values of identifiers used by older releases."""
if self.setting("ID_KEY"):
return []
identifiers = []
for id_key in self.LEGACY_ID_KEYS:
try:
identifier = self.get_user_id_for_key(details, response, id_key)
except AuthResponseError as error:
if error.code != "missing_claim":
raise
continue
value = str(identifier)
if value not in identifiers:
identifiers.append(value)
return identifiers
[docs]
def get_user_id(self, details, response):
"""Return a unique ID for the current user, by default from server
response or details."""
id_key = self.id_key()
if self.REQUIRES_USER_ID or self.setting("ID_KEY"):
return self.get_user_id_for_key(details, response, id_key)
if details:
user_id = details.get(id_key)
if user_id:
return user_id
return response.get(id_key)
[docs]
def get_user_id_from_sources(
self,
*sources: Mapping[str, Any] | None,
id_key: str | None = None,
):
"""Return the selected user ID from mappings or fail clearly.
Sources are searched in order for the configured or explicitly passed
ID key. Missing, ``None``, and empty-string values are rejected.
"""
if id_key is None:
id_key = self.id_key()
for source in sources:
if source is not None:
user_id = source.get(id_key)
if user_id is not None and user_id != "":
return user_id
raise AuthResponseError(
self, claim=id_key, code="missing_claim", stage="user_info"
)
[docs]
def get_user_details(self, response) -> dict[str, Any]:
"""Return provider-supplied user details in a known internal structure.
Leave name conversion to the social_names pipeline step.
The returned dictionary can contain:
``username``
Username, if any.
``email``
User email, if any.
``fullname``
User full name, if any.
``first_name``
User first name, if any.
``last_name``
User last name, if any.
"""
raise NotImplementedError("Implement in subclass")
[docs]
def get_refresh_token(self, extra_data: dict[str, Any]) -> str | None:
"""Select a stored renewal credential, or None when renewal is unavailable.
Backends that exchange an access token should override this method.
"""
token = extra_data.get("refresh_token")
return token if isinstance(token, str) and token else None
[docs]
def get_refresh_token_kwargs(self, extra_data: dict[str, Any]) -> dict[str, Any]:
"""Return default refresh arguments from stored account credentials."""
return {}
[docs]
def get_user_names(self, fullname="", first_name="", last_name=""):
warnings.warn(
"BaseAuth.get_user_names() is deprecated. Return provider-supplied "
"names from get_user_details() and use the "
"social_core.pipeline.social_auth.social_names pipeline step.",
DeprecationWarning,
stacklevel=2,
)
return normalize_user_names(fullname, first_name, last_name)
[docs]
def get_user(self, user_id):
"""
Return user with given ID from the User model used by this backend.
This is called by django.contrib.auth.middleware.
"""
return self.strategy.get_user(user_id)
[docs]
def continue_pipeline(
self, partial: PartialMixin
) -> UserProtocol | HttpResponseProtocol | None:
"""Continue previous halted pipeline"""
with self._partial_pipeline_context(partial):
return self.strategy.authenticate(
self, *partial.args, pipeline_index=partial.next_step, **partial.kwargs
)
[docs]
def continue_disconnect_pipeline(
self, partial: PartialMixin
) -> dict | HttpResponseProtocol:
"""Continue a halted disconnect with its effective request data."""
with self._partial_pipeline_context(partial, pipeline_type="disconnect"):
return self.disconnect(
*partial.args, pipeline_index=partial.next_step, **partial.kwargs
)
@contextmanager
def _partial_pipeline_context(
self, partial: PartialMixin, pipeline_type: str = "authentication"
) -> Iterator[None]:
if partial.pipeline_type != pipeline_type:
raise AuthPolicyError(
self, code="authentication_disallowed", stage="callback"
)
previous_data = self.data
with self.strategy.pipeline_request_data(partial.request_data):
self.data = self.strategy.request_data()
try:
yield
finally:
self.data = previous_data
[docs]
def validate_partial_pipeline(
self, partial: PartialMixin, user: UserProtocol | None = None
) -> None:
"""Validate backend-specific requirements before resuming a pipeline."""
if self.ASSOCIATION_ONLY:
user = self.require_association_user(user, stage="callback")
key = self.association_user_id_key(partial.pipeline_type)
if partial.kwargs.get(key) != str(user.id):
raise AuthSessionError(
self,
"Association user mismatch",
code="user_mismatch",
stage="callback",
)
[docs]
def uses_redirect(self) -> bool:
"""Return True if this provider uses redirect url method,
otherwise return false."""
return True
[docs]
def request( # noqa: PLR0913
self,
url: str,
*,
method: Literal["GET", "POST", "DELETE"] = "GET",
headers: Mapping[str, str | bytes] | None = None,
data: dict | None = None,
json: dict | None = None,
auth: tuple[str, str] | AuthBase | None = None,
params: dict | None = None,
timeout: float | None = None,
stage: ErrorStage = "user_info",
) -> Response:
headers = {} if headers is None else dict(headers)
proxies = self.setting("PROXIES")
verify = self.setting("VERIFY_SSL", True)
if timeout is None:
timeout = (
self.setting("REQUESTS_TIMEOUT")
or self.setting("URLOPEN_TIMEOUT")
or 5.0
)
if self.SEND_USER_AGENT and "User-Agent" not in headers:
headers["User-Agent"] = self.setting("USER_AGENT") or user_agent()
try:
response = requests.request(
method,
url,
headers=headers,
data=data,
json=json,
auth=auth,
params=params,
timeout=timeout,
proxies=proxies,
verify=verify,
)
except requests.exceptions.SSLError as error:
raise AuthProviderError(self, code="tls_error", stage=stage) from error
except requests.Timeout as error:
raise AuthProviderError(self, code="timeout", stage=stage) from error
except (
requests.exceptions.InvalidURL,
requests.exceptions.InvalidSchema,
requests.exceptions.MissingSchema,
) as error:
raise AuthConfigurationError(
self, code="invalid_setting", parameter="url", stage=stage
) from error
except requests.ConnectionError as error:
raise AuthProviderError(
self, code="connection_failed", stage=stage
) from error
except requests.HTTPError as error:
raise http_error(self, error, stage=stage) from error
except requests.RequestException as error:
raise AuthProviderError(self, stage=stage) from error
try:
response.raise_for_status()
except requests.HTTPError as error:
raise http_error(self, error, stage=stage) from error
return response
[docs]
def get_json( # noqa: PLR0913, PLR0917
self,
url: str,
method: Literal["GET", "POST", "DELETE"] = "GET",
headers: Mapping[str, str | bytes] | None = None,
data: dict | None = None,
json: dict | None = None,
auth: tuple[str, str] | AuthBase | None = None,
params: dict | None = None,
timeout: float | None = None,
stage: ErrorStage = "user_info",
) -> dict[Any, Any]:
response = self.request(
url,
method=method,
headers=headers,
data=data,
json=json,
auth=auth,
params=params,
timeout=timeout,
stage=stage,
)
try:
return response.json()
except ValueError as error:
raise AuthResponseError(
self, code="malformed_response", stage=stage
) from error
[docs]
def get_querystring(self, url, *args, **kwargs) -> dict[str, str]:
return parse_qs(self.request(url, *args, **kwargs).text)
[docs]
def get_key_and_secret(self) -> tuple[str, str]:
"""Return tuple with Consumer Key and Consumer Secret for current
service provider. Must return (key, secret), order *must* be respected.
"""
return cast("str", self.setting("KEY")), cast("str", self.setting("SECRET"))
[docs]
def get_key_and_secret_basic_auth(self) -> bytes:
"""Generate HTTP Basic Authentication header value from KEY and SECRET.
Returns:
Basic authentication value in the format b"Basic <base64-encoded-credentials>"
"""
key, secret = self.get_key_and_secret()
credentials = f"{key}:{secret}".encode()
encoded = base64.b64encode(credentials)
return b"Basic " + encoded