feat: 完成权益到期与续期闭环

This commit is contained in:
2026-08-03 14:32:38 +08:00
parent 20ed4875b8
commit b2b56f43bb
18 changed files with 792 additions and 20 deletions

View File

@@ -0,0 +1,52 @@
"""add entitlement renewal idempotency and expiry index
Revision ID: 0026_entitlement_renewal
Revises: 0025_recent_practice_review
"""
from __future__ import annotations
from alembic import context, op
import sqlalchemy as sa
revision = "0026_entitlement_renewal"
down_revision = "0025_recent_practice_review"
branch_labels = None
depends_on = None
def upgrade() -> None:
inspector = None if context.is_offline_mode() else sa.inspect(op.get_bind())
log_columns = set() if inspector is None else {
column["name"] for column in inspector.get_columns("sys_user_entitlement_log")
}
if "request_key" not in log_columns:
op.add_column("sys_user_entitlement_log", sa.Column("request_key", sa.String(160), nullable=True))
log_indexes = set() if inspector is None else {
index["name"] for index in inspector.get_indexes("sys_user_entitlement_log")
}
if "ux_entitlement_log_request_key" not in log_indexes:
op.create_index(
"ux_entitlement_log_request_key",
"sys_user_entitlement_log",
["request_key"],
unique=True,
)
entitlement_indexes = set() if inspector is None else {
index["name"] for index in inspector.get_indexes("sys_user_entitlement")
}
if "ix_entitlement_status_expiry" not in entitlement_indexes:
op.create_index(
"ix_entitlement_status_expiry",
"sys_user_entitlement",
["status", "expired_at", "user_id"],
)
def downgrade() -> None:
op.drop_index("ix_entitlement_status_expiry", table_name="sys_user_entitlement")
op.drop_index("ux_entitlement_log_request_key", table_name="sys_user_entitlement_log")
op.drop_column("sys_user_entitlement_log", "request_key")

View File

@@ -9,7 +9,12 @@ from app.core.responses import api_success
from app.models.admin import Admin
from app.models.entitlement import EntitlementPlan
from app.models.user import User
from app.schemas.admin import EntitlementPlanSaveRequest, UserEntitlementAssignRequest
from app.schemas.admin import (
EntitlementPlanSaveRequest,
UserEntitlementAssignRequest,
UserEntitlementBatchRenewRequest,
UserEntitlementRenewRequest,
)
from app.services.admin_service import OperationLogService
from app.services.entitlement_service import EntitlementService, entitlement_dict, plan_dict
from app.services.topic_session_service import TopicSessionService
@@ -97,6 +102,79 @@ def assign_user_entitlement(
return api_success(entitlement_dict(view))
@router.post("/user/{user_id}/entitlement/renew")
def renew_user_entitlement(
user_id: int,
payload: UserEntitlementRenewRequest,
db: Session = Depends(get_db),
current_admin: Admin = Depends(get_current_admin),
) -> dict:
user = db.get(User, user_id)
if user is None or user.is_deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="用户不存在")
EntitlementService.renew_user_plan(
db,
user=user,
extension_days=payload.extensionDays,
request_key=f"{payload.idempotencyKey}:{user.id}",
operated_by=current_admin.id,
remark=payload.remark,
)
OperationLogService.write(
db,
admin_id=current_admin.id,
module="entitlement",
action="renew_user_plan",
target_id=user.id,
)
db.commit()
view = EntitlementService.active_entitlement(
db,
user,
monthly_topic_used=TopicSessionService.monthly_used_count(db, user.id),
)
return api_success(entitlement_dict(view))
@router.post("/user/entitlement/batch-renew")
def batch_renew_user_entitlement(
payload: UserEntitlementBatchRenewRequest,
db: Session = Depends(get_db),
current_admin: Admin = Depends(get_current_admin),
) -> dict:
succeeded: list[dict] = []
failed: list[dict] = []
for user_id in dict.fromkeys(payload.userIds):
user = db.get(User, user_id)
if user is None or user.is_deleted:
failed.append({"userId": user_id, "reason": "用户不存在"})
continue
try:
EntitlementService.renew_user_plan(
db,
user=user,
extension_days=payload.extensionDays,
request_key=f"{payload.idempotencyKey}:{user.id}",
operated_by=current_admin.id,
remark=payload.remark,
)
succeeded.append({"userId": user.id, "name": user.name})
except HTTPException as exc:
failed.append({"userId": user.id, "name": user.name, "reason": str(exc.detail)})
OperationLogService.write(
db,
admin_id=current_admin.id,
module="entitlement",
action="batch_renew_user_plan",
target_id=None,
result="SUCCESS" if not failed else "PARTIAL",
)
db.commit()
return api_success(
{"succeeded": succeeded, "failed": failed, "successCount": len(succeeded), "failureCount": len(failed)}
)
def _apply_plan_payload(plan: EntitlementPlan, payload: EntitlementPlanSaveRequest) -> None:
plan.name = payload.name.strip()
plan.plan_type = payload.planType

View File

@@ -1,8 +1,10 @@
from __future__ import annotations
import re
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from io import BytesIO
from typing import Annotated
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile, status
from fastapi.responses import StreamingResponse
@@ -45,6 +47,8 @@ class AdminGenerateReportRequest(BaseModel):
@router.get("/user/list")
def list_users(
keyword: str = Query(default=""),
planId: Annotated[int | None, Query(gt=0)] = None,
entitlementStatus: Annotated[str, Query(pattern="^(all|active|expiring7|expiring30|expired|default)$")] = "all",
page: int = Query(default=1, ge=1),
pageSize: int = Query(default=20, ge=10, le=100),
db: Session = Depends(get_db),
@@ -54,6 +58,51 @@ def list_users(
if keyword:
like = f"%{keyword}%"
query = query.where((User.phone.like(like)) | (User.name.like(like)))
now = datetime.now(UTC).replace(tzinfo=None)
active_user_ids = select(UserEntitlement.user_id).join(
EntitlementPlan, EntitlementPlan.id == UserEntitlement.plan_id
).where(
UserEntitlement.status == "active",
EntitlementPlan.status == 1,
EntitlementPlan.plan_type != "teacher",
(UserEntitlement.effective_at.is_(None)) | (UserEntitlement.effective_at <= now),
(UserEntitlement.expired_at.is_(None)) | (UserEntitlement.expired_at >= now),
)
if planId is not None:
query = query.where(User.id.in_(active_user_ids.where(UserEntitlement.plan_id == planId)))
if entitlementStatus == "active":
query = query.where(User.id.in_(active_user_ids))
elif entitlementStatus == "expiring7":
query = query.where(
User.id.in_(
active_user_ids.where(
UserEntitlement.expired_at.is_not(None),
UserEntitlement.expired_at <= now + timedelta(days=7),
)
)
)
elif entitlementStatus == "expiring30":
query = query.where(
User.id.in_(
active_user_ids.where(
UserEntitlement.expired_at.is_not(None),
UserEntitlement.expired_at > now + timedelta(days=7),
UserEntitlement.expired_at <= now + timedelta(days=30),
)
)
)
elif entitlementStatus == "default":
query = query.where(User.id.not_in(active_user_ids))
elif entitlementStatus == "expired":
expired_user_ids = select(UserEntitlement.user_id).join(
EntitlementPlan, EntitlementPlan.id == UserEntitlement.plan_id
).where(
UserEntitlement.status.in_(("active", "expired")),
UserEntitlement.expired_at.is_not(None),
UserEntitlement.expired_at < now,
EntitlementPlan.plan_type != "teacher",
)
query = query.where(User.id.not_in(active_user_ids), User.id.in_(expired_user_ids))
total = db.scalar(select(func.count()).select_from(query.order_by(None).subquery())) or 0
users = db.scalars(query.offset((page - 1) * pageSize).limit(pageSize)).all()
entitlements = _entitlement_views(db, users)
@@ -457,13 +506,32 @@ def _entitlement_views(db: Session, users: list[User]) -> dict[int, dict]:
return {}
counts = _monthly_topic_counts(db, user_ids)
explicit = _active_entitlement_rows(db, user_ids)
previous_expired = _latest_expired_entitlement_rows(db, user_ids)
default_plan = EntitlementService.default_plan(db)
result: dict[int, dict] = {}
for user in users:
if user.id in explicit:
entitlement, plan = explicit[user.id]
view = view_from_plan(plan, monthly_topic_used=counts.get(user.id, 0), entitlement=entitlement, source="assigned")
else:
view = EntitlementService.active_entitlement(db, user, monthly_topic_used=counts.get(user.id, 0))
if default_plan is None:
view = EntitlementService.active_entitlement(db, user, monthly_topic_used=counts.get(user.id, 0))
else:
view = view_from_plan(
default_plan,
monthly_topic_used=counts.get(user.id, 0),
entitlement=None,
source="default",
)
if user.id in previous_expired:
expired_entitlement, expired_plan = previous_expired[user.id]
view = replace(
view,
source="default_after_expiry",
lifecycle_status="expired_fallback",
previous_plan_name=expired_plan.name,
previous_expired_at=expired_entitlement.expired_at,
)
result[user.id] = entitlement_dict(view)
return result
@@ -488,6 +556,26 @@ def _active_entitlement_rows(db: Session, user_ids: list[int]) -> dict[int, tupl
return result
def _latest_expired_entitlement_rows(db: Session, user_ids: list[int]) -> dict[int, tuple[UserEntitlement, EntitlementPlan]]:
now = datetime.now(UTC).replace(tzinfo=None)
rows = db.execute(
select(UserEntitlement, EntitlementPlan)
.join(EntitlementPlan, EntitlementPlan.id == UserEntitlement.plan_id)
.where(
UserEntitlement.user_id.in_(user_ids),
UserEntitlement.status.in_(("active", "expired")),
UserEntitlement.expired_at.is_not(None),
UserEntitlement.expired_at < now,
EntitlementPlan.plan_type != "teacher",
)
.order_by(UserEntitlement.expired_at.desc(), UserEntitlement.id.desc())
).all()
result: dict[int, tuple[UserEntitlement, EntitlementPlan]] = {}
for entitlement, plan in rows:
result.setdefault(entitlement.user_id, (entitlement, plan))
return result
def _user_metrics(db: Session, *, user: User, monthly_topic_limit: int | None) -> dict:
now = datetime.now(UTC).replace(tzinfo=None)
month_start = now.replace(day=1, hour=0, minute=0, second=0, microsecond=0)

View File

@@ -2,7 +2,7 @@ from __future__ import annotations
from datetime import datetime
from sqlalchemy import BigInteger, DateTime, ForeignKey, Integer, String, Text
from sqlalchemy import BigInteger, DateTime, ForeignKey, Index, Integer, String, Text
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base, TimestampMixin
@@ -30,6 +30,7 @@ class EntitlementPlan(Base, TimestampMixin):
class UserEntitlement(Base, TimestampMixin):
__tablename__ = "sys_user_entitlement"
__table_args__ = (Index("ix_entitlement_status_expiry", "status", "expired_at", "user_id"),)
id: Mapped[int] = mapped_column(PRIMARY_KEY_TYPE, primary_key=True, autoincrement=True)
user_id: Mapped[int] = mapped_column(ForeignKey("sys_user.id"), index=True, nullable=False)
@@ -50,6 +51,7 @@ class UserEntitlementLog(Base):
from_plan_id: Mapped[int | None] = mapped_column(BigInteger, nullable=True)
to_plan_id: Mapped[int | None] = mapped_column(BigInteger, nullable=True)
action: Mapped[str] = mapped_column(String(30), nullable=False)
request_key: Mapped[str | None] = mapped_column(String(160), unique=True, nullable=True)
detail_json: Mapped[str | None] = mapped_column(Text, nullable=True)
operated_by: Mapped[int | None] = mapped_column(BigInteger, nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime, nullable=False)

View File

@@ -104,6 +104,16 @@ class UserEntitlementAssignRequest(BaseModel):
remark: str | None = Field(default=None, max_length=255)
class UserEntitlementRenewRequest(BaseModel):
extensionDays: int = Field(ge=1, le=3650)
idempotencyKey: str = Field(min_length=8, max_length=100)
remark: str | None = Field(default=None, max_length=255)
class UserEntitlementBatchRenewRequest(UserEntitlementRenewRequest):
userIds: list[int] = Field(min_length=1, max_length=200)
class KnowledgeSaveRequest(BaseModel):
name: str = Field(min_length=1, max_length=100)
feishuSpaceId: str = Field(min_length=1, max_length=100)

View File

@@ -1,7 +1,8 @@
from __future__ import annotations
import json
from dataclasses import dataclass
import math
from dataclasses import dataclass, replace
from datetime import UTC, datetime, timedelta
from fastapi import HTTPException, status
@@ -32,6 +33,10 @@ class EntitlementView:
effective_at: datetime | None = None
expired_at: datetime | None = None
source: str = "legacy"
lifecycle_status: str = "legacy"
days_until_expiry: int | None = None
previous_plan_name: str | None = None
previous_expired_at: datetime | None = None
@property
def monthly_topic_remaining(self) -> int | None:
@@ -92,7 +97,30 @@ class EntitlementService:
plan = EntitlementService.default_plan(db)
if plan is not None:
return view_from_plan(plan, monthly_topic_used=monthly_topic_used, entitlement=None, source="default")
previous = db.execute(
select(UserEntitlement, EntitlementPlan)
.join(EntitlementPlan, EntitlementPlan.id == UserEntitlement.plan_id)
.where(
UserEntitlement.user_id == user.id,
UserEntitlement.status.in_(("active", "expired")),
UserEntitlement.expired_at.is_not(None),
UserEntitlement.expired_at < now,
EntitlementPlan.plan_type != "teacher",
)
.order_by(UserEntitlement.expired_at.desc(), UserEntitlement.id.desc())
.limit(1)
).first()
view = view_from_plan(plan, monthly_topic_used=monthly_topic_used, entitlement=None, source="default")
if previous:
expired_entitlement, expired_plan = previous
return replace(
view,
source="default_after_expiry",
lifecycle_status="expired_fallback",
previous_plan_name=expired_plan.name,
previous_expired_at=expired_entitlement.expired_at,
)
return view
return EntitlementView(
plan_id=None,
@@ -108,8 +136,127 @@ class EntitlementService:
allow_share_draft=True,
deduct_quota=True,
source="legacy",
lifecycle_status="legacy",
)
@staticmethod
def expire_due_entitlements(db: Session, *, batch_size: int = 500, now: datetime | None = None) -> int:
current_time = now or _now()
total = 0
while True:
due = list(
db.scalars(
select(UserEntitlement)
.where(
UserEntitlement.status == "active",
UserEntitlement.expired_at.is_not(None),
UserEntitlement.expired_at < current_time,
)
.order_by(UserEntitlement.id.asc())
.limit(batch_size)
)
)
if not due:
return total
for entitlement in due:
entitlement.status = "expired"
db.add(entitlement)
db.add(
UserEntitlementLog(
user_id=entitlement.user_id,
entitlement_id=entitlement.id,
from_plan_id=entitlement.plan_id,
to_plan_id=None,
action="expire",
detail_json=json.dumps(
{"expiredAt": entitlement.expired_at.isoformat() if entitlement.expired_at else None},
ensure_ascii=False,
),
operated_by=None,
created_at=current_time,
)
)
db.flush()
total += len(due)
@staticmethod
def renew_user_plan(
db: Session,
*,
user: User,
extension_days: int,
request_key: str,
operated_by: int | None,
remark: str | None = None,
) -> UserEntitlement:
existing_log = db.scalar(select(UserEntitlementLog).where(UserEntitlementLog.request_key == request_key))
if existing_log is not None and existing_log.entitlement_id is not None:
existing = db.get(UserEntitlement, existing_log.entitlement_id)
if existing is not None:
return existing
db.scalar(select(User.id).where(User.id == user.id).with_for_update())
existing_log = db.scalar(select(UserEntitlementLog).where(UserEntitlementLog.request_key == request_key))
if existing_log is not None and existing_log.entitlement_id is not None:
existing = db.get(UserEntitlement, existing_log.entitlement_id)
if existing is not None:
return existing
now = _now()
current = db.scalar(
select(UserEntitlement)
.where(UserEntitlement.user_id == user.id, UserEntitlement.status.in_(("active", "expired")))
.order_by(UserEntitlement.created_at.desc(), UserEntitlement.id.desc())
.limit(1)
)
if current is None:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="该用户尚未分配专属权益,请先选择权益版本")
plan = db.get(EntitlementPlan, current.plan_id)
if plan is None or plan.status != 1 or plan.plan_type == "teacher":
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="原权益版本已停用,无法续期,请重新分配权益")
if current.status == "active" and current.expired_at is None:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="当前权益为长期有效,无需续期")
previous_expired_at = current.expired_at
base_time = max(now, previous_expired_at) if previous_expired_at is not None else now
renewed_expired_at = base_time + timedelta(days=extension_days)
current.status = "replaced"
db.add(current)
renewed = UserEntitlement(
user_id=user.id,
plan_id=plan.id,
status="active",
effective_at=now,
expired_at=renewed_expired_at,
assigned_by=operated_by,
remark=remark,
)
db.add(renewed)
db.flush()
db.add(
UserEntitlementLog(
user_id=user.id,
entitlement_id=renewed.id,
from_plan_id=plan.id,
to_plan_id=plan.id,
action="renew",
request_key=request_key,
detail_json=json.dumps(
{
"extensionDays": extension_days,
"previousExpiredAt": previous_expired_at.isoformat() if previous_expired_at else None,
"renewedExpiredAt": renewed_expired_at.isoformat(),
"sourceEntitlementId": current.id,
"remark": remark,
},
ensure_ascii=False,
),
operated_by=operated_by,
created_at=now,
)
)
return renewed
@staticmethod
def assign_user_plan(
db: Session,
@@ -217,6 +364,10 @@ def entitlement_dict(view: EntitlementView) -> dict:
"effectiveAt": view.effective_at,
"expiredAt": view.expired_at,
"source": view.source,
"lifecycleStatus": view.lifecycle_status,
"daysUntilExpiry": view.days_until_expiry,
"previousPlanName": view.previous_plan_name,
"previousExpiredAt": view.previous_expired_at,
}
@@ -240,6 +391,17 @@ def view_from_plan(
entitlement: UserEntitlement | None,
source: str,
) -> EntitlementView:
days_until_expiry = _days_until_expiry(entitlement.expired_at) if entitlement else None
lifecycle_status = "default"
if entitlement is not None:
if entitlement.expired_at is None:
lifecycle_status = "long_term"
elif days_until_expiry is not None and days_until_expiry <= 7:
lifecycle_status = "expiring_7"
elif days_until_expiry is not None and days_until_expiry <= 30:
lifecycle_status = "expiring_30"
else:
lifecycle_status = "active"
return EntitlementView(
plan_id=plan.id,
name=plan.name,
@@ -256,8 +418,16 @@ def view_from_plan(
effective_at=entitlement.effective_at if entitlement else None,
expired_at=entitlement.expired_at if entitlement else None,
source=source,
lifecycle_status=lifecycle_status,
days_until_expiry=days_until_expiry,
)
def _now() -> datetime:
return datetime.now(UTC).replace(tzinfo=None)
def _days_until_expiry(expired_at: datetime | None) -> int | None:
if expired_at is None:
return None
return max(0, math.ceil((expired_at - _now()).total_seconds() / 86400))

View File

@@ -10,6 +10,7 @@ from app.core.database import SessionLocal
from app.models.knowledge import KnowledgeRetrievalCandidate, KnowledgeRetrievalLog
from app.models.logs import LogRetentionPolicy
from app.services.redis_client import get_sync_redis_client
from app.services.entitlement_service import EntitlementService
class MaintenanceService:
@@ -31,6 +32,8 @@ class MaintenanceService:
if not lock_acquired:
return
with SessionLocal() as db:
EntitlementService.expire_due_entitlements(db)
db.commit()
policy = db.scalar(select(LogRetentionPolicy).order_by(LogRetentionPolicy.id).limit(1))
if not policy or not policy.enabled or not policy.retention_days:
return

View File

@@ -1,16 +1,16 @@
from __future__ import annotations
from datetime import UTC, datetime
from datetime import UTC, datetime, timedelta
import pytest
from fastapi import HTTPException
from sqlalchemy import create_engine, event
from sqlalchemy import create_engine, event, func, select
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool
from app.models import Base
from app.models.chat import ChatMessage, ChatSession, TopicSession
from app.models.entitlement import EntitlementPlan
from app.models.entitlement import EntitlementPlan, UserEntitlement, UserEntitlementLog
from app.models.user import User
from app.services.chat_service import ChatService
from app.services.entitlement_service import EntitlementService, entitlement_dict
@@ -113,6 +113,112 @@ def test_assign_user_plan_replaces_previous_active_plan():
assert view.monthly_topic_remaining == 86
def test_expired_entitlement_falls_back_and_keeps_previous_plan_context():
with _db() as db:
user, _session = _seed_user_session(db)
db.add_all(
[
EntitlementPlan(id=10, name="基础版", plan_type="basic", monthly_topic_limit=30, status=1, sort_order=10),
EntitlementPlan(id=20, name="深度陪伴版", plan_type="deep", monthly_topic_limit=90, status=1, sort_order=20),
UserEntitlement(
id=100,
user_id=user.id,
plan_id=20,
status="active",
effective_at=_now() - timedelta(days=40),
expired_at=_now() - timedelta(days=1),
),
]
)
db.commit()
view = EntitlementService.active_entitlement(db, user)
expired_count = EntitlementService.expire_due_entitlements(db)
db.commit()
assert view.plan_id == 10
assert view.lifecycle_status == "expired_fallback"
assert view.previous_plan_name == "深度陪伴版"
assert view.previous_expired_at is not None
assert expired_count == 1
assert db.get(UserEntitlement, 100).status == "expired"
assert db.scalar(select(UserEntitlementLog).where(UserEntitlementLog.action == "expire")) is not None
def test_renew_user_plan_extends_from_current_expiry_and_is_idempotent():
with _db() as db:
user, _session = _seed_user_session(db)
current_expiry = _now() + timedelta(days=5)
db.add(EntitlementPlan(id=20, name="深度陪伴版", plan_type="deep", monthly_topic_limit=90, status=1))
db.add(
UserEntitlement(
id=100,
user_id=user.id,
plan_id=20,
status="active",
effective_at=_now() - timedelta(days=20),
expired_at=current_expiry,
)
)
db.commit()
renewed = EntitlementService.renew_user_plan(
db,
user=user,
extension_days=30,
request_key="renew-test-key:1",
operated_by=99,
remark="测试续期",
)
db.commit()
duplicate = EntitlementService.renew_user_plan(
db,
user=user,
extension_days=30,
request_key="renew-test-key:1",
operated_by=99,
remark="重复请求",
)
db.commit()
assert renewed.id == duplicate.id
assert renewed.expired_at == current_expiry + timedelta(days=30)
assert db.get(UserEntitlement, 100).status == "replaced"
assert db.scalar(select(func.count(UserEntitlementLog.id)).where(UserEntitlementLog.action == "renew")) == 1
def test_renew_expired_user_plan_restores_same_plan_from_now():
with _db() as db:
user, _session = _seed_user_session(db)
db.add(EntitlementPlan(id=20, name="深度陪伴版", plan_type="deep", monthly_topic_limit=90, status=1))
db.add(
UserEntitlement(
id=100,
user_id=user.id,
plan_id=20,
status="expired",
effective_at=_now() - timedelta(days=40),
expired_at=_now() - timedelta(days=1),
)
)
db.commit()
renewed = EntitlementService.renew_user_plan(
db,
user=user,
extension_days=30,
request_key="restore-test-key:1",
operated_by=99,
)
db.commit()
view = EntitlementService.active_entitlement(db, user)
assert renewed.expired_at is not None
assert timedelta(days=29) < renewed.expired_at - _now() <= timedelta(days=30)
assert view.plan_id == 20
assert view.source == "assigned"
def test_monthly_topic_quota_blocks_new_topic_but_allows_existing_topic():
with _db() as db:
user, session = _seed_user_session(db)