diff --git a/app/auth.py b/app/auth.py index aaae52a..8c69b5f 100644 --- a/app/auth.py +++ b/app/auth.py @@ -11,7 +11,10 @@ def check_token_domain(db_path: str, token: str, domain: str) -> bool: # 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: - userid = db.User.get_id(db_path, username) + try: + userid = db.User.get_id(db_path, username) + except ValueError: + return None pwsalt, pwhash = db.User.get_password_hash(db_path, userid) diff --git a/app/db.py b/app/db.py index 17d4080..2b094d4 100644 --- a/app/db.py +++ b/app/db.py @@ -243,6 +243,9 @@ class User: for domain in domains: domain = domain.strip() + if len(domain) < 1: + continue + cursor.execute(''' INSERT INTO Domains (domain, userid) VALUES (?, ?) ''', (domain, userid)) @@ -270,15 +273,15 @@ class User: 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") + if not token_name.isalnum(): + raise ValueError("Token name has to be alphanumeric") + token_name_list = User.get_token_names(db_path, userid) if len(token_name_list) >= 5: diff --git a/app/main.py b/app/main.py index 698f836..a0e4299 100644 --- a/app/main.py +++ b/app/main.py @@ -49,19 +49,15 @@ def homepage(): # Handle login @app.route('/login', methods=["POST"]) def login(): - if "user" not in request.form or "pass" not in request.form: - return redirect("/"), 400 + userid = auth.user_login(DB_PATH, request.form.get("user"), request.form.get("pass")) - if request.form["user"].strip() == "" or request.form["pass"].strip() == "": - return redirect("/"), 400 - - userid = auth.user_login(DB_PATH, request.form["user"], request.form["pass"]) if userid is None: + flash("Invalid user or password") return redirect("/") session["USERID"] = userid - return redirect("/dashboard", code=302) + return redirect("/dashboard") # Log out by clearing session data @@ -76,7 +72,9 @@ def logout(): def dashboard(): if session.get("USERID") is None: return redirect("/") + userid = session.get("USERID") + username = db.User.get_username(DB_PATH, userid) domains = db.User.get_domains(DB_PATH, userid) tokens = db.User.get_token_names(DB_PATH, userid) @@ -96,11 +94,14 @@ def generate_token(): return redirect("/") if "token" not in request.form: - return redirect("/dashboard") - if request.form["token"].strip() == "": + flash("Token name not supplied") return redirect("/dashboard") - token = db.User.generate_token(DB_PATH, session.get("USERID"), request.form["token"]) + try: + token = db.User.generate_token(DB_PATH, session.get("USERID"), request.form["token"]) + except ValueError as e: + flash(str(e)) + token = None if token is not None: flash(token) @@ -115,11 +116,13 @@ def revoke_token(): return redirect("/") if "token" not in request.form: - return redirect("/dashboard") - if request.form["token"].strip() == "": + flash("Token name not supplied") return redirect("/dashboard") - db.User.revoke_token(DB_PATH, session.get("USERID"), request.form["token"]) + try: + db.User.revoke_token(DB_PATH, session.get("USERID"), request.form["token"]) + except ValueError as e: + flash(str(e)) return redirect("/dashboard") @@ -133,6 +136,7 @@ def create_user(): return redirect("/dashboard") if 'username' not in request.form or 'password' not in request.form: + flash('Required fields not supplied') return redirect("/dashboard") username = request.form.get("username").strip() @@ -147,9 +151,21 @@ def create_user(): domains = collapse_exp.sub(' ', domains.strip()) domains = domains.split(" ") - is_admin = request.form.get("is_admin") is not None + is_admin = request.form.get("is_admin") + if is_admin is not None: + is_admin = is_admin == "true" + else: + is_admin = False - flash(db.User.new(DB_PATH, username, password, domains, is_admin, request.form.get('email'))) + email = request.form.get("email") + if email is not None: + if email == "": + email = None + + try: + db.User.new(DB_PATH, username, password, domains, is_admin, email) + except ValueError as e: + flash(str(e)) return redirect("/dashboard") @@ -164,10 +180,11 @@ def update_user(): userid = request.form.get("userid") if userid is None: + flash("UserID not supplied") return redirect("/dashboard") - try: - userid = int(userid) - except ValueError: + + if not db.User.exists(DB_PATH, userid): + flash("User does not exist") return redirect("/dashboard") domains = request.form.get('domains') @@ -175,11 +192,17 @@ def update_user(): collapse_exp = re.compile(r'\s+') domains = collapse_exp.sub(' ', domains.strip()) domains = domains.split(" ") - db.User.set_domains(DB_PATH, userid, domains) + try: + db.User.set_domains(DB_PATH, userid, domains) + except ValueError as e: + flash(str(e)) password = request.form.get("password") if password is not None: - db.User.set_password(DB_PATH, userid, password) + try: + db.User.set_password(DB_PATH, userid, password) + except ValueError as e: + flash(str(e)) is_admin = request.form.get("is_admin") if is_admin is not None: @@ -206,9 +229,9 @@ def delete_user(): userid = request.form.get("userid") if userid is None: return redirect("/dashboard") - try: - userid = int(userid) - except ValueError: + + if db.User.exists(DB_PATH, userid): + flash("User does not exst") return redirect("/dashboard") db.User.delete(DB_PATH, userid) @@ -223,17 +246,15 @@ def change_password(): return redirect("/") if "pass" not in request.form or "pass-new" not in request.form or "pass-rep" not in request.form: - flash("password change failed") - return redirect("/dashboard") - if request.form["pass"] == "" or request.form["pass-new"] == "" or request.form["pass-rep"] == "": - flash("password change failed") + flash("Required values not supplied") return redirect("/dashboard") oldpass = request.form["pass"] newpass = request.form["pass-new"] + reppass = request.form["pass-rep"] - if newpass != request.form["pass-rep"]: - flash("passwords do not match") + if newpass != reppass: + flash("Passwords do not match") return redirect("/dashboard") try: @@ -242,7 +263,7 @@ def change_password(): flash(str(e)) return redirect("/dashboard") - flash("password changed successfully") + flash("Password changed successfully") return redirect("/dashboard") @@ -262,10 +283,10 @@ def update_addr(): if 'ip' not in request.args: ip = request.remote_addr - elif not validate_ip(request.args['ip']): - return jsonify({"status": "400", "code": "ip-error", "comment": "Invalid IP"}), 400 - else: + elif validate_ip(request.args['ip']): ip = request.args['ip'] + else: + return jsonify({"status": "400", "code": "ip-error", "comment": "Invalid IP"}), 400 if 'token' not in request.args: return jsonify({"status": "400", "code": "token-error", "comment": "Missing or invalid token"}), 400 diff --git a/app/templates/admin.html b/app/templates/admin.html index f2c3d91..50affcd 100644 --- a/app/templates/admin.html +++ b/app/templates/admin.html @@ -36,7 +36,7 @@