Files
X-Financial/server/tests/test_expense_claim_action_protocol.py

323 lines
11 KiB
Python
Raw Normal View History

from __future__ import annotations
from concurrent.futures import ThreadPoolExecutor
from datetime import UTC, datetime
from decimal import Decimal
import pytest
from sqlalchemy import create_engine, func, select
from sqlalchemy.orm import Session, sessionmaker
from app.api.deps import CurrentUserContext
from app.db.base import Base
from app.models.approval_action import ApprovalActionLedger
from app.models.audit_log import AuditLog
from app.models.employee import Employee
from app.models.expense_case import BusinessEvent
from app.models.financial_record import ExpenseClaim
from app.services.approval_action_protocol import (
ApprovalActionConflictError,
ApprovalActionProtocol,
)
from app.services.expense_claims import ExpenseClaimService
def _manager_user() -> CurrentUserContext:
return CurrentUserContext(
username="manager-action@example.com",
name="李经理",
role_codes=["manager"],
is_admin=False,
)
def _finance_user() -> CurrentUserContext:
return CurrentUserContext(
username="finance-action@example.com",
name="王财务",
role_codes=["finance"],
is_admin=False,
)
def _seed_claim(
db: Session,
*,
claim_id: str = "claim-action-1",
manager_email: str = "manager-action@example.com",
) -> ExpenseClaim:
manager = Employee(
id=f"manager-{claim_id}",
employee_no=f"M-{claim_id}",
name="李经理",
email=manager_email,
)
employee = Employee(
id=f"employee-{claim_id}",
employee_no=f"E-{claim_id}",
name="张三",
email=f"employee-{claim_id}@example.com",
manager=manager,
)
claim = ExpenseClaim(
id=claim_id,
claim_no=f"EXP-{claim_id}",
employee=employee,
employee_name="张三",
department_name="市场部",
expense_type="transport",
reason="客户拜访",
location="上海",
amount=Decimal("88.00"),
currency="CNY",
invoice_count=1,
occurred_at=datetime(2026, 7, 16, tzinfo=UTC),
submitted_at=datetime(2026, 7, 16, tzinfo=UTC),
status="submitted",
approval_stage="直属领导审批",
risk_flags_json=[],
)
db.add(claim)
db.commit()
return claim
@pytest.fixture
def session_factory() -> sessionmaker[Session]:
engine = create_engine("sqlite+pysqlite:///:memory:")
Base.metadata.create_all(engine)
factory = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
try:
yield factory
finally:
engine.dispose()
def test_approve_replay_persists_one_ledger_event_and_audit(
session_factory: sessionmaker[Session],
) -> None:
with session_factory() as db:
claim = _seed_claim(db)
service = ExpenseClaimService(db)
first = service.approve_claim(
claim.id,
_manager_user(),
opinion="同意",
request_id="approve-retry-1",
expected_status="submitted",
expected_approval_stage="直属领导审批",
)
replay = service.approve_claim(
claim.id,
_manager_user(),
opinion="同意",
request_id="approve-retry-1",
expected_status="submitted",
expected_approval_stage="直属领导审批",
)
assert first is not None and replay is not None
assert replay.approval_stage == "财务审批"
assert db.scalar(select(func.count()).select_from(ApprovalActionLedger)) == 1
assert db.scalar(select(func.count()).select_from(BusinessEvent)) == 1
assert (
db.scalar(
select(func.count())
.select_from(AuditLog)
.where(AuditLog.request_id == "approve-retry-1")
)
== 1
)
def test_request_id_payload_mismatch_and_stale_preconditions_return_conflict(
session_factory: sessionmaker[Session],
) -> None:
with session_factory() as db:
claim = _seed_claim(db, claim_id="claim-action-conflict")
service = ExpenseClaimService(db)
with pytest.raises(ApprovalActionConflictError, match="单据状态已从"):
service.approve_claim(
claim.id,
_manager_user(),
opinion="同意",
request_id="approve-stale-1",
expected_status="draft",
expected_approval_stage="直属领导审批",
)
service.approve_claim(
claim.id,
_manager_user(),
opinion="同意",
request_id="approve-conflict-1",
expected_status="submitted",
expected_approval_stage="直属领导审批",
)
with pytest.raises(ApprovalActionConflictError, match="已用于另一项"):
service.approve_claim(
claim.id,
_manager_user(),
opinion="改为有条件通过",
request_id="approve-conflict-1",
expected_status="submitted",
expected_approval_stage="直属领导审批",
)
assert (
db.scalar(
select(func.count())
.select_from(ApprovalActionLedger)
.where(ApprovalActionLedger.request_id == "approve-stale-1")
)
== 0
)
def test_action_failure_rolls_back_ledger_claim_event_and_audit(
session_factory: sessionmaker[Session],
monkeypatch: pytest.MonkeyPatch,
) -> None:
with session_factory() as db:
claim = _seed_claim(db, claim_id="claim-action-rollback")
service = ExpenseClaimService(db)
def fail_completion(*args, **kwargs):
raise RuntimeError("ledger completion failed")
monkeypatch.setattr(ApprovalActionProtocol, "complete", fail_completion)
with pytest.raises(RuntimeError, match="ledger completion failed"):
service.approve_claim(
claim.id,
_manager_user(),
opinion="同意",
request_id="approve-rollback-1",
expected_status="submitted",
expected_approval_stage="直属领导审批",
)
db.expire_all()
persisted = db.get(ExpenseClaim, claim.id)
assert persisted is not None
assert persisted.status == "submitted"
assert persisted.approval_stage == "直属领导审批"
assert db.scalar(select(func.count()).select_from(ApprovalActionLedger)) == 0
assert db.scalar(select(func.count()).select_from(BusinessEvent)) == 0
assert db.scalar(select(func.count()).select_from(AuditLog)) == 0
def test_legacy_stage_repair_cannot_commit_inside_action_protocol(
session_factory: sessionmaker[Session],
) -> None:
with session_factory() as db:
claim = _seed_claim(db, claim_id="claim-action-stage-repair")
claim.approval_stage = "预算管理者审批"
claim.risk_flags_json = [
{
"source": "manual_approval",
"event_type": "expense_claim_approval",
"previous_approval_stage": "直属领导审批",
"next_approval_stage": "预算管理者审批",
"operator": "李经理",
"next_approver_name": "李经理",
}
]
db.commit()
admin_user = CurrentUserContext(
username="admin-action@example.com",
name="审批管理员",
role_codes=["admin"],
is_admin=True,
)
with pytest.raises(ApprovalActionConflictError, match="审批节点已从"):
ExpenseClaimService(db).approve_claim(
claim.id,
admin_user,
opinion="同意",
request_id="approve-stage-repair-1",
expected_status="submitted",
expected_approval_stage="预算管理者审批",
)
db.expire_all()
persisted = db.get(ExpenseClaim, claim.id)
assert persisted is not None
assert persisted.approval_stage == "预算管理者审批"
assert db.scalar(select(func.count()).select_from(ApprovalActionLedger)) == 0
def test_return_and_pay_actions_use_the_same_protocol(
session_factory: sessionmaker[Session],
) -> None:
with session_factory() as db:
returned_claim = _seed_claim(db, claim_id="claim-action-return")
paid_claim = _seed_claim(
db,
claim_id="claim-action-pay",
manager_email="manager-pay-action@example.com",
)
paid_claim.status = "pending_payment"
paid_claim.approval_stage = "待付款"
db.commit()
returned = ExpenseClaimService(db).return_claim(
returned_claim.id,
_manager_user(),
reason="请补充材料",
request_id="return-action-1",
expected_status="submitted",
expected_approval_stage="直属领导审批",
)
paid = ExpenseClaimService(db).mark_claim_paid(
paid_claim.id,
_finance_user(),
request_id="pay-action-1",
expected_status="pending_payment",
expected_approval_stage="待付款",
)
assert returned is not None and returned.status == "returned"
assert paid is not None and paid.status == "paid"
ledgers = list(
db.scalars(select(ApprovalActionLedger).order_by(ApprovalActionLedger.action)).all()
)
assert [(item.action, item.result_status) for item in ledgers] == [
("pay", "paid"),
("return", "returned"),
]
def test_concurrent_identical_request_executes_once(tmp_path) -> None:
engine = create_engine(
f"sqlite+pysqlite:///{tmp_path / 'approval-action.db'}",
connect_args={"check_same_thread": False},
)
Base.metadata.create_all(engine)
factory = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
with factory() as db:
_seed_claim(db, claim_id="claim-action-concurrent")
def approve() -> str:
with factory() as db:
result = ExpenseClaimService(db).approve_claim(
"claim-action-concurrent",
_manager_user(),
opinion="同意",
request_id="approve-concurrent-1",
expected_status="submitted",
expected_approval_stage="直属领导审批",
)
assert result is not None
return str(result.approval_stage)
try:
with ThreadPoolExecutor(max_workers=2) as pool:
results = list(pool.map(lambda _: approve(), range(2)))
assert results == ["财务审批", "财务审批"]
with factory() as db:
assert db.scalar(select(func.count()).select_from(ApprovalActionLedger)) == 1
assert db.scalar(select(func.count()).select_from(BusinessEvent)) == 1
finally:
engine.dispose()