diff --git a/core/api.py b/core/api.py index 08aefa6ff..5fe1c16ab 100644 --- a/core/api.py +++ b/core/api.py @@ -2,9 +2,10 @@ from annotated_types import Ge, Le, MinLen from django.conf import settings +from django.contrib.auth import login from django.db.models import F from django.http import HttpResponse -from ninja import File, Query +from ninja import File, Query, Status from ninja.security import SessionAuth from ninja_extra import ControllerBase, api_controller, paginate, route from ninja_extra.exceptions import PermissionDenied @@ -18,6 +19,7 @@ from core.schemas import ( FamilyGodfatherSchema, GroupSchema, + LoginSchema, MarkdownSchema, SithFileSchema, UploadedFileSchema, @@ -28,6 +30,7 @@ UserSchema, ) from core.templatetags.renderer import markdown +from core.views.forms import LoginForm from counter.utils import is_logged_in_counter @@ -166,3 +169,30 @@ def get_family_graph( ] ), } + + +@api_controller("/auth") +class AuthController(ControllerBase): + @route.post( + "/login", + auth=None, + response={200: UserSchema, 401: dict[str, list[str]]}, + ) + def login(self, body: LoginSchema): + """Authenticate a user by username, email or AE account id. + + Returns the user's data on success, 401 on failure. + """ + if self.context.request.user.is_authenticated: + raise PermissionDenied + + login_form = LoginForm( + self.context.request, + data={"username": body.identifier, "password": body.password}, + ) + if not login_form.is_valid(): + return Status(401, login_form.errors) + + user = login_form.get_user() + login(self.context.request, user) + return user diff --git a/core/schemas.py b/core/schemas.py index 325664e93..c9c6d7580 100644 --- a/core/schemas.py +++ b/core/schemas.py @@ -155,6 +155,11 @@ class MarkdownSchema(Schema): text: str +class LoginSchema(Schema): + identifier: str + password: str + + class FamilyGodfatherSchema(Schema): godfather: int godchild: int diff --git a/core/tests/test_auth_api.py b/core/tests/test_auth_api.py new file mode 100644 index 000000000..66b49a6ed --- /dev/null +++ b/core/tests/test_auth_api.py @@ -0,0 +1,58 @@ +import pytest +from django.contrib.auth.hashers import make_password +from django.test import Client +from django.urls import reverse +from model_bakery import baker + +from core.models import User +from core.schemas import UserSchema +from counter.models import Customer + + +def post_login(client: Client, identifier: str, password: str): + return client.post( + reverse("api:login"), + data={"identifier": identifier, "password": password}, + content_type="application/json", + ) + + +@pytest.fixture() +def user(db) -> User: + return baker.make(User, password=make_password("plop")) + + +@pytest.mark.django_db +@pytest.mark.parametrize( + "identifier_getter", + [ + lambda user: user.username, + lambda user: user.email, + lambda user: Customer.get_or_create(user)[0].account_id, + ], +) +def test_api_login_success(client: Client, user: User, identifier_getter): + response = post_login(client, identifier_getter(user), "plop") + + assert response.status_code == 200 + assert response.json() == UserSchema.model_validate(user).model_dump(mode="json") + assert int(client.session["_auth_user_id"]) == user.id + + +@pytest.mark.django_db +def test_api_login_fail_invalid_credentials(client: Client, user: User): + response = post_login(client, user.username, "wrong-password") + + assert response.status_code == 401 + assert "_auth_user_id" not in client.session + + +@pytest.mark.django_db +def test_api_login_fail_if_already_authenticated(client: Client, user: User): + already_logged_user = baker.make(User) + client.force_login(already_logged_user) + + response = post_login(client, user.username, "plop") + + assert response.status_code == 403 + assert int(client.session["_auth_user_id"]) == already_logged_user.id