diff --git a/internal/handler/ai.go b/internal/handler/ai.go index 1f1bbb7..67abab1 100644 --- a/internal/handler/ai.go +++ b/internal/handler/ai.go @@ -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) } } diff --git a/internal/handler/handler.go b/internal/handler/handler.go index 7c6aece..39eaf3e 100644 --- a/internal/handler/handler.go +++ b/internal/handler/handler.go @@ -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 ── diff --git a/internal/handler/media.go b/internal/handler/media.go index 38c5cd5..2706577 100644 --- a/internal/handler/media.go +++ b/internal/handler/media.go @@ -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 diff --git a/internal/handler/play_profile.go b/internal/handler/play_profile.go index d7eff40..edda468 100644 --- a/internal/handler/play_profile.go +++ b/internal/handler/play_profile.go @@ -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), + }) + } +} diff --git a/internal/handler/playback.go b/internal/handler/playback.go index 86475f8..5b661c3 100644 --- a/internal/handler/playback.go +++ b/internal/handler/playback.go @@ -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) } } diff --git a/internal/handler/playback_extra.go b/internal/handler/playback_extra.go index 604da4c..048efce 100644 --- a/internal/handler/playback_extra.go +++ b/internal/handler/playback_extra.go @@ -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)) diff --git a/internal/handler/search_extra.go b/internal/handler/search_extra.go index ada2f84..0375f4f 100644 --- a/internal/handler/search_extra.go +++ b/internal/handler/search_extra.go @@ -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"), diff --git a/internal/handler/streaming.go b/internal/handler/streaming.go index 59e7200..6e9b04e 100644 --- a/internal/handler/streaming.go +++ b/internal/handler/streaming.go @@ -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 diff --git a/internal/handler/visibility.go b/internal/handler/visibility.go new file mode 100644 index 0000000..2692c96 --- /dev/null +++ b/internal/handler/visibility.go @@ -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)) +} diff --git a/internal/handler/watch_history.go b/internal/handler/watch_history.go index 9b86392..e6f3e8d 100644 --- a/internal/handler/watch_history.go +++ b/internal/handler/watch_history.go @@ -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)) diff --git a/internal/repository/repository.go b/internal/repository/repository.go index 425a038..46e0b5a 100644 --- a/internal/repository/repository.go +++ b/internal/repository/repository.go @@ -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) diff --git a/internal/service/media.go b/internal/service/media.go index 0c56ee8..cd81d85 100644 --- a/internal/service/media.go +++ b/internal/service/media.go @@ -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. diff --git a/internal/service/media_visibility_test.go b/internal/service/media_visibility_test.go new file mode 100644 index 0000000..d0c369b --- /dev/null +++ b/internal/service/media_visibility_test.go @@ -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 +} diff --git a/internal/service/play_profile.go b/internal/service/play_profile.go index 299f18f..599435d 100644 --- a/internal/service/play_profile.go +++ b/internal/service/play_profile.go @@ -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") diff --git a/internal/service/play_profile_test.go b/internal/service/play_profile_test.go new file mode 100644 index 0000000..6d44f99 --- /dev/null +++ b/internal/service/play_profile_test.go @@ -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") + } +} diff --git a/web/src/api/client.ts b/web/src/api/client.ts index 1d8e851..280a7d0 100644 --- a/web/src/api/client.ts +++ b/web/src/api/client.ts @@ -1,6 +1,7 @@ import axios, { AxiosError, type InternalAxiosRequestConfig } from 'axios' import { useAuthStore } from '../stores/auth' +import { getActivePlayProfileId, getActivePlayProfilePinToken } from '../stores/playProfile' // Single axios instance used by every API helper. Adds the JWT to outgoing // requests and routes 401s back to the login page. @@ -31,6 +32,15 @@ api.interceptors.request.use((config) => { config.headers = config.headers ?? {} config.headers.Authorization = `Bearer ${token}` } + const activeProfileId = getActivePlayProfileId() + if (activeProfileId) { + config.headers = config.headers ?? {} + config.headers['X-Play-Profile-ID'] = activeProfileId + const pinToken = getActivePlayProfilePinToken() + if (pinToken) { + config.headers['X-Play-Profile-PIN-Token'] = pinToken + } + } return config }) @@ -91,16 +101,25 @@ const tokenQuery = () => { return `token=${encodeURIComponent(t)}` } +const profileQuery = () => { + const id = getActivePlayProfileId() + if (!id) return '' + const pinToken = getActivePlayProfilePinToken() + return `&profile_id=${encodeURIComponent(id)}${ + pinToken ? `&profile_pin_token=${encodeURIComponent(pinToken)}` : '' + }` +} + // streamURL returns a direct-play URL for