feat: harden auth and study workflows
- replace plaintext login and unlock requests with RSA-OAEP/AES-GCM encrypted payloads - add login challenge replay protection, production RSA key validation, and auth tests - wire compose to environment-driven dev/prod settings without committing local secrets - update setup-config smoke scripts and Postman docs for encrypted login - add visit schedule migrations/tests and update study/subject setup workflows
This commit is contained in:
@@ -0,0 +1,43 @@
|
||||
"""replace global visit window with per-visit schedule
|
||||
|
||||
Revision ID: 20260508_01
|
||||
Revises: 20260331_01
|
||||
Create Date: 2026-05-08 11:30:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
|
||||
revision: str = "20260508_01"
|
||||
down_revision: Union[str, None] = "20260331_01"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"studies",
|
||||
sa.Column(
|
||||
"visit_schedule",
|
||||
postgresql.JSONB(astext_type=sa.Text()),
|
||||
nullable=False,
|
||||
server_default=sa.text("'[]'::jsonb"),
|
||||
),
|
||||
)
|
||||
op.drop_column("studies", "visit_window_end_offset")
|
||||
op.drop_column("studies", "visit_window_start_offset")
|
||||
op.drop_column("studies", "visit_total")
|
||||
op.drop_column("studies", "visit_interval_days")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.add_column("studies", sa.Column("visit_interval_days", sa.Integer(), nullable=True))
|
||||
op.add_column("studies", sa.Column("visit_total", sa.Integer(), nullable=True))
|
||||
op.add_column("studies", sa.Column("visit_window_start_offset", sa.Integer(), nullable=True))
|
||||
op.add_column("studies", sa.Column("visit_window_end_offset", sa.Integer(), nullable=True))
|
||||
op.drop_column("studies", "visit_schedule")
|
||||
@@ -0,0 +1,47 @@
|
||||
"""remove summary and objective note fields
|
||||
|
||||
Revision ID: 20260508_02
|
||||
Revises: 20260508_01
|
||||
Create Date: 2026-05-08 15:55:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = "20260508_02"
|
||||
down_revision: Union[str, None] = "20260508_01"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
bind.execute(
|
||||
sa.text(
|
||||
"""
|
||||
UPDATE study_setup_configs
|
||||
SET published_project_snapshot = published_project_snapshot - 'summary_note' - 'objective_note'
|
||||
WHERE published_project_snapshot IS NOT NULL
|
||||
"""
|
||||
)
|
||||
)
|
||||
bind.execute(
|
||||
sa.text(
|
||||
"""
|
||||
UPDATE study_setup_config_versions
|
||||
SET published_project_snapshot = published_project_snapshot - 'summary_note' - 'objective_note'
|
||||
WHERE published_project_snapshot IS NOT NULL
|
||||
"""
|
||||
)
|
||||
)
|
||||
op.drop_column("studies", "objective_note")
|
||||
op.drop_column("studies", "summary_note")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.add_column("studies", sa.Column("summary_note", sa.Text(), nullable=True))
|
||||
op.add_column("studies", sa.Column("objective_note", sa.Text(), nullable=True))
|
||||
@@ -0,0 +1,26 @@
|
||||
"""add subject baseline date
|
||||
|
||||
Revision ID: 20260508_03
|
||||
Revises: 20260508_02
|
||||
Create Date: 2026-05-08 16:35:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = "20260508_03"
|
||||
down_revision: Union[str, None] = "20260508_02"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("subjects", sa.Column("baseline_date", sa.Date(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("subjects", "baseline_date")
|
||||
+64
-30
@@ -1,12 +1,13 @@
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from fastapi import File, UploadFile
|
||||
from pydantic import BaseModel, EmailStr, Field
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from pathlib import Path
|
||||
import uuid
|
||||
|
||||
from app.core.config import settings
|
||||
from app.core.login_crypto import create_login_challenge, decrypt_login_payload, get_public_key_pem
|
||||
from app.core.security import create_access_token, decode_token_allow_expired, oauth2_scheme, verify_password
|
||||
from app.core.deps import get_current_user, get_db_session
|
||||
from app.crud import user as user_crud
|
||||
@@ -16,8 +17,16 @@ from fastapi.responses import FileResponse
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
email: EmailStr
|
||||
password: str = Field(min_length=1)
|
||||
key_id: str = Field(min_length=1)
|
||||
challenge: str = Field(min_length=16)
|
||||
ciphertext: str = Field(min_length=1)
|
||||
|
||||
|
||||
class LoginKeyResponse(BaseModel):
|
||||
key_id: str
|
||||
public_key: str
|
||||
challenge: str
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
class ExtendResponse(BaseModel):
|
||||
@@ -25,9 +34,8 @@ class ExtendResponse(BaseModel):
|
||||
expiresAt: datetime
|
||||
|
||||
|
||||
class UnlockRequest(BaseModel):
|
||||
email: EmailStr
|
||||
password: str = Field(min_length=1)
|
||||
class UnlockRequest(LoginRequest):
|
||||
pass
|
||||
|
||||
|
||||
class UnlockResponse(BaseModel):
|
||||
@@ -40,6 +48,42 @@ AVATAR_ROOT = Path(__file__).resolve().parent.parent.parent / "uploads" / "avata
|
||||
AVATAR_ROOT.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
||||
def issue_user_token(db_user) -> Token:
|
||||
session_start = datetime.now(timezone.utc)
|
||||
access_token = create_access_token(
|
||||
user_id=str(db_user.id),
|
||||
role=db_user.role.value if hasattr(db_user.role, "value") else db_user.role,
|
||||
expires_minutes=None,
|
||||
session_start=session_start,
|
||||
)
|
||||
return Token(access_token=access_token, token_type="bearer")
|
||||
|
||||
|
||||
async def authenticate_encrypted_password(payload: LoginRequest, db: AsyncSession):
|
||||
decrypted = decrypt_login_payload(
|
||||
key_id=payload.key_id,
|
||||
challenge=payload.challenge,
|
||||
ciphertext=payload.ciphertext,
|
||||
)
|
||||
if not decrypted:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="无法验证登录凭据",
|
||||
)
|
||||
db_user = await user_crud.get_by_email(db, decrypted.email)
|
||||
if not db_user:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="账号不存在",
|
||||
)
|
||||
if not verify_password(decrypted.password, db_user.password_hash):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="密码错误",
|
||||
)
|
||||
return db_user
|
||||
|
||||
|
||||
@router.post("/register", status_code=status.HTTP_201_CREATED)
|
||||
async def register(
|
||||
payload: UserRegisterRequest,
|
||||
@@ -54,35 +98,29 @@ async def register(
|
||||
return {"message": "注册成功,等待管理员审核"}
|
||||
|
||||
|
||||
@router.get("/login-key", response_model=LoginKeyResponse)
|
||||
async def get_login_key() -> LoginKeyResponse:
|
||||
challenge = create_login_challenge()
|
||||
return LoginKeyResponse(
|
||||
key_id=settings.LOGIN_RSA_KEY_ID,
|
||||
public_key=get_public_key_pem(),
|
||||
challenge=challenge.value,
|
||||
expires_at=challenge.expires_at,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/login", response_model=Token)
|
||||
async def login_for_access_token(
|
||||
payload: LoginRequest, db: AsyncSession = Depends(get_db_session)
|
||||
) -> Token:
|
||||
db_user = await user_crud.get_by_email(db, payload.email)
|
||||
if not db_user:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="账号不存在",
|
||||
)
|
||||
if not verify_password(payload.password, db_user.password_hash):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="密码错误",
|
||||
)
|
||||
db_user = await authenticate_encrypted_password(payload, db)
|
||||
if db_user.status != UserStatus.ACTIVE:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="账号未审核或不可用",
|
||||
)
|
||||
|
||||
session_start = datetime.now(timezone.utc)
|
||||
access_token = create_access_token(
|
||||
user_id=str(db_user.id),
|
||||
role=db_user.role.value if hasattr(db_user.role, "value") else db_user.role,
|
||||
expires_minutes=None,
|
||||
session_start=session_start,
|
||||
)
|
||||
return Token(access_token=access_token, token_type="bearer")
|
||||
return issue_user_token(db_user)
|
||||
|
||||
|
||||
@router.get("/me", response_model=UserRead)
|
||||
@@ -133,11 +171,7 @@ async def unlock_session(
|
||||
payload: UnlockRequest,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> UnlockResponse:
|
||||
db_user = await user_crud.get_by_email(db, payload.email)
|
||||
if not db_user:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="账号不存在")
|
||||
if not verify_password(payload.password, db_user.password_hash):
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="密码错误")
|
||||
db_user = await authenticate_encrypted_password(payload, db)
|
||||
if db_user.status != UserStatus.ACTIVE:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="账号已停用")
|
||||
session_start = datetime.now(timezone.utc)
|
||||
|
||||
@@ -283,13 +283,8 @@ def _build_project_publish_snapshot(study) -> ProjectPublishSnapshot:
|
||||
plan_end_date=_to_date_text(getattr(study, "plan_end_date", None)),
|
||||
planned_site_count=getattr(study, "planned_site_count", None),
|
||||
planned_enrollment_count=getattr(study, "planned_enrollment_count", None),
|
||||
summary_note=getattr(study, "summary_note", None) or "",
|
||||
objective_note=getattr(study, "objective_note", None) or "",
|
||||
status=getattr(study, "status", None) or "",
|
||||
visit_interval_days=getattr(study, "visit_interval_days", None),
|
||||
visit_total=getattr(study, "visit_total", None),
|
||||
visit_window_start_offset=getattr(study, "visit_window_start_offset", None),
|
||||
visit_window_end_offset=getattr(study, "visit_window_end_offset", None),
|
||||
visit_schedule=getattr(study, "visit_schedule", None) or [],
|
||||
)
|
||||
|
||||
|
||||
@@ -360,17 +355,9 @@ def _ensure_study_timeline_valid(
|
||||
*,
|
||||
plan_start_date: date | None,
|
||||
plan_end_date: date | None,
|
||||
visit_window_start_offset: int | None,
|
||||
visit_window_end_offset: int | None,
|
||||
) -> None:
|
||||
if plan_start_date and plan_end_date and plan_start_date > plan_end_date:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="项目计划结束日期不能早于开始日期")
|
||||
if (
|
||||
visit_window_start_offset is not None
|
||||
and visit_window_end_offset is not None
|
||||
and visit_window_start_offset > visit_window_end_offset
|
||||
):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="访视窗口开始偏移不能晚于结束偏移")
|
||||
|
||||
|
||||
def _resolve_version_label(record: Any) -> str:
|
||||
@@ -649,8 +636,6 @@ async def create_study(
|
||||
_ensure_study_timeline_valid(
|
||||
plan_start_date=study_in.plan_start_date,
|
||||
plan_end_date=study_in.plan_end_date,
|
||||
visit_window_start_offset=study_in.visit_window_start_offset,
|
||||
visit_window_end_offset=study_in.visit_window_end_offset,
|
||||
)
|
||||
study_in.code = study_in.code.strip()
|
||||
if not study_in.code:
|
||||
@@ -746,21 +731,9 @@ async def update_study(
|
||||
|
||||
next_plan_start = study_in.plan_start_date if "plan_start_date" in study_in.model_fields_set else study.plan_start_date
|
||||
next_plan_end = study_in.plan_end_date if "plan_end_date" in study_in.model_fields_set else study.plan_end_date
|
||||
next_window_start = (
|
||||
study_in.visit_window_start_offset
|
||||
if "visit_window_start_offset" in study_in.model_fields_set
|
||||
else study.visit_window_start_offset
|
||||
)
|
||||
next_window_end = (
|
||||
study_in.visit_window_end_offset
|
||||
if "visit_window_end_offset" in study_in.model_fields_set
|
||||
else study.visit_window_end_offset
|
||||
)
|
||||
_ensure_study_timeline_valid(
|
||||
plan_start_date=next_plan_start,
|
||||
plan_end_date=next_plan_end,
|
||||
visit_window_start_offset=next_window_start,
|
||||
visit_window_end_offset=next_window_end,
|
||||
)
|
||||
|
||||
updated = await study_crud.update(db, study, study_in)
|
||||
|
||||
@@ -173,9 +173,9 @@ async def update_subject(
|
||||
await _ensure_subject_active(db, subject)
|
||||
old_status = subject.status
|
||||
updated = await subject_crud.update_subject(db, subject, subject_in)
|
||||
# auto-generate visits when enrolled
|
||||
if subject_in.status and subject_in.status == "ENROLLED" and (subject_in.enrollment_date or updated.enrollment_date):
|
||||
await subject_crud.generate_default_visits(db, updated)
|
||||
# 基线/治疗日期是访视计划的唯一推算基准,不能用入组日期替代。
|
||||
if updated.baseline_date:
|
||||
await subject_crud.sync_visits_from_baseline(db, updated)
|
||||
detail = None
|
||||
if subject_in.status and subject_in.status != old_status:
|
||||
detail = f"参与者 {updated.subject_no} 状态 {old_status} -> {subject_in.status}"
|
||||
|
||||
@@ -90,15 +90,12 @@ async def create_visit(
|
||||
if next_visit_code == "V1" and visit_in.planned_date:
|
||||
study = await study_crud.get(db, study_id)
|
||||
if study:
|
||||
await visit_crud.create_followup_visits(
|
||||
await visit_crud.create_scheduled_visits(
|
||||
db,
|
||||
study_id=study_id,
|
||||
subject=subject,
|
||||
base_date=visit_in.planned_date,
|
||||
visit_total=study.visit_total,
|
||||
visit_interval_days=study.visit_interval_days,
|
||||
window_start_offset=study.visit_window_start_offset,
|
||||
window_end_offset=study.visit_window_end_offset,
|
||||
visit_schedule=study.visit_schedule,
|
||||
)
|
||||
await audit_crud.log_action(
|
||||
db,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from functools import lru_cache
|
||||
from typing import Literal
|
||||
from typing import Literal, Optional
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
@@ -19,6 +19,11 @@ class Settings(BaseSettings):
|
||||
JWT_EXPIRE_MINUTES: int = 60
|
||||
JWT_EXTEND_GRACE_SECONDS: int = 120
|
||||
ABSOLUTE_SESSION_MAX_HOURS: int = 8
|
||||
LOGIN_RSA_PRIVATE_KEY: Optional[str] = None
|
||||
LOGIN_RSA_PUBLIC_KEY: Optional[str] = None
|
||||
LOGIN_RSA_KEY_ID: str = "default"
|
||||
LOGIN_CHALLENGE_TTL_SECONDS: int = 120
|
||||
LOGIN_CHALLENGE_MAX_ACTIVE: int = 1000
|
||||
|
||||
|
||||
@lru_cache
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
import base64
|
||||
import json
|
||||
import secrets
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import padding, rsa
|
||||
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey, RSAPublicKey
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
from pydantic import BaseModel, EmailStr, Field, ValidationError
|
||||
|
||||
from app.core.config import settings
|
||||
|
||||
|
||||
class DecryptedLoginPayload(BaseModel):
|
||||
email: EmailStr
|
||||
password: str = Field(min_length=1, max_length=72)
|
||||
challenge: str = Field(min_length=16)
|
||||
|
||||
|
||||
class EncryptedLoginEnvelope(BaseModel):
|
||||
encrypted_key: str = Field(min_length=1)
|
||||
iv: str = Field(min_length=1)
|
||||
data: str = Field(min_length=1)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LoginChallenge:
|
||||
value: str
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
_private_key: RSAPrivateKey | None = None
|
||||
_public_key_pem: str | None = None
|
||||
_challenges: dict[str, LoginChallenge] = {}
|
||||
|
||||
|
||||
def _normalize_pem(value: str) -> str:
|
||||
return value.replace("\\n", "\n").strip()
|
||||
|
||||
|
||||
def _load_or_create_private_key() -> RSAPrivateKey:
|
||||
configured_key = settings.LOGIN_RSA_PRIVATE_KEY
|
||||
if configured_key:
|
||||
key = serialization.load_pem_private_key(_normalize_pem(configured_key).encode("utf-8"), password=None)
|
||||
if not isinstance(key, RSAPrivateKey):
|
||||
raise ValueError("LOGIN_RSA_PRIVATE_KEY 必须是 RSA 私钥")
|
||||
return key
|
||||
if settings.ENV == "production":
|
||||
raise ValueError("生产环境必须配置 LOGIN_RSA_PRIVATE_KEY")
|
||||
return rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
|
||||
|
||||
def get_private_key() -> RSAPrivateKey:
|
||||
global _private_key
|
||||
if _private_key is None:
|
||||
_private_key = _load_or_create_private_key()
|
||||
return _private_key
|
||||
|
||||
|
||||
def validate_login_crypto_configuration() -> None:
|
||||
if settings.ENV == "production" and not settings.LOGIN_RSA_PRIVATE_KEY:
|
||||
raise ValueError("生产环境必须配置 LOGIN_RSA_PRIVATE_KEY")
|
||||
|
||||
|
||||
def get_public_key_pem() -> str:
|
||||
global _public_key_pem
|
||||
if _public_key_pem is not None:
|
||||
return _public_key_pem
|
||||
if settings.LOGIN_RSA_PUBLIC_KEY:
|
||||
if not settings.LOGIN_RSA_PRIVATE_KEY:
|
||||
raise ValueError("配置 LOGIN_RSA_PUBLIC_KEY 时必须同时配置 LOGIN_RSA_PRIVATE_KEY")
|
||||
_public_key_pem = _normalize_pem(settings.LOGIN_RSA_PUBLIC_KEY)
|
||||
return _public_key_pem
|
||||
public_key: RSAPublicKey = get_private_key().public_key()
|
||||
_public_key_pem = public_key.public_bytes(
|
||||
encoding=serialization.Encoding.PEM,
|
||||
format=serialization.PublicFormat.SubjectPublicKeyInfo,
|
||||
).decode("utf-8")
|
||||
return _public_key_pem
|
||||
|
||||
|
||||
def create_login_challenge() -> LoginChallenge:
|
||||
prune_expired_challenges()
|
||||
while len(_challenges) >= settings.LOGIN_CHALLENGE_MAX_ACTIVE:
|
||||
oldest = next(iter(_challenges))
|
||||
_challenges.pop(oldest, None)
|
||||
value = secrets.token_urlsafe(32)
|
||||
expires_at = datetime.now(timezone.utc) + timedelta(seconds=settings.LOGIN_CHALLENGE_TTL_SECONDS)
|
||||
challenge = LoginChallenge(value=value, expires_at=expires_at)
|
||||
_challenges[value] = challenge
|
||||
return challenge
|
||||
|
||||
|
||||
def prune_expired_challenges() -> None:
|
||||
now = datetime.now(timezone.utc)
|
||||
expired = [value for value, challenge in _challenges.items() if challenge.expires_at <= now]
|
||||
for value in expired:
|
||||
_challenges.pop(value, None)
|
||||
|
||||
|
||||
def consume_login_challenge(value: str) -> bool:
|
||||
prune_expired_challenges()
|
||||
challenge = _challenges.pop(value, None)
|
||||
if not challenge:
|
||||
return False
|
||||
return challenge.expires_at > datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def decrypt_login_payload(*, key_id: str, challenge: str, ciphertext: str) -> DecryptedLoginPayload | None:
|
||||
if key_id != settings.LOGIN_RSA_KEY_ID:
|
||||
return None
|
||||
if not consume_login_challenge(challenge):
|
||||
return None
|
||||
try:
|
||||
envelope_data = json.loads(base64.b64decode(ciphertext, validate=True).decode("utf-8"))
|
||||
envelope = EncryptedLoginEnvelope.model_validate(envelope_data)
|
||||
encrypted_key = base64.b64decode(envelope.encrypted_key, validate=True)
|
||||
iv = base64.b64decode(envelope.iv, validate=True)
|
||||
encrypted_payload = base64.b64decode(envelope.data, validate=True)
|
||||
aes_key = get_private_key().decrypt(
|
||||
encrypted_key,
|
||||
padding.OAEP(
|
||||
mgf=padding.MGF1(algorithm=hashes.SHA256()),
|
||||
algorithm=hashes.SHA256(),
|
||||
label=None,
|
||||
),
|
||||
)
|
||||
decrypted = AESGCM(aes_key).decrypt(iv, encrypted_payload, None)
|
||||
payload: Any = json.loads(decrypted.decode("utf-8"))
|
||||
parsed = DecryptedLoginPayload.model_validate(payload)
|
||||
except (ValueError, TypeError, json.JSONDecodeError, ValidationError):
|
||||
return None
|
||||
if parsed.challenge != challenge:
|
||||
return None
|
||||
return parsed
|
||||
@@ -30,14 +30,9 @@ async def create(db: AsyncSession, study_in: StudyCreate, *, created_by: uuid.UU
|
||||
planned_enrollment_count=study_in.planned_enrollment_count,
|
||||
enrollment_monthly_goal_note=study_in.enrollment_monthly_goal_note,
|
||||
enrollment_stage_breakdown=study_in.enrollment_stage_breakdown,
|
||||
summary_note=study_in.summary_note,
|
||||
objective_note=study_in.objective_note,
|
||||
phase=study_in.phase,
|
||||
status=study_in.status,
|
||||
visit_interval_days=study_in.visit_interval_days,
|
||||
visit_total=study_in.visit_total,
|
||||
visit_window_start_offset=study_in.visit_window_start_offset,
|
||||
visit_window_end_offset=study_in.visit_window_end_offset,
|
||||
visit_schedule=[item.model_dump() for item in study_in.visit_schedule],
|
||||
created_by=created_by,
|
||||
)
|
||||
db.add(study)
|
||||
|
||||
+26
-44
@@ -1,5 +1,5 @@
|
||||
import uuid
|
||||
from datetime import date, timedelta
|
||||
from datetime import date
|
||||
from typing import Sequence
|
||||
|
||||
from sqlalchemy import delete as sa_delete, select, update as sa_update
|
||||
@@ -37,6 +37,10 @@ def _validate_subject_date_chain(
|
||||
raise ValueError("完成日期不能早于入组日期")
|
||||
|
||||
|
||||
def _should_sync_visits(previous_baseline_date: date | None, next_baseline_date: date | None) -> bool:
|
||||
return next_baseline_date is not None and previous_baseline_date != next_baseline_date
|
||||
|
||||
|
||||
async def _validate_site(db: AsyncSession, study_id: uuid.UUID, site_id: uuid.UUID) -> None:
|
||||
result = await db.execute(select(Site).where(Site.id == site_id))
|
||||
site = result.scalar_one_or_none()
|
||||
@@ -51,7 +55,7 @@ async def create_subject(db: AsyncSession, study_id: uuid.UUID, subject_in: Subj
|
||||
_validate_subject_date_chain(
|
||||
screening_date=subject_in.screening_date,
|
||||
consent_date=subject_in.consent_date,
|
||||
enrollment_date=None,
|
||||
enrollment_date=subject_in.enrollment_date,
|
||||
completion_date=None,
|
||||
)
|
||||
subject = Subject(
|
||||
@@ -61,7 +65,8 @@ async def create_subject(db: AsyncSession, study_id: uuid.UUID, subject_in: Subj
|
||||
status="SCREENING",
|
||||
screening_date=subject_in.screening_date,
|
||||
consent_date=subject_in.consent_date,
|
||||
enrollment_date=None,
|
||||
enrollment_date=subject_in.enrollment_date,
|
||||
baseline_date=subject_in.baseline_date,
|
||||
completion_date=None,
|
||||
drop_reason=None,
|
||||
)
|
||||
@@ -69,15 +74,7 @@ async def create_subject(db: AsyncSession, study_id: uuid.UUID, subject_in: Subj
|
||||
await db.commit()
|
||||
await db.refresh(subject)
|
||||
|
||||
# initial visit: Screening (V0)
|
||||
await visit_crud.create_visit(
|
||||
db,
|
||||
study_id=study_id,
|
||||
visit_in=None,
|
||||
subject=subject,
|
||||
visit_code="V0",
|
||||
planned_date=subject.screening_date,
|
||||
)
|
||||
await sync_visits_from_baseline(db, subject)
|
||||
return subject
|
||||
|
||||
|
||||
@@ -103,49 +100,26 @@ async def list_subjects(
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
async def generate_default_visits(db: AsyncSession, subject: Subject) -> None:
|
||||
# Baseline + Follow-up visits based on enrollment_date
|
||||
if not subject.enrollment_date:
|
||||
return
|
||||
async def sync_visits_from_baseline(db: AsyncSession, subject: Subject) -> None:
|
||||
# 基线/治疗日期是访视计划的唯一推算基准,不能用入组日期替代。
|
||||
result = await db.execute(select(Study).where(Study.id == subject.study_id))
|
||||
study = result.scalar_one_or_none()
|
||||
if not study:
|
||||
return
|
||||
|
||||
visit_total = study.visit_total or 3
|
||||
visit_interval_days = study.visit_interval_days or 30
|
||||
window_start_offset = study.visit_window_start_offset
|
||||
window_end_offset = study.visit_window_end_offset
|
||||
baseline_date = subject.enrollment_date
|
||||
|
||||
result = await db.execute(select(Visit.visit_code).where(Visit.subject_id == subject.id))
|
||||
existing_codes = {row[0] for row in result.all()}
|
||||
if "V1" not in existing_codes:
|
||||
window_start = baseline_date + timedelta(days=window_start_offset) if window_start_offset is not None else None
|
||||
window_end = baseline_date + timedelta(days=window_end_offset) if window_end_offset is not None else None
|
||||
await visit_crud.create_visit(
|
||||
db,
|
||||
study_id=subject.study_id,
|
||||
visit_in=None,
|
||||
subject=subject,
|
||||
visit_code="V1",
|
||||
planned_date=baseline_date,
|
||||
window_start=window_start,
|
||||
window_end=window_end,
|
||||
)
|
||||
|
||||
await visit_crud.create_followup_visits(
|
||||
await visit_crud.create_scheduled_visits(
|
||||
db,
|
||||
study_id=subject.study_id,
|
||||
subject=subject,
|
||||
base_date=baseline_date,
|
||||
visit_total=visit_total,
|
||||
visit_interval_days=visit_interval_days,
|
||||
window_start_offset=window_start_offset,
|
||||
window_end_offset=window_end_offset,
|
||||
base_date=subject.baseline_date,
|
||||
visit_schedule=study.visit_schedule,
|
||||
)
|
||||
|
||||
|
||||
async def generate_default_visits(db: AsyncSession, subject: Subject) -> None:
|
||||
await sync_visits_from_baseline(db, subject)
|
||||
|
||||
|
||||
async def update_subject(db: AsyncSession, subject: Subject, subject_in: SubjectUpdate) -> Subject:
|
||||
update_data = subject_in.model_dump(exclude_unset=True)
|
||||
next_screening_date = subject.screening_date
|
||||
@@ -169,6 +143,14 @@ async def update_subject(db: AsyncSession, subject: Subject, subject_in: Subject
|
||||
return subject
|
||||
|
||||
|
||||
def should_generate_visits_after_subject_update(
|
||||
*,
|
||||
previous_baseline_date: date | None,
|
||||
next_baseline_date: date | None,
|
||||
) -> bool:
|
||||
return _should_sync_visits(previous_baseline_date, next_baseline_date)
|
||||
|
||||
|
||||
async def delete_subject(db: AsyncSession, subject: Subject) -> None:
|
||||
subject_id = subject.id
|
||||
await db.execute(sa_delete(Visit).where(Visit.subject_id == subject_id))
|
||||
|
||||
+99
-24
@@ -5,11 +5,59 @@ from typing import Sequence
|
||||
from sqlalchemy import and_, func, or_, select, update as sa_update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.study import Study
|
||||
from app.models.subject import Subject
|
||||
from app.models.visit import Visit
|
||||
from app.schemas.visit import VisitCreate, VisitUpdate
|
||||
|
||||
|
||||
def _visit_schedule_order_map(visit_schedule: list[dict] | None) -> dict[str, int]:
|
||||
order_map: dict[str, int] = {}
|
||||
for index, item in enumerate(visit_schedule or []):
|
||||
code = str(item.get("visit_code") or "").strip()
|
||||
if code and code not in order_map:
|
||||
order_map[code] = index
|
||||
return order_map
|
||||
|
||||
|
||||
def sort_visits_for_display(visits: Sequence[Visit], visit_schedule: list[dict] | None = None) -> list[Visit]:
|
||||
order_map = _visit_schedule_order_map(visit_schedule)
|
||||
fallback_start = len(order_map)
|
||||
return sorted(
|
||||
visits,
|
||||
key=lambda visit: (
|
||||
order_map.get((visit.visit_code or "").strip(), fallback_start),
|
||||
visit.planned_date is None,
|
||||
visit.planned_date or date.max,
|
||||
visit.visit_code or "",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def build_visit_schedule_dates(visit_schedule: list[dict] | None, base_date: date | None) -> list[dict]:
|
||||
if not visit_schedule:
|
||||
return []
|
||||
rows: list[dict] = []
|
||||
for item in visit_schedule:
|
||||
code = str(item.get("visit_code") or "").strip()
|
||||
if not code:
|
||||
continue
|
||||
baseline_offset_days = int(item.get("baseline_offset_days", 0))
|
||||
window_before_days = int(item.get("window_before_days", 0))
|
||||
window_after_days = int(item.get("window_after_days", 0))
|
||||
planned_date = base_date + timedelta(days=baseline_offset_days) if base_date else None
|
||||
rows.append(
|
||||
{
|
||||
"visit_code": code,
|
||||
"baseline_offset_days": baseline_offset_days,
|
||||
"planned_date": planned_date,
|
||||
"window_start": planned_date - timedelta(days=window_before_days) if planned_date else None,
|
||||
"window_end": planned_date + timedelta(days=window_after_days) if planned_date else None,
|
||||
}
|
||||
)
|
||||
return rows
|
||||
|
||||
|
||||
async def create_visit(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
@@ -20,14 +68,16 @@ async def create_visit(
|
||||
planned_date,
|
||||
window_start: date | None = None,
|
||||
window_end: date | None = None,
|
||||
actual_date: date | None = None,
|
||||
status: str = "PLANNED",
|
||||
) -> Visit:
|
||||
visit = Visit(
|
||||
study_id=study_id,
|
||||
subject_id=subject.id,
|
||||
visit_code=visit_code,
|
||||
planned_date=planned_date,
|
||||
actual_date=None,
|
||||
status="PLANNED",
|
||||
actual_date=actual_date,
|
||||
status=status,
|
||||
window_start=window_start if window_start is not None else (visit_in.window_start if visit_in else None),
|
||||
window_end=window_end if window_end is not None else (visit_in.window_end if visit_in else None),
|
||||
notes=None,
|
||||
@@ -39,8 +89,13 @@ async def create_visit(
|
||||
|
||||
|
||||
async def list_visits(db: AsyncSession, subject_id: uuid.UUID) -> Sequence[Visit]:
|
||||
result = await db.execute(select(Visit).where(Visit.subject_id == subject_id).order_by(Visit.planned_date))
|
||||
return result.scalars().all()
|
||||
result = await db.execute(select(Visit).where(Visit.subject_id == subject_id))
|
||||
visits = result.scalars().all()
|
||||
if not visits:
|
||||
return []
|
||||
study_result = await db.execute(select(Study.visit_schedule).where(Study.id == visits[0].study_id))
|
||||
visit_schedule = study_result.scalar_one_or_none() or []
|
||||
return sort_visits_for_display(visits, visit_schedule)
|
||||
|
||||
|
||||
async def mark_overdue_as_lost(db: AsyncSession, subject_id: uuid.UUID) -> None:
|
||||
@@ -150,32 +205,50 @@ async def list_lost_visits(
|
||||
return result.all()
|
||||
|
||||
|
||||
async def create_followup_visits(
|
||||
async def create_scheduled_visits(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
study_id: uuid.UUID,
|
||||
subject: Subject,
|
||||
base_date: date,
|
||||
visit_total: int | None,
|
||||
visit_interval_days: int | None,
|
||||
window_start_offset: int | None,
|
||||
window_end_offset: int | None,
|
||||
base_date: date | None,
|
||||
visit_schedule: list[dict] | None,
|
||||
) -> Sequence[Visit]:
|
||||
if not visit_total or not visit_interval_days or visit_total < 2:
|
||||
if not visit_schedule:
|
||||
return []
|
||||
|
||||
result = await db.execute(select(Visit.visit_code).where(Visit.subject_id == subject.id))
|
||||
existing_codes = {row[0] for row in result.all()}
|
||||
result = await db.execute(select(Visit).where(Visit.subject_id == subject.id))
|
||||
existing_by_code = {visit.visit_code: visit for visit in result.scalars().all()}
|
||||
created: list[Visit] = []
|
||||
for index in range(2, visit_total + 1):
|
||||
code = f"V{index}"
|
||||
if code in existing_codes:
|
||||
for item in build_visit_schedule_dates(visit_schedule, base_date):
|
||||
code = item["visit_code"]
|
||||
is_baseline_visit = item["baseline_offset_days"] == 0 and base_date is not None
|
||||
existing = existing_by_code.get(code)
|
||||
if existing:
|
||||
if is_baseline_visit:
|
||||
await db.execute(
|
||||
sa_update(Visit)
|
||||
.where(Visit.id == existing.id)
|
||||
.values(
|
||||
planned_date=item["planned_date"],
|
||||
window_start=item["window_start"],
|
||||
window_end=item["window_end"],
|
||||
actual_date=base_date,
|
||||
status="DONE",
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
elif existing.status == "PLANNED" and existing.actual_date is None:
|
||||
await db.execute(
|
||||
sa_update(Visit)
|
||||
.where(Visit.id == existing.id)
|
||||
.values(
|
||||
planned_date=item["planned_date"],
|
||||
window_start=item["window_start"],
|
||||
window_end=item["window_end"],
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
continue
|
||||
planned_date = base_date + timedelta(days=visit_interval_days * (index - 1))
|
||||
window_start = (
|
||||
planned_date + timedelta(days=window_start_offset) if window_start_offset is not None else None
|
||||
)
|
||||
window_end = planned_date + timedelta(days=window_end_offset) if window_end_offset is not None else None
|
||||
created.append(
|
||||
await create_visit(
|
||||
db,
|
||||
@@ -183,9 +256,11 @@ async def create_followup_visits(
|
||||
visit_in=None,
|
||||
subject=subject,
|
||||
visit_code=code,
|
||||
planned_date=planned_date,
|
||||
window_start=window_start,
|
||||
window_end=window_end,
|
||||
planned_date=item["planned_date"],
|
||||
window_start=item["window_start"],
|
||||
window_end=item["window_end"],
|
||||
actual_date=base_date if is_baseline_visit else None,
|
||||
status="DONE" if is_baseline_visit else "PLANNED",
|
||||
)
|
||||
)
|
||||
return created
|
||||
|
||||
@@ -12,6 +12,7 @@ from sqlalchemy import text
|
||||
from app.api.v1.router import api_router
|
||||
from app.core.config import settings
|
||||
from app.core.exceptions import register_exception_handlers
|
||||
from app.core.login_crypto import validate_login_crypto_configuration
|
||||
from app.crud.user import ensure_admin_exists
|
||||
from app.db.base import Base
|
||||
from app.db.session import SessionLocal, engine
|
||||
@@ -76,6 +77,7 @@ async def _ensure_legacy_primary_keys(conn) -> None:
|
||||
|
||||
|
||||
def create_app() -> FastAPI:
|
||||
validate_login_crypto_configuration()
|
||||
app = FastAPI(
|
||||
title="CTMS 后端 API",
|
||||
description="临床试验项目管理系统后端接口文档",
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import uuid
|
||||
from datetime import date, datetime
|
||||
|
||||
from sqlalchemy import Boolean, Date, DateTime, ForeignKey, Integer, String, Text, func
|
||||
from sqlalchemy import Boolean, Date, DateTime, ForeignKey, Integer, JSON, String, Text, func
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
@@ -32,14 +32,9 @@ class Study(Base):
|
||||
planned_enrollment_count: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
enrollment_monthly_goal_note: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
enrollment_stage_breakdown: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
summary_note: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
objective_note: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
phase: Mapped[str | None] = mapped_column(String(50), nullable=True)
|
||||
status: Mapped[str] = mapped_column(String(20), nullable=False, default="DRAFT")
|
||||
is_locked: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
visit_interval_days: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
visit_total: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
visit_window_start_offset: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
visit_window_end_offset: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
visit_schedule: Mapped[list[dict]] = mapped_column(JSON, nullable=False, default=list)
|
||||
created_by: Mapped[uuid.UUID | None] = 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())
|
||||
|
||||
@@ -20,6 +20,7 @@ class Subject(Base):
|
||||
screening_date: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
consent_date: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
enrollment_date: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
baseline_date: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
completion_date: Mapped[date | None] = mapped_column(Date, nullable=True)
|
||||
drop_reason: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
|
||||
|
||||
@@ -2,12 +2,36 @@ import uuid
|
||||
from datetime import date, datetime
|
||||
from typing import Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
StudyStatus = Literal["DRAFT", "ACTIVE", "CLOSED"]
|
||||
|
||||
|
||||
class StudyCreate(BaseModel):
|
||||
class VisitScheduleItem(BaseModel):
|
||||
visit_code: str = Field(min_length=1, max_length=50)
|
||||
baseline_offset_days: int = Field(ge=0, le=3650)
|
||||
window_before_days: int = Field(ge=0, le=365)
|
||||
window_after_days: int = Field(ge=0, le=365)
|
||||
|
||||
|
||||
class VisitScheduleMixin(BaseModel):
|
||||
visit_schedule: list[VisitScheduleItem] = Field(default_factory=list)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_visit_schedule(self):
|
||||
codes: set[str] = set()
|
||||
for index, item in enumerate(self.visit_schedule):
|
||||
code = item.visit_code.strip()
|
||||
if not code:
|
||||
raise ValueError(f"第 {index + 1} 行访视编号不能为空")
|
||||
if code in codes:
|
||||
raise ValueError(f"访视编号重复:{code}")
|
||||
codes.add(code)
|
||||
item.visit_code = code
|
||||
return self
|
||||
|
||||
|
||||
class StudyCreate(VisitScheduleMixin):
|
||||
name: str = Field(min_length=1)
|
||||
code: str = Field(min_length=1)
|
||||
project_full_name: Optional[str] = None
|
||||
@@ -28,14 +52,8 @@ class StudyCreate(BaseModel):
|
||||
planned_enrollment_count: Optional[int] = None
|
||||
enrollment_monthly_goal_note: Optional[str] = None
|
||||
enrollment_stage_breakdown: Optional[str] = None
|
||||
summary_note: Optional[str] = None
|
||||
objective_note: Optional[str] = None
|
||||
phase: Optional[str] = None
|
||||
status: StudyStatus = "DRAFT"
|
||||
visit_interval_days: Optional[int] = None
|
||||
visit_total: Optional[int] = None
|
||||
visit_window_start_offset: Optional[int] = None
|
||||
visit_window_end_offset: Optional[int] = None
|
||||
|
||||
|
||||
class StudyUpdate(BaseModel):
|
||||
@@ -59,15 +77,25 @@ class StudyUpdate(BaseModel):
|
||||
planned_enrollment_count: Optional[int] = None
|
||||
enrollment_monthly_goal_note: Optional[str] = None
|
||||
enrollment_stage_breakdown: Optional[str] = None
|
||||
summary_note: Optional[str] = None
|
||||
objective_note: Optional[str] = None
|
||||
phase: Optional[str] = None
|
||||
status: Optional[StudyStatus] = None
|
||||
is_locked: Optional[bool] = None
|
||||
visit_interval_days: Optional[int] = None
|
||||
visit_total: Optional[int] = None
|
||||
visit_window_start_offset: Optional[int] = None
|
||||
visit_window_end_offset: Optional[int] = None
|
||||
visit_schedule: Optional[list[VisitScheduleItem]] = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_visit_schedule(self):
|
||||
if self.visit_schedule is None:
|
||||
return self
|
||||
codes: set[str] = set()
|
||||
for index, item in enumerate(self.visit_schedule):
|
||||
code = item.visit_code.strip()
|
||||
if not code:
|
||||
raise ValueError(f"第 {index + 1} 行访视编号不能为空")
|
||||
if code in codes:
|
||||
raise ValueError(f"访视编号重复:{code}")
|
||||
codes.add(code)
|
||||
item.visit_code = code
|
||||
return self
|
||||
|
||||
|
||||
class StudyRead(BaseModel):
|
||||
@@ -92,15 +120,10 @@ class StudyRead(BaseModel):
|
||||
planned_enrollment_count: Optional[int]
|
||||
enrollment_monthly_goal_note: Optional[str]
|
||||
enrollment_stage_breakdown: Optional[str]
|
||||
summary_note: Optional[str]
|
||||
objective_note: Optional[str]
|
||||
phase: Optional[str]
|
||||
status: StudyStatus
|
||||
is_locked: bool
|
||||
visit_interval_days: Optional[int]
|
||||
visit_total: Optional[int]
|
||||
visit_window_start_offset: Optional[int]
|
||||
visit_window_end_offset: Optional[int]
|
||||
visit_schedule: list[VisitScheduleItem]
|
||||
created_by: Optional[uuid.UUID]
|
||||
created_at: datetime
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import uuid
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
ConfirmStatus = Literal["待确认", "已确认", "退回"]
|
||||
@@ -143,6 +143,13 @@ class SetupProjectionSummary(BaseModel):
|
||||
skipped_items: list[SetupProjectionSkippedItem] = Field(default_factory=list)
|
||||
|
||||
|
||||
class VisitScheduleItem(BaseModel):
|
||||
visit_code: str = Field(min_length=1, max_length=50)
|
||||
baseline_offset_days: int = Field(ge=0, le=3650)
|
||||
window_before_days: int = Field(ge=0, le=365)
|
||||
window_after_days: int = Field(ge=0, le=365)
|
||||
|
||||
|
||||
class ProjectPublishSnapshot(BaseModel):
|
||||
code: str = ""
|
||||
name: str = ""
|
||||
@@ -162,13 +169,21 @@ class ProjectPublishSnapshot(BaseModel):
|
||||
plan_end_date: str = ""
|
||||
planned_site_count: int | None = None
|
||||
planned_enrollment_count: int | None = None
|
||||
summary_note: str = ""
|
||||
objective_note: str = ""
|
||||
status: str = ""
|
||||
visit_interval_days: int | None = None
|
||||
visit_total: int | None = None
|
||||
visit_window_start_offset: int | None = None
|
||||
visit_window_end_offset: int | None = None
|
||||
visit_schedule: list[VisitScheduleItem] = Field(default_factory=list)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_visit_schedule(self):
|
||||
codes: set[str] = set()
|
||||
for index, item in enumerate(self.visit_schedule):
|
||||
code = item.visit_code.strip()
|
||||
if not code:
|
||||
raise ValueError(f"第 {index + 1} 行访视编号不能为空")
|
||||
if code in codes:
|
||||
raise ValueError(f"访视编号重复:{code}")
|
||||
codes.add(code)
|
||||
item.visit_code = code
|
||||
return self
|
||||
|
||||
|
||||
class StudySetupConfigRead(BaseModel):
|
||||
|
||||
@@ -10,12 +10,15 @@ class SubjectCreate(BaseModel):
|
||||
subject_no: str
|
||||
screening_date: Optional[date] = None
|
||||
consent_date: Optional[date] = None
|
||||
enrollment_date: Optional[date] = None
|
||||
baseline_date: Optional[date] = None
|
||||
|
||||
|
||||
class SubjectUpdate(BaseModel):
|
||||
status: Optional[str] = None
|
||||
consent_date: Optional[date] = None
|
||||
enrollment_date: Optional[date] = None
|
||||
baseline_date: Optional[date] = None
|
||||
completion_date: Optional[date] = None
|
||||
drop_reason: Optional[str] = None
|
||||
|
||||
@@ -29,6 +32,7 @@ class SubjectRead(BaseModel):
|
||||
screening_date: Optional[date]
|
||||
consent_date: Optional[date]
|
||||
enrollment_date: Optional[date]
|
||||
baseline_date: Optional[date]
|
||||
completion_date: Optional[date]
|
||||
drop_reason: Optional[str]
|
||||
created_at: datetime
|
||||
|
||||
@@ -23,7 +23,7 @@ class UserDisplay(BaseModel):
|
||||
|
||||
|
||||
class _PasswordValidator(BaseModel):
|
||||
password: Optional[str] = Field(default=None, min_length=8)
|
||||
password: Optional[str] = Field(default=None, min_length=8, max_length=72)
|
||||
|
||||
@field_validator("password")
|
||||
@classmethod
|
||||
@@ -36,7 +36,7 @@ class _PasswordValidator(BaseModel):
|
||||
|
||||
|
||||
class UserRegisterRequest(_PasswordValidator):
|
||||
password: str = Field(min_length=8)
|
||||
password: str = Field(min_length=8, max_length=72)
|
||||
email: EmailStr
|
||||
full_name: str = Field(min_length=1)
|
||||
role: RegisterRole
|
||||
@@ -44,7 +44,7 @@ class UserRegisterRequest(_PasswordValidator):
|
||||
|
||||
|
||||
class UserCreate(_PasswordValidator):
|
||||
password: str = Field(min_length=8)
|
||||
password: str = Field(min_length=8, max_length=72)
|
||||
email: EmailStr
|
||||
full_name: str = Field(min_length=1)
|
||||
role: UserRole
|
||||
|
||||
@@ -2,6 +2,7 @@ fastapi==0.104.1
|
||||
uvicorn[standard]==0.24.0.post1
|
||||
sqlalchemy==2.0.23
|
||||
asyncpg==0.29.0
|
||||
aiosqlite==0.20.0
|
||||
pydantic-settings==2.1.0
|
||||
python-jose[cryptography]==3.3.0
|
||||
passlib[bcrypt]==1.7.4
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
@@ -6,6 +7,9 @@ import urllib.error
|
||||
import urllib.request
|
||||
|
||||
import asyncpg
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import padding
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
|
||||
BASE = os.getenv("BASE_URL", "http://localhost:8000")
|
||||
@@ -66,13 +70,53 @@ def db_fetch(sql: str, *args) -> list[dict]:
|
||||
return asyncio.run(_db_fetch(sql, *args))
|
||||
|
||||
|
||||
def main() -> int:
|
||||
print(f"[config] BASE={BASE} EMAIL={ADMIN_EMAIL}")
|
||||
status, login = request_json(
|
||||
def encrypted_login(email: str, password: str) -> tuple[int, dict]:
|
||||
status, login_key = request_json("/api/v1/auth/login-key")
|
||||
if status != 200:
|
||||
return status, login_key
|
||||
public_key = serialization.load_pem_public_key(login_key["public_key"].encode("utf-8"))
|
||||
plaintext = json.dumps(
|
||||
{
|
||||
"email": email,
|
||||
"password": password,
|
||||
"challenge": login_key["challenge"],
|
||||
},
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
aes_key = AESGCM.generate_key(bit_length=256)
|
||||
iv = os.urandom(12)
|
||||
encrypted_data = AESGCM(aes_key).encrypt(iv, plaintext, None)
|
||||
encrypted_key = public_key.encrypt(
|
||||
aes_key,
|
||||
padding.OAEP(
|
||||
mgf=padding.MGF1(algorithm=hashes.SHA256()),
|
||||
algorithm=hashes.SHA256(),
|
||||
label=None,
|
||||
),
|
||||
)
|
||||
return request_json(
|
||||
"/api/v1/auth/login",
|
||||
method="POST",
|
||||
payload={"email": ADMIN_EMAIL, "password": ADMIN_PASSWORD},
|
||||
payload={
|
||||
"key_id": login_key["key_id"],
|
||||
"challenge": login_key["challenge"],
|
||||
"ciphertext": base64.b64encode(
|
||||
json.dumps(
|
||||
{
|
||||
"encrypted_key": base64.b64encode(encrypted_key).decode("ascii"),
|
||||
"iv": base64.b64encode(iv).decode("ascii"),
|
||||
"data": base64.b64encode(encrypted_data).decode("ascii"),
|
||||
},
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
).decode("ascii"),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
print(f"[config] BASE={BASE} EMAIL={ADMIN_EMAIL}")
|
||||
status, login = encrypted_login(ADMIN_EMAIL, ADMIN_PASSWORD)
|
||||
assert_or_exit(status == 200, f"登录失败 status={status} body={login}")
|
||||
token = login["access_token"]
|
||||
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
import pytest
|
||||
|
||||
from app.core import login_crypto
|
||||
from app.core.config import settings
|
||||
from app.main import create_app
|
||||
|
||||
|
||||
def reset_login_crypto_state():
|
||||
login_crypto._private_key = None
|
||||
login_crypto._public_key_pem = None
|
||||
login_crypto._challenges.clear()
|
||||
|
||||
|
||||
def test_production_requires_configured_login_private_key(monkeypatch):
|
||||
reset_login_crypto_state()
|
||||
monkeypatch.setattr(settings, "ENV", "production")
|
||||
monkeypatch.setattr(settings, "LOGIN_RSA_PRIVATE_KEY", None)
|
||||
|
||||
with pytest.raises(ValueError, match="生产环境必须配置 LOGIN_RSA_PRIVATE_KEY"):
|
||||
create_app()
|
||||
|
||||
reset_login_crypto_state()
|
||||
|
||||
|
||||
def test_login_challenge_cache_has_max_active_limit(monkeypatch):
|
||||
reset_login_crypto_state()
|
||||
monkeypatch.setattr(settings, "ENV", "test")
|
||||
monkeypatch.setattr(settings, "LOGIN_CHALLENGE_MAX_ACTIVE", 2)
|
||||
|
||||
first = login_crypto.create_login_challenge()
|
||||
second = login_crypto.create_login_challenge()
|
||||
third = login_crypto.create_login_challenge()
|
||||
|
||||
assert len(login_crypto._challenges) == 2
|
||||
assert not login_crypto.consume_login_challenge(first.value)
|
||||
assert login_crypto.consume_login_challenge(second.value)
|
||||
assert login_crypto.consume_login_challenge(third.value)
|
||||
|
||||
reset_login_crypto_state()
|
||||
@@ -1,7 +1,15 @@
|
||||
import pytest
|
||||
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 httpx import AsyncClient
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||
from sqlalchemy.ext.compiler import compiles
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
|
||||
from app.main import create_app
|
||||
from app.core.deps import get_db_session
|
||||
@@ -12,6 +20,55 @@ from app.models.user import User, UserRole, UserStatus
|
||||
TEST_DATABASE_URL = "sqlite+aiosqlite:///:memory:"
|
||||
|
||||
|
||||
@compiles(UUID, "sqlite")
|
||||
def compile_uuid_for_sqlite(_type, _compiler, **_kw):
|
||||
return "CHAR(32)"
|
||||
|
||||
|
||||
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
|
||||
login_key = key_resp.json()
|
||||
public_key = serialization.load_pem_public_key(login_key["public_key"].encode("utf-8"))
|
||||
plaintext = json.dumps(
|
||||
{
|
||||
"email": email,
|
||||
"password": password,
|
||||
"challenge": login_key["challenge"],
|
||||
},
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
aes_key = AESGCM.generate_key(bit_length=256)
|
||||
iv = os.urandom(12)
|
||||
encrypted_data = AESGCM(aes_key).encrypt(iv, plaintext, None)
|
||||
encrypted_key = public_key.encrypt(
|
||||
aes_key,
|
||||
padding.OAEP(
|
||||
mgf=padding.MGF1(algorithm=hashes.SHA256()),
|
||||
algorithm=hashes.SHA256(),
|
||||
label=None,
|
||||
),
|
||||
)
|
||||
return {
|
||||
"key_id": login_key["key_id"],
|
||||
"challenge": login_key["challenge"],
|
||||
"ciphertext": base64.b64encode(
|
||||
json.dumps(
|
||||
{
|
||||
"encrypted_key": base64.b64encode(encrypted_key).decode("ascii"),
|
||||
"iv": base64.b64encode(iv).decode("ascii"),
|
||||
"data": base64.b64encode(encrypted_data).decode("ascii"),
|
||||
},
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
).decode("ascii"),
|
||||
}
|
||||
|
||||
|
||||
async def encrypted_login(client: AsyncClient, email: str, password: str):
|
||||
return await client.post("/api/v1/auth/login", json=await encrypted_auth_payload(client, email, password))
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def client_and_db():
|
||||
engine = create_async_engine(TEST_DATABASE_URL, future=True)
|
||||
@@ -87,7 +144,7 @@ async def test_login_blocked_before_approval(client_and_db):
|
||||
"department": "Safety",
|
||||
}
|
||||
await client.post("/api/v1/auth/register", json=payload)
|
||||
resp = await client.post("/api/v1/auth/login", json={"email": payload["email"], "password": payload["password"]})
|
||||
resp = await encrypted_login(client, payload["email"], payload["password"])
|
||||
assert resp.status_code == 401
|
||||
assert "账号未审核" in resp.json().get("detail", "")
|
||||
|
||||
@@ -107,7 +164,7 @@ async def test_admin_can_approve_user(client_and_db):
|
||||
user = await user_crud.get_by_email(session, payload["email"])
|
||||
user_id = user.id
|
||||
|
||||
admin_login = await client.post("/api/v1/auth/login", json={"email": "admin@test.com", "password": "admin123"})
|
||||
admin_login = await encrypted_login(client, "admin@test.com", "admin123")
|
||||
token = admin_login.json()["access_token"]
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
|
||||
@@ -131,3 +188,71 @@ async def test_admin_role_cannot_register(client_and_db):
|
||||
}
|
||||
resp = await client.post("/api/v1/auth/register", json=payload)
|
||||
assert resp.status_code in (400, 422)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_plaintext_login_is_rejected(client_and_db):
|
||||
client, _ = client_and_db
|
||||
resp = await client.post("/api/v1/auth/login", json={"email": "admin@test.com", "password": "admin123"})
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_login_challenge_cannot_be_reused(client_and_db):
|
||||
client, _ = client_and_db
|
||||
key_resp = await client.get("/api/v1/auth/login-key")
|
||||
login_key = key_resp.json()
|
||||
public_key = serialization.load_pem_public_key(login_key["public_key"].encode("utf-8"))
|
||||
plaintext = json.dumps(
|
||||
{
|
||||
"email": "admin@test.com",
|
||||
"password": "admin123",
|
||||
"challenge": login_key["challenge"],
|
||||
},
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
aes_key = AESGCM.generate_key(bit_length=256)
|
||||
iv = os.urandom(12)
|
||||
encrypted_data = AESGCM(aes_key).encrypt(iv, plaintext, None)
|
||||
encrypted_key = public_key.encrypt(
|
||||
aes_key,
|
||||
padding.OAEP(
|
||||
mgf=padding.MGF1(algorithm=hashes.SHA256()),
|
||||
algorithm=hashes.SHA256(),
|
||||
label=None,
|
||||
)
|
||||
)
|
||||
ciphertext = base64.b64encode(
|
||||
json.dumps(
|
||||
{
|
||||
"encrypted_key": base64.b64encode(encrypted_key).decode("ascii"),
|
||||
"iv": base64.b64encode(iv).decode("ascii"),
|
||||
"data": base64.b64encode(encrypted_data).decode("ascii"),
|
||||
},
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
).decode("ascii")
|
||||
payload = {
|
||||
"key_id": login_key["key_id"],
|
||||
"challenge": login_key["challenge"],
|
||||
"ciphertext": ciphertext,
|
||||
}
|
||||
|
||||
first = await client.post("/api/v1/auth/login", json=payload)
|
||||
second = await client.post("/api/v1/auth/login", json=payload)
|
||||
|
||||
assert first.status_code == 200
|
||||
assert second.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unlock_requires_encrypted_password(client_and_db):
|
||||
client, _ = client_and_db
|
||||
plaintext = await client.post("/api/v1/auth/unlock", json={"email": "admin@test.com", "password": "admin123"})
|
||||
encrypted = await client.post(
|
||||
"/api/v1/auth/unlock",
|
||||
json=await encrypted_auth_payload(client, "admin@test.com", "admin123"),
|
||||
)
|
||||
|
||||
assert plaintext.status_code == 422
|
||||
assert encrypted.status_code == 200
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
from datetime import date
|
||||
from types import SimpleNamespace
|
||||
|
||||
from app.crud.subject import should_generate_visits_after_subject_update
|
||||
from app.crud.visit import build_visit_schedule_dates, sort_visits_for_display
|
||||
from app.schemas.study import StudyUpdate
|
||||
|
||||
|
||||
def test_build_visit_schedule_dates_uses_per_visit_windows():
|
||||
baseline_date = date(2026, 5, 8)
|
||||
|
||||
visits = build_visit_schedule_dates(
|
||||
[
|
||||
{
|
||||
"visit_code": "基线访视",
|
||||
"baseline_offset_days": 0,
|
||||
"window_before_days": 0,
|
||||
"window_after_days": 0,
|
||||
},
|
||||
{
|
||||
"visit_code": "V1",
|
||||
"baseline_offset_days": 7,
|
||||
"window_before_days": 2,
|
||||
"window_after_days": 2,
|
||||
},
|
||||
{
|
||||
"visit_code": "V2",
|
||||
"baseline_offset_days": 15,
|
||||
"window_before_days": 2,
|
||||
"window_after_days": 2,
|
||||
},
|
||||
{
|
||||
"visit_code": "V3",
|
||||
"baseline_offset_days": 23,
|
||||
"window_before_days": 3,
|
||||
"window_after_days": 3,
|
||||
},
|
||||
],
|
||||
baseline_date,
|
||||
)
|
||||
|
||||
assert [
|
||||
(visit["visit_code"], visit["baseline_offset_days"], visit["planned_date"], visit["window_start"], visit["window_end"])
|
||||
for visit in visits
|
||||
] == [
|
||||
("基线访视", 0, date(2026, 5, 8), date(2026, 5, 8), date(2026, 5, 8)),
|
||||
("V1", 7, date(2026, 5, 15), date(2026, 5, 13), date(2026, 5, 17)),
|
||||
("V2", 15, date(2026, 5, 23), date(2026, 5, 21), date(2026, 5, 25)),
|
||||
("V3", 23, date(2026, 5, 31), date(2026, 5, 28), date(2026, 6, 3)),
|
||||
]
|
||||
|
||||
|
||||
def test_build_visit_schedule_dates_keeps_visit_structure_without_baseline_date():
|
||||
visits = build_visit_schedule_dates(
|
||||
[
|
||||
{
|
||||
"visit_code": "基线访视",
|
||||
"baseline_offset_days": 0,
|
||||
"window_before_days": 0,
|
||||
"window_after_days": 0,
|
||||
},
|
||||
{
|
||||
"visit_code": "V1",
|
||||
"baseline_offset_days": 7,
|
||||
"window_before_days": 2,
|
||||
"window_after_days": 2,
|
||||
},
|
||||
],
|
||||
None,
|
||||
)
|
||||
|
||||
assert [
|
||||
(visit["visit_code"], visit["planned_date"], visit["window_start"], visit["window_end"])
|
||||
for visit in visits
|
||||
] == [
|
||||
("基线访视", None, None, None),
|
||||
("V1", None, None, None),
|
||||
]
|
||||
|
||||
|
||||
def test_study_update_allows_multiple_visits_with_same_baseline_offset():
|
||||
payload = StudyUpdate(
|
||||
visit_schedule=[
|
||||
{
|
||||
"visit_code": "筛选访视",
|
||||
"baseline_offset_days": 0,
|
||||
"window_before_days": 0,
|
||||
"window_after_days": 0,
|
||||
},
|
||||
{
|
||||
"visit_code": "基线访视",
|
||||
"baseline_offset_days": 0,
|
||||
"window_before_days": 0,
|
||||
"window_after_days": 0,
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
assert [item.visit_code for item in payload.visit_schedule or []] == ["筛选访视", "基线访视"]
|
||||
|
||||
|
||||
def test_sort_visits_for_display_follows_configured_visit_order_without_dates():
|
||||
visits = [
|
||||
SimpleNamespace(visit_code="V1", planned_date=None),
|
||||
SimpleNamespace(visit_code="V2", planned_date=None),
|
||||
SimpleNamespace(visit_code="基线访视", planned_date=None),
|
||||
]
|
||||
visit_schedule = [
|
||||
{"visit_code": "基线访视"},
|
||||
{"visit_code": "V1"},
|
||||
{"visit_code": "V2"},
|
||||
]
|
||||
|
||||
assert [visit.visit_code for visit in sort_visits_for_display(visits, visit_schedule)] == ["基线访视", "V1", "V2"]
|
||||
|
||||
|
||||
def test_sort_visits_for_display_does_not_infer_business_order():
|
||||
visits = [
|
||||
SimpleNamespace(visit_code="基线访视", planned_date=None),
|
||||
SimpleNamespace(visit_code="V1", planned_date=None),
|
||||
SimpleNamespace(visit_code="V2", planned_date=None),
|
||||
]
|
||||
visit_schedule = [
|
||||
{"visit_code": "V1"},
|
||||
{"visit_code": "V2"},
|
||||
{"visit_code": "基线访视"},
|
||||
]
|
||||
|
||||
assert [visit.visit_code for visit in sort_visits_for_display(visits, visit_schedule)] == ["V1", "V2", "基线访视"]
|
||||
|
||||
|
||||
def test_should_generate_visits_when_baseline_date_is_set_or_changed():
|
||||
assert should_generate_visits_after_subject_update(
|
||||
previous_baseline_date=None,
|
||||
next_baseline_date=date(2026, 5, 8),
|
||||
)
|
||||
assert should_generate_visits_after_subject_update(
|
||||
previous_baseline_date=date(2026, 5, 8),
|
||||
next_baseline_date=date(2026, 5, 9),
|
||||
)
|
||||
assert not should_generate_visits_after_subject_update(
|
||||
previous_baseline_date=None,
|
||||
next_baseline_date=None,
|
||||
)
|
||||
Reference in New Issue
Block a user