This commit is contained in:
Arc
2025-09-16 22:45:55 +01:00
parent 94ebc22dcc
commit 416d170996
+16 -32
View File
@@ -6,6 +6,7 @@ from typing import Any, Literal
from urllib.parse import urlencode from urllib.parse import urlencode
import httpx import httpx
from httpx import HTTPStatusError
from loguru import logger from loguru import logger
from pydantic import BaseModel, Field, ValidationError from pydantic import BaseModel, Field, ValidationError
@@ -43,24 +44,20 @@ class StripeCheckoutOptions(BaseModel):
line_item_name: str | None = None line_item_name: str | None = None
# === NEW: Direct-debit subscription options === # === Direct-debit subscription options ===
class StripeRecurringOptions(BaseModel): class StripeRecurringOptions(BaseModel):
class Config: class Config:
extra = "ignore" extra = "ignore"
# You will pass one of these (prefer price_id). We DO NOT create prices here.
price_id: str | None = None price_id: str | None = None
price_lookup_key: str | None = None # convenient if you use lookup keys in Stripe price_lookup_key: str | None = None
# Direct-debit rails to allow in Checkout (defaults to UK Bacs only)
# Use ["sepa_debit"] for EUR, or ["us_bank_account"] for US ACH.
payment_method_types: list[str] = Field(default_factory=lambda: ["bacs_debit"]) payment_method_types: list[str] = Field(default_factory=lambda: ["bacs_debit"])
# Optional niceties
success_url: str | None = None success_url: str | None = None
cancel_url: str | None = None
metadata: dict[str, str] = Field(default_factory=dict) metadata: dict[str, str] = Field(default_factory=dict)
customer_email: str | None = None customer_email: str | None = None
trial_days: int | None = None # Stripe supports trials on subs trial_days: int | None = None
class StripeCreateInvoiceOptions(BaseModel): class StripeCreateInvoiceOptions(BaseModel):
@@ -70,7 +67,6 @@ class StripeCreateInvoiceOptions(BaseModel):
fiat_method: FiatMethod = "checkout" fiat_method: FiatMethod = "checkout"
terminal: StripeTerminalOptions | None = None terminal: StripeTerminalOptions | None = None
checkout: StripeCheckoutOptions | None = None checkout: StripeCheckoutOptions | None = None
# NEW: when present we do mode=subscription with DD
recurring: StripeRecurringOptions | None = None recurring: StripeRecurringOptions | None = None
@@ -136,7 +132,6 @@ class StripeWallet(FiatProvider):
if not opts: if not opts:
return FiatInvoiceResponse(ok=False, error_message="Invalid Stripe options") return FiatInvoiceResponse(ok=False, error_message="Invalid Stripe options")
# Direct-debit subscriptions via Checkout (mode=subscription)
if opts.recurring is not None: if opts.recurring is not None:
return await self._create_subscription_checkout_session( return await self._create_subscription_checkout_session(
payment_hash, memo, opts payment_hash, memo, opts
@@ -196,7 +191,7 @@ class StripeWallet(FiatProvider):
r.raise_for_status() r.raise_for_status()
return r.json() return r.json()
# ---------- One-off Checkout (existing) ---------- # ---------- One-off Checkout ----------
async def _create_checkout_invoice( async def _create_checkout_invoice(
self, self,
amount_cents: int, amount_cents: int,
@@ -250,7 +245,7 @@ class StripeWallet(FiatProvider):
ok=False, error_message=f"Unable to connect to {self.endpoint}." ok=False, error_message=f"Unable to connect to {self.endpoint}."
) )
# ---------- Terminal (existing) ---------- # ---------- Terminal ----------
async def _create_terminal_invoice( async def _create_terminal_invoice(
self, self,
amount_cents: int, amount_cents: int,
@@ -293,7 +288,7 @@ class StripeWallet(FiatProvider):
ok=False, error_message=f"Unable to connect to {self.endpoint}." ok=False, error_message=f"Unable to connect to {self.endpoint}."
) )
# ---------- NEW: Direct-debit subscription via Checkout ---------- # ---------- Subscription Checkout ----------
async def _create_subscription_checkout_session( async def _create_subscription_checkout_session(
self, self,
payment_hash: str, payment_hash: str,
@@ -301,7 +296,6 @@ class StripeWallet(FiatProvider):
opts: StripeCreateInvoiceOptions, opts: StripeCreateInvoiceOptions,
) -> FiatInvoiceResponse: ) -> FiatInvoiceResponse:
rc = opts.recurring or StripeRecurringOptions() rc = opts.recurring or StripeRecurringOptions()
# Resolve a price_id (prefer explicit price_id; else lookup_key)
try: try:
price_id = rc.price_id price_id = rc.price_id
if not price_id and rc.price_lookup_key: if not price_id and rc.price_lookup_key:
@@ -315,26 +309,24 @@ class StripeWallet(FiatProvider):
or settings.stripe_payment_success_url or settings.stripe_payment_success_url
or "https://lnbits.com" or "https://lnbits.com"
) )
cancel_url = rc.cancel_url or success_url
form_data: list[tuple[str, str]] = [ form_data: list[tuple[str, str]] = [
("mode", "subscription"), ("mode", "subscription"),
("success_url", success_url), ("success_url", success_url),
("cancel_url", cancel_url),
("payment_method_collection", "always"),
("metadata[payment_hash]", payment_hash), ("metadata[payment_hash]", payment_hash),
("line_items[0][price]", price_id), ("line_items[0][price]", price_id),
("line_items[0][quantity]", "1"), ("line_items[0][quantity]", "1"),
] ]
# Allow only direct-debit rails (default: ["bacs_debit"])
for t in rc.payment_method_types:
form_data.append(("payment_method_types[]", t))
if rc.trial_days: if rc.trial_days:
form_data.append(("subscription_data[trial_period_days]", str(rc.trial_days))) form_data.append(("subscription_data[trial_period_days]", str(rc.trial_days)))
if rc.customer_email: if rc.customer_email:
form_data.append(("customer_email", rc.customer_email)) form_data.append(("customer_email", rc.customer_email))
# Attach arbitrary metadata (helps link invoices to your user)
form_data += self._encode_metadata("metadata", rc.metadata) form_data += self._encode_metadata("metadata", rc.metadata)
r = await self.client.post( r = await self.client.post(
@@ -351,21 +343,19 @@ class StripeWallet(FiatProvider):
) )
return FiatInvoiceResponse(ok=True, checking_id=session_id, payment_request=url) return FiatInvoiceResponse(ok=True, checking_id=session_id, payment_request=url)
except HTTPStatusError as e:
body = e.response.text if e.response is not None else "<no body>"
logger.warning(f"Stripe subscription 400: {body}")
return FiatInvoiceResponse(ok=False, error_message=body)
except json.JSONDecodeError: except json.JSONDecodeError:
return FiatInvoiceResponse(ok=False, error_message="Server error: invalid json response") return FiatInvoiceResponse(ok=False, error_message="Server error: invalid json response")
except Exception as exc: except Exception as exc:
logger.warning(exc) logger.warning(exc)
return FiatInvoiceResponse(ok=False, error_message=f"Unable to connect to {self.endpoint}.") return FiatInvoiceResponse(ok=False, error_message=f"Unable to connect to {self.endpoint}.")
# ---------- NEW: Fetch helpers (no creation) ---------- # ---------- Helpers ----------
async def _get_price_id_by_lookup_key(self, lookup_key: str) -> str | None: async def _get_price_id_by_lookup_key(self, lookup_key: str) -> str | None:
"""
Return the active price id for a given lookup_key, or None.
Tip: in Stripe dashboard set a unique lookup_key on your recurring price.
"""
# Stripe allows filtering prices by lookup_keys[]=<key>&active=true
params = {"active": "true", "expand[]": "data.product", "limit": "1"} params = {"active": "true", "expand[]": "data.product", "limit": "1"}
# passing array param:
qs = urlencode(params) + f"&lookup_keys[]={lookup_key}" qs = urlencode(params) + f"&lookup_keys[]={lookup_key}"
r = await self.client.get(f"/v1/prices?{qs}") r = await self.client.get(f"/v1/prices?{qs}")
r.raise_for_status() r.raise_for_status()
@@ -376,18 +366,13 @@ class StripeWallet(FiatProvider):
return items[0].get("id") return items[0].get("id")
async def list_prices_for_product(self, product_id: str) -> list[dict]: async def list_prices_for_product(self, product_id: str) -> list[dict]:
"""
List active recurring prices for a given product (handy for admin UI).
"""
qs = urlencode({"product": product_id, "active": "true", "limit": "100"}) qs = urlencode({"product": product_id, "active": "true", "limit": "100"})
r = await self.client.get(f"/v1/prices?{qs}") r = await self.client.get(f"/v1/prices?{qs}")
r.raise_for_status() r.raise_for_status()
data = r.json() data = r.json()
return (data or {}).get("data") or [] return (data or {}).get("data") or []
# ---------- utils ----------
def _normalize_stripe_id(self, checking_id: str) -> str: def _normalize_stripe_id(self, checking_id: str) -> str:
"""Remove our internal prefix so Stripe sees a real id."""
return ( return (
checking_id.replace("fiat_stripe_", "", 1) checking_id.replace("fiat_stripe_", "", 1)
if checking_id.startswith("fiat_stripe_") if checking_id.startswith("fiat_stripe_")
@@ -395,7 +380,6 @@ class StripeWallet(FiatProvider):
) )
def _status_from_checkout_session(self, data: dict) -> FiatPaymentStatus: def _status_from_checkout_session(self, data: dict) -> FiatPaymentStatus:
# For one-offs, "paid" means done; subs rely on webhooks (invoice.paid)
if data.get("payment_status") == "paid": if data.get("payment_status") == "paid":
return FiatPaymentSuccessStatus() return FiatPaymentSuccessStatus()