From df02fd1166d1b224d143d60d3cd83c08841b2f22 Mon Sep 17 00:00:00 2001 From: ShukeBta Date: Sun, 7 Jun 2026 18:30:27 +0800 Subject: [PATCH] fix: harden bot accounts and download handling --- internal/database/database.go | 41 ++++- internal/database/database_test.go | 47 +++++ internal/handler/emby.go | 107 ++++++++++- internal/handler/emby_test.go | 60 +++++++ internal/handler/profile.go | 22 ++- internal/handler/profile_test.go | 34 ++++ internal/repository/repository.go | 3 + internal/service/auth_user_limits_test.go | 11 ++ internal/service/bot_features_test.go | 38 ++++ internal/service/download_clients.go | 78 ++++++-- internal/service/download_clients_test.go | 82 +++++++++ internal/service/downloads.go | 61 ++++++- internal/service/downloads_test.go | 134 ++++++++++++++ internal/service/episode_parser.go | 9 + internal/service/episode_parser_test.go | 1 + internal/service/qbittorrent.go | 18 +- internal/service/subscription_test.go | 89 ++++++++- internal/service/telegram_api.go | 12 +- internal/service/telegram_api_test.go | 52 ++++++ internal/service/telegram_bot.go | 198 +++++++++++++++++---- internal/service/telegram_bot_user_test.go | 137 ++++++++++++++ internal/service/telegram_menu.go | 154 +++++++++++++--- web/src/api/client.ts | 40 ++++- web/src/pages/AdminPage.tsx | 10 +- web/src/pages/DownloadClientsPage.tsx | 40 +++-- web/src/pages/ProfilePage.tsx | 31 +++- 26 files changed, 1373 insertions(+), 136 deletions(-) create mode 100644 internal/database/database_test.go create mode 100644 internal/service/download_clients_test.go diff --git a/internal/database/database.go b/internal/database/database.go index 52fb04f..8915aa8 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -72,7 +72,46 @@ func buildDSN(cfg *config.Config) string { // AutoMigrate creates tables for every model registered in the model package. func AutoMigrate(db *gorm.DB) error { - return db.AutoMigrate(model.AllModels()...) + if err := db.AutoMigrate(model.AllModels()...); err != nil { + return err + } + return enforceTelegramBindingOneToOne(db) +} + +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 created_at ASC, id ASC) AS rn + FROM telegram_bindings + WHERE deleted_at IS NULL + ) + 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 + }) } // zapStdLogger adapts a *zap.Logger to GORM's tiny logger interface. diff --git a/internal/database/database_test.go b/internal/database/database_test.go new file mode 100644 index 0000000..634ee15 --- /dev/null +++ b/internal/database/database_test.go @@ -0,0 +1,47 @@ +package database + +import ( + "testing" + "time" + + "github.com/glebarez/sqlite" + "gorm.io/gorm" + + "github.com/ShukeBta/MediaStationGo/internal/model" +) + +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) + } + 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") + } +} diff --git a/internal/handler/emby.go b/internal/handler/emby.go index 44a1628..0862898 100644 --- a/internal/handler/emby.go +++ b/internal/handler/emby.go @@ -6,6 +6,8 @@ package handler import ( + "bytes" + "encoding/json" "errors" "io" "net/http" @@ -104,9 +106,11 @@ func embyPingHandler(_ *service.Container) gin.HandlerFunc { // ─── Users / Auth ──────────────────────────────────────────────────────────── type embyAuthByNameReq struct { - Username string `json:"Username"` - Pw string `json:"Pw"` - Password string `json:"Password"` + Username string `json:"Username"` + Pw string `json:"Pw"` + Password string `json:"Password"` + PasswordMd5 string `json:"PasswordMd5"` + PasswordSha1 string `json:"PasswordSha1"` } func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) { @@ -116,12 +120,10 @@ func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) { if err := c.ShouldBindJSON(&body); err != nil && !errors.Is(err, io.EOF) { return req, err } - req.Username = firstStringFromMap(body, "Username", "username", "Name", "name") - req.Pw = firstStringFromMap(body, "Pw", "pw") - req.Password = firstStringFromMap(body, "Password", "password") + fillEmbyAuthFromMap(&req, body) } - if req.Username == "" || (req.Pw == "" && req.Password == "") { + if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") { _ = c.Request.ParseForm() if req.Username == "" { req.Username = firstFormValue(c, "Username", "username", "Name", "name") @@ -132,6 +134,12 @@ func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) { if req.Password == "" { req.Password = firstFormValue(c, "Password", "password") } + if req.PasswordMd5 == "" { + req.PasswordMd5 = firstFormValue(c, "PasswordMd5", "passwordMd5", "password_md5") + } + if req.PasswordSha1 == "" { + req.PasswordSha1 = firstFormValue(c, "PasswordSha1", "passwordSha1", "password_sha1") + } } if req.Username == "" { @@ -143,9 +151,88 @@ func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) { if req.Password == "" { req.Password = firstQueryValue(c, "Password", "password") } + if req.PasswordMd5 == "" { + req.PasswordMd5 = firstQueryValue(c, "PasswordMd5", "passwordMd5", "password_md5") + } + if req.PasswordSha1 == "" { + req.PasswordSha1 = firstQueryValue(c, "PasswordSha1", "passwordSha1", "password_sha1") + } + if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") { + fillEmbyAuthFromRawBody(c, &req) + } return req, nil } +func fillEmbyAuthFromMap(req *embyAuthByNameReq, body map[string]any) { + if req.Username == "" { + req.Username = firstStringFromMap(body, "Username", "username", "UserName", "userName", "Name", "name", "LoginName", "loginName") + } + if req.Pw == "" { + req.Pw = firstStringFromMap(body, "Pw", "pw", "PW") + } + if req.Password == "" { + req.Password = firstStringFromMap(body, "Password", "password", "Pass", "pass", "Pwd", "pwd") + } + if req.PasswordMd5 == "" { + req.PasswordMd5 = firstStringFromMap(body, "PasswordMd5", "passwordMd5", "password_md5") + } + if req.PasswordSha1 == "" { + req.PasswordSha1 = firstStringFromMap(body, "PasswordSha1", "passwordSha1", "password_sha1") + } +} + +func fillEmbyAuthFromRawBody(c *gin.Context, req *embyAuthByNameReq) { + if c.Request == nil || c.Request.Body == nil { + return + } + raw, err := io.ReadAll(io.LimitReader(c.Request.Body, 1<<20)) + if err != nil { + return + } + c.Request.Body = io.NopCloser(bytes.NewReader(raw)) + raw = bytes.TrimSpace(raw) + if len(raw) == 0 { + return + } + if bytes.HasPrefix(raw, []byte("{")) { + var body map[string]any + if err := json.Unmarshal(raw, &body); err == nil { + fillEmbyAuthFromMap(req, body) + } + return + } + if values, err := url.ParseQuery(string(raw)); err == nil { + fillEmbyAuthFromValues(req, values) + } +} + +func fillEmbyAuthFromValues(req *embyAuthByNameReq, values url.Values) { + if req.Username == "" { + req.Username = firstValue(values, "Username", "username", "UserName", "userName", "Name", "name", "LoginName", "loginName") + } + if req.Pw == "" { + req.Pw = firstValue(values, "Pw", "pw", "PW") + } + if req.Password == "" { + req.Password = firstValue(values, "Password", "password", "Pass", "pass", "Pwd", "pwd") + } + if req.PasswordMd5 == "" { + req.PasswordMd5 = firstValue(values, "PasswordMd5", "passwordMd5", "password_md5") + } + if req.PasswordSha1 == "" { + req.PasswordSha1 = firstValue(values, "PasswordSha1", "passwordSha1", "password_sha1") + } +} + +func firstValue(values url.Values, keys ...string) string { + for _, key := range keys { + if value := strings.TrimSpace(values.Get(key)); value != "" { + return value + } + } + return "" +} + func firstStringFromMap(body map[string]any, keys ...string) string { if len(body) == 0 { return "" @@ -196,6 +283,10 @@ func embyAuthByNameHandler(svc *service.Container) gin.HandlerFunc { password = req.Password } if strings.TrimSpace(req.Username) == "" || password == "" { + if req.PasswordMd5 != "" || req.PasswordSha1 != "" { + embyError(c, http.StatusBadRequest, "plain password required") + return + } embyError(c, http.StatusBadRequest, "missing username or password") return } @@ -909,7 +1000,7 @@ func registerEmbyRoutes(r *gin.Engine, jwtSecret string, svc *service.Container) // 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. embyLoginLimiter := middleware.NewRateLimiter(30, 1*time.Minute) - for _, path := range []string{"/Users/AuthenticateByName", "/users/authenticatebyname"} { + for _, path := range []string{"/Users/AuthenticateByName", "/Users/authenticatebyname", "/users/AuthenticateByName", "/users/authenticatebyname"} { grp.POST(path, middleware.RateLimit(embyLoginLimiter), embyAuthByNameHandler(svc)) } for _, path := range []string{"/Users/Public", "/users/public"} { diff --git a/internal/handler/emby_test.go b/internal/handler/emby_test.go index 8683f4e..c7c2d04 100644 --- a/internal/handler/emby_test.go +++ b/internal/handler/emby_test.go @@ -1,6 +1,7 @@ package handler import ( + "context" "encoding/json" "net/http" "net/http/httptest" @@ -55,6 +56,65 @@ func TestParseEmbyAuthByNameReqAcceptsFormBody(t *testing.T) { } } +func TestParseEmbyAuthByNameReqAcceptsJSONWithoutContentType(t *testing.T) { + gin.SetMode(gin.TestMode) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/emby/users/authenticatebyname", strings.NewReader(`{"UserName":"carol","PW":"secret"}`)) + + req, err := parseEmbyAuthByNameReq(c) + if err != nil { + t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err) + } + if req.Username != "carol" || req.Pw != "secret" { + t.Fatalf("unexpected request: %#v", req) + } +} + +func TestEmbyAuthenticateByNameAcceptsCaseVariantUsernameAndPath(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{}, &model.UserPermission{}, &model.RefreshToken{}, &model.Setting{}); err != nil { + t.Fatalf("migrate: %v", err) + } + repos := repository.New(db) + cfg := &config.Config{} + cfg.Secrets.JWTSecret = "test-secret" + log := zap.NewNop() + permissions := service.NewPermissionService(log, repos) + auth := service.NewAuthService(cfg, log, repos, service.NewTokenService(cfg, log, repos), permissions) + if _, _, err := auth.Register(context.Background(), "viewer", "secret-pass"); err != nil { + t.Fatalf("register: %v", err) + } + + router := gin.New() + registerEmbyRoutes(router, cfg.Secrets.JWTSecret, &service.Container{ + Repo: repos, + Auth: auth, + Emby: service.NewEmbyService(cfg, log, repos), + Audit: service.NewAuditService(log, repos), + }) + + req := httptest.NewRequest(http.MethodPost, "/emby/users/authenticatebyname", strings.NewReader(`{"Username":"Viewer","Pw":"secret-pass"}`)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String()) + } + var payload map[string]any + if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil { + t.Fatalf("decode response: %v", err) + } + if payload["AccessToken"] == "" { + t.Fatalf("missing AccessToken: %#v", payload) + } +} + func TestEmbyWithRequestAddressUsesHost(t *testing.T) { gin.SetMode(gin.TestMode) w := httptest.NewRecorder() diff --git a/internal/handler/profile.go b/internal/handler/profile.go index 659e771..54c51e6 100644 --- a/internal/handler/profile.go +++ b/internal/handler/profile.go @@ -5,6 +5,7 @@ import ( "context" "errors" "net/http" + "strings" "github.com/gin-gonic/gin" @@ -26,7 +27,12 @@ func updateProfileHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - if hideAdultChanged { + usernameChanged, err := profileUsernameChanged(c.Request.Context(), svc, userID, patch) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + if hideAdultChanged || usernameChanged { if err := svc.Auth.VerifyPassword(c.Request.Context(), userID, patch.Password); err != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "需要输入当前账号密码确认"}) return @@ -59,6 +65,20 @@ func profileHideAdultChanged(ctx context.Context, svc *service.Container, userID return user.HideAdult != *patch.HideAdult, nil } +func profileUsernameChanged(ctx context.Context, svc *service.Container, userID string, patch service.ProfileUpdate) (bool, error) { + if patch.Username == nil { + return false, nil + } + user, err := svc.Repo.User.FindByID(ctx, userID) + if err != nil { + return false, err + } + if user == nil { + return false, errors.New("user not found") + } + return user.Username != strings.TrimSpace(*patch.Username), nil +} + type adminUpdateRoleReq struct { Role string `json:"role" binding:"required"` } diff --git a/internal/handler/profile_test.go b/internal/handler/profile_test.go index 448527a..f22bd0b 100644 --- a/internal/handler/profile_test.go +++ b/internal/handler/profile_test.go @@ -44,3 +44,37 @@ func TestProfileHideAdultRequiresPasswordOnlyWhenChanged(t *testing.T) { t.Fatal("changed hide_adult value should require password") } } + +func TestProfileUsernameChangeRequiresPasswordOnlyWhenChanged(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.User{}); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + user := &model.User{Username: "viewer", PasswordHash: "hash", Role: "user", HideAdult: true} + if err := repos.User.Create(t.Context(), user); err != nil { + t.Fatal(err) + } + svc := &service.Container{Repo: repos} + + same := " viewer " + changed, err := profileUsernameChanged(t.Context(), svc, user.ID, service.ProfileUpdate{Username: &same}) + if err != nil { + t.Fatalf("same username returned error: %v", err) + } + if changed { + t.Fatal("same username after trimming should not require password") + } + + next := "renamed" + changed, err = profileUsernameChanged(t.Context(), svc, user.ID, service.ProfileUpdate{Username: &next}) + if err != nil { + t.Fatalf("changed username returned error: %v", err) + } + if !changed { + t.Fatal("changed username should require password") + } +} diff --git a/internal/repository/repository.go b/internal/repository/repository.go index 1d3e4e4..6a266c1 100644 --- a/internal/repository/repository.go +++ b/internal/repository/repository.go @@ -112,6 +112,9 @@ func (r *UserRepository) ReleaseDeletedUsername(ctx context.Context, username st func (r *UserRepository) FindByUsername(ctx context.Context, username string) (*model.User, error) { var u model.User err := r.db.WithContext(ctx).Where("username = ?", username).First(&u).Error + if errors.Is(err, gorm.ErrRecordNotFound) && username != "" { + err = r.db.WithContext(ctx).Where("LOWER(username) = LOWER(?)", username).First(&u).Error + } if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } diff --git a/internal/service/auth_user_limits_test.go b/internal/service/auth_user_limits_test.go index 6d787e9..b931450 100644 --- a/internal/service/auth_user_limits_test.go +++ b/internal/service/auth_user_limits_test.go @@ -188,6 +188,17 @@ func TestAdminResetPasswordAllowsLoginWithNewPassword(t *testing.T) { } } +func TestLoginAcceptsUsernameCaseVariant(t *testing.T) { + ctx := context.Background() + _, auth, _, _ := newAuthTestServices(t) + if _, _, err := auth.Register(ctx, "viewer", "password"); err != nil { + t.Fatalf("register: %v", err) + } + if _, err := auth.Login(ctx, "Viewer", "password"); err != nil { + t.Fatalf("case variant username should login: %v", err) + } +} + func TestLoginKeepsOnlyConfiguredActiveRefreshTokens(t *testing.T) { ctx := context.Background() repos, auth, _, _ := newAuthTestServices(t) diff --git a/internal/service/bot_features_test.go b/internal/service/bot_features_test.go index ac17c01..651ab4e 100644 --- a/internal/service/bot_features_test.go +++ b/internal/service/bot_features_test.go @@ -467,6 +467,44 @@ func TestBotRedeemRegisterRequiresAllowedTelegramUser(t *testing.T) { } } +func TestBotRedeemRegisterCodeCreatesOnlyOneAccount(t *testing.T) { + ctx := context.Background() + repos, bot := newBotTestService(t) + code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "") + if err != nil { + t.Fatal(err) + } + channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9201,9202"}`} + + first := &TelegramMessage{From: TelegramUser{ID: 9201, Username: "first"}, Chat: TelegramChat{ID: 9201, Type: "private"}} + reply, err := bot.executeCommand(ctx, channel, first, "/redeem_register "+code.Code) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(reply.Text, "兑换成功") { + t.Fatalf("first redeem should succeed, got %q", reply.Text) + } + + second := &TelegramMessage{From: TelegramUser{ID: 9202, Username: "second"}, Chat: TelegramChat{ID: 9202, Type: "private"}} + reply, err = bot.executeCommand(ctx, channel, second, "/redeem_register "+code.Code) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(reply.Text, "兑换码已被使用") && !strings.Contains(reply.Text, "兑换码刚刚被使用") { + t.Fatalf("second redeem should be rejected as used, got %q", reply.Text) + } + var users int64 + if err := repos.DB.Model(&model.User{}).Count(&users).Error; err != nil { + t.Fatal(err) + } + if users != 1 { + t.Fatalf("one register code must create exactly one user, got %d", users) + } + if binding := bot.telegramBinding(ctx, 9202); binding != nil { + t.Fatal("second telegram user must not be bound by an already-used register code") + } +} + func TestBotAdminCodeAndUserCommands(t *testing.T) { ctx := context.Background() repos, bot := newBotTestService(t) diff --git a/internal/service/download_clients.go b/internal/service/download_clients.go index e8e3799..e065e6f 100644 --- a/internal/service/download_clients.go +++ b/internal/service/download_clients.go @@ -9,6 +9,7 @@ import ( "errors" "fmt" "net/http" + "net/url" "strings" "time" @@ -52,17 +53,22 @@ func (s *DownloadClientService) List(ctx context.Context) ([]model.DownloadClien // Create inserts a new client. func (s *DownloadClientService) Create(ctx context.Context, in DownloadClientInput) (*model.DownloadClient, error) { - if err := validateClient(in); err != nil { + normalized, err := normalizeDownloadClientInput(in) + if err != nil { return nil, err } + s.markManaged(ctx) c := &model.DownloadClient{ - Name: strings.TrimSpace(in.Name), - Type: in.Type, - Host: strings.TrimSpace(in.Host), - Username: in.Username, - Password: in.Password, - IsDefault: in.IsDefault, - Enabled: in.Enabled, + Name: normalized.Name, + Type: normalized.Type, + Host: normalized.Host, + Username: normalized.Username, + Password: normalized.Password, + IsDefault: normalized.IsDefault, + Enabled: normalized.Enabled, + } + if normalized.IsDefault { + _ = s.repo.DownloadClient.ClearDefault(ctx) } if err := s.repo.DownloadClient.Create(ctx, c); err != nil { return nil, err @@ -72,20 +78,22 @@ func (s *DownloadClientService) Create(ctx context.Context, in DownloadClientInp // Update applies a patch. func (s *DownloadClientService) Update(ctx context.Context, id string, in DownloadClientInput) (*model.DownloadClient, error) { - if err := validateClient(in); err != nil { + normalized, err := normalizeDownloadClientInput(in) + if err != nil { return nil, err } + s.markManaged(ctx) patch := map[string]any{ - "name": strings.TrimSpace(in.Name), - "type": in.Type, - "host": strings.TrimSpace(in.Host), - "username": in.Username, - "is_default": in.IsDefault, - "enabled": in.Enabled, + "name": normalized.Name, + "type": normalized.Type, + "host": normalized.Host, + "username": normalized.Username, + "is_default": normalized.IsDefault, + "enabled": normalized.Enabled, } // Only overwrite the password when the caller actually sent one. - if in.Password != "" { - patch["password"] = in.Password + if normalized.Password != "" { + patch["password"] = normalized.Password } // Fetch existing row, apply patch via Save existing, err := s.repo.DownloadClient.FindByID(ctx, id) @@ -95,6 +103,9 @@ func (s *DownloadClientService) Update(ctx context.Context, id string, in Downlo if existing == nil { return nil, errors.New("client not found") } + if normalized.IsDefault { + _ = s.repo.DownloadClient.ClearDefault(ctx) + } existing.Name = patch["name"].(string) existing.Type = patch["type"].(string) existing.Host = patch["host"].(string) @@ -112,6 +123,7 @@ func (s *DownloadClientService) Update(ctx context.Context, id string, in Downlo // Delete removes one client. func (s *DownloadClientService) Delete(ctx context.Context, id string) error { + s.markManaged(ctx) return s.repo.DownloadClient.Delete(ctx, id) } @@ -119,6 +131,9 @@ func (s *DownloadClientService) Delete(ctx context.Context, id string) error { // /api/v2/auth/login for qBittorrent, /jsonrpc for Aria2, and the // Transmission RPC URL otherwise. func (s *DownloadClientService) Test(ctx context.Context, id string) error { + ctx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + c, err := s.repo.DownloadClient.FindByID(ctx, id) if err != nil { return err @@ -189,3 +204,32 @@ func validateClient(in DownloadClientInput) error { } return nil } + +func normalizeDownloadClientInput(in DownloadClientInput) (DownloadClientInput, error) { + in.Name = strings.TrimSpace(in.Name) + in.Type = strings.TrimSpace(in.Type) + in.Host = strings.TrimSpace(in.Host) + in.Username = strings.TrimSpace(in.Username) + if err := validateClient(in); err != nil { + return in, err + } + if !strings.Contains(in.Host, "://") { + in.Host = "http://" + in.Host + } + parsed, err := url.Parse(in.Host) + if err != nil || parsed.Host == "" { + return in, errors.New("host must be a valid http(s) URL") + } + if parsed.Scheme != "http" && parsed.Scheme != "https" { + return in, errors.New("host only supports http or https") + } + in.Host = strings.TrimRight(parsed.String(), "/") + return in, nil +} + +func (s *DownloadClientService) markManaged(ctx context.Context) { + if s == nil || s.repo == nil || s.repo.Setting == nil { + return + } + _ = s.repo.Setting.Set(ctx, settingDownloadClientsManaged, "true") +} diff --git a/internal/service/download_clients_test.go b/internal/service/download_clients_test.go new file mode 100644 index 0000000..35214b1 --- /dev/null +++ b/internal/service/download_clients_test.go @@ -0,0 +1,82 @@ +package service + +import ( + "testing" + + "github.com/glebarez/sqlite" + "go.uber.org/zap" + "gorm.io/gorm" + + "github.com/ShukeBta/MediaStationGo/internal/model" + "github.com/ShukeBta/MediaStationGo/internal/repository" +) + +func TestDownloadClientCreateNormalizesHostAndClearsDefault(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + svc := NewDownloadClientService(zap.NewNop(), repos) + + first, err := svc.Create(t.Context(), DownloadClientInput{ + Name: "qB old", + Type: "qbittorrent", + Host: "http://127.0.0.1:8080/", + IsDefault: true, + Enabled: true, + }) + if err != nil { + t.Fatal(err) + } + second, err := svc.Create(t.Context(), DownloadClientInput{ + Name: "qB NAS", + Type: "qbittorrent", + Host: "172.17.0.1:8085", + IsDefault: true, + Enabled: true, + }) + if err != nil { + t.Fatal(err) + } + if second.Host != "http://172.17.0.1:8085" { + t.Fatalf("host = %q, want normalized http URL", second.Host) + } + refreshedFirst, err := repos.DownloadClient.FindByID(t.Context(), first.ID) + if err != nil { + t.Fatal(err) + } + if refreshedFirst == nil || refreshedFirst.IsDefault { + t.Fatalf("old default should be cleared, got %#v", refreshedFirst) + } + refreshedSecond, err := repos.DownloadClient.FindByID(t.Context(), second.ID) + if err != nil { + t.Fatal(err) + } + if refreshedSecond == nil || !refreshedSecond.IsDefault { + t.Fatalf("new default should be active, got %#v", refreshedSecond) + } +} + +func TestDownloadClientRejectsUnsupportedHostScheme(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil { + t.Fatal(err) + } + svc := NewDownloadClientService(zap.NewNop(), repository.New(db)) + + if _, err := svc.Create(t.Context(), DownloadClientInput{ + Name: "bad", + Type: "qbittorrent", + Host: "ftp://127.0.0.1:8080", + Enabled: true, + }); err == nil { + t.Fatal("expected unsupported scheme error") + } +} diff --git a/internal/service/downloads.go b/internal/service/downloads.go index b35c3e3..df6caa8 100644 --- a/internal/service/downloads.go +++ b/internal/service/downloads.go @@ -50,6 +50,8 @@ type DownloadService struct { var torrentEpisodeToken = regexp.MustCompile(`(?i)e\d{1,3}`) +const settingDownloadClientsManaged = "download_clients.managed" + // ErrDownloadAlreadyExists tells callers that the requested resource is already // tracked locally or present in qBittorrent. Subscriptions treat this as a // successful dedup hit, not as a retryable enqueue failure. @@ -160,6 +162,7 @@ func (d *DownloadService) Stop() { func (d *DownloadService) ReloadConfig(ctx context.Context) error { cfg := QBitConfig{} hasConfiguredClients := false + managedByDownloadClients := false // Path 1: download_clients 表 if d.repo.DownloadClient != nil { @@ -170,12 +173,16 @@ func (d *DownloadService) ReloadConfig(ctx context.Context) error { cfg.Password = c.Password } } + if d.repo.Setting != nil { + managedRaw, _ := d.repo.Setting.Get(ctx, settingDownloadClientsManaged) + managedByDownloadClients = strings.EqualFold(strings.TrimSpace(managedRaw), "true") + } // Path 2: legacy Setting 表。 // 仅在旧部署“从未使用过 download_clients 表”时回退。只要操作员曾经 // 配置过下载器,删除/禁用全部下载器就表示应停止投递,不能再偷偷用 // qbittorrent.* 旧设置继续往下载器添加任务。 - if cfg.BaseURL == "" && !hasConfiguredClients { + if cfg.BaseURL == "" && !hasConfiguredClients && !managedByDownloadClients { get := func(k string) string { v, _ := d.repo.Setting.Get(ctx, k) return v @@ -212,6 +219,10 @@ func (d *DownloadService) AddDownloadWithMeta(ctx context.Context, userID, urlSt if existing, ok := d.findExistingDownloadTask(ctx, title); ok { return existing, ErrDownloadAlreadyExists } + _ = d.ReloadConfig(ctx) + if !d.qb.IsConfigured() { + return nil, errors.New("no default downloader configured") + } if d.torrentExistsByIdentity(ctx, title) { task, err := d.createTask(ctx, userID, urlStr, savePath, meta) if err != nil { @@ -249,14 +260,22 @@ func (d *DownloadService) localMediaAlreadyExists(ctx context.Context, title str if !d.repo.DB.Migrator().HasTable(&model.Media{}) { return false } - query := availabilityQuery(title, "") - if query == "" { + queries := localAvailabilityTitleCandidates(title) + if len(queries) == 0 { return false } - like := "%" + query + "%" var rows []model.Media - if err := d.repo.DB.WithContext(ctx). - Where("title LIKE ? OR original_name LIKE ? OR path LIKE ?", like, like, like). + db := d.repo.DB.WithContext(ctx).Model(&model.Media{}) + for i, query := range queries { + like := "%" + query + "%" + clause := "title LIKE ? OR original_name LIKE ? OR path LIKE ?" + if i == 0 { + db = db.Where(clause, like, like, like) + } else { + db = db.Or(clause, like, like, like) + } + } + if err := db. Order("season_num asc, episode_num asc, created_at desc"). Limit(200). Find(&rows).Error; err != nil || len(rows) == 0 { @@ -295,6 +314,36 @@ func (d *DownloadService) localMediaAlreadyExists(ctx context.Context, title str return false } +func localAvailabilityTitleCandidates(title string) []string { + seen := map[string]struct{}{} + out := make([]string, 0, 6) + add := func(value string) { + value = strings.TrimSpace(value) + if value == "" { + return + } + if _, ok := seen[value]; ok { + return + } + seen[value] = struct{}{} + out = append(out, value) + } + add(availabilityQuery(title, "")) + if cleaned, _ := CleanQuery(title); cleaned != "" { + for _, candidate := range titleCandidates(cleaned) { + add(candidate) + fields := strings.Fields(candidate) + for i := len(fields) - 1; i >= 1; i-- { + prefix := strings.Join(fields[:i], " ") + if containsCJK(prefix) { + add(prefix) + } + } + } + } + return out +} + func (d *DownloadService) findExistingDownloadTask(ctx context.Context, title string) (*model.DownloadTask, bool) { key := downloadTaskIdentityKey(title) if key == "" || d == nil || d.repo == nil || d.repo.Download == nil { diff --git a/internal/service/downloads_test.go b/internal/service/downloads_test.go index 2cd9cae..692fd98 100644 --- a/internal/service/downloads_test.go +++ b/internal/service/downloads_test.go @@ -51,6 +51,24 @@ func TestPublicDownloadTitleUsesMagnetDisplayName(t *testing.T) { } } +func configureTestDefaultQB(t *testing.T, repos *repository.Container, baseURL string) { + t.Helper() + if err := repos.DownloadClient.Create(t.Context(), &model.DownloadClient{ + Name: "qB test", + Type: "qbittorrent", + Host: baseURL, + Username: "admin", + Password: "admin", + IsDefault: true, + Enabled: true, + }); err != nil { + t.Fatalf("create default qB client: %v", err) + } + if err := repos.Setting.Set(t.Context(), settingDownloadClientsManaged, "true"); err != nil { + t.Fatalf("mark download clients managed: %v", err) + } +} + func TestAddDownloadWithMetaSkipsExistingTaskBeforeQBAdd(t *testing.T) { var addCalls int32 qb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -295,3 +313,119 @@ func TestReloadConfigDoesNotFallbackToLegacyAfterClientDisabled(t *testing.T) { t.Fatalf("qb add calls = %d, want 0", got) } } + +func TestAddDownloadWithMetaFailsClosedWhenNoDownloaderConfigured(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + + task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:cccccccccccccccccccccccccccccccccccccccc&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{ + Title: "Movie 2026 1080p", + }) + if err == nil { + t.Fatal("expected no downloader configured error") + } + if task != nil { + t.Fatalf("task = %#v, want nil", task) + } + rows, err := repos.Download.List(t.Context()) + if err != nil { + t.Fatal(err) + } + if len(rows) != 0 { + t.Fatalf("download rows = %d, want 0", len(rows)) + } +} + +func TestReloadConfigManagedModeDoesNotFallbackToLegacyWithoutRows(t *testing.T) { + var addCalls int32 + qb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/v2/auth/login": + _, _ = w.Write([]byte("Ok.")) + case "/api/v2/torrents/info": + _, _ = w.Write([]byte(`[]`)) + case "/api/v2/torrents/add": + atomic.AddInt32(&addCalls, 1) + _, _ = w.Write([]byte("Ok.")) + default: + http.NotFound(w, r) + } + })) + defer qb.Close() + + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.DownloadClient{}, &model.DownloadTask{}, &model.Setting{}); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + if err := repos.Setting.Set(t.Context(), "qbittorrent.url", qb.URL); err != nil { + t.Fatal(err) + } + if err := repos.Setting.Set(t.Context(), "qbittorrent.username", "admin"); err != nil { + t.Fatal(err) + } + if err := repos.Setting.Set(t.Context(), "qbittorrent.password", "admin"); err != nil { + t.Fatal(err) + } + if err := repos.Setting.Set(t.Context(), settingDownloadClientsManaged, "true"); err != nil { + t.Fatal(err) + } + + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + _, err = svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:dddddddddddddddddddddddddddddddddddddddd&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{ + Title: "Movie 2026 1080p", + }) + if err == nil { + t.Fatal("expected managed mode to reject missing default downloader") + } + if got := atomic.LoadInt32(&addCalls); got != 0 { + t.Fatalf("qb add calls = %d, want 0", got) + } +} + +func TestAddDownloadWithMetaSkipsExistingLocalEpisodeWithReleaseGroup(t *testing.T) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.Media{}, &model.DownloadTask{}, &model.Setting{}, &model.DownloadClient{}); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + if err := db.Create(&model.Media{ + Title: "凡人修仙传", + Path: "/media/动漫/国漫/凡人修仙传/Season 01/凡人修仙传 - S01E146.mkv", + SeasonNum: 1, + EpisodeNum: 146, + }).Error; err != nil { + t.Fatal(err) + } + + svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + task, err := svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:eeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeee&dn=%5BMagicStar%5D+%E5%87%A1%E4%BA%BA%E4%BF%AE%E4%BB%99%E4%BC%A0+%E5%B9%B4%E7%95%AA+-+146+%5B1080p%5D", "/downloads", DownloadTaskMeta{ + Title: "[MagicStar] 凡人修仙传 年番 - 146 [1080p][WEB-DL]", + }) + if !errors.Is(err, ErrMediaAlreadyInLibrary) { + t.Fatalf("err = %v, want ErrMediaAlreadyInLibrary", err) + } + if task != nil { + t.Fatalf("task = %#v, want nil", task) + } + rows, err := repos.Download.List(t.Context()) + if err != nil { + t.Fatal(err) + } + if len(rows) != 0 { + t.Fatalf("download rows = %d, want 0", len(rows)) + } +} diff --git a/internal/service/episode_parser.go b/internal/service/episode_parser.go index d2e2e9a..2f4e662 100644 --- a/internal/service/episode_parser.go +++ b/internal/service/episode_parser.go @@ -25,6 +25,7 @@ var ( patNxE = regexp.MustCompile(`(\d{1,2})x(\d{1,3})`) patEP = regexp.MustCompile(`(?i)(?:^|[^a-z])(?:e|ep)\.?\s*(\d{1,3})(?:[^0-9]|$)`) patCN = regexp.MustCompile(`第\s*(\d{1,3})\s*[集话話期]`) + patDashEpisode = regexp.MustCompile(`[\s._-][-–—]\s*(\d{1,3})(?:\s*(?:v\d+)?)?(?:\s*[\[\(._-]|$)`) patSeasonFolder = regexp.MustCompile(`(?i)(?:^|[^a-z])(?:s|season)\.?\s*(\d{1,2})(?:[^0-9]|$)|第\s*(\d{1,2})\s*季`) // patCNSeason 匹配中文季/部标记,支持阿拉伯数字与中文数字(如「第二季」「第2部」)。 patCNSeason = regexp.MustCompile(`第\s*[0-9一二三四五六七八九十百零两]+\s*[季部]`) @@ -61,6 +62,14 @@ func ParseEpisode(path string) (season, episode int) { episode = mustAtoi(m[1]) return } + if m := patDashEpisode.FindStringSubmatch(name); len(m) >= 2 { + season = seasonFromParents(path) + if season == 0 { + season = 1 + } + episode = mustAtoi(m[1]) + return + } return 0, 0 } diff --git a/internal/service/episode_parser_test.go b/internal/service/episode_parser_test.go index d93f69f..d61de4d 100644 --- a/internal/service/episode_parser_test.go +++ b/internal/service/episode_parser_test.go @@ -13,6 +13,7 @@ func TestParseEpisode(t *testing.T) { {"Friends 10x24 - The One Where.mkv", 10, 24}, {"Some Anime - EP05 [1080p].mkv", 1, 5}, {"Some Anime - E12.mkv", 1, 12}, + {"[MagicStar] 凡人修仙传 年番 - 146 [1080p].mkv", 1, 146}, {`Some Show/Season 02/Some Show - EP03.mkv`, 2, 3}, {`Some Show/S02/Some Show - E04.mkv`, 2, 4}, {`剧集/第2季/剧集 第05集.mkv`, 2, 5}, diff --git a/internal/service/qbittorrent.go b/internal/service/qbittorrent.go index 8d07ca6..f553140 100644 --- a/internal/service/qbittorrent.go +++ b/internal/service/qbittorrent.go @@ -69,11 +69,11 @@ var ( qbitAddVerifyInterval = 800 * time.Millisecond ) -// NewQBitClient builds a fresh client, applying default URL if blank. +// NewQBitClient builds a fresh client. A blank URL intentionally stays blank: +// an unconfigured downloader must fail closed instead of silently trying a +// localhost qBittorrent instance. func NewQBitClient(log *zap.Logger, cfg QBitConfig) *QBitClient { - if cfg.BaseURL == "" { - cfg.BaseURL = "http://localhost:8080" - } + cfg.BaseURL = strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/") jar, _ := cookiejar.New(nil) return &QBitClient{ log: log, @@ -86,11 +86,18 @@ func NewQBitClient(log *zap.Logger, cfg QBitConfig) *QBitClient { func (q *QBitClient) Configure(cfg QBitConfig) { q.mu.Lock() defer q.mu.Unlock() + cfg.BaseURL = strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/") q.cfg = cfg jar, _ := cookiejar.New(nil) q.client.Jar = jar } +func (q *QBitClient) IsConfigured() bool { + q.mu.Lock() + defer q.mu.Unlock() + return strings.TrimSpace(q.cfg.BaseURL) != "" +} + // Login performs POST /api/v2/auth/login. func (q *QBitClient) Login(ctx context.Context) error { if q.cfg.BaseURL == "" { @@ -505,6 +512,9 @@ func (q *QBitClient) SetLocation(ctx context.Context, hash, location string) err // ensureAuth makes sure we have a valid SID cookie. Cheap on the happy // path; logs in transparently otherwise. func (q *QBitClient) ensureAuth(ctx context.Context) error { + if strings.TrimSpace(q.cfg.BaseURL) == "" { + return errors.New("qbittorrent base url not configured") + } u, err := url.Parse(q.cfg.BaseURL) if err != nil { return err diff --git a/internal/service/subscription_test.go b/internal/service/subscription_test.go index ea380ee..1793bd5 100644 --- a/internal/service/subscription_test.go +++ b/internal/service/subscription_test.go @@ -397,12 +397,12 @@ func TestSubscriptionRunOneDeduplicatesDuplicateRSSGUIDInSameFeed(t *testing.T) if err != nil { t.Fatal(err) } - if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}); err != nil { + if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{}); err != nil { t.Fatal(err) } repos := repository.New(db) + configureTestDefaultQB(t, repos, qb.URL) downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) - downloads.qb.Configure(QBitConfig{BaseURL: qb.URL, Username: "admin", Password: "admin"}) svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop())) sub := &model.Subscription{ @@ -477,12 +477,12 @@ func TestSubscriptionRunOneSkipsSameEpisodeAddedEarlierInFeed(t *testing.T) { if err != nil { t.Fatal(err) } - if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}); err != nil { + if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{}); err != nil { t.Fatal(err) } repos := repository.New(db) + configureTestDefaultQB(t, repos, qb.URL) downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) - downloads.qb.Configure(QBitConfig{BaseURL: qb.URL, Username: "admin", Password: "admin"}) svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop())) sub := &model.Subscription{ @@ -515,6 +515,87 @@ func TestSubscriptionRunOneSkipsSameEpisodeAddedEarlierInFeed(t *testing.T) { } } +func TestSubscriptionRunOneDoesNotUseDeletedDownloader(t *testing.T) { + rss := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/rss+xml") + _, _ = w.Write([]byte(` + + + Deleted Downloader Show S01E01 1080p + deleted-downloader-episode-1 + magnet:?xt=urn:btih:cccccccccccccccccccccccccccccccccccccccc&dn=Deleted+Downloader+Show+S01E01 + +`)) + })) + defer rss.Close() + + var qbCalls int32 + qb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + atomic.AddInt32(&qbCalls, 1) + switch r.URL.Path { + case "/api/v2/auth/login": + _, _ = w.Write([]byte("Ok.")) + case "/api/v2/torrents/info": + _, _ = w.Write([]byte(`[]`)) + case "/api/v2/torrents/add": + _, _ = w.Write([]byte("Ok.")) + default: + http.NotFound(w, r) + } + })) + defer qb.Close() + + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&model.Subscription{}, &model.Setting{}, &model.DownloadTask{}, &model.Media{}, &model.DownloadClient{}); err != nil { + t.Fatal(err) + } + repos := repository.New(db) + client := &model.DownloadClient{Name: "qB deleted", Type: "qbittorrent", Host: qb.URL, Username: "admin", Password: "admin", IsDefault: true, Enabled: true} + if err := repos.DownloadClient.Create(t.Context(), client); err != nil { + t.Fatal(err) + } + if err := repos.Setting.Set(t.Context(), settingDownloadClientsManaged, "true"); err != nil { + t.Fatal(err) + } + if err := repos.DownloadClient.Delete(t.Context(), client.ID); err != nil { + t.Fatal(err) + } + + downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil) + svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop())) + sub := &model.Subscription{ + Name: "Deleted Downloader Show 自动订阅", + FeedURL: rss.URL, + Filter: "Deleted Downloader Show", + MediaType: "tv", + SavePath: "/downloads/tv", + } + if err := repos.Subscription.Create(t.Context(), sub); err != nil { + t.Fatal(err) + } + + queued, err := svc.runOne(t.Context(), sub) + if err != nil { + t.Fatal(err) + } + if queued != 0 { + t.Fatalf("queued = %d, want 0 when default downloader was deleted", queued) + } + if got := atomic.LoadInt32(&qbCalls); got != 0 { + t.Fatalf("qB calls = %d, want 0 after downloader deletion", got) + } + rows, err := repos.Download.List(t.Context()) + if err != nil { + t.Fatal(err) + } + if len(rows) != 0 { + t.Fatalf("download rows = %d, want 0", len(rows)) + } +} + func TestMatchesSubscriptionRulesUserExcludeWords(t *testing.T) { sub := &model.Subscription{ExcludeWords: "10bit,dolby vision,杜比"} cases := []struct { diff --git a/internal/service/telegram_api.go b/internal/service/telegram_api.go index f7c8ce2..bfddde8 100644 --- a/internal/service/telegram_api.go +++ b/internal/service/telegram_api.go @@ -78,12 +78,14 @@ func telegramHTTPClients(timeout time.Duration, cfg map[string]string) []*http.C if proxyURL, err := telegramAutoProxyURL(cfg); err == nil { addProxy(proxyURL) } - for _, proxyRaw := range telegramFallbackProxyCandidates() { - proxyURL, err := normalizeProxyURL(proxyRaw, "http") - if err != nil || proxyURL == nil { - continue + if telegramAPIBaseURL(cfg) == defaultTelegramAPIBaseURL { + for _, proxyRaw := range telegramFallbackProxyCandidates() { + proxyURL, err := normalizeProxyURL(proxyRaw, "http") + if err != nil || proxyURL == nil { + continue + } + addProxy(proxyURL) } - addProxy(proxyURL) } transport := NewExternalTransport() transport.Proxy = nil diff --git a/internal/service/telegram_api_test.go b/internal/service/telegram_api_test.go index aa47179..1dcd3c8 100644 --- a/internal/service/telegram_api_test.go +++ b/internal/service/telegram_api_test.go @@ -1,11 +1,16 @@ package service import ( + "context" + "encoding/json" "errors" "net/http" + "net/http/httptest" "strings" "testing" "time" + + "github.com/ShukeBta/MediaStationGo/internal/model" ) func TestTelegramMethodURLUsesCustomAPIBase(t *testing.T) { @@ -131,6 +136,53 @@ func telegramClientProxyString(t *testing.T, client *http.Client) string { return proxyURL.String() } +func TestTelegramReplyAutoDeletesSentMessage(t *testing.T) { + requests := make(chan string, 4) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch { + case strings.HasSuffix(r.URL.Path, "/sendMessage"): + requests <- "sendMessage" + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"ok":true,"result":{"message_id":777}}`)) + case strings.HasSuffix(r.URL.Path, "/deleteMessage"): + requests <- "deleteMessage" + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"ok":true,"result":true}`)) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + cfg, _ := json.Marshal(map[string]string{ + "bot_token": "123456:ABC-def", + "api_base_url": server.URL, + "auto_delete_seconds": "0", + }) + _, bot := newBotTestService(t) + channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: string(cfg)} + if err := bot.reply(context.Background(), channel, 42, telegramCommandReply{Text: "hello"}); err != nil { + t.Fatalf("reply: %v", err) + } + waitForTelegramMethod(t, requests, "sendMessage") + waitForTelegramMethod(t, requests, "deleteMessage") +} + +func waitForTelegramMethod(t *testing.T, requests <-chan string, want string) { + t.Helper() + deadline := time.After(2 * time.Second) + for { + select { + case got := <-requests: + if got == want { + return + } + case <-deadline: + t.Fatalf("timed out waiting for telegram %s", want) + } + } +} + func TestTelegramCommandFiltering(t *testing.T) { if telegramIsCommandText("今天看什么") { t.Fatal("plain chat message should not be treated as command") diff --git a/internal/service/telegram_bot.go b/internal/service/telegram_bot.go index 384b02c..dd3ab51 100644 --- a/internal/service/telegram_bot.go +++ b/internal/service/telegram_bot.go @@ -184,6 +184,7 @@ func (s *TelegramBotService) HandleWebhook(ctx context.Context, body []byte) err s.log.Error("reply failed", zap.Error(err)) } } + s.deleteTelegramSourceMessage(channel, msg.Chat.ID, msg.MessageID) return nil } } @@ -217,6 +218,7 @@ func (s *TelegramBotService) HandleWebhook(ctx context.Context, body []byte) err if err != nil { s.log.Error("command failed", zap.Error(err)) _ = s.reply(ctx, channel, msg.Chat.ID, telegramCommandReply{Text: "命令执行失败: " + err.Error()}) + s.deleteTelegramSourceMessage(channel, msg.Chat.ID, msg.MessageID) return nil } @@ -224,6 +226,7 @@ func (s *TelegramBotService) HandleWebhook(ctx context.Context, body []byte) err if err := s.reply(ctx, channel, msg.Chat.ID, reply); err != nil { s.log.Error("reply failed", zap.Error(err)) } + s.deleteTelegramSourceMessage(channel, msg.Chat.ID, msg.MessageID) } return nil @@ -287,14 +290,23 @@ func (s *TelegramBotService) cmdStart(ctx context.Context, msg *TelegramMessage, if username == "" || password == "" { return telegramCommandReply{Text: "绑定格式不正确,请使用:\n/start 用户名 密码\n或:/start 用户名-密码"} } + existingBinding := s.telegramBinding(ctx, msg.From.ID) user, err := s.repo.User.FindByUsername(ctx, username) if err != nil || user == nil { + if existingBinding != nil { + _ = s.unbindTelegramUser(ctx, msg.From.ID) + return telegramCommandReply{Text: "当前绑定的媒体账号信息已失效,已自动解绑。请使用新的用户名和密码重新绑定。"} + } return telegramCommandReply{Text: "未找到此用户,请联系管理员注册。"} } if !user.IsActive { return telegramCommandReply{Text: "此账号已被禁用,请联系管理员。"} } if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)); err != nil { + if existingBinding != nil && existingBinding.UserID == user.ID { + _ = s.unbindTelegramUser(ctx, msg.From.ID) + return telegramCommandReply{Text: "当前绑定账号的密码已失效,已自动解绑。请使用新密码重新绑定。"} + } return telegramCommandReply{Text: "账号或密码错误。"} } if err := s.upsertTelegramBinding(ctx, msg, user.ID); err != nil { @@ -774,16 +786,18 @@ func telegramPollingRequest(ctx context.Context, clients []*http.Client, pollURL // ── Message Sending ── +const defaultTelegramMessageDeleteDelay = 120 * time.Second + +type telegramSendMessageResponse struct { + OK bool `json:"ok"` + Result struct { + MessageID int `json:"message_id"` + } `json:"result"` +} + // reply 通过 Telegram Bot API 发送回复消息。 func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyChannel, chatID int, reply telegramCommandReply) error { - cfg := map[string]string{} - if channel != nil { - configStr := channel.Config - if s.crypto != nil && configStr != "" { - configStr = s.crypto.Decrypt(configStr) - } - _ = json.Unmarshal([]byte(configStr), &cfg) - } + cfg := s.telegramChannelConfig(channel) if strings.TrimSpace(cfg["bot_token"]) == "" { return fmt.Errorf("bot_token not configured") } @@ -807,7 +821,73 @@ func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyCha } payload["reply_markup"] = map[string]interface{}{"inline_keyboard": keyboard} } - return telegramPostJSON(ctx, cfg, "sendMessage", payload, 15*time.Second) + var sent telegramSendMessageResponse + if err := telegramPostJSONDecode(ctx, cfg, "sendMessage", payload, 15*time.Second, &sent); err != nil { + return err + } + if sent.Result.MessageID > 0 { + s.scheduleTelegramMessageDelete(cfg, chatID, sent.Result.MessageID) + } + return nil +} + +func (s *TelegramBotService) deleteTelegramSourceMessage(channel *model.NotifyChannel, chatID, messageID int) { + if messageID <= 0 { + return + } + s.scheduleTelegramMessageDelete(s.telegramChannelConfig(channel), chatID, messageID) +} + +func (s *TelegramBotService) scheduleTelegramMessageDelete(cfg map[string]string, chatID, messageID int) { + if chatID == 0 || messageID <= 0 || strings.TrimSpace(cfg["bot_token"]) == "" { + return + } + delay := telegramMessageDeleteDelay(cfg) + if delay < 0 { + return + } + cfgCopy := make(map[string]string, len(cfg)) + for k, v := range cfg { + cfgCopy[k] = v + } + go func() { + if delay > 0 { + timer := time.NewTimer(delay) + defer timer.Stop() + <-timer.C + } + deleteCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + err := telegramPostJSON(deleteCtx, cfgCopy, "deleteMessage", map[string]interface{}{ + "chat_id": strconv.Itoa(chatID), + "message_id": messageID, + }, 10*time.Second) + if err != nil && s.log != nil { + s.log.Debug("telegram deleteMessage failed", + zap.Int("chat_id", chatID), + zap.Int("message_id", messageID), + zap.Error(sanitizeTelegramError(err)), + ) + } + }() +} + +func telegramMessageDeleteDelay(cfg map[string]string) time.Duration { + for _, key := range []string{"auto_delete_seconds", "message_delete_seconds", "delete_after_seconds"} { + raw := strings.TrimSpace(cfg[key]) + if raw == "" { + continue + } + seconds, err := strconv.Atoi(raw) + if err != nil { + continue + } + if seconds < 0 { + return -1 + } + return time.Duration(seconds) * time.Second + } + return defaultTelegramMessageDeleteDelay } // findChannelByChatID 根据 chat_id 查找已配置的通知渠道。 @@ -881,13 +961,17 @@ func (s *TelegramBotService) handleCallback(ctx context.Context, cb *TelegramCal if data == "adult_toggle" { reply := s.cmdHideAdult(ctx, &msg, nil) if reply.Text != "" { - return s.reply(ctx, channel, cb.Message.Chat.ID, reply) + err := s.reply(ctx, channel, cb.Message.Chat.ID, reply) + s.deleteTelegramSourceMessage(channel, cb.Message.Chat.ID, cb.Message.MessageID) + return err } return nil } if reply, handled := s.handleMenuCallback(ctx, channel, &msg, data); handled { if reply.Text != "" { - return s.reply(ctx, channel, cb.Message.Chat.ID, reply) + err := s.reply(ctx, channel, cb.Message.Chat.ID, reply) + s.deleteTelegramSourceMessage(channel, cb.Message.Chat.ID, cb.Message.MessageID) + return err } } return nil @@ -921,6 +1005,15 @@ func (s *TelegramBotService) telegramBinding(ctx context.Context, telegramUserID return &binding } +func (s *TelegramBotService) unbindTelegramUser(ctx context.Context, telegramUserID int) error { + if s == nil || s.repo == nil || s.repo.DB == nil || telegramUserID == 0 { + return nil + } + return s.repo.DB.WithContext(ctx).Unscoped(). + Where("telegram_user_id = ?", int64(telegramUserID)). + Delete(&model.TelegramBinding{}).Error +} + func (s *TelegramBotService) telegramUserIsAdmin(ctx context.Context, channel *model.NotifyChannel, telegramUserID int) bool { if s.telegramUserIDConfigured(channel, telegramUserID) { return true @@ -1072,35 +1165,47 @@ func (s *TelegramBotService) upsertTelegramBinding(ctx context.Context, msg *Tel if msg.From.Username != "" { name = "@" + strings.TrimSpace(msg.From.Username) } - var existing model.TelegramBinding - err := s.repo.DB.WithContext(ctx).Where("telegram_user_id = ?", int64(msg.From.ID)).First(&existing).Error - if err == nil { - if existing.UserID != userID { - if err := s.ensureTelegramAccountBindingAvailable(ctx, userID, int64(msg.From.ID)); err != nil { + telegramUserID := int64(msg.From.ID) + return s.repo.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var existing model.TelegramBinding + err := tx.Where("telegram_user_id = ?", telegramUserID).First(&existing).Error + if err == nil { + if existing.UserID != userID { + if err := s.ensureTelegramAccountBindingAvailableTx(ctx, tx, userID, telegramUserID); err != nil { + return err + } + } + if err := tx.Model(&existing).Updates(map[string]any{ + "telegram_name": name, + "chat_id": telegramBindingChatIDForMessage(msg, &existing), + "user_id": userID, + }).Error; telegramBindingUniqueErr(err) { + return errTelegramAccountAlreadyBound + } else if err != nil { return err } + return nil } - return s.repo.DB.WithContext(ctx).Model(&existing).Updates(map[string]any{ - "telegram_name": name, - "chat_id": telegramBindingChatIDForMessage(msg, &existing), - "user_id": userID, + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + if err := tx.Unscoped().Where("telegram_user_id = ?", telegramUserID).Delete(&model.TelegramBinding{}).Error; err != nil { + return err + } + if err := s.ensureTelegramAccountBindingAvailableTx(ctx, tx, userID, telegramUserID); err != nil { + return err + } + err = tx.Create(&model.TelegramBinding{ + TelegramUserID: telegramUserID, + TelegramName: name, + ChatID: telegramBindingChatIDForMessage(msg, nil), + UserID: userID, }).Error - } - if err != nil && err != gorm.ErrRecordNotFound { + if telegramBindingUniqueErr(err) { + return errTelegramAccountAlreadyBound + } return err - } - if err := s.repo.DB.WithContext(ctx).Unscoped().Where("telegram_user_id = ?", int64(msg.From.ID)).Delete(&model.TelegramBinding{}).Error; err != nil { - return err - } - if err := s.ensureTelegramAccountBindingAvailable(ctx, userID, int64(msg.From.ID)); err != nil { - return err - } - return s.repo.DB.WithContext(ctx).Create(&model.TelegramBinding{ - TelegramUserID: int64(msg.From.ID), - TelegramName: name, - ChatID: telegramBindingChatIDForMessage(msg, nil), - UserID: userID, - }).Error + }) } func telegramBindingChatIDForMessage(msg *TelegramMessage, existing *model.TelegramBinding) int64 { @@ -1127,8 +1232,12 @@ func telegramPrivateChatIDFromBinding(binding model.TelegramBinding) int64 { } func (s *TelegramBotService) ensureTelegramAccountBindingAvailable(ctx context.Context, userID string, telegramUserID int64) error { + return s.ensureTelegramAccountBindingAvailableTx(ctx, s.repo.DB.WithContext(ctx), userID, telegramUserID) +} + +func (s *TelegramBotService) ensureTelegramAccountBindingAvailableTx(ctx context.Context, tx *gorm.DB, userID string, telegramUserID int64) error { var bound model.TelegramBinding - err := s.repo.DB.WithContext(ctx). + err := tx.WithContext(ctx). Where("user_id = ? AND telegram_user_id <> ?", userID, telegramUserID). First(&bound).Error if errors.Is(err, gorm.ErrRecordNotFound) { @@ -1137,13 +1246,26 @@ func (s *TelegramBotService) ensureTelegramAccountBindingAvailable(ctx context.C if err != nil { return err } - if user, _ := s.repo.User.FindByID(ctx, bound.UserID); user == nil { - _ = s.repo.DB.WithContext(ctx).Unscoped().Delete(&model.TelegramBinding{}, "id = ?", bound.ID).Error + var user model.User + if err := tx.WithContext(ctx).Where("id = ?", bound.UserID).First(&user).Error; errors.Is(err, gorm.ErrRecordNotFound) { + _ = tx.WithContext(ctx).Unscoped().Delete(&model.TelegramBinding{}, "id = ?", bound.ID).Error return nil + } else if err != nil { + return err } return errTelegramAccountAlreadyBound } +func telegramBindingUniqueErr(err error) bool { + if err == nil { + return false + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "idx_telegram_bindings_user_id_active") || + strings.Contains(msg, "telegram_bindings.user_id") || + (strings.Contains(msg, "unique") && strings.Contains(msg, "telegram_bindings")) +} + func parseStartCredentials(args []string) (string, string) { if len(args) >= 2 { return strings.TrimSpace(args[0]), strings.TrimSpace(strings.Join(args[1:], " ")) diff --git a/internal/service/telegram_bot_user_test.go b/internal/service/telegram_bot_user_test.go index 5f76cf7..ca92824 100644 --- a/internal/service/telegram_bot_user_test.go +++ b/internal/service/telegram_bot_user_test.go @@ -2,6 +2,7 @@ package service import ( "encoding/json" + "errors" "strings" "testing" @@ -207,6 +208,142 @@ func TestTelegramStartRejectsAccountAlreadyBoundToAnotherTelegram(t *testing.T) } } +func TestTelegramStartUnbindsWhenBoundPasswordChanged(t *testing.T) { + ctx := t.Context() + repos, auth, _, _ := newAuthTestServices(t) + user, _, err := auth.Register(ctx, "viewer", "old-password") + if err != nil { + t.Fatalf("register: %v", err) + } + if err := repos.DB.Create(&model.TelegramBinding{ + TelegramUserID: 20003, + TelegramName: "@viewer", + ChatID: 20003, + UserID: user.ID, + }).Error; err != nil { + t.Fatalf("create binding: %v", err) + } + if err := auth.ResetPassword(ctx, user.ID, "new-password"); err != nil { + t.Fatalf("reset password: %v", err) + } + if err := repos.DB.AutoMigrate(&model.NotifyChannel{}); err != nil { + t.Fatalf("migrate notify channel: %v", err) + } + cfgJSON, _ := json.Marshal(map[string]string{"admin_user_ids": "20003"}) + if err := repos.DB.Create(&model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: string(cfgJSON)}).Error; err != nil { + t.Fatalf("create notify channel: %v", err) + } + + bot := NewTelegramBotService(zap.NewNop(), repos, nil, auth) + msg := &TelegramMessage{ + From: TelegramUser{ID: 20003, Username: "viewer", FirstName: "Viewer"}, + Chat: TelegramChat{ID: 20003, Type: "private"}, + } + reply := bot.cmdStart(ctx, msg, []string{"viewer", "old-password"}) + + if !strings.Contains(reply.Text, "已自动解绑") { + t.Fatalf("expected auto unbind reply, got %q", reply.Text) + } + if binding := bot.telegramBinding(ctx, 20003); binding != nil { + t.Fatal("stale binding should be removed after password mismatch") + } +} + +func TestTelegramSelfSetNameRequiresCurrentPassword(t *testing.T) { + ctx := t.Context() + repos, auth, _, _ := newAuthTestServices(t) + user, _, err := auth.Register(ctx, "viewer", "old-password") + if err != nil { + t.Fatalf("register: %v", err) + } + if err := repos.DB.Create(&model.TelegramBinding{ + TelegramUserID: 20004, + TelegramName: "@viewer", + ChatID: 20004, + UserID: user.ID, + }).Error; err != nil { + t.Fatalf("create binding: %v", err) + } + + bot := NewTelegramBotService(zap.NewNop(), repos, nil, auth) + msg := &TelegramMessage{From: TelegramUser{ID: 20004, Username: "viewer"}, Chat: TelegramChat{ID: 20004, Type: "private"}} + if reply := bot.selfSetName(ctx, msg, "renamed"); !strings.Contains(reply.Text, "当前密码 新用户名") { + t.Fatalf("expected usage reply, got %q", reply.Text) + } + if reply := bot.selfSetName(ctx, msg, "old-password renamed"); !strings.Contains(reply.Text, "用户名已修改") { + t.Fatalf("expected rename success, got %q", reply.Text) + } + updated, _ := repos.User.FindByID(ctx, user.ID) + if updated == nil || updated.Username != "renamed" { + t.Fatalf("username not updated: %#v", updated) + } +} + +func TestTelegramSelfSetPassWrongCurrentPasswordUnbinds(t *testing.T) { + ctx := t.Context() + repos, auth, _, _ := newAuthTestServices(t) + user, _, err := auth.Register(ctx, "viewer", "old-password") + if err != nil { + t.Fatalf("register: %v", err) + } + if err := repos.DB.Create(&model.TelegramBinding{ + TelegramUserID: 20005, + TelegramName: "@viewer", + ChatID: 20005, + UserID: user.ID, + }).Error; err != nil { + t.Fatalf("create binding: %v", err) + } + + bot := NewTelegramBotService(zap.NewNop(), repos, nil, auth) + msg := &TelegramMessage{From: TelegramUser{ID: 20005, Username: "viewer"}, Chat: TelegramChat{ID: 20005, Type: "private"}} + reply := bot.selfSetPass(ctx, msg, "wrong-password new-password") + + if !strings.Contains(reply.Text, "已自动解绑") { + t.Fatalf("expected auto unbind reply, got %q", reply.Text) + } + if binding := bot.telegramBinding(ctx, 20005); binding != nil { + t.Fatal("binding should be removed after wrong current password") + } + if _, err := auth.Login(ctx, "viewer", "old-password"); err != nil { + t.Fatalf("old password should remain valid after failed change: %v", err) + } +} + +func TestTelegramSelfSetPassChangesPasswordWithCurrentPassword(t *testing.T) { + ctx := t.Context() + repos, auth, _, _ := newAuthTestServices(t) + user, _, err := auth.Register(ctx, "viewer", "old-password") + if err != nil { + t.Fatalf("register: %v", err) + } + if err := repos.DB.Create(&model.TelegramBinding{ + TelegramUserID: 20006, + TelegramName: "@viewer", + ChatID: 20006, + UserID: user.ID, + }).Error; err != nil { + t.Fatalf("create binding: %v", err) + } + + bot := NewTelegramBotService(zap.NewNop(), repos, nil, auth) + msg := &TelegramMessage{From: TelegramUser{ID: 20006, Username: "viewer"}, Chat: TelegramChat{ID: 20006, Type: "private"}} + reply := bot.selfSetPass(ctx, msg, "old-password new-password") + + if !strings.Contains(reply.Text, "密码已修改") { + t.Fatalf("expected password change success, got %q", reply.Text) + } + if _, err := auth.Login(ctx, "viewer", "old-password"); !errors.Is(err, ErrInvalidCredentials) { + t.Fatalf("old password should fail, got %v", err) + } + if _, err := auth.Login(ctx, "viewer", "new-password"); err != nil { + t.Fatalf("new password should login: %v", err) + } + if binding := bot.telegramBinding(ctx, 20006); binding == nil { + t.Fatal("successful password change should keep telegram binding") + } +} + func TestTelegramBindingFromGroupStoresPrivateUserChatID(t *testing.T) { ctx := t.Context() repos, auth, _, _ := newAuthTestServices(t) diff --git a/internal/service/telegram_menu.go b/internal/service/telegram_menu.go index 853d9a4..6873b8b 100644 --- a/internal/service/telegram_menu.go +++ b/internal/service/telegram_menu.go @@ -3,17 +3,24 @@ package service import ( "context" "encoding/json" + "errors" "fmt" "strconv" "strings" "time" "github.com/ShukeBta/MediaStationGo/internal/model" + "gorm.io/gorm" ) // pendingTTL bounds how long a button-initiated text prompt stays valid. const pendingTTL = 5 * time.Minute +var ( + errRegistrationCodeAlreadyUsed = errors.New("registration code already used") + errRegistrationCodeExpired = errors.New("registration code expired") +) + func (s *TelegramBotService) setPending(userID int64, kind string) { s.pendingMu.Lock() s.pending[userID] = pendingInput{Kind: kind, CreatedAt: time.Now()} @@ -130,10 +137,10 @@ func (s *TelegramBotService) handleMenuCallback(ctx context.Context, channel *mo return s.replyDevices(ctx, msg), true case data == "act_setname": s.setPending(int64(msg.From.ID), "setname") - return telegramCommandReply{Text: "请发送新的用户名。"}, true + return telegramCommandReply{Text: "请发送:当前密码 新用户名。"}, true case data == "act_setpass": s.setPending(int64(msg.From.ID), "setpass") - return telegramCommandReply{Text: "请发送新的密码(至少 6 位)。"}, true + return telegramCommandReply{Text: "请发送:当前密码 新密码(新密码至少 6 位)。"}, true case strings.HasPrefix(data, "kick:"): return s.replyKick(ctx, msg, strings.TrimPrefix(data, "kick:")), true } @@ -261,15 +268,15 @@ func (s *TelegramBotService) cmdKick(ctx context.Context, msg *TelegramMessage, } func (s *TelegramBotService) cmdSetName(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply { - if len(args) == 0 { - return telegramCommandReply{Text: "请发送:/setname 新用户名"} + if len(args) < 2 { + return telegramCommandReply{Text: "请发送:/setname 当前密码 新用户名"} } return s.selfSetName(ctx, msg, strings.Join(args, " ")) } func (s *TelegramBotService) cmdSetPass(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply { - if len(args) == 0 { - return telegramCommandReply{Text: "请发送:/setpass 新密码"} + if len(args) < 2 { + return telegramCommandReply{Text: "请发送:/setpass 当前密码 新密码"} } return s.selfSetPass(ctx, msg, strings.Join(args, " ")) } @@ -373,15 +380,22 @@ func (s *TelegramBotService) replyKick(ctx context.Context, msg *TelegramMessage return s.replyDevices(ctx, msg) } -func (s *TelegramBotService) selfSetName(ctx context.Context, msg *TelegramMessage, newName string) telegramCommandReply { +func (s *TelegramBotService) selfSetName(ctx context.Context, msg *TelegramMessage, input string) telegramCommandReply { user := s.boundUser(ctx, msg.From.ID) if user == nil { return telegramCommandReply{Text: "请先绑定账号。"} } + currentPassword, newName := splitCurrentPasswordAndValue(input) + if currentPassword == "" || newName == "" { + return telegramCommandReply{Text: "请发送:当前密码 新用户名。"} + } newName = strings.TrimSpace(newName) if len(newName) < 2 || strings.ContainsAny(newName, " \t\n") { return telegramCommandReply{Text: "用户名至少 2 位且不能含空格,请重试。"} } + if reply, ok := s.verifyTelegramSelfPassword(ctx, msg, user, currentPassword); !ok { + return reply + } if existing, _ := s.repo.User.FindByUsername(ctx, newName); existing != nil && existing.ID != user.ID { return telegramCommandReply{Text: "该用户名已被占用,请换一个。"} } @@ -391,16 +405,24 @@ func (s *TelegramBotService) selfSetName(ctx context.Context, msg *TelegramMessa return telegramCommandReply{Text: fmt.Sprintf("用户名已修改为 %s。请用新用户名登录。", newName)} } -func (s *TelegramBotService) selfSetPass(ctx context.Context, msg *TelegramMessage, newPass string) telegramCommandReply { +func (s *TelegramBotService) selfSetPass(ctx context.Context, msg *TelegramMessage, input string) telegramCommandReply { user := s.boundUser(ctx, msg.From.ID) if user == nil { return telegramCommandReply{Text: "请先绑定账号。"} } + currentPassword, newPass := splitCurrentPasswordAndValue(input) + if currentPassword == "" || newPass == "" { + return telegramCommandReply{Text: "请发送:当前密码 新密码。"} + } newPass = strings.TrimSpace(newPass) if s.auth == nil { return telegramCommandReply{Text: "服务暂不可用。"} } - if err := s.auth.ResetPassword(ctx, user.ID, newPass); err != nil { + if err := s.auth.ChangePassword(ctx, user.ID, currentPassword, newPass); err != nil { + if errors.Is(err, ErrInvalidCredentials) { + _ = s.unbindTelegramUser(ctx, msg.From.ID) + return telegramCommandReply{Text: "当前密码验证失败,绑定已自动解绑。请用新密码重新绑定账号。"} + } return telegramCommandReply{Text: "修改失败:" + err.Error()} } if s.device != nil { @@ -409,6 +431,28 @@ func (s *TelegramBotService) selfSetPass(ctx context.Context, msg *TelegramMessa return telegramCommandReply{Text: "密码已修改,请用新密码重新登录第三方客户端。"} } +func splitCurrentPasswordAndValue(input string) (string, string) { + fields := strings.Fields(strings.TrimSpace(input)) + if len(fields) < 2 { + return "", "" + } + return fields[0], strings.TrimSpace(strings.Join(fields[1:], " ")) +} + +func (s *TelegramBotService) verifyTelegramSelfPassword(ctx context.Context, msg *TelegramMessage, user *model.User, currentPassword string) (telegramCommandReply, bool) { + if s.auth == nil { + return telegramCommandReply{Text: "服务暂不可用。"}, false + } + if err := s.auth.VerifyPassword(ctx, user.ID, currentPassword); err != nil { + if errors.Is(err, ErrInvalidCredentials) { + _ = s.unbindTelegramUser(ctx, msg.From.ID) + return telegramCommandReply{Text: "当前密码验证失败,绑定已自动解绑。请用新密码重新绑定账号。"}, false + } + return telegramCommandReply{Text: "验证失败:" + err.Error()}, false + } + return telegramCommandReply{}, true +} + // ── 兑换码流程 ─────────────────────────────────────────────────────────────── func (s *TelegramBotService) redeemRegisterFlow(ctx context.Context, channel *model.NotifyChannel, msg *TelegramMessage, raw string) telegramCommandReply { @@ -430,30 +474,98 @@ func (s *TelegramBotService) redeemRegisterFlow(ctx context.Context, channel *mo return telegramCommandReply{Text: fmt.Sprintf("当前 Telegram 已绑定账号 %s,无需再用注册码。", u.Username)} } } - // Generate a memorable default account from the code; users can rename via - //「改用户名/改密码」afterwards. We avoid asking for two more text turns here. - username := "u" + strings.ToLower(rc.Code[:8]) - password := randomCode(10) - user, _, err := s.auth.Register(ctx, username, password) + user, password, claimedCode, err := s.createUserFromRegistrationCode(ctx, rc.Code) if err != nil { + if errors.Is(err, errRegistrationCodeAlreadyUsed) { + return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"} + } + if errors.Is(err, errRegistrationCodeExpired) { + return telegramCommandReply{Text: "兑换码已过期。"} + } + if errors.Is(err, ErrUserLimitReached) { + return telegramCommandReply{Text: "注册失败:用户数量已达授权上限。"} + } return telegramCommandReply{Text: "注册失败:" + err.Error()} } - if err := s.repo.RegCode.MarkUsed(ctx, rc.ID, user.ID); err != nil { - // Code was raced; roll back the just-created account to avoid free signups. - _ = s.repo.User.Delete(ctx, user.ID) + if claimedCode == nil { return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"} } - if rc.DurationDays > 0 { - _ = s.applyRenewal(ctx, user.ID, rc.DurationDays) - } _ = s.upsertTelegramBinding(ctx, msg, user.ID) return telegramCommandReply{ Text: fmt.Sprintf("兑换成功并已创建账号:\n用户名:%s\n密码:%s\n到期:%s\n\n请尽快用「改用户名/改密码」修改为你自己的凭据。", - username, password, formatExpiry(s.userExpiry(ctx, user.ID))), + user.Username, password, formatExpiry(s.userExpiry(ctx, user.ID))), Buttons: [][]telegramInlineButton{{{Text: "⬅️ 返回菜单", Data: "menu_main"}}}, } } +func (s *TelegramBotService) createUserFromRegistrationCode(ctx context.Context, rawCode string) (*model.User, string, *model.RegistrationCode, error) { + code := strings.TrimSpace(rawCode) + if code == "" { + return nil, "", nil, errRegistrationCodeAlreadyUsed + } + password := randomCode(10) + var created model.User + var claimed model.RegistrationCode + err := s.repo.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where("code = ? AND kind = ? AND used_at IS NULL", code, model.RegistrationCodeRegister). + First(&claimed).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return errRegistrationCodeAlreadyUsed + } + return err + } + if claimed.IsExpired() { + return errRegistrationCodeExpired + } + var count int64 + if err := tx.Model(&model.User{}).Count(&count).Error; err != nil { + return err + } + if count >= LicensedMaxUsers(ctx, s.repo) { + return ErrUserLimitReached + } + hash, err := hashPassword(password) + if err != nil { + return err + } + codePrefix := strings.ToLower(claimed.Code) + if len(codePrefix) > 8 { + codePrefix = codePrefix[:8] + } + created = model.User{ + Username: "u" + codePrefix, + PasswordHash: hash, + Role: "user", + Tier: "free", + HideAdult: true, + ExpiredAt: renewExpiry(nil, claimed.DurationDays), + } + if err := tx.Create(&created).Error; err != nil { + return err + } + if err := tx.Create(DefaultPermissions(created.ID)).Error; err != nil { + return err + } + now := time.Now() + res := tx.Model(&model.RegistrationCode{}). + Where("id = ? AND used_at IS NULL", claimed.ID). + Updates(map[string]any{"used_by_user_id": created.ID, "used_at": &now}) + if res.Error != nil { + return res.Error + } + if res.RowsAffected == 0 { + return errRegistrationCodeAlreadyUsed + } + claimed.UsedByUserID = created.ID + claimed.UsedAt = &now + return nil + }) + if err != nil { + return nil, "", nil, err + } + return &created, password, &claimed, nil +} + func (s *TelegramBotService) redeemRenewFlow(ctx context.Context, msg *TelegramMessage, raw string) telegramCommandReply { user := s.boundUser(ctx, msg.From.ID) if user == nil { diff --git a/web/src/api/client.ts b/web/src/api/client.ts index 280a7d0..4ab14ca 100644 --- a/web/src/api/client.ts +++ b/web/src/api/client.ts @@ -12,19 +12,31 @@ export const api = axios.create({ // Flag to prevent multiple simultaneous refresh attempts let isRefreshing = false -let refreshSubscribers: Array<(token: string) => void> = [] +let refreshSubscribers: Array<{ + resolve: (token: string) => void + reject: (error: unknown) => void +}> = [] // Subscribe to token refresh -function subscribeTokenRefresh(callback: (token: string) => void) { - refreshSubscribers.push(callback) +function subscribeTokenRefresh(resolve: (token: string) => void, reject: (error: unknown) => void) { + refreshSubscribers.push({ resolve, reject }) } // Notify all subscribers about new token function onTokenRefreshed(newToken: string) { - refreshSubscribers.forEach(callback => callback(newToken)) + refreshSubscribers.forEach((subscriber) => subscriber.resolve(newToken)) refreshSubscribers = [] } +function onTokenRefreshFailed(error: unknown) { + refreshSubscribers.forEach((subscriber) => subscriber.reject(error)) + refreshSubscribers = [] +} + +function isRefreshRequest(config?: InternalAxiosRequestConfig | null): boolean { + return Boolean(config?.url?.includes('/auth/refresh')) +} + // Add auth token to requests api.interceptors.request.use((config) => { const token = useAuthStore.getState().token @@ -51,16 +63,21 @@ api.interceptors.response.use( const originalRequest = err.config as InternalAxiosRequestConfig & { _retry?: boolean } // If 401 and not already retried - if (err.response?.status === 401 && originalRequest && !originalRequest._retry) { + if ( + err.response?.status === 401 && + originalRequest && + !originalRequest._retry && + !isRefreshRequest(originalRequest) + ) { if (isRefreshing) { // Wait for token refresh to complete - return new Promise((resolve) => { + return new Promise((resolve, reject) => { subscribeTokenRefresh((token: string) => { if (originalRequest.headers) { originalRequest.headers.Authorization = `Bearer ${token}` } resolve(api(originalRequest)) - }) + }, reject) }) } @@ -80,10 +97,17 @@ api.interceptors.response.use( } } catch (refreshError) { isRefreshing = false - refreshSubscribers = [] + onTokenRefreshFailed(refreshError) + useAuthStore.getState().logout() + if (typeof window !== 'undefined' && window.location.pathname !== '/login') { + window.location.href = '/login' + } + return Promise.reject(refreshError) } // Refresh failed, logout + isRefreshing = false + onTokenRefreshFailed(err) useAuthStore.getState().logout() if (typeof window !== 'undefined' && window.location.pathname !== '/login') { window.location.href = '/login' diff --git a/web/src/pages/AdminPage.tsx b/web/src/pages/AdminPage.tsx index a784a0f..f0f5041 100644 --- a/web/src/pages/AdminPage.tsx +++ b/web/src/pages/AdminPage.tsx @@ -1,7 +1,7 @@ import { FormEvent, useEffect, useState } from 'react' import { useSearchParams } from 'react-router-dom' import toast from 'react-hot-toast' -import { KeyRound, Pencil, Plus, ShieldCheck, Trash2, UserCheck, UserX, X } from 'lucide-react' +import { KeyRound, Loader2, Pencil, Plus, ShieldCheck, Trash2, UserCheck, UserX, X } from 'lucide-react' import { adminAPI } from '../api/admin' import { libraryAPI } from '../api/library' @@ -188,6 +188,7 @@ function UsersPanel() { const [password, setPassword] = useState('') const [editingID, setEditingID] = useState(null) const [editingUsername, setEditingUsername] = useState('') + const [resettingPasswordID, setResettingPasswordID] = useState(null) const refresh = async () => { const [nextUsers, nextLicense] = await Promise.all([ adminAPI.listUsers(), @@ -243,6 +244,7 @@ function UsersPanel() { } const resetPassword = async (u: User) => { + if (resettingPasswordID) return const nextPassword = await requestPassword({ title: `重置 ${u.username} 的密码`, message: '请输入新的临时密码,至少 6 位。保存后该用户可立即使用新密码登录 Web、Bot 与第三方客户端。', @@ -253,6 +255,7 @@ function UsersPanel() { toast.error('新密码至少 6 位') return } + setResettingPasswordID(u.id) try { await adminAPI.resetUserPassword(u.id, nextPassword) toast.success('密码已重置') @@ -261,6 +264,8 @@ function UsersPanel() { (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '重置密码失败' toast.error(msg) + } finally { + setResettingPasswordID(null) } } @@ -396,9 +401,10 @@ function UsersPanel() { @@ -152,8 +166,9 @@ export function ProfilePage() { autoComplete="new-password" /> -