fix: harden bot accounts and download handling

This commit is contained in:
ShukeBta
2026-06-07 18:30:27 +08:00
parent 7e37126f7c
commit df02fd1166
26 changed files with 1373 additions and 136 deletions
+40 -1
View File
@@ -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.
+47
View File
@@ -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")
}
}
+99 -8
View File
@@ -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"} {
+60
View File
@@ -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()
+21 -1
View File
@@ -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"`
}
+34
View File
@@ -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")
}
}
+3
View File
@@ -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
}
+11
View File
@@ -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)
+38
View File
@@ -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)
+61 -17
View File
@@ -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")
}
+82
View File
@@ -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")
}
}
+55 -6
View File
@@ -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 {
+134
View File
@@ -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))
}
}
+9
View File
@@ -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
}
+1
View File
@@ -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},
+14 -4
View File
@@ -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
+85 -4
View File
@@ -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(`<?xml version="1.0"?>
<rss><channel>
<item>
<title>Deleted Downloader Show S01E01 1080p</title>
<guid>deleted-downloader-episode-1</guid>
<link>magnet:?xt=urn:btih:cccccccccccccccccccccccccccccccccccccccc&amp;dn=Deleted+Downloader+Show+S01E01</link>
</item>
</channel></rss>`))
}))
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 {
+7 -5
View File
@@ -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
+52
View File
@@ -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")
+160 -38
View File
@@ -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<code>/start 用户名 密码</code>\n或:<code>/start 用户名-密码</code>"}
}
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:], " "))
+137
View File
@@ -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)
+133 -21
View File
@@ -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: "请发送新的<b>用户名</b>。"}, true
return telegramCommandReply{Text: "请发送:<code>当前密码 新用户名</code>。"}, true
case data == "act_setpass":
s.setPending(int64(msg.From.ID), "setpass")
return telegramCommandReply{Text: "请发送新的<b>密码</b>(至少 6 位)。"}, true
return telegramCommandReply{Text: "请发送:<code>当前密码 新密码</code>(新密码至少 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: "请发送:<code>/setname 新用户名</code>"}
if len(args) < 2 {
return telegramCommandReply{Text: "请发送:<code>/setname 当前密码 新用户名</code>"}
}
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: "请发送:<code>/setpass 新密码</code>"}
if len(args) < 2 {
return telegramCommandReply{Text: "请发送:<code>/setpass 当前密码 新密码</code>"}
}
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: "请发送:<code>当前密码 新用户名</code>。"}
}
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("用户名已修改为 <b>%s</b>。请用新用户名登录。", 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: "请发送:<code>当前密码 新密码</code>。"}
}
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 已绑定账号 <b>%s</b>,无需再用注册码。", 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用户名:<b>%s</b>\n密码:<b>%s</b>\n到期:<b>%s</b>\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 {
+32 -8
View File
@@ -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'
+8 -2
View File
@@ -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<string | null>(null)
const [editingUsername, setEditingUsername] = useState('')
const [resettingPasswordID, setResettingPasswordID] = useState<string | null>(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() {
<button
className="rounded-lg border border-amber-400/40 px-2 py-1 text-xs text-amber-500 hover:bg-amber-400/10"
title="重置密码"
disabled={resettingPasswordID === u.id}
onClick={() => resetPassword(u)}
>
<KeyRound size={12} />
{resettingPasswordID === u.id ? <Loader2 size={12} className="animate-spin" /> : <KeyRound size={12} />}
</button>
<button
className={
+27 -13
View File
@@ -18,6 +18,7 @@ export function DownloadClientsPage() {
const [loading, setLoading] = useState(true)
const [editing, setEditing] = useState<DownloadClient | null>(null)
const [showForm, setShowForm] = useState(false)
const [testing, setTesting] = useState<Record<string, boolean>>({})
const refresh = async () => {
setLoading(true)
@@ -33,15 +34,17 @@ export function DownloadClientsPage() {
}, [])
const onTest = async (id: string) => {
if (testing[id]) return
setTesting((current) => ({ ...current, [id]: true }))
try {
const r = await downloadClientsAPI.test(id)
if (r.ok) toast.success('连接成功')
else toast.error(r.error ?? '连接失败')
} catch (err: unknown) {
const msg =
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
'测试失败'
const msg = apiErrorMessage(err, '测试失败')
toast.error(msg)
} finally {
setTesting((current) => ({ ...current, [id]: false }))
}
}
@@ -52,9 +55,7 @@ export function DownloadClientsPage() {
toast.success('已删除')
await refresh()
} catch (err: unknown) {
const msg =
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
'删除失败'
const msg = apiErrorMessage(err, '删除失败')
toast.error(msg)
}
}
@@ -126,9 +127,15 @@ export function DownloadClientsPage() {
<div className="flex shrink-0 gap-2">
<button
onClick={() => onTest(c.id)}
disabled={testing[c.id]}
className="rounded-lg border border-gray-200 px-2 py-1 text-xs text-ink-100 hover:border-primary-400/40 hover:text-brand-500"
>
<Send size={12} className="inline" /> 测试
{testing[c.id] ? (
<Loader2 size={12} className="inline animate-spin" />
) : (
<Send size={12} className="inline" />
)}{' '}
测试
</button>
<button
onClick={() => {
@@ -157,7 +164,7 @@ export function DownloadClientsPage() {
onClose={() => setShowForm(false)}
onSaved={async () => {
setShowForm(false)
await refresh()
refresh().catch((err: unknown) => toast.error(apiErrorMessage(err, '刷新下载器列表失败')))
}}
/>
)}
@@ -187,18 +194,17 @@ function ClientFormModal({
const onSubmit = async (e: FormEvent) => {
e.preventDefault()
if (saving) return
setSaving(true)
try {
if (editing) await downloadClientsAPI.update(editing.id, form)
else await downloadClientsAPI.create(form)
toast.success('已保存')
await onSaved()
setSaving(false)
onSaved()
} catch (err: unknown) {
const msg =
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
'保存失败'
const msg = apiErrorMessage(err, '保存失败')
toast.error(msg)
} finally {
setSaving(false)
}
}
@@ -316,6 +322,14 @@ function ClientFormModal({
)
}
function apiErrorMessage(err: unknown, fallback: string): string {
const data = (err as { response?: { data?: { error?: string; message?: string } } })?.response?.data
if (data?.error) return data.error
if (data?.message) return data.message
if ((err as { code?: string })?.code === 'ECONNABORTED') return '请求超时,请检查服务或网络'
return fallback
}
function Field({ label, children }: { label: string; children: React.ReactNode }) {
return (
<label className="block">
+23 -8
View File
@@ -1,6 +1,6 @@
import { FormEvent, useState } from 'react'
import toast from 'react-hot-toast'
import { EyeOff, KeyRound, Save } from 'lucide-react'
import { EyeOff, KeyRound, Loader2, Save } from 'lucide-react'
import { authAPI } from '../api/auth'
import { profileAPI } from '../api/profile'
@@ -18,16 +18,23 @@ export function ProfilePage() {
const [hideAdult, setHideAdult] = useState(Boolean(user?.hide_adult))
const [oldPwd, setOldPwd] = useState('')
const [newPwd, setNewPwd] = useState('')
const [savingProfile, setSavingProfile] = useState(false)
const [savingPassword, setSavingPassword] = useState(false)
const onProfile = async (e: FormEvent) => {
e.preventDefault()
if (savingProfile) return
setSavingProfile(true)
try {
let password: string | undefined
const hideAdultChanged = hideAdult !== Boolean(user?.hide_adult)
if (hideAdultChanged) {
const usernameChanged = username.trim() !== (user?.username ?? '')
if (hideAdultChanged || usernameChanged) {
const input = await requestPassword({
title: hideAdult ? '隐藏成人目录' : '取消隐藏成人目录',
message: '此设置会同步影响 Web 与 Emby/Jellyfin/Infuse 等第三方客户端,请输入当前账号密码确认。',
title: usernameChanged ? '修改用户名' : hideAdult ? '隐藏成人目录' : '取消隐藏成人目录',
message: usernameChanged
? '修改用户名后需要使用新用户名登录,请输入当前账号密码确认。'
: '此设置会同步影响 Web 与 Emby/Jellyfin/Infuse 等第三方客户端,请输入当前账号密码确认。',
confirmText: '保存设置',
})
if (!input) return
@@ -50,11 +57,15 @@ export function ProfilePage() {
const msg =
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '保存失败'
toast.error(msg)
} finally {
setSavingProfile(false)
}
}
const onPwd = async (e: FormEvent) => {
e.preventDefault()
if (savingPassword) return
setSavingPassword(true)
try {
await authAPI.changePassword(oldPwd, newPwd)
toast.success('密码已更新')
@@ -65,6 +76,8 @@ export function ProfilePage() {
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
'密码更新失败'
toast.error(msg)
} finally {
setSavingPassword(false)
}
}
@@ -124,8 +137,9 @@ export function ProfilePage() {
onChange={(e) => setHideAdult(e.target.checked)}
/>
</label>
<button type="submit" className="neon-button">
<Save size={16} /> 保存
<button type="submit" disabled={savingProfile} className="neon-button">
{savingProfile ? <Loader2 size={16} className="animate-spin" /> : <Save size={16} />}
保存
</button>
</form>
@@ -152,8 +166,9 @@ export function ProfilePage() {
autoComplete="new-password"
/>
</Field>
<button type="submit" className="neon-button">
<KeyRound size={16} /> 更新密码
<button type="submit" disabled={savingPassword} className="neon-button">
{savingPassword ? <Loader2 size={16} className="animate-spin" /> : <KeyRound size={16} />}
更新密码
</button>
</form>
</div>