mirror of
https://github.com/truewhile/MeBox.git
synced 2026-10-04 20:46: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.
|
// AutoMigrate creates tables for every model registered in the model package.
|
||||||
func AutoMigrate(db *gorm.DB) error {
|
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.
|
// 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
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -104,9 +106,11 @@ func embyPingHandler(_ *service.Container) gin.HandlerFunc {
|
|||||||
// ─── Users / Auth ────────────────────────────────────────────────────────────
|
// ─── Users / Auth ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
type embyAuthByNameReq struct {
|
type embyAuthByNameReq struct {
|
||||||
Username string `json:"Username"`
|
Username string `json:"Username"`
|
||||||
Pw string `json:"Pw"`
|
Pw string `json:"Pw"`
|
||||||
Password string `json:"Password"`
|
Password string `json:"Password"`
|
||||||
|
PasswordMd5 string `json:"PasswordMd5"`
|
||||||
|
PasswordSha1 string `json:"PasswordSha1"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) {
|
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) {
|
if err := c.ShouldBindJSON(&body); err != nil && !errors.Is(err, io.EOF) {
|
||||||
return req, err
|
return req, err
|
||||||
}
|
}
|
||||||
req.Username = firstStringFromMap(body, "Username", "username", "Name", "name")
|
fillEmbyAuthFromMap(&req, body)
|
||||||
req.Pw = firstStringFromMap(body, "Pw", "pw")
|
|
||||||
req.Password = firstStringFromMap(body, "Password", "password")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Username == "" || (req.Pw == "" && req.Password == "") {
|
if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") {
|
||||||
_ = c.Request.ParseForm()
|
_ = c.Request.ParseForm()
|
||||||
if req.Username == "" {
|
if req.Username == "" {
|
||||||
req.Username = firstFormValue(c, "Username", "username", "Name", "name")
|
req.Username = firstFormValue(c, "Username", "username", "Name", "name")
|
||||||
@@ -132,6 +134,12 @@ func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) {
|
|||||||
if req.Password == "" {
|
if req.Password == "" {
|
||||||
req.Password = firstFormValue(c, "Password", "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 == "" {
|
if req.Username == "" {
|
||||||
@@ -143,9 +151,88 @@ func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) {
|
|||||||
if req.Password == "" {
|
if req.Password == "" {
|
||||||
req.Password = firstQueryValue(c, "Password", "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
|
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 {
|
func firstStringFromMap(body map[string]any, keys ...string) string {
|
||||||
if len(body) == 0 {
|
if len(body) == 0 {
|
||||||
return ""
|
return ""
|
||||||
@@ -196,6 +283,10 @@ func embyAuthByNameHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
password = req.Password
|
password = req.Password
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(req.Username) == "" || 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")
|
embyError(c, http.StatusBadRequest, "missing username or password")
|
||||||
return
|
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
|
// 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.
|
// IP, so a low limit would throttle legitimate logins into 429s.
|
||||||
embyLoginLimiter := middleware.NewRateLimiter(30, 1*time.Minute)
|
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))
|
grp.POST(path, middleware.RateLimit(embyLoginLimiter), embyAuthByNameHandler(svc))
|
||||||
}
|
}
|
||||||
for _, path := range []string{"/Users/Public", "/users/public"} {
|
for _, path := range []string{"/Users/Public", "/users/public"} {
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package handler
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"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) {
|
func TestEmbyWithRequestAddressUsesHost(t *testing.T) {
|
||||||
gin.SetMode(gin.TestMode)
|
gin.SetMode(gin.TestMode)
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"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()})
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
return
|
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 {
|
if err := svc.Auth.VerifyPassword(c.Request.Context(), userID, patch.Password); err != nil {
|
||||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "需要输入当前账号密码确认"})
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "需要输入当前账号密码确认"})
|
||||||
return
|
return
|
||||||
@@ -59,6 +65,20 @@ func profileHideAdultChanged(ctx context.Context, svc *service.Container, userID
|
|||||||
return user.HideAdult != *patch.HideAdult, nil
|
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 {
|
type adminUpdateRoleReq struct {
|
||||||
Role string `json:"role" binding:"required"`
|
Role string `json:"role" binding:"required"`
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -44,3 +44,37 @@ func TestProfileHideAdultRequiresPasswordOnlyWhenChanged(t *testing.T) {
|
|||||||
t.Fatal("changed hide_adult value should require password")
|
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) {
|
func (r *UserRepository) FindByUsername(ctx context.Context, username string) (*model.User, error) {
|
||||||
var u model.User
|
var u model.User
|
||||||
err := r.db.WithContext(ctx).Where("username = ?", username).First(&u).Error
|
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) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return nil, nil
|
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) {
|
func TestLoginKeepsOnlyConfiguredActiveRefreshTokens(t *testing.T) {
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
repos, auth, _, _ := newAuthTestServices(t)
|
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) {
|
func TestBotAdminCodeAndUserCommands(t *testing.T) {
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
repos, bot := newBotTestService(t)
|
repos, bot := newBotTestService(t)
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"net/url"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -52,17 +53,22 @@ func (s *DownloadClientService) List(ctx context.Context) ([]model.DownloadClien
|
|||||||
|
|
||||||
// Create inserts a new client.
|
// Create inserts a new client.
|
||||||
func (s *DownloadClientService) Create(ctx context.Context, in DownloadClientInput) (*model.DownloadClient, error) {
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
s.markManaged(ctx)
|
||||||
c := &model.DownloadClient{
|
c := &model.DownloadClient{
|
||||||
Name: strings.TrimSpace(in.Name),
|
Name: normalized.Name,
|
||||||
Type: in.Type,
|
Type: normalized.Type,
|
||||||
Host: strings.TrimSpace(in.Host),
|
Host: normalized.Host,
|
||||||
Username: in.Username,
|
Username: normalized.Username,
|
||||||
Password: in.Password,
|
Password: normalized.Password,
|
||||||
IsDefault: in.IsDefault,
|
IsDefault: normalized.IsDefault,
|
||||||
Enabled: in.Enabled,
|
Enabled: normalized.Enabled,
|
||||||
|
}
|
||||||
|
if normalized.IsDefault {
|
||||||
|
_ = s.repo.DownloadClient.ClearDefault(ctx)
|
||||||
}
|
}
|
||||||
if err := s.repo.DownloadClient.Create(ctx, c); err != nil {
|
if err := s.repo.DownloadClient.Create(ctx, c); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -72,20 +78,22 @@ func (s *DownloadClientService) Create(ctx context.Context, in DownloadClientInp
|
|||||||
|
|
||||||
// Update applies a patch.
|
// Update applies a patch.
|
||||||
func (s *DownloadClientService) Update(ctx context.Context, id string, in DownloadClientInput) (*model.DownloadClient, error) {
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
s.markManaged(ctx)
|
||||||
patch := map[string]any{
|
patch := map[string]any{
|
||||||
"name": strings.TrimSpace(in.Name),
|
"name": normalized.Name,
|
||||||
"type": in.Type,
|
"type": normalized.Type,
|
||||||
"host": strings.TrimSpace(in.Host),
|
"host": normalized.Host,
|
||||||
"username": in.Username,
|
"username": normalized.Username,
|
||||||
"is_default": in.IsDefault,
|
"is_default": normalized.IsDefault,
|
||||||
"enabled": in.Enabled,
|
"enabled": normalized.Enabled,
|
||||||
}
|
}
|
||||||
// Only overwrite the password when the caller actually sent one.
|
// Only overwrite the password when the caller actually sent one.
|
||||||
if in.Password != "" {
|
if normalized.Password != "" {
|
||||||
patch["password"] = in.Password
|
patch["password"] = normalized.Password
|
||||||
}
|
}
|
||||||
// Fetch existing row, apply patch via Save
|
// Fetch existing row, apply patch via Save
|
||||||
existing, err := s.repo.DownloadClient.FindByID(ctx, id)
|
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 {
|
if existing == nil {
|
||||||
return nil, errors.New("client not found")
|
return nil, errors.New("client not found")
|
||||||
}
|
}
|
||||||
|
if normalized.IsDefault {
|
||||||
|
_ = s.repo.DownloadClient.ClearDefault(ctx)
|
||||||
|
}
|
||||||
existing.Name = patch["name"].(string)
|
existing.Name = patch["name"].(string)
|
||||||
existing.Type = patch["type"].(string)
|
existing.Type = patch["type"].(string)
|
||||||
existing.Host = patch["host"].(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.
|
// Delete removes one client.
|
||||||
func (s *DownloadClientService) Delete(ctx context.Context, id string) error {
|
func (s *DownloadClientService) Delete(ctx context.Context, id string) error {
|
||||||
|
s.markManaged(ctx)
|
||||||
return s.repo.DownloadClient.Delete(ctx, id)
|
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
|
// /api/v2/auth/login for qBittorrent, /jsonrpc for Aria2, and the
|
||||||
// Transmission RPC URL otherwise.
|
// Transmission RPC URL otherwise.
|
||||||
func (s *DownloadClientService) Test(ctx context.Context, id string) error {
|
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)
|
c, err := s.repo.DownloadClient.FindByID(ctx, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -189,3 +204,32 @@ func validateClient(in DownloadClientInput) error {
|
|||||||
}
|
}
|
||||||
return nil
|
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}`)
|
var torrentEpisodeToken = regexp.MustCompile(`(?i)e\d{1,3}`)
|
||||||
|
|
||||||
|
const settingDownloadClientsManaged = "download_clients.managed"
|
||||||
|
|
||||||
// ErrDownloadAlreadyExists tells callers that the requested resource is already
|
// ErrDownloadAlreadyExists tells callers that the requested resource is already
|
||||||
// tracked locally or present in qBittorrent. Subscriptions treat this as a
|
// tracked locally or present in qBittorrent. Subscriptions treat this as a
|
||||||
// successful dedup hit, not as a retryable enqueue failure.
|
// 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 {
|
func (d *DownloadService) ReloadConfig(ctx context.Context) error {
|
||||||
cfg := QBitConfig{}
|
cfg := QBitConfig{}
|
||||||
hasConfiguredClients := false
|
hasConfiguredClients := false
|
||||||
|
managedByDownloadClients := false
|
||||||
|
|
||||||
// Path 1: download_clients 表
|
// Path 1: download_clients 表
|
||||||
if d.repo.DownloadClient != nil {
|
if d.repo.DownloadClient != nil {
|
||||||
@@ -170,12 +173,16 @@ func (d *DownloadService) ReloadConfig(ctx context.Context) error {
|
|||||||
cfg.Password = c.Password
|
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 表。
|
// Path 2: legacy Setting 表。
|
||||||
// 仅在旧部署“从未使用过 download_clients 表”时回退。只要操作员曾经
|
// 仅在旧部署“从未使用过 download_clients 表”时回退。只要操作员曾经
|
||||||
// 配置过下载器,删除/禁用全部下载器就表示应停止投递,不能再偷偷用
|
// 配置过下载器,删除/禁用全部下载器就表示应停止投递,不能再偷偷用
|
||||||
// qbittorrent.* 旧设置继续往下载器添加任务。
|
// qbittorrent.* 旧设置继续往下载器添加任务。
|
||||||
if cfg.BaseURL == "" && !hasConfiguredClients {
|
if cfg.BaseURL == "" && !hasConfiguredClients && !managedByDownloadClients {
|
||||||
get := func(k string) string {
|
get := func(k string) string {
|
||||||
v, _ := d.repo.Setting.Get(ctx, k)
|
v, _ := d.repo.Setting.Get(ctx, k)
|
||||||
return v
|
return v
|
||||||
@@ -212,6 +219,10 @@ func (d *DownloadService) AddDownloadWithMeta(ctx context.Context, userID, urlSt
|
|||||||
if existing, ok := d.findExistingDownloadTask(ctx, title); ok {
|
if existing, ok := d.findExistingDownloadTask(ctx, title); ok {
|
||||||
return existing, ErrDownloadAlreadyExists
|
return existing, ErrDownloadAlreadyExists
|
||||||
}
|
}
|
||||||
|
_ = d.ReloadConfig(ctx)
|
||||||
|
if !d.qb.IsConfigured() {
|
||||||
|
return nil, errors.New("no default downloader configured")
|
||||||
|
}
|
||||||
if d.torrentExistsByIdentity(ctx, title) {
|
if d.torrentExistsByIdentity(ctx, title) {
|
||||||
task, err := d.createTask(ctx, userID, urlStr, savePath, meta)
|
task, err := d.createTask(ctx, userID, urlStr, savePath, meta)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -249,14 +260,22 @@ func (d *DownloadService) localMediaAlreadyExists(ctx context.Context, title str
|
|||||||
if !d.repo.DB.Migrator().HasTable(&model.Media{}) {
|
if !d.repo.DB.Migrator().HasTable(&model.Media{}) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
query := availabilityQuery(title, "")
|
queries := localAvailabilityTitleCandidates(title)
|
||||||
if query == "" {
|
if len(queries) == 0 {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
like := "%" + query + "%"
|
|
||||||
var rows []model.Media
|
var rows []model.Media
|
||||||
if err := d.repo.DB.WithContext(ctx).
|
db := d.repo.DB.WithContext(ctx).Model(&model.Media{})
|
||||||
Where("title LIKE ? OR original_name LIKE ? OR path LIKE ?", like, like, like).
|
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").
|
Order("season_num asc, episode_num asc, created_at desc").
|
||||||
Limit(200).
|
Limit(200).
|
||||||
Find(&rows).Error; err != nil || len(rows) == 0 {
|
Find(&rows).Error; err != nil || len(rows) == 0 {
|
||||||
@@ -295,6 +314,36 @@ func (d *DownloadService) localMediaAlreadyExists(ctx context.Context, title str
|
|||||||
return false
|
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) {
|
func (d *DownloadService) findExistingDownloadTask(ctx context.Context, title string) (*model.DownloadTask, bool) {
|
||||||
key := downloadTaskIdentityKey(title)
|
key := downloadTaskIdentityKey(title)
|
||||||
if key == "" || d == nil || d.repo == nil || d.repo.Download == nil {
|
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) {
|
func TestAddDownloadWithMetaSkipsExistingTaskBeforeQBAdd(t *testing.T) {
|
||||||
var addCalls int32
|
var addCalls int32
|
||||||
qb := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
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)
|
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})`)
|
patNxE = regexp.MustCompile(`(\d{1,2})x(\d{1,3})`)
|
||||||
patEP = regexp.MustCompile(`(?i)(?:^|[^a-z])(?:e|ep)\.?\s*(\d{1,3})(?:[^0-9]|$)`)
|
patEP = regexp.MustCompile(`(?i)(?:^|[^a-z])(?:e|ep)\.?\s*(\d{1,3})(?:[^0-9]|$)`)
|
||||||
patCN = regexp.MustCompile(`第\s*(\d{1,3})\s*[集话話期]`)
|
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*季`)
|
patSeasonFolder = regexp.MustCompile(`(?i)(?:^|[^a-z])(?:s|season)\.?\s*(\d{1,2})(?:[^0-9]|$)|第\s*(\d{1,2})\s*季`)
|
||||||
// patCNSeason 匹配中文季/部标记,支持阿拉伯数字与中文数字(如「第二季」「第2部」)。
|
// patCNSeason 匹配中文季/部标记,支持阿拉伯数字与中文数字(如「第二季」「第2部」)。
|
||||||
patCNSeason = regexp.MustCompile(`第\s*[0-9一二三四五六七八九十百零两]+\s*[季部]`)
|
patCNSeason = regexp.MustCompile(`第\s*[0-9一二三四五六七八九十百零两]+\s*[季部]`)
|
||||||
@@ -61,6 +62,14 @@ func ParseEpisode(path string) (season, episode int) {
|
|||||||
episode = mustAtoi(m[1])
|
episode = mustAtoi(m[1])
|
||||||
return
|
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
|
return 0, 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ func TestParseEpisode(t *testing.T) {
|
|||||||
{"Friends 10x24 - The One Where.mkv", 10, 24},
|
{"Friends 10x24 - The One Where.mkv", 10, 24},
|
||||||
{"Some Anime - EP05 [1080p].mkv", 1, 5},
|
{"Some Anime - EP05 [1080p].mkv", 1, 5},
|
||||||
{"Some Anime - E12.mkv", 1, 12},
|
{"Some Anime - E12.mkv", 1, 12},
|
||||||
|
{"[MagicStar] 凡人修仙传 年番 - 146 [1080p].mkv", 1, 146},
|
||||||
{`Some Show/Season 02/Some Show - EP03.mkv`, 2, 3},
|
{`Some Show/Season 02/Some Show - EP03.mkv`, 2, 3},
|
||||||
{`Some Show/S02/Some Show - E04.mkv`, 2, 4},
|
{`Some Show/S02/Some Show - E04.mkv`, 2, 4},
|
||||||
{`剧集/第2季/剧集 第05集.mkv`, 2, 5},
|
{`剧集/第2季/剧集 第05集.mkv`, 2, 5},
|
||||||
|
|||||||
@@ -69,11 +69,11 @@ var (
|
|||||||
qbitAddVerifyInterval = 800 * time.Millisecond
|
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 {
|
func NewQBitClient(log *zap.Logger, cfg QBitConfig) *QBitClient {
|
||||||
if cfg.BaseURL == "" {
|
cfg.BaseURL = strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/")
|
||||||
cfg.BaseURL = "http://localhost:8080"
|
|
||||||
}
|
|
||||||
jar, _ := cookiejar.New(nil)
|
jar, _ := cookiejar.New(nil)
|
||||||
return &QBitClient{
|
return &QBitClient{
|
||||||
log: log,
|
log: log,
|
||||||
@@ -86,11 +86,18 @@ func NewQBitClient(log *zap.Logger, cfg QBitConfig) *QBitClient {
|
|||||||
func (q *QBitClient) Configure(cfg QBitConfig) {
|
func (q *QBitClient) Configure(cfg QBitConfig) {
|
||||||
q.mu.Lock()
|
q.mu.Lock()
|
||||||
defer q.mu.Unlock()
|
defer q.mu.Unlock()
|
||||||
|
cfg.BaseURL = strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/")
|
||||||
q.cfg = cfg
|
q.cfg = cfg
|
||||||
jar, _ := cookiejar.New(nil)
|
jar, _ := cookiejar.New(nil)
|
||||||
q.client.Jar = jar
|
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.
|
// Login performs POST /api/v2/auth/login.
|
||||||
func (q *QBitClient) Login(ctx context.Context) error {
|
func (q *QBitClient) Login(ctx context.Context) error {
|
||||||
if q.cfg.BaseURL == "" {
|
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
|
// ensureAuth makes sure we have a valid SID cookie. Cheap on the happy
|
||||||
// path; logs in transparently otherwise.
|
// path; logs in transparently otherwise.
|
||||||
func (q *QBitClient) ensureAuth(ctx context.Context) error {
|
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)
|
u, err := url.Parse(q.cfg.BaseURL)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -397,12 +397,12 @@ func TestSubscriptionRunOneDeduplicatesDuplicateRSSGUIDInSameFeed(t *testing.T)
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
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)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
repos := repository.New(db)
|
repos := repository.New(db)
|
||||||
|
configureTestDefaultQB(t, repos, qb.URL)
|
||||||
downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
|
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()))
|
svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
|
||||||
|
|
||||||
sub := &model.Subscription{
|
sub := &model.Subscription{
|
||||||
@@ -477,12 +477,12 @@ func TestSubscriptionRunOneSkipsSameEpisodeAddedEarlierInFeed(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
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)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
repos := repository.New(db)
|
repos := repository.New(db)
|
||||||
|
configureTestDefaultQB(t, repos, qb.URL)
|
||||||
downloads := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
|
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()))
|
svc := NewSubscriptionService(nil, zap.NewNop(), repos, downloads, nil, NewHub(zap.NewNop()))
|
||||||
|
|
||||||
sub := &model.Subscription{
|
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) {
|
func TestMatchesSubscriptionRulesUserExcludeWords(t *testing.T) {
|
||||||
sub := &model.Subscription{ExcludeWords: "10bit,dolby vision,杜比"}
|
sub := &model.Subscription{ExcludeWords: "10bit,dolby vision,杜比"}
|
||||||
cases := []struct {
|
cases := []struct {
|
||||||
|
|||||||
@@ -78,12 +78,14 @@ func telegramHTTPClients(timeout time.Duration, cfg map[string]string) []*http.C
|
|||||||
if proxyURL, err := telegramAutoProxyURL(cfg); err == nil {
|
if proxyURL, err := telegramAutoProxyURL(cfg); err == nil {
|
||||||
addProxy(proxyURL)
|
addProxy(proxyURL)
|
||||||
}
|
}
|
||||||
for _, proxyRaw := range telegramFallbackProxyCandidates() {
|
if telegramAPIBaseURL(cfg) == defaultTelegramAPIBaseURL {
|
||||||
proxyURL, err := normalizeProxyURL(proxyRaw, "http")
|
for _, proxyRaw := range telegramFallbackProxyCandidates() {
|
||||||
if err != nil || proxyURL == nil {
|
proxyURL, err := normalizeProxyURL(proxyRaw, "http")
|
||||||
continue
|
if err != nil || proxyURL == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
addProxy(proxyURL)
|
||||||
}
|
}
|
||||||
addProxy(proxyURL)
|
|
||||||
}
|
}
|
||||||
transport := NewExternalTransport()
|
transport := NewExternalTransport()
|
||||||
transport.Proxy = nil
|
transport.Proxy = nil
|
||||||
|
|||||||
@@ -1,11 +1,16 @@
|
|||||||
package service
|
package service
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestTelegramMethodURLUsesCustomAPIBase(t *testing.T) {
|
func TestTelegramMethodURLUsesCustomAPIBase(t *testing.T) {
|
||||||
@@ -131,6 +136,53 @@ func telegramClientProxyString(t *testing.T, client *http.Client) string {
|
|||||||
return proxyURL.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) {
|
func TestTelegramCommandFiltering(t *testing.T) {
|
||||||
if telegramIsCommandText("今天看什么") {
|
if telegramIsCommandText("今天看什么") {
|
||||||
t.Fatal("plain chat message should not be treated as command")
|
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.log.Error("reply failed", zap.Error(err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
s.deleteTelegramSourceMessage(channel, msg.Chat.ID, msg.MessageID)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -217,6 +218,7 @@ func (s *TelegramBotService) HandleWebhook(ctx context.Context, body []byte) err
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
s.log.Error("command failed", zap.Error(err))
|
s.log.Error("command failed", zap.Error(err))
|
||||||
_ = s.reply(ctx, channel, msg.Chat.ID, telegramCommandReply{Text: "命令执行失败: " + err.Error()})
|
_ = s.reply(ctx, channel, msg.Chat.ID, telegramCommandReply{Text: "命令执行失败: " + err.Error()})
|
||||||
|
s.deleteTelegramSourceMessage(channel, msg.Chat.ID, msg.MessageID)
|
||||||
return nil
|
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 {
|
if err := s.reply(ctx, channel, msg.Chat.ID, reply); err != nil {
|
||||||
s.log.Error("reply failed", zap.Error(err))
|
s.log.Error("reply failed", zap.Error(err))
|
||||||
}
|
}
|
||||||
|
s.deleteTelegramSourceMessage(channel, msg.Chat.ID, msg.MessageID)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -287,14 +290,23 @@ func (s *TelegramBotService) cmdStart(ctx context.Context, msg *TelegramMessage,
|
|||||||
if username == "" || password == "" {
|
if username == "" || password == "" {
|
||||||
return telegramCommandReply{Text: "绑定格式不正确,请使用:\n<code>/start 用户名 密码</code>\n或:<code>/start 用户名-密码</code>"}
|
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)
|
user, err := s.repo.User.FindByUsername(ctx, username)
|
||||||
if err != nil || user == nil {
|
if err != nil || user == nil {
|
||||||
|
if existingBinding != nil {
|
||||||
|
_ = s.unbindTelegramUser(ctx, msg.From.ID)
|
||||||
|
return telegramCommandReply{Text: "当前绑定的媒体账号信息已失效,已自动解绑。请使用新的用户名和密码重新绑定。"}
|
||||||
|
}
|
||||||
return telegramCommandReply{Text: "未找到此用户,请联系管理员注册。"}
|
return telegramCommandReply{Text: "未找到此用户,请联系管理员注册。"}
|
||||||
}
|
}
|
||||||
if !user.IsActive {
|
if !user.IsActive {
|
||||||
return telegramCommandReply{Text: "此账号已被禁用,请联系管理员。"}
|
return telegramCommandReply{Text: "此账号已被禁用,请联系管理员。"}
|
||||||
}
|
}
|
||||||
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)); err != nil {
|
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: "账号或密码错误。"}
|
return telegramCommandReply{Text: "账号或密码错误。"}
|
||||||
}
|
}
|
||||||
if err := s.upsertTelegramBinding(ctx, msg, user.ID); err != nil {
|
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 ──
|
// ── 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 发送回复消息。
|
// reply 通过 Telegram Bot API 发送回复消息。
|
||||||
func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyChannel, chatID int, reply telegramCommandReply) error {
|
func (s *TelegramBotService) reply(ctx context.Context, channel *model.NotifyChannel, chatID int, reply telegramCommandReply) error {
|
||||||
cfg := map[string]string{}
|
cfg := s.telegramChannelConfig(channel)
|
||||||
if channel != nil {
|
|
||||||
configStr := channel.Config
|
|
||||||
if s.crypto != nil && configStr != "" {
|
|
||||||
configStr = s.crypto.Decrypt(configStr)
|
|
||||||
}
|
|
||||||
_ = json.Unmarshal([]byte(configStr), &cfg)
|
|
||||||
}
|
|
||||||
if strings.TrimSpace(cfg["bot_token"]) == "" {
|
if strings.TrimSpace(cfg["bot_token"]) == "" {
|
||||||
return fmt.Errorf("bot_token not configured")
|
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}
|
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 查找已配置的通知渠道。
|
// findChannelByChatID 根据 chat_id 查找已配置的通知渠道。
|
||||||
@@ -881,13 +961,17 @@ func (s *TelegramBotService) handleCallback(ctx context.Context, cb *TelegramCal
|
|||||||
if data == "adult_toggle" {
|
if data == "adult_toggle" {
|
||||||
reply := s.cmdHideAdult(ctx, &msg, nil)
|
reply := s.cmdHideAdult(ctx, &msg, nil)
|
||||||
if reply.Text != "" {
|
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
|
return nil
|
||||||
}
|
}
|
||||||
if reply, handled := s.handleMenuCallback(ctx, channel, &msg, data); handled {
|
if reply, handled := s.handleMenuCallback(ctx, channel, &msg, data); handled {
|
||||||
if reply.Text != "" {
|
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
|
return nil
|
||||||
@@ -921,6 +1005,15 @@ func (s *TelegramBotService) telegramBinding(ctx context.Context, telegramUserID
|
|||||||
return &binding
|
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 {
|
func (s *TelegramBotService) telegramUserIsAdmin(ctx context.Context, channel *model.NotifyChannel, telegramUserID int) bool {
|
||||||
if s.telegramUserIDConfigured(channel, telegramUserID) {
|
if s.telegramUserIDConfigured(channel, telegramUserID) {
|
||||||
return true
|
return true
|
||||||
@@ -1072,35 +1165,47 @@ func (s *TelegramBotService) upsertTelegramBinding(ctx context.Context, msg *Tel
|
|||||||
if msg.From.Username != "" {
|
if msg.From.Username != "" {
|
||||||
name = "@" + strings.TrimSpace(msg.From.Username)
|
name = "@" + strings.TrimSpace(msg.From.Username)
|
||||||
}
|
}
|
||||||
var existing model.TelegramBinding
|
telegramUserID := int64(msg.From.ID)
|
||||||
err := s.repo.DB.WithContext(ctx).Where("telegram_user_id = ?", int64(msg.From.ID)).First(&existing).Error
|
return s.repo.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||||
if err == nil {
|
var existing model.TelegramBinding
|
||||||
if existing.UserID != userID {
|
err := tx.Where("telegram_user_id = ?", telegramUserID).First(&existing).Error
|
||||||
if err := s.ensureTelegramAccountBindingAvailable(ctx, userID, int64(msg.From.ID)); err != nil {
|
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 err
|
||||||
}
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
return s.repo.DB.WithContext(ctx).Model(&existing).Updates(map[string]any{
|
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
"telegram_name": name,
|
return err
|
||||||
"chat_id": telegramBindingChatIDForMessage(msg, &existing),
|
}
|
||||||
"user_id": userID,
|
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
|
}).Error
|
||||||
}
|
if telegramBindingUniqueErr(err) {
|
||||||
if err != nil && err != gorm.ErrRecordNotFound {
|
return errTelegramAccountAlreadyBound
|
||||||
|
}
|
||||||
return err
|
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 {
|
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 {
|
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
|
var bound model.TelegramBinding
|
||||||
err := s.repo.DB.WithContext(ctx).
|
err := tx.WithContext(ctx).
|
||||||
Where("user_id = ? AND telegram_user_id <> ?", userID, telegramUserID).
|
Where("user_id = ? AND telegram_user_id <> ?", userID, telegramUserID).
|
||||||
First(&bound).Error
|
First(&bound).Error
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
@@ -1137,13 +1246,26 @@ func (s *TelegramBotService) ensureTelegramAccountBindingAvailable(ctx context.C
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if user, _ := s.repo.User.FindByID(ctx, bound.UserID); user == nil {
|
var user model.User
|
||||||
_ = s.repo.DB.WithContext(ctx).Unscoped().Delete(&model.TelegramBinding{}, "id = ?", bound.ID).Error
|
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
|
return nil
|
||||||
|
} else if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
return errTelegramAccountAlreadyBound
|
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) {
|
func parseStartCredentials(args []string) (string, string) {
|
||||||
if len(args) >= 2 {
|
if len(args) >= 2 {
|
||||||
return strings.TrimSpace(args[0]), strings.TrimSpace(strings.Join(args[1:], " "))
|
return strings.TrimSpace(args[0]), strings.TrimSpace(strings.Join(args[1:], " "))
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package service
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"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) {
|
func TestTelegramBindingFromGroupStoresPrivateUserChatID(t *testing.T) {
|
||||||
ctx := t.Context()
|
ctx := t.Context()
|
||||||
repos, auth, _, _ := newAuthTestServices(t)
|
repos, auth, _, _ := newAuthTestServices(t)
|
||||||
|
|||||||
@@ -3,17 +3,24 @@ package service
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||||
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
// pendingTTL bounds how long a button-initiated text prompt stays valid.
|
// pendingTTL bounds how long a button-initiated text prompt stays valid.
|
||||||
const pendingTTL = 5 * time.Minute
|
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) {
|
func (s *TelegramBotService) setPending(userID int64, kind string) {
|
||||||
s.pendingMu.Lock()
|
s.pendingMu.Lock()
|
||||||
s.pending[userID] = pendingInput{Kind: kind, CreatedAt: time.Now()}
|
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
|
return s.replyDevices(ctx, msg), true
|
||||||
case data == "act_setname":
|
case data == "act_setname":
|
||||||
s.setPending(int64(msg.From.ID), "setname")
|
s.setPending(int64(msg.From.ID), "setname")
|
||||||
return telegramCommandReply{Text: "请发送新的<b>用户名</b>。"}, true
|
return telegramCommandReply{Text: "请发送:<code>当前密码 新用户名</code>。"}, true
|
||||||
case data == "act_setpass":
|
case data == "act_setpass":
|
||||||
s.setPending(int64(msg.From.ID), "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:"):
|
case strings.HasPrefix(data, "kick:"):
|
||||||
return s.replyKick(ctx, msg, strings.TrimPrefix(data, "kick:")), true
|
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 {
|
func (s *TelegramBotService) cmdSetName(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
|
||||||
if len(args) == 0 {
|
if len(args) < 2 {
|
||||||
return telegramCommandReply{Text: "请发送:<code>/setname 新用户名</code>"}
|
return telegramCommandReply{Text: "请发送:<code>/setname 当前密码 新用户名</code>"}
|
||||||
}
|
}
|
||||||
return s.selfSetName(ctx, msg, strings.Join(args, " "))
|
return s.selfSetName(ctx, msg, strings.Join(args, " "))
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *TelegramBotService) cmdSetPass(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
|
func (s *TelegramBotService) cmdSetPass(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
|
||||||
if len(args) == 0 {
|
if len(args) < 2 {
|
||||||
return telegramCommandReply{Text: "请发送:<code>/setpass 新密码</code>"}
|
return telegramCommandReply{Text: "请发送:<code>/setpass 当前密码 新密码</code>"}
|
||||||
}
|
}
|
||||||
return s.selfSetPass(ctx, msg, strings.Join(args, " "))
|
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)
|
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)
|
user := s.boundUser(ctx, msg.From.ID)
|
||||||
if user == nil {
|
if user == nil {
|
||||||
return telegramCommandReply{Text: "请先绑定账号。"}
|
return telegramCommandReply{Text: "请先绑定账号。"}
|
||||||
}
|
}
|
||||||
|
currentPassword, newName := splitCurrentPasswordAndValue(input)
|
||||||
|
if currentPassword == "" || newName == "" {
|
||||||
|
return telegramCommandReply{Text: "请发送:<code>当前密码 新用户名</code>。"}
|
||||||
|
}
|
||||||
newName = strings.TrimSpace(newName)
|
newName = strings.TrimSpace(newName)
|
||||||
if len(newName) < 2 || strings.ContainsAny(newName, " \t\n") {
|
if len(newName) < 2 || strings.ContainsAny(newName, " \t\n") {
|
||||||
return telegramCommandReply{Text: "用户名至少 2 位且不能含空格,请重试。"}
|
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 {
|
if existing, _ := s.repo.User.FindByUsername(ctx, newName); existing != nil && existing.ID != user.ID {
|
||||||
return telegramCommandReply{Text: "该用户名已被占用,请换一个。"}
|
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)}
|
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)
|
user := s.boundUser(ctx, msg.From.ID)
|
||||||
if user == nil {
|
if user == nil {
|
||||||
return telegramCommandReply{Text: "请先绑定账号。"}
|
return telegramCommandReply{Text: "请先绑定账号。"}
|
||||||
}
|
}
|
||||||
|
currentPassword, newPass := splitCurrentPasswordAndValue(input)
|
||||||
|
if currentPassword == "" || newPass == "" {
|
||||||
|
return telegramCommandReply{Text: "请发送:<code>当前密码 新密码</code>。"}
|
||||||
|
}
|
||||||
newPass = strings.TrimSpace(newPass)
|
newPass = strings.TrimSpace(newPass)
|
||||||
if s.auth == nil {
|
if s.auth == nil {
|
||||||
return telegramCommandReply{Text: "服务暂不可用。"}
|
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()}
|
return telegramCommandReply{Text: "修改失败:" + err.Error()}
|
||||||
}
|
}
|
||||||
if s.device != nil {
|
if s.device != nil {
|
||||||
@@ -409,6 +431,28 @@ func (s *TelegramBotService) selfSetPass(ctx context.Context, msg *TelegramMessa
|
|||||||
return telegramCommandReply{Text: "密码已修改,请用新密码重新登录第三方客户端。"}
|
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 {
|
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)}
|
return telegramCommandReply{Text: fmt.Sprintf("当前 Telegram 已绑定账号 <b>%s</b>,无需再用注册码。", u.Username)}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// Generate a memorable default account from the code; users can rename via
|
user, password, claimedCode, err := s.createUserFromRegistrationCode(ctx, rc.Code)
|
||||||
//「改用户名/改密码」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)
|
|
||||||
if err != nil {
|
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()}
|
return telegramCommandReply{Text: "注册失败:" + err.Error()}
|
||||||
}
|
}
|
||||||
if err := s.repo.RegCode.MarkUsed(ctx, rc.ID, user.ID); err != nil {
|
if claimedCode == nil {
|
||||||
// Code was raced; roll back the just-created account to avoid free signups.
|
|
||||||
_ = s.repo.User.Delete(ctx, user.ID)
|
|
||||||
return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"}
|
return telegramCommandReply{Text: "兑换码刚刚被使用,请换一个。"}
|
||||||
}
|
}
|
||||||
if rc.DurationDays > 0 {
|
|
||||||
_ = s.applyRenewal(ctx, user.ID, rc.DurationDays)
|
|
||||||
}
|
|
||||||
_ = s.upsertTelegramBinding(ctx, msg, user.ID)
|
_ = s.upsertTelegramBinding(ctx, msg, user.ID)
|
||||||
return telegramCommandReply{
|
return telegramCommandReply{
|
||||||
Text: fmt.Sprintf("兑换成功并已创建账号:\n用户名:<b>%s</b>\n密码:<b>%s</b>\n到期:<b>%s</b>\n\n请尽快用「改用户名/改密码」修改为你自己的凭据。",
|
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"}}},
|
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 {
|
func (s *TelegramBotService) redeemRenewFlow(ctx context.Context, msg *TelegramMessage, raw string) telegramCommandReply {
|
||||||
user := s.boundUser(ctx, msg.From.ID)
|
user := s.boundUser(ctx, msg.From.ID)
|
||||||
if user == nil {
|
if user == nil {
|
||||||
|
|||||||
+32
-8
@@ -12,19 +12,31 @@ export const api = axios.create({
|
|||||||
|
|
||||||
// Flag to prevent multiple simultaneous refresh attempts
|
// Flag to prevent multiple simultaneous refresh attempts
|
||||||
let isRefreshing = false
|
let isRefreshing = false
|
||||||
let refreshSubscribers: Array<(token: string) => void> = []
|
let refreshSubscribers: Array<{
|
||||||
|
resolve: (token: string) => void
|
||||||
|
reject: (error: unknown) => void
|
||||||
|
}> = []
|
||||||
|
|
||||||
// Subscribe to token refresh
|
// Subscribe to token refresh
|
||||||
function subscribeTokenRefresh(callback: (token: string) => void) {
|
function subscribeTokenRefresh(resolve: (token: string) => void, reject: (error: unknown) => void) {
|
||||||
refreshSubscribers.push(callback)
|
refreshSubscribers.push({ resolve, reject })
|
||||||
}
|
}
|
||||||
|
|
||||||
// Notify all subscribers about new token
|
// Notify all subscribers about new token
|
||||||
function onTokenRefreshed(newToken: string) {
|
function onTokenRefreshed(newToken: string) {
|
||||||
refreshSubscribers.forEach(callback => callback(newToken))
|
refreshSubscribers.forEach((subscriber) => subscriber.resolve(newToken))
|
||||||
refreshSubscribers = []
|
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
|
// Add auth token to requests
|
||||||
api.interceptors.request.use((config) => {
|
api.interceptors.request.use((config) => {
|
||||||
const token = useAuthStore.getState().token
|
const token = useAuthStore.getState().token
|
||||||
@@ -51,16 +63,21 @@ api.interceptors.response.use(
|
|||||||
const originalRequest = err.config as InternalAxiosRequestConfig & { _retry?: boolean }
|
const originalRequest = err.config as InternalAxiosRequestConfig & { _retry?: boolean }
|
||||||
|
|
||||||
// If 401 and not already retried
|
// 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) {
|
if (isRefreshing) {
|
||||||
// Wait for token refresh to complete
|
// Wait for token refresh to complete
|
||||||
return new Promise((resolve) => {
|
return new Promise((resolve, reject) => {
|
||||||
subscribeTokenRefresh((token: string) => {
|
subscribeTokenRefresh((token: string) => {
|
||||||
if (originalRequest.headers) {
|
if (originalRequest.headers) {
|
||||||
originalRequest.headers.Authorization = `Bearer ${token}`
|
originalRequest.headers.Authorization = `Bearer ${token}`
|
||||||
}
|
}
|
||||||
resolve(api(originalRequest))
|
resolve(api(originalRequest))
|
||||||
})
|
}, reject)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -80,10 +97,17 @@ api.interceptors.response.use(
|
|||||||
}
|
}
|
||||||
} catch (refreshError) {
|
} catch (refreshError) {
|
||||||
isRefreshing = false
|
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
|
// Refresh failed, logout
|
||||||
|
isRefreshing = false
|
||||||
|
onTokenRefreshFailed(err)
|
||||||
useAuthStore.getState().logout()
|
useAuthStore.getState().logout()
|
||||||
if (typeof window !== 'undefined' && window.location.pathname !== '/login') {
|
if (typeof window !== 'undefined' && window.location.pathname !== '/login') {
|
||||||
window.location.href = '/login'
|
window.location.href = '/login'
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import { FormEvent, useEffect, useState } from 'react'
|
import { FormEvent, useEffect, useState } from 'react'
|
||||||
import { useSearchParams } from 'react-router-dom'
|
import { useSearchParams } from 'react-router-dom'
|
||||||
import toast from 'react-hot-toast'
|
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 { adminAPI } from '../api/admin'
|
||||||
import { libraryAPI } from '../api/library'
|
import { libraryAPI } from '../api/library'
|
||||||
@@ -188,6 +188,7 @@ function UsersPanel() {
|
|||||||
const [password, setPassword] = useState('')
|
const [password, setPassword] = useState('')
|
||||||
const [editingID, setEditingID] = useState<string | null>(null)
|
const [editingID, setEditingID] = useState<string | null>(null)
|
||||||
const [editingUsername, setEditingUsername] = useState('')
|
const [editingUsername, setEditingUsername] = useState('')
|
||||||
|
const [resettingPasswordID, setResettingPasswordID] = useState<string | null>(null)
|
||||||
const refresh = async () => {
|
const refresh = async () => {
|
||||||
const [nextUsers, nextLicense] = await Promise.all([
|
const [nextUsers, nextLicense] = await Promise.all([
|
||||||
adminAPI.listUsers(),
|
adminAPI.listUsers(),
|
||||||
@@ -243,6 +244,7 @@ function UsersPanel() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
const resetPassword = async (u: User) => {
|
const resetPassword = async (u: User) => {
|
||||||
|
if (resettingPasswordID) return
|
||||||
const nextPassword = await requestPassword({
|
const nextPassword = await requestPassword({
|
||||||
title: `重置 ${u.username} 的密码`,
|
title: `重置 ${u.username} 的密码`,
|
||||||
message: '请输入新的临时密码,至少 6 位。保存后该用户可立即使用新密码登录 Web、Bot 与第三方客户端。',
|
message: '请输入新的临时密码,至少 6 位。保存后该用户可立即使用新密码登录 Web、Bot 与第三方客户端。',
|
||||||
@@ -253,6 +255,7 @@ function UsersPanel() {
|
|||||||
toast.error('新密码至少 6 位')
|
toast.error('新密码至少 6 位')
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
setResettingPasswordID(u.id)
|
||||||
try {
|
try {
|
||||||
await adminAPI.resetUserPassword(u.id, nextPassword)
|
await adminAPI.resetUserPassword(u.id, nextPassword)
|
||||||
toast.success('密码已重置')
|
toast.success('密码已重置')
|
||||||
@@ -261,6 +264,8 @@ function UsersPanel() {
|
|||||||
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
|
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
|
||||||
'重置密码失败'
|
'重置密码失败'
|
||||||
toast.error(msg)
|
toast.error(msg)
|
||||||
|
} finally {
|
||||||
|
setResettingPasswordID(null)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -396,9 +401,10 @@ function UsersPanel() {
|
|||||||
<button
|
<button
|
||||||
className="rounded-lg border border-amber-400/40 px-2 py-1 text-xs text-amber-500 hover:bg-amber-400/10"
|
className="rounded-lg border border-amber-400/40 px-2 py-1 text-xs text-amber-500 hover:bg-amber-400/10"
|
||||||
title="重置密码"
|
title="重置密码"
|
||||||
|
disabled={resettingPasswordID === u.id}
|
||||||
onClick={() => resetPassword(u)}
|
onClick={() => resetPassword(u)}
|
||||||
>
|
>
|
||||||
<KeyRound size={12} />
|
{resettingPasswordID === u.id ? <Loader2 size={12} className="animate-spin" /> : <KeyRound size={12} />}
|
||||||
</button>
|
</button>
|
||||||
<button
|
<button
|
||||||
className={
|
className={
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ export function DownloadClientsPage() {
|
|||||||
const [loading, setLoading] = useState(true)
|
const [loading, setLoading] = useState(true)
|
||||||
const [editing, setEditing] = useState<DownloadClient | null>(null)
|
const [editing, setEditing] = useState<DownloadClient | null>(null)
|
||||||
const [showForm, setShowForm] = useState(false)
|
const [showForm, setShowForm] = useState(false)
|
||||||
|
const [testing, setTesting] = useState<Record<string, boolean>>({})
|
||||||
|
|
||||||
const refresh = async () => {
|
const refresh = async () => {
|
||||||
setLoading(true)
|
setLoading(true)
|
||||||
@@ -33,15 +34,17 @@ export function DownloadClientsPage() {
|
|||||||
}, [])
|
}, [])
|
||||||
|
|
||||||
const onTest = async (id: string) => {
|
const onTest = async (id: string) => {
|
||||||
|
if (testing[id]) return
|
||||||
|
setTesting((current) => ({ ...current, [id]: true }))
|
||||||
try {
|
try {
|
||||||
const r = await downloadClientsAPI.test(id)
|
const r = await downloadClientsAPI.test(id)
|
||||||
if (r.ok) toast.success('连接成功')
|
if (r.ok) toast.success('连接成功')
|
||||||
else toast.error(r.error ?? '连接失败')
|
else toast.error(r.error ?? '连接失败')
|
||||||
} catch (err: unknown) {
|
} catch (err: unknown) {
|
||||||
const msg =
|
const msg = apiErrorMessage(err, '测试失败')
|
||||||
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
|
|
||||||
'测试失败'
|
|
||||||
toast.error(msg)
|
toast.error(msg)
|
||||||
|
} finally {
|
||||||
|
setTesting((current) => ({ ...current, [id]: false }))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -52,9 +55,7 @@ export function DownloadClientsPage() {
|
|||||||
toast.success('已删除')
|
toast.success('已删除')
|
||||||
await refresh()
|
await refresh()
|
||||||
} catch (err: unknown) {
|
} catch (err: unknown) {
|
||||||
const msg =
|
const msg = apiErrorMessage(err, '删除失败')
|
||||||
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
|
|
||||||
'删除失败'
|
|
||||||
toast.error(msg)
|
toast.error(msg)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -126,9 +127,15 @@ export function DownloadClientsPage() {
|
|||||||
<div className="flex shrink-0 gap-2">
|
<div className="flex shrink-0 gap-2">
|
||||||
<button
|
<button
|
||||||
onClick={() => onTest(c.id)}
|
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"
|
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>
|
||||||
<button
|
<button
|
||||||
onClick={() => {
|
onClick={() => {
|
||||||
@@ -157,7 +164,7 @@ export function DownloadClientsPage() {
|
|||||||
onClose={() => setShowForm(false)}
|
onClose={() => setShowForm(false)}
|
||||||
onSaved={async () => {
|
onSaved={async () => {
|
||||||
setShowForm(false)
|
setShowForm(false)
|
||||||
await refresh()
|
refresh().catch((err: unknown) => toast.error(apiErrorMessage(err, '刷新下载器列表失败')))
|
||||||
}}
|
}}
|
||||||
/>
|
/>
|
||||||
)}
|
)}
|
||||||
@@ -187,18 +194,17 @@ function ClientFormModal({
|
|||||||
|
|
||||||
const onSubmit = async (e: FormEvent) => {
|
const onSubmit = async (e: FormEvent) => {
|
||||||
e.preventDefault()
|
e.preventDefault()
|
||||||
|
if (saving) return
|
||||||
setSaving(true)
|
setSaving(true)
|
||||||
try {
|
try {
|
||||||
if (editing) await downloadClientsAPI.update(editing.id, form)
|
if (editing) await downloadClientsAPI.update(editing.id, form)
|
||||||
else await downloadClientsAPI.create(form)
|
else await downloadClientsAPI.create(form)
|
||||||
toast.success('已保存')
|
toast.success('已保存')
|
||||||
await onSaved()
|
setSaving(false)
|
||||||
|
onSaved()
|
||||||
} catch (err: unknown) {
|
} catch (err: unknown) {
|
||||||
const msg =
|
const msg = apiErrorMessage(err, '保存失败')
|
||||||
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
|
|
||||||
'保存失败'
|
|
||||||
toast.error(msg)
|
toast.error(msg)
|
||||||
} finally {
|
|
||||||
setSaving(false)
|
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 }) {
|
function Field({ label, children }: { label: string; children: React.ReactNode }) {
|
||||||
return (
|
return (
|
||||||
<label className="block">
|
<label className="block">
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
import { FormEvent, useState } from 'react'
|
import { FormEvent, useState } from 'react'
|
||||||
import toast from 'react-hot-toast'
|
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 { authAPI } from '../api/auth'
|
||||||
import { profileAPI } from '../api/profile'
|
import { profileAPI } from '../api/profile'
|
||||||
@@ -18,16 +18,23 @@ export function ProfilePage() {
|
|||||||
const [hideAdult, setHideAdult] = useState(Boolean(user?.hide_adult))
|
const [hideAdult, setHideAdult] = useState(Boolean(user?.hide_adult))
|
||||||
const [oldPwd, setOldPwd] = useState('')
|
const [oldPwd, setOldPwd] = useState('')
|
||||||
const [newPwd, setNewPwd] = useState('')
|
const [newPwd, setNewPwd] = useState('')
|
||||||
|
const [savingProfile, setSavingProfile] = useState(false)
|
||||||
|
const [savingPassword, setSavingPassword] = useState(false)
|
||||||
|
|
||||||
const onProfile = async (e: FormEvent) => {
|
const onProfile = async (e: FormEvent) => {
|
||||||
e.preventDefault()
|
e.preventDefault()
|
||||||
|
if (savingProfile) return
|
||||||
|
setSavingProfile(true)
|
||||||
try {
|
try {
|
||||||
let password: string | undefined
|
let password: string | undefined
|
||||||
const hideAdultChanged = hideAdult !== Boolean(user?.hide_adult)
|
const hideAdultChanged = hideAdult !== Boolean(user?.hide_adult)
|
||||||
if (hideAdultChanged) {
|
const usernameChanged = username.trim() !== (user?.username ?? '')
|
||||||
|
if (hideAdultChanged || usernameChanged) {
|
||||||
const input = await requestPassword({
|
const input = await requestPassword({
|
||||||
title: hideAdult ? '隐藏成人目录' : '取消隐藏成人目录',
|
title: usernameChanged ? '修改用户名' : hideAdult ? '隐藏成人目录' : '取消隐藏成人目录',
|
||||||
message: '此设置会同步影响 Web 与 Emby/Jellyfin/Infuse 等第三方客户端,请输入当前账号密码确认。',
|
message: usernameChanged
|
||||||
|
? '修改用户名后需要使用新用户名登录,请输入当前账号密码确认。'
|
||||||
|
: '此设置会同步影响 Web 与 Emby/Jellyfin/Infuse 等第三方客户端,请输入当前账号密码确认。',
|
||||||
confirmText: '保存设置',
|
confirmText: '保存设置',
|
||||||
})
|
})
|
||||||
if (!input) return
|
if (!input) return
|
||||||
@@ -50,11 +57,15 @@ export function ProfilePage() {
|
|||||||
const msg =
|
const msg =
|
||||||
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '保存失败'
|
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? '保存失败'
|
||||||
toast.error(msg)
|
toast.error(msg)
|
||||||
|
} finally {
|
||||||
|
setSavingProfile(false)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const onPwd = async (e: FormEvent) => {
|
const onPwd = async (e: FormEvent) => {
|
||||||
e.preventDefault()
|
e.preventDefault()
|
||||||
|
if (savingPassword) return
|
||||||
|
setSavingPassword(true)
|
||||||
try {
|
try {
|
||||||
await authAPI.changePassword(oldPwd, newPwd)
|
await authAPI.changePassword(oldPwd, newPwd)
|
||||||
toast.success('密码已更新')
|
toast.success('密码已更新')
|
||||||
@@ -65,6 +76,8 @@ export function ProfilePage() {
|
|||||||
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
|
(err as { response?: { data?: { error?: string } } })?.response?.data?.error ??
|
||||||
'密码更新失败'
|
'密码更新失败'
|
||||||
toast.error(msg)
|
toast.error(msg)
|
||||||
|
} finally {
|
||||||
|
setSavingPassword(false)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -124,8 +137,9 @@ export function ProfilePage() {
|
|||||||
onChange={(e) => setHideAdult(e.target.checked)}
|
onChange={(e) => setHideAdult(e.target.checked)}
|
||||||
/>
|
/>
|
||||||
</label>
|
</label>
|
||||||
<button type="submit" className="neon-button">
|
<button type="submit" disabled={savingProfile} className="neon-button">
|
||||||
<Save size={16} /> 保存
|
{savingProfile ? <Loader2 size={16} className="animate-spin" /> : <Save size={16} />}
|
||||||
|
保存
|
||||||
</button>
|
</button>
|
||||||
</form>
|
</form>
|
||||||
|
|
||||||
@@ -152,8 +166,9 @@ export function ProfilePage() {
|
|||||||
autoComplete="new-password"
|
autoComplete="new-password"
|
||||||
/>
|
/>
|
||||||
</Field>
|
</Field>
|
||||||
<button type="submit" className="neon-button">
|
<button type="submit" disabled={savingPassword} className="neon-button">
|
||||||
<KeyRound size={16} /> 更新密码
|
{savingPassword ? <Loader2 size={16} className="animate-spin" /> : <KeyRound size={16} />}
|
||||||
|
更新密码
|
||||||
</button>
|
</button>
|
||||||
</form>
|
</form>
|
||||||
</div>
|
</div>
|
||||||
|
|||||||
Reference in New Issue
Block a user