Skip to content
Draft
Show file tree
Hide file tree
Changes from 1 commit
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
Next Next commit
Add types using typemonkey
  • Loading branch information
melizeche committed Jul 18, 2025
commit 6383d0bfc03a6a9b8082682e21c51034d5147132
8 changes: 4 additions & 4 deletions authentik/core/api/groups.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,15 +77,15 @@ def get_users_obj(self, instance: Group) -> list[GroupMemberSerializer] | None:
return None
return GroupMemberSerializer(instance.users, many=True).data

def validate_parent(self, parent: Group | None):
def validate_parent(self, parent: Group | None) -> None:
"""Validate group parent (if set), ensuring the parent isn't itself"""
if not self.instance or not parent:
return parent
if str(parent.group_uuid) == str(self.instance.group_uuid):
raise ValidationError(_("Cannot set group as parent of itself."))
return parent

def validate_is_superuser(self, superuser: bool):
def validate_is_superuser(self, superuser: bool) -> bool:
"""Ensure that the user creating this group has permissions to set the superuser flag"""
request: Request = self.context.get("request", None)
if not request:
Expand Down Expand Up @@ -210,15 +210,15 @@ def get_queryset(self):
OpenApiParameter("include_users", bool, default=True),
]
)
def list(self, request, *args, **kwargs):
def list(self, request: Request, *args, **kwargs) -> Response:
return super().list(request, *args, **kwargs)

@extend_schema(
parameters=[
OpenApiParameter("include_users", bool, default=True),
]
)
def retrieve(self, request, *args, **kwargs):
def retrieve(self, request: Request, *args, **kwargs) -> Response:
return super().retrieve(request, *args, **kwargs)

@permission_required("authentik_core.add_user_to_group")
Expand Down
3 changes: 2 additions & 1 deletion authentik/core/api/providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from django.utils.translation import gettext_lazy as _
from django_filters.filters import BooleanFilter
from django_filters.filterset import FilterSet
from model_utils.managers import InheritanceQuerySet
from rest_framework import mixins
from rest_framework.fields import ReadOnlyField, SerializerMethodField
from rest_framework.viewsets import GenericViewSet
Expand Down Expand Up @@ -99,5 +100,5 @@ class ProviderViewSet(
"application__name",
]

def get_queryset(self): # pragma: no cover
def get_queryset(self) -> InheritanceQuerySet: # pragma: no cover
return Provider.objects.select_subclasses()
3 changes: 2 additions & 1 deletion authentik/core/api/sources.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from collections.abc import Iterable

from drf_spectacular.utils import OpenApiResponse, extend_schema
from model_utils.managers import InheritanceQuerySet
from rest_framework import mixins
from rest_framework.decorators import action
from rest_framework.exceptions import ValidationError
Expand Down Expand Up @@ -88,7 +89,7 @@ class SourceViewSet(
search_fields = ["slug", "name"]
filterset_fields = ["slug", "name", "managed", "pbm_uuid"]

def get_queryset(self): # pragma: no cover
def get_queryset(self) -> InheritanceQuerySet: # pragma: no cover
return Source.objects.select_subclasses()

@permission_required("authentik_core.change_source")
Expand Down
7 changes: 4 additions & 3 deletions authentik/core/api/tokens.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

from typing import Any

from django.db.models.query import QuerySet
from django.utils.timezone import now
from drf_spectacular.utils import OpenApiResponse, extend_schema, inline_serializer
from guardian.shortcuts import assign_perm, get_anonymous_user
Expand Down Expand Up @@ -41,7 +42,7 @@ def __init__(self, *args, **kwargs) -> None:
if SERIALIZER_CONTEXT_BLUEPRINT in self.context:
self.fields["key"] = CharField(required=False)

def validate_user(self, user: User):
def validate_user(self, user: User) -> User:
"""Ensure user of token cannot be changed"""
if self.instance and self.instance.user_id:
if user.pk != self.instance.user_id:
Expand Down Expand Up @@ -138,13 +139,13 @@ class TokenViewSet(UsedByMixin, ModelViewSet):
owner_field = "user"
rbac_allow_create_without_perm = True

def get_queryset(self):
def get_queryset(self) -> QuerySet:
user = self.request.user if self.request else get_anonymous_user()
if user.is_superuser:
return super().get_queryset()
return super().get_queryset().filter(user=user.pk)

def perform_create(self, serializer: TokenSerializer):
def perform_create(self, serializer: TokenSerializer) -> Token:
if not self.request.user.is_superuser:
instance = serializer.save(
user=self.request.user,
Expand Down
13 changes: 8 additions & 5 deletions authentik/core/api/users.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
UUIDFilter,
)
from django_filters.filterset import FilterSet
from djangoql.schema import BoolField, StrField
from drf_spectacular.types import OpenApiTypes
from drf_spectacular.utils import (
OpenApiParameter,
Expand Down Expand Up @@ -72,8 +73,10 @@
Token,
TokenIntents,
User,
UserQuerySet,
UserTypes,
)
from authentik.enterprise.search.fields import ChoiceSearchField, JSONSearchField
from authentik.events.models import Event, EventAction
from authentik.flows.exceptions import FlowNonApplicableException
from authentik.flows.models import FlowToken
Expand Down Expand Up @@ -349,7 +352,7 @@ class UsersFilter(FilterSet):
queryset=Group.objects.all().order_by("name"),
)

def filter_is_superuser(self, queryset, name, value):
def filter_is_superuser(self, queryset: UserQuerySet, name: str, value: bool) -> UserQuerySet:
if value:
return queryset.filter(ak_groups__is_superuser=True).distinct()
return queryset.exclude(ak_groups__is_superuser=True).distinct()
Expand Down Expand Up @@ -395,7 +398,7 @@ class UserViewSet(UsedByMixin, ModelViewSet):
filterset_class = UsersFilter
search_fields = ["username", "name", "is_active", "email", "uuid", "attributes"]

def get_ql_fields(self):
def get_ql_fields(self) -> list[StrField | BoolField | ChoiceSearchField | JSONSearchField]:
from djangoql.schema import BoolField, StrField

from authentik.enterprise.search.fields import ChoiceSearchField, JSONSearchField
Expand All @@ -410,7 +413,7 @@ def get_ql_fields(self):
JSONSearchField(User, "attributes", suggest_nested=False),
]

def get_queryset(self):
def get_queryset(self) -> UserQuerySet:
base_qs = User.objects.all().exclude_anonymous()
if self.serializer_class(context={"request": self.request})._should_include_groups:
base_qs = base_qs.prefetch_related("ak_groups")
Expand All @@ -421,10 +424,10 @@ def get_queryset(self):
OpenApiParameter("include_groups", bool, default=True),
]
)
def list(self, request, *args, **kwargs):
def list(self, request: Request, *args, **kwargs) -> Response:
return super().list(request, *args, **kwargs)

def _create_recovery_link(self, for_email=False) -> tuple[str, Token]:
def _create_recovery_link(self, for_email: bool = False) -> tuple[str, Token]:
"""Create a recovery link (when the current brand has a recovery flow set),
that can either be shown to an admin or sent to the user directly"""
brand: Brand = self.request._request.brand
Expand Down
6 changes: 3 additions & 3 deletions authentik/core/api/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@ class JSONExtension(OpenApiSerializerFieldExtension):

target_class = "authentik.core.api.utils.JSONDictField"

def map_serializer_field(self, auto_schema, direction):
def map_serializer_field(self, auto_schema, direction: str) -> dict[str, str]:
return build_basic_type(OpenApiTypes.OBJECT)


Expand All @@ -52,7 +52,7 @@ class ModelSerializer(BaseModelSerializer):
serializer_field_mapping = BaseModelSerializer.serializer_field_mapping.copy()
serializer_field_mapping[models.JSONField] = JSONDictField

def create(self, validated_data):
def create(self, validated_data: dict[str, Any]):
instance = super().create(validated_data)

request = self.context.get("request")
Expand All @@ -61,7 +61,7 @@ def create(self, validated_data):

return instance

def update(self, instance: Model, validated_data):
def update(self, instance: Model, validated_data: dict[str, Any]):
raise_errors_on_nested_writes("update", self, validated_data)
info = model_meta.get_field_info(instance)

Expand Down
7 changes: 5 additions & 2 deletions authentik/core/middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,13 +7,16 @@

from django.contrib.auth.models import AnonymousUser
from django.core.exceptions import ImproperlyConfigured
from django.core.handlers.wsgi import WSGIRequest
from django.http import HttpRequest, HttpResponse
from django.utils.deprecation import MiddlewareMixin
from django.utils.functional import SimpleLazyObject
from django.utils.translation import override
from sentry_sdk.api import set_tag
from structlog.contextvars import STRUCTLOG_KEY_PREFIX

from authentik.core.models import User

SESSION_KEY_IMPERSONATE_USER = "authentik/impersonate/user"
SESSION_KEY_IMPERSONATE_ORIGINAL_USER = "authentik/impersonate/original_user"
RESPONSE_HEADER_ID = "X-authentik-id"
Expand All @@ -25,7 +28,7 @@
CTX_AUTH_VIA = ContextVar[str | None](STRUCTLOG_KEY_PREFIX + KEY_AUTH_VIA, default=None)


def get_user(request):
def get_user(request: WSGIRequest) -> AnonymousUser | User:
if not hasattr(request, "_cached_user"):
user = None
if (authenticated_session := request.session.get("authenticatedsession", None)) is not None:
Expand All @@ -46,7 +49,7 @@ async def aget_user(request):


class AuthenticationMiddleware(MiddlewareMixin):
def process_request(self, request):
def process_request(self, request: WSGIRequest):
if not hasattr(request, "session"):
raise ImproperlyConfigured(
"The Django authentication middleware requires session "
Expand Down
38 changes: 23 additions & 15 deletions authentik/core/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
from django.contrib.auth.models import AbstractUser
from django.contrib.auth.models import UserManager as DjangoUserManager
from django.contrib.sessions.base_session import AbstractBaseSession
from django.core.handlers.wsgi import WSGIRequest
from django.db import models
from django.db.models import Q, QuerySet, options
from django.db.models.constants import LOOKUP_SEP
Expand All @@ -22,6 +23,7 @@
from guardian.conf import settings
from guardian.mixins import GuardianUserMixin
from model_utils.managers import InheritanceManager
from rest_framework.request import Request
from rest_framework.serializers import Serializer
from structlog.stdlib import get_logger

Expand Down Expand Up @@ -137,7 +139,7 @@ def update_or_create_attributes(


class GroupQuerySet(QuerySet):
def with_children_recursive(self):
def with_children_recursive(self) -> "GroupQuerySet":
"""Recursively get all groups that have the current queryset as parents
or are indirectly related."""

Expand Down Expand Up @@ -210,7 +212,7 @@ class Meta:
("disable_group_superuser", _("Disable superuser status")),
]

def __str__(self):
def __str__(self) -> str:
return f"Group {self.name}"

@property
Expand Down Expand Up @@ -241,15 +243,15 @@ def children_recursive(self: Self | QuerySet["Group"]) -> QuerySet["Group"]:
class UserQuerySet(models.QuerySet):
"""User queryset"""

def exclude_anonymous(self):
def exclude_anonymous(self) -> "UserQuerySet":
"""Exclude anonymous user"""
return self.exclude(**{User.USERNAME_FIELD: settings.ANONYMOUS_USER_NAME})


class UserManager(DjangoUserManager):
"""User manager that doesn't assign is_superuser and is_staff"""

def get_queryset(self):
def get_queryset(self) -> UserQuerySet:
"""Create special user queryset"""
return UserQuerySet(self.model, using=self._db)

Expand Down Expand Up @@ -295,7 +297,7 @@ class Meta:
models.Index(fields=["type"]),
]

def __str__(self):
def __str__(self) -> str:
return self.username

@staticmethod
Expand Down Expand Up @@ -360,7 +362,13 @@ def is_staff(self) -> bool:
"""superuser == staff user"""
return self.is_superuser # type: ignore

def set_password(self, raw_password, signal=True, sender=None, request=None):
def set_password(
self,
raw_password: str,
signal: bool = True,
sender: None = None,
request: WSGIRequest | Request | None = None,
) -> None:
if self.pk and signal:
from authentik.core.signals import password_changed

Expand Down Expand Up @@ -479,7 +487,7 @@ def serializer(self) -> type[Serializer]:
"""Get serializer for this model"""
raise NotImplementedError

def __str__(self):
def __str__(self) -> str:
return str(self.name)


Expand Down Expand Up @@ -611,7 +619,7 @@ def backchannel_provider_for[T: Provider](self, provider_type: type[T], **kwargs
)
return getattr(providers.first(), provider_type._meta.model_name)

def __str__(self):
def __str__(self) -> str:
return str(self.name)

class Meta:
Expand All @@ -631,7 +639,7 @@ class Meta:
verbose_name_plural = _("Application Entitlements")
unique_together = (("app", "name"),)

def __str__(self):
def __str__(self) -> str:
return f"Application Entitlement {self.name} for app {self.app_id}"

@property
Expand All @@ -640,7 +648,7 @@ def serializer(self) -> type[Serializer]:

return ApplicationEntitlementSerializer

def supported_policy_binding_targets(self):
def supported_policy_binding_targets(self) -> list[str]:
return ["group", "user"]


Expand Down Expand Up @@ -812,7 +820,7 @@ def get_base_group_properties(self, **kwargs) -> dict[str, Any | dict[str, Any]]
return {}
raise NotImplementedError

def __str__(self):
def __str__(self) -> str:
return str(self.name)

class Meta:
Expand Down Expand Up @@ -895,7 +903,7 @@ class Meta:
models.Index(fields=["expiring", "expires"]),
]

def expire_action(self, *args, **kwargs):
def expire_action(self, *args, **kwargs) -> tuple[int, dict[str, int]]:
"""Handler which is called when this object is expired. By
default the object is deleted. This is less efficient compared
to bulk deleting objects, but classes like Token() need to change
Expand Down Expand Up @@ -958,7 +966,7 @@ class Meta:
("set_token_key", _("Set a token's key")),
]

def __str__(self):
def __str__(self) -> str:
description = f"{self.identifier}"
if self.expiring:
description += f" (expires={self.expires})"
Expand Down Expand Up @@ -1023,7 +1031,7 @@ def evaluate(self, user: User | None, request: HttpRequest | None, **kwargs) ->
except Exception as exc:
raise PropertyMappingExpressionException(exc, self) from exc

def __str__(self):
def __str__(self) -> str:
return f"Property Mapping {self.name}"

class Meta:
Expand Down Expand Up @@ -1051,7 +1059,7 @@ class Meta:
]
default_permissions = []

def __str__(self):
def __str__(self) -> str:
return self.session_key

class Keys(StrEnum):
Expand Down
Loading
Loading