Build reusable bot framework
This commit is contained in:
224
core/identity.py
Normal file
224
core/identity.py
Normal file
@@ -0,0 +1,224 @@
|
||||
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
|
||||
Reference in New Issue
Block a user