auth.py 3.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113
  1. """登录与注销接口。"""
  2. from __future__ import annotations
  3. from datetime import datetime, timezone
  4. from flask import request
  5. from sqlalchemy import select
  6. from dms.api.v1.blueprint import api_v1
  7. from dms.common.enums import AuditAction, AuditResult, EnabledStatus
  8. from dms.common.errors import InvalidArgumentError, InvalidCredentialsError, UserDisabledError
  9. from dms.common.response import success_response
  10. from dms.database.transaction import for_update, transaction
  11. from dms.extensions import db
  12. from dms.models import User
  13. from dms.security.auth_context import get_auth_context
  14. from dms.security.decorators import bearer_auth_required
  15. from dms.security.jwt_tokens import issue_access_token
  16. from dms.security.passwords import verify_password
  17. from dms.services.audit_service import auth_audit, client_ip, record_login_failure
  18. from dms.services.user_service import user_summary
  19. def _login_payload() -> tuple[str, str, bool]:
  20. if not request.is_json:
  21. raise InvalidArgumentError("请求体必须是JSON对象")
  22. body = request.get_json(silent=True)
  23. if not isinstance(body, dict):
  24. raise InvalidArgumentError("请求体必须是JSON对象")
  25. expected = {"username", "password", "keepSignedIn"}
  26. if set(body) != expected:
  27. raise InvalidArgumentError("请求体必须且只能包含username、password和keepSignedIn")
  28. username = body["username"]
  29. password = body["password"]
  30. keep_signed_in = body["keepSignedIn"]
  31. if not isinstance(username, str) or not username.strip():
  32. raise InvalidArgumentError("username不能为空")
  33. if not isinstance(password, str) or not password:
  34. raise InvalidArgumentError("password不能为空")
  35. if not isinstance(keep_signed_in, bool):
  36. raise InvalidArgumentError("keepSignedIn必须是布尔值")
  37. return username.strip(), password, keep_signed_in
  38. @api_v1.post("/auth/login")
  39. def login():
  40. username, password, keep_signed_in = _login_payload()
  41. user = db.session.scalar(
  42. select(User).where(User.username == username, User.is_deleted.is_(False))
  43. )
  44. password_valid = verify_password(
  45. user.password_hash if user is not None else None,
  46. password,
  47. )
  48. if user is None or not password_valid:
  49. record_login_failure(
  50. user=user,
  51. attempted_username=username,
  52. failure_reason="INVALID_CREDENTIALS",
  53. )
  54. raise InvalidCredentialsError()
  55. if user.status != EnabledStatus.ENABLED.value:
  56. record_login_failure(
  57. user=user,
  58. attempted_username=username,
  59. failure_reason="USER_DISABLED",
  60. )
  61. raise UserDisabledError()
  62. token, expires_in = issue_access_token(user, keep_signed_in=keep_signed_in)
  63. with transaction() as session:
  64. user.last_login_at = datetime.now(timezone.utc).replace(tzinfo=None)
  65. user.last_login_ip = client_ip()
  66. user.row_version += 1
  67. session.add(
  68. auth_audit(
  69. action=AuditAction.LOGIN,
  70. result=AuditResult.SUCCESS,
  71. user=user,
  72. detail={"keepSignedIn": keep_signed_in},
  73. )
  74. )
  75. return success_response(
  76. {
  77. "accessToken": token,
  78. "tokenType": "Bearer",
  79. "expiresIn": expires_in,
  80. "user": user_summary(user),
  81. }
  82. )
  83. @api_v1.post("/auth/logout")
  84. @bearer_auth_required
  85. def logout():
  86. context = get_auth_context()
  87. with transaction() as session:
  88. user = session.scalar(for_update(select(User).where(User.id == context.user_id)))
  89. if user is None or user.is_deleted:
  90. raise InvalidCredentialsError()
  91. user.auth_version += 1
  92. user.row_version += 1
  93. session.add(
  94. auth_audit(
  95. action=AuditAction.LOGOUT,
  96. result=AuditResult.SUCCESS,
  97. user=user,
  98. detail={"authVersionIncremented": True},
  99. )
  100. )
  101. return success_response()