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