import sqlite3 from os import urandom from hashlib import sha256 def init_db(path: str) -> None: with sqlite3.connect(path) as connection: cursor = connection.cursor() cursor.execute(''' CREATE TABLE IF NOT EXISTS Tokens ( TKHASH TEXT PRIMARY KEY NOT NULL, USERID INTEGER NOT NULL, NAME TEXT NOT NULL ) ''') cursor.execute(''' CREATE TABLE IF NOT EXISTS Domains ( DOMAIN TEXT PRIMARY KEY NOT NULL, USERID INTEGER NOT NULL ) ''') cursor.execute(''' CREATE TABLE IF NOT EXISTS Users ( USERID INTEGER PRIMARY KEY AUTOINCREMENT NOT NULL, USERNAME TEXT NOT NULL UNIQUE, PWHASH TEXT NOT NULL, PWSALT TEXT NOT NULL, EMAIL TEXT IS_ADMIN INTEGER NOT NULL DEFAULT 0 ) ''') connection.commit() cursor.close() def check_token_domain(db_path: str, token: str, domain: str) -> bool: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() token_hash = sha256(token.encode('utf-8')).hexdigest() cursor.execute(''' SELECT COUNT(*) FROM Tokens JOIN Domains ON Tokens.userid = Domains.userid WHERE tkhash = ? AND domain = ? ''', (token_hash, domain)) return cursor.fetchone()[0] == 1 def user_login(db_path: str, username: str, password: str) -> int | None: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT pwsalt FROM Users WHERE username = ? ''', (username,)) password_salt = cursor.fetchone() if password_salt is None: return None password_salt = password_salt[0] password_hash = sha256( (password_salt + password).encode('utf-8') ).hexdigest() cursor.execute(''' SELECT userid FROM Users WHERE username = ? AND pwhash = ? ''', (username, password_hash)) result = cursor.fetchone() if result is None: return None return result[0] def get_username(db_path: str, userid: int) -> str | None: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT username FROM Users WHERE userid = ? ''', (userid, )) username = cursor.fetchone() if username is None: return None return username[0] def get_user_domains(db_path: str, userid: int) -> list[str]: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT domain FROM Domains WHERE userid = ? ''', (userid, )) domains = cursor.fetchall() domains = map(lambda x: x[0], domains) return domains def get_user_tokens(db_path: str, userid: int) -> list[str]: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT name FROM Tokens WHERE userid = ? ''', (userid, )) tokens = cursor.fetchall() tokens = map(lambda x: x[0], tokens) return tokens def generate_user_token(db_path: str, userid: int, token_name: str) -> str | None: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT COUNT(*) FROM Users WHERE userid = ? ''', (userid, )) if cursor.fetchone()[0] != 1: return None cursor.execute(''' SELECT COUNT(*) FROM Tokens WHERE userid = ? ''', (userid, )) if cursor.fetchone()[0] >= 5: return None cursor.execute(''' SELECT COUNT(*) FROM Tokens WHERE userid = ? AND name = ? ''', (userid, token_name)) if cursor.fetchone()[0] != 0: return None if not token_name.isalnum(): return None if len(token_name) >= 16: return None for _ in range(10): try: token = urandom(32).hex() token_hash = sha256(token.encode("utf-8")).hexdigest() cursor.execute(''' INSERT INTO Tokens (tkhash, userid, name) VALUES (?, ?, ?) ''', (token_hash, userid, token_name)) return token except sqlite3.IntegrityError: pass return None def revoke_user_token(db_path: str, userid: int, token_name: str) -> None: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' DELETE FROM Tokens WHERE userid = ? AND name = ? ''', (userid, token_name)) def change_user_password(db_path: str, userid: int, oldpass: str, newpass: str) -> bool: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT COUNT(*) FROM Users WHERE userid = ? ''', (userid, )) if cursor.fetchone()[0] != 1: return False cursor.execute(''' SELECT pwsalt FROM Users WHERE userid = ? ''', (userid, )) salt = cursor.fetchone() if salt is None: return False salt = salt[0] oldpass_hash = sha256( (salt + oldpass).encode('utf-8')).hexdigest() cursor.execute(''' SELECT COUNT(*) FROM Users WHERE userid = ? AND pwhash = ? ''', (userid, oldpass_hash)) if cursor.fetchone()[0] != 1: return False return set_user_password(db_path, userid, newpass) def set_user_password(db_path: str, userid: int, newpass: str) -> bool: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT COUNT(*) FROM Users WHERE userid = ? ''', (userid, )) if cursor.fetchone()[0] != 1: return False salt = urandom(16).hex() newpass_hash = sha256( (salt + newpass).encode('utf-8')).hexdigest() cursor.execute(''' UPDATE Users SET pwsalt = ?, pwhash = ? WHERE userid = ? ''', (salt, newpass_hash, userid)) return True def is_admin(db_path: str, userid: int) -> bool: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT COUNT(*) FROM Users WHERE userid = ? AND IS_ADMIN = 1 ''', (userid, )) return cursor.fetchone()[0] == 1 def get_users(db_path: str) -> dict: with sqlite3.connect(db_path) as connection: cursor = connection.cursor() cursor.execute(''' SELECT userid, username, is_admin FROM Users ''') users = cursor.fetchall() users = list(map(lambda x: {"userid": x[0], "username": x[1], "domains": [], "is_admin": bool(x[2])}, users)) for user in users: cursor.execute(''' SELECT domain FROM Domains WHERE userid = ? ''', (user["userid"], )) domains = cursor.fetchall() domains = list(map(lambda x: x[0], domains)) user['domains'] = domains return users