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 # )