feat(auth): wire auth.py to UserModel DB, drop fake_users_db, add require_role

This commit is contained in:
2026-08-06 16:22:58 +00:00
parent 4c983d4ab6
commit 437dbfe148
+94 -39
View File
@@ -1,83 +1,138 @@
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from typing import Optional from typing import Optional
from fastapi import Depends, HTTPException, status from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer from fastapi.security import OAuth2PasswordBearer
from jose import JWTError, jwt from jose import JWTError, jwt
from passlib.context import CryptContext from passlib.context import CryptContext
from pydantic import BaseModel from pydantic import BaseModel
from sqlalchemy.orm import Session
from app.database import get_db
from app.models import UserModel
SECRET_KEY = "change-me-in-production" SECRET_KEY = "change-me-in-production"
ALGORITHM = "HS256" ALGORITHM = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES = 120 ACCESS_TOKEN_EXPIRE_MINUTES = 120
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/auth/token") oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/auth/token")
# ---------- Pydantic schemas ----------
class Token(BaseModel): class Token(BaseModel):
access_token: str access_token: str
token_type: str = "bearer" token_type: str = "bearer"
class TokenData(BaseModel): class TokenData(BaseModel):
username: Optional[str] = None username: Optional[str] = None
class User(BaseModel):
class UserOut(BaseModel):
id: int
username: str username: str
disabled: bool = False email: Optional[str] = None
role: str = "viewer" role: str
disabled: bool
force_password_change: bool
created_at: datetime
fake_users_db = { class Config:
"admin": { from_attributes = True
"username": "admin",
"hashed_password": "$2b$12$j2rrsvzYhC9lQ4w6WJ1wPeY9CKEMMvmFo0xSg6u40qCMgfHdCqkfG",
"disabled": False,
"role": "admin",
},
"viewer": {
"username": "viewer",
"hashed_password": "$2b$12$DA7Nn4MVSr1m3Q0P6x1Qe.i6yd0qJ7Yx1C2VYLRNvKcJsteVEh9W6",
"disabled": False,
"role": "viewer",
},
}
def verify_password(plain_password: str, hashed_password: str) -> bool:
return pwd_context.verify(plain_password, hashed_password)
def get_user(username: str): # ---------- Crypto helpers ----------
user = fake_users_db.get(username)
if not user: def verify_password(plain: str, hashed: str) -> bool:
return None return pwd_context.verify(plain, hashed)
return User(username=user["username"], disabled=user["disabled"], role=user["role"])
def hash_password(plain: str) -> str:
return pwd_context.hash(plain)
def authenticate_user(username: str, password: str):
user_dict = fake_users_db.get(username)
if not user_dict:
return None
if not pwd_context.verify(password, user_dict["hashed_password"]):
return None
return User(username=user_dict["username"], disabled=user_dict["disabled"], role=user_dict["role"])
def create_access_token(data: dict, expires_delta: Optional[timedelta] = None) -> str: def create_access_token(data: dict, expires_delta: Optional[timedelta] = None) -> str:
to_encode = data.copy() to_encode = data.copy()
expire = datetime.now(timezone.utc) + (expires_delta or timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)) expire = datetime.now(timezone.utc) + (
expires_delta or timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
)
to_encode.update({"exp": expire}) to_encode.update({"exp": expire})
return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM) return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
def get_current_user(token: str = Depends(oauth2_scheme)) -> User:
credentials_exception = HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Could not validate credentials", headers={"WWW-Authenticate": "Bearer"}) # ---------- DB helpers ----------
def get_user_by_username(db: Session, username: str) -> Optional[UserModel]:
return db.query(UserModel).filter(UserModel.username == username).first()
def authenticate_user(db: Session, username: str, password: str) -> Optional[UserModel]:
user = get_user_by_username(db, username)
if not user or not verify_password(password, user.hashed_password):
return None
return user
def ensure_default_admin(db: Session) -> None:
"""Seed a default admin user if the users table is empty (first-run bootstrap)."""
if db.query(UserModel).count() == 0:
db.add(UserModel(
username="admin",
hashed_password=hash_password("admin"),
role="admin",
disabled=False,
force_password_change=False,
))
db.add(UserModel(
username="viewer",
hashed_password=hash_password("viewer"),
role="viewer",
disabled=False,
force_password_change=False,
))
db.commit()
# ---------- FastAPI dependencies ----------
def get_current_user(
token: str = Depends(oauth2_scheme),
db: Session = Depends(get_db),
) -> UserModel:
credentials_exception = HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Could not validate credentials",
headers={"WWW-Authenticate": "Bearer"},
)
try: try:
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM]) payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
username: str = payload.get("sub") username: str = payload.get("sub")
if username is None: if not username:
raise credentials_exception raise credentials_exception
token_data = TokenData(username=username)
except JWTError: except JWTError:
raise credentials_exception raise credentials_exception
user = get_user(token_data.username)
if user is None: user = get_user_by_username(db, username)
if user is None or user.disabled:
raise credentials_exception raise credentials_exception
return user return user
def get_current_active_user(current_user: User = Depends(get_current_user)) -> User:
def get_current_active_user(current_user: UserModel = Depends(get_current_user)) -> UserModel:
if current_user.disabled: if current_user.disabled:
raise HTTPException(status_code=400, detail="Inactive user") raise HTTPException(status_code=400, detail="Inactive user")
return current_user return current_user
def require_role(*roles: str):
"""Dependency factory: require one of the given roles."""
def _check(current_user: UserModel = Depends(get_current_active_user)) -> UserModel:
if current_user.role not in roles:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"Required role: {' or '.join(roles)}",
)
return current_user
return _check