91 lines
3.0 KiB
Python
91 lines
3.0 KiB
Python
from datetime import datetime, timezone
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException
|
|
from sqlalchemy.orm import Session
|
|
|
|
from auth import create_access_token, get_current_user, hash_password, verify_password
|
|
from database import get_db
|
|
from models import User, UserSettings
|
|
from schemas import TokenResponse, UserLogin, UserOut, UserRegister, UserUpdate
|
|
|
|
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
|
|
|
|
|
def utc_now_iso() -> str:
|
|
return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
|
|
|
|
|
|
@router.post("/register", response_model=UserOut)
|
|
def register(data: UserRegister, db: Session = Depends(get_db)):
|
|
if db.query(User).filter(User.username == data.username).first():
|
|
raise HTTPException(status_code=400, detail="用户名已存在")
|
|
|
|
user = User(
|
|
username=data.username,
|
|
password_hash=hash_password(data.password),
|
|
created_at=utc_now_iso(),
|
|
)
|
|
db.add(user)
|
|
db.flush()
|
|
|
|
settings = UserSettings(
|
|
user_id=user.id,
|
|
daily_target=20,
|
|
master_required_count=3,
|
|
weak_wrong_threshold=3,
|
|
)
|
|
db.add(settings)
|
|
db.commit()
|
|
db.refresh(user)
|
|
return user
|
|
|
|
|
|
@router.post("/login", response_model=TokenResponse)
|
|
def login(data: UserLogin, db: Session = Depends(get_db)):
|
|
user = db.query(User).filter(User.username == data.username).first()
|
|
if not user or not verify_password(data.password, user.password_hash):
|
|
raise HTTPException(status_code=401, detail="用户名或密码错误")
|
|
|
|
token = create_access_token(user.id, user.username, remember=data.remember)
|
|
return TokenResponse(access_token=token)
|
|
|
|
|
|
@router.get("/me", response_model=UserOut)
|
|
def me(current_user: User = Depends(get_current_user)):
|
|
return current_user
|
|
|
|
|
|
@router.patch("/me", response_model=UserOut)
|
|
def update_me(
|
|
data: UserUpdate,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
):
|
|
if data.username is None and data.new_password is None:
|
|
raise HTTPException(status_code=400, detail="请填写要修改的内容")
|
|
|
|
if data.username is not None:
|
|
username = data.username.strip()
|
|
if not username:
|
|
raise HTTPException(status_code=400, detail="用户名不能为空")
|
|
exists = (
|
|
db.query(User)
|
|
.filter(User.username == username, User.id != current_user.id)
|
|
.first()
|
|
)
|
|
if exists:
|
|
raise HTTPException(status_code=400, detail="用户名已存在")
|
|
current_user.username = username
|
|
|
|
if data.new_password is not None:
|
|
if not data.current_password:
|
|
raise HTTPException(status_code=400, detail="修改密码需要先输入当前密码")
|
|
if not verify_password(data.current_password, current_user.password_hash):
|
|
raise HTTPException(status_code=400, detail="当前密码不正确")
|
|
current_user.password_hash = hash_password(data.new_password)
|
|
|
|
db.add(current_user)
|
|
db.commit()
|
|
db.refresh(current_user)
|
|
return current_user
|