122 lines
3.8 KiB
Python
122 lines
3.8 KiB
Python
from fastapi import Depends, HTTPException, Request, status
|
|
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
|
|
from sqlalchemy.orm import Session
|
|
from typing import List
|
|
from app.config.database import get_db
|
|
from app.config.security import security
|
|
from app.models.auth.user_model import User
|
|
from app.models.auth.access_model import Access
|
|
|
|
security_scheme = HTTPBearer(auto_error=False)
|
|
|
|
def get_current_user(
|
|
request: Request,
|
|
credentials: HTTPAuthorizationCredentials = Depends(security_scheme),
|
|
db: Session = Depends(get_db)
|
|
) -> User:
|
|
token = credentials.credentials if credentials else request.cookies.get("access_token")
|
|
if not token:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Not authenticated"
|
|
)
|
|
try:
|
|
payload = security.verify_access_token(token)
|
|
user_id = payload.get("sub")
|
|
if user_id is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Invalid token payload"
|
|
)
|
|
except HTTPException:
|
|
raise
|
|
except Exception:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Could not validate credentials"
|
|
)
|
|
|
|
user = db.query(User).filter(User.id == user_id).first()
|
|
if user is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="User not found"
|
|
)
|
|
|
|
if user.status != "active":
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="User is inactive"
|
|
)
|
|
|
|
return user
|
|
|
|
def require_active_user(current_user: User = Depends(get_current_user)) -> User:
|
|
if current_user.status != "active":
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="Inactive user"
|
|
)
|
|
return current_user
|
|
|
|
def has_access(user: User, access_code: str) -> bool:
|
|
if not user.role:
|
|
return False
|
|
|
|
user_access_codes = {ra.access.access_code for ra in user.role.role_accesses}
|
|
|
|
return access_code in user_access_codes
|
|
|
|
def can_access(user: User, access_code: str, db: Session) -> bool:
|
|
if not user.role:
|
|
return False
|
|
|
|
user_access_codes = {ra.access.access_code for ra in user.role.role_accesses}
|
|
|
|
if access_code in user_access_codes:
|
|
return True
|
|
requested_access = db.query(Access).filter(
|
|
Access.access_code == access_code
|
|
).first()
|
|
|
|
if not requested_access:
|
|
return False
|
|
|
|
current = requested_access
|
|
while current.parent:
|
|
if current.parent.access_code in user_access_codes:
|
|
return True
|
|
current = current.parent
|
|
|
|
return False
|
|
|
|
def get_user_accesses(user: User) -> List[str]:
|
|
if not user.role:
|
|
return []
|
|
|
|
return [ra.access.access_code for ra in user.role.role_accesses]
|
|
|
|
def require_access(access_code: str):
|
|
def check_permission(current_user: User = Depends(get_current_user)) -> bool:
|
|
if not has_access(current_user, access_code):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail=f"Insufficient permissions. Required: {access_code}"
|
|
)
|
|
return True
|
|
|
|
return check_permission
|
|
|
|
def require_access_hierarchical(access_code: str):
|
|
def check_permission(
|
|
current_user: User = Depends(get_current_user),
|
|
db: Session = Depends(get_db)
|
|
) -> bool:
|
|
if not can_access(current_user, access_code, db):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail=f"Insufficient permissions. Required: {access_code}"
|
|
)
|
|
return True
|
|
|
|
return check_permission |