from __future__ import annotations
import contextlib
import functools
import hmac
import logging
import re
import time
import unicodedata
from dataclasses import dataclass
from importlib import import_module
from typing import TYPE_CHECKING, Any, cast
from urllib.parse import parse_qs as battery_parse_qs
from urllib.parse import unquote, urlencode, urlparse, urlunparse
import requests
import social_core
from social_core.pipeline.utils import is_dict_type, to_plain_dict
from .exceptions import (
AuthCanceled,
AuthConfigurationError,
AuthCredentialError,
AuthProviderError,
ErrorStage,
SocialAuthBaseException,
)
if TYPE_CHECKING:
from collections.abc import Collection
from .backends.base import BaseAuth
from .storage import PartialMixin, UserProtocol
from .strategy import BaseStrategy, HttpResponseProtocol
SETTING_PREFIX = "SOCIAL_AUTH"
DEFAULT_REDIRECT_SCHEMES = ("http", "https")
PARTIAL_TOKEN_SESSION_NAME = "partial_pipeline_token"
PARTIAL_TOKEN_PENDING_SESSION_NAME = "partial_pipeline_pending_token"
PARTIAL_TOKEN_PENDING_REQUEST_SESSION_NAME = "partial_pipeline_pending_request"
PARTIAL_TOKEN_PENDING_CONFIRMATION_SESSION_NAME = (
"partial_pipeline_pending_confirmation"
)
PARTIAL_PIPELINE_ALLOW_EXTERNAL_RESUME = "allow_external_resume"
social_logger = logging.getLogger("social")
[docs]
@dataclass
class PartialPipelineResult:
partial: PartialMixin | None = None
response: HttpResponseProtocol | None = None
halt: bool = False
[docs]
@dataclass
class PartialPipelineSelection:
token: str | None = None
owns_token: bool = False
pending_resume: bool = False
[docs]
def normalize_user_names(
fullname: str | None = "",
first_name: str | None = "",
last_name: str | None = "",
*,
firstlast_from_full: bool = True,
full_from_firstlast: bool = True,
) -> tuple[str, str, str]:
"""Fill missing name representations without replacing supplied names."""
fullname = (fullname or "").strip()
first_name = (first_name or "").strip()
last_name = (last_name or "").strip()
if firstlast_from_full and fullname and not (first_name or last_name):
first_name, _, last_name = fullname.partition(" ")
first_name, last_name = first_name.strip(), last_name.strip()
if full_from_firstlast and not fullname:
if first_name and re.search(
rf"(?<!\S){re.escape(first_name)}(?!\S)", last_name
):
fullname = last_name
else:
fullname = f"{first_name} {last_name}".strip()
return fullname, first_name, last_name
[docs]
def module_member(name):
mod, member = name.rsplit(".", 1)
module = import_module(mod)
return getattr(module, member)
[docs]
def user_agent() -> str:
"""Builds a simple User-Agent string to send in requests"""
return f"social-auth-{social_core.__version__}"
[docs]
def url_add_parameters(
url: str, params: dict[str, str] | None, _unquote_query: bool = False
) -> str:
"""Adds parameters to URL, parameter will be repeated if already present"""
if params:
fragments = list(urlparse(url))
value = parse_qs(fragments[4])
value.update(params)
fragments[4] = urlencode(value)
if _unquote_query:
fragments[4] = unquote(fragments[4])
url = urlunparse(fragments)
return url
[docs]
def to_setting_name(*names: str) -> str:
return "_".join([name.upper().replace("-", "_") for name in names if name])
[docs]
def setting_name(*names: str) -> str:
return to_setting_name(*((SETTING_PREFIX, *names)))
[docs]
def normalize_redirect_schemes(schemes: Collection[str]) -> set[str]:
"""URI schemes are case-insensitive, and urlparse lowercases them."""
return {scheme.lower() for scheme in schemes}
[docs]
def sanitize_redirect(
hosts: list[str],
redirect_to: str | Any,
allowed_schemes: Collection[str] | None = None,
) -> str | None:
"""
Given a list of hostnames and an untrusted URL to redirect to,
this method tests it to make sure it isn't garbage/harmful
and returns it, else returns None, similar as how's it done
on django.contrib.auth.views.
``allowed_schemes`` defaults to http and https. Deployments that need to
hand control back to a native application can add a private-use URI scheme
(RFC 8252) through the ``ALLOWED_REDIRECT_SCHEMES`` setting.
"""
# Avoid redirect on evil URLs like ///evil.com and URLs containing
# backslashes or control characters that browsers may normalize.
if (
not redirect_to
or not isinstance(redirect_to, str)
or redirect_to.startswith("///")
or "\\" in redirect_to
or any(unicodedata.category(char)[0] == "C" for char in redirect_to)
):
return None
schemes = (
set(DEFAULT_REDIRECT_SCHEMES)
if allowed_schemes is None
else normalize_redirect_schemes(allowed_schemes)
)
try:
parsed_url = urlparse(redirect_to)
scheme = parsed_url.scheme
if scheme:
if scheme not in schemes:
return None
if scheme not in DEFAULT_REDIRECT_SCHEMES:
# A private-use URI scheme is registered by a single native
# application, so the scheme itself is the trust boundary and
# the host allowlist does not apply.
return redirect_to
if not parsed_url.netloc:
return None
# Don't redirect to a host that's not in the list
netloc = parsed_url.netloc or hosts[0]
except (IndexError, TypeError, AttributeError, ValueError):
return None
return redirect_to if netloc in hosts else None
[docs]
def user_is_authenticated(user: UserProtocol | None) -> bool:
if user and hasattr(user, "is_authenticated"):
if callable(user.is_authenticated):
authenticated = user.is_authenticated()
else:
authenticated = user.is_authenticated
elif user:
authenticated = True
else:
authenticated = False
return authenticated
[docs]
def user_is_active(user: UserProtocol | None) -> bool:
if user and hasattr(user, "is_active"):
is_active = user.is_active() if callable(user.is_active) else user.is_active
elif user:
is_active = True
else:
is_active = False
return is_active
# This slugify version was borrowed from django revision a61dbd6
[docs]
def slugify(value):
"""Converts to lowercase, removes non-word characters (alphanumerics
and underscores) and converts spaces to hyphens. Also strips leading
and trailing whitespace."""
value = (
unicodedata.normalize("NFKD", str(value))
.encode("ascii", "ignore")
.decode("ascii")
)
value = re.sub(r"[^\w\s-]", "", value).strip().lower()
return re.sub(r"[-\s]+", "-", value)
[docs]
def first(func, items):
"""Return the first item in the list for what func returns True"""
for item in items:
if func(item):
return item
return None
[docs]
def parse_qs(value):
"""Like urlparse.parse_qs but transform list values to single items"""
return drop_lists(battery_parse_qs(value))
[docs]
def get_querystring(url: str):
return parse_qs(urlparse(url).query)
[docs]
def drop_lists(value):
out = {}
for key, val in value.items():
val = val[0]
if isinstance(key, bytes):
key = str(key, "utf-8")
if isinstance(val, bytes):
val = str(val, "utf-8")
out[key] = val
return out
def _partial_pipeline_matches_request(
backend: BaseAuth,
partial: PartialMixin | None,
request_data: dict[str, Any],
pipeline_type: str,
) -> bool:
if (
not partial
or partial.backend != backend.name
or partial.pipeline_type != pipeline_type
):
return False
# Normally when resuming a pipeline, request_data will be empty. We only
# need to check for a uid match if new data was provided (i.e. if current
# request specifies the ID_KEY).
id_key = backend.id_key()
if id_key and id_key in request_data:
id_from_partial = partial.kwargs.get("uid")
id_from_request = request_data.get(id_key)
return id_from_partial == id_from_request
return True
def _extend_partial_pipeline(
partial: PartialMixin,
request_data: dict[str, Any],
user: UserProtocol | None,
kwargs: dict[str, Any],
) -> PartialMixin:
if user: # don't update user if it's None
kwargs.setdefault("user", user)
partial.request_data = request_data
kwargs.pop("request", None)
partial.extend_kwargs(kwargs)
return partial
def _select_partial_pipeline_token(
request_token: str | None,
session_token: str | None,
pending_token: str | None,
confirmation_requested: bool,
) -> PartialPipelineSelection:
if confirmation_requested and pending_token:
selected_token = request_token or pending_token
pending_resume = selected_token == pending_token
return PartialPipelineSelection(
token=selected_token,
owns_token=pending_resume,
pending_resume=pending_resume,
)
if request_token and request_token == session_token:
return PartialPipelineSelection(token=request_token, owns_token=True)
if request_token:
return PartialPipelineSelection(token=request_token)
return PartialPipelineSelection(token=session_token, owns_token=bool(session_token))
def _partial_pipeline_requires_confirmation(
partial: PartialMixin,
request_token: str | None,
request_data: dict[str, Any],
pending_resume: bool,
) -> bool:
return bool(
not pending_resume
and partial.data.get(PARTIAL_PIPELINE_ALLOW_EXTERNAL_RESUME)
and (request_token or request_data)
)
def _confirmed_partial_pipeline_request_data(
backend: BaseAuth,
request_data: dict[str, Any],
) -> dict[str, Any] | None:
if not backend.strategy.partial_pipeline_external_resume_confirmed(
backend, request_data
):
return None
pending_request_data = backend.strategy.from_session_value(
backend.strategy.session_get(PARTIAL_TOKEN_PENDING_REQUEST_SESSION_NAME, {})
or {}
)
return {**pending_request_data, **request_data}
def _external_partial_pipeline_result(
backend: BaseAuth,
partial: PartialMixin,
selected_token: str,
request_data: dict[str, Any],
) -> PartialPipelineResult:
response = backend.strategy.partial_pipeline_external_resume_confirmation(
backend, partial, request_data
)
if response is None:
return PartialPipelineResult(halt=True)
backend.strategy.session_set(PARTIAL_TOKEN_PENDING_SESSION_NAME, selected_token)
backend.strategy.session_set(
PARTIAL_TOKEN_PENDING_REQUEST_SESSION_NAME,
backend.strategy.to_session_value(
to_plain_dict(request_data) if is_dict_type(request_data) else request_data
),
)
return PartialPipelineResult(response=response)
[docs]
def partial_pipeline_result(
backend: BaseAuth,
user: UserProtocol | None = None,
partial_token: str | None = None,
*args,
pipeline_type: str = "authentication",
**kwargs,
) -> PartialPipelineResult:
request_data = backend.strategy.get_request_data()
partial_argument_name = cast(
"str", backend.setting("PARTIAL_PIPELINE_TOKEN_NAME", "partial_token")
)
request_token = cast(
"str | None", partial_token or request_data.get(partial_argument_name)
)
session_token = backend.strategy.session_get(PARTIAL_TOKEN_SESSION_NAME, None)
pending_token = backend.strategy.session_get(
PARTIAL_TOKEN_PENDING_SESSION_NAME, None
)
confirmation_parameter = backend.setting(
"PARTIAL_PIPELINE_EXTERNAL_RESUME_CONFIRMATION_PARAMETER",
"partial_pipeline_confirm",
)
confirmation_requested = (
bool(confirmation_parameter) and confirmation_parameter in request_data
)
selection = _select_partial_pipeline_token(
request_token=request_token,
session_token=session_token,
pending_token=pending_token,
confirmation_requested=confirmation_requested,
)
if not selection.token:
return PartialPipelineResult()
result = PartialPipelineResult(halt=bool(request_token or confirmation_requested))
effective_request_data = request_data
if selection.pending_resume:
confirmed_request_data = _confirmed_partial_pipeline_request_data(
backend, request_data
)
if confirmed_request_data is None:
return PartialPipelineResult(halt=True)
effective_request_data = confirmed_request_data
partial: PartialMixin | None = backend.strategy.partial_load(selection.token)
partial_matches = _partial_pipeline_matches_request(
backend, partial, effective_request_data, pipeline_type
)
if partial and partial_matches:
backend.validate_partial_pipeline(partial, user)
if _partial_pipeline_requires_confirmation(
partial,
request_token,
effective_request_data,
selection.pending_resume,
):
result = _external_partial_pipeline_result(
backend, partial, selection.token, effective_request_data
)
elif selection.owns_token:
result = PartialPipelineResult(
partial=_extend_partial_pipeline(
partial, effective_request_data, user, kwargs
)
)
elif partial.data.get(PARTIAL_PIPELINE_ALLOW_EXTERNAL_RESUME):
result = _external_partial_pipeline_result(
backend, partial, selection.token, effective_request_data
)
else:
result = PartialPipelineResult(halt=True)
elif selection.owns_token:
backend.strategy.clean_partial_pipeline(selection.token)
return result
[docs]
def partial_pipeline_data(
backend: BaseAuth,
user: UserProtocol | None = None,
partial_token: str | None = None,
*args,
**kwargs,
) -> PartialMixin | None:
return partial_pipeline_result(
backend, user, partial_token, *args, **kwargs
).partial
[docs]
def build_absolute_uri(host_url: str, path: str | None = None) -> str:
"""Build absolute URI with given (optional) path"""
path = path or ""
if path.startswith(("http://", "https://")):
return path
if host_url.endswith("/") and path.startswith("/"):
path = path[1:]
return host_url + path
[docs]
def constant_time_compare(val1: str | bytes, val2: str | bytes) -> bool:
"""Compare two values and prevent timing attacks for cryptographic use."""
if isinstance(val1, str):
val1 = val1.encode("utf-8")
if isinstance(val2, str):
val2 = val2.encode("utf-8")
return hmac.compare_digest(val1, val2)
[docs]
def get_allowed_redirect_schemes(backend: BaseAuth) -> set[str]:
return normalize_redirect_schemes(
cast(
"Collection[str]",
backend.setting("ALLOWED_REDIRECT_SCHEMES", DEFAULT_REDIRECT_SCHEMES),
)
)
[docs]
def is_private_use_redirect(
value: str | None, allowed_schemes: Collection[str] | None = None
) -> bool:
"""
Whether ``value`` uses a non-web scheme that has been explicitly allowed.
URI construction helpers such as ``build_absolute_uri`` only preserve http
and https, so callers must check this before turning a redirect candidate
into an absolute URI.
"""
if not value or not isinstance(value, str) or not allowed_schemes:
return False
try:
scheme = urlparse(value).scheme
except ValueError:
return False
return (
bool(scheme)
and scheme not in DEFAULT_REDIRECT_SCHEMES
and scheme in normalize_redirect_schemes(allowed_schemes)
)
[docs]
def is_url(value: str | None, allowed_schemes: Collection[str] | None = None) -> bool:
if value is None:
return False
if value.startswith(("http://", "https://", "/")):
return True
return is_private_use_redirect(value, allowed_schemes)
[docs]
def setting_url(backend: BaseAuth, *names: str | None) -> str | None:
allowed_schemes = get_allowed_redirect_schemes(backend)
for name in names:
# Name can actually None, value or setting name
if not name:
continue
if is_url(name, allowed_schemes):
return name
value = backend.setting(name)
if is_url(value, allowed_schemes):
return value
return None
[docs]
def provider_error(
backend: BaseAuth,
data: object,
*,
stage: ErrorStage = "callback",
status_code: int | None = None,
retry_after: str | None = None,
) -> SocialAuthBaseException | None:
"""Classify structured protocol errors without inspecting descriptions."""
if not isinstance(data, dict):
return None
provider_code = data.get("error")
detail = data.get("error_description", "")
if isinstance(provider_code, dict):
detail = provider_code.get("message", provider_code.get("error_msg", ""))
provider_code = provider_code.get("code", provider_code.get("error_code"))
if not isinstance(provider_code, (str, int)) or not provider_code:
return None
fields: dict[str, Any] = {
"stage": stage,
"provider_code": provider_code,
"status_code": status_code,
"retry_after": retry_after,
}
if provider_code in {"access_denied", "user_denied", "cancelled", "canceled"}:
return AuthCanceled(backend, detail, **fields)
if provider_code in {
"invalid_client",
"unauthorized_client",
"unsupported_grant_type",
"invalid_scope",
}:
return AuthConfigurationError(backend, detail, code="invalid_setting", **fields)
if provider_code in {"invalid_grant", "bad_verification_code", "invalid_token"}:
code = (
"credential_rejected"
if provider_code == "invalid_token"
else "reauthentication_required"
if stage == "refresh"
else "authorization_code_rejected"
)
return AuthCredentialError(
backend, detail, code=code, source="provider_response", **fields
)
code = (
"unavailable"
if provider_code in {"server_error", "temporarily_unavailable"}
else "rate_limited"
if status_code == 429
else "unavailable"
if status_code is not None and status_code >= 500
else "http_error"
)
return AuthProviderError(backend, detail, code=code, **fields)
[docs]
def http_error(
backend: BaseAuth, error: requests.HTTPError, *, stage: ErrorStage = "user_info"
) -> SocialAuthBaseException:
"""Normalize a Requests HTTP failure, including errors without responses."""
response = error.response
status = response.status_code if response is not None else None
retry_after = response.headers.get("Retry-After") if response is not None else None
if response is not None:
try:
data = response.json()
except ValueError:
data = None
classified = provider_error(
backend, data, stage=stage, status_code=status, retry_after=retry_after
)
if classified is not None:
return classified
code = (
"rate_limited"
if status == 429
else "unavailable"
if status is not None and status >= 500
else "http_error"
)
return AuthProviderError(
backend, code=code, stage=stage, status_code=status, retry_after=retry_after
)
[docs]
def handle_http_errors(func):
"""Normalize raw HTTP errors from integrations outside BaseAuth.request."""
@functools.wraps(func)
def wrapper(*args, **kwargs):
try:
return func(*args, **kwargs)
except requests.HTTPError as error:
raise http_error(args[0], error, stage="callback") from error
return wrapper
[docs]
@contextlib.contextmanager
def wrap_access_token_error(backend: BaseAuth, *, stage: ErrorStage = "token_exchange"):
"""Normalize token request HTTP errors without guessing their cause."""
try:
yield
except requests.HTTPError as error:
raise http_error(backend, error, stage=stage) from error
[docs]
def append_slash(url: str) -> str:
"""Make sure we append a slash at the end of the URL otherwise we
have issues with urljoin Example:
>>> urlparse.urljoin('http://www.example.com/api/v3', 'user/1/')
'http://www.example.com/api/user/1/'
"""
if url and not url.endswith("/"):
url = f"{url}/"
return url
[docs]
def get_strategy(strategy: str, storage: str, *args, **kwargs) -> BaseStrategy:
Strategy = module_member(strategy)
Storage = module_member(storage)
return Strategy(Storage, *args, **kwargs)
[docs]
class cache:
"""
Cache decorator that caches the return value of a method for a
specified time.
It maintains a cache per class and method arguments, so subclasses have a
different cache entry for the same cached method.
Call ``method.invalidate()`` to clear all entries, or
``method.invalidate(instance, *args, **kwargs)`` to clear one entry.
Call ``method.refresh(instance, *args, **kwargs)`` to replace one entry
only after the underlying method succeeds. Failed refreshes propagate
their exception and preserve the existing value and expiry time.
"""
def __init__(self, ttl: int) -> None:
self.ttl = ttl
self.cache: dict[
tuple[type, tuple[Any, ...], tuple[tuple[str, Any], ...]], Any
] = {}
def __call__(self, fn):
def refresh(this, *args, **kwargs):
cached_value = fn(this, *args, **kwargs)
cache_key = (this.__class__, args, tuple(sorted(kwargs.items())))
self.cache[cache_key] = (time.time(), cached_value)
return cached_value
def wrapped(this, *args, **kwargs):
now = time.time()
last_updated = None
cached_value = None
cache_key = (this.__class__, args, tuple(sorted(kwargs.items())))
if cache_key in self.cache:
last_updated, cached_value = self.cache[cache_key]
# ignoring this type issue is safe; if cached_value is returned, last_updated
# is also set, but the type checker doesn't know it.
if not cached_value or not last_updated or now - last_updated > self.ttl:
try:
cached_value = fn(this, *args, **kwargs)
self.cache[cache_key] = (now, cached_value)
# pylint: disable-next=broad-exception-caught
except Exception:
# Use previously cached value when call fails, if available
if not cached_value:
raise
return cached_value
cast("Any", wrapped).invalidate = self._invalidate
cast("Any", wrapped).refresh = refresh
return wrapped
def _invalidate(
self, this: object | None = None, *args: Any, **kwargs: Any
) -> None:
if this is None:
self.cache.clear()
else:
cache_key = (this.__class__, args, tuple(sorted(kwargs.items())))
self.cache.pop(cache_key, None)