125 lines
4.8 KiB
Python
125 lines
4.8 KiB
Python
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, **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, "user_id": str(getattr(user, "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)
|