login system
This commit is contained in:
+210
-3
@@ -1,13 +1,41 @@
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from functools import wraps
|
||||
from pathlib import Path
|
||||
|
||||
from flask import Flask, jsonify, request, send_from_directory
|
||||
import jwt
|
||||
from dotenv import load_dotenv
|
||||
from flask import Flask, g, jsonify, request, send_from_directory
|
||||
from flask_cors import CORS
|
||||
from werkzeug.security import check_password_hash, generate_password_hash
|
||||
|
||||
from db_py import count_tests, del_config, get_all_tests, get_config, set_config
|
||||
from db_py import (
|
||||
count_tests,
|
||||
count_users,
|
||||
create_user,
|
||||
del_config,
|
||||
get_all_tests,
|
||||
get_all_users,
|
||||
get_config,
|
||||
get_user_by_id,
|
||||
get_user_by_username,
|
||||
set_config,
|
||||
clear_users,
|
||||
)
|
||||
from scanner import full_scan, is_scan_in_progress, resolve_runtime_path, scan_results_only
|
||||
|
||||
BASE_DIR = Path(__file__).resolve().parent
|
||||
load_dotenv(BASE_DIR / ".env")
|
||||
|
||||
PORT = int(os.getenv("PORT", "3001"))
|
||||
JWT_SECRET = os.getenv("JWT_SECRET")
|
||||
JWT_ALGORITHM = os.getenv("JWT_ALGORITHM", "HS256")
|
||||
JWT_EXPIRES_HOURS = int(os.getenv("JWT_EXPIRES_HOURS", "8"))
|
||||
DEFAULT_ADMIN_USERNAME = os.getenv("DEFAULT_ADMIN_USERNAME")
|
||||
DEFAULT_ADMIN_PASSWORD = os.getenv("DEFAULT_ADMIN_PASSWORD")
|
||||
DEFAULT_VIEWER_USERNAME = os.getenv("DEFAULT_VIEWER_USERNAME")
|
||||
DEFAULT_VIEWER_PASSWORD = os.getenv("DEFAULT_VIEWER_PASSWORD")
|
||||
|
||||
ALLOWED_KEYS = {
|
||||
"target_dir",
|
||||
"results_dir",
|
||||
@@ -21,13 +49,110 @@ ALLOWED_KEYS = {
|
||||
"smb_domain",
|
||||
}
|
||||
|
||||
BASE_DIR = Path(__file__).resolve().parent
|
||||
DIST_DIR = BASE_DIR.parent / "dashboard" / "dist"
|
||||
|
||||
app = Flask(__name__, static_folder=str(DIST_DIR), static_url_path="")
|
||||
CORS(app)
|
||||
|
||||
|
||||
def _make_token(user):
|
||||
now = datetime.now(timezone.utc)
|
||||
payload = {
|
||||
"sub": str(user["id"]),
|
||||
"username": user["username"],
|
||||
"role": user["role"],
|
||||
"iat": int(now.timestamp()),
|
||||
"exp": int((now + timedelta(hours=JWT_EXPIRES_HOURS)).timestamp()),
|
||||
}
|
||||
return jwt.encode(payload, JWT_SECRET, algorithm=JWT_ALGORITHM)
|
||||
|
||||
|
||||
def _decode_token(token):
|
||||
try:
|
||||
return jwt.decode(token, JWT_SECRET, algorithms=[JWT_ALGORITHM])
|
||||
except jwt.InvalidTokenError:
|
||||
return None
|
||||
|
||||
|
||||
def _extract_bearer_token():
|
||||
auth_header = request.headers.get("Authorization", "")
|
||||
if not auth_header.lower().startswith("bearer "):
|
||||
return None
|
||||
return auth_header[7:].strip() or None
|
||||
|
||||
|
||||
def require_auth(fn):
|
||||
@wraps(fn)
|
||||
def wrapper(*args, **kwargs):
|
||||
token = _extract_bearer_token()
|
||||
if not token:
|
||||
return jsonify({"error": "Authentication required"}), 401
|
||||
|
||||
payload = _decode_token(token)
|
||||
if payload is None:
|
||||
return jsonify({"error": "Invalid or expired token"}), 401
|
||||
|
||||
try:
|
||||
user_id = int(payload.get("sub"))
|
||||
except (TypeError, ValueError):
|
||||
return jsonify({"error": "Invalid token subject"}), 401
|
||||
|
||||
user = get_user_by_id(user_id)
|
||||
if not user or not user.get("is_active"):
|
||||
return jsonify({"error": "User is not authorized"}), 401
|
||||
|
||||
g.current_user = {
|
||||
"id": user["id"],
|
||||
"username": user["username"],
|
||||
"role": user["role"],
|
||||
}
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def require_role(required_role):
|
||||
def decorator(fn):
|
||||
@wraps(fn)
|
||||
def wrapper(*args, **kwargs):
|
||||
current_user = getattr(g, "current_user", None)
|
||||
if not current_user:
|
||||
return jsonify({"error": "Authentication required"}), 401
|
||||
if current_user.get("role") != required_role:
|
||||
return jsonify({"error": "Forbidden"}), 403
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def _create_default_user_if_missing(username, password, role):
|
||||
if not username or not password:
|
||||
print(f"[server] Skipping default {role} seed: username/password not configured")
|
||||
return
|
||||
|
||||
existing = get_user_by_username(username)
|
||||
if existing:
|
||||
return
|
||||
|
||||
create_user(
|
||||
username,
|
||||
generate_password_hash(password),
|
||||
role=role,
|
||||
is_active=1,
|
||||
)
|
||||
print(f"[server] Created default {role} user: {username}")
|
||||
|
||||
|
||||
def _ensure_default_users():
|
||||
if count_users() == 0:
|
||||
print("[server] No users found. Seeding default accounts...")
|
||||
|
||||
_create_default_user_if_missing(DEFAULT_ADMIN_USERNAME, DEFAULT_ADMIN_PASSWORD, "admin")
|
||||
_create_default_user_if_missing(DEFAULT_VIEWER_USERNAME, DEFAULT_VIEWER_PASSWORD, "viewer")
|
||||
|
||||
|
||||
def _apply_smb_env_from_config():
|
||||
mapping = {
|
||||
"SMB_USERNAME": get_config("smb_username"),
|
||||
@@ -43,6 +168,7 @@ def _apply_smb_env_from_config():
|
||||
|
||||
|
||||
@app.get("/api/tests")
|
||||
@require_auth
|
||||
def get_tests_route():
|
||||
completed = request.args.get("completed")
|
||||
interference = request.args.get("interference")
|
||||
@@ -90,6 +216,7 @@ def get_tests_route():
|
||||
|
||||
|
||||
@app.get("/api/stats")
|
||||
@require_auth
|
||||
def get_stats_route():
|
||||
tests = get_all_tests()
|
||||
types = ["COE", "P2P", "P3P"]
|
||||
@@ -197,11 +324,84 @@ def get_stats_route():
|
||||
|
||||
|
||||
@app.get("/api/scan-status")
|
||||
@require_auth
|
||||
def get_scan_status_route():
|
||||
return jsonify({"scanning": is_scan_in_progress()})
|
||||
|
||||
|
||||
@app.post("/api/auth/login")
|
||||
def auth_login_route():
|
||||
body = request.get_json(silent=True)
|
||||
if not isinstance(body, dict):
|
||||
return jsonify({"error": "Request body must be a JSON object"}), 400
|
||||
|
||||
username = (body.get("username") or "").strip()
|
||||
password = body.get("password") or ""
|
||||
|
||||
if not username or not password:
|
||||
return jsonify({"error": "Username and password are required"}), 400
|
||||
|
||||
user = get_user_by_username(username)
|
||||
if not user or not user.get("is_active"):
|
||||
return jsonify({"error": "Invalid username or password"}), 401
|
||||
|
||||
if not check_password_hash(user["password_hash"], password):
|
||||
return jsonify({"error": "Invalid username or password"}), 401
|
||||
|
||||
token = _make_token(user)
|
||||
return jsonify(
|
||||
{
|
||||
"token": token,
|
||||
"user": {
|
||||
"id": user["id"],
|
||||
"username": user["username"],
|
||||
"role": user["role"],
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@app.get("/api/auth/me")
|
||||
@require_auth
|
||||
def auth_me_route():
|
||||
return jsonify({"user": g.current_user})
|
||||
|
||||
|
||||
@app.get("/api/users")
|
||||
@require_auth
|
||||
@require_role("admin")
|
||||
def list_users_route():
|
||||
return jsonify(get_all_users())
|
||||
|
||||
|
||||
@app.post("/api/users")
|
||||
@require_auth
|
||||
@require_role("admin")
|
||||
def create_user_route():
|
||||
body = request.get_json(silent=True)
|
||||
if not isinstance(body, dict):
|
||||
return jsonify({"error": "Request body must be a JSON object"}), 400
|
||||
|
||||
username = (body.get("username") or "").strip()
|
||||
password = body.get("password") or ""
|
||||
role = (body.get("role") or "viewer").strip().lower()
|
||||
|
||||
if not username or not password:
|
||||
return jsonify({"error": "Username and password are required"}), 400
|
||||
|
||||
if role not in {"admin", "viewer"}:
|
||||
return jsonify({"error": "Role must be admin or viewer"}), 400
|
||||
|
||||
if get_user_by_username(username):
|
||||
return jsonify({"error": "User already exists"}), 409
|
||||
|
||||
user_id = create_user(username, generate_password_hash(password), role=role, is_active=1)
|
||||
return jsonify({"id": user_id, "username": username, "role": role, "is_active": 1}), 201
|
||||
|
||||
|
||||
@app.get("/api/config")
|
||||
@require_auth
|
||||
@require_role("admin")
|
||||
def get_config_route():
|
||||
config = {}
|
||||
for key in ALLOWED_KEYS:
|
||||
@@ -210,6 +410,8 @@ def get_config_route():
|
||||
|
||||
|
||||
@app.post("/api/config")
|
||||
@require_auth
|
||||
@require_role("admin")
|
||||
def set_config_route():
|
||||
updates = request.get_json(silent=True)
|
||||
if not isinstance(updates, dict):
|
||||
@@ -249,6 +451,8 @@ def set_config_route():
|
||||
|
||||
|
||||
@app.post("/api/config/rescan")
|
||||
@require_auth
|
||||
@require_role("admin")
|
||||
def rescan_route():
|
||||
_apply_smb_env_from_config()
|
||||
target_dir = resolve_runtime_path(get_config("target_dir"))
|
||||
@@ -267,6 +471,8 @@ def rescan_route():
|
||||
|
||||
|
||||
@app.post("/api/config/rescan-results")
|
||||
@require_auth
|
||||
@require_role("admin")
|
||||
def rescan_results_route():
|
||||
_apply_smb_env_from_config()
|
||||
results_dir = resolve_runtime_path(get_config("results_dir"))
|
||||
@@ -295,6 +501,7 @@ def static_or_spa(path=""):
|
||||
|
||||
|
||||
def bootstrap():
|
||||
_ensure_default_users()
|
||||
_apply_smb_env_from_config()
|
||||
target_dir = resolve_runtime_path(get_config("target_dir"))
|
||||
results_dir = resolve_runtime_path(get_config("results_dir"))
|
||||
|
||||
Reference in New Issue
Block a user