v0.163.0
 1import logging
 2
 3from plain.auth.requests import get_request_user
 4from plain.auth.views import AuthView
 5from plain.http import RedirectResponse, Response
 6from plain.templates.views import TemplateView
 7from plain.views import View
 8
 9from .exceptions import (
10    OAuthError,
11)
12from .providers import get_oauth_provider_instance
13
14logger = logging.getLogger(__name__)
15
16
17class OAuthLoginView(View):
18    def post(self) -> Response:
19        request = self.request
20        provider = self.url_kwargs["provider"]
21        if get_request_user(request):
22            return RedirectResponse("/", status_code=302)
23
24        provider_instance = get_oauth_provider_instance(provider_key=provider)
25        return provider_instance.handle_login_request(request=request)
26
27
28class OAuthCallbackView(TemplateView):
29    """
30    The callback view is used for signup, login, and connect.
31    """
32
33    template_name = "oauth/error.html"
34
35    def get(self) -> Response:
36        provider = self.url_kwargs["provider"]
37        provider_instance = get_oauth_provider_instance(provider_key=provider)
38        try:
39            return provider_instance.handle_callback_request(request=self.request)
40        except OAuthError as e:
41            logger.warning("OAuth error: %s", e.message)
42            self.oauth_error = e
43
44            return self.render(status_code=400)
45
46    def get_template_names(self) -> list[str]:
47        names = []
48        if (
49            oauth_error := getattr(self, "oauth_error", None)
50        ) and oauth_error.template_name:
51            names.append(oauth_error.template_name)
52        names.append(self.template_name)
53        return names
54
55    def get_template_context(self) -> dict:
56        context = super().get_template_context()
57        context["oauth_error"] = getattr(self, "oauth_error", None)
58        return context
59
60
61class OAuthConnectView(AuthView):
62    login_required = True
63
64    def post(self) -> Response:
65        request = self.request
66        provider = self.url_kwargs["provider"]
67        provider_instance = get_oauth_provider_instance(provider_key=provider)
68        return provider_instance.handle_connect_request(request=request)
69
70
71class OAuthDisconnectView(AuthView):
72    login_required = True
73
74    def post(self) -> Response:
75        request = self.request
76        provider = self.url_kwargs["provider"]
77        provider_instance = get_oauth_provider_instance(provider_key=provider)
78        # try:
79        return provider_instance.handle_disconnect_request(request=request)
80        # except OAuthCannotDisconnectError:
81        #     return render(
82        #         request,
83        #         "oauth/error.html",
84        #         {
85        #             "oauth_error": "This connection can't be removed. You must have a usable password or at least one active connection."
86        #         },
87        #         status_code=400,
88        #     )