mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-29 19:36:36 +08:00
初始化
初始化项目
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
# MediaStationGo
|
||||
# MMTL (My Movie and TV Library)
|
||||
|
||||
<p align="center">
|
||||
<img src="web/public/brand/mgo-emby-icon.svg" width="96" height="96" alt="MediaStationGo Logo" />
|
||||
<img src="web/public/brand/logo-192.png" width="96" height="96" alt="MMTL Logo" />
|
||||
</p>
|
||||
|
||||
<h3 align="center">适合 NAS、家庭共享和多端播放的私人媒体中心</h3>
|
||||
|
||||
+2
-2
@@ -1,7 +1,7 @@
|
||||
# MediaStationGo
|
||||
# MMTL (My Movie and TV Library)
|
||||
|
||||
<p align="center">
|
||||
<img src="web/public/brand/mgo-emby-icon.svg" width="96" height="96" alt="MediaStationGo Logo" />
|
||||
<img src="web/public/brand/logo-192.png" width="96" height="96" alt="MMTL Logo" />
|
||||
</p>
|
||||
|
||||
<h3 align="center">A lightweight, polished, NAS-friendly private media center</h3>
|
||||
|
||||
+2
-11
@@ -26,7 +26,6 @@ import (
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/config"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/database"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/handler"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
@@ -63,7 +62,7 @@ func main() {
|
||||
defer func() { _ = logger.Sync() }()
|
||||
|
||||
appVersion := effectiveVersion(version)
|
||||
logger.Info("starting MediaStationGo",
|
||||
logger.Info("starting MMTL",
|
||||
zap.String("version", appVersion),
|
||||
zap.Int("port", cfg.App.Port),
|
||||
zap.String("data_dir", cfg.App.DataDir),
|
||||
@@ -95,12 +94,6 @@ func main() {
|
||||
applyCPUThreadLimit(cfg, logger)
|
||||
services := service.NewWithVersion(cfg, logger, repos, appVersion)
|
||||
|
||||
if repaired, err := services.RepairCloudPathMetadata(context.Background()); err != nil {
|
||||
logger.Warn("cloud path metadata repair failed", zap.Error(err))
|
||||
} else if repaired > 0 {
|
||||
logger.Info("cloud path metadata repair completed", zap.Int("media_count", repaired))
|
||||
}
|
||||
|
||||
// 一次性清洗历史脏数据: 老版本把单集 episode id / 单集名写进整剧字段, 导致
|
||||
// 同一部剧被拆成多张单集卡。清空被污染的字段并重置为 pending(借后续重刮修正)。
|
||||
if cleaned, err := services.NormalizePollutedEpisodeMetadata(context.Background()); err != nil {
|
||||
@@ -143,8 +136,6 @@ func main() {
|
||||
}
|
||||
}()
|
||||
go services.Boot()
|
||||
go handler.RunLicenseHeartbeatLoop(services.Context(), services)
|
||||
go services.TelegramBot.StartPolling(context.Background())
|
||||
|
||||
// Graceful shutdown.
|
||||
stop := make(chan os.Signal, 1)
|
||||
@@ -158,5 +149,5 @@ func main() {
|
||||
logger.Error("graceful shutdown failed", zap.Error(err))
|
||||
}
|
||||
services.Close()
|
||||
logger.Info("MediaStationGo stopped")
|
||||
logger.Info("MMTL stopped")
|
||||
}
|
||||
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
@@ -54,106 +53,6 @@ func TestOpenSQLiteWithNilLoggerConfiguresPool(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnforceTelegramBindingOneToOneCleansDuplicatesAndAddsIndex(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.TelegramBinding{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
createdAt := time.Now().Add(-time.Hour)
|
||||
rows := []model.TelegramBinding{
|
||||
{TelegramUserID: 10001, ChatID: 10001, UserID: "user-1"},
|
||||
{TelegramUserID: 10002, ChatID: 10002, UserID: "user-1"},
|
||||
}
|
||||
for i := range rows {
|
||||
rows[i].CreatedAt = createdAt.Add(time.Duration(i) * time.Minute)
|
||||
if err := db.Create(&rows[i]).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := enforceTelegramBindingOneToOne(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := db.Model(&model.TelegramBinding{}).Where("user_id = ?", "user-1").Count(&count).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Fatalf("active bindings for user-1 = %d, want 1", count)
|
||||
}
|
||||
var kept model.TelegramBinding
|
||||
if err := db.First(&kept, "user_id = ?", "user-1").Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if kept.TelegramUserID != 10002 {
|
||||
t.Fatalf("kept telegram binding = %d, want newest 10002", kept.TelegramUserID)
|
||||
}
|
||||
if err := db.Create(&model.TelegramBinding{TelegramUserID: 10003, ChatID: 10003, UserID: "user-1"}).Error; err == nil {
|
||||
t.Fatal("expected unique index to reject another active binding for the same user")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureSubscriptionIdentityUniquenessArchivesDuplicatesAndAddsIndex(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.Subscription{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
createdAt := time.Now().Add(-time.Hour)
|
||||
rows := []model.Subscription{
|
||||
{UserID: "user-1", Name: "Example", FeedURL: "site-search://search?keyword=Example", Filter: "Example", Resolution: "1080p", Priority: 50},
|
||||
{UserID: "user-1", Name: "Example", FeedURL: "site-search://search?keyword=Example", Filter: "Example", Resolution: "1080p", Priority: 50},
|
||||
}
|
||||
for i := range rows {
|
||||
rows[i].CreatedAt = createdAt.Add(time.Duration(i) * time.Minute)
|
||||
if err := db.Select("*").Omit("DeletedAt").Create(&rows[i]).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := ensureSubscriptionIdentityUniqueness(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var active []model.Subscription
|
||||
if err := db.Where("archived_at IS NULL").Order("created_at asc").Find(&active).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(active) != 1 || active[0].ID != rows[0].ID {
|
||||
t.Fatalf("active subscriptions = %#v, want earliest row only", active)
|
||||
}
|
||||
var archived model.Subscription
|
||||
if err := db.First(&archived, "id = ?", rows[1].ID).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if archived.ArchivedAt == nil || archived.ArchiveReason != duplicateSubscriptionMigrationReason {
|
||||
t.Fatalf("duplicate was not archived by migration: %#v", archived)
|
||||
}
|
||||
|
||||
duplicate := rows[0]
|
||||
duplicate.ID = ""
|
||||
duplicate.CreatedAt = time.Time{}
|
||||
duplicate.UpdatedAt = time.Time{}
|
||||
model.RefreshSubscriptionIdentity(&duplicate)
|
||||
if err := db.Select("*").Omit("DeletedAt").Create(&duplicate).Error; err == nil {
|
||||
t.Fatal("expected active identity index to reject a duplicate rule")
|
||||
}
|
||||
different := rows[0]
|
||||
different.ID = ""
|
||||
different.CreatedAt = time.Time{}
|
||||
different.UpdatedAt = time.Time{}
|
||||
different.Resolution = "2160p"
|
||||
model.RefreshSubscriptionIdentity(&different)
|
||||
if err := db.Select("*").Omit("DeletedAt").Create(&different).Error; err != nil {
|
||||
t.Fatalf("different rule should be allowed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsurePerformanceIndexesCreatesHotPathIndexes(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
|
||||
@@ -14,12 +14,6 @@ func AutoMigrate(db *gorm.DB) error {
|
||||
if err := ensurePostgresColumnCompatibility(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := enforceTelegramBindingOneToOne(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureSubscriptionIdentityUniqueness(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensurePerformanceIndexes(db); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -44,7 +38,6 @@ func ensurePostgresColumnCompatibility(db *gorm.DB) error {
|
||||
`ALTER TABLE playback_histories ALTER COLUMN media_id TYPE varchar(128)`,
|
||||
`ALTER TABLE favorites ALTER COLUMN media_id TYPE varchar(128)`,
|
||||
`ALTER TABLE playlist_items ALTER COLUMN media_id TYPE varchar(128)`,
|
||||
`ALTER TABLE strm_records ALTER COLUMN media_id TYPE varchar(128)`,
|
||||
}
|
||||
for _, stmt := range statements {
|
||||
if err := db.Exec(stmt).Error; err != nil {
|
||||
@@ -84,39 +77,3 @@ func ensurePerformanceIndexes(db *gorm.DB) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func enforceTelegramBindingOneToOne(db *gorm.DB) error {
|
||||
if !db.Migrator().HasTable(&model.TelegramBinding{}) {
|
||||
return nil
|
||||
}
|
||||
return db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Exec(`
|
||||
DELETE FROM telegram_bindings
|
||||
WHERE deleted_at IS NULL
|
||||
AND user_id IN (
|
||||
SELECT user_id
|
||||
FROM telegram_bindings
|
||||
WHERE deleted_at IS NULL
|
||||
GROUP BY user_id
|
||||
HAVING COUNT(*) > 1
|
||||
)
|
||||
AND id NOT IN (
|
||||
SELECT id
|
||||
FROM (
|
||||
SELECT id,
|
||||
ROW_NUMBER() OVER (PARTITION BY user_id ORDER BY updated_at DESC, created_at DESC, id DESC) AS rn
|
||||
FROM telegram_bindings
|
||||
WHERE deleted_at IS NULL
|
||||
) AS ranked_bindings
|
||||
WHERE rn = 1
|
||||
)
|
||||
`).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Exec(`
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_telegram_bindings_user_id_active
|
||||
ON telegram_bindings(user_id)
|
||||
WHERE deleted_at IS NULL
|
||||
`).Error
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,54 +0,0 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
)
|
||||
|
||||
const duplicateSubscriptionMigrationReason = "迁移合并重复订阅规则"
|
||||
|
||||
func ensureSubscriptionIdentityUniqueness(db *gorm.DB) error {
|
||||
if db == nil || !db.Migrator().HasTable(&model.Subscription{}) {
|
||||
return nil
|
||||
}
|
||||
return db.Transaction(func(tx *gorm.DB) error {
|
||||
var rows []model.Subscription
|
||||
if err := tx.Unscoped().Order("created_at asc, id asc").Find(&rows).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
seen := make(map[string]string)
|
||||
for i := range rows {
|
||||
row := &rows[i]
|
||||
key := model.RefreshSubscriptionIdentity(row)
|
||||
updates := map[string]any{"identity_key": key}
|
||||
if !row.DeletedAt.Valid && row.ArchivedAt == nil {
|
||||
activeKey := row.UserID + "\x00" + key
|
||||
if _, duplicate := seen[activeKey]; duplicate {
|
||||
archivedAt := row.UpdatedAt
|
||||
if archivedAt.IsZero() {
|
||||
archivedAt = row.CreatedAt
|
||||
}
|
||||
if archivedAt.IsZero() {
|
||||
archivedAt = time.Now()
|
||||
}
|
||||
updates["enabled"] = false
|
||||
updates["archived_at"] = &archivedAt
|
||||
updates["archive_reason"] = duplicateSubscriptionMigrationReason
|
||||
} else {
|
||||
seen[activeKey] = row.ID
|
||||
}
|
||||
}
|
||||
if err := tx.Unscoped().Model(&model.Subscription{}).Where("id = ?", row.ID).Updates(updates).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return tx.Exec(`
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_subscriptions_user_identity_active
|
||||
ON subscriptions(user_id, identity_key)
|
||||
WHERE deleted_at IS NULL AND archived_at IS NULL AND identity_key <> ''
|
||||
`).Error
|
||||
})
|
||||
}
|
||||
@@ -68,8 +68,6 @@ func targetLooksLikeBootstrapOnly(target *gorm.DB) (bool, error) {
|
||||
&model.Favorite{},
|
||||
&model.Playlist{},
|
||||
&model.PlaylistItem{},
|
||||
&model.DownloadTask{},
|
||||
&model.Subscription{},
|
||||
} {
|
||||
if !target.Migrator().HasTable(m) {
|
||||
continue
|
||||
|
||||
@@ -43,7 +43,6 @@ func createUserHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
refreshLicenseCapacityBestEffort(c.Request.Context(), svc)
|
||||
u, _, err := svc.Auth.Register(c.Request.Context(), req.Username, req.Password)
|
||||
if err != nil {
|
||||
writeUserMutationError(c, svc, err)
|
||||
@@ -227,8 +226,7 @@ func writeUserMutationError(c *gin.Context, svc *service.Container, err error) {
|
||||
case errors.Is(err, service.ErrUsernameTaken):
|
||||
c.JSON(http.StatusConflict, gin.H{"error": "username already taken"})
|
||||
case errors.Is(err, service.ErrUserLimitReached):
|
||||
maxUsers := service.LicensedMaxUsers(c.Request.Context(), svc.Repo)
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "user limit reached", "max_users": maxUsers})
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "user limit reached", "max_users": service.UserLimit})
|
||||
default:
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
}
|
||||
|
||||
@@ -59,9 +59,6 @@ func updateSettingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
if req.Key == "transcode.hw_enabled" || req.Key == "transcode.hw_accel" || req.Key == "transcoder.hardware_accel" || req.Key == "transcoder.encoder" {
|
||||
svc.Transcoder.StopAll()
|
||||
}
|
||||
if req.Key == "cloud.auto_sync_enabled" && !service.ParseBoolSetting(req.Value, false) && svc.Scan != nil {
|
||||
_ = svc.Scan.CancelAllCloudScans()
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,79 +0,0 @@
|
||||
// Package handler — AI integration endpoints.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/middleware"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type smartSearchReq struct {
|
||||
Query string `json:"query" binding:"required"`
|
||||
}
|
||||
|
||||
func smartSearchHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req smartSearchReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
intent, err := svc.AI.SmartSearch(c.Request.Context(), req.Query)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
// Run the actual library search using the cleaned query so the
|
||||
// caller can render local + external results in one round-trip.
|
||||
items, _ := svc.Media.SearchMediaVisible(c.Request.Context(), intent.Query, 60, mediaVisibilityForRequest(c, svc))
|
||||
external := service.SearchExternalMedia(
|
||||
c.Request.Context(),
|
||||
intent.Query,
|
||||
intent.Year,
|
||||
intent.Type,
|
||||
svc.TMDb,
|
||||
svc.Douban,
|
||||
svc.Bangumi,
|
||||
)
|
||||
service.EnrichExternalMediaAvailability(c.Request.Context(), svc.Repo, external)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"intent": intent,
|
||||
"items": items,
|
||||
"external_items": external,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func aiRecommendHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
hist, err := svc.Playback.RecentHistory(c.Request.Context(), toString(uid), 10)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
titles := make([]string, 0, len(hist))
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
for _, h := range hist {
|
||||
if h.Media != nil && visibility.Allows(h.Media) && strings.TrimSpace(h.Media.Title) != "" {
|
||||
titles = append(titles, h.Media.Title)
|
||||
}
|
||||
}
|
||||
out, err := svc.AI.Recommend(c.Request.Context(), titles, 8)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"titles": out})
|
||||
}
|
||||
}
|
||||
|
||||
func aiStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, svc.AI.Status(c.Request.Context()))
|
||||
}
|
||||
}
|
||||
@@ -1,147 +0,0 @@
|
||||
// Package handler — multi-turn AI assistant chat endpoints.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/middleware"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func listAssistantSessionsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
rows, err := svc.Assistant.ListSessions(
|
||||
c.Request.Context(), toString(uid), role == "admin",
|
||||
)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, rows)
|
||||
}
|
||||
}
|
||||
|
||||
type createSessionReq struct {
|
||||
Title string `json:"title"`
|
||||
}
|
||||
|
||||
func createAssistantSessionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req createSessionReq
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
sess, err := svc.Assistant.CreateSession(c.Request.Context(), toString(uid), req.Title)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, sess)
|
||||
}
|
||||
}
|
||||
|
||||
func getAssistantSessionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
view, err := svc.Assistant.GetSession(
|
||||
c.Request.Context(), c.Param("id"), toString(uid), role == "admin",
|
||||
)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, view)
|
||||
}
|
||||
}
|
||||
|
||||
func deleteAssistantSessionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
if err := svc.Assistant.DeleteSession(
|
||||
c.Request.Context(), c.Param("id"), toString(uid), role == "admin",
|
||||
); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
type chatReq struct {
|
||||
SessionID string `json:"session_id" binding:"required"`
|
||||
Message string `json:"message" binding:"required"`
|
||||
}
|
||||
|
||||
func assistantChatHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req chatReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
view, err := svc.Assistant.Chat(
|
||||
c.Request.Context(), req.SessionID, toString(uid), req.Message, role == "admin",
|
||||
)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, view)
|
||||
}
|
||||
}
|
||||
|
||||
type executeReq struct {
|
||||
SessionID string `json:"session_id" binding:"required"`
|
||||
Action map[string]interface{} `json:"action" binding:"required"`
|
||||
}
|
||||
|
||||
func assistantExecuteHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req executeReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
opID, err := svc.Assistant.Execute(
|
||||
c.Request.Context(), req.SessionID, toString(uid), req.Action,
|
||||
)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"op_id": opID})
|
||||
}
|
||||
}
|
||||
|
||||
func assistantUndoHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Assistant.Undo(c.Request.Context(), c.Param("op_id")); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
|
||||
func assistantHistoryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
rows, err := svc.Assistant.History(
|
||||
c.Request.Context(), toString(uid), role == "admin",
|
||||
)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": rows})
|
||||
}
|
||||
}
|
||||
@@ -62,7 +62,6 @@ func registerHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
refreshLicenseCapacityBestEffort(c.Request.Context(), svc)
|
||||
u, tokens, err := svc.Auth.Register(c.Request.Context(), req.Username, req.Password)
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrUsernameTaken) {
|
||||
|
||||
@@ -1,160 +0,0 @@
|
||||
// Package handler — cloud-disk (网盘) endpoints: directory browsing, QR-code
|
||||
// login, media import and 302 playback redirects.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// cloudListHandler browses a configured cloud disk directory.
|
||||
func cloudListHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
typ := c.Param("type")
|
||||
if !service.IsAdminCloudConfigurable(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider", "items": []any{}})
|
||||
return
|
||||
}
|
||||
dir := c.Query("dir")
|
||||
entries, err := svc.StorageCfg.CloudList(c.Request.Context(), typ, dir)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"error": err.Error(), "items": []any{}})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": entries})
|
||||
}
|
||||
}
|
||||
|
||||
func cloudMkdirHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
typ := c.Param("type")
|
||||
if !service.IsAdminCloudConfigurable(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||
return
|
||||
}
|
||||
var in struct {
|
||||
Dir string `json:"dir"`
|
||||
Name string `json:"name" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
entry, err := svc.StorageCfg.CloudMkdir(c.Request.Context(), typ, in.Dir, in.Name)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"entry": entry})
|
||||
}
|
||||
}
|
||||
|
||||
func cloudRenameHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
typ := c.Param("type")
|
||||
if !service.IsAdminCloudConfigurable(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||
return
|
||||
}
|
||||
var in struct {
|
||||
Ref string `json:"ref" binding:"required"`
|
||||
Name string `json:"name" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
entry, err := svc.StorageCfg.CloudRename(c.Request.Context(), typ, in.Ref, in.Name)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"entry": entry})
|
||||
}
|
||||
}
|
||||
|
||||
// cloudImportHandler turns a cloud file into a playable 302-backed media item.
|
||||
func cloudImportHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
typ := c.Param("type")
|
||||
if !service.IsAdminCloudConfigurable(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||
return
|
||||
}
|
||||
var in struct {
|
||||
Ref string `json:"ref" binding:"required"`
|
||||
Name string `json:"name"`
|
||||
Size int64 `json:"size"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
m, err := svc.StorageCfg.CloudImport(c.Request.Context(), typ, in.Ref, in.Name, in.Size)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, m)
|
||||
}
|
||||
}
|
||||
|
||||
func cloudScanAllHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc.Scan == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "scanner unavailable"})
|
||||
return
|
||||
}
|
||||
statuses, err := svc.Scan.StartAllCloudLibraryScans()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusAccepted, gin.H{
|
||||
"items": statuses,
|
||||
"scan_queued": true,
|
||||
"message": "已开始扫描所有启用的网盘媒体库",
|
||||
"resume_message": "中断后再次点击扫描会重新遍历,但已入库媒体会去重更新,只补齐缺失项。",
|
||||
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func cloudScanCancelHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc.Scan == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "scanner unavailable"})
|
||||
return
|
||||
}
|
||||
libraryID := strings.TrimSpace(c.Query("library_id"))
|
||||
provider := strings.TrimSpace(c.Query("provider"))
|
||||
cancelled := 0
|
||||
if libraryID != "" {
|
||||
if svc.Scan.CancelCloudScan(libraryID) {
|
||||
cancelled = 1
|
||||
}
|
||||
} else if provider != "" {
|
||||
cancelled = svc.Scan.CancelCloudScansForProvider(provider)
|
||||
} else {
|
||||
cancelled = svc.Scan.CancelAllCloudScans()
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"cancelled": cancelled,
|
||||
"message": "已发送中断信号;正在等待当前网盘请求返回后停止",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func cloudScanStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if svc.Scan == nil {
|
||||
c.JSON(http.StatusOK, gin.H{"items": []service.CloudScanStatus{}})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": svc.Scan.CloudScanStatuses()})
|
||||
}
|
||||
}
|
||||
@@ -1,149 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
|
||||
)
|
||||
|
||||
// cloudMountHandler creates or reuses a cloud:// media library for a cloud
|
||||
// directory, then queues a recursive import scan. The scan runs outside the
|
||||
// request so large 115/OpenList folders do not make the UI report a timeout.
|
||||
func cloudMountHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
typ := c.Param("type")
|
||||
if !service.IsAdminCloudConfigurable(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||
return
|
||||
}
|
||||
var in struct {
|
||||
Dir string `json:"dir"`
|
||||
DirPath string `json:"dir_path"`
|
||||
Name string `json:"name"`
|
||||
MediaType string `json:"media_type"`
|
||||
}
|
||||
_ = c.ShouldBindJSON(&in)
|
||||
if !cloud.IsCloudType(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||
return
|
||||
}
|
||||
if _, err := svc.StorageCfg.CloudProvider(c.Request.Context(), typ); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
path := service.BuildCloudLibraryPath(typ, in.Dir, in.DirPath)
|
||||
name := strings.TrimSpace(in.Name)
|
||||
if name == "" {
|
||||
name = cloudMountLibraryName(typ, strings.TrimSpace(in.Dir), strings.TrimSpace(in.DirPath))
|
||||
}
|
||||
mediaType := strings.TrimSpace(in.MediaType)
|
||||
if mediaType == "" || strings.EqualFold(mediaType, "auto") {
|
||||
displayDir := strings.TrimSpace(in.DirPath)
|
||||
if displayDir == "" {
|
||||
displayDir = strings.TrimSpace(in.Dir)
|
||||
}
|
||||
mediaType = service.InferCloudMountMediaType(displayDir, name)
|
||||
}
|
||||
libs, err := svc.Repo.Library.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
var lib *model.Library
|
||||
alreadyMounted := false
|
||||
if conflict := service.FindCloudMountConflict(libs, typ, in.Dir, in.DirPath); conflict != nil {
|
||||
lib = &conflict.Library
|
||||
alreadyMounted = conflict.Exact
|
||||
if conflict.Nested {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"library": lib,
|
||||
"skipped": true,
|
||||
"reason": "cloud mount overlaps an existing mounted parent/child directory",
|
||||
"conflict_library": conflict.Library,
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
if lib == nil {
|
||||
lib = &model.Library{Name: name, Path: path, Type: mediaType, Enabled: true}
|
||||
if err := svc.Repo.Library.Create(c.Request.Context(), lib); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
} else if alreadyMounted {
|
||||
updates := map[string]any{}
|
||||
if path != "" && path != lib.Path {
|
||||
updates["path"] = path
|
||||
lib.Path = path
|
||||
}
|
||||
if mediaType != "" && mediaType != lib.Type {
|
||||
updates["type"] = mediaType
|
||||
lib.Type = mediaType
|
||||
}
|
||||
currentDisplayName, _ := service.CloudLibraryDisplayName(*lib)
|
||||
if name != "" && name != lib.Name && (currentDisplayName == "" || currentDisplayName != name || strings.Contains(lib.Name, " · ")) {
|
||||
updates["name"] = name
|
||||
lib.Name = name
|
||||
}
|
||||
if len(updates) > 0 {
|
||||
if err := svc.Repo.DB.WithContext(c.Request.Context()).Model(&model.Library{}).Where("id = ?", lib.ID).Updates(updates).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
if svc.Scan != nil {
|
||||
libID := lib.ID
|
||||
if svc.WSHub != nil {
|
||||
svc.WSHub.Publish("scan", gin.H{
|
||||
"library_id": libID,
|
||||
"cloud": true,
|
||||
"queued": true,
|
||||
"stage": "queued",
|
||||
"message": "云盘扫描已加入后台队列,会递归扫描并自动加入媒体库",
|
||||
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
||||
})
|
||||
}
|
||||
_, _, _ = svc.Scan.StartCloudLibraryScan(libID, false)
|
||||
}
|
||||
c.JSON(http.StatusAccepted, gin.H{
|
||||
"library": lib,
|
||||
"already_mounted": alreadyMounted,
|
||||
"scan_queued": svc.Scan != nil,
|
||||
"message": "挂载后会后台递归扫描,发现的媒体会自动加入当前媒体库",
|
||||
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func cloudMountLibraryName(typ, dir, displayDir string) string {
|
||||
base := service.CloudMountProviderLabel(typ)
|
||||
displayDir = strings.Trim(strings.TrimSpace(strings.ReplaceAll(displayDir, "\\", "/")), "/")
|
||||
if displayDir != "" {
|
||||
parts := strings.Split(displayDir, "/")
|
||||
for i := len(parts) - 1; i >= 0; i-- {
|
||||
if part := strings.TrimSpace(parts[i]); part != "" {
|
||||
return part
|
||||
}
|
||||
}
|
||||
}
|
||||
if dir == "" || dir == "0" {
|
||||
return base
|
||||
}
|
||||
dir = strings.Trim(strings.TrimSpace(strings.ReplaceAll(dir, "\\", "/")), "/")
|
||||
if dir == "" {
|
||||
return base
|
||||
}
|
||||
parts := strings.Split(dir, "/")
|
||||
for i := len(parts) - 1; i >= 0; i-- {
|
||||
if part := strings.TrimSpace(parts[i]); part != "" {
|
||||
return part
|
||||
}
|
||||
}
|
||||
return base
|
||||
}
|
||||
@@ -1,284 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"path"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
|
||||
)
|
||||
|
||||
type cloudPlaybackRequest struct {
|
||||
svc *service.Container
|
||||
c *gin.Context
|
||||
typ string
|
||||
ref string
|
||||
link *cloud.DirectLink
|
||||
resolveStart time.Time
|
||||
resolveDur time.Duration
|
||||
}
|
||||
|
||||
// cloudPlayHandler resolves a cloud file to its direct link and either issues a
|
||||
// 302 redirect (true offload — host does not stream the bytes) or, when the
|
||||
// provider requires authenticated headers, reverse-proxies the response.
|
||||
func cloudPlayHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
typ := c.Param("type")
|
||||
ref := c.Query("ref")
|
||||
if !service.IsAdminCloudConfigurable(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||
return
|
||||
}
|
||||
if ref == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "ref required"})
|
||||
return
|
||||
}
|
||||
if !enforceScopedCloudPlaybackToken(c, svc, typ, ref) {
|
||||
return
|
||||
}
|
||||
serveCloudResolvedLink(svc, c, typ, ref)
|
||||
}
|
||||
}
|
||||
|
||||
func serveCloudResolvedLink(svc *service.Container, c *gin.Context, typ, ref string) {
|
||||
if isCloudImageRef(ref) && svc != nil && svc.ImageProxy != nil {
|
||||
if svc.ImageProxy.ServeCloudCached(c.Writer, c.Request, typ+":"+ref) {
|
||||
return
|
||||
}
|
||||
}
|
||||
if svc == nil || svc.StorageCfg == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "cloud storage service unavailable"})
|
||||
return
|
||||
}
|
||||
resolveStart := time.Now()
|
||||
link, err := svc.StorageCfg.CloudResolve(c.Request.Context(), typ, ref, c.Request.UserAgent())
|
||||
resolveDur := time.Since(resolveStart)
|
||||
if err != nil {
|
||||
logCloudPlayback(svc, "cloud playback resolve failed",
|
||||
append(cloudPlaybackLogFields(typ, ref, nil, resolveDur), zap.Error(err))...)
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if isCloudImageRef(ref) && svc.ImageProxy != nil {
|
||||
if err := svc.ImageProxy.ServeCloudResolved(c.Request.Context(), c.Writer, c.Request, typ+":"+ref, link); err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
}
|
||||
return
|
||||
}
|
||||
if isCloudImageRef(ref) {
|
||||
c.Header("Cache-Control", "public, max-age=2592000, immutable")
|
||||
}
|
||||
if !link.Proxy {
|
||||
// Pure offload: send the client straight to the cloud CDN.
|
||||
setRedirectNoStoreHeaders(c)
|
||||
logCloudPlayback(svc, "cloud playback redirect",
|
||||
append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
|
||||
zap.String("mode", "redirect"),
|
||||
zap.Int("status", http.StatusFound),
|
||||
zap.String("method", c.Request.Method),
|
||||
zap.String("range", c.GetHeader("Range")),
|
||||
)...)
|
||||
c.Redirect(http.StatusFound, link.URL)
|
||||
return
|
||||
}
|
||||
proxyCloudResolvedLink(cloudPlaybackRequest{
|
||||
svc: svc,
|
||||
c: c,
|
||||
typ: typ,
|
||||
ref: ref,
|
||||
link: link,
|
||||
resolveStart: resolveStart,
|
||||
resolveDur: resolveDur,
|
||||
})
|
||||
}
|
||||
|
||||
func proxyCloudResolvedLink(playback cloudPlaybackRequest) {
|
||||
c := playback.c
|
||||
clientMethod := playback.c.Request.Method
|
||||
if clientMethod == "" {
|
||||
clientMethod = http.MethodGet
|
||||
}
|
||||
upstreamMethod := clientMethod
|
||||
req, err := http.NewRequestWithContext(c.Request.Context(), upstreamMethod, playback.link.URL, nil)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
for k, v := range playback.link.Headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
if rng := c.GetHeader("Range"); rng != "" {
|
||||
req.Header.Set("Range", rng)
|
||||
}
|
||||
if accept := c.GetHeader("Accept"); accept != "" {
|
||||
req.Header.Set("Accept", accept)
|
||||
}
|
||||
if c.GetHeader("Accept-Encoding") == "" {
|
||||
req.Header.Set("Accept-Encoding", "identity")
|
||||
}
|
||||
upstreamStart := time.Now()
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
upstreamHeaderDur := time.Since(upstreamStart)
|
||||
if err != nil {
|
||||
logCloudPlayback(playback.svc, "cloud playback proxy upstream failed",
|
||||
append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
|
||||
zap.String("mode", "proxy"),
|
||||
zap.String("method", clientMethod),
|
||||
zap.String("upstream_method", upstreamMethod),
|
||||
zap.String("range", c.GetHeader("Range")),
|
||||
zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
|
||||
zap.Error(err),
|
||||
)...)
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
for _, h := range []string{"Content-Type", "Content-Length", "Content-Range", "Accept-Ranges", "ETag", "Last-Modified"} {
|
||||
if v := resp.Header.Get(h); v != "" {
|
||||
c.Header(h, v)
|
||||
}
|
||||
}
|
||||
if c.Writer.Header().Get("Accept-Ranges") == "" {
|
||||
c.Header("Accept-Ranges", "bytes")
|
||||
}
|
||||
if resp.StatusCode >= 400 {
|
||||
handleCloudProxyError(playback, req, resp, clientMethod, upstreamMethod, upstreamHeaderDur)
|
||||
return
|
||||
}
|
||||
streamCloudProxyResponse(playback, req, resp, clientMethod, upstreamMethod, upstreamHeaderDur)
|
||||
}
|
||||
|
||||
func handleCloudProxyError(playback cloudPlaybackRequest, req *http.Request, resp *http.Response, clientMethod, upstreamMethod string, upstreamHeaderDur time.Duration) {
|
||||
c := playback.c
|
||||
c.Header("Cache-Control", "no-store")
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
fields := append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
|
||||
zap.String("mode", "proxy"),
|
||||
zap.String("method", clientMethod),
|
||||
zap.String("upstream_method", upstreamMethod),
|
||||
zap.String("range", c.GetHeader("Range")),
|
||||
zap.String("upstream_range", req.Header.Get("Range")),
|
||||
zap.Int("status", resp.StatusCode),
|
||||
zap.String("content_range", resp.Header.Get("Content-Range")),
|
||||
zap.String("content_length", resp.Header.Get("Content-Length")),
|
||||
zap.String("upstream_error_body", strings.TrimSpace(string(body))),
|
||||
zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
|
||||
zap.Int64("total_ms", durationMilliseconds(time.Since(playback.resolveStart))),
|
||||
)
|
||||
logCloudPlayback(playback.svc, "cloud playback proxy upstream returned error", fields...)
|
||||
c.Status(resp.StatusCode)
|
||||
if clientMethod != http.MethodHead && len(body) > 0 {
|
||||
_, _ = c.Writer.Write(body)
|
||||
}
|
||||
}
|
||||
|
||||
func streamCloudProxyResponse(playback cloudPlaybackRequest, req *http.Request, resp *http.Response, clientMethod, upstreamMethod string, upstreamHeaderDur time.Duration) {
|
||||
c := playback.c
|
||||
c.Status(resp.StatusCode)
|
||||
var copied int64
|
||||
var copyErr error
|
||||
streamStart := time.Now()
|
||||
if c.Request.Method != http.MethodHead {
|
||||
copied, copyErr = io.Copy(c.Writer, resp.Body)
|
||||
}
|
||||
fields := append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
|
||||
zap.String("mode", "proxy"),
|
||||
zap.String("method", clientMethod),
|
||||
zap.String("upstream_method", upstreamMethod),
|
||||
zap.String("range", c.GetHeader("Range")),
|
||||
zap.String("upstream_range", req.Header.Get("Range")),
|
||||
zap.Int("status", resp.StatusCode),
|
||||
zap.String("content_range", resp.Header.Get("Content-Range")),
|
||||
zap.String("content_length", resp.Header.Get("Content-Length")),
|
||||
zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
|
||||
zap.Int64("stream_ms", durationMilliseconds(time.Since(streamStart))),
|
||||
zap.Int64("total_ms", durationMilliseconds(time.Since(playback.resolveStart))),
|
||||
zap.Int64("bytes", copied),
|
||||
)
|
||||
if copyErr != nil {
|
||||
logCloudPlayback(playback.svc, "cloud playback proxy copy failed", append(fields, zap.Error(copyErr))...)
|
||||
return
|
||||
}
|
||||
logCloudPlayback(playback.svc, "cloud playback proxy finished", fields...)
|
||||
}
|
||||
|
||||
func isCloudImageRef(ref string) bool {
|
||||
ref = strings.ToLower(strings.TrimSpace(ref))
|
||||
for _, suffix := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp", ".tbn"} {
|
||||
if strings.HasSuffix(ref, suffix) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func logCloudPlayback(svc *service.Container, msg string, fields ...zap.Field) {
|
||||
if svc == nil || svc.Log == nil {
|
||||
return
|
||||
}
|
||||
svc.Log.Info(msg, fields...)
|
||||
}
|
||||
|
||||
func cloudPlaybackLogFields(typ, ref string, link *cloud.DirectLink, resolveDur time.Duration) []zap.Field {
|
||||
refHash, refExt := cloudPlaybackRefFingerprint(ref)
|
||||
fields := []zap.Field{
|
||||
zap.String("provider", strings.TrimSpace(typ)),
|
||||
zap.String("ref_hash", refHash),
|
||||
zap.String("ref_ext", refExt),
|
||||
zap.Int64("resolve_ms", durationMilliseconds(resolveDur)),
|
||||
}
|
||||
if link != nil {
|
||||
fields = append(fields,
|
||||
zap.String("target_host", cloudPlaybackLinkHost(link.URL)),
|
||||
zap.Bool("headers_required", len(link.Headers) > 0),
|
||||
zap.Strings("header_names", cloudPlaybackHeaderNames(link.Headers)),
|
||||
)
|
||||
}
|
||||
return fields
|
||||
}
|
||||
|
||||
func cloudPlaybackRefFingerprint(ref string) (string, string) {
|
||||
ref = strings.TrimSpace(ref)
|
||||
sum := sha256.Sum256([]byte(ref))
|
||||
ext := strings.ToLower(path.Ext(strings.Trim(strings.ReplaceAll(ref, "\\", "/"), "/")))
|
||||
return hex.EncodeToString(sum[:])[:12], ext
|
||||
}
|
||||
|
||||
func cloudPlaybackLinkHost(raw string) string {
|
||||
u, err := url.Parse(strings.TrimSpace(raw))
|
||||
if err != nil || u.Host == "" {
|
||||
return ""
|
||||
}
|
||||
return u.Host
|
||||
}
|
||||
|
||||
func cloudPlaybackHeaderNames(headers map[string]string) []string {
|
||||
if len(headers) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]string, 0, len(headers))
|
||||
for key := range headers {
|
||||
if key = strings.TrimSpace(key); key != "" {
|
||||
out = append(out, key)
|
||||
}
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func durationMilliseconds(d time.Duration) int64 {
|
||||
if d <= 0 {
|
||||
return 0
|
||||
}
|
||||
return d.Milliseconds()
|
||||
}
|
||||
@@ -1,55 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
|
||||
)
|
||||
|
||||
func TestProxyCloudResolvedLinkUsesHEADWithoutSyntheticRange(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
var upstreamMethod, upstreamRange string
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
upstreamMethod = r.Method
|
||||
upstreamRange = r.Header.Get("Range")
|
||||
w.Header().Set("Content-Type", "video/mp4")
|
||||
w.Header().Set("Content-Length", "123456")
|
||||
w.Header().Set("Accept-Ranges", "bytes")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodHead, "/api/cloud/play/openlist?ref=movie", nil)
|
||||
|
||||
proxyCloudResolvedLink(cloudPlaybackRequest{
|
||||
c: c,
|
||||
typ: "openlist",
|
||||
ref: "movie",
|
||||
link: &cloud.DirectLink{
|
||||
URL: upstream.URL + "/movie.mp4",
|
||||
Proxy: true,
|
||||
},
|
||||
})
|
||||
|
||||
if upstreamMethod != http.MethodHead {
|
||||
t.Fatalf("upstream method = %q, want HEAD", upstreamMethod)
|
||||
}
|
||||
if upstreamRange != "" {
|
||||
t.Fatalf("upstream Range = %q, want empty", upstreamRange)
|
||||
}
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rec.Code)
|
||||
}
|
||||
if got := rec.Header().Get("Content-Length"); got != "123456" {
|
||||
t.Fatalf("Content-Length = %q, want full upstream length", got)
|
||||
}
|
||||
if rec.Body.Len() != 0 {
|
||||
t.Fatalf("HEAD response body length = %d, want 0", rec.Body.Len())
|
||||
}
|
||||
}
|
||||
@@ -1,49 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
|
||||
)
|
||||
|
||||
// cloud115QRStartHandler begins a 115 QR-code login and returns the session +
|
||||
// QR image URL for the frontend to render.
|
||||
func cloud115QRStartHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if c.Param("type") != cloud.Type115 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "qr login is only supported for 115"})
|
||||
return
|
||||
}
|
||||
sess, err := cloud.QRStart(c.Request.Context(), nil)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, sess)
|
||||
}
|
||||
}
|
||||
|
||||
// cloud115QRPollHandler polls a 115 QR session; on confirmation it returns the
|
||||
// session cookie so the frontend can save it as the storage credential.
|
||||
func cloud115QRPollHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if c.Param("type") != cloud.Type115 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "qr login is only supported for 115"})
|
||||
return
|
||||
}
|
||||
var sess cloud.QRSession
|
||||
if err := c.ShouldBindJSON(&sess); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
st, err := cloud.QRPoll(c.Request.Context(), nil, &sess)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, st)
|
||||
}
|
||||
}
|
||||
@@ -1,157 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/config"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
|
||||
)
|
||||
|
||||
var handlerTestJPEG = []byte{0xff, 0xd8, 0xff, 0xe0, 0x00, 0x10, 'J', 'F', 'I', 'F', 0x00, 0xff, 0xd9}
|
||||
|
||||
func TestCloudMountLibraryNameDefaultsToDirectoryBaseName(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
provider string
|
||||
dir string
|
||||
displayDir string
|
||||
want string
|
||||
}{
|
||||
{name: "openlist directory", provider: "openlist", dir: "/国产剧", displayDir: "/国产剧", want: "国产剧"},
|
||||
{name: "nested directory", provider: "openlist", dir: "id-123", displayDir: "剧集/国产剧", want: "国产剧"},
|
||||
{name: "provider root", provider: "openlist", dir: "", displayDir: "", want: "OpenList"},
|
||||
{name: "115 root id", provider: "cloud115", dir: "0", displayDir: "", want: "115 网盘"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := cloudMountLibraryName(tt.provider, tt.dir, tt.displayDir); got != tt.want {
|
||||
t.Fatalf("cloudMountLibraryName() = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudPlaybackDiagnosticsDoNotExposeRawRefOrURL(t *testing.T) {
|
||||
rawRef := "/剧集/国产剧/很长的敏感文件名.S01E01.mkv"
|
||||
refHash, refExt := cloudPlaybackRefFingerprint(rawRef)
|
||||
if refHash == "" || strings.Contains(rawRef, refHash) {
|
||||
t.Fatalf("ref hash should be a short fingerprint, got %q", refHash)
|
||||
}
|
||||
if refExt != ".mkv" {
|
||||
t.Fatalf("ref ext = %q, want .mkv", refExt)
|
||||
}
|
||||
if host := cloudPlaybackLinkHost("https://cdn.example.test/movie.mkv?token=secret"); host != "cdn.example.test" {
|
||||
t.Fatalf("host = %q, want cdn.example.test", host)
|
||||
}
|
||||
names := cloudPlaybackHeaderNames(map[string]string{
|
||||
"Authorization": "Bearer secret",
|
||||
"Cookie": "sid=secret",
|
||||
})
|
||||
if got := strings.Join(names, ","); got != "Authorization,Cookie" {
|
||||
t.Fatalf("header names = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAdminCloudHandlersRejectQuarkBrowsing(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
router.GET("/admin/cloud/:type/list", cloudListHandler(nil))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/admin/cloud/quark/list?dir=0", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
|
||||
}
|
||||
if !strings.Contains(w.Body.String(), "unsupported cloud provider") {
|
||||
t.Fatalf("body = %s, want unsupported cloud provider", w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudPlayRejectsQuarkProvider(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
router.GET("/api/cloud/play/:type", cloudPlayHandler(nil))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/cloud/play/quark?ref=file-1", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
|
||||
}
|
||||
if !strings.Contains(w.Body.String(), "unsupported cloud provider") {
|
||||
t.Fatalf("body = %s, want unsupported cloud provider", w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudArtworkProxyServesCachedImageWithoutCloudResolve(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "image/jpeg")
|
||||
_, _ = w.Write(handlerTestJPEG)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
imageProxy := service.NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, zap.NewNop())
|
||||
stableKey := "openlist:/Anime/JianLai/poster.jpg"
|
||||
if err := imageProxy.PrefetchCloudResolved(t.Context(), stableKey, &cloud.DirectLink{URL: upstream.URL + "/poster.jpg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.GET("/api/img/cloud/:type", cloudArtworkProxyHandler(&service.Container{ImageProxy: imageProxy}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/img/cloud/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s, want 200", w.Code, w.Body.String())
|
||||
}
|
||||
if got := w.Body.Bytes(); !bytes.Equal(got, handlerTestJPEG) {
|
||||
t.Fatalf("body = %q, want cached poster", got)
|
||||
}
|
||||
if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "max-age=2592000") {
|
||||
t.Fatalf("cache-control = %q, want long static cache", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloudArtworkProxyAcceptsCachedTBNImage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "image/jpeg")
|
||||
_, _ = w.Write(handlerTestJPEG)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
imageProxy := service.NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, zap.NewNop())
|
||||
stableKey := "openlist:/Movies/Movie.tbn"
|
||||
if err := imageProxy.PrefetchCloudResolved(t.Context(), stableKey, &cloud.DirectLink{URL: upstream.URL + "/Movie.tbn"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.GET("/api/img/cloud/:type", cloudArtworkProxyHandler(&service.Container{ImageProxy: imageProxy}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/img/cloud/openlist?ref=%2FMovies%2FMovie.tbn", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s, want 200", w.Code, w.Body.String())
|
||||
}
|
||||
if got := w.Body.Bytes(); !bytes.Equal(got, handlerTestJPEG) {
|
||||
t.Fatalf("body = %q, want cached tbn poster", got)
|
||||
}
|
||||
}
|
||||
@@ -1,48 +0,0 @@
|
||||
// Package handler — TMDb discovery endpoints.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// trendingHandler 返回 TMDb 当日热门列表。
|
||||
//
|
||||
// 当本机无法连接 TMDb(GFW / 代理未配 / API key 无效)时,TMDb 调用会
|
||||
// 在 15 秒后超时;这种情况下不应该让首页显示 500 错误,而是把空列表
|
||||
// 直接返回——前端按 items.length === 0 渲染"暂无推荐"即可。
|
||||
func trendingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
items, err := svc.Discover.Trending(c.Request.Context())
|
||||
if err != nil {
|
||||
svc.Log.Warn("discover trending failed (returning empty list)", zap.Error(err))
|
||||
c.JSON(http.StatusOK, gin.H{"items": []service.Match{}, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
if items == nil {
|
||||
items = []service.Match{}
|
||||
}
|
||||
svc.Discover.WarmMatchArtwork(items)
|
||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||
}
|
||||
}
|
||||
|
||||
func popularHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
items, err := svc.Discover.Popular(c.Request.Context())
|
||||
if err != nil {
|
||||
svc.Log.Warn("discover popular failed (returning empty list)", zap.Error(err))
|
||||
c.JSON(http.StatusOK, gin.H{"items": []service.Match{}, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
if items == nil {
|
||||
items = []service.Match{}
|
||||
}
|
||||
svc.Discover.WarmMatchArtwork(items)
|
||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||
}
|
||||
}
|
||||
@@ -1,338 +0,0 @@
|
||||
// Package handler — multi-section discover endpoints.
|
||||
//
|
||||
// The Vue DiscoverView paginates a configurable list of "sections"
|
||||
// (trending day/week, popular movies, top rated, etc.) and asks the
|
||||
// backend for a feed keyed by section name. We mirror that surface so
|
||||
// the React DiscoverPage can render the same rails without a rewrite.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type discoverSectionDef struct {
|
||||
Key string
|
||||
Label string
|
||||
Provider string
|
||||
}
|
||||
|
||||
var discoverSectionCatalog = []discoverSectionDef{
|
||||
{Key: "tmdb_trending_day", Label: "TMDb 今日趋势", Provider: "tmdb"},
|
||||
{Key: "tmdb_trending_week", Label: "TMDb 本周热门", Provider: "tmdb"},
|
||||
{Key: "tmdb_latest_movie", Label: "TMDb 最新电影", Provider: "tmdb"},
|
||||
{Key: "tmdb_latest_tv", Label: "TMDb 最新剧集", Provider: "tmdb"},
|
||||
{Key: "tmdb_popular_movie", Label: "TMDb 热门电影", Provider: "tmdb"},
|
||||
{Key: "tmdb_popular_tv", Label: "TMDb 热门剧集", Provider: "tmdb"},
|
||||
{Key: "tmdb_top_rated_movie", Label: "TMDb 高分电影", Provider: "tmdb"},
|
||||
{Key: "tmdb_upcoming_movie", Label: "TMDb 即将上映", Provider: "tmdb"},
|
||||
{Key: "douban_hot_movie", Label: "豆瓣热门电影", Provider: "douban"},
|
||||
{Key: "douban_hot_tv", Label: "豆瓣热门剧集", Provider: "douban"},
|
||||
{Key: "douban_top_movie", Label: "豆瓣高分电影", Provider: "douban"},
|
||||
{Key: "bangumi_calendar", Label: "Bangumi 每日放送", Provider: "bangumi"},
|
||||
}
|
||||
|
||||
const discoverFeedSectionTimeout = 20 * time.Second
|
||||
const discoverFeedBangumiTimeout = 30 * time.Second
|
||||
const discoverFeedSlowSectionThreshold = 2 * time.Second
|
||||
|
||||
// discoverSectionsHandler returns the catalog of sections the UI can
|
||||
// pick from. The names match the upstream Vue UI so existing settings
|
||||
// keep working.
|
||||
func discoverSectionsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
sections := make([]gin.H, 0, len(discoverSectionCatalog))
|
||||
for _, section := range enabledDiscoverSections(c.Request.Context(), svc) {
|
||||
sections = append(sections, gin.H{"key": section.Key, "label": section.Label, "provider": section.Provider})
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"sections": sections})
|
||||
}
|
||||
}
|
||||
|
||||
// discoverFeedHandler resolves one or more section keys (?sections=a,b)
|
||||
// to TMDb / Douban / Bangumi rails and returns the joined results keyed by
|
||||
// section name. Unknown keys are silently dropped so URL typos don't break
|
||||
// the page.
|
||||
func discoverFeedHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
rawSections := c.Query("sections")
|
||||
if strings.TrimSpace(rawSections) == "" {
|
||||
rawSections = strings.Join(defaultDiscoverSectionKeys(c.Request.Context(), svc), ",")
|
||||
}
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
keys := strings.Split(rawSections, ",")
|
||||
out := gin.H{}
|
||||
meta := gin.H{}
|
||||
artworkItems := []service.ExternalMediaResult{}
|
||||
for _, raw := range keys {
|
||||
k := strings.TrimSpace(raw)
|
||||
if k == "" {
|
||||
continue
|
||||
}
|
||||
if provider := discoverSectionProvider(k); provider != "" && !discoverProviderEnabled(c.Request.Context(), svc, provider) {
|
||||
out[k] = []service.ExternalMediaResult{}
|
||||
meta[k] = gin.H{"page": page, "has_next": false, "disabled": true}
|
||||
continue
|
||||
}
|
||||
sectionTimeout := discoverSectionTimeout(k)
|
||||
sectionCtx, cancel := context.WithTimeout(c.Request.Context(), sectionTimeout)
|
||||
started := time.Now()
|
||||
items, err := discoverSectionItems(sectionCtx, svc, k, page)
|
||||
elapsed := time.Since(started)
|
||||
cancel()
|
||||
metaEntry := gin.H{"page": page, "has_next": false, "duration_ms": elapsed.Milliseconds()}
|
||||
if err != nil {
|
||||
logDiscoverFetchFailed(svc, k, page, elapsed, sectionTimeout, err)
|
||||
if cached, ok := cachedDiscoverSection(svc, k, page); ok {
|
||||
items = cached
|
||||
metaEntry["stale"] = true
|
||||
metaEntry["warning"] = discoverFeedStaleMessage(err)
|
||||
} else if fallbackItems, fallbackKey, ok := fallbackDiscoverSectionItems(c.Request.Context(), svc, k, page); ok {
|
||||
items = fallbackItems
|
||||
metaEntry["fallback"] = fallbackKey
|
||||
metaEntry["warning"] = discoverFeedFallbackMessage(fallbackKey, err)
|
||||
rememberDiscoverSection(svc, k, page, items)
|
||||
} else {
|
||||
metaEntry["error"] = discoverFeedErrorMessage(err)
|
||||
items = nil
|
||||
}
|
||||
} else {
|
||||
logDiscoverFetchSlow(svc, k, page, elapsed, len(items))
|
||||
rememberDiscoverSection(svc, k, page, items)
|
||||
}
|
||||
artworkItems = append(artworkItems, items...)
|
||||
out[k] = items
|
||||
metaEntry["has_next"] = discoverSectionHasNext(k, len(items))
|
||||
meta[k] = metaEntry
|
||||
}
|
||||
out["_meta"] = meta
|
||||
if svc != nil && svc.Discover != nil {
|
||||
svc.Discover.WarmExternalArtwork(artworkItems)
|
||||
}
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
|
||||
func cachedDiscoverSection(svc *service.Container, key string, page int) ([]service.ExternalMediaResult, bool) {
|
||||
if svc == nil || svc.Discover == nil {
|
||||
return nil, false
|
||||
}
|
||||
return svc.Discover.CachedSection(key, page)
|
||||
}
|
||||
|
||||
func rememberDiscoverSection(svc *service.Container, key string, page int, items []service.ExternalMediaResult) {
|
||||
if svc == nil || svc.Discover == nil {
|
||||
return
|
||||
}
|
||||
svc.Discover.RememberSection(key, page, items)
|
||||
}
|
||||
|
||||
func fallbackDiscoverSectionItems(parent context.Context, svc *service.Container, key string, page int) ([]service.ExternalMediaResult, string, bool) {
|
||||
fallbackKey := fallbackDiscoverSectionKey(key)
|
||||
if fallbackKey == "" || svc == nil || svc.Discover == nil {
|
||||
return nil, "", false
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(parent, discoverSectionTimeout(fallbackKey))
|
||||
defer cancel()
|
||||
items, err := discoverSectionItems(ctx, svc, fallbackKey, page)
|
||||
if err != nil || len(items) == 0 {
|
||||
return nil, fallbackKey, false
|
||||
}
|
||||
if svc.Log != nil {
|
||||
svc.Log.Info("discover section fallback used",
|
||||
zap.String("section", key),
|
||||
zap.String("fallback_section", fallbackKey),
|
||||
zap.Int("page", page),
|
||||
zap.Int("items", len(items)))
|
||||
}
|
||||
return items, fallbackKey, true
|
||||
}
|
||||
|
||||
func fallbackDiscoverSectionKey(key string) string {
|
||||
switch key {
|
||||
case "douban_hot_movie":
|
||||
return "tmdb_popular_movie"
|
||||
case "douban_hot_tv":
|
||||
return "tmdb_popular_tv"
|
||||
case "douban_top_movie":
|
||||
return "tmdb_top_rated_movie"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func logDiscoverFetchFailed(svc *service.Container, key string, page int, elapsed, timeout time.Duration, err error) {
|
||||
if svc == nil || svc.Log == nil || err == nil {
|
||||
return
|
||||
}
|
||||
svc.Log.Warn("discover section fetch failed",
|
||||
zap.String("section", key),
|
||||
zap.String("provider", discoverSectionProvider(key)),
|
||||
zap.Int("page", page),
|
||||
zap.Duration("duration", elapsed),
|
||||
zap.Int64("duration_ms", elapsed.Milliseconds()),
|
||||
zap.Duration("timeout", timeout),
|
||||
zap.Error(err))
|
||||
}
|
||||
|
||||
func logDiscoverFetchSlow(svc *service.Container, key string, page int, elapsed time.Duration, itemCount int) {
|
||||
if svc == nil || svc.Log == nil || elapsed < discoverFeedSlowSectionThreshold {
|
||||
return
|
||||
}
|
||||
svc.Log.Info("discover section fetch slow",
|
||||
zap.String("section", key),
|
||||
zap.String("provider", discoverSectionProvider(key)),
|
||||
zap.Int("page", page),
|
||||
zap.Int("items", itemCount),
|
||||
zap.Duration("duration", elapsed),
|
||||
zap.Int64("duration_ms", elapsed.Milliseconds()),
|
||||
zap.Duration("slow_threshold", discoverFeedSlowSectionThreshold))
|
||||
}
|
||||
|
||||
func discoverFeedErrorMessage(err error) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
if errors.Is(err, context.DeadlineExceeded) {
|
||||
return "推荐源响应超时,已跳过本次加载"
|
||||
}
|
||||
var timeout interface{ Timeout() bool }
|
||||
if errors.As(err, &timeout) && timeout.Timeout() {
|
||||
return "推荐源响应超时,已跳过本次加载"
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
if strings.Contains(msg, "timeout") || strings.Contains(msg, "deadline exceeded") || strings.Contains(msg, "context deadline exceeded") {
|
||||
return "推荐源响应超时,已跳过本次加载"
|
||||
}
|
||||
return "推荐源暂时不可用,已跳过本次加载"
|
||||
}
|
||||
|
||||
func discoverFeedStaleMessage(err error) string {
|
||||
if discoverFeedErrorMessage(err) == "推荐源响应超时,已跳过本次加载" {
|
||||
return "推荐源响应超时,已显示上次成功结果"
|
||||
}
|
||||
return "推荐源暂时不可用,已显示上次成功结果"
|
||||
}
|
||||
|
||||
func discoverFeedFallbackMessage(fallbackKey string, err error) string {
|
||||
if strings.TrimSpace(fallbackKey) == "" {
|
||||
return discoverFeedErrorMessage(err)
|
||||
}
|
||||
return "推荐源暂时不可用,已显示同类备用榜单"
|
||||
}
|
||||
|
||||
func discoverSectionTimeout(key string) time.Duration {
|
||||
if key == "bangumi_calendar" {
|
||||
return discoverFeedBangumiTimeout
|
||||
}
|
||||
return discoverFeedSectionTimeout
|
||||
}
|
||||
|
||||
func enabledDiscoverSections(ctx context.Context, svc *service.Container) []discoverSectionDef {
|
||||
sections := make([]discoverSectionDef, 0, len(discoverSectionCatalog))
|
||||
for _, section := range discoverSectionCatalog {
|
||||
if !discoverProviderEnabled(ctx, svc, section.Provider) {
|
||||
continue
|
||||
}
|
||||
sections = append(sections, section)
|
||||
}
|
||||
return sections
|
||||
}
|
||||
|
||||
func defaultDiscoverSectionKeys(ctx context.Context, svc *service.Container) []string {
|
||||
preferred := []string{"tmdb_trending_day", "tmdb_latest_movie", "tmdb_latest_tv", "douban_hot_movie", "douban_hot_tv", "bangumi_calendar"}
|
||||
enabled := map[string]struct{}{}
|
||||
for _, section := range enabledDiscoverSections(ctx, svc) {
|
||||
enabled[section.Key] = struct{}{}
|
||||
}
|
||||
out := make([]string, 0, len(preferred))
|
||||
for _, key := range preferred {
|
||||
if _, ok := enabled[key]; ok {
|
||||
out = append(out, key)
|
||||
}
|
||||
}
|
||||
if len(out) > 0 {
|
||||
return out
|
||||
}
|
||||
for _, section := range enabledDiscoverSections(ctx, svc) {
|
||||
out = append(out, section.Key)
|
||||
if len(out) >= 4 {
|
||||
break
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func discoverSectionProvider(key string) string {
|
||||
for _, section := range discoverSectionCatalog {
|
||||
if section.Key == key {
|
||||
return section.Provider
|
||||
}
|
||||
}
|
||||
switch key {
|
||||
case "trending_day", "trending_week", "latest_movie", "latest_tv", "popular_movie", "popular_tv", "top_rated_movie", "upcoming_movie":
|
||||
return "tmdb"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func discoverProviderEnabled(ctx context.Context, svc *service.Container, provider string) bool {
|
||||
if svc == nil || svc.APIConfig == nil || strings.TrimSpace(provider) == "" {
|
||||
return true
|
||||
}
|
||||
cfg, err := svc.APIConfig.Get(ctx, provider)
|
||||
if err != nil || cfg == nil {
|
||||
return true
|
||||
}
|
||||
return cfg.Enabled
|
||||
}
|
||||
|
||||
func discoverSectionItems(ctx context.Context, svc *service.Container, k string, page int) ([]service.ExternalMediaResult, error) {
|
||||
switch k {
|
||||
case "tmdb_trending_day", "tmdb_trending_week", "tmdb_latest_movie", "tmdb_latest_tv", "tmdb_popular_movie", "tmdb_popular_tv", "tmdb_top_rated_movie", "tmdb_upcoming_movie",
|
||||
"trending_day", "trending_week", "latest_movie", "latest_tv", "popular_movie", "popular_tv", "top_rated_movie", "upcoming_movie":
|
||||
return svc.Discover.TMDbSection(ctx, k, page)
|
||||
case "douban_hot_movie", "douban_hot_tv", "douban_top_movie":
|
||||
if svc.Douban == nil {
|
||||
return []service.ExternalMediaResult{}, nil
|
||||
}
|
||||
return svc.Douban.Discover(ctx, k, page)
|
||||
case "bangumi_calendar":
|
||||
if svc.Bangumi == nil {
|
||||
return []service.ExternalMediaResult{}, nil
|
||||
}
|
||||
if page > 1 {
|
||||
return []service.ExternalMediaResult{}, nil
|
||||
}
|
||||
return svc.Bangumi.Calendar(ctx)
|
||||
default:
|
||||
return []service.ExternalMediaResult{}, nil
|
||||
}
|
||||
}
|
||||
|
||||
func discoverSectionHasNext(key string, itemCount int) bool {
|
||||
if itemCount <= 0 {
|
||||
return false
|
||||
}
|
||||
switch discoverSectionProvider(key) {
|
||||
case "tmdb":
|
||||
return itemCount >= 20
|
||||
case "douban":
|
||||
return itemCount >= 24
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -1,183 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zaptest/observer"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func TestDiscoverProviderEnabledHonorsAPIConfigToggle(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
apiConfig := service.NewAPIConfigService(zap.NewNop(), repos, service.NewCryptoService("", zap.NewNop()))
|
||||
enabled := false
|
||||
if _, err := apiConfig.Update(t.Context(), "douban", service.APIConfigPatch{Enabled: &enabled}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc := &service.Container{APIConfig: apiConfig}
|
||||
|
||||
if discoverProviderEnabled(t.Context(), svc, "douban") {
|
||||
t.Fatal("disabled API config should disable discover provider")
|
||||
}
|
||||
if !discoverProviderEnabled(t.Context(), svc, "missing-provider") {
|
||||
t.Fatal("missing API config should keep discover provider available")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverFetchFailureLogIncludesDiagnostics(t *testing.T) {
|
||||
core, observed := observer.New(zap.WarnLevel)
|
||||
logger := zap.New(core)
|
||||
|
||||
logDiscoverFetchFailed(
|
||||
&service.Container{Log: logger},
|
||||
"tmdb_latest_movie",
|
||||
2,
|
||||
1500*time.Millisecond,
|
||||
discoverSectionTimeout("tmdb_latest_movie"),
|
||||
context.DeadlineExceeded,
|
||||
)
|
||||
|
||||
entries := observed.FilterMessage("discover section fetch failed").All()
|
||||
if len(entries) != 1 {
|
||||
t.Fatalf("expected one failure log entry, got %d", len(entries))
|
||||
}
|
||||
fields := entries[0].ContextMap()
|
||||
if fields["section"] != "tmdb_latest_movie" || fields["provider"] != "tmdb" {
|
||||
t.Fatalf("unexpected section/provider fields: %#v", fields)
|
||||
}
|
||||
if fields["page"] != int64(2) && fields["page"] != 2 {
|
||||
t.Fatalf("page field missing or wrong: %#v", fields["page"])
|
||||
}
|
||||
if fields["duration_ms"] != int64(1500) && fields["duration_ms"] != 1500 {
|
||||
t.Fatalf("duration_ms field missing or wrong: %#v", fields["duration_ms"])
|
||||
}
|
||||
if _, ok := fields["timeout"]; !ok {
|
||||
t.Fatalf("timeout field missing: %#v", fields)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverSectionTimeoutRaisesBangumiBudget(t *testing.T) {
|
||||
if got := discoverSectionTimeout("bangumi_calendar"); got != discoverFeedBangumiTimeout {
|
||||
t.Fatalf("bangumi timeout = %s, want %s", got, discoverFeedBangumiTimeout)
|
||||
}
|
||||
if got := discoverSectionTimeout("tmdb_latest_movie"); got != discoverFeedSectionTimeout {
|
||||
t.Fatalf("tmdb timeout = %s, want %s", got, discoverFeedSectionTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverSlowFetchLogIncludesSectionTiming(t *testing.T) {
|
||||
core, observed := observer.New(zap.InfoLevel)
|
||||
logger := zap.New(core)
|
||||
|
||||
logDiscoverFetchSlow(&service.Container{Log: logger}, "douban_hot_movie", 1, discoverFeedSlowSectionThreshold-time.Millisecond, 24)
|
||||
if got := observed.FilterMessage("discover section fetch slow").Len(); got != 0 {
|
||||
t.Fatalf("fast section should not log, got %d entries", got)
|
||||
}
|
||||
|
||||
logDiscoverFetchSlow(&service.Container{Log: logger}, "douban_hot_movie", 1, discoverFeedSlowSectionThreshold, 24)
|
||||
entries := observed.FilterMessage("discover section fetch slow").All()
|
||||
if len(entries) != 1 {
|
||||
t.Fatalf("expected one slow log entry, got %d", len(entries))
|
||||
}
|
||||
fields := entries[0].ContextMap()
|
||||
if fields["section"] != "douban_hot_movie" || fields["provider"] != "douban" {
|
||||
t.Fatalf("unexpected section/provider fields: %#v", fields)
|
||||
}
|
||||
if fields["items"] != int64(24) && fields["items"] != 24 {
|
||||
t.Fatalf("items field missing or wrong: %#v", fields["items"])
|
||||
}
|
||||
if _, ok := fields["duration_ms"]; !ok {
|
||||
t.Fatalf("duration_ms field missing: %#v", fields)
|
||||
}
|
||||
if _, ok := fields["slow_threshold"]; !ok {
|
||||
t.Fatalf("slow_threshold field missing: %#v", fields)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverFeedErrorMessageHidesTechnicalTimeout(t *testing.T) {
|
||||
for _, err := range []error{
|
||||
context.DeadlineExceeded,
|
||||
errors.New("timeout of 30000ms exceeded"),
|
||||
} {
|
||||
got := discoverFeedErrorMessage(err)
|
||||
if got != "推荐源响应超时,已跳过本次加载" {
|
||||
t.Fatalf("message for %q = %q", err, got)
|
||||
}
|
||||
}
|
||||
if got := discoverFeedErrorMessage(errors.New("upstream 503")); got != "推荐源暂时不可用,已跳过本次加载" {
|
||||
t.Fatalf("generic message = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultDiscoverSectionKeysSkipDisabledProviders(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
apiConfig := service.NewAPIConfigService(zap.NewNop(), repos, service.NewCryptoService("", zap.NewNop()))
|
||||
disabled := false
|
||||
for _, provider := range []string{"douban", "bangumi"} {
|
||||
if _, err := apiConfig.Update(t.Context(), provider, service.APIConfigPatch{Enabled: &disabled}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
svc := &service.Container{APIConfig: apiConfig}
|
||||
|
||||
keys := defaultDiscoverSectionKeys(t.Context(), svc)
|
||||
for _, key := range keys {
|
||||
switch discoverSectionProvider(key) {
|
||||
case "douban", "bangumi":
|
||||
t.Fatalf("disabled provider key %q should not be selected by default; keys=%v", key, keys)
|
||||
}
|
||||
}
|
||||
if len(keys) == 0 {
|
||||
t.Fatal("default keys should keep enabled providers")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultDiscoverSectionKeysIncludeLatestTMDbRails(t *testing.T) {
|
||||
keys := defaultDiscoverSectionKeys(t.Context(), &service.Container{})
|
||||
keySet := map[string]struct{}{}
|
||||
for _, key := range keys {
|
||||
keySet[key] = struct{}{}
|
||||
}
|
||||
for _, key := range []string{"tmdb_latest_movie", "tmdb_latest_tv"} {
|
||||
if _, ok := keySet[key]; !ok {
|
||||
t.Fatalf("default discover keys should include %q: %v", key, keys)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFallbackDiscoverSectionKeyUsesTMDbForDoubanRails(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"douban_hot_movie": "tmdb_popular_movie",
|
||||
"douban_hot_tv": "tmdb_popular_tv",
|
||||
"douban_top_movie": "tmdb_top_rated_movie",
|
||||
"tmdb_latest_tv": "",
|
||||
}
|
||||
for key, want := range cases {
|
||||
if got := fallbackDiscoverSectionKey(key); got != want {
|
||||
t.Fatalf("fallbackDiscoverSectionKey(%q) = %q, want %q", key, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,274 +0,0 @@
|
||||
// Package handler — 下载客户端管理 HTTP 端点。
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// DownloadClientHandler 处理下载客户端的 CRUD 操作。
|
||||
type DownloadClientHandler struct {
|
||||
svc *service.Container
|
||||
log *zap.Logger
|
||||
}
|
||||
|
||||
// NewDownloadClientHandler 创建下载客户端处理器。
|
||||
func NewDownloadClientHandler(svc *service.Container, log *zap.Logger) *DownloadClientHandler {
|
||||
return &DownloadClientHandler{svc: svc, log: log}
|
||||
}
|
||||
|
||||
// downloadClientCreateRequest 创建下载客户端请求体。
|
||||
type downloadClientCreateRequest struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Type string `json:"type" binding:"required,oneof=qbittorrent transmission aria2"`
|
||||
Host string `json:"host" binding:"required"`
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
IsDefault bool `json:"is_default"`
|
||||
Extra map[string]string `json:"extra,omitempty"`
|
||||
}
|
||||
|
||||
// downloadClientUpdateRequest 更新下载客户端请求体。
|
||||
type downloadClientUpdateRequest struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type" binding:"omitempty,oneof=qbittorrent transmission aria2"`
|
||||
Host string `json:"host"`
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
IsDefault *bool `json:"is_default"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Extra map[string]string `json:"extra,omitempty"`
|
||||
}
|
||||
|
||||
// Create 创建新的下载客户端。
|
||||
func (h *DownloadClientHandler) Create(c *gin.Context) {
|
||||
var req downloadClientCreateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
_ = h.svc.Repo.Setting.Set(ctx, "download_clients.managed", "true")
|
||||
normalizedHost, err := service.NormalizeDownloadClientHost(req.Type, req.Host)
|
||||
if err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 加密密码
|
||||
password := req.Password
|
||||
if password != "" && h.svc.Crypto != nil {
|
||||
password = h.svc.Crypto.Encrypt(password)
|
||||
}
|
||||
|
||||
// 加密 Extra 配置
|
||||
extraStr := ""
|
||||
if len(req.Extra) > 0 {
|
||||
extraJSON, _ := json.Marshal(req.Extra)
|
||||
extraStr = string(extraJSON)
|
||||
if h.svc.Crypto != nil {
|
||||
extraStr = h.svc.Crypto.Encrypt(extraStr)
|
||||
}
|
||||
}
|
||||
|
||||
// 如果设为默认,先清除其他默认
|
||||
if req.IsDefault {
|
||||
_ = h.svc.Repo.DownloadClient.ClearDefault(ctx)
|
||||
}
|
||||
|
||||
client := &model.DownloadClient{
|
||||
Name: req.Name,
|
||||
Type: req.Type,
|
||||
Host: normalizedHost,
|
||||
Username: req.Username,
|
||||
Password: password,
|
||||
IsDefault: req.IsDefault,
|
||||
Enabled: true,
|
||||
Extra: extraStr,
|
||||
}
|
||||
|
||||
if err := h.svc.Repo.DownloadClient.Create(ctx, client); err != nil {
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "创建失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.svc.Downloads.ReloadConfig(ctx); err != nil {
|
||||
h.log.Warn("failed to reload download clients after create", zap.Error(err))
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "客户端已保存,但运行时重载失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
Success(c, client)
|
||||
}
|
||||
|
||||
// List 返回所有下载客户端。
|
||||
func (h *DownloadClientHandler) List(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
clients, err := h.svc.Repo.DownloadClient.List(ctx)
|
||||
if err != nil {
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "查询失败")
|
||||
return
|
||||
}
|
||||
Success(c, clients)
|
||||
}
|
||||
|
||||
// Get 返回指定下载客户端详情。
|
||||
func (h *DownloadClientHandler) Get(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
ctx := c.Request.Context()
|
||||
client, err := h.svc.Repo.DownloadClient.FindByID(ctx, id)
|
||||
if err != nil || client == nil {
|
||||
Error(c, http.StatusNotFound, ErrNotFound, "客户端不存在")
|
||||
return
|
||||
}
|
||||
Success(c, client)
|
||||
}
|
||||
|
||||
// Update 更新下载客户端。
|
||||
func (h *DownloadClientHandler) Update(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
var req downloadClientUpdateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
client, err := h.svc.Repo.DownloadClient.FindByID(ctx, id)
|
||||
if err != nil || client == nil {
|
||||
Error(c, http.StatusNotFound, ErrNotFound, "客户端不存在")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name != "" {
|
||||
client.Name = req.Name
|
||||
}
|
||||
if req.Type != "" {
|
||||
client.Type = req.Type
|
||||
}
|
||||
if req.Host != "" {
|
||||
clientType := client.Type
|
||||
if req.Type != "" {
|
||||
clientType = req.Type
|
||||
}
|
||||
normalizedHost, err := service.NormalizeDownloadClientHost(clientType, req.Host)
|
||||
if err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||
return
|
||||
}
|
||||
client.Host = normalizedHost
|
||||
}
|
||||
if req.Username != "" {
|
||||
client.Username = req.Username
|
||||
}
|
||||
if req.Password != "" {
|
||||
if h.svc.Crypto != nil {
|
||||
client.Password = h.svc.Crypto.Encrypt(req.Password)
|
||||
} else {
|
||||
client.Password = req.Password
|
||||
}
|
||||
}
|
||||
if req.IsDefault != nil && *req.IsDefault {
|
||||
_ = h.svc.Repo.DownloadClient.ClearDefault(ctx)
|
||||
client.IsDefault = *req.IsDefault
|
||||
}
|
||||
if req.Enabled != nil {
|
||||
client.Enabled = *req.Enabled
|
||||
}
|
||||
if len(req.Extra) > 0 {
|
||||
extraJSON, _ := json.Marshal(req.Extra)
|
||||
extraStr := string(extraJSON)
|
||||
if h.svc.Crypto != nil {
|
||||
client.Extra = h.svc.Crypto.Encrypt(extraStr)
|
||||
} else {
|
||||
client.Extra = extraStr
|
||||
}
|
||||
}
|
||||
normalizedHost, err := service.NormalizeDownloadClientHost(client.Type, client.Host)
|
||||
if err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||
return
|
||||
}
|
||||
client.Host = normalizedHost
|
||||
|
||||
if err := h.svc.Repo.DownloadClient.Update(ctx, client); err != nil {
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "更新失败")
|
||||
return
|
||||
}
|
||||
clearLegacyQBitSettingsIfNoDefault(c.Request.Context(), h.svc)
|
||||
|
||||
if err := h.svc.Downloads.ReloadConfig(ctx); err != nil {
|
||||
h.log.Warn("failed to reload download clients after update", zap.Error(err))
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "客户端已更新,但运行时重载失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
Success(c, client)
|
||||
}
|
||||
|
||||
// Delete 删除下载客户端。
|
||||
func (h *DownloadClientHandler) Delete(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
ctx := c.Request.Context()
|
||||
_ = h.svc.Repo.Setting.Set(ctx, "download_clients.managed", "true")
|
||||
|
||||
_, err := h.svc.Repo.DownloadClient.FindByID(ctx, id)
|
||||
if err != nil {
|
||||
Error(c, http.StatusNotFound, ErrNotFound, "客户端不存在")
|
||||
return
|
||||
}
|
||||
|
||||
if delErr := h.svc.Repo.DownloadClient.Delete(ctx, id); delErr != nil {
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "删除失败")
|
||||
return
|
||||
}
|
||||
|
||||
clearLegacyQBitSettingsIfNoDefault(c.Request.Context(), h.svc)
|
||||
if err := h.svc.Downloads.ReloadConfig(ctx); err != nil {
|
||||
h.log.Warn("failed to reload download clients after delete", zap.Error(err))
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "客户端已删除,但运行时重载失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
SuccessWithMessage(c, "已删除", nil)
|
||||
}
|
||||
|
||||
func clearLegacyQBitSettingsIfNoDefault(ctx context.Context, svc *service.Container) {
|
||||
if svc == nil || svc.Repo == nil || svc.Repo.DownloadClient == nil || svc.Repo.Setting == nil {
|
||||
return
|
||||
}
|
||||
defaultClient, err := svc.Repo.DownloadClient.FindDefault(ctx)
|
||||
if err != nil || defaultClient != nil {
|
||||
return
|
||||
}
|
||||
_ = svc.Repo.Setting.Set(ctx, "qbittorrent.url", "")
|
||||
_ = svc.Repo.Setting.Set(ctx, "qbittorrent.username", "")
|
||||
_ = svc.Repo.Setting.Set(ctx, "qbittorrent.password", "")
|
||||
}
|
||||
|
||||
// Test 测试下载客户端连接。
|
||||
func (h *DownloadClientHandler) Test(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
ctx := c.Request.Context()
|
||||
|
||||
client, err := h.svc.Repo.DownloadClient.FindByID(ctx, id)
|
||||
if err != nil || client == nil {
|
||||
Error(c, http.StatusNotFound, ErrNotFound, "客户端不存在")
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.svc.DownloadMgr.TestConnection(ctx, client); err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrExternal, "连接测试失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
SuccessWithMessage(c, "连接成功", nil)
|
||||
}
|
||||
@@ -1,108 +0,0 @@
|
||||
// Package handler — download client (qBittorrent / Aria2 / Transmission)
|
||||
// configuration endpoints.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func listDownloadClientsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
rows, err := svc.DownloadClients.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if rows == nil {
|
||||
rows = []model.DownloadClient{}
|
||||
}
|
||||
c.JSON(http.StatusOK, rows)
|
||||
}
|
||||
}
|
||||
|
||||
func createDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var in service.DownloadClientInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
row, err := svc.DownloadClients.Create(c.Request.Context(), in)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
// 让真正发起下载的 DownloadService 立刻读到新的 qb 配置,
|
||||
// 避免保存后还要重启进程才能生效。
|
||||
if err := svc.Downloads.ReloadConfig(c.Request.Context()); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "client saved but runtime reload failed: " + err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, row)
|
||||
}
|
||||
}
|
||||
|
||||
func updateDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var in service.DownloadClientInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
row, err := svc.DownloadClients.Update(c.Request.Context(), c.Param("id"), in)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := svc.Downloads.ReloadConfig(c.Request.Context()); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "client updated but runtime reload failed: " + err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, row)
|
||||
}
|
||||
}
|
||||
|
||||
func deleteDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.DownloadClients.Delete(c.Request.Context(), c.Param("id")); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := svc.Downloads.ReloadConfig(c.Request.Context()); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "client deleted but runtime reload failed: " + err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
func testDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.DownloadClients.Test(c.Request.Context(), c.Param("id")); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": false, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
|
||||
func aria2StatsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
clientID := c.Query("client_id")
|
||||
if clientID == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "client_id required"})
|
||||
return
|
||||
}
|
||||
out, err := svc.DownloadClients.Aria2GlobalStats(c.Request.Context(), clientID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
@@ -1,258 +0,0 @@
|
||||
// Package handler — download manager endpoints.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/middleware"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type addDownloadReq struct {
|
||||
URL string `json:"url" binding:"required"`
|
||||
SavePath string `json:"save_path"`
|
||||
Title string `json:"title"`
|
||||
PosterURL string `json:"poster_url"`
|
||||
BackdropURL string `json:"backdrop_url"`
|
||||
Overview string `json:"overview"`
|
||||
MediaType string `json:"media_type"`
|
||||
MediaCategory string `json:"media_category"`
|
||||
SourceCategory string `json:"source_category"`
|
||||
}
|
||||
|
||||
// resolvePTDownloadURL 把站点搜索结果里的"详情/获取签名"URL 解析成 qb 能直接
|
||||
// 拉到 .torrent 文件的真实下载 URL。
|
||||
//
|
||||
// 链路:
|
||||
//
|
||||
// 1. 拿 URL 的 host,到 sites 表里找 base_url 同源的站点。
|
||||
// 2. 如果站点的 type 是已知 PT 框架(mteam/nexusphp/unit3d/...),
|
||||
// 就用对应适配器的 GetDownloadURL,传入从 URL 里 parse 出来的 id。
|
||||
// 3. 任一步失败都直接返回原 URL,让 qb 自己去拉(保持向后兼容)。
|
||||
//
|
||||
// 这一步存在的意义:M-Team 等站点的搜索结果里 download_url 是
|
||||
// /api/torrent/genDlToken?id=xxx,需要带 x-api-key 才能调用,qb 自己
|
||||
// 是没法识别这种 PT 专属端点的。
|
||||
func resolvePTDownloadURL(ctx context.Context, svc *service.Container, raw string, log *zap.Logger) string {
|
||||
if raw == "" || svc == nil || svc.Site == nil {
|
||||
return raw
|
||||
}
|
||||
resolved := svc.Site.ResolveDownloadURL(ctx, raw)
|
||||
if resolved == raw {
|
||||
return raw
|
||||
}
|
||||
log.Info("resolved PT download URL",
|
||||
zap.String("from", redactDownloadURL(raw)),
|
||||
zap.String("to", redactDownloadURL(resolved)))
|
||||
return resolved
|
||||
}
|
||||
|
||||
func addDownloadHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req addDownloadReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
// 把站点搜索 URL 转换成真实可下载 URL(M-Team 走 genDlToken 等)。
|
||||
realURL := resolvePTDownloadURL(c.Request.Context(), svc, req.URL, svc.Log)
|
||||
fallbackTitle := req.Title
|
||||
if strings.TrimSpace(fallbackTitle) == "" {
|
||||
fallbackTitle = realURL
|
||||
}
|
||||
meta := enrichDownloadTaskMeta(c.Request.Context(), svc, service.DownloadTaskMeta{
|
||||
Title: req.Title,
|
||||
PosterURL: req.PosterURL,
|
||||
BackdropURL: req.BackdropURL,
|
||||
Overview: req.Overview,
|
||||
MediaType: req.MediaType,
|
||||
MediaCategory: req.MediaCategory,
|
||||
SourceCategory: req.SourceCategory,
|
||||
}, fallbackTitle, req.MediaType)
|
||||
t, err := svc.Downloads.AddDownloadWithMeta(c.Request.Context(), uid.(string), realURL, req.SavePath, meta)
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrMediaAlreadyInLibrary) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": "media already exists in library"})
|
||||
return
|
||||
}
|
||||
if errors.Is(err, service.ErrDownloadAlreadyExists) {
|
||||
c.JSON(http.StatusOK, t)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
svc.Audit.Record(c.Request.Context(), uid.(string), "download.add", redactDownloadURL(realURL), c.ClientIP(), "")
|
||||
c.JSON(http.StatusCreated, t)
|
||||
}
|
||||
}
|
||||
|
||||
func listDownloadsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
rows, live, err := svc.Downloads.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
rows = visibleDownloadRows(c, rows)
|
||||
enrichAndPersistDownloadRows(c.Request.Context(), svc, rows)
|
||||
if !isAdminRequest(c) {
|
||||
live = visibleLiveTorrents(rows, live)
|
||||
}
|
||||
taskViews, torrentViews := service.DownloadViews(rows, live)
|
||||
enrichDownloadTorrentViews(c.Request.Context(), svc, torrentViews)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"tasks": taskViews,
|
||||
"torrents": torrentViews,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func visibleDownloadRows(c *gin.Context, rows []model.DownloadTask) []model.DownloadTask {
|
||||
if isAdminRequest(c) {
|
||||
return rows
|
||||
}
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
userID, _ := uid.(string)
|
||||
filtered := make([]model.DownloadTask, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
if row.UserID == userID {
|
||||
filtered = append(filtered, row)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func visibleLiveTorrents(rows []model.DownloadTask, live []service.QBitTorrent) []service.QBitTorrent {
|
||||
if len(rows) == 0 || len(live) == 0 {
|
||||
return nil
|
||||
}
|
||||
filtered := make([]service.QBitTorrent, 0, len(live))
|
||||
for _, torrent := range live {
|
||||
matchedByID := false
|
||||
for _, row := range rows {
|
||||
if strings.TrimSpace(row.ExternalID) != "" && strings.EqualFold(row.ExternalID, torrent.Hash) &&
|
||||
(strings.TrimSpace(row.DownloadClientID) == "" || row.DownloadClientID == torrent.ClientID) {
|
||||
filtered = append(filtered, torrent)
|
||||
matchedByID = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if matchedByID {
|
||||
continue
|
||||
}
|
||||
torrentTitle := normalizeTitle(torrent.Name)
|
||||
if torrentTitle == "" {
|
||||
continue
|
||||
}
|
||||
for _, row := range rows {
|
||||
if strings.TrimSpace(row.DownloadClientID) != "" && strings.TrimSpace(torrent.ClientID) != "" && row.DownloadClientID != torrent.ClientID {
|
||||
continue
|
||||
}
|
||||
rowTitle := normalizeTitle(row.Title)
|
||||
if rowTitle == "" {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(torrentTitle, rowTitle) || strings.Contains(rowTitle, torrentTitle) {
|
||||
filtered = append(filtered, torrent)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func isAdminRequest(c *gin.Context) bool {
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
return role == "admin"
|
||||
}
|
||||
|
||||
func normalizeTitle(title string) string {
|
||||
title = strings.ToLower(title)
|
||||
var b strings.Builder
|
||||
for _, r := range title {
|
||||
if (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') || r > 127 {
|
||||
b.WriteRune(r)
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func redactDownloadURL(raw string) string {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return ""
|
||||
}
|
||||
if strings.HasPrefix(strings.ToLower(raw), "magnet:") {
|
||||
return "magnet:?xt=***"
|
||||
}
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil || u.Host == "" {
|
||||
return "[redacted-download-url]"
|
||||
}
|
||||
u.RawQuery = ""
|
||||
u.Fragment = ""
|
||||
base := u.String()
|
||||
if base == "" {
|
||||
return u.Scheme + "://" + u.Host
|
||||
}
|
||||
return base
|
||||
}
|
||||
|
||||
func deleteDownloadHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
hash := c.Param("hash")
|
||||
withFiles := c.Query("delete_files") == "true"
|
||||
if err := svc.Downloads.Delete(c.Request.Context(), hash, withFiles, c.Query("client_id")); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
type relocateDownloadReq struct {
|
||||
Hash string `json:"hash" binding:"required"`
|
||||
Location string `json:"location" binding:"required"`
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
|
||||
// relocateDownloadHandler moves a torrent's data to a new directory while
|
||||
// keeping it seeding (qBittorrent setLocation).
|
||||
func relocateDownloadHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req relocateDownloadReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := svc.Downloads.RelocateTorrent(c.Request.Context(), req.Hash, req.Location, req.ClientID); err != nil {
|
||||
status := http.StatusInternalServerError
|
||||
if errors.Is(err, service.ErrDownloadOperationUnsupported) {
|
||||
status = http.StatusBadRequest
|
||||
}
|
||||
c.JSON(status, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"hash": strings.TrimSpace(req.Hash), "location": strings.TrimSpace(req.Location)})
|
||||
}
|
||||
}
|
||||
|
||||
func reloadDownloadConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Downloads.ReloadConfig(c.Request.Context()); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
@@ -1,112 +0,0 @@
|
||||
// Package handler — pause/resume/organize on individual download tasks
|
||||
// and a thin sync-trigger surface used by the Vue UI's auto-sync toggle.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func downloadPauseHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Downloads.PauseDownloadTask(c.Request.Context(), c.Param("id")); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
|
||||
func downloadResumeHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Downloads.ResumeDownloadTask(c.Request.Context(), c.Param("id")); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
|
||||
// downloadOrganizeOneHandler runs the file organizer for one task.
|
||||
// It looks up the task, then delegates to OrganizerService.OrganizePath().
|
||||
func downloadOrganizeOneHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var t model.DownloadTask
|
||||
if err := svc.Repo.DB.WithContext(c.Request.Context()).
|
||||
Where("id = ?", c.Param("id")).First(&t).Error; err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "task not found"})
|
||||
return
|
||||
}
|
||||
if t.SavePath == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "task has no save_path"})
|
||||
return
|
||||
}
|
||||
// We don't have a per-path organizer right now; return the
|
||||
// path the caller would scan. The general OrganizeAll endpoint
|
||||
// (below) is the supported workflow.
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"ok": true,
|
||||
"path": t.SavePath,
|
||||
"note": "use POST /api/download/organize to bulk-organize",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// downloadOrganizeAllHandler triggers a bulk re-organize. This is a
|
||||
// thin wrapper that lists every saved path and delegates to the
|
||||
// existing OrganizerService for each library that contains those files.
|
||||
func downloadOrganizeAllHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// Walk each library through the unified organize pipeline so rename,
|
||||
// scan, scrape and task reporting stay identical to manual organize.
|
||||
libs, err := svc.Repo.Library.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
results := make([]any, 0, len(libs))
|
||||
for _, l := range libs {
|
||||
resp, err := organizePipeline(svc).Run(c.Request.Context(), service.OrganizePipelineRequest{
|
||||
Scope: service.OrganizeScopeLibrary,
|
||||
Trigger: service.OrganizeTriggerManual,
|
||||
TaskName: "批量整理媒体库:" + l.Name,
|
||||
LibraryID: l.ID,
|
||||
})
|
||||
if err != nil {
|
||||
results = append(results, gin.H{"library": l.Name, "error": err.Error()})
|
||||
continue
|
||||
}
|
||||
results = append(results, gin.H{"library": l.Name, "result": resp.Result})
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"results": results})
|
||||
}
|
||||
}
|
||||
|
||||
// downloadSyncHandler triggers the qBittorrent reload + immediate poll.
|
||||
func downloadSyncHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Downloads.ReloadConfig(c.Request.Context()); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
|
||||
// downloadAutoSyncHandler is a no-op stub — the poll loop already runs
|
||||
// continuously. Returning 200 keeps the Vue UI's toggle happy.
|
||||
func downloadAutoSyncHandler(_ *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true, "auto_sync": true})
|
||||
}
|
||||
}
|
||||
|
||||
// downloadTasksAliasHandler is the alias used by the Vue UI; it
|
||||
// returns the same shape as listDownloadsHandler but at /download/tasks.
|
||||
func downloadTasksAliasHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return listDownloadsHandler(svc)
|
||||
}
|
||||
@@ -1,46 +0,0 @@
|
||||
// Package handler — duplicate-file finder.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func listDuplicatesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libraryID := c.Query("library_id")
|
||||
report, err := svc.Duplicate.Current(c.Request.Context(), libraryID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, report)
|
||||
}
|
||||
}
|
||||
|
||||
func detectDuplicatesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libraryID := c.Query("library_id")
|
||||
report, err := svc.Duplicate.Detect(c.Request.Context(), libraryID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, report)
|
||||
}
|
||||
}
|
||||
|
||||
func unmarkDuplicatesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libraryID := c.Query("library_id")
|
||||
n, err := svc.Duplicate.Unmark(c.Request.Context(), libraryID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"unmarked": n})
|
||||
}
|
||||
}
|
||||
@@ -40,11 +40,6 @@ func embyItemImageHandler(svc *service.Container) gin.HandlerFunc {
|
||||
embyServePlaceholderImage(c)
|
||||
return
|
||||
}
|
||||
if typ, ref, ok := service.ParseCloudArtworkURL(raw); ok {
|
||||
c.Request = req
|
||||
serveCloudResolvedLink(svc, c, typ, ref)
|
||||
return
|
||||
}
|
||||
if svc.ImageProxy == nil {
|
||||
embyServePlaceholderImage(c)
|
||||
return
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -19,7 +18,6 @@ import (
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
|
||||
)
|
||||
|
||||
func TestEmbyItemImageServesWithoutAPIAuth(t *testing.T) {
|
||||
@@ -92,63 +90,6 @@ func TestEmbyItemImageServesWithoutAPIAuth(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbyItemImageServesCachedCloudArtworkWithoutResolve(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "image/jpeg")
|
||||
_, _ = w.Write(handlerTestJPEG)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatalf("open db: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.Media{}); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
|
||||
cfg := &config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}
|
||||
imageProxy := service.NewImageProxy(cfg, zap.NewNop())
|
||||
ref := "/Movies/Cloud Movie/poster.jpg"
|
||||
if err := imageProxy.PrefetchCloudResolved(t.Context(), "openlist:"+ref, &cloud.DirectLink{URL: upstream.URL + "/poster.jpg"}); err != nil {
|
||||
t.Fatalf("prefetch cloud poster: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
if err := db.Create(&model.Media{
|
||||
Base: model.Base{ID: "cloud-media-1"},
|
||||
Title: "Cloud Poster Test",
|
||||
Path: "cloud://openlist/Movies/Cloud Movie/movie.mkv",
|
||||
PosterURL: service.CloudArtworkURL("openlist", ref),
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create media: %v", err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
registerEmbyRoutes(router, "test-secret", &service.Container{
|
||||
Repo: repos,
|
||||
Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
|
||||
ImageProxy: imageProxy,
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/Items/cloud-media-1/Images/Primary", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
if got := w.Body.Bytes(); !bytes.Equal(got, handlerTestJPEG) {
|
||||
t.Fatalf("body = %q, want cached cloud poster", got)
|
||||
}
|
||||
if location := w.Header().Get("Location"); location != "" {
|
||||
t.Fatalf("expected direct cached image response, got redirect to %q", location)
|
||||
}
|
||||
if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "max-age=2592000") {
|
||||
t.Fatalf("image Cache-Control = %q, want long browser cache", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbyMissingItemImageReturnsTransparentPlaceholder(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
|
||||
@@ -1,223 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"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/ShukeBta/MediaStationGo/internal/config"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/middleware"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func TestEmbyVideoStreamUsesSTRMWhenRedirectProxyDisabled(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.AllModels()...); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
|
||||
t.Fatalf("set cloud playback mode: %v", err)
|
||||
}
|
||||
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackSTRMEnabledSettingKey, "true"); err != nil {
|
||||
t.Fatalf("enable strm playback: %v", err)
|
||||
}
|
||||
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackRedirectEnabledSettingKey, "false"); err != nil {
|
||||
t.Fatalf("disable redirect playback: %v", err)
|
||||
}
|
||||
if err := repos.User.Create(t.Context(), &model.User{
|
||||
Base: model.Base{ID: "user-1"},
|
||||
Username: "tester",
|
||||
PasswordHash: "x",
|
||||
Role: "admin",
|
||||
Tier: "plus",
|
||||
IsActive: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
|
||||
if err := repos.Library.Create(t.Context(), &lib); err != nil {
|
||||
t.Fatalf("create library: %v", err)
|
||||
}
|
||||
if err := db.Create(&model.Media{
|
||||
Base: model.Base{ID: "cloud-1"},
|
||||
LibraryID: lib.ID,
|
||||
Title: "Cloud Movie",
|
||||
Path: "cloud://openlist/Movies/Movie.mkv",
|
||||
STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
|
||||
Container: "mkv",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create media: %v", err)
|
||||
}
|
||||
|
||||
const secret = "test-secret"
|
||||
router := gin.New()
|
||||
cfg := &config.Config{Secrets: config.SecretsConfig{JWTSecret: secret}}
|
||||
registerEmbyRoutes(router, secret, &service.Container{
|
||||
Repo: repos,
|
||||
Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
|
||||
Stream: service.NewStreamService(cfg, zap.NewNop(), repos, nil),
|
||||
})
|
||||
|
||||
token := signedTestToken(t, secret)
|
||||
req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream?api_key="+token, nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusFound {
|
||||
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
loc := w.Header().Get("Location")
|
||||
if !strings.Contains(loc, "/api/stream/cloud-1") || !strings.Contains(loc, "api_key=") {
|
||||
t.Fatalf("STRM mode should redirect /Videos fallback to tokenized /api/stream, got %q", loc)
|
||||
}
|
||||
if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
|
||||
t.Fatalf("STRM fallback redirect Cache-Control = %q, want no-store", got)
|
||||
}
|
||||
if strings.Contains(loc, "/api/cloud/play/") {
|
||||
t.Fatalf("STRM mode should not expose cloud play directly from /Videos fallback: %q", loc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbyVideoStreamIssuesTokenForSessionFallbackSTRMRedirect(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.AllModels()...); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
|
||||
t.Fatalf("set cloud playback mode: %v", err)
|
||||
}
|
||||
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackSTRMEnabledSettingKey, "true"); err != nil {
|
||||
t.Fatalf("enable strm playback: %v", err)
|
||||
}
|
||||
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackRedirectEnabledSettingKey, "false"); err != nil {
|
||||
t.Fatalf("disable redirect playback: %v", err)
|
||||
}
|
||||
user := model.User{
|
||||
Base: model.Base{ID: "user-1"},
|
||||
Username: "tester",
|
||||
PasswordHash: "x",
|
||||
Role: "admin",
|
||||
Tier: "plus",
|
||||
IsActive: true,
|
||||
}
|
||||
if err := repos.User.Create(t.Context(), &user); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
|
||||
if err := repos.Library.Create(t.Context(), &lib); err != nil {
|
||||
t.Fatalf("create library: %v", err)
|
||||
}
|
||||
if err := db.Create(&model.Media{
|
||||
Base: model.Base{ID: "cloud-1"},
|
||||
LibraryID: lib.ID,
|
||||
Title: "Cloud Movie",
|
||||
Path: "cloud://openlist/Movies/Movie.mkv",
|
||||
STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
|
||||
Container: "mkv",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create media: %v", err)
|
||||
}
|
||||
|
||||
const secret = "test-secret"
|
||||
cfg := &config.Config{Secrets: config.SecretsConfig{JWTSecret: secret}}
|
||||
svc := &service.Container{
|
||||
Repo: repos,
|
||||
Auth: service.NewAuthService(cfg, zap.NewNop(), repos, nil, nil),
|
||||
Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
|
||||
Stream: service.NewStreamService(cfg, zap.NewNop(), repos, nil),
|
||||
}
|
||||
router := gin.New()
|
||||
router.GET("/videos/:id/stream", func(c *gin.Context) {
|
||||
c.Set(middleware.CtxUserID, user.ID)
|
||||
c.Set(middleware.CtxUserRole, user.Role)
|
||||
embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy)(c)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusFound {
|
||||
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
loc := w.Header().Get("Location")
|
||||
if !strings.Contains(loc, "/api/stream/cloud-1") || !strings.Contains(loc, "api_key=") {
|
||||
t.Fatalf("session fallback redirect should include api_key for /api/stream, got %q", loc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbyVideoStreamRedirectKeepsMediaBrowserAuthorizationToken(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.AllModels()...); err != nil {
|
||||
t.Fatalf("migrate: %v", err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
if err := repos.User.Create(t.Context(), &model.User{
|
||||
Base: model.Base{ID: "user-1"},
|
||||
Username: "tester",
|
||||
PasswordHash: "x",
|
||||
Role: "admin",
|
||||
Tier: "plus",
|
||||
IsActive: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
|
||||
if err := repos.Library.Create(t.Context(), &lib); err != nil {
|
||||
t.Fatalf("create library: %v", err)
|
||||
}
|
||||
if err := db.Create(&model.Media{
|
||||
Base: model.Base{ID: "cloud-1"},
|
||||
LibraryID: lib.ID,
|
||||
Title: "Cloud Movie",
|
||||
Path: "cloud://openlist/Movies/Movie.mkv",
|
||||
STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
|
||||
Container: "mkv",
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create media: %v", err)
|
||||
}
|
||||
|
||||
const secret = "test-secret"
|
||||
router := gin.New()
|
||||
registerEmbyRoutes(router, secret, &service.Container{
|
||||
Repo: repos,
|
||||
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
|
||||
Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
|
||||
})
|
||||
|
||||
token := signedTestToken(t, secret)
|
||||
req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream", nil)
|
||||
req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="PC", Token="`+token+`"`)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusFound {
|
||||
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
loc := w.Header().Get("Location")
|
||||
if !strings.Contains(loc, "/api/cloud/play/openlist?") || !strings.Contains(loc, "token=") {
|
||||
t.Fatalf("redirect Location should target tokenized cloud play endpoint, got %q", loc)
|
||||
}
|
||||
}
|
||||
@@ -20,9 +20,6 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
|
||||
api.GET("/version", versionInfo)
|
||||
api.GET("/public/ui-config", publicUIConfigHandler(svc))
|
||||
|
||||
// Telegram Bot webhook — called by Telegram servers, no auth.
|
||||
api.POST("/telegram/webhook", telegramWebhookHandler(svc))
|
||||
|
||||
registerPublicAuthRoutes(api, svc, log)
|
||||
|
||||
registerAuthenticatedRoutes(api, cfg, svc)
|
||||
@@ -90,42 +87,6 @@ func testApiConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return h.TestApiConfig
|
||||
}
|
||||
|
||||
// ─── Download Client Handler 包装 ─────────────────────────────────────────────
|
||||
|
||||
func getDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
|
||||
h := NewDownloadClientHandler(svc, svc.Log)
|
||||
return h.Get
|
||||
}
|
||||
|
||||
// ─── Notify Channel Handler 包装 ──────────────────────────────────────────────
|
||||
|
||||
func getNotifyChannelTypesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
h := NewNotifyHandler(svc, svc.Log)
|
||||
return h.GetTypes
|
||||
}
|
||||
|
||||
func getNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
|
||||
h := NewNotifyHandler(svc, svc.Log)
|
||||
return h.Get
|
||||
}
|
||||
|
||||
// ─── Scheduler Handler 包装 ──────────────────────────────────────────────────
|
||||
|
||||
func schedulerListTasksHandler(svc *service.Container) gin.HandlerFunc {
|
||||
h := NewSchedulerHandler(svc, svc.Log)
|
||||
return h.ListTasks
|
||||
}
|
||||
|
||||
func schedulerRunTaskHandler(svc *service.Container) gin.HandlerFunc {
|
||||
h := NewSchedulerHandler(svc, svc.Log)
|
||||
return h.RunTask
|
||||
}
|
||||
|
||||
func schedulerGetStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
h := NewSchedulerHandler(svc, svc.Log)
|
||||
return h.GetStatus
|
||||
}
|
||||
|
||||
// ─── SSE Handler ──────────────────────────────────────────────────────────────
|
||||
|
||||
func sseHandler(svc *service.Container) gin.HandlerFunc {
|
||||
|
||||
@@ -1,167 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
const (
|
||||
licenseServerURLSetting = "license.server_url"
|
||||
licenseHMACSecretSetting = "license.hmac_secret" // #nosec G101 -- setting key name, not the HMAC secret value.
|
||||
licenseDeviceIDSetting = "license.device_id"
|
||||
licenseDeviceNameSetting = "license.device_name"
|
||||
)
|
||||
|
||||
type licenseActivateReq struct {
|
||||
Key string `json:"key" binding:"required"`
|
||||
// DeviceID is accepted for wire compatibility with older web clients but is
|
||||
// intentionally ignored. Licensing binds to this MediaStationGo server
|
||||
// instance, not to the browser that opened the admin page.
|
||||
DeviceID string `json:"device_id"`
|
||||
DeviceName string `json:"device_name"`
|
||||
}
|
||||
|
||||
type licenseServerSignedResp struct {
|
||||
Valid bool `json:"valid"`
|
||||
LicenseType string `json:"license_type"`
|
||||
ExpiryDate *string `json:"expiry_date"`
|
||||
MaxDevices int `json:"max_devices"`
|
||||
MaxUsers *int `json:"max_users"`
|
||||
DaysRemaining *int `json:"days_remaining"`
|
||||
NextHeartbeat string `json:"next_heartbeat"`
|
||||
Signature string `json:"signature"`
|
||||
SignatureAlg string `json:"signature_alg"`
|
||||
LegacySignature bool `json:"-"`
|
||||
}
|
||||
|
||||
type licenseServerStatusResp struct {
|
||||
Valid bool `json:"valid"`
|
||||
LicenseType *string `json:"license_type"`
|
||||
ExpiryDate *string `json:"expiry_date"`
|
||||
MaxDevices int `json:"max_devices"`
|
||||
MaxUsers *int `json:"max_users"`
|
||||
UnlimitedUsers bool `json:"unlimited_users"`
|
||||
DaysRemaining *int `json:"days_remaining"`
|
||||
DeviceName string `json:"device_name"`
|
||||
IsActive bool `json:"is_active"`
|
||||
HeartbeatRequested bool `json:"heartbeat_requested"`
|
||||
}
|
||||
|
||||
func licenseActivateHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req licenseActivateReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
client, err := newLicenseClient(c.Request.Context(), svc)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
deviceID, err := ensureLicenseDeviceID(c.Request.Context(), svc, "")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
deviceName := strings.TrimSpace(req.DeviceName)
|
||||
if deviceName == "" {
|
||||
deviceName = defaultLicenseDeviceName()
|
||||
}
|
||||
_ = svc.Repo.Setting.Set(c.Request.Context(), licenseDeviceNameSetting, deviceName)
|
||||
|
||||
payload := map[string]any{
|
||||
"key": strings.TrimSpace(req.Key),
|
||||
"fingerprint": deviceID,
|
||||
"device_name": deviceName,
|
||||
"instance_id": deviceID,
|
||||
}
|
||||
var upstream licenseServerSignedResp
|
||||
if err := client.post(c.Request.Context(), "/api/v1/activate", payload, &upstream); err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := client.verifySigned(&upstream); err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
state := licenseStateFromSigned(upstream, deviceID, deviceName)
|
||||
state.LicenseKey = strings.TrimSpace(req.Key)
|
||||
if err := persistLicenseState(c.Request.Context(), svc, state); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, licenseActivationView(state))
|
||||
}
|
||||
}
|
||||
|
||||
func licenseStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
state, _ := loadLicenseState(c.Request.Context(), svc)
|
||||
client, err := newLicenseClient(c.Request.Context(), svc)
|
||||
hasLicenseKey := strings.TrimSpace(state.LicenseKey) != ""
|
||||
if err == nil && hasLicenseKey {
|
||||
deviceID, idErr := ensureLicenseDeviceID(c.Request.Context(), svc, state.DeviceID)
|
||||
if idErr == nil {
|
||||
deviceName, _ := svc.Repo.Setting.Get(c.Request.Context(), licenseDeviceNameSetting)
|
||||
if strings.TrimSpace(deviceName) == "" {
|
||||
deviceName = defaultLicenseDeviceName()
|
||||
_ = svc.Repo.Setting.Set(c.Request.Context(), licenseDeviceNameSetting, deviceName)
|
||||
}
|
||||
var signed licenseServerSignedResp
|
||||
if heartbeatErr := client.post(c.Request.Context(), "/api/v1/heartbeat", licenseHeartbeatPayload(state, deviceID, deviceName), &signed); heartbeatErr == nil && client.verifySigned(&signed) == nil {
|
||||
nextState := licenseStateFromSigned(signed, deviceID, deviceName)
|
||||
nextState.LicenseKey = state.LicenseKey
|
||||
state = nextState
|
||||
if refreshed, ok, _, refreshErr := refreshLicenseServerStatus(c.Request.Context(), client, state, deviceID); refreshErr == nil && ok {
|
||||
refreshed.LicenseKey = state.LicenseKey
|
||||
state = refreshed
|
||||
}
|
||||
_ = persistLicenseState(c.Request.Context(), svc, state)
|
||||
} else {
|
||||
if refreshed, ok, _, getErr := refreshLicenseServerStatus(c.Request.Context(), client, state, deviceID); getErr == nil && ok {
|
||||
state = refreshed
|
||||
_ = persistLicenseState(c.Request.Context(), svc, state)
|
||||
} else if getErr == nil {
|
||||
state.Valid = false
|
||||
_ = persistLicenseState(c.Request.Context(), svc, state)
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if err == nil && state.Valid {
|
||||
deviceID, idErr := ensureLicenseDeviceID(c.Request.Context(), svc, state.DeviceID)
|
||||
if idErr == nil {
|
||||
if refreshed, ok, _, getErr := refreshLicenseServerStatus(c.Request.Context(), client, state, deviceID); getErr == nil && ok {
|
||||
state = refreshed
|
||||
_ = persistLicenseState(c.Request.Context(), svc, state)
|
||||
} else if getErr == nil {
|
||||
state.Valid = false
|
||||
_ = persistLicenseState(c.Request.Context(), svc, state)
|
||||
}
|
||||
}
|
||||
}
|
||||
active := state.Valid && !licenseStateExpired(state.ExpiryDate)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"active": active,
|
||||
"message": licenseStatusMessage(active, err),
|
||||
"max_users": licenseStatusMaxUsers(state),
|
||||
"unlimited_users": state.Valid && !licenseStateExpired(state.ExpiryDate) && state.UnlimitedUsers,
|
||||
"activation": licenseActivationView(state),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func licenseHeartbeatHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
state, err := sendLicenseHeartbeat(c.Request.Context(), svc)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, licenseActivationView(state))
|
||||
}
|
||||
}
|
||||
@@ -1,229 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"crypto/x509"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type licenseClient struct {
|
||||
baseURL string
|
||||
hmacSecret string
|
||||
ed25519PublicKey ed25519.PublicKey
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func newLicenseClient(ctx context.Context, svc *service.Container) (*licenseClient, error) {
|
||||
baseURL, _ := svc.Repo.Setting.Get(ctx, licenseServerURLSetting)
|
||||
if strings.TrimSpace(baseURL) == "" {
|
||||
baseURL = svc.Cfg.License.ServerURL
|
||||
}
|
||||
secret, _ := svc.Repo.Setting.Get(ctx, licenseHMACSecretSetting)
|
||||
if strings.TrimSpace(secret) == "" {
|
||||
secret = svc.Cfg.License.HMACSecret
|
||||
}
|
||||
publicKeyRaw, _ := svc.Repo.Setting.Get(ctx, "license.public_key")
|
||||
if strings.TrimSpace(publicKeyRaw) == "" {
|
||||
publicKeyRaw = svc.Cfg.License.PublicKey
|
||||
}
|
||||
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
|
||||
if baseURL == "" {
|
||||
return nil, errors.New("license server url not configured")
|
||||
}
|
||||
secret = strings.TrimSpace(secret)
|
||||
publicKey, err := parseLicenseEd25519PublicKey(publicKeyRaw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if secret == "" && len(publicKey) == 0 {
|
||||
return nil, errors.New("license public key or hmac secret not configured")
|
||||
}
|
||||
return &licenseClient{
|
||||
baseURL: baseURL,
|
||||
hmacSecret: secret,
|
||||
ed25519PublicKey: publicKey,
|
||||
httpClient: &http.Client{Timeout: 15 * time.Second},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *licenseClient) post(ctx context.Context, path string, payload any, out any) error {
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+path, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
return c.do(req, out)
|
||||
}
|
||||
|
||||
func (c *licenseClient) get(ctx context.Context, path string, out any) error {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+path, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return c.do(req, out)
|
||||
}
|
||||
|
||||
func (c *licenseClient) do(req *http.Request, out any) error {
|
||||
resp, err := c.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
data, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
var er struct {
|
||||
Error string `json:"error"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
_ = json.Unmarshal(data, &er)
|
||||
if er.Message != "" {
|
||||
return fmt.Errorf("license server: %s", er.Message)
|
||||
}
|
||||
if er.Error != "" {
|
||||
return fmt.Errorf("license server: %s", er.Error)
|
||||
}
|
||||
return fmt.Errorf("license server http %d", resp.StatusCode)
|
||||
}
|
||||
return json.Unmarshal(data, out)
|
||||
}
|
||||
|
||||
func (c *licenseClient) verifySigned(resp *licenseServerSignedResp) error {
|
||||
if strings.TrimSpace(resp.Signature) == "" {
|
||||
return errors.New("license server signature missing")
|
||||
}
|
||||
switch strings.ToLower(strings.TrimSpace(resp.SignatureAlg)) {
|
||||
case "ed25519":
|
||||
return c.verifyEd25519Signed(*resp)
|
||||
case "", "hmac", "hmac-sha256":
|
||||
return c.verifyHMACSigned(resp)
|
||||
default:
|
||||
return fmt.Errorf("unsupported license signature algorithm %q", resp.SignatureAlg)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *licenseClient) verifyEd25519Signed(resp licenseServerSignedResp) error {
|
||||
if len(c.ed25519PublicKey) != ed25519.PublicKeySize {
|
||||
return errors.New("license Ed25519 public key not configured")
|
||||
}
|
||||
signature, err := base64.StdEncoding.DecodeString(strings.TrimSpace(resp.Signature))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
payload, err := json.Marshal(licenseSignedPayload(resp))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ed25519.Verify(c.ed25519PublicKey, payload, signature) {
|
||||
return errors.New("license server signature verification failed")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *licenseClient) verifyHMACSigned(resp *licenseServerSignedResp) error {
|
||||
if c.hmacSecret == "" {
|
||||
return errors.New("license hmac secret not configured")
|
||||
}
|
||||
payload, err := json.Marshal(licenseSignedPayload(*resp))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
mac := hmac.New(sha256.New, []byte(c.hmacSecret))
|
||||
_, _ = mac.Write(payload)
|
||||
expected := hex.EncodeToString(mac.Sum(nil))
|
||||
if !hmac.Equal([]byte(expected), []byte(resp.Signature)) {
|
||||
if c.verifyLegacySigned(*resp) {
|
||||
resp.LegacySignature = true
|
||||
return nil
|
||||
}
|
||||
return errors.New("license server signature verification failed")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *licenseClient) verifyLegacySigned(resp licenseServerSignedResp) bool {
|
||||
unsigned := struct {
|
||||
Valid bool `json:"valid"`
|
||||
LicenseType string `json:"license_type"`
|
||||
ExpiryDate *string `json:"expiry_date"`
|
||||
MaxDevices int `json:"max_devices"`
|
||||
DaysRemaining *int `json:"days_remaining"`
|
||||
NextHeartbeat string `json:"next_heartbeat"`
|
||||
}{
|
||||
Valid: resp.Valid,
|
||||
LicenseType: resp.LicenseType,
|
||||
ExpiryDate: resp.ExpiryDate,
|
||||
MaxDevices: resp.MaxDevices,
|
||||
DaysRemaining: resp.DaysRemaining,
|
||||
NextHeartbeat: resp.NextHeartbeat,
|
||||
}
|
||||
payload, err := json.Marshal(unsigned)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
mac := hmac.New(sha256.New, []byte(c.hmacSecret))
|
||||
_, _ = mac.Write(payload)
|
||||
expected := hex.EncodeToString(mac.Sum(nil))
|
||||
return hmac.Equal([]byte(expected), []byte(resp.Signature))
|
||||
}
|
||||
|
||||
func parseLicenseEd25519PublicKey(encoded string) (ed25519.PublicKey, error) {
|
||||
encoded = strings.TrimSpace(encoded)
|
||||
if encoded == "" {
|
||||
return nil, nil
|
||||
}
|
||||
raw, err := base64.StdEncoding.DecodeString(encoded)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode license Ed25519 public key: %w", err)
|
||||
}
|
||||
if key, err := x509.ParsePKIXPublicKey(raw); err == nil {
|
||||
if publicKey, ok := key.(ed25519.PublicKey); ok && len(publicKey) == ed25519.PublicKeySize {
|
||||
return publicKey, nil
|
||||
}
|
||||
return nil, errors.New("license public key is not an Ed25519 PKIX public key")
|
||||
}
|
||||
if len(raw) == ed25519.PublicKeySize {
|
||||
return ed25519.PublicKey(raw), nil
|
||||
}
|
||||
return nil, errors.New("license public key must be base64 PKIX or raw 32-byte Ed25519 public key")
|
||||
}
|
||||
|
||||
func licenseSignedPayload(resp licenseServerSignedResp) any {
|
||||
return struct {
|
||||
Valid bool `json:"valid"`
|
||||
LicenseType string `json:"license_type"`
|
||||
ExpiryDate *string `json:"expiry_date"`
|
||||
MaxDevices int `json:"max_devices"`
|
||||
MaxUsers *int `json:"max_users"`
|
||||
DaysRemaining *int `json:"days_remaining"`
|
||||
NextHeartbeat string `json:"next_heartbeat"`
|
||||
}{
|
||||
Valid: resp.Valid,
|
||||
LicenseType: resp.LicenseType,
|
||||
ExpiryDate: resp.ExpiryDate,
|
||||
MaxDevices: resp.MaxDevices,
|
||||
MaxUsers: resp.MaxUsers,
|
||||
DaysRemaining: resp.DaysRemaining,
|
||||
NextHeartbeat: resp.NextHeartbeat,
|
||||
}
|
||||
}
|
||||
@@ -1,172 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
const (
|
||||
licenseHeartbeatInterval = 12 * time.Hour
|
||||
licenseHeartbeatCheckInterval = 30 * time.Minute
|
||||
licenseHeartbeatStartupDelay = 2 * time.Minute
|
||||
)
|
||||
|
||||
func refreshLicenseCapacityBestEffort(ctx context.Context, svc *service.Container) {
|
||||
if svc == nil || svc.Repo == nil || svc.Repo.Setting == nil {
|
||||
return
|
||||
}
|
||||
_, _, _ = maybeSendLicenseHeartbeat(ctx, svc, 0)
|
||||
}
|
||||
|
||||
// RunLicenseHeartbeatLoop keeps the license server aware of active deployments.
|
||||
// The loop checks periodically, but only sends when the last stored heartbeat is
|
||||
// older than licenseHeartbeatInterval.
|
||||
func RunLicenseHeartbeatLoop(ctx context.Context, svc *service.Container) {
|
||||
if svc == nil {
|
||||
return
|
||||
}
|
||||
run := func(interval time.Duration) {
|
||||
state, sent, err := maybeSendLicenseHeartbeat(ctx, svc, interval)
|
||||
logLicenseHeartbeatResult(svc, state, sent, err)
|
||||
}
|
||||
runStartup := func() {
|
||||
state, sent, err := maybeSendStartupLicenseHeartbeat(ctx, svc)
|
||||
logLicenseHeartbeatResult(svc, state, sent, err)
|
||||
}
|
||||
|
||||
timer := time.NewTimer(licenseHeartbeatStartupDelay)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-timer.C:
|
||||
runStartup()
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(licenseHeartbeatCheckInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
run(licenseHeartbeatInterval)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func logLicenseHeartbeatResult(svc *service.Container, state service.LicenseActivationState, sent bool, err error) {
|
||||
if svc == nil || svc.Log == nil {
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
svc.Log.Warn("license heartbeat failed", zap.Error(err))
|
||||
return
|
||||
}
|
||||
if sent {
|
||||
svc.Log.Info("license heartbeat sent", zap.String("device_id", state.DeviceID))
|
||||
}
|
||||
}
|
||||
|
||||
func maybeSendStartupLicenseHeartbeat(ctx context.Context, svc *service.Container) (service.LicenseActivationState, bool, error) {
|
||||
state, err := loadLicenseState(ctx, svc)
|
||||
if err != nil {
|
||||
return state, false, nil
|
||||
}
|
||||
if strings.TrimSpace(state.LicenseKey) == "" {
|
||||
return state, false, nil
|
||||
}
|
||||
return maybeSendLicenseHeartbeat(ctx, svc, 0)
|
||||
}
|
||||
|
||||
func maybeSendLicenseHeartbeat(ctx context.Context, svc *service.Container, interval time.Duration) (service.LicenseActivationState, bool, error) {
|
||||
state, err := loadLicenseState(ctx, svc)
|
||||
if err != nil {
|
||||
return state, false, nil
|
||||
}
|
||||
if !licenseHeartbeatEligible(state) {
|
||||
return state, false, nil
|
||||
}
|
||||
if !licenseHeartbeatDue(state, interval) {
|
||||
client, clientErr := newLicenseClient(ctx, svc)
|
||||
if clientErr != nil {
|
||||
return state, false, nil
|
||||
}
|
||||
deviceID, idErr := ensureLicenseDeviceID(ctx, svc, state.DeviceID)
|
||||
if idErr != nil {
|
||||
return state, false, idErr
|
||||
}
|
||||
refreshed, ok, requested, refreshErr := refreshLicenseServerStatus(ctx, client, state, deviceID)
|
||||
if refreshErr == nil && ok {
|
||||
state = refreshed
|
||||
_ = persistLicenseState(ctx, svc, state)
|
||||
}
|
||||
if !requested {
|
||||
return state, false, nil
|
||||
}
|
||||
}
|
||||
next, err := sendLicenseHeartbeat(ctx, svc)
|
||||
if err != nil {
|
||||
return state, false, err
|
||||
}
|
||||
return next, true, nil
|
||||
}
|
||||
|
||||
func licenseHeartbeatEligible(state service.LicenseActivationState) bool {
|
||||
return strings.TrimSpace(state.LicenseKey) != ""
|
||||
}
|
||||
|
||||
func licenseHeartbeatDue(state service.LicenseActivationState, interval time.Duration) bool {
|
||||
if interval <= 0 {
|
||||
return true
|
||||
}
|
||||
updatedAt := strings.TrimSpace(state.UpdatedAt)
|
||||
if updatedAt == "" {
|
||||
return true
|
||||
}
|
||||
for _, layout := range []string{time.RFC3339, "2006-01-02 15:04:05", "2006-01-02"} {
|
||||
if t, err := time.Parse(layout, updatedAt); err == nil {
|
||||
return time.Since(t) >= interval
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func sendLicenseHeartbeat(ctx context.Context, svc *service.Container) (service.LicenseActivationState, error) {
|
||||
client, err := newLicenseClient(ctx, svc)
|
||||
if err != nil {
|
||||
return service.LicenseActivationState{}, err
|
||||
}
|
||||
oldState, _ := loadLicenseState(ctx, svc)
|
||||
deviceID, err := ensureLicenseDeviceID(ctx, svc, oldState.DeviceID)
|
||||
if err != nil {
|
||||
return service.LicenseActivationState{}, err
|
||||
}
|
||||
deviceName, _ := svc.Repo.Setting.Get(ctx, licenseDeviceNameSetting)
|
||||
if strings.TrimSpace(deviceName) == "" {
|
||||
deviceName = defaultLicenseDeviceName()
|
||||
_ = svc.Repo.Setting.Set(ctx, licenseDeviceNameSetting, deviceName)
|
||||
}
|
||||
var upstream licenseServerSignedResp
|
||||
if err := client.post(ctx, "/api/v1/heartbeat", licenseHeartbeatPayload(oldState, deviceID, deviceName), &upstream); err != nil {
|
||||
return service.LicenseActivationState{}, err
|
||||
}
|
||||
if err := client.verifySigned(&upstream); err != nil {
|
||||
return service.LicenseActivationState{}, err
|
||||
}
|
||||
state := licenseStateFromSigned(upstream, deviceID, deviceName)
|
||||
state.LicenseKey = oldState.LicenseKey
|
||||
if refreshed, ok, _, refreshErr := refreshLicenseServerStatus(ctx, client, state, deviceID); refreshErr == nil && ok {
|
||||
refreshed.LicenseKey = state.LicenseKey
|
||||
state = refreshed
|
||||
}
|
||||
if err := persistLicenseState(ctx, svc, state); err != nil {
|
||||
return service.LicenseActivationState{}, err
|
||||
}
|
||||
return state, nil
|
||||
}
|
||||
@@ -1,211 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func refreshLicenseServerStatus(ctx context.Context, client *licenseClient, state service.LicenseActivationState, deviceID string) (service.LicenseActivationState, bool, bool, error) {
|
||||
var upstream licenseServerStatusResp
|
||||
if err := client.get(ctx, "/api/v1/status/"+url.PathEscape(deviceID), &upstream); err != nil {
|
||||
return state, false, false, err
|
||||
}
|
||||
if !upstream.Valid {
|
||||
state.Valid = false
|
||||
state.UpdatedAt = time.Now().Format(time.RFC3339)
|
||||
return state, false, upstream.HeartbeatRequested, nil
|
||||
}
|
||||
applyLicenseStatus(&state, upstream, deviceID)
|
||||
return state, true, upstream.HeartbeatRequested, nil
|
||||
}
|
||||
|
||||
func ensureLicenseDeviceID(ctx context.Context, svc *service.Container, candidate string) (string, error) {
|
||||
if strings.TrimSpace(candidate) != "" {
|
||||
return strings.TrimSpace(candidate), svc.Repo.Setting.Set(ctx, licenseDeviceIDSetting, strings.TrimSpace(candidate))
|
||||
}
|
||||
existing, err := svc.Repo.Setting.Get(ctx, licenseDeviceIDSetting)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if strings.TrimSpace(existing) != "" {
|
||||
return strings.TrimSpace(existing), nil
|
||||
}
|
||||
var buf [16]byte
|
||||
if _, err := rand.Read(buf[:]); err != nil {
|
||||
return "", err
|
||||
}
|
||||
id := "msgo-" + hex.EncodeToString(buf[:])
|
||||
return id, svc.Repo.Setting.Set(ctx, licenseDeviceIDSetting, id)
|
||||
}
|
||||
|
||||
func defaultLicenseDeviceName() string {
|
||||
host, _ := os.Hostname()
|
||||
if strings.TrimSpace(host) == "" {
|
||||
return "MediaStationGo Server"
|
||||
}
|
||||
return "MediaStationGo - " + host
|
||||
}
|
||||
|
||||
func licenseStateFromSigned(resp licenseServerSignedResp, deviceID, deviceName string) service.LicenseActivationState {
|
||||
expiry := ""
|
||||
if resp.ExpiryDate != nil {
|
||||
expiry = *resp.ExpiryDate
|
||||
}
|
||||
return service.LicenseActivationState{
|
||||
Valid: resp.Valid,
|
||||
LicenseType: resp.LicenseType,
|
||||
ExpiryDate: expiry,
|
||||
MaxDevices: resp.MaxDevices,
|
||||
MaxUsers: resp.MaxUsers,
|
||||
UnlimitedUsers: !resp.LegacySignature && resp.MaxUsers == nil,
|
||||
DaysRemaining: resp.DaysRemaining,
|
||||
NextHeartbeat: resp.NextHeartbeat,
|
||||
DeviceID: deviceID,
|
||||
DeviceName: deviceName,
|
||||
UpdatedAt: time.Now().Format(time.RFC3339),
|
||||
}
|
||||
}
|
||||
|
||||
func licenseHeartbeatPayload(state service.LicenseActivationState, deviceID, deviceName string) map[string]any {
|
||||
payload := map[string]any{
|
||||
"fingerprint": deviceID,
|
||||
"instance_id": deviceID,
|
||||
"device_name": deviceName,
|
||||
}
|
||||
if key := strings.TrimSpace(state.LicenseKey); key != "" {
|
||||
payload["key"] = key
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
func licenseStatusMaxUsers(state service.LicenseActivationState) any {
|
||||
active := state.Valid && !licenseStateExpired(state.ExpiryDate)
|
||||
if active {
|
||||
if state.UnlimitedUsers {
|
||||
return nil
|
||||
}
|
||||
if state.MaxUsers != nil && *state.MaxUsers > 0 {
|
||||
return *state.MaxUsers
|
||||
}
|
||||
return service.LicensedUserLimit
|
||||
}
|
||||
return service.OpenSourceUserLimit
|
||||
}
|
||||
|
||||
func applyLicenseStatus(state *service.LicenseActivationState, upstream licenseServerStatusResp, deviceID string) {
|
||||
state.Valid = upstream.Valid
|
||||
if upstream.LicenseType != nil {
|
||||
state.LicenseType = *upstream.LicenseType
|
||||
}
|
||||
if upstream.ExpiryDate != nil {
|
||||
state.ExpiryDate = *upstream.ExpiryDate
|
||||
} else {
|
||||
state.ExpiryDate = ""
|
||||
}
|
||||
if upstream.MaxDevices > 0 {
|
||||
state.MaxDevices = upstream.MaxDevices
|
||||
}
|
||||
state.MaxUsers = upstream.MaxUsers
|
||||
state.UnlimitedUsers = upstream.UnlimitedUsers
|
||||
state.DaysRemaining = upstream.DaysRemaining
|
||||
if upstream.DeviceName != "" {
|
||||
state.DeviceName = upstream.DeviceName
|
||||
}
|
||||
state.DeviceID = deviceID
|
||||
state.UpdatedAt = time.Now().Format(time.RFC3339)
|
||||
}
|
||||
|
||||
func persistLicenseState(ctx context.Context, svc *service.Container, state service.LicenseActivationState) error {
|
||||
data, err := json.Marshal(state)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.TrimSpace(state.DeviceID) != "" {
|
||||
_ = svc.Repo.Setting.Set(ctx, licenseDeviceIDSetting, strings.TrimSpace(state.DeviceID))
|
||||
}
|
||||
return svc.Repo.Setting.Set(ctx, service.LicenseSettingActivation, string(data))
|
||||
}
|
||||
|
||||
func loadLicenseState(ctx context.Context, svc *service.Container) (service.LicenseActivationState, error) {
|
||||
raw, err := svc.Repo.Setting.Get(ctx, service.LicenseSettingActivation)
|
||||
if err != nil || raw == "" {
|
||||
return service.LicenseActivationState{}, err
|
||||
}
|
||||
var state service.LicenseActivationState
|
||||
if err := json.Unmarshal([]byte(raw), &state); err != nil {
|
||||
return service.LicenseActivationState{}, err
|
||||
}
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func licenseActivationView(state service.LicenseActivationState) gin.H {
|
||||
updatedAt := state.UpdatedAt
|
||||
if strings.TrimSpace(updatedAt) == "" {
|
||||
updatedAt = time.Now().Format(time.RFC3339)
|
||||
}
|
||||
return gin.H{
|
||||
"id": state.DeviceID,
|
||||
"key_id": state.LicenseType,
|
||||
"key": maskLicenseKey(state.LicenseKey),
|
||||
"device_id": state.DeviceID,
|
||||
"device_name": state.DeviceName,
|
||||
"plan": state.LicenseType,
|
||||
"max_activations": state.MaxDevices,
|
||||
"max_users": state.MaxUsers,
|
||||
"unlimited_users": state.UnlimitedUsers,
|
||||
"expires_at": emptyAsNil(state.ExpiryDate),
|
||||
"valid": state.Valid && !licenseStateExpired(state.ExpiryDate),
|
||||
"heartbeat_at": updatedAt,
|
||||
"created_at": updatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func maskLicenseKey(key string) string {
|
||||
key = strings.TrimSpace(key)
|
||||
if key == "" {
|
||||
return ""
|
||||
}
|
||||
if len(key) <= 8 {
|
||||
return key
|
||||
}
|
||||
return key[:5] + "..." + key[len(key)-4:]
|
||||
}
|
||||
|
||||
func licenseStatusMessage(active bool, clientErr error) string {
|
||||
if active {
|
||||
return "已激活"
|
||||
}
|
||||
if clientErr != nil && !strings.Contains(clientErr.Error(), "not configured") {
|
||||
return clientErr.Error()
|
||||
}
|
||||
return "开源版:最多 20 个用户"
|
||||
}
|
||||
|
||||
func licenseStateExpired(expiry string) bool {
|
||||
if expiry == "" {
|
||||
return false
|
||||
}
|
||||
for _, layout := range []string{time.RFC3339, "2006-01-02 15:04:05", "2006-01-02"} {
|
||||
if t, err := time.Parse(layout, expiry); err == nil {
|
||||
return time.Now().After(t)
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func emptyAsNil(v string) any {
|
||||
if strings.TrimSpace(v) == "" {
|
||||
return nil
|
||||
}
|
||||
return v
|
||||
}
|
||||
@@ -1,445 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"crypto/ed25519"
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/config"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestLicenseStatusMaxUsersUsesLicensedLimit(t *testing.T) {
|
||||
maxUsers := 25
|
||||
state := service.LicenseActivationState{Valid: true, MaxUsers: &maxUsers}
|
||||
|
||||
if got := licenseStatusMaxUsers(state); got != maxUsers {
|
||||
t.Fatalf("expected licensed max users %d, got %#v", maxUsers, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseStatusMaxUsersAllowsUnlimited(t *testing.T) {
|
||||
state := service.LicenseActivationState{Valid: true, UnlimitedUsers: true}
|
||||
|
||||
if got := licenseStatusMaxUsers(state); got != nil {
|
||||
t.Fatalf("expected unlimited max users to be nil, got %#v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseStatusMaxUsersFallsBackToOpenSourceLimit(t *testing.T) {
|
||||
state := service.LicenseActivationState{}
|
||||
|
||||
if got := licenseStatusMaxUsers(state); got != service.OpenSourceUserLimit {
|
||||
t.Fatalf("expected open-source max users %d, got %#v", service.OpenSourceUserLimit, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyLicenseStatusReflectsEditedLimitAndClearsExpiry(t *testing.T) {
|
||||
maxUsers := 60
|
||||
licenseType := "subscription"
|
||||
state := service.LicenseActivationState{
|
||||
Valid: true,
|
||||
LicenseType: "enterprise",
|
||||
ExpiryDate: "2026-01-01",
|
||||
MaxDevices: 2,
|
||||
UnlimitedUsers: true,
|
||||
}
|
||||
|
||||
applyLicenseStatus(&state, licenseServerStatusResp{
|
||||
Valid: true,
|
||||
LicenseType: &licenseType,
|
||||
ExpiryDate: nil,
|
||||
MaxDevices: 5,
|
||||
MaxUsers: &maxUsers,
|
||||
UnlimitedUsers: false,
|
||||
DeviceName: "Edited Device",
|
||||
}, "device-1")
|
||||
|
||||
if !state.Valid || state.LicenseType != "subscription" || state.ExpiryDate != "" || state.MaxDevices != 5 {
|
||||
t.Fatalf("status fields were not fully refreshed: %+v", state)
|
||||
}
|
||||
if state.MaxUsers == nil || *state.MaxUsers != 60 || state.UnlimitedUsers {
|
||||
t.Fatalf("user limit was not refreshed from status: %+v", state)
|
||||
}
|
||||
if state.DeviceID != "device-1" || state.DeviceName != "Edited Device" {
|
||||
t.Fatalf("device fields were not refreshed: %+v", state)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyLicenseStatusReflectsUnlimitedUsers(t *testing.T) {
|
||||
maxUsers := 30
|
||||
state := service.LicenseActivationState{Valid: true, MaxUsers: &maxUsers}
|
||||
|
||||
applyLicenseStatus(&state, licenseServerStatusResp{
|
||||
Valid: true,
|
||||
MaxUsers: nil,
|
||||
UnlimitedUsers: true,
|
||||
}, "device-1")
|
||||
|
||||
if state.MaxUsers != nil || !state.UnlimitedUsers {
|
||||
t.Fatalf("unlimited status should clear previous finite user limit: %+v", state)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRefreshLicenseServerStatusReflectsEditedLimitAndHeartbeatRequest(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/api/v1/status/device-1" {
|
||||
t.Fatalf("unexpected path %s", r.URL.Path)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{
|
||||
"valid": true,
|
||||
"license_type": "subscription",
|
||||
"max_devices": 5,
|
||||
"max_users": 60,
|
||||
"unlimited_users": false,
|
||||
"device_name": "NAS",
|
||||
"heartbeat_requested": true,
|
||||
"is_active": true
|
||||
}`))
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
state := service.LicenseActivationState{Valid: true, UnlimitedUsers: true}
|
||||
client := &licenseClient{baseURL: upstream.URL, httpClient: upstream.Client()}
|
||||
|
||||
refreshed, ok, requested, err := refreshLicenseServerStatus(t.Context(), client, state, "device-1")
|
||||
if err != nil {
|
||||
t.Fatalf("refresh status: %v", err)
|
||||
}
|
||||
if !ok || !requested {
|
||||
t.Fatalf("expected valid status with requested heartbeat, ok=%v requested=%v", ok, requested)
|
||||
}
|
||||
if refreshed.MaxUsers == nil || *refreshed.MaxUsers != 60 || refreshed.UnlimitedUsers {
|
||||
t.Fatalf("edited user limit was not reflected: %+v", refreshed)
|
||||
}
|
||||
if refreshed.MaxDevices != 5 || refreshed.DeviceName != "NAS" {
|
||||
t.Fatalf("server status fields were not applied: %+v", refreshed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseHeartbeatPayloadIncludesStoredLicenseKey(t *testing.T) {
|
||||
payload := licenseHeartbeatPayload(service.LicenseActivationState{
|
||||
LicenseKey: "MS-ABCD-EFGH-JKLM-NPQR",
|
||||
}, "device-1", "NAS")
|
||||
|
||||
if payload["fingerprint"] != "device-1" || payload["instance_id"] != "device-1" || payload["device_name"] != "NAS" {
|
||||
t.Fatalf("heartbeat identity payload is wrong: %#v", payload)
|
||||
}
|
||||
if payload["key"] != "MS-ABCD-EFGH-JKLM-NPQR" {
|
||||
t.Fatalf("heartbeat should include stored license key for server-side backfill: %#v", payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseClientVerifiesEd25519Signature(t *testing.T) {
|
||||
publicKey, privateKey, err := ed25519.GenerateKey(rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp := licenseServerSignedResp{
|
||||
Valid: true,
|
||||
LicenseType: "subscription",
|
||||
MaxDevices: 2,
|
||||
NextHeartbeat: time.Now().Add(time.Hour).Format(time.RFC3339),
|
||||
SignatureAlg: "ed25519",
|
||||
}
|
||||
resp.Signature = signLicenseTestPayloadEd25519(privateKey, resp)
|
||||
|
||||
client := &licenseClient{ed25519PublicKey: publicKey}
|
||||
if err := client.verifySigned(&resp); err != nil {
|
||||
t.Fatalf("verify ed25519 signature: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseClientRejectsEd25519WithoutPublicKey(t *testing.T) {
|
||||
_, privateKey, err := ed25519.GenerateKey(rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp := licenseServerSignedResp{
|
||||
Valid: true,
|
||||
LicenseType: "subscription",
|
||||
MaxDevices: 2,
|
||||
NextHeartbeat: time.Now().Add(time.Hour).Format(time.RFC3339),
|
||||
SignatureAlg: "ed25519",
|
||||
}
|
||||
resp.Signature = signLicenseTestPayloadEd25519(privateKey, resp)
|
||||
|
||||
client := &licenseClient{}
|
||||
if err := client.verifySigned(&resp); err == nil || !strings.Contains(err.Error(), "public key") {
|
||||
t.Fatalf("expected missing public key error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseHeartbeatDueUsesTwelveHourWindow(t *testing.T) {
|
||||
state := service.LicenseActivationState{
|
||||
Valid: true,
|
||||
UpdatedAt: time.Now().Add(-11 * time.Hour).Format(time.RFC3339),
|
||||
}
|
||||
if licenseHeartbeatDue(state, 12*time.Hour) {
|
||||
t.Fatalf("heartbeat should not be due before interval")
|
||||
}
|
||||
|
||||
state.UpdatedAt = time.Now().Add(-13 * time.Hour).Format(time.RFC3339)
|
||||
if !licenseHeartbeatDue(state, 12*time.Hour) {
|
||||
t.Fatalf("heartbeat should be due after interval")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartupLicenseHeartbeatIgnoresTwelveHourWindow(t *testing.T) {
|
||||
heartbeatCount := 0
|
||||
maxUsers := 40
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/api/v1/heartbeat":
|
||||
heartbeatCount++
|
||||
resp := licenseServerSignedResp{
|
||||
Valid: true,
|
||||
LicenseType: "subscription",
|
||||
MaxDevices: 2,
|
||||
MaxUsers: &maxUsers,
|
||||
NextHeartbeat: time.Now().Add(time.Hour).Format(time.RFC3339),
|
||||
}
|
||||
resp.Signature = signLicenseTestPayload("test-secret", resp)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(resp)
|
||||
case "/api/v1/status/device-1":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{
|
||||
"valid": true,
|
||||
"license_type": "subscription",
|
||||
"max_devices": 2,
|
||||
"max_users": 40,
|
||||
"unlimited_users": false,
|
||||
"device_name": "NAS",
|
||||
"is_active": true
|
||||
}`))
|
||||
default:
|
||||
t.Fatalf("unexpected upstream path %s", r.URL.Path)
|
||||
}
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc := newLicenseHandlerTestService(t)
|
||||
if err := svc.Repo.Setting.Set(t.Context(), licenseServerURLSetting, upstream.URL); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.Repo.Setting.Set(t.Context(), licenseHMACSecretSetting, "test-secret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
state := service.LicenseActivationState{
|
||||
Valid: true,
|
||||
LicenseKey: "MS-ABCD-EFGH-JKLM-NPQR",
|
||||
DeviceID: "device-1",
|
||||
DeviceName: "NAS",
|
||||
UpdatedAt: time.Now().Format(time.RFC3339),
|
||||
}
|
||||
if err := persistLicenseState(t.Context(), svc, state); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
refreshed, sent, err := maybeSendStartupLicenseHeartbeat(t.Context(), svc)
|
||||
if err != nil {
|
||||
t.Fatalf("startup heartbeat: %v", err)
|
||||
}
|
||||
if !sent || heartbeatCount != 1 {
|
||||
t.Fatalf("startup heartbeat should be sent once, sent=%v count=%d", sent, heartbeatCount)
|
||||
}
|
||||
if refreshed.MaxUsers == nil || *refreshed.MaxUsers != 40 {
|
||||
t.Fatalf("startup heartbeat should refresh licensed user capacity, got %+v", refreshed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartupLicenseHeartbeatSkipsStateWithoutStoredKey(t *testing.T) {
|
||||
heartbeatCount := 0
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
heartbeatCount++
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc := newLicenseHandlerTestService(t)
|
||||
if err := svc.Repo.Setting.Set(t.Context(), licenseServerURLSetting, upstream.URL); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.Repo.Setting.Set(t.Context(), licenseHMACSecretSetting, "test-secret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := persistLicenseState(t.Context(), svc, service.LicenseActivationState{
|
||||
Valid: true,
|
||||
DeviceID: "device-1",
|
||||
UpdatedAt: time.Now().Format(time.RFC3339),
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, sent, err := maybeSendStartupLicenseHeartbeat(t.Context(), svc)
|
||||
if err != nil {
|
||||
t.Fatalf("startup heartbeat should skip without error, got %v", err)
|
||||
}
|
||||
if sent || heartbeatCount != 0 {
|
||||
t.Fatalf("startup heartbeat without stored license key should be skipped, sent=%v count=%d", sent, heartbeatCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseHeartbeatEligibleRequiresActivationState(t *testing.T) {
|
||||
if licenseHeartbeatEligible(service.LicenseActivationState{DeviceID: "device-only"}) {
|
||||
t.Fatalf("device id alone should not trigger automatic license heartbeat")
|
||||
}
|
||||
if !licenseHeartbeatEligible(service.LicenseActivationState{LicenseKey: "MS-KEY"}) {
|
||||
t.Fatalf("stored license key should trigger automatic license heartbeat")
|
||||
}
|
||||
if licenseHeartbeatEligible(service.LicenseActivationState{Valid: true}) {
|
||||
t.Fatalf("valid state without stored license key should not trigger automatic license heartbeat")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseStatusSkipsUnlicensedHeartbeatWithDefaultServer(t *testing.T) {
|
||||
upstreamCalls := 0
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
upstreamCalls++
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc := newLicenseHandlerTestService(t)
|
||||
svc.Cfg.License.ServerURL = upstream.URL
|
||||
svc.Cfg.License.HMACSecret = "test-secret"
|
||||
|
||||
router := gin.New()
|
||||
router.GET("/license/status", licenseStatusHandler(svc))
|
||||
req := httptest.NewRequest(http.MethodGet, "/license/status", nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
if upstreamCalls != 0 {
|
||||
t.Fatalf("unlicensed status should not contact license server, got %d calls", upstreamCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseActivateBindsServerInstanceNotBrowserFingerprint(t *testing.T) {
|
||||
var upstreamFingerprint string
|
||||
maxUsers := 60
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var payload map[string]any
|
||||
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode upstream payload: %v", err)
|
||||
}
|
||||
upstreamFingerprint, _ = payload["fingerprint"].(string)
|
||||
resp := licenseServerSignedResp{
|
||||
Valid: true,
|
||||
LicenseType: "subscription",
|
||||
MaxDevices: 3,
|
||||
MaxUsers: &maxUsers,
|
||||
NextHeartbeat: time.Now().Add(time.Hour).Format(time.RFC3339),
|
||||
}
|
||||
resp.Signature = signLicenseTestPayload("test-secret", resp)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
svc := newLicenseHandlerTestService(t)
|
||||
if err := svc.Repo.Setting.Set(t.Context(), licenseServerURLSetting, upstream.URL); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := svc.Repo.Setting.Set(t.Context(), licenseHMACSecretSetting, "test-secret"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.POST("/license/activate", licenseActivateHandler(svc))
|
||||
req := httptest.NewRequest(http.MethodPost, "/license/activate", strings.NewReader(`{
|
||||
"key": "MS-ABCD-EFGH-JKLM-NPQR",
|
||||
"device_id": "browser-fingerprint",
|
||||
"device_name": ""
|
||||
}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("activate status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
if upstreamFingerprint == "" || upstreamFingerprint == "browser-fingerprint" || !strings.HasPrefix(upstreamFingerprint, "msgo-") {
|
||||
t.Fatalf("activation should use server-generated msgo id, got %q", upstreamFingerprint)
|
||||
}
|
||||
stored, err := svc.Repo.Setting.Get(t.Context(), licenseDeviceIDSetting)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if stored != upstreamFingerprint {
|
||||
t.Fatalf("stored device id = %q, upstream fingerprint = %q", stored, upstreamFingerprint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewLicenseClientRequiresSignatureVerifier(t *testing.T) {
|
||||
svc := newLicenseHandlerTestService(t)
|
||||
if err := svc.Repo.Setting.Set(t.Context(), licenseServerURLSetting, "http://127.0.0.1:8001"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := newLicenseClient(t.Context(), svc); err == nil || !strings.Contains(err.Error(), "public key or hmac secret") {
|
||||
t.Fatalf("expected missing signature verifier error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func newLicenseHandlerTestService(t *testing.T) *service.Container {
|
||||
t.Helper()
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.Setting{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return &service.Container{
|
||||
Cfg: &config.Config{},
|
||||
Repo: repository.New(db),
|
||||
}
|
||||
}
|
||||
|
||||
func signLicenseTestPayload(secret string, resp licenseServerSignedResp) string {
|
||||
unsigned := struct {
|
||||
Valid bool `json:"valid"`
|
||||
LicenseType string `json:"license_type"`
|
||||
ExpiryDate *string `json:"expiry_date"`
|
||||
MaxDevices int `json:"max_devices"`
|
||||
MaxUsers *int `json:"max_users"`
|
||||
DaysRemaining *int `json:"days_remaining"`
|
||||
NextHeartbeat string `json:"next_heartbeat"`
|
||||
}{
|
||||
Valid: resp.Valid,
|
||||
LicenseType: resp.LicenseType,
|
||||
ExpiryDate: resp.ExpiryDate,
|
||||
MaxDevices: resp.MaxDevices,
|
||||
MaxUsers: resp.MaxUsers,
|
||||
DaysRemaining: resp.DaysRemaining,
|
||||
NextHeartbeat: resp.NextHeartbeat,
|
||||
}
|
||||
payload, _ := json.Marshal(unsigned)
|
||||
mac := hmac.New(sha256.New, []byte(secret))
|
||||
_, _ = mac.Write(payload)
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
func signLicenseTestPayloadEd25519(privateKey ed25519.PrivateKey, resp licenseServerSignedResp) string {
|
||||
payload, _ := json.Marshal(licenseSignedPayload(resp))
|
||||
return base64.StdEncoding.EncodeToString(ed25519.Sign(privateKey, payload))
|
||||
}
|
||||
@@ -31,7 +31,6 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
libs = service.FilterDeprecatedNativeCloudLibraries(libs)
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("all") == "1")
|
||||
if !includeHidden {
|
||||
@@ -44,9 +43,6 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
libs = filtered
|
||||
} else {
|
||||
libs = service.FilterMergedCloudAutoCategoryLibraries(libs)
|
||||
libs = service.NormalizeCloudLibraryDisplayNames(libs)
|
||||
}
|
||||
c.JSON(http.StatusOK, libs)
|
||||
}
|
||||
@@ -63,23 +59,18 @@ func getLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
libs := service.FilterDeprecatedNativeCloudLibraries([]model.Library{*lib})
|
||||
if len(libs) == 0 {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
role, _ := c.Get(middleware.CtxUserRole)
|
||||
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("all") == "1")
|
||||
if includeHidden {
|
||||
c.JSON(http.StatusOK, service.NormalizeCloudLibraryDisplayNames(libs)[0])
|
||||
return
|
||||
if !includeHidden {
|
||||
libs := service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, []model.Library{*lib})
|
||||
if len(libs) == 0 || !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, libs[0], mediaVisibilityForRequest(c, svc)) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, libs[0])
|
||||
} else {
|
||||
c.JSON(http.StatusOK, lib)
|
||||
}
|
||||
libs = service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, libs)
|
||||
if len(libs) == 0 || !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, libs[0], mediaVisibilityForRequest(c, svc)) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, libs[0])
|
||||
}
|
||||
}
|
||||
|
||||
@@ -152,11 +143,6 @@ func updateLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
func deleteLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), id); err == nil && lib != nil {
|
||||
if _, ok := service.ParseCloudLibraryMount(lib.Path); ok && svc.Scan != nil {
|
||||
_ = svc.Scan.CancelCloudScan(id)
|
||||
}
|
||||
}
|
||||
if err := svc.Media.DeleteLibrary(c.Request.Context(), id); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
|
||||
@@ -23,47 +23,8 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
if _, ok := service.ParseCloudLibraryMount(lib.Path); ok {
|
||||
status, started, startErr := svc.Scan.StartCloudLibraryScan(id, true)
|
||||
if startErr != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": startErr.Error()})
|
||||
return
|
||||
}
|
||||
if !started {
|
||||
c.JSON(http.StatusAccepted, gin.H{
|
||||
"library_id": id,
|
||||
"queued": true,
|
||||
"cloud": true,
|
||||
"already_running": true,
|
||||
"stage": status.Stage,
|
||||
"state": status.State,
|
||||
"message": "该云盘媒体库正在后台扫描,请在任务面板查看进度",
|
||||
"estimate_message": "页面关闭不会中断扫描",
|
||||
})
|
||||
return
|
||||
}
|
||||
task := startScanHTTPTask(svc, "云盘扫描队列", lib.Name, lib.Path)
|
||||
if svc.WSHub != nil {
|
||||
svc.WSHub.Publish("scan", gin.H{
|
||||
"library_id": id,
|
||||
"cloud": true,
|
||||
"queued": true,
|
||||
"stage": "queued",
|
||||
"message": "云盘扫描已加入后台队列,会递归扫描并自动加入媒体库",
|
||||
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
||||
})
|
||||
}
|
||||
finishHTTPTask(task, nil, "queued", "云盘扫描已加入后台队列", map[string]int64{"queued": 1}, nil)
|
||||
c.JSON(http.StatusAccepted, gin.H{
|
||||
"library_id": id,
|
||||
"visited": 0,
|
||||
"added": 0,
|
||||
"updated": 0,
|
||||
"probed": 0,
|
||||
"queued": true,
|
||||
"cloud": true,
|
||||
"message": "云盘扫描已在后台运行,发现的媒体会自动加入当前媒体库;若已开启自动刮削,会在扫描后补齐元数据",
|
||||
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
||||
})
|
||||
// 云盘扫描已随网盘后端移除;此处保留空分支以兼容既有客户端。
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "cloud 网盘后端已移除,无法扫描云盘媒体库"})
|
||||
return
|
||||
}
|
||||
finishScan, ok := svc.Scan.TryBeginLocalScan(id)
|
||||
|
||||
@@ -110,86 +110,6 @@ func TestListLibrariesHidesAdultDirectoriesUnlessAdminRequestsAll(t *testing.T)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListLibrariesIncludeHiddenNormalizesCloudDisplayNames(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
cloud := model.Library{Name: "OpenList · 国产剧", Path: service.BuildCloudLibraryPath("openlist", "/国产剧", "/国产剧"), Type: "tv", Enabled: true}
|
||||
if err := repos.Library.Create(t.Context(), &cloud); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc := &service.Container{
|
||||
Repo: repos,
|
||||
Media: service.NewMediaService(&config.Config{}, zap.NewNop(), repos),
|
||||
}
|
||||
|
||||
all := requestLibraries(t, svc, "admin", "admin", "/api/libraries?include_hidden=1")
|
||||
if len(all) != 1 {
|
||||
t.Fatalf("include_hidden list = %#v, want one library", all)
|
||||
}
|
||||
if all[0].Name != "国产剧" {
|
||||
t.Fatalf("cloud display name = %q, want stripped directory name", all[0].Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListLibrariesShowsAutoCategoryLibraries(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.Library{}, &model.Media{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
root := model.Library{Name: "OpenList", Path: "cloud://openlist", Type: "movie", Enabled: true}
|
||||
auto := model.Library{Name: "欧美剧", Path: service.BuildCloudAutoCategoryLibraryPath("openlist", "电视剧/欧美剧"), Type: "tv", Enabled: true}
|
||||
for _, lib := range []*model.Library{&root, &auto} {
|
||||
if err := repos.Library.Create(t.Context(), lib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
svc := &service.Container{
|
||||
Repo: repos,
|
||||
Media: service.NewMediaService(&config.Config{}, zap.NewNop(), repos),
|
||||
}
|
||||
|
||||
all := requestLibraries(t, svc, "admin", "admin", "/api/libraries?include_hidden=1")
|
||||
if len(all) != 2 {
|
||||
t.Fatalf("include_hidden list = %#v, want root plus auto category library", all)
|
||||
}
|
||||
ids := map[string]bool{}
|
||||
for _, lib := range all {
|
||||
ids[lib.ID] = true
|
||||
}
|
||||
if !ids[root.ID] || !ids[auto.ID] {
|
||||
t.Fatalf("include_hidden list = %#v, want root %s and auto category %s", all, root.ID, auto.ID)
|
||||
}
|
||||
|
||||
visible := requestLibraries(t, svc, "user-1", "user", "/api/libraries")
|
||||
if len(visible) != 2 {
|
||||
t.Fatalf("visible list = %#v, want root plus auto category library", visible)
|
||||
}
|
||||
ids = map[string]bool{}
|
||||
for _, lib := range visible {
|
||||
ids[lib.ID] = true
|
||||
}
|
||||
if !ids[root.ID] || !ids[auto.ID] {
|
||||
t.Fatalf("visible list = %#v, want root %s and auto category %s", visible, root.ID, auto.ID)
|
||||
}
|
||||
|
||||
got := requestLibrary(t, svc, "user-1", "user", "/api/libraries/"+auto.ID, auto.ID)
|
||||
if got.ID != auto.ID || got.Name != auto.Name {
|
||||
t.Fatalf("auto category detail = %#v, want accessible category library", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetLibraryAllowsEmptyLibrary(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
@@ -336,41 +256,6 @@ func TestListLibrarySeriesDoesNotTruncateLargeEpisodeLibraries(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestScanLibraryHandlerSurfacesCloudQueueStartFailure(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.Library{}, &model.LibraryRoot{}, &model.Media{}, &model.Setting{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
lib := model.Library{Name: "旧夸克云盘", Path: service.BuildCloudLibraryPath(service.LegacyQuarkProvider, "archive", "archive"), Type: "movie", Enabled: true}
|
||||
if err := repos.Library.Create(t.Context(), &lib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
log := zap.NewNop()
|
||||
svc := &service.Container{
|
||||
Log: log,
|
||||
Repo: repos,
|
||||
Scan: service.NewScannerService(&config.Config{}, log, repos, service.NewHub(log), nil, nil),
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
c.Params = gin.Params{{Key: "id", Value: lib.ID}}
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/api/libraries/"+lib.ID+"/scan", nil)
|
||||
|
||||
scanLibraryHandler(svc)(c)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status=%d body=%s, want bad request", w.Code, w.Body.String())
|
||||
}
|
||||
if !strings.Contains(w.Body.String(), "deprecated") {
|
||||
t.Fatalf("body=%s, want cloud queue start error", w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestScrapeOptionsFromRequestPreservesEpisodeImagesFalse(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
@@ -1,188 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/url"
|
||||
"path"
|
||||
"strings"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func enrichAndPersistDownloadRows(ctx context.Context, svc *service.Container, rows []model.DownloadTask) {
|
||||
for i := range rows {
|
||||
before := rows[i]
|
||||
enrichDownloadRow(ctx, svc, &rows[i])
|
||||
updates := downloadTaskMetadataUpdates(before, rows[i])
|
||||
if len(updates) == 0 {
|
||||
continue
|
||||
}
|
||||
if err := svc.Repo.DB.WithContext(ctx).Model(&model.DownloadTask{}).Where("id = ?", rows[i].ID).Updates(updates).Error; err != nil {
|
||||
svc.Log.Debug("download artwork backfill failed", zap.String("id", rows[i].ID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func enrichDownloadRow(ctx context.Context, svc *service.Container, row *model.DownloadTask) {
|
||||
if strings.TrimSpace(row.Title) == "" {
|
||||
row.Title = downloadDisplayTitle(row.URL)
|
||||
}
|
||||
if !downloadTaskNeedsMetadata(*row) {
|
||||
return
|
||||
}
|
||||
meta := enrichDownloadTaskMeta(ctx, svc, service.DownloadTaskMeta{
|
||||
Title: row.Title,
|
||||
PosterURL: row.PosterURL,
|
||||
BackdropURL: row.BackdropURL,
|
||||
Overview: row.Overview,
|
||||
OriginalName: row.OriginalName,
|
||||
OriginalLanguage: row.OriginalLanguage,
|
||||
Year: row.Year,
|
||||
Rating: row.Rating,
|
||||
Genres: row.Genres,
|
||||
}, firstNonEmptyString(row.Title, row.URL), row.MediaType)
|
||||
row.Title = firstNonEmptyString(row.Title, meta.Title)
|
||||
row.PosterURL = meta.PosterURL
|
||||
row.BackdropURL = meta.BackdropURL
|
||||
row.Overview = meta.Overview
|
||||
row.OriginalName = meta.OriginalName
|
||||
row.OriginalLanguage = meta.OriginalLanguage
|
||||
row.Year = meta.Year
|
||||
row.Rating = meta.Rating
|
||||
row.Genres = meta.Genres
|
||||
}
|
||||
|
||||
func downloadTaskNeedsMetadata(row model.DownloadTask) bool {
|
||||
return strings.TrimSpace(row.PosterURL) == "" || strings.TrimSpace(row.BackdropURL) == "" ||
|
||||
strings.TrimSpace(row.Overview) == "" || row.Year <= 0 || row.Rating <= 0 ||
|
||||
strings.TrimSpace(row.Genres) == "" || strings.TrimSpace(row.OriginalName) == "" ||
|
||||
strings.TrimSpace(row.OriginalLanguage) == ""
|
||||
}
|
||||
|
||||
func downloadTaskMetadataUpdates(before, after model.DownloadTask) map[string]any {
|
||||
updates := map[string]any{}
|
||||
if before.Title != after.Title {
|
||||
updates["title"] = after.Title
|
||||
}
|
||||
if before.PosterURL != after.PosterURL {
|
||||
updates["poster_url"] = after.PosterURL
|
||||
}
|
||||
if before.BackdropURL != after.BackdropURL {
|
||||
updates["backdrop_url"] = after.BackdropURL
|
||||
}
|
||||
if before.Overview != after.Overview {
|
||||
updates["overview"] = after.Overview
|
||||
}
|
||||
if before.OriginalName != after.OriginalName {
|
||||
updates["original_name"] = after.OriginalName
|
||||
}
|
||||
if before.OriginalLanguage != after.OriginalLanguage {
|
||||
updates["original_language"] = after.OriginalLanguage
|
||||
}
|
||||
if before.Year != after.Year {
|
||||
updates["year"] = after.Year
|
||||
}
|
||||
if before.Rating != after.Rating {
|
||||
updates["rating"] = after.Rating
|
||||
}
|
||||
if before.Genres != after.Genres {
|
||||
updates["genres"] = after.Genres
|
||||
}
|
||||
return updates
|
||||
}
|
||||
|
||||
func enrichDownloadTorrentViews(ctx context.Context, svc *service.Container, views []service.DownloadTorrentView) {
|
||||
cache := map[string]displayMetadata{}
|
||||
for i := range views {
|
||||
if strings.TrimSpace(views[i].PosterURL) != "" && strings.TrimSpace(views[i].BackdropURL) != "" {
|
||||
continue
|
||||
}
|
||||
query := firstNonEmptyString(views[i].Title, views[i].Name)
|
||||
cacheKey, _ := metadataSearchQuery(query)
|
||||
meta, ok := cache[cacheKey]
|
||||
if !ok {
|
||||
meta = lookupDisplayMetadata(ctx, svc, query, "", "")
|
||||
cache[cacheKey] = meta
|
||||
}
|
||||
if strings.TrimSpace(views[i].PosterURL) == "" {
|
||||
views[i].PosterURL = meta.PosterURL
|
||||
}
|
||||
if strings.TrimSpace(views[i].BackdropURL) == "" {
|
||||
views[i].BackdropURL = meta.BackdropURL
|
||||
}
|
||||
if strings.TrimSpace(views[i].Overview) == "" {
|
||||
views[i].Overview = meta.Overview
|
||||
}
|
||||
if (views[i].Title == "" || views[i].Title == views[i].Name) && meta.Title != "" {
|
||||
views[i].Title = meta.Title
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func enrichDownloadTaskMeta(ctx context.Context, svc *service.Container, meta service.DownloadTaskMeta, fallbackTitle, mediaType string) service.DownloadTaskMeta {
|
||||
if strings.TrimSpace(meta.Title) == "" {
|
||||
meta.Title = strings.TrimSpace(downloadDisplayTitle(fallbackTitle))
|
||||
}
|
||||
// 仅当展示字段都齐时才跳过查询(含媒体富通知所需的年份/评分/类型/原始信息)。
|
||||
if strings.TrimSpace(meta.PosterURL) != "" && strings.TrimSpace(meta.BackdropURL) != "" &&
|
||||
strings.TrimSpace(meta.Overview) != "" && meta.Year > 0 && meta.Rating > 0 &&
|
||||
strings.TrimSpace(meta.Genres) != "" && strings.TrimSpace(meta.OriginalName) != "" &&
|
||||
strings.TrimSpace(meta.OriginalLanguage) != "" {
|
||||
return meta
|
||||
}
|
||||
found := lookupDisplayMetadata(ctx, svc, meta.Title, fallbackTitle, mediaType)
|
||||
if strings.TrimSpace(meta.Title) == "" {
|
||||
meta.Title = found.Title
|
||||
}
|
||||
if strings.TrimSpace(meta.PosterURL) == "" {
|
||||
meta.PosterURL = found.PosterURL
|
||||
}
|
||||
if strings.TrimSpace(meta.BackdropURL) == "" {
|
||||
meta.BackdropURL = found.BackdropURL
|
||||
}
|
||||
if strings.TrimSpace(meta.Overview) == "" {
|
||||
meta.Overview = found.Overview
|
||||
}
|
||||
if strings.TrimSpace(meta.OriginalName) == "" {
|
||||
meta.OriginalName = found.OriginalName
|
||||
}
|
||||
if strings.TrimSpace(meta.OriginalLanguage) == "" {
|
||||
meta.OriginalLanguage = found.OriginalLanguage
|
||||
}
|
||||
if meta.Year <= 0 && found.Year > 0 {
|
||||
meta.Year = found.Year
|
||||
}
|
||||
if meta.Rating <= 0 && found.Rating > 0 {
|
||||
meta.Rating = found.Rating
|
||||
}
|
||||
if strings.TrimSpace(meta.Genres) == "" {
|
||||
meta.Genres = found.Genres
|
||||
}
|
||||
return meta
|
||||
}
|
||||
|
||||
func downloadDisplayTitle(raw string) string {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return ""
|
||||
}
|
||||
if u, err := url.Parse(raw); err == nil {
|
||||
if dn := strings.TrimSpace(u.Query().Get("dn")); dn != "" {
|
||||
if decoded, err := url.QueryUnescape(dn); err == nil && strings.TrimSpace(decoded) != "" {
|
||||
return strings.TrimSpace(decoded)
|
||||
}
|
||||
return dn
|
||||
}
|
||||
if u.Host != "" {
|
||||
base := path.Base(u.Path)
|
||||
if base != "." && base != "/" && base != "" {
|
||||
base = strings.TrimSuffix(base, path.Ext(base))
|
||||
return strings.TrimSpace(base)
|
||||
}
|
||||
}
|
||||
}
|
||||
return raw
|
||||
}
|
||||
@@ -1,100 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func enrichSubscriptionArtwork(ctx context.Context, svc *service.Container, sub *model.Subscription) {
|
||||
if svc == nil || sub == nil {
|
||||
return
|
||||
}
|
||||
// 已有图片且媒体展示字段也齐全时才跳过;否则仍需查一次补 年份/评分/类型 等
|
||||
// 富通知字段(老订阅只存了 poster/overview 的情况)。
|
||||
if strings.TrimSpace(sub.PosterURL) != "" && strings.TrimSpace(sub.BackdropURL) != "" &&
|
||||
sub.Year > 0 && sub.Rating > 0 && strings.TrimSpace(sub.Genres) != "" {
|
||||
return
|
||||
}
|
||||
meta := lookupDisplayMetadata(ctx, svc, sub.Name, sub.Filter, sub.MediaType)
|
||||
if meta.Title == "" && meta.PosterURL == "" && meta.BackdropURL == "" && meta.Overview == "" {
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(sub.Source) == "" {
|
||||
sub.Source = meta.Source
|
||||
}
|
||||
if strings.TrimSpace(sub.PosterURL) == "" {
|
||||
sub.PosterURL = meta.PosterURL
|
||||
}
|
||||
if strings.TrimSpace(sub.BackdropURL) == "" {
|
||||
sub.BackdropURL = meta.BackdropURL
|
||||
}
|
||||
if strings.TrimSpace(sub.Overview) == "" {
|
||||
sub.Overview = meta.Overview
|
||||
}
|
||||
if strings.TrimSpace(sub.OriginalName) == "" {
|
||||
sub.OriginalName = meta.OriginalName
|
||||
}
|
||||
if strings.TrimSpace(sub.OriginalLanguage) == "" {
|
||||
sub.OriginalLanguage = meta.OriginalLanguage
|
||||
}
|
||||
if sub.Year <= 0 && meta.Year > 0 {
|
||||
sub.Year = meta.Year
|
||||
}
|
||||
if sub.Rating <= 0 && meta.Rating > 0 {
|
||||
sub.Rating = meta.Rating
|
||||
}
|
||||
if strings.TrimSpace(sub.Genres) == "" {
|
||||
sub.Genres = meta.Genres
|
||||
}
|
||||
}
|
||||
|
||||
func enrichAndPersistSubscriptions(ctx context.Context, svc *service.Container, items []model.Subscription) {
|
||||
for i := range items {
|
||||
before := items[i]
|
||||
enrichSubscriptionArtwork(ctx, svc, &items[i])
|
||||
updates := subscriptionMetadataUpdates(before, items[i])
|
||||
if len(updates) == 0 {
|
||||
continue
|
||||
}
|
||||
if err := svc.Repo.DB.WithContext(ctx).Model(&model.Subscription{}).Where("id = ?", items[i].ID).Updates(updates).Error; err != nil {
|
||||
svc.Log.Debug("subscription artwork backfill failed", zap.String("id", items[i].ID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func subscriptionMetadataUpdates(before, after model.Subscription) map[string]any {
|
||||
updates := map[string]any{}
|
||||
if before.Source != after.Source {
|
||||
updates["source"] = after.Source
|
||||
}
|
||||
if before.PosterURL != after.PosterURL {
|
||||
updates["poster_url"] = after.PosterURL
|
||||
}
|
||||
if before.BackdropURL != after.BackdropURL {
|
||||
updates["backdrop_url"] = after.BackdropURL
|
||||
}
|
||||
if before.Overview != after.Overview {
|
||||
updates["overview"] = after.Overview
|
||||
}
|
||||
if before.OriginalName != after.OriginalName {
|
||||
updates["original_name"] = after.OriginalName
|
||||
}
|
||||
if before.OriginalLanguage != after.OriginalLanguage {
|
||||
updates["original_language"] = after.OriginalLanguage
|
||||
}
|
||||
if before.Year != after.Year {
|
||||
updates["year"] = after.Year
|
||||
}
|
||||
if before.Rating != after.Rating {
|
||||
updates["rating"] = after.Rating
|
||||
}
|
||||
if before.Genres != after.Genres {
|
||||
updates["genres"] = after.Genres
|
||||
}
|
||||
return updates
|
||||
}
|
||||
@@ -1,27 +0,0 @@
|
||||
// Package handler — notification test endpoint.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type notifyTestReq struct {
|
||||
Title string `json:"title" binding:"required"`
|
||||
Body string `json:"body" binding:"required"`
|
||||
}
|
||||
|
||||
func notifyTestHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req notifyTestReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
svc.Notifier.Send(c.Request.Context(), req.Title, req.Body, "test")
|
||||
c.JSON(http.StatusOK, gin.H{"message": "notification dispatched"})
|
||||
}
|
||||
}
|
||||
@@ -1,77 +0,0 @@
|
||||
// Package handler — notify channel CRUD + per-channel test endpoint.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func listNotifyChannelsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
rows, err := svc.NotifyChannels.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if rows == nil {
|
||||
c.JSON(http.StatusOK, []struct{}{})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, rows)
|
||||
}
|
||||
}
|
||||
|
||||
func createNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var in service.ChannelInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
row, err := svc.NotifyChannels.Create(c.Request.Context(), in)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, row)
|
||||
}
|
||||
}
|
||||
|
||||
func updateNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var in service.ChannelInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
row, err := svc.NotifyChannels.Update(c.Request.Context(), c.Param("id"), in)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, row)
|
||||
}
|
||||
}
|
||||
|
||||
func deleteNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.NotifyChannels.Delete(c.Request.Context(), c.Param("id")); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
func testNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.NotifyChannels.Test(c.Request.Context(), c.Param("id")); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"message": "test sent"})
|
||||
}
|
||||
}
|
||||
@@ -1,196 +0,0 @@
|
||||
// Package handler — 通知渠道管理 HTTP 端点。
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// NotifyHandler 处理通知渠道的 CRUD 操作。
|
||||
type NotifyHandler struct {
|
||||
svc *service.Container
|
||||
log *zap.Logger
|
||||
}
|
||||
|
||||
// NewNotifyHandler 创建通知渠道处理器。
|
||||
func NewNotifyHandler(svc *service.Container, log *zap.Logger) *NotifyHandler {
|
||||
return &NotifyHandler{svc: svc, log: log}
|
||||
}
|
||||
|
||||
// notifyCreateRequest 创建通知渠道请求体。
|
||||
type notifyCreateRequest struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Type string `json:"type" binding:"required,oneof=telegram wechat bark webhook email"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Config map[string]string `json:"config" binding:"required"`
|
||||
Events []string `json:"events"`
|
||||
}
|
||||
|
||||
// notifyUpdateRequest 更新通知渠道请求体。
|
||||
type notifyUpdateRequest struct {
|
||||
Name string `json:"name"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Config map[string]string `json:"config"`
|
||||
Events []string `json:"events"`
|
||||
}
|
||||
|
||||
// Create 创建新的通知渠道。
|
||||
func (h *NotifyHandler) Create(c *gin.Context) {
|
||||
var req notifyCreateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
|
||||
// 验证配置
|
||||
if err := h.svc.Notify.ValidateChannelConfig(req.Type, req.Config); err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, "配置验证失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 加密配置
|
||||
configJSON, _ := json.Marshal(req.Config)
|
||||
configStr := string(configJSON)
|
||||
if h.svc.Crypto != nil {
|
||||
configStr = h.svc.Crypto.Encrypt(configStr)
|
||||
}
|
||||
|
||||
// 序列化事件列表
|
||||
eventsJSON, _ := json.Marshal(req.Events)
|
||||
eventsStr := string(eventsJSON)
|
||||
|
||||
channel := &model.NotifyChannel{
|
||||
Name: req.Name,
|
||||
Type: req.Type,
|
||||
Enabled: req.Enabled,
|
||||
Config: configStr,
|
||||
Events: eventsStr,
|
||||
}
|
||||
|
||||
if err := h.svc.Repo.NotifyChannel.Create(ctx, channel); err != nil {
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "创建失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
Success(c, channel)
|
||||
}
|
||||
|
||||
// List 返回所有通知渠道。
|
||||
func (h *NotifyHandler) List(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
channels, err := h.svc.Repo.NotifyChannel.List(ctx)
|
||||
if err != nil {
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "查询失败")
|
||||
return
|
||||
}
|
||||
Success(c, channels)
|
||||
}
|
||||
|
||||
// Get 返回指定通知渠道详情。
|
||||
func (h *NotifyHandler) Get(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
ctx := c.Request.Context()
|
||||
channel, err := h.svc.Repo.NotifyChannel.FindByID(ctx, id)
|
||||
if err != nil || channel == nil {
|
||||
Error(c, http.StatusNotFound, ErrNotFound, "通知渠道不存在")
|
||||
return
|
||||
}
|
||||
Success(c, channel)
|
||||
}
|
||||
|
||||
// Update 更新通知渠道。
|
||||
func (h *NotifyHandler) Update(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
var req notifyUpdateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
channel, err := h.svc.Repo.NotifyChannel.FindByID(ctx, id)
|
||||
if err != nil || channel == nil {
|
||||
Error(c, http.StatusNotFound, ErrNotFound, "通知渠道不存在")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name != "" {
|
||||
channel.Name = req.Name
|
||||
}
|
||||
if req.Enabled != nil {
|
||||
channel.Enabled = *req.Enabled
|
||||
}
|
||||
|
||||
// 更新配置
|
||||
if len(req.Config) > 0 {
|
||||
if err := h.svc.Notify.ValidateChannelConfig(channel.Type, req.Config); err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrInvalidParams, "配置验证失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
configJSON, _ := json.Marshal(req.Config)
|
||||
configStr := string(configJSON)
|
||||
if h.svc.Crypto != nil {
|
||||
configStr = h.svc.Crypto.Encrypt(configStr)
|
||||
}
|
||||
channel.Config = configStr
|
||||
}
|
||||
|
||||
// 更新事件列表
|
||||
if req.Events != nil {
|
||||
eventsJSON, _ := json.Marshal(req.Events)
|
||||
channel.Events = string(eventsJSON)
|
||||
}
|
||||
|
||||
if err := h.svc.Repo.NotifyChannel.Update(ctx, channel); err != nil {
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "更新失败")
|
||||
return
|
||||
}
|
||||
|
||||
Success(c, channel)
|
||||
}
|
||||
|
||||
// Delete 删除通知渠道。
|
||||
func (h *NotifyHandler) Delete(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
ctx := c.Request.Context()
|
||||
|
||||
_, err := h.svc.Repo.NotifyChannel.FindByID(ctx, id)
|
||||
if err != nil {
|
||||
Error(c, http.StatusNotFound, ErrNotFound, "通知渠道不存在")
|
||||
return
|
||||
}
|
||||
|
||||
if delErr := h.svc.Repo.NotifyChannel.Delete(ctx, id); delErr != nil {
|
||||
Error(c, http.StatusInternalServerError, ErrInternal, "删除失败")
|
||||
return
|
||||
}
|
||||
|
||||
SuccessWithMessage(c, "已删除", nil)
|
||||
}
|
||||
|
||||
// Test 发送测试通知。
|
||||
func (h *NotifyHandler) Test(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
ctx := c.Request.Context()
|
||||
|
||||
if err := h.svc.Notify.SendTest(ctx, id); err != nil {
|
||||
Error(c, http.StatusBadRequest, ErrExternal, "测试通知发送失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
SuccessWithMessage(c, "测试通知已发送", nil)
|
||||
}
|
||||
|
||||
// GetTypes 返回支持的通知渠道类型列表。
|
||||
func (h *NotifyHandler) GetTypes(c *gin.Context) {
|
||||
types := h.svc.Notify.GetProviderTypes()
|
||||
Success(c, types)
|
||||
}
|
||||
@@ -275,63 +275,6 @@ func TestScopedPlaybackTokenCannotStreamAnotherMedia(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestScopedPlaybackTokenCanFollowCloudRedirectForSameMedia(t *testing.T) {
|
||||
router, svc, _ := newPlaybackScopeTestRouter(t)
|
||||
user, err := svc.Repo.User.FindByID(t.Context(), "user-1")
|
||||
if err != nil || user == nil {
|
||||
t.Fatalf("find user: %v", err)
|
||||
}
|
||||
playToken, err := svc.Auth.IssueExternalPlaybackToken(user, "media-1", 2*60*60)
|
||||
if err != nil {
|
||||
t.Fatalf("issue playback token: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://nas.local/api/stream/media-1?token="+url.QueryEscape(playToken), nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusFound {
|
||||
t.Fatalf("stream status = %d body=%s, want 302", w.Code, w.Body.String())
|
||||
}
|
||||
loc := w.Header().Get("Location")
|
||||
if !strings.Contains(loc, "/api/cloud/play/openlist?") ||
|
||||
!strings.Contains(loc, "media_id=media-1") ||
|
||||
!strings.Contains(loc, "token="+url.QueryEscape(playToken)) {
|
||||
t.Fatalf("redirect Location should carry scoped token and media_id, got %q", loc)
|
||||
}
|
||||
|
||||
req = httptest.NewRequest(http.MethodGet, loc, nil)
|
||||
w = httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code == http.StatusForbidden || w.Code == http.StatusUnauthorized {
|
||||
t.Fatalf("cloud redirect rejected scoped token: status=%d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
if w.Code != http.StatusServiceUnavailable {
|
||||
t.Fatalf("cloud redirect status = %d body=%s, want storage service fallback 503", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestScopedPlaybackTokenCannotRetargetCloudRef(t *testing.T) {
|
||||
router, svc, _ := newPlaybackScopeTestRouter(t)
|
||||
user, err := svc.Repo.User.FindByID(t.Context(), "user-1")
|
||||
if err != nil || user == nil {
|
||||
t.Fatalf("find user: %v", err)
|
||||
}
|
||||
playToken, err := svc.Auth.IssueExternalPlaybackToken(user, "media-1", 2*60*60)
|
||||
if err != nil {
|
||||
t.Fatalf("issue playback token: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://nas.local/api/cloud/play/openlist?ref=other&media_id=media-1&token="+url.QueryEscape(playToken), nil)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusForbidden {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestScopedPlaybackTokenCannotCallRegularAPI(t *testing.T) {
|
||||
router, svc, _ := newPlaybackScopeTestRouter(t)
|
||||
api := router.Group("/api")
|
||||
@@ -464,6 +407,5 @@ func newPlaybackScopeTestRouter(t *testing.T) (*gin.Engine, *service.Container,
|
||||
api.GET("/playback/:id/external-url", externalURLHandler(svc))
|
||||
api.GET("/playback/:id/external-players", externalPlayersHandler(svc))
|
||||
api.GET("/stream/:id", streamHandler(svc))
|
||||
api.GET("/cloud/play/:type", cloudPlayHandler(svc))
|
||||
return router, svc, cfg.Secrets.JWTSecret
|
||||
}
|
||||
|
||||
@@ -2,19 +2,13 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type recycleBatchReq struct {
|
||||
MediaIDs []string `json:"media_ids"`
|
||||
}
|
||||
|
||||
func deleteMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Media.SoftDelete(c.Request.Context(), c.Param("id")); err != nil {
|
||||
@@ -25,17 +19,6 @@ func deleteMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func listRecycleHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
items, err := svc.Media.ListRecycleBin(c.Request.Context(), 200)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||
}
|
||||
}
|
||||
|
||||
func restoreMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Media.RestoreDeleted(c.Request.Context(), c.Param("id")); err != nil {
|
||||
@@ -46,22 +29,6 @@ func restoreMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func restoreMediaBatchHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req recycleBatchReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
applied, errorsOut := runRecycleBatch(c, compactManualScrapeIDs(req.MediaIDs), svc.Media.RestoreDeleted)
|
||||
if applied == 0 && len(errorsOut) > 0 {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": strings.Join(errorsOut, "\n")})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"applied": applied, "errors": errorsOut})
|
||||
}
|
||||
}
|
||||
|
||||
func purgeMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Media.PurgeDeleted(c.Request.Context(), c.Param("id")); err != nil {
|
||||
@@ -71,35 +38,3 @@ func purgeMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
func purgeMediaBatchHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req recycleBatchReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
applied, errorsOut := runRecycleBatch(c, compactManualScrapeIDs(req.MediaIDs), svc.Media.PurgeDeleted)
|
||||
if applied == 0 && len(errorsOut) > 0 {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": strings.Join(errorsOut, "\n")})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"applied": applied, "errors": errorsOut})
|
||||
}
|
||||
}
|
||||
|
||||
func runRecycleBatch(c *gin.Context, ids []string, action func(context.Context, string) error) (int, []string) {
|
||||
if len(ids) == 0 {
|
||||
return 0, []string{"media_ids required"}
|
||||
}
|
||||
applied := 0
|
||||
errorsOut := make([]string, 0)
|
||||
for _, id := range ids {
|
||||
if err := action(c.Request.Context(), id); err != nil {
|
||||
errorsOut = append(errorsOut, id+": "+err.Error())
|
||||
continue
|
||||
}
|
||||
applied++
|
||||
}
|
||||
return applied, errorsOut
|
||||
}
|
||||
|
||||
@@ -1,84 +0,0 @@
|
||||
// Package handler — batch repair+rescrape endpoint.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// repairAndRescrapeAllHandler 触发"全库修复+重刮"流程:先从媒体路径中的
|
||||
// {tmdb-N}/{bangumi-N} 等占位符回填缺失的外部 ID, 再对所有媒体库重刮一遍。
|
||||
//
|
||||
// 路由: POST /api/admin/media/repair-rescrape (需 admin)
|
||||
// 异步执行, 立即返回 202;通过 WS hub "scrape" topic 推送进度。
|
||||
func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
options, err := scrapeOptionsFromRequest(c, true)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
|
||||
return
|
||||
}
|
||||
task := startScrapeHTTPTask(svc, "全库修复并重刮", "", "")
|
||||
go func(options service.ScrapeOptions) {
|
||||
result, err := svc.RepairAndRescrapeAllLibraries(context.Background(), options)
|
||||
metrics := map[string]int64{
|
||||
"repaired": int64(result.Repaired),
|
||||
"reclassified": int64(result.Reclassified),
|
||||
"libraries": int64(result.Libraries),
|
||||
"matched": int64(result.Matched),
|
||||
"processed": int64(result.Processed),
|
||||
"errors": int64(result.Errors),
|
||||
"reset": int64(result.Reset),
|
||||
}
|
||||
stage := "completed"
|
||||
message := "全库修复并重刮完成"
|
||||
if err != nil {
|
||||
stage = "scrape"
|
||||
message = "全库修复并重刮失败"
|
||||
}
|
||||
finishHTTPTask(task, err, stage, message, metrics, nil)
|
||||
}(options)
|
||||
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
|
||||
}
|
||||
}
|
||||
|
||||
// repairAndRescrapeLibraryHandler 触发"单库修复+重刮":只对路径参数指定的
|
||||
// 媒体库回填占位符外部 ID 并重刮, 不影响其它库。
|
||||
//
|
||||
// 路由: POST /api/admin/libraries/:id/repair-rescrape (需 admin)
|
||||
// 异步执行, 立即返回 202;通过 WS hub "scrape" topic 推送进度。
|
||||
func repairAndRescrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libraryID := c.Param("id")
|
||||
options, err := scrapeOptionsFromRequest(c, true)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
|
||||
return
|
||||
}
|
||||
task := startScrapeHTTPTask(svc, "媒体库修复并重刮", "", "")
|
||||
go func(options service.ScrapeOptions) {
|
||||
result, err := svc.RepairAndRescrapeLibrary(context.Background(), libraryID, options)
|
||||
metrics := map[string]int64{
|
||||
"repaired": int64(result.Repaired),
|
||||
"reclassified": int64(result.Reclassified),
|
||||
"libraries": int64(result.Libraries),
|
||||
"matched": int64(result.Matched),
|
||||
"processed": int64(result.Processed),
|
||||
"errors": int64(result.Errors),
|
||||
"reset": int64(result.Reset),
|
||||
}
|
||||
stage := "completed"
|
||||
message := "媒体库修复并重刮完成"
|
||||
if err != nil {
|
||||
stage = "scrape"
|
||||
message = "媒体库修复并重刮失败"
|
||||
}
|
||||
finishHTTPTask(task, err, stage, message, metrics, nil)
|
||||
}(options)
|
||||
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
|
||||
}
|
||||
}
|
||||
@@ -14,17 +14,10 @@ func registerAdminRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.
|
||||
admin.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret), middleware.AdminRequired())
|
||||
registerAdminUserRoutes(admin, svc)
|
||||
registerAdminPermissionRoutes(admin, svc)
|
||||
registerAdminStorageRoutes(admin, svc)
|
||||
registerAdminCloudRoutes(admin, svc)
|
||||
registerAdminDownloadClientRoutes(admin, svc)
|
||||
registerAdminSystemRoutes(admin, svc)
|
||||
registerAdminBackupRoutes(admin, svc)
|
||||
registerAdminNotificationRoutes(admin, svc)
|
||||
registerAdminTelegramRoutes(admin, svc)
|
||||
registerAdminOrganizerRoutes(admin, svc)
|
||||
registerAdminRepairRoutes(admin, svc)
|
||||
registerAdminAPIConfigRoutes(admin, svc)
|
||||
registerAdminSchedulerRoutes(admin, svc)
|
||||
registerAdminRecognitionWordRoutes(admin, svc)
|
||||
}
|
||||
|
||||
@@ -48,39 +41,7 @@ func registerAdminPermissionRoutes(admin *gin.RouterGroup, svc *service.Containe
|
||||
admin.POST("/users/:id/permissions/reset", resetUserPermissionsHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminStorageRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.GET("/storage/status", listStorageConfigsHandler(svc))
|
||||
admin.GET("/storage/:type", getStorageConfigHandler(svc))
|
||||
admin.PUT("/storage/:type", saveStorageConfigHandler(svc))
|
||||
admin.POST("/storage/:type/test", testStorageConfigHandler(svc))
|
||||
admin.POST("/storage/:type/logout", logoutStorageConfigHandler(svc))
|
||||
admin.POST("/storage/:type/upload-local", storageUploadLocalHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminCloudRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.POST("/cloud/scan-all", cloudScanAllHandler(svc))
|
||||
admin.POST("/cloud/scan/cancel", cloudScanCancelHandler(svc))
|
||||
admin.GET("/cloud/scan/status", cloudScanStatusHandler(svc))
|
||||
admin.GET("/cloud/:type/list", cloudListHandler(svc))
|
||||
admin.POST("/cloud/:type/mkdir", cloudMkdirHandler(svc))
|
||||
admin.PUT("/cloud/:type/rename", cloudRenameHandler(svc))
|
||||
admin.POST("/cloud/:type/import", cloudImportHandler(svc))
|
||||
admin.POST("/cloud/:type/mount", cloudMountHandler(svc))
|
||||
admin.POST("/cloud/:type/qr/start", cloud115QRStartHandler(svc))
|
||||
admin.POST("/cloud/:type/qr/poll", cloud115QRPollHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminDownloadClientRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.GET("/download/clients", listDownloadClientsHandler(svc))
|
||||
admin.POST("/download/clients", createDownloadClientHandler(svc))
|
||||
admin.PUT("/download/clients/:id", updateDownloadClientHandler(svc))
|
||||
admin.DELETE("/download/clients/:id", deleteDownloadClientHandler(svc))
|
||||
admin.POST("/download/clients/:id/test", testDownloadClientHandler(svc))
|
||||
admin.GET("/download/aria2/stats", aria2StatsHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminSystemRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.POST("/system/scheduler/:name/trigger", schedulerTriggerHandler(svc))
|
||||
admin.GET("/system/update", systemUpdateStatusHandler(svc))
|
||||
admin.POST("/system/update/check", systemUpdateCheckHandler(svc))
|
||||
admin.POST("/system/update/apply", systemUpdateApplyHandler(svc))
|
||||
@@ -93,22 +54,6 @@ func registerAdminBackupRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.POST("/backups/restore", restoreBackupHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminNotificationRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.POST("/notify/test", notifyTestHandler(svc))
|
||||
admin.GET("/notify/channels", listNotifyChannelsHandler(svc))
|
||||
admin.POST("/notify/channels", createNotifyChannelHandler(svc))
|
||||
admin.PUT("/notify/channels/:id", updateNotifyChannelHandler(svc))
|
||||
admin.DELETE("/notify/channels/:id", deleteNotifyChannelHandler(svc))
|
||||
admin.POST("/notify/channels/:id/test", testNotifyChannelHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminTelegramRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.GET("/telegram/webhook", telegramGetWebhookHandler(svc))
|
||||
admin.POST("/telegram/webhook", telegramSetWebhookHandler(svc))
|
||||
admin.POST("/telegram/polling/start", telegramStartPollingHandler(svc))
|
||||
admin.POST("/telegram/polling/stop", telegramStopPollingHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminOrganizerRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.POST("/media/:id/organize", organizeMediaHandler(svc))
|
||||
admin.POST("/libraries/:id/organize", organizeLibraryHandler(svc))
|
||||
@@ -116,11 +61,6 @@ func registerAdminOrganizerRoutes(admin *gin.RouterGroup, svc *service.Container
|
||||
admin.POST("/organize/source", organizeDirectoryHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminRepairRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.POST("/media/repair-rescrape", repairAndRescrapeAllHandler(svc))
|
||||
admin.POST("/libraries/:id/repair-rescrape", repairAndRescrapeLibraryHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminAPIConfigRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.GET("/api-configs", listAPIConfigsHandler(svc))
|
||||
admin.GET("/api-configs/:provider", getAPIConfigHandler(svc))
|
||||
@@ -129,11 +69,6 @@ func registerAdminAPIConfigRoutes(admin *gin.RouterGroup, svc *service.Container
|
||||
admin.POST("/api-configs/:provider/test", testAPIConfigHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminSchedulerRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.GET("/scheduler", schedulerStatusHandler(svc))
|
||||
admin.POST("/scheduler/:name/run", schedulerRunHandler(svc))
|
||||
}
|
||||
|
||||
func registerAdminRecognitionWordRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||
admin.GET("/recognition-words", getRecognitionWordsHandler(svc))
|
||||
admin.PUT("/recognition-words", saveRecognitionWordsHandler(svc))
|
||||
|
||||
@@ -26,17 +26,9 @@ func TestAdminRouteSurfacesAreRegistered(t *testing.T) {
|
||||
for _, want := range []string{
|
||||
"GET /api/admin/users",
|
||||
"GET /api/admin/users/:id/permissions",
|
||||
"GET /api/admin/storage/status",
|
||||
"GET /api/admin/cloud/:type/list",
|
||||
"GET /api/admin/download/clients",
|
||||
"POST /api/admin/system/scheduler/:name/trigger",
|
||||
"POST /api/admin/backups",
|
||||
"GET /api/admin/notify/channels",
|
||||
"GET /api/admin/telegram/webhook",
|
||||
"GET /api/admin/organize/sources",
|
||||
"POST /api/admin/media/repair-rescrape",
|
||||
"GET /api/admin/api-configs",
|
||||
"POST /api/admin/scheduler/:name/run",
|
||||
} {
|
||||
if !routes[want] {
|
||||
t.Fatalf("%s route is not registered", want)
|
||||
|
||||
@@ -19,26 +19,15 @@ func registerAuthenticatedRoutes(api *gin.RouterGroup, cfg *config.Config, svc *
|
||||
registerAuthedMediaRoutes(authed, svc)
|
||||
registerAuthedPlaybackAndProxyRoutes(authed, svc)
|
||||
registerAuthedCollectionRoutes(authed, svc)
|
||||
registerAuthedDownloadRoutes(authed, svc)
|
||||
registerAuthedSubscriptionRoutes(authed, svc)
|
||||
registerAuthedStatsDiscoveryAndAIRoutes(authed, svc)
|
||||
registerAuthedFileRoutes(authed, svc)
|
||||
registerAuthedDLNARoutes(authed, svc)
|
||||
registerAuthedSTRMRoutes(authed, svc)
|
||||
registerAuthedDuplicateRoutes(authed, svc)
|
||||
registerAuthedSiteRoutes(authed, svc)
|
||||
registerAuthedRecycleAndRealtimeRoutes(authed, svc)
|
||||
registerAuthedSchedulerRoutes(authed, svc)
|
||||
registerAuthedUISurfaceRoutes(authed, svc)
|
||||
registerAuthedSearchRoutes(authed, svc)
|
||||
registerAuthedSystemExtraRoutes(authed, svc)
|
||||
registerAuthedStatsExtraRoutes(authed, svc)
|
||||
registerAuthedSitesExtraRoutes(authed, svc)
|
||||
registerAuthedSubscriptionExtraRoutes(authed, svc)
|
||||
registerAuthedPlaylistExtraRoutes(authed, svc)
|
||||
registerAuthedDLNAControlRoutes(authed, svc)
|
||||
registerAuthedFavoriteAndMediaActionRoutes(authed, svc)
|
||||
registerAuthedPlaybackExtraRoutes(authed, svc)
|
||||
registerAuthedDownloadOpsRoutes(authed, svc)
|
||||
registerAuthedAssistantRoutes(authed, svc)
|
||||
}
|
||||
|
||||
@@ -14,10 +14,6 @@ func registerAuthedUserAndLicenseRoutes(authed *gin.RouterGroup, svc *service.Co
|
||||
authed.POST("/me/logout", logoutHandler(svc))
|
||||
|
||||
authed.GET("/auth/permissions", getMyPermissionsHandler(svc))
|
||||
|
||||
authed.GET("/license/status", middleware.AdminRequired(), licenseStatusHandler(svc))
|
||||
authed.POST("/license/activate", middleware.AdminRequired(), licenseActivateHandler(svc))
|
||||
authed.POST("/license/heartbeat", middleware.AdminRequired(), licenseHeartbeatHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedLibraryRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
@@ -65,11 +61,6 @@ func registerAuthedPlaybackAndProxyRoutes(authed *gin.RouterGroup, svc *service.
|
||||
authed.GET("/hls/:id/:seg", hlsSegmentHandler(svc))
|
||||
authed.DELETE("/hls/:id", stopTranscodeHandler(svc))
|
||||
|
||||
authed.GET("/cloud/play/:type", cloudPlayHandler(svc))
|
||||
authed.HEAD("/cloud/play/:type", cloudPlayHandler(svc))
|
||||
|
||||
authed.GET("/img/cloud/:type", cloudArtworkProxyHandler(svc))
|
||||
authed.HEAD("/img/cloud/:type", cloudArtworkProxyHandler(svc))
|
||||
authed.GET("/img", imageProxyHandler(svc))
|
||||
}
|
||||
|
||||
|
||||
@@ -16,18 +16,8 @@ func registerAuthedUISurfaceRoutes(authed *gin.RouterGroup, svc *service.Contain
|
||||
authed.DELETE("/watch-history", historyDeleteHandler(svc))
|
||||
authed.DELETE("/watch-history/:id", historyDeleteOneHandler(svc))
|
||||
|
||||
authed.GET("/discover/sections", requirePermission(svc, "can_view_discover"), discoverSectionsHandler(svc))
|
||||
authed.GET("/discover/feed", requirePermission(svc, "can_view_discover"), discoverFeedHandler(svc))
|
||||
|
||||
authed.GET("/system/info", systemInfoHandler(svc))
|
||||
authed.GET("/system/status", systemStatusHandler(svc))
|
||||
authed.GET("/system/scheduler", systemSchedulerHandler(svc))
|
||||
|
||||
authed.GET("/stats/overview", statsOverviewHandler(svc))
|
||||
authed.GET("/stats/trend", statsTrendHandler(svc))
|
||||
authed.GET("/stats/top-content", statsTopContentHandler(svc))
|
||||
authed.GET("/stats/libraries", statsLibrariesHandler(svc))
|
||||
authed.GET("/stats/monitor", statsMonitorHandler(svc))
|
||||
|
||||
authed.GET("/play-profiles", listPlayProfilesHandler(svc))
|
||||
authed.POST("/play-profiles", createPlayProfileHandler(svc))
|
||||
@@ -40,7 +30,6 @@ func registerAuthedSearchRoutes(authed *gin.RouterGroup, svc *service.Container)
|
||||
authed.GET("/search", searchUnifiedHandler(svc))
|
||||
authed.GET("/search/advanced", searchAdvancedHandler(svc))
|
||||
authed.GET("/search/tmdb", searchTMDbHandler(svc))
|
||||
authed.GET("/search/sites", searchSitesHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedSystemExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
@@ -55,16 +44,6 @@ func registerAuthedStatsExtraRoutes(authed *gin.RouterGroup, svc *service.Contai
|
||||
authed.POST("/stats/play", statsPlayHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedSitesExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/sites/:id/resource", requirePermission(svc, "can_manage_sites"), siteResourceHandler(svc))
|
||||
authed.GET("/sites/:id/userdata", requirePermission(svc, "can_manage_sites"), siteUserdataHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedSubscriptionExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.PUT("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), updateSubscriptionHandler(svc))
|
||||
authed.POST("/subscriptions/:id/search", requirePermission(svc, "can_manage_subscriptions"), searchSubscriptionHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedPlaylistExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.POST("/playlists/:id/reorder", reorderPlaylistHandler(svc))
|
||||
authed.DELETE("/playlists/:id/items/by-id/:item_id", deletePlaylistItemByIDHandler(svc))
|
||||
@@ -94,24 +73,3 @@ func registerAuthedPlaybackExtraRoutes(authed *gin.RouterGroup, svc *service.Con
|
||||
authed.GET("/playback/:id/external-url", externalURLHandler(svc))
|
||||
authed.GET("/playback/transcode/:job_id/status", transcodeStatusHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedDownloadOpsRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.POST("/download/:id/pause", requirePermission(svc, "can_manage_downloads"), downloadPauseHandler(svc))
|
||||
authed.POST("/download/:id/resume", requirePermission(svc, "can_manage_downloads"), downloadResumeHandler(svc))
|
||||
authed.POST("/download/:id/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeOneHandler(svc))
|
||||
authed.POST("/download/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeAllHandler(svc))
|
||||
authed.POST("/download/sync", requirePermission(svc, "can_manage_downloads"), downloadSyncHandler(svc))
|
||||
authed.POST("/download/start-auto-sync", requirePermission(svc, "can_manage_downloads"), downloadAutoSyncHandler(svc))
|
||||
authed.GET("/download/tasks", requirePermission(svc, "can_manage_downloads"), downloadTasksAliasHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedAssistantRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/admin/assistant/sessions", listAssistantSessionsHandler(svc))
|
||||
authed.POST("/admin/assistant/sessions", createAssistantSessionHandler(svc))
|
||||
authed.GET("/admin/assistant/session/:id", getAssistantSessionHandler(svc))
|
||||
authed.DELETE("/admin/assistant/session/:id", deleteAssistantSessionHandler(svc))
|
||||
authed.POST("/admin/assistant/chat", assistantChatHandler(svc))
|
||||
authed.POST("/admin/assistant/execute", assistantExecuteHandler(svc))
|
||||
authed.POST("/admin/assistant/undo/:op_id", assistantUndoHandler(svc))
|
||||
authed.GET("/admin/assistant/history", assistantHistoryHandler(svc))
|
||||
}
|
||||
|
||||
@@ -7,35 +7,6 @@ import (
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func registerAuthedDownloadRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/downloads", requirePermission(svc, "can_manage_downloads"), listDownloadsHandler(svc))
|
||||
authed.POST("/downloads", requirePermission(svc, "can_manage_downloads"), addDownloadHandler(svc))
|
||||
authed.DELETE("/downloads/:hash", requirePermission(svc, "can_manage_downloads"), deleteDownloadHandler(svc))
|
||||
authed.POST("/downloads/relocate", requirePermission(svc, "can_manage_downloads"), relocateDownloadHandler(svc))
|
||||
authed.POST("/downloads/reload", requirePermission(svc, "can_manage_downloads"), reloadDownloadConfigHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedSubscriptionRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionsHandler(svc))
|
||||
authed.GET("/subscriptions/history", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionHistoryHandler(svc))
|
||||
authed.POST("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), createSubscriptionHandler(svc))
|
||||
authed.DELETE("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), deleteSubscriptionHandler(svc))
|
||||
authed.POST("/subscriptions/:id/restore", requirePermission(svc, "can_manage_subscriptions"), restoreSubscriptionHandler(svc))
|
||||
authed.POST("/subscriptions/:id/run", requirePermission(svc, "can_manage_subscriptions"), runSubscriptionHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedStatsDiscoveryAndAIRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/stats", statsHandler(svc))
|
||||
authed.GET("/tasks", middleware.AdminRequired(), tasksHandler(svc))
|
||||
|
||||
authed.GET("/discover/trending", requirePermission(svc, "can_view_discover"), trendingHandler(svc))
|
||||
authed.GET("/discover/popular", requirePermission(svc, "can_view_discover"), popularHandler(svc))
|
||||
|
||||
authed.GET("/ai/status", requirePermission(svc, "can_use_ai"), aiStatusHandler(svc))
|
||||
authed.POST("/ai/search", requirePermission(svc, "can_use_ai"), smartSearchHandler(svc))
|
||||
authed.GET("/ai/recommend", requirePermission(svc, "can_use_ai"), aiRecommendHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedFileRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/files", middleware.AdminRequired(), browseFilesHandler(svc))
|
||||
authed.POST("/files/folders", middleware.AdminRequired(), createFolderHandler(svc))
|
||||
@@ -49,46 +20,7 @@ func registerAuthedDLNARoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.POST("/dlna/cast", dlnaCastHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedSTRMRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.PUT("/media/:id/strm", middleware.AdminRequired(), setSTRMHandler(svc))
|
||||
authed.DELETE("/media/:id/strm", middleware.AdminRequired(), clearSTRMHandler(svc))
|
||||
authed.GET("/strm/output-presets", middleware.AdminRequired(), listSTRMOutputPresetsHandler(svc))
|
||||
authed.POST("/strm/import", middleware.AdminRequired(), importSTRMHandler(svc))
|
||||
authed.POST("/strm/generate", middleware.AdminRequired(), generateSTRMHandler(svc))
|
||||
authed.POST("/strm/generate-from-tree", middleware.AdminRequired(), generateSTRMFromTreeHandler(svc))
|
||||
authed.POST("/strm/repair", middleware.AdminRequired(), repairSTRMHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedDuplicateRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/duplicates", middleware.AdminRequired(), listDuplicatesHandler(svc))
|
||||
authed.POST("/duplicates/scan", middleware.AdminRequired(), detectDuplicatesHandler(svc))
|
||||
authed.POST("/duplicates/unmark", middleware.AdminRequired(), unmarkDuplicatesHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedSiteRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
siteHandler := NewSiteHandler(svc)
|
||||
authed.GET("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.ListSites)
|
||||
authed.GET("/sites/types", requirePermission(svc, "can_manage_sites"), siteHandler.GetSiteTypes)
|
||||
authed.GET("/sites/auth-types", requirePermission(svc, "can_manage_sites"), siteHandler.GetAuthTypes)
|
||||
authed.POST("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.CreateSite)
|
||||
authed.GET("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.GetSite)
|
||||
authed.PUT("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.UpdateSite)
|
||||
authed.DELETE("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.DeleteSite)
|
||||
authed.POST("/sites/:id/test", requirePermission(svc, "can_manage_sites"), siteHandler.TestSite)
|
||||
authed.GET("/sites/search", requirePermission(svc, "can_manage_sites"), siteSearchHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedRecycleAndRealtimeRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/recycle", middleware.AdminRequired(), listRecycleHandler(svc))
|
||||
authed.POST("/recycle/restore", middleware.AdminRequired(), restoreMediaBatchHandler(svc))
|
||||
authed.POST("/recycle/purge", middleware.AdminRequired(), purgeMediaBatchHandler(svc))
|
||||
|
||||
authed.GET("/ws", wsHandler(svc))
|
||||
authed.GET("/events", sseHandler(svc))
|
||||
}
|
||||
|
||||
func registerAuthedSchedulerRoutes(authed *gin.RouterGroup, svc *service.Container) {
|
||||
authed.GET("/scheduler/tasks", schedulerListTasksHandler(svc))
|
||||
authed.POST("/scheduler/tasks/:id/run", middleware.AdminRequired(), schedulerRunTaskHandler(svc))
|
||||
authed.GET("/scheduler/status", schedulerGetStatusHandler(svc))
|
||||
}
|
||||
|
||||
@@ -23,22 +23,16 @@ func TestAuthenticatedRouteSurfacesAreRegistered(t *testing.T) {
|
||||
routes[route.Method+" "+route.Path] = true
|
||||
}
|
||||
|
||||
for _, want := range []string{
|
||||
"GET /api/me",
|
||||
"GET /api/auth/permissions",
|
||||
"GET /api/libraries",
|
||||
"GET /api/media",
|
||||
"GET /api/stream/:id",
|
||||
"GET /api/storage",
|
||||
"GET /api/downloads",
|
||||
"GET /api/subscriptions",
|
||||
"GET /api/sites/search",
|
||||
"GET /api/watch-history",
|
||||
"GET /api/discover/feed",
|
||||
"GET /api/playback/:id/info",
|
||||
"GET /api/download/tasks",
|
||||
"GET /api/admin/assistant/history",
|
||||
} {
|
||||
for _, want := range []string{
|
||||
"GET /api/me",
|
||||
"GET /api/auth/permissions",
|
||||
"GET /api/libraries",
|
||||
"GET /api/media",
|
||||
"GET /api/stream/:id",
|
||||
"GET /api/storage",
|
||||
"GET /api/watch-history",
|
||||
"GET /api/playback/:id/info",
|
||||
} {
|
||||
if !routes[want] {
|
||||
t.Fatalf("%s route is not registered", want)
|
||||
}
|
||||
|
||||
@@ -1,42 +0,0 @@
|
||||
// Package handler — scheduled jobs admin page.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func schedulerStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"jobs": svc.Scheduler.Status()})
|
||||
}
|
||||
}
|
||||
|
||||
func schedulerRunHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
name := c.Param("name")
|
||||
if !triggerSchedulerJob(c, svc, name) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusAccepted, gin.H{"ok": true, "message": "任务已在后台触发"})
|
||||
}
|
||||
}
|
||||
|
||||
func triggerSchedulerJob(c *gin.Context, svc *service.Container, name string) bool {
|
||||
if err := svc.Scheduler.RunNowAsync(c.Request.Context(), name); err != nil {
|
||||
switch {
|
||||
case errors.Is(err, service.ErrSchedulerJobAlreadyRunning):
|
||||
c.JSON(http.StatusConflict, gin.H{"error": "任务正在运行,请稍后到实时任务查看进度"})
|
||||
case errors.Is(err, service.ErrSchedulerJobNotFound):
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "任务不存在"})
|
||||
default:
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
}
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -1,42 +0,0 @@
|
||||
// Package handler — 定时任务管理 HTTP 端点。
|
||||
package handler
|
||||
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// SchedulerHandler 处理定时任务的查询和管理操作。
|
||||
type SchedulerHandler struct {
|
||||
svc *service.Container
|
||||
log *zap.Logger
|
||||
}
|
||||
|
||||
// NewSchedulerHandler 创建定时任务处理器。
|
||||
func NewSchedulerHandler(svc *service.Container, log *zap.Logger) *SchedulerHandler {
|
||||
return &SchedulerHandler{svc: svc, log: log}
|
||||
}
|
||||
|
||||
// ListTasks 返回所有定时任务列表。
|
||||
func (h *SchedulerHandler) ListTasks(c *gin.Context) {
|
||||
tasks := h.svc.Scheduler.Status()
|
||||
Success(c, tasks)
|
||||
}
|
||||
|
||||
// RunTask 手动触发指定任务。
|
||||
func (h *SchedulerHandler) RunTask(c *gin.Context) {
|
||||
name := c.Param("id")
|
||||
if !triggerSchedulerJob(c, h.svc, name) {
|
||||
return
|
||||
}
|
||||
|
||||
SuccessWithMessage(c, "任务已触发后台执行", nil)
|
||||
}
|
||||
|
||||
// GetStatus 返回调度器运行状态。
|
||||
func (h *SchedulerHandler) GetStatus(c *gin.Context) {
|
||||
status := h.svc.Scheduler.Status()
|
||||
Success(c, status)
|
||||
}
|
||||
@@ -83,24 +83,3 @@ func searchTMDbHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusOK, gin.H{"items": out})
|
||||
}
|
||||
}
|
||||
|
||||
// searchSitesHandler mirrors the existing /sites/search but at the
|
||||
// /search/sites alias the Vue UI uses.
|
||||
func searchSitesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
keyword := c.Query("keyword")
|
||||
if keyword == "" {
|
||||
keyword = c.Query("q")
|
||||
}
|
||||
if keyword == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "keyword required"})
|
||||
return
|
||||
}
|
||||
results, err := svc.Site.Search(c.Request.Context(), keyword)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": results})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,170 +0,0 @@
|
||||
// Package handler — PT 站点管理 HTTP 处理。
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// SiteHandler 站点管理 CRUD。
|
||||
type SiteHandler struct {
|
||||
svc *service.Container
|
||||
}
|
||||
|
||||
// NewSiteHandler 创建站点管理 Handler。
|
||||
func NewSiteHandler(svc *service.Container) *SiteHandler {
|
||||
return &SiteHandler{svc: svc}
|
||||
}
|
||||
|
||||
// ListSites 列出所有站点。
|
||||
func (h *SiteHandler) ListSites(c *gin.Context) {
|
||||
sites, err := h.svc.Site.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 1, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
if sites == nil {
|
||||
sites = []model.Site{}
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": sites})
|
||||
}
|
||||
|
||||
// GetSite 获取单个站点详情。
|
||||
func (h *SiteHandler) GetSite(c *gin.Context) {
|
||||
site, err := h.svc.Site.FindByID(c.Request.Context(), c.Param("id"))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 1, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
if site == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"code": 1, "message": "site not found"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": site})
|
||||
}
|
||||
|
||||
// CreateSite 创建站点。
|
||||
func (h *SiteHandler) CreateSite(c *gin.Context) {
|
||||
var body map[string]any
|
||||
if err := c.ShouldBindJSON(&body); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
var site model.Site
|
||||
// 手动映射字段(因为敏感字段有 json:"-" 标签,不会自动反序列化)
|
||||
if v, ok := body["name"].(string); ok {
|
||||
site.Name = v
|
||||
}
|
||||
if v, ok := body["url"].(string); ok {
|
||||
site.URL = v
|
||||
}
|
||||
if v, ok := body["type"].(string); ok {
|
||||
site.Type = v
|
||||
}
|
||||
if v, ok := body["auth_type"].(string); ok {
|
||||
site.AuthType = v
|
||||
}
|
||||
if v, ok := body["api_key"].(string); ok {
|
||||
site.APIKey = v
|
||||
}
|
||||
if v, ok := body["cookie"].(string); ok {
|
||||
site.Cookie = v
|
||||
}
|
||||
if v, ok := body["auth_header"].(string); ok {
|
||||
site.AuthHeader = v
|
||||
}
|
||||
if v, ok := body["enabled"].(bool); ok {
|
||||
site.Enabled = v
|
||||
}
|
||||
if v, ok := body["is_default"].(bool); ok {
|
||||
site.IsDefault = v
|
||||
}
|
||||
if v, ok := body["extra"].(string); ok {
|
||||
site.Extra = v
|
||||
}
|
||||
// 高级设置字段
|
||||
if v, ok := body["user_agent"].(string); ok {
|
||||
site.UserAgent = v
|
||||
}
|
||||
if v, ok := body["rss_url"].(string); ok {
|
||||
site.RSSURL = v
|
||||
}
|
||||
if v, ok := body["timeout"].(float64); ok {
|
||||
site.Timeout = int(v)
|
||||
}
|
||||
if v, ok := body["priority"].(float64); ok {
|
||||
site.Priority = int(v)
|
||||
}
|
||||
if v, ok := body["use_proxy"].(bool); ok {
|
||||
site.UseProxy = v
|
||||
}
|
||||
if v, ok := body["rate_limit"].(bool); ok {
|
||||
site.RateLimit = v
|
||||
}
|
||||
if v, ok := body["browser_emulation"].(bool); ok {
|
||||
site.BrowserEmulation = v
|
||||
}
|
||||
if v, ok := body["downloader"].(string); ok {
|
||||
site.Downloader = v
|
||||
}
|
||||
|
||||
if err := h.svc.Site.Create(c.Request.Context(), &site); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, gin.H{"code": 0, "message": "ok", "data": site})
|
||||
}
|
||||
|
||||
// UpdateSite 更新站点。
|
||||
func (h *SiteHandler) UpdateSite(c *gin.Context) {
|
||||
patch := make(map[string]any)
|
||||
if err := c.ShouldBindJSON(&patch); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
// 字段名映射:前端使用蛇形命名,GORM 会正确映射到列名
|
||||
// 无需额外处理,Updates 直接使用 patch 中的 key-value
|
||||
if err := h.svc.Site.Update(c.Request.Context(), c.Param("id"), patch); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok"})
|
||||
}
|
||||
|
||||
// DeleteSite 删除站点。
|
||||
func (h *SiteHandler) DeleteSite(c *gin.Context) {
|
||||
if err := h.svc.Site.Delete(c.Request.Context(), c.Param("id")); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 1, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok"})
|
||||
}
|
||||
|
||||
// TestSite 测试站点连通性。
|
||||
func (h *SiteHandler) TestSite(c *gin.Context) {
|
||||
ok, msg, err := h.svc.Site.TestConnection(c.Request.Context(), c.Param("id"))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"code": 1, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
if !ok {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": msg})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "message": msg})
|
||||
}
|
||||
|
||||
// GetSiteTypes 返回支持的站点类型列表。
|
||||
func (h *SiteHandler) GetSiteTypes(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": model.SiteTypes()})
|
||||
}
|
||||
|
||||
// GetAuthTypes 返回支持的认证方式列表。
|
||||
func (h *SiteHandler) GetAuthTypes(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": model.AuthTypes()})
|
||||
}
|
||||
@@ -1,164 +0,0 @@
|
||||
// Package handler — site management (PT/BT tracker CRUD + cross-site search).
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// ─── CRUD ────────────────────────────────────────────────────────────────────
|
||||
|
||||
func listSitesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
sites, err := svc.Site.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": sites})
|
||||
}
|
||||
}
|
||||
|
||||
func getSiteHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
site, err := svc.Site.FindByID(c.Request.Context(), c.Param("id"))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if site == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, site)
|
||||
}
|
||||
}
|
||||
|
||||
type createSiteReq struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
BaseURL string `json:"base_url" binding:"required"`
|
||||
SiteType string `json:"site_type"`
|
||||
AuthType string `json:"auth_type"`
|
||||
Cookie string `json:"cookie"`
|
||||
APIKey string `json:"api_key"`
|
||||
AuthHeader string `json:"auth_header"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
RSSURL string `json:"rss_url"`
|
||||
Timeout int `json:"timeout"`
|
||||
Priority int `json:"priority"`
|
||||
UseProxy bool `json:"use_proxy"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Downloader string `json:"downloader"`
|
||||
}
|
||||
|
||||
func createSiteHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req createSiteReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
enabled := true
|
||||
if req.Enabled != nil {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
// Pack fields not in the core model into Extra JSON.
|
||||
extraMap := map[string]any{}
|
||||
if req.UserAgent != "" {
|
||||
extraMap["user_agent"] = req.UserAgent
|
||||
}
|
||||
if req.RSSURL != "" {
|
||||
extraMap["rss_url"] = req.RSSURL
|
||||
}
|
||||
if req.Timeout > 0 {
|
||||
extraMap["timeout"] = req.Timeout
|
||||
}
|
||||
if req.Priority > 0 {
|
||||
extraMap["priority"] = req.Priority
|
||||
}
|
||||
extraMap["use_proxy"] = req.UseProxy
|
||||
if req.Downloader != "" {
|
||||
extraMap["downloader"] = req.Downloader
|
||||
}
|
||||
extraJSON, _ := json.Marshal(extraMap)
|
||||
|
||||
site := &model.Site{
|
||||
Name: req.Name,
|
||||
URL: req.BaseURL,
|
||||
Type: req.SiteType,
|
||||
AuthType: req.AuthType,
|
||||
Cookie: req.Cookie,
|
||||
APIKey: req.APIKey,
|
||||
AuthHeader: req.AuthHeader,
|
||||
Extra: string(extraJSON),
|
||||
Enabled: enabled,
|
||||
}
|
||||
if err := svc.Site.Create(c.Request.Context(), site); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, site)
|
||||
}
|
||||
}
|
||||
|
||||
func updateSiteHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var patch map[string]any
|
||||
if err := c.ShouldBindJSON(&patch); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := svc.Site.Update(c.Request.Context(), c.Param("id"), patch); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
// Return the updated row.
|
||||
site, _ := svc.Site.FindByID(c.Request.Context(), c.Param("id"))
|
||||
c.JSON(http.StatusOK, site)
|
||||
}
|
||||
}
|
||||
|
||||
func deleteSiteHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Site.Delete(c.Request.Context(), c.Param("id")); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── Connection test ─────────────────────────────────────────────────────────
|
||||
|
||||
func testSiteHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
ok, msg, err := svc.Site.TestConnection(c.Request.Context(), c.Param("id"))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"success": ok, "message": msg})
|
||||
}
|
||||
}
|
||||
|
||||
// ─── Cross-site search ───────────────────────────────────────────────────────
|
||||
|
||||
func siteSearchHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
keyword := c.Query("keyword")
|
||||
if keyword == "" {
|
||||
keyword = c.Query("q")
|
||||
}
|
||||
results, err := svc.Site.Search(c.Request.Context(), keyword)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": results, "total": len(results)})
|
||||
}
|
||||
}
|
||||
@@ -1,62 +0,0 @@
|
||||
// Package handler — extra site endpoints used by the Vue UI:
|
||||
//
|
||||
// GET /sites/:id/resource → keyword search scoped to one site
|
||||
// GET /sites/:id/userdata → cookie-derived user info (stubbed)
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// siteResourceHandler runs a search restricted to a single site.
|
||||
//
|
||||
// We reuse the full SiteService.Search() and post-filter by site_id;
|
||||
// it's not the hottest path so we trade simplicity for speed here.
|
||||
func siteResourceHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
keyword := c.Query("keyword")
|
||||
if keyword == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "keyword required"})
|
||||
return
|
||||
}
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
items, err := svc.Site.SearchSite(c.Request.Context(), c.Param("id"), keyword, page)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": items, "total": len(items)})
|
||||
}
|
||||
}
|
||||
|
||||
// siteUserdataHandler returns whatever the site exposes about the
|
||||
// authenticated user (upload/download stats, ratio, etc.). This is a
|
||||
// stub: we report the cookie length so the UI can confirm a login is
|
||||
// present, but full per-site parsing is out of scope here.
|
||||
func siteUserdataHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
s, err := svc.Site.FindByID(c.Request.Context(), c.Param("id"))
|
||||
if err != nil || s == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "site not found"})
|
||||
return
|
||||
}
|
||||
loginStatus := "unknown"
|
||||
if s.LastError == "ok" {
|
||||
loginStatus = "ok"
|
||||
} else if s.LastError != "" {
|
||||
loginStatus = "fail"
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"site_id": s.ID,
|
||||
"name": s.Name,
|
||||
"cookie_set": len(s.Cookie) > 0,
|
||||
"login_status": loginStatus,
|
||||
"note": "userdata parsing not implemented; stub",
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,154 +0,0 @@
|
||||
// Package handler — stats / dashboard endpoints.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"crypto/sha1"
|
||||
"encoding/hex"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func statsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
snap, err := svc.Stats.Compute(c.Request.Context(), svc.Cfg.App.DataDir)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := applyStatsVisibility(c, svc, snap); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, snap)
|
||||
}
|
||||
}
|
||||
|
||||
func applyStatsVisibility(c *gin.Context, svc *service.Container, snap *service.Snapshot) error {
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
libs, err := svc.Repo.Library.List(c.Request.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
libs = service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, libs)
|
||||
var visibleLibraries int64
|
||||
activeLibraryIDs := make([]string, 0, len(libs))
|
||||
for _, lib := range libs {
|
||||
if !lib.Enabled {
|
||||
continue
|
||||
}
|
||||
activeLibraryIDs = append(activeLibraryIDs, lib.ID)
|
||||
if service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, lib, visibility) {
|
||||
visibleLibraries++
|
||||
}
|
||||
}
|
||||
snap.Libraries = visibleLibraries
|
||||
cacheKey := visibleStatsCacheKey(visibility, activeLibraryIDs)
|
||||
if svc.Cache != nil {
|
||||
var cached visibleStatsCacheValue
|
||||
if svc.Cache.GetJSON(c.Request.Context(), cacheKey, &cached) {
|
||||
snap.Libraries = cached.Libraries
|
||||
snap.MediaCount = cached.MediaCount
|
||||
snap.TotalSizeBytes = cached.TotalSizeBytes
|
||||
snap.TotalSeconds = cached.TotalSeconds
|
||||
snap.RecentlyAdded = cached.RecentlyAdded
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
q := applyMediaVisibilityQuery(svc.Repo.DB.WithContext(c.Request.Context()).Model(&model.Media{}), visibility)
|
||||
q = applyActiveLibraryQuery(q, activeLibraryIDs)
|
||||
if err := q.Count(&snap.MediaCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
type sumRow struct {
|
||||
Size int64
|
||||
Seconds int64
|
||||
}
|
||||
var sum sumRow
|
||||
if err := applyActiveLibraryQuery(applyMediaVisibilityQuery(svc.Repo.DB.WithContext(c.Request.Context()).Model(&model.Media{}), visibility), activeLibraryIDs).
|
||||
Select("COALESCE(SUM(size_bytes),0) as size, COALESCE(SUM(duration_sec),0) as seconds").
|
||||
Scan(&sum).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
snap.TotalSizeBytes = sum.Size
|
||||
snap.TotalSeconds = sum.Seconds
|
||||
|
||||
var recent []model.Media
|
||||
if err := applyActiveLibraryQuery(applyMediaVisibilityQuery(svc.Repo.DB.WithContext(c.Request.Context()).Model(&model.Media{}), visibility), activeLibraryIDs).
|
||||
Order("created_at desc").
|
||||
Limit(12).
|
||||
Find(&recent).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
snap.RecentlyAdded = recent
|
||||
if svc.Cache != nil {
|
||||
svc.Cache.SetJSON(c.Request.Context(), cacheKey, visibleStatsCacheValue{
|
||||
Libraries: snap.Libraries,
|
||||
MediaCount: snap.MediaCount,
|
||||
TotalSizeBytes: snap.TotalSizeBytes,
|
||||
TotalSeconds: snap.TotalSeconds,
|
||||
RecentlyAdded: snap.RecentlyAdded,
|
||||
}, 10*time.Second)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type visibleStatsCacheValue struct {
|
||||
Libraries int64 `json:"libraries"`
|
||||
MediaCount int64 `json:"media_count"`
|
||||
TotalSizeBytes int64 `json:"total_size_bytes"`
|
||||
TotalSeconds int64 `json:"total_seconds"`
|
||||
RecentlyAdded []model.Media `json:"recently_added"`
|
||||
}
|
||||
|
||||
func visibleStatsCacheKey(visibility service.MediaVisibility, activeLibraryIDs []string) string {
|
||||
allowed := append([]string(nil), visibility.AllowedLibraryIDs...)
|
||||
hidden := append([]string(nil), visibility.HiddenLibraryIDs...)
|
||||
active := append([]string(nil), activeLibraryIDs...)
|
||||
sort.Strings(allowed)
|
||||
sort.Strings(hidden)
|
||||
sort.Strings(active)
|
||||
sum := sha1.Sum([]byte(strings.Join([]string{
|
||||
"visible",
|
||||
strings.Join(active, ","),
|
||||
strings.Join(allowed, ","),
|
||||
strings.Join(hidden, ","),
|
||||
boolString(visibility.IncludeNSFW),
|
||||
}, "|")))
|
||||
return "stats:visible:" + hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func boolString(value bool) string {
|
||||
if value {
|
||||
return "1"
|
||||
}
|
||||
return "0"
|
||||
}
|
||||
|
||||
func applyMediaVisibilityQuery(q *gorm.DB, visibility service.MediaVisibility) *gorm.DB {
|
||||
if !visibility.IncludeNSFW {
|
||||
q = q.Where("nsfw = ?", false)
|
||||
}
|
||||
if len(visibility.HiddenLibraryIDs) > 0 {
|
||||
q = q.Where("library_id NOT IN ?", visibility.HiddenLibraryIDs)
|
||||
}
|
||||
if len(visibility.AllowedLibraryIDs) > 0 {
|
||||
q = q.Where("library_id IN ?", visibility.AllowedLibraryIDs)
|
||||
}
|
||||
return q
|
||||
}
|
||||
|
||||
func applyActiveLibraryQuery(q *gorm.DB, libraryIDs []string) *gorm.DB {
|
||||
if len(libraryIDs) == 0 {
|
||||
return q.Where("1 = 0")
|
||||
}
|
||||
return q.Where("library_id IN ?", libraryIDs)
|
||||
}
|
||||
@@ -1,180 +0,0 @@
|
||||
// Package handler — richer dashboard statistics endpoints.
|
||||
//
|
||||
// /api/stats already returns the basic snapshot. The Vue admin
|
||||
// dashboard also uses:
|
||||
//
|
||||
// /api/stats/overview — counts + total size + total seconds
|
||||
// /api/stats/trend — daily play count over last N days
|
||||
// /api/stats/top-content — top played media (by play count)
|
||||
// /api/stats/libraries — per-library item count + size
|
||||
// /api/stats/monitor — live CPU/mem/disk
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func statsOverviewHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
snap, err := svc.Stats.Compute(c.Request.Context(), svc.Cfg.App.DataDir)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := applyStatsVisibility(c, svc, snap); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"libraries": snap.Libraries,
|
||||
"media_count": snap.MediaCount,
|
||||
"users_count": snap.UsersCount,
|
||||
"total_size": snap.TotalSizeBytes,
|
||||
"total_seconds": snap.TotalSeconds,
|
||||
"generated_at": snap.GeneratedAt,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// statsTrendHandler returns play counts per day for the last N days.
|
||||
func statsTrendHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
days, _ := strconv.Atoi(c.DefaultQuery("days", "14"))
|
||||
if days <= 0 || days > 90 {
|
||||
days = 14
|
||||
}
|
||||
// Use the playback_history table; one row per (user, media)
|
||||
// per day if we group by date(watched_at).
|
||||
type bucket struct {
|
||||
Day string `json:"day"`
|
||||
Count int64 `json:"count"`
|
||||
}
|
||||
out := make([]bucket, 0, days)
|
||||
now := time.Now().UTC()
|
||||
for i := days - 1; i >= 0; i-- {
|
||||
start := now.AddDate(0, 0, -i).Truncate(24 * time.Hour)
|
||||
end := start.Add(24 * time.Hour)
|
||||
var n int64
|
||||
_ = svc.Repo.DB.Model(&model.PlaybackHistory{}).
|
||||
Where("watched_at >= ? AND watched_at < ?", start, end).
|
||||
Count(&n).Error
|
||||
out = append(out, bucket{
|
||||
Day: start.Format("2006-01-02"),
|
||||
Count: n,
|
||||
})
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"trend": out, "days": days})
|
||||
}
|
||||
}
|
||||
|
||||
// statsTopContentHandler returns the most-watched media items.
|
||||
func statsTopContentHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "10"))
|
||||
if limit <= 0 || limit > 50 {
|
||||
limit = 10
|
||||
}
|
||||
type row struct {
|
||||
MediaID string `json:"media_id"`
|
||||
PlayCount int64 `json:"play_count"`
|
||||
LastPlayed time.Time `json:"last_played"`
|
||||
}
|
||||
var rows []row
|
||||
_ = svc.Repo.DB.Table("playback_histories").
|
||||
Select("media_id, COUNT(*) as play_count, MAX(watched_at) as last_played").
|
||||
Group("media_id").
|
||||
Order("play_count desc").
|
||||
Limit(limit).
|
||||
Scan(&rows).Error
|
||||
// Hydrate media titles in a single query.
|
||||
ids := make([]string, 0, len(rows))
|
||||
for _, r := range rows {
|
||||
ids = append(ids, r.MediaID)
|
||||
}
|
||||
mIdx := map[string]model.Media{}
|
||||
if len(ids) > 0 {
|
||||
var media []model.Media
|
||||
_ = svc.Repo.DB.Where("id IN ?", ids).Find(&media).Error
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
for _, m := range media {
|
||||
if !visibility.Allows(&m) {
|
||||
continue
|
||||
}
|
||||
mIdx[m.ID] = m
|
||||
}
|
||||
}
|
||||
out := make([]gin.H, 0, len(rows))
|
||||
for _, r := range rows {
|
||||
media, ok := mIdx[r.MediaID]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
out = append(out, gin.H{
|
||||
"media": media,
|
||||
"play_count": r.PlayCount,
|
||||
"last_played": r.LastPlayed,
|
||||
})
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": out})
|
||||
}
|
||||
}
|
||||
|
||||
// statsLibrariesHandler returns per-library counts + size.
|
||||
func statsLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
libs, err := svc.Repo.Library.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
libs = service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, libs)
|
||||
out := make([]gin.H, 0, len(libs))
|
||||
visibility := mediaVisibilityForRequest(c, svc)
|
||||
for _, l := range libs {
|
||||
if !l.Enabled {
|
||||
continue
|
||||
}
|
||||
if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, l, visibility) {
|
||||
continue
|
||||
}
|
||||
libraryIDs, err := service.MergedLibraryIDsForLibrary(c.Request.Context(), svc.Repo, l.ID)
|
||||
if err != nil || len(libraryIDs) == 0 {
|
||||
libraryIDs = []string{l.ID}
|
||||
}
|
||||
var count int64
|
||||
var size int64
|
||||
_ = applyMediaVisibilityQuery(svc.Repo.DB.Model(&model.Media{}), visibility).
|
||||
Where("library_id IN ?", libraryIDs).
|
||||
Count(&count).Error
|
||||
_ = applyMediaVisibilityQuery(svc.Repo.DB.Model(&model.Media{}), visibility).
|
||||
Where("library_id IN ?", libraryIDs).
|
||||
Select("COALESCE(SUM(size_bytes),0)").Row().Scan(&size)
|
||||
out = append(out, gin.H{
|
||||
"library": l,
|
||||
"item_count": count,
|
||||
"total_size": size,
|
||||
})
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"libraries": out})
|
||||
}
|
||||
}
|
||||
|
||||
// statsMonitorHandler returns live system resource usage; this is just
|
||||
// the Hardware portion of the snapshot but with a snappy schema.
|
||||
func statsMonitorHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
snap, err := svc.Stats.Compute(c.Request.Context(), svc.Cfg.App.DataDir)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, snap.Hardware)
|
||||
}
|
||||
}
|
||||
@@ -1,115 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/middleware"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func TestStatsSnapshotHidesAdultRecentlyAddedForUser(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.Setting{}, &model.PlayProfile{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
viewer := &model.User{Username: "viewer", PasswordHash: "hash", Role: "user", HideAdult: true}
|
||||
if err := repos.User.Create(t.Context(), viewer); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
safe := model.Library{Name: "电影", Path: "/media/movie", Type: "movie", Enabled: true}
|
||||
adult := model.Library{Name: "9KG", Path: "/media/9KG", Type: "movie", Enabled: true}
|
||||
if err := repos.Library.Create(t.Context(), &safe); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := repos.Library.Create(t.Context(), &adult); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := repos.Setting.Set(t.Context(), service.AdultLibraryIDsSettingKey, `["`+adult.ID+`"]`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.Create(&model.Media{LibraryID: safe.ID, Title: "普通电影", Path: "/media/movie/a.mkv", SizeBytes: 100, DurationSec: 10}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.Create(&model.Media{LibraryID: adult.ID, Title: "成人影片", Path: "/media/9KG/a.mkv", SizeBytes: 200, DurationSec: 20}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc := &service.Container{Repo: repos}
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
c.Set(middleware.CtxUserID, viewer.ID)
|
||||
c.Request = httptest.NewRequest("GET", "/api/stats", nil)
|
||||
|
||||
snap := &service.Snapshot{}
|
||||
if err := applyStatsVisibility(c, svc, snap); err != nil {
|
||||
t.Fatalf("applyStatsVisibility: %v", err)
|
||||
}
|
||||
if snap.MediaCount != 1 || snap.TotalSizeBytes != 100 || snap.TotalSeconds != 10 {
|
||||
t.Fatalf("stats should only include visible media, got count=%d size=%d seconds=%d", snap.MediaCount, snap.TotalSizeBytes, snap.TotalSeconds)
|
||||
}
|
||||
if len(snap.RecentlyAdded) != 1 || snap.RecentlyAdded[0].LibraryID != safe.ID {
|
||||
t.Fatalf("recently added should hide adult library, got %#v", snap.RecentlyAdded)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStatsLibrariesCountsMergedCloudLibraryItems(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.Setting{}, &model.PlayProfile{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
local := model.Library{Name: "国产电影", Path: "/media/国产电影", Type: "movie", Enabled: true}
|
||||
cloud := model.Library{Name: "OpenList · 国产电影", Path: service.BuildCloudLibraryPath("openlist", "/国产电影", "/国产电影"), Type: "movie", Enabled: true}
|
||||
for _, lib := range []*model.Library{&local, &cloud} {
|
||||
if err := repos.Library.Create(t.Context(), lib); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := db.Create(&[]model.Media{
|
||||
{LibraryID: local.ID, Title: "本地版本", Path: "/media/国产电影/local.mkv", SizeBytes: 100},
|
||||
{LibraryID: cloud.ID, Title: "云盘版本", Path: "cloud://openlist/国产电影/cloud.mkv", SizeBytes: 200},
|
||||
}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc := &service.Container{Repo: repos}
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
c.Request = httptest.NewRequest("GET", "/api/stats/libraries", nil)
|
||||
|
||||
statsLibrariesHandler(svc)(c)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
var payload struct {
|
||||
Libraries []struct {
|
||||
ItemCount int64 `json:"item_count"`
|
||||
TotalSize int64 `json:"total_size"`
|
||||
} `json:"libraries"`
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(payload.Libraries) != 1 {
|
||||
t.Fatalf("libraries = %#v, want one merged display library", payload.Libraries)
|
||||
}
|
||||
if payload.Libraries[0].ItemCount != 2 || payload.Libraries[0].TotalSize != 300 {
|
||||
t.Fatalf("merged stats = %#v, want count=2 size=300", payload.Libraries[0])
|
||||
}
|
||||
}
|
||||
@@ -1,145 +0,0 @@
|
||||
// Package handler — external storage config endpoints.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// listStorageConfigsHandler returns the status overview used by the
|
||||
// admin storage panel: every persisted backend with secrets redacted.
|
||||
func listStorageConfigsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
rows, err := svc.StorageCfg.List(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": rows})
|
||||
}
|
||||
}
|
||||
|
||||
// getStorageConfigHandler returns one config (with the decrypted body).
|
||||
func getStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !service.IsAdminStorageConfigurable(c.Param("type")) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
|
||||
return
|
||||
}
|
||||
row, err := svc.StorageCfg.Get(c.Request.Context(), c.Param("type"))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if row == nil {
|
||||
c.JSON(http.StatusOK, gin.H{"type": c.Param("type"), "config": gin.H{}})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, row)
|
||||
}
|
||||
}
|
||||
|
||||
// saveStorageConfigHandler upserts the config row; the caller passes
|
||||
// the type via URL and the body as a JSON object.
|
||||
func saveStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !service.IsAdminStorageConfigurable(c.Param("type")) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
|
||||
return
|
||||
}
|
||||
var in service.StorageInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
in.Type = c.Param("type")
|
||||
row, err := svc.StorageCfg.Save(c.Request.Context(), in)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if in.Enabled != nil && !*in.Enabled && svc.Scan != nil {
|
||||
_ = svc.Scan.CancelCloudScansForProvider(in.Type)
|
||||
}
|
||||
c.JSON(http.StatusOK, row)
|
||||
}
|
||||
}
|
||||
|
||||
// testStorageConfigHandler probes an unsaved config.
|
||||
func testStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !service.IsAdminStorageConfigurable(c.Param("type")) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"ok": false, "error": "unsupported storage type"})
|
||||
return
|
||||
}
|
||||
var in service.StorageInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
in.Type = c.Param("type")
|
||||
if err := svc.StorageCfg.Test(c.Request.Context(), in); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": false, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
if service.IsAdminCloudConfigurable(in.Type) {
|
||||
enabled := true
|
||||
if _, err := svc.StorageCfg.Save(c.Request.Context(), service.StorageInput{
|
||||
Type: in.Type,
|
||||
Config: in.Config,
|
||||
Enabled: &enabled,
|
||||
}); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": false, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
|
||||
func logoutStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
typ := c.Param("type")
|
||||
if !service.IsAdminStorageConfigurable(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
|
||||
return
|
||||
}
|
||||
row, err := svc.StorageCfg.Logout(c.Request.Context(), typ)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if svc.Scan != nil {
|
||||
_ = svc.Scan.CancelCloudScansForProvider(typ)
|
||||
}
|
||||
c.JSON(http.StatusOK, row)
|
||||
}
|
||||
}
|
||||
|
||||
func storageUploadLocalHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !service.IsAdminStorageConfigurable(c.Param("type")) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
|
||||
return
|
||||
}
|
||||
var req service.CloudUploadInput
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
req.Type = c.Param("type")
|
||||
res, err := svc.StorageCfg.UploadLocal(c.Request.Context(), req)
|
||||
if err != nil && res == nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"result": res, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"result": res})
|
||||
}
|
||||
}
|
||||
@@ -1,101 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"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/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func TestStorageConfigHandlersRejectQuark(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
router.PUT("/admin/storage/:type", saveStorageConfigHandler(nil))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPut, "/admin/storage/quark", strings.NewReader(`{"type":"quark","config":{"cookie":"x"}}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
|
||||
}
|
||||
if !strings.Contains(w.Body.String(), "unsupported storage type") {
|
||||
t.Fatalf("body = %s, want unsupported storage type", w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestStorageConfigCloudTestSuccessSavesAndEnablesProvider(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
openlist := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/api/fs/list" {
|
||||
t.Fatalf("unexpected openlist path %s", r.URL.Path)
|
||||
}
|
||||
if r.Header.Get("Authorization") != "openlist-token" {
|
||||
t.Fatalf("authorization = %q, want openlist-token", r.Header.Get("Authorization"))
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"code": 200,
|
||||
"data": map[string]any{"content": []any{}, "total": 0},
|
||||
})
|
||||
}))
|
||||
defer openlist.Close()
|
||||
|
||||
db, err := gorm.Open(sqlite.Open("file:storage_config_cloud_test?mode=memory&cache=shared"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.StorageConfig{}, &model.Setting{}, &model.Library{}, &model.Media{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
repos := repository.New(db)
|
||||
log := zap.NewNop()
|
||||
storage := service.NewStorageConfigService(log, repos, service.NewCryptoService("", log))
|
||||
enabled := false
|
||||
if _, err := storage.Save(t.Context(), service.StorageInput{
|
||||
Type: "openlist",
|
||||
Config: map[string]any{
|
||||
"server": openlist.URL,
|
||||
"token": "openlist-token",
|
||||
},
|
||||
Enabled: &enabled,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.POST("/admin/storage/:type/test", testStorageConfigHandler(&service.Container{StorageCfg: storage}))
|
||||
body := `{"type":"openlist","config":{"server":"` + openlist.URL + `","token":"openlist-token"}}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/admin/storage/openlist/test", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d body=%s, want 200", w.Code, w.Body.String())
|
||||
}
|
||||
if !strings.Contains(w.Body.String(), `"ok":true`) {
|
||||
t.Fatalf("body = %s, want ok true", w.Body.String())
|
||||
}
|
||||
view, err := storage.Get(t.Context(), "openlist")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if view == nil || !view.Enabled {
|
||||
t.Fatalf("openlist enabled = %#v, want enabled after successful test", view)
|
||||
}
|
||||
if _, err := storage.CloudProvider(t.Context(), "openlist"); err != nil {
|
||||
t.Fatalf("cloud provider after test: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -85,42 +85,6 @@ func imageProxyHandler(svc *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func cloudArtworkProxyHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
typ := c.Param("type")
|
||||
ref := c.Query("ref")
|
||||
if !service.IsAdminCloudConfigurable(typ) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||
return
|
||||
}
|
||||
if ref == "" || !isCloudImageRef(ref) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "image ref required"})
|
||||
return
|
||||
}
|
||||
if svc == nil || svc.ImageProxy == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "image proxy unavailable"})
|
||||
return
|
||||
}
|
||||
stableKey := typ + ":" + ref
|
||||
if svc.ImageProxy.ServeCloudCached(c.Writer, c.Request, stableKey) {
|
||||
return
|
||||
}
|
||||
if svc.StorageCfg == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "cloud storage service unavailable"})
|
||||
return
|
||||
}
|
||||
link, err := svc.StorageCfg.CloudResolve(c.Request.Context(), typ, ref, c.Request.UserAgent())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := svc.ImageProxy.ServeCloudResolved(c.Request.Context(), c.Writer, c.Request, stableKey, link); err != nil {
|
||||
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type scrapeRequest struct {
|
||||
EpisodeArtwork *bool `json:"episode_artwork"`
|
||||
EpisodeImages *bool `json:"episode_images"`
|
||||
|
||||
@@ -1,437 +0,0 @@
|
||||
// Package handler — STRM (URL-as-file) admin endpoints.
|
||||
//
|
||||
// Setting a media row's strm_url makes the stream handler issue a 302
|
||||
// redirect to that URL instead of opening a local file. This lets the
|
||||
// operator expose WebDAV / Alist / S3 / HTTP direct links as ordinary
|
||||
// MediaStationGo entries.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/middleware"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type strmReq struct {
|
||||
URL string `json:"url" binding:"required"`
|
||||
}
|
||||
|
||||
func setSTRMHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req strmReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
url := strings.TrimSpace(req.URL)
|
||||
if !strings.HasPrefix(url, "http://") && !strings.HasPrefix(url, "https://") {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "url must start with http:// or https://"})
|
||||
return
|
||||
}
|
||||
mediaID := c.Param("id")
|
||||
m, err := svc.Repo.Media.FindByID(c.Request.Context(), mediaID)
|
||||
if err != nil || m == nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
|
||||
return
|
||||
}
|
||||
if err := svc.Repo.DB.WithContext(c.Request.Context()).
|
||||
Model(&model.Media{}).
|
||||
Where("id = ?", mediaID).
|
||||
Update("strm_url", url).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"strm_url": url})
|
||||
}
|
||||
}
|
||||
|
||||
func clearSTRMHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Repo.DB.WithContext(c.Request.Context()).
|
||||
Model(&model.Media{}).
|
||||
Where("id = ?", c.Param("id")).
|
||||
Update("strm_url", "").Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
// importSTRMHandler creates a media row directly from a (library_id, title, url)
|
||||
// tuple — useful for adding a streaming-only entry without an on-disk file.
|
||||
type importSTRMReq struct {
|
||||
LibraryID string `json:"library_id" binding:"required"`
|
||||
Title string `json:"title" binding:"required"`
|
||||
URL string `json:"url" binding:"required"`
|
||||
}
|
||||
|
||||
func importSTRMHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req importSTRMReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
url := strings.TrimSpace(req.URL)
|
||||
if !strings.HasPrefix(url, "http://") && !strings.HasPrefix(url, "https://") {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "url must start with http:// or https://"})
|
||||
return
|
||||
}
|
||||
m := &model.Media{
|
||||
LibraryID: req.LibraryID,
|
||||
Title: req.Title,
|
||||
Path: url,
|
||||
STRMURL: url,
|
||||
Container: "strm",
|
||||
}
|
||||
if err := svc.Repo.Media.Upsert(c.Request.Context(), m); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusCreated, m)
|
||||
}
|
||||
}
|
||||
|
||||
type generateSTRMReq struct {
|
||||
LibraryID string `json:"library_id"`
|
||||
OutputDir string `json:"output_dir"`
|
||||
BaseURL string `json:"base_url"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Overwrite bool `json:"overwrite"`
|
||||
IncludeLocal *bool `json:"include_local"`
|
||||
PreserveTree bool `json:"preserve_tree"`
|
||||
Refresh bool `json:"refresh_library"`
|
||||
ScrapeAfter bool `json:"scrape_after"`
|
||||
}
|
||||
|
||||
type generateSTRMTreeReq struct {
|
||||
Provider string `json:"provider"`
|
||||
TreeText string `json:"tree_text"`
|
||||
Paths []string `json:"paths"`
|
||||
SourceRoot string `json:"source_root"`
|
||||
OutputPrefix string `json:"output_prefix"`
|
||||
OutputDir string `json:"output_dir"`
|
||||
BaseURL string `json:"base_url"`
|
||||
Overwrite bool `json:"overwrite"`
|
||||
Cleanup bool `json:"cleanup"`
|
||||
DryRun bool `json:"dry_run"`
|
||||
BatchLimit int `json:"batch_limit"`
|
||||
RecognizeRename bool `json:"recognize_rename"`
|
||||
TransferSubtitles bool `json:"transfer_subtitles"`
|
||||
MissingOnly bool `json:"missing_only"`
|
||||
RefreshLibrary bool `json:"refresh_library"`
|
||||
ScrapeAfter bool `json:"scrape_after"`
|
||||
}
|
||||
|
||||
type repairSTRMReq struct {
|
||||
OutputDir string `json:"output_dir" binding:"required"`
|
||||
BaseURL string `json:"base_url"`
|
||||
DryRun bool `json:"dry_run"`
|
||||
RefreshLibrary bool `json:"refresh_library"`
|
||||
ScrapeAfter bool `json:"scrape_after"`
|
||||
}
|
||||
|
||||
func listSTRMOutputPresetsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
items, err := service.STRMOutputPresets(c.Request.Context(), svc.Repo)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||
}
|
||||
}
|
||||
|
||||
func generateSTRMHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req generateSTRMReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
strmSvc := svc.STRM
|
||||
if strmSvc == nil {
|
||||
strmSvc = service.NewSTRMService(svc.Log, svc.Repo, svc.Cfg)
|
||||
}
|
||||
baseURL := strings.TrimRight(strings.TrimSpace(req.BaseURL), "/")
|
||||
if baseURL == "" {
|
||||
baseURL = strings.TrimRight(absoluteRequestURL(c, "/"), "/")
|
||||
}
|
||||
includeLocal := true
|
||||
if req.IncludeLocal != nil {
|
||||
includeLocal = *req.IncludeLocal
|
||||
}
|
||||
options := service.GenerateSTRMOptions{
|
||||
LibraryID: req.LibraryID,
|
||||
OutputDir: req.OutputDir,
|
||||
BaseURL: baseURL,
|
||||
Enabled: req.Enabled,
|
||||
Overwrite: req.Overwrite,
|
||||
IncludeLocal: includeLocal,
|
||||
PreserveTree: req.PreserveTree,
|
||||
PlaybackToken: strmPlaybackTokenForRequest(c, svc),
|
||||
}
|
||||
var res *service.GenerateSTRMResult
|
||||
var err error
|
||||
if strings.TrimSpace(req.LibraryID) == "*" {
|
||||
res, err = strmSvc.GenerateForAllLibraries(c.Request.Context(), options)
|
||||
} else {
|
||||
res, err = strmSvc.GenerateForLibrary(c.Request.Context(), options)
|
||||
}
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if req.Refresh {
|
||||
res.Refresh = queueSTRMRefreshAfterChanges(c.Request.Context(), svc, res.OutputDir, strmRefreshQueueOptions{
|
||||
TaskName: "STRM 生成后刷新媒体库",
|
||||
Changed: strmGenerationChanged(res),
|
||||
ScrapeAfter: req.ScrapeAfter,
|
||||
})
|
||||
}
|
||||
c.JSON(http.StatusOK, res)
|
||||
}
|
||||
}
|
||||
|
||||
func generateSTRMFromTreeHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req generateSTRMTreeReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
strmSvc := svc.STRM
|
||||
if strmSvc == nil {
|
||||
strmSvc = service.NewSTRMService(svc.Log, svc.Repo, svc.Cfg)
|
||||
}
|
||||
baseURL := strings.TrimRight(strings.TrimSpace(req.BaseURL), "/")
|
||||
if baseURL == "" {
|
||||
baseURL = strings.TrimRight(absoluteRequestURL(c, "/"), "/")
|
||||
}
|
||||
res, err := strmSvc.GenerateFromTree(c.Request.Context(), service.GenerateSTRMTreeOptions{
|
||||
Provider: req.Provider,
|
||||
TreeText: req.TreeText,
|
||||
Paths: req.Paths,
|
||||
SourceRoot: req.SourceRoot,
|
||||
OutputPrefix: req.OutputPrefix,
|
||||
OutputDir: req.OutputDir,
|
||||
BaseURL: baseURL,
|
||||
Overwrite: req.Overwrite,
|
||||
Cleanup: req.Cleanup,
|
||||
DryRun: req.DryRun,
|
||||
BatchLimit: req.BatchLimit,
|
||||
RecognizeRename: req.RecognizeRename,
|
||||
TransferSubtitles: req.TransferSubtitles,
|
||||
MissingOnly: req.MissingOnly,
|
||||
})
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if req.RefreshLibrary && !req.DryRun {
|
||||
res.Refresh = queueSTRMRefreshAfterChanges(c.Request.Context(), svc, res.OutputDir, strmRefreshQueueOptions{
|
||||
TaskName: "STRM 目录树生成后刷新媒体库",
|
||||
Changed: strmGenerationChanged(res),
|
||||
ScrapeAfter: req.ScrapeAfter,
|
||||
})
|
||||
}
|
||||
c.JSON(http.StatusOK, res)
|
||||
}
|
||||
}
|
||||
|
||||
func repairSTRMHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req repairSTRMReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
strmSvc := svc.STRM
|
||||
if strmSvc == nil {
|
||||
strmSvc = service.NewSTRMService(svc.Log, svc.Repo, svc.Cfg)
|
||||
}
|
||||
baseURL := strings.TrimRight(strings.TrimSpace(req.BaseURL), "/")
|
||||
if baseURL == "" {
|
||||
baseURL = strings.TrimRight(absoluteRequestURL(c, "/"), "/")
|
||||
}
|
||||
res, err := strmSvc.RepairFiles(c.Request.Context(), service.RepairSTRMOptions{
|
||||
OutputDir: req.OutputDir,
|
||||
BaseURL: baseURL,
|
||||
DryRun: req.DryRun,
|
||||
})
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if req.RefreshLibrary && !req.DryRun {
|
||||
res.Refresh = queueSTRMRefreshAfterChanges(c.Request.Context(), svc, res.OutputDir, strmRefreshQueueOptions{
|
||||
TaskName: "STRM 修复后刷新媒体库",
|
||||
Changed: res.Repaired > 0,
|
||||
ScrapeAfter: req.ScrapeAfter,
|
||||
})
|
||||
}
|
||||
c.JSON(http.StatusOK, res)
|
||||
}
|
||||
}
|
||||
|
||||
func strmGenerationChanged(res *service.GenerateSTRMResult) bool {
|
||||
return res != nil && (res.Generated > 0 || res.Updated > 0 || res.Cleaned > 0)
|
||||
}
|
||||
|
||||
type strmRefreshQueueOptions struct {
|
||||
TaskName string
|
||||
Changed bool
|
||||
ScrapeAfter bool
|
||||
}
|
||||
|
||||
type strmRefreshRunOptions struct {
|
||||
ScrapeAfter bool
|
||||
}
|
||||
|
||||
func queueSTRMRefreshAfterChanges(ctx context.Context, svc *service.Container, outputDir string, options strmRefreshQueueOptions) *service.STRMRefreshResult {
|
||||
refresh := &service.STRMRefreshResult{Requested: true, ScrapeRequested: options.ScrapeAfter}
|
||||
if !options.Changed {
|
||||
refresh.Reason = "no strm changes"
|
||||
refresh.ScrapeReason = strmRefreshScrapeSkipReason(refresh)
|
||||
return refresh
|
||||
}
|
||||
if svc == nil || svc.Scan == nil {
|
||||
refresh.Reason = "scanner unavailable"
|
||||
refresh.ScrapeReason = strmRefreshScrapeSkipReason(refresh)
|
||||
return refresh
|
||||
}
|
||||
targets, err := service.FindSTRMRefreshTargets(ctx, svc.Repo, outputDir)
|
||||
if err != nil {
|
||||
refresh.Reason = err.Error()
|
||||
refresh.ScrapeReason = strmRefreshScrapeSkipReason(refresh)
|
||||
return refresh
|
||||
}
|
||||
if len(targets) == 0 {
|
||||
refresh.Reason = "no matching local library"
|
||||
refresh.ScrapeReason = strmRefreshScrapeSkipReason(refresh)
|
||||
return refresh
|
||||
}
|
||||
refresh.Targets = targets
|
||||
if options.ScrapeAfter && svc.Scraper == nil {
|
||||
refresh.ScrapeReason = "scraper unavailable"
|
||||
}
|
||||
for _, target := range targets {
|
||||
key := target.LibraryID
|
||||
if target.RootID != "" {
|
||||
key += ":" + target.RootID
|
||||
}
|
||||
finishScan, ok := svc.Scan.TryBeginLocalScan(key)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
refresh.Queued = true
|
||||
runOptions := strmRefreshRunOptions{ScrapeAfter: options.ScrapeAfter && svc.Scraper != nil}
|
||||
if runOptions.ScrapeAfter {
|
||||
refresh.ScrapeQueued = true
|
||||
}
|
||||
task := startScanHTTPTask(svc, options.TaskName, target.Name, target.Path)
|
||||
go runSTRMRefreshScan(svc, target, task, finishScan, runOptions)
|
||||
}
|
||||
if !refresh.Queued {
|
||||
refresh.Reason = "matching library already scanning"
|
||||
refresh.ScrapeReason = strmRefreshScrapeSkipReason(refresh)
|
||||
}
|
||||
return refresh
|
||||
}
|
||||
|
||||
func runSTRMRefreshScan(svc *service.Container, target service.STRMRefreshTarget, task *service.TaskHandle, finish func(), options strmRefreshRunOptions) {
|
||||
defer finish()
|
||||
var (
|
||||
res *service.ScanResult
|
||||
err error
|
||||
)
|
||||
if target.RootID != "" {
|
||||
res, err = svc.Scan.ScanLibraryRoot(context.Background(), target.LibraryID, target.RootID)
|
||||
} else {
|
||||
res, err = svc.Scan.ScanLibrary(context.Background(), target.LibraryID)
|
||||
}
|
||||
if err != nil {
|
||||
finishHTTPTask(task, err, "scan", "STRM 刷新媒体库失败", scanTaskMetrics(res), scanTaskDetails(res, 20))
|
||||
return
|
||||
}
|
||||
if !options.ScrapeAfter {
|
||||
finishHTTPTask(task, nil, "completed", "STRM 刷新媒体库结束", scanTaskMetrics(res), scanTaskDetails(res, 20))
|
||||
return
|
||||
}
|
||||
scrape, reclassified, scrapeErr := runSTRMRefreshScrape(svc, target, task)
|
||||
metrics := strmRefreshTaskMetrics(res, scrape, reclassified)
|
||||
if scrapeErr != nil {
|
||||
finishHTTPTask(task, scrapeErr, "scrape", "STRM 刷新媒体库完成,刮削失败", metrics, scanTaskDetails(res, 20))
|
||||
return
|
||||
}
|
||||
finishHTTPTask(task, nil, "completed", "STRM 刷新媒体库和刮削结束", metrics, scanTaskDetails(res, 20))
|
||||
}
|
||||
|
||||
func runSTRMRefreshScrape(svc *service.Container, target service.STRMRefreshTarget, task *service.TaskHandle) (service.EnrichLibraryResult, int, error) {
|
||||
if task != nil {
|
||||
task.Update(service.TaskUpdate{
|
||||
Stage: "scrape",
|
||||
SourcePath: target.Path,
|
||||
Message: "正在刮削 STRM 媒体库",
|
||||
})
|
||||
}
|
||||
result, err := svc.Scraper.EnrichLibraryDetailedWithOptions(context.Background(), target.LibraryID, service.ScrapeOptions{RetryNoMatch: true})
|
||||
reclassified := 0
|
||||
if result.Processed > 0 {
|
||||
reclassified = reclassifyLibraryAfterScrape(context.Background(), svc, target.LibraryID)
|
||||
}
|
||||
return result, reclassified, err
|
||||
}
|
||||
|
||||
func strmRefreshTaskMetrics(scan *service.ScanResult, scrape service.EnrichLibraryResult, reclassified int) map[string]int64 {
|
||||
metrics := scanTaskMetrics(scan)
|
||||
if metrics == nil {
|
||||
metrics = map[string]int64{}
|
||||
}
|
||||
metrics["scrape_matched"] = int64(scrape.Matched)
|
||||
metrics["scrape_processed"] = int64(scrape.Processed)
|
||||
metrics["scrape_candidates"] = int64(scrape.Candidates)
|
||||
if scrape.Failed > 0 {
|
||||
metrics["scrape_failed"] = int64(scrape.Failed)
|
||||
}
|
||||
if reclassified > 0 {
|
||||
metrics["scrape_reclassified"] = int64(reclassified)
|
||||
}
|
||||
return metrics
|
||||
}
|
||||
|
||||
func strmRefreshScrapeSkipReason(refresh *service.STRMRefreshResult) string {
|
||||
if refresh == nil || !refresh.ScrapeRequested {
|
||||
return ""
|
||||
}
|
||||
if refresh.Reason != "" {
|
||||
return "refresh not queued: " + refresh.Reason
|
||||
}
|
||||
return "refresh not queued"
|
||||
}
|
||||
|
||||
func strmPlaybackTokenForRequest(c *gin.Context, svc *service.Container) string {
|
||||
if svc == nil || svc.Auth == nil || svc.Repo == nil || svc.Repo.User == nil {
|
||||
return ""
|
||||
}
|
||||
uid := middleware.GetUserID(c)
|
||||
if uid == "" {
|
||||
return ""
|
||||
}
|
||||
u, err := svc.Repo.User.FindByID(c.Request.Context(), uid)
|
||||
if err != nil || u == nil {
|
||||
return ""
|
||||
}
|
||||
token, err := svc.Auth.IssueEmbyToken(u)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return token
|
||||
}
|
||||
@@ -1,47 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func TestSTRMRefreshTaskMetricsIncludesScanAndScrape(t *testing.T) {
|
||||
metrics := strmRefreshTaskMetrics(
|
||||
&service.ScanResult{Visited: 3, Added: 2, Updated: 1, ErrorCount: 1},
|
||||
service.EnrichLibraryResult{Matched: 4, Processed: 5, Candidates: 6, Failed: 1},
|
||||
2,
|
||||
)
|
||||
|
||||
want := map[string]int64{
|
||||
"visited": 3,
|
||||
"added": 2,
|
||||
"updated": 1,
|
||||
"errors": 1,
|
||||
"scrape_matched": 4,
|
||||
"scrape_processed": 5,
|
||||
"scrape_candidates": 6,
|
||||
"scrape_failed": 1,
|
||||
"scrape_reclassified": 2,
|
||||
}
|
||||
for key, value := range want {
|
||||
if metrics[key] != value {
|
||||
t.Fatalf("metrics[%q] = %d, want %d in %#v", key, metrics[key], value, metrics)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSTRMRefreshScrapeSkipReasonRequiresScrapeRequest(t *testing.T) {
|
||||
if got := strmRefreshScrapeSkipReason(&service.STRMRefreshResult{Requested: true, Reason: "no strm changes"}); got != "" {
|
||||
t.Fatalf("skip reason without scrape request = %q, want empty", got)
|
||||
}
|
||||
|
||||
got := strmRefreshScrapeSkipReason(&service.STRMRefreshResult{
|
||||
Requested: true,
|
||||
ScrapeRequested: true,
|
||||
Reason: "no matching local library",
|
||||
})
|
||||
if got != "refresh not queued: no matching local library" {
|
||||
t.Fatalf("skip reason = %q", got)
|
||||
}
|
||||
}
|
||||
@@ -1,22 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func writeSubscriptionConflict(c *gin.Context, err error) bool {
|
||||
if !errors.Is(err, service.ErrSubscriptionAlreadyExists) {
|
||||
return false
|
||||
}
|
||||
body := gin.H{"error": "subscription already exists"}
|
||||
if existingID := service.SubscriptionAlreadyExistsID(err); existingID != "" {
|
||||
body["existing_id"] = existingID
|
||||
}
|
||||
c.JSON(http.StatusConflict, body)
|
||||
return true
|
||||
}
|
||||
@@ -1,32 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func TestWriteSubscriptionConflictReturns409WithExistingID(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
err := &service.SubscriptionAlreadyExistsError{ExistingID: "existing-sub"}
|
||||
if !writeSubscriptionConflict(c, err) {
|
||||
t.Fatal("expected conflict to be handled")
|
||||
}
|
||||
if w.Code != http.StatusConflict {
|
||||
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
var body map[string]any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if body["existing_id"] != "existing-sub" {
|
||||
t.Fatalf("body = %#v", body)
|
||||
}
|
||||
}
|
||||
@@ -1,210 +0,0 @@
|
||||
// Package handler — subscription update + per-subscription site search.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type subscriptionPatchReq struct {
|
||||
Name *string `json:"name"`
|
||||
FeedURL *string `json:"feed_url"`
|
||||
Filter *string `json:"filter"`
|
||||
MediaType *string `json:"media_type"`
|
||||
MediaCategory *string `json:"media_category"`
|
||||
SavePath *string `json:"save_path"`
|
||||
SearchMode *string `json:"search_mode"`
|
||||
IMDBID *string `json:"imdb_id"`
|
||||
Source *string `json:"source"`
|
||||
PosterURL *string `json:"poster_url"`
|
||||
BackdropURL *string `json:"backdrop_url"`
|
||||
Overview *string `json:"overview"`
|
||||
OriginalName *string `json:"original_name"`
|
||||
Year *int `json:"year"`
|
||||
Rating *float32 `json:"rating"`
|
||||
Genres *string `json:"genres"`
|
||||
Resolution *string `json:"resolution"`
|
||||
Quality *string `json:"quality"`
|
||||
Effects *string `json:"effects"`
|
||||
ReleaseGroups *string `json:"release_groups"`
|
||||
ExcludeWords *string `json:"exclude_words"`
|
||||
MinSeeders *int `json:"min_seeders"`
|
||||
MaxSeeders *int `json:"max_seeders"`
|
||||
MinSizeGB *float64 `json:"min_size_gb"`
|
||||
MaxSizeGB *float64 `json:"max_size_gb"`
|
||||
FreeOnly *bool `json:"free_only"`
|
||||
WashEnabled *bool `json:"wash_enabled"`
|
||||
WashPriority *string `json:"wash_priority"`
|
||||
TotalEpisodes *int `json:"total_episodes"`
|
||||
Priority *int `json:"priority"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// updateSubscriptionHandler patches a subscription row.
|
||||
func updateSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var patch subscriptionPatchReq
|
||||
if err := c.ShouldBindJSON(&patch); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
updates := subscriptionPatchUpdates(patch)
|
||||
if len(updates) == 0 {
|
||||
c.Status(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
if err := svc.Subscription.Update(c.Request.Context(), c.Param("id"), updates); err != nil {
|
||||
logSubscriptionWarn(svc, "subscription update failed",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.String("subscription_id", c.Param("id")),
|
||||
zap.Error(err))
|
||||
if writeSubscriptionConflict(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
logSubscriptionInfo(svc, "subscription updated",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.String("subscription_id", c.Param("id")),
|
||||
zap.Strings("fields", subscriptionUpdateFieldNames(updates)))
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
func subscriptionPatchUpdates(patch subscriptionPatchReq) map[string]any {
|
||||
updates := map[string]any{}
|
||||
if patch.Name != nil {
|
||||
updates["name"] = *patch.Name
|
||||
}
|
||||
if patch.FeedURL != nil {
|
||||
updates["feed_url"] = *patch.FeedURL
|
||||
}
|
||||
if patch.Filter != nil {
|
||||
updates["filter"] = *patch.Filter
|
||||
}
|
||||
if patch.MediaType != nil {
|
||||
updates["media_type"] = *patch.MediaType
|
||||
}
|
||||
if patch.MediaCategory != nil {
|
||||
updates["media_category"] = *patch.MediaCategory
|
||||
}
|
||||
if patch.SavePath != nil {
|
||||
updates["save_path"] = *patch.SavePath
|
||||
}
|
||||
if patch.SearchMode != nil {
|
||||
updates["search_mode"] = *patch.SearchMode
|
||||
}
|
||||
if patch.IMDBID != nil {
|
||||
updates["imdb_id"] = *patch.IMDBID
|
||||
}
|
||||
if patch.Source != nil {
|
||||
updates["source"] = *patch.Source
|
||||
}
|
||||
if patch.PosterURL != nil {
|
||||
updates["poster_url"] = *patch.PosterURL
|
||||
}
|
||||
if patch.BackdropURL != nil {
|
||||
updates["backdrop_url"] = *patch.BackdropURL
|
||||
}
|
||||
if patch.Overview != nil {
|
||||
updates["overview"] = *patch.Overview
|
||||
}
|
||||
if patch.OriginalName != nil {
|
||||
updates["original_name"] = *patch.OriginalName
|
||||
}
|
||||
if patch.Year != nil {
|
||||
updates["year"] = *patch.Year
|
||||
}
|
||||
if patch.Rating != nil {
|
||||
updates["rating"] = *patch.Rating
|
||||
}
|
||||
if patch.Genres != nil {
|
||||
updates["genres"] = *patch.Genres
|
||||
}
|
||||
if patch.Resolution != nil {
|
||||
updates["resolution"] = *patch.Resolution
|
||||
}
|
||||
if patch.Quality != nil {
|
||||
updates["quality"] = *patch.Quality
|
||||
}
|
||||
if patch.Effects != nil {
|
||||
updates["effects"] = *patch.Effects
|
||||
}
|
||||
if patch.ReleaseGroups != nil {
|
||||
updates["release_groups"] = *patch.ReleaseGroups
|
||||
}
|
||||
if patch.ExcludeWords != nil {
|
||||
updates["exclude_words"] = *patch.ExcludeWords
|
||||
}
|
||||
if patch.MinSeeders != nil {
|
||||
updates["min_seeders"] = *patch.MinSeeders
|
||||
}
|
||||
if patch.MaxSeeders != nil {
|
||||
updates["max_seeders"] = *patch.MaxSeeders
|
||||
}
|
||||
if patch.MinSizeGB != nil {
|
||||
updates["min_size_gb"] = *patch.MinSizeGB
|
||||
}
|
||||
if patch.MaxSizeGB != nil {
|
||||
updates["max_size_gb"] = *patch.MaxSizeGB
|
||||
}
|
||||
if patch.FreeOnly != nil {
|
||||
updates["free_only"] = *patch.FreeOnly
|
||||
}
|
||||
if patch.WashEnabled != nil {
|
||||
updates["wash_enabled"] = *patch.WashEnabled
|
||||
}
|
||||
if patch.WashPriority != nil {
|
||||
updates["wash_priority"] = *patch.WashPriority
|
||||
}
|
||||
if patch.TotalEpisodes != nil {
|
||||
updates["total_episodes"] = *patch.TotalEpisodes
|
||||
}
|
||||
if patch.Priority != nil {
|
||||
updates["priority"] = *patch.Priority
|
||||
}
|
||||
if patch.Enabled != nil {
|
||||
updates["enabled"] = *patch.Enabled
|
||||
}
|
||||
return updates
|
||||
}
|
||||
|
||||
func subscriptionUpdateFieldNames(updates map[string]any) []string {
|
||||
names := make([]string, 0, len(updates))
|
||||
for name := range updates {
|
||||
names = append(names, name)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// searchSubscriptionHandler runs a one-off keyword search against the
|
||||
// configured tracker sites for the given subscription. We treat the
|
||||
// subscription's filter as the search term; this lets the UI preview
|
||||
// what would be queued without actually downloading anything.
|
||||
func searchSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var sub model.Subscription
|
||||
err := svc.Repo.DB.WithContext(c.Request.Context()).
|
||||
Where("id = ?", c.Param("id")).First(&sub).Error
|
||||
if err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "subscription not found"})
|
||||
return
|
||||
}
|
||||
keyword := sub.Filter
|
||||
if keyword == "" {
|
||||
keyword = sub.Name
|
||||
}
|
||||
results, err := svc.Site.Search(c.Request.Context(), keyword)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": results, "subscription": sub})
|
||||
}
|
||||
}
|
||||
@@ -1,54 +0,0 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/middleware"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
func subscriptionRequestUserID(c *gin.Context) string {
|
||||
if c == nil {
|
||||
return ""
|
||||
}
|
||||
if uid, ok := c.Get(middleware.CtxUserID); ok {
|
||||
if userID, ok := uid.(string); ok {
|
||||
return userID
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func subscriptionFeedKind(feedURL string) string {
|
||||
raw := strings.TrimSpace(feedURL)
|
||||
if raw == "" {
|
||||
return "empty"
|
||||
}
|
||||
lower := strings.ToLower(raw)
|
||||
if strings.HasPrefix(lower, "site-search://") {
|
||||
return "site-search"
|
||||
}
|
||||
parsed, err := url.Parse(raw)
|
||||
if err == nil && parsed.Scheme != "" {
|
||||
return parsed.Scheme
|
||||
}
|
||||
return "unknown"
|
||||
}
|
||||
|
||||
func logSubscriptionInfo(svc *service.Container, msg string, fields ...zap.Field) {
|
||||
if svc == nil || svc.Log == nil {
|
||||
return
|
||||
}
|
||||
svc.Log.Info(msg, fields...)
|
||||
}
|
||||
|
||||
func logSubscriptionWarn(svc *service.Container, msg string, fields ...zap.Field) {
|
||||
if svc == nil || svc.Log == nil {
|
||||
return
|
||||
}
|
||||
svc.Log.Warn(msg, fields...)
|
||||
}
|
||||
@@ -1,219 +0,0 @@
|
||||
// Package handler — RSS subscription endpoints.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/middleware"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
type subscriptionReq struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
FeedURL string `json:"feed_url" binding:"required"`
|
||||
Filter string `json:"filter"`
|
||||
MediaType string `json:"media_type"`
|
||||
MediaCategory string `json:"media_category"`
|
||||
SavePath string `json:"save_path"`
|
||||
SearchMode string `json:"search_mode"`
|
||||
IMDBID string `json:"imdb_id"`
|
||||
Source string `json:"source"`
|
||||
PosterURL string `json:"poster_url"`
|
||||
BackdropURL string `json:"backdrop_url"`
|
||||
Overview string `json:"overview"`
|
||||
OriginalName string `json:"original_name"`
|
||||
Year int `json:"year"`
|
||||
Resolution string `json:"resolution"`
|
||||
Quality string `json:"quality"`
|
||||
Effects string `json:"effects"`
|
||||
ReleaseGroups string `json:"release_groups"`
|
||||
ExcludeWords string `json:"exclude_words"`
|
||||
MinSeeders int `json:"min_seeders"`
|
||||
MaxSeeders int `json:"max_seeders"`
|
||||
MinSizeGB float64 `json:"min_size_gb"`
|
||||
MaxSizeGB float64 `json:"max_size_gb"`
|
||||
FreeOnly bool `json:"free_only"`
|
||||
WashEnabled bool `json:"wash_enabled"`
|
||||
WashPriority string `json:"wash_priority"`
|
||||
TotalEpisodes int `json:"total_episodes"`
|
||||
Priority int `json:"priority"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
|
||||
func createSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req subscriptionReq
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
uid, _ := c.Get(middleware.CtxUserID)
|
||||
enabled := true
|
||||
if req.Enabled != nil {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
s := &model.Subscription{
|
||||
UserID: uid.(string),
|
||||
Name: req.Name,
|
||||
FeedURL: req.FeedURL,
|
||||
Filter: req.Filter,
|
||||
MediaType: req.MediaType,
|
||||
MediaCategory: req.MediaCategory,
|
||||
SavePath: req.SavePath,
|
||||
SearchMode: req.SearchMode,
|
||||
IMDBID: req.IMDBID,
|
||||
Source: req.Source,
|
||||
PosterURL: req.PosterURL,
|
||||
BackdropURL: req.BackdropURL,
|
||||
Overview: req.Overview,
|
||||
OriginalName: req.OriginalName,
|
||||
Year: req.Year,
|
||||
Resolution: req.Resolution,
|
||||
Quality: req.Quality,
|
||||
Effects: req.Effects,
|
||||
ReleaseGroups: req.ReleaseGroups,
|
||||
ExcludeWords: req.ExcludeWords,
|
||||
MinSeeders: req.MinSeeders,
|
||||
MaxSeeders: req.MaxSeeders,
|
||||
MinSizeGB: req.MinSizeGB,
|
||||
MaxSizeGB: req.MaxSizeGB,
|
||||
FreeOnly: req.FreeOnly,
|
||||
WashEnabled: req.WashEnabled,
|
||||
WashPriority: req.WashPriority,
|
||||
TotalEpisodes: req.TotalEpisodes,
|
||||
Priority: req.Priority,
|
||||
Enabled: enabled,
|
||||
}
|
||||
enrichSubscriptionArtwork(c.Request.Context(), svc, s)
|
||||
if err := svc.Subscription.Create(c.Request.Context(), s); err != nil {
|
||||
logSubscriptionWarn(svc, "subscription create failed",
|
||||
zap.String("user_id", s.UserID),
|
||||
zap.String("name", req.Name),
|
||||
zap.String("feed_kind", subscriptionFeedKind(req.FeedURL)),
|
||||
zap.Bool("enabled", enabled),
|
||||
zap.Error(err))
|
||||
if writeSubscriptionConflict(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
logSubscriptionInfo(svc, "subscription created",
|
||||
zap.String("user_id", s.UserID),
|
||||
zap.String("subscription_id", s.ID),
|
||||
zap.String("name", s.Name),
|
||||
zap.String("feed_kind", subscriptionFeedKind(s.FeedURL)),
|
||||
zap.String("media_type", s.MediaType),
|
||||
zap.String("media_category", s.MediaCategory),
|
||||
zap.Bool("enabled", s.Enabled),
|
||||
zap.Bool("wash_enabled", s.WashEnabled))
|
||||
enriched := []model.Subscription{*s}
|
||||
svc.Subscription.EnrichManagementProgress(c.Request.Context(), enriched)
|
||||
*s = enriched[0]
|
||||
c.JSON(http.StatusCreated, s)
|
||||
}
|
||||
}
|
||||
|
||||
func listSubscriptionsHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
items, err := svc.Subscription.List(c.Request.Context())
|
||||
if err != nil {
|
||||
logSubscriptionWarn(svc, "subscription list failed",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.Error(err))
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
svc.Subscription.EnrichManagementProgress(c.Request.Context(), items)
|
||||
go enrichAndPersistSubscriptions(context.Background(), svc, append([]model.Subscription(nil), items...))
|
||||
logSubscriptionInfo(svc, "subscription list returned",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.Int("count", len(items)),
|
||||
zap.Bool("history", false))
|
||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||
}
|
||||
}
|
||||
|
||||
func listSubscriptionHistoryHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
items, err := svc.Subscription.History(c.Request.Context())
|
||||
if err != nil {
|
||||
logSubscriptionWarn(svc, "subscription history list failed",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.Error(err))
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
svc.Subscription.EnrichManagementProgress(c.Request.Context(), items)
|
||||
logSubscriptionInfo(svc, "subscription list returned",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.Int("count", len(items)),
|
||||
zap.Bool("history", true))
|
||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||
}
|
||||
}
|
||||
|
||||
func deleteSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if err := svc.Subscription.Delete(c.Request.Context(), c.Param("id")); err != nil {
|
||||
logSubscriptionWarn(svc, "subscription delete failed",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.String("subscription_id", c.Param("id")),
|
||||
zap.Error(err))
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
logSubscriptionInfo(svc, "subscription deleted",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.String("subscription_id", c.Param("id")))
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
func runSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
n, err := svc.Subscription.RunNow(c.Request.Context(), c.Param("id"))
|
||||
if err != nil {
|
||||
logSubscriptionWarn(svc, "subscription run now failed",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.String("subscription_id", c.Param("id")),
|
||||
zap.Error(err))
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
logSubscriptionInfo(svc, "subscription run now completed",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.String("subscription_id", c.Param("id")),
|
||||
zap.Int("queued", n))
|
||||
c.JSON(http.StatusOK, gin.H{"queued": n})
|
||||
}
|
||||
}
|
||||
|
||||
func restoreSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
sub, err := svc.Subscription.Restore(c.Request.Context(), c.Param("id"))
|
||||
if err != nil {
|
||||
logSubscriptionWarn(svc, "subscription restore failed",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.String("subscription_id", c.Param("id")),
|
||||
zap.Error(err))
|
||||
if writeSubscriptionConflict(c, err) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
logSubscriptionInfo(svc, "subscription restored",
|
||||
zap.String("user_id", subscriptionRequestUserID(c)),
|
||||
zap.String("subscription_id", sub.ID),
|
||||
zap.String("name", sub.Name))
|
||||
enriched := []model.Subscription{*sub}
|
||||
svc.Subscription.EnrichManagementProgress(c.Request.Context(), enriched)
|
||||
c.JSON(http.StatusOK, enriched[0])
|
||||
}
|
||||
}
|
||||
@@ -58,13 +58,12 @@ func schemaHandler(_ *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"groups": []gin.H{
|
||||
{
|
||||
"key": "general",
|
||||
"label": "常规",
|
||||
"items": []gin.H{
|
||||
{"key": "tmdb.language", "type": "select", "label": "TMDb 元数据语言"},
|
||||
{"key": "app.server_url", "type": "text", "label": "公开访问域名 / STRM 域名"},
|
||||
{"key": "transcode.enabled", "type": "toggle", "label": "启用转码"},
|
||||
{
|
||||
"key": "general",
|
||||
"label": "常规",
|
||||
"items": []gin.H{
|
||||
{"key": "tmdb.language", "type": "select", "label": "TMDb 元数据语言"},
|
||||
{"key": "transcode.enabled", "type": "toggle", "label": "启用转码"},
|
||||
{"key": "transcode.hw_accel", "type": "select", "label": "硬件编码器"},
|
||||
{"key": "transcode.hw_enabled", "type": "toggle", "label": "启用硬件加速"},
|
||||
{"key": "transcode.max_jobs", "type": "number", "label": "最大并发"},
|
||||
@@ -121,19 +120,10 @@ func schemaHandler(_ *service.Container) gin.HandlerFunc {
|
||||
{"key": "qbittorrent.password", "type": "text"},
|
||||
{"key": "qbittorrent.savepath", "type": "text"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"key": "license",
|
||||
"label": "授权服务",
|
||||
"items": []gin.H{
|
||||
{"key": "license.server_url", "type": "text", "label": "License Server 地址"},
|
||||
{"key": "license.public_key", "type": "text", "label": "Ed25519 验签公钥"},
|
||||
{"key": "license.hmac_secret", "type": "text", "label": "HMAC 签名密钥(旧版兼容)"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"key": "system-update",
|
||||
"label": "系统更新",
|
||||
},
|
||||
{
|
||||
"key": "system-update",
|
||||
"label": "系统更新",
|
||||
"items": []gin.H{
|
||||
{"key": "system.update.image", "type": "text", "label": "应用镜像"},
|
||||
{"key": "system.update.compose_dir", "type": "text", "label": "Docker Compose 安装目录"},
|
||||
@@ -145,16 +135,6 @@ func schemaHandler(_ *service.Container) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// schedulerTriggerHandler is the alternate path for /admin/scheduler/:name/run.
|
||||
func schedulerTriggerHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !triggerSchedulerJob(c, svc, c.Param("name")) {
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusAccepted, gin.H{"ok": true, "message": "任务已在后台触发"})
|
||||
}
|
||||
}
|
||||
|
||||
// ─── SSE ticket store ───────────────────────────────────────────────────────
|
||||
//
|
||||
// The Vue UI's SSE event stream wants a one-time signed ticket so the
|
||||
|
||||
@@ -63,11 +63,3 @@ func systemStatusHandler(svc *service.Container) gin.HandlerFunc {
|
||||
c.JSON(http.StatusOK, out)
|
||||
}
|
||||
}
|
||||
|
||||
// systemSchedulerHandler is the read-only (non-admin) variant of
|
||||
// /admin/scheduler — handy on user-facing dashboards.
|
||||
func systemSchedulerHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"jobs": svc.Scheduler.Status()})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,4 +7,4 @@ func finishHTTPTask(task *service.TaskHandle, err error, stage, message string,
|
||||
return
|
||||
}
|
||||
task.Finish(err, service.TaskUpdate{Stage: stage, Message: message, Metrics: metrics, Details: details})
|
||||
}
|
||||
}
|
||||
@@ -1,40 +0,0 @@
|
||||
// Package handler — live tasks board.
|
||||
//
|
||||
// /api/tasks aggregates running ffmpeg transcodes, qBittorrent torrents
|
||||
// and recent scrape progress into a single snapshot suitable for the
|
||||
// React Tasks panel. The panel can layer this REST snapshot on top of
|
||||
// live WS events for instant updates.
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
const tasksLiveTorrentSnapshotMaxAge = 30 * time.Second
|
||||
|
||||
func tasksHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var transcodes []service.ActiveJob
|
||||
if svc.Transcoder != nil {
|
||||
transcodes = svc.Transcoder.Active()
|
||||
}
|
||||
var torrents []service.QBitTorrent
|
||||
if svc.Downloads != nil {
|
||||
torrents = svc.Downloads.LiveTorrentSnapshot(tasksLiveTorrentSnapshotMaxAge)
|
||||
}
|
||||
background := service.TaskSnapshot{}
|
||||
if svc.Tasks != nil {
|
||||
background = svc.Tasks.Snapshot()
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"transcodes": transcodes,
|
||||
"torrents": torrents,
|
||||
"background_tasks": background,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,104 +0,0 @@
|
||||
// Package handler — Telegram Bot Webhook 端点。
|
||||
//
|
||||
// 接收 Telegram Bot API 推送的 update 消息,路由到 TelegramBotService 处理。
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||
)
|
||||
|
||||
// telegramWebhookHandler 处理 Telegram Bot 的 Webhook 回调。
|
||||
//
|
||||
// 路由:POST /api/telegram/webhook (无需认证,由 Telegram 服务器调用)
|
||||
func telegramWebhookHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
body, err := io.ReadAll(c.Request.Body)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "cannot read body"})
|
||||
return
|
||||
}
|
||||
|
||||
if err := svc.TelegramBot.HandleWebhook(c.Request.Context(), body); err != nil {
|
||||
svc.Log.Error("telegram webhook error", zap.Error(err))
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
}
|
||||
}
|
||||
|
||||
// telegramSetWebhookHandler 管理员手动设置/更新 Webhook URL。
|
||||
//
|
||||
// 路由:POST /api/admin/telegram/webhook (需 admin 认证)
|
||||
func telegramSetWebhookHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
var req struct {
|
||||
BotToken string `json:"bot_token" binding:"required"`
|
||||
WebhookURL string `json:"webhook_url" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
if err := svc.TelegramBot.SetWebhook(c.Request.Context(), req.BotToken, req.WebhookURL); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{"message": "webhook set successfully", "url": req.WebhookURL})
|
||||
}
|
||||
}
|
||||
|
||||
// telegramGetWebhookHandler 获取当前 Webhook 配置信息。
|
||||
//
|
||||
// 路由:GET /api/admin/telegram/webhook (需 admin 认证)
|
||||
func telegramGetWebhookHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
botToken := c.Query("bot_token")
|
||||
if botToken == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "bot_token is required"})
|
||||
return
|
||||
}
|
||||
|
||||
info, err := svc.TelegramBot.GetWebhookInfo(c.Request.Context(), botToken)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, info)
|
||||
}
|
||||
}
|
||||
|
||||
// telegramStartPollingHandler 启动 Telegram 长轮询。
|
||||
//
|
||||
// 路由:POST /api/admin/telegram/polling/start (需 admin 认证)
|
||||
func telegramStartPollingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
setupCtx, cancel := context.WithTimeout(svc.Context(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
result := svc.TelegramBot.StartPolling(setupCtx)
|
||||
c.JSON(http.StatusOK, result)
|
||||
}
|
||||
}
|
||||
|
||||
// telegramStopPollingHandler 停止 Telegram 长轮询。
|
||||
//
|
||||
// 路由:POST /api/admin/telegram/polling/stop (需 admin 认证)
|
||||
func telegramStopPollingHandler(svc *service.Container) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
stopped := svc.TelegramBot.StopPolling()
|
||||
c.JSON(http.StatusOK, gin.H{"message": "polling stopped", "stopped": stopped})
|
||||
}
|
||||
}
|
||||
@@ -1,128 +0,0 @@
|
||||
package helper
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// FlareSolverrRequest represents a request to FlareSolverr.
|
||||
type FlareSolverrRequest struct {
|
||||
Cmd string `json:"cmd"`
|
||||
URL string `json:"url"`
|
||||
Session string `json:"session,omitempty"`
|
||||
MaxTimeout int `json:"maxTimeout,omitempty"`
|
||||
Proxy *FlareSolverrProxy `json:"proxy,omitempty"`
|
||||
Cookies []FlareSolverrCookie `json:"cookies,omitempty"`
|
||||
}
|
||||
|
||||
// FlareSolverrProxy represents proxy config for FlareSolverr.
|
||||
type FlareSolverrProxy struct {
|
||||
URL string `json:"url"`
|
||||
Username string `json:"username,omitempty"`
|
||||
Password string `json:"password,omitempty"`
|
||||
}
|
||||
|
||||
// FlareSolverrCookie represents a cookie for FlareSolverr.
|
||||
type FlareSolverrCookie struct {
|
||||
Name string `json:"name"`
|
||||
Value string `json:"value"`
|
||||
Domain string `json:"domain,omitempty"`
|
||||
Path string `json:"path,omitempty"`
|
||||
}
|
||||
|
||||
// FlareSolverrResponse represents FlareSolverr's response.
|
||||
type FlareSolverrResponse struct {
|
||||
Status string `json:"status"`
|
||||
Message string `json:"message"`
|
||||
Solution *FlareSolverrSolution `json:"solution,omitempty"`
|
||||
}
|
||||
|
||||
// FlareSolverrSolution contains the solved challenge result.
|
||||
type FlareSolverrSolution struct {
|
||||
URL string `json:"url"`
|
||||
Status int `json:"status"`
|
||||
Headers map[string]string `json:"headers"`
|
||||
Cookies []FlareSolverrCookie `json:"cookies"`
|
||||
UserAgent string `json:"userAgent"`
|
||||
Response string `json:"response"`
|
||||
}
|
||||
|
||||
// FetchURLWithFlareSolverr uses FlareSolverr to fetch a URL,
|
||||
// bypassing Cloudflare/WAF challenges.
|
||||
func FetchURLWithFlareSolverr(flareSolverrURL string, targetURL string, cookieStr string, timeout int, proxyURL string, log *zap.Logger) (string, error) {
|
||||
if flareSolverrURL == "" {
|
||||
return "", fmt.Errorf("FlareSolverr URL not configured")
|
||||
}
|
||||
if timeout <= 0 {
|
||||
timeout = 60
|
||||
}
|
||||
|
||||
var cookies []FlareSolverrCookie
|
||||
if cookieStr != "" {
|
||||
cookies = parseCookiesForFlareSolverr(cookieStr)
|
||||
}
|
||||
|
||||
reqBody := FlareSolverrRequest{
|
||||
Cmd: "request.get",
|
||||
URL: targetURL,
|
||||
MaxTimeout: timeout * 1000,
|
||||
Cookies: cookies,
|
||||
}
|
||||
if proxyURL != "" {
|
||||
reqBody.Proxy = &FlareSolverrProxy{URL: proxyURL}
|
||||
}
|
||||
|
||||
jsonBody, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to marshal FlareSolverr request: %w", err)
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: time.Duration(timeout+10) * time.Second}
|
||||
resp, err := client.Post(flareSolverrURL, "application/json", strings.NewReader(string(jsonBody)))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("FlareSolverr request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to read FlareSolverr response: %w", err)
|
||||
}
|
||||
|
||||
var fsResp FlareSolverrResponse
|
||||
if err := json.Unmarshal(body, &fsResp); err != nil {
|
||||
return "", fmt.Errorf("failed to parse FlareSolverr response: %w", err)
|
||||
}
|
||||
|
||||
if fsResp.Status != "ok" {
|
||||
return "", fmt.Errorf("FlareSolverr error: %s", fsResp.Message)
|
||||
}
|
||||
|
||||
if fsResp.Solution != nil {
|
||||
return fsResp.Solution.Response, nil
|
||||
}
|
||||
return "", fmt.Errorf("FlareSolverr returned no solution")
|
||||
}
|
||||
|
||||
// parseCookiesForFlareSolverr converts a cookie header string to FlareSolverr format.
|
||||
func parseCookiesForFlareSolverr(cookieStr string) []FlareSolverrCookie {
|
||||
var cookies []FlareSolverrCookie
|
||||
parts := strings.Split(cookieStr, ";")
|
||||
for _, part := range parts {
|
||||
part = strings.TrimSpace(part)
|
||||
kv := strings.SplitN(part, "=", 2)
|
||||
if len(kv) == 2 {
|
||||
cookies = append(cookies, FlareSolverrCookie{
|
||||
Name: kv[0],
|
||||
Value: kv[1],
|
||||
})
|
||||
}
|
||||
}
|
||||
return cookies
|
||||
}
|
||||
@@ -6,10 +6,11 @@ import (
|
||||
"net/http"
|
||||
"net/url"
|
||||
"time"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
)
|
||||
|
||||
// defaultUserAgent 是默认浏览器 User-Agent(用于 HTTP 请求头)。
|
||||
const defaultUserAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/122.0.0.0 Safari/537.36"
|
||||
|
||||
// NewSiteHTTPClient builds an http.Client honoring per-site policies:
|
||||
// - timeout (seconds, defaults to 15)
|
||||
// - proxy via HTTP(S)_PROXY environment variables when site.UseProxy is on
|
||||
@@ -42,7 +43,7 @@ func NewSiteHTTPClient(timeoutSeconds int, useProxy bool) *http.Client {
|
||||
// These mimic a real Chrome browser to avoid WAF/bot detection.
|
||||
func HTTPHeaderPresets() map[string]string {
|
||||
return map[string]string{
|
||||
"User-Agent": model.DefaultUserAgent,
|
||||
"User-Agent": defaultUserAgent,
|
||||
"Accept": "text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7",
|
||||
"Accept-Language": "zh-CN,zh;q=0.9,en;q=0.8",
|
||||
"Accept-Encoding": "gzip, deflate, br",
|
||||
|
||||
@@ -1,48 +0,0 @@
|
||||
package helper
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// GetPageSource fetches a page with browser-like headers.
|
||||
// Returns (pageSource, cookies, error).
|
||||
func GetPageSource(url string, site *model.Site, timeout int, log *zap.Logger) (string, string, error) {
|
||||
client := NewSiteHTTPClient(timeout, site.UseProxy)
|
||||
|
||||
req, err := http.NewRequest("GET", url, nil)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
|
||||
headers := HTTPHeaderPresets()
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
|
||||
ApplySiteAuthHeaders(req, site)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
|
||||
cookies := ""
|
||||
for _, c := range resp.Cookies() {
|
||||
if cookies != "" {
|
||||
cookies += "; "
|
||||
}
|
||||
cookies += c.Name + "=" + c.Value
|
||||
}
|
||||
|
||||
return string(body), cookies, nil
|
||||
}
|
||||
@@ -1,50 +0,0 @@
|
||||
package helper
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
)
|
||||
|
||||
// ApplySiteAuthHeaders applies authentication headers based on site config.
|
||||
func ApplySiteAuthHeaders(req *http.Request, site *model.Site) {
|
||||
switch site.AuthType {
|
||||
case "cookie":
|
||||
if site.Cookie != "" {
|
||||
req.Header.Set("Cookie", site.Cookie)
|
||||
}
|
||||
case "api_key":
|
||||
if site.APIKey != "" {
|
||||
if isYemaPTSite(site) {
|
||||
req.Header.Set("Authorization", site.APIKey)
|
||||
} else {
|
||||
req.Header.Set("x-api-key", site.APIKey)
|
||||
}
|
||||
}
|
||||
case "auth_header":
|
||||
if site.AuthHeader != "" {
|
||||
req.Header.Set("Authorization", site.AuthHeader)
|
||||
}
|
||||
}
|
||||
|
||||
if site.UserAgent != "" {
|
||||
req.Header.Set("User-Agent", site.UserAgent)
|
||||
}
|
||||
}
|
||||
|
||||
func isYemaPTSite(site *model.Site) bool {
|
||||
if site == nil {
|
||||
return false
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(site.Type), "yemapt") {
|
||||
return true
|
||||
}
|
||||
u, err := url.Parse(strings.TrimSpace(site.URL))
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
host := strings.ToLower(u.Hostname())
|
||||
return host == "yemapt.org" || strings.HasSuffix(host, ".yemapt.org")
|
||||
}
|
||||
@@ -1,96 +0,0 @@
|
||||
package helper
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// TestSiteConnectivity performs a site connectivity test with browser-like headers.
|
||||
// If flareSolverrURL is non-empty AND the site has BrowserEmulation turned on,
|
||||
// it will attempt to use FlareSolverr first.
|
||||
// Returns (ok, message, error).
|
||||
func TestSiteConnectivity(site *model.Site, flareSolverrURL string, timeout int, log *zap.Logger) (bool, string, error) {
|
||||
// Try FlareSolverr first when (a) globally enabled and (b) the site
|
||||
// asked for browser emulation. This matches the contract used by the
|
||||
// search path (see service.SiteService.siteModelToConfig).
|
||||
useFlare := flareSolverrURL != "" && site.BrowserEmulation
|
||||
if useFlare {
|
||||
log.Info("Trying FlareSolverr for site test", zap.String("url", site.URL))
|
||||
body, err := FetchURLWithFlareSolverr(flareSolverrURL, site.URL, site.Cookie, timeout, "", log)
|
||||
if err == nil {
|
||||
// Successfully got page via FlareSolverr
|
||||
if IsCloudflareChallenge(body) {
|
||||
return false, "站点被 Cloudflare/WAF 拦截,但 FlareSolverr 未能完全解决", nil
|
||||
}
|
||||
return true, "连接成功 (via FlareSolverr)", nil
|
||||
}
|
||||
log.Warn("FlareSolverr failed, falling back to direct request", zap.Error(err))
|
||||
// Fall through to direct request
|
||||
}
|
||||
|
||||
// Direct HTTP request with browser-like headers. Honors HTTP(S)_PROXY
|
||||
// when the site has UseProxy enabled — this makes the "use proxy"
|
||||
// checkbox in the UI actually do something.
|
||||
client := NewSiteHTTPClient(timeout, site.UseProxy)
|
||||
client.CheckRedirect = func(req *http.Request, via []*http.Request) error {
|
||||
if len(via) >= 10 {
|
||||
return fmt.Errorf("too many redirects")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
req, err := http.NewRequest("GET", site.URL, nil)
|
||||
if err != nil {
|
||||
return false, err.Error(), nil
|
||||
}
|
||||
|
||||
headers := HTTPHeaderPresets()
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
|
||||
ApplySiteAuthHeaders(req, site)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return false, err.Error(), nil
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
bodyStr := string(body)
|
||||
|
||||
if IsCloudflareChallenge(bodyStr) {
|
||||
log.Warn("Cloudflare challenge detected", zap.String("url", site.URL))
|
||||
return false, "站点被 Cloudflare/WAF 拦截,请配置 FlareSolverr 或浏览器模拟", nil
|
||||
}
|
||||
|
||||
// Evaluate status code (mirror the reference Python project's semantics:
|
||||
// 200 → success, 3xx → success/redirect-to-self, 401/403 → failure with
|
||||
// hint to check credentials, 4xx/5xx → failure with raw status text).
|
||||
switch {
|
||||
case resp.StatusCode >= 200 && resp.StatusCode < 300:
|
||||
return true, fmt.Sprintf("连接成功 (%s)", resp.Status), nil
|
||||
case resp.StatusCode == 301 || resp.StatusCode == 302 || resp.StatusCode == 307 || resp.StatusCode == 308:
|
||||
loc := resp.Header.Get("Location")
|
||||
if loc == "" {
|
||||
loc = "(unknown)"
|
||||
}
|
||||
// Most PT sites redirect logged-out users to login; treat as failure.
|
||||
return false, fmt.Sprintf("未登录或 Cookie 失效(重定向至 %s)", loc), nil
|
||||
case resp.StatusCode == 401:
|
||||
return false, "未授权(HTTP 401),请检查 API Key / Cookie", nil
|
||||
case resp.StatusCode == 403:
|
||||
return false, "认证失败(HTTP 403),请检查 Cookie / API Key 或站点是否需要浏览器模拟", nil
|
||||
case resp.StatusCode == 429:
|
||||
return false, "请求被限流(HTTP 429),请稍后再试", nil
|
||||
case resp.StatusCode == 503:
|
||||
return false, "服务暂时不可用(HTTP 503)", nil
|
||||
default:
|
||||
return false, resp.Status, nil
|
||||
}
|
||||
}
|
||||
@@ -106,13 +106,13 @@ func TestAuthRequiredSyncsAccessTokenCookieFromBearer(t *testing.T) {
|
||||
},
|
||||
})
|
||||
|
||||
router := gin.New()
|
||||
router.Use(AuthRequired(secret))
|
||||
router.GET("/api/discover/feed", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
})
|
||||
router := gin.New()
|
||||
router.Use(AuthRequired(secret))
|
||||
router.GET("/api/test-auth-cookie", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/discover/feed", nil)
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/test-auth-cookie", nil)
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
// Package model 定义下载客户端配置数据模型。
|
||||
package model
|
||||
|
||||
// DownloadClient 下载客户端配置(qBittorrent / Transmission / Aria2)。
|
||||
// 支持多客户端并行运行,一个标记为默认客户端。
|
||||
type DownloadClient struct {
|
||||
Base
|
||||
Name string `gorm:"size:128;not null" json:"name"`
|
||||
Type string `gorm:"size:32;not null" json:"type"` // qbittorrent / transmission / aria2
|
||||
Host string `gorm:"size:512;not null" json:"host"` // http://host:port
|
||||
Username string `gorm:"size:256" json:"username"`
|
||||
Password string `gorm:"size:1024" json:"-"` // AES加密存储
|
||||
IsDefault bool `gorm:"default:false" json:"is_default"`
|
||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||
Extra string `gorm:"type:text" json:"-"` // JSON配置, AES加密
|
||||
}
|
||||
@@ -1,81 +0,0 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
// DownloadTask 是待处理(或已完成)的 torrent / HTTP 下载。
|
||||
type DownloadTask struct {
|
||||
Base
|
||||
UserID string `gorm:"index;size:36" json:"user_id"`
|
||||
SubscriptionID string `gorm:"index;size:36" json:"subscription_id,omitempty"`
|
||||
DownloadClientID string `gorm:"index;size:36" json:"download_client_id,omitempty"`
|
||||
ExternalID string `gorm:"index;size:128" json:"external_id,omitempty"`
|
||||
Source string `gorm:"size:32;not null" json:"source"` // qbittorrent / transmission / aria2 / http
|
||||
URL string `gorm:"size:2048;not null" json:"-"`
|
||||
Title string `gorm:"size:512" json:"title,omitempty"`
|
||||
PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
|
||||
BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
|
||||
Overview string `gorm:"type:text" json:"overview,omitempty"`
|
||||
SavePath string `gorm:"size:1024" json:"save_path"`
|
||||
MediaType string `gorm:"size:16" json:"media_type,omitempty"`
|
||||
MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
|
||||
// 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。
|
||||
OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
|
||||
OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
|
||||
Year int `json:"year,omitempty"`
|
||||
Rating float32 `json:"rating,omitempty"`
|
||||
Genres string `gorm:"size:255" json:"genres,omitempty"` // comma separated
|
||||
Status string `gorm:"size:32;default:queued" json:"status"`
|
||||
Progress float32 `json:"progress"`
|
||||
|
||||
// AllowExistingLibrary is true for subscription wash/upgrade tasks that are
|
||||
// allowed to replace an existing library item after download completion.
|
||||
AllowExistingLibrary bool `gorm:"default:false" json:"allow_existing_library,omitempty"`
|
||||
}
|
||||
|
||||
// Subscription 是自动化规则,轮询 RSS 源并将匹配种子排队到配置的下载客户端。
|
||||
type Subscription struct {
|
||||
Base
|
||||
UserID string `gorm:"index;size:36" json:"user_id"`
|
||||
IdentityKey string `gorm:"size:64" json:"-"`
|
||||
Name string `gorm:"size:128;not null" json:"name"`
|
||||
FeedURL string `gorm:"size:2048;not null" json:"feed_url"`
|
||||
Filter string `gorm:"size:512" json:"filter"`
|
||||
MediaType string `gorm:"size:16" json:"media_type,omitempty"`
|
||||
MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
|
||||
SavePath string `gorm:"size:1024" json:"save_path,omitempty"`
|
||||
SearchMode string `gorm:"size:16;default:keyword" json:"search_mode,omitempty"` // keyword / imdb
|
||||
IMDBID string `gorm:"size:32" json:"imdb_id,omitempty"`
|
||||
Source string `gorm:"size:32" json:"source,omitempty"`
|
||||
PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
|
||||
BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
|
||||
Overview string `gorm:"type:text" json:"overview,omitempty"`
|
||||
// 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。
|
||||
OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
|
||||
OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
|
||||
Year int `json:"year,omitempty"`
|
||||
Rating float32 `json:"rating,omitempty"`
|
||||
Genres string `gorm:"size:255" json:"genres,omitempty"` // comma separated
|
||||
Resolution string `gorm:"size:32" json:"resolution,omitempty"` // 2160p / 1080p / 720p / best
|
||||
Quality string `gorm:"size:64" json:"quality,omitempty"` // remux / bluray / web-dl / hdtv
|
||||
Effects string `gorm:"size:128" json:"effects,omitempty"` // hdr,dolby-vision,atmos
|
||||
ReleaseGroups string `gorm:"size:255" json:"release_groups,omitempty"` // comma separated
|
||||
ExcludeWords string `gorm:"size:255" json:"exclude_words,omitempty"` // comma separated
|
||||
MinSeeders int `gorm:"default:0" json:"min_seeders,omitempty"`
|
||||
MaxSeeders int `gorm:"default:0" json:"max_seeders,omitempty"`
|
||||
MinSizeGB float64 `gorm:"default:0" json:"min_size_gb,omitempty"`
|
||||
MaxSizeGB float64 `gorm:"default:0" json:"max_size_gb,omitempty"`
|
||||
FreeOnly bool `gorm:"default:false" json:"free_only,omitempty"`
|
||||
WashEnabled bool `gorm:"default:false" json:"wash_enabled"`
|
||||
WashPriority string `gorm:"size:32" json:"wash_priority,omitempty"` // balanced / resolution / quality / effects / seeders
|
||||
TotalEpisodes int `gorm:"default:0" json:"total_episodes,omitempty"`
|
||||
Priority int `gorm:"default:50" json:"priority,omitempty"` // lower is earlier when schedulers sort later
|
||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||
LastRunAt *time.Time `json:"last_run_at,omitempty"`
|
||||
ArchivedAt *time.Time `gorm:"index" json:"archived_at,omitempty"`
|
||||
ArchiveReason string `gorm:"size:255" json:"archive_reason,omitempty"`
|
||||
|
||||
DownloadedEpisodes int `gorm:"-" json:"downloaded_episodes,omitempty"`
|
||||
LocalMediaCount int `gorm:"-" json:"local_media_count,omitempty"`
|
||||
MissingEpisodes []int `gorm:"-" json:"missing_episodes,omitempty"`
|
||||
InLibrary bool `gorm:"-" json:"in_library"`
|
||||
}
|
||||
@@ -41,23 +41,13 @@ func AllModels() []interface{} {
|
||||
&Favorite{},
|
||||
&Playlist{},
|
||||
&PlaylistItem{},
|
||||
&DownloadTask{},
|
||||
&Subscription{},
|
||||
&Setting{},
|
||||
&Site{},
|
||||
&AccessLog{},
|
||||
&APIConfig{},
|
||||
&UserPermission{},
|
||||
&RefreshToken{},
|
||||
&ApiConfig{},
|
||||
&DownloadClient{},
|
||||
&NotifyChannel{},
|
||||
&TelegramBinding{},
|
||||
&STRMRecord{},
|
||||
&PlayProfile{},
|
||||
&StorageConfig{},
|
||||
&AssistantSession{},
|
||||
&AssistantMessage{},
|
||||
&RegistrationCode{},
|
||||
&SignIn{},
|
||||
&UserDevice{},
|
||||
|
||||
@@ -19,7 +19,6 @@ func TestMediaReferenceFieldsAllowVirtualEmbyIDs(t *testing.T) {
|
||||
{name: "playback_history_media_id", model: &PlaybackHistory{}, fieldName: "MediaID"},
|
||||
{name: "favorite_media_id", model: &Favorite{}, fieldName: "MediaID"},
|
||||
{name: "playlist_item_media_id", model: &PlaylistItem{}, fieldName: "MediaID"},
|
||||
{name: "strm_record_media_id", model: &STRMRecord{}, fieldName: "MediaID"},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
|
||||
@@ -1,14 +0,0 @@
|
||||
// Package model 定义通知渠道配置数据模型。
|
||||
package model
|
||||
|
||||
// NotifyChannel 通知渠道配置。
|
||||
// 支持多种通知渠道:telegram / wechat / bark / webhook / email。
|
||||
// Events 字段存储 JSON array,表示该渠道订阅的事件类型。
|
||||
type NotifyChannel struct {
|
||||
Base
|
||||
Name string `gorm:"size:128;not null" json:"name"`
|
||||
Type string `gorm:"size:32;not null" json:"type"` // telegram / wechat / bark / webhook / email
|
||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||
Config string `gorm:"type:text" json:"-"` // JSON配置, AES加密
|
||||
Events string `gorm:"type:text" json:"events"` // 订阅的事件列表, JSON array
|
||||
}
|
||||
@@ -28,9 +28,8 @@ type UserPermission struct {
|
||||
CanRescrape bool `gorm:"default:false" json:"can_rescrape"`
|
||||
CanUseAI bool `gorm:"default:false" json:"can_use_ai"`
|
||||
CanCaptureFrames bool `gorm:"default:false" json:"can_capture_frames"`
|
||||
CanManageDownloads bool `gorm:"default:false" json:"can_manage_downloads"`
|
||||
CanViewDiscover bool `gorm:"default:false" json:"can_view_discover"`
|
||||
CanManageSubscriptions bool `gorm:"default:false" json:"can_manage_subscriptions"`
|
||||
CanManageDownloads bool `gorm:"default:false" json:"can_manage_downloads"`
|
||||
CanManageSubscriptions bool `gorm:"default:false" json:"can_manage_subscriptions"`
|
||||
CanManageSites bool `gorm:"default:false" json:"can_manage_sites"`
|
||||
CanUseAIAssistant bool `gorm:"default:false" json:"can_use_ai_assistant"`
|
||||
CanManageUsers bool `gorm:"default:false" json:"can_manage_users"`
|
||||
@@ -65,9 +64,8 @@ func NewDefaultPermission(userID string) *UserPermission {
|
||||
CanRescrape: false,
|
||||
CanUseAI: false,
|
||||
CanCaptureFrames: false,
|
||||
CanManageDownloads: false,
|
||||
CanViewDiscover: false,
|
||||
CanManageSubscriptions: false,
|
||||
CanManageDownloads: false,
|
||||
CanManageSubscriptions: false,
|
||||
CanManageSites: false,
|
||||
CanUseAIAssistant: false,
|
||||
CanManageUsers: false,
|
||||
@@ -90,9 +88,8 @@ func (p *UserPermission) PermissionMap() map[string]bool {
|
||||
"can_rescrape": p.CanRescrape,
|
||||
"can_use_ai": p.CanUseAI,
|
||||
"can_capture_frames": p.CanCaptureFrames,
|
||||
"can_manage_downloads": p.CanManageDownloads,
|
||||
"can_view_discover": p.CanViewDiscover,
|
||||
"can_manage_subscriptions": p.CanManageSubscriptions,
|
||||
"can_manage_downloads": p.CanManageDownloads,
|
||||
"can_manage_subscriptions": p.CanManageSubscriptions,
|
||||
"can_manage_sites": p.CanManageSites,
|
||||
"can_use_ai_assistant": p.CanUseAIAssistant,
|
||||
"can_manage_users": p.CanManageUsers,
|
||||
|
||||
@@ -1,57 +0,0 @@
|
||||
// Package model — PT 站点配置数据模型。
|
||||
package model
|
||||
|
||||
import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// Site PT 站点配置。
|
||||
type Site struct {
|
||||
Base
|
||||
Name string `gorm:"size:128;not null" json:"name"`
|
||||
Type string `gorm:"size:32;not null" json:"type"` // nexusphp / gazelle / unit3d / mteam / yemapt / discuz / custom_rss
|
||||
URL string `gorm:"size:512;not null" json:"url"`
|
||||
AuthType string `gorm:"size:32;not null" json:"auth_type"` // cookie / api_key / auth_header
|
||||
Cookie string `gorm:"type:text" json:"-"` // AES 加密
|
||||
APIKey string `gorm:"type:text" json:"-"` // AES 加密
|
||||
AuthHeader string `gorm:"type:text" json:"-"` // AES 加密
|
||||
|
||||
// ── 高级设置 ──────────────────────────────────────────────────────────
|
||||
UserAgent string `gorm:"size:500" json:"user_agent"` // 自定义 User-Agent
|
||||
RSSURL string `gorm:"size:1000" json:"rss_url"` // RSS 订阅地址
|
||||
Timeout int `gorm:"default:15" json:"timeout"` // 请求超时(秒), 0=不限制
|
||||
Priority int `gorm:"default:50" json:"priority"` // 优先级, 越小越优先
|
||||
UseProxy bool `gorm:"default:false" json:"use_proxy"` // 是否使用代理
|
||||
RateLimit bool `gorm:"default:false" json:"rate_limit"` // 是否限制访问频率
|
||||
BrowserEmulation bool `gorm:"default:false" json:"browser_emulation"` // 浏览器仿真(防爬)
|
||||
|
||||
// ── 状态与统计 ────────────────────────────────────────────────────────
|
||||
LoginStatus string `gorm:"size:20;default:unknown" json:"login_status"` // unknown / ok / fail
|
||||
UploadBytes int64 `gorm:"default:0" json:"upload_bytes"` // 上传字节统计
|
||||
DownloadBytes int64 `gorm:"default:0" json:"download_bytes"` // 下载字节统计
|
||||
|
||||
// ── 关联下载器 ────────────────────────────────────────────────────────
|
||||
Downloader string `gorm:"size:50" json:"downloader"` // qbittorrent / transmission / aria2
|
||||
|
||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||
IsDefault bool `gorm:"default:false" json:"is_default"`
|
||||
Extra string `gorm:"type:text" json:"-"` // JSON 扩展配置, AES 加密
|
||||
LastError string `gorm:"size:1024" json:"last_error"`
|
||||
LastCheckAt *time.Time `json:"last_check_at"`
|
||||
}
|
||||
|
||||
// SiteType 返回支持的站点类型列表。
|
||||
func SiteTypes() []string {
|
||||
return []string{"nexusphp", "gazelle", "unit3d", "mteam", "yemapt", "discuz", "custom_rss"}
|
||||
}
|
||||
|
||||
// AuthTypes 返回支持的认证方式列表。
|
||||
func AuthTypes() []string {
|
||||
return []string{"cookie", "api_key", "auth_header"}
|
||||
}
|
||||
|
||||
// DefaultUserAgent 是默认浏览器 User-Agent(用于 HTTP 请求头)。
|
||||
const DefaultUserAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/122.0.0.0 Safari/537.36"
|
||||
|
||||
// SiteTypes 返回支持的站点类型列表(用于前端下拉)。
|
||||
// 注意:此函数供 API 返回类型列表使用,不在此处添加新类型。
|
||||
@@ -14,23 +14,3 @@ type StorageConfig struct {
|
||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||
LastError string `gorm:"size:512" json:"last_error,omitempty"`
|
||||
}
|
||||
|
||||
// AssistantSession groups a multi-turn chat with the AI assistant.
|
||||
type AssistantSession struct {
|
||||
Base
|
||||
UserID string `gorm:"index;size:36;not null" json:"user_id"`
|
||||
Title string `gorm:"size:255" json:"title,omitempty"`
|
||||
}
|
||||
|
||||
// AssistantMessage is one entry in an AssistantSession transcript.
|
||||
//
|
||||
// Role is "user" | "assistant" | "system". The optional OperationID
|
||||
// links a message to an action the assistant proposed (so the UI can
|
||||
// offer Undo).
|
||||
type AssistantMessage struct {
|
||||
Base
|
||||
SessionID string `gorm:"index;size:36;not null" json:"session_id"`
|
||||
Role string `gorm:"size:16;not null" json:"role"`
|
||||
Content string `gorm:"type:text;not null" json:"content"`
|
||||
OperationID string `gorm:"size:36" json:"operation_id,omitempty"`
|
||||
}
|
||||
|
||||
@@ -1,36 +0,0 @@
|
||||
// Package model — STRM 文件记录数据模型。
|
||||
package model
|
||||
|
||||
// STRMRecord STRM 文件记录。
|
||||
// 外部存储以"文件"形式入库,URL 指向实际资源。
|
||||
type STRMRecord struct {
|
||||
Base
|
||||
Title string `gorm:"size:512;not null;index" json:"title"`
|
||||
URL string `gorm:"size:2048;not null" json:"url"` // STRM 文件指向的 URL
|
||||
FilePath string `gorm:"size:1024;not null" json:"file_path"` // 本地 STRM 文件路径
|
||||
Protocol string `gorm:"size:32;not null" json:"protocol"` // webdav / alist / s3 / http / https
|
||||
FileSize int64 `json:"file_size"`
|
||||
MediaID string `gorm:"size:128;index" json:"media_id"` // 关联媒体 ID
|
||||
MediaType string `gorm:"size:16" json:"media_type"` // movie / series
|
||||
SeasonNum int `json:"season_num"`
|
||||
EpisodeNum int `json:"episode_num"`
|
||||
}
|
||||
|
||||
// AllowedSTRMProtocols 协议白名单。
|
||||
var AllowedSTRMProtocols = []string{
|
||||
"webdav", "davs",
|
||||
"alist", "alists",
|
||||
"openlist", "openlists",
|
||||
"s3",
|
||||
"http", "https",
|
||||
}
|
||||
|
||||
// IsAllowedProtocol 检查协议是否在白名单中。
|
||||
func IsAllowedProtocol(protocol string) bool {
|
||||
for _, p := range AllowedSTRMProtocols {
|
||||
if p == protocol {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -1,90 +0,0 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type subscriptionIdentityPayload struct {
|
||||
Name string `json:"name"`
|
||||
FeedURL string `json:"feed_url"`
|
||||
Filter string `json:"filter"`
|
||||
MediaType string `json:"media_type"`
|
||||
MediaCategory string `json:"media_category"`
|
||||
SavePath string `json:"save_path"`
|
||||
SearchMode string `json:"search_mode"`
|
||||
IMDBID string `json:"imdb_id"`
|
||||
Resolution string `json:"resolution"`
|
||||
Quality string `json:"quality"`
|
||||
Effects string `json:"effects"`
|
||||
ReleaseGroups string `json:"release_groups"`
|
||||
ExcludeWords string `json:"exclude_words"`
|
||||
MinSeeders int `json:"min_seeders"`
|
||||
MaxSeeders int `json:"max_seeders"`
|
||||
MinSizeGB float64 `json:"min_size_gb"`
|
||||
MaxSizeGB float64 `json:"max_size_gb"`
|
||||
FreeOnly bool `json:"free_only"`
|
||||
WashEnabled bool `json:"wash_enabled"`
|
||||
WashPriority string `json:"wash_priority"`
|
||||
Priority int `json:"priority"`
|
||||
}
|
||||
|
||||
// SubscriptionIdentityKey identifies one functional subscription rule. It
|
||||
// intentionally excludes display metadata and runtime/archive state because
|
||||
// those fields are backfilled or changed automatically after creation.
|
||||
func SubscriptionIdentityKey(sub *Subscription) string {
|
||||
if sub == nil {
|
||||
return ""
|
||||
}
|
||||
payload := subscriptionIdentityPayload{
|
||||
Name: subscriptionIdentityFold(sub.Name),
|
||||
FeedURL: strings.TrimSpace(sub.FeedURL),
|
||||
Filter: strings.TrimSpace(sub.Filter),
|
||||
MediaType: subscriptionIdentityFold(sub.MediaType),
|
||||
MediaCategory: subscriptionIdentityFold(sub.MediaCategory),
|
||||
SavePath: strings.TrimSpace(sub.SavePath),
|
||||
SearchMode: subscriptionIdentityFold(sub.SearchMode),
|
||||
IMDBID: subscriptionIdentityFold(sub.IMDBID),
|
||||
Resolution: subscriptionIdentityFold(sub.Resolution),
|
||||
Quality: subscriptionIdentityFold(sub.Quality),
|
||||
Effects: subscriptionIdentityList(sub.Effects),
|
||||
ReleaseGroups: subscriptionIdentityList(sub.ReleaseGroups),
|
||||
ExcludeWords: subscriptionIdentityList(sub.ExcludeWords),
|
||||
MinSeeders: sub.MinSeeders,
|
||||
MaxSeeders: sub.MaxSeeders,
|
||||
MinSizeGB: sub.MinSizeGB,
|
||||
MaxSizeGB: sub.MaxSizeGB,
|
||||
FreeOnly: sub.FreeOnly,
|
||||
WashEnabled: sub.WashEnabled,
|
||||
WashPriority: subscriptionIdentityFold(sub.WashPriority),
|
||||
Priority: sub.Priority,
|
||||
}
|
||||
raw, _ := json.Marshal(payload)
|
||||
sum := sha256.Sum256(raw)
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func RefreshSubscriptionIdentity(sub *Subscription) string {
|
||||
key := SubscriptionIdentityKey(sub)
|
||||
if sub != nil {
|
||||
sub.IdentityKey = key
|
||||
}
|
||||
return key
|
||||
}
|
||||
|
||||
func subscriptionIdentityFold(value string) string {
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
}
|
||||
|
||||
func subscriptionIdentityList(value string) string {
|
||||
parts := strings.Split(value, ",")
|
||||
out := make([]string, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
if part = subscriptionIdentityFold(part); part != "" {
|
||||
out = append(out, part)
|
||||
}
|
||||
}
|
||||
return strings.Join(out, ",")
|
||||
}
|
||||
@@ -1,45 +0,0 @@
|
||||
package model
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestSubscriptionIdentityKeyNormalizesEquivalentRules(t *testing.T) {
|
||||
base := Subscription{
|
||||
Name: " Example Show ",
|
||||
FeedURL: "site-search://search?keyword=Example",
|
||||
Filter: " Example.*Show ",
|
||||
MediaType: "TV",
|
||||
MediaCategory: " 欧美剧 ",
|
||||
SearchMode: "KEYWORD",
|
||||
Resolution: "1080P",
|
||||
Effects: " HDR, Atmos ",
|
||||
ExcludeWords: " CAM, TS ",
|
||||
WashPriority: "Balanced",
|
||||
Priority: 50,
|
||||
}
|
||||
equivalent := base
|
||||
equivalent.Name = "example show"
|
||||
equivalent.MediaType = "tv"
|
||||
equivalent.MediaCategory = "欧美剧"
|
||||
equivalent.SearchMode = "keyword"
|
||||
equivalent.Resolution = "1080p"
|
||||
equivalent.Effects = "hdr,atmos"
|
||||
equivalent.ExcludeWords = "cam,ts"
|
||||
equivalent.WashPriority = "balanced"
|
||||
if got, want := SubscriptionIdentityKey(&equivalent), SubscriptionIdentityKey(&base); got != want {
|
||||
t.Fatalf("equivalent keys differ: %s != %s", got, want)
|
||||
}
|
||||
|
||||
differentRule := base
|
||||
differentRule.Resolution = "2160p"
|
||||
if SubscriptionIdentityKey(&differentRule) == SubscriptionIdentityKey(&base) {
|
||||
t.Fatal("different resolution should produce a different identity")
|
||||
}
|
||||
|
||||
displayOnly := base
|
||||
displayOnly.PosterURL = "https://example.test/poster.jpg"
|
||||
displayOnly.Overview = "new overview"
|
||||
displayOnly.TotalEpisodes = 24
|
||||
if SubscriptionIdentityKey(&displayOnly) != SubscriptionIdentityKey(&base) {
|
||||
t.Fatal("display/runtime metadata should not change identity")
|
||||
}
|
||||
}
|
||||
@@ -1,12 +0,0 @@
|
||||
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"`
|
||||
}
|
||||
@@ -1,64 +0,0 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
)
|
||||
|
||||
// AssistantRepository persists model.AssistantSession + AssistantMessage records.
|
||||
type AssistantRepository struct{ db *gorm.DB }
|
||||
|
||||
// ─── Session ────────────────────────────────────────────────────────────
|
||||
|
||||
// CreateSession inserts a new chat session.
|
||||
func (r *AssistantRepository) CreateSession(ctx context.Context, s *model.AssistantSession) error {
|
||||
return r.db.WithContext(ctx).Create(s).Error
|
||||
}
|
||||
|
||||
// FindSession returns a session by ID, or (nil, nil).
|
||||
func (r *AssistantRepository) FindSession(ctx context.Context, id string) (*model.AssistantSession, error) {
|
||||
var s model.AssistantSession
|
||||
err := r.db.WithContext(ctx).Where("id = ?", id).First(&s).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &s, nil
|
||||
}
|
||||
|
||||
// ListSessions returns sessions for a user, or all when userID is empty.
|
||||
func (r *AssistantRepository) ListSessions(ctx context.Context, userID string) ([]model.AssistantSession, error) {
|
||||
q := r.db.WithContext(ctx).Model(&model.AssistantSession{})
|
||||
if userID != "" {
|
||||
q = q.Where("user_id = ?", userID)
|
||||
}
|
||||
var rows []model.AssistantSession
|
||||
err := q.Order("created_at desc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// DeleteSession soft-deletes a session (cascade handled by GORM hooks if set).
|
||||
func (r *AssistantRepository) DeleteSession(ctx context.Context, id string) error {
|
||||
return r.db.WithContext(ctx).Delete(&model.AssistantSession{}, "id = ?", id).Error
|
||||
}
|
||||
|
||||
// ─── Message ────────────────────────────────────────────────────────────
|
||||
|
||||
// AppendMessage inserts a new message into a session.
|
||||
func (r *AssistantRepository) AppendMessage(ctx context.Context, m *model.AssistantMessage) error {
|
||||
return r.db.WithContext(ctx).Create(m).Error
|
||||
}
|
||||
|
||||
// ListMessages returns all messages for a session in chronological order.
|
||||
func (r *AssistantRepository) ListMessages(ctx context.Context, sessionID string) ([]model.AssistantMessage, error) {
|
||||
var rows []model.AssistantMessage
|
||||
err := r.db.WithContext(ctx).Where("session_id = ?", sessionID).
|
||||
Order("created_at asc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
@@ -1,105 +0,0 @@
|
||||
// Package repository 实现下载客户端配置的数据访问层。
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
)
|
||||
|
||||
// DownloadClientRepository persists model.DownloadClient records.
|
||||
type DownloadClientRepository struct{ db *gorm.DB }
|
||||
|
||||
// Create inserts a new download client.
|
||||
func (r *DownloadClientRepository) Create(ctx context.Context, c *model.DownloadClient) error {
|
||||
return r.db.WithContext(ctx).Create(c).Error
|
||||
}
|
||||
|
||||
// FindByID returns the download client by ID, or (nil, nil) when absent.
|
||||
func (r *DownloadClientRepository) FindByID(ctx context.Context, id string) (*model.DownloadClient, error) {
|
||||
var c model.DownloadClient
|
||||
err := r.db.WithContext(ctx).Where("id = ?", id).First(&c).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &c, nil
|
||||
}
|
||||
|
||||
// FindDefault returns the default download client, or (nil, nil).
|
||||
func (r *DownloadClientRepository) FindDefault(ctx context.Context) (*model.DownloadClient, error) {
|
||||
var c model.DownloadClient
|
||||
err := r.db.WithContext(ctx).
|
||||
Where("is_default = ? AND enabled = ?", true, true).
|
||||
Order("created_at asc").
|
||||
First(&c).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &c, nil
|
||||
}
|
||||
|
||||
// List returns all download clients ordered by creation time.
|
||||
func (r *DownloadClientRepository) List(ctx context.Context) ([]model.DownloadClient, error) {
|
||||
var rows []model.DownloadClient
|
||||
err := r.db.WithContext(ctx).Order("created_at asc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// ListEnabled returns all enabled download clients.
|
||||
func (r *DownloadClientRepository) ListEnabled(ctx context.Context) ([]model.DownloadClient, error) {
|
||||
var rows []model.DownloadClient
|
||||
err := r.db.WithContext(ctx).Where("enabled = ?", true).Order("created_at asc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// HasAnyIncludingDeleted reports whether the operator has ever configured a
|
||||
// download client. This distinguishes legacy-only installations from systems
|
||||
// where deleting/disabling all clients is an intentional "stop downloads"
|
||||
// action, even though rows are soft-deleted.
|
||||
func (r *DownloadClientRepository) HasAnyIncludingDeleted(ctx context.Context) (bool, error) {
|
||||
var n int64
|
||||
err := r.db.WithContext(ctx).Unscoped().Model(&model.DownloadClient{}).Count(&n).Error
|
||||
return n > 0, err
|
||||
}
|
||||
|
||||
// Update persists changes to a download client.
|
||||
func (r *DownloadClientRepository) Update(ctx context.Context, c *model.DownloadClient) error {
|
||||
return r.db.WithContext(ctx).Save(c).Error
|
||||
}
|
||||
|
||||
// Delete removes a download client (soft-delete).
|
||||
func (r *DownloadClientRepository) Delete(ctx context.Context, id string) error {
|
||||
return r.db.WithContext(ctx).Delete(&model.DownloadClient{}, "id = ?", id).Error
|
||||
}
|
||||
|
||||
// ClearDefault unsets the default flag for all clients.
|
||||
func (r *DownloadClientRepository) ClearDefault(ctx context.Context) error {
|
||||
return r.db.WithContext(ctx).Model(&model.DownloadClient{}).
|
||||
Where("is_default = ?", true).Update("is_default", false).Error
|
||||
}
|
||||
|
||||
// SetDefault sets a specific client as default and clears others.
|
||||
func (r *DownloadClientRepository) SetDefault(ctx context.Context, id string) error {
|
||||
now := time.Now()
|
||||
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&model.DownloadClient{}).
|
||||
Where("is_default = ?", true).Update("is_default", false).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&model.DownloadClient{}).
|
||||
Where("id = ?", id).Updates(map[string]any{
|
||||
"is_default": true,
|
||||
"updated_at": now,
|
||||
}).Error
|
||||
})
|
||||
}
|
||||
@@ -1,24 +0,0 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
)
|
||||
|
||||
// DownloadRepository persists model.DownloadTask records.
|
||||
type DownloadRepository struct{ db *gorm.DB }
|
||||
|
||||
// Create inserts a new download task.
|
||||
func (r *DownloadRepository) Create(ctx context.Context, t *model.DownloadTask) error {
|
||||
return r.db.WithContext(ctx).Create(t).Error
|
||||
}
|
||||
|
||||
// List returns all download tasks (admin view).
|
||||
func (r *DownloadRepository) List(ctx context.Context) ([]model.DownloadTask, error) {
|
||||
var rows []model.DownloadTask
|
||||
err := r.db.WithContext(ctx).Order("created_at desc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
@@ -1,74 +0,0 @@
|
||||
// Package repository 实现通知渠道配置的数据访问层。
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||
)
|
||||
|
||||
// NotifyChannelRepository persists model.NotifyChannel records.
|
||||
type NotifyChannelRepository struct{ db *gorm.DB }
|
||||
|
||||
// Create inserts a new notification channel.
|
||||
func (r *NotifyChannelRepository) Create(ctx context.Context, c *model.NotifyChannel) error {
|
||||
return r.db.WithContext(ctx).Create(c).Error
|
||||
}
|
||||
|
||||
// FindByID returns the notification channel by ID, or (nil, nil) when absent.
|
||||
func (r *NotifyChannelRepository) FindByID(ctx context.Context, id string) (*model.NotifyChannel, error) {
|
||||
var c model.NotifyChannel
|
||||
err := r.db.WithContext(ctx).Where("id = ?", id).First(&c).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &c, nil
|
||||
}
|
||||
|
||||
// List returns all notification channels ordered by creation time.
|
||||
func (r *NotifyChannelRepository) List(ctx context.Context) ([]model.NotifyChannel, error) {
|
||||
var rows []model.NotifyChannel
|
||||
err := r.db.WithContext(ctx).Order("created_at asc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// ListEnabled returns all enabled notification channels.
|
||||
func (r *NotifyChannelRepository) ListEnabled(ctx context.Context) ([]model.NotifyChannel, error) {
|
||||
var rows []model.NotifyChannel
|
||||
err := r.db.WithContext(ctx).Where("enabled = ?", true).Order("created_at asc").Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// ListByEvent returns all enabled channels that subscribe to the given event type.
|
||||
func (r *NotifyChannelRepository) ListByEvent(ctx context.Context, eventType string) ([]model.NotifyChannel, error) {
|
||||
var rows []model.NotifyChannel
|
||||
// Events is a JSON array stored as text; use LIKE for simple matching.
|
||||
// This works for exact event type matches within the JSON array.
|
||||
err := r.db.WithContext(ctx).
|
||||
Where("enabled = ? AND events LIKE ?", true, "%\""+eventType+"\"%").
|
||||
Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// ListByType returns all channels of a given type (telegram/wechat/bark/webhook/email).
|
||||
func (r *NotifyChannelRepository) ListByType(ctx context.Context, channelType string) ([]model.NotifyChannel, error) {
|
||||
var rows []model.NotifyChannel
|
||||
err := r.db.WithContext(ctx).Where("type = ?", channelType).Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// Update persists changes to a notification channel.
|
||||
func (r *NotifyChannelRepository) Update(ctx context.Context, c *model.NotifyChannel) error {
|
||||
return r.db.WithContext(ctx).Save(c).Error
|
||||
}
|
||||
|
||||
// Delete removes a notification channel (soft-delete).
|
||||
func (r *NotifyChannelRepository) Delete(ctx context.Context, id string) error {
|
||||
return r.db.WithContext(ctx).Delete(&model.NotifyChannel{}, "id = ?", id).Error
|
||||
}
|
||||
@@ -17,20 +17,12 @@ type Container struct {
|
||||
History *HistoryRepository
|
||||
Favorite *FavoriteRepository
|
||||
Playlist *PlaylistRepository
|
||||
Download *DownloadRepository
|
||||
Subscription *SubscriptionRepository
|
||||
Setting *SettingRepository
|
||||
Log *AccessLogRepository
|
||||
Permission *PermissionRepository
|
||||
RefreshToken *RefreshTokenRepository
|
||||
ApiConfig *ApiConfigRepository
|
||||
DownloadClient *DownloadClientRepository
|
||||
NotifyChannel *NotifyChannelRepository
|
||||
Site *SiteRepository
|
||||
STRM *STRMRepository
|
||||
PlayProfile *PlayProfileRepository
|
||||
StorageConfig *StorageConfigRepository
|
||||
Assistant *AssistantRepository
|
||||
RegCode *RegistrationCodeRepository
|
||||
SignIn *SignInRepository
|
||||
UserDevice *UserDeviceRepository
|
||||
@@ -47,20 +39,12 @@ func New(db *gorm.DB) *Container {
|
||||
History: &HistoryRepository{db: db},
|
||||
Favorite: &FavoriteRepository{db: db},
|
||||
Playlist: &PlaylistRepository{db: db},
|
||||
Download: &DownloadRepository{db: db},
|
||||
Subscription: &SubscriptionRepository{db: db},
|
||||
Setting: &SettingRepository{db: db},
|
||||
Log: &AccessLogRepository{db: db},
|
||||
Permission: &PermissionRepository{db: db},
|
||||
RefreshToken: &RefreshTokenRepository{db: db},
|
||||
ApiConfig: &ApiConfigRepository{db: db},
|
||||
DownloadClient: &DownloadClientRepository{db: db},
|
||||
NotifyChannel: &NotifyChannelRepository{db: db},
|
||||
Site: &SiteRepository{db: db},
|
||||
STRM: &STRMRepository{db: db},
|
||||
PlayProfile: &PlayProfileRepository{db: db},
|
||||
StorageConfig: &StorageConfigRepository{db: db},
|
||||
Assistant: &AssistantRepository{db: db},
|
||||
RegCode: &RegistrationCodeRepository{db: db},
|
||||
SignIn: &SignInRepository{db: db},
|
||||
UserDevice: &UserDeviceRepository{db: db},
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user