fix(bot): keep private replies out of group chats

This commit is contained in:
ShukeBta
2026-06-07 15:43:47 +08:00
parent afd42a5985
commit eb5ceba366
5 changed files with 234 additions and 5 deletions
@@ -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
+7 -2
View File
@@ -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
+119
View File
@@ -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)
}
}
+30 -3
View File
@@ -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)
}
}