225 lines
7.1 KiB
Python
225 lines
7.1 KiB
Python
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
|