完善邮件验证与密码重置安全流程
This commit is contained in:
@@ -0,0 +1,94 @@
|
||||
"""add email verification settings
|
||||
|
||||
Revision ID: 20260629_01
|
||||
Revises: 20260608_01
|
||||
Create Date: 2026-06-29 11:20:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
|
||||
revision: str = "20260629_01"
|
||||
down_revision: Union[str, None] = "20260608_01"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _table_exists(inspector: sa.Inspector, table_name: str) -> bool:
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def _uuid_column() -> sa.Column:
|
||||
return sa.Column(
|
||||
"id",
|
||||
postgresql.UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
inspector = sa.inspect(bind)
|
||||
|
||||
smtp_security = postgresql.ENUM("NONE", "SSL", "STARTTLS", name="smtp_security", create_type=False)
|
||||
email_purpose = postgresql.ENUM("REGISTER", "PASSWORD_RESET", name="email_verification_purpose", create_type=False)
|
||||
smtp_security.create(bind, checkfirst=True)
|
||||
email_purpose.create(bind, checkfirst=True)
|
||||
|
||||
if not _table_exists(inspector, "system_email_settings"):
|
||||
op.create_table(
|
||||
"system_email_settings",
|
||||
_uuid_column(),
|
||||
sa.Column("smtp_host", sa.String(length=255), nullable=False),
|
||||
sa.Column("smtp_port", sa.Integer(), nullable=False),
|
||||
sa.Column("smtp_security", smtp_security, server_default="SSL", nullable=False),
|
||||
sa.Column("smtp_username", sa.String(length=255), nullable=False),
|
||||
sa.Column("smtp_password_encrypted", sa.Text(), nullable=True),
|
||||
sa.Column("sender_email", sa.String(length=255), nullable=False),
|
||||
sa.Column("sender_name", sa.String(length=255), nullable=True),
|
||||
sa.Column("allowed_register_domain", sa.String(length=255), nullable=False),
|
||||
sa.Column("verification_code_ttl_minutes", sa.Integer(), nullable=False),
|
||||
sa.Column("send_cooldown_seconds", sa.Integer(), nullable=False),
|
||||
sa.Column("max_verify_attempts", sa.Integer(), nullable=False),
|
||||
sa.Column("updated_by", postgresql.UUID(as_uuid=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["updated_by"], ["users.id"]),
|
||||
)
|
||||
|
||||
if not _table_exists(inspector, "email_verification_codes"):
|
||||
op.create_table(
|
||||
"email_verification_codes",
|
||||
_uuid_column(),
|
||||
sa.Column("email", sa.String(length=255), nullable=False),
|
||||
sa.Column("purpose", email_purpose, server_default="REGISTER", nullable=False),
|
||||
sa.Column("code_hash", sa.String(length=255), nullable=False),
|
||||
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("verified_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("attempt_count", sa.Integer(), server_default="0", nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_email_verification_codes_email",
|
||||
"email_verification_codes",
|
||||
["email"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
inspector = sa.inspect(bind)
|
||||
if _table_exists(inspector, "email_verification_codes"):
|
||||
op.drop_index("ix_email_verification_codes_email", table_name="email_verification_codes")
|
||||
op.drop_table("email_verification_codes")
|
||||
if _table_exists(inspector, "system_email_settings"):
|
||||
op.drop_table("system_email_settings")
|
||||
sa.Enum("REGISTER", "PASSWORD_RESET", name="email_verification_purpose").drop(bind, checkfirst=True)
|
||||
sa.Enum("NONE", "SSL", "STARTTLS", name="smtp_security").drop(bind, checkfirst=True)
|
||||
@@ -0,0 +1,40 @@
|
||||
"""default multiple register email domains
|
||||
|
||||
Revision ID: 20260629_02
|
||||
Revises: 20260629_01
|
||||
Create Date: 2026-06-29 12:30:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = "20260629_02"
|
||||
down_revision: Union[str, None] = "20260629_01"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
inspector = sa.inspect(bind)
|
||||
if "system_email_settings" in inspector.get_table_names():
|
||||
op.execute(
|
||||
"UPDATE system_email_settings "
|
||||
"SET allowed_register_domain = 'huapont.cn,qq.com' "
|
||||
"WHERE allowed_register_domain IN ('huapont.cn', '@huapont.cn')"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
inspector = sa.inspect(bind)
|
||||
if "system_email_settings" in inspector.get_table_names():
|
||||
op.execute(
|
||||
"UPDATE system_email_settings "
|
||||
"SET allowed_register_domain = 'huapont.cn' "
|
||||
"WHERE allowed_register_domain = 'huapont.cn,qq.com'"
|
||||
)
|
||||
@@ -0,0 +1,123 @@
|
||||
"""email settings per register domain
|
||||
|
||||
Revision ID: 20260629_03
|
||||
Revises: 20260629_02
|
||||
Create Date: 2026-06-29 14:20:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
import uuid
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
|
||||
revision: str = "20260629_03"
|
||||
down_revision: Union[str, None] = "20260629_02"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _table_exists(inspector: sa.Inspector, table_name: str) -> bool:
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def _column_exists(inspector: sa.Inspector, table_name: str, column_name: str) -> bool:
|
||||
return any(column["name"] == column_name for column in inspector.get_columns(table_name))
|
||||
|
||||
|
||||
def _normalize_domains(value: str | None) -> list[str]:
|
||||
domains: list[str] = []
|
||||
for item in (value or "huapont.cn,qq.com").replace(",", ",").split(","):
|
||||
domain = item.strip().lower().lstrip("@")
|
||||
if domain and domain not in domains:
|
||||
domains.append(domain)
|
||||
return domains or ["huapont.cn", "qq.com"]
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
inspector = sa.inspect(bind)
|
||||
if not _table_exists(inspector, "system_email_settings"):
|
||||
return
|
||||
|
||||
if not _column_exists(inspector, "system_email_settings", "register_domain"):
|
||||
op.add_column("system_email_settings", sa.Column("register_domain", sa.String(length=255), nullable=True))
|
||||
|
||||
rows = bind.execute(sa.text("SELECT * FROM system_email_settings ORDER BY created_at ASC")).mappings().all()
|
||||
for row in rows:
|
||||
domains = _normalize_domains(row.get("allowed_register_domain"))
|
||||
primary_domain = domains[0]
|
||||
bind.execute(
|
||||
sa.text(
|
||||
"UPDATE system_email_settings "
|
||||
"SET register_domain = :register_domain, allowed_register_domain = :register_domain "
|
||||
"WHERE id = :id"
|
||||
),
|
||||
{"register_domain": primary_domain, "id": row["id"]},
|
||||
)
|
||||
for domain in domains[1:]:
|
||||
exists = bind.execute(
|
||||
sa.text("SELECT 1 FROM system_email_settings WHERE register_domain = :register_domain"),
|
||||
{"register_domain": domain},
|
||||
).first()
|
||||
if exists:
|
||||
continue
|
||||
bind.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO system_email_settings (
|
||||
id, register_domain, smtp_host, smtp_port, smtp_security, smtp_username,
|
||||
smtp_password_encrypted, sender_email, sender_name, allowed_register_domain,
|
||||
verification_code_ttl_minutes, send_cooldown_seconds, max_verify_attempts,
|
||||
updated_by, created_at, updated_at
|
||||
) VALUES (
|
||||
:id, :register_domain, :smtp_host, :smtp_port, :smtp_security, :smtp_username,
|
||||
:smtp_password_encrypted, :sender_email, :sender_name, :allowed_register_domain,
|
||||
:verification_code_ttl_minutes, :send_cooldown_seconds, :max_verify_attempts,
|
||||
:updated_by, now(), now()
|
||||
)
|
||||
"""
|
||||
),
|
||||
{
|
||||
"id": uuid.uuid4(),
|
||||
"register_domain": domain,
|
||||
"smtp_host": row["smtp_host"],
|
||||
"smtp_port": row["smtp_port"],
|
||||
"smtp_security": row["smtp_security"],
|
||||
"smtp_username": row["smtp_username"],
|
||||
"smtp_password_encrypted": row["smtp_password_encrypted"],
|
||||
"sender_email": row["sender_email"],
|
||||
"sender_name": row["sender_name"],
|
||||
"allowed_register_domain": domain,
|
||||
"verification_code_ttl_minutes": row["verification_code_ttl_minutes"],
|
||||
"send_cooldown_seconds": row["send_cooldown_seconds"],
|
||||
"max_verify_attempts": row["max_verify_attempts"],
|
||||
"updated_by": row["updated_by"],
|
||||
},
|
||||
)
|
||||
|
||||
op.alter_column("system_email_settings", "register_domain", existing_type=sa.String(length=255), nullable=False)
|
||||
op.create_unique_constraint(
|
||||
"uq_system_email_settings_register_domain",
|
||||
"system_email_settings",
|
||||
["register_domain"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
inspector = sa.inspect(bind)
|
||||
if not _table_exists(inspector, "system_email_settings"):
|
||||
return
|
||||
uniques = {item["name"] for item in inspector.get_unique_constraints("system_email_settings")}
|
||||
if "uq_system_email_settings_register_domain" in uniques:
|
||||
op.drop_constraint(
|
||||
"uq_system_email_settings_register_domain",
|
||||
"system_email_settings",
|
||||
type_="unique",
|
||||
)
|
||||
if _column_exists(inspector, "system_email_settings", "register_domain"):
|
||||
op.drop_column("system_email_settings", "register_domain")
|
||||
@@ -0,0 +1,26 @@
|
||||
"""add password reset email purpose
|
||||
|
||||
Revision ID: 20260629_04
|
||||
Revises: 20260629_03
|
||||
Create Date: 2026-06-29 16:50:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260629_04"
|
||||
down_revision: Union[str, None] = "20260629_03"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("ALTER TYPE email_verification_purpose ADD VALUE IF NOT EXISTS 'PASSWORD_RESET'")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# PostgreSQL enum values cannot be removed safely without recreating the type.
|
||||
pass
|
||||
@@ -0,0 +1,26 @@
|
||||
"""add password reset link email purpose
|
||||
|
||||
Revision ID: 20260630_01
|
||||
Revises: 20260629_04
|
||||
Create Date: 2026-06-30 10:30:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "20260630_01"
|
||||
down_revision: Union[str, None] = "20260629_04"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("ALTER TYPE email_verification_purpose ADD VALUE IF NOT EXISTS 'PASSWORD_RESET_LINK'")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# PostgreSQL enum values cannot be removed safely without recreating the type.
|
||||
pass
|
||||
@@ -0,0 +1,76 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.deps import get_db_session, require_roles
|
||||
from app.schemas.email_settings import (
|
||||
EmailDomainCreateRequest,
|
||||
EmailSettingsListResponse,
|
||||
EmailSettingsRead,
|
||||
EmailSettingsUpdate,
|
||||
EmailTestRequest,
|
||||
)
|
||||
from app.services import email_service
|
||||
|
||||
router = APIRouter(prefix="/email-settings")
|
||||
|
||||
|
||||
@router.get("/", response_model=EmailSettingsListResponse)
|
||||
async def read_email_settings(
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
current_user=Depends(require_roles(["ADMIN"])),
|
||||
) -> EmailSettingsListResponse:
|
||||
items = [email_service.to_email_settings_read(row) for row in await email_service.list_email_settings(db)]
|
||||
return EmailSettingsListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.post("/", response_model=EmailSettingsRead)
|
||||
async def create_email_domain(
|
||||
payload: EmailDomainCreateRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
current_user=Depends(require_roles(["ADMIN"])),
|
||||
) -> EmailSettingsRead:
|
||||
settings_row = await email_service.create_email_domain(db, payload.register_domain, updated_by=current_user.id)
|
||||
return email_service.to_email_settings_read(settings_row)
|
||||
|
||||
|
||||
@router.get("/{register_domain}", response_model=EmailSettingsRead)
|
||||
async def read_email_domain_settings(
|
||||
register_domain: str,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
current_user=Depends(require_roles(["ADMIN"])),
|
||||
) -> EmailSettingsRead:
|
||||
return email_service.to_email_settings_read(await email_service.get_email_settings(db, register_domain))
|
||||
|
||||
|
||||
@router.put("/{register_domain}", response_model=EmailSettingsRead)
|
||||
async def update_email_settings(
|
||||
register_domain: str,
|
||||
payload: EmailSettingsUpdate,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
current_user=Depends(require_roles(["ADMIN"])),
|
||||
) -> EmailSettingsRead:
|
||||
settings_row = await email_service.upsert_email_settings(db, register_domain, payload, updated_by=current_user.id)
|
||||
return email_service.to_email_settings_read(settings_row)
|
||||
|
||||
|
||||
@router.delete("/{register_domain}")
|
||||
async def delete_email_domain(
|
||||
register_domain: str,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
current_user=Depends(require_roles(["ADMIN"])),
|
||||
):
|
||||
await email_service.delete_email_domain(db, register_domain)
|
||||
return {"message": "邮箱后缀已删除"}
|
||||
|
||||
|
||||
@router.post("/{register_domain}/test")
|
||||
async def test_email_settings(
|
||||
register_domain: str,
|
||||
payload: EmailTestRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
current_user=Depends(require_roles(["ADMIN"])),
|
||||
):
|
||||
await email_service.send_test_email(db, register_domain, str(payload.recipient_email))
|
||||
return {"message": "测试邮件已发送"}
|
||||
+107
-2
@@ -1,7 +1,7 @@
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from fastapi import File, UploadFile
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, EmailStr, Field
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from pathlib import Path
|
||||
import uuid
|
||||
@@ -12,7 +12,19 @@ from app.core.security import create_access_token, decode_token_allow_expired, o
|
||||
from app.core.deps import get_current_user, get_db_session
|
||||
from app.crud import user as user_crud
|
||||
from app.models.user import UserStatus
|
||||
from app.schemas.email_settings import (
|
||||
EmailCodeResponse,
|
||||
EmailCodeVerifyResponse,
|
||||
PasswordResetCodeVerifyRequest,
|
||||
PasswordResetCodeVerifyResponse,
|
||||
PasswordResetLinkSendRequest,
|
||||
PasswordResetRequest,
|
||||
PasswordResetTokenRequest,
|
||||
RegisterEmailCodeSendRequest,
|
||||
RegisterEmailCodeVerifyRequest,
|
||||
)
|
||||
from app.schemas.user import Token, UserRead, UserRegisterRequest, UserSelfUpdate, UserUpdate
|
||||
from app.services import email_service
|
||||
from fastapi.responses import FileResponse
|
||||
|
||||
|
||||
@@ -39,6 +51,14 @@ class ExtendResponse(BaseModel):
|
||||
expiresAt: datetime
|
||||
|
||||
|
||||
class EmailAvailabilityResponse(BaseModel):
|
||||
available: bool
|
||||
|
||||
|
||||
class EmailDomainsResponse(BaseModel):
|
||||
items: list[str]
|
||||
|
||||
|
||||
router = APIRouter()
|
||||
AVATAR_ROOT = Path(__file__).resolve().parent.parent.parent / "uploads" / "avatars"
|
||||
AVATAR_ROOT.mkdir(parents=True, exist_ok=True)
|
||||
@@ -108,6 +128,90 @@ def ensure_user_active(db_user) -> None:
|
||||
)
|
||||
|
||||
|
||||
@router.get("/email-domains", response_model=EmailDomainsResponse)
|
||||
async def read_email_domains(
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> EmailDomainsResponse:
|
||||
rows = await email_service.list_email_settings(db)
|
||||
return EmailDomainsResponse(items=[row.register_domain for row in rows])
|
||||
|
||||
|
||||
@router.get("/register/email-availability", response_model=EmailAvailabilityResponse)
|
||||
async def check_register_email_availability(
|
||||
email: EmailStr = Query(...),
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> EmailAvailabilityResponse:
|
||||
existing = await user_crud.get_by_email(db, str(email))
|
||||
return EmailAvailabilityResponse(available=existing is None)
|
||||
|
||||
|
||||
@router.post("/register/email-code/send", response_model=EmailCodeResponse)
|
||||
async def send_register_email_code(
|
||||
payload: RegisterEmailCodeSendRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> EmailCodeResponse:
|
||||
await email_service.send_register_code(db, str(payload.email))
|
||||
return EmailCodeResponse(message="验证码已发送")
|
||||
|
||||
|
||||
@router.post("/register/email-code/verify", response_model=EmailCodeVerifyResponse)
|
||||
async def verify_register_email_code(
|
||||
payload: RegisterEmailCodeVerifyRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> EmailCodeVerifyResponse:
|
||||
await email_service.verify_register_code(db, str(payload.email), payload.code)
|
||||
return EmailCodeVerifyResponse(verified=True)
|
||||
|
||||
|
||||
@router.post("/password-reset/email-code/send", response_model=EmailCodeResponse)
|
||||
async def send_password_reset_email_code(
|
||||
payload: PasswordResetLinkSendRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> EmailCodeResponse:
|
||||
await email_service.send_password_reset_code(db, str(payload.email))
|
||||
return EmailCodeResponse(message="验证码发送成功,请查收邮箱")
|
||||
|
||||
|
||||
@router.post("/password-reset/email-code/verify", response_model=PasswordResetCodeVerifyResponse)
|
||||
async def verify_password_reset_email_code(
|
||||
payload: PasswordResetCodeVerifyRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> PasswordResetCodeVerifyResponse:
|
||||
reset_token = await email_service.verify_password_reset_code(db, str(payload.email), payload.code)
|
||||
return PasswordResetCodeVerifyResponse(verified=True, reset_token=reset_token)
|
||||
|
||||
|
||||
@router.post("/password-reset-link/send", response_model=EmailCodeResponse)
|
||||
async def send_password_reset_link(
|
||||
payload: PasswordResetLinkSendRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> EmailCodeResponse:
|
||||
await email_service.send_password_reset_link(
|
||||
db,
|
||||
str(payload.email),
|
||||
frontend_origin=settings.FRONTEND_PUBLIC_URL,
|
||||
)
|
||||
return EmailCodeResponse(message="如果账号存在,重置链接已发送")
|
||||
|
||||
|
||||
@router.post("/password-reset", response_model=EmailCodeResponse)
|
||||
async def reset_password(
|
||||
payload: PasswordResetRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> EmailCodeResponse:
|
||||
await email_service.reset_password_with_code(db, str(payload.email), payload.code, payload.password)
|
||||
return EmailCodeResponse(message="密码已重置,请返回登录")
|
||||
|
||||
|
||||
@router.post("/password-reset-link", response_model=EmailCodeResponse)
|
||||
async def reset_password_with_link(
|
||||
payload: PasswordResetTokenRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> EmailCodeResponse:
|
||||
await email_service.reset_password_with_token(db, payload.token, payload.password)
|
||||
return EmailCodeResponse(message="密码已重置,请返回登录")
|
||||
|
||||
|
||||
@router.post("/register", status_code=status.HTTP_201_CREATED)
|
||||
async def register(
|
||||
payload: UserRegisterRequest,
|
||||
@@ -116,6 +220,7 @@ async def register(
|
||||
existing = await user_crud.get_by_email(db, payload.email)
|
||||
if existing:
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="邮箱已注册")
|
||||
await email_service.ensure_register_email_verified(db, str(payload.email))
|
||||
await user_crud.create_registered_user(db, payload)
|
||||
return {"message": "注册成功,请登录"}
|
||||
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.v1 import auth, users, studies, sites, members, attachments, audit_logs, dashboard, subjects, visits, aes, finance_dashboard, fees_contracts, drug_shipments, material_equipments, project_milestones, startup, precautions, subject_histories, subject_pds, study_subject_pds, faq_categories, faqs, documents, etmf, overview, notifications, monitoring_visit_issues, api_permissions, permission_monitoring, permission_templates, system_permissions, study_active_roles
|
||||
from app.api.v1 import auth, users, admin_email_settings, studies, sites, members, attachments, audit_logs, dashboard, subjects, visits, aes, finance_dashboard, fees_contracts, drug_shipments, material_equipments, project_milestones, startup, precautions, subject_histories, subject_pds, study_subject_pds, faq_categories, faqs, documents, etmf, overview, notifications, monitoring_visit_issues, api_permissions, permission_monitoring, permission_templates, system_permissions, study_active_roles
|
||||
|
||||
|
||||
api_router = APIRouter()
|
||||
api_router.include_router(auth.router, prefix="/auth", tags=["auth"])
|
||||
api_router.include_router(admin_email_settings.router, prefix="/admin", tags=["admin"])
|
||||
api_router.include_router(users.router, prefix="/users", tags=["users"])
|
||||
api_router.include_router(studies.router, prefix="/studies", tags=["studies"])
|
||||
api_router.include_router(overview.router, prefix="/studies/{study_id}", tags=["overview"])
|
||||
|
||||
@@ -24,6 +24,8 @@ class Settings(BaseSettings):
|
||||
LOGIN_RSA_KEY_ID: str = "default"
|
||||
LOGIN_CHALLENGE_TTL_SECONDS: int = 120
|
||||
LOGIN_CHALLENGE_MAX_ACTIVE: int = 1000
|
||||
SETTINGS_ENCRYPTION_KEY: Optional[str] = None
|
||||
FRONTEND_PUBLIC_URL: str = "http://localhost:8888"
|
||||
IP2REGION_XDB_PATH: Optional[str] = None
|
||||
IP2REGION_IPV6_XDB_PATH: Optional[str] = None
|
||||
|
||||
|
||||
@@ -42,3 +42,4 @@ from app.models.permission_access_log import PermissionAccessLog # noqa: F401
|
||||
from app.models.permission_metric_snapshot import PermissionMetricSnapshot # noqa: F401
|
||||
from app.models.permission_template import PermissionTemplate, PermissionTemplateVersion # noqa: F401
|
||||
from app.models.security_access_log import SecurityAccessLog # noqa: F401
|
||||
from app.models.email_settings import EmailVerificationCode, SystemEmailSettings # noqa: F401
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import DateTime, Enum, ForeignKey, Integer, String, Text, UniqueConstraint, func
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.db.base_class import Base
|
||||
|
||||
|
||||
class SmtpSecurity(str, enum.Enum):
|
||||
NONE = "NONE"
|
||||
SSL = "SSL"
|
||||
STARTTLS = "STARTTLS"
|
||||
|
||||
|
||||
class EmailVerificationPurpose(str, enum.Enum):
|
||||
REGISTER = "REGISTER"
|
||||
PASSWORD_RESET = "PASSWORD_RESET"
|
||||
PASSWORD_RESET_LINK = "PASSWORD_RESET_LINK"
|
||||
|
||||
|
||||
class SystemEmailSettings(Base):
|
||||
__tablename__ = "system_email_settings"
|
||||
__table_args__ = (UniqueConstraint("register_domain", name="uq_system_email_settings_register_domain"),)
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
register_domain: Mapped[str] = mapped_column(String(255), nullable=False, default="huapont.cn")
|
||||
smtp_host: Mapped[str] = mapped_column(String(255), nullable=False, default="")
|
||||
smtp_port: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||
smtp_security: Mapped[SmtpSecurity] = mapped_column(
|
||||
Enum(SmtpSecurity, name="smtp_security"),
|
||||
nullable=False,
|
||||
default=SmtpSecurity.SSL,
|
||||
server_default=SmtpSecurity.SSL.value,
|
||||
)
|
||||
smtp_username: Mapped[str] = mapped_column(String(255), nullable=False, default="")
|
||||
smtp_password_encrypted: Mapped[Optional[str]] = mapped_column(Text, nullable=True)
|
||||
sender_email: Mapped[str] = mapped_column(String(255), nullable=False, default="")
|
||||
sender_name: Mapped[Optional[str]] = mapped_column(String(255), nullable=True)
|
||||
allowed_register_domain: Mapped[str] = mapped_column(String(255), nullable=False, default="huapont.cn")
|
||||
verification_code_ttl_minutes: Mapped[int] = mapped_column(Integer, nullable=False, default=10)
|
||||
send_cooldown_seconds: Mapped[int] = mapped_column(Integer, nullable=False, default=60)
|
||||
max_verify_attempts: Mapped[int] = mapped_column(Integer, nullable=False, default=5)
|
||||
updated_by: Mapped[Optional[uuid.UUID]] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now()
|
||||
)
|
||||
|
||||
|
||||
class EmailVerificationCode(Base):
|
||||
__tablename__ = "email_verification_codes"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
email: Mapped[str] = mapped_column(String(255), nullable=False, index=True)
|
||||
purpose: Mapped[EmailVerificationPurpose] = mapped_column(
|
||||
Enum(EmailVerificationPurpose, name="email_verification_purpose"),
|
||||
nullable=False,
|
||||
default=EmailVerificationPurpose.REGISTER,
|
||||
server_default=EmailVerificationPurpose.REGISTER.value,
|
||||
)
|
||||
code_hash: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
verified_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
attempt_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
|
||||
@@ -0,0 +1,128 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, EmailStr, Field, field_validator
|
||||
|
||||
SmtpSecurityValue = Literal["NONE", "SSL", "STARTTLS"]
|
||||
|
||||
|
||||
class EmailSettingsRead(BaseModel):
|
||||
register_domain: str = "huapont.cn"
|
||||
smtp_host: str = ""
|
||||
smtp_port: int = 465
|
||||
smtp_security: SmtpSecurityValue = "SSL"
|
||||
smtp_username: str = ""
|
||||
smtp_password_configured: bool = False
|
||||
sender_email: str = ""
|
||||
sender_name: Optional[str] = "CTMS 系统"
|
||||
allowed_register_domain: str = "huapont.cn,qq.com"
|
||||
verification_code_ttl_minutes: int = 10
|
||||
send_cooldown_seconds: int = 60
|
||||
max_verify_attempts: int = 5
|
||||
updated_at: Optional[datetime] = None
|
||||
|
||||
|
||||
class EmailSettingsUpdate(BaseModel):
|
||||
smtp_host: str = Field(default="", max_length=255)
|
||||
smtp_port: int = Field(ge=1, le=65535)
|
||||
smtp_security: SmtpSecurityValue = "SSL"
|
||||
smtp_username: str = Field(default="", max_length=255)
|
||||
smtp_password: Optional[str] = Field(default=None, max_length=500)
|
||||
sender_email: str = Field(default="", max_length=255)
|
||||
sender_name: Optional[str] = Field(default="CTMS 系统", max_length=255)
|
||||
allowed_register_domain: str = Field(default="", max_length=255)
|
||||
verification_code_ttl_minutes: int = Field(default=10, ge=1, le=60)
|
||||
send_cooldown_seconds: int = Field(default=60, ge=10, le=3600)
|
||||
max_verify_attempts: int = Field(default=5, ge=1, le=20)
|
||||
|
||||
@field_validator("allowed_register_domain")
|
||||
@classmethod
|
||||
def normalize_domain(cls, value: str) -> str:
|
||||
domains = []
|
||||
for item in value.replace(",", ",").split(","):
|
||||
domain = item.strip().lower().lstrip("@")
|
||||
if domain and domain not in domains:
|
||||
domains.append(domain)
|
||||
if not domains:
|
||||
raise ValueError("至少配置一个允许注册邮箱域名")
|
||||
return ",".join(domains)
|
||||
|
||||
|
||||
class EmailTestRequest(BaseModel):
|
||||
recipient_email: EmailStr
|
||||
|
||||
|
||||
class EmailDomainCreateRequest(BaseModel):
|
||||
register_domain: str = Field(min_length=1, max_length=255)
|
||||
|
||||
@field_validator("register_domain")
|
||||
@classmethod
|
||||
def normalize_domain(cls, value: str) -> str:
|
||||
return value.strip().lower().lstrip("@")
|
||||
|
||||
|
||||
class EmailSettingsListResponse(BaseModel):
|
||||
items: list[EmailSettingsRead]
|
||||
total: int
|
||||
|
||||
|
||||
class RegisterEmailCodeSendRequest(BaseModel):
|
||||
email: EmailStr
|
||||
|
||||
|
||||
class RegisterEmailCodeVerifyRequest(BaseModel):
|
||||
email: EmailStr
|
||||
code: str = Field(min_length=4, max_length=12)
|
||||
|
||||
|
||||
class PasswordResetLinkSendRequest(BaseModel):
|
||||
email: EmailStr
|
||||
|
||||
|
||||
class PasswordResetRequest(BaseModel):
|
||||
email: EmailStr
|
||||
code: str = Field(min_length=4, max_length=12)
|
||||
password: str = Field(min_length=8, max_length=72)
|
||||
|
||||
@field_validator("password")
|
||||
@classmethod
|
||||
def validate_password_strength(cls, value: str) -> str:
|
||||
has_letter = any(ch.isalpha() for ch in value)
|
||||
has_number = any(ch.isdigit() for ch in value)
|
||||
if not has_letter or not has_number:
|
||||
raise ValueError("密码需至少 8 位且包含字母和数字")
|
||||
return value
|
||||
|
||||
|
||||
class PasswordResetCodeVerifyRequest(BaseModel):
|
||||
email: EmailStr
|
||||
code: str = Field(min_length=4, max_length=12)
|
||||
|
||||
|
||||
class PasswordResetTokenRequest(BaseModel):
|
||||
token: str = Field(min_length=20, max_length=200)
|
||||
password: str = Field(min_length=8, max_length=72)
|
||||
|
||||
@field_validator("password")
|
||||
@classmethod
|
||||
def validate_password_strength(cls, value: str) -> str:
|
||||
has_letter = any(ch.isalpha() for ch in value)
|
||||
has_number = any(ch.isdigit() for ch in value)
|
||||
if not has_letter or not has_number:
|
||||
raise ValueError("密码需至少 8 位且包含字母和数字")
|
||||
return value
|
||||
|
||||
|
||||
class EmailCodeResponse(BaseModel):
|
||||
message: str
|
||||
|
||||
|
||||
class EmailCodeVerifyResponse(BaseModel):
|
||||
verified: bool
|
||||
|
||||
|
||||
class PasswordResetCodeVerifyResponse(BaseModel):
|
||||
verified: bool
|
||||
reset_token: str
|
||||
@@ -0,0 +1,563 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import random
|
||||
import secrets
|
||||
import smtplib
|
||||
import ssl
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from email.message import EmailMessage
|
||||
from email.utils import formataddr
|
||||
|
||||
from cryptography.fernet import Fernet, InvalidToken
|
||||
from fastapi import HTTPException, status
|
||||
from sqlalchemy import delete, desc, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.config import settings
|
||||
from app.core.security import hash_password, verify_password
|
||||
from app.crud import user as user_crud
|
||||
from app.models.email_settings import EmailVerificationCode, EmailVerificationPurpose, SmtpSecurity, SystemEmailSettings
|
||||
from app.schemas.email_settings import EmailSettingsRead, EmailSettingsUpdate
|
||||
|
||||
|
||||
def _normalize_email(email: str) -> str:
|
||||
return email.strip().lower()
|
||||
|
||||
|
||||
def _parse_allowed_domains(value: str) -> list[str]:
|
||||
domains: list[str] = []
|
||||
for item in (value or "").replace(",", ",").split(","):
|
||||
domain = item.strip().lower().lstrip("@")
|
||||
if domain and domain not in domains:
|
||||
domains.append(domain)
|
||||
return domains or ["huapont.cn", "qq.com"]
|
||||
|
||||
|
||||
def _normalize_domain(value: str) -> str:
|
||||
return value.strip().lower().lstrip("@")
|
||||
|
||||
|
||||
def _domain_from_email(email: str) -> str:
|
||||
parts = _normalize_email(email).rsplit("@", 1)
|
||||
return parts[1] if len(parts) == 2 else ""
|
||||
|
||||
|
||||
def _as_aware(value: datetime) -> datetime:
|
||||
return value if value.tzinfo else value.replace(tzinfo=timezone.utc)
|
||||
|
||||
|
||||
def _get_fernet() -> Fernet:
|
||||
source = settings.SETTINGS_ENCRYPTION_KEY or settings.JWT_SECRET_KEY
|
||||
try:
|
||||
return Fernet(source.encode("utf-8"))
|
||||
except Exception:
|
||||
key = base64.urlsafe_b64encode(hashlib.sha256(source.encode("utf-8")).digest())
|
||||
return Fernet(key)
|
||||
|
||||
|
||||
def encrypt_secret(value: str) -> str:
|
||||
return _get_fernet().encrypt(value.encode("utf-8")).decode("utf-8")
|
||||
|
||||
|
||||
def decrypt_secret(value: str | None) -> str | None:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
return _get_fernet().decrypt(value.encode("utf-8")).decode("utf-8")
|
||||
except InvalidToken as exc:
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="邮件授权码解密失败") from exc
|
||||
|
||||
|
||||
async def list_email_settings(db: AsyncSession) -> list[SystemEmailSettings]:
|
||||
result = await db.execute(select(SystemEmailSettings).order_by(SystemEmailSettings.register_domain.asc()))
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def get_email_settings(db: AsyncSession, register_domain: str | None = None) -> SystemEmailSettings | None:
|
||||
if register_domain:
|
||||
result = await db.execute(
|
||||
select(SystemEmailSettings).where(SystemEmailSettings.register_domain == _normalize_domain(register_domain))
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
result = await db.execute(select(SystemEmailSettings).order_by(SystemEmailSettings.created_at.asc()))
|
||||
return result.scalars().first()
|
||||
|
||||
|
||||
async def get_email_settings_for_email(db: AsyncSession, email: str) -> SystemEmailSettings | None:
|
||||
domain = _domain_from_email(email)
|
||||
return await get_email_settings(db, domain) if domain else None
|
||||
|
||||
|
||||
def to_email_settings_read(settings_row: SystemEmailSettings | None) -> EmailSettingsRead:
|
||||
if not settings_row:
|
||||
return EmailSettingsRead()
|
||||
return EmailSettingsRead(
|
||||
register_domain=settings_row.register_domain,
|
||||
smtp_host=settings_row.smtp_host,
|
||||
smtp_port=settings_row.smtp_port,
|
||||
smtp_security=settings_row.smtp_security.value,
|
||||
smtp_username=settings_row.smtp_username,
|
||||
smtp_password_configured=bool(settings_row.smtp_password_encrypted),
|
||||
sender_email=settings_row.sender_email,
|
||||
sender_name=settings_row.sender_name,
|
||||
allowed_register_domain=settings_row.allowed_register_domain,
|
||||
verification_code_ttl_minutes=settings_row.verification_code_ttl_minutes,
|
||||
send_cooldown_seconds=settings_row.send_cooldown_seconds,
|
||||
max_verify_attempts=settings_row.max_verify_attempts,
|
||||
updated_at=settings_row.updated_at,
|
||||
)
|
||||
|
||||
|
||||
async def upsert_email_settings(
|
||||
db: AsyncSession,
|
||||
register_domain: str,
|
||||
payload: EmailSettingsUpdate,
|
||||
*,
|
||||
updated_by: uuid.UUID,
|
||||
) -> SystemEmailSettings:
|
||||
register_domain = _normalize_domain(register_domain)
|
||||
settings_row = await get_email_settings(db, register_domain)
|
||||
if not settings_row:
|
||||
settings_row = SystemEmailSettings(
|
||||
register_domain=register_domain,
|
||||
smtp_host=payload.smtp_host,
|
||||
smtp_port=payload.smtp_port,
|
||||
smtp_security=SmtpSecurity(payload.smtp_security),
|
||||
smtp_username=payload.smtp_username,
|
||||
smtp_password_encrypted=encrypt_secret(payload.smtp_password) if payload.smtp_password else None,
|
||||
sender_email=str(payload.sender_email),
|
||||
sender_name=payload.sender_name,
|
||||
allowed_register_domain=register_domain,
|
||||
verification_code_ttl_minutes=payload.verification_code_ttl_minutes,
|
||||
send_cooldown_seconds=payload.send_cooldown_seconds,
|
||||
max_verify_attempts=payload.max_verify_attempts,
|
||||
updated_by=updated_by,
|
||||
)
|
||||
db.add(settings_row)
|
||||
else:
|
||||
settings_row.smtp_host = payload.smtp_host
|
||||
settings_row.smtp_port = payload.smtp_port
|
||||
settings_row.smtp_security = SmtpSecurity(payload.smtp_security)
|
||||
settings_row.smtp_username = payload.smtp_username
|
||||
if payload.smtp_password:
|
||||
settings_row.smtp_password_encrypted = encrypt_secret(payload.smtp_password)
|
||||
settings_row.sender_email = str(payload.sender_email)
|
||||
settings_row.sender_name = payload.sender_name
|
||||
settings_row.allowed_register_domain = register_domain
|
||||
settings_row.verification_code_ttl_minutes = payload.verification_code_ttl_minutes
|
||||
settings_row.send_cooldown_seconds = payload.send_cooldown_seconds
|
||||
settings_row.max_verify_attempts = payload.max_verify_attempts
|
||||
settings_row.updated_by = updated_by
|
||||
await db.commit()
|
||||
await db.refresh(settings_row)
|
||||
return settings_row
|
||||
|
||||
|
||||
async def create_email_domain(db: AsyncSession, register_domain: str, *, updated_by: uuid.UUID) -> SystemEmailSettings:
|
||||
register_domain = _normalize_domain(register_domain)
|
||||
existing = await get_email_settings(db, register_domain)
|
||||
if existing:
|
||||
return existing
|
||||
settings_row = SystemEmailSettings(
|
||||
register_domain=register_domain,
|
||||
smtp_host="",
|
||||
smtp_port=465,
|
||||
smtp_security=SmtpSecurity.SSL,
|
||||
smtp_username="",
|
||||
smtp_password_encrypted=None,
|
||||
sender_email="",
|
||||
sender_name="CTMS 系统",
|
||||
allowed_register_domain=register_domain,
|
||||
verification_code_ttl_minutes=10,
|
||||
send_cooldown_seconds=60,
|
||||
max_verify_attempts=5,
|
||||
updated_by=updated_by,
|
||||
)
|
||||
db.add(settings_row)
|
||||
await db.commit()
|
||||
await db.refresh(settings_row)
|
||||
return settings_row
|
||||
|
||||
|
||||
async def delete_email_domain(db: AsyncSession, register_domain: str) -> None:
|
||||
await db.execute(
|
||||
delete(SystemEmailSettings).where(SystemEmailSettings.register_domain == _normalize_domain(register_domain))
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
|
||||
def ensure_email_settings_ready(settings_row: SystemEmailSettings | None) -> SystemEmailSettings:
|
||||
if not settings_row:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="邮件服务尚未配置")
|
||||
if not settings_row.smtp_host or not settings_row.smtp_username or not settings_row.sender_email:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="邮件服务配置不完整")
|
||||
if not settings_row.smtp_password_encrypted:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="SMTP 授权码尚未配置")
|
||||
return settings_row
|
||||
|
||||
|
||||
def ensure_allowed_register_email(email: str, settings_row: SystemEmailSettings) -> None:
|
||||
domain = settings_row.register_domain
|
||||
if not _normalize_email(email).endswith(f"@{domain}"):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=f"仅允许 @{domain} 邮箱注册")
|
||||
|
||||
|
||||
def _send_email_sync(
|
||||
settings_row: SystemEmailSettings,
|
||||
*,
|
||||
to_email: str,
|
||||
subject: str,
|
||||
body: str,
|
||||
) -> None:
|
||||
password = decrypt_secret(settings_row.smtp_password_encrypted)
|
||||
if not password:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="SMTP 授权码尚未配置")
|
||||
|
||||
message = EmailMessage()
|
||||
message["Subject"] = subject
|
||||
message["From"] = formataddr((settings_row.sender_name or "", settings_row.sender_email))
|
||||
message["To"] = to_email
|
||||
message.set_content(body)
|
||||
|
||||
if settings_row.smtp_security == SmtpSecurity.SSL:
|
||||
with smtplib.SMTP_SSL(settings_row.smtp_host, settings_row.smtp_port, timeout=15) as smtp:
|
||||
smtp.login(settings_row.smtp_username, password)
|
||||
smtp.send_message(message)
|
||||
else:
|
||||
with smtplib.SMTP(settings_row.smtp_host, settings_row.smtp_port, timeout=15) as smtp:
|
||||
if settings_row.smtp_security == SmtpSecurity.STARTTLS:
|
||||
smtp.starttls()
|
||||
smtp.login(settings_row.smtp_username, password)
|
||||
smtp.send_message(message)
|
||||
|
||||
|
||||
async def send_email(
|
||||
settings_row: SystemEmailSettings,
|
||||
*,
|
||||
to_email: str,
|
||||
subject: str,
|
||||
body: str,
|
||||
) -> None:
|
||||
try:
|
||||
await asyncio.to_thread(_send_email_sync, settings_row, to_email=to_email, subject=subject, body=body)
|
||||
except HTTPException:
|
||||
raise
|
||||
except ssl.SSLError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=(
|
||||
"邮件发送失败:SMTP 安全协议或端口不匹配,"
|
||||
"请确认当前后缀的 SMTP 地址、端口和安全协议。"
|
||||
"SSL 通常使用 465,STARTTLS 通常使用 587。"
|
||||
),
|
||||
) from exc
|
||||
except smtplib.SMTPAuthenticationError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="邮件发送失败:SMTP 账号或授权码不正确",
|
||||
) from exc
|
||||
except (smtplib.SMTPRecipientsRefused, smtplib.SMTPDataError) as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="邮箱不存在或无法接收邮件,请检查邮箱地址",
|
||||
) from exc
|
||||
except (smtplib.SMTPConnectError, TimeoutError, OSError) as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="邮件发送失败:无法连接 SMTP 服务器,请检查 SMTP 地址、端口和网络连通性",
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=f"邮件发送失败:{exc}") from exc
|
||||
|
||||
|
||||
async def send_test_email(db: AsyncSession, register_domain: str, recipient_email: str) -> None:
|
||||
settings_row = ensure_email_settings_ready(await get_email_settings(db, register_domain))
|
||||
await send_email(
|
||||
settings_row,
|
||||
to_email=recipient_email,
|
||||
subject="CTMS 邮件服务测试",
|
||||
body="这是一封来自 CTMS 系统的邮件服务测试邮件。收到此邮件说明 SMTP 配置可用。",
|
||||
)
|
||||
|
||||
|
||||
async def send_register_code(db: AsyncSession, email: str) -> None:
|
||||
email = _normalize_email(email)
|
||||
settings_row = ensure_email_settings_ready(await get_email_settings_for_email(db, email))
|
||||
ensure_allowed_register_email(email, settings_row)
|
||||
if await user_crud.get_by_email(db, email):
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="邮箱已注册")
|
||||
|
||||
await _send_verification_code(
|
||||
db,
|
||||
settings_row=settings_row,
|
||||
email=email,
|
||||
purpose=EmailVerificationPurpose.REGISTER,
|
||||
subject="CTMS 注册邮箱验证码",
|
||||
body_factory=lambda code: (
|
||||
f"您的 CTMS 注册验证码是:{code}\n\n"
|
||||
f"验证码 {settings_row.verification_code_ttl_minutes} 分钟内有效。"
|
||||
"如非本人操作,请忽略此邮件。"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def send_password_reset_code(db: AsyncSession, email: str) -> None:
|
||||
email = _normalize_email(email)
|
||||
db_user = await user_crud.get_by_email(db, email)
|
||||
if not db_user:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="该邮箱未注册")
|
||||
settings_row = ensure_email_settings_ready(await get_email_settings_for_email(db, email))
|
||||
await _send_verification_code(
|
||||
db,
|
||||
settings_row=settings_row,
|
||||
email=email,
|
||||
purpose=EmailVerificationPurpose.PASSWORD_RESET,
|
||||
subject="CTMS 密码重置验证码",
|
||||
body_factory=lambda code: (
|
||||
f"您的 CTMS 密码重置验证码是:{code}\n\n"
|
||||
f"验证码 {settings_row.verification_code_ttl_minutes} 分钟内有效。"
|
||||
"如非本人操作,请尽快联系系统管理员。"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def send_password_reset_link(db: AsyncSession, email: str, *, frontend_origin: str) -> None:
|
||||
email = _normalize_email(email)
|
||||
db_user = await user_crud.get_by_email(db, email)
|
||||
if not db_user:
|
||||
return
|
||||
settings_row = ensure_email_settings_ready(await get_email_settings_for_email(db, email))
|
||||
now = datetime.now(timezone.utc)
|
||||
latest_result = await db.execute(
|
||||
select(EmailVerificationCode)
|
||||
.where(
|
||||
EmailVerificationCode.email == email,
|
||||
EmailVerificationCode.purpose == EmailVerificationPurpose.PASSWORD_RESET_LINK,
|
||||
)
|
||||
.order_by(desc(EmailVerificationCode.created_at))
|
||||
.limit(1)
|
||||
)
|
||||
latest = latest_result.scalar_one_or_none()
|
||||
if latest and latest.created_at and _as_aware(latest.created_at) + timedelta(seconds=settings_row.send_cooldown_seconds) > now:
|
||||
raise HTTPException(status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail="重置邮件发送过于频繁,请稍后再试")
|
||||
|
||||
token = secrets.token_urlsafe(32)
|
||||
token_hash = _hash_reset_token(token)
|
||||
record = EmailVerificationCode(
|
||||
email=email,
|
||||
purpose=EmailVerificationPurpose.PASSWORD_RESET_LINK,
|
||||
code_hash=token_hash,
|
||||
expires_at=now + timedelta(minutes=settings_row.verification_code_ttl_minutes),
|
||||
)
|
||||
db.add(record)
|
||||
await db.commit()
|
||||
|
||||
origin = frontend_origin.rstrip("/") if frontend_origin else ""
|
||||
reset_link = f"{origin}/forgot-password?token={token}" if origin else f"/forgot-password?token={token}"
|
||||
await send_email(
|
||||
settings_row,
|
||||
to_email=email,
|
||||
subject="CTMS 密码重置链接",
|
||||
body=(
|
||||
"请点击以下链接重置您的 CTMS 登录密码:\n\n"
|
||||
f"{reset_link}\n\n"
|
||||
f"链接 {settings_row.verification_code_ttl_minutes} 分钟内有效,且只能使用一次。"
|
||||
"如非本人操作,请忽略此邮件并尽快联系系统管理员。"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _hash_reset_token(token: str) -> str:
|
||||
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
async def _send_verification_code(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
settings_row: SystemEmailSettings,
|
||||
email: str,
|
||||
purpose: EmailVerificationPurpose,
|
||||
subject: str,
|
||||
body_factory,
|
||||
) -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
latest_result = await db.execute(
|
||||
select(EmailVerificationCode)
|
||||
.where(
|
||||
EmailVerificationCode.email == email,
|
||||
EmailVerificationCode.purpose == purpose,
|
||||
)
|
||||
.order_by(desc(EmailVerificationCode.created_at))
|
||||
.limit(1)
|
||||
)
|
||||
latest = latest_result.scalar_one_or_none()
|
||||
if latest and latest.created_at and _as_aware(latest.created_at) + timedelta(seconds=settings_row.send_cooldown_seconds) > now:
|
||||
raise HTTPException(status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail="验证码发送过于频繁,请稍后再试")
|
||||
|
||||
code = f"{random.SystemRandom().randint(0, 999999):06d}"
|
||||
record = EmailVerificationCode(
|
||||
email=email,
|
||||
purpose=purpose,
|
||||
code_hash=hash_password(code),
|
||||
expires_at=now + timedelta(minutes=settings_row.verification_code_ttl_minutes),
|
||||
)
|
||||
db.add(record)
|
||||
await db.commit()
|
||||
|
||||
await send_email(
|
||||
settings_row,
|
||||
to_email=email,
|
||||
subject=subject,
|
||||
body=body_factory(code),
|
||||
)
|
||||
|
||||
|
||||
async def verify_register_code(db: AsyncSession, email: str, code: str) -> bool:
|
||||
email = _normalize_email(email)
|
||||
settings_row = ensure_email_settings_ready(await get_email_settings_for_email(db, email))
|
||||
ensure_allowed_register_email(email, settings_row)
|
||||
await _verify_email_code(
|
||||
db,
|
||||
settings_row=settings_row,
|
||||
email=email,
|
||||
code=code,
|
||||
purpose=EmailVerificationPurpose.REGISTER,
|
||||
mark_verified=True,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
async def verify_password_reset_code(db: AsyncSession, email: str, code: str) -> str:
|
||||
email = _normalize_email(email)
|
||||
if not await user_crud.get_by_email(db, email):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="验证码错误或已过期")
|
||||
settings_row = ensure_email_settings_ready(await get_email_settings_for_email(db, email))
|
||||
await _verify_email_code(
|
||||
db,
|
||||
settings_row=settings_row,
|
||||
email=email,
|
||||
code=code,
|
||||
purpose=EmailVerificationPurpose.PASSWORD_RESET,
|
||||
mark_verified=True,
|
||||
commit=False,
|
||||
)
|
||||
now = datetime.now(timezone.utc)
|
||||
token = secrets.token_urlsafe(32)
|
||||
db.add(
|
||||
EmailVerificationCode(
|
||||
email=email,
|
||||
purpose=EmailVerificationPurpose.PASSWORD_RESET_LINK,
|
||||
code_hash=_hash_reset_token(token),
|
||||
expires_at=now + timedelta(minutes=settings_row.verification_code_ttl_minutes),
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
return token
|
||||
|
||||
|
||||
async def reset_password_with_code(db: AsyncSession, email: str, code: str, password: str) -> None:
|
||||
email = _normalize_email(email)
|
||||
db_user = await user_crud.get_by_email(db, email)
|
||||
if not db_user:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="验证码错误或已过期")
|
||||
settings_row = ensure_email_settings_ready(await get_email_settings_for_email(db, email))
|
||||
await _verify_email_code(
|
||||
db,
|
||||
settings_row=settings_row,
|
||||
email=email,
|
||||
code=code,
|
||||
purpose=EmailVerificationPurpose.PASSWORD_RESET,
|
||||
mark_verified=True,
|
||||
commit=False,
|
||||
)
|
||||
db_user.password_hash = hash_password(password)
|
||||
await db.commit()
|
||||
|
||||
|
||||
async def reset_password_with_token(db: AsyncSession, token: str, password: str) -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
token_hash = _hash_reset_token(token)
|
||||
result = await db.execute(
|
||||
select(EmailVerificationCode)
|
||||
.where(
|
||||
EmailVerificationCode.code_hash == token_hash,
|
||||
EmailVerificationCode.purpose == EmailVerificationPurpose.PASSWORD_RESET_LINK,
|
||||
EmailVerificationCode.verified_at.is_(None),
|
||||
)
|
||||
.order_by(desc(EmailVerificationCode.created_at))
|
||||
.limit(1)
|
||||
.with_for_update()
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if not record or record.verified_at is not None or _as_aware(record.expires_at) < now:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="重置链接无效或已过期")
|
||||
|
||||
db_user = await user_crud.get_by_email(db, record.email)
|
||||
if not db_user:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="重置链接无效或已过期")
|
||||
|
||||
record.verified_at = now
|
||||
db_user.password_hash = hash_password(password)
|
||||
await db.commit()
|
||||
|
||||
|
||||
async def _verify_email_code(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
settings_row: SystemEmailSettings,
|
||||
email: str,
|
||||
code: str,
|
||||
purpose: EmailVerificationPurpose,
|
||||
mark_verified: bool,
|
||||
commit: bool = True,
|
||||
) -> EmailVerificationCode:
|
||||
now = datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(EmailVerificationCode)
|
||||
.where(
|
||||
EmailVerificationCode.email == email,
|
||||
EmailVerificationCode.purpose == purpose,
|
||||
)
|
||||
.order_by(desc(EmailVerificationCode.created_at))
|
||||
.limit(1)
|
||||
.with_for_update()
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if (
|
||||
not record
|
||||
or record.verified_at is not None
|
||||
or _as_aware(record.expires_at) < now
|
||||
or record.attempt_count >= settings_row.max_verify_attempts
|
||||
):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="验证码错误或已过期")
|
||||
record.attempt_count += 1
|
||||
if not verify_password(code, record.code_hash):
|
||||
await db.commit()
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="验证码错误或已过期")
|
||||
if mark_verified:
|
||||
record.verified_at = now
|
||||
if commit:
|
||||
await db.commit()
|
||||
return record
|
||||
|
||||
|
||||
async def ensure_register_email_verified(db: AsyncSession, email: str) -> None:
|
||||
email = _normalize_email(email)
|
||||
now = datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(EmailVerificationCode)
|
||||
.where(
|
||||
EmailVerificationCode.email == email,
|
||||
EmailVerificationCode.purpose == EmailVerificationPurpose.REGISTER,
|
||||
EmailVerificationCode.verified_at.is_not(None),
|
||||
EmailVerificationCode.expires_at >= now,
|
||||
)
|
||||
.order_by(desc(EmailVerificationCode.verified_at))
|
||||
.limit(1)
|
||||
)
|
||||
if not result.scalar_one_or_none():
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="请先完成邮箱验证码校验")
|
||||
@@ -3,10 +3,12 @@ import pytest_asyncio
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import padding
|
||||
from fastapi import HTTPException
|
||||
from httpx import AsyncClient
|
||||
from sqlalchemy import UUID as SA_UUID, select
|
||||
from sqlalchemy.dialects.postgresql import UUID as PG_UUID
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||
from datetime import datetime, timedelta, timezone
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
@@ -14,9 +16,17 @@ import os
|
||||
from app.main import create_app
|
||||
from app.core.config import settings
|
||||
from app.core.deps import get_db_session
|
||||
from app.core.security import hash_password
|
||||
from app.core.security import hash_password, verify_password
|
||||
from app.crud import user as user_crud
|
||||
from app.db.base_class import Base
|
||||
from app.models.audit_log import AuditLog
|
||||
from app.models.email_settings import (
|
||||
EmailVerificationCode,
|
||||
EmailVerificationPurpose,
|
||||
SmtpSecurity,
|
||||
SystemEmailSettings,
|
||||
)
|
||||
from app.services import email_service
|
||||
from tests.conftest import GUID
|
||||
from app.models.study_member import StudyMember
|
||||
from app.models.user import User, UserStatus
|
||||
@@ -25,6 +35,41 @@ from app.schemas.user import UserRegisterRequest
|
||||
TEST_DATABASE_URL = "sqlite+aiosqlite:///:memory:"
|
||||
|
||||
|
||||
async def mark_register_email_verified(SessionLocal, email: str) -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
async with SessionLocal() as session:
|
||||
session.add(
|
||||
EmailVerificationCode(
|
||||
email=email,
|
||||
purpose=EmailVerificationPurpose.REGISTER,
|
||||
code_hash=hash_password("123456"),
|
||||
expires_at=now + timedelta(minutes=10),
|
||||
verified_at=now,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def add_email_settings(SessionLocal, domain: str) -> None:
|
||||
async with SessionLocal() as session:
|
||||
session.add(
|
||||
SystemEmailSettings(
|
||||
register_domain=domain,
|
||||
smtp_host="smtp.test",
|
||||
smtp_port=465,
|
||||
smtp_security=SmtpSecurity.SSL,
|
||||
smtp_username="mailer",
|
||||
smtp_password_encrypted="encrypted",
|
||||
sender_email=f"mailer@{domain}",
|
||||
allowed_register_domain=domain,
|
||||
verification_code_ttl_minutes=10,
|
||||
send_cooldown_seconds=60,
|
||||
max_verify_attempts=5,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def encrypted_auth_payload(client: AsyncClient, email: str, password: str) -> dict:
|
||||
key_resp = await client.get("/api/v1/auth/login-key")
|
||||
assert key_resp.status_code == 200
|
||||
@@ -114,7 +159,7 @@ async def client_and_db():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_creates_pending_user(client_and_db):
|
||||
async def test_register_creates_active_user(client_and_db):
|
||||
client, SessionLocal = client_and_db
|
||||
payload = {
|
||||
"email": "newuser@test.com",
|
||||
@@ -122,12 +167,13 @@ async def test_register_creates_pending_user(client_and_db):
|
||||
"full_name": "New User",
|
||||
"clinical_department": "Clinical",
|
||||
}
|
||||
await mark_register_email_verified(SessionLocal, payload["email"])
|
||||
resp = await client.post("/api/v1/auth/register", json=payload)
|
||||
assert resp.status_code == 201
|
||||
async with SessionLocal() as session:
|
||||
user = await user_crud.get_by_email(session, payload["email"])
|
||||
assert user is not None
|
||||
assert user.status == UserStatus.PENDING
|
||||
assert user.status == UserStatus.ACTIVE
|
||||
assert user.is_admin is False
|
||||
member_rows = (
|
||||
await session.execute(select(StudyMember).where(StudyMember.user_id == user.id))
|
||||
@@ -135,6 +181,201 @@ async def test_register_creates_pending_user(client_and_db):
|
||||
assert member_rows == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_rejects_duplicate_email(client_and_db):
|
||||
client, SessionLocal = client_and_db
|
||||
payload = {
|
||||
"email": "duplicate@test.com",
|
||||
"password": "Password123",
|
||||
"full_name": "Duplicate User",
|
||||
"clinical_department": "Clinical",
|
||||
}
|
||||
await mark_register_email_verified(SessionLocal, payload["email"])
|
||||
first_resp = await client.post("/api/v1/auth/register", json=payload)
|
||||
assert first_resp.status_code == 201
|
||||
|
||||
second_resp = await client.post("/api/v1/auth/register", json=payload)
|
||||
assert second_resp.status_code == 409
|
||||
assert second_resp.json()["detail"] == "邮箱已注册"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_requires_verified_email_code(client_and_db):
|
||||
client, _ = client_and_db
|
||||
payload = {
|
||||
"email": "unverified@test.com",
|
||||
"password": "Password123",
|
||||
"full_name": "Unverified User",
|
||||
"clinical_department": "Clinical",
|
||||
}
|
||||
resp = await client.post("/api/v1/auth/register", json=payload)
|
||||
assert resp.status_code == 400
|
||||
assert resp.json()["detail"] == "请先完成邮箱验证码校验"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_email_availability(client_and_db):
|
||||
client, SessionLocal = client_and_db
|
||||
available_resp = await client.get(
|
||||
"/api/v1/auth/register/email-availability",
|
||||
params={"email": "available@test.com"},
|
||||
)
|
||||
assert available_resp.status_code == 200
|
||||
assert available_resp.json() == {"available": True}
|
||||
|
||||
payload = {
|
||||
"email": "used@test.com",
|
||||
"password": "Password123",
|
||||
"full_name": "Used User",
|
||||
"clinical_department": "Clinical",
|
||||
}
|
||||
await mark_register_email_verified(SessionLocal, payload["email"])
|
||||
await client.post("/api/v1/auth/register", json=payload)
|
||||
used_resp = await client.get(
|
||||
"/api/v1/auth/register/email-availability",
|
||||
params={"email": payload["email"]},
|
||||
)
|
||||
assert used_resp.status_code == 200
|
||||
assert used_resp.json() == {"available": False}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_reset_code_is_one_time_and_ignores_link_records(client_and_db):
|
||||
_, SessionLocal = client_and_db
|
||||
email = "reset-once@test.com"
|
||||
now = datetime.now(timezone.utc)
|
||||
await add_email_settings(SessionLocal, "test.com")
|
||||
async with SessionLocal() as session:
|
||||
user = User(
|
||||
email=email,
|
||||
password_hash=hash_password("OldPassword123"),
|
||||
full_name="Reset Once",
|
||||
clinical_department="Clinical",
|
||||
status=UserStatus.ACTIVE,
|
||||
)
|
||||
session.add(user)
|
||||
session.add_all(
|
||||
[
|
||||
EmailVerificationCode(
|
||||
email=email,
|
||||
purpose=EmailVerificationPurpose.PASSWORD_RESET,
|
||||
code_hash=hash_password("123456"),
|
||||
expires_at=now + timedelta(minutes=10),
|
||||
created_at=now,
|
||||
),
|
||||
EmailVerificationCode(
|
||||
email=email,
|
||||
purpose=EmailVerificationPurpose.PASSWORD_RESET_LINK,
|
||||
code_hash="a" * 64,
|
||||
expires_at=now + timedelta(minutes=10),
|
||||
created_at=now + timedelta(seconds=1),
|
||||
),
|
||||
]
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
await email_service.reset_password_with_code(session, email, "123456", "NewPassword123")
|
||||
await session.refresh(user)
|
||||
assert verify_password("NewPassword123", user.password_hash)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await email_service.reset_password_with_code(session, email, "123456", "OtherPassword123")
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_reset_code_verification_issues_one_time_reset_token(client_and_db):
|
||||
client, SessionLocal = client_and_db
|
||||
email = "verified-reset@test.com"
|
||||
now = datetime.now(timezone.utc)
|
||||
await add_email_settings(SessionLocal, "test.com")
|
||||
async with SessionLocal() as session:
|
||||
session.add(
|
||||
User(
|
||||
email=email,
|
||||
password_hash=hash_password("OldPassword123"),
|
||||
full_name="Verified Reset",
|
||||
clinical_department="Clinical",
|
||||
status=UserStatus.ACTIVE,
|
||||
)
|
||||
)
|
||||
session.add(
|
||||
EmailVerificationCode(
|
||||
email=email,
|
||||
purpose=EmailVerificationPurpose.PASSWORD_RESET,
|
||||
code_hash=hash_password("654321"),
|
||||
expires_at=now + timedelta(minutes=10),
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
verify_response = await client.post(
|
||||
"/api/v1/auth/password-reset/email-code/verify",
|
||||
json={"email": email, "code": "654321"},
|
||||
)
|
||||
assert verify_response.status_code == 200
|
||||
assert verify_response.json()["verified"] is True
|
||||
reset_token = verify_response.json()["reset_token"]
|
||||
|
||||
reused_code_response = await client.post(
|
||||
"/api/v1/auth/password-reset/email-code/verify",
|
||||
json={"email": email, "code": "654321"},
|
||||
)
|
||||
assert reused_code_response.status_code == 400
|
||||
|
||||
reset_response = await client.post(
|
||||
"/api/v1/auth/password-reset-link",
|
||||
json={"token": reset_token, "password": "NewPassword123"},
|
||||
)
|
||||
assert reset_response.status_code == 200
|
||||
reused_token_response = await client.post(
|
||||
"/api/v1/auth/password-reset-link",
|
||||
json={"token": reset_token, "password": "OtherPassword123"},
|
||||
)
|
||||
assert reused_token_response.status_code == 400
|
||||
|
||||
async with SessionLocal() as session:
|
||||
user = await user_crud.get_by_email(session, email)
|
||||
assert user is not None
|
||||
assert verify_password("NewPassword123", user.password_hash)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_reset_code_send_reports_missing_account(client_and_db):
|
||||
client, _ = client_and_db
|
||||
response = await client.post(
|
||||
"/api/v1/auth/password-reset/email-code/send",
|
||||
json={"email": "missing@test.com"},
|
||||
)
|
||||
assert response.status_code == 404
|
||||
assert response.json()["detail"] == "该邮箱未注册"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_reset_link_uses_configured_frontend_url(client_and_db, monkeypatch):
|
||||
client, _ = client_and_db
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
async def fake_send_password_reset_link(db, email: str, *, frontend_origin: str) -> None:
|
||||
captured["email"] = email
|
||||
captured["frontend_origin"] = frontend_origin
|
||||
|
||||
monkeypatch.setattr(email_service, "send_password_reset_link", fake_send_password_reset_link)
|
||||
monkeypatch.setattr(settings, "FRONTEND_PUBLIC_URL", "https://ctms.example.com")
|
||||
|
||||
response = await client.post(
|
||||
"/api/v1/auth/password-reset-link/send",
|
||||
json={"email": "victim@test.com"},
|
||||
headers={"Origin": "https://attacker.example"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured == {
|
||||
"email": "victim@test.com",
|
||||
"frontend_origin": "https://ctms.example.com",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_root_returns_service_metadata(client_and_db):
|
||||
client, _ = client_and_db
|
||||
@@ -150,44 +391,41 @@ async def test_root_returns_service_metadata(client_and_db):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_login_blocked_before_approval(client_and_db):
|
||||
async def test_registered_user_can_login_after_email_verification(client_and_db):
|
||||
client, SessionLocal = client_and_db
|
||||
payload = {
|
||||
"email": "pending@test.com",
|
||||
"email": "active-register@test.com",
|
||||
"password": "Password123",
|
||||
"full_name": "Pending User",
|
||||
"full_name": "Active Registered User",
|
||||
"clinical_department": "Safety",
|
||||
}
|
||||
await mark_register_email_verified(SessionLocal, payload["email"])
|
||||
await client.post("/api/v1/auth/register", json=payload)
|
||||
resp = await encrypted_login(client, payload["email"], payload["password"])
|
||||
assert resp.status_code == 401
|
||||
assert "账号未审核" in resp.json().get("detail", "")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["access_token"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_can_approve_user(client_and_db):
|
||||
async def test_admin_created_user_is_active_by_default(client_and_db):
|
||||
client, SessionLocal = client_and_db
|
||||
payload = {
|
||||
"email": "approve@test.com",
|
||||
"password": "Password123",
|
||||
"full_name": "Approve Target",
|
||||
"clinical_department": "Supply",
|
||||
}
|
||||
await client.post("/api/v1/auth/register", json=payload)
|
||||
async with SessionLocal() as session:
|
||||
user = await user_crud.get_by_email(session, payload["email"])
|
||||
user_id = user.id
|
||||
|
||||
admin_login = await encrypted_login(client, "admin@test.com", "admin123")
|
||||
token = admin_login.json()["access_token"]
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
payload = {
|
||||
"email": "admin-created@test.com",
|
||||
"password": "Password123",
|
||||
"full_name": "Admin Created",
|
||||
"clinical_department": "Supply",
|
||||
}
|
||||
|
||||
resp = await client.post(f"/api/v1/admin/users/{user_id}/approve", json={"action": "approve"}, headers=headers)
|
||||
assert resp.status_code == 200
|
||||
resp = await client.post("/api/v1/users/", json=payload, headers=headers)
|
||||
|
||||
assert resp.status_code == 201
|
||||
assert resp.json()["status"] == "ACTIVE"
|
||||
async with SessionLocal() as session:
|
||||
refreshed = await user_crud.get_by_email(session, payload["email"])
|
||||
assert refreshed.status == UserStatus.ACTIVE
|
||||
assert refreshed.approved_by is not None
|
||||
user = await user_crud.get_by_email(session, payload["email"])
|
||||
assert user.status == UserStatus.ACTIVE
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -248,6 +486,43 @@ async def test_admin_users_list_filters_by_keyword_and_status(client_and_db):
|
||||
assert combined_data["items"][0]["email"] == "pending-filter@test.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_delete_user_with_audit_history_returns_400(client_and_db):
|
||||
client, SessionLocal = client_and_db
|
||||
async with SessionLocal() as session:
|
||||
user = User(
|
||||
email="audited-delete@test.com",
|
||||
password_hash=hash_password("Password123"),
|
||||
full_name="Audited Delete",
|
||||
clinical_department="Medical",
|
||||
status=UserStatus.ACTIVE,
|
||||
)
|
||||
session.add(user)
|
||||
await session.flush()
|
||||
session.add(
|
||||
AuditLog(
|
||||
entity_type="user",
|
||||
entity_id=user.id,
|
||||
action="LOGIN",
|
||||
operator_id=user.id,
|
||||
operator_role="USER",
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
user_id = user.id
|
||||
|
||||
admin_login = await encrypted_login(client, "admin@test.com", "admin123")
|
||||
token = admin_login.json()["access_token"]
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
|
||||
resp = await client.delete(f"/api/v1/users/{user_id}", headers=headers)
|
||||
|
||||
assert resp.status_code == 400
|
||||
assert resp.json()["detail"] == "该账号已有审计或权限访问记录,请停用账号以保留历史追溯"
|
||||
async with SessionLocal() as session:
|
||||
assert await user_crud.get_by_id(session, user_id) is not None
|
||||
|
||||
|
||||
def test_register_request_does_not_expose_role_input():
|
||||
assert "role" not in UserRegisterRequest.model_fields
|
||||
|
||||
@@ -313,17 +588,21 @@ async def test_dev_login_rejects_inactive_users(client_and_db):
|
||||
client, SessionLocal = client_and_db
|
||||
original_env = settings.ENV
|
||||
settings.ENV = "development"
|
||||
payload = {
|
||||
"email": "pending-dev-login@test.com",
|
||||
"password": "Password123",
|
||||
"full_name": "Pending Dev Login",
|
||||
"clinical_department": "Clinical",
|
||||
}
|
||||
try:
|
||||
await client.post("/api/v1/auth/register", json=payload)
|
||||
async with SessionLocal() as session:
|
||||
session.add(
|
||||
User(
|
||||
email="pending-dev-login@test.com",
|
||||
password_hash=hash_password("Password123"),
|
||||
full_name="Pending Dev Login",
|
||||
clinical_department="Clinical",
|
||||
status=UserStatus.PENDING,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
resp = await client.post(
|
||||
"/api/v1/auth/dev-login",
|
||||
json={"email": payload["email"], "password": payload["password"]},
|
||||
json={"email": "pending-dev-login@test.com", "password": "Password123"},
|
||||
)
|
||||
finally:
|
||||
settings.ENV = original_env
|
||||
|
||||
Reference in New Issue
Block a user