import os import re import uuid import core.postgres as postgres import core.users as users DISCORD_PROVIDER = "discord" DEFAULT_ENROLLMENT_MODE = "allowlist" _PROVIDER_PATTERN = re.compile(r"^[a-z0-9][a-z0-9_.-]{0,49}$") def _normalizeProvider(provider): provider = provider.strip().lower() if isinstance(provider, str) else "" if not _PROVIDER_PATTERN.fullmatch(provider): raise ValueError("provider must be a valid provider name") return provider def _normalizeProviderUserID(providerUserID): if providerUserID is None: raise ValueError("provider user ID is required") providerUserID = str(providerUserID).strip() if not providerUserID or len(providerUserID) > 255: raise ValueError("provider user ID must be between 1 and 255 characters") return providerUserID def _normalizeDisplayName(displayName): if displayName is None: return None if not isinstance(displayName, str): raise ValueError("display name must be a string") displayName = displayName.strip() if not displayName: return None if len(displayName) > 255: raise ValueError("display name must be at most 255 characters") return displayName def getDiscordEnrollmentMode(): mode = os.getenv("DISCORD_ENROLLMENT_MODE", DEFAULT_ENROLLMENT_MODE) mode = mode.strip().lower() return "open" if mode == "open" else DEFAULT_ENROLLMENT_MODE def getDiscordAllowlist(): configured = os.getenv("DISCORD_ALLOWLIST") if configured is None: configured = os.getenv("DISCORD_ALLOWED_USER_IDS", "") return { value for value in re.split(r"[\s,]+", configured) if value } def isDiscordEnrollmentAllowed(discordID): discordID = _normalizeProviderUserID(discordID) if getDiscordEnrollmentMode() == "open": return True return discordID in getDiscordAllowlist() def getProviderIdentity(provider, providerUserID): provider = _normalizeProvider(provider) providerUserID = _normalizeProviderUserID(providerUserID) return postgres.select_one( "provider_identities", {"provider": provider, "provider_user_id": providerUserID}, ) def getUserForIdentity(provider, providerUserID): identity = getProviderIdentity(provider, providerUserID) if not identity: return None return users.getUser(identity["user_uuid"]) def listUserIdentities(userUUID): return postgres.select("provider_identities", {"user_uuid": userUUID}) def linkProviderIdentity(userUUID, provider, providerUserID, displayName=None): if not users.doesUserUUIDExist(userUUID): raise ValueError("user does not exist") provider = _normalizeProvider(provider) providerUserID = _normalizeProviderUserID(providerUserID) displayName = _normalizeDisplayName(displayName) existing = getProviderIdentity(provider, providerUserID) if existing: if str(existing["user_uuid"]) != str(userUUID): raise ValueError("provider identity is already linked to another user") if existing.get("display_name") != displayName: updated = postgres.update( "provider_identities", {"display_name": displayName}, {"id": existing["id"]}, ) return updated[0] return existing return postgres.insert( "provider_identities", { "id": str(uuid.uuid4()), "user_uuid": str(userUUID), "provider": provider, "provider_user_id": providerUserID, "display_name": displayName, }, ) def unlinkProviderIdentity(userUUID, provider, providerUserID): provider = _normalizeProvider(provider) providerUserID = _normalizeProviderUserID(providerUserID) deleted = postgres.delete( "provider_identities", { "user_uuid": userUUID, "provider": provider, "provider_user_id": providerUserID, }, ) return bool(deleted) def linkDiscordIdentity(userUUID, discordID, displayName=None): return linkProviderIdentity( userUUID, DISCORD_PROVIDER, discordID, displayName=displayName, ) def linkDiscordUser(discordID, userUUID=None, username=None, displayName=None): if bool(userUUID) == bool(username): raise ValueError("provide either user UUID or username") if username: userUUID = users.getUserUUID(username) if not userUUID or not users.doesUserUUIDExist(userUUID): raise ValueError("user does not exist") return linkDiscordIdentity(userUUID, discordID, displayName=displayName) def getDiscordUser(discordID): return getUserForIdentity(DISCORD_PROVIDER, discordID) def getOrCreateDiscordUser(discordID, displayName=None, timezoneName=None): discordID = _normalizeProviderUserID(discordID) displayName = _normalizeDisplayName(displayName) if timezoneName is None: timezoneName = users.getDefaultTimezone() else: timezoneName = users.normalizeTimezone(timezoneName) existing = getProviderIdentity(DISCORD_PROVIDER, discordID) if existing: if displayName is not None and existing.get("display_name") != displayName: postgres.update( "provider_identities", {"display_name": displayName}, {"id": existing["id"]}, ) return users.getUser(existing["user_uuid"]) if not isDiscordEnrollmentAllowed(discordID): return None with postgres.get_cursor() as cursor: cursor.execute( "SELECT pg_advisory_xact_lock(hashtext(%(lock_key)s))", {"lock_key": f"discord:{discordID}"}, ) cursor.execute( """ SELECT users.* FROM provider_identities JOIN users ON users.id = provider_identities.user_uuid WHERE provider = %(provider)s AND provider_user_id = %(provider_user_id)s """, { "provider": DISCORD_PROVIDER, "provider_user_id": discordID, }, ) existingUser = cursor.fetchone() if existingUser: return dict(existingUser) userUUID = str(uuid.uuid4()) cursor.execute( """ INSERT INTO users (id, username, password_hashed, timezone) VALUES (%(id)s, NULL, NULL, %(timezone)s) RETURNING * """, {"id": userUUID, "timezone": timezoneName}, ) userRecord = dict(cursor.fetchone()) cursor.execute( """ INSERT INTO provider_identities ( id, user_uuid, provider, provider_user_id, display_name ) VALUES ( %(id)s, %(user_uuid)s, %(provider)s, %(provider_user_id)s, %(display_name)s ) """, { "id": str(uuid.uuid4()), "user_uuid": userUUID, "provider": DISCORD_PROVIDER, "provider_user_id": discordID, "display_name": displayName, }, ) return userRecord