diff --git a/internal/handler/emby_playstate_handlers.go b/internal/handler/emby_playstate_handlers.go index f9a3ee2..d8e528c 100644 --- a/internal/handler/emby_playstate_handlers.go +++ b/internal/handler/emby_playstate_handlers.go @@ -43,7 +43,10 @@ func embyPlayingProgressHandler(svc *service.Container) gin.HandlerFunc { c.Status(http.StatusUnauthorized) return } - _ = svc.Emby.RecordProgress(c.Request.Context(), uid, req.ItemId, req.PositionTicks, req.RunTimeTicks) + if err := svc.Emby.RecordProgress(c.Request.Context(), uid, req.ItemId, req.PositionTicks, req.RunTimeTicks); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } stopped := strings.Contains(strings.ToLower(c.FullPath()+" "+c.Request.URL.Path), "stopped") if svc.Sessions != nil { svc.Sessions.RecordPlayback(c.Request.Context(), uid, "", diff --git a/internal/service/emby_compat.go b/internal/service/emby_compat.go index a775e12..6c12ac1 100644 --- a/internal/service/emby_compat.go +++ b/internal/service/emby_compat.go @@ -153,7 +153,14 @@ func (e *EmbyService) Items(ctx context.Context, p ItemsParams) (map[string]any, if mount == nil || acct == nil { return emptyItemsEnvelope(p.StartIndex), nil } - return e.remote.RemoteItems(ctx, mount, acct, p) + out, err := e.remote.RemoteItems(ctx, mount, acct, p) + if err != nil { + return nil, err + } + if err := e.mergeRemoteUserData(ctx, p.UserID, out); err != nil { + return nil, err + } + return out, nil } // 全局搜索:无 ParentId 且带搜索词 → 聚合本地 + 全部远程。 if p.ParentID == "" && p.SearchTerm != "" { @@ -274,6 +281,9 @@ func (e *EmbyService) aggregatedSearch(ctx context.Context, p ItemsParams) (map[ } continue } + if err := e.mergeRemoteUserData(ctx, p.UserID, remote); err != nil { + return nil, err + } if raw, ok := remote["Items"].([]any); ok { results = append(results, remoteResult{items: raw}) } else if rawMap, ok := remote["Items"].([]map[string]any); ok { diff --git a/internal/service/emby_items_detail.go b/internal/service/emby_items_detail.go index 8182b03..c78613f 100644 --- a/internal/service/emby_items_detail.go +++ b/internal/service/emby_items_detail.go @@ -18,7 +18,14 @@ func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[str if mount == nil || acct == nil { return nil, nil } - return e.remote.RemoteItem(ctx, mount, acct, remoteID) + out, err := e.remote.RemoteItem(ctx, mount, acct, remoteID) + if err != nil || out == nil { + return out, err + } + if err := e.mergeRemoteUserData(ctx, userID, out); err != nil { + return nil, err + } + return out, nil } if lib, err := e.repo.Library.FindByID(ctx, mediaID); err != nil { return nil, err @@ -91,7 +98,14 @@ func (e *EmbyService) LatestItems(ctx context.Context, userID, parentID string, if mount == nil || acct == nil { return nil, nil } - return e.remote.RemoteLatest(ctx, mount, acct, remoteParent, limit) + out, err := e.remote.RemoteLatest(ctx, mount, acct, remoteParent, limit) + if err != nil { + return nil, err + } + if err := e.mergeRemoteUserData(ctx, userID, out); err != nil { + return nil, err + } + return out, nil } cacheKey := e.embyLatestCacheKey(userID, parentID, limit) var cached embyLatestCacheValue @@ -172,27 +186,46 @@ func (e *EmbyService) ResumeItems(ctx context.Context, userID string, limit int) if len(hist) == 0 { return map[string]any{"Items": []any{}, "TotalRecordCount": 0}, nil } - ids := make([]string, 0, len(hist)) - posByID := map[string]int64{} + + localIDs := make([]string, 0, len(hist)) for _, h := range hist { - ids = append(ids, h.MediaID) - posByID[h.MediaID] = h.PositionMs - } - var medias []model.Media - q := e.repo.DB.WithContext(ctx).Where("id IN ?", ids) - q = e.applyUserMediaVisibility(ctx, q, userID) - if err := q.Find(&medias).Error; err != nil { - return nil, err + if !IsEmbyRemoteID(h.MediaID) { + localIDs = append(localIDs, h.MediaID) + } } byID := map[string]*model.Media{} - for i := range medias { - byID[medias[i].ID] = &medias[i] + if len(localIDs) > 0 { + var medias []model.Media + q := e.repo.DB.WithContext(ctx).Where("id IN ?", localIDs) + q = e.applyUserMediaVisibility(ctx, q, userID) + if err := q.Find(&medias).Error; err != nil { + return nil, err + } + for i := range medias { + byID[medias[i].ID] = &medias[i] + } } + items := make([]map[string]any, 0, len(hist)) for _, h := range hist { if m, ok := byID[h.MediaID]; ok { - items = append(items, e.itemPayload(ctx, m, false, posByID[h.MediaID])) + items = append(items, e.itemPayload(ctx, m, false, h.PositionMs)) + continue } + if e.remote == nil || !IsEmbyRemoteID(h.MediaID) { + continue + } + mountID, remoteID, _ := DecodeEmbyRemoteID(h.MediaID) + mount, acct, err := e.remote.ResolveMount(ctx, mountID) + if err != nil || mount == nil || acct == nil { + continue + } + item, err := e.remote.RemoteItem(ctx, mount, acct, remoteID) + if err != nil || item == nil { + continue + } + item["UserData"] = mergedRemoteUserData(item["UserData"], &h) + items = append(items, item) } return map[string]any{"Items": items, "TotalRecordCount": len(items)}, nil } diff --git a/internal/service/emby_playback.go b/internal/service/emby_playback.go index b97d927..25f393f 100644 --- a/internal/service/emby_playback.go +++ b/internal/service/emby_playback.go @@ -31,6 +31,9 @@ func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID, userID string) if out == nil { return nil, ErrEmbyRemoteNotFound } + if err := e.mergeRemoteUserData(ctx, userID, out); err != nil { + return nil, err + } out["PlaySessionId"] = fmt.Sprintf("remote-%s-%d", mountID, time.Now().Unix()) return out, nil } diff --git a/internal/service/emby_user_data.go b/internal/service/emby_user_data.go index c656f85..835dbd7 100644 --- a/internal/service/emby_user_data.go +++ b/internal/service/emby_user_data.go @@ -49,9 +49,13 @@ func (e *EmbyService) MarkPlayed(ctx context.Context, userID, mediaID string, pl return nil } if !played { - return e.repo.DB.WithContext(ctx). + err := e.repo.DB.WithContext(ctx). Where("user_id = ? AND media_id = ?", userID, mediaID). Delete(&model.PlaybackHistory{}).Error + if err == nil { + e.invalidateEmbyItemsCache(ctx) + } + return err } m, err := e.repo.Media.FindByID(ctx, mediaID) if err != nil || m == nil { @@ -61,7 +65,7 @@ func (e *EmbyService) MarkPlayed(ctx context.Context, userID, mediaID string, pl if dur <= 0 { dur = 1 } - return e.repo.History.Upsert(ctx, &model.PlaybackHistory{ + err = e.repo.History.Upsert(ctx, &model.PlaybackHistory{ UserID: userID, MediaID: mediaID, PositionMs: dur, @@ -69,6 +73,10 @@ func (e *EmbyService) MarkPlayed(ctx context.Context, userID, mediaID string, pl WatchedAt: time.Now(), Completed: true, }) + if err == nil { + e.invalidateEmbyItemsCache(ctx) + } + return err } // RecordProgress 记录播放进度(来自 Emby 客户端的 /Sessions/Playing/Progress)。 @@ -82,7 +90,7 @@ func (e *EmbyService) RecordProgress(ctx context.Context, userID, mediaID string } } completed := dur > 0 && pos >= dur*9/10 - return e.repo.History.Upsert(ctx, &model.PlaybackHistory{ + err := e.repo.History.Upsert(ctx, &model.PlaybackHistory{ UserID: userID, MediaID: mediaID, PositionMs: pos, @@ -90,6 +98,115 @@ func (e *EmbyService) RecordProgress(ctx context.Context, userID, mediaID string WatchedAt: time.Now(), Completed: completed, }) + if err == nil { + e.invalidateEmbyItemsCache(ctx) + } + return err +} + +// mergeRemoteUserData applies the current MMTL user's locally recorded playback +// state to remote Emby payloads. Remote metadata remains authoritative unless the +// user has played the item through MMTL. +func (e *EmbyService) mergeRemoteUserData(ctx context.Context, userID string, payload any) error { + if strings.TrimSpace(userID) == "" || payload == nil { + return nil + } + items := remoteItemMaps(payload) + ids := make([]string, 0, len(items)) + seen := make(map[string]struct{}, len(items)) + for _, item := range items { + id, _ := item["Id"].(string) + if !IsEmbyRemoteID(id) { + continue + } + if _, ok := seen[id]; !ok { + ids = append(ids, id) + seen[id] = struct{}{} + } + } + if len(ids) == 0 { + return nil + } + var histories []model.PlaybackHistory + if err := e.repo.DB.WithContext(ctx).Where("user_id = ? AND media_id IN ?", userID, ids).Find(&histories).Error; err != nil { + return err + } + byMediaID := make(map[string]*model.PlaybackHistory, len(histories)) + for i := range histories { + byMediaID[histories[i].MediaID] = &histories[i] + } + for _, item := range items { + id, _ := item["Id"].(string) + if h := byMediaID[id]; h != nil { + item["UserData"] = mergedRemoteUserData(item["UserData"], h) + } + } + return nil +} + +func remoteItemMaps(payload any) []map[string]any { + items := make([]map[string]any, 0) + var visit func(any) + visit = func(value any) { + switch typed := value.(type) { + case map[string]any: + if _, ok := typed["Id"].(string); ok { + items = append(items, typed) + } + if nested, ok := typed["Items"]; ok { + visit(nested) + } + case []any: + for _, value := range typed { + visit(value) + } + case []map[string]any: + for _, value := range typed { + visit(value) + } + } + } + visit(payload) + return items +} + +func mergedRemoteUserData(raw any, history *model.PlaybackHistory) map[string]any { + userData := map[string]any{} + if existing, ok := raw.(map[string]any); ok { + for key, value := range existing { + userData[key] = value + } + } + duration := history.DurationMs + position := history.PositionMs + percentage := float64(0) + if duration > 0 { + percentage = float64(position) / float64(duration) * 100 + } + userData["PlaybackPositionTicks"] = position * 10_000 + userData["Played"] = history.Completed + userData["PlayedPercentage"] = percentage + if history.Completed { + playCount := 0 + switch value := userData["PlayCount"].(type) { + case int: + playCount = value + case int64: + playCount = int(value) + case float64: + playCount = int(value) + } + if playCount < 1 { + userData["PlayCount"] = 1 + } + } + return userData +} + +func (e *EmbyService) invalidateEmbyItemsCache(ctx context.Context) { + if e.cache != nil { + e.cache.DeletePrefix(ctx, "media:emby:") + } } func splitCSV(s string) []string { diff --git a/internal/service/emby_user_data_test.go b/internal/service/emby_user_data_test.go new file mode 100644 index 0000000..43b6ab3 --- /dev/null +++ b/internal/service/emby_user_data_test.go @@ -0,0 +1,81 @@ +package service + +import ( + "testing" + + "github.com/ShukeBta/MMTL/internal/model" +) + +func TestMergedRemoteUserData(t *testing.T) { + tests := []struct { + name string + raw any + history model.PlaybackHistory + position int64 + played bool + percent float64 + count int + preserve any + }{ + { + name: "in-progress preserves remote fields", + raw: map[string]any{ + "PlayCount": 2, + "Custom": "remote-value", + }, + history: model.PlaybackHistory{PositionMs: 25_000, DurationMs: 100_000}, + position: 250_000_000, + played: false, + percent: 25, + count: 2, + preserve: "remote-value", + }, + { + name: "completed ensures a play count", + raw: map[string]any{"PlayCount": 0}, + history: model.PlaybackHistory{PositionMs: 100_000, DurationMs: 100_000, Completed: true}, + position: 1_000_000_000, + played: true, + percent: 100, + count: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + out := mergedRemoteUserData(tt.raw, &tt.history) + if got := out["PlaybackPositionTicks"]; got != tt.position { + t.Fatalf("PlaybackPositionTicks = %#v, want %d", got, tt.position) + } + if got := out["Played"]; got != tt.played { + t.Fatalf("Played = %#v, want %t", got, tt.played) + } + if got := out["PlayedPercentage"]; got != tt.percent { + t.Fatalf("PlayedPercentage = %#v, want %v", got, tt.percent) + } + if got := out["PlayCount"]; got != tt.count { + t.Fatalf("PlayCount = %#v, want %d", got, tt.count) + } + if tt.preserve != nil && out["Custom"] != tt.preserve { + t.Fatalf("Custom = %#v, want %#v", out["Custom"], tt.preserve) + } + }) + } +} + +func TestRemoteItemMapsFindsEnvelopeItems(t *testing.T) { + remoteID := EncodeEmbyRemoteID("mount-1", "item-1") + payload := map[string]any{ + "Items": []any{ + map[string]any{"Id": remoteID}, + map[string]any{"Id": "local-item"}, + }, + } + items := remoteItemMaps(payload) + if len(items) != 2 { + t.Fatalf("item count = %d, want 2", len(items)) + } + if items[0]["Id"] != remoteID { + t.Fatalf("first item ID = %#v, want %q", items[0]["Id"], remoteID) + } +}