fix: enforce adult profile pin visibility

This commit is contained in:
ShukeBta
2026-05-29 15:16:20 +08:00
parent 931db7a242
commit 633a8cf715
21 changed files with 839 additions and 39 deletions
+3 -2
View File
@@ -29,7 +29,7 @@ func smartSearchHandler(svc *service.Container) gin.HandlerFunc {
}
// Run the actual library search using the cleaned query so the
// caller can render local + external results in one round-trip.
items, _ := svc.Media.SearchMedia(c.Request.Context(), intent.Query, 60)
items, _ := svc.Media.SearchMediaVisible(c.Request.Context(), intent.Query, 60, mediaVisibilityForRequest(c, svc))
external := service.SearchExternalMedia(
c.Request.Context(),
intent.Query,
@@ -57,8 +57,9 @@ func aiRecommendHandler(svc *service.Container) gin.HandlerFunc {
return
}
titles := make([]string, 0, len(hist))
visibility := mediaVisibilityForRequest(c, svc)
for _, h := range hist {
if h.Media != nil && strings.TrimSpace(h.Media.Title) != "" {
if h.Media != nil && visibility.Allows(h.Media) && strings.TrimSpace(h.Media.Title) != "" {
titles = append(titles, h.Media.Title)
}
}
+1
View File
@@ -200,6 +200,7 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
authed.GET("/play-profiles", listPlayProfilesHandler(svc))
authed.POST("/play-profiles", createPlayProfileHandler(svc))
authed.PUT("/play-profiles/:id", updatePlayProfileHandler(svc))
authed.POST("/play-profiles/:id/verify-pin", verifyPlayProfilePINHandler(svc))
authed.DELETE("/play-profiles/:id", deletePlayProfileHandler(svc))
// ── Search aliases ──
+12 -3
View File
@@ -83,7 +83,7 @@ func listMediaHandler(svc *service.Container) gin.HandlerFunc {
id := c.Param("id")
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("page_size", "50"))
items, total, err := svc.Media.ListMedia(c.Request.Context(), id, page, size)
items, total, err := svc.Media.ListMediaVisible(c.Request.Context(), id, page, size, mediaVisibilityForRequest(c, svc))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -108,6 +108,10 @@ func getMediaHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
if !mediaVisibleForRequest(c, svc, m) {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
c.JSON(http.StatusOK, m)
}
}
@@ -116,7 +120,7 @@ func searchMediaHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
q := c.Query("q")
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "50"))
items, err := svc.Media.SearchMedia(c.Request.Context(), q, limit)
items, err := svc.Media.SearchMediaVisible(c.Request.Context(), q, limit, mediaVisibilityForRequest(c, svc))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -127,7 +131,12 @@ func searchMediaHandler(svc *service.Container) gin.HandlerFunc {
func streamHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
err := svc.Stream.ServeFile(c.Writer, c.Request, c.Param("id"))
m, err := svc.Media.GetMedia(c.Request.Context(), c.Param("id"))
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
err = svc.Stream.ServeFile(c.Writer, c.Request, c.Param("id"))
if errors.Is(err, service.ErrMediaNotFound) {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
+38
View File
@@ -5,7 +5,9 @@
package handler
import (
"errors"
"net/http"
"time"
"github.com/gin-gonic/gin"
@@ -13,6 +15,10 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/service"
)
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.
func listPlayProfilesHandler(svc *service.Container) gin.HandlerFunc {
@@ -84,3 +90,35 @@ func deletePlayProfileHandler(svc *service.Container) gin.HandlerFunc {
c.Status(http.StatusNoContent)
}
}
func verifyPlayProfilePINHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req verifyPlayProfilePINReq
_ = c.ShouldBindJSON(&req)
uid, _ := c.Get(middleware.CtxUserID)
profile, err := svc.PlayProfiles.VerifyPIN(c.Request.Context(), c.Param("id"), toString(uid), req.PIN)
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 errors.Is(err, service.ErrPlayProfilePINInvalid) {
c.JSON(http.StatusUnauthorized, gin.H{"error": "PIN 错误"})
return
}
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
expiresAt := time.Now().Add(12 * time.Hour)
token := signPlayProfilePINToken(svc, toString(uid), profile.ID, expiresAt)
c.JSON(http.StatusOK, gin.H{
"profile": profile,
"token": token,
"expires_at": expiresAt.Format(time.RFC3339),
})
}
}
+24 -2
View File
@@ -44,7 +44,14 @@ func recentHistoryHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"items": items})
visibility := mediaVisibilityForRequest(c, svc)
filtered := make([]service.HistoryItem, 0, len(items))
for _, item := range items {
if item.Media == nil || visibility.Allows(item.Media) {
filtered = append(filtered, item)
}
}
c.JSON(http.StatusOK, gin.H{"items": filtered})
}
}
@@ -72,7 +79,14 @@ func listFavouritesHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"items": items})
visibility := mediaVisibilityForRequest(c, svc)
filtered := make([]any, 0, len(items))
for i := range items {
if visibility.Allows(&items[i]) {
filtered = append(filtered, items[i])
}
}
c.JSON(http.StatusOK, gin.H{"items": filtered})
}
}
@@ -127,6 +141,14 @@ func getPlaylistHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusForbidden, gin.H{"error": "forbidden"})
return
}
visibility := mediaVisibilityForRequest(c, svc)
filtered := detail.Items[:0]
for i := range detail.Items {
if visibility.Allows(&detail.Items[i]) {
filtered = append(filtered, detail.Items[i])
}
}
detail.Items = filtered
c.JSON(http.StatusOK, detail)
}
}
+28 -7
View File
@@ -24,14 +24,16 @@ import (
func playbackInfoHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || m == nil {
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
return
}
token := externalPlaybackToken(c, svc)
profileQuery := externalProfileQuery(c)
c.JSON(http.StatusOK, gin.H{
"media": m,
"stream_url": "/api/stream/" + m.ID,
"hls_url": "/api/hls/" + m.ID + "/index.m3u8",
"stream_url": "/api/stream/" + m.ID + "?token=" + url.QueryEscape(token) + profileQuery,
"hls_url": "/api/hls/" + m.ID + "/index.m3u8?token=" + url.QueryEscape(token) + profileQuery,
})
}
}
@@ -67,12 +69,12 @@ func playbackProgressHandler(svc *service.Container) gin.HandlerFunc {
func externalPlayersHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || m == nil {
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
return
}
token := externalPlaybackToken(c, svc)
streamURL := absoluteRequestURL(c, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token))
streamURL := absoluteRequestURL(c, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c))
escapedStream := url.QueryEscape(streamURL)
c.JSON(http.StatusOK, gin.H{
"url": streamURL,
@@ -92,18 +94,37 @@ func externalPlayersHandler(svc *service.Container) gin.HandlerFunc {
func externalURLHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || m == nil {
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
return
}
token := externalPlaybackToken(c, svc)
c.JSON(http.StatusOK, gin.H{
"url": absoluteRequestURL(c, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)),
"url": absoluteRequestURL(c, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c)),
"token": token,
})
}
}
func externalProfileQuery(c *gin.Context) string {
profileID := strings.TrimSpace(c.GetHeader("X-Play-Profile-ID"))
if profileID == "" {
profileID = strings.TrimSpace(c.Query("profile_id"))
}
if profileID == "" {
return ""
}
query := "&profile_id=" + url.QueryEscape(profileID)
pinToken := strings.TrimSpace(c.GetHeader("X-Play-Profile-PIN-Token"))
if pinToken == "" {
pinToken = strings.TrimSpace(c.Query("profile_pin_token"))
}
if pinToken != "" {
query += "&profile_pin_token=" + url.QueryEscape(pinToken)
}
return query
}
func externalPlaybackToken(c *gin.Context, svc *service.Container) string {
uid, _ := c.Get(middleware.CtxUserID)
u, err := svc.Repo.User.FindByID(c.Request.Context(), toString(uid))
+3 -3
View File
@@ -21,7 +21,7 @@ func searchUnifiedHandler(svc *service.Container) gin.HandlerFunc {
if limit <= 0 || limit > 200 {
limit = 30
}
items, err := svc.Media.SearchMedia(c.Request.Context(), q, limit)
items, err := svc.Media.SearchMediaVisible(c.Request.Context(), q, limit, mediaVisibilityForRequest(c, svc))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -41,13 +41,13 @@ func searchAdvancedHandler(svc *service.Container) gin.HandlerFunc {
if limit <= 0 || limit > 200 {
limit = 30
}
items, err := svc.Media.SearchMedia(c.Request.Context(), q, limit)
items, err := svc.Media.SearchMediaVisible(c.Request.Context(), q, limit, mediaVisibilityForRequest(c, svc))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{
"items": items,
"items": items,
"filters": gin.H{
"year": c.Query("year"),
"type": c.Query("type"),
+12 -2
View File
@@ -13,7 +13,12 @@ import (
func hlsPlaylistHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
err := svc.Stream.ServeHLSPlaylist(c.Writer, c.Request, c.Param("id"))
m, err := svc.Media.GetMedia(c.Request.Context(), c.Param("id"))
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
err = svc.Stream.ServeHLSPlaylist(c.Writer, c.Request, c.Param("id"))
if errors.Is(err, service.ErrMediaNotFound) {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
@@ -35,7 +40,12 @@ func hlsPlaylistHandler(svc *service.Container) gin.HandlerFunc {
func hlsSegmentHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
err := svc.Stream.ServeHLSSegment(c.Writer, c.Request, c.Param("id"), c.Param("seg"))
m, err := svc.Media.GetMedia(c.Request.Context(), c.Param("id"))
if err != nil || m == nil || !mediaVisibleForRequest(c, svc, m) {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
err = svc.Stream.ServeHLSSegment(c.Writer, c.Request, c.Param("id"), c.Param("seg"))
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": err.Error()})
return
+160
View File
@@ -0,0 +1,160 @@
package handler
import (
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func mediaVisibilityForRequest(c *gin.Context, svc *service.Container) service.MediaVisibility {
adultEnabled := settingBool(c, svc, "adult.enabled", false)
visibility := service.MediaVisibility{IncludeNSFW: adultEnabled}
profile, locked := selectedPlayProfile(c, svc)
if locked {
return service.MediaVisibility{
IncludeNSFW: false,
AllowedLibraryIDs: []string{"__locked__"},
}
}
if profile == nil {
return visibility
}
visibility.IncludeNSFW = adultEnabled && profile.AllowAdult
visibility.AllowedLibraryIDs = profileAllowedLibraryIDs(*profile)
return visibility
}
func selectedPlayProfile(c *gin.Context, svc *service.Container) (*model.PlayProfile, bool) {
if svc == nil || svc.Repo == nil || svc.Repo.PlayProfile == nil {
return nil, false
}
userID := currentUserID(c)
if userID == "" {
return nil, false
}
profileID := strings.TrimSpace(c.GetHeader("X-Play-Profile-ID"))
if profileID == "" {
profileID = strings.TrimSpace(c.Query("profile_id"))
}
if profileID != "" {
profile, err := svc.Repo.PlayProfile.FindByID(c.Request.Context(), profileID)
if err == nil && profile != nil && profile.UserID == userID {
if profile.RequirePIN && !validPlayProfilePINToken(c, svc, userID, profile.ID) {
return nil, true
}
return profile, false
}
}
rows, err := svc.Repo.PlayProfile.ListByUser(c.Request.Context(), userID)
if err != nil {
return nil, false
}
for i := range rows {
if rows[i].IsDefault {
if rows[i].RequirePIN && !validPlayProfilePINToken(c, svc, userID, rows[i].ID) {
return nil, true
}
return &rows[i], false
}
}
return nil, false
}
func mediaVisibleForRequest(c *gin.Context, svc *service.Container, media *model.Media) bool {
return mediaVisibilityForRequest(c, svc).Allows(media)
}
func settingBool(c *gin.Context, svc *service.Container, key string, fallback bool) bool {
if svc == nil || svc.Repo == nil || svc.Repo.Setting == nil {
return fallback
}
value, err := svc.Repo.Setting.Get(c.Request.Context(), key)
if err != nil {
return fallback
}
switch strings.ToLower(strings.TrimSpace(value)) {
case "1", "true", "yes", "on", "enabled", "启用", "开启":
return true
case "0", "false", "no", "off", "disabled", "禁用", "关闭", "":
return false
default:
return fallback
}
}
func currentUserID(c *gin.Context) string {
uid, _ := c.Get(middleware.CtxUserID)
return toString(uid)
}
func profileAllowedLibraryIDs(profile model.PlayProfile) []string {
if strings.TrimSpace(profile.AllowedLibraryIDs) == "" {
return nil
}
var ids []string
if err := json.Unmarshal([]byte(profile.AllowedLibraryIDs), &ids); err != nil {
return nil
}
return ids
}
func signPlayProfilePINToken(svc *service.Container, userID, profileID string, expiresAt time.Time) string {
if svc == nil || svc.Cfg == nil {
return ""
}
payload := fmt.Sprintf("%s|%s|%d", userID, profileID, expiresAt.Unix())
encodedPayload := base64.RawURLEncoding.EncodeToString([]byte(payload))
signature := playProfilePINSignature(svc.Cfg.Secrets.JWTSecret, encodedPayload)
if signature == "" {
return ""
}
return encodedPayload + "." + signature
}
func validPlayProfilePINToken(c *gin.Context, svc *service.Container, userID, profileID string) bool {
token := strings.TrimSpace(c.GetHeader("X-Play-Profile-PIN-Token"))
if token == "" {
token = strings.TrimSpace(c.Query("profile_pin_token"))
}
parts := strings.Split(token, ".")
if len(parts) != 2 || parts[0] == "" || parts[1] == "" {
return false
}
expectedSignature := playProfilePINSignature(svc.Cfg.Secrets.JWTSecret, parts[0])
if expectedSignature == "" || !hmac.Equal([]byte(expectedSignature), []byte(parts[1])) {
return false
}
payloadBytes, err := base64.RawURLEncoding.DecodeString(parts[0])
if err != nil {
return false
}
fields := strings.Split(string(payloadBytes), "|")
if len(fields) != 3 || fields[0] != userID || fields[1] != profileID {
return false
}
expiresUnix, err := strconv.ParseInt(fields[2], 10, 64)
if err != nil {
return false
}
return time.Now().Unix() <= expiresUnix
}
func playProfilePINSignature(secret, encodedPayload string) string {
if strings.TrimSpace(secret) == "" || encodedPayload == "" {
return ""
}
mac := hmac.New(sha256.New, []byte(secret))
_, _ = mac.Write([]byte(encodedPayload))
return base64.RawURLEncoding.EncodeToString(mac.Sum(nil))
}
+16 -6
View File
@@ -3,11 +3,11 @@
// The base /history GET / POST routes already exist; these add the three
// auxiliary surfaces the React WatchHistoryPage needs:
//
// GET /api/watch-history paginated list (admin sees every user)
// GET /api/watch-history/stats aggregate watch time + completion
// GET /api/watch-history/continue resume rail (incomplete only)
// DELETE /api/watch-history clear (?media_item_id= optional)
// DELETE /api/watch-history/:id remove one row
// GET /api/watch-history paginated list (admin sees every user)
// GET /api/watch-history/stats aggregate watch time + completion
// GET /api/watch-history/continue resume rail (incomplete only)
// DELETE /api/watch-history clear (?media_item_id= optional)
// DELETE /api/watch-history/:id remove one row
package handler
import (
@@ -36,7 +36,14 @@ func historyListHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, items)
visibility := mediaVisibilityForRequest(c, svc)
filtered := make([]service.HistoryItem, 0, len(items))
for _, item := range items {
if item.Media == nil || visibility.Allows(item.Media) {
filtered = append(filtered, item)
}
}
c.JSON(http.StatusOK, filtered)
}
}
@@ -109,6 +116,9 @@ func historyContinueHandler(svc *service.Container) gin.HandlerFunc {
}
mIdx := make(map[string]model.Media, len(media))
for _, m := range media {
if !mediaVisibleForRequest(c, svc, &m) {
continue
}
mIdx[m.ID] = m
}
out := make([]gin.H, 0, len(rows))
+27
View File
@@ -207,6 +207,23 @@ func (r *LibraryRepository) Delete(ctx context.Context, id string) error {
// MediaRepository persists model.Media records.
type MediaRepository struct{ db *gorm.DB }
// MediaQueryFilter is applied to user-facing media queries so NSFW items and
// profile-restricted libraries are filtered in SQL instead of only in React.
type MediaQueryFilter struct {
IncludeNSFW bool
AllowedLibraryIDs []string
}
func applyMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB {
if !filter.IncludeNSFW {
q = q.Where("nsfw = ?", false)
}
if len(filter.AllowedLibraryIDs) > 0 {
q = q.Where("library_id IN ?", filter.AllowedLibraryIDs)
}
return q
}
// Upsert inserts or updates a media row keyed by Path (unique index).
//
// 重要:当一条行已经存在时,scanner 重扫只应该刷新文件级元数据
@@ -336,9 +353,14 @@ func (r *MediaRepository) FindByID(ctx context.Context, id string) (*model.Media
// ListByLibrary returns paginated media items for a library.
func (r *MediaRepository) ListByLibrary(ctx context.Context, libraryID string, offset, limit int) ([]model.Media, int64, error) {
return r.ListByLibraryFiltered(ctx, libraryID, offset, limit, MediaQueryFilter{IncludeNSFW: true})
}
func (r *MediaRepository) ListByLibraryFiltered(ctx context.Context, libraryID string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
var items []model.Media
var total int64
q := r.db.WithContext(ctx).Model(&model.Media{}).Where("library_id = ?", libraryID)
q = applyMediaQueryFilter(q, filter)
if err := q.Count(&total).Error; err != nil {
return nil, 0, err
}
@@ -349,8 +371,13 @@ func (r *MediaRepository) ListByLibrary(ctx context.Context, libraryID string, o
// Search runs a LIKE search against the title field. Empty query returns the
// most recently added items.
func (r *MediaRepository) Search(ctx context.Context, query string, limit int) ([]model.Media, error) {
return r.SearchFiltered(ctx, query, limit, MediaQueryFilter{IncludeNSFW: true})
}
func (r *MediaRepository) SearchFiltered(ctx context.Context, query string, limit int, filter MediaQueryFilter) ([]model.Media, error) {
var items []model.Media
q := r.db.WithContext(ctx).Model(&model.Media{}).Limit(limit)
q = applyMediaQueryFilter(q, filter)
if query != "" {
like := "%" + query + "%"
q = q.Where("title LIKE ? OR original_name LIKE ?", like, like)
+39 -2
View File
@@ -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.
+78
View File
@@ -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
}
+33 -1
View File
@@ -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")
+61
View File
@@ -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")
}
}