feat: set default storage by inline keyboard

This commit is contained in:
krau
2025-02-19 12:23:12 +08:00
parent 692e970772
commit c4eb824457
11 changed files with 157 additions and 32 deletions
+2 -1
View File
@@ -9,6 +9,7 @@ import (
"github.com/krau/SaveAny-Bot/config" "github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/dao" "github.com/krau/SaveAny-Bot/dao"
"github.com/krau/SaveAny-Bot/logger" "github.com/krau/SaveAny-Bot/logger"
"github.com/krau/SaveAny-Bot/storage"
) )
func InitAll() { func InitAll() {
@@ -18,7 +19,7 @@ func InitAll() {
} }
logger.InitLogger() logger.InitLogger()
logger.L.Info("Starting SaveAny-Bot...") logger.L.Info("Starting SaveAny-Bot...")
storage.LoadStorages()
common.Init() common.Init()
dao.Init() dao.Init()
bot.Init() bot.Init()
+76 -4
View File
@@ -35,8 +35,8 @@ func RegisterHandlers(dispatcher dispatcher.Dispatcher) {
} }
dispatcher.AddHandler(handlers.NewMessage(linkRegexFilter, handleLinkMessage)) dispatcher.AddHandler(handlers.NewMessage(linkRegexFilter, handleLinkMessage))
dispatcher.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix("add"), AddToQueue)) dispatcher.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix("add"), AddToQueue))
dispatcher.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix("set_default"), setDefaultStorage))
dispatcher.AddHandler(handlers.NewMessage(filters.Message.Media, handleFileMessage)) dispatcher.AddHandler(handlers.NewMessage(filters.Message.Media, handleFileMessage))
// dispatcher.AddHandler(handlers.NewMessage(filters.Message.Text, handleConversation))
} }
const noPermissionText string = ` const noPermissionText string = `
@@ -69,7 +69,6 @@ Save Any Bot - 转存你的 Telegram 文件
/silent - 开关静默模式 /silent - 开关静默模式
/storage - 设置默认存储位置 /storage - 设置默认存储位置
/save [自定义文件名] - 保存文件 /save [自定义文件名] - 保存文件
/path <存储类型> <路径> - 更改文件保存路径
静默模式: 开启后 Bot 直接保存到收到的文件到默认位置, 不再询问 静默模式: 开启后 Bot 直接保存到收到的文件到默认位置, 不再询问
@@ -196,11 +195,82 @@ func saveCmd(ctx *ext.Context, update *ext.Update) error {
ReplyMessageID: replied.ID, ReplyMessageID: replied.ID,
ReplyChatID: update.GetUserChat().GetID(), ReplyChatID: update.GetUserChat().GetID(),
FileMessageID: msg.ID, FileMessageID: msg.ID,
UserID: user.ChatID,
}) })
} }
func storageCmd(ctx *ext.Context, update *ext.Update) error { func storageCmd(ctx *ext.Context, update *ext.Update) error {
// TODO: Implement user, err := dao.GetUserByChatID(update.GetUserChat().GetID())
if err != nil {
logger.L.Errorf("Failed to get user: %s", err)
ctx.Reply(update, ext.ReplyTextString("获取用户失败"), nil)
return dispatcher.EndGroups
}
storages := storage.GetUserStorages(user.ChatID)
if len(storages) == 0 {
ctx.Reply(update, ext.ReplyTextString("无可用的存储"), nil)
return dispatcher.EndGroups
}
ctx.Reply(update, ext.ReplyTextString("请选择要设为默认的存储位置"), &ext.ReplyOpts{
Markup: getSetDefaultStorageMarkup(user.ChatID, storages),
})
return dispatcher.EndGroups
}
func setDefaultStorage(ctx *ext.Context, update *ext.Update) error {
args := strings.Split(string(update.CallbackQuery.Data), " ")
userID, _ := strconv.Atoi(args[1])
storageNameHash := args[2]
if userID != int(update.CallbackQuery.GetUserID()) {
ctx.AnswerCallback(&tg.MessagesSetBotCallbackAnswerRequest{
QueryID: update.CallbackQuery.QueryID,
Alert: true,
Message: "你没有权限",
CacheTime: 5,
})
return dispatcher.EndGroups
}
storageName := storageHashName[storageNameHash]
selectedStorage, err := storage.GetStorageByName(storageName)
if err != nil {
logger.L.Errorf("failed to get storage: %s", err)
ctx.AnswerCallback(&tg.MessagesSetBotCallbackAnswerRequest{
QueryID: update.CallbackQuery.QueryID,
Alert: true,
Message: "获取指定存储失败",
CacheTime: 5,
})
return dispatcher.EndGroups
}
user, err := dao.GetUserByChatID(int64(userID))
if err != nil {
logger.L.Errorf("Failed to get user: %s", err)
ctx.AnswerCallback(&tg.MessagesSetBotCallbackAnswerRequest{
QueryID: update.CallbackQuery.QueryID,
Alert: true,
Message: "获取用户失败",
CacheTime: 5,
})
return dispatcher.EndGroups
}
user.DefaultStorage = storageName
if err := dao.UpdateUser(user); err != nil {
logger.L.Errorf("Failed to update user: %s", err)
ctx.AnswerCallback(&tg.MessagesSetBotCallbackAnswerRequest{
QueryID: update.CallbackQuery.QueryID,
Alert: true,
Message: "更新用户失败",
CacheTime: 5,
})
return dispatcher.EndGroups
}
ctx.EditMessage(update.EffectiveChat().GetID(), &tg.MessagesEditMessageRequest{
Message: fmt.Sprintf("已将 %s (%s) 设为默认存储位置", selectedStorage.Name(), selectedStorage.Type()),
ID: update.CallbackQuery.GetMsgID(),
})
return dispatcher.EndGroups return dispatcher.EndGroups
} }
@@ -272,11 +342,12 @@ func handleFileMessage(ctx *ext.Context, update *ext.Update) error {
ReplyMessageID: msg.ID, ReplyMessageID: msg.ID,
ReplyChatID: update.GetUserChat().GetID(), ReplyChatID: update.GetUserChat().GetID(),
FileMessageID: update.EffectiveMessage.ID, FileMessageID: update.EffectiveMessage.ID,
UserID: user.ChatID,
}) })
} }
func AddToQueue(ctx *ext.Context, update *ext.Update) error { func AddToQueue(ctx *ext.Context, update *ext.Update) error {
if !slice.Contain(config.Cfg.Telegram.Admins, update.CallbackQuery.UserID) { if !slice.Contain(config.Cfg.GetUsersID(), update.CallbackQuery.UserID) {
ctx.AnswerCallback(&tg.MessagesSetBotCallbackAnswerRequest{ ctx.AnswerCallback(&tg.MessagesSetBotCallbackAnswerRequest{
QueryID: update.CallbackQuery.QueryID, QueryID: update.CallbackQuery.QueryID,
Alert: true, Alert: true,
@@ -339,6 +410,7 @@ func AddToQueue(ctx *ext.Context, update *ext.Update) error {
ReplyMessageID: record.ReplyMessageID, ReplyMessageID: record.ReplyMessageID,
FileMessageID: record.MessageID, FileMessageID: record.MessageID,
ReplyChatID: record.ReplyChatID, ReplyChatID: record.ReplyChatID,
UserID: update.EffectiveUser().GetID(),
}) })
entityBuilder := entity.Builder{} entityBuilder := entity.Builder{}
+20
View File
@@ -68,6 +68,26 @@ func getSelectStorageMarkup(userChatID int64, fileChatID, fileMessageID int) (*t
return markup, nil return markup, nil
} }
func getSetDefaultStorageMarkup(userChatID int64, storages []storage.Storage) *tg.ReplyInlineMarkup {
buttons := make([]tg.KeyboardButtonClass, 0)
for _, storage := range storages {
nameHash := common.HashString(storage.Name())
storageHashName[nameHash] = storage.Name()
buttons = append(buttons, &tg.KeyboardButtonCallback{
Text: storage.Name(),
Data: []byte(fmt.Sprintf("set_default %d %s", userChatID, nameHash)),
})
}
markup := &tg.ReplyInlineMarkup{}
for i := 0; i < len(buttons); i += 3 {
row := tg.KeyboardButtonRow{}
row.Buttons = buttons[i:min(i+3, len(buttons))]
markup.Rows = append(markup.Rows, row)
}
return markup
}
func FileFromMedia(media tg.MessageMediaClass, customFileName string) (*types.File, error) { func FileFromMedia(media tg.MessageMediaClass, customFileName string) (*types.File, error) {
switch media := media.(type) { switch media := media.(type) {
case *tg.MessageMediaDocument: case *tg.MessageMediaDocument:
+21
View File
@@ -26,3 +26,24 @@ func (c *Config) GetStorageNamesByUserID(userID int64) []string {
} }
return nil return nil
} }
func (c *Config) GetUsersID() []int64 {
var ids []int64
for _, user := range c.Users {
ids = append(ids, user.ID)
}
return ids
}
func (c *Config) HasStorage(userID int64, storageName string) bool {
for _, user := range c.Users {
if user.ID == userID {
if user.Blacklist {
return !slice.Contain(user.Storages, storageName)
} else {
return slice.Contain(user.Storages, storageName)
}
}
}
return false
}
+10
View File
@@ -101,6 +101,16 @@ func Init() error {
if Cfg.Telegram.Admins != nil { if Cfg.Telegram.Admins != nil {
fmt.Println("警告: 你正在使用旧版 Telegram 管理员配置, 该配置下的用户将可用所有存储.\ntelegram.admins 未来版本将会被废弃, 请参考新的配置文件模板, 使用 users 配置替代.") fmt.Println("警告: 你正在使用旧版 Telegram 管理员配置, 该配置下的用户将可用所有存储.\ntelegram.admins 未来版本将会被废弃, 请参考新的配置文件模板, 使用 users 配置替代.")
for _, admin := range Cfg.Telegram.Admins { for _, admin := range Cfg.Telegram.Admins {
found := false
for _, user := range Cfg.Users {
if user.ID == admin {
found = true
break
}
}
if found {
continue
}
Cfg.Users = append(Cfg.Users, userConfig{ Cfg.Users = append(Cfg.Users, userConfig{
ID: admin, ID: admin,
Storages: []string{}, Storages: []string{},
+1 -1
View File
@@ -41,7 +41,7 @@ func processPendingTask(task *types.Task) error {
task.StoragePath = task.File.FileName task.StoragePath = task.File.FileName
} }
taskStorage, err := storage.GetStorageByName(task.StorageName) taskStorage, err := storage.GetStorageByUserIDAndName(task.UserID, task.StorageName)
if err != nil { if err != nil {
return err return err
} }
-9
View File
@@ -24,15 +24,6 @@ type Alist struct {
config config.AlistStorageConfig config config.AlistStorageConfig
} }
var ConfigurableItems = []string{
"url",
"username",
"password",
"base_path",
"token_exp",
"token",
}
func (a *Alist) Init(cfg config.StorageConfig) error { func (a *Alist) Init(cfg config.StorageConfig) error {
alistConfig, ok := cfg.(*config.AlistStorageConfig) alistConfig, ok := cfg.(*config.AlistStorageConfig)
if !ok { if !ok {
-4
View File
@@ -15,10 +15,6 @@ type Local struct {
config config.LocalStorageConfig config config.LocalStorageConfig
} }
var ConfigurableItems = []string{
"base_path",
}
func (l *Local) Init(cfg config.StorageConfig) error { func (l *Local) Init(cfg config.StorageConfig) error {
localConfig, ok := cfg.(*config.LocalStorageConfig) localConfig, ok := cfg.(*config.LocalStorageConfig)
if !ok { if !ok {
+22 -9
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"github.com/krau/SaveAny-Bot/config" "github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/logger"
"github.com/krau/SaveAny-Bot/storage/alist" "github.com/krau/SaveAny-Bot/storage/alist"
"github.com/krau/SaveAny-Bot/storage/local" "github.com/krau/SaveAny-Bot/storage/local"
"github.com/krau/SaveAny-Bot/storage/webdav" "github.com/krau/SaveAny-Bot/storage/webdav"
@@ -44,6 +45,19 @@ func GetStorageByName(name string) (Storage, error) {
return storage, nil return storage, nil
} }
// 检查 user 是否可用指定的 storage, 若不可用则返回未找到错误
func GetStorageByUserIDAndName(chatID int64, name string) (Storage, error) {
if name == "" {
return nil, fmt.Errorf("storage name is required")
}
if !config.Cfg.HasStorage(chatID, name) {
return nil, fmt.Errorf("storage %s not found for user %d", name, chatID)
}
return GetStorageByName(name)
}
func GetUserStorages(chatID int64) []Storage { func GetUserStorages(chatID int64) []Storage {
var storages []Storage var storages []Storage
for _, name := range config.Cfg.GetStorageNamesByUserID(chatID) { for _, name := range config.Cfg.GetStorageNamesByUserID(chatID) {
@@ -78,14 +92,13 @@ func NewStorage(cfg config.StorageConfig) (Storage, error) {
return storage, nil return storage, nil
} }
func GetStorageConfigurableItems(storageType types.StorageType) []string { func LoadStorages() {
switch storageType { logger.L.Info("Loading storages")
case types.StorageTypeAlist: for _, storage := range config.Cfg.Storages {
return alist.ConfigurableItems _, err := GetStorageByName(storage.GetName())
case types.StorageTypeLocal: if err != nil {
return local.ConfigurableItems logger.L.Errorf("Failed to load storage %s: %v", storage.GetName(), err)
case types.StorageTypeWebdav: }
return webdav.ConfigurableItems
} }
return nil logger.L.Infof("Successfully loaded %d storages", len(Storages))
} }
-2
View File
@@ -18,8 +18,6 @@ type Webdav struct {
client *gowebdav.Client client *gowebdav.Client
} }
var ConfigurableItems = []string{"url", "username", "password", "base_path"}
func (w *Webdav) Init(cfg config.StorageConfig) error { func (w *Webdav) Init(cfg config.StorageConfig) error {
webdavConfig, ok := cfg.(*config.WebdavStorageConfig) webdavConfig, ok := cfg.(*config.WebdavStorageConfig)
if !ok { if !ok {
+5 -2
View File
@@ -43,10 +43,13 @@ type Task struct {
StoragePath string StoragePath string
StartTime time.Time StartTime time.Time
FileMessageID int FileMessageID int
FileChatID int64 FileChatID int64
// to track the reply message
ReplyMessageID int ReplyMessageID int
ReplyChatID int64 ReplyChatID int64
// to track the user
UserID int64
} }
func (t Task) String() string { func (t Task) String() string {