Skip to content

Commit 786c1db

Browse files
committed
Type annotations on plain-auth
1 parent 97c8bbe commit 786c1db

7 files changed

Lines changed: 49 additions & 30 deletions

File tree

plain-auth/plain/auth/requests.py

Lines changed: 4 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,24 +1,20 @@
11
from __future__ import annotations
22

3-
from typing import TYPE_CHECKING
3+
from typing import TYPE_CHECKING, Any
44
from weakref import WeakKeyDictionary
55

66
if TYPE_CHECKING:
77
from plain.http import Request
88

9-
from .sessions import get_user_model
9+
_request_users: WeakKeyDictionary[Request, Any | None] = WeakKeyDictionary()
1010

11-
User = get_user_model()
1211

13-
_request_users: WeakKeyDictionary[Request, User | None] = WeakKeyDictionary()
14-
15-
16-
def set_request_user(request: Request, user: User | None) -> None:
12+
def set_request_user(request: Request, user: Any | None) -> None:
1713
"""Store the authenticated user for this request."""
1814
_request_users[request] = user
1915

2016

21-
def get_request_user(request: Request) -> User | None:
17+
def get_request_user(request: Request) -> Any | None:
2218
"""
2319
Get the authenticated user for this request, if any.
2420

plain-auth/plain/auth/sessions.py

Lines changed: 15 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,8 @@
1+
from __future__ import annotations
2+
13
import hmac
4+
from collections.abc import Generator
5+
from typing import TYPE_CHECKING, Any
26

37
from plain.exceptions import ImproperlyConfigured
48
from plain.models import models_registry
@@ -9,18 +13,21 @@
913

1014
from .requests import get_request_user, set_request_user
1115

16+
if TYPE_CHECKING:
17+
from plain.http import Request
18+
1219
USER_ID_SESSION_KEY = "_auth_user_id"
1320
USER_HASH_SESSION_KEY = "_auth_user_hash"
1421

1522

16-
def get_session_auth_hash(user):
23+
def get_session_auth_hash(user: Any) -> str:
1724
"""
1825
Return an HMAC of the password field.
1926
"""
2027
return _get_session_auth_hash(user)
2128

2229

23-
def update_session_auth_hash(request, user):
30+
def update_session_auth_hash(request: Request, user: Any) -> None:
2431
"""
2532
Updating a user's password (for example) logs out all sessions for the user.
2633
@@ -36,12 +43,12 @@ def update_session_auth_hash(request, user):
3643
session[USER_HASH_SESSION_KEY] = get_session_auth_hash(user)
3744

3845

39-
def get_session_auth_fallback_hash(user):
46+
def get_session_auth_fallback_hash(user: Any) -> Generator[str, None, None]:
4047
for fallback_secret in settings.SECRET_KEY_FALLBACKS:
4148
yield _get_session_auth_hash(user, secret=fallback_secret)
4249

4350

44-
def _get_session_auth_hash(user, secret=None):
51+
def _get_session_auth_hash(user: Any, secret: str | None = None) -> str:
4552
key_salt = "plain.auth.get_session_auth_hash"
4653
return salted_hmac(
4754
key_salt,
@@ -51,7 +58,7 @@ def _get_session_auth_hash(user, secret=None):
5158
).hexdigest()
5259

5360

54-
def login(request, user):
61+
def login(request: Request, user: Any) -> None:
5562
"""
5663
Persist a user id and a backend in the request. This way a user doesn't
5764
have to reauthenticate on every request. Note that data set during
@@ -87,7 +94,7 @@ def login(request, user):
8794
set_request_user(request, user)
8895

8996

90-
def logout(request):
97+
def logout(request: Request) -> None:
9198
"""
9299
Remove the authenticated user's ID from the request and flush their session
93100
data.
@@ -99,7 +106,7 @@ def logout(request):
99106
set_request_user(request, None)
100107

101108

102-
def get_user_model():
109+
def get_user_model() -> type[Any]:
103110
"""
104111
Return the User model that is active in this project.
105112
"""
@@ -115,7 +122,7 @@ def get_user_model():
115122
)
116123

117124

118-
def get_user(request):
125+
def get_user(request: Request) -> Any | None:
119126
"""
120127
Return the user model instance associated with the given request session.
121128
If no user is retrieved, return None.

plain-auth/plain/auth/templates.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,20 @@
1+
from __future__ import annotations
2+
3+
from typing import TYPE_CHECKING, Any
4+
15
from jinja2 import pass_context
26

37
from plain.templates import register_template_global
48

59
from .requests import get_request_user
610

11+
if TYPE_CHECKING:
12+
from jinja2.runtime import Context
13+
714

815
@register_template_global
916
@pass_context
10-
def get_current_user(context):
17+
def get_current_user(context: Context) -> Any | None:
1118
"""Get the authenticated user for the current request."""
1219
request = context.get("request")
1320
assert request is not None, "No request in template context"

plain-auth/plain/auth/test.py

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,7 @@
1+
from __future__ import annotations
2+
13
from http.cookies import SimpleCookie
4+
from typing import TYPE_CHECKING, Any
25

36
from plain.http.request import Request
47
from plain.runtime import settings
@@ -8,8 +11,11 @@
811
from .requests import set_request_user
912
from .sessions import get_user, login, logout
1013

14+
if TYPE_CHECKING:
15+
from plain.test.client import Client
16+
1117

12-
def login_client(client, user):
18+
def login_client(client: Client, user: Any) -> None:
1319
"""Log a user into a test client."""
1420
request = Request()
1521
if client.session:
@@ -32,7 +38,7 @@ def login_client(client, user):
3238
client.cookies[session_cookie].update(cookie_data)
3339

3440

35-
def logout_client(client):
41+
def logout_client(client: Client) -> None:
3642
"""Log out a user from a test client."""
3743
request = Request()
3844
if client.session:

plain-auth/plain/auth/utils.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,12 @@
1+
from __future__ import annotations
2+
3+
from typing import Any
4+
15
from plain.urls import NoReverseMatch, reverse
26
from plain.utils.functional import Promise
37

48

5-
def resolve_url(to, *args, **kwargs):
9+
def resolve_url(to: Any, *args: Any, **kwargs: Any) -> str:
610
"""
711
Return a URL appropriate for the arguments passed.
812

plain-auth/plain/auth/views.py

Lines changed: 8 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
from __future__ import annotations
22

33
from functools import cached_property
4-
from typing import TYPE_CHECKING
4+
from typing import TYPE_CHECKING, Any
55
from urllib.parse import urlparse, urlunparse
66

77
from plain.exceptions import PermissionDenied
@@ -23,13 +23,9 @@
2323
if TYPE_CHECKING:
2424
from plain.http import Request
2525

26-
from .sessions import get_user_model
27-
28-
User = get_user_model()
29-
3026

3127
class LoginRequired(Exception):
32-
def __init__(self, login_url=None, redirect_field_name="next"):
28+
def __init__(self, login_url: str | None = None, redirect_field_name: str = "next"):
3329
self.login_url = login_url or settings.AUTH_LOGIN_URL
3430
self.redirect_field_name = redirect_field_name
3531

@@ -42,7 +38,7 @@ class AuthViewMixin(SessionViewMixin):
4238
request: Request
4339

4440
@cached_property
45-
def user(self) -> User | None:
41+
def user(self) -> Any | None:
4642
"""Get the authenticated user for this request."""
4743
from .requests import get_request_user
4844

@@ -116,12 +112,14 @@ def get_response(self) -> Response:
116112

117113

118114
class LogoutView(View):
119-
def post(self):
115+
def post(self) -> ResponseRedirect:
120116
logout(self.request)
121117
return ResponseRedirect("/")
122118

123119

124-
def redirect_to_login(next, login_url=None, redirect_field_name="next"):
120+
def redirect_to_login(
121+
next: str, login_url: str | None = None, redirect_field_name: str = "next"
122+
) -> ResponseRedirect:
125123
"""
126124
Redirect the user to the login page, passing the given 'next' page.
127125
"""
@@ -133,4 +131,4 @@ def redirect_to_login(next, login_url=None, redirect_field_name="next"):
133131
querystring[redirect_field_name] = next
134132
login_url_parts[4] = querystring.urlencode(safe="/")
135133

136-
return ResponseRedirect(urlunparse(login_url_parts))
134+
return ResponseRedirect(str(urlunparse(login_url_parts)))

scripts/type-validate

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ from pathlib import Path
1414
# Directories that must maintain 100% type annotation coverage
1515
FULLY_TYPED_DIRS = [
1616
"plain/plain",
17+
"plain-auth/plain/auth",
1718
]
1819

1920

0 commit comments

Comments
 (0)