From eb5ceba366ac0154dbd1c15e6979d6332b72a612 Mon Sep 17 00:00:00 2001 From: ShukeBta Date: Sun, 7 Jun 2026 15:43:47 +0800 Subject: [PATCH] fix(bot): keep private replies out of group chats --- internal/repository/download_client_repo.go | 10 ++ internal/service/downloads.go | 9 +- internal/service/downloads_test.go | 119 ++++++++++++++++++++ internal/service/telegram_bot.go | 33 +++++- internal/service/telegram_bot_user_test.go | 68 +++++++++++ 5 files changed, 234 insertions(+), 5 deletions(-) diff --git a/internal/repository/download_client_repo.go b/internal/repository/download_client_repo.go index c16251f..b04bbd2 100644 --- a/internal/repository/download_client_repo.go +++ b/internal/repository/download_client_repo.go @@ -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 diff --git a/internal/service/downloads.go b/internal/service/downloads.go index 78c73d6..b18ee35 100644 --- a/internal/service/downloads.go +++ b/internal/service/downloads.go @@ -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 diff --git a/internal/service/downloads_test.go b/internal/service/downloads_test.go index e2cca52..c7f475f 100644 --- a/internal/service/downloads_test.go +++ b/internal/service/downloads_test.go @@ -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) + } +} diff --git a/internal/service/telegram_bot.go b/internal/service/telegram_bot.go index a05f173..384b02c 100644 --- a/internal/service/telegram_bot.go +++ b/internal/service/telegram_bot.go @@ -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). diff --git a/internal/service/telegram_bot_user_test.go b/internal/service/telegram_bot_user_test.go index 5c6a7e2..5f76cf7 100644 --- a/internal/service/telegram_bot_user_test.go +++ b/internal/service/telegram_bot_user_test.go @@ -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) + } +}