mirror of
https://github.com/truewhile/MeBox.git
synced 2026-10-02 20:26:36 +08:00
fix: enforce adult profile pin visibility
This commit is contained in:
@@ -23,6 +23,29 @@ type MediaService struct {
|
||||
repo *repository.Container
|
||||
}
|
||||
|
||||
type MediaVisibility struct {
|
||||
IncludeNSFW bool
|
||||
AllowedLibraryIDs []string
|
||||
}
|
||||
|
||||
func (v MediaVisibility) Allows(media *model.Media) bool {
|
||||
if media == nil {
|
||||
return false
|
||||
}
|
||||
if !v.IncludeNSFW && media.NSFW {
|
||||
return false
|
||||
}
|
||||
if len(v.AllowedLibraryIDs) == 0 {
|
||||
return true
|
||||
}
|
||||
for _, id := range v.AllowedLibraryIDs {
|
||||
if id == media.LibraryID {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// NewMediaService is the constructor.
|
||||
func NewMediaService(cfg *config.Config, log *zap.Logger, repo *repository.Container) *MediaService {
|
||||
return &MediaService{cfg: cfg, log: log, repo: repo}
|
||||
@@ -143,6 +166,10 @@ func (s *MediaService) DeleteLibrary(ctx context.Context, id string) error {
|
||||
|
||||
// ListMedia paginates media items inside a library.
|
||||
func (s *MediaService) ListMedia(ctx context.Context, libraryID string, page, pageSize int) ([]model.Media, int64, error) {
|
||||
return s.ListMediaVisible(ctx, libraryID, page, pageSize, MediaVisibility{IncludeNSFW: true})
|
||||
}
|
||||
|
||||
func (s *MediaService) ListMediaVisible(ctx context.Context, libraryID string, page, pageSize int, visibility MediaVisibility) ([]model.Media, int64, error) {
|
||||
if pageSize <= 0 {
|
||||
pageSize = 50
|
||||
}
|
||||
@@ -152,15 +179,25 @@ func (s *MediaService) ListMedia(ctx context.Context, libraryID string, page, pa
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
return s.repo.Media.ListByLibrary(ctx, libraryID, (page-1)*pageSize, pageSize)
|
||||
return s.repo.Media.ListByLibraryFiltered(ctx, libraryID, (page-1)*pageSize, pageSize, repository.MediaQueryFilter{
|
||||
IncludeNSFW: visibility.IncludeNSFW,
|
||||
AllowedLibraryIDs: visibility.AllowedLibraryIDs,
|
||||
})
|
||||
}
|
||||
|
||||
// SearchMedia performs a simple LIKE search across titles.
|
||||
func (s *MediaService) SearchMedia(ctx context.Context, query string, limit int) ([]model.Media, error) {
|
||||
return s.SearchMediaVisible(ctx, query, limit, MediaVisibility{IncludeNSFW: true})
|
||||
}
|
||||
|
||||
func (s *MediaService) SearchMediaVisible(ctx context.Context, query string, limit int, visibility MediaVisibility) ([]model.Media, error) {
|
||||
if limit <= 0 || limit > 200 {
|
||||
limit = 50
|
||||
}
|
||||
return s.repo.Media.Search(ctx, query, limit)
|
||||
return s.repo.Media.SearchFiltered(ctx, query, limit, repository.MediaQueryFilter{
|
||||
IncludeNSFW: visibility.IncludeNSFW,
|
||||
AllowedLibraryIDs: visibility.AllowedLibraryIDs,
|
||||
})
|
||||
}
|
||||
|
||||
// GetMedia returns a single media row.
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/config"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestMediaVisibilityFiltersNSFWAndLibraries(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
svc := NewMediaService(&config.Config{}, zap.NewNop(), repos)
|
||||
|
||||
libA := model.Library{Name: "电影", Path: "/media/movies", Type: "movie", Enabled: true}
|
||||
libB := model.Library{Name: "成人", Path: "/media/adult", Type: "movie", Enabled: true}
|
||||
if err := db.Create(&libA).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.Create(&libB).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rows := []model.Media{
|
||||
{LibraryID: libA.ID, Title: "普通电影", Path: "/media/movies/a.mkv"},
|
||||
{LibraryID: libA.ID, Title: "成人电影", Path: "/media/movies/b.mkv", NSFW: true},
|
||||
{LibraryID: libB.ID, Title: "限制媒体库电影", Path: "/media/adult/c.mkv"},
|
||||
}
|
||||
if err := db.Create(&rows).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
items, err := svc.SearchMediaVisible(t.Context(), "电影", 20, MediaVisibility{IncludeNSFW: false})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := sortedMediaTitles(items); !slices.Equal(got, []string{"普通电影", "限制媒体库电影"}) {
|
||||
t.Fatalf("NSFW-filtered search = %#v", got)
|
||||
}
|
||||
|
||||
items, err = svc.SearchMediaVisible(t.Context(), "电影", 20, MediaVisibility{
|
||||
IncludeNSFW: true,
|
||||
AllowedLibraryIDs: []string{libA.ID},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := sortedMediaTitles(items); !slices.Equal(got, []string{"成人电影", "普通电影"}) {
|
||||
t.Fatalf("library-filtered search = %#v", got)
|
||||
}
|
||||
|
||||
listed, total, err := svc.ListMediaVisible(t.Context(), libA.ID, 1, 20, MediaVisibility{IncludeNSFW: false})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if total != 1 || len(listed) != 1 || listed[0].Title != "普通电影" {
|
||||
t.Fatalf("NSFW-filtered list total=%d rows=%#v", total, sortedMediaTitles(listed))
|
||||
}
|
||||
}
|
||||
|
||||
func sortedMediaTitles(rows []model.Media) []string {
|
||||
out := make([]string, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
out = append(out, row.Title)
|
||||
}
|
||||
slices.Sort(out)
|
||||
return out
|
||||
}
|
||||
@@ -29,6 +29,12 @@ type PlayProfileService struct {
|
||||
repo *repository.Container
|
||||
}
|
||||
|
||||
var (
|
||||
ErrPlayProfileNotFound = errors.New("profile not found")
|
||||
ErrPlayProfileForbidden = errors.New("profile forbidden")
|
||||
ErrPlayProfilePINInvalid = errors.New("pin invalid")
|
||||
)
|
||||
|
||||
// NewPlayProfileService is the constructor.
|
||||
func NewPlayProfileService(log *zap.Logger, repo *repository.Container) *PlayProfileService {
|
||||
return &PlayProfileService{log: log, repo: repo}
|
||||
@@ -138,7 +144,7 @@ func (s *PlayProfileService) Update(ctx context.Context, id string, in PlayProfi
|
||||
return nil, err
|
||||
}
|
||||
if row == nil {
|
||||
return nil, errors.New("profile not found")
|
||||
return nil, ErrPlayProfileNotFound
|
||||
}
|
||||
if err := validateProfileInput(in, false); err != nil {
|
||||
return nil, err
|
||||
@@ -158,6 +164,8 @@ func (s *PlayProfileService) Update(ctx context.Context, id string, in PlayProfi
|
||||
}
|
||||
if in.RequirePIN && in.PIN != "" {
|
||||
patch["pin_hash"] = hashPIN(in.PIN)
|
||||
} else if in.RequirePIN && row.PINHash == "" {
|
||||
return nil, errors.New("pin required")
|
||||
}
|
||||
if !in.RequirePIN {
|
||||
patch["pin_hash"] = ""
|
||||
@@ -183,6 +191,27 @@ func (s *PlayProfileService) Delete(ctx context.Context, id string) error {
|
||||
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)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if row == nil {
|
||||
return nil, ErrPlayProfileNotFound
|
||||
}
|
||||
if row.UserID != userID {
|
||||
return nil, ErrPlayProfileForbidden
|
||||
}
|
||||
if row.RequirePIN {
|
||||
if row.PINHash == "" || hashPIN(pin) != row.PINHash {
|
||||
return nil, ErrPlayProfilePINInvalid
|
||||
}
|
||||
}
|
||||
view := toProfileView(*row)
|
||||
return &view, nil
|
||||
}
|
||||
|
||||
// TouchActive bumps the LastActiveAt timestamp; called by the player
|
||||
// when a profile is selected.
|
||||
func (s *PlayProfileService) TouchActive(ctx context.Context, id string) error {
|
||||
@@ -201,6 +230,9 @@ func validateProfileInput(in PlayProfileInput, requireUser bool) error {
|
||||
if requireUser && strings.TrimSpace(in.UserID) == "" {
|
||||
return errors.New("user_id required")
|
||||
}
|
||||
if requireUser && in.RequirePIN && strings.TrimSpace(in.PIN) == "" {
|
||||
return errors.New("pin required")
|
||||
}
|
||||
if in.RequirePIN && in.PIN != "" {
|
||||
if len(in.PIN) < 4 || len(in.PIN) > 8 {
|
||||
return errors.New("pin must be 4-8 characters")
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestPlayProfileVerifyPIN(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))
|
||||
profile, err := service.Create(t.Context(), PlayProfileInput{
|
||||
UserID: "user-1",
|
||||
Name: "成人模式",
|
||||
AllowAdult: true,
|
||||
RequirePIN: true,
|
||||
PIN: "1234",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if _, err := service.VerifyPIN(t.Context(), profile.ID, "user-1", "0000"); !errors.Is(err, ErrPlayProfilePINInvalid) {
|
||||
t.Fatalf("wrong PIN error = %v", err)
|
||||
}
|
||||
if _, err := service.VerifyPIN(t.Context(), profile.ID, "user-2", "1234"); !errors.Is(err, ErrPlayProfileForbidden) {
|
||||
t.Fatalf("wrong owner error = %v", err)
|
||||
}
|
||||
if verified, err := service.VerifyPIN(t.Context(), profile.ID, "user-1", "1234"); err != nil || verified.ID != profile.ID {
|
||||
t.Fatalf("verify PIN got profile=%v err=%v", verified, err)
|
||||
}
|
||||
}
|
||||
|
||||
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))
|
||||
if _, err := service.Create(t.Context(), PlayProfileInput{
|
||||
UserID: "user-1",
|
||||
Name: "锁定模式",
|
||||
RequirePIN: true,
|
||||
}); err == nil {
|
||||
t.Fatal("expected PIN-required profile create to fail without PIN")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user