1from __future__ import annotations
2
3from typing import TYPE_CHECKING, Any
4
5import psycopg
6from app.users.models import User
7from plain.exceptions import ValidationError
8from plain.postgres import transaction, types
9from plain.utils import timezone
10
11from plain import postgres
12
13from .exceptions import OAuthUserAlreadyExistsError
14
15if TYPE_CHECKING:
16 from .providers import OAuthToken, OAuthUser
17
18__all__ = ["OAuthConnection"]
19
20
21@postgres.register_model
22class OAuthConnection(postgres.Model):
23 created_at = types.DateTimeField(create_now=True)
24 updated_at = types.DateTimeField(create_now=True, update_now=True)
25
26 user = types.ForeignKeyField(
27 "users.User",
28 on_delete=postgres.CASCADE,
29 )
30
31 # The key used to refer to this provider type (in settings)
32 provider_key = types.TextField(max_length=100)
33
34 # The unique ID of the user on the provider's system
35 provider_user_id = types.TextField(max_length=100)
36
37 # Token data
38 access_token = types.EncryptedTextField(max_length=2000)
39 refresh_token = types.EncryptedTextField(
40 max_length=2000, required=False, default=""
41 )
42 access_token_expires_at = types.DateTimeField(required=False, allow_null=True)
43 refresh_token_expires_at = types.DateTimeField(required=False, allow_null=True)
44
45 query: postgres.QuerySet[OAuthConnection] = postgres.QuerySet()
46
47 model_options = postgres.Options(
48 indexes=[
49 postgres.Index(
50 name="plainoauth_oauthconnection_user_id_idx", fields=["user"]
51 ),
52 ],
53 constraints=[
54 postgres.UniqueConstraint(
55 fields=["provider_key", "provider_user_id"],
56 name="plainoauth_oauthconnection_unique_provider_key_user_id",
57 )
58 ],
59 ordering=("provider_key",),
60 )
61
62 def __str__(self) -> str:
63 return f"{self.provider_key}[{self.user}:{self.provider_user_id}]"
64
65 def refresh_access_token(self) -> None:
66 from .providers import OAuthToken, get_oauth_provider_instance
67
68 provider_instance = get_oauth_provider_instance(provider_key=self.provider_key)
69 oauth_token = OAuthToken(
70 access_token=self.access_token,
71 refresh_token=self.refresh_token,
72 access_token_expires_at=self.access_token_expires_at,
73 refresh_token_expires_at=self.refresh_token_expires_at,
74 )
75 refreshed_oauth_token = provider_instance.refresh_oauth_token(
76 oauth_token=oauth_token
77 )
78 self.set_token_fields(refreshed_oauth_token)
79 self.update()
80
81 def set_token_fields(self, oauth_token: OAuthToken) -> None:
82 self.access_token = oauth_token.access_token
83 self.refresh_token = oauth_token.refresh_token
84 self.access_token_expires_at = oauth_token.access_token_expires_at
85 self.refresh_token_expires_at = oauth_token.refresh_token_expires_at
86
87 def set_user_fields(self, oauth_user: OAuthUser) -> None:
88 self.provider_user_id = oauth_user.provider_id
89
90 def access_token_expired(self) -> bool:
91 return (
92 self.access_token_expires_at is not None
93 and self.access_token_expires_at < timezone.now()
94 )
95
96 def refresh_token_expired(self) -> bool:
97 return (
98 self.refresh_token_expires_at is not None
99 and self.refresh_token_expires_at < timezone.now()
100 )
101
102 @classmethod
103 def get_or_create_user(
104 cls, *, provider_key: str, oauth_token: OAuthToken, oauth_user: OAuthUser
105 ) -> OAuthConnection:
106 try:
107 connection = cls.query.get(
108 provider_key=provider_key,
109 provider_user_id=oauth_user.provider_id,
110 )
111 connection.set_token_fields(oauth_token)
112 connection.update()
113 return connection
114 except cls.DoesNotExist:
115 # If email needs to be unique, then we expect
116 # that to be taken care of on the user model itself
117 with transaction.atomic():
118 try:
119 with transaction.atomic():
120 user = User(
121 **oauth_user.user_model_fields,
122 )
123 user.create()
124 except (psycopg.IntegrityError, ValidationError):
125 raise OAuthUserAlreadyExistsError(
126 provider_key=provider_key,
127 user_model_fields=oauth_user.user_model_fields,
128 )
129
130 return cls.connect(
131 user=user,
132 provider_key=provider_key,
133 oauth_token=oauth_token,
134 oauth_user=oauth_user,
135 )
136
137 @classmethod
138 def connect(
139 cls,
140 *,
141 user: Any,
142 provider_key: str,
143 oauth_token: OAuthToken,
144 oauth_user: OAuthUser,
145 ) -> OAuthConnection:
146 """
147 Connect will either create a new connection or update an existing connection
148 """
149 try:
150 connection = cls.query.get(
151 user=user,
152 provider_key=provider_key,
153 provider_user_id=oauth_user.provider_id,
154 )
155 except cls.DoesNotExist:
156 # Create our own instance (not using get_or_create)
157 # so that any created signals contain the token fields too
158 connection = cls(
159 user=user,
160 provider_key=provider_key,
161 provider_user_id=oauth_user.provider_id,
162 )
163
164 connection.set_user_fields(oauth_user)
165 connection.set_token_fields(oauth_token)
166 if connection._state.adding:
167 connection.create()
168 else:
169 connection.update()
170
171 return connection