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:
Cheng Zhou
2026-05-08 22:13:12 +08:00
parent a7bbcaa5dc
commit 74feca4467
47 changed files with 2423 additions and 534 deletions
+64 -30
View File
@@ -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)
+1 -28
View File
@@ -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)
+3 -3
View File
@@ -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}"
+2 -5
View File
@@ -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,
+6 -1
View File
@@ -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
+138
View File
@@ -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
+1 -6
View File
@@ -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
View File
@@ -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
View File
@@ -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
+2
View File
@@ -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="临床试验项目管理系统后端接口文档",
+2 -7
View File
@@ -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())
+1
View File
@@ -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())
+43 -20
View File
@@ -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
+22 -7
View File
@@ -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):
+4
View File
@@ -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
+3 -3
View File
@@ -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