diff --git a/app/auth.py b/app/auth.py index 669d5ec..aaae52a 100644 --- a/app/auth.py +++ b/app/auth.py @@ -1,460 +1,32 @@ -import sqlite3 -from os import urandom +import db from hashlib import sha256 - -# Initialize database -def init_db(path: str) -> None: - with sqlite3.connect(path) as connection: - cursor = connection.cursor() - - # Auth token table - cursor.execute(''' - CREATE TABLE IF NOT EXISTS Tokens ( - TKHASH TEXT PRIMARY KEY NOT NULL, - USERID INTEGER NOT NULL, - NAME TEXT NOT NULL - ) - ''') - - # Registered domains table - cursor.execute(''' - CREATE TABLE IF NOT EXISTS Domains ( - DOMAIN TEXT PRIMARY KEY NOT NULL, - USERID INTEGER NOT NULL - ) - ''') - - # Users table - 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 - ) - ''') - - # Create a new admin if there are no accounts - cursor.execute(''' - SELECT COUNT(*) FROM Users - ''') - - if cursor.fetchone()[0] == 0: - username = "root" - password = urandom(8).hex() - - create_user(path, username, password, is_admin=True) - - print(f"No users detected, created new admin\nUser: {username}\nPass: {password}", flush=True) - - # Check if access token can update domain's ip def check_token_domain(db_path: str, token: str, domain: str) -> bool: - with sqlite3.connect(db_path) as connection: - cursor = connection.cursor() + userid = db.Token.get_user(db_path, token) - 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 + domain = domain.strip() + return domain in db.User.get_domains(db_path, userid) # Verify user's passowrd for login. Return UserID on success, None on failure def user_login(db_path: str, username: str, password: str) -> int | None: - with sqlite3.connect(db_path) as connection: - cursor = connection.cursor() + userid = db.User.get_id(db_path, username) - # Get salt - cursor.execute(''' - SELECT pwsalt FROM Users - WHERE username = ? - ''', (username,)) + pwsalt, pwhash = db.User.get_password_hash(db_path, userid) - password_salt = cursor.fetchone() - if password_salt is None: - return None - password_salt = password_salt[0] + password_hash = sha256( (pwsalt + password).encode('utf-8')).hexdigest() - # Calculate hash - password_hash = sha256( (password_salt + password).encode('utf-8') ).hexdigest() - - # Check - 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] - - -# Get username of UserID. Return username if user exists, None if not -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] - - -# Get list of domains assigned to UserID -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 - - -# Get list of tokens assigned to UserID -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 - - -# Generate a new user token with a name, assigned to UserID -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 - - -# Revoke access token -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)) + if pwhash == password_hash: + return userid + return None # Change password of user (by user) -def change_user_password(db_path: str, userid: int, oldpass: str, newpass: str) -> bool: - with sqlite3.connect(db_path) as connection: - cursor = connection.cursor() +def change_user_password(db_path: str, userid: int, oldpass: str, newpass: str) -> None: + username = db.User.get_username(db_path, userid) - cursor.execute(''' - SELECT COUNT(*) FROM Users - WHERE userid = ? - ''', (userid, )) + if user_login(db_path, username, oldpass) is None: + raise ValueError("Incorrect password") - 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) - - -# Set password of user (by admin) -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 - - -# Check if user is admin -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 - - -# Get list of users as dictionaries -def get_users(db_path: str) -> dict: - with sqlite3.connect(db_path) as connection: - cursor = connection.cursor() - - cursor.execute(''' - SELECT userid, username, is_admin, email FROM Users - ''') - - users = cursor.fetchall() - - users = list(map(lambda x: {"userid": x[0], "username": x[1], "domains": [], "email": x[3], "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 - - -# Create a new user -def create_user(db_path: str, username: str, password: str, domains: list = [], is_admin: bool = False, email: str = None): - with sqlite3.connect(db_path) as connection: - cursor = connection.cursor() - - # Sanitize username ??? - username.strip() - - # Sanitize email ??? - if email is not None: - email.strip() - - # Check if username is free - cursor.execute(''' - SELECT COUNT(*) FROM Users - WHERE username = ? - ''', (username, )) - - if cursor.fetchone()[0] != 0: - return "Username exists" - - # Sanitize 2 electric boogaloo - if username == "": - return "Username is empty" - - if not username.isalnum(): - return "Username is not alphanumeric" - - if password == "": - return "Password empty" - - if len(password) < 8: - return "Password is shorter than 8 characters" - - # Generate salt - salt = urandom(16).hex() - password_hash = sha256( (salt + password).encode('utf-8')).hexdigest() - - # Add user to the database - cursor.execute(''' - INSERT INTO Users (username, pwsalt, pwhash, email, is_admin) VALUES (?, ?, ?, ?, ?) - ''', (username, salt, password_hash, email, is_admin)) - - # Get new user's ID - cursor.execute(''' - SELECT userid FROM Users - WHERE username = ? - ''', (username, )) - - userid = cursor.fetchone()[0] - - # Add domains - for domain in domains: - domain = domain.strip() - - if domain == "": - continue - - cursor.execute(''' - SELECT COUNT(*) FROM Domains - WHERE domain = ? - ''', (domain, )) - - if cursor.fetchone()[0] != 0: - continue - - cursor.execute(''' - INSERT INTO Domains (userid, domain) VALUES (?, ?) - ''', (userid, domain)) - - -# Update user's field -def update_user(db_path: str, userid: int, domains: list[str] = None, password: str = None, is_admin: bool = None, email: str = None) -> tuple[bool, str]: - with sqlite3.connect(db_path) as connection: - cursor = connection.cursor() - - # Check if user exists - cursor.execute(''' - SELECT COUNT(*) FROM Users - WHERE userid = ? - ''', (userid, )) - - if cursor.fetchone()[0] != 1: - return False, "Uerid not found" - - if domains is not None: - # crude check if domains is iterable - for domain in domains: - break - - # Delete all domains and add the fresh list - cursor.execute(''' - DELETE FROM Domains - WHERE userid = ? - ''', (userid, )) - - for domain in domains: - cursor.execute(''' - INSERT INTO Domains (userid, domain) VALUES (?, ?) - ''', (userid, domain)) - - # Update passowrd - if password is not None: - set_user_password(db_path, userid, password) - - # Update if user is admin - if is_admin is not None: - cursor.execute(''' - UPDATE Users - SET - is_admin = ? - WHERE userid = ? - ''', (is_admin, userid)) - - # Update email ??? - if email is not None: - if email == "": - email = None - - cursor.execute(''' - UPDATE Users - SET - email = ? - WHERE userid = ? - ''', (email, userid)) - - -# Delete user adn all data -def delete_user(db_path: str, userid: int): - with sqlite3.connect(db_path) as connection: - cursor = connection.cursor() - - cursor.execute('DELETE FROM Domains WHERE userid = ?', (userid,)) - cursor.execute('DELETE FROM Tokens WHERE userid = ?', (userid,)) - cursor.execute('DELETE FROM Users WHERE userid = ?', (userid,)) + db.User.set_password(db_path, userid, newpass) diff --git a/app/db.py b/app/db.py new file mode 100644 index 0000000..17d4080 --- /dev/null +++ b/app/db.py @@ -0,0 +1,429 @@ +import sqlite3 +from os import urandom +from hashlib import sha256 + +# Initialize database +def init(db_path: str) -> None: + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + # Auth token table + cursor.execute(''' + CREATE TABLE IF NOT EXISTS Tokens ( + TKHASH TEXT PRIMARY KEY NOT NULL, + USERID INTEGER NOT NULL, + NAME TEXT NOT NULL + ) + ''') + + # Registered domains table + cursor.execute(''' + CREATE TABLE IF NOT EXISTS Domains ( + DOMAIN TEXT PRIMARY KEY NOT NULL, + USERID INTEGER NOT NULL + ) + ''') + + # Users table + 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 + ) + ''') + + # Create a new admin if there are no accounts + cursor.execute(''' + SELECT COUNT(*) FROM Users + ''') + + if cursor.fetchone()[0] == 0: + username = "root" + password = urandom(8).hex() + + User.new(db_path, username, password, is_admin=True) + + print(f"No users detected, created new admin\nUser: {username}\nPass: {password}", flush=True) + + +def list_users(db_path: str) -> list[dict]: + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT userid, username, is_admin, email FROM Users + ''') + + users = cursor.fetchall() + + users = list(map(lambda x: {"userid": x[0], "username": x[1], "domains": [], "email": x[3], "is_admin": 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 + + +# Collection of getters and setters +class User: + def exists(db_path: str, userid: int) -> bool: + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT COUNT(*) FROM Users + WHERE userid = ? + ''', (userid, )) + + return cursor.fetchone()[0] == 1 + + + def new(db_path: str, username: str, password: str, domains: list = [], is_admin: bool = False, email: str = None) -> int: + username = username.strip() + + if not username.isalnum(): + raise ValueError("Username has to be alphanumeric") + + if len(username) < 1: + raise ValueError("Username has to be at least 1 character long") + + if len(password) > 16: + raise ValueError("Username has to be at most 16 characters long") + + if len(password) < 8: + raise ValueError("Password has to be at least 8 characters long") + + if len(password) > 128: + raise ValueError("Password has to be at most 128 characters long") + + if email is not None: + email = email.strip() + + if len(email) > 128: + raise ValueError("Email has to be at most 128 characters long, be reasonable") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + pwsalt = urandom(16).hex() + pwhash = sha256( (pwsalt + password).encode('utf-8') ).hexdigest() + + cursor.execute(''' + INSERT INTO Users (username, pwsalt, pwhash, email, is_admin) VALUES (?, ?, ?, ?, ?) + ''', (username, pwsalt, pwhash, email, is_admin)) + + # Get new user's ID + cursor.execute(''' + SELECT userid FROM Users + WHERE username = ? + ''', (username, )) + + userid = cursor.fetchone()[0] + + User.set_domains(db_path, userid, domains) + + return userid + + + def delete(db_path: str, userid: int): + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute('DELETE FROM Domains WHERE userid = ?', (userid, )) + cursor.execute('DELETE FROM Tokens WHERE userid = ?', (userid, )) + cursor.execute('DELETE FROM Users WHERE userid = ?', (userid, )) + + + def is_admin(db_path: str, userid: int) -> bool: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT is_admin FROM Users + WHERE userid = ? + ''', (userid, )) + + return cursor.fetchone()[0] + + def set_admin(db_path: str, userid: int, is_admin: bool): + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + UPDATE Users + SET is_admin = ? + WHERE userid = ? + ''', (is_admin, userid)) + + def get_username(db_path: str, userid: int) -> str: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT username FROM Users + WHERE userid = ? + ''', (userid, )) + + return cursor.fetchone()[0] + + + def get_id(db_path: str, username: str) -> int: + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT userid FROM Users + WHERE username = ? + ''', (username, )) + + userid = cursor.fetchone() + + if userid is None: + raise ValueError(f"User {username} does not exist") + + return userid[0] + + def get_domains(db_path: str, userid: int) -> list[str]: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT domain FROM Domains + WHERE userid = ? + ''', (userid, )) + + domains = cursor.fetchall() + domains = list(map(lambda x: x[0], domains)) + + return domains + + + def set_domains(db_path: str, userid: int, domains: list[str]) -> None: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + #TODO domain regex + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + DELETE FROM Domains + WHERE userid = ? + ''', (userid, )) + + for domain in domains: + domain = domain.strip() + + cursor.execute(''' + INSERT INTO Domains (domain, userid) VALUES (?, ?) + ''', (domain, userid)) + + + def get_token_names(db_path: str, userid: int) -> list[str]: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT name FROM Tokens + WHERE userid = ? + ''', (userid, )) + + token_names = cursor.fetchall() + token_names = list(map(lambda x: x[0], token_names)) + + return token_names + + + def generate_token(db_path: str, userid: int, token_name: str) -> str | None: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + if not token_name.isalnum(): + raise ValueError("Token name has to be alphanumeric") + + if len(token_name) < 1: + raise ValueError("Token name has to be at least 1 character long") + + if len(token_name) > 16: + raise ValueError("Token name has to be at most 16 characters long") + + token_name_list = User.get_token_names(db_path, userid) + + if len(token_name_list) >= 5: + raise ValueError("One user can have at most 5 tokens") + + if token_name in token_name_list: + raise ValueError("Token name is not unique") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + # Try 10 times in case of collisions + for _ in range(10): + token = urandom(32).hex() + token_hash = sha256(token.encode("utf-8")).hexdigest() + + try: + cursor.execute(''' + INSERT INTO Tokens (tkhash, userid, name) VALUES (?, ?, ?) + ''', (token_hash, userid, token_name)) + except sqlite3.IntegrityError: + continue + + return token + + return None + + + def revoke_token(db_path: str, userid: int, token_name: str) -> None: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + if token_name not in User.get_token_names(db_path, userid): + raise ValueError("Token name not found") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + DELETE FROM Tokens + WHERE userid = ? + AND name = ? + ''', (userid, token_name)) + + + def get_password_hash(db_path: str, userid: int) -> tuple[str, str]: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT pwsalt, pwhash FROM Users + WHERE userid = ? + ''', (userid, )) + + return cursor.fetchone() + + + def set_password(db_path: str, userid: int, password: str) -> None: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + if len(password) < 8: + raise ValueError("Password has to be at least 8 characters long") + + if len(password) > 128: + raise ValueError("Password has to be at most 128 characters long") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + pwsalt = urandom(16).hex() + pwhash = sha256( (pwsalt + password).encode('utf-8') ).hexdigest() + + cursor.execute(''' + UPdATE Users + SET pwsalt = ?, + pwhash = ? + WHERE userid = ? + ''', (pwsalt, pwhash, userid)) + + def get_email(db_path: str, userid: int) -> str: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + SELECT email FROM Users + WHERE userid = ? + ''', (userid, )) + + return cursor.fetchone()[0] + + def set_email(db_path: str, userid: int, email: str | None) -> None: + if not User.exists(db_path, userid): + raise ValueError(f"User {userid} does not exist") + + if email is not None: + email = email.strip() + + if len(email) > 128: + raise ValueError("Email has to be at most 128 characters long, be reasonable") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + cursor.execute(''' + UPDATE Users + SET email = ? + WHERE userid = ? + ''', (email, userid, )) + + +# Collection of getters +class Token: + def exists(db_path: str, token: str) -> bool: + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + tkhash = sha256(token.encode('utf-8')).hexdigest() + + cursor.execute(''' + SELECT COUNT(*) FROM Tokens + WHERE tkhash = ? + ''', (tkhash, )) + + return cursor.fetchone()[0] == 1 + + + def get_user(db_path: str, token: str) -> int: + if not Token.exists(db_path, userid): + raise ValueError("Token does not exist") + + with sqlite3.connect(db_path) as connection: + cursor = connection.cursor() + + tkhash = sha256(token.encode('utf-8')).hexdigest() + + cursor.execute(''' + SELECT userid FROM Tokens + WHERE tkhash = ? + ''', (tkhash, )) + + return cursor.fetcone()[0]