mirror of
https://github.com/truewhile/MeBox.git
synced 2026-10-01 20:16:36 +08:00
优化
This commit is contained in:
@@ -15,6 +15,8 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha1"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
@@ -69,6 +71,7 @@ type EmbyRemoteService struct {
|
||||
repo *repository.Container
|
||||
crypto *CryptoService
|
||||
http *http.Client
|
||||
cache *RuntimeCacheService
|
||||
}
|
||||
|
||||
// NewEmbyRemoteService 构造远程 Emby 聚合服务。
|
||||
@@ -85,6 +88,32 @@ func NewEmbyRemoteService(cfg *config.Config, log *zap.Logger, repo *repository.
|
||||
}
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) SetRuntimeCache(cache *RuntimeCacheService) *EmbyRemoteService {
|
||||
if r != nil {
|
||||
r.cache = cache
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) remoteMediaCacheTTL() time.Duration {
|
||||
seconds := 15
|
||||
if r != nil && r.cfg != nil && r.cfg.Cache.MediaTTLSeconds > 0 {
|
||||
seconds = r.cfg.Cache.MediaTTLSeconds
|
||||
}
|
||||
return time.Duration(seconds) * time.Second
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) remoteCacheKey(parts ...string) string {
|
||||
sum := sha1.Sum([]byte(strings.Join(parts, "|")))
|
||||
return "media:embyremote:" + hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func (r *EmbyRemoteService) invalidateRemoteMediaCache(ctx context.Context) {
|
||||
if r != nil && r.cache != nil {
|
||||
r.cache.DeletePrefix(ctx, "media:embyremote:")
|
||||
}
|
||||
}
|
||||
|
||||
// ListAccounts 返回全部启用的远程 Emby 挂载账号。
|
||||
func (r *EmbyRemoteService) ListAccounts(ctx context.Context) ([]model.StrmAccount, error) {
|
||||
accounts, err := r.repo.StrmAccount.List(ctx)
|
||||
@@ -143,6 +172,7 @@ func (r *EmbyRemoteService) CreateMount(ctx context.Context, m *model.EmbyMount)
|
||||
if err := r.repo.EmbyMount.Create(ctx, m); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
return m, nil
|
||||
}
|
||||
|
||||
@@ -172,6 +202,7 @@ func (r *EmbyRemoteService) CreateMounts(ctx context.Context, mounts []*model.Em
|
||||
if err := r.repo.EmbyMount.CreateInBatches(ctx, fresh, 50); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
return len(fresh), nil
|
||||
}
|
||||
|
||||
@@ -187,12 +218,17 @@ func (r *EmbyRemoteService) UpdateMount(ctx context.Context, id string, m *model
|
||||
if err := r.repo.EmbyMount.Update(ctx, existing); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
return existing, nil
|
||||
}
|
||||
|
||||
// DeleteMount 删除挂载。
|
||||
func (r *EmbyRemoteService) DeleteMount(ctx context.Context, id string) error {
|
||||
return r.repo.EmbyMount.Delete(ctx, id)
|
||||
err := r.repo.EmbyMount.Delete(ctx, id)
|
||||
if err == nil {
|
||||
r.invalidateRemoteMediaCache(ctx)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// FullMountAccount 把账号的全部远程媒体库(View)挂载进来(幂等,已存在跳过)。
|
||||
@@ -407,12 +443,12 @@ func (r *EmbyRemoteService) doGet(ctx context.Context, acct *model.StrmAccount,
|
||||
// token 失效:清空后重认证重试一次。
|
||||
cfg.Token = ""
|
||||
if acct != nil {
|
||||
raw := map[string]string{}
|
||||
_ = json.Unmarshal([]byte(acct.Config), &raw)
|
||||
delete(raw, "api_key")
|
||||
enc, _ := json.Marshal(raw)
|
||||
acct.Config = string(enc)
|
||||
_ = r.repo.StrmAccount.Update(ctx, acct)
|
||||
raw := map[string]string{}
|
||||
_ = json.Unmarshal([]byte(acct.Config), &raw)
|
||||
delete(raw, "api_key")
|
||||
enc, _ := json.Marshal(raw)
|
||||
acct.Config = string(enc)
|
||||
_ = r.repo.StrmAccount.Update(ctx, acct)
|
||||
}
|
||||
continue
|
||||
}
|
||||
@@ -457,6 +493,11 @@ func (r *EmbyRemoteService) RemoteViews(ctx context.Context, acct *model.StrmAcc
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cacheKey := r.remoteCacheKey("views", acct.ID, r.remoteUserID(cfg))
|
||||
var cached []map[string]any
|
||||
if r.cache != nil && r.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
return cached, nil
|
||||
}
|
||||
q := url.Values{"api_key": {cfg.Token}}
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
@@ -464,6 +505,12 @@ func (r *EmbyRemoteService) RemoteViews(ctx context.Context, acct *model.StrmAcc
|
||||
if err := r.doGet(ctx, acct, cfg, "/Users/"+url.PathEscape(r.remoteUserID(cfg))+"/Views", q, &body); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if body.Items == nil {
|
||||
body.Items = []map[string]any{}
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, body.Items, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return body.Items, nil
|
||||
}
|
||||
|
||||
@@ -859,4 +906,4 @@ func (r *EmbyRemoteService) doMutate(ctx context.Context, acct *model.StrmAccoun
|
||||
return fmt.Errorf("远程 Emby 状态同步失败(%d): %s", resp.StatusCode, strings.TrimSpace(string(data)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,8 @@ package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
)
|
||||
|
||||
// rewriteSubtitleDeliveryURLs 只应改动字幕轨道的 DeliveryUrl,其余媒体流不动。
|
||||
@@ -47,8 +49,51 @@ func TestRewriteSubtitleDeliveryURLsFallsBackIndexOne(t *testing.T) {
|
||||
}
|
||||
rewriteSubtitleDeliveryURLs(src, "/Videos/embyremote~acct-1~item-1", &EmbyRemoteConfig{})
|
||||
streams := src["MediaStreams"].([]any)
|
||||
want := "/Videos/embyremote~acct-1~item-1/Subtitles/1/Stream"
|
||||
if got := streams[0].(map[string]any)["DeliveryUrl"]; got != want {
|
||||
t.Fatalf("subtitle DeliveryUrl = %v, want %v", got, want)
|
||||
want := "/Videos/embyremote~acct-1~item-1/Subtitles/1/Stream"
|
||||
if got := streams[0].(map[string]any)["DeliveryUrl"]; got != want {
|
||||
t.Fatalf("subtitle DeliveryUrl = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMapRemoteItemToMediaExtractsCodecsAndContainer(t *testing.T) {
|
||||
r := &EmbyRemoteService{}
|
||||
item := map[string]any{
|
||||
"Id": "item-100",
|
||||
"Name": "Test Movie",
|
||||
"Container": "mkv",
|
||||
"MediaStreams": []any{
|
||||
map[string]any{
|
||||
"Type": "Video",
|
||||
"Codec": "h264",
|
||||
"Width": 1920,
|
||||
"Height": 1080,
|
||||
},
|
||||
map[string]any{
|
||||
"Type": "Audio",
|
||||
"Codec": "aac",
|
||||
},
|
||||
},
|
||||
"MediaSources": []any{
|
||||
map[string]any{
|
||||
"Container": "mkv",
|
||||
"Size": int64(104857600),
|
||||
},
|
||||
},
|
||||
}
|
||||
media := r.MapRemoteItemToMedia(t.Context(), nil, &model.StrmAccount{Base: model.Base{ID: "acct-1"}}, &EmbyRemoteConfig{}, item)
|
||||
if media.Container != "mkv" {
|
||||
t.Fatalf("media.Container = %v, want mkv", media.Container)
|
||||
}
|
||||
if media.VideoCodec != "h264" {
|
||||
t.Fatalf("media.VideoCodec = %v, want h264", media.VideoCodec)
|
||||
}
|
||||
if media.AudioCodec != "aac" {
|
||||
t.Fatalf("media.AudioCodec = %v, want aac", media.AudioCodec)
|
||||
}
|
||||
if media.Width != 1920 || media.Height != 1080 {
|
||||
t.Fatalf("resolution = %dx%d, want 1920x1080", media.Width, media.Height)
|
||||
}
|
||||
if media.SizeBytes != 104857600 {
|
||||
t.Fatalf("size = %d, want 104857600", media.SizeBytes)
|
||||
}
|
||||
}
|
||||
@@ -191,6 +191,59 @@ func (r *EmbyRemoteService) MapRemoteItemToMedia(ctx context.Context, mount *mod
|
||||
media.DoubanID = v
|
||||
}
|
||||
}
|
||||
media.Container = remoteItemString(item, "Container")
|
||||
media.Width = remoteItemInt(item, "Width")
|
||||
media.Height = remoteItemInt(item, "Height")
|
||||
|
||||
extractStreamInfo := func(streams []any) {
|
||||
for _, s := range streams {
|
||||
sm, ok := s.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
typ := remoteItemString(sm, "Type")
|
||||
if strings.EqualFold(typ, "Video") {
|
||||
if media.VideoCodec == "" {
|
||||
media.VideoCodec = remoteItemString(sm, "Codec")
|
||||
}
|
||||
if media.Width == 0 {
|
||||
media.Width = remoteItemInt(sm, "Width")
|
||||
}
|
||||
if media.Height == 0 {
|
||||
media.Height = remoteItemInt(sm, "Height")
|
||||
}
|
||||
} else if strings.EqualFold(typ, "Audio") {
|
||||
if media.AudioCodec == "" {
|
||||
media.AudioCodec = remoteItemString(sm, "Codec")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if streams, ok := item["MediaStreams"].([]any); ok {
|
||||
extractStreamInfo(streams)
|
||||
} else if streams, ok := item["MediaStreams"].([]map[string]any); ok {
|
||||
anyStreams := make([]any, len(streams))
|
||||
for i, v := range streams {
|
||||
anyStreams[i] = v
|
||||
}
|
||||
extractStreamInfo(anyStreams)
|
||||
}
|
||||
|
||||
if sources, ok := item["MediaSources"].([]any); ok && len(sources) > 0 {
|
||||
if sourceMap, ok := sources[0].(map[string]any); ok {
|
||||
if media.Container == "" {
|
||||
media.Container = remoteItemString(sourceMap, "Container")
|
||||
}
|
||||
if media.SizeBytes == 0 {
|
||||
media.SizeBytes = remoteItemInt64(sourceMap, "Size")
|
||||
}
|
||||
if streams, ok := sourceMap["MediaStreams"].([]any); ok && (media.VideoCodec == "" || media.AudioCodec == "") {
|
||||
extractStreamInfo(streams)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
switch remoteItemString(item, "Type") {
|
||||
case "Episode":
|
||||
media.SeasonNum = remoteItemInt(item, "ParentIndexNumber")
|
||||
@@ -221,13 +274,21 @@ func (r *EmbyRemoteService) RemoteLibraryMedia(ctx context.Context, mount *model
|
||||
if itemTypes == "" {
|
||||
itemTypes = "Movie,Series" // 未知类型时两者都取(前端自行按 episode-like 分组)
|
||||
}
|
||||
cacheKey := r.remoteCacheKey("library-media", acct.ID, mount.ID, remoteViewID, itemTypes, strconv.Itoa(offset), strconv.Itoa(limit))
|
||||
var cached struct {
|
||||
Items []model.Media `json:"items"`
|
||||
TotalRecordCount int64 `json:"total_record_count"`
|
||||
}
|
||||
if r.cache != nil && r.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
return cached.Items, cached.TotalRecordCount, nil
|
||||
}
|
||||
q := url.Values{}
|
||||
q.Set("ParentId", remoteViewID)
|
||||
q.Set("IncludeItemTypes", itemTypes)
|
||||
q.Set("Recursive", "false")
|
||||
q.Set("StartIndex", strconv.Itoa(offset))
|
||||
q.Set("Limit", strconv.Itoa(limit))
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources")
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
TotalRecordCount int64 `json:"TotalRecordCount"`
|
||||
@@ -240,6 +301,12 @@ func (r *EmbyRemoteService) RemoteLibraryMedia(ctx context.Context, mount *model
|
||||
RewriteEmbyRemoteIDs(it, mount.ID) // 嵌套/关联 ID 一并伪装
|
||||
items = append(items, r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it))
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, struct {
|
||||
Items []model.Media `json:"items"`
|
||||
TotalRecordCount int64 `json:"total_record_count"`
|
||||
}{Items: items, TotalRecordCount: body.TotalRecordCount}, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return items, body.TotalRecordCount, nil
|
||||
}
|
||||
|
||||
@@ -250,7 +317,7 @@ func (r *EmbyRemoteService) RemoteMediaDetail(ctx context.Context, mount *model.
|
||||
return nil, err
|
||||
}
|
||||
path := "/Users/" + url.PathEscape(r.remoteUserID(cfg)) + "/Items/" + url.PathEscape(remoteID)
|
||||
path += "?Fields=Overview,Genres,ProviderIds,People,Studios,Path"
|
||||
path += "?Fields=Overview,Genres,ProviderIds,People,Studios,Path,MediaStreams,MediaSources"
|
||||
var out map[string]any
|
||||
if err := r.doGet(ctx, acct, cfg, path, nil, &out); err != nil {
|
||||
return nil, err
|
||||
@@ -311,7 +378,7 @@ func (r *EmbyRemoteService) remoteEpisodesOf(ctx context.Context, mount *model.E
|
||||
q.Set("Recursive", "true")
|
||||
q.Set("StartIndex", "0")
|
||||
q.Set("Limit", "500")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage")
|
||||
q.Set("Fields", "Overview,Genres,ProviderIds,Path,SeriesPrimaryImage,MediaStreams,MediaSources")
|
||||
var body struct {
|
||||
Items []map[string]any `json:"Items"`
|
||||
TotalRecordCount int64 `json:"TotalRecordCount"`
|
||||
@@ -334,6 +401,11 @@ func (r *EmbyRemoteService) RemoteSeriesCards(ctx context.Context, mount *model.
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cacheKey := r.remoteCacheKey("series-cards", acct.ID, mount.ID, remoteViewID)
|
||||
var cached []SeriesCard
|
||||
if r.cache != nil && r.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
return cached, nil
|
||||
}
|
||||
q := url.Values{}
|
||||
q.Set("ParentId", remoteViewID)
|
||||
q.Set("IncludeItemTypes", "Series")
|
||||
@@ -361,6 +433,9 @@ func (r *EmbyRemoteService) RemoteSeriesCards(ctx context.Context, mount *model.
|
||||
}
|
||||
cards = append(cards, SeriesCard{Key: m.ID, Rep: m, LinkMedia: m, Count: count})
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, cards, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return cards, nil
|
||||
}
|
||||
|
||||
@@ -370,6 +445,11 @@ func (r *EmbyRemoteService) RemoteLatestCards(ctx context.Context, mount *model.
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cacheKey := r.remoteCacheKey("latest-cards", acct.ID, mount.ID, remoteViewID, strconv.Itoa(limit))
|
||||
var cached []SeriesCard
|
||||
if r.cache != nil && r.cache.GetJSON(ctx, cacheKey, &cached) {
|
||||
return cached, nil
|
||||
}
|
||||
items, err := r.RemoteLatest(ctx, mount, acct, remoteViewID, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -379,6 +459,9 @@ func (r *EmbyRemoteService) RemoteLatestCards(ctx context.Context, mount *model.
|
||||
m := r.MapRemoteItemToMedia(ctx, mount, acct, cfg, it)
|
||||
cards = append(cards, SeriesCard{Key: m.ID, Rep: m, LinkMedia: m, Count: 0})
|
||||
}
|
||||
if r.cache != nil {
|
||||
r.cache.SetJSON(ctx, cacheKey, cards, r.remoteMediaCacheTTL())
|
||||
}
|
||||
return cards, nil
|
||||
}
|
||||
|
||||
@@ -449,6 +532,8 @@ func remoteItemInt(item map[string]any, key string) int {
|
||||
return int(v)
|
||||
case int:
|
||||
return v
|
||||
case int64:
|
||||
return int(v)
|
||||
case string:
|
||||
n, _ := strconv.Atoi(v)
|
||||
return n
|
||||
@@ -463,6 +548,8 @@ func remoteItemInt64(item map[string]any, key string) int64 {
|
||||
switch v := item[key].(type) {
|
||||
case float64:
|
||||
return int64(v)
|
||||
case int64:
|
||||
return v
|
||||
case int:
|
||||
return int64(v)
|
||||
case string:
|
||||
@@ -572,6 +659,7 @@ func remoteItemTypeOf(m *model.Media) string {
|
||||
}
|
||||
return "Movie"
|
||||
}
|
||||
|
||||
// ─── 供 handler 层使用的远程 View 条目取值(导出薄封装) ──────────────────────
|
||||
|
||||
// RemoteItemIDString 提取远程 View 条目的 Id。
|
||||
@@ -581,7 +669,9 @@ func RemoteItemIDString(item map[string]any) string { return remoteItemString(it
|
||||
func RemoteItemNameString(item map[string]any) string { return remoteItemString(item, "Name") }
|
||||
|
||||
// RemoteItemCollectionType 提取远程 View 条目的 CollectionType。
|
||||
func RemoteItemCollectionType(item map[string]any) string { return remoteItemString(item, "CollectionType") }
|
||||
func RemoteItemCollectionType(item map[string]any) string {
|
||||
return remoteItemString(item, "CollectionType")
|
||||
}
|
||||
|
||||
// RemoteItemChildCount 提取远程 View 条目的 ChildCount。
|
||||
func RemoteItemChildCount(item map[string]any) int { return remoteItemInt(item, "ChildCount") }
|
||||
|
||||
@@ -103,7 +103,7 @@ func (b *serviceContainerBuilder) initContentServices() {
|
||||
b.c.DLNA = NewDLNAService(b.log)
|
||||
b.c.Storage = NewStorageService(b.log, b.repos)
|
||||
b.c.Emby = NewEmbyService(b.cfg, b.log, b.repos)
|
||||
b.c.EmbyRemote = NewEmbyRemoteService(b.cfg, b.log, b.repos, b.c.Crypto)
|
||||
b.c.EmbyRemote = NewEmbyRemoteService(b.cfg, b.log, b.repos, b.c.Crypto).SetRuntimeCache(b.c.Cache)
|
||||
b.c.Emby.SetEmbyRemote(b.c.EmbyRemote)
|
||||
b.c.Backup = NewBackupService(b.cfg, b.log, b.repos.DB)
|
||||
b.c.Media = NewMediaService(b.cfg, b.log, b.repos).SetRuntimeCache(b.c.Cache)
|
||||
|
||||
@@ -86,9 +86,9 @@ func playableSTRMTarget(ctx context.Context, repo *repository.Container, raw str
|
||||
return STRMPlaybackEnabled(ctx, repo)
|
||||
}
|
||||
|
||||
// isStrmMediaRow 判断媒体行是否为 .strm(远程直链)媒体:STRMURL 非空、
|
||||
// IsStrmMediaRow 判断媒体行是否为 .strm(远程直链)媒体:STRMURL 非空、
|
||||
// container=strm 或路径以 .strm 结尾。strm 媒体只能直连播放,禁止转码。
|
||||
func isStrmMediaRow(m *model.Media) bool {
|
||||
func IsStrmMediaRow(m *model.Media) bool {
|
||||
if m == nil {
|
||||
return false
|
||||
}
|
||||
@@ -101,6 +101,10 @@ func isStrmMediaRow(m *model.Media) bool {
|
||||
return strings.HasSuffix(strings.ToLower(strings.TrimSpace(m.Path)), ".strm")
|
||||
}
|
||||
|
||||
func isStrmMediaRow(m *model.Media) bool {
|
||||
return IsStrmMediaRow(m)
|
||||
}
|
||||
|
||||
func isHTTPPlaybackTarget(raw string) bool {
|
||||
u, err := url.Parse(strings.TrimSpace(raw))
|
||||
if err != nil || u == nil || !u.IsAbs() {
|
||||
|
||||
@@ -609,11 +609,26 @@ func (s *StrmService) ClearCanceledUploadTasks(ctx context.Context) (int64, erro
|
||||
return s.repo.StrmUpload.ClearCanceled(ctx)
|
||||
}
|
||||
|
||||
// ClearDoneUploadTasks 清空全部已完成上传记录,返回删除数量。
|
||||
func (s *StrmService) ClearDoneUploadTasks(ctx context.Context) (int64, error) {
|
||||
return s.repo.StrmUpload.ClearDone(ctx)
|
||||
}
|
||||
|
||||
// ClearFinishedUploadTasks 清空全部已完成与失败的上传记录,返回删除数量。
|
||||
func (s *StrmService) ClearFinishedUploadTasks(ctx context.Context) (int64, error) {
|
||||
return s.repo.StrmUpload.ClearFinished(ctx)
|
||||
}
|
||||
|
||||
// RetryAllFailedDownloadTasks 批量重试所有失败下载任务,返回重新入队数量。
|
||||
func (s *StrmService) RetryAllFailedDownloadTasks(ctx context.Context) (int64, error) {
|
||||
return s.repo.StrmDownload.RetryAllFailed(ctx)
|
||||
}
|
||||
|
||||
// RetryAllFailedUploadTasks 批量重试所有失败上传任务,返回重新入队数量。
|
||||
func (s *StrmService) RetryAllFailedUploadTasks(ctx context.Context) (int64, error) {
|
||||
return s.repo.StrmUpload.RetryAllFailed(ctx)
|
||||
}
|
||||
|
||||
// CancelPendingDownloadTasks 批量取消所有排队下载任务,返回取消数量。
|
||||
func (s *StrmService) CancelPendingDownloadTasks(ctx context.Context) (int64, error) {
|
||||
return s.repo.StrmDownload.CancelPending(ctx)
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/ShukeBta/MMTL/internal/model"
|
||||
"github.com/ShukeBta/MMTL/internal/repository"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
func TestIs115Blocked(t *testing.T) {
|
||||
@@ -44,3 +49,58 @@ func TestIsHTTPDownloadFailure(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStrmUploadTasksClearAndRetry(t *testing.T) {
|
||||
db := newServiceTestDB(t, &model.StrmUploadTask{})
|
||||
repos := repository.New(db)
|
||||
svc := NewStrmService(nil, zap.NewNop(), repos, nil)
|
||||
ctx := context.Background()
|
||||
|
||||
tasks := []*model.StrmUploadTask{
|
||||
{Base: model.Base{ID: "task-pending"}, Status: model.StrmTaskPending, FileName: "1.nfo"},
|
||||
{Base: model.Base{ID: "task-running"}, Status: model.StrmTaskRunning, FileName: "2.nfo"},
|
||||
{Base: model.Base{ID: "task-done"}, Status: model.StrmTaskDone, FileName: "3.nfo"},
|
||||
{Base: model.Base{ID: "task-failed"}, Status: model.StrmTaskFailed, FileName: "4.nfo", Error: "some error", RetryCount: 3},
|
||||
{Base: model.Base{ID: "task-canceled"}, Status: model.StrmTaskCanceled, FileName: "5.nfo"},
|
||||
}
|
||||
for _, task := range tasks {
|
||||
if err := db.Create(task).Error; err != nil {
|
||||
t.Fatalf("failed to insert task: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 1. RetryAllFailedUploadTasks
|
||||
retried, err := svc.RetryAllFailedUploadTasks(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("RetryAllFailedUploadTasks failed: %v", err)
|
||||
}
|
||||
if retried != 1 {
|
||||
t.Fatalf("expected 1 retried task, got %d", retried)
|
||||
}
|
||||
var failedTask model.StrmUploadTask
|
||||
if err := db.First(&failedTask, "id = ?", "task-failed").Error; err != nil {
|
||||
t.Fatalf("failed to get task-failed: %v", err)
|
||||
}
|
||||
if failedTask.Status != model.StrmTaskPending || failedTask.Error != "" || failedTask.RetryCount != 0 {
|
||||
t.Fatalf("task-failed was not reset properly: %+v", failedTask)
|
||||
}
|
||||
|
||||
// 再次改为 failed 以便测试 ClearFinished
|
||||
db.Model(&model.StrmUploadTask{}).Where("id = ?", "task-failed").Updates(map[string]any{"status": model.StrmTaskFailed})
|
||||
|
||||
// 2. ClearFinishedUploadTasks 应删除 done, failed, canceled 三条历史记录
|
||||
deleted, err := svc.ClearFinishedUploadTasks(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ClearFinishedUploadTasks failed: %v", err)
|
||||
}
|
||||
if deleted != 3 {
|
||||
t.Fatalf("expected 3 deleted tasks (done, failed, canceled), got %d", deleted)
|
||||
}
|
||||
|
||||
// 验证剩余的任务只有 pending 和 running
|
||||
var count int64
|
||||
db.Model(&model.StrmUploadTask{}).Count(&count)
|
||||
if count != 2 {
|
||||
t.Fatalf("expected 2 remaining tasks, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user