mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-29 11:36:36 +08:00
添加新功能,完善项目
This commit is contained in:
@@ -20,6 +20,12 @@ type settingReq struct {
|
||||
Value string `json:"value"`
|
||||
}
|
||||
|
||||
// maskedSettingKeys 里的设置值绝不能被完整下发:它们是可用于对外操作的凭据。
|
||||
// 下发脱敏值,保存时再靠 isMaskedSettingValue 还原为「保持原值」。
|
||||
var maskedSettingKeys = map[string]bool{
|
||||
service.SettingTelegramBotToken: true,
|
||||
}
|
||||
|
||||
func listSettingsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
settings, err := svc.Repo.Setting.All(c.Request.Context())
|
||||
@@ -27,10 +33,21 @@ func listSettingsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
for i := range settings {
|
||||
if maskedSettingKeys[settings[i].Key] {
|
||||
settings[i].Value = service.MaskSecret(settings[i].Value)
|
||||
}
|
||||
}
|
||||
c.JSON(http.StatusOK, settings)
|
||||
}
|
||||
}
|
||||
|
||||
// isMaskedSettingValue 识别「前端把脱敏值原样提交回来」的情况。此时必须保留
|
||||
// 已存的真实值,否则一次保存就会把凭据覆盖成 ***。
|
||||
func isMaskedSettingValue(value string) bool {
|
||||
return strings.Contains(value, "***")
|
||||
}
|
||||
|
||||
func updateSettingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req settingReq
|
||||
@@ -38,6 +55,11 @@ func updateSettingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
// 脱敏值回传 == 用户没改这个凭据,保留库里已存的真实值。
|
||||
if maskedSettingKeys[req.Key] && isMaskedSettingValue(req.Value) {
|
||||
c.Status(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
oldValue := ""
|
||||
if req.Key == service.AdultLibraryIDsSettingKey {
|
||||
oldValue, _ = svc.Repo.Setting.Get(c.Request.Context(), req.Key)
|
||||
|
||||
@@ -0,0 +1,173 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/middleware"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
// 设备管理接口。
|
||||
//
|
||||
// 普通用户只能操作自己的设备(路由挂在 /me 下,用户 ID 始终取自会话);
|
||||
// 管理员通过 /admin/users/:id/devices 代管任意用户。两组接口共用同一份
|
||||
// DeviceService,因此「谁上线过、谁被踢掉」只有一处事实来源。
|
||||
|
||||
// deviceListPayload 是设备列表的下发形状。Fingerprint 不外发:它是防共享
|
||||
// 判定用的内部标识,暴露出去只会方便伪造。
|
||||
type deviceListPayload struct {
|
||||
Devices []devicePayload `json:"devices"`
|
||||
}
|
||||
|
||||
type devicePayload struct {
|
||||
ID string `json:"id"`
|
||||
DeviceID string `json:"device_id"`
|
||||
DeviceName string `json:"device_name,omitempty"`
|
||||
Client string `json:"client,omitempty"`
|
||||
LastIP string `json:"last_ip,omitempty"`
|
||||
LastSeenAt string `json:"last_seen_at,omitempty"`
|
||||
LastPlayAt string `json:"last_play_at,omitempty"`
|
||||
Kicked bool `json:"kicked"`
|
||||
Online bool `json:"online"`
|
||||
Playing bool `json:"playing"`
|
||||
Warnings int `json:"warnings"`
|
||||
}
|
||||
|
||||
func toDevicePayload(d model.UserDevice) devicePayload {
|
||||
out := devicePayload{
|
||||
ID: d.ID,
|
||||
DeviceID: d.DeviceID,
|
||||
DeviceName: d.DeviceName,
|
||||
Client: d.Client,
|
||||
LastIP: d.LastIP,
|
||||
Kicked: d.Kicked,
|
||||
Online: d.Online,
|
||||
Playing: d.Playing,
|
||||
Warnings: d.Warnings,
|
||||
}
|
||||
if !d.LastSeenAt.IsZero() {
|
||||
out.LastSeenAt = d.LastSeenAt.Format("2006-01-02T15:04:05Z07:00")
|
||||
}
|
||||
if d.LastPlayAt != nil && !d.LastPlayAt.IsZero() {
|
||||
out.LastPlayAt = d.LastPlayAt.Format("2006-01-02T15:04:05Z07:00")
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func deviceListResponse(devices []model.UserDevice) deviceListPayload {
|
||||
items := make([]devicePayload, 0, len(devices))
|
||||
for _, d := range devices {
|
||||
items = append(items, toDevicePayload(d))
|
||||
}
|
||||
return deviceListPayload{Devices: items}
|
||||
}
|
||||
|
||||
// myDevicesHandler 返回当前会话用户的设备列表。
|
||||
func myDevicesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID := sessionUserID(c)
|
||||
devices, err := svc.Device.ListDevices(c.Request.Context(), userID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, deviceListResponse(devices))
|
||||
}
|
||||
}
|
||||
|
||||
// myKickDeviceHandler 踢掉当前用户的一台设备。
|
||||
func myKickDeviceHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID := sessionUserID(c)
|
||||
deviceID := strings.TrimSpace(c.Param("deviceID"))
|
||||
if deviceID == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "device id required"})
|
||||
return
|
||||
}
|
||||
if err := svc.Device.KickDevice(c.Request.Context(), userID, deviceID); err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
// myKickAllDevicesHandler 踢掉当前用户的全部设备。
|
||||
func myKickAllDevicesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID := sessionUserID(c)
|
||||
if err := svc.Device.KickAllDevices(c.Request.Context(), userID); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
// adminUserDevicesHandler 返回指定用户的设备列表。
|
||||
func adminUserDevicesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID := strings.TrimSpace(c.Param("id"))
|
||||
// FindByID 对「不存在」返回 (nil, nil),必须判空而不是判 error,
|
||||
// 否则「用户不存在」会伪装成「该用户没有设备」的空列表。
|
||||
user, err := svc.Repo.User.FindByID(c.Request.Context(), userID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if user == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "user not found"})
|
||||
return
|
||||
}
|
||||
devices, err := svc.Device.ListDevices(c.Request.Context(), userID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, deviceListResponse(devices))
|
||||
}
|
||||
}
|
||||
|
||||
// adminKickUserDeviceHandler 由管理员踢掉指定用户的一台设备。
|
||||
func adminKickUserDeviceHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID := strings.TrimSpace(c.Param("id"))
|
||||
deviceID := strings.TrimSpace(c.Param("deviceID"))
|
||||
if deviceID == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "device id required"})
|
||||
return
|
||||
}
|
||||
if err := svc.Device.KickDevice(c.Request.Context(), userID, deviceID); err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
// adminKickAllUserDevicesHandler 由管理员踢掉指定用户的全部设备。
|
||||
func adminKickAllUserDevicesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID := strings.TrimSpace(c.Param("id"))
|
||||
if err := svc.Device.KickAllDevices(c.Request.Context(), userID); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
// sessionUserID 读取会话用户 ID。调用方路由都挂在鉴权中间件之后,因此这里
|
||||
// 只做类型断言兜底,不做权限判断。
|
||||
func sessionUserID(c *gin.Context) string {
|
||||
if v, ok := c.Get(middleware.CtxUserID); ok {
|
||||
if s, ok := v.(string); ok {
|
||||
return s
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/middleware"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
// newDeviceTestEnv 搭一个只挂设备/Telegram 路由的最小环境,并预置两个用户,
|
||||
// 用于验证「只能操作自己的设备」这条边界。
|
||||
func newDeviceTestEnv(t *testing.T) (*gin.Engine, *service.Container) {
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.User{}, &model.UserDevice{}, &model.Setting{}); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
svc := &service.Container{Repo: repos, Log: zap.NewNop()}
|
||||
svc.Device = service.NewDeviceService(zap.NewNop(), repos)
|
||||
svc.Device.SetSessionTracker(service.NewSessionTrackerService(zap.NewNop()))
|
||||
svc.Telegram = service.NewTelegramService(zap.NewNop(), repos)
|
||||
|
||||
const secret = "test-secret"
|
||||
router := gin.New()
|
||||
authed := router.Group("/api", func(c *gin.Context) {
|
||||
// 测试里直接注入会话身份,绕开真实 JWT 解析。
|
||||
if uid := c.GetHeader("X-Test-User"); uid != "" {
|
||||
c.Set(middleware.CtxUserID, uid)
|
||||
c.Set(middleware.CtxUserRole, c.GetHeader("X-Test-Role"))
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
authed.GET("/me/devices", myDevicesHandler(svc))
|
||||
authed.POST("/me/devices/kick-all", myKickAllDevicesHandler(svc))
|
||||
authed.POST("/me/devices/:deviceID/kick", myKickDeviceHandler(svc))
|
||||
authed.GET("/me/telegram", getTelegramStatusHandler(svc))
|
||||
authed.POST("/me/telegram/bind-code", startTelegramBindHandler(svc))
|
||||
authed.DELETE("/me/telegram", unbindTelegramHandler(svc))
|
||||
authed.GET("/admin/users/:id/devices", adminUserDevicesHandler(svc))
|
||||
authed.POST("/admin/users/:id/devices/:deviceID/kick", adminKickUserDeviceHandler(svc))
|
||||
_ = secret
|
||||
return router, svc
|
||||
}
|
||||
|
||||
func seedDeviceUsers(t *testing.T, svc *service.Container) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
for _, id := range []string{"user-a", "user-b"} {
|
||||
if err := svc.Repo.User.Create(ctx, &model.User{
|
||||
Base: model.Base{ID: id}, Username: id, PasswordHash: "x", Role: "user", IsActive: true,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func doJSON(t *testing.T, router *gin.Engine, method, path, userID, role string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(method, path, nil)
|
||||
req.Header.Set("X-Test-User", userID)
|
||||
if role != "" {
|
||||
req.Header.Set("X-Test-Role", role)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
return w
|
||||
}
|
||||
|
||||
// /me/devices 只能返回调用者自己的设备。
|
||||
func TestMyDevicesScopedToCaller(t *testing.T) {
|
||||
router, svc := newDeviceTestEnv(t)
|
||||
seedDeviceUsers(t, svc)
|
||||
ctx := context.Background()
|
||||
|
||||
svc.Device.RecordLogin(ctx, "user-a", "dev-a", "A-Phone", "Infuse", "1.1.1.1")
|
||||
svc.Device.RecordLogin(ctx, "user-b", "dev-b", "B-Phone", "Infuse", "2.2.2.2")
|
||||
|
||||
w := doJSON(t, router, http.MethodGet, "/api/me/devices", "user-a", "user")
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
var payload struct {
|
||||
Devices []struct {
|
||||
DeviceID string `json:"device_id"`
|
||||
} `json:"devices"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(payload.Devices) != 1 || payload.Devices[0].DeviceID != "dev-a" {
|
||||
t.Fatalf("devices = %+v, want only dev-a", payload.Devices)
|
||||
}
|
||||
// 设备指纹属于内部判定标识,不能下发。
|
||||
if strings.Contains(w.Body.String(), "fingerprint") {
|
||||
t.Fatalf("response must not expose fingerprint: %s", w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// 踢别人的设备必须失败:/me 路由用会话身份,deviceID 属于他人时查不到。
|
||||
func TestKickForeignDeviceFails(t *testing.T) {
|
||||
router, svc := newDeviceTestEnv(t)
|
||||
seedDeviceUsers(t, svc)
|
||||
svc.Device.RecordLogin(context.Background(), "user-b", "dev-b", "B-Phone", "Infuse", "2.2.2.2")
|
||||
|
||||
w := doJSON(t, router, http.MethodPost, "/api/me/devices/dev-b/kick", "user-a", "user")
|
||||
if w.Code == http.StatusNoContent {
|
||||
t.Fatal("user-a must not be able to kick user-b's device")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMyKickOwnDeviceSucceeds(t *testing.T) {
|
||||
router, svc := newDeviceTestEnv(t)
|
||||
seedDeviceUsers(t, svc)
|
||||
svc.Device.RecordLogin(context.Background(), "user-a", "dev-a", "A-Phone", "Infuse", "1.1.1.1")
|
||||
|
||||
w := doJSON(t, router, http.MethodPost, "/api/me/devices/dev-a/kick", "user-a", "user")
|
||||
if w.Code != http.StatusNoContent {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// 管理员接口对不存在的用户返回 404,避免把「用户不存在」和「用户没有设备」
|
||||
// 混成同一个空列表。
|
||||
func TestAdminDevicesUnknownUserReturns404(t *testing.T) {
|
||||
router, svc := newDeviceTestEnv(t)
|
||||
seedDeviceUsers(t, svc)
|
||||
|
||||
w := doJSON(t, router, http.MethodGet, "/api/admin/users/nope/devices", "admin-1", "admin")
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Fatalf("status = %d, want 404", w.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTelegramStatusAndBindCode(t *testing.T) {
|
||||
router, svc := newDeviceTestEnv(t)
|
||||
seedDeviceUsers(t, svc)
|
||||
|
||||
w := doJSON(t, router, http.MethodGet, "/api/me/telegram", "user-a", "user")
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", w.Code)
|
||||
}
|
||||
if !strings.Contains(w.Body.String(), `"bound":false`) {
|
||||
t.Fatalf("body = %s, want bound=false", w.Body.String())
|
||||
}
|
||||
|
||||
w = doJSON(t, router, http.MethodPost, "/api/me/telegram/bind-code", "user-a", "user")
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
var code struct {
|
||||
Code string `json:"code"`
|
||||
ExpiresIn int `json:"expires_in_seconds"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &code); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(code.Code) != 6 {
|
||||
t.Fatalf("code = %q, want 6 chars", code.Code)
|
||||
}
|
||||
if code.ExpiresIn <= 0 {
|
||||
t.Fatalf("expires_in_seconds = %d, want > 0", code.ExpiresIn)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAdminSettingsMasksBotToken(t *testing.T) {
|
||||
router, svc := newDeviceTestEnv(t)
|
||||
ctx := context.Background()
|
||||
if err := svc.Repo.Setting.Set(ctx, service.SettingTelegramBotToken, "123456:AAHsecretTOKEN"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
router.GET("/api/admin/settings", listSettingsHandler(svc))
|
||||
router.PUT("/api/admin/settings", updateSettingHandler(svc))
|
||||
|
||||
w := doJSON(t, router, http.MethodGet, "/api/admin/settings", "admin-1", "admin")
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", w.Code)
|
||||
}
|
||||
if strings.Contains(w.Body.String(), "AAHsecretTOKEN") {
|
||||
t.Fatalf("bot token leaked: %s", w.Body.String())
|
||||
}
|
||||
if !strings.Contains(w.Body.String(), "***") {
|
||||
t.Fatalf("bot token should be masked: %s", w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// 把脱敏值原样提交回来时,必须保留库里真实 Token —— 否则一次保存就把凭据毁掉。
|
||||
func TestSavingMaskedTokenKeepsRealValue(t *testing.T) {
|
||||
router, svc := newDeviceTestEnv(t)
|
||||
ctx := context.Background()
|
||||
const real = "123456:AAHsecretTOKEN"
|
||||
if err := svc.Repo.Setting.Set(ctx, service.SettingTelegramBotToken, real); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
router.PUT("/api/admin/settings", updateSettingHandler(svc))
|
||||
|
||||
body := `{"key":"telegram.bot_token","value":"12***EN"}`
|
||||
req := httptest.NewRequest(http.MethodPut, "/api/admin/settings", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-Test-User", "admin-1")
|
||||
req.Header.Set("X-Test-Role", "admin")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusNoContent {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
got, err := svc.Repo.Setting.Get(ctx, service.SettingTelegramBotToken)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != real {
|
||||
t.Fatalf("stored token = %q, want the original value preserved", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
// Emby 发现类接口的 handler:NextUp / Similar / Genres。
|
||||
//
|
||||
// 这三个接口此前返回空列表,导致第三方客户端首页「接下来播放」、详情页
|
||||
// 「相似推荐」、按类型浏览全部为空白。它们必须始终返回 200 + 合法信封,
|
||||
// 因为客户端在首页刷新时会并发请求,任何 4xx/5xx 都会被判定为服务端异常。
|
||||
|
||||
func embyNextUpHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
userID := embyScopedUserID(c)
|
||||
if userID == "" {
|
||||
c.JSON(http.StatusOK, embyEmptyItemsPayload())
|
||||
return
|
||||
}
|
||||
limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), ""))
|
||||
out, err := svc.Emby.NextUp(c.Request.Context(), userID, limit)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, embyEmptyItemsPayload())
|
||||
return
|
||||
}
|
||||
embyAttachRequestTokenToMediaSources(c, out)
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
|
||||
func embySimilarHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
mediaID := strings.TrimSpace(c.Param("id"))
|
||||
if mediaID == "" {
|
||||
c.JSON(http.StatusOK, embyEmptyItemsPayload())
|
||||
return
|
||||
}
|
||||
limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), ""))
|
||||
out, err := svc.Emby.SimilarItems(c.Request.Context(), mediaID, embyEffectiveUserID(c), limit)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, embyEmptyItemsPayload())
|
||||
return
|
||||
}
|
||||
embyAttachRequestTokenToMediaSources(c, out)
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
|
||||
func embyGenresHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
parentID := firstQueryValue(c, "ParentId", "parentId", "parentid")
|
||||
out, err := svc.Emby.Genres(c.Request.Context(), embyEffectiveUserID(c), parentID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, embyEmptyItemsPayload())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
|
||||
// embyScopedUserID 解析「按用户请求」的 Emby 接口的生效用户。
|
||||
//
|
||||
// 路由上带 :userId 时(/Users/{uid}/Shows/NextUp),只允许查询自己:客户端
|
||||
// 偶尔会带着别人的 id 请求,直接采信等于开放他人观看历史的读取。管理员同样
|
||||
// 按自己处理,避免出现一条无人使用的越权路径。
|
||||
func embyScopedUserID(c *gin.Context) string {
|
||||
caller := embyEffectiveUserID(c)
|
||||
requested := strings.TrimSpace(c.Param("userId"))
|
||||
if requested == "" {
|
||||
return caller
|
||||
}
|
||||
if requested == caller {
|
||||
return caller
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// embyEmptyItemsPayload 与 embyEmptyItemsHandler 保持同一形状。
|
||||
func embyEmptyItemsPayload() gin.H {
|
||||
return gin.H{"Items": []any{}, "TotalRecordCount": 0}
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/config"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
// newEmbyDiscoveryEnv 搭一个跑在内存库上的 Emby 路由环境。
|
||||
func newEmbyDiscoveryEnv(t *testing.T) (*gin.Engine, *service.Container, string) {
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(
|
||||
&model.User{}, &model.Library{}, &model.Media{}, &model.PlaybackHistory{},
|
||||
&model.Setting{}, &model.Favorite{}, &model.UserDevice{},
|
||||
); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
cfg := &config.Config{}
|
||||
cfg.Secrets.JWTSecret = "test-secret"
|
||||
|
||||
svc := &service.Container{Repo: repos, Log: zap.NewNop()}
|
||||
svc.Emby = service.NewEmbyService(cfg, zap.NewNop(), repos).
|
||||
SetDiscovery(service.NewMediaDiscoveryService(zap.NewNop(), repos))
|
||||
|
||||
const userID = "user-1"
|
||||
if err := repos.User.Create(context.Background(), &model.User{
|
||||
Base: model.Base{ID: userID}, Username: "tester", PasswordHash: "x",
|
||||
Role: "user", IsActive: true,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
registerEmbyRoutes(router, cfg.Secrets.JWTSecret, svc)
|
||||
return router, svc, userID
|
||||
}
|
||||
|
||||
func seedEmbyLibrary(t *testing.T, svc *service.Container, typ string) string {
|
||||
t.Helper()
|
||||
lib := &model.Library{Name: "库-" + typ, Path: "/media/" + typ, Type: typ, Enabled: true}
|
||||
if err := svc.Repo.Library.Create(context.Background(), lib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return lib.ID
|
||||
}
|
||||
|
||||
func embyGet(t *testing.T, router *gin.Engine, path, token string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodGet, path, nil)
|
||||
if token != "" {
|
||||
req.Header.Set("X-Emby-Token", token)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
return w
|
||||
}
|
||||
|
||||
func decodeItemsEnvelope(t *testing.T, body []byte) []map[string]any {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
TotalRecordCount int64 `json:"TotalRecordCount"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
t.Fatalf("decode %s: %v", string(body), err)
|
||||
}
|
||||
return payload.Items
|
||||
}
|
||||
|
||||
// NextUp 必须真的返回下一集,而不是空数组。
|
||||
func TestEmbyNextUpReturnsNextEpisode(t *testing.T) {
|
||||
router, svc, userID := newEmbyDiscoveryEnv(t)
|
||||
libID := seedEmbyLibrary(t, svc, "tv")
|
||||
watchedAt := time.Now().Add(-time.Hour)
|
||||
|
||||
for episode, watched := range map[int]bool{1: true, 2: false, 3: false} {
|
||||
m := &model.Media{
|
||||
LibraryID: libID, SeriesID: "series-1", Title: "剧一",
|
||||
SeasonNum: 1, EpisodeNum: episode,
|
||||
Path: "/media/tv/S1E" + string(rune('0'+episode)) + ".mkv",
|
||||
}
|
||||
if err := svc.Repo.DB.Create(m).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if watched {
|
||||
h := &model.PlaybackHistory{
|
||||
UserID: userID, MediaID: m.ID, PositionMs: 1000, DurationMs: 2000,
|
||||
WatchedAt: watchedAt, Completed: false,
|
||||
}
|
||||
if err := svc.Repo.DB.Create(h).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
w := embyGet(t, router, "/emby/Shows/NextUp", signedTestToken(t, "test-secret"))
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
items := decodeItemsEnvelope(t, w.Body.Bytes())
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("items = %d, want 1 (body=%s)", len(items), w.Body.String())
|
||||
}
|
||||
if index, ok := items[0]["IndexNumber"].(float64); !ok || int(index) != 2 {
|
||||
t.Fatalf("IndexNumber = %v, want 2 (body=%s)", items[0]["IndexNumber"], w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// 没有历史时必须返回合法空信封,不能 404/500。
|
||||
func TestEmbyNextUpEmptyWithoutHistory(t *testing.T) {
|
||||
router, _, _ := newEmbyDiscoveryEnv(t)
|
||||
|
||||
w := embyGet(t, router, "/emby/Shows/NextUp", signedTestToken(t, "test-secret"))
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
if items := decodeItemsEnvelope(t, w.Body.Bytes()); len(items) != 0 {
|
||||
t.Fatalf("items = %d, want 0", len(items))
|
||||
}
|
||||
}
|
||||
|
||||
// 小写别名路由同样要走到真实实现(客户端路径大小写并不统一)。
|
||||
func TestEmbyNextUpLowercaseAlias(t *testing.T) {
|
||||
router, svc, _ := newEmbyDiscoveryEnv(t)
|
||||
_ = seedEmbyLibrary(t, svc, "tv")
|
||||
|
||||
w := embyGet(t, router, "/emby/shows/nextup", signedTestToken(t, "test-secret"))
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// Similar 对不存在的条目返回空列表(客户端详情页会无条件请求)。
|
||||
func TestEmbySimilarUnknownItemReturnsEmpty(t *testing.T) {
|
||||
router, _, _ := newEmbyDiscoveryEnv(t)
|
||||
|
||||
w := embyGet(t, router, "/emby/Items/does-not-exist/Similar", signedTestToken(t, "test-secret"))
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
if items := decodeItemsEnvelope(t, w.Body.Bytes()); len(items) != 0 {
|
||||
t.Fatalf("items = %d, want 0", len(items))
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbySimilarReturnsCandidates(t *testing.T) {
|
||||
router, svc, _ := newEmbyDiscoveryEnv(t)
|
||||
libID := seedEmbyLibrary(t, svc, "movie")
|
||||
|
||||
source := &model.Media{
|
||||
LibraryID: libID, Title: "源片", Genres: "Action", Year: 2010, Rating: 8,
|
||||
Path: "/media/movie/source.mkv",
|
||||
}
|
||||
if err := svc.Repo.DB.Create(source).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
other := &model.Media{
|
||||
LibraryID: libID, Title: "同类片", Genres: "Action", Year: 2011, Rating: 8,
|
||||
Path: "/media/movie/other.mkv",
|
||||
}
|
||||
if err := svc.Repo.DB.Create(other).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
w := embyGet(t, router, "/emby/Items/"+source.ID+"/Similar", signedTestToken(t, "test-secret"))
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
items := decodeItemsEnvelope(t, w.Body.Bytes())
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("items = %d, want 1 (body=%s)", len(items), w.Body.String())
|
||||
}
|
||||
if name, _ := items[0]["Name"].(string); name != "同类片" {
|
||||
t.Fatalf("Name = %q, want 同类片", name)
|
||||
}
|
||||
}
|
||||
|
||||
// Genres 必须返回真实类型与计数。
|
||||
func TestEmbyGenresReturnsCounts(t *testing.T) {
|
||||
router, svc, _ := newEmbyDiscoveryEnv(t)
|
||||
libID := seedEmbyLibrary(t, svc, "movie")
|
||||
|
||||
for i, genres := range []string{"Action,Drama", "Action"} {
|
||||
m := &model.Media{
|
||||
LibraryID: libID, Title: "片" + string(rune('A'+i)), Genres: genres,
|
||||
Path: "/media/movie/m" + string(rune('0'+i)) + ".mkv",
|
||||
}
|
||||
if err := svc.Repo.DB.Create(m).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
w := embyGet(t, router, "/emby/Genres", signedTestToken(t, "test-secret"))
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
items := decodeItemsEnvelope(t, w.Body.Bytes())
|
||||
if len(items) != 2 {
|
||||
t.Fatalf("items = %d, want 2 (body=%s)", len(items), w.Body.String())
|
||||
}
|
||||
// 排序按计数降序:Action(2) 在前。
|
||||
if name, _ := items[0]["Name"].(string); name != "Action" {
|
||||
t.Fatalf("first Name = %q, want Action", name)
|
||||
}
|
||||
if count, ok := items[0]["ItemCount"].(float64); !ok || int(count) != 2 {
|
||||
t.Fatalf("ItemCount = %v, want 2", items[0]["ItemCount"])
|
||||
}
|
||||
if id, _ := items[0]["Id"].(string); len(id) == 0 {
|
||||
t.Fatal("genre item must carry a stable Id")
|
||||
}
|
||||
}
|
||||
|
||||
// 按别人的 userId 请求 NextUp 不允许泄露他人历史。
|
||||
func TestEmbyNextUpRejectsForeignUserID(t *testing.T) {
|
||||
router, _, _ := newEmbyDiscoveryEnv(t)
|
||||
|
||||
w := embyGet(t, router, "/emby/Users/someone-else/Shows/NextUp", signedTestToken(t, "test-secret"))
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
if items := decodeItemsEnvelope(t, w.Body.Bytes()); len(items) != 0 {
|
||||
t.Fatalf("items = %d, want 0", len(items))
|
||||
}
|
||||
}
|
||||
@@ -196,15 +196,15 @@ func registerEmbyAuthenticatedItemRoutes(auth *gin.RouterGroup, svc *service.Con
|
||||
auth.GET("/Shows/:id/Episodes", embyShowEpisodesHandler(svc))
|
||||
auth.GET("/Users/:userId/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
|
||||
auth.GET("/Users/:userId/Shows/:id/Episodes", embyShowEpisodesHandler(svc))
|
||||
auth.GET("/Shows/NextUp", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Users/:userId/Shows/NextUp", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Shows/NextUp", embyNextUpHandler(svc))
|
||||
auth.GET("/Users/:userId/Shows/NextUp", embyNextUpHandler(svc))
|
||||
auth.GET("/MediaSegments/:id", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Artists", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Persons", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Genres", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Genres", embyGenresHandler(svc))
|
||||
auth.GET("/Shows/Upcoming", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Users/:userId/Shows/Upcoming", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Items/:id/Similar", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Items/:id/Similar", embySimilarHandler(svc))
|
||||
auth.GET("/Items/:id/ThumbnailSet", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/Items/:id/ThemeMedia", embyThemeMediaHandler(svc))
|
||||
auth.GET("/Users/:userId/Items/:id/SpecialFeatures", embyEmptyItemsHandler(svc))
|
||||
|
||||
@@ -40,15 +40,15 @@ func registerLowercaseEmbyItemRoutes(auth *gin.RouterGroup, svc *service.Contain
|
||||
auth.GET("/shows/:id/episodes", embyShowEpisodesHandler(svc))
|
||||
auth.GET("/users/:userId/shows/:id/seasons", embyShowSeasonsHandler(svc))
|
||||
auth.GET("/users/:userId/shows/:id/episodes", embyShowEpisodesHandler(svc))
|
||||
auth.GET("/shows/nextup", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/users/:userId/shows/nextup", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/shows/nextup", embyNextUpHandler(svc))
|
||||
auth.GET("/users/:userId/shows/nextup", embyNextUpHandler(svc))
|
||||
auth.GET("/mediasegments/:id", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/artists", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/persons", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/genres", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/genres", embyGenresHandler(svc))
|
||||
auth.GET("/shows/upcoming", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/users/:userId/shows/upcoming", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/items/:id/similar", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/items/:id/similar", embySimilarHandler(svc))
|
||||
auth.GET("/items/:id/thumbnailset", embyEmptyItemsHandler(svc))
|
||||
auth.GET("/items/:id/thememedia", embyThemeMediaHandler(svc))
|
||||
auth.GET("/users/:userId/items/:id/specialfeatures", embyEmptyItemsHandler(svc))
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
// 媒体库筛选面板接口:facets 提供可选项,random 提供「随便看看」。
|
||||
//
|
||||
// 两者都走与列表完全相同的可见性判定(mediaVisibilityForRequest)与筛选解析
|
||||
// (parseLibraryFilters),因此不会出现「列表里有、facets 里没有」或「随机跳
|
||||
// 到了筛选条件之外的条目」这类不一致。
|
||||
|
||||
func libraryFacetsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libraryID := c.Param("id")
|
||||
facets, err := svc.Media.LibraryFacets(
|
||||
c.Request.Context(),
|
||||
libraryID,
|
||||
mediaVisibilityForRequest(c, svc),
|
||||
svc.Discovery,
|
||||
)
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, facets)
|
||||
}
|
||||
}
|
||||
|
||||
func libraryRandomHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libraryID := c.Param("id")
|
||||
filters := parseLibraryFilters(c)
|
||||
// 随机只取一条,因此不带分页参数;未观看筛选仍需要会话用户。
|
||||
media, err := svc.Media.RandomMedia(
|
||||
c.Request.Context(),
|
||||
libraryID,
|
||||
mediaVisibilityForRequest(c, svc),
|
||||
filters,
|
||||
)
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
}
|
||||
if media == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "no media matches the current filters"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, media)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,283 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/middleware"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
func newLibraryFilterEnv(t *testing.T) (*gin.Engine, *service.Container, string, string) {
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(
|
||||
&model.User{}, &model.Library{}, &model.Media{}, &model.PlaybackHistory{}, &model.Setting{},
|
||||
); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
svc := &service.Container{Repo: repos, Log: zap.NewNop()}
|
||||
svc.Media = service.NewMediaService(nil, zap.NewNop(), repos)
|
||||
svc.Discovery = service.NewMediaDiscoveryService(zap.NewNop(), repos)
|
||||
|
||||
const userID = "user-1"
|
||||
if err := repos.User.Create(context.Background(), &model.User{
|
||||
Base: model.Base{ID: userID}, Username: "tester", PasswordHash: "x", Role: "user", IsActive: true,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
lib := &model.Library{Name: "电影", Path: "/media/movies", Type: "movie", Enabled: true}
|
||||
if err := repos.Library.Create(context.Background(), lib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
authed := router.Group("/api", func(c *gin.Context) {
|
||||
c.Set(middleware.CtxUserID, userID)
|
||||
c.Set(middleware.CtxUserRole, "user")
|
||||
c.Next()
|
||||
})
|
||||
authed.GET("/libraries/:id/media", listMediaHandler(svc))
|
||||
authed.GET("/libraries/:id/facets", libraryFacetsHandler(svc))
|
||||
authed.GET("/libraries/:id/random", libraryRandomHandler(svc))
|
||||
return router, svc, userID, lib.ID
|
||||
}
|
||||
|
||||
func seedLibraryMedia(t *testing.T, svc *service.Container, rows ...*model.Media) {
|
||||
t.Helper()
|
||||
for _, row := range rows {
|
||||
if err := svc.Repo.DB.Create(row).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func getJSON(t *testing.T, router *gin.Engine, path string) (int, []byte) {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodGet, path, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
return w.Code, w.Body.Bytes()
|
||||
}
|
||||
|
||||
func TestLibraryFacetsReturnGenresAndYears(t *testing.T) {
|
||||
router, svc, _, libID := newLibraryFilterEnv(t)
|
||||
seedLibraryMedia(t, svc,
|
||||
&model.Media{LibraryID: libID, Title: "A", Genres: "Action,Drama", Year: 1999, Path: "/a.mkv"},
|
||||
&model.Media{LibraryID: libID, Title: "B", Genres: "Action", Year: 2021, Path: "/b.mkv"},
|
||||
)
|
||||
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/facets")
|
||||
if code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", code, body)
|
||||
}
|
||||
var facets struct {
|
||||
Genres []struct {
|
||||
Name string `json:"name"`
|
||||
Count int `json:"count"`
|
||||
} `json:"genres"`
|
||||
YearMin int `json:"year_min"`
|
||||
YearMax int `json:"year_max"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &facets); err != nil {
|
||||
t.Fatalf("decode %s: %v", body, err)
|
||||
}
|
||||
if facets.YearMin != 1999 || facets.YearMax != 2021 {
|
||||
t.Fatalf("year range = %d..%d, want 1999..2021", facets.YearMin, facets.YearMax)
|
||||
}
|
||||
if len(facets.Genres) != 2 {
|
||||
t.Fatalf("genres = %+v, want 2 entries", facets.Genres)
|
||||
}
|
||||
if facets.Genres[0].Name != "Action" || facets.Genres[0].Count != 2 {
|
||||
t.Fatalf("first genre = %+v, want Action:2", facets.Genres[0])
|
||||
}
|
||||
}
|
||||
|
||||
// 空库时 facets 必须返回空数组而不是 null,前端无需额外判空。
|
||||
func TestLibraryFacetsEmptyLibrary(t *testing.T) {
|
||||
router, _, _, libID := newLibraryFilterEnv(t)
|
||||
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/facets")
|
||||
if code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", code, body)
|
||||
}
|
||||
if !containsSubstring(string(body), `"genres":[]`) {
|
||||
t.Fatalf("body = %s, want genres:[]", body)
|
||||
}
|
||||
}
|
||||
|
||||
// 列表筛选:按类型过滤后只返回命中的条目。
|
||||
func TestListMediaAppliesGenreFilter(t *testing.T) {
|
||||
router, svc, _, libID := newLibraryFilterEnv(t)
|
||||
seedLibraryMedia(t, svc,
|
||||
&model.Media{LibraryID: libID, Title: "动作", Genres: "Action", Path: "/a.mkv"},
|
||||
&model.Media{LibraryID: libID, Title: "喜剧", Genres: "Comedy", Path: "/b.mkv"},
|
||||
)
|
||||
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/media?group_versions=0&genre=Action")
|
||||
if code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", code, body)
|
||||
}
|
||||
var payload struct {
|
||||
Items []model.Media `json:"items"`
|
||||
Total int64 `json:"total"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
t.Fatalf("decode %s: %v", body, err)
|
||||
}
|
||||
if payload.Total != 1 || len(payload.Items) != 1 || payload.Items[0].Title != "动作" {
|
||||
t.Fatalf("payload = %+v, want only 动作", payload)
|
||||
}
|
||||
}
|
||||
|
||||
// 筛选条件必须进入缓存键:先请求未筛选列表、再筛选时不能命中旧缓存。
|
||||
func TestListMediaFilterBypassesUnfilteredCache(t *testing.T) {
|
||||
router, svc, _, libID := newLibraryFilterEnv(t)
|
||||
seedLibraryMedia(t, svc,
|
||||
&model.Media{LibraryID: libID, Title: "动作", Genres: "Action", Path: "/a.mkv"},
|
||||
&model.Media{LibraryID: libID, Title: "喜剧", Genres: "Comedy", Path: "/b.mkv"},
|
||||
)
|
||||
|
||||
// 先拉全量(可能写缓存),再拉筛选结果。
|
||||
if code, body := getJSON(t, router, "/api/libraries/"+libID+"/media?group_versions=0"); code != http.StatusOK {
|
||||
t.Fatalf("unfiltered status = %d body=%s", code, body)
|
||||
}
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/media?group_versions=0&genre=Comedy")
|
||||
if code != http.StatusOK {
|
||||
t.Fatalf("filtered status = %d body=%s", code, body)
|
||||
}
|
||||
var payload struct {
|
||||
Total int64 `json:"total"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if payload.Total != 1 {
|
||||
t.Fatalf("total = %d, want 1 (filtered response must not be served from the unfiltered cache)", payload.Total)
|
||||
}
|
||||
}
|
||||
|
||||
// 未观看筛选:已看完的不出现,看了一半的仍出现。
|
||||
func TestListMediaUnwatchedFilter(t *testing.T) {
|
||||
router, svc, userID, libID := newLibraryFilterEnv(t)
|
||||
seedLibraryMedia(t, svc,
|
||||
&model.Media{Base: model.Base{ID: "m-done"}, LibraryID: libID, Title: "看完", Path: "/a.mkv"},
|
||||
&model.Media{Base: model.Base{ID: "m-half"}, LibraryID: libID, Title: "看一半", Path: "/b.mkv"},
|
||||
)
|
||||
for _, h := range []*model.PlaybackHistory{
|
||||
{UserID: userID, MediaID: "m-done", Completed: true},
|
||||
{UserID: userID, MediaID: "m-half", Completed: false},
|
||||
} {
|
||||
if err := svc.Repo.DB.Create(h).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/media?group_versions=0&unwatched=1")
|
||||
if code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", code, body)
|
||||
}
|
||||
var payload struct {
|
||||
Items []model.Media `json:"items"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(payload.Items) != 1 || payload.Items[0].Title != "看一半" {
|
||||
t.Fatalf("items = %+v, want only 看一半", payload.Items)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLibraryRandomReturnsMedia(t *testing.T) {
|
||||
router, svc, _, libID := newLibraryFilterEnv(t)
|
||||
seedLibraryMedia(t, svc,
|
||||
&model.Media{LibraryID: libID, Title: "唯一", Genres: "Action", Path: "/a.mkv"},
|
||||
)
|
||||
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/random")
|
||||
if code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", code, body)
|
||||
}
|
||||
var media model.Media
|
||||
if err := json.Unmarshal(body, &media); err != nil {
|
||||
t.Fatalf("decode %s: %v", body, err)
|
||||
}
|
||||
if media.Title != "唯一" {
|
||||
t.Fatalf("title = %q, want 唯一", media.Title)
|
||||
}
|
||||
}
|
||||
|
||||
// 筛选后没有命中时返回 404,前端据此提示「没有符合条件的媒体」。
|
||||
func TestLibraryRandomEmptyResultIs404(t *testing.T) {
|
||||
router, svc, _, libID := newLibraryFilterEnv(t)
|
||||
seedLibraryMedia(t, svc,
|
||||
&model.Media{LibraryID: libID, Title: "动作", Genres: "Action", Path: "/a.mkv"},
|
||||
)
|
||||
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/random?genre=Nonexistent")
|
||||
if code != http.StatusNotFound {
|
||||
t.Fatalf("status = %d body=%s, want 404", code, body)
|
||||
}
|
||||
}
|
||||
|
||||
// axios 默认把数组序列化为 genre[]=Action 格式;后端必须把它当作 genre=Action 处理。
|
||||
func TestListMediaAcceptsBracketGenreParam(t *testing.T) {
|
||||
router, svc, _, libID := newLibraryFilterEnv(t)
|
||||
seedLibraryMedia(t, svc,
|
||||
&model.Media{LibraryID: libID, Title: "动作", Genres: "Action", Path: "/a.mkv"},
|
||||
&model.Media{LibraryID: libID, Title: "喜剧", Genres: "Comedy", Path: "/b.mkv"},
|
||||
)
|
||||
|
||||
// genre[]=Action — axios bracket format without custom paramsSerializer
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/media?group_versions=0&genre[]=Action")
|
||||
if code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", code, body)
|
||||
}
|
||||
var payload struct {
|
||||
Items []model.Media `json:"items"`
|
||||
Total int64 `json:"total"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
t.Fatalf("decode %s: %v", body, err)
|
||||
}
|
||||
if payload.Total != 1 || len(payload.Items) != 1 || payload.Items[0].Title != "动作" {
|
||||
t.Fatalf("payload = %+v, want only 动作 for genre[]=Action", payload)
|
||||
}
|
||||
}
|
||||
|
||||
// 随机也遵守筛选:只命中 Action 时,带 Comedy 筛选必须 404。
|
||||
func TestLibraryRandomHonoursFilters(t *testing.T) {
|
||||
router, svc, _, libID := newLibraryFilterEnv(t)
|
||||
seedLibraryMedia(t, svc,
|
||||
&model.Media{LibraryID: libID, Title: "A", Genres: "Action", Year: 2001, Path: "/a.mkv"},
|
||||
&model.Media{LibraryID: libID, Title: "B", Genres: "Comedy", Year: 2002, Path: "/b.mkv"},
|
||||
)
|
||||
|
||||
code, body := getJSON(t, router, "/api/libraries/"+libID+"/random?genre=Comedy&year_min=2002")
|
||||
if code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", code, body)
|
||||
}
|
||||
var media model.Media
|
||||
if err := json.Unmarshal(body, &media); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if media.Title != "B" {
|
||||
t.Fatalf("title = %q, want B", media.Title)
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@ package handler
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"math"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -407,6 +408,92 @@ func deleteLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// parseLibraryFilters 解析媒体库列表的筛选查询参数。
|
||||
//
|
||||
// 全部参数都是可选的:缺省时返回零值,`MediaListFilters.empty()` 为真,列表
|
||||
// 行为与此前完全一致(不引入任何默认筛选)。
|
||||
//
|
||||
// 参数约定:
|
||||
// - genre=Action&genre=Comedy 类型多选(或关系,整词匹配)
|
||||
// - year_min / year_max 年份区间,0 或非法值表示不限
|
||||
// - rating_min 评分下限(浮点)
|
||||
// - unwatched=1 仅显示未看完;用户 ID 取自会话
|
||||
func parseLibraryFilters(c *gin.Context) service.MediaListFilters {
|
||||
filters := service.MediaListFilters{
|
||||
Genres: parseRepeatedQueryValues(c, "genre"),
|
||||
YearMin: parseNonNegativeInt(firstQueryValue(c, "year_min", "yearMin")),
|
||||
YearMax: parseNonNegativeInt(firstQueryValue(c, "year_max", "yearMax")),
|
||||
RatingMin: parseNonNegativeFloat(firstQueryValue(c, "rating_min", "ratingMin")),
|
||||
}
|
||||
if isTruthyQuery(firstQueryValue(c, "unwatched", "unwatched_only", "unwatchedOnly")) {
|
||||
filters.Unwatched = true
|
||||
filters.UserID = toString(mustSessionUserID(c))
|
||||
}
|
||||
return filters
|
||||
}
|
||||
|
||||
// parseRepeatedQueryValues 读取可重复出现的查询参数,去重并丢弃空值。
|
||||
// 同时接受 key[] 括号格式(axios 1.x 默认序列化方式)作为向后兼容回退,
|
||||
// 在前端 paramsSerializer 未正确配置时不会静默返回空结果。
|
||||
func parseRepeatedQueryValues(c *gin.Context, key string) []string {
|
||||
raw := c.QueryArray(key)
|
||||
if len(raw) == 0 {
|
||||
// fallback: axios bracket format (e.g. genre[]=Action&genre[]=Comedy)
|
||||
raw = c.QueryArray(key + "[]")
|
||||
}
|
||||
if len(raw) == 0 {
|
||||
return nil
|
||||
}
|
||||
seen := make(map[string]struct{}, len(raw))
|
||||
out := make([]string, 0, len(raw))
|
||||
for _, value := range raw {
|
||||
// 客户端可能把多值拼成一次逗号分隔,两种形式都要接受。
|
||||
for _, part := range strings.Split(value, ",") {
|
||||
trimmed := strings.TrimSpace(part)
|
||||
if trimmed == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[trimmed]; ok {
|
||||
continue
|
||||
}
|
||||
seen[trimmed] = struct{}{}
|
||||
out = append(out, trimmed)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func parseNonNegativeInt(raw string) int {
|
||||
value, err := strconv.Atoi(strings.TrimSpace(raw))
|
||||
if err != nil || value < 0 {
|
||||
return 0
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func parseNonNegativeFloat(raw string) float64 {
|
||||
value, err := strconv.ParseFloat(strings.TrimSpace(raw), 64)
|
||||
if err != nil || value < 0 || math.IsNaN(value) {
|
||||
return 0
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func isTruthyQuery(raw string) bool {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "1", "true", "yes", "on":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// mustSessionUserID 取会话用户 ID,缺失时返回空串(筛选逻辑会忽略它)。
|
||||
func mustSessionUserID(c *gin.Context) any {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
return uid
|
||||
}
|
||||
|
||||
func listMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
@@ -450,9 +537,10 @@ func listMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
if sortSpec.Field == "last_played" {
|
||||
history = mediaHistoryMap(c, svc)
|
||||
}
|
||||
filters := parseLibraryFilters(c)
|
||||
groupVersions := c.DefaultQuery("group_versions", "1") != "0"
|
||||
if !groupVersions {
|
||||
items, total, err := svc.Media.ListMediaVisible(ctx, id, page, size, mediaVisibilityForRequest(c, svc))
|
||||
items, total, err := svc.Media.ListMediaVisibleFiltered(ctx, id, page, size, mediaVisibilityForRequest(c, svc), filters)
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
@@ -468,7 +556,7 @@ func listMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
})
|
||||
return
|
||||
}
|
||||
grouped, err := svc.Media.GroupedMediaVisible(ctx, id, mediaVisibilityForRequest(c, svc))
|
||||
grouped, err := svc.Media.GroupedMediaVisibleFiltered(ctx, id, mediaVisibilityForRequest(c, svc), filters)
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
|
||||
@@ -116,8 +116,13 @@ func registerAdminUserRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.PATCH("/users/:id/role", adminUpdateRoleHandler(svc))
|
||||
admin.PATCH("/users/:id/libraries", updateUserLibrariesHandler(svc))
|
||||
admin.DELETE("/users/:id", deleteUserHandler(svc))
|
||||
// 设备代管:管理员可查看并踢掉任意用户的设备。
|
||||
admin.GET("/users/:id/devices", adminUserDevicesHandler(svc))
|
||||
admin.POST("/users/:id/devices/kick-all", adminKickAllUserDevicesHandler(svc))
|
||||
admin.POST("/users/:id/devices/:deviceID/kick", adminKickUserDeviceHandler(svc))
|
||||
admin.GET("/settings", listSettingsHandler(svc))
|
||||
admin.PUT("/settings", updateSettingHandler(svc))
|
||||
admin.POST("/telegram/test", testTelegramHandler(svc))
|
||||
admin.POST("/adult/test-scraper", testAdultScraperHandler(svc))
|
||||
admin.GET("/logs", recentLogsHandler(svc))
|
||||
}
|
||||
|
||||
@@ -19,6 +19,16 @@ func registerAuthedUserAndLicenseRoutes(authed *gin.RouterGroup, svc *service.Co
|
||||
authed.GET("/me/temporary-password", temporaryPasswordHandler(svc))
|
||||
authed.POST("/me/temporary-password", temporaryPasswordHandler(svc))
|
||||
|
||||
// 设备管理:路由挂在 /me 下,用户 ID 一律取自会话,天然只能管自己的设备。
|
||||
authed.GET("/me/devices", myDevicesHandler(svc))
|
||||
authed.POST("/me/devices/kick-all", myKickAllDevicesHandler(svc))
|
||||
authed.POST("/me/devices/:deviceID/kick", myKickDeviceHandler(svc))
|
||||
|
||||
// Telegram 通知绑定:一次性码 + Bot /bind <code>。
|
||||
authed.GET("/me/telegram", getTelegramStatusHandler(svc))
|
||||
authed.POST("/me/telegram/bind-code", startTelegramBindHandler(svc))
|
||||
authed.DELETE("/me/telegram", unbindTelegramHandler(svc))
|
||||
|
||||
authed.GET("/auth/permissions", getMyPermissionsHandler(svc))
|
||||
}
|
||||
|
||||
@@ -38,6 +48,8 @@ func registerAuthedLibraryRoutes(authed *gin.RouterGroup, svc *service.Container
|
||||
authed.POST("/libraries/:id/scrape", middleware.AdminRequired(), scrapeLibraryHandler(svc))
|
||||
|
||||
authed.GET("/libraries/:id/media", listMediaHandler(svc))
|
||||
authed.GET("/libraries/:id/facets", libraryFacetsHandler(svc))
|
||||
authed.GET("/libraries/:id/random", libraryRandomHandler(svc))
|
||||
authed.GET("/libraries/:id/series", listLibrarySeriesHandler(svc))
|
||||
authed.GET("/libraries/:id/series/episodes", listLibrarySeriesEpisodesHandler(svc))
|
||||
authed.GET("/libraries/:id/seasons", listSeasonsHandler(svc))
|
||||
|
||||
@@ -15,7 +15,7 @@ func registerAuthedUISurfaceRoutes(authed *gin.RouterGroup, svc *service.Contain
|
||||
authed.PUT("/danmaku/settings", updateDanmakuSettingsHandler(svc))
|
||||
|
||||
authed.GET("/watch-history", historyListHandler(svc))
|
||||
authed.GET("/watch-history/stats", historyStatsHandler(svc))
|
||||
authed.GET("/watch-history/stats", requirePermission(svc, "can_view_history"), historyStatsHandler(svc))
|
||||
authed.GET("/watch-history/continue", historyContinueHandler(svc))
|
||||
authed.DELETE("/watch-history", historyDeleteHandler(svc))
|
||||
authed.DELETE("/watch-history/:id", historyDeleteOneHandler(svc))
|
||||
|
||||
@@ -124,7 +124,9 @@ func listLibrarySeriesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
}
|
||||
items, total, err := svc.Media.ListLibrarySeriesCards(c.Request.Context(), libID, mediaVisibilityForRequest(c, svc))
|
||||
items, total, err := svc.Media.ListLibrarySeriesCardsFiltered(
|
||||
c.Request.Context(), libID, mediaVisibilityForRequest(c, svc), parseLibraryFilters(c),
|
||||
)
|
||||
if err != nil {
|
||||
writeInternalOrCanceled(c, err)
|
||||
return
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
// Telegram 绑定与测试接口。
|
||||
//
|
||||
// 绑定刻意做成「网页生成一次性码 → 用户在 Bot 里发 /bind <码>」:服务端不需要
|
||||
// 用户手工填写 chat id,也不需要站点暴露 Bot 命令以外任何能力。
|
||||
|
||||
type telegramBindCodePayload struct {
|
||||
Code string `json:"code"`
|
||||
ExpiresIn int `json:"expires_in_seconds"`
|
||||
}
|
||||
|
||||
// startTelegramBindHandler 生成一次性绑定码。同一用户重复调用时旧码作废。
|
||||
func startTelegramBindHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc.Telegram == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "telegram service unavailable"})
|
||||
return
|
||||
}
|
||||
userID := sessionUserID(c)
|
||||
code, err := svc.Telegram.StartBind(c.Request.Context(), userID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, telegramBindCodePayload{Code: code, ExpiresIn: 300})
|
||||
}
|
||||
}
|
||||
|
||||
// getTelegramStatusHandler 返回绑定状态与脱敏会话 ID。
|
||||
func getTelegramStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc.Telegram == nil {
|
||||
c.JSON(http.StatusOK, gin.H{"bound": false})
|
||||
return
|
||||
}
|
||||
bound, masked := svc.Telegram.Status(c.Request.Context(), sessionUserID(c))
|
||||
payload := gin.H{"bound": bound}
|
||||
if bound {
|
||||
payload["chat_id_masked"] = masked
|
||||
}
|
||||
c.JSON(http.StatusOK, payload)
|
||||
}
|
||||
}
|
||||
|
||||
// unbindTelegramHandler 解除当前用户的 Telegram 绑定。
|
||||
func unbindTelegramHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc.Telegram == nil {
|
||||
c.Status(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
if err := svc.Telegram.Unbind(c.Request.Context(), sessionUserID(c)); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
// testTelegramHandler 向管理员会话发送一条测试消息,用于验证 Token/会话 ID。
|
||||
func testTelegramHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc.Telegram == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "telegram service unavailable"})
|
||||
return
|
||||
}
|
||||
if !svc.Telegram.Configured(c.Request.Context()) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"success": false,
|
||||
"error": "请先启用 Telegram 通知并填写 Bot Token 与管理员 Chat ID",
|
||||
})
|
||||
return
|
||||
}
|
||||
if err := svc.Telegram.SendToAdminChecked(c.Request.Context(), "✅ MeBox 测试消息:通知通道工作正常。"); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"success": false, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"success": true})
|
||||
}
|
||||
}
|
||||
@@ -11,8 +11,11 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -48,11 +51,16 @@ func historyListHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
|
||||
// historyStatsHandler returns aggregate watch time + completion counts
|
||||
// for the caller. Used by the WatchHistoryPage hero card.
|
||||
// for the caller. Used by the WatchHistoryPage hero card and the dedicated
|
||||
// personal statistics page.
|
||||
//
|
||||
// 统计口径全部来自 PlaybackHistory 本身,不新增统计表:position_ms 是「已看
|
||||
// 时长」的近似值,足以支撑趋势图;精确到秒的播放时长另有会话统计负责。
|
||||
func historyStatsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
userID := toString(uid)
|
||||
ctx := c.Request.Context()
|
||||
|
||||
var total int64
|
||||
_ = svc.Repo.DB.Model(&model.PlaybackHistory{}).
|
||||
@@ -77,16 +85,196 @@ func historyStatsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
last = &lastT
|
||||
}
|
||||
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
daily, byType, recent := historyStatsBreakdowns(ctx, svc, userID, visibility)
|
||||
|
||||
inProgress := total - completed
|
||||
if inProgress < 0 {
|
||||
inProgress = 0
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"total": total,
|
||||
"completed": completed,
|
||||
"watched_ms": watchedMs,
|
||||
"watched_hours": float64(watchedMs) / 1000.0 / 3600.0,
|
||||
"last_watched": last,
|
||||
"total": total,
|
||||
"completed": completed,
|
||||
"in_progress": inProgress,
|
||||
"watched_ms": watchedMs,
|
||||
"watched_hours": float64(watchedMs) / 1000.0 / 3600.0,
|
||||
"last_watched": last,
|
||||
"daily": daily,
|
||||
"by_library_type": byType,
|
||||
"recent": recent,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// historyStatsDailyDays 是趋势图回看的天数。
|
||||
const historyStatsDailyDays = 30
|
||||
|
||||
// historyStatsRecentLimit 是「最近看过」返回的条数。
|
||||
const historyStatsRecentLimit = 8
|
||||
|
||||
type historyDailyStat struct {
|
||||
Day string `json:"day"`
|
||||
WatchMs int64 `json:"watch_ms"`
|
||||
Plays int64 `json:"plays"`
|
||||
}
|
||||
|
||||
type historyTypeStat struct {
|
||||
Type string `json:"type"`
|
||||
WatchMs int64 `json:"watch_ms"`
|
||||
Count int64 `json:"count"`
|
||||
}
|
||||
|
||||
// historyStatsBreakdowns 产出每日趋势、按媒体库类型分布与最近记录。
|
||||
//
|
||||
// 分桶在 Go 里做而不是用 SQL 的日期函数:SQLite 的 strftime 与 PostgreSQL 的
|
||||
// to_char 语法不同,写两份 SQL 会在方言差异上长期出错,而历史行数受用户规模
|
||||
// 约束(每人一行一部媒体),一次全量读取是可以接受的。
|
||||
//
|
||||
// visibility 控制哪些媒体对调用者可见(播放档案、成人锁等)。
|
||||
func historyStatsBreakdowns(ctx context.Context, svc *service.Container, userID string, visibility service.MediaVisibility) ([]historyDailyStat, []historyTypeStat, []map[string]any) {
|
||||
daily := make([]historyDailyStat, 0, historyStatsDailyDays)
|
||||
byType := make([]historyTypeStat, 0)
|
||||
recent := make([]map[string]any, 0, historyStatsRecentLimit)
|
||||
|
||||
var rows []model.PlaybackHistory
|
||||
if err := svc.Repo.DB.WithContext(ctx).
|
||||
Where("user_id = ?", userID).
|
||||
Order("watched_at desc").
|
||||
Find(&rows).Error; err != nil || len(rows) == 0 {
|
||||
return daily, byType, recent
|
||||
}
|
||||
|
||||
// 每日趋势:只回看最近 N 天,且按「本地日」分桶,避免跨时区偏移。
|
||||
now := time.Now()
|
||||
cutoff := now.AddDate(0, 0, -(historyStatsDailyDays - 1))
|
||||
startOfDay := func(t time.Time) time.Time {
|
||||
local := t.In(time.Local)
|
||||
return time.Date(local.Year(), local.Month(), local.Day(), 0, 0, 0, 0, time.Local)
|
||||
}
|
||||
buckets := make(map[string]*historyDailyStat, historyStatsDailyDays)
|
||||
for i := 0; i < historyStatsDailyDays; i++ {
|
||||
day := startOfDay(cutoff).AddDate(0, 0, i).Format("2006-01-02")
|
||||
buckets[day] = &historyDailyStat{Day: day}
|
||||
}
|
||||
for _, r := range rows {
|
||||
watched := r.WatchedAt.In(time.Local)
|
||||
if watched.Before(startOfDay(cutoff)) {
|
||||
continue
|
||||
}
|
||||
if bucket, ok := buckets[watched.Format("2006-01-02")]; ok {
|
||||
bucket.WatchMs += r.PositionMs
|
||||
bucket.Plays++
|
||||
}
|
||||
}
|
||||
for i := 0; i < historyStatsDailyDays; i++ {
|
||||
day := startOfDay(cutoff).AddDate(0, 0, i).Format("2006-01-02")
|
||||
if bucket, ok := buckets[day]; ok && bucket.Plays > 0 {
|
||||
daily = append(daily, *bucket)
|
||||
}
|
||||
}
|
||||
|
||||
mediaIDs := make([]string, 0, len(rows))
|
||||
for _, r := range rows {
|
||||
mediaIDs = append(mediaIDs, r.MediaID)
|
||||
}
|
||||
var medias []model.Media
|
||||
_ = svc.Repo.DB.WithContext(ctx).Where("id IN ?", mediaIDs).Find(&medias).Error
|
||||
mediaByID := make(map[string]*model.Media, len(medias))
|
||||
for i := range medias {
|
||||
mediaByID[medias[i].ID] = &medias[i]
|
||||
}
|
||||
|
||||
libraryTypes := make(map[string]string)
|
||||
var libraries []model.Library
|
||||
if svc.Repo.Library != nil {
|
||||
if libs, err := svc.Repo.Library.List(ctx); err == nil {
|
||||
libraries = libs
|
||||
}
|
||||
}
|
||||
for _, lib := range libraries {
|
||||
libraryTypes[lib.ID] = lib.Type
|
||||
}
|
||||
|
||||
typeAcc := make(map[string]*historyTypeStat)
|
||||
order := make([]string, 0, 4)
|
||||
for _, r := range rows {
|
||||
media := mediaByID[r.MediaID]
|
||||
var key string
|
||||
if media == nil {
|
||||
// 媒体记录已删除(含 Emby 远程缓存失效):计入 "other" 桶而非丢弃,
|
||||
// 这样类型分布总数才能与播放历史总数吻合。
|
||||
key = "other"
|
||||
} else {
|
||||
// 如果调用者的可见性策略排除了该媒体,则跳过统计(visibility leak fix)。
|
||||
if !visibility.Allows(media) {
|
||||
continue
|
||||
}
|
||||
key = strings.TrimSpace(libraryTypes[media.LibraryID])
|
||||
if key == "" {
|
||||
key = "other"
|
||||
}
|
||||
}
|
||||
acc, ok := typeAcc[key]
|
||||
if !ok {
|
||||
acc = &historyTypeStat{Type: key}
|
||||
typeAcc[key] = acc
|
||||
order = append(order, key)
|
||||
}
|
||||
acc.WatchMs += r.PositionMs
|
||||
acc.Count++
|
||||
}
|
||||
// 顺序按观看时长降序,让「我主要在看什么」一眼可见。
|
||||
for _, key := range order {
|
||||
byType = append(byType, *typeAcc[key])
|
||||
}
|
||||
sort.SliceStable(byType, func(i, j int) bool {
|
||||
if byType[i].WatchMs != byType[j].WatchMs {
|
||||
return byType[i].WatchMs > byType[j].WatchMs
|
||||
}
|
||||
return byType[i].Type < byType[j].Type
|
||||
})
|
||||
|
||||
for _, r := range rows {
|
||||
if len(recent) >= historyStatsRecentLimit {
|
||||
break
|
||||
}
|
||||
entry := map[string]any{"history": r}
|
||||
if media := mediaByID[r.MediaID]; media != nil {
|
||||
// 可见性检查:隐藏库或受档案限制的媒体不进入最近记录(visibility leak fix)。
|
||||
if !visibility.Allows(media) {
|
||||
continue
|
||||
}
|
||||
entry["media"] = media
|
||||
} else if svc.EmbyRemote != nil && service.IsEmbyRemoteID(r.MediaID) {
|
||||
// 尝试从 Emby 远端补全媒体详情,与 historyContinueHandler 保持相同策略。
|
||||
mountID, remoteID, _ := service.DecodeEmbyRemoteID(r.MediaID)
|
||||
mount, acct, resolveErr := svc.EmbyRemote.ResolveMount(ctx, mountID)
|
||||
if resolveErr == nil && mount != nil && acct != nil {
|
||||
remoteMedia, detailErr := svc.EmbyRemote.RemoteMediaDetail(ctx, mount, acct, remoteID)
|
||||
if detailErr == nil && remoteMedia != nil {
|
||||
if !visibility.Allows(remoteMedia) {
|
||||
continue
|
||||
}
|
||||
entry["media"] = *remoteMedia
|
||||
} else {
|
||||
// 无法获取 Emby 媒体详情,跳过此条记录。
|
||||
continue
|
||||
}
|
||||
} else {
|
||||
// 挂载不可用,跳过。
|
||||
continue
|
||||
}
|
||||
} else {
|
||||
// 媒体记录不存在且无法 Emby 补全,跳过。
|
||||
continue
|
||||
}
|
||||
recent = append(recent, entry)
|
||||
}
|
||||
|
||||
return daily, byType, recent
|
||||
}
|
||||
|
||||
// historyContinueHandler returns "Continue Watching" rows: incomplete
|
||||
// items, most recent first.
|
||||
func historyContinueHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
@@ -0,0 +1,349 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/truewhile/MeBox/internal/middleware"
|
||||
"github.com/truewhile/MeBox/internal/model"
|
||||
"github.com/truewhile/MeBox/internal/repository"
|
||||
"github.com/truewhile/MeBox/internal/service"
|
||||
)
|
||||
|
||||
func newHistoryStatsEnv(t *testing.T) (*gin.Engine, *service.Container, string) {
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.PlaybackHistory{}, &model.UserPermission{}); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
svc := &service.Container{Repo: repos, Log: zap.NewNop()}
|
||||
svc.Permissions = service.NewPermissionService(zap.NewNop(), repos)
|
||||
|
||||
const userID = "user-1"
|
||||
if err := repos.User.Create(context.Background(), &model.User{
|
||||
Base: model.Base{ID: userID}, Username: "tester", PasswordHash: "x", Role: "user", IsActive: true,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
authed := router.Group("/api", func(c *gin.Context) {
|
||||
c.Set(middleware.CtxUserID, userID)
|
||||
c.Next()
|
||||
})
|
||||
authed.GET("/watch-history/stats", historyStatsHandler(svc))
|
||||
return router, svc, userID
|
||||
}
|
||||
|
||||
type historyStatsPayload struct {
|
||||
Total int64 `json:"total"`
|
||||
Completed int64 `json:"completed"`
|
||||
InProgress int64 `json:"in_progress"`
|
||||
WatchedMs int64 `json:"watched_ms"`
|
||||
WatchedHours float64 `json:"watched_hours"`
|
||||
Daily []struct {
|
||||
Day string `json:"day"`
|
||||
WatchMs int64 `json:"watch_ms"`
|
||||
Plays int64 `json:"plays"`
|
||||
} `json:"daily"`
|
||||
ByLibraryType []struct {
|
||||
Type string `json:"type"`
|
||||
WatchMs int64 `json:"watch_ms"`
|
||||
Count int64 `json:"count"`
|
||||
} `json:"by_library_type"`
|
||||
Recent []struct {
|
||||
Media *model.Media `json:"media"`
|
||||
} `json:"recent"`
|
||||
}
|
||||
|
||||
func fetchHistoryStats(t *testing.T, router *gin.Engine) historyStatsPayload {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/watch-history/stats", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
var payload historyStatsPayload
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("decode %s: %v", w.Body.String(), err)
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
// 新字段必须提供每日聚合、库类型分布与在看数量,供个人统计页绘图。
|
||||
func TestHistoryStatsIncludesDailyAndTypes(t *testing.T) {
|
||||
router, svc, userID := newHistoryStatsEnv(t)
|
||||
ctx := context.Background()
|
||||
|
||||
movieLib := &model.Library{Name: "电影", Path: "/media/movies", Type: "movie", Enabled: true}
|
||||
if err := svc.Repo.Library.Create(ctx, movieLib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tvLib := &model.Library{Name: "剧集", Path: "/media/tv", Type: "tv", Enabled: true}
|
||||
if err := svc.Repo.Library.Create(ctx, tvLib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
yesterday := time.Now().Add(-24 * time.Hour)
|
||||
today := time.Now().Add(-time.Hour)
|
||||
|
||||
rows := []struct {
|
||||
media *model.Media
|
||||
watchedAt time.Time
|
||||
position int64
|
||||
completed bool
|
||||
}{
|
||||
{
|
||||
media: &model.Media{LibraryID: movieLib.ID, Title: "电影A", Path: "/media/movies/a.mkv"},
|
||||
watchedAt: yesterday, position: 60000, completed: true,
|
||||
},
|
||||
{
|
||||
media: &model.Media{LibraryID: tvLib.ID, Title: "剧B", Path: "/media/tv/b.mkv"},
|
||||
watchedAt: today, position: 30000, completed: false,
|
||||
},
|
||||
}
|
||||
for _, row := range rows {
|
||||
if err := svc.Repo.DB.Create(row.media).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
h := &model.PlaybackHistory{
|
||||
UserID: userID, MediaID: row.media.ID, PositionMs: row.position,
|
||||
DurationMs: 120000, WatchedAt: row.watchedAt, Completed: row.completed,
|
||||
}
|
||||
if err := svc.Repo.DB.Create(h).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
payload := fetchHistoryStats(t, router)
|
||||
|
||||
if payload.Total != 2 {
|
||||
t.Fatalf("total = %d, want 2", payload.Total)
|
||||
}
|
||||
if payload.Completed != 1 {
|
||||
t.Fatalf("completed = %d, want 1", payload.Completed)
|
||||
}
|
||||
if payload.InProgress != 1 {
|
||||
t.Fatalf("in_progress = %d, want 1", payload.InProgress)
|
||||
}
|
||||
if payload.WatchedMs != 90000 {
|
||||
t.Fatalf("watched_ms = %d, want 90000", payload.WatchedMs)
|
||||
}
|
||||
if len(payload.Daily) != 2 {
|
||||
t.Fatalf("daily = %+v, want 2 days", payload.Daily)
|
||||
}
|
||||
if len(payload.ByLibraryType) != 2 {
|
||||
t.Fatalf("by_library_type = %+v, want 2 entries", payload.ByLibraryType)
|
||||
}
|
||||
if len(payload.Recent) != 2 {
|
||||
t.Fatalf("recent = %d entries, want 2", len(payload.Recent))
|
||||
}
|
||||
if payload.Recent[0].Media == nil || payload.Recent[0].Media.Title != "剧B" {
|
||||
t.Fatalf("recent[0] = %+v, want the most recent entry (剧B)", payload.Recent[0])
|
||||
}
|
||||
}
|
||||
|
||||
// 没有任何播放记录时,新字段要返回空数组而不是 null,前端无需额外判空。
|
||||
func TestHistoryStatsEmptyProvidesEmptyArrays(t *testing.T) {
|
||||
router, _, _ := newHistoryStatsEnv(t)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/watch-history/stats", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
body := w.Body.String()
|
||||
for _, field := range []string{`"daily":[]`, `"by_library_type":[]`, `"recent":[]`} {
|
||||
if !containsSubstring(body, field) {
|
||||
t.Fatalf("body = %s, want %s", body, field)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func containsSubstring(haystack, needle string) bool {
|
||||
for i := 0; i+len(needle) <= len(haystack); i++ {
|
||||
if haystack[i:i+len(needle)] == needle {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// TestHistoryStatsBreakdownsVisibilityFilter validates that historyStatsBreakdowns
|
||||
// respects the caller's MediaVisibility: media in a hidden library must be absent
|
||||
// from both the recent list and the by_library_type buckets.
|
||||
func TestHistoryStatsBreakdownsVisibilityFilter(t *testing.T) {
|
||||
_, svc, userID := newHistoryStatsEnv(t)
|
||||
ctx := context.Background()
|
||||
|
||||
allowedLib := &model.Library{Name: "允许库", Path: "/media/allowed", Type: "movie", Enabled: true}
|
||||
hiddenLib := &model.Library{Name: "隐藏库", Path: "/media/hidden", Type: "tv", Enabled: true}
|
||||
if err := svc.Repo.Library.Create(ctx, allowedLib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.Repo.Library.Create(ctx, hiddenLib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
allowedMedia := &model.Media{LibraryID: allowedLib.ID, Title: "允许媒体", Path: "/media/allowed/a.mkv"}
|
||||
hiddenMedia := &model.Media{LibraryID: hiddenLib.ID, Title: "隐藏媒体", Path: "/media/hidden/b.mkv"}
|
||||
if err := svc.Repo.DB.Create(allowedMedia).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.Repo.DB.Create(hiddenMedia).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
for _, mid := range []string{allowedMedia.ID, hiddenMedia.ID} {
|
||||
h := &model.PlaybackHistory{
|
||||
UserID: userID, MediaID: mid, PositionMs: 10000,
|
||||
DurationMs: 100000, WatchedAt: now, Completed: false,
|
||||
}
|
||||
if err := svc.Repo.DB.Create(h).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// visibility that hides hiddenLib
|
||||
vis := service.MediaVisibility{
|
||||
HiddenLibraryIDs: []string{hiddenLib.ID},
|
||||
}
|
||||
|
||||
_, byType, recent := historyStatsBreakdowns(ctx, svc, userID, vis)
|
||||
|
||||
// recent must contain only the allowed media
|
||||
for _, entry := range recent {
|
||||
m, ok := entry["media"]
|
||||
if !ok {
|
||||
t.Fatal("recent entry missing media field")
|
||||
}
|
||||
switch med := m.(type) {
|
||||
case *model.Media:
|
||||
if med.LibraryID == hiddenLib.ID {
|
||||
t.Fatalf("hidden media appeared in recent: %s", med.Title)
|
||||
}
|
||||
case model.Media:
|
||||
if med.LibraryID == hiddenLib.ID {
|
||||
t.Fatalf("hidden media appeared in recent: %s", med.Title)
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(recent) != 1 {
|
||||
t.Fatalf("recent length = %d, want 1 (hidden entry must be excluded)", len(recent))
|
||||
}
|
||||
|
||||
// by_library_type must not contain the hidden library's type ("tv")
|
||||
for _, bt := range byType {
|
||||
if bt.Type == "tv" {
|
||||
t.Fatalf("hidden library type 'tv' appeared in by_library_type (count=%d)", bt.Count)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestHistoryStatsBreakdownsNilMediaCountsAsOther confirms that a history row
|
||||
// whose media has been deleted (nil lookup) is counted under the "other" type
|
||||
// bucket rather than silently dropped.
|
||||
func TestHistoryStatsBreakdownsNilMediaCountsAsOther(t *testing.T) {
|
||||
_, svc, userID := newHistoryStatsEnv(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Insert a history row whose media_id does not correspond to any Media row.
|
||||
ghost := &model.PlaybackHistory{
|
||||
UserID: userID,
|
||||
MediaID: "ghost-media-id",
|
||||
PositionMs: 5000,
|
||||
DurationMs: 50000,
|
||||
WatchedAt: time.Now(),
|
||||
Completed: false,
|
||||
}
|
||||
if err := svc.Repo.DB.Create(ghost).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
vis := service.MediaVisibility{} // unrestricted
|
||||
_, byType, _ := historyStatsBreakdowns(ctx, svc, userID, vis)
|
||||
|
||||
var otherEntry *historyTypeStat
|
||||
for i := range byType {
|
||||
if byType[i].Type == "other" {
|
||||
otherEntry = &byType[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
if otherEntry == nil {
|
||||
t.Fatalf("expected 'other' bucket for nil-media history row, got %+v", byType)
|
||||
}
|
||||
if otherEntry.Count != 1 {
|
||||
t.Fatalf("other.Count = %d, want 1", otherEntry.Count)
|
||||
}
|
||||
}
|
||||
|
||||
// TestHistoryStatsPermissionDeny checks that a user without can_view_history
|
||||
// receives HTTP 403 from the gated route.
|
||||
func TestHistoryStatsPermissionDeny(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(
|
||||
&model.User{}, &model.Library{}, &model.Media{},
|
||||
&model.PlaybackHistory{}, &model.UserPermission{},
|
||||
); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
svc := &service.Container{Repo: repos, Log: zap.NewNop()}
|
||||
svc.Permissions = service.NewPermissionService(zap.NewNop(), repos)
|
||||
|
||||
const userID = "user-noperm"
|
||||
if err := repos.User.Create(context.Background(), &model.User{
|
||||
Base: model.Base{ID: userID}, Username: "noperm", PasswordHash: "x", Role: "user", IsActive: true,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Explicitly deny can_view_history for this user.
|
||||
// First seed defaults (Effective will create the row with defaults), then
|
||||
// update to deny via Save which uses an explicit map update path in the repo.
|
||||
if _, err := svc.Permissions.Effective(context.Background(), userID); err != nil {
|
||||
t.Fatalf("seed permissions: %v", err)
|
||||
}
|
||||
denyPerm := &model.UserPermission{UserID: userID, CanViewHistory: false}
|
||||
if err := svc.Permissions.Save(context.Background(), userID, denyPerm); err != nil {
|
||||
t.Fatalf("save permission: %v", err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
authed := router.Group("/api", func(c *gin.Context) {
|
||||
c.Set(middleware.CtxUserID, userID)
|
||||
c.Set(middleware.CtxUserRole, "user")
|
||||
c.Next()
|
||||
})
|
||||
authed.GET("/watch-history/stats", requirePermission(svc, "can_view_history"), historyStatsHandler(svc))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/watch-history/stats", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusForbidden {
|
||||
t.Fatalf("status = %d, want 403 for user without can_view_history", w.Code)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user