from datetime import UTC, date, datetime import pytest from sqlalchemy import create_engine, event, select from sqlalchemy.orm import Session, sessionmaker from zbt.domains.enrollment.models import Policy from zbt.domains.surrender.models import ( RefundAttempt, SurrenderRequest, SurrenderTimelineEvent, ) from zbt.infrastructure.mysql.core_models import ( CoreBase, PolicyRecord, RefundAttemptRecord, SurrenderRequestRecord, SurrenderTimelineEventRecord, ) from zbt.infrastructure.mysql.refund_settlement import ( SqlAlchemyRefundSettlementWriter, ) NOW = datetime(2026, 7, 30, 10, 0, tzinfo=UTC) ATTEMPT_ID = "01KZ0000000000000000000201" REQUEST_ID = "01KZ0000000000000000000202" POLICY_ID = "01KZ0000000000000000000203" def session_factory() -> sessionmaker[Session]: engine = create_engine("sqlite+pysqlite:///:memory:") CoreBase.metadata.create_all(engine) factory = sessionmaker(bind=engine, expire_on_commit=False) with factory.begin() as session: session.add( PolicyRecord( id=POLICY_ID, policy_no="POL-20260730-TX01", order_id="01KZ0000000000000000000204", h5_user_id="01KZ0000000000000000000205", product_id="01KZ0000000000000000000206", product_version_id="01KZ0000000000000000000207", plan_id="01KZ0000000000000000000208", premium_cents=36_500, currency="CNY", coverage_start=date(2026, 7, 1), coverage_end=date(2027, 6, 30), status="SURRENDERING", issued_at=NOW, ) ) session.add( SurrenderRequestRecord( id=REQUEST_ID, request_no="ZBT-TB-20260730-2001", h5_user_id="01KZ0000000000000000000205", policy_id=POLICY_ID, idempotency_key="transaction-test", reason="保障调整", status="REFUND_PENDING", refund_amount_cents=30_000, currency="CNY", rule_version="surrender-v1", calculation_json={}, workflow_thread_id="transaction-thread", risk_reasons_json=[], workflow_plan_json=[], created_at=NOW, ) ) session.add( RefundAttemptRecord( id=ATTEMPT_ID, surrender_request_id=REQUEST_ID, attempt_no=1, refund_no="ZBT-RF-20260730-2001", provider="INTERNAL_REFUND", idempotency_key="transaction-refund", amount_cents=30_000, currency="CNY", status="SUBMITTED", created_at=NOW, updated_at=NOW, ) ) return factory def settlement_values() -> tuple[ RefundAttempt, SurrenderRequest, Policy, SurrenderTimelineEvent, ]: attempt = RefundAttempt( id=ATTEMPT_ID, surrender_request_id=REQUEST_ID, attempt_no=1, refund_no="ZBT-RF-20260730-2001", provider="INTERNAL_REFUND", idempotency_key="transaction-refund", amount_cents=30_000, currency="CNY", status="SUCCEEDED", provider_refund_no=None, error_code=None, error_message=None, created_at=NOW, updated_at=NOW, ) request = SurrenderRequest( id=REQUEST_ID, request_no="ZBT-TB-20260730-2001", user_id="01KZ0000000000000000000205", policy_id=POLICY_ID, idempotency_key="transaction-test", reason="保障调整", status="REFUNDED", refund_amount_cents=30_000, currency="CNY", rule_version="surrender-v1", calculation={}, workflow_thread_id="transaction-thread", created_at=NOW, completed_at=NOW, ) policy = Policy( id=POLICY_ID, policy_no="POL-20260730-TX01", order_id="01KZ0000000000000000000204", user_id="01KZ0000000000000000000205", product_id="01KZ0000000000000000000206", product_version_id="01KZ0000000000000000000207", plan_id="01KZ0000000000000000000208", premium_cents=36_500, currency="CNY", coverage_start=date(2026, 7, 1), coverage_end=date(2027, 6, 30), status="SURRENDERED", issued_at=NOW, ) timeline = SurrenderTimelineEvent( id="01KZ0000000000000000000209", surrender_request_id=REQUEST_ID, event_type="REFUND_SUCCEEDED", title="退款成功,保单已终止", detail={"refund_no": attempt.refund_no}, actor_type="SYSTEM", actor_id=None, idempotency_key=f"{REQUEST_ID}:REFUND_SUCCEEDED", created_at=NOW, ) return attempt, request, policy, timeline def test_refund_settlement_rolls_back_all_states_when_flush_fails() -> None: factory = session_factory() writer = SqlAlchemyRefundSettlementWriter(factory) attempt, request, policy, timeline = settlement_values() def fail_flush( session: Session, flush_context: object, instances: object, ) -> None: del session, flush_context, instances raise RuntimeError("模拟时间线写入失败") event.listen(factory.class_, "before_flush", fail_flush) try: with pytest.raises(RuntimeError, match="模拟时间线写入失败"): writer.complete( attempt=attempt, request=request, policy=policy, timeline=timeline, ) finally: event.remove(factory.class_, "before_flush", fail_flush) with factory() as session: assert session.get(RefundAttemptRecord, ATTEMPT_ID).status == "SUBMITTED" assert session.get(SurrenderRequestRecord, REQUEST_ID).status == ( "REFUND_PENDING" ) assert session.get(PolicyRecord, POLICY_ID).status == "SURRENDERING" assert session.scalars(select(SurrenderTimelineEventRecord)).all() == []