mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-28 11:16:37 +08:00
fix: harden bot accounts and download handling
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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"} {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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&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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:], " "))
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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'
|
||||
|
||||
@@ -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={
|
||||
|
||||
@@ -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">
|
||||
|
||||
@@ -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>
|
||||
|
||||
Reference in New Issue
Block a user