authentication.py 5.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149
  1. from __future__ import annotations
  2. import functools
  3. import inspect
  4. from collections.abc import Callable, Sequence
  5. from typing import Any, ParamSpec
  6. from urllib.parse import urlencode
  7. from starlette._utils import is_async_callable
  8. from starlette.exceptions import HTTPException
  9. from starlette.requests import HTTPConnection, Request
  10. from starlette.responses import RedirectResponse
  11. from starlette.websockets import WebSocket
  12. _P = ParamSpec("_P")
  13. def has_required_scope(conn: HTTPConnection, scopes: Sequence[str]) -> bool:
  14. for scope in scopes:
  15. if scope not in conn.auth.scopes:
  16. return False
  17. return True
  18. def requires(
  19. scopes: str | Sequence[str],
  20. status_code: int = 403,
  21. redirect: str | None = None,
  22. ) -> Callable[[Callable[_P, Any]], Callable[_P, Any]]:
  23. scopes_list = [scopes] if isinstance(scopes, str) else list(scopes)
  24. def decorator(
  25. func: Callable[_P, Any],
  26. ) -> Callable[_P, Any]:
  27. sig = inspect.signature(func)
  28. for idx, parameter in enumerate(sig.parameters.values()):
  29. if parameter.name == "request" or parameter.name == "websocket":
  30. type_ = parameter.name
  31. break
  32. else:
  33. raise Exception(f'No "request" or "websocket" argument on function "{func}"')
  34. if type_ == "websocket":
  35. # Handle websocket functions. (Always async)
  36. @functools.wraps(func)
  37. async def websocket_wrapper(*args: _P.args, **kwargs: _P.kwargs) -> None:
  38. websocket = kwargs.get("websocket", args[idx] if idx < len(args) else None)
  39. assert isinstance(websocket, WebSocket), (
  40. "Parameter with name 'websocket' is required to be of type 'WebSocket'"
  41. f" not '{type(websocket).__name__}'"
  42. )
  43. if not has_required_scope(websocket, scopes_list):
  44. await websocket.close()
  45. else:
  46. await func(*args, **kwargs)
  47. return websocket_wrapper
  48. elif is_async_callable(func):
  49. # Handle async request/response functions.
  50. @functools.wraps(func)
  51. async def async_wrapper(*args: _P.args, **kwargs: _P.kwargs) -> Any:
  52. request = kwargs.get("request", args[idx] if idx < len(args) else None)
  53. assert isinstance(request, Request), (
  54. f"Parameter with name 'request' is required to be of type 'Request' not '{type(request).__name__}'"
  55. )
  56. if not has_required_scope(request, scopes_list):
  57. if redirect is not None:
  58. orig_request_qparam = urlencode({"next": str(request.url)})
  59. next_url = f"{request.url_for(redirect)}?{orig_request_qparam}"
  60. return RedirectResponse(url=next_url, status_code=303)
  61. raise HTTPException(status_code=status_code)
  62. return await func(*args, **kwargs)
  63. return async_wrapper
  64. else:
  65. # Handle sync request/response functions.
  66. @functools.wraps(func)
  67. def sync_wrapper(*args: _P.args, **kwargs: _P.kwargs) -> Any:
  68. request = kwargs.get("request", args[idx] if idx < len(args) else None)
  69. assert isinstance(request, Request), (
  70. f"Parameter with name 'request' is required to be of type 'Request' not '{type(request).__name__}'"
  71. )
  72. if not has_required_scope(request, scopes_list):
  73. if redirect is not None:
  74. orig_request_qparam = urlencode({"next": str(request.url)})
  75. next_url = f"{request.url_for(redirect)}?{orig_request_qparam}"
  76. return RedirectResponse(url=next_url, status_code=303)
  77. raise HTTPException(status_code=status_code)
  78. return func(*args, **kwargs)
  79. return sync_wrapper
  80. return decorator
  81. class AuthenticationError(Exception):
  82. pass
  83. class AuthenticationBackend:
  84. async def authenticate(self, conn: HTTPConnection) -> tuple[AuthCredentials, BaseUser] | None:
  85. raise NotImplementedError() # pragma: no cover
  86. class AuthCredentials:
  87. def __init__(self, scopes: Sequence[str] | None = None):
  88. self.scopes = [] if scopes is None else list(scopes)
  89. class BaseUser:
  90. @property
  91. def is_authenticated(self) -> bool:
  92. raise NotImplementedError() # pragma: no cover
  93. @property
  94. def display_name(self) -> str:
  95. raise NotImplementedError() # pragma: no cover
  96. @property
  97. def identity(self) -> str:
  98. raise NotImplementedError() # pragma: no cover
  99. class SimpleUser(BaseUser):
  100. def __init__(self, username: str) -> None:
  101. self.username = username
  102. @property
  103. def is_authenticated(self) -> bool:
  104. return True
  105. @property
  106. def display_name(self) -> str:
  107. return self.username
  108. class UnauthenticatedUser(BaseUser):
  109. @property
  110. def is_authenticated(self) -> bool:
  111. return False
  112. @property
  113. def display_name(self) -> str:
  114. return ""