343 lines
11 KiB
Python
343 lines
11 KiB
Python
"""JWT authentication dependencies for FastAPI."""
|
|
from fastapi import Depends, HTTPException, status
|
|
from fastapi.security import HTTPBearer
|
|
from starlette.requests import Request
|
|
from app.utils import decode_access_token
|
|
import jwt
|
|
from sqlalchemy.orm import Session
|
|
from app.database import get_db
|
|
from app.models import RevokedToken, User
|
|
import re
|
|
import datetime
|
|
|
|
security = HTTPBearer()
|
|
|
|
|
|
# In-memory revoked token store (non-persistent). Use DB for production.
|
|
# In-memory fallback (không bền vững). Giữ lại để chặn ngay trong tiến trình hiện tại.
|
|
REVOKED_TOKENS: set[str] = set()
|
|
LATEST_JTI_PER_USER: dict[int, str] = {}
|
|
|
|
def revoke_token(token: str) -> None:
|
|
REVOKED_TOKENS.add(token)
|
|
|
|
def set_latest_jti(user_id: int, jti: str) -> None:
|
|
LATEST_JTI_PER_USER[user_id] = jti
|
|
|
|
def revoke_token_db(db: Session, token: str) -> None:
|
|
"""Decode token and persist its jti to DB for revocation."""
|
|
payload = decode_access_token(token)
|
|
jti = payload.get("jti")
|
|
sub = payload.get("sub")
|
|
exp = payload.get("exp")
|
|
expires_at = None
|
|
if isinstance(exp, (int, float)):
|
|
expires_at = datetime.datetime.utcfromtimestamp(exp)
|
|
elif isinstance(exp, datetime.datetime):
|
|
expires_at = exp
|
|
|
|
if not jti:
|
|
# không có jti thì vẫn dùng in-memory để tránh bỏ sót
|
|
REVOKED_TOKENS.add(token)
|
|
return
|
|
|
|
if not db.query(RevokedToken).filter(RevokedToken.jti == jti).first():
|
|
db.add(RevokedToken(
|
|
jti=jti,
|
|
user_id=int(sub) if sub is not None else None,
|
|
expires_at=expires_at
|
|
))
|
|
db.commit()
|
|
|
|
async def get_current_user(request: Request, db: Session = Depends(get_db)):
|
|
"""Dependency to extract and validate JWT token from request headers.
|
|
|
|
Returns the decoded token payload with user info.
|
|
|
|
Raises:
|
|
HTTPException: If token is invalid, expired, or missing
|
|
"""
|
|
auth_header = request.headers.get("Authorization")
|
|
if not auth_header:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Missing authorization header",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
|
|
try:
|
|
scheme, token = auth_header.split()
|
|
if scheme.lower() != "bearer":
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Invalid authentication scheme",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
except ValueError:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Invalid authorization header",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
|
|
# Reject revoked tokens (in-memory fallback)
|
|
if token in REVOKED_TOKENS:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Token has been revoked",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
|
|
try:
|
|
payload = decode_access_token(token)
|
|
# Check DB revocation via jti
|
|
jti = payload.get("jti")
|
|
if jti and db.query(RevokedToken).filter(RevokedToken.jti == jti).first():
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Token has been revoked",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
|
|
user_id: str = payload.get("sub")
|
|
if user_id is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Invalid token",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
try:
|
|
uid_int = int(user_id)
|
|
latest = LATEST_JTI_PER_USER.get(uid_int)
|
|
if latest and jti and jti != latest:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Token is not the latest",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
u = db.query(User).filter(User.id == uid_int).first()
|
|
if u:
|
|
if payload.get("role_id") is None:
|
|
payload["role_id"] = getattr(u, "role_id", None)
|
|
if payload.get("unit_id") is None:
|
|
payload["unit_id"] = getattr(u, "unit_id", None)
|
|
if not payload.get("email"):
|
|
payload["email"] = getattr(u, "email", None)
|
|
if not payload.get("fullname"):
|
|
payload["fullname"] = getattr(u, "fullname", None)
|
|
if payload.get("status") is None:
|
|
payload["status"] = getattr(u, "status", None)
|
|
except HTTPException:
|
|
raise
|
|
except Exception:
|
|
pass
|
|
return payload
|
|
except jwt.ExpiredSignatureError:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Token has expired",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
except jwt.InvalidTokenError:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Invalid token",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
# ---- RBAC helpers ----
|
|
|
|
# Role IDs
|
|
ROLE_ADMIN = 1
|
|
ROLE_DIRECTOR = 2
|
|
ROLE_MANAGER = 3
|
|
ROLE_LEADER_PM = 4
|
|
ROLE_ADMIN_LIMITED = 9
|
|
|
|
def _role_id(user: dict) -> int | None:
|
|
try:
|
|
rid = user.get("role_id")
|
|
return int(rid) if rid is not None else None
|
|
except Exception:
|
|
return None
|
|
|
|
def is_admin(user: dict) -> bool:
|
|
return _role_id(user) in (ROLE_ADMIN, ROLE_ADMIN_LIMITED)
|
|
|
|
def is_director(user: dict) -> bool:
|
|
return _role_id(user) == ROLE_DIRECTOR
|
|
|
|
def is_manager(user: dict) -> bool:
|
|
return _role_id(user) == ROLE_MANAGER
|
|
|
|
def is_leader_pm(user: dict) -> bool:
|
|
return _role_id(user) == ROLE_LEADER_PM
|
|
|
|
def get_user_email(user: dict) -> str:
|
|
return (user.get("email") or "").strip()
|
|
|
|
def get_user_unit_id(user: dict) -> int | None:
|
|
try:
|
|
uid = user.get("unit_id")
|
|
return int(uid) if uid is not None else None
|
|
except Exception:
|
|
return None
|
|
# RBAC
|
|
ROLE_ADMIN = 1
|
|
ROLE_DIRECTOR = 2
|
|
ROLE_MANAGER = 3
|
|
ROLE_LEADER_PM = 4
|
|
ROLE_ADMIN_LIMITED = 9
|
|
|
|
def _role_id(user: dict) -> int | None:
|
|
try:
|
|
rid = user.get("role_id")
|
|
return int(rid) if rid is not None else None
|
|
except Exception:
|
|
return None
|
|
|
|
# def is_admin(user: dict) -> bool:
|
|
# return _role_id(user) in (ROLE_ADMIN, ROLE_ADMIN_LIMITED)
|
|
|
|
# def is_director(user: dict) -> bool:
|
|
# return _role_id(user) == ROLE_DIRECTOR
|
|
|
|
# def is_manager(user: dict) -> bool:
|
|
# return _role_id(user) == ROLE_MANAGER
|
|
|
|
# def is_leader_pm(user: dict) -> bool:
|
|
# return _role_id(user) == ROLE_LEADER_PM
|
|
|
|
# def get_user_email(user: dict) -> str:
|
|
# return (user.get("email") or "").strip()
|
|
|
|
def _method_path(request: Request) -> tuple[str, str]:
|
|
return request.method.upper(), request.url.path.lower()
|
|
|
|
def _path_match(path: str, patterns: list[str]) -> bool:
|
|
for p in patterns:
|
|
q = p.lower()
|
|
if "*" in q:
|
|
regex = "^" + re.escape(q).replace("\\*", ".*") + "$"
|
|
if re.match(regex, path):
|
|
return True
|
|
else:
|
|
if path == q or path.startswith(q + "/"):
|
|
return True
|
|
return False
|
|
|
|
DIRECTOR_GET = [
|
|
"/months",
|
|
"/units",
|
|
"/jira/tickets-overview",
|
|
"/scorecard*",
|
|
"/overview/nac",
|
|
"/overview/compliance",
|
|
"/months/*/compliance",
|
|
"/overview/smart-ir",
|
|
"/pentest/get-overview",
|
|
"/pentest/vuln-lastest-12month",
|
|
"/pentest/category",
|
|
"/pentest/category-vuln",
|
|
"/pentest/get-top-vulns",
|
|
"/pentest/get-target-detail",
|
|
"/pentest/download-target-checklist",
|
|
"/months/*/files",
|
|
"/files/download",
|
|
"/jira/tickets-week",
|
|
"/jira/tickets-12-month-lastest",
|
|
"/months/*/root-access",
|
|
"/months/*/attack-statics",
|
|
"/months/*/logsource",
|
|
"/months/*/nac-smartir",
|
|
"/months/*/compliance",
|
|
"/posts",
|
|
"/categories",
|
|
"/documents",
|
|
"/documents/download",
|
|
"/months/*/pam-tsc",
|
|
"/hardening/*",
|
|
"/pentest/log_history/*",
|
|
"/pentest-service*",
|
|
"/logsource/*/comments",
|
|
]
|
|
DIRECTOR_POST = [
|
|
"/auth/login",
|
|
"/auth/logout",
|
|
"/pentest-service*",
|
|
]
|
|
|
|
MANAGER_GET = DIRECTOR_GET + [
|
|
"/pentest/update-vuln-status",
|
|
]
|
|
MANAGER_POST = DIRECTOR_POST
|
|
|
|
LEADER_GET = MANAGER_GET
|
|
LEADER_POST = MANAGER_POST
|
|
|
|
def enforce_rbac(request: Request, user: dict = Depends(get_current_user)):
|
|
method, path = _method_path(request)
|
|
|
|
# Cho phép tất cả các nhóm người dùng đã đăng nhập truy cập vào /pentest-service
|
|
if "pentest-service" in path:
|
|
return
|
|
|
|
rid = _role_id(user)
|
|
if rid == ROLE_ADMIN:
|
|
return
|
|
if rid == ROLE_ADMIN_LIMITED:
|
|
deny_get = [
|
|
# "/pentest/category-vuln",
|
|
"/pentest/get-target-detail",
|
|
"/pentest/download-target-checklist",
|
|
]
|
|
deny_post = [
|
|
"/users",
|
|
"/roles",
|
|
]
|
|
deny_put = [
|
|
"/users/*",
|
|
"/roles/*",
|
|
]
|
|
deny_delete = [
|
|
"/users/*",
|
|
"/roles/*",
|
|
]
|
|
if method == "GET" and _path_match(path, deny_get):
|
|
raise HTTPException(status_code=403, detail="Forbidden")
|
|
if method == "POST" and _path_match(path, deny_post):
|
|
raise HTTPException(status_code=403, detail="Forbidden")
|
|
if method == "PUT" and _path_match(path, deny_put):
|
|
raise HTTPException(status_code=403, detail="Forbidden")
|
|
if method == "DELETE" and _path_match(path, deny_delete):
|
|
raise HTTPException(status_code=403, detail="Forbidden")
|
|
return
|
|
|
|
# Cho phép các method khác GET cho pentest-service
|
|
NONGET_ALLOWED_ALL_ROLES = [
|
|
"/systems/search-leak",
|
|
"/hardening/convert/upload",
|
|
"/pentest-service*",
|
|
"/posts/*/comments", # Cho phép tất cả người dùng đã đăng nhập thêm comment
|
|
"/logsource/*/comments", # Cho phép tất cả người dùng đã đăng nhập thêm bình luận LogSource
|
|
]
|
|
|
|
if method != "GET":
|
|
if _path_match(path, NONGET_ALLOWED_ALL_ROLES):
|
|
return
|
|
raise HTTPException(status_code=403, detail="Forbidden")
|
|
|
|
if is_director(user):
|
|
allowed = DIRECTOR_GET if method == "GET" else DIRECTOR_POST
|
|
elif is_manager(user):
|
|
allowed = MANAGER_GET if method == "GET" else MANAGER_POST
|
|
elif is_leader_pm(user):
|
|
allowed = LEADER_GET if method == "GET" else LEADER_POST
|
|
else:
|
|
# Nếu không thuộc các nhóm trên nhưng là /pentest-service thì đã return ở trên
|
|
raise HTTPException(status_code=403, detail="Forbidden")
|
|
|
|
if not _path_match(path, allowed):
|
|
raise HTTPException(status_code=403, detail="Forbidden")
|
|
return
|
|
|