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