1
0

auth.py 2.9 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364
  1. from fastapi import APIRouter, Depends, HTTPException, status
  2. from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm
  3. from datetime import timedelta
  4. from typing import Annotated, Any
  5. from app.db.session import get_db
  6. from app.crud import get_user_by_username, create_user
  7. from app.schemas import Token, UserCreate, UserResponse
  8. from app.core.security import create_access_token, ACCESS_TOKEN_EXPIRE_MINUTES
  9. from app.core.hashing import Hasher
  10. import logging
  11. logger = logging.getLogger(__name__)
  12. router = APIRouter()
  13. oauth2_scheme = OAuth2PasswordBearer(tokenUrl="api/v1/auth/login")
  14. @router.post("/login", response_model=Token)
  15. def login_for_access_token(form_data: Annotated[OAuth2PasswordRequestForm, Depends()], db: Any = Depends(get_db)):
  16. logger.debug(f"Login attempt for user: {form_data.username}")
  17. # Explicitly check for empty credentials (though OAuth2PasswordRequestForm should handle it)
  18. if not form_data.username or not form_data.password:
  19. logger.warning(f"Empty credentials provided for user: {form_data.username}")
  20. raise HTTPException(
  21. status_code=status.HTTP_400_BAD_REQUEST,
  22. detail="Username and password are required",
  23. )
  24. try:
  25. user = get_user_by_username(db, form_data.username)
  26. if not user or not Hasher.verify_password(form_data.password, user.password_hash):
  27. logger.warning(f"Failed login attempt for user: {form_data.username}")
  28. raise HTTPException(
  29. status_code=status.HTTP_401_UNAUTHORIZED,
  30. detail="用户名或密码错误",
  31. headers={"WWW-Authenticate": "Bearer"},
  32. )
  33. logger.info(f"Successful login for user: {form_data.username}")
  34. access_token_expires = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
  35. access_token = create_access_token(
  36. subject=user.username, expires_delta=access_token_expires
  37. )
  38. return {"access_token": access_token, "token_type": "bearer"}
  39. except HTTPException:
  40. raise
  41. except Exception as e:
  42. logger.error(f"Error during login for user {form_data.username}: {str(e)}", exc_info=True)
  43. # Re-raise to be caught by global exception handler, but we've logged it
  44. raise
  45. @router.post("/register", response_model=UserResponse)
  46. def register(user: UserCreate, db: Any = Depends(get_db)):
  47. if len(user.password) < 8:
  48. raise HTTPException(status_code=400, detail="密码至少需要 8 个字符")
  49. db_user = get_user_by_username(db, user.username)
  50. if db_user:
  51. raise HTTPException(status_code=400, detail="用户名已被注册")
  52. if user.email:
  53. from app.db.client import fetch_one
  54. if fetch_one(db.execute("SELECT id FROM users WHERE email = ?", [user.email])):
  55. raise HTTPException(status_code=400, detail="该邮箱已被注册")
  56. return create_user(db=db, user=user.model_copy(update={"role": "user"}))