初始化

初始化项目
This commit is contained in:
truewhile
2026-08-23 22:12:32 +08:00
parent 0bcb1fec87
commit 71bf60c69c
631 changed files with 2121 additions and 76050 deletions
+2 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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")
}
-101
View File
@@ -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 {
-43
View File
@@ -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
+1 -3
View File
@@ -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()})
}
-3
View File
@@ -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)
}
}
-79
View File
@@ -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()))
}
}
-147
View File
@@ -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})
}
}
-1
View File
@@ -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) {
-160
View File
@@ -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()})
}
}
-149
View File
@@ -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
}
-284
View File
@@ -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()
}
-55
View File
@@ -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())
}
}
-49
View File
@@ -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)
}
}
-157
View File
@@ -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)
}
}
-48
View File
@@ -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})
}
}
-338
View File
@@ -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
}
}
-183
View File
@@ -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)
}
}
}
-274
View File
@@ -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)
}
-108
View File
@@ -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)
}
}
-258
View File
@@ -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)
}
}
-112
View File
@@ -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)
}
-46
View File
@@ -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})
}
}
-5
View File
@@ -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
-59
View File
@@ -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)
}
}
-39
View File
@@ -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 {
-167
View File
@@ -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))
}
}
-229
View File
@@ -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,
}
}
-172
View File
@@ -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
}
-211
View File
@@ -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
}
-445
View File
@@ -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))
}
+9 -23
View File
@@ -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
+2 -41
View File
@@ -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)
-115
View File
@@ -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
}
-27
View File
@@ -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"})
}
}
-77
View File
@@ -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"})
}
}
-196
View File
@@ -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)
}
-58
View File
@@ -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
}
-65
View File
@@ -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
}
-84
View File
@@ -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"})
}
}
-65
View File
@@ -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))
-8
View File
@@ -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)
-11
View File
@@ -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))
}
+10 -16
View File
@@ -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)
}
-42
View File
@@ -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
}
-42
View File
@@ -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)
}
-21
View File
@@ -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})
}
}
-170
View File
@@ -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()})
}
-164
View File
@@ -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)})
}
}
-62
View File
@@ -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",
})
}
}
-154
View File
@@ -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)
}
-180
View File
@@ -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)
}
}
-115
View File
@@ -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])
}
}
-145
View File
@@ -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})
}
}
-101
View File
@@ -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)
}
}
-36
View File
@@ -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"`
-437
View File
@@ -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
}
-47
View File
@@ -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)
}
}
-22
View File
@@ -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)
}
}
-210
View File
@@ -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})
}
}
-54
View File
@@ -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...)
}
-219
View File
@@ -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])
}
}
+10 -30
View File
@@ -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
-8
View File
@@ -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()})
}
}
+1 -1
View File
@@ -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})
}
}
-40
View File
@@ -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,
})
}
}
-104
View File
@@ -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})
}
}
-128
View File
@@ -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
}
+4 -3
View File
@@ -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",
-48
View File
@@ -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
}
-50
View File
@@ -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")
}
-96
View File
@@ -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
}
}
+6 -6
View File
@@ -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)
-16
View File
@@ -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加密
}
-81
View File
@@ -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"`
}
-10
View File
@@ -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{},
-1
View File
@@ -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 {
-14
View File
@@ -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
}
+6 -9
View File
@@ -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,
-57
View File
@@ -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 返回类型列表使用,不在此处添加新类型。
-20
View File
@@ -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"`
}
-36
View File
@@ -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
}
-90
View File
@@ -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")
}
}
-12
View File
@@ -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"`
}
-64
View File
@@ -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
}
-105
View File
@@ -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
}
-16
View File
@@ -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