Merge pull request #3586 from deqiying/codex/fix-subscription-revoke-soft-delete

修复订阅撤销操作实际上是软删除的bug
This commit is contained in:
Wesley Liddick
2026-07-01 15:38:45 +08:00
committed by GitHub
14 changed files with 365 additions and 27 deletions
@@ -249,8 +249,9 @@ func (h *SubscriptionHandler) ResetQuota(c *gin.Context) {
response.Success(c, dto.UserSubscriptionFromServiceAdmin(sub))
}
// Revoke handles revoking a subscription
// DELETE /api/v1/admin/subscriptions/:id
// Revoke handles revoking a subscription.
// POST /api/v1/admin/subscriptions/:id/revoke
// DELETE /api/v1/admin/subscriptions/:id is kept for backward compatibility.
func (h *SubscriptionHandler) Revoke(c *gin.Context) {
subscriptionID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil {
+1
View File
@@ -755,6 +755,7 @@ func userSubscriptionFromServiceBase(sub *service.UserSubscription) UserSubscrip
MonthlyUsageUSD: sub.MonthlyUsageUSD,
CreatedAt: sub.CreatedAt,
UpdatedAt: sub.UpdatedAt,
RevokedAt: sub.DeletedAt,
User: UserFromServiceShallow(sub.User),
Group: GroupFromServiceShallow(sub.Group),
}
+3 -2
View File
@@ -594,8 +594,9 @@ type UserSubscription struct {
WeeklyUsageUSD float64 `json:"weekly_usage_usd"`
MonthlyUsageUSD float64 `json:"monthly_usage_usd"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
RevokedAt *time.Time `json:"revoked_at,omitempty"`
User *User `json:"user,omitempty"`
Group *Group `json:"group,omitempty"`
@@ -18,6 +18,7 @@ const (
billingBalanceKeyPrefix = "billing:balance:"
billingSubKeyPrefix = "billing:sub:"
billingRateLimitKeyPrefix = "apikey:rate:"
subCacheInvalidateChannel = "subscription:cache:invalidate"
billingCacheTTL = 5 * time.Minute
billingCacheJitter = 30 * time.Second
rateLimitCacheTTL = 7 * 24 * time.Hour // 7 days matches the longest window
@@ -256,6 +257,45 @@ func (c *billingCache) InvalidateSubscriptionCache(ctx context.Context, userID,
return c.rdb.Del(ctx, key).Err()
}
func (c *billingCache) PublishSubscriptionCacheInvalidation(ctx context.Context, cacheKey string) error {
return c.rdb.Publish(ctx, subCacheInvalidateChannel, cacheKey).Err()
}
func (c *billingCache) SubscribeSubscriptionCacheInvalidation(ctx context.Context, handler func(cacheKey string)) error {
pubsub := c.rdb.Subscribe(ctx, subCacheInvalidateChannel)
_, err := pubsub.Receive(ctx)
if err != nil {
_ = pubsub.Close()
return fmt.Errorf("subscribe to subscription cache invalidation: %w", err)
}
go func() {
defer func() {
if err := pubsub.Close(); err != nil {
log.Printf("Warning: failed to close subscription cache invalidation pubsub: %v", err)
}
}()
ch := pubsub.Channel()
for {
select {
case <-ctx.Done():
return
case msg, ok := <-ch:
if !ok {
return
}
if msg != nil {
handler(msg.Payload)
}
}
}
}()
return nil
}
func (c *billingCache) GetAPIKeyRateLimit(ctx context.Context, keyID int64) (*service.APIKeyRateLimitCacheData, error) {
key := billingRateLimitKey(keyID)
result, err := c.rdb.HGetAll(ctx, key).Result()
@@ -6,6 +6,9 @@ import (
dbent "github.com/Wei-Shaw/sub2api/ent"
"github.com/Wei-Shaw/sub2api/ent/group"
"github.com/Wei-Shaw/sub2api/ent/predicate"
"github.com/Wei-Shaw/sub2api/ent/schema/mixins"
"github.com/Wei-Shaw/sub2api/ent/user"
"github.com/Wei-Shaw/sub2api/ent/usersubscription"
"github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
"github.com/Wei-Shaw/sub2api/internal/service"
@@ -194,6 +197,7 @@ func (r *userSubscriptionRepository) ListByGroupID(ctx context.Context, groupID
func (r *userSubscriptionRepository) List(ctx context.Context, params pagination.PaginationParams, userID, groupID *int64, status, platform, sortBy, sortOrder string) ([]service.UserSubscription, *pagination.PaginationResult, error) {
client := clientFromContext(ctx, r.client)
q := client.UserSubscription.Query()
includeSoftDeleted := status == "" || status == service.SubscriptionStatusRevoked
if userID != nil {
q = q.Where(usersubscription.UserIDEQ(*userID))
}
@@ -201,7 +205,11 @@ func (r *userSubscriptionRepository) List(ctx context.Context, params pagination
q = q.Where(usersubscription.GroupIDEQ(*groupID))
}
if platform != "" {
q = q.Where(usersubscription.HasGroupWith(group.PlatformEQ(platform)))
groupPredicates := []predicate.Group{group.PlatformEQ(platform)}
if includeSoftDeleted {
groupPredicates = append(groupPredicates, group.DeletedAtIsNil())
}
q = q.Where(usersubscription.HasGroupWith(groupPredicates...))
}
// Status filtering with real-time expiration check
@@ -224,20 +232,29 @@ func (r *userSubscriptionRepository) List(ctx context.Context, params pagination
),
),
)
case service.SubscriptionStatusRevoked:
// Revoked is a DTO/API display state backed by user_subscriptions.deleted_at.
q = q.Where(usersubscription.DeletedAtNotNil())
case "":
// No filter
// No filter. Use SkipSoftDelete below so admin "all status" includes revoked history.
default:
// Other status (e.g., revoked)
// Other persisted status.
q = q.Where(usersubscription.StatusEQ(status))
}
total, err := q.Clone().Count(ctx)
queryCtx := ctx
if includeSoftDeleted {
queryCtx = mixins.SkipSoftDelete(ctx)
}
total, err := q.Clone().Count(queryCtx)
if err != nil {
return nil, nil, err
}
// Apply sorting
q = q.WithUser().WithGroup().WithAssignedByUser()
if !includeSoftDeleted {
q = q.WithUser().WithGroup().WithAssignedByUser()
}
// Determine sort field
var field string
@@ -260,12 +277,19 @@ func (r *userSubscriptionRepository) List(ctx context.Context, params pagination
subs, err := q.
Offset(params.Offset()).
Limit(params.Limit()).
All(ctx)
All(queryCtx)
if err != nil {
return nil, nil, err
}
return userSubscriptionEntitiesToService(subs), paginationResultFromTotal(int64(total), params), nil
result := userSubscriptionEntitiesToService(subs)
if includeSoftDeleted {
if err := r.attachUserSubscriptionRelations(ctx, result); err != nil {
return nil, nil, err
}
}
return result, paginationResultFromTotal(int64(total), params), nil
}
func (r *userSubscriptionRepository) ExistsByUserIDAndGroupID(ctx context.Context, userID, groupID int64) (bool, error) {
@@ -425,17 +449,91 @@ func (r *userSubscriptionRepository) DeleteByGroupID(ctx context.Context, groupI
return int64(n), err
}
func (r *userSubscriptionRepository) attachUserSubscriptionRelations(ctx context.Context, subs []service.UserSubscription) error {
if len(subs) == 0 {
return nil
}
userIDs := make([]int64, 0, len(subs))
groupIDs := make([]int64, 0, len(subs))
assignedByIDs := make([]int64, 0, len(subs))
for i := range subs {
userIDs = append(userIDs, subs[i].UserID)
groupIDs = append(groupIDs, subs[i].GroupID)
if subs[i].AssignedBy != nil {
assignedByIDs = append(assignedByIDs, *subs[i].AssignedBy)
}
}
client := clientFromContext(ctx, r.client)
users, err := client.User.Query().Where(user.IDIn(uniqueInt64s(userIDs)...)).All(ctx)
if err != nil {
return err
}
userByID := make(map[int64]*service.User, len(users))
for _, u := range users {
userByID[u.ID] = userEntityToService(u)
}
groups, err := client.Group.Query().Where(group.IDIn(uniqueInt64s(groupIDs)...)).All(ctx)
if err != nil {
return err
}
groupByID := make(map[int64]*service.Group, len(groups))
for _, g := range groups {
groupByID[g.ID] = groupEntityToService(g)
}
assignedByID := map[int64]*service.User{}
if len(assignedByIDs) > 0 {
assignedUsers, err := client.User.Query().Where(user.IDIn(uniqueInt64s(assignedByIDs)...)).All(ctx)
if err != nil {
return err
}
assignedByID = make(map[int64]*service.User, len(assignedUsers))
for _, u := range assignedUsers {
assignedByID[u.ID] = userEntityToService(u)
}
}
for i := range subs {
subs[i].User = userByID[subs[i].UserID]
subs[i].Group = groupByID[subs[i].GroupID]
if subs[i].AssignedBy != nil {
subs[i].AssignedByUser = assignedByID[*subs[i].AssignedBy]
}
}
return nil
}
func uniqueInt64s(values []int64) []int64 {
seen := make(map[int64]struct{}, len(values))
out := make([]int64, 0, len(values))
for _, v := range values {
if _, ok := seen[v]; ok {
continue
}
seen[v] = struct{}{}
out = append(out, v)
}
return out
}
func userSubscriptionEntityToService(m *dbent.UserSubscription) *service.UserSubscription {
if m == nil {
return nil
}
status := m.Status
if m.DeletedAt != nil {
status = service.SubscriptionStatusRevoked
}
out := &service.UserSubscription{
ID: m.ID,
UserID: m.UserID,
GroupID: m.GroupID,
StartsAt: m.StartsAt,
ExpiresAt: m.ExpiresAt,
Status: m.Status,
Status: status,
DailyWindowStart: m.DailyWindowStart,
WeeklyWindowStart: m.WeeklyWindowStart,
MonthlyWindowStart: m.MonthlyWindowStart,
@@ -447,6 +545,7 @@ func userSubscriptionEntityToService(m *dbent.UserSubscription) *service.UserSub
Notes: derefString(m.Notes),
CreatedAt: m.CreatedAt,
UpdatedAt: m.UpdatedAt,
DeletedAt: m.DeletedAt,
}
if m.Edges.User != nil {
out.User = userEntityToService(m.Edges.User)
@@ -326,6 +326,61 @@ func (s *UserSubscriptionRepoSuite) TestList_FilterByStatus() {
s.Require().Equal(service.SubscriptionStatusExpired, subs[0].Status)
}
func (s *UserSubscriptionRepoSuite) TestList_IncludesRevokedWhenStatusEmpty() {
user1 := s.mustCreateUser("allstatus1@test.com", service.RoleUser)
user2 := s.mustCreateUser("allstatus2@test.com", service.RoleUser)
user3 := s.mustCreateUser("allstatus3@test.com", service.RoleUser)
group1 := s.mustCreateGroup("g-allstatus-1")
group2 := s.mustCreateGroup("g-allstatus-2")
group3 := s.mustCreateGroup("g-allstatus-3")
s.mustCreateSubscription(user1.ID, group1.ID, nil)
s.mustCreateSubscription(user2.ID, group2.ID, func(c *dbent.UserSubscriptionCreate) {
c.SetStatus(service.SubscriptionStatusExpired)
c.SetExpiresAt(time.Now().Add(-24 * time.Hour))
})
revoked := s.mustCreateSubscription(user3.ID, group3.ID, nil)
s.Require().NoError(s.repo.Delete(s.ctx, revoked.ID))
subs, pag, err := s.repo.List(s.ctx, pagination.PaginationParams{Page: 1, PageSize: 10}, nil, nil, "", "", "", "")
s.Require().NoError(err)
s.Require().Len(subs, 3)
s.Require().Equal(int64(3), pag.Total)
var gotRevoked *service.UserSubscription
for i := range subs {
if subs[i].ID == revoked.ID {
gotRevoked = &subs[i]
break
}
}
s.Require().NotNil(gotRevoked, "all status should include soft-deleted subscription")
s.Require().Equal(service.SubscriptionStatusRevoked, gotRevoked.Status)
s.Require().NotNil(gotRevoked.DeletedAt)
s.Require().NotNil(gotRevoked.User)
s.Require().NotNil(gotRevoked.Group)
}
func (s *UserSubscriptionRepoSuite) TestList_FilterByRevokedStatus() {
user1 := s.mustCreateUser("revokedfilter1@test.com", service.RoleUser)
user2 := s.mustCreateUser("revokedfilter2@test.com", service.RoleUser)
group1 := s.mustCreateGroup("g-revoked-1")
group2 := s.mustCreateGroup("g-revoked-2")
active := s.mustCreateSubscription(user1.ID, group1.ID, nil)
revoked := s.mustCreateSubscription(user2.ID, group2.ID, nil)
s.Require().NoError(s.repo.Delete(s.ctx, revoked.ID))
subs, pag, err := s.repo.List(s.ctx, pagination.PaginationParams{Page: 1, PageSize: 10}, nil, nil, service.SubscriptionStatusRevoked, "", "", "")
s.Require().NoError(err)
s.Require().Len(subs, 1)
s.Require().Equal(int64(1), pag.Total)
s.Require().Equal(revoked.ID, subs[0].ID)
s.Require().NotEqual(active.ID, subs[0].ID)
s.Require().Equal(service.SubscriptionStatusRevoked, subs[0].Status)
s.Require().NotNil(subs[0].DeletedAt)
}
// --- Usage tracking ---
func (s *UserSubscriptionRepoSuite) TestIncrementUsage() {
+1
View File
@@ -561,6 +561,7 @@ func registerSubscriptionRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
subscriptions.POST("/bulk-assign", h.Admin.Subscription.BulkAssign)
subscriptions.POST("/:id/extend", h.Admin.Subscription.Extend)
subscriptions.POST("/:id/reset-quota", h.Admin.Subscription.ResetQuota)
subscriptions.POST("/:id/revoke", h.Admin.Subscription.Revoke)
subscriptions.DELETE("/:id", h.Admin.Subscription.Revoke)
}
@@ -96,6 +96,11 @@ type apiKeyRateLimitLoader interface {
GetRateLimitData(ctx context.Context, keyID int64) (*APIKeyRateLimitData, error)
}
type subscriptionCacheInvalidationPubSub interface {
PublishSubscriptionCacheInvalidation(ctx context.Context, cacheKey string) error
SubscribeSubscriptionCacheInvalidation(ctx context.Context, handler func(cacheKey string)) error
}
// BillingCacheService 计费缓存服务
// 负责余额和订阅数据的缓存管理,提供高性能的计费资格检查
type BillingCacheService struct {
@@ -525,6 +530,28 @@ func (s *BillingCacheService) InvalidateSubscription(ctx context.Context, userID
return nil
}
func (s *BillingCacheService) PublishSubscriptionCacheInvalidation(ctx context.Context, cacheKey string) error {
if s.cache == nil {
return nil
}
pubsub, ok := s.cache.(subscriptionCacheInvalidationPubSub)
if !ok {
return nil
}
return pubsub.PublishSubscriptionCacheInvalidation(ctx, cacheKey)
}
func (s *BillingCacheService) SubscribeSubscriptionCacheInvalidation(ctx context.Context, handler func(cacheKey string)) error {
if s.cache == nil {
return nil
}
pubsub, ok := s.cache.(subscriptionCacheInvalidationPubSub)
if !ok {
return nil
}
return pubsub.SubscribeSubscriptionCacheInvalidation(ctx, handler)
}
// InvalidateAPIKeyRateLimit invalidates the Redis rate-limit usage cache for an API key.
func (s *BillingCacheService) InvalidateAPIKeyRateLimit(ctx context.Context, keyID int64) error {
if s.cache == nil {
@@ -107,6 +107,8 @@ const (
SubscriptionStatusActive = domain.SubscriptionStatusActive
SubscriptionStatusExpired = domain.SubscriptionStatusExpired
SubscriptionStatusSuspended = domain.SubscriptionStatusSuspended
// SubscriptionStatusRevoked 是 soft-deleted 订阅的 API 展示态,不写入 status 字段。
SubscriptionStatusRevoked = "revoked"
)
// LinuxDoConnectSyntheticEmailDomain 是 LinuxDo Connect 用户的合成邮箱后缀(RFC 保留域名)。
@@ -0,0 +1,76 @@
//go:build unit
package service
import (
"context"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
)
type revokeCacheUserSubRepoStub struct {
userSubRepoNoop
sub *UserSubscription
deleted bool
getActiveCalls int
}
func (r *revokeCacheUserSubRepoStub) GetByID(_ context.Context, id int64) (*UserSubscription, error) {
if r.sub == nil || r.sub.ID != id || r.deleted {
return nil, ErrSubscriptionNotFound
}
cp := *r.sub
return &cp, nil
}
func (r *revokeCacheUserSubRepoStub) Delete(_ context.Context, id int64) error {
if r.sub == nil || r.sub.ID != id || r.deleted {
return ErrSubscriptionNotFound
}
r.deleted = true
return nil
}
func (r *revokeCacheUserSubRepoStub) GetActiveByUserIDAndGroupID(_ context.Context, userID, groupID int64) (*UserSubscription, error) {
r.getActiveCalls++
if r.deleted || r.sub == nil || r.sub.UserID != userID || r.sub.GroupID != groupID {
return nil, ErrSubscriptionNotFound
}
cp := *r.sub
return &cp, nil
}
func TestRevokeSubscription_InvalidatesL1CacheSynchronously(t *testing.T) {
repo := &revokeCacheUserSubRepoStub{
sub: &UserSubscription{
ID: 1,
UserID: 10,
GroupID: 20,
Status: SubscriptionStatusActive,
ExpiresAt: time.Now().Add(time.Hour),
},
}
svc := NewSubscriptionService(groupRepoNoop{}, repo, nil, nil, &config.Config{
SubscriptionCache: config.SubscriptionCacheConfig{
L1Size: 16,
L1TTLSeconds: 60,
},
})
t.Cleanup(svc.Stop)
_, err := svc.GetActiveSubscription(context.Background(), 10, 20)
require.NoError(t, err)
svc.subCacheL1.Wait()
require.Equal(t, 1, repo.getActiveCalls)
err = svc.RevokeSubscription(context.Background(), 1)
require.NoError(t, err)
_, err = svc.GetActiveSubscription(context.Background(), 10, 20)
require.ErrorIs(t, err, ErrSubscriptionNotFound)
require.Equal(t, 2, repo.getActiveCalls, "撤销后应回源确认订阅已不存在,不能命中旧 L1")
}
@@ -65,6 +65,7 @@ func NewSubscriptionService(groupRepo GroupRepository, userSubRepo UserSubscript
}
svc.initSubCache(cfg)
svc.initMaintenanceQueue(cfg)
svc.StartSubCacheInvalidationSubscriber(context.Background())
return svc
}
@@ -142,6 +143,48 @@ func (s *SubscriptionService) InvalidateSubCache(userID, groupID int64) {
s.subCacheL1.Del(subCacheKey(userID, groupID))
}
// InvalidateSubCacheSync 失效订阅 L1 缓存并等待 Ristretto 删除操作生效。
func (s *SubscriptionService) InvalidateSubCacheSync(userID, groupID int64) {
s.invalidateSubCacheKeySync(subCacheKey(userID, groupID))
}
func (s *SubscriptionService) invalidateSubCacheKeySync(key string) {
if s.subCacheL1 == nil {
return
}
s.subCacheL1.Del(key)
s.subCacheL1.Wait()
}
// StartSubCacheInvalidationSubscriber 启动跨实例订阅 L1 缓存失效订阅。
func (s *SubscriptionService) StartSubCacheInvalidationSubscriber(ctx context.Context) {
if s.billingCacheService == nil || s.subCacheL1 == nil {
return
}
if err := s.billingCacheService.SubscribeSubscriptionCacheInvalidation(ctx, func(cacheKey string) {
s.invalidateSubCacheKeySync(cacheKey)
}); err != nil {
log.Printf("Warning: failed to start subscription cache invalidation subscriber: %v", err)
}
}
func (s *SubscriptionService) invalidateSubscriptionCaches(userID, groupID int64) error {
s.InvalidateSubCacheSync(userID, groupID)
if s.billingCacheService == nil {
return nil
}
cacheCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := s.billingCacheService.InvalidateSubscription(cacheCtx, userID, groupID); err != nil {
return fmt.Errorf("invalidate billing subscription cache: %w", err)
}
if err := s.billingCacheService.PublishSubscriptionCacheInvalidation(cacheCtx, subCacheKey(userID, groupID)); err != nil {
return fmt.Errorf("publish subscription cache invalidation: %w", err)
}
return nil
}
// AssignSubscriptionInput 分配订阅输入
type AssignSubscriptionInput struct {
UserID int64
@@ -528,15 +571,8 @@ func (s *SubscriptionService) RevokeSubscription(ctx context.Context, subscripti
return err
}
// 失效订阅缓存
s.InvalidateSubCache(sub.UserID, sub.GroupID)
if s.billingCacheService != nil {
userID, groupID := sub.UserID, sub.GroupID
go func() {
cacheCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_ = s.billingCacheService.InvalidateSubscription(cacheCtx, userID, groupID)
}()
if err := s.invalidateSubscriptionCaches(sub.UserID, sub.GroupID); err != nil {
return err
}
return nil
@@ -779,10 +815,7 @@ func (s *SubscriptionService) AdminResetQuota(ctx context.Context, subscriptionI
// Invalidate L1 ristretto cache. Ristretto's Del() is asynchronous by design,
// so call Wait() immediately after to flush pending operations and guarantee
// the deleted key is not returned on the very next Get() call.
s.InvalidateSubCache(sub.UserID, sub.GroupID)
if s.subCacheL1 != nil {
s.subCacheL1.Wait()
}
s.InvalidateSubCacheSync(sub.UserID, sub.GroupID)
if s.billingCacheService != nil {
_ = s.billingCacheService.InvalidateSubscription(ctx, sub.UserID, sub.GroupID)
}
@@ -25,6 +25,7 @@ type UserSubscription struct {
CreatedAt time.Time
UpdatedAt time.Time
DeletedAt *time.Time
User *User
Group *Group
+1 -1
View File
@@ -117,7 +117,7 @@ export async function extend(
* @returns Success confirmation
*/
export async function revoke(id: number): Promise<{ message: string }> {
const { data } = await apiClient.delete<{ message: string }>(`/admin/subscriptions/${id}`)
const { data } = await apiClient.post<{ message: string }>(`/admin/subscriptions/${id}/revoke`)
return data
}
+1
View File
@@ -1612,6 +1612,7 @@ export interface UserSubscription {
monthly_window_start: string | null
created_at: string
updated_at: string
revoked_at?: string | null
expires_at: string | null
user?: User
group?: Group