import json import logging from datetime import UTC, datetime import stripe from infrasynth.shared.enums import SubscriptionStatus from .base import BasePaymentGateway, CheckoutSessionResult, WebhookResult logger = logging.getLogger(__name__) class StripeGateway(BasePaymentGateway): """Stripe payment gateway. Configuration keys (read from ``PaymentGateway.config``): - ``api_key`` (required for most operations) - ``webhook_secret`` (required to verify inbound webhooks) """ gateway_slug = "stripe" STATUS_MAP = { "active": SubscriptionStatus.ACTIVE, "trialing": SubscriptionStatus.TRIALING, "past_due": SubscriptionStatus.PAST_DUE, "unpaid": SubscriptionStatus.PAST_DUE, "canceled": SubscriptionStatus.CANCELLED, "incomplete": SubscriptionStatus.PAST_DUE, "incomplete_expired": SubscriptionStatus.EXPIRED, } def __init__(self, config: dict | None = None): config = {str(key).lower(): value for key, value in (config or {}).items()} self.api_key = config.get("api_key") self.webhook_secret = config.get("webhook_secret") if self.api_key: stripe.api_key = self.api_key def create_checkout_session(self, plan, user=None, *, tenant=None, **kwargs) -> CheckoutSessionResult: self._require_credentials() if not plan.external_id: raise ValueError("Plan has no external price ID configured for Stripe") session = stripe.checkout.Session.create( mode="subscription", line_items=[{"price": plan.external_id, "quantity": 1}], success_url=kwargs.get("success_url") or "https://example.com/success", cancel_url=kwargs.get("cancel_url") or "https://example.com/cancel", customer_email=str(getattr(user, "email", "") or ""), metadata={ "plan_slug": plan.slug, "app_slug": getattr(getattr(plan, "app", None), "slug", ""), "user_id": str(getattr(user, "pk", "")), "tenant_id": str(getattr(tenant, "pk", "")), }, ) return CheckoutSessionResult( session_id=session.id, checkout_url=session.url or "", client_secret=session.client_secret or "", ) def handle_webhook(self, payload, headers) -> WebhookResult: if not self.webhook_secret: raise ValueError("Stripe webhook secret not configured") signature_header = headers.get("Stripe-Signature", "") raw_payload = json.dumps(payload) if isinstance(payload, dict) else payload event = stripe.Webhook.construct_event(raw_payload, signature_header, self.webhook_secret) return WebhookResult( event_type=event["type"], is_handled=True, data=event["data"]["object"], ) def cancel_subscription(self, subscription) -> bool: self._require_credentials() if not subscription.external_id: return False stripe.Subscription.cancel(subscription.external_id) return True def sync_subscription(self, subscription) -> dict: self._require_credentials() if not subscription.external_id: return {} data = stripe.Subscription.retrieve(subscription.external_id) return { "status": self.STATUS_MAP.get(data.get("status"), data.get("status")), "current_period_start": self._to_datetime(data.get("current_period_start")), "current_period_end": self._to_datetime(data.get("current_period_end")), "cancel_at_period_end": data.get("cancel_at_period_end", False), "cancelled_at": self._to_datetime(data.get("canceled_at")), "trial_end": self._to_datetime(data.get("trial_end")), "metadata": data.get("metadata") or {}, } def get_invoice(self, invoice) -> dict: self._require_credentials() if not invoice.external_id: return {} data = stripe.Invoice.retrieve(invoice.external_id) return { "external_id": data.get("id"), "status": data.get("status"), "amount": (data.get("amount_due") or 0) / 100, "currency": (data.get("currency") or "usd").upper(), "paid_at": self._to_datetime(data.get("paid_at")), "line_items": [ { "description": item.get("description"), "amount": (item.get("amount") or 0) / 100, "quantity": item.get("quantity"), } for item in data.get("lines", {}).get("data", []) ], } def health_check(self) -> bool: return bool(self.api_key) def _require_credentials(self) -> None: if not self.api_key: raise ValueError("Stripe API key not configured") @staticmethod def _to_datetime(timestamp) -> datetime | None: if not timestamp: return None return datetime.fromtimestamp(timestamp, tz=UTC)