mirror of
https://github.com/truewhile/MeBox.git
synced 2026-10-01 12:06:38 +08:00
fix: isolate play profiles per user
This commit is contained in:
@@ -196,7 +196,7 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
|
||||
authed.GET("/stats/libraries", statsLibrariesHandler(svc))
|
||||
authed.GET("/stats/monitor", statsMonitorHandler(svc))
|
||||
|
||||
// Multi-persona play profiles (caller-scoped, admins via ?all=true).
|
||||
// Multi-persona play profiles (caller-scoped).
|
||||
authed.GET("/play-profiles", listPlayProfilesHandler(svc))
|
||||
authed.POST("/play-profiles", createPlayProfileHandler(svc))
|
||||
authed.PUT("/play-profiles/:id", updatePlayProfileHandler(svc))
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
// Package handler — multi-persona play profile CRUD endpoints.
|
||||
//
|
||||
// Non-admin users see / mutate only their own profiles. Admins see
|
||||
// every profile so they can manage child accounts, etc.
|
||||
// Every user sees / mutates only their own profiles. Admin user
|
||||
// management belongs in a separate admin surface, not this switcher.
|
||||
package handler
|
||||
|
||||
import (
|
||||
@@ -19,21 +19,10 @@ type verifyPlayProfilePINReq struct {
|
||||
PIN string `json:"pin"`
|
||||
}
|
||||
|
||||
// listPlayProfilesHandler returns the caller's profiles, or every
|
||||
// profile when the caller is an admin AND ?all=true is set.
|
||||
// listPlayProfilesHandler returns only the caller's own profiles.
|
||||
func listPlayProfilesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
if c.Query("all") == "true" && role == "admin" {
|
||||
rows, err := svc.PlayProfiles.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, rows)
|
||||
return
|
||||
}
|
||||
rows, err := svc.PlayProfiles.ListByUser(c.Request.Context(), toString(uid))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
@@ -50,13 +39,13 @@ func createPlayProfileHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
// Default the user_id to the caller; admins can override.
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
if in.UserID == "" || role != "admin" {
|
||||
in.UserID = toString(uid)
|
||||
}
|
||||
in.UserID = toString(uid)
|
||||
row, err := svc.PlayProfiles.Create(c.Request.Context(), in)
|
||||
if errors.Is(err, service.ErrPlayProfileLimit) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "每个用户最多只能创建 3 个观影 Profile"})
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -72,7 +61,16 @@ func updatePlayProfileHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
row, err := svc.PlayProfiles.Update(c.Request.Context(), c.Param("id"), in)
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
row, err := svc.PlayProfiles.UpdateForUser(c.Request.Context(), c.Param("id"), toString(uid), in)
|
||||
if errors.Is(err, service.ErrPlayProfileNotFound) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "profile not found"})
|
||||
return
|
||||
}
|
||||
if errors.Is(err, service.ErrPlayProfileForbidden) {
|
||||
c.JSON(http.StatusForbidden, gin.H{"error": "profile forbidden"})
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -83,7 +81,14 @@ func updatePlayProfileHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
func deletePlayProfileHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.PlayProfiles.Delete(c.Request.Context(), c.Param("id")); err != nil {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
if err := svc.PlayProfiles.DeleteForUser(c.Request.Context(), c.Param("id"), toString(uid)); errors.Is(err, service.ErrPlayProfileNotFound) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "profile not found"})
|
||||
return
|
||||
} else if errors.Is(err, service.ErrPlayProfileForbidden) {
|
||||
c.JSON(http.StatusForbidden, gin.H{"error": "profile forbidden"})
|
||||
return
|
||||
} else if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
@@ -45,6 +45,14 @@ func (r *PlayProfileRepository) ListByUser(ctx context.Context, userID string) (
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// CountByUser returns the number of active profiles owned by a user.
|
||||
func (r *PlayProfileRepository) CountByUser(ctx context.Context, userID string) (int64, error) {
|
||||
var count int64
|
||||
err := r.db.WithContext(ctx).Model(&model.PlayProfile{}).
|
||||
Where("user_id = ?", userID).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
// Update applies a partial update to a profile row.
|
||||
func (r *PlayProfileRepository) Update(ctx context.Context, id string, patch map[string]any) error {
|
||||
return r.db.WithContext(ctx).Model(&model.PlayProfile{}).
|
||||
|
||||
@@ -122,6 +122,7 @@ func (e *EmbyService) FindUser(ctx context.Context, id string) (map[string]any,
|
||||
}
|
||||
|
||||
func (e *EmbyService) userPayload(u *model.User) map[string]any {
|
||||
canDownload := u.Role == "admin"
|
||||
return map[string]any{
|
||||
"Id": u.ID,
|
||||
"Name": u.Username,
|
||||
@@ -153,9 +154,9 @@ func (e *EmbyService) userPayload(u *model.User) map[string]any {
|
||||
"EnableVideoPlaybackTranscoding": true,
|
||||
"EnablePlaybackRemuxing": true,
|
||||
"EnableLiveTvAccess": false,
|
||||
"EnableContentDownloading": true,
|
||||
"EnableSyncTranscoding": true,
|
||||
"EnableMediaConversion": true,
|
||||
"EnableContentDownloading": canDownload,
|
||||
"EnableSyncTranscoding": canDownload,
|
||||
"EnableMediaConversion": canDownload,
|
||||
"EnableAllChannels": true,
|
||||
"EnableAllFolders": true,
|
||||
"EnableAllDevices": true,
|
||||
|
||||
@@ -127,13 +127,47 @@ func TestEmbyRootItemsExposeLibraries(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbyUserPolicyDisablesDownloadsForViewers(t *testing.T) {
|
||||
svc := newTestEmbyService(t)
|
||||
viewer := &model.User{Username: "viewer", Role: "user", Tier: "free", IsActive: true}
|
||||
admin := &model.User{Username: "admin", Role: "admin", Tier: "plus", IsActive: true}
|
||||
if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
|
||||
t.Fatalf("create viewer: %v", err)
|
||||
}
|
||||
if err := svc.repo.User.Create(t.Context(), admin); err != nil {
|
||||
t.Fatalf("create admin: %v", err)
|
||||
}
|
||||
|
||||
viewerPayload, err := svc.FindUser(t.Context(), viewer.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("viewer payload: %v", err)
|
||||
}
|
||||
adminPayload, err := svc.FindUser(t.Context(), admin.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("admin payload: %v", err)
|
||||
}
|
||||
viewerPolicy := viewerPayload["Policy"].(map[string]any)
|
||||
adminPolicy := adminPayload["Policy"].(map[string]any)
|
||||
if viewerPolicy["EnableMediaPlayback"] != true {
|
||||
t.Fatalf("viewer must keep playback enabled: %#v", viewerPolicy)
|
||||
}
|
||||
if viewerPolicy["EnableContentDownloading"] != false ||
|
||||
viewerPolicy["EnableSyncTranscoding"] != false ||
|
||||
viewerPolicy["EnableMediaConversion"] != false {
|
||||
t.Fatalf("viewer must not be allowed to download/sync media: %#v", viewerPolicy)
|
||||
}
|
||||
if adminPolicy["EnableContentDownloading"] != true {
|
||||
t.Fatalf("admin should keep downloading capability: %#v", adminPolicy)
|
||||
}
|
||||
}
|
||||
|
||||
func newTestEmbyService(t *testing.T) *EmbyService {
|
||||
t.Helper()
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}, &model.Favorite{}, &model.PlaybackHistory{}); err != nil {
|
||||
if err := db.AutoMigrate(&model.Library{}, &model.Series{}, &model.Media{}, &model.Favorite{}, &model.PlaybackHistory{}, &model.User{}); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
|
||||
@@ -29,10 +29,13 @@ type PlayProfileService struct {
|
||||
repo *repository.Container
|
||||
}
|
||||
|
||||
const MaxPlayProfilesPerUser = 3
|
||||
|
||||
var (
|
||||
ErrPlayProfileNotFound = errors.New("profile not found")
|
||||
ErrPlayProfileForbidden = errors.New("profile forbidden")
|
||||
ErrPlayProfilePINInvalid = errors.New("pin invalid")
|
||||
ErrPlayProfileLimit = errors.New("profile limit reached")
|
||||
)
|
||||
|
||||
// NewPlayProfileService is the constructor.
|
||||
@@ -108,6 +111,13 @@ func (s *PlayProfileService) Create(ctx context.Context, in PlayProfileInput) (*
|
||||
if err := validateProfileInput(in, true); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
count, err := s.repo.PlayProfile.CountByUser(ctx, in.UserID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if count >= MaxPlayProfilesPerUser {
|
||||
return nil, ErrPlayProfileLimit
|
||||
}
|
||||
libsBlob, _ := json.Marshal(in.AllowedLibraryIDs)
|
||||
p := &model.PlayProfile{
|
||||
UserID: in.UserID,
|
||||
@@ -137,6 +147,21 @@ func (s *PlayProfileService) Create(ctx context.Context, in PlayProfileInput) (*
|
||||
return &v, nil
|
||||
}
|
||||
|
||||
// UpdateForUser applies a patch only when the profile belongs to userID.
|
||||
func (s *PlayProfileService) UpdateForUser(ctx context.Context, id, userID string, in PlayProfileInput) (*ProfileView, error) {
|
||||
row, err := s.repo.PlayProfile.FindByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if row == nil {
|
||||
return nil, ErrPlayProfileNotFound
|
||||
}
|
||||
if row.UserID != userID {
|
||||
return nil, ErrPlayProfileForbidden
|
||||
}
|
||||
return s.updateExisting(ctx, row, in)
|
||||
}
|
||||
|
||||
// Update applies a patch to an existing profile.
|
||||
func (s *PlayProfileService) Update(ctx context.Context, id string, in PlayProfileInput) (*ProfileView, error) {
|
||||
row, err := s.repo.PlayProfile.FindByID(ctx, id)
|
||||
@@ -146,6 +171,10 @@ func (s *PlayProfileService) Update(ctx context.Context, id string, in PlayProfi
|
||||
if row == nil {
|
||||
return nil, ErrPlayProfileNotFound
|
||||
}
|
||||
return s.updateExisting(ctx, row, in)
|
||||
}
|
||||
|
||||
func (s *PlayProfileService) updateExisting(ctx context.Context, row *model.PlayProfile, in PlayProfileInput) (*ProfileView, error) {
|
||||
if err := validateProfileInput(in, false); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -175,10 +204,10 @@ func (s *PlayProfileService) Update(ctx context.Context, id string, in PlayProfi
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if err := s.repo.PlayProfile.Update(ctx, id, patch); err != nil {
|
||||
if err := s.repo.PlayProfile.Update(ctx, row.ID, patch); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
row, err = s.repo.PlayProfile.FindByID(ctx, id)
|
||||
row, err := s.repo.PlayProfile.FindByID(ctx, row.ID)
|
||||
if err != nil || row == nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -191,6 +220,21 @@ func (s *PlayProfileService) Delete(ctx context.Context, id string) error {
|
||||
return s.repo.PlayProfile.Delete(ctx, id)
|
||||
}
|
||||
|
||||
// DeleteForUser removes a profile only when it belongs to userID.
|
||||
func (s *PlayProfileService) DeleteForUser(ctx context.Context, id, userID string) error {
|
||||
row, err := s.repo.PlayProfile.FindByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if row == nil {
|
||||
return ErrPlayProfileNotFound
|
||||
}
|
||||
if row.UserID != userID {
|
||||
return ErrPlayProfileForbidden
|
||||
}
|
||||
return s.repo.PlayProfile.Delete(ctx, id)
|
||||
}
|
||||
|
||||
// VerifyPIN validates that the caller can switch to a PIN-protected profile.
|
||||
func (s *PlayProfileService) VerifyPIN(ctx context.Context, id, userID, pin string) (*ProfileView, error) {
|
||||
row, err := s.repo.PlayProfile.FindByID(ctx, id)
|
||||
|
||||
@@ -11,7 +11,8 @@ import (
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestPlayProfileVerifyPIN(t *testing.T) {
|
||||
func newPlayProfileTestService(t *testing.T) *PlayProfileService {
|
||||
t.Helper()
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -19,7 +20,11 @@ func TestPlayProfileVerifyPIN(t *testing.T) {
|
||||
if err := db.AutoMigrate(&model.PlayProfile{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
service := NewPlayProfileService(zap.NewNop(), repository.New(db))
|
||||
return NewPlayProfileService(zap.NewNop(), repository.New(db))
|
||||
}
|
||||
|
||||
func TestPlayProfileVerifyPIN(t *testing.T) {
|
||||
service := newPlayProfileTestService(t)
|
||||
profile, err := service.Create(t.Context(), PlayProfileInput{
|
||||
UserID: "user-1",
|
||||
Name: "成人模式",
|
||||
@@ -43,14 +48,7 @@ func TestPlayProfileVerifyPIN(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestPlayProfileCreateRequiresPINWhenEnabled(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.PlayProfile{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
service := NewPlayProfileService(zap.NewNop(), repository.New(db))
|
||||
service := newPlayProfileTestService(t)
|
||||
if _, err := service.Create(t.Context(), PlayProfileInput{
|
||||
UserID: "user-1",
|
||||
Name: "锁定模式",
|
||||
@@ -59,3 +57,65 @@ func TestPlayProfileCreateRequiresPINWhenEnabled(t *testing.T) {
|
||||
t.Fatal("expected PIN-required profile create to fail without PIN")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlayProfileCreateLimitIsPerUser(t *testing.T) {
|
||||
service := newPlayProfileTestService(t)
|
||||
|
||||
for i := 1; i <= MaxPlayProfilesPerUser; i++ {
|
||||
if _, err := service.Create(t.Context(), PlayProfileInput{
|
||||
UserID: "user-1",
|
||||
Name: "模式 " + string(rune('0'+i)),
|
||||
}); err != nil {
|
||||
t.Fatalf("create profile %d: %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := service.Create(t.Context(), PlayProfileInput{
|
||||
UserID: "user-1",
|
||||
Name: "超限模式",
|
||||
}); !errors.Is(err, ErrPlayProfileLimit) {
|
||||
t.Fatalf("expected limit error, got %v", err)
|
||||
}
|
||||
|
||||
if _, err := service.Create(t.Context(), PlayProfileInput{
|
||||
UserID: "user-2",
|
||||
Name: "另一个用户的模式",
|
||||
}); err != nil {
|
||||
t.Fatalf("different user should have independent limit: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlayProfileUpdateDeleteRequireOwner(t *testing.T) {
|
||||
service := newPlayProfileTestService(t)
|
||||
profile, err := service.Create(t.Context(), PlayProfileInput{
|
||||
UserID: "user-1",
|
||||
Name: "私人模式",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if _, err := service.UpdateForUser(t.Context(), profile.ID, "user-2", PlayProfileInput{
|
||||
Name: "越权修改",
|
||||
}); !errors.Is(err, ErrPlayProfileForbidden) {
|
||||
t.Fatalf("expected forbidden update, got %v", err)
|
||||
}
|
||||
|
||||
if err := service.DeleteForUser(t.Context(), profile.ID, "user-2"); !errors.Is(err, ErrPlayProfileForbidden) {
|
||||
t.Fatalf("expected forbidden delete, got %v", err)
|
||||
}
|
||||
|
||||
updated, err := service.UpdateForUser(t.Context(), profile.ID, "user-1", PlayProfileInput{
|
||||
Name: "已修改",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("owner update failed: %v", err)
|
||||
}
|
||||
if updated.Name != "已修改" || updated.UserID != "user-1" {
|
||||
t.Fatalf("unexpected updated profile: %+v", updated)
|
||||
}
|
||||
|
||||
if err := service.DeleteForUser(t.Context(), profile.ID, "user-1"); err != nil {
|
||||
t.Fatalf("owner delete failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user