v0.165.0
  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