fix: secure adult visibility and telegram bot access

This commit is contained in:
ShukeBta
2026-05-30 01:39:11 +08:00
parent db65e54c45
commit b5e11b6938
22 changed files with 913 additions and 114 deletions
+22 -3
View File
@@ -318,7 +318,11 @@ func embyFallbackUser(id string) gin.H {
func embyViewsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
out, err := svc.Emby.Views(c.Request.Context())
uid := c.Param("userId")
if uid == "" {
uid = embyUserID(c)
}
out, err := svc.Emby.Views(c.Request.Context(), uid)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -335,8 +339,13 @@ func embyVirtualFoldersHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
uid := embyUserID(c)
visibility := service.UserDefaultMediaVisibility(c.Request.Context(), svc.Repo, uid)
out := make([]gin.H, 0, len(libs))
for _, lib := range libs {
if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, lib, visibility) {
continue
}
collectionType := "movies"
switch lib.Type {
case "tv", "anime", "variety":
@@ -538,7 +547,11 @@ func embyShowEpisodesHandler(svc *service.Container) gin.HandlerFunc {
func embyPlaybackInfoHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
out, err := svc.Emby.PlaybackInfo(c.Request.Context(), c.Param("id"))
uid := c.Param("userId")
if uid == "" {
uid = embyUserID(c)
}
out, err := svc.Emby.PlaybackInfo(c.Request.Context(), c.Param("id"), uid)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -555,8 +568,14 @@ func embyPlaybackInfoHandler(svc *service.Container) gin.HandlerFunc {
// 直接代理到我们的 /api/stream/{id}(同一个 ServeFile)。
func embyVideoStreamHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := embyUserID(c)
item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
if err != nil || item == nil {
c.Status(http.StatusNotFound)
return
}
// 直接调用 Stream service 写入 response
err := svc.Stream.ServeFile(c.Writer, c.Request, c.Param("id"))
err = svc.Stream.ServeFile(c.Writer, c.Request, c.Param("id"))
if err != nil {
c.Status(http.StatusNotFound)
}
+43 -1
View File
@@ -19,6 +19,11 @@ type verifyPlayProfilePINReq struct {
PIN string `json:"pin"`
}
type deletePlayProfileReq struct {
PIN string `json:"pin"`
Password string `json:"password"`
}
// listPlayProfilesHandler returns only the caller's own profiles.
func listPlayProfilesHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
@@ -82,7 +87,44 @@ func updatePlayProfileHandler(svc *service.Container) gin.HandlerFunc {
func deletePlayProfileHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid, _ := c.Get(middleware.CtxUserID)
if err := svc.PlayProfiles.DeleteForUser(c.Request.Context(), c.Param("id"), toString(uid)); errors.Is(err, service.ErrPlayProfileNotFound) {
userID := toString(uid)
var req deletePlayProfileReq
_ = c.ShouldBindJSON(&req)
profile, err := svc.Repo.PlayProfile.FindByID(c.Request.Context(), c.Param("id"))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if profile == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "profile not found"})
return
}
if profile.UserID != userID {
c.JSON(http.StatusForbidden, gin.H{"error": "profile forbidden"})
return
}
verified := false
if profile.RequirePIN && req.PIN != "" {
if _, err := svc.PlayProfiles.VerifyPIN(c.Request.Context(), profile.ID, userID, req.PIN); err == nil {
verified = true
}
}
if !verified && req.Password != "" {
if err := svc.Auth.VerifyPassword(c.Request.Context(), userID, req.Password); err == nil {
verified = true
}
}
if !verified {
if profile.RequirePIN {
c.JSON(http.StatusUnauthorized, gin.H{"error": "删除此 Profile 需要输入 PIN 或当前账号密码"})
} else {
c.JSON(http.StatusUnauthorized, gin.H{"error": "删除此 Profile 需要输入当前账号密码"})
}
return
}
if err := svc.PlayProfiles.DeleteForUser(c.Request.Context(), profile.ID, userID); errors.Is(err, service.ErrPlayProfileNotFound) {
c.JSON(http.StatusNotFound, gin.H{"error": "profile not found"})
return
} else if errors.Is(err, service.ErrPlayProfileForbidden) {
+13 -1
View File
@@ -2,6 +2,7 @@
package handler
import (
"errors"
"net/http"
"github.com/gin-gonic/gin"
@@ -18,8 +19,19 @@ func updateProfileHandler(svc *service.Container) gin.HandlerFunc {
return
}
uid, _ := c.Get(middleware.CtxUserID)
u, err := svc.Profile.UpdateProfile(c.Request.Context(), uid.(string), patch)
userID := uid.(string)
if patch.HideAdult != nil {
if err := svc.Auth.VerifyPassword(c.Request.Context(), userID, patch.Password); err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "需要输入当前账号密码确认"})
return
}
}
u, err := svc.Profile.UpdateProfile(c.Request.Context(), userID, patch)
if err != nil {
if errors.Is(err, service.ErrUsernameTaken) {
c.JSON(http.StatusConflict, gin.H{"error": "username already taken"})
return
}
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
+6 -12
View File
@@ -4,7 +4,6 @@ import (
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"strconv"
"strings"
@@ -18,8 +17,10 @@ import (
)
func mediaVisibilityForRequest(c *gin.Context, svc *service.Container) service.MediaVisibility {
adultEnabled := settingBool(c, svc, "adult.enabled", false)
visibility := service.MediaVisibility{IncludeNSFW: adultEnabled}
userID := currentUserID(c)
adultEnabled := service.AdultContentEnabled(c.Request.Context(), svc.Repo)
userHidesAdult := service.UserHidesAdult(c.Request.Context(), svc.Repo, userID)
visibility := service.UserDefaultMediaVisibility(c.Request.Context(), svc.Repo, userID)
profile, locked := selectedPlayProfile(c, svc)
if locked {
return service.MediaVisibility{
@@ -30,7 +31,7 @@ func mediaVisibilityForRequest(c *gin.Context, svc *service.Container) service.M
if profile == nil {
return visibility
}
visibility.IncludeNSFW = adultEnabled && profile.AllowAdult
visibility.IncludeNSFW = adultEnabled && profile.AllowAdult && !userHidesAdult
visibility.AllowedLibraryIDs = profileAllowedLibraryIDs(*profile)
return visibility
}
@@ -99,14 +100,7 @@ func currentUserID(c *gin.Context) string {
}
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
return service.DecodeAllowedLibraryIDs(profile.AllowedLibraryIDs)
}
func signPlayProfilePINToken(svc *service.Container, userID, profileID string, expiresAt time.Time) string {
+2
View File
@@ -40,6 +40,7 @@ type User struct {
Nickname string `gorm:"size:128" json:"nickname,omitempty"`
Email string `gorm:"size:128" json:"email,omitempty"`
AvatarURL string `gorm:"size:255" json:"avatar_url,omitempty"`
HideAdult bool `gorm:"default:false" json:"hide_adult"`
ForcePasswordReset bool `gorm:"default:false" json:"force_password_reset"`
IsActive bool `gorm:"default:true" json:"is_active"`
LastLoginAt *time.Time `json:"last_login_at,omitempty"`
@@ -318,6 +319,7 @@ func AllModels() []interface{} {
&ApiConfig{},
&DownloadClient{},
&NotifyChannel{},
&TelegramBinding{},
&STRMRecord{},
&PlayProfile{},
&StorageConfig{},
+12
View File
@@ -0,0 +1,12 @@
package model
// TelegramBinding links a Telegram account to a local MediaStationGo user.
// The binding is password-verified when /start is used, then reused for
// low-risk self-service actions such as toggling adult-library visibility.
type TelegramBinding struct {
Base
TelegramUserID int64 `gorm:"uniqueIndex;not null" json:"telegram_user_id"`
TelegramName string `gorm:"size:128" json:"telegram_name,omitempty"`
ChatID int64 `gorm:"index" json:"chat_id"`
UserID string `gorm:"index;size:36;not null" json:"user_id"`
}
+20
View File
@@ -159,6 +159,9 @@ func (s *AuthService) Login(ctx context.Context, username, password string) (*Lo
// ChangePassword updates the user password if the old one matches.
func (s *AuthService) ChangePassword(ctx context.Context, userID, oldPwd, newPwd string) error {
if strings.TrimSpace(newPwd) == "" || len(newPwd) < 6 {
return errors.New("new password must be at least 6 characters")
}
u, err := s.repo.User.FindByID(ctx, userID)
if err != nil {
return err
@@ -176,6 +179,23 @@ func (s *AuthService) ChangePassword(ctx context.Context, userID, oldPwd, newPwd
return s.repo.User.UpdatePassword(ctx, userID, hash)
}
// VerifyPassword checks a user's current password without mutating account
// state. It is used for sensitive self-service actions such as hiding adult
// libraries or deleting play profiles.
func (s *AuthService) VerifyPassword(ctx context.Context, userID, password string) error {
u, err := s.repo.User.FindByID(ctx, userID)
if err != nil {
return err
}
if u == nil || strings.TrimSpace(password) == "" {
return ErrInvalidCredentials
}
if err := bcrypt.CompareHashAndPassword([]byte(u.PasswordHash), []byte(password)); err != nil {
return ErrInvalidCredentials
}
return nil
}
// IssueToken signs a JWT for the given user (60min validity, includes tier).
func (s *AuthService) IssueToken(u *model.User) (string, error) {
claims := Claims{
+104 -21
View File
@@ -169,13 +169,17 @@ func (e *EmbyService) userPayload(u *model.User) map[string]any {
// ─── Views / MediaFolders ────────────────────────────────────────────────────
// Views 返回 Emby 中"虚拟根目录"——每个 library 一个条目。
func (e *EmbyService) Views(ctx context.Context) (map[string]any, error) {
func (e *EmbyService) Views(ctx context.Context, userID string) (map[string]any, error) {
libs, err := e.repo.Library.List(ctx)
if err != nil {
return nil, err
}
visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
items := make([]map[string]any, 0, len(libs))
for _, l := range libs {
if !LibraryVisibleForUser(ctx, e.repo, l, visibility) {
continue
}
items = append(items, e.libraryAsView(&l))
}
return map[string]any{"Items": items, "TotalRecordCount": len(items)}, nil
@@ -290,16 +294,16 @@ func (e *EmbyService) Items(ctx context.Context, p ItemsParams) (map[string]any,
}
if p.ParentID == "" && p.SearchTerm == "" && !p.Recursive && len(p.IncludeItemTypes) == 0 {
return e.Views(ctx)
return e.Views(ctx, p.UserID)
}
if season, ok, err := e.findSeasonGroup(ctx, p.ParentID); err != nil {
if season, ok, err := e.findSeasonGroup(ctx, p.ParentID, p.UserID); err != nil {
return nil, err
} else if ok {
return e.episodeItems(ctx, season.Episodes, p)
}
if series, ok, err := e.findSeriesGroup(ctx, p.ParentID); err != nil {
if series, ok, err := e.findSeriesGroup(ctx, p.ParentID, p.UserID); err != nil {
return nil, err
} else if ok {
if p.Recursive || containsItemType(p.IncludeItemTypes, "Episode") {
@@ -329,6 +333,7 @@ func (e *EmbyService) Items(ctx context.Context, p ItemsParams) (map[string]any,
func (e *EmbyService) mediaItems(ctx context.Context, p ItemsParams) (map[string]any, error) {
q := e.repo.DB.WithContext(ctx).Model(&model.Media{})
q = e.applyUserMediaVisibility(ctx, q, p.UserID)
if p.ParentID != "" {
q = q.Where("library_id = ? OR series_id = ?", p.ParentID, p.ParentID)
}
@@ -375,6 +380,7 @@ func (e *EmbyService) mediaItems(ctx context.Context, p ItemsParams) (map[string
}
func (e *EmbyService) episodeItems(ctx context.Context, rows []model.Media, p ItemsParams) (map[string]any, error) {
rows = e.filterMediaRowsForUser(ctx, rows, p.UserID)
if p.SearchTerm != "" {
filtered := rows[:0]
needle := strings.ToLower(p.SearchTerm)
@@ -428,14 +434,14 @@ func (e *EmbyService) payloadsForMedia(ctx context.Context, rows []model.Media,
// Item 单条目详情。
func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[string]any, error) {
if strings.HasPrefix(mediaID, embyVirtualSeasonPrefix) {
if season, ok, err := e.findSeasonGroup(ctx, mediaID); err != nil {
if season, ok, err := e.findSeasonGroup(ctx, mediaID, userID); err != nil {
return nil, err
} else if ok {
return e.seasonPayload(season), nil
}
}
if strings.HasPrefix(mediaID, embyVirtualSeriesPrefix) {
if series, ok, err := e.findSeriesGroup(ctx, mediaID); err != nil {
if series, ok, err := e.findSeriesGroup(ctx, mediaID, userID); err != nil {
return nil, err
} else if ok {
return e.seriesPayload(series), nil
@@ -446,13 +452,16 @@ func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[str
return nil, err
}
if m == nil {
if series, ok, err := e.findSeriesGroup(ctx, mediaID); err != nil {
if series, ok, err := e.findSeriesGroup(ctx, mediaID, userID); err != nil {
return nil, err
} else if ok {
return e.seriesPayload(series), nil
}
return nil, nil
}
if !UserDefaultMediaVisibility(ctx, e.repo, userID).Allows(m) {
return nil, nil
}
fav := false
pos := int64(0)
if userID != "" {
@@ -477,6 +486,7 @@ func (e *EmbyService) LatestItems(ctx context.Context, userID, parentID string,
limit = 20
}
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("deleted_at IS NULL")
q = e.applyUserMediaVisibility(ctx, q, userID)
if parentID != "" {
if episodic, err := e.libraryIsEpisodic(ctx, parentID); err == nil && episodic {
resp, err := e.seriesItemsForLibrary(ctx, parentID, ItemsParams{
@@ -540,7 +550,9 @@ func (e *EmbyService) ResumeItems(ctx context.Context, userID string, limit int)
posByID[h.MediaID] = h.PositionMs
}
var medias []model.Media
if err := e.repo.DB.WithContext(ctx).Where("id IN ?", ids).Find(&medias).Error; err != nil {
q := e.repo.DB.WithContext(ctx).Where("id IN ?", ids)
q = e.applyUserMediaVisibility(ctx, q, userID)
if err := q.Find(&medias).Error; err != nil {
return nil, err
}
// 维持时间倒序
@@ -638,6 +650,7 @@ func (e *EmbyService) itemPayload(m *model.Media, fav bool, posMs int64) map[str
func (e *EmbyService) seriesItemsForLibrary(ctx context.Context, libraryID string, p ItemsParams) (map[string]any, error) {
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("season_num > 0 OR episode_num > 0")
q = e.applyUserMediaVisibility(ctx, q, p.UserID)
if libraryID != "" {
q = q.Where("library_id = ?", libraryID)
}
@@ -677,12 +690,13 @@ func (e *EmbyService) libraryIsEpisodic(ctx context.Context, libraryID string) (
return count > 0, err
}
func (e *EmbyService) findSeriesGroup(ctx context.Context, id string) (embySeriesGroup, bool, error) {
func (e *EmbyService) findSeriesGroup(ctx context.Context, id, userID string) (embySeriesGroup, bool, error) {
if strings.TrimSpace(id) == "" {
return embySeriesGroup{}, false, nil
}
var rows []model.Media
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).Where("season_num > 0 OR episode_num > 0")
q = e.applyUserMediaVisibility(ctx, q, userID)
if !strings.HasPrefix(id, embyVirtualSeriesPrefix) {
q = q.Where("series_id = ?", id)
}
@@ -716,13 +730,15 @@ func (e *EmbyService) findSeriesGroup(ctx context.Context, id string) (embySerie
return embySeriesGroup{}, false, nil
}
func (e *EmbyService) findSeasonGroup(ctx context.Context, id string) (embySeasonGroup, bool, error) {
func (e *EmbyService) findSeasonGroup(ctx context.Context, id, userID string) (embySeasonGroup, bool, error) {
if strings.TrimSpace(id) == "" || !strings.HasPrefix(id, embyVirtualSeasonPrefix) {
return embySeasonGroup{}, false, nil
}
var rows []model.Media
if err := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("season_num > 0 OR episode_num > 0").
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("season_num > 0 OR episode_num > 0")
q = e.applyUserMediaVisibility(ctx, q, userID)
if err := q.
Order("season_num asc, episode_num asc, created_at asc").
Find(&rows).Error; err != nil {
return embySeasonGroup{}, false, err
@@ -900,14 +916,14 @@ func (e *EmbyService) ImageURL(ctx context.Context, id, imageType string) (strin
return backdrop
}
if strings.HasPrefix(id, embyVirtualSeasonPrefix) {
if season, ok, err := e.findSeasonGroup(ctx, id); err != nil {
if season, ok, err := e.findSeasonGroup(ctx, id, ""); err != nil {
return "", err
} else if ok {
return pick(season.Series.PosterURL, season.Series.BackdropURL), nil
}
}
if strings.HasPrefix(id, embyVirtualSeriesPrefix) {
if series, ok, err := e.findSeriesGroup(ctx, id); err != nil {
if series, ok, err := e.findSeriesGroup(ctx, id, ""); err != nil {
return "", err
} else if ok {
return pick(series.PosterURL, series.BackdropURL), nil
@@ -920,7 +936,7 @@ func (e *EmbyService) ImageURL(ctx context.Context, id, imageType string) (strin
if err != nil {
return "", err
}
if series, ok, err := e.findSeriesGroup(ctx, id); err != nil {
if series, ok, err := e.findSeriesGroup(ctx, id, ""); err != nil {
return "", err
} else if ok {
return pick(series.PosterURL, series.BackdropURL), nil
@@ -1050,6 +1066,66 @@ func emptyUserData() map[string]any {
}
}
func (e *EmbyService) applyUserMediaVisibility(ctx context.Context, q *gorm.DB, userID string) *gorm.DB {
visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
if !visibility.IncludeNSFW {
q = q.Where("nsfw = ?", false)
if hidden := e.hiddenLibraryIDs(ctx, visibility); len(hidden) > 0 {
q = q.Where("library_id NOT IN ?", hidden)
}
}
if len(visibility.AllowedLibraryIDs) > 0 {
q = q.Where("library_id IN ?", visibility.AllowedLibraryIDs)
}
return q
}
func (e *EmbyService) filterMediaRowsForUser(ctx context.Context, rows []model.Media, userID string) []model.Media {
visibility := UserDefaultMediaVisibility(ctx, e.repo, userID)
if visibility.IncludeNSFW && len(visibility.AllowedLibraryIDs) == 0 {
return rows
}
allowed := map[string]bool{}
for _, id := range visibility.AllowedLibraryIDs {
allowed[id] = true
}
hiddenLibraries := map[string]bool{}
for _, id := range e.hiddenLibraryIDs(ctx, visibility) {
hiddenLibraries[id] = true
}
out := rows[:0]
for _, row := range rows {
if row.NSFW && !visibility.IncludeNSFW {
continue
}
if hiddenLibraries[row.LibraryID] {
continue
}
if len(allowed) > 0 && !allowed[row.LibraryID] {
continue
}
out = append(out, row)
}
return out
}
func (e *EmbyService) hiddenLibraryIDs(ctx context.Context, visibility MediaVisibility) []string {
if visibility.IncludeNSFW {
return nil
}
libs, err := e.repo.Library.List(ctx)
if err != nil {
return nil
}
ids := make([]string, 0)
for _, lib := range libs {
if !LibraryVisibleForUser(ctx, e.repo, lib, visibility) {
ids = append(ids, lib.ID)
}
}
return ids
}
func minInt(a, b int) int {
if a < b {
return a
@@ -1067,8 +1143,8 @@ func maxInt(a, b int) int {
// ─── Playback ────────────────────────────────────────────────────────────────
// PlaybackInfo returns a PlaybackInfoResponse usable by Emby clients.
func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID string) (map[string]any, error) {
m, err := e.playableMedia(ctx, mediaID)
func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID, userID string) (map[string]any, error) {
m, err := e.playableMedia(ctx, mediaID, userID)
if err != nil || m == nil {
return nil, err
}
@@ -1078,18 +1154,25 @@ func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID string) (map[str
}, nil
}
func (e *EmbyService) playableMedia(ctx context.Context, id string) (*model.Media, error) {
if season, ok, err := e.findSeasonGroup(ctx, id); err != nil {
func (e *EmbyService) playableMedia(ctx context.Context, id, userID string) (*model.Media, error) {
if season, ok, err := e.findSeasonGroup(ctx, id, userID); err != nil {
return nil, err
} else if ok && len(season.Episodes) > 0 {
return &season.Episodes[0], nil
}
if series, ok, err := e.findSeriesGroup(ctx, id); err != nil {
if series, ok, err := e.findSeriesGroup(ctx, id, userID); err != nil {
return nil, err
} else if ok && len(series.Episodes) > 0 {
return &series.Episodes[0], nil
}
return e.repo.Media.FindByID(ctx, id)
m, err := e.repo.Media.FindByID(ctx, id)
if err != nil || m == nil {
return m, err
}
if !UserDefaultMediaVisibility(ctx, e.repo, userID).Allows(m) {
return nil, nil
}
return m, nil
}
// mediaSource 是 /Items 与 /PlaybackInfo 共享的 MediaSource 结构。
+39 -1
View File
@@ -87,7 +87,7 @@ func TestEmbyItemsExposeSeriesSeasonEpisodeHierarchy(t *testing.T) {
t.Fatalf("latest should be grouped by series: %#v", latest)
}
playback, err := svc.PlaybackInfo(t.Context(), seriesID)
playback, err := svc.PlaybackInfo(t.Context(), seriesID, "user-1")
if err != nil {
t.Fatalf("series playback fallback: %v", err)
}
@@ -161,6 +161,44 @@ func TestEmbyUserPolicyDisablesDownloadsForViewers(t *testing.T) {
}
}
func TestEmbyHidesAdultLibrariesForUserLock(t *testing.T) {
svc := newTestEmbyService(t)
viewer := &model.User{Username: "viewer", Role: "user", Tier: "free", IsActive: true, HideAdult: true}
if err := svc.repo.User.Create(t.Context(), viewer); err != nil {
t.Fatalf("create viewer: %v", err)
}
safe := model.Library{Name: "电影", Path: `/media/movies`, Type: "movie", Enabled: true}
adult := model.Library{Name: "9KG 成人", Path: `/media/9KG`, Type: "movie", Enabled: true}
if err := svc.repo.Library.Create(t.Context(), &safe); err != nil {
t.Fatalf("create safe library: %v", err)
}
if err := svc.repo.Library.Create(t.Context(), &adult); err != nil {
t.Fatalf("create adult library: %v", err)
}
if err := svc.repo.DB.Create(&model.Media{LibraryID: safe.ID, Title: "安全电影", Path: `/media/movies/a.mkv`}).Error; err != nil {
t.Fatalf("create safe media: %v", err)
}
if err := svc.repo.DB.Create(&model.Media{LibraryID: adult.ID, Title: "成人电影", Path: `/media/9KG/a.mkv`}).Error; err != nil {
t.Fatalf("create adult media: %v", err)
}
root, err := svc.Items(t.Context(), ItemsParams{UserID: viewer.ID, Limit: 50})
if err != nil {
t.Fatalf("root items: %v", err)
}
items := root["Items"].([]map[string]any)
if len(items) != 1 || items[0]["Name"] != "电影" {
t.Fatalf("adult library should be hidden: %#v", items)
}
adultItems, err := svc.Items(t.Context(), ItemsParams{UserID: viewer.ID, ParentID: adult.ID, Limit: 50})
if err != nil {
t.Fatalf("adult items: %v", err)
}
if got := adultItems["TotalRecordCount"]; got != int64(0) {
t.Fatalf("adult media should be hidden, total=%#v payload=%#v", got, adultItems)
}
}
func newTestEmbyService(t *testing.T) *EmbyService {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+29
View File
@@ -26,8 +26,12 @@ func NewProfileService(log *zap.Logger, repo *repository.Container) *ProfileServ
// ProfileUpdate is the patch object accepted by UpdateProfile. Empty
// fields are ignored so the same payload can be reused across screens.
type ProfileUpdate struct {
Username *string `json:"username,omitempty"`
Nickname *string `json:"nickname,omitempty"`
Email *string `json:"email,omitempty"`
AvatarURL *string `json:"avatar_url,omitempty"`
HideAdult *bool `json:"hide_adult,omitempty"`
Password string `json:"password,omitempty"`
}
// UpdateProfile applies a non-credential patch to the user.
@@ -35,7 +39,29 @@ func (p *ProfileService) UpdateProfile(ctx context.Context, userID string, patch
if userID == "" {
return nil, errors.New("missing user id")
}
current, err := p.repo.User.FindByID(ctx, userID)
if err != nil {
return nil, err
}
if current == nil {
return nil, errors.New("user not found")
}
updates := map[string]any{}
if patch.Username != nil {
v := strings.TrimSpace(*patch.Username)
if v == "" {
return nil, errors.New("username required")
}
if existing, err := p.repo.User.FindByUsername(ctx, v); err != nil {
return nil, err
} else if existing != nil && existing.ID != userID {
return nil, ErrUsernameTaken
}
updates["username"] = v
}
if patch.Nickname != nil {
updates["nickname"] = strings.TrimSpace(*patch.Nickname)
}
if patch.Email != nil {
v := strings.TrimSpace(*patch.Email)
updates["email"] = v
@@ -43,6 +69,9 @@ func (p *ProfileService) UpdateProfile(ctx context.Context, userID string, patch
if patch.AvatarURL != nil {
updates["avatar_url"] = strings.TrimSpace(*patch.AvatarURL)
}
if patch.HideAdult != nil {
updates["hide_adult"] = *patch.HideAdult
}
if len(updates) > 0 {
if err := p.repo.DB.Model(&model.User{}).Where("id = ?", userID).
Updates(updates).Error; err != nil {
+275 -39
View File
@@ -17,6 +17,8 @@ import (
"time"
"go.uber.org/zap"
"golang.org/x/crypto/bcrypt"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
@@ -24,8 +26,9 @@ import (
// TelegramUpdate 是 Telegram Bot API 推送的 update 对象。
type TelegramUpdate struct {
UpdateID int `json:"update_id"`
Message *TelegramMessage `json:"message,omitempty"`
UpdateID int `json:"update_id"`
Message *TelegramMessage `json:"message,omitempty"`
CallbackQuery *TelegramCallbackQuery `json:"callback_query,omitempty"`
}
// TelegramMessage 是 Telegram 消息对象。
@@ -37,6 +40,13 @@ type TelegramMessage struct {
Date int `json:"date"`
}
type TelegramCallbackQuery struct {
ID string `json:"id"`
From TelegramUser `json:"from"`
Message *TelegramMessage `json:"message,omitempty"`
Data string `json:"data,omitempty"`
}
// TelegramUser 是 Telegram 用户对象。
type TelegramUser struct {
ID int `json:"id"`
@@ -50,6 +60,16 @@ type TelegramChat struct {
Type string `json:"type"`
}
type telegramCommandReply struct {
Text string
Buttons [][]telegramInlineButton
}
type telegramInlineButton struct {
Text string `json:"text"`
Data string `json:"callback_data"`
}
// TelegramBotService 处理 Telegram Bot 的交互命令。
type TelegramBotService struct {
log *zap.Logger
@@ -77,6 +97,10 @@ func (s *TelegramBotService) HandleWebhook(ctx context.Context, body []byte) err
return fmt.Errorf("invalid update: %w", err)
}
if update.CallbackQuery != nil {
return s.handleCallback(ctx, update.CallbackQuery)
}
if update.Message == nil || update.Message.Text == "" {
return nil
}
@@ -97,11 +121,11 @@ func (s *TelegramBotService) HandleWebhook(ctx context.Context, body []byte) err
reply, err := s.executeCommand(ctx, channel, msg, text)
if err != nil {
s.log.Error("command failed", zap.Error(err))
_ = s.reply(ctx, channel, msg.Chat.ID, "命令执行失败: "+err.Error())
_ = s.reply(ctx, channel, msg.Chat.ID, telegramCommandReply{Text: "命令执行失败: " + err.Error()})
return nil
}
if reply != "" {
if reply.Text != "" {
if err := s.reply(ctx, channel, msg.Chat.ID, reply); err != nil {
s.log.Error("reply failed", zap.Error(err))
}
@@ -111,57 +135,111 @@ func (s *TelegramBotService) HandleWebhook(ctx context.Context, body []byte) err
}
// executeCommand 解析命令并执行。
func (s *TelegramBotService) executeCommand(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, text string) (string, error) {
func (s *TelegramBotService) executeCommand(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, text string) (telegramCommandReply, error) {
parts := strings.Fields(text)
if len(parts) == 0 {
return "", nil
return telegramCommandReply{}, nil
}
cmd := strings.ToLower(parts[0])
args := parts[1:]
if msg.Chat.Type != "" && msg.Chat.Type != "private" && !s.telegramChatAllowed(channel, msg.Chat.ID) {
return telegramCommandReply{Text: "此群组/频道未绑定到 Bot 管理入口,请在通知渠道里填写「命令群组/频道 Chat ID」。"}, nil
}
switch cmd {
case "/start":
return s.cmdStart(msg), nil
return s.cmdStart(ctx, msg, args), nil
case "/help":
return s.cmdHelp(), nil
return telegramCommandReply{Text: s.cmdHelp(ctx, msg)}, nil
case "/hideadult", "/hide_adult", "/adult":
return s.cmdHideAdult(ctx, msg, args), nil
case "/status":
if !s.telegramUserIsAdmin(ctx, msg.From.ID) {
return telegramCommandReply{Text: "此命令仅管理员可用。普通用户只能使用 /start 绑定账号,并通过按钮隐藏成人目录。"}, nil
}
return s.cmdStatus(ctx)
case "/search":
if !s.telegramUserIsAdmin(ctx, msg.From.ID) {
return telegramCommandReply{Text: "此命令仅管理员可用。"}, nil
}
return s.cmdSearch(ctx, args)
case "/downloads":
if !s.telegramUserIsAdmin(ctx, msg.From.ID) {
return telegramCommandReply{Text: "此命令仅管理员可用。"}, nil
}
return s.cmdDownloads(ctx)
case "/stats":
if !s.telegramUserIsAdmin(ctx, msg.From.ID) {
return telegramCommandReply{Text: "此命令仅管理员可用。"}, nil
}
return s.cmdStats(ctx)
default:
return fmt.Sprintf("未知命令: %s\n\n输入 /help 查看可用命令列表。", cmd), nil
return telegramCommandReply{Text: fmt.Sprintf("未知命令: %s\n\n输入 /help 查看可用命令列表。", cmd)}, nil
}
}
// cmdStart 处理 /start 命令。
func (s *TelegramBotService) cmdStart(msg *TelegramMessage) string {
func (s *TelegramBotService) cmdStart(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
name := msg.From.FirstName
if msg.From.Username != "" {
name = "@" + msg.From.Username
}
return fmt.Sprintf(
"<b>欢迎使用 MediaStationGo</b>\n\n"+
"你好 %s!你已成功连接到媒体中心。\n\n"+
"<b>可用命令:</b>\n"+
"📊 /status — 系统运行状态\n"+
"🔍 /search 关键词 — 搜索媒体\n"+
"📥 /downloads — 下载进度\n"+
"📈 /stats — 媒体库统计\n"+
"❓ /help — 帮助信息",
name,
)
if len(args) == 0 {
if binding := s.telegramBinding(ctx, msg.From.ID); binding != nil {
user, _ := s.repo.User.FindByID(ctx, binding.UserID)
status := "未隐藏"
if user != nil && user.HideAdult {
status = "已隐藏"
}
return telegramCommandReply{
Text: fmt.Sprintf("<b>MediaStationGo 已绑定</b>\n\n你好 %s,当前账号:<b>%s</b>\n成人目录:<b>%s</b>", name, userNameOrFallback(user), status),
Buttons: [][]telegramInlineButton{{{
Text: map[bool]string{true: "显示成人目录", false: "隐藏成人目录"}[user != nil && user.HideAdult],
Data: "adult_toggle",
}}},
}
}
return telegramCommandReply{Text: "<b>欢迎使用 MediaStationGo</b>\n\n普通用户请先绑定账号:\n<code>/start 用户名 密码</code>\n或:<code>/start 用户名-密码</code>\n\n如果没有账号,请联系管理员注册。"}
}
username, password := parseStartCredentials(args)
if username == "" || password == "" {
return telegramCommandReply{Text: "绑定格式不正确,请使用:\n<code>/start 用户名 密码</code>\n或:<code>/start 用户名-密码</code>"}
}
user, err := s.repo.User.FindByUsername(ctx, username)
if err != nil || user == nil {
return telegramCommandReply{Text: "未找到此用户,请联系管理员注册。"}
}
if !user.IsActive {
return telegramCommandReply{Text: "此账号已被禁用,请联系管理员。"}
}
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)); err != nil {
return telegramCommandReply{Text: "账号或密码错误。"}
}
if err := s.upsertTelegramBinding(ctx, msg, user.ID); err != nil {
return telegramCommandReply{Text: "绑定失败:" + err.Error()}
}
return telegramCommandReply{
Text: fmt.Sprintf("绑定成功:<b>%s</b>\n\n普通用户只能使用此 Bot 管理自己的成人目录隐藏状态;系统状态、搜索、下载和统计命令仅管理员可用。", user.Username),
Buttons: [][]telegramInlineButton{{{
Text: map[bool]string{true: "显示成人目录", false: "隐藏成人目录"}[user.HideAdult],
Data: "adult_toggle",
}}},
}
}
// cmdHelp 处理 /help 命令。
func (s *TelegramBotService) cmdHelp() string {
func (s *TelegramBotService) cmdHelp(ctx context.Context, msg *TelegramMessage) string {
if !s.telegramUserIsAdmin(ctx, msg.From.ID) {
return "<b>MediaStationGo 用户命令</b>\n\n" +
"<b>/start 用户名 密码</b> — 绑定账号\n" +
"<b>/hideadult on|off</b> — 隐藏或显示成人目录\n\n" +
"系统状态、搜索、下载列表与统计命令仅管理员可用。"
}
return "<b>MediaStationGo 命令列表</b>\n\n" +
"<b>/start</b> — 开始使用\n" +
"<b>/help</b> — 帮助信息\n" +
"<b>/hideadult on|off</b> — 隐藏/显示当前绑定账号的成人目录\n" +
"<b>/status</b> — 系统运行状态\n" +
"<b>/search 关键词</b> — 搜索媒体库\n" +
"<b>/downloads</b> — 下载列表\n" +
@@ -174,7 +252,43 @@ func (s *TelegramBotService) cmdHelp() string {
}
// cmdStatus 处理 /status 命令。
func (s *TelegramBotService) cmdStatus(ctx context.Context) (string, error) {
func (s *TelegramBotService) cmdHideAdult(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
binding := s.telegramBinding(ctx, msg.From.ID)
if binding == nil {
return telegramCommandReply{Text: "请先绑定账号:<code>/start 用户名 密码</code>"}
}
user, err := s.repo.User.FindByID(ctx, binding.UserID)
if err != nil || user == nil {
return telegramCommandReply{Text: "绑定用户不存在,请重新 /start 绑定。"}
}
next := true
if len(args) > 0 {
switch strings.ToLower(strings.TrimSpace(args[0])) {
case "off", "false", "0", "show", "显示", "关闭":
next = false
case "on", "true", "1", "hide", "隐藏", "开启":
next = true
default:
next = !user.HideAdult
}
} else {
next = !user.HideAdult
}
if err := s.repo.User.UpdateFields(ctx, user.ID, map[string]any{"hide_adult": next}); err != nil {
return telegramCommandReply{Text: "更新失败:" + err.Error()}
}
status := map[bool]string{true: "已隐藏", false: "已显示"}[next]
return telegramCommandReply{
Text: "成人目录" + status + "。此设置会同步影响网页与第三方客户端。",
Buttons: [][]telegramInlineButton{{{
Text: map[bool]string{true: "显示成人目录", false: "隐藏成人目录"}[next],
Data: "adult_toggle",
}}},
}
}
// cmdStatus 处理 /status 命令。
func (s *TelegramBotService) cmdStatus(ctx context.Context) (telegramCommandReply, error) {
var mediaCount int64
s.repo.DB.Model(&model.Media{}).Count(&mediaCount)
@@ -182,18 +296,18 @@ func (s *TelegramBotService) cmdStatus(ctx context.Context) (string, error) {
s.repo.DB.Raw("SELECT COALESCE(SUM(size_bytes), 0) FROM media").Scan(&totalSize)
totalSizeGB := float64(totalSize) / 1024 / 1024 / 1024
return fmt.Sprintf(
return telegramCommandReply{Text: fmt.Sprintf(
"<b>系统运行状态</b>\n\n"+
"🎬 媒体总数: <b>%d</b>\n"+
"💾 存储占用: <b>%.1f GB</b>",
mediaCount, totalSizeGB,
), nil
)}, nil
}
// cmdSearch 处理 /search 命令。
func (s *TelegramBotService) cmdSearch(ctx context.Context, args []string) (string, error) {
func (s *TelegramBotService) cmdSearch(ctx context.Context, args []string) (telegramCommandReply, error) {
if len(args) == 0 {
return "请提供搜索关键词\n例: <code>/search 哥斯拉</code>", nil
return telegramCommandReply{Text: "请提供搜索关键词\n例: <code>/search 哥斯拉</code>"}, nil
}
keyword := strings.Join(args, " ")
@@ -202,11 +316,11 @@ func (s *TelegramBotService) cmdSearch(ctx context.Context, args []string) (stri
Order("year DESC").Limit(8).
Find(&results).Error
if err != nil {
return "", err
return telegramCommandReply{}, err
}
if len(results) == 0 {
return fmt.Sprintf("未找到与 <b>%s</b> 相关的媒体", keyword), nil
return telegramCommandReply{Text: fmt.Sprintf("未找到与 <b>%s</b> 相关的媒体", keyword)}, nil
}
var sb strings.Builder
@@ -223,11 +337,11 @@ func (s *TelegramBotService) cmdSearch(ctx context.Context, args []string) (stri
sb.WriteString(fmt.Sprintf("%d. <b>%s</b>%s%s — %s\n", i+1, m.Title, year, ep, formatSize(m.SizeBytes)))
}
return sb.String(), nil
return telegramCommandReply{Text: sb.String()}, nil
}
// cmdDownloads 处理 /downloads 命令。
func (s *TelegramBotService) cmdDownloads(ctx context.Context) (string, error) {
func (s *TelegramBotService) cmdDownloads(ctx context.Context) (telegramCommandReply, error) {
type Row struct {
Title string
Status string
@@ -236,11 +350,11 @@ func (s *TelegramBotService) cmdDownloads(ctx context.Context) (string, error) {
if err := s.repo.DB.Raw(
"SELECT COALESCE(NULLIF(title,''),'下载任务') as title, COALESCE(status,'unknown') as status FROM download_tasks ORDER BY created_at DESC LIMIT 8",
).Scan(&rows).Error; err != nil {
return "", err
return telegramCommandReply{}, err
}
if len(rows) == 0 {
return "当前没有下载任务。", nil
return telegramCommandReply{Text: "当前没有下载任务。"}, nil
}
var sb strings.Builder
@@ -265,11 +379,11 @@ func (s *TelegramBotService) cmdDownloads(ctx context.Context) (string, error) {
sb.WriteString(fmt.Sprintf("%s %s\n", icon, name))
}
return sb.String(), nil
return telegramCommandReply{Text: sb.String()}, nil
}
// cmdStats 处理 /stats 命令。
func (s *TelegramBotService) cmdStats(ctx context.Context) (string, error) {
func (s *TelegramBotService) cmdStats(ctx context.Context) (telegramCommandReply, error) {
var totalMedia int64
s.repo.DB.Model(&model.Media{}).Count(&totalMedia)
@@ -307,7 +421,7 @@ func (s *TelegramBotService) cmdStats(ctx context.Context) (string, error) {
}
}
return sb.String(), nil
return telegramCommandReply{Text: sb.String()}, nil
}
// ── Polling ──
@@ -421,7 +535,7 @@ func (s *TelegramBotService) pollLoop(ctx context.Context, botToken string) {
// ── Message Sending ──
// reply 通过 Telegram Bot API 发送回复消息。
func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyChannel, chatID int, text string) error {
func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyChannel, chatID int, reply telegramCommandReply) error {
botToken := ""
if channel != nil {
configStr := channel.Config
@@ -439,9 +553,23 @@ func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyCha
payload := map[string]interface{}{
"chat_id": strconv.Itoa(chatID),
"text": text,
"text": reply.Text,
"parse_mode": "HTML",
}
if len(reply.Buttons) > 0 {
keyboard := make([][]map[string]string, 0, len(reply.Buttons))
for _, row := range reply.Buttons {
buttons := make([]map[string]string, 0, len(row))
for _, button := range row {
buttons = append(buttons, map[string]string{
"text": button.Text,
"callback_data": button.Data,
})
}
keyboard = append(keyboard, buttons)
}
payload["reply_markup"] = map[string]interface{}{"inline_keyboard": keyboard}
}
body, _ := json.Marshal(payload)
apiURL := fmt.Sprintf("https://api.telegram.org/bot%s/sendMessage", botToken)
@@ -482,13 +610,121 @@ func (s *TelegramBotService) findChannelByChatID(ctx context.Context, chatID int
}
var cfg map[string]string
json.Unmarshal([]byte(configStr), &cfg)
if cfg["chat_id"] == target {
if cfg["chat_id"] == target || cfg["command_chat_id"] == target {
return &ch
}
}
if len(channels) == 1 && channels[0].Enabled {
return &channels[0]
}
return nil
}
func (s *TelegramBotService) handleCallback(ctx context.Context, cb *TelegramCallbackQuery) error {
if cb == nil || cb.Message == nil {
return nil
}
channel := s.findChannelByChatID(ctx, cb.Message.Chat.ID)
switch strings.TrimSpace(cb.Data) {
case "adult_toggle":
msg := *cb.Message
msg.From = cb.From
reply := s.cmdHideAdult(ctx, &msg, nil)
if reply.Text != "" {
return s.reply(ctx, channel, cb.Message.Chat.ID, reply)
}
}
return nil
}
func (s *TelegramBotService) telegramBinding(ctx context.Context, telegramUserID int) *model.TelegramBinding {
if telegramUserID == 0 {
return nil
}
var binding model.TelegramBinding
err := s.repo.DB.WithContext(ctx).Where("telegram_user_id = ?", int64(telegramUserID)).First(&binding).Error
if err != nil {
return nil
}
return &binding
}
func (s *TelegramBotService) telegramUserIsAdmin(ctx context.Context, telegramUserID int) bool {
binding := s.telegramBinding(ctx, telegramUserID)
if binding == nil {
return false
}
user, err := s.repo.User.FindByID(ctx, binding.UserID)
return err == nil && user != nil && user.Role == "admin" && user.IsActive
}
func (s *TelegramBotService) telegramChatAllowed(channel *model.NotifyChannel, chatID int) bool {
if channel == nil {
return false
}
configStr := channel.Config
if s.crypto != nil && configStr != "" {
configStr = s.crypto.Decrypt(configStr)
}
var cfg map[string]string
if err := json.Unmarshal([]byte(configStr), &cfg); err != nil {
return false
}
target := strconv.Itoa(chatID)
commandChatID := strings.TrimSpace(cfg["command_chat_id"])
if commandChatID != "" {
return commandChatID == target
}
return strings.TrimSpace(cfg["chat_id"]) == target
}
func (s *TelegramBotService) upsertTelegramBinding(ctx context.Context, msg *TelegramMessage, userID string) error {
name := strings.TrimSpace(msg.From.FirstName)
if msg.From.Username != "" {
name = "@" + strings.TrimSpace(msg.From.Username)
}
var existing model.TelegramBinding
err := s.repo.DB.WithContext(ctx).Where("telegram_user_id = ?", int64(msg.From.ID)).First(&existing).Error
if err == nil {
return s.repo.DB.WithContext(ctx).Model(&existing).Updates(map[string]any{
"telegram_name": name,
"chat_id": int64(msg.Chat.ID),
"user_id": userID,
}).Error
}
if err != nil && err != gorm.ErrRecordNotFound {
return err
}
return s.repo.DB.WithContext(ctx).Create(&model.TelegramBinding{
TelegramUserID: int64(msg.From.ID),
TelegramName: name,
ChatID: int64(msg.Chat.ID),
UserID: userID,
}).Error
}
func parseStartCredentials(args []string) (string, string) {
if len(args) >= 2 {
return strings.TrimSpace(args[0]), strings.TrimSpace(strings.Join(args[1:], " "))
}
if len(args) == 1 {
raw := strings.TrimSpace(args[0])
for _, sep := range []string{"-", ":", ":"} {
if parts := strings.SplitN(raw, sep, 2); len(parts) == 2 {
return strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1])
}
}
}
return "", ""
}
func userNameOrFallback(user *model.User) string {
if user == nil || strings.TrimSpace(user.Username) == "" {
return "未知用户"
}
return user.Username
}
// ── Webhook Management ──
// SetWebhook 注册 Telegram Bot Webhook URL。
+129
View File
@@ -0,0 +1,129 @@
package service
import (
"context"
"encoding/json"
"strings"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// AdultContentEnabled reads the global Adult / NSFW switch.
func AdultContentEnabled(ctx context.Context, repo *repository.Container) bool {
if repo == nil || repo.Setting == nil {
return false
}
value, err := repo.Setting.Get(ctx, "adult.enabled")
if err != nil {
return false
}
switch strings.ToLower(strings.TrimSpace(value)) {
case "1", "true", "yes", "on", "enabled", "启用", "开启":
return true
default:
return false
}
}
// UserHidesAdult reports whether a user's own lock overrides all profiles.
func UserHidesAdult(ctx context.Context, repo *repository.Container, userID string) bool {
if strings.TrimSpace(userID) == "" || repo == nil || repo.User == nil {
return false
}
user, err := repo.User.FindByID(ctx, userID)
return err == nil && user != nil && user.HideAdult
}
// UserDefaultMediaVisibility is the visibility policy used by clients that
// cannot pass a web play-profile token, notably Emby/Jellyfin-compatible apps.
func UserDefaultMediaVisibility(ctx context.Context, repo *repository.Container, userID string) MediaVisibility {
visibility := MediaVisibility{IncludeNSFW: AdultContentEnabled(ctx, repo)}
if repo == nil {
return visibility
}
if UserHidesAdult(ctx, repo, userID) {
visibility.IncludeNSFW = false
}
if userID == "" || repo.PlayProfile == nil {
return visibility
}
rows, err := repo.PlayProfile.ListByUser(ctx, userID)
if err != nil {
return visibility
}
for _, row := range rows {
if !row.IsDefault {
continue
}
visibility.IncludeNSFW = visibility.IncludeNSFW && row.AllowAdult
visibility.AllowedLibraryIDs = DecodeAllowedLibraryIDs(row.AllowedLibraryIDs)
break
}
return visibility
}
// DecodeAllowedLibraryIDs normalises a PlayProfile allowed-library JSON string.
func DecodeAllowedLibraryIDs(raw string) []string {
if strings.TrimSpace(raw) == "" {
return nil
}
var ids []string
if err := json.Unmarshal([]byte(raw), &ids); err != nil {
return nil
}
out := ids[:0]
for _, id := range ids {
if strings.TrimSpace(id) != "" {
out = append(out, strings.TrimSpace(id))
}
}
return out
}
// LibraryVisibleForUser applies profile library limits and adult-directory
// hiding to a library card/folder.
func LibraryVisibleForUser(ctx context.Context, repo *repository.Container, lib model.Library, visibility MediaVisibility) bool {
if len(visibility.AllowedLibraryIDs) > 0 {
found := false
for _, id := range visibility.AllowedLibraryIDs {
if id == lib.ID {
found = true
break
}
}
if !found {
return false
}
}
if visibility.IncludeNSFW {
return true
}
if LibraryLooksAdult(lib) {
return false
}
if repo != nil && repo.DB != nil {
var count int64
_ = repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("library_id = ? AND nsfw = ?", lib.ID, true).
Count(&count).Error
if count > 0 {
return false
}
}
return true
}
// LibraryLooksAdult catches adult-only roots even before all rows are scraped.
func LibraryLooksAdult(lib model.Library) bool {
text := strings.ToLower(strings.TrimSpace(lib.Name + " " + lib.Path + " " + lib.Type))
if text == "" {
return false
}
for _, token := range []string{"成人", "限制级", "nsfw", "adult", "jav", "javdb", "javbus", "9kg", "里番", "番号"} {
if strings.Contains(text, token) {
return true
}
}
return false
}