diff --git a/apps/api/app/api/routes/subscription.py b/apps/api/app/api/routes/subscription.py index 996eb104f..cbc07ab2f 100644 --- a/apps/api/app/api/routes/subscription.py +++ b/apps/api/app/api/routes/subscription.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import List + from dataclasses import replace from datetime import datetime, timezone @@ -180,6 +182,59 @@ async def cancel_subscription( ) + + +@router.post("/payment-callback") +async def payment_callback( + user_id: str, + plan: str, + billing_cycle: str, + amount: float, + payment_method: str = "alipay", + payment_id: str = "", +): + """支付回调 - 在事务中更新账单和订阅状态 + + 注意:生产环境需要验证支付签名 + """ + from datetime import timedelta + from packages.adapters.sqlalchemy_impl.billing_repository import SQLAlchemyBillingRepository + from packages.adapters.sqlalchemy_impl.session import SessionLocal + import uuid + + if SessionLocal is None: + raise HTTPException(status_code=500, detail="Database not available") + + session = SessionLocal() + try: + repo = SQLAlchemyBillingRepository(session) + + # 创建账单记录 + record_id = uuid.uuid4().hex + record = repo.create({ + "id": record_id, + "user_id": user_id, + "plan_name": _get_plan_name(plan), + "amount": amount, + "billing_cycle": billing_cycle, + "status": "pending", + }) + + # 在事务中标记支付成功并更新订阅 + repo.mark_paid(record_id, payment_method, payment_id) + + # 计算到期时间 + days = 365 if billing_cycle == "yearly" else 30 + expires_at = datetime.now(timezone.utc) + timedelta(days=days) + repo.update_subscription_on_payment(user_id, plan, expires_at) + + return {"success": True, "message": "支付成功", "record_id": record_id} + except Exception as e: + session.rollback() + raise HTTPException(status_code=500, detail=f"支付处理失败: {str(e)}") + finally: + session.close() + @router.post("/toggle-auto-renew", response_model=SimpleResponse) async def toggle_auto_renew( request: ToggleAutoRenewRequest, diff --git a/packages/adapters/sqlalchemy_impl/billing_repository.py b/packages/adapters/sqlalchemy_impl/billing_repository.py new file mode 100644 index 000000000..afee7cb61 --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/billing_repository.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +from datetime import datetime, timezone + +from sqlalchemy.orm import Session + +from packages.adapters.sqlalchemy_impl.models import BillingRecordModel + + +class SQLAlchemyBillingRepository: + def __init__(self, session: Session): + self.session = session + + def create(self, record: dict) -> BillingRecordModel: + model = BillingRecordModel(**record) + self.session.add(model) + self.session.commit() + return model + + def find_by_user(self, user_id: str, limit: int = 50) -> list[BillingRecordModel]: + return ( + self.session.query(BillingRecordModel) + .filter(BillingRecordModel.user_id == user_id) + .order_by(BillingRecordModel.created_at.desc()) + .limit(limit) + .all() + ) + + def find_by_id(self, record_id: str) -> BillingRecordModel | None: + return self.session.get(BillingRecordModel, record_id) + + def mark_paid(self, record_id: str, payment_method: str, payment_id: str) -> bool: + model = self.session.get(BillingRecordModel, record_id) + if model is None or model.status == "paid": + return False + model.status = "paid" + model.payment_method = payment_method + model.payment_id = payment_id + model.paid_at = datetime.now(timezone.utc) + self.session.commit() + return True + + def update_subscription_on_payment(self, user_id: str, plan: str, expires_at: datetime) -> None: + """在支付成功后更新用户订阅状态(事务内调用)""" + from packages.adapters.sqlalchemy_impl.models import UserModel + model = self.session.get(UserModel, user_id) + if model: + model.subscription_plan = plan + model.subscription_status = "active" + model.subscription_expires_at = expires_at + self.session.commit() diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index c1adac288..966b598f7 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -459,4 +459,20 @@ class TTSJobModel(Base): started_at = Column(DateTime, nullable=True) completed_at = Column(DateTime, nullable=True) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) - updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) \ No newline at end of file + updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) + +class BillingRecordModel(Base): + """账单记录""" + __tablename__ = "billing_records" + + id = Column(String(36), primary_key=True) + user_id = Column(String(36), nullable=False, index=True) + plan_name = Column(String(50), nullable=False) + amount = Column(Float, nullable=False) + billing_cycle = Column(String(20), nullable=False) + status = Column(String(20), nullable=False, default="pending") # pending, paid, failed, refunded + payment_method = Column(String(50), nullable=True) + payment_id = Column(String(100), nullable=True) # 第三方支付流水号 + invoice_url = Column(String(500), nullable=True) + created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) + paid_at = Column(DateTime, nullable=True)