diff --git a/internal/handler/emby_auth.go b/internal/handler/emby_auth.go index cfda629..ffe32ac 100644 --- a/internal/handler/emby_auth.go +++ b/internal/handler/emby_auth.go @@ -52,20 +52,23 @@ 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()) - } - } + recordEmbySessionActivity(c, svc, embyUserID(c), embyContextUserName(c)) c.Next() } } +func recordEmbySessionActivity(c *gin.Context, svc *service.Container, userID, userName string) { + if c == nil || svc == nil || svc.Sessions == nil || strings.TrimSpace(userID) == "" { + return + } + clientInfo := embyClientInfoFromRequest(c) + svc.Sessions.RecordActivity(c.Request.Context(), userID, userName, + clientInfo.DeviceID, + clientInfo.DeviceName, + clientInfo.Client, + c.ClientIP()) +} + func embyContextUserName(c *gin.Context) string { if c == nil { return "" @@ -167,15 +170,25 @@ func embyClientInfoFromRequest(c *gin.Context) embyClientInfo { info := embyClientInfo{ DeviceID: firstNonEmptyHeaderString( firstHeaderValue(c, "X-Emby-Device-Id", "X-Emby-DeviceId", "X-MediaBrowser-Device-Id", "X-MediaBrowser-DeviceId"), + c.Query("DeviceId"), + c.Query("DeviceID"), + c.Query("deviceId"), + c.Query("deviceID"), auth["DeviceId"], auth["DeviceID"], ), DeviceName: firstNonEmptyHeaderString( firstHeaderValue(c, "X-Emby-Device-Name", "X-Emby-DeviceName", "X-MediaBrowser-Device-Name", "X-MediaBrowser-DeviceName"), + c.Query("Device"), + c.Query("DeviceName"), + c.Query("device"), + c.Query("deviceName"), auth["Device"], ), Client: firstNonEmptyHeaderString( firstHeaderValue(c, "X-Emby-Client", "X-MediaBrowser-Client"), + c.Query("Client"), + c.Query("client"), auth["Client"], ), } diff --git a/internal/handler/emby_misc_routes_test.go b/internal/handler/emby_misc_routes_test.go index 52858e6..9c6d527 100644 --- a/internal/handler/emby_misc_routes_test.go +++ b/internal/handler/emby_misc_routes_test.go @@ -10,6 +10,7 @@ import ( "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" "github.com/gorilla/websocket" + "go.uber.org/zap" "gorm.io/gorm" "github.com/ShukeBta/MediaStationGo/internal/model" @@ -179,3 +180,54 @@ func TestEmbyWebSocketRouteUpgradesForOfficialClients(t *testing.T) { t.Fatalf("expected websocket upgrade, got resp=%#v", resp) } } + +func TestEmbyWebSocketRefreshesRealtimeActivity(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.User{}); 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: "viewer", + PasswordHash: "x", + Role: "admin", + Tier: "plus", + IsActive: true, + }); err != nil { + t.Fatalf("create user: %v", err) + } + const secret = "test-secret" + tracker := service.NewSessionTrackerService(zap.NewNop()) + router := gin.New() + registerEmbyRoutes(router, secret, &service.Container{ + Repo: repos, + Sessions: tracker, + }) + server := httptest.NewServer(router) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/embywebsocket?deviceId=device-1&device=Windows&client=Emby" + header := http.Header{"X-Emby-Token": []string{signedTestToken(t, secret)}} + conn, resp, err := websocket.DefaultDialer.Dial(wsURL, header) + if err != nil { + status := 0 + if resp != nil { + status = resp.StatusCode + } + t.Fatalf("websocket dial failed status=%d err=%v", status, err) + } + defer conn.Close() + + sessions := tracker.List(t.Context()) + if len(sessions) != 1 { + t.Fatalf("sessions = %#v, want websocket heartbeat session", sessions) + } + if sessions[0].DeviceID != "device-1" || sessions[0].DeviceName != "Windows" || sessions[0].Client != "Emby" { + t.Fatalf("websocket did not refresh client session: %#v", sessions[0]) + } +} diff --git a/internal/handler/emby_routes.go b/internal/handler/emby_routes.go index 70f7e3a..0f533c6 100644 --- a/internal/handler/emby_routes.go +++ b/internal/handler/emby_routes.go @@ -16,7 +16,7 @@ func registerEmbyRoutes(r *gin.Engine, jwtSecret string, svc *service.Container) grp.Use(embyNoStoreHeaders()) registerEmbyRootRoutes(grp, prefix, svc) - registerEmbyPublicRoutes(grp, svc) + registerEmbyPublicRoutes(grp, jwtSecret, svc) registerEmbyPublicImageRoutes(grp, svc) // 鉴权后端点 @@ -46,10 +46,10 @@ func registerEmbyRootRoutes(grp *gin.RouterGroup, prefix string, svc *service.Co grp.HEAD("/", embyRootHandler(svc)) } -func registerEmbyPublicRoutes(grp *gin.RouterGroup, svc *service.Container) { +func registerEmbyPublicRoutes(grp *gin.RouterGroup, jwtSecret string, svc *service.Container) { registerEmbyPublicSystemRoutes(grp, svc) - registerEmbyPublicSessionRoutes(grp, svc) - registerEmbyPublicClientRoutes(grp, svc) + registerEmbyPublicSessionRoutes(grp, jwtSecret, svc) + registerEmbyPublicClientRoutes(grp, jwtSecret, svc) } func registerEmbyPublicSystemRoutes(grp *gin.RouterGroup, svc *service.Container) { @@ -68,11 +68,13 @@ func registerEmbyPublicSystemRoutes(grp *gin.RouterGroup, svc *service.Container } } -func registerEmbyPublicSessionRoutes(grp *gin.RouterGroup, svc *service.Container) { - registerEmbyPostRoutes(grp, svc, []string{ +func registerEmbyPublicSessionRoutes(grp *gin.RouterGroup, jwtSecret string, svc *service.Container) { + for _, path := range []string{ "/Sessions/Capabilities", "/Sessions/Capabilities/Full", "/sessions/capabilities", "/sessions/capabilities/full", - }, embyNoContentHandler) + } { + grp.POST(path, embySessionCapabilitiesHandler(svc, jwtSecret)) + } // 30/min per IP: many Emby clients sit behind a single NAT/reverse-proxy // IP, so a low limit would throttle legitimate logins into 429s. @@ -84,17 +86,19 @@ func registerEmbyPublicSessionRoutes(grp *gin.RouterGroup, svc *service.Containe registerEmbyGetRoutes(grp, svc, []string{"/Users/Public", "/users/public"}, embyPublicUsersHandler) } -func registerEmbyPublicClientRoutes(grp *gin.RouterGroup, svc *service.Container) { +func registerEmbyPublicClientRoutes(grp *gin.RouterGroup, jwtSecret string, svc *service.Container) { registerEmbyGetRoutes(grp, svc, []string{"/Branding/Configuration", "/branding/configuration"}, embyBrandingConfigHandler) registerEmbyGetHeadRoutes(grp, svc, []string{"/Branding/Css", "/branding/css"}, embyBrandingCSSHandler) registerEmbyGetRoutes(grp, svc, []string{"/Localization/Options", "/localization/options"}, embyLocalizationOptionsHandler) registerEmbyGetRoutes(grp, svc, []string{"/Localization/Cultures", "/Localization/cultures", "/localization/cultures"}, embyLocalizationCulturesHandler) registerEmbyGetHeadRoutes(grp, svc, []string{"/CustomCssJS/Scripts", "/customcssjs/scripts"}, embyCustomCSSJSScriptsHandler) for _, path := range []string{"/embywebsocket", "/EmbyWebSocket"} { - grp.GET(path, embyWebSocketHandler(svc)) + grp.GET(path, embyWebSocketHandler(svc, jwtSecret)) grp.HEAD(path, embyNoContentHandler(svc)) } - registerEmbyPostRoutes(grp, svc, []string{"/Sessions/Logout", "/sessions/logout"}, embySessionLogoutHandler) + for _, path := range []string{"/Sessions/Logout", "/sessions/logout"} { + grp.POST(path, embySessionLogoutHandler(svc, jwtSecret)) + } grp.GET("/DisplayPreferences/:id", embyDisplayPreferencesHandler(svc)) grp.POST("/DisplayPreferences/:id", embySaveDisplayPreferencesHandler(svc)) grp.GET("/displaypreferences/:id", embyDisplayPreferencesHandler(svc)) diff --git a/internal/handler/emby_session_routes_test.go b/internal/handler/emby_session_routes_test.go index 0e7f494..2d0e78d 100644 --- a/internal/handler/emby_session_routes_test.go +++ b/internal/handler/emby_session_routes_test.go @@ -228,6 +228,188 @@ func TestEmbyAuthenticatedRequestRefreshesRealtimeUserActivity(t *testing.T) { } } +func TestEmbySessionCapabilitiesRefreshesRealtimeActivity(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.User{}); 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: "viewer", + PasswordHash: "x", + Role: "admin", + Tier: "plus", + IsActive: true, + }); err != nil { + t.Fatalf("create user: %v", err) + } + + const secret = "test-secret" + tracker := service.NewSessionTrackerService(zap.NewNop()) + router := gin.New() + registerEmbyRoutes(router, secret, &service.Container{ + Repo: repos, + Sessions: tracker, + }) + + req := httptest.NewRequest(http.MethodPost, "/emby/Sessions/Capabilities/Full?deviceId=phone-1&device=iPhone&client=Infuse", strings.NewReader(`{}`)) + req.Header.Set("X-Emby-Token", signedTestToken(t, secret)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + if w.Code != http.StatusNoContent { + t.Fatalf("capabilities status: %d body=%s", w.Code, w.Body.String()) + } + sessions := tracker.List(t.Context()) + if len(sessions) != 1 { + t.Fatalf("sessions = %#v, want one heartbeat session", sessions) + } + if sessions[0].DeviceID != "phone-1" || sessions[0].DeviceName != "iPhone" || sessions[0].Client != "Infuse" || sessions[0].UserName != "viewer" { + t.Fatalf("capabilities did not refresh client session: %#v", sessions[0]) + } +} + +func TestEmbySessionCapabilitiesIgnoresScopedPlaybackToken(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.User{}); 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: "viewer", + PasswordHash: "x", + Role: "admin", + Tier: "plus", + IsActive: true, + }); err != nil { + t.Fatalf("create user: %v", err) + } + + const secret = "test-secret" + tracker := service.NewSessionTrackerService(zap.NewNop()) + router := gin.New() + registerEmbyRoutes(router, secret, &service.Container{ + Repo: repos, + Sessions: tracker, + }) + + req := httptest.NewRequest(http.MethodPost, "/emby/Sessions/Capabilities/Full?deviceId=phone-1&device=iPhone&client=Infuse", strings.NewReader(`{}`)) + req.Header.Set("X-Emby-Token", signedTestTokenWithPurpose(t, secret, service.ExternalPlaybackTokenPurpose)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + if w.Code != http.StatusNoContent { + t.Fatalf("capabilities status: %d body=%s", w.Code, w.Body.String()) + } + if sessions := tracker.List(t.Context()); len(sessions) != 0 { + t.Fatalf("scoped external playback token must not create realtime session: %#v", sessions) + } +} + +func TestEmbyLogoutRemovesRealtimeSessionFromPublicRoute(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.User{}); 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: "viewer", + PasswordHash: "x", + Role: "admin", + Tier: "plus", + IsActive: true, + }); err != nil { + t.Fatalf("create user: %v", err) + } + + const secret = "test-secret" + tracker := service.NewSessionTrackerService(zap.NewNop()) + tracker.RecordActivity(t.Context(), "user-1", "viewer", "phone-1", "iPhone", "Infuse", "192.0.2.10") + tracker.RecordActivity(t.Context(), "user-1", "viewer", "tv-1", "Apple TV", "Yamby", "192.0.2.11") + router := gin.New() + registerEmbyRoutes(router, secret, &service.Container{ + Repo: repos, + Sessions: tracker, + }) + + req := httptest.NewRequest(http.MethodPost, "/emby/Sessions/Logout", 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.StatusNoContent { + t.Fatalf("logout status: %d body=%s", w.Code, w.Body.String()) + } + sessions := tracker.List(t.Context()) + if len(sessions) != 1 { + t.Fatalf("sessions after logout = %#v, want only the other device", sessions) + } + if sessions[0].DeviceID != "tv-1" { + t.Fatalf("remaining session = %#v, want tv-1", sessions[0]) + } +} + +func TestEmbyLogoutIgnoresScopedPlaybackToken(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.User{}); 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: "viewer", + PasswordHash: "x", + Role: "admin", + Tier: "plus", + IsActive: true, + }); err != nil { + t.Fatalf("create user: %v", err) + } + + const secret = "test-secret" + tracker := service.NewSessionTrackerService(zap.NewNop()) + tracker.RecordActivity(t.Context(), "user-1", "viewer", "phone-1", "iPhone", "Infuse", "192.0.2.10") + router := gin.New() + registerEmbyRoutes(router, secret, &service.Container{ + Repo: repos, + Sessions: tracker, + }) + + req := httptest.NewRequest(http.MethodPost, "/emby/Sessions/Logout", nil) + req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="iPhone", DeviceId="phone-1", Token="`+signedTestTokenWithPurpose(t, secret, service.ExternalPlaybackTokenPurpose)+`"`) + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + if w.Code != http.StatusNoContent { + t.Fatalf("logout status: %d body=%s", w.Code, w.Body.String()) + } + sessions := tracker.List(t.Context()) + if len(sessions) != 1 || sessions[0].DeviceID != "phone-1" { + t.Fatalf("scoped external playback token must not logout realtime session: %#v", sessions) + } +} + func TestEmbyUppercaseSessionCapabilitiesRouteNoContent(t *testing.T) { gin.SetMode(gin.TestMode) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) diff --git a/internal/handler/emby_sessions.go b/internal/handler/emby_sessions.go index a52540b..8bcab26 100644 --- a/internal/handler/emby_sessions.go +++ b/internal/handler/emby_sessions.go @@ -2,9 +2,12 @@ package handler import ( "net/http" + "strings" "github.com/gin-gonic/gin" + "github.com/golang-jwt/jwt/v5" + "github.com/ShukeBta/MediaStationGo/internal/middleware" "github.com/ShukeBta/MediaStationGo/internal/service" ) @@ -47,13 +50,58 @@ func embySessionsHandler(svc *service.Container) gin.HandlerFunc { } } -func embySessionLogoutHandler(svc *service.Container) gin.HandlerFunc { +func embySessionLogoutHandler(svc *service.Container, jwtSecret string) gin.HandlerFunc { return func(c *gin.Context) { if svc.Sessions != nil { - uid := embyUserID(c) + uid, _ := embyPublicSessionIdentity(c, svc, jwtSecret) clientInfo := embyClientInfoFromRequest(c) svc.Sessions.Logout(c.Request.Context(), uid, clientInfo.DeviceID, c.ClientIP()) } c.Status(http.StatusNoContent) } } + +func embySessionCapabilitiesHandler(svc *service.Container, jwtSecret string) gin.HandlerFunc { + return func(c *gin.Context) { + recordEmbyPublicSessionActivity(c, svc, jwtSecret) + c.Status(http.StatusNoContent) + } +} + +func recordEmbyPublicSessionActivity(c *gin.Context, svc *service.Container, jwtSecret string) { + uid, username := embyPublicSessionIdentity(c, svc, jwtSecret) + recordEmbySessionActivity(c, svc, uid, username) +} + +func embyPublicSessionIdentity(c *gin.Context, svc *service.Container, jwtSecret string) (string, string) { + if uid := embyUserID(c); uid != "" { + return uid, embyContextUserName(c) + } + token := embyRequestToken(c) + if strings.TrimSpace(token) == "" || strings.TrimSpace(jwtSecret) == "" { + return "", "" + } + claims := &middleware.Claims{} + parsed, err := jwt.ParseWithClaims(token, claims, func(t *jwt.Token) (interface{}, error) { + if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, jwt.ErrTokenSignatureInvalid + } + return []byte(jwtSecret), nil + }) + if err != nil || !parsed.Valid { + return "", "" + } + if strings.TrimSpace(claims.Purpose) != "" { + return "", "" + } + uid := strings.TrimSpace(claims.UserID) + if uid == "" { + return "", "" + } + if svc != nil && svc.Repo != nil && svc.Repo.User != nil { + if user, err := svc.Repo.User.FindByID(c.Request.Context(), uid); err == nil && user != nil { + return uid, user.Username + } + } + return uid, "" +} diff --git a/internal/handler/emby_static.go b/internal/handler/emby_static.go index 20740d8..d4b1501 100644 --- a/internal/handler/emby_static.go +++ b/internal/handler/emby_static.go @@ -16,8 +16,9 @@ func embyNoContentHandler(_ *service.Container) gin.HandlerFunc { } } -func embyWebSocketHandler(_ *service.Container) gin.HandlerFunc { +func embyWebSocketHandler(svc *service.Container, jwtSecret string) gin.HandlerFunc { return func(c *gin.Context) { + recordEmbyPublicSessionActivity(c, svc, jwtSecret) if !websocket.IsWebSocketUpgrade(c.Request) { c.Status(http.StatusNoContent) return @@ -45,6 +46,7 @@ func embyWebSocketHandler(_ *service.Container) gin.HandlerFunc { case <-done: return case <-ticker.C: + recordEmbyPublicSessionActivity(c, svc, jwtSecret) _ = conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { return diff --git a/internal/handler/emby_test_helpers_test.go b/internal/handler/emby_test_helpers_test.go index e707787..a1bc770 100644 --- a/internal/handler/emby_test_helpers_test.go +++ b/internal/handler/emby_test_helpers_test.go @@ -10,11 +10,17 @@ import ( ) func signedTestToken(t *testing.T, secret string) string { + t.Helper() + return signedTestTokenWithPurpose(t, secret, "") +} + +func signedTestTokenWithPurpose(t *testing.T, secret, purpose string) string { t.Helper() claims := middleware.Claims{ - UserID: "user-1", - Role: "admin", - Tier: "plus", + UserID: "user-1", + Role: "admin", + Tier: "plus", + Purpose: purpose, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)), Issuer: "mediastationgo-test", diff --git a/internal/service/bot_unbind_test.go b/internal/service/bot_unbind_test.go index aa471de..ad88a50 100644 --- a/internal/service/bot_unbind_test.go +++ b/internal/service/bot_unbind_test.go @@ -6,6 +6,8 @@ import ( "testing" "time" + "go.uber.org/zap" + "github.com/ShukeBta/MediaStationGo/internal/model" ) @@ -61,8 +63,9 @@ func TestBotAdminUnbindInactiveAndInvalidBindings(t *testing.T) { recentTime := time.Now().Add(-2 * 24 * time.Hour) admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true, LastLoginAt: &oldTime} oldUser := &model.User{Username: "old", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &oldTime} + realtimeUser := &model.User{Username: "realtime", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &oldTime} recentUser := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &recentTime} - for _, user := range []*model.User{admin, oldUser, recentUser} { + for _, user := range []*model.User{admin, oldUser, realtimeUser, recentUser} { if err := repos.User.Create(ctx, user); err != nil { t.Fatal(err) } @@ -71,6 +74,7 @@ func TestBotAdminUnbindInactiveAndInvalidBindings(t *testing.T) { {TelegramUserID: 9501, TelegramName: "@root", ChatID: 9501, UserID: admin.ID}, {TelegramUserID: 9502, TelegramName: "@old", ChatID: 9502, UserID: oldUser.ID}, {TelegramUserID: 9503, TelegramName: "@recent", ChatID: 9503, UserID: recentUser.ID}, + {TelegramUserID: 9505, TelegramName: "@realtime", ChatID: 9505, UserID: realtimeUser.ID}, {TelegramUserID: 9504, TelegramName: "@ghost", ChatID: 9504, UserID: "missing-user"}, } { row := binding @@ -80,6 +84,11 @@ func TestBotAdminUnbindInactiveAndInvalidBindings(t *testing.T) { } channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9501"}`} msg := &TelegramMessage{From: TelegramUser{ID: 9501, Username: "root"}, Chat: TelegramChat{ID: 9501, Type: "private"}} + tracker := NewSessionTrackerService(zap.NewNop()) + tracker.RecordActivity(ctx, realtimeUser.ID, realtimeUser.Username, "phone-1", "iPhone", "Infuse", "192.0.2.10") + device := NewDeviceService(zap.NewNop(), repos) + device.SetSessionTracker(tracker) + bot.SetDeviceService(device) reply, err := bot.executeCommand(ctx, channel, msg, "/unbind_inactive 30") if err != nil { @@ -97,6 +106,9 @@ func TestBotAdminUnbindInactiveAndInvalidBindings(t *testing.T) { if binding := bot.telegramBinding(ctx, 9503); binding == nil { t.Fatal("recent user binding should remain") } + if binding := bot.telegramBinding(ctx, 9505); binding == nil { + t.Fatal("realtime active user binding should remain") + } reply, err = bot.executeCommand(ctx, channel, msg, "/unbind_duplicates") if err != nil { diff --git a/internal/service/session_tracker.go b/internal/service/session_tracker.go index 0b6f77e..f42a523 100644 --- a/internal/service/session_tracker.go +++ b/internal/service/session_tracker.go @@ -70,6 +70,7 @@ type SessionTrackerService struct { mu sync.RWMutex sessions map[string]RealtimeSession + activity map[string]time.Time now func() time.Time } @@ -77,6 +78,7 @@ func NewSessionTrackerService(log *zap.Logger) *SessionTrackerService { return &SessionTrackerService{ log: log, sessions: make(map[string]RealtimeSession), + activity: make(map[string]time.Time), now: time.Now, } } @@ -130,6 +132,10 @@ func (s *SessionTrackerService) Logout(ctx context.Context, userID, deviceID, re remoteEndPoint = strings.TrimSpace(remoteEndPoint) s.mu.Lock() defer s.mu.Unlock() + if s.activity == nil { + s.activity = make(map[string]time.Time) + } + s.activity[userID] = s.now() for key, sess := range s.sessions { if sess.UserID != userID { continue @@ -208,6 +214,17 @@ func (s *SessionTrackerService) activityByUser(ctx context.Context) map[string]u sessions := s.List(ctx) now := s.now() out := make(map[string]userRealtimeActivity) + for userID, lastActivity := range s.userActivitySnapshot() { + if strings.TrimSpace(userID) == "" { + continue + } + a := out[userID] + if a.LastActivityAt == nil || lastActivity.After(*a.LastActivityAt) { + t := lastActivity + a.LastActivityAt = &t + } + out[userID] = a + } seenDevices := make(map[string]map[string]struct{}) for _, sess := range sessions { if strings.TrimSpace(sess.UserID) == "" { @@ -231,6 +248,16 @@ func (s *SessionTrackerService) activityByUser(ctx context.Context) map[string]u return out } +func (s *SessionTrackerService) userActivitySnapshot() map[string]time.Time { + s.mu.RLock() + defer s.mu.RUnlock() + out := make(map[string]time.Time, len(s.activity)) + for userID, lastActivity := range s.activity { + out[userID] = lastActivity + } + return out +} + func (s *SessionTrackerService) upsert(ctx context.Context, in realtimeSessionInput) { userID := strings.TrimSpace(in.UserID) if userID == "" { @@ -248,6 +275,10 @@ func (s *SessionTrackerService) upsert(ctx context.Context, in realtimeSessionIn s.mu.Lock() defer s.mu.Unlock() s.pruneLocked(now) + if s.activity == nil { + s.activity = make(map[string]time.Time) + } + s.activity[userID] = now existing, existed := s.sessions[key] if strings.TrimSpace(in.UserName) == "" { in.UserName = existing.UserName @@ -315,6 +346,11 @@ func (s *SessionTrackerService) pruneLocked(now time.Time) { delete(s.sessions, key) } } + for userID, lastActivity := range s.activity { + if lastActivity.Before(expiresBefore) { + delete(s.activity, userID) + } + } } func fallbackSessionDeviceID(deviceName, client, remoteEndPoint string) string { diff --git a/internal/service/session_tracker_test.go b/internal/service/session_tracker_test.go index 0eb0d42..964e653 100644 --- a/internal/service/session_tracker_test.go +++ b/internal/service/session_tracker_test.go @@ -77,6 +77,26 @@ func TestActivityRefreshKeepsPlaybackState(t *testing.T) { } } +func TestLogoutKeepsRealtimeLastActivityWithoutOnlineSession(t *testing.T) { + tracker := NewSessionTrackerService(zap.NewNop()) + now := time.Date(2026, 6, 21, 12, 30, 0, 0, time.UTC) + tracker.now = func() time.Time { return now } + old := now.Add(-8 * time.Hour) + users := []model.User{{Base: model.Base{ID: "u1"}, Username: "viewer", LastLoginAt: &old}} + + tracker.RecordActivity(t.Context(), "u1", "viewer", "dev-1", "iPhone", "Infuse", "10.0.0.8") + now = now.Add(time.Minute) + tracker.Logout(t.Context(), "u1", "dev-1", "10.0.0.8") + tracker.ApplyToUsers(t.Context(), users) + + if users[0].LastLoginAt == nil || !users[0].LastLoginAt.Equal(now) { + t.Fatalf("last_login_at = %v, want logout activity %v", users[0].LastLoginAt, now) + } + if users[0].RealtimeOnline || users[0].RealtimeDeviceCount != 0 { + t.Fatalf("logged-out user should keep last activity but no online devices, online=%v devices=%d", users[0].RealtimeOnline, users[0].RealtimeDeviceCount) + } +} + 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} @@ -107,6 +127,31 @@ func TestBotDevicesIncludesRealtimeSessionOnlyDevices(t *testing.T) { } } +func TestBotUserInfoUsesRealtimeLastLogin(t *testing.T) { + repos, bot := newBotTestService(t) + now := time.Date(2026, 6, 21, 13, 45, 0, 0, time.UTC) + old := now.Add(-6 * time.Hour) + user := model.User{Base: model.Base{ID: "u1"}, Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &old} + if err := repos.User.Create(t.Context(), &user); err != nil { + t.Fatal(err) + } + tracker := NewSessionTrackerService(zap.NewNop()) + 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.cmdMgoUserInfo(t.Context(), []string{"viewer"}) + + if !strings.Contains(reply.Text, "最后登录:2026-06-21 13:45") { + t.Fatalf("reply should use realtime last login, got %q", reply.Text) + } + if !strings.Contains(reply.Text, "设备:1") { + t.Fatalf("reply should count realtime device, got %q", reply.Text) + } +} + func TestRealtimeRecentLoginProtectsCleanupCandidate(t *testing.T) { repos := newSessionTrackerTestRepos(t) now := time.Date(2026, 6, 21, 12, 0, 0, 0, time.UTC) diff --git a/internal/service/telegram_mgo_users.go b/internal/service/telegram_mgo_users.go index eb4f685..552ed3b 100644 --- a/internal/service/telegram_mgo_users.go +++ b/internal/service/telegram_mgo_users.go @@ -43,6 +43,7 @@ func (s *TelegramBotService) cmdMgoUserInfo(ctx context.Context, args []string) if user == nil { return telegramCommandReply{Text: "未找到用户。"} } + s.applyRealtimeUserActivity(ctx, user) devices, _ := s.listUserDevices(ctx, user.ID) var historyCount int64 _ = s.repo.DB.WithContext(ctx).Model(&model.PlaybackHistory{}).Where("user_id = ?", user.ID).Count(&historyCount).Error @@ -60,6 +61,15 @@ func (s *TelegramBotService) cmdMgoUserInfo(ctx context.Context, args []string) )} } +func (s *TelegramBotService) applyRealtimeUserActivity(ctx context.Context, user *model.User) { + if s == nil || user == nil || s.device == nil || s.device.sessions == nil { + return + } + users := []model.User{*user} + s.device.sessions.ApplyToUsers(ctx, users) + *user = users[0] +} + func (s *TelegramBotService) cmdMgoDeleteUser(ctx context.Context, args []string) telegramCommandReply { if len(args) < 2 || !strings.EqualFold(args[len(args)-1], "confirm") { return telegramCommandReply{Text: "删除用户需要确认:/rmemby 用户名 confirm 或 /urm 用户名 confirm"} diff --git a/internal/service/telegram_unbind.go b/internal/service/telegram_unbind.go index 41e067c..d496eee 100644 --- a/internal/service/telegram_unbind.go +++ b/internal/service/telegram_unbind.go @@ -119,6 +119,7 @@ func (s *TelegramBotService) cmdUnbindInactive(ctx context.Context, args []strin return telegramCommandReply{Text: "读取用户失败:" + err.Error()} } cutoff := time.Now().Add(-time.Duration(days) * 24 * time.Hour) + recentWindow := time.Duration(days) * 24 * time.Hour var userIDs []string var done []string for _, user := range users { @@ -132,6 +133,9 @@ func (s *TelegramBotService) cmdUnbindInactive(ctx context.Context, args []strin if lastActive.IsZero() || lastActive.After(cutoff) { continue } + if s.device != nil && s.device.UserRecentlyActive(ctx, user.ID, recentWindow) { + continue + } var count int64 _ = s.repo.DB.WithContext(ctx).Model(&model.TelegramBinding{}).Where("user_id = ?", user.ID).Count(&count).Error if count == 0 {