家庭组 权益修改
Build docker and publish / build (20.15.1) (push) Successful in 8m16s

This commit is contained in:
2026-03-04 22:02:42 -08:00
parent 3594097d47
commit 4349a7ea2f
28 changed files with 960 additions and 96 deletions
@@ -0,0 +1,88 @@
package common
import (
"context"
modelUser "github.com/perfect-panel/server/internal/model/user"
"github.com/perfect-panel/server/pkg/xerr"
"github.com/pkg/errors"
"gorm.io/gorm"
)
const (
EntitlementSourceSelf = "self"
EntitlementSourceFamilyOwner = "family_owner"
)
type EntitlementContext struct {
EffectiveUserID int64
Source string
OwnerUserID int64
ReadOnly bool
}
type familyEntitlementRelation struct {
Role uint8 `gorm:"column:role"`
FamilyStatus uint8 `gorm:"column:family_status"`
OwnerUserID int64 `gorm:"column:owner_user_id"`
}
func ResolveEntitlementUser(ctx context.Context, db *gorm.DB, currentUserID int64) (*EntitlementContext, error) {
entitlement := buildEntitlementContext(currentUserID, nil)
if currentUserID <= 0 {
return entitlement, nil
}
var relation familyEntitlementRelation
err := db.WithContext(ctx).
Table("user_family_member").
Select("user_family_member.role, user_family.status AS family_status, user_family.owner_user_id").
Joins("JOIN user_family ON user_family.id = user_family_member.family_id AND user_family.deleted_at IS NULL").
Where("user_family_member.user_id = ? AND user_family_member.deleted_at IS NULL AND user_family_member.status = ?", currentUserID, modelUser.FamilyMemberActive).
First(&relation).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return entitlement, nil
}
return nil, errors.Wrapf(xerr.NewErrCode(xerr.DatabaseQueryError), "query family entitlement relation failed")
}
return buildEntitlementContext(currentUserID, &relation), nil
}
func DenyIfFamilyMemberReadonly(ctx context.Context, db *gorm.DB, currentUserID int64) error {
entitlement, err := ResolveEntitlementUser(ctx, db, currentUserID)
if err != nil {
return err
}
return denyReadonlyEntitlement(entitlement)
}
func buildEntitlementContext(currentUserID int64, relation *familyEntitlementRelation) *EntitlementContext {
entitlement := &EntitlementContext{
EffectiveUserID: currentUserID,
Source: EntitlementSourceSelf,
}
if relation == nil {
return entitlement
}
if relation.Role == modelUser.FamilyRoleMember &&
relation.FamilyStatus == modelUser.FamilyStatusActive &&
relation.OwnerUserID > 0 &&
relation.OwnerUserID != currentUserID {
return &EntitlementContext{
EffectiveUserID: relation.OwnerUserID,
Source: EntitlementSourceFamilyOwner,
OwnerUserID: relation.OwnerUserID,
ReadOnly: true,
}
}
return entitlement
}
func denyReadonlyEntitlement(entitlement *EntitlementContext) error {
if entitlement != nil && entitlement.ReadOnly {
return errors.Wrapf(xerr.NewErrCode(xerr.FamilyOwnerOperationForbidden), "family member operation is forbidden")
}
return nil
}
@@ -0,0 +1,78 @@
package common
import (
stderrors "errors"
"testing"
modelUser "github.com/perfect-panel/server/internal/model/user"
"github.com/perfect-panel/server/pkg/xerr"
pkgerrors "github.com/pkg/errors"
"github.com/stretchr/testify/require"
)
func extractFamilyEntitlementCode(err error) uint32 {
if err == nil {
return 0
}
var codeErr *xerr.CodeError
if stderrors.As(pkgerrors.Cause(err), &codeErr) {
return codeErr.GetErrCode()
}
return 0
}
func TestBuildEntitlementContext(t *testing.T) {
t.Run("default self entitlement", func(t *testing.T) {
entitlement := buildEntitlementContext(1001, nil)
require.Equal(t, int64(1001), entitlement.EffectiveUserID)
require.Equal(t, EntitlementSourceSelf, entitlement.Source)
require.Equal(t, int64(0), entitlement.OwnerUserID)
require.False(t, entitlement.ReadOnly)
})
t.Run("active family member uses owner entitlement", func(t *testing.T) {
entitlement := buildEntitlementContext(1001, &familyEntitlementRelation{
Role: modelUser.FamilyRoleMember,
FamilyStatus: modelUser.FamilyStatusActive,
OwnerUserID: 2001,
})
require.Equal(t, int64(2001), entitlement.EffectiveUserID)
require.Equal(t, EntitlementSourceFamilyOwner, entitlement.Source)
require.Equal(t, int64(2001), entitlement.OwnerUserID)
require.True(t, entitlement.ReadOnly)
})
t.Run("owner relation keeps self entitlement", func(t *testing.T) {
entitlement := buildEntitlementContext(2001, &familyEntitlementRelation{
Role: modelUser.FamilyRoleOwner,
FamilyStatus: modelUser.FamilyStatusActive,
OwnerUserID: 2001,
})
require.Equal(t, int64(2001), entitlement.EffectiveUserID)
require.Equal(t, EntitlementSourceSelf, entitlement.Source)
require.False(t, entitlement.ReadOnly)
})
t.Run("disabled family keeps self entitlement", func(t *testing.T) {
entitlement := buildEntitlementContext(1001, &familyEntitlementRelation{
Role: modelUser.FamilyRoleMember,
FamilyStatus: 0,
OwnerUserID: 2001,
})
require.Equal(t, int64(1001), entitlement.EffectiveUserID)
require.Equal(t, EntitlementSourceSelf, entitlement.Source)
require.False(t, entitlement.ReadOnly)
})
}
func TestDenyReadonlyEntitlement(t *testing.T) {
require.NoError(t, denyReadonlyEntitlement(&EntitlementContext{ReadOnly: false}))
err := denyReadonlyEntitlement(&EntitlementContext{
Source: EntitlementSourceFamilyOwner,
ReadOnly: true,
})
require.Error(t, err)
require.Equal(t, xerr.FamilyOwnerOperationForbidden, extractFamilyEntitlementCode(err))
}
+289
View File
@@ -0,0 +1,289 @@
package common
import (
"context"
"encoding/json"
"fmt"
"net/url"
"strings"
"sync"
"time"
"github.com/perfect-panel/server/internal/svc"
"github.com/perfect-panel/server/pkg/kutt"
)
const inviteShortLinkCachePrefix = "cache:invite:short_link:"
type inviteLinkCustomData struct {
ShareURL string `json:"shareUrl"`
Domain string `json:"domain"`
}
type InviteLinkResolver struct {
ctx context.Context
svcCtx *svc.ServiceContext
createShortLink func(ctx context.Context, targetURL, domain string) (string, error)
}
func NewInviteLinkResolver(ctx context.Context, svcCtx *svc.ServiceContext) *InviteLinkResolver {
resolver := &InviteLinkResolver{
ctx: ctx,
svcCtx: svcCtx,
}
resolver.createShortLink = func(ctx context.Context, targetURL, domain string) (string, error) {
client := kutt.NewClient(svcCtx.Config.Kutt.ApiURL, svcCtx.Config.Kutt.ApiKey)
link, err := client.CreateShortLink(ctx, &kutt.CreateLinkRequest{
Target: targetURL,
Reuse: true,
Domain: domain,
})
if err != nil {
return "", err
}
shortLink := strings.TrimSpace(link.Link)
if strings.HasPrefix(shortLink, "http://") {
shortLink = strings.Replace(shortLink, "http://", "https://", 1)
}
return shortLink, nil
}
return resolver
}
func (r *InviteLinkResolver) ResolveInviteLink(referCode string) string {
normalizedCode := strings.TrimSpace(referCode)
if normalizedCode == "" {
return ""
}
longLink := r.buildLongInviteLink(normalizedCode)
if !r.canUseKutt() || longLink == "" {
return longLink
}
if cached := r.getCachedShortLink(normalizedCode); cached != "" {
return cached
}
shortLink, err := r.generateShortLinkWithTimeout(normalizedCode, 1500*time.Millisecond)
if err != nil || strings.TrimSpace(shortLink) == "" {
return longLink
}
r.cacheShortLink(normalizedCode, shortLink)
return shortLink
}
func (r *InviteLinkResolver) ResolveInviteLinksBatch(referCodes []string, maxGenerate, maxConcurrency int, timeout time.Duration) map[string]string {
result := make(map[string]string)
uniqueCodes := uniqueReferCodes(referCodes)
if len(uniqueCodes) == 0 {
return result
}
for _, referCode := range uniqueCodes {
result[referCode] = r.buildLongInviteLink(referCode)
}
if !r.canUseKutt() {
return result
}
toGenerate := make([]string, 0, len(uniqueCodes))
for _, referCode := range uniqueCodes {
if cached := r.getCachedShortLink(referCode); cached != "" {
result[referCode] = cached
continue
}
toGenerate = append(toGenerate, referCode)
}
if maxGenerate > 0 && len(toGenerate) > maxGenerate {
toGenerate = toGenerate[:maxGenerate]
}
if len(toGenerate) == 0 {
return result
}
if maxConcurrency <= 0 {
maxConcurrency = 1
}
if timeout <= 0 {
timeout = 1500 * time.Millisecond
}
limiter := make(chan struct{}, maxConcurrency)
var waitGroup sync.WaitGroup
var mutex sync.Mutex
for _, referCode := range toGenerate {
waitGroup.Add(1)
currentCode := referCode
go func() {
defer waitGroup.Done()
limiter <- struct{}{}
defer func() { <-limiter }()
shortLink, err := r.generateShortLinkWithTimeout(currentCode, timeout)
if err != nil || strings.TrimSpace(shortLink) == "" {
return
}
mutex.Lock()
result[currentCode] = shortLink
mutex.Unlock()
r.cacheShortLink(currentCode, shortLink)
}()
}
waitGroup.Wait()
return result
}
func (r *InviteLinkResolver) canUseKutt() bool {
if r == nil || r.svcCtx == nil {
return false
}
if !r.svcCtx.Config.Kutt.Enable {
return false
}
if strings.TrimSpace(r.svcCtx.Config.Kutt.ApiURL) == "" || strings.TrimSpace(r.svcCtx.Config.Kutt.ApiKey) == "" {
return false
}
return r.createShortLink != nil
}
func (r *InviteLinkResolver) resolveShareURLAndDomain() (string, string) {
if r == nil || r.svcCtx == nil {
return "", ""
}
shareURL := strings.TrimSpace(r.svcCtx.Config.Kutt.TargetURL)
domain := strings.TrimSpace(r.svcCtx.Config.Kutt.Domain)
customData := strings.TrimSpace(r.svcCtx.Config.Site.CustomData)
if customData == "" {
return shareURL, domain
}
var parsedData inviteLinkCustomData
if err := json.Unmarshal([]byte(customData), &parsedData); err != nil {
return shareURL, domain
}
if strings.TrimSpace(parsedData.ShareURL) != "" {
shareURL = strings.TrimSpace(parsedData.ShareURL)
}
if strings.TrimSpace(parsedData.Domain) != "" {
domain = strings.TrimSpace(parsedData.Domain)
}
return shareURL, domain
}
func (r *InviteLinkResolver) buildLongInviteLink(referCode string) string {
normalizedCode := strings.TrimSpace(referCode)
if normalizedCode == "" {
return ""
}
shareURL, _ := r.resolveShareURLAndDomain()
if shareURL == "" {
return ""
}
parsedURL, err := url.Parse(shareURL)
if err != nil {
return fallbackLongInviteLink(shareURL, normalizedCode)
}
queryValues := parsedURL.Query()
queryValues.Set("ic", normalizedCode)
parsedURL.RawQuery = queryValues.Encode()
return parsedURL.String()
}
func (r *InviteLinkResolver) generateShortLinkWithTimeout(referCode string, timeout time.Duration) (string, error) {
longLink := r.buildLongInviteLink(referCode)
if longLink == "" {
return "", nil
}
_, domain := r.resolveShareURLAndDomain()
requestCtx := r.ctx
var cancel context.CancelFunc
if timeout > 0 {
requestCtx, cancel = context.WithTimeout(r.ctx, timeout)
defer cancel()
}
shortLink, err := r.createShortLink(requestCtx, longLink, domain)
if err != nil {
return "", err
}
return strings.TrimSpace(shortLink), nil
}
func (r *InviteLinkResolver) getCachedShortLink(referCode string) string {
if r == nil || r.svcCtx == nil || r.svcCtx.Redis == nil {
return ""
}
cacheKey := inviteShortLinkCachePrefix + referCode
shortLink, err := r.svcCtx.Redis.Get(r.ctx, cacheKey).Result()
if err != nil {
return ""
}
return strings.TrimSpace(shortLink)
}
func (r *InviteLinkResolver) cacheShortLink(referCode, shortLink string) {
if r == nil || r.svcCtx == nil || r.svcCtx.Redis == nil {
return
}
if strings.TrimSpace(referCode) == "" || strings.TrimSpace(shortLink) == "" {
return
}
cacheKey := inviteShortLinkCachePrefix + referCode
_ = r.svcCtx.Redis.Set(r.ctx, cacheKey, shortLink, 0).Err()
}
func uniqueReferCodes(referCodes []string) []string {
uniqueCodes := make([]string, 0, len(referCodes))
seen := make(map[string]struct{}, len(referCodes))
for _, referCode := range referCodes {
normalizedCode := strings.TrimSpace(referCode)
if normalizedCode == "" {
continue
}
if _, exists := seen[normalizedCode]; exists {
continue
}
seen[normalizedCode] = struct{}{}
uniqueCodes = append(uniqueCodes, normalizedCode)
}
return uniqueCodes
}
func fallbackLongInviteLink(baseURL, referCode string) string {
normalizedBase := strings.TrimSpace(baseURL)
normalizedCode := strings.TrimSpace(referCode)
if normalizedBase == "" || normalizedCode == "" {
return ""
}
separator := "?"
if strings.Contains(normalizedBase, "?") {
separator = "&"
}
trimmedBase := strings.TrimRight(normalizedBase, "?&")
return fmt.Sprintf("%s%sic=%s", trimmedBase, separator, url.QueryEscape(normalizedCode))
}
@@ -0,0 +1,145 @@
package common
import (
"context"
"errors"
"net/url"
"testing"
"github.com/alicebob/miniredis/v2"
"github.com/perfect-panel/server/internal/config"
"github.com/perfect-panel/server/internal/svc"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/require"
)
func buildInviteResolverForTest(t *testing.T, cfg config.Config) (*InviteLinkResolver, *miniredis.Miniredis) {
t.Helper()
redisServer, err := miniredis.Run()
require.NoError(t, err)
t.Cleanup(func() {
redisServer.Close()
})
redisClient := redis.NewClient(&redis.Options{
Addr: redisServer.Addr(),
DB: 0,
})
t.Cleanup(func() {
_ = redisClient.Close()
})
serviceCtx := &svc.ServiceContext{
Config: cfg,
Redis: redisClient,
}
resolver := NewInviteLinkResolver(context.Background(), serviceCtx)
return resolver, redisServer
}
func TestInviteLinkResolverResolveInviteLink(t *testing.T) {
t.Run("kutt disabled returns long link", func(t *testing.T) {
cfg := config.Config{}
cfg.Kutt.TargetURL = "https://example.com/register"
resolver, _ := buildInviteResolverForTest(t, cfg)
link := resolver.ResolveInviteLink("abc123")
require.Equal(t, "https://example.com/register?ic=abc123", link)
})
t.Run("cache hit returns cached short link", func(t *testing.T) {
cfg := config.Config{}
cfg.Kutt.Enable = true
cfg.Kutt.ApiURL = "https://kutt.local/api/v2"
cfg.Kutt.ApiKey = "token"
cfg.Kutt.TargetURL = "https://example.com/register"
resolver, redisServer := buildInviteResolverForTest(t, cfg)
redisServer.Set(inviteShortLinkCachePrefix+"abc123", "https://sho.rt/cached")
called := 0
resolver.createShortLink = func(ctx context.Context, targetURL, domain string) (string, error) {
called++
return "", errors.New("should not call createShortLink on cache hit")
}
link := resolver.ResolveInviteLink("abc123")
require.Equal(t, "https://sho.rt/cached", link)
require.Equal(t, 0, called)
})
t.Run("cache miss kutt success returns short link and writes cache", func(t *testing.T) {
cfg := config.Config{}
cfg.Kutt.Enable = true
cfg.Kutt.ApiURL = "https://kutt.local/api/v2"
cfg.Kutt.ApiKey = "token"
cfg.Kutt.TargetURL = "https://example.com/register"
resolver, _ := buildInviteResolverForTest(t, cfg)
resolver.createShortLink = func(ctx context.Context, targetURL, domain string) (string, error) {
return "https://sho.rt/new", nil
}
link := resolver.ResolveInviteLink("abc123")
require.Equal(t, "https://sho.rt/new", link)
cached := resolver.getCachedShortLink("abc123")
require.Equal(t, "https://sho.rt/new", cached)
})
t.Run("kutt failure falls back to long link", func(t *testing.T) {
cfg := config.Config{}
cfg.Kutt.Enable = true
cfg.Kutt.ApiURL = "https://kutt.local/api/v2"
cfg.Kutt.ApiKey = "token"
cfg.Kutt.TargetURL = "https://example.com/register"
resolver, _ := buildInviteResolverForTest(t, cfg)
resolver.createShortLink = func(ctx context.Context, targetURL, domain string) (string, error) {
return "", errors.New("kutt request failed")
}
link := resolver.ResolveInviteLink("abc123")
require.Equal(t, "https://example.com/register?ic=abc123", link)
})
t.Run("long link preserves existing query string", func(t *testing.T) {
cfg := config.Config{}
cfg.Kutt.TargetURL = "https://example.com/register?channel=ios"
resolver, _ := buildInviteResolverForTest(t, cfg)
link := resolver.ResolveInviteLink("abc123")
parsed, err := url.Parse(link)
require.NoError(t, err)
require.Equal(t, "https", parsed.Scheme)
require.Equal(t, "example.com", parsed.Host)
require.Equal(t, "/register", parsed.Path)
require.Equal(t, "ios", parsed.Query().Get("channel"))
require.Equal(t, "abc123", parsed.Query().Get("ic"))
})
t.Run("kutt target preserves existing query string", func(t *testing.T) {
cfg := config.Config{}
cfg.Kutt.Enable = true
cfg.Kutt.ApiURL = "https://kutt.local/api/v2"
cfg.Kutt.ApiKey = "token"
cfg.Kutt.TargetURL = "https://example.com/register?channel=ios"
resolver, _ := buildInviteResolverForTest(t, cfg)
capturedTargetURL := ""
resolver.createShortLink = func(ctx context.Context, targetURL, domain string) (string, error) {
capturedTargetURL = targetURL
return "https://sho.rt/query", nil
}
link := resolver.ResolveInviteLink("abc123")
require.Equal(t, "https://sho.rt/query", link)
parsed, err := url.Parse(capturedTargetURL)
require.NoError(t, err)
require.Equal(t, "ios", parsed.Query().Get("channel"))
require.Equal(t, "abc123", parsed.Query().Get("ic"))
})
}