Update MediaStationGo branding and deployment docs

This commit is contained in:
ShukeBta
2026-06-17 16:01:16 +08:00
parent 3368ac2947
commit ab4637beed
62 changed files with 1035 additions and 599 deletions
+1 -1
View File
@@ -166,7 +166,7 @@ type AIConfig struct {
MaxConcurrent int `mapstructure:"max_concurrent"`
}
// LicenseConfig configures the optional MediaStationLicenseServer bridge.
// LicenseConfig configures the optional MediaStationGo license server bridge.
type LicenseConfig struct {
ServerURL string `mapstructure:"server_url"`
HMACSecret string `mapstructure:"hmac_secret"`
+1 -1
View File
@@ -645,7 +645,7 @@ func embyFallbackUser(id string) gin.H {
}
return gin.H{
"Id": id,
"Name": "MediaStation",
"Name": "MediaStationGo",
"ServerId": "mediastation-go-001",
"HasPassword": true,
"HasConfiguredPassword": true,
+2 -6
View File
@@ -1,7 +1,7 @@
// Package handler wires the HTTP routes to the service container.
//
// All routes are mounted under /api/* (matching the original MediaStation
// surface) so the frontend dev-server can proxy a single prefix.
// All routes are mounted under /api/* so the frontend dev-server can proxy a
// single prefix.
package handler
import (
@@ -162,7 +162,3 @@ func sseHandler(svc *service.Container) gin.HandlerFunc {
}
}
}
+25 -7
View File
@@ -124,15 +124,33 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
})
return
}
task := startScanHTTPTask(svc, "手动扫描入库", lib.Name, lib.Path)
res, err := svc.Scan.ScanLibrary(c.Request.Context(), id)
if err != nil {
finishHTTPTask(task, err, "scan", "手动扫描入库失败", nil, nil)
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
finishScan, ok := svc.Scan.TryBeginLocalScan(id)
if !ok {
c.JSON(http.StatusAccepted, gin.H{
"library_id": id,
"queued": true,
"already_running": true,
"message": "该媒体库正在后台扫描,请在任务面板查看进度",
"estimate_message": "页面关闭不会中断扫描",
})
return
}
finishHTTPTask(task, nil, "completed", "手动扫描入库结束", scanTaskMetrics(res), scanTaskDetails(res, 20))
c.JSON(http.StatusOK, res)
task := startScanHTTPTask(svc, "手动扫描入库", lib.Name, lib.Path)
go func(libraryID string, task *service.TaskHandle, finish func()) {
defer finish()
res, err := svc.Scan.ScanLibrary(context.Background(), libraryID)
if err != nil {
finishHTTPTask(task, err, "scan", "手动扫描入库失败", scanTaskMetrics(res), scanTaskDetails(res, 20))
return
}
finishHTTPTask(task, nil, "completed", "手动扫描入库结束", scanTaskMetrics(res), scanTaskDetails(res, 20))
}(id, task, finishScan)
c.JSON(http.StatusAccepted, gin.H{
"library_id": id,
"queued": true,
"message": "本地媒体库扫描已在后台运行,页面关闭不会中断",
"estimate_message": "可在右上角任务面板查看扫描进度",
})
}
}
+1 -1
View File
@@ -23,7 +23,7 @@ func registerAuthenticatedRoutes(api *gin.RouterGroup, cfg *config.Config, svc *
// Permissions.
authed.GET("/auth/permissions", getMyPermissionsHandler(svc))
// License activation bridge (admin only; talks to MediaStationLicenseServer).
// License activation bridge (admin only; talks to the configured license server).
authed.GET("/license/status", middleware.AdminRequired(), licenseStatusHandler(svc))
authed.POST("/license/activate", middleware.AdminRequired(), licenseActivateHandler(svc))
authed.POST("/license/heartbeat", middleware.AdminRequired(), licenseHeartbeatHandler(svc))
+18 -5
View File
@@ -16,12 +16,12 @@ const (
RegistrationCodeRenew = "renew"
)
// RegistrationCode 是一次性兑换码。管理员生成后发给用户,用户通过 Bot 兑换:
// RegistrationCode 是兑换码。管理员生成后发给用户,用户通过 Bot 兑换:
// - register:创建并绑定一个新账号;兑换时按 DurationDays 设置账号有效期。
// - renew:给当前绑定账号延长 DurationDays 天有效期。
//
// 兑换成功后记录 UsedByUserID + UsedAt,之后不可再用。ExpiresAt 是兑换码本身
// 的有效期(过期后即使未使用也不能再兑换)。
// MaxUses 控制最多可兑换次数,旧数据/零值按 1 次处理。UsedAt 表示达到最大
// 次数后的耗尽时间;ExpiresAt 是兑换码本身的有效期。
type RegistrationCode struct {
Base
Code string `gorm:"uniqueIndex;size:32;not null" json:"code"`
@@ -29,6 +29,8 @@ type RegistrationCode struct {
DurationDays int `gorm:"default:0" json:"duration_days"` // 账号有效期天数;0 表示永久
CreatedByID string `gorm:"size:36" json:"created_by_id,omitempty"`
UsedByUserID string `gorm:"index;size:36" json:"used_by_user_id,omitempty"`
MaxUses int `gorm:"default:1" json:"max_uses"`
UsedCount int `gorm:"default:0" json:"used_count"`
UsedAt *time.Time `json:"used_at,omitempty"`
ExpiresAt *time.Time `json:"expires_at,omitempty"` // 兑换码本身的有效期
}
@@ -41,8 +43,19 @@ func (c *RegistrationCode) BeforeCreate(_ *gorm.DB) error {
return nil
}
// IsUsed 报告兑换码是否已被使用。
func (c *RegistrationCode) IsUsed() bool { return c.UsedAt != nil }
// EffectiveMaxUses returns the configured max uses, treating legacy zero values
// as one-use codes.
func (c *RegistrationCode) EffectiveMaxUses() int {
if c == nil || c.MaxUses <= 0 {
return 1
}
return c.MaxUses
}
// IsUsed 报告兑换码是否已耗尽。
func (c *RegistrationCode) IsUsed() bool {
return c != nil && (c.UsedAt != nil || c.UsedCount >= c.EffectiveMaxUses())
}
// IsExpired 报告兑换码自身是否过期(与账号有效期无关)。
func (c *RegistrationCode) IsExpired() bool {
+10 -6
View File
@@ -33,14 +33,18 @@ func (r *RegistrationCodeRepository) FindByCode(ctx context.Context, code string
return &c, nil
}
// MarkUsed atomically marks an unused, unexpired code as consumed by userID.
// It returns gorm.ErrRecordNotFound when the code was already used so callers
// can avoid double-spend races.
// MarkUsed atomically consumes one use of a redeemable code. It returns
// gorm.ErrRecordNotFound when the code is exhausted so callers can avoid
// double-spend races.
func (r *RegistrationCodeRepository) MarkUsed(ctx context.Context, id, userID string) error {
now := time.Now()
res := r.db.WithContext(ctx).Model(&model.RegistrationCode{}).
Where("id = ? AND used_at IS NULL", id).
Updates(map[string]any{"used_by_user_id": userID, "used_at": &now})
Where("id = ? AND used_at IS NULL AND used_count < CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END", id).
Updates(map[string]any{
"used_by_user_id": userID,
"used_count": gorm.Expr("used_count + 1"),
"used_at": gorm.Expr("CASE WHEN used_count + 1 >= CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END THEN ? ELSE used_at END", now),
})
if res.Error != nil {
return res.Error
}
@@ -64,7 +68,7 @@ func (r *RegistrationCodeRepository) List(ctx context.Context, limit int) ([]mod
func (r *RegistrationCodeRepository) CountUnused(ctx context.Context) (int64, error) {
var n int64
err := r.db.WithContext(ctx).Model(&model.RegistrationCode{}).
Where("used_at IS NULL").Count(&n).Error
Where("used_at IS NULL AND used_count < CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END").Count(&n).Error
return n, err
}
+2 -2
View File
@@ -4,8 +4,8 @@
// transparently encrypts the api_key column on write and decrypts it on
// read so values stored on disk are useless without the JWT secret.
//
// On first read it seeds the table with the providers MediaStation
// supports today (TMDb / Bangumi / TheTVDB / Fanart / OpenAI / Douban).
// On first read it seeds the table with the providers supported by
// MediaStationGo today (TMDb / Bangumi / TheTVDB / Fanart / OpenAI / Douban).
package service
import (
+1 -1
View File
@@ -46,7 +46,7 @@ var (
const MaxUsers = OpenSourceUserLimit
// SeedAdmin makes sure at least one admin user exists. It mirrors the
// MediaStation behaviour: if no admin row is found we create
// legacy default behaviour: if no admin row is found we create
// `admin / admin123` (overridable through ADMIN_INITIAL_PASSWORD) and warn.
func (s *AuthService) SeedAdmin(ctx context.Context) error {
n, err := s.repo.User.CountAdmins(ctx)
+8
View File
@@ -113,13 +113,21 @@ func (s *TelegramBotService) consumeOpenRegSlot(ctx context.Context) {
// sets the account validity granted on redeem (0 = permanent). validDays sets
// how long the code itself stays redeemable (0 = never expires).
func (s *TelegramBotService) generateCode(ctx context.Context, kind string, durationDays, validDays int, createdBy string) (*model.RegistrationCode, error) {
return s.generateCodeWithUses(ctx, kind, durationDays, validDays, 1, createdBy)
}
func (s *TelegramBotService) generateCodeWithUses(ctx context.Context, kind string, durationDays, validDays, maxUses int, createdBy string) (*model.RegistrationCode, error) {
if kind != model.RegistrationCodeRegister && kind != model.RegistrationCodeRenew {
kind = model.RegistrationCodeRegister
}
if maxUses <= 0 {
maxUses = 1
}
code := &model.RegistrationCode{
Code: randomCode(12),
Kind: kind,
DurationDays: durationDays,
MaxUses: maxUses,
CreatedByID: createdBy,
}
if validDays > 0 {
+57
View File
@@ -176,6 +176,40 @@ func TestRegistrationCodeRedeemOnce(t *testing.T) {
}
}
func TestRegistrationCodeCanBeGeneratedForMultipleUses(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCodeWithUses(ctx, model.RegistrationCodeRenew, 30, 0, 2, "")
if err != nil {
t.Fatal(err)
}
rc, msg := bot.lookupRedeemableCode(ctx, code.Code, model.RegistrationCodeRenew)
if rc == nil {
t.Fatalf("expected valid code, got msg=%q", msg)
}
if err := repos.RegCode.MarkUsed(ctx, rc.ID, "user-1"); err != nil {
t.Fatal(err)
}
rc, msg = bot.lookupRedeemableCode(ctx, code.Code, model.RegistrationCodeRenew)
if rc == nil {
t.Fatalf("code should remain redeemable after first use, got msg=%q", msg)
}
if err := repos.RegCode.MarkUsed(ctx, rc.ID, "user-2"); err != nil {
t.Fatal(err)
}
if _, msg := bot.lookupRedeemableCode(ctx, code.Code, model.RegistrationCodeRenew); msg == "" {
t.Fatal("code should be exhausted after max uses")
}
var used model.RegistrationCode
if err := repos.DB.Where("id = ?", code.ID).First(&used).Error; err != nil {
t.Fatal(err)
}
if used.UsedCount != 2 || used.UsedAt == nil {
t.Fatalf("expected exhausted code with used_count=2, got %+v", used)
}
}
func TestRenewalClearsExpiry(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
@@ -976,6 +1010,29 @@ func TestBotAdminCodeAndUserCommands(t *testing.T) {
}
}
func TestBotGroupMenuShowsAdminActionsOnlyForAdmins(t *testing.T) {
ctx := context.Background()
_, bot := newBotTestService(t)
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9301","group_chat_id":"-1001"}`}
adminMsg := &TelegramMessage{From: TelegramUser{ID: 9301, Username: "admin"}, Chat: TelegramChat{ID: -1001, Type: "group"}}
reply, err := bot.executeCommand(ctx, channel, adminMsg, "/menu")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "管理员入口") || len(reply.Buttons) == 0 {
t.Fatalf("admin group menu should expose management actions, got %#v", reply)
}
userMsg := &TelegramMessage{From: TelegramUser{ID: 9302, Username: "user"}, Chat: TelegramChat{ID: -1001, Type: "group"}}
reply, err = bot.executeCommand(ctx, channel, userMsg, "/menu")
if err != nil {
t.Fatal(err)
}
if strings.Contains(reply.Text, "管理员入口") {
t.Fatalf("non-admin group menu must not expose management actions, got %#v", reply)
}
}
func TestBotAdminUnbindMultipleUsers(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
+25
View File
@@ -48,6 +48,7 @@ type DownloadService struct {
scanner *ScannerService
site *SiteService
tasks *TaskTrackerService
notify *NotifyChannelService
mu sync.Mutex
stopCh chan struct{}
@@ -71,6 +72,10 @@ func (d *DownloadService) SetTaskTracker(tasks *TaskTrackerService) {
d.tasks = tasks
}
func (d *DownloadService) SetNotifyChannels(notify *NotifyChannelService) {
d.notify = notify
}
var torrentEpisodeToken = regexp.MustCompile(`(?i)e\d{1,3}`)
const settingDownloadClientsManaged = "download_clients.managed"
@@ -1201,6 +1206,7 @@ func downloadTaskNeedsCompletion(task model.DownloadTask) bool {
// Media rows is too late for freshly-downloaded files: they usually have not
// been scanned into the library yet.
func (d *DownloadService) onTorrentComplete(ctx context.Context, torrent QBitTorrent) {
d.notifyDownloadComplete(torrent)
if d.organizer == nil {
return
}
@@ -1268,6 +1274,25 @@ func (d *DownloadService) onTorrentComplete(ctx context.Context, torrent QBitTor
zap.Int("errors", len(res.Errors)))
}
func (d *DownloadService) notifyDownloadComplete(torrent QBitTorrent) {
if d == nil || d.notify == nil {
return
}
name := strings.TrimSpace(torrent.Name)
if name == "" {
name = strings.TrimSpace(filepath.Base(torrent.ContentPath))
}
if name == "" {
name = "下载任务"
}
body := fmt.Sprintf("任务:%s\n保存路径:%s\nHash:%s", name, firstNonEmpty(torrent.ContentPath, torrent.SavePath), torrent.Hash)
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
d.notify.Broadcast(ctx, "MediaStationGo 下载完成", body, EventDownloadComplete)
}()
}
func (d *DownloadService) downloadOrganizeTaskName(torrent QBitTorrent, allowReplace bool) string {
name := strings.TrimSpace(torrent.Name)
if name == "" {
+204 -5
View File
@@ -490,10 +490,10 @@ func (e *EmbyService) mediaItems(ctx context.Context, p ItemsParams) (map[string
return emptyItemsEnvelope(p.StartIndex), nil
}
if filterBySeasonNumbers && containsItemType(p.IncludeItemTypes, "Movie") && !containsItemType(p.IncludeItemTypes, "Episode") {
q = q.Where("season_num = 0 AND episode_num = 0")
q = e.filterMovieItems(ctx, q)
}
if filterBySeasonNumbers && containsItemType(p.IncludeItemTypes, "Episode") && !containsItemType(p.IncludeItemTypes, "Movie") {
q = q.Where("season_num > 0 OR episode_num > 0")
q = e.filterEpisodeItems(ctx, q)
}
var total int64
@@ -519,10 +519,23 @@ func (e *EmbyService) mediaItems(ctx context.Context, p ItemsParams) (map[string
}
}
fetchLimit := p.Limit
fetchOffset := p.StartIndex
if fetchLimit > 0 && e.shouldCollapseMediaVersions(ctx, p) {
// Duplicates across merged local/cloud libraries collapse into one Emby
// item with multiple MediaSources. Fetch a wider window so duplicates do
// not consume the whole requested page.
fetchOffset = 0
fetchLimit = p.StartIndex + maxInt(p.Limit*4, p.Limit)
}
var rows []model.Media
if err := q.Order(order).Offset(p.StartIndex).Limit(p.Limit).Find(&rows).Error; err != nil {
if err := q.Order(order).Offset(fetchOffset).Limit(fetchLimit).Find(&rows).Error; err != nil {
return nil, err
}
if e.shouldCollapseMediaVersions(ctx, p) {
rows = e.collapseMediaVersionRows(ctx, rows)
rows = pageSlice(rows, p.StartIndex, p.Limit)
}
items, err := e.payloadsForMedia(ctx, rows, p.UserID)
if err != nil {
return nil, err
@@ -610,6 +623,7 @@ func (e *EmbyService) episodeItems(ctx context.Context, rows []model.Media, p It
}
func (e *EmbyService) payloadsForMedia(ctx context.Context, rows []model.Media, userID string) ([]map[string]any, error) {
rows = e.collapseMediaVersionRows(ctx, rows)
userFavs := map[string]bool{}
userPos := map[string]int64{}
if userID != "" && len(rows) > 0 {
@@ -643,6 +657,44 @@ func (e *EmbyService) payloadsForMedia(ctx context.Context, rows []model.Media,
return items, nil
}
func (e *EmbyService) shouldCollapseMediaVersions(ctx context.Context, p ItemsParams) bool {
if containsItemType(p.IncludeItemTypes, "Series") || containsItemType(p.IncludeItemTypes, "Season") {
return false
}
if containsItemType(p.IncludeItemTypes, "Episode") && !containsItemType(p.IncludeItemTypes, "Movie") {
return true
}
if p.ParentID == "" {
return true
}
episodic, err := e.libraryIsEpisodic(ctx, p.ParentID)
return err == nil && !episodic
}
func (e *EmbyService) collapseMediaVersionRows(ctx context.Context, rows []model.Media) []model.Media {
if len(rows) < 2 {
return rows
}
out := make([]model.Media, 0, len(rows))
indexByKey := make(map[string]int, len(rows))
for _, row := range rows {
key := e.mediaVersionKey(ctx, &row)
if key == "" {
out = append(out, row)
continue
}
if idx, ok := indexByKey[key]; ok {
if preferMediaVersion(row, out[idx]) {
out[idx] = row
}
continue
}
indexByKey[key] = len(out)
out = append(out, row)
}
return out
}
// Item 单条目详情。
func (e *EmbyService) Item(ctx context.Context, mediaID, userID string) (map[string]any, error) {
if lib, err := e.repo.Library.FindByID(ctx, mediaID); err != nil {
@@ -908,7 +960,7 @@ func (e *EmbyService) itemPayload(ctx context.Context, m *model.Media, fav bool,
"Played": played,
"PlayedPercentage": pct,
},
"MediaSources": []map[string]any{e.mediaSource(ctx, m, true, false)},
"MediaSources": e.mediaSourcesForItem(ctx, m, true, false),
}
}
@@ -984,6 +1036,35 @@ func embyLibraryTypeIsEpisodic(typ string) bool {
}
}
func (e *EmbyService) filterMovieItems(ctx context.Context, q *gorm.DB) *gorm.DB {
episodicIDs := e.episodicLibraryIDs(ctx)
if len(episodicIDs) == 0 {
return q
}
return q.Where("(media.season_num = 0 AND media.episode_num = 0) OR media.library_id NOT IN ?", episodicIDs)
}
func (e *EmbyService) filterEpisodeItems(ctx context.Context, q *gorm.DB) *gorm.DB {
episodicIDs := e.episodicLibraryIDs(ctx)
if len(episodicIDs) == 0 {
return q.Where("1 = 0")
}
return q.Where("media.library_id IN ? AND (media.season_num > 0 OR media.episode_num > 0)", episodicIDs)
}
func (e *EmbyService) episodicLibraryIDs(ctx context.Context) []string {
if e == nil || e.repo == nil || e.repo.DB == nil {
return nil
}
var ids []string
if err := e.repo.DB.WithContext(ctx).Model(&model.Library{}).
Where("LOWER(type) IN ?", []string{"tv", "anime", "variety"}).
Pluck("id", &ids).Error; err != nil {
return nil
}
return ids
}
func (e *EmbyService) rememberSeriesGroup(group embySeriesGroup) {
if e == nil || strings.TrimSpace(group.ID) == "" {
return
@@ -1712,7 +1793,7 @@ func (e *EmbyService) PlaybackInfo(ctx context.Context, mediaID, userID string)
}
e.ensureCloudTrackMetadata(ctx, m)
return map[string]any{
"MediaSources": []map[string]any{e.mediaSource(ctx, m, false, e.directPlayOnly(ctx))},
"MediaSources": e.mediaSourcesForItem(ctx, m, false, e.directPlayOnly(ctx)),
"PlaySessionId": fmt.Sprintf("%s-%d", m.ID, time.Now().Unix()),
}, nil
}
@@ -1936,6 +2017,124 @@ func (e *EmbyService) mediaSource(ctx context.Context, m *model.Media, asEmbedde
return src
}
func (e *EmbyService) mediaSourcesForItem(ctx context.Context, m *model.Media, asEmbedded, directOnly bool) []map[string]any {
siblings := e.mediaVersionSiblings(ctx, m)
if len(siblings) == 0 {
return []map[string]any{e.mediaSource(ctx, m, asEmbedded, directOnly)}
}
sources := make([]map[string]any, 0, len(siblings))
for i := range siblings {
media := siblings[i]
sources = append(sources, e.mediaSource(ctx, &media, asEmbedded, directOnly))
}
return sources
}
func (e *EmbyService) mediaVersionSiblings(ctx context.Context, m *model.Media) []model.Media {
if e == nil || e.repo == nil || e.repo.DB == nil || m == nil || strings.TrimSpace(m.ID) == "" {
return nil
}
libraryIDs := e.mergedLibraryIDs(ctx, m.LibraryID)
if len(libraryIDs) == 0 {
libraryIDs = []string{m.LibraryID}
}
q := e.repo.DB.WithContext(ctx).Model(&model.Media{}).
Where("library_id IN ?", libraryIDs).
Where("season_num = ? AND episode_num = ?", m.SeasonNum, m.EpisodeNum)
if m.TMDbID > 0 {
q = q.Where("tm_db_id = ?", m.TMDbID)
} else if m.BangumiID > 0 {
q = q.Where("bangumi_id = ?", m.BangumiID)
} else {
title := strings.TrimSpace(m.Title)
if title == "" {
title = strings.TrimSpace(m.OriginalName)
}
if title == "" {
return []model.Media{*m}
}
q = q.Where("LOWER(title) = ?", strings.ToLower(title))
if m.Year > 0 {
q = q.Where("year = ?", m.Year)
}
}
var rows []model.Media
if err := q.Find(&rows).Error; err != nil || len(rows) == 0 {
return []model.Media{*m}
}
rows = e.collapseExactPathRows(rows)
sort.SliceStable(rows, func(i, j int) bool {
if rows[i].ID == m.ID {
return true
}
if rows[j].ID == m.ID {
return false
}
return preferMediaVersion(rows[i], rows[j])
})
return rows
}
func (e *EmbyService) collapseExactPathRows(rows []model.Media) []model.Media {
if len(rows) < 2 {
return rows
}
out := rows[:0]
seen := map[string]struct{}{}
for _, row := range rows {
path := strings.TrimSpace(row.Path)
if path != "" {
if _, ok := seen[path]; ok {
continue
}
seen[path] = struct{}{}
}
out = append(out, row)
}
return out
}
func (e *EmbyService) mediaVersionKey(ctx context.Context, m *model.Media) string {
if e == nil || m == nil {
return ""
}
ids := e.mergedLibraryIDs(ctx, m.LibraryID)
sort.Strings(ids)
libraryGroup := strings.Join(ids, ",")
if libraryGroup == "" {
libraryGroup = strings.TrimSpace(m.LibraryID)
}
if m.TMDbID > 0 {
return fmt.Sprintf("%s|tmdb:%d|s:%d|e:%d", libraryGroup, m.TMDbID, m.SeasonNum, m.EpisodeNum)
}
if m.BangumiID > 0 {
return fmt.Sprintf("%s|bangumi:%d|s:%d|e:%d", libraryGroup, m.BangumiID, m.SeasonNum, m.EpisodeNum)
}
title := strings.ToLower(strings.TrimSpace(m.Title))
if title == "" {
title = strings.ToLower(strings.TrimSpace(m.OriginalName))
}
if title == "" {
return ""
}
return fmt.Sprintf("%s|title:%s|y:%d|s:%d|e:%d", libraryGroup, title, m.Year, m.SeasonNum, m.EpisodeNum)
}
func preferMediaVersion(candidate, current model.Media) bool {
candidateCloud := strings.TrimSpace(candidate.STRMURL) != "" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(candidate.Path)), "cloud://")
currentCloud := strings.TrimSpace(current.STRMURL) != "" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(current.Path)), "cloud://")
if candidateCloud != currentCloud {
return !candidateCloud
}
if candidate.Width != current.Width {
return candidate.Width > current.Width
}
if candidate.SizeBytes != current.SizeBytes {
return candidate.SizeBytes > current.SizeBytes
}
return candidate.CreatedAt.After(current.CreatedAt)
}
func embySTRMStreamURL(mediaID string) string {
return "/api/stream/" + url.PathEscape(strings.TrimSpace(mediaID))
}
+79 -1
View File
@@ -228,7 +228,7 @@ func TestEmbyCloudAnimeUsesSeriesNameFromChineseSeasonFolder(t *testing.T) {
func TestEmbyMovieLibrarySeasonNumbersStayMovies(t *testing.T) {
svc := newTestEmbyService(t)
lib := model.Library{Name: "动画电影", Path: `/media/movies/animation`, Type: "movie", Enabled: true}
lib := model.Library{Name: "动画电影", Path: `/media/movies/animation`, Type: "Movie", Enabled: true}
if err := svc.repo.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
@@ -276,6 +276,84 @@ func TestEmbyMovieLibrarySeasonNumbersStayMovies(t *testing.T) {
if item["Type"] != "Movie" || item["ParentId"] != lib.ID {
t.Fatalf("direct item should stay Movie, got %#v", item)
}
rootMovies, err := svc.Items(t.Context(), ItemsParams{IncludeItemTypes: []string{"Movie"}, Recursive: true, Limit: 50})
if err != nil {
t.Fatalf("root movie query: %v", err)
}
rootItems := rootMovies["Items"].([]map[string]any)
if len(rootItems) != 1 || rootItems[0]["Id"] != media.ID || rootItems[0]["Type"] != "Movie" {
t.Fatalf("root movie query should include movie-library item despite season numbers, got %#v", rootItems)
}
rootEpisodes, err := svc.Items(t.Context(), ItemsParams{IncludeItemTypes: []string{"Episode"}, Recursive: true, Limit: 50})
if err != nil {
t.Fatalf("root episode query: %v", err)
}
if len(rootEpisodes["Items"].([]map[string]any)) != 0 {
t.Fatalf("root episode query should not expose movie-library item, got %#v", rootEpisodes)
}
}
func TestEmbyMergedLocalCloudMovieVersionsShareMediaSources(t *testing.T) {
svc := newTestEmbyService(t)
local := model.Library{Name: "国产电影", Path: `/media/国产电影`, Type: "movie", Enabled: true}
cloud := model.Library{Name: "OpenList · 国产电影", Path: BuildCloudLibraryPath("openlist", "/国产电影", "/国产电影"), Type: "movie", Enabled: true}
for _, lib := range []*model.Library{&local, &cloud} {
if err := svc.repo.Library.Create(t.Context(), lib); err != nil {
t.Fatalf("create library: %v", err)
}
}
for _, media := range []model.Media{
{
Base: model.Base{ID: "local-version", CreatedAt: time.Now()},
LibraryID: local.ID,
Title: "流浪地球",
Year: 2019,
Path: `/media/国产电影/流浪地球.2019.1080p.mkv`,
Container: "mkv",
Width: 1920,
},
{
Base: model.Base{ID: "cloud-version", CreatedAt: time.Now().Add(time.Minute)},
LibraryID: cloud.ID,
Title: "流浪地球",
Year: 2019,
Path: `cloud://openlist/国产电影/流浪地球.2019.2160p.mkv`,
Container: "mkv",
STRMURL: "https://example.invalid/cloud",
Width: 3840,
},
} {
if err := svc.repo.DB.Create(&media).Error; err != nil {
t.Fatalf("create media: %v", err)
}
}
items, err := svc.Items(t.Context(), ItemsParams{ParentID: local.ID, IncludeItemTypes: []string{"Movie"}, Recursive: true, Limit: 10})
if err != nil {
t.Fatalf("items: %v", err)
}
rows := items["Items"].([]map[string]any)
if len(rows) != 1 {
t.Fatalf("merged local/cloud versions should show as one item, got %#v", rows)
}
if rows[0]["Id"] != "local-version" {
t.Fatalf("local media should be the representative item, got %#v", rows[0])
}
sources := rows[0]["MediaSources"].([]map[string]any)
if len(sources) != 2 {
t.Fatalf("merged item should expose two media sources, got %#v", sources)
}
playback, err := svc.PlaybackInfo(t.Context(), "local-version", "user-1")
if err != nil {
t.Fatalf("playback: %v", err)
}
playSources := playback["MediaSources"].([]map[string]any)
if len(playSources) != 2 {
t.Fatalf("playback should expose local and cloud versions, got %#v", playSources)
}
}
func TestEmbyRootItemsExposeLibraries(t *testing.T) {
+2 -1
View File
@@ -21,6 +21,7 @@ const (
EventDownloadComplete = "download_complete"
EventScrapeFailed = "scrape_failed"
EventSystemAlert = "system_alert"
EventLibraryIngest = "library_ingest"
)
// NotifyEvent 是通知事件的数据结构。
@@ -43,7 +44,7 @@ type NotifyService struct {
repo *repository.Container
crypto *CryptoService
mu sync.RWMutex
mu sync.RWMutex
providers map[string]NotifyProvider // type -> provider
}
+62 -1
View File
@@ -30,7 +30,7 @@ import (
)
// videoExtensions lists the file extensions treated as media. Matches the
// MediaStation Python defaults.
// legacy Python defaults.
var videoExtensions = map[string]struct{}{
".mkv": {},
".mp4": {},
@@ -57,6 +57,7 @@ type ScannerService struct {
scraper *ScraperService
storage *StorageConfigService
cache *RuntimeCacheService
notify *NotifyChannelService
imageProxy *ImageProxy
@@ -78,6 +79,8 @@ type ScannerService struct {
localMediaProbeQueue chan localMediaProbeTask
localMediaProbeMu sync.Mutex
localMediaProbing map[string]struct{}
localScanMu sync.Mutex
localScans map[string]struct{}
}
// NewScannerService is the constructor.
@@ -102,6 +105,7 @@ func NewScannerService(
cloudMediaProbeBackoff: make(map[string]time.Time),
localMediaProbeQueue: make(chan localMediaProbeTask, 1024),
localMediaProbing: make(map[string]struct{}),
localScans: make(map[string]struct{}),
}
}
@@ -126,6 +130,12 @@ func (s *ScannerService) SetRuntimeCache(cache *RuntimeCacheService) {
}
}
func (s *ScannerService) SetNotifyChannels(notify *NotifyChannelService) {
if s != nil {
s.notify = notify
}
}
// SetImageProxy lets cloud scans warm sidecar poster/backdrop files into the
// local image cache. This keeps library opening fast without forcing the UI or
// Emby clients to resolve/download every cloud poster on demand.
@@ -281,6 +291,7 @@ type ScanResult struct {
}
var ErrCloudScanAlreadyRunning = errors.New("cloud scan already running")
var ErrLocalScanAlreadyRunning = errors.New("local scan already running")
const maxScanErrorDetails = 20
@@ -594,6 +605,7 @@ func (s *ScannerService) beginCloudScan(ctx context.Context, lib *model.Library,
"errors": current.status.Errors,
})
}
s.notifyScanFinished(lib, res, err, true)
}
return runCtx, finish, nil
}
@@ -832,6 +844,27 @@ func (s *ScannerService) ScanLibraryWithoutAutoScrape(ctx context.Context, libra
return s.scanLibrary(ctx, libraryID, false)
}
func (s *ScannerService) TryBeginLocalScan(libraryID string) (func(), bool) {
if s == nil || strings.TrimSpace(libraryID) == "" {
return func() {}, true
}
s.localScanMu.Lock()
if s.localScans == nil {
s.localScans = make(map[string]struct{})
}
if _, ok := s.localScans[libraryID]; ok {
s.localScanMu.Unlock()
return nil, false
}
s.localScans[libraryID] = struct{}{}
s.localScanMu.Unlock()
return func() {
s.localScanMu.Lock()
delete(s.localScans, libraryID)
s.localScanMu.Unlock()
}, true
}
func (s *ScannerService) scanLibrary(ctx context.Context, libraryID string, autoScrape bool) (*ScanResult, error) {
lib, err := s.repo.Library.FindByID(ctx, libraryID)
if err != nil || lib == nil {
@@ -940,6 +973,7 @@ func (s *ScannerService) scanLibrary(ctx context.Context, libraryID string, auto
"error_count": res.ErrorCount,
"errors": res.Errors,
})
s.notifyScanFinished(lib, res, nil, false)
s.invalidateMediaCache(ctx)
s.maybeGenerateSTRMAfterScan(lib.ID)
@@ -951,6 +985,33 @@ func (s *ScannerService) scanLibrary(ctx context.Context, libraryID string, auto
return res, nil
}
func (s *ScannerService) notifyScanFinished(lib *model.Library, res *ScanResult, err error, cloud bool) {
if s == nil || s.notify == nil || lib == nil || res == nil {
return
}
if err != nil {
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
s.notify.Broadcast(ctx, "MediaStationGo 扫描异常", fmt.Sprintf("媒体库:%s\n错误:%s", lib.Name, err.Error()), EventSystemAlert)
}()
return
}
if res.Added+res.Updated <= 0 {
return
}
source := "本地媒体库"
if cloud {
source = "网盘媒体库"
}
body := fmt.Sprintf("%s:%s\n新增:%d\n更新:%d\n跳过:%d\n移除:%d", source, lib.Name, res.Added, res.Updated, res.Skipped, res.Removed)
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
s.notify.Broadcast(ctx, "MediaStationGo 入库完成", body, EventLibraryIngest)
}()
}
// IngestPath ingests a single file into the given library without walking the
// whole tree. Used by the watcher for incremental, event-driven additions so
// adding one new file no longer triggers a full library re-scan (减少硬盘损耗).
+24
View File
@@ -38,6 +38,7 @@ type ScraperService struct {
fanart *FanartProvider
adult *AdultProvider
hub *Hub
notify *NotifyChannelService
}
// NewScraperService is the constructor.
@@ -66,6 +67,12 @@ func (s *ScraperService) SetDouban(douban *DoubanProvider) {
s.douban = douban
}
func (s *ScraperService) SetNotifyChannels(notify *NotifyChannelService) {
if s != nil {
s.notify = notify
}
}
// yearPattern extracts a 4-digit year (1900-2099).
var yearPattern = regexp.MustCompile(`(?:^|[^\d])(19\d{2}|20\d{2})(?:[^\d]|$)`)
@@ -780,6 +787,7 @@ func (s *ScraperService) EnrichLibrary(ctx context.Context, libraryID string, re
}
if err := s.EnrichOne(ctx, &rows[i]); err != nil {
s.log.Warn("enrich failed", zap.String("media", rows[i].ID), zap.Error(err))
s.notifyScrapeFailed(rows[i], err)
continue
}
processed++
@@ -805,6 +813,22 @@ func (s *ScraperService) EnrichLibrary(ctx context.Context, libraryID string, re
return matched, nil
}
func (s *ScraperService) notifyScrapeFailed(m model.Media, err error) {
if s == nil || s.notify == nil || err == nil {
return
}
body := strings.TrimSpace(m.Title)
if body == "" {
body = m.Path
}
body = "媒体:" + body + "\n错误:" + err.Error()
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
s.notify.Broadcast(ctx, "MediaStationGo 刮削失败", body, EventScrapeFailed)
}()
}
func (s *ScraperService) scrapeDelay(ctx context.Context) time.Duration {
minMS := s.scrapeDelaySetting(ctx, "scrape.delay_min_ms", defaultScrapeDelayMinMS)
maxMS := s.scrapeDelaySetting(ctx, "scrape.delay_max_ms", defaultScrapeDelayMaxMS)
+4
View File
@@ -129,6 +129,8 @@ func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Cont
backup := NewBackupService(cfg, log, repos.DB)
notifier := NewNotifierService(log, repos)
notifyChannels := NewNotifyChannelService(log, repos)
scanner.SetNotifyChannels(notifyChannels)
scraper.SetNotifyChannels(notifyChannels)
playProfiles := NewPlayProfileService(log, repos)
permissions := NewPermissionService(log, repos)
storageCfg := NewStorageConfigService(log, repos, crypto)
@@ -166,8 +168,10 @@ func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Cont
downloads.SetScanner(scanner)
downloads.SetTaskTracker(tasks)
downloads.SetOrganizePipeline(organizePipeline)
downloads.SetNotifyChannels(notifyChannels)
subscription := NewSubscriptionService(cfg, log, repos, downloads, siteSvc, hub)
subscription.SetScraper(scraper)
subscription.SetNotifyChannels(notifyChannels)
// 让图片代理把媒体库根目录视为可读的本地图片位置:海报/封面等
// sidecar 资源就存放在这些(用户自定义、任意)目录下,否则会被
+17 -17
View File
@@ -218,23 +218,23 @@ func (s *SiteService) FindByID(ctx context.Context, id string) (*model.Site, err
// upload_bytes, download_bytes are excluded to prevent injection.
var siteUpdatableFields = map[string]bool{
"name": true,
"url": true,
"type": true,
"auth_type": true,
"api_key": true,
"cookie": true,
"auth_header": true,
"user_agent": true,
"rss_url": true,
"timeout": true,
"priority": true,
"use_proxy": true,
"rate_limit": true,
"url": true,
"type": true,
"auth_type": true,
"api_key": true,
"cookie": true,
"auth_header": true,
"user_agent": true,
"rss_url": true,
"timeout": true,
"priority": true,
"use_proxy": true,
"rate_limit": true,
"browser_emulation": true,
"downloader": true,
"enabled": true,
"is_default": true,
"extra": true,
"downloader": true,
"enabled": true,
"is_default": true,
"extra": true,
}
// Update applies a partial patch to an existing site.
@@ -270,7 +270,7 @@ func (s *SiteService) Delete(ctx context.Context, id string) error {
// TestConnection tries to reach the site's base URL with the configured
// credentials and reports success/failure.
//
// 测试逻辑(与参考项目 ShukeBta/MediaStation 对齐):
// 测试逻辑(与旧版参考实现对齐):
//
// 1. 优先调用对应站点适配器的 Authenticate(),让 PT 站点(M-Team / UNIT3D /
// Gazelle 等)使用各自的开放 API 验证,而不是去拉首页 HTML——后者通常
+1 -1
View File
@@ -119,7 +119,7 @@ func buildRequest(ctx context.Context, method, rawURL string, cfg SiteConfig, bo
req.Header.Set("Cookie", cfg.Cookie)
}
case "api_key":
// 与参考项目(ShukeBta/MediaStation)的 ApplySiteAuthHeaders 对齐:
// 与旧版参考实现的 ApplySiteAuthHeaders 对齐:
// M-Team / UNIT3D 等开放 API 的 PT 站点都使用 `x-api-key` 头部,
// 不要再为 mteam 单独走 Authorization: Bearer,否则服务端会 401。
if cfg.APIKey != "" {
+4 -4
View File
@@ -30,7 +30,7 @@ func (a *MTeamAdapter) Authenticate(ctx context.Context, cfg SiteConfig) error {
if strings.TrimSpace(cfg.APIKey) == "" {
return fmt.Errorf("M-Team 需要填写 API Access Token(控制台 → 实验室 → 存取令牌),不能使用 Cookie 访问开放 API")
}
// 与 ShukeBta/MediaStation 参考实现对齐:
// 与旧版参考实现对齐:
// 用 camelCase 参数(pageNumber / pageSize),同时接受 code 为字符串 "0"
// 或数值 0;兼容 M-Team v3 API 不同版本的返回。
u := cfg.URL + "/api/torrent/search"
@@ -181,8 +181,8 @@ func (a *MTeamAdapter) GetDetail(ctx context.Context, cfg SiteConfig, id string)
// POST /api/torrent/genDlToken?id={tid} (带 x-api-key)
// → {"code":"0","data":"https://api.m-team.cc/api/rss/dlv2?sign=..."}
//
// 拿到的 sign URL 可被任何下载客户端无认证地直接 GET。这是参考项目
// (ShukeBta/MediaStation) 的 _download_torrent_file 方法的子集。
// 拿到的 sign URL 可被任何下载客户端无认证地直接 GET。这是旧版参考实现
// _download_torrent_file 方法的子集。
func (a *MTeamAdapter) GetDownloadURL(ctx context.Context, cfg SiteConfig, id string) (string, error) {
u := cfg.URL + "/api/torrent/genDlToken?id=" + id
// genDlToken 是 POST 但参数走 query string;body 留空。
@@ -220,7 +220,7 @@ func (a *MTeamAdapter) GetDownloadURL(ctx context.Context, cfg SiteConfig, id st
// parseMTeamJSON 解析 MTeam v3 JSON 响应。
//
// 响应结构(与 ShukeBta/MediaStation 参考项目一致):
// 响应结构(与旧版参考实现一致):
//
// {
// "code": "0", // 字符串 "0" 表示成功
+22
View File
@@ -34,6 +34,7 @@ type SubscriptionService struct {
site *SiteService
scraper *ScraperService
hub *Hub
notify *NotifyChannelService
stop chan struct{}
}
@@ -54,6 +55,10 @@ func (s *SubscriptionService) SetScraper(scraper *ScraperService) {
s.scraper = scraper
}
func (s *SubscriptionService) SetNotifyChannels(notify *NotifyChannelService) {
s.notify = notify
}
// Start runs the polling loop in the background.
func (s *SubscriptionService) Start(ctx context.Context) {
go s.loop(ctx)
@@ -272,6 +277,7 @@ func (s *SubscriptionService) runOne(ctx context.Context, sub *model.Subscriptio
"name": sub.Name,
"queued": queued,
})
s.notifySubscriptionHit(sub, queued, nil)
}
return queued, nil
}
@@ -377,6 +383,7 @@ func (s *SubscriptionService) runSiteSearch(ctx context.Context, sub *model.Subs
"keyword": keyword,
"resources": resources,
})
s.notifySubscriptionHit(sub, queued, resources)
return queued, nil
}
if lastEnqueueErr != nil {
@@ -385,6 +392,21 @@ func (s *SubscriptionService) runSiteSearch(ctx context.Context, sub *model.Subs
return 0, nil
}
func (s *SubscriptionService) notifySubscriptionHit(sub *model.Subscription, queued int, resources []string) {
if s == nil || s.notify == nil || sub == nil || queued <= 0 {
return
}
body := fmt.Sprintf("订阅:%s\n新增资源:%d", sub.Name, queued)
if len(resources) > 0 {
body += "\n资源:\n- " + strings.Join(resources, "\n- ")
}
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
s.notify.Broadcast(ctx, "MediaStationGo 订阅命中新资源", body, EventSubscriptionHit)
}()
}
func (s *SubscriptionService) archiveCompletedSubscription(ctx context.Context, sub *model.Subscription, availability LocalAvailability) error {
if s == nil || s.repo == nil || s.repo.Subscription == nil || sub == nil {
return nil
+1 -1
View File
@@ -4,7 +4,7 @@
// converts SRT to WebVTT on the fly so the browser <track> element can
// load them directly.
//
// External-subtitle discovery rules (matching MediaStation Python defaults):
// External-subtitle discovery rules (matching the legacy Python defaults):
//
// 1. Same directory, same basename, different extension.
// 2. Same directory, ".sub/" or "subs/" subdirectory.
+6 -15
View File
@@ -272,7 +272,7 @@ func TestTelegramReplyAutoDeletesSentMessage(t *testing.T) {
waitForTelegramMethod(t, requests, "deleteMessage")
}
func TestTelegramGroupCommandSendsPanelPrivately(t *testing.T) {
func TestTelegramGroupCommandSendsPanelInGroup(t *testing.T) {
var payloads []struct {
ChatID any `json:"chat_id"`
Text string `json:"text"`
@@ -319,23 +319,14 @@ func TestTelegramGroupCommandSendsPanelPrivately(t *testing.T) {
if err := bot.HandleWebhook(t.Context(), update); err != nil {
t.Fatalf("handle webhook: %v", err)
}
if len(payloads) != 2 {
if len(payloads) != 1 {
t.Fatalf("sendMessage count = %d, payloads=%#v", len(payloads), payloads)
}
if got := fmt.Sprint(payloads[0].ChatID); got != "9002" {
t.Fatalf("first message should be private to requester, chat_id=%s payload=%#v", got, payloads[0])
if got := fmt.Sprint(payloads[0].ChatID); got != "-100123" {
t.Fatalf("message should stay in group, chat_id=%s payload=%#v", got, payloads[0])
}
if payloads[0].ReplyMarkup == nil {
t.Fatalf("private panel should include inline keyboard: %#v", payloads[0])
}
if got := fmt.Sprint(payloads[1].ChatID); got != "-100123" {
t.Fatalf("second message should be group ack, chat_id=%s payload=%#v", got, payloads[1])
}
if payloads[1].ReplyMarkup != nil {
t.Fatalf("group ack must not expose buttons: %#v", payloads[1])
}
if !strings.Contains(payloads[1].Text, "私聊") {
t.Fatalf("group ack should explain private delivery, got %q", payloads[1].Text)
if strings.Contains(payloads[0].Text, "管理员入口") {
t.Fatalf("normal group user must not see admin panel: %#v", payloads[0])
}
}
+1 -14
View File
@@ -981,20 +981,7 @@ func (s *TelegramBotService) replyForMessage(ctx context.Context, channel *model
if strings.TrimSpace(reply.Text) == "" {
return nil
}
if !telegramIsGroupChat(msg.Chat.Type) {
return s.reply(ctx, channel, msg.Chat.ID, reply)
}
if err := s.reply(ctx, channel, msg.From.ID, reply); err != nil {
if s.log != nil {
s.log.Warn("telegram private reply from group failed",
zap.Int("group_chat_id", msg.Chat.ID),
zap.Int("telegram_user_id", msg.From.ID),
zap.Error(sanitizeTelegramError(err)),
)
}
return s.reply(ctx, channel, msg.Chat.ID, telegramCommandReply{Text: telegramGroupPrivateDeliveryFailedHint()})
}
return s.reply(ctx, channel, msg.Chat.ID, telegramCommandReply{Text: telegramGroupPrivateDeliverySentHint()})
return s.reply(ctx, channel, msg.Chat.ID, reply)
}
func (s *TelegramBotService) deleteTelegramSourceMessage(channel *model.NotifyChannel, chatID, messageID int) {
+22 -9
View File
@@ -165,7 +165,7 @@ func TestTelegramGroupHidesAdminPanelFromRegularUsers(t *testing.T) {
}
}
func TestTelegramGroupAdminMenuDoesNotExposeButtonsInGroup(t *testing.T) {
func TestTelegramGroupAdminMenuExposesButtonsOnlyToAdmins(t *testing.T) {
ctx := t.Context()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
@@ -179,19 +179,32 @@ func TestTelegramGroupAdminMenuDoesNotExposeButtonsInGroup(t *testing.T) {
}
menu := bot.mainMenu(ctx, channel, msg)
if telegramReplyHasButtonPrefix(menu, "adm_") {
t.Fatalf("admin group menu must not expose admin buttons publicly: %#v", menu.Buttons)
if !telegramReplyHasButtonPrefix(menu, "adm_") {
t.Fatalf("admin group menu should expose admin buttons, got %#v", menu.Buttons)
}
if !strings.Contains(menu.Text, "请私聊 Bot") {
t.Fatalf("admin group menu should tell admins to use private chat, got %q", menu.Text)
if !strings.Contains(menu.Text, "管理员入口") {
t.Fatalf("admin group menu should label admin section, got %q", menu.Text)
}
reply, handled := bot.handleMenuCallback(ctx, channel, msg, "adm_users")
if !handled {
t.Fatal("admin callback should be handled")
}
if !strings.Contains(reply.Text, "请私聊 Bot") || telegramReplyHasButtonPrefix(reply, "adm_") {
t.Fatalf("group admin callback should not render admin panel publicly: %#v", reply)
if !strings.Contains(reply.Text, "用户管理") {
t.Fatalf("group admin callback should render admin panel, got %#v", reply)
}
normal := &TelegramMessage{
From: TelegramUser{ID: 9002, Username: "viewer", FirstName: "Viewer"},
Chat: TelegramChat{ID: -100123, Type: "group"},
}
normalMenu := bot.mainMenu(ctx, channel, normal)
if telegramReplyHasButtonPrefix(normalMenu, "adm_") || strings.Contains(normalMenu.Text, "管理员入口") {
t.Fatalf("normal group user must not see admin controls: %#v", normalMenu)
}
normalReply, handled := bot.handleMenuCallback(ctx, channel, normal, "adm_users")
if !handled || normalReply.Text != "" || len(normalReply.Buttons) != 0 {
t.Fatalf("normal group user must not use admin callbacks: %#v handled=%v", normalReply, handled)
}
reply, err := bot.executeCommand(ctx, channel, msg, "/users")
@@ -201,8 +214,8 @@ func TestTelegramGroupAdminMenuDoesNotExposeButtonsInGroup(t *testing.T) {
if !strings.Contains(reply.Text, "用户管理") {
t.Fatalf("bound group admin text command should run, got %q", reply.Text)
}
if len(reply.Buttons) != 0 {
t.Fatalf("group admin text command must not expose inline buttons publicly: %#v", reply.Buttons)
if len(reply.Buttons) == 0 {
t.Fatalf("group admin text command should expose admin action buttons: %#v", reply.Buttons)
}
}
+6 -3
View File
@@ -25,12 +25,12 @@ func (s *TelegramBotService) telegramCommandDefinitions(ctx context.Context, cha
return []telegramCommandDefinition{
{Aliases: []string{"/start"}, GroupAllowed: true, Handle: func(args []string) (telegramCommandReply, error) {
if len(args) == 0 {
return s.mainMenu(ctx, channel, telegramPrivateMessageForUser(msg)), nil
return s.mainMenu(ctx, channel, msg), nil
}
return s.cmdStart(ctx, msg, args), nil
}},
{Aliases: []string{"/menu"}, GroupAllowed: true, Handle: func(args []string) (telegramCommandReply, error) {
return s.mainMenu(ctx, channel, telegramPrivateMessageForUser(msg)), nil
return s.mainMenu(ctx, channel, msg), nil
}},
{Aliases: []string{"/cancel"}, GroupAllowed: true, Handle: func(args []string) (telegramCommandReply, error) {
s.takePending(int64(msg.From.ID))
@@ -176,7 +176,9 @@ func (s *TelegramBotService) executeCommand(ctx context.Context, channel *model.
}
reply, err := def.Handle(args)
if telegramIsGroupChat(msg.Chat.Type) && def.AdminOnly && !def.GroupAllowed {
reply.Buttons = nil
if !s.telegramUserIsAdmin(ctx, channel, msg.From.ID) {
reply.Buttons = nil
}
}
return reply, err
}
@@ -273,6 +275,7 @@ func registerTelegramBotCommands(ctx context.Context, cfg map[string]string) err
if err := telegramSetBotCommands(ctx, cfg, telegramGroupBotCommandMenu(), map[string]interface{}{"type": "all_group_chats"}); err != nil {
return err
}
_ = telegramSetBotCommands(ctx, cfg, telegramAdminBotCommandMenu(), map[string]interface{}{"type": "all_chat_administrators"})
adminCommands := telegramAdminBotCommandMenu()
for _, adminID := range telegramConfiguredUserIDs(cfg["admin_user_ids"]) {
+39 -12
View File
@@ -80,7 +80,19 @@ func (s *TelegramBotService) mainMenu(ctx context.Context, channel *model.Notify
)
}
if isAdmin {
header += "\n\n" + telegramGroupPrivateAdminHint()
header += "\n\n<b>管理员入口</b>"
rows = append(rows,
[]telegramInlineButton{{Text: "—— 管理员 ——", Data: "noop"}},
[]telegramInlineButton{
{Text: "📊 容量/状态", Data: "adm_capacity"},
{Text: "👥 用户管理", Data: "adm_users"},
},
[]telegramInlineButton{
{Text: "🔓 开注设置", Data: "adm_openreg"},
{Text: "🎟 生成兑换码", Data: "adm_gencode"},
},
[]telegramInlineButton{{Text: "⚙️ 设备策略", Data: "adm_devicepolicy"}},
)
}
return telegramCommandReply{Text: header, Buttons: rows}
}
@@ -190,10 +202,7 @@ func (s *TelegramBotService) handleMenuCallback(ctx context.Context, channel *mo
}
// ── 管理员专属 ──
if isGroup {
if isAdmin {
return telegramCommandReply{Text: telegramGroupPrivateAdminHint()}, true
}
if isGroup && !isAdmin {
return telegramCommandReply{}, true
}
if !isAdmin {
@@ -557,7 +566,7 @@ func (s *TelegramBotService) createUserFromRegistrationCode(ctx context.Context,
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).
if err := tx.Where("code = ? AND kind = ? AND used_at IS NULL AND used_count < CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END", code, model.RegistrationCodeRegister).
First(&claimed).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errRegistrationCodeAlreadyUsed
@@ -598,8 +607,12 @@ func (s *TelegramBotService) createUserFromRegistrationCode(ctx context.Context,
}
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})
Where("id = ? AND used_at IS NULL AND used_count < CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END", claimed.ID).
Updates(map[string]any{
"used_by_user_id": created.ID,
"used_count": gorm.Expr("used_count + 1"),
"used_at": gorm.Expr("CASE WHEN used_count + 1 >= CASE WHEN max_uses > 0 THEN max_uses ELSE 1 END THEN ? ELSE used_at END", now),
})
if res.Error != nil {
return res.Error
}
@@ -607,7 +620,10 @@ func (s *TelegramBotService) createUserFromRegistrationCode(ctx context.Context,
return errRegistrationCodeAlreadyUsed
}
claimed.UsedByUserID = created.ID
claimed.UsedAt = &now
claimed.UsedCount++
if claimed.UsedCount >= claimed.EffectiveMaxUses() {
claimed.UsedAt = &now
}
return nil
})
if err != nil {
@@ -713,7 +729,7 @@ func (s *TelegramBotService) replyGenCode(ctx context.Context, msg *TelegramMess
func (s *TelegramBotService) cmdGenCode(ctx context.Context, msg *TelegramMessage, args []string) telegramCommandReply {
if len(args) < 2 {
return telegramCommandReply{Text: "用法:<code>/gencode register|renew 天数 [有效天数]</code>\n示例:<code>/gencode register 30</code>、<code>/gencode renew 90 7</code>"}
return telegramCommandReply{Text: "用法:<code>/gencode register|renew 天数 [有效天数] [可用次数]</code>\n示例:<code>/gencode register 30</code>、<code>/gencode renew 90 7 5</code>"}
}
kind := strings.ToLower(strings.TrimSpace(args[0]))
switch kind {
@@ -735,11 +751,18 @@ func (s *TelegramBotService) cmdGenCode(ctx context.Context, msg *TelegramMessag
return telegramCommandReply{Text: "有效天数必须是非负整数。"}
}
}
maxUses := 1
if len(args) > 3 {
maxUses, err = strconv.Atoi(args[3])
if err != nil || maxUses <= 0 {
return telegramCommandReply{Text: "可用次数必须是正整数。"}
}
}
createdBy := ""
if u := s.boundUser(ctx, msg.From.ID); u != nil {
createdBy = u.ID
}
code, err := s.generateCode(ctx, kind, days, validDays, createdBy)
code, err := s.generateCodeWithUses(ctx, kind, days, validDays, maxUses, createdBy)
if err != nil {
return telegramCommandReply{Text: "生成失败:" + err.Error()}
}
@@ -752,7 +775,11 @@ func (s *TelegramBotService) cmdGenCode(ctx context.Context, msg *TelegramMessag
if validDays > 0 && code.ExpiresAt != nil {
valid = "有效至 " + code.ExpiresAt.Format("2006-01-02 15:04")
}
return telegramCommandReply{Text: fmt.Sprintf("已生成%s(%s,%s):\n\n<code>%s</code>", kindLabel, dur, valid, code.Code)}
uses := "单次使用"
if code.EffectiveMaxUses() > 1 {
uses = fmt.Sprintf("最多 %d 次", code.EffectiveMaxUses())
}
return telegramCommandReply{Text: fmt.Sprintf("已生成%s(%s,%s,%s):\n\n<code>%s</code>", kindLabel, dur, valid, uses, code.Code)}
}
func (s *TelegramBotService) replyUserList(ctx context.Context) telegramCommandReply {