feat(approval): add task workflow and waiver decisions
This commit is contained in:
237
server/tests/test_approval_task_concurrency_postgres.py
Normal file
237
server/tests/test_approval_task_concurrency_postgres.py
Normal file
@@ -0,0 +1,237 @@
|
||||
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
|
||||
Reference in New Issue
Block a user