Files
Lamont/core/identity.py
Chelsea Lee fbdf33e894
Some checks failed
CI / test (push) Has been cancelled
CI / compose-smoke (push) Has been cancelled
Build reusable bot framework
2026-07-19 21:53:24 -05:00

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