Spaces:
Sleeping
Sleeping
| from fastapi import APIRouter, Depends, HTTPException, status | |
| from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm | |
| from sqlalchemy.orm import Session | |
| from datetime import timedelta | |
| from database import get_db | |
| from models.user import User | |
| from schemas.user import UserCreate, UserResponse, Token | |
| from utils.auth import verify_password, get_password_hash, create_access_token, decode_token, ACCESS_TOKEN_EXPIRE_MINUTES | |
| router = APIRouter() | |
| oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login") | |
| async def get_current_user(token: str = Depends(oauth2_scheme), db: Session = Depends(get_db)): | |
| """获取当前用户""" | |
| from jose import JWTError, jwt | |
| from utils.auth import SECRET_KEY, ALGORITHM | |
| credentials_exception = HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Could not validate credentials", | |
| headers={"WWW-Authenticate": "Bearer"}, | |
| ) | |
| try: | |
| payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM]) | |
| username: str = payload.get("sub") | |
| if username is None: | |
| raise credentials_exception | |
| except JWTError as e: | |
| print(f"JWTError: {e}") | |
| raise credentials_exception | |
| user = db.query(User).filter(User.username == username).first() | |
| if user is None: | |
| raise credentials_exception | |
| return user | |
| async def get_current_superadmin(current_user: User = Depends(get_current_user)): | |
| """获取当前超级管理员或admin用户""" | |
| if not current_user.is_superadmin and current_user.username != 'admin': | |
| raise HTTPException( | |
| status_code=status.HTTP_403_FORBIDDEN, | |
| detail="Not enough permissions" | |
| ) | |
| return current_user | |
| def is_admin_user(current_user: User) -> bool: | |
| """判断是否为admin用户,admin用户可以查看所有数据""" | |
| return current_user.username == 'admin' | |
| def register(user: UserCreate, db: Session = Depends(get_db), current_user: User = Depends(get_current_superadmin)): | |
| """注册用户(仅超级管理员可注册)""" | |
| # 检查用户名是否已存在 | |
| db_user = db.query(User).filter(User.username == user.username).first() | |
| if db_user: | |
| raise HTTPException( | |
| status_code=status.HTTP_400_BAD_REQUEST, | |
| detail="Username already registered" | |
| ) | |
| # 检查邮箱是否已存在 | |
| db_user = db.query(User).filter(User.email == user.email).first() | |
| if db_user: | |
| raise HTTPException( | |
| status_code=status.HTTP_400_BAD_REQUEST, | |
| detail="Email already registered" | |
| ) | |
| # 创建新用户 | |
| hashed_password = get_password_hash(user.password) | |
| db_user = User( | |
| username=user.username, | |
| email=user.email, | |
| password=hashed_password, | |
| company_code=user.company_code, | |
| is_superadmin=user.is_superadmin or False | |
| ) | |
| db.add(db_user) | |
| db.commit() | |
| db.refresh(db_user) | |
| return db_user | |
| from fastapi import APIRouter, Depends, HTTPException, status, Request | |
| from pydantic import BaseModel | |
| class PasswordChange(BaseModel): | |
| old_password: str | |
| new_password: str | |
| async def login(request: Request, db: Session = Depends(get_db)): | |
| """用户登录""" | |
| # 尝试解析JSON数据 | |
| try: | |
| data = await request.json() | |
| username = data.get("username") | |
| password = data.get("password") | |
| except Exception: | |
| # 如果JSON解析失败,尝试解析表单数据 | |
| form_data = await request.form() | |
| username = form_data.get("username") | |
| password = form_data.get("password") | |
| # 查找用户 | |
| user = db.query(User).filter(User.username == username).first() | |
| if not user or not verify_password(password, user.password): | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Incorrect username or password", | |
| headers={"WWW-Authenticate": "Bearer"}, | |
| ) | |
| # 创建访问令牌 | |
| access_token_expires = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES) | |
| access_token = create_access_token( | |
| data={"sub": user.username}, expires_delta=access_token_expires | |
| ) | |
| return {"access_token": access_token, "token_type": "bearer"} | |
| def get_me(current_user: User = Depends(get_current_user)): | |
| """获取当前用户信息""" | |
| return current_user | |
| def change_password(password_data: PasswordChange, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)): | |
| """修改当前用户密码""" | |
| # 验证原密码 | |
| if not verify_password(password_data.old_password, current_user.password): | |
| raise HTTPException( | |
| status_code=status.HTTP_400_BAD_REQUEST, | |
| detail="原密码不正确" | |
| ) | |
| # 更新密码 | |
| hashed_password = get_password_hash(password_data.new_password) | |
| current_user.password = hashed_password | |
| db.commit() | |
| db.refresh(current_user) | |
| return {"detail": "密码修改成功"} | |
| def get_users(current_user: User = Depends(get_current_superadmin), db: Session = Depends(get_db)): | |
| """获取用户列表(仅超级管理员可访问)""" | |
| users = db.query(User).all() | |
| return users | |
| def update_user(user_id: int, user: UserCreate, current_user: User = Depends(get_current_superadmin), db: Session = Depends(get_db)): | |
| """更新用户信息(仅超级管理员可访问)""" | |
| db_user = db.query(User).filter(User.id == user_id).first() | |
| if not db_user: | |
| raise HTTPException( | |
| status_code=status.HTTP_404_NOT_FOUND, | |
| detail="User not found" | |
| ) | |
| # 检查是否为admin账号,admin账号的用户名不能修改 | |
| if db_user.username == "admin" and user.username and user.username != "admin": | |
| raise HTTPException( | |
| status_code=status.HTTP_403_FORBIDDEN, | |
| detail="Admin username cannot be changed" | |
| ) | |
| # 更新用户信息 | |
| if user.username: | |
| db_user.username = user.username | |
| if user.email: | |
| db_user.email = user.email | |
| if user.password: | |
| db_user.password = get_password_hash(user.password) | |
| if user.company_code: | |
| db_user.company_code = user.company_code | |
| if hasattr(user, "is_superadmin"): | |
| db_user.is_superadmin = user.is_superadmin | |
| db.commit() | |
| db.refresh(db_user) | |
| return db_user | |
| def delete_user(user_id: int, current_user: User = Depends(get_current_superadmin), db: Session = Depends(get_db)): | |
| """删除用户(仅超级管理员可访问)""" | |
| db_user = db.query(User).filter(User.id == user_id).first() | |
| if not db_user: | |
| raise HTTPException( | |
| status_code=status.HTTP_404_NOT_FOUND, | |
| detail="User not found" | |
| ) | |
| # 检查是否为admin账号,admin账号永久不可删除 | |
| if db_user.username == "admin": | |
| raise HTTPException( | |
| status_code=status.HTTP_403_FORBIDDEN, | |
| detail="Admin account cannot be deleted" | |
| ) | |
| db.delete(db_user) | |
| db.commit() | |
| return {"detail": "User deleted successfully"} | |