352 lines
10 KiB
Python
352 lines
10 KiB
Python
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()))
|