from sqlalchemy import select from app.domains.enrollment.models import ( AsyncTask, EnrollmentDraft, EnrollmentOrder, OutboxEvent, PaymentTransaction, Policy, Quote, UserConfirmation, ) from app.domains.enrollment.repository import EnrollmentRepository from app.infrastructure.mysql.core_models import ( AsyncTaskRecord, EnrollmentDraftRecord, EnrollmentOrderRecord, OutboxEventRecord, PaymentTransactionRecord, PolicyRecord, QuoteRecord, UserConfirmationRecord, ) from app.infrastructure.mysql.repositories import SessionFactory class SqlAlchemyEnrollmentRepository(EnrollmentRepository): def __init__(self, session_factory: SessionFactory) -> None: self._session_factory = session_factory def save_quote(self, quote: Quote) -> None: with self._session_factory() as session: session.merge( QuoteRecord( id=quote.id, h5_user_id=quote.user_id, product_id=quote.product_id, product_version_id=quote.product_version_id, plan_id=quote.plan_id, insured_age=quote.insured_age, insured_region_code=quote.insured_region_code, occupation_code=quote.occupation_code, relationship=quote.relationship, premium_cents=quote.premium_cents, currency=quote.currency, rule_version=quote.rule_version, rate_version=quote.rate_version, status=quote.status, expires_at=quote.expires_at, created_at=quote.created_at, ) ) session.commit() def get_quote(self, quote_id: str) -> Quote | None: with self._session_factory() as session: record = session.get(QuoteRecord, quote_id) return self._quote(record) if record else None def save_draft(self, draft: EnrollmentDraft) -> None: with self._session_factory() as session: session.merge( EnrollmentDraftRecord( id=draft.id, h5_user_id=draft.user_id, quote_id=draft.quote_id, applicant_json=draft.applicant, insured_json=draft.insured, contact_json=draft.contact, status=draft.status, expires_at=draft.expires_at, created_at=draft.created_at, ) ) session.commit() def get_draft(self, draft_id: str) -> EnrollmentDraft | None: with self._session_factory() as session: record = session.get(EnrollmentDraftRecord, draft_id) return self._draft(record) if record else None def save_confirmation(self, confirmation: UserConfirmation) -> None: with self._session_factory() as session: session.merge( UserConfirmationRecord( id=confirmation.id, h5_user_id=confirmation.user_id, draft_id=confirmation.draft_id, token_hash=confirmation.token_hash, status=confirmation.status, expires_at=confirmation.expires_at, created_at=confirmation.created_at, ) ) session.commit() def get_confirmation_by_token_hash(self, token_hash: str) -> UserConfirmation | None: with self._session_factory() as session: record = session.scalar( select(UserConfirmationRecord).where( UserConfirmationRecord.token_hash == token_hash ) ) return self._confirmation(record) if record else None def save_order(self, order: EnrollmentOrder) -> None: with self._session_factory() as session: session.merge( EnrollmentOrderRecord( id=order.id, order_no=order.order_no, h5_user_id=order.user_id, quote_id=order.quote_id, draft_id=order.draft_id, confirmation_id=order.confirmation_id, idempotency_key=order.idempotency_key, product_id=order.product_id, product_version_id=order.product_version_id, plan_id=order.plan_id, applicant_snapshot_json=order.applicant_snapshot, insured_snapshot_json=order.insured_snapshot, amount_cents=order.amount_cents, currency=order.currency, status=order.status, created_at=order.created_at, ) ) session.commit() def get_order(self, order_id: str) -> EnrollmentOrder | None: with self._session_factory() as session: record = session.get(EnrollmentOrderRecord, order_id) return self._order(record) if record else None def find_order_by_idempotency( self, user_id: str, idempotency_key: str ) -> EnrollmentOrder | None: with self._session_factory() as session: record = session.scalar( select(EnrollmentOrderRecord).where( EnrollmentOrderRecord.h5_user_id == user_id, EnrollmentOrderRecord.idempotency_key == idempotency_key, ) ) return self._order(record) if record else None def list_orders(self, user_id: str | None = None) -> list[EnrollmentOrder]: with self._session_factory() as session: statement = select(EnrollmentOrderRecord).order_by( EnrollmentOrderRecord.created_at.desc() ) if user_id is not None: statement = statement.where(EnrollmentOrderRecord.h5_user_id == user_id) return [self._order(record) for record in session.scalars(statement)] def save_payment(self, payment: PaymentTransaction) -> None: with self._session_factory() as session: session.merge( PaymentTransactionRecord( id=payment.id, payment_no=payment.payment_no, order_id=payment.order_id, h5_user_id=payment.user_id, transaction_type=payment.transaction_type, provider=payment.provider, idempotency_key=payment.idempotency_key, amount_cents=payment.amount_cents, currency=payment.currency, status=payment.status, provider_transaction_no=payment.provider_transaction_no, succeeded_at=payment.succeeded_at, created_at=payment.created_at, ) ) session.commit() def get_payment(self, payment_id: str) -> PaymentTransaction | None: with self._session_factory() as session: record = session.get(PaymentTransactionRecord, payment_id) return self._payment(record) if record else None def get_payment_by_no(self, payment_no: str) -> PaymentTransaction | None: with self._session_factory() as session: record = session.scalar( select(PaymentTransactionRecord).where( PaymentTransactionRecord.payment_no == payment_no ) ) return self._payment(record) if record else None def get_payment_by_order(self, order_id: str) -> PaymentTransaction | None: with self._session_factory() as session: record = session.scalar( select(PaymentTransactionRecord).where( PaymentTransactionRecord.order_id == order_id ) ) return self._payment(record) if record else None def find_payment_by_idempotency( self, order_id: str, idempotency_key: str ) -> PaymentTransaction | None: with self._session_factory() as session: record = session.scalar( select(PaymentTransactionRecord).where( PaymentTransactionRecord.order_id == order_id, PaymentTransactionRecord.idempotency_key == idempotency_key, ) ) return self._payment(record) if record else None def save_policy(self, policy: Policy) -> None: with self._session_factory() as session: session.merge( PolicyRecord( id=policy.id, policy_no=policy.policy_no, order_id=policy.order_id, h5_user_id=policy.user_id, product_id=policy.product_id, product_version_id=policy.product_version_id, plan_id=policy.plan_id, premium_cents=policy.premium_cents, currency=policy.currency, coverage_start=policy.coverage_start, coverage_end=policy.coverage_end, status=policy.status, issued_at=policy.issued_at, ) ) session.commit() def get_policy_by_order(self, order_id: str) -> Policy | None: with self._session_factory() as session: record = session.scalar(select(PolicyRecord).where(PolicyRecord.order_id == order_id)) return self._policy(record) if record else None def list_policies(self, user_id: str | None = None) -> list[Policy]: with self._session_factory() as session: statement = select(PolicyRecord).order_by(PolicyRecord.issued_at.desc()) if user_id is not None: statement = statement.where(PolicyRecord.h5_user_id == user_id) return [self._policy(record) for record in session.scalars(statement)] def save_outbox_event(self, event: OutboxEvent) -> None: with self._session_factory() as session: session.merge( OutboxEventRecord( id=event.id, aggregate_type=event.aggregate_type, aggregate_id=event.aggregate_id, event_type=event.event_type, payload_json=event.payload, status=event.status, created_at=event.created_at, ) ) session.commit() def save_task(self, task: AsyncTask) -> None: with self._session_factory() as session: session.merge( AsyncTaskRecord( id=task.id, task_type=task.task_type, business_key=task.business_key, idempotency_key=task.idempotency_key, payload_json=task.payload, status=task.status, attempt_count=task.attempt_count, created_at=task.created_at, ) ) session.commit() def get_task_by_idempotency(self, idempotency_key: str) -> AsyncTask | None: with self._session_factory() as session: record = session.scalar( select(AsyncTaskRecord).where(AsyncTaskRecord.idempotency_key == idempotency_key) ) if record is None: return None return AsyncTask( id=record.id, task_type=record.task_type, business_key=record.business_key, idempotency_key=record.idempotency_key, payload=dict(record.payload_json), status=record.status, attempt_count=record.attempt_count, created_at=record.created_at, ) @staticmethod def _quote(record: QuoteRecord) -> Quote: return Quote( id=record.id, user_id=record.h5_user_id, product_id=record.product_id, product_version_id=record.product_version_id, plan_id=record.plan_id, insured_age=record.insured_age, insured_region_code=record.insured_region_code, occupation_code=record.occupation_code, relationship=record.relationship, premium_cents=record.premium_cents, currency=record.currency, rule_version=record.rule_version, rate_version=record.rate_version, status=record.status, expires_at=record.expires_at, created_at=record.created_at, ) @staticmethod def _draft(record: EnrollmentDraftRecord) -> EnrollmentDraft: return EnrollmentDraft( id=record.id, user_id=record.h5_user_id, quote_id=record.quote_id, applicant=dict(record.applicant_json), insured=dict(record.insured_json), contact=dict(record.contact_json), status=record.status, expires_at=record.expires_at, created_at=record.created_at, ) @staticmethod def _confirmation(record: UserConfirmationRecord) -> UserConfirmation: return UserConfirmation( id=record.id, user_id=record.h5_user_id, draft_id=record.draft_id, token_hash=record.token_hash, status=record.status, expires_at=record.expires_at, created_at=record.created_at, ) @staticmethod def _order(record: EnrollmentOrderRecord) -> EnrollmentOrder: return EnrollmentOrder( id=record.id, order_no=record.order_no, user_id=record.h5_user_id, quote_id=record.quote_id, draft_id=record.draft_id, confirmation_id=record.confirmation_id, idempotency_key=record.idempotency_key, product_id=record.product_id, product_version_id=record.product_version_id, plan_id=record.plan_id, applicant_snapshot=dict(record.applicant_snapshot_json), insured_snapshot=dict(record.insured_snapshot_json), amount_cents=record.amount_cents, currency=record.currency, status=record.status, created_at=record.created_at, ) @staticmethod def _payment(record: PaymentTransactionRecord) -> PaymentTransaction: return PaymentTransaction( id=record.id, payment_no=record.payment_no, order_id=record.order_id, user_id=record.h5_user_id, transaction_type=record.transaction_type, provider=record.provider, idempotency_key=record.idempotency_key, amount_cents=record.amount_cents, currency=record.currency, status=record.status, provider_transaction_no=record.provider_transaction_no, succeeded_at=record.succeeded_at, created_at=record.created_at, ) @staticmethod def _policy(record: PolicyRecord) -> Policy: return Policy( id=record.id, policy_no=record.policy_no, order_id=record.order_id, user_id=record.h5_user_id, product_id=record.product_id, product_version_id=record.product_version_id, plan_id=record.plan_id, premium_cents=record.premium_cents, currency=record.currency, coverage_start=record.coverage_start, coverage_end=record.coverage_end, status=record.status, issued_at=record.issued_at, )