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:
+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
|
||||
|
||||
Reference in New Issue
Block a user