from __future__ import annotations import os import uuid from concurrent.futures import ThreadPoolExecutor from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest from sqlalchemy import create_engine, func, select from sqlalchemy.engine import make_url from sqlalchemy.orm import Session, sessionmaker from app.api.deps import CurrentUserContext from app.db.base import Base from app.models.approval_task import ApprovalTask, ApprovalTaskEvent from app.models.employee import Employee from app.models.financial_record import ExpenseClaim from app.models.role import Role from app.schemas.approval_task import ApprovalTaskAssignmentAction from app.services.approval_task_actions import ApprovalTaskActionService from app.services.approval_task_lifecycle import ApprovalTaskLifecycleService from app.services.approval_task_protocol import ApprovalTaskVersionConflictError from app.services.expense_cases import ExpenseCaseService DATABASE_URL = os.environ.get("MIGRATION_TEST_DATABASE_URL", "").strip() TENANT_ID = "tenant-approval-task-concurrency" def test_concurrent_identical_delegation_replays_one_immutable_event() -> None: database_url = _require_disposable_database_url() engine = create_engine(database_url, pool_pre_ping=True) Base.metadata.create_all(bind=engine) factory = sessionmaker(bind=engine, expire_on_commit=False) suffix = uuid.uuid4().hex[:12] try: task_id, manager, delegate_ids = _seed_case(factory, suffix=suffix, delegates=1) payload = ApprovalTaskAssignmentAction( request_id=f"delegate-concurrent-{suffix}", expected_task_version=1, reason="并发幂等委托验证", target_employee_id=delegate_ids[0], expires_at=datetime.now(UTC) + timedelta(days=1), ) def delegate_once(): with factory() as db: return ApprovalTaskActionService(db).assign( task_id, manager, action="delegate", payload=payload, ) with ThreadPoolExecutor(max_workers=2) as pool: futures = (pool.submit(delegate_once), pool.submit(delegate_once)) responses = [future.result(timeout=10) for future in futures] assert sorted(response.replayed for response in responses) == [False, True] assert responses[0].task.model_dump(mode="json") == responses[1].task.model_dump( mode="json" ) with factory() as db: task = db.get(ApprovalTask, task_id) assert task is not None and task.version == 2 assert task.assignee_employee_id == delegate_ids[0] assert ( db.scalar( select(func.count()) .select_from(ApprovalTaskEvent) .where( ApprovalTaskEvent.task_id == task_id, ApprovalTaskEvent.request_id == payload.request_id, ) ) == 1 ) finally: engine.dispose() def test_concurrent_distinct_delegations_allow_only_expected_version_winner() -> None: database_url = _require_disposable_database_url() engine = create_engine(database_url, pool_pre_ping=True) Base.metadata.create_all(bind=engine) factory = sessionmaker(bind=engine, expire_on_commit=False) suffix = uuid.uuid4().hex[:12] try: task_id, manager, delegate_ids = _seed_case(factory, suffix=suffix, delegates=2) def delegate_once(index: int): with factory() as db: try: return ApprovalTaskActionService(db).assign( task_id, manager, action="delegate", payload=ApprovalTaskAssignmentAction( request_id=f"delegate-race-{suffix}-{index}", expected_task_version=1, reason="并发版本竞争验证", target_employee_id=delegate_ids[index], expires_at=datetime.now(UTC) + timedelta(days=1), ), ) except ApprovalTaskVersionConflictError as error: return error with ThreadPoolExecutor(max_workers=2) as pool: futures = (pool.submit(delegate_once, 0), pool.submit(delegate_once, 1)) responses = [future.result(timeout=10) for future in futures] assert sum(not isinstance(item, Exception) for item in responses) == 1 assert sum(isinstance(item, ApprovalTaskVersionConflictError) for item in responses) == 1 with factory() as db: task = db.get(ApprovalTask, task_id) assert task is not None and task.version == 2 assert task.assignee_employee_id in set(delegate_ids) assert ( db.scalar( select(func.count()) .select_from(ApprovalTaskEvent) .where( ApprovalTaskEvent.task_id == task_id, ApprovalTaskEvent.event_type == "task_delegated", ) ) == 1 ) finally: engine.dispose() def _seed_case( factory: sessionmaker[Session], *, suffix: str, delegates: int, ) -> tuple[str, CurrentUserContext, list[str]]: with factory() as db: role = db.scalar(select(Role).where(Role.role_code == "manager")) if role is None: role = Role( id="role-appr-task-concur-manager", role_code="manager", name="并发审批经理", ) manager = Employee( id=f"manager-{suffix}", employee_no=f"M-{suffix}", name="并发审批经理", email=f"manager-{suffix}@example.com", roles=[role], ) delegate_rows = [ Employee( id=f"delegate-{suffix}-{index}", employee_no=f"D-{suffix}-{index}", name=f"委托审批人{index + 1}", email=f"delegate-{suffix}-{index}@example.com", roles=[role], ) for index in range(delegates) ] claimant = Employee( id=f"claimant-{suffix}", employee_no=f"E-{suffix}", name="并发报销申请人", email=f"claimant-{suffix}@example.com", manager=manager, ) occurred_at = datetime.now(UTC) - timedelta(hours=1) claim = ExpenseClaim( id=f"claim-{suffix}", claim_no=f"RE-CONCURRENT-{suffix}", employee=claimant, employee_name=claimant.name, department_name="并发验证部", expense_type="transport", reason="审批任务并发验证", location="上海", amount=Decimal("88.00"), currency="CNY", invoice_count=1, occurred_at=occurred_at, submitted_at=occurred_at, status="submitted", approval_stage="直属领导审批", risk_flags_json=[], ) db.add_all([claim, *delegate_rows]) db.flush() _, business_event = ExpenseCaseService(db).record_claim_event( claim, event_type="claim_submitted", actor_id=claimant.email, tenant_id=TENANT_ID, idempotency_key=f"submit-{suffix}", previous_status="draft", previous_approval_stage="待提交", ) task = ApprovalTaskLifecycleService(db).ensure_root_task( claim, tenant_id=TENANT_ID, entered_at=business_event.occurred_at, entered_at_source="workflow_event", business_event=business_event, request_id=f"node-enter-{suffix}", ) assert task is not None db.commit() return ( task.id, CurrentUserContext( username=manager.email, name=manager.name, role_codes=["manager"], is_admin=False, tenant_id=TENANT_ID, employee_id=manager.id, employee_no=manager.employee_no, ), [row.id for row in delegate_rows], ) def _require_disposable_database_url() -> str: if not DATABASE_URL: pytest.skip("仅在显式配置 MIGRATION_TEST_DATABASE_URL 时运行 PostgreSQL 并发测试") parsed = make_url(DATABASE_URL) host = str(parsed.host or "").replace("_", "-").lower() database = str(parsed.database or "").replace("_", "-").lower() if not host.startswith(("migration-probe", "disposable-probe")): raise RuntimeError("并发测试数据库主机必须使用 disposable 前缀") if not database.startswith(("migration-probe", "disposable-probe")): raise RuntimeError("并发测试数据库名必须使用 disposable 前缀") return DATABASE_URL