| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113 |
- """登录与注销接口。"""
- from __future__ import annotations
- from datetime import datetime, timezone
- from flask import request
- from sqlalchemy import select
- from dms.api.v1.blueprint import api_v1
- from dms.common.enums import AuditAction, AuditResult, EnabledStatus
- from dms.common.errors import InvalidArgumentError, InvalidCredentialsError, UserDisabledError
- from dms.common.response import success_response
- from dms.database.transaction import for_update, transaction
- from dms.extensions import db
- from dms.models import User
- from dms.security.auth_context import get_auth_context
- from dms.security.decorators import bearer_auth_required
- from dms.security.jwt_tokens import issue_access_token
- from dms.security.passwords import verify_password
- from dms.services.audit_service import auth_audit, client_ip, record_login_failure
- from dms.services.user_service import user_summary
- def _login_payload() -> tuple[str, str, bool]:
- if not request.is_json:
- raise InvalidArgumentError("请求体必须是JSON对象")
- body = request.get_json(silent=True)
- if not isinstance(body, dict):
- raise InvalidArgumentError("请求体必须是JSON对象")
- expected = {"username", "password", "keepSignedIn"}
- if set(body) != expected:
- raise InvalidArgumentError("请求体必须且只能包含username、password和keepSignedIn")
- username = body["username"]
- password = body["password"]
- keep_signed_in = body["keepSignedIn"]
- if not isinstance(username, str) or not username.strip():
- raise InvalidArgumentError("username不能为空")
- if not isinstance(password, str) or not password:
- raise InvalidArgumentError("password不能为空")
- if not isinstance(keep_signed_in, bool):
- raise InvalidArgumentError("keepSignedIn必须是布尔值")
- return username.strip(), password, keep_signed_in
- @api_v1.post("/auth/login")
- def login():
- username, password, keep_signed_in = _login_payload()
- user = db.session.scalar(
- select(User).where(User.username == username, User.is_deleted.is_(False))
- )
- password_valid = verify_password(
- user.password_hash if user is not None else None,
- password,
- )
- if user is None or not password_valid:
- record_login_failure(
- user=user,
- attempted_username=username,
- failure_reason="INVALID_CREDENTIALS",
- )
- raise InvalidCredentialsError()
- if user.status != EnabledStatus.ENABLED.value:
- record_login_failure(
- user=user,
- attempted_username=username,
- failure_reason="USER_DISABLED",
- )
- raise UserDisabledError()
- token, expires_in = issue_access_token(user, keep_signed_in=keep_signed_in)
- with transaction() as session:
- user.last_login_at = datetime.now(timezone.utc).replace(tzinfo=None)
- user.last_login_ip = client_ip()
- user.row_version += 1
- session.add(
- auth_audit(
- action=AuditAction.LOGIN,
- result=AuditResult.SUCCESS,
- user=user,
- detail={"keepSignedIn": keep_signed_in},
- )
- )
- return success_response(
- {
- "accessToken": token,
- "tokenType": "Bearer",
- "expiresIn": expires_in,
- "user": user_summary(user),
- }
- )
- @api_v1.post("/auth/logout")
- @bearer_auth_required
- def logout():
- context = get_auth_context()
- with transaction() as session:
- user = session.scalar(for_update(select(User).where(User.id == context.user_id)))
- if user is None or user.is_deleted:
- raise InvalidCredentialsError()
- user.auth_version += 1
- user.row_version += 1
- session.add(
- auth_audit(
- action=AuditAction.LOGOUT,
- result=AuditResult.SUCCESS,
- user=user,
- detail={"authVersionIncremented": True},
- )
- )
- return success_response()
|