diff --git a/internal/handler/active_user.go b/internal/handler/active_user.go index 1f70fd0..662507c 100644 --- a/internal/handler/active_user.go +++ b/internal/handler/active_user.go @@ -10,6 +10,8 @@ import ( "github.com/ShukeBta/MediaStationGo/internal/service" ) +const embyCtxUserName = "emby_user_name" + func activeUserRequired(svc *service.Container) gin.HandlerFunc { return func(c *gin.Context) { uid, _ := c.Get(middleware.CtxUserID) @@ -72,6 +74,7 @@ func activeEmbyUserRequired(svc *service.Container) gin.HandlerFunc { c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"Code": 40303, "Message": "User account has expired"}) return } + c.Set(embyCtxUserName, u.Username) c.Next() } } diff --git a/internal/handler/emby_auth.go b/internal/handler/emby_auth.go index 8866b3e..cfda629 100644 --- a/internal/handler/emby_auth.go +++ b/internal/handler/emby_auth.go @@ -8,6 +8,7 @@ import ( "github.com/gin-gonic/gin" "github.com/ShukeBta/MediaStationGo/internal/middleware" + "github.com/ShukeBta/MediaStationGo/internal/service" ) // embyError 返回 Emby 风格的错误(顶层 Code/Message)。 @@ -49,6 +50,34 @@ func embyAuthRequiredWithSessionFallback(secret string) gin.HandlerFunc { } } +func embyRealtimeSessionActivity(svc *service.Container) gin.HandlerFunc { + return func(c *gin.Context) { + if svc != nil && svc.Sessions != nil { + if uid := embyUserID(c); uid != "" { + clientInfo := embyClientInfoFromRequest(c) + svc.Sessions.RecordActivity(c.Request.Context(), uid, embyContextUserName(c), + clientInfo.DeviceID, + clientInfo.DeviceName, + clientInfo.Client, + c.ClientIP()) + } + } + c.Next() + } +} + +func embyContextUserName(c *gin.Context) string { + if c == nil { + return "" + } + if value, ok := c.Get(embyCtxUserName); ok { + if username, ok := value.(string); ok { + return strings.TrimSpace(username) + } + } + return "" +} + func embyRememberCompatSession(c *gin.Context, token string) { token = strings.TrimSpace(token) if token == "" { diff --git a/internal/handler/emby_routes.go b/internal/handler/emby_routes.go index 967ee12..70f7e3a 100644 --- a/internal/handler/emby_routes.go +++ b/internal/handler/emby_routes.go @@ -20,7 +20,7 @@ func registerEmbyRoutes(r *gin.Engine, jwtSecret string, svc *service.Container) registerEmbyPublicImageRoutes(grp, svc) // 鉴权后端点 - auth := grp.Group("", embyAuthRequiredWithSessionFallback(jwtSecret), activeEmbyUserRequired(svc)) + auth := grp.Group("", embyAuthRequiredWithSessionFallback(jwtSecret), activeEmbyUserRequired(svc), embyRealtimeSessionActivity(svc)) registerEmbyAuthenticatedRoutes(auth, prefix, svc) } } diff --git a/internal/handler/emby_session_routes_test.go b/internal/handler/emby_session_routes_test.go index 25c00f5..0e7f494 100644 --- a/internal/handler/emby_session_routes_test.go +++ b/internal/handler/emby_session_routes_test.go @@ -7,6 +7,7 @@ import ( "net/http/httptest" "strings" "testing" + "time" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" @@ -153,6 +154,80 @@ func TestEmbyCompatSessionAllowsSameClientRequestsWithoutToken(t *testing.T) { } } +func TestEmbyAuthenticatedRequestRefreshesRealtimeUserActivity(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 sqlDB, err := db.DB(); err == nil { + sqlDB.SetMaxOpenConns(1) + } + if err := db.AutoMigrate(model.AllModels()...); err != nil { + t.Fatalf("migrate: %v", err) + } + repos := repository.New(db) + oldLogin := time.Now().Add(-6 * time.Hour) + if err := repos.User.Create(t.Context(), &model.User{ + Base: model.Base{ID: "user-1"}, + Username: "viewer", + Role: "admin", + Tier: "plus", + IsActive: true, + LastLoginAt: &oldLogin, + }); err != nil { + t.Fatalf("create user: %v", err) + } + + const secret = "test-secret" + log := zap.NewNop() + tracker := service.NewSessionTrackerService(log) + svc := &service.Container{ + Repo: repos, + Emby: service.NewEmbyService(&config.Config{}, log, repos), + Sessions: tracker, + } + router := gin.New() + registerEmbyRoutes(router, secret, svc) + router.GET("/admin/users", listUsersHandler(svc)) + + req := httptest.NewRequest(http.MethodGet, "/emby/Users/Me", nil) + req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="iPhone", DeviceId="phone-1", Token="`+signedTestToken(t, secret)+`"`) + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("me status: %d body=%s", w.Code, w.Body.String()) + } + + sessions := tracker.List(t.Context()) + if len(sessions) != 1 { + t.Fatalf("sessions = %#v, want one realtime session", sessions) + } + if sessions[0].UserID != "user-1" || sessions[0].DeviceID != "phone-1" || sessions[0].Client != "Infuse" { + t.Fatalf("session did not capture client info: %#v", sessions[0]) + } + + req = httptest.NewRequest(http.MethodGet, "/admin/users", nil) + w = httptest.NewRecorder() + router.ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("admin users status: %d body=%s", w.Code, w.Body.String()) + } + var users []model.User + if err := json.Unmarshal(w.Body.Bytes(), &users); err != nil { + t.Fatalf("decode users: %v", err) + } + if len(users) != 1 { + t.Fatalf("users = %#v", users) + } + if users[0].LastLoginAt == nil || !users[0].LastLoginAt.After(oldLogin) { + t.Fatalf("last_login_at = %v, want realtime value after %v", users[0].LastLoginAt, oldLogin) + } + if !users[0].RealtimeOnline || users[0].RealtimeDeviceCount != 1 { + t.Fatalf("realtime flags online=%v devices=%d", users[0].RealtimeOnline, users[0].RealtimeDeviceCount) + } +} + func TestEmbyUppercaseSessionCapabilitiesRouteNoContent(t *testing.T) { gin.SetMode(gin.TestMode) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) diff --git a/internal/service/session_tracker.go b/internal/service/session_tracker.go index 52f53f7..0b6f77e 100644 --- a/internal/service/session_tracker.go +++ b/internal/service/session_tracker.go @@ -53,6 +53,7 @@ type realtimeSessionInput struct { RuntimeTicks int64 IsPlaying bool IsPaused bool + PlaybackUpdate bool } type userRealtimeActivity struct { @@ -81,6 +82,10 @@ func NewSessionTrackerService(log *zap.Logger) *SessionTrackerService { } func (s *SessionTrackerService) RecordLogin(ctx context.Context, userID, userName, deviceID, deviceName, client, remoteEndPoint string) { + s.RecordActivity(ctx, userID, userName, deviceID, deviceName, client, remoteEndPoint) +} + +func (s *SessionTrackerService) RecordActivity(ctx context.Context, userID, userName, deviceID, deviceName, client, remoteEndPoint string) { if s == nil { return } @@ -109,6 +114,7 @@ func (s *SessionTrackerService) RecordPlayback(ctx context.Context, userID, user PositionTicks: positionTicks, RuntimeTicks: runtimeTicks, IsPlaying: !stopped, + PlaybackUpdate: true, }) } @@ -242,7 +248,7 @@ func (s *SessionTrackerService) upsert(ctx context.Context, in realtimeSessionIn s.mu.Lock() defer s.mu.Unlock() s.pruneLocked(now) - existing := s.sessions[key] + existing, existed := s.sessions[key] if strings.TrimSpace(in.UserName) == "" { in.UserName = existing.UserName } @@ -256,7 +262,23 @@ func (s *SessionTrackerService) upsert(ctx context.Context, in realtimeSessionIn in.RemoteEndPoint = existing.RemoteEndPoint } lastPlaybackAt := existing.LastPlaybackAt - if in.ItemID != "" || in.IsPlaying { + itemID := existing.ItemID + positionTicks := existing.PositionTicks + runtimeTicks := existing.RuntimeTicks + isPlaying := existing.IsPlaying + isPaused := existing.IsPaused + if in.PlaybackUpdate { + itemID = firstNonEmptyString(in.ItemID, existing.ItemID) + if in.ItemID != "" || in.PositionTicks != 0 { + positionTicks = in.PositionTicks + } + if in.ItemID != "" || in.RuntimeTicks != 0 { + runtimeTicks = in.RuntimeTicks + } + isPlaying = in.IsPlaying + isPaused = in.IsPaused + } + if in.PlaybackUpdate && (in.ItemID != "" || in.IsPlaying) { t := now lastPlaybackAt = &t } @@ -269,13 +291,21 @@ func (s *SessionTrackerService) upsert(ctx context.Context, in realtimeSessionIn Client: in.Client, RemoteEndPoint: in.RemoteEndPoint, LastActivityAt: now, - ItemID: firstNonEmptyString(in.ItemID, existing.ItemID), - PositionTicks: in.PositionTicks, - RuntimeTicks: in.RuntimeTicks, - IsPlaying: in.IsPlaying, - IsPaused: in.IsPaused, + ItemID: itemID, + PositionTicks: positionTicks, + RuntimeTicks: runtimeTicks, + IsPlaying: isPlaying, + IsPaused: isPaused, LastPlaybackAt: lastPlaybackAt, } + if !existed && s.log != nil { + s.log.Debug("realtime session started", + zap.String("user_id", userID), + zap.String("device_id", in.DeviceID), + zap.String("client", in.Client), + zap.String("remote", in.RemoteEndPoint), + ) + } } func (s *SessionTrackerService) pruneLocked(now time.Time) { diff --git a/internal/service/session_tracker_test.go b/internal/service/session_tracker_test.go index bcc2125..0eb0d42 100644 --- a/internal/service/session_tracker_test.go +++ b/internal/service/session_tracker_test.go @@ -1,6 +1,7 @@ package service import ( + "strings" "testing" "time" @@ -56,6 +57,56 @@ func TestDeviceListMergesRealtimeSessions(t *testing.T) { } } +func TestActivityRefreshKeepsPlaybackState(t *testing.T) { + tracker := NewSessionTrackerService(zap.NewNop()) + now := time.Date(2026, 6, 21, 11, 30, 0, 0, time.UTC) + tracker.now = func() time.Time { return now } + tracker.RecordPlayback(t.Context(), "u1", "viewer", "dev-1", "Apple TV", "Yamby", "10.0.0.8", "media-1", 123, 456, false) + now = now.Add(time.Minute) + tracker.RecordActivity(t.Context(), "u1", "viewer", "dev-1", "Apple TV", "Yamby", "10.0.0.8") + + sessions := tracker.List(t.Context()) + if len(sessions) != 1 { + t.Fatalf("sessions = %#v, want one", sessions) + } + if !sessions[0].IsPlaying || sessions[0].ItemID != "media-1" || sessions[0].PositionTicks != 123 || sessions[0].RuntimeTicks != 456 { + t.Fatalf("activity refresh should keep playback state, got %#v", sessions[0]) + } + if !sessions[0].LastActivityAt.Equal(now) { + t.Fatalf("last activity = %v, want %v", sessions[0].LastActivityAt, now) + } +} + +func TestBotDevicesIncludesRealtimeSessionOnlyDevices(t *testing.T) { + repos, bot := newBotTestService(t) + user := model.User{Base: model.Base{ID: "u1"}, Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true} + if err := repos.User.Create(t.Context(), &user); err != nil { + t.Fatal(err) + } + if err := repos.DB.Create(&model.TelegramBinding{TelegramUserID: 9103, ChatID: 9103, UserID: user.ID}).Error; err != nil { + t.Fatal(err) + } + tracker := NewSessionTrackerService(zap.NewNop()) + now := time.Date(2026, 6, 21, 11, 0, 0, 0, time.UTC) + tracker.now = func() time.Time { return now } + tracker.RecordActivity(t.Context(), user.ID, user.Username, "dev-1", "Apple TV", "Yamby", "10.0.0.8") + device := NewDeviceService(zap.NewNop(), repos) + device.SetSessionTracker(tracker) + bot.SetDeviceService(device) + + reply := bot.replyDevices(t.Context(), &TelegramMessage{ + From: TelegramUser{ID: 9103, Username: "viewer"}, + Chat: TelegramChat{ID: 9103, Type: "private"}, + }) + + if !strings.Contains(reply.Text, "Apple TV / Yamby") || !strings.Contains(reply.Text, "在线") { + t.Fatalf("reply should include realtime online device, got %q", reply.Text) + } + if !strings.Contains(reply.Text, "06-21 11:00") { + t.Fatalf("reply should use realtime last seen time, got %q", reply.Text) + } +} + func TestRealtimeRecentLoginProtectsCleanupCandidate(t *testing.T) { repos := newSessionTrackerTestRepos(t) now := time.Date(2026, 6, 21, 12, 0, 0, 0, time.UTC)