Build reusable bot framework
This commit is contained in:
351
core/api_keys.py
Normal file
351
core/api_keys.py
Normal file
@@ -0,0 +1,351 @@
|
||||
import datetime
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import uuid
|
||||
|
||||
import core.postgres as postgres
|
||||
import core.users as users
|
||||
|
||||
|
||||
USER_KEY_TYPE = "user"
|
||||
SERVICE_KEY_TYPE = "service"
|
||||
DEFAULT_SERVICE_SCOPES = (
|
||||
"discord:session",
|
||||
"outbox:claim",
|
||||
"outbox:deliver",
|
||||
)
|
||||
_TOKEN_PREFIX_LENGTH = 20
|
||||
_MIN_TOKEN_LENGTH = 32
|
||||
|
||||
|
||||
def _utcNow():
|
||||
return datetime.datetime.now(datetime.timezone.utc)
|
||||
|
||||
|
||||
def _hashToken(token):
|
||||
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _getPrefix(token):
|
||||
return token[:_TOKEN_PREFIX_LENGTH]
|
||||
|
||||
|
||||
def _normalizeScopes(scopes):
|
||||
if scopes is None:
|
||||
return []
|
||||
if isinstance(scopes, str):
|
||||
scopes = re.split(r"[\s,]+", scopes)
|
||||
if not isinstance(scopes, (list, tuple, set)):
|
||||
raise ValueError("scopes must be a list or comma-separated string")
|
||||
|
||||
normalized = []
|
||||
for scope in scopes:
|
||||
if not isinstance(scope, str) or not scope.strip():
|
||||
continue
|
||||
scope = scope.strip()
|
||||
if len(scope) > 100:
|
||||
raise ValueError("service scopes must be at most 100 characters")
|
||||
if scope not in normalized:
|
||||
normalized.append(scope)
|
||||
return normalized
|
||||
|
||||
|
||||
def _normalizeExpiry(expiresAt):
|
||||
if expiresAt is None:
|
||||
return None
|
||||
if isinstance(expiresAt, str):
|
||||
try:
|
||||
expiresAt = datetime.datetime.fromisoformat(expiresAt.replace("Z", "+00:00"))
|
||||
except ValueError as error:
|
||||
raise ValueError("expires_at must be an ISO-8601 datetime") from error
|
||||
if not isinstance(expiresAt, datetime.datetime):
|
||||
raise ValueError("expires_at must be a datetime")
|
||||
if expiresAt.tzinfo is None:
|
||||
expiresAt = expiresAt.replace(tzinfo=datetime.timezone.utc)
|
||||
expiresAt = expiresAt.astimezone(datetime.timezone.utc)
|
||||
if expiresAt <= _utcNow():
|
||||
raise ValueError("expires_at must be in the future")
|
||||
return expiresAt
|
||||
|
||||
|
||||
def _validateSecret(token):
|
||||
if not isinstance(token, str) or len(token) < _MIN_TOKEN_LENGTH:
|
||||
raise ValueError(f"API keys must be at least {_MIN_TOKEN_LENGTH} characters")
|
||||
return token
|
||||
|
||||
|
||||
def _keyID(value):
|
||||
try:
|
||||
return str(uuid.UUID(str(value)))
|
||||
except (TypeError, ValueError, AttributeError):
|
||||
return None
|
||||
|
||||
|
||||
def _publicKey(record, includeSecret=None):
|
||||
if not record:
|
||||
return None
|
||||
public = dict(record)
|
||||
public.pop("key_hash", None)
|
||||
scopes = public.get("scopes")
|
||||
if isinstance(scopes, str):
|
||||
public["scopes"] = json.loads(scopes)
|
||||
if includeSecret is not None:
|
||||
public["key"] = includeSecret
|
||||
return public
|
||||
|
||||
|
||||
def createApiKey(
|
||||
name,
|
||||
keyType,
|
||||
userUUID=None,
|
||||
serviceName=None,
|
||||
scopes=None,
|
||||
expiresAt=None,
|
||||
secret=None,
|
||||
):
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
raise ValueError("API key name is required")
|
||||
name = name.strip()
|
||||
if len(name) > 255:
|
||||
raise ValueError("API key name must be at most 255 characters")
|
||||
if keyType not in {USER_KEY_TYPE, SERVICE_KEY_TYPE}:
|
||||
raise ValueError("key_type must be user or service")
|
||||
|
||||
normalizedScopes = _normalizeScopes(scopes)
|
||||
if keyType == USER_KEY_TYPE:
|
||||
if not userUUID or not users.doesUserUUIDExist(userUUID):
|
||||
raise ValueError("user does not exist")
|
||||
if serviceName is not None:
|
||||
raise ValueError("user API keys cannot have a service name")
|
||||
if normalizedScopes:
|
||||
raise ValueError("user API keys cannot have service scopes")
|
||||
else:
|
||||
if userUUID is not None:
|
||||
raise ValueError("service API keys cannot have a user")
|
||||
if not isinstance(serviceName, str) or not serviceName.strip():
|
||||
raise ValueError("service name is required")
|
||||
serviceName = serviceName.strip()
|
||||
if len(serviceName) > 255:
|
||||
raise ValueError("service name must be at most 255 characters")
|
||||
|
||||
expiresAt = _normalizeExpiry(expiresAt)
|
||||
if secret is None:
|
||||
secret = f"llmbot_{keyType}_{secrets.token_urlsafe(32)}"
|
||||
secret = _validateSecret(secret)
|
||||
|
||||
keyData = {
|
||||
"id": str(uuid.uuid4()),
|
||||
"name": name,
|
||||
"key_type": keyType,
|
||||
"user_uuid": str(userUUID) if userUUID else None,
|
||||
"service_name": serviceName,
|
||||
"key_prefix": _getPrefix(secret),
|
||||
"key_hash": _hashToken(secret),
|
||||
"scopes": json.dumps(normalizedScopes),
|
||||
"expires_at": expiresAt,
|
||||
}
|
||||
with postgres.get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO api_keys (
|
||||
id, name, key_type, user_uuid, service_name,
|
||||
key_prefix, key_hash, scopes, expires_at
|
||||
) VALUES (
|
||||
%(id)s, %(name)s, %(key_type)s, %(user_uuid)s, %(service_name)s,
|
||||
%(key_prefix)s, %(key_hash)s, %(scopes)s::jsonb, %(expires_at)s
|
||||
)
|
||||
RETURNING *
|
||||
""",
|
||||
keyData,
|
||||
)
|
||||
record = dict(cursor.fetchone())
|
||||
return _publicKey(record, includeSecret=secret)
|
||||
|
||||
|
||||
def createUserApiKey(userUUID, name, expiresAt=None):
|
||||
return createApiKey(
|
||||
name,
|
||||
USER_KEY_TYPE,
|
||||
userUUID=userUUID,
|
||||
expiresAt=expiresAt,
|
||||
)
|
||||
|
||||
|
||||
def createServiceApiKey(serviceName, name, scopes, expiresAt=None):
|
||||
return createApiKey(
|
||||
name,
|
||||
SERVICE_KEY_TYPE,
|
||||
serviceName=serviceName,
|
||||
scopes=scopes,
|
||||
expiresAt=expiresAt,
|
||||
)
|
||||
|
||||
|
||||
def getApiKey(keyID):
|
||||
keyID = _keyID(keyID)
|
||||
if not keyID:
|
||||
return None
|
||||
return _publicKey(postgres.select_one("api_keys", {"id": keyID}))
|
||||
|
||||
|
||||
def listApiKeys(userUUID=None, serviceName=None, includeRevoked=False):
|
||||
if userUUID is not None and serviceName is not None:
|
||||
raise ValueError("filter by either user or service, not both")
|
||||
|
||||
clauses = []
|
||||
params = {}
|
||||
if userUUID is not None:
|
||||
clauses.append("user_uuid = %(user_uuid)s")
|
||||
params["user_uuid"] = str(userUUID)
|
||||
if serviceName is not None:
|
||||
clauses.append("service_name = %(service_name)s")
|
||||
params["service_name"] = serviceName
|
||||
if not includeRevoked:
|
||||
clauses.append("revoked_at IS NULL")
|
||||
|
||||
query = "SELECT * FROM api_keys"
|
||||
if clauses:
|
||||
query += " WHERE " + " AND ".join(clauses)
|
||||
query += " ORDER BY created_at DESC"
|
||||
return [_publicKey(record) for record in postgres.execute(query, params)]
|
||||
|
||||
|
||||
def listUserApiKeys(userUUID, includeRevoked=False):
|
||||
return listApiKeys(userUUID=userUUID, includeRevoked=includeRevoked)
|
||||
|
||||
|
||||
def revokeApiKey(keyID, userUUID=None, serviceName=None):
|
||||
keyID = _keyID(keyID)
|
||||
if not keyID:
|
||||
return False
|
||||
clauses = ["id = %(id)s", "revoked_at IS NULL"]
|
||||
params = {"id": keyID}
|
||||
if userUUID is not None:
|
||||
clauses.append("user_uuid = %(user_uuid)s")
|
||||
params["user_uuid"] = str(userUUID)
|
||||
if serviceName is not None:
|
||||
clauses.append("service_name = %(service_name)s")
|
||||
params["service_name"] = serviceName
|
||||
|
||||
records = postgres.execute(
|
||||
"UPDATE api_keys SET revoked_at = CURRENT_TIMESTAMP "
|
||||
f"WHERE {' AND '.join(clauses)} RETURNING id",
|
||||
params,
|
||||
)
|
||||
return bool(records)
|
||||
|
||||
|
||||
def revokeUserApiKey(userUUID, keyID):
|
||||
return revokeApiKey(keyID, userUUID=userUUID)
|
||||
|
||||
|
||||
def authenticateApiKey(token, requiredScopes=None):
|
||||
if not isinstance(token, str) or len(token) < _MIN_TOKEN_LENGTH:
|
||||
return None
|
||||
|
||||
records = postgres.select(
|
||||
"api_keys",
|
||||
{
|
||||
"key_prefix": _getPrefix(token),
|
||||
"key_hash": _hashToken(token),
|
||||
"revoked_at": None,
|
||||
},
|
||||
)
|
||||
if not records:
|
||||
return None
|
||||
|
||||
record = records[0]
|
||||
if not hmac.compare_digest(record["key_hash"], _hashToken(token)):
|
||||
return None
|
||||
expiresAt = record.get("expires_at")
|
||||
if expiresAt is not None:
|
||||
if expiresAt.tzinfo is None:
|
||||
expiresAt = expiresAt.replace(tzinfo=datetime.timezone.utc)
|
||||
if expiresAt <= _utcNow():
|
||||
return None
|
||||
|
||||
scopes = record.get("scopes") or []
|
||||
if isinstance(scopes, str):
|
||||
scopes = json.loads(scopes)
|
||||
requiredScopes = _normalizeScopes(requiredScopes)
|
||||
if requiredScopes and "*" not in scopes:
|
||||
if not set(requiredScopes).issubset(set(scopes)):
|
||||
return None
|
||||
|
||||
postgres.update(
|
||||
"api_keys",
|
||||
{"last_used_at": _utcNow()},
|
||||
{"id": record["id"]},
|
||||
)
|
||||
return {
|
||||
"type": record["key_type"],
|
||||
"authentication": "api_key",
|
||||
"user_uuid": record.get("user_uuid"),
|
||||
"service_name": record.get("service_name"),
|
||||
"api_key_id": record["id"],
|
||||
"scopes": scopes,
|
||||
"can_manage_api_keys": False,
|
||||
}
|
||||
|
||||
|
||||
def bootstrapServiceApiKey():
|
||||
secret = os.getenv("BOT_API_KEY")
|
||||
if not secret:
|
||||
return None
|
||||
_validateSecret(secret)
|
||||
|
||||
serviceName = "discord-bot"
|
||||
keyName = "BOT_API_KEY"
|
||||
scopes = _normalizeScopes(os.getenv("BOT_API_KEY_SCOPES"))
|
||||
if not scopes:
|
||||
scopes = list(DEFAULT_SERVICE_SCOPES)
|
||||
tokenHash = _hashToken(secret)
|
||||
|
||||
with postgres.get_cursor() as cursor:
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE api_keys
|
||||
SET revoked_at = CURRENT_TIMESTAMP
|
||||
WHERE key_type = 'service'
|
||||
AND service_name = %(service_name)s
|
||||
AND name = %(name)s
|
||||
AND key_hash != %(key_hash)s
|
||||
AND revoked_at IS NULL
|
||||
""",
|
||||
{
|
||||
"service_name": serviceName,
|
||||
"name": keyName,
|
||||
"key_hash": tokenHash,
|
||||
},
|
||||
)
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO api_keys (
|
||||
id, name, key_type, service_name,
|
||||
key_prefix, key_hash, scopes
|
||||
) VALUES (
|
||||
%(id)s, %(name)s, 'service', %(service_name)s,
|
||||
%(key_prefix)s, %(key_hash)s, %(scopes)s::jsonb
|
||||
)
|
||||
ON CONFLICT (key_hash) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
service_name = EXCLUDED.service_name,
|
||||
scopes = EXCLUDED.scopes,
|
||||
expires_at = NULL,
|
||||
revoked_at = NULL
|
||||
RETURNING *
|
||||
""",
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"name": keyName,
|
||||
"service_name": serviceName,
|
||||
"key_prefix": _getPrefix(secret),
|
||||
"key_hash": tokenHash,
|
||||
"scopes": json.dumps(scopes),
|
||||
},
|
||||
)
|
||||
return _publicKey(dict(cursor.fetchone()))
|
||||
Reference in New Issue
Block a user