mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-29 03:26:37 +08:00
fix(bot): keep private replies out of group chats
This commit is contained in:
@@ -59,6 +59,16 @@ func (r *DownloadClientRepository) ListEnabled(ctx context.Context) ([]model.Dow
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// HasAnyIncludingDeleted reports whether the operator has ever configured a
|
||||
// download client. This distinguishes legacy-only installations from systems
|
||||
// where deleting/disabling all clients is an intentional "stop downloads"
|
||||
// action, even though rows are soft-deleted.
|
||||
func (r *DownloadClientRepository) HasAnyIncludingDeleted(ctx context.Context) (bool, error) {
|
||||
var n int64
|
||||
err := r.db.WithContext(ctx).Unscoped().Model(&model.DownloadClient{}).Count(&n).Error
|
||||
return n > 0, err
|
||||
}
|
||||
|
||||
// Update persists changes to a download client.
|
||||
func (r *DownloadClientRepository) Update(ctx context.Context, c *model.DownloadClient) error {
|
||||
return r.db.WithContext(ctx).Save(c).Error
|
||||
|
||||
@@ -149,9 +149,11 @@ func (d *DownloadService) Stop() {
|
||||
// 默认 qb,但实际下载链路读的还是 Setting 表,导致一直连不上。
|
||||
func (d *DownloadService) ReloadConfig(ctx context.Context) error {
|
||||
cfg := QBitConfig{}
|
||||
hasConfiguredClients := false
|
||||
|
||||
// Path 1: download_clients 表
|
||||
if d.repo.DownloadClient != nil {
|
||||
hasConfiguredClients, _ = d.repo.DownloadClient.HasAnyIncludingDeleted(ctx)
|
||||
if c, err := d.repo.DownloadClient.FindDefault(ctx); err == nil && c != nil && c.Type == "qbittorrent" {
|
||||
cfg.BaseURL = strings.TrimRight(c.Host, "/")
|
||||
cfg.Username = c.Username
|
||||
@@ -159,8 +161,11 @@ func (d *DownloadService) ReloadConfig(ctx context.Context) error {
|
||||
}
|
||||
}
|
||||
|
||||
// Path 2: legacy Setting 表(仅在 client 表未配置时回退)
|
||||
if cfg.BaseURL == "" {
|
||||
// Path 2: legacy Setting 表。
|
||||
// 仅在旧部署“从未使用过 download_clients 表”时回退。只要操作员曾经
|
||||
// 配置过下载器,删除/禁用全部下载器就表示应停止投递,不能再偷偷用
|
||||
// qbittorrent.* 旧设置继续往下载器添加任务。
|
||||
if cfg.BaseURL == "" && !hasConfiguredClients {
|
||||
get := func(k string) string {
|
||||
v, _ := d.repo.Setting.Get(ctx, k)
|
||||
return v
|
||||
|
||||
@@ -104,3 +104,122 @@ func TestAddDownloadWithMetaSkipsExistingTaskBeforeQBAdd(t *testing.T) {
|
||||
t.Fatalf("qb add calls = %d, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReloadConfigDoesNotFallbackToLegacyAfterClientDeleted(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":
|
||||
if atomic.LoadInt32(&addCalls) > 0 {
|
||||
_, _ = w.Write([]byte(`[{"hash":"abc123","name":"Movie 2026 1080p","state":"downloading","progress":0.1}]`))
|
||||
return
|
||||
}
|
||||
_, _ = 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)
|
||||
}
|
||||
client := &model.DownloadClient{Name: "qB", 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.DownloadClient.Delete(t.Context(), client.ID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
|
||||
if err := svc.ReloadConfig(t.Context()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{
|
||||
Title: "Movie 2026 1080p",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected add to fail when the configured downloader was deleted")
|
||||
}
|
||||
if got := atomic.LoadInt32(&addCalls); got != 0 {
|
||||
t.Fatalf("qb add calls = %d, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReloadConfigDoesNotFallbackToLegacyAfterClientDisabled(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)
|
||||
}
|
||||
client := &model.DownloadClient{Name: "qB", 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)
|
||||
}
|
||||
client.Enabled = false
|
||||
if err := repos.DownloadClient.Update(t.Context(), client); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
svc := NewDownloadService(zap.NewNop(), repos, NewHub(zap.NewNop()), nil)
|
||||
if err := svc.ReloadConfig(t.Context()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = svc.AddDownloadWithMeta(t.Context(), "u1", "magnet:?xt=urn:btih:bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb&dn=Movie+2026+1080p", "/downloads", DownloadTaskMeta{
|
||||
Title: "Movie 2026 1080p",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected add to fail when the configured downloader was disabled")
|
||||
}
|
||||
if got := atomic.LoadInt32(&addCalls); got != 0 {
|
||||
t.Fatalf("qb add calls = %d, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -107,6 +107,10 @@ func (s *TelegramBotService) NotifyUserByID(ctx context.Context, userID, text st
|
||||
if err := s.repo.DB.WithContext(ctx).Where("user_id = ?", userID).First(&binding).Error; err != nil {
|
||||
return
|
||||
}
|
||||
targetChatID := telegramPrivateChatIDFromBinding(binding)
|
||||
if targetChatID == 0 {
|
||||
return
|
||||
}
|
||||
channel := s.findChannelByChatID(ctx, int(binding.ChatID))
|
||||
if channel == nil {
|
||||
channels, err := s.repo.NotifyChannel.ListByType(ctx, "telegram")
|
||||
@@ -115,7 +119,7 @@ func (s *TelegramBotService) NotifyUserByID(ctx context.Context, userID, text st
|
||||
}
|
||||
channel = &channels[0]
|
||||
}
|
||||
_ = s.reply(ctx, channel, int(binding.ChatID), telegramCommandReply{Text: text})
|
||||
_ = s.reply(ctx, channel, int(targetChatID), telegramCommandReply{Text: text})
|
||||
}
|
||||
|
||||
// NewTelegramBotService 创建 Telegram Bot 服务。
|
||||
@@ -1078,7 +1082,7 @@ func (s *TelegramBotService) upsertTelegramBinding(ctx context.Context, msg *Tel
|
||||
}
|
||||
return s.repo.DB.WithContext(ctx).Model(&existing).Updates(map[string]any{
|
||||
"telegram_name": name,
|
||||
"chat_id": int64(msg.Chat.ID),
|
||||
"chat_id": telegramBindingChatIDForMessage(msg, &existing),
|
||||
"user_id": userID,
|
||||
}).Error
|
||||
}
|
||||
@@ -1094,11 +1098,34 @@ func (s *TelegramBotService) upsertTelegramBinding(ctx context.Context, msg *Tel
|
||||
return s.repo.DB.WithContext(ctx).Create(&model.TelegramBinding{
|
||||
TelegramUserID: int64(msg.From.ID),
|
||||
TelegramName: name,
|
||||
ChatID: int64(msg.Chat.ID),
|
||||
ChatID: telegramBindingChatIDForMessage(msg, nil),
|
||||
UserID: userID,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func telegramBindingChatIDForMessage(msg *TelegramMessage, existing *model.TelegramBinding) int64 {
|
||||
if msg == nil {
|
||||
if existing != nil {
|
||||
return existing.ChatID
|
||||
}
|
||||
return 0
|
||||
}
|
||||
if msg.Chat.Type == "" || msg.Chat.Type == "private" {
|
||||
return int64(msg.Chat.ID)
|
||||
}
|
||||
if existing != nil && existing.ChatID > 0 {
|
||||
return existing.ChatID
|
||||
}
|
||||
return int64(msg.From.ID)
|
||||
}
|
||||
|
||||
func telegramPrivateChatIDFromBinding(binding model.TelegramBinding) int64 {
|
||||
if binding.ChatID > 0 {
|
||||
return binding.ChatID
|
||||
}
|
||||
return binding.TelegramUserID
|
||||
}
|
||||
|
||||
func (s *TelegramBotService) ensureTelegramAccountBindingAvailable(ctx context.Context, userID string, telegramUserID int64) error {
|
||||
var bound model.TelegramBinding
|
||||
err := s.repo.DB.WithContext(ctx).
|
||||
|
||||
@@ -206,3 +206,71 @@ func TestTelegramStartRejectsAccountAlreadyBoundToAnotherTelegram(t *testing.T)
|
||||
t.Fatal("second telegram account must not be bound")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTelegramBindingFromGroupStoresPrivateUserChatID(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
repos, auth, _, _ := newAuthTestServices(t)
|
||||
user, _, err := auth.Register(ctx, "viewer", "secret-pass")
|
||||
if err != nil {
|
||||
t.Fatalf("register: %v", err)
|
||||
}
|
||||
bot := NewTelegramBotService(zap.NewNop(), repos, nil, auth)
|
||||
msg := &TelegramMessage{
|
||||
From: TelegramUser{ID: 21001, Username: "viewer", FirstName: "Viewer"},
|
||||
Chat: TelegramChat{ID: -100123456, Type: "group"},
|
||||
}
|
||||
|
||||
if err := bot.upsertTelegramBinding(ctx, msg, user.ID); err != nil {
|
||||
t.Fatalf("upsert binding: %v", err)
|
||||
}
|
||||
binding := bot.telegramBinding(ctx, 21001)
|
||||
if binding == nil {
|
||||
t.Fatal("binding should be created")
|
||||
}
|
||||
if binding.ChatID != 21001 {
|
||||
t.Fatalf("group binding must store private user chat id, got %d", binding.ChatID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTelegramBindingFromGroupPreservesExistingPrivateChatID(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
repos, auth, _, _ := newAuthTestServices(t)
|
||||
user, _, err := auth.Register(ctx, "viewer", "secret-pass")
|
||||
if err != nil {
|
||||
t.Fatalf("register: %v", err)
|
||||
}
|
||||
if err := repos.DB.Create(&model.TelegramBinding{
|
||||
TelegramUserID: 21002,
|
||||
TelegramName: "@viewer",
|
||||
ChatID: 987654,
|
||||
UserID: user.ID,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("seed binding: %v", err)
|
||||
}
|
||||
bot := NewTelegramBotService(zap.NewNop(), repos, nil, auth)
|
||||
msg := &TelegramMessage{
|
||||
From: TelegramUser{ID: 21002, Username: "viewer", FirstName: "Viewer"},
|
||||
Chat: TelegramChat{ID: -100123456, Type: "supergroup"},
|
||||
}
|
||||
|
||||
if err := bot.upsertTelegramBinding(ctx, msg, user.ID); err != nil {
|
||||
t.Fatalf("upsert binding: %v", err)
|
||||
}
|
||||
binding := bot.telegramBinding(ctx, 21002)
|
||||
if binding == nil {
|
||||
t.Fatal("binding should exist")
|
||||
}
|
||||
if binding.ChatID != 987654 {
|
||||
t.Fatalf("group command must not overwrite existing private chat id, got %d", binding.ChatID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTelegramPrivateNotifyChatIDFallsBackFromLegacyGroupBinding(t *testing.T) {
|
||||
binding := model.TelegramBinding{
|
||||
TelegramUserID: 21003,
|
||||
ChatID: -100123456,
|
||||
}
|
||||
if got := telegramPrivateChatIDFromBinding(binding); got != 21003 {
|
||||
t.Fatalf("legacy group binding should notify private user chat, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user