Coverage for src/ai_lls_lib/payment/webhook_processor.py: 89%

198 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-08-24 12:44 +0000

1"""Stripe webhook event processing.""" 

2 

3import base64 

4import json 

5import logging 

6from datetime import UTC, datetime 

7from decimal import Decimal 

8from typing import Any 

9 

10try: 

11 import stripe 

12 

13 HAS_STRIPE = True 

14except ImportError: 

15 stripe = None # type: ignore[assignment] 

16 HAS_STRIPE = False 

17 

18try: 

19 from botocore.exceptions import ClientError 

20 

21 HAS_BOTOCORE = True 

22except ImportError: 

23 HAS_BOTOCORE = False 

24 

25from .credit_manager import CreditManager 

26 

27logger = logging.getLogger(__name__) 

28 

29# Subscriptions grant no credits. An active subscription_status makes check_and_deduct 

30# return unlimited=True and skip deduction entirely, so any balance granted here could 

31# never be spent while the subscription is active (issue #134). 

32 

33 

34class WebhookProcessor: 

35 """Process Stripe webhook events.""" 

36 

37 def __init__(self, webhook_secret: str, credit_manager: CreditManager): 

38 """Initialize with webhook secret and credit manager.""" 

39 self.webhook_secret = webhook_secret 

40 self.credit_manager = credit_manager 

41 

42 def verify_and_parse(self, payload: str, signature: str) -> dict[str, Any]: 

43 """Verify webhook signature and parse event.""" 

44 if not HAS_STRIPE or not stripe: 

45 raise ImportError("stripe package not installed") 

46 

47 try: 

48 event = stripe.Webhook.construct_event(payload, signature, self.webhook_secret) 

49 return dict(event) 

50 except ValueError as e: 

51 logger.error(f"Invalid webhook payload: {e}") 

52 raise 

53 except stripe.error.SignatureVerificationError as e: 

54 logger.error(f"Invalid webhook signature: {e}") 

55 raise 

56 

57 def process_event(self, event: dict[str, Any]) -> dict[str, Any]: 

58 """ 

59 Process a verified webhook event. 

60 Returns response data. 

61 """ 

62 event_type = event.get("type") 

63 event_data = event.get("data", {}).get("object", {}) 

64 

65 logger.info(f"Processing webhook event: {event_type}") 

66 

67 if event_type == "payment_intent.succeeded": 

68 return self._handle_payment_intent_succeeded(event_data) 

69 

70 elif event_type == "checkout.session.completed": 

71 return self._handle_checkout_completed(event_data) 

72 

73 elif event_type == "customer.subscription.created": 

74 return self._handle_subscription_created(event_data) 

75 

76 elif event_type == "customer.subscription.updated": 

77 return self._handle_subscription_updated(event_data) 

78 

79 elif event_type == "customer.subscription.deleted": 

80 return self._handle_subscription_deleted(event_data) 

81 

82 elif event_type == "invoice.payment_succeeded": 

83 return self._handle_invoice_paid(event_data) 

84 

85 elif event_type == "invoice.payment_failed": 

86 return self._handle_invoice_failed(event_data) 

87 

88 elif event_type == "charge.dispute.created": 

89 return self._handle_dispute_created(event_data) 

90 

91 else: 

92 logger.info(f"Unhandled event type: {event_type}") 

93 return {"message": f"Event {event_type} received but not processed"} 

94 

95 def process_eventbridge_event(self, event: dict[str, Any]) -> dict[str, Any]: 

96 """Process a Stripe event delivered via EventBridge. 

97 

98 EventBridge events are pre-verified by AWS, so no signature 

99 verification is needed. The Stripe event is in event["detail"]. 

100 Falls back to parsing event["body"] for direct webhook calls. 

101 """ 

102 detail = event.get("detail", {}) 

103 

104 if detail and "type" in detail: 

105 return self.process_event(detail) 

106 

107 # Fallback: direct webhook call with body 

108 body = event.get("body", "") 

109 if event.get("isBase64Encoded") and body: 

110 body = base64.b64decode(body).decode("utf-8") 

111 

112 if body: 

113 stripe_event = json.loads(body) 

114 return self.process_event(stripe_event) 

115 

116 logger.warning("EventBridge event has no detail or body") 

117 return {"message": "No event data found"} 

118 

119 def _handle_checkout_completed(self, session: dict[str, Any]) -> dict[str, Any]: 

120 """Handle successful checkout session for credit purchase.""" 

121 metadata = session.get("metadata", {}) 

122 user_id = metadata.get("user_id") 

123 

124 if not user_id: 

125 logger.error("No user_id in checkout session metadata") 

126 return {"error": "Missing user_id"} 

127 

128 if session.get("mode") == "payment": 

129 credits = int(metadata.get("credits", 0)) 

130 # Stripe delivers events at least once. Key the grant on the payment intent so a 

131 # redelivered session, or the matching payment_intent.succeeded event, cannot grant 

132 # the same purchase twice. Fall back to the session id if no intent is attached. 

133 idempotency_key = session.get("payment_intent") or session.get("id", "") 

134 

135 if credits > 0 and idempotency_key: 

136 granted = self.credit_manager.idempotent_credit_grant( 

137 user_id, credits, idempotency_key 

138 ) 

139 if not granted: 

140 logger.info(f"Checkout {idempotency_key} already processed for {user_id}") 

141 return {"message": "Checkout already processed, credits not added"} 

142 new_balance = self.credit_manager.get_balance(user_id) 

143 logger.info( 

144 f"Added {credits} credits to user {user_id}, new balance: {new_balance}" 

145 ) 

146 return {"credits_added": credits, "new_balance": new_balance} 

147 

148 return {"message": "Checkout processed"} 

149 

150 def _handle_subscription_created(self, subscription: dict[str, Any]) -> dict[str, Any]: 

151 """Handle new subscription creation: record state and the first billing period.""" 

152 metadata = subscription.get("metadata", {}) 

153 user_id = metadata.get("user_id") 

154 customer_id = subscription.get("customer") 

155 subscription_id = subscription.get("id") 

156 status = str(subscription.get("status", "unknown")) 

157 

158 if not user_id: 

159 logger.warning("No user_id in subscription metadata") 

160 return {"subscription_id": subscription_id, "status": status} 

161 

162 self.credit_manager.set_subscription_state( 

163 user_id=user_id, 

164 status=status, 

165 stripe_customer_id=customer_id, 

166 stripe_subscription_id=subscription_id, 

167 ) 

168 

169 logger.info(f"Created subscription {subscription_id} for user {user_id} ({status})") 

170 

171 # Record the initial billing period so renewals can be detected 

172 current_period_start = subscription.get("current_period_start") 

173 if current_period_start and self.credit_manager.table: 

174 try: 

175 self.credit_manager.table.update_item( 

176 Key={"user_id": user_id}, 

177 UpdateExpression="SET last_credited_period = :period", 

178 ExpressionAttributeValues={ 

179 ":period": Decimal(str(current_period_start)), 

180 }, 

181 ) 

182 except Exception as e: 

183 logger.error(f"Error setting last_credited_period for {user_id}: {e}") 

184 

185 return { 

186 "subscription_id": subscription_id, 

187 "status": status, 

188 } 

189 

190 def _handle_subscription_updated(self, subscription: dict[str, Any]) -> dict[str, Any]: 

191 """Handle subscription updates and record billing-period renewals. 

192 

193 When a billing period changes (detected via current_period_start) the new 

194 period is recorded with a conditional update so duplicate deliveries are inert. 

195 """ 

196 metadata = subscription.get("metadata", {}) 

197 user_id = metadata.get("user_id") 

198 subscription_id = subscription.get("id") 

199 status = str(subscription.get("status", "unknown")) 

200 

201 if not user_id: 

202 return {"subscription_id": subscription_id, "status": status} 

203 

204 # Always update subscription status 

205 self.credit_manager.set_subscription_state( 

206 user_id=user_id, status=status, stripe_subscription_id=subscription_id 

207 ) 

208 

209 renewed = False 

210 current_period_start = subscription.get("current_period_start") 

211 

212 if status == "active" and current_period_start and self.credit_manager.table: 

213 renewed = self._record_billing_period(user_id, current_period_start) 

214 

215 logger.info( 

216 f"Updated subscription {subscription_id} status to {status}" 

217 + (f", new billing period {current_period_start}" if renewed else "") 

218 ) 

219 

220 return { 

221 "subscription_id": subscription_id, 

222 "status": status, 

223 } 

224 

225 def _record_billing_period(self, user_id: str, current_period_start: int) -> bool: 

226 """Record a new billing period on the credits item. 

227 

228 Uses a conditional update on last_credited_period so duplicate webhook 

229 deliveries and concurrent handlers cannot advance it twice. 

230 

231 Returns True if a new period was recorded, False if already recorded. 

232 """ 

233 if not self.credit_manager.table: 

234 return False 

235 

236 try: 

237 # Read current last_credited_period 

238 response = self.credit_manager.table.get_item(Key={"user_id": user_id}) 

239 item = response.get("Item", {}) 

240 last_period = item.get("last_credited_period") 

241 

242 period_decimal = Decimal(str(current_period_start)) 

243 

244 # Check if this is a new period 

245 if last_period is not None and period_decimal <= Decimal(str(last_period)): 

246 logger.info(f"Period {current_period_start} already recorded for {user_id}") 

247 return False 

248 

249 # Conditional update to prevent race conditions 

250 if last_period is None: 

251 condition = "attribute_not_exists(last_credited_period)" 

252 expr_values = {":new_period": period_decimal} 

253 else: 

254 condition = "last_credited_period = :prev_period" 

255 expr_values = { 

256 ":new_period": period_decimal, 

257 ":prev_period": Decimal(str(last_period)), 

258 } 

259 

260 self.credit_manager.table.update_item( 

261 Key={"user_id": user_id}, 

262 UpdateExpression="SET last_credited_period = :new_period", 

263 ConditionExpression=condition, 

264 ExpressionAttributeValues=expr_values, 

265 ) 

266 

267 logger.info(f"Recorded billing period {current_period_start} for {user_id}") 

268 return True 

269 

270 except ClientError as e: 

271 if e.response["Error"]["Code"] == "ConditionalCheckFailedException": 

272 logger.info( 

273 f"Billing period {current_period_start} already recorded for {user_id} " 

274 "(concurrent delivery)" 

275 ) 

276 return False 

277 logger.error(f"Error recording billing period for {user_id}: {e}") 

278 raise 

279 except Exception as e: 

280 logger.error(f"Error recording billing period for {user_id}: {e}") 

281 return False 

282 

283 def _handle_subscription_deleted(self, subscription: dict[str, Any]) -> dict[str, Any]: 

284 """Handle subscription cancellation with timestamp.""" 

285 metadata = subscription.get("metadata", {}) 

286 user_id = metadata.get("user_id") 

287 subscription_id = subscription.get("id") 

288 

289 if user_id: 

290 self.credit_manager.set_subscription_state( 

291 user_id=user_id, status="cancelled", stripe_subscription_id=subscription_id 

292 ) 

293 

294 # Set cancellation timestamp 

295 if self.credit_manager.table: 

296 try: 

297 self.credit_manager.table.update_item( 

298 Key={"user_id": user_id}, 

299 UpdateExpression="SET subscription_cancelled_at = :now", 

300 ExpressionAttributeValues={ 

301 ":now": datetime.now(UTC).isoformat(), 

302 }, 

303 ) 

304 except Exception as e: 

305 logger.error(f"Error setting cancelled_at for {user_id}: {e}") 

306 

307 logger.info(f"Cancelled subscription {subscription_id} for user {user_id}") 

308 

309 return {"subscription_id": subscription_id, "status": "cancelled"} 

310 

311 def _handle_invoice_paid(self, invoice: dict[str, Any]) -> dict[str, Any]: 

312 """Handle successful subscription payment.""" 

313 customer_id = invoice.get("customer") 

314 amount = invoice.get("amount_paid", 0) / 100.0 

315 logger.info(f"Invoice paid: ${amount} from customer {customer_id}") 

316 return {"amount_paid": amount} 

317 

318 def _handle_invoice_failed(self, invoice: dict[str, Any]) -> dict[str, Any]: 

319 """Handle failed subscription payment.""" 

320 customer_id = invoice.get("customer") 

321 logger.warning(f"Invoice payment failed for customer {customer_id}") 

322 return {"status": "payment_failed"} 

323 

324 def _handle_payment_intent_succeeded(self, payment_intent: dict[str, Any]) -> dict[str, Any]: 

325 """Handle successful payment intent with idempotent credit grant.""" 

326 metadata = payment_intent.get("metadata", {}) 

327 user_id = metadata.get("user_id") 

328 

329 if not user_id: 

330 logger.error("No user_id in payment_intent metadata") 

331 return {"error": "Missing user_id"} 

332 

333 # Check if this is a verification charge ($1) 

334 if metadata.get("type") == "verification": 

335 logger.info(f"Verification charge completed for user {user_id}") 

336 return {"type": "verification", "status": "completed"} 

337 

338 # Get credits from metadata (set during payment creation) 

339 credits = int(metadata.get("credits", 0)) 

340 payment_intent_id = payment_intent.get("id", "") 

341 

342 if credits > 0 and payment_intent_id: 

343 granted = self.credit_manager.idempotent_credit_grant( 

344 user_id, credits, payment_intent_id 

345 ) 

346 if granted: 

347 new_balance = self.credit_manager.get_balance(user_id) 

348 logger.info( 

349 f"Added {credits} credits to user {user_id}, new balance: {new_balance}" 

350 ) 

351 return {"credits_added": credits, "new_balance": new_balance} 

352 else: 

353 logger.info(f"Payment {payment_intent_id} already processed for {user_id}") 

354 return {"message": "Payment already processed, credits not added"} 

355 

356 return {"message": "Payment processed"} 

357 

358 def _handle_dispute_created(self, dispute: dict[str, Any]) -> dict[str, Any]: 

359 """Handle charge dispute (mark account as disputed).""" 

360 charge_id = dispute.get("charge") 

361 

362 if not charge_id: 

363 logger.error("No charge_id in dispute") 

364 return {"error": "Missing charge_id"} 

365 

366 amount = dispute.get("amount", 0) / 100.0 

367 reason = dispute.get("reason", "unknown") 

368 

369 logger.warning(f"Dispute created for charge {charge_id}: ${amount}, reason: {reason}") 

370 

371 return {"dispute_id": dispute.get("id"), "status": "created", "amount": amount}