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
@@ -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
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
+1
View File
@@ -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
+48 -4
View File
@@ -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"]
+39
View File
@@ -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()
+127 -2
View File
@@ -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
+144
View File
@@ -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,
)