refactor: split modules and harden scraping workflows

This commit is contained in:
ShukeBta
2026-06-24 11:59:18 +08:00
parent efca3cbe69
commit 192f35d9fa
470 changed files with 47957 additions and 33839 deletions
+11
View File
@@ -7,3 +7,14 @@
Dockerfile text eol=lf
*.ps1 text eol=crlf
# GitHub Linguist: keep repository language stats focused on product code
# (Go backend + React/TypeScript frontend + Docker packaging). Deployment
# helpers, generated lock files, and static brand assets are still tracked but
# should not appear as primary project languages.
scripts/** linguist-vendored
docker-entrypoint.sh linguist-vendored
web/package-lock.json linguist-generated
web/*.config.js linguist-vendored
web/public/** linguist-vendored
web/src/**/*.css linguist-vendored
+4 -19
View File
@@ -85,24 +85,9 @@ EXPOSE 8080
HEALTHCHECK --interval=30s --timeout=5s --start-period=15s --retries=3 \
CMD busybox wget -q --spider http://127.0.0.1:8080/api/health || exit 1
# Tiny entrypoint that lets us swap to a different UID/GID via PUID/PGID
# (handy on NAS deployments where bind-mounted volumes belong to a non-root
# user). When PUID == 0 we skip su-exec entirely and run as root.
RUN printf '#!/bin/sh\n\
PUID=${PUID:-$(id -u mediastation)}\n\
PGID=${PGID:-$(id -g mediastation)}\n\
if [ "$PUID" != "$(id -u mediastation)" ] || [ "$PGID" != "$(id -g mediastation)" ]; then\n\
deluser mediastation 2>/dev/null || true\n\
delgroup mediastation 2>/dev/null || true\n\
addgroup -g "$PGID" -S mediastation\n\
adduser -u "$PUID" -G mediastation -S mediastation\n\
fi\n\
chown -R mediastation:mediastation /data /cache 2>/dev/null || true\n\
chown mediastation:mediastation /media 2>/dev/null || true\n\
if [ "$PUID" = "0" ]; then\n\
exec mediastation-go\n\
fi\n\
exec su-exec mediastation mediastation-go\n' > /entrypoint.sh \
&& chmod +x /entrypoint.sh
# Tiny entrypoint that lets us run as a NAS host UID/GID via PUID/PGID without
# rewriting /etc/passwd or /etc/group on every container start.
COPY docker-entrypoint.sh /entrypoint.sh
RUN chmod +x /entrypoint.sh
CMD ["/entrypoint.sh"]
+5 -5
View File
@@ -218,14 +218,14 @@ func serveSPA(r *gin.Engine, webDir string) {
assets.Static("/", filepath.Join(webDir, "assets"))
brand := r.Group("/brand")
brand.Use(func(c *gin.Context) {
c.Header("Cache-Control", "public, max-age=86400")
setNoCacheHeaders(c)
c.Next()
})
brand.Static("/", filepath.Join(webDir, "brand"))
for _, icon := range []string{"/favicon.ico", "/favicon.svg"} {
iconPath := filepath.Join(webDir, strings.TrimPrefix(icon, "/"))
r.GET(icon, serveNoCacheFile(iconPath))
r.HEAD(icon, serveNoCacheFile(iconPath))
for _, rootFile := range []string{"/favicon.ico", "/favicon.svg", "/artwork-cache-sw.js"} {
filePath := filepath.Join(webDir, strings.TrimPrefix(rootFile, "/"))
r.GET(rootFile, serveNoCacheFile(filePath))
r.HEAD(rootFile, serveNoCacheFile(filePath))
}
r.NoRoute(func(c *gin.Context) {
path := c.Request.URL.Path
+19
View File
@@ -64,6 +64,9 @@ func TestServeSPAServesAssetsImmutableAndBypassesAPIRoutes(t *testing.T) {
if err := os.WriteFile(filepath.Join(webDir, "brand", "mediastationgo-logo.svg"), []byte("<svg></svg>"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(webDir, "artwork-cache-sw.js"), []byte("self.addEventListener('fetch', () => {})"), 0o644); err != nil {
t.Fatal(err)
}
router := gin.New()
serveSPA(router, webDir)
@@ -84,10 +87,26 @@ func TestServeSPAServesAssetsImmutableAndBypassesAPIRoutes(t *testing.T) {
if brandResp.Code != http.StatusOK {
t.Fatalf("brand asset status = %d, want 200", brandResp.Code)
}
if got := brandResp.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
t.Fatalf("brand asset Cache-Control = %q, want no-store", got)
}
if strings.Contains(brandResp.Body.String(), "index") {
t.Fatalf("brand asset should not serve SPA index: %q", brandResp.Body.String())
}
swReq := httptest.NewRequest(http.MethodGet, "/artwork-cache-sw.js", nil)
swResp := httptest.NewRecorder()
router.ServeHTTP(swResp, swReq)
if swResp.Code != http.StatusOK {
t.Fatalf("service worker status = %d, want 200", swResp.Code)
}
if got := swResp.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
t.Fatalf("service worker Cache-Control = %q, want no-store", got)
}
if strings.Contains(swResp.Body.String(), "index") {
t.Fatalf("service worker should not serve SPA index: %q", swResp.Body.String())
}
for _, path := range []string{
"/api/missing",
"/emby",
+28
View File
@@ -0,0 +1,28 @@
#!/bin/sh
set -eu
run_uid="${PUID:-$(id -u mediastation 2>/dev/null || echo 1000)}"
run_gid="${PGID:-$(id -g mediastation 2>/dev/null || echo 1000)}"
case "$run_uid" in
''|*[!0-9]*)
echo "PUID must be a numeric uid, got: $run_uid" >&2
exit 1
;;
esac
case "$run_gid" in
''|*[!0-9]*)
echo "PGID must be a numeric gid, got: $run_gid" >&2
exit 1
;;
esac
if [ "$run_uid" = "0" ]; then
exec mediastation-go
fi
chown -R "$run_uid:$run_gid" /data /cache 2>/dev/null || true
chown "$run_uid:$run_gid" /media 2>/dev/null || true
exec su-exec "$run_uid:$run_gid" mediastation-go
+32 -733
View File
@@ -3,42 +3,25 @@
package database
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"reflect"
"sort"
"strings"
"time"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"gorm.io/gorm/logger"
"gorm.io/gorm/schema"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// Open initialises the configured GORM database. database.type=auto chooses
// PostgreSQL when database.dsn is present (the Docker Compose default) and
// otherwise falls back to SQLite for old/bare-metal installs.
// PostgreSQL when database.dsn is present and otherwise falls back to SQLite.
func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
gormLogger := logger.New(
zapStdLogger{log: log},
logger.Config{
SlowThreshold: 0,
LogLevel: logger.Warn,
IgnoreRecordNotFoundError: true,
Colorful: false,
},
)
if cfg == nil {
return nil, errors.New("database config is required")
}
dialect := normalizeDatabaseType(cfg.Database.Type)
if dialect == "auto" {
dialect = effectiveAutoDatabaseType(cfg)
@@ -48,7 +31,7 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
return nil, err
}
db, err := gorm.Open(dialector, &gorm.Config{
Logger: gormLogger,
Logger: newGormLogger(log),
PrepareStmt: true,
DisableForeignKeyConstraintWhenMigrating: false,
})
@@ -58,9 +41,31 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
if dialect == "sqlite" {
installSQLiteWriteGate(db)
}
if err := configureConnectionPool(db, cfg); err != nil {
return nil, err
}
return db, nil
}
func newGormLogger(log *zap.Logger) logger.Interface {
if log == nil {
log = zap.NewNop()
}
return logger.New(
zapStdLogger{log: log},
logger.Config{
SlowThreshold: 0,
LogLevel: logger.Warn,
IgnoreRecordNotFoundError: true,
Colorful: false,
},
)
}
func configureConnectionPool(db *gorm.DB, cfg *config.Config) error {
sqlDB, err := db.DB()
if err != nil {
return nil, fmt.Errorf("gorm sqldb: %w", err)
return fmt.Errorf("gorm sqldb: %w", err)
}
if cfg.Database.MaxOpenConns > 0 {
sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns)
@@ -68,7 +73,7 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
if cfg.Database.MaxIdleConns > 0 {
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
}
return db, nil
return nil
}
func normalizeDatabaseType(value string) string {
@@ -106,718 +111,12 @@ func databaseDialector(cfg *config.Config, dialect string) (gorm.Dialector, erro
}
}
// MigrateSQLiteToCurrentIfNeeded copies an existing SQLite database into
// PostgreSQL. Redis is not migrated because it is a rebuildable cache, not a
// source of truth.
const sqliteMigrationCompleteSettingKey = "database.sqlite_migration_complete"
func MigrateSQLiteToCurrentIfNeeded(cfg *config.Config, target *gorm.DB, log *zap.Logger) error {
if cfg == nil || target == nil || target.Dialector == nil || target.Dialector.Name() != "postgres" {
return nil
}
sqlitePath, err := sqliteMigrationSourcePath(cfg, log)
if err != nil {
return err
}
if sqlitePath == "" {
return nil
}
if complete, err := sqliteMigrationMarkedComplete(target); err != nil {
return err
} else if complete {
if log != nil {
log.Info("skip sqlite to postgres migration: already completed")
}
return nil
}
src, err := openSQLiteMigrationSource(cfg, sqlitePath)
if err != nil {
return fmt.Errorf("open sqlite migration source: %w", err)
}
sqlDB, err := src.DB()
if err == nil {
defer sqlDB.Close()
}
started := time.Now()
if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target, log); err != nil {
return err
}
copied, err := copyModelTables(src, target, 500)
if err != nil {
return err
}
if err := markSQLiteMigrationComplete(target); err != nil {
return err
}
if log != nil {
log.Info("sqlite data migrated to postgres",
zap.String("source", sqlitePath),
zap.Int64("rows", copied),
zap.Duration("duration", time.Since(started)))
}
return nil
}
func openSQLiteMigrationSource(cfg *config.Config, sqlitePath string) (*gorm.DB, error) {
srcCfg := *cfg
srcCfg.Database.Type = "sqlite"
srcCfg.Database.DBPath = sqlitePath
return gorm.Open(sqlite.Open(buildSQLiteDSN(&srcCfg)), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
}
func sqliteMigrationSourcePath(cfg *config.Config, log *zap.Logger) (string, error) {
configured := strings.TrimSpace(cfg.Database.DBPath)
if configured != "" {
exists, err := regularFileExists(configured)
if err != nil {
return "", fmt.Errorf("stat sqlite migration source: %w", err)
}
if exists {
return configured, nil
}
}
fallback := filepath.Join(strings.TrimSpace(cfg.App.DataDir), "mediastation.db")
if fallback == "" || sameCleanPath(configured, fallback) {
return "", nil
}
exists, err := regularFileExists(fallback)
if err != nil {
return "", fmt.Errorf("stat default sqlite migration source: %w", err)
}
if !exists {
return "", nil
}
if log != nil && configured != "" {
log.Warn("configured sqlite migration source not found; using data-dir default",
zap.String("configured", configured),
zap.String("fallback", fallback))
}
return fallback, nil
}
func regularFileExists(path string) (bool, error) {
if strings.TrimSpace(path) == "" {
return false, nil
}
info, err := os.Stat(path)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return false, nil
}
return false, err
}
return !info.IsDir(), nil
}
func sameCleanPath(a, b string) bool {
if a == "" || b == "" {
return false
}
return filepath.Clean(a) == filepath.Clean(b)
}
func resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target *gorm.DB, log *zap.Logger) error {
hasRows, err := sqliteSourceHasMigratableRows(src)
if err != nil {
return err
}
if !hasRows {
return nil
}
bootstrapOnly, err := targetLooksLikeBootstrapOnly(target)
if err != nil || !bootstrapOnly {
return err
}
for i := len(model.AllModels()) - 1; i >= 0; i-- {
m := model.AllModels()[i]
if !target.Migrator().HasTable(m) {
continue
}
if err := target.Session(&gorm.Session{AllowGlobalUpdate: true}).Unscoped().Delete(m).Error; err != nil {
return fmt.Errorf("clear bootstrap target table %T: %w", m, err)
}
}
if log != nil {
log.Warn("cleared bootstrap postgres rows before sqlite migration")
}
return nil
}
func sqliteSourceHasMigratableRows(src *gorm.DB) (bool, error) {
for _, table := range []string{"users", "libraries", "media", "settings"} {
exists, err := sqliteTableExists(src, table)
if err != nil {
return false, err
}
if !exists {
continue
}
var count int64
if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&count).Error; err != nil {
return false, fmt.Errorf("count sqlite table %s: %w", table, err)
}
if count > 0 {
return true, nil
}
}
return false, nil
}
func targetLooksLikeBootstrapOnly(target *gorm.DB) (bool, error) {
for _, m := range []any{
&model.Library{},
&model.Series{},
&model.Media{},
&model.PlaybackHistory{},
&model.Favorite{},
&model.Playlist{},
&model.PlaylistItem{},
&model.DownloadTask{},
&model.Subscription{},
} {
if !target.Migrator().HasTable(m) {
continue
}
var count int64
if err := target.Unscoped().Model(m).Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return false, nil
}
}
var userCount int64
if !target.Migrator().HasTable(&model.User{}) {
return true, nil
}
if err := target.Model(&model.User{}).Count(&userCount).Error; err != nil {
return false, err
}
if userCount == 0 {
return true, nil
}
if userCount != 1 {
return false, nil
}
var user model.User
if err := target.Unscoped().Where("username = ?", "admin").First(&user).Error; err != nil {
return false, nil
}
return user.Role == "admin", nil
}
func sqliteMigrationMarkedComplete(db *gorm.DB) (bool, error) {
var value string
err := db.Raw("SELECT value FROM "+quoteIdent("settings")+" WHERE "+quoteIdent("key")+" = ?", sqliteMigrationCompleteSettingKey).Scan(&value).Error
if err != nil {
return false, fmt.Errorf("check sqlite migration marker: %w", err)
}
return strings.EqualFold(strings.TrimSpace(value), "true"), nil
}
func markSQLiteMigrationComplete(db *gorm.DB) error {
now := time.Now()
if err := db.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "key"}},
DoUpdates: clause.AssignmentColumns([]string{"value", "updated_at"}),
}).Create(&model.Setting{
Key: sqliteMigrationCompleteSettingKey,
Value: "true",
UpdatedAt: now,
}).Error; err != nil {
return fmt.Errorf("mark sqlite migration complete: %w", err)
}
return nil
}
func copyModelTables(src, target *gorm.DB, batchSize int) (int64, error) {
if batchSize <= 0 {
batchSize = 500
}
var copied int64
for _, m := range model.AllModels() {
table, err := modelTableName(src, m)
if err != nil {
return copied, err
}
primaryColumns, err := modelPrimaryColumns(src, m)
if err != nil {
return copied, fmt.Errorf("inspect model %T primary keys: %w", m, err)
}
exists, err := sqliteTableExists(src, table)
if err != nil {
return copied, err
}
if !exists {
continue
}
var sourceCount int64
if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&sourceCount).Error; err != nil {
return copied, fmt.Errorf("count sqlite table %s: %w", table, err)
}
if sourceCount == 0 {
continue
}
var targetCount int64
if err := target.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&targetCount).Error; err != nil {
return copied, fmt.Errorf("count target table %s: %w", table, err)
}
modelType := reflect.TypeOf(m)
if modelType.Kind() != reflect.Ptr {
return copied, fmt.Errorf("model %T is not a pointer", m)
}
sliceType := reflect.SliceOf(modelType.Elem())
slicePtr := reflect.New(sliceType)
if err := src.Unscoped().Find(slicePtr.Interface()).Error; err != nil {
return copied, fmt.Errorf("read sqlite table %s: %w", table, err)
}
filtered := slicePtr.Elem()
if targetCount > 0 {
primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns)
if err != nil {
return copied, err
}
filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet)
}
if filtered.Len() == 0 {
continue
}
filteredPtr := reflect.New(filtered.Type())
filteredPtr.Elem().Set(filtered)
if err := target.Clauses(clause.OnConflict{DoNothing: true}).CreateInBatches(filteredPtr.Interface(), batchSize).Error; err != nil {
return copied, fmt.Errorf("copy sqlite table %s: %w", table, err)
}
copied += int64(filtered.Len())
}
return copied, nil
}
func modelPrimaryColumns(db *gorm.DB, m any) ([]string, error) {
stmt := &gorm.Statement{DB: db}
if err := stmt.Parse(m); err != nil {
return nil, err
}
var cols []string
for _, field := range stmt.Schema.PrimaryFields {
cols = append(cols, field.DBName)
}
if len(cols) == 0 {
return nil, fmt.Errorf("no primary key columns")
}
return cols, nil
}
func targetPrimaryKeySet(target *gorm.DB, table string, primaryColumns []string) (map[string]struct{}, error) {
if len(primaryColumns) != 1 {
return nil, nil
}
var values []string
if err := target.Raw("SELECT " + quoteIdent(primaryColumns[0]) + " FROM " + quoteIdent(table)).Scan(&values).Error; err != nil {
return nil, fmt.Errorf("read target primary keys for table %s: %w", table, err)
}
set := make(map[string]struct{}, len(values))
for _, value := range values {
set[value] = struct{}{}
}
return set, nil
}
func filterRowsMissingInTarget(target *gorm.DB, table string, primaryColumns []string, rows reflect.Value, primaryKeySet map[string]struct{}) reflect.Value {
if rows.Kind() != reflect.Slice || rows.Len() == 0 || len(primaryColumns) == 0 {
return rows
}
out := reflect.MakeSlice(rows.Type(), 0, rows.Len())
for i := 0; i < rows.Len(); i++ {
row := rows.Index(i)
keys, ok := rowPrimaryKeys(row, primaryColumns)
if !ok {
out = reflect.Append(out, row)
continue
}
if primaryKeySet != nil {
if _, exists := primaryKeySet[fmt.Sprint(keys[primaryColumns[0]])]; !exists {
out = reflect.Append(out, row)
}
continue
}
if !targetHasPrimaryKey(target, table, keys) {
out = reflect.Append(out, row)
}
}
return out
}
func rowPrimaryKeys(row reflect.Value, primaryColumns []string) (map[string]any, bool) {
if row.Kind() == reflect.Pointer {
if row.IsNil() {
return nil, false
}
row = row.Elem()
}
if row.Kind() != reflect.Struct {
return nil, false
}
keys := make(map[string]any, len(primaryColumns))
for _, column := range primaryColumns {
value, ok := fieldByDBName(row, column)
if !ok || value.IsZero() {
return nil, false
}
keys[column] = value.Interface()
}
return keys, true
}
func fieldByDBName(row reflect.Value, column string) (reflect.Value, bool) {
rowType := row.Type()
for i := 0; i < row.NumField(); i++ {
fieldType := rowType.Field(i)
field := row.Field(i)
if fieldType.Anonymous {
if value, ok := fieldByDBName(field, column); ok {
return value, true
}
}
if columnNameForStructField(fieldType) == column {
if field.Kind() == reflect.Pointer && field.IsNil() {
return reflect.Value{}, false
}
return field, field.CanInterface()
}
}
return reflect.Value{}, false
}
func columnNameForStructField(field reflect.StructField) string {
if field.PkgPath != "" && !field.Anonymous {
return ""
}
tag := field.Tag.Get("gorm")
settings := schema.ParseTagSetting(tag, ";")
if column := settings["COLUMN"]; column != "" {
return column
}
return schema.NamingStrategy{}.ColumnName("", field.Name)
}
func targetHasPrimaryKey(target *gorm.DB, table string, keys map[string]any) bool {
where := make([]string, 0, len(keys))
args := make([]any, 0, len(keys))
for _, column := range sortedMapKeys(keys) {
where = append(where, quoteIdent(column)+" = ?")
args = append(args, keys[column])
}
var count int64
err := target.Raw("SELECT COUNT(1) FROM "+quoteIdent(table)+" WHERE "+strings.Join(where, " AND "), args...).Scan(&count).Error
return err == nil && count > 0
}
func sortedMapKeys(m map[string]any) []string {
keys := make([]string, 0, len(m))
for key := range m {
keys = append(keys, key)
}
sort.Strings(keys)
return keys
}
func sqliteTableExists(db *gorm.DB, table string) (bool, error) {
var count int64
if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'table' AND name = ?`, table).Scan(&count).Error; err != nil {
return false, fmt.Errorf("inspect sqlite table %s: %w", table, err)
}
return count > 0, nil
}
func modelTableName(db *gorm.DB, m any) (string, error) {
stmt := &gorm.Statement{DB: db}
if err := stmt.Parse(m); err != nil {
return "", err
}
return stmt.Schema.Table, nil
}
func quoteIdent(value string) string {
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
}
func installSQLiteWriteGate(db *gorm.DB) {
if db == nil {
return
}
const lockedKey = "mediastation:sqlite_write_locked"
gate := newSQLiteWriteGate()
lock := func(tx *gorm.DB) {
ctx := context.Background()
if tx.Statement != nil && tx.Statement.Context != nil {
ctx = tx.Statement.Context
}
if err := gate.Lock(ctx); err != nil {
_ = tx.AddError(err)
return
}
tx.InstanceSet(lockedKey, struct{}{})
}
unlock := func(tx *gorm.DB) {
if _, ok := tx.InstanceGet(lockedKey); ok {
gate.Unlock()
}
}
_ = db.Callback().Create().Before("gorm:create").Register("mediastation:sqlite_write_lock", lock)
_ = db.Callback().Create().After("gorm:create").Register("mediastation:sqlite_write_unlock", unlock)
_ = db.Callback().Update().Before("gorm:update").Register("mediastation:sqlite_write_lock", lock)
_ = db.Callback().Update().After("gorm:update").Register("mediastation:sqlite_write_unlock", unlock)
_ = db.Callback().Delete().Before("gorm:delete").Register("mediastation:sqlite_write_lock", lock)
_ = db.Callback().Delete().After("gorm:delete").Register("mediastation:sqlite_write_unlock", unlock)
_ = db.Callback().Raw().Before("gorm:raw").Register("mediastation:sqlite_write_lock", lock)
_ = db.Callback().Raw().After("gorm:raw").Register("mediastation:sqlite_write_unlock", unlock)
}
// sqliteWriteGate 串行化进程内的 SQLite 写操作,避免多连接写竞争触发
// SQLITE_BUSY。Lock 尊重语句自身的 context:此前用 sync.Mutex 时,一条
// 长写语句(如 FTS 回填批次)会让登录等关键写操作无限期排队——客户端
// 早已超时断开,goroutine 还挂在互斥锁上。现在等待方可随 context 取消
// 及时失败,不再把整个进程的写路径拖死。
type sqliteWriteGate struct {
ch chan struct{}
}
func newSQLiteWriteGate() *sqliteWriteGate {
return &sqliteWriteGate{ch: make(chan struct{}, 1)}
}
func (g *sqliteWriteGate) Lock(ctx context.Context) error {
select {
case g.ch <- struct{}{}:
return nil
default:
}
if ctx == nil {
ctx = context.Background()
}
select {
case g.ch <- struct{}{}:
return nil
case <-ctx.Done():
return ctx.Err()
}
}
func (g *sqliteWriteGate) Unlock() {
select {
case <-g.ch:
default:
}
}
func buildSQLiteDSN(cfg *config.Config) string {
dbPath := cfg.Database.DBPath
if !filepath.IsAbs(dbPath) {
// keep as-is to respect user-provided relative paths.
dbPath = filepath.Clean(dbPath)
}
dsn := dbPath + "?_pragma=foreign_keys(1)"
if cfg.Database.WALMode {
dsn += "&_pragma=journal_mode(WAL)"
}
if cfg.Database.BusyTimeout > 0 {
dsn += fmt.Sprintf("&_pragma=busy_timeout(%d)", cfg.Database.BusyTimeout)
}
if cfg.Database.CacheSize != 0 {
dsn += fmt.Sprintf("&_pragma=cache_size(%d)", cfg.Database.CacheSize)
}
return dsn
}
// AutoMigrate creates tables for every model registered in the model package.
func AutoMigrate(db *gorm.DB) error {
if err := db.AutoMigrate(model.AllModels()...); err != nil {
return err
}
if err := ensurePostgresColumnCompatibility(db); err != nil {
return err
}
if err := enforceTelegramBindingOneToOne(db); err != nil {
return err
}
if err := ensurePerformanceIndexes(db); err != nil {
return err
}
if isSQLite(db) {
return ensureMediaSearchIndex(db)
}
return nil
}
func ensurePostgresColumnCompatibility(db *gorm.DB) error {
if !isPostgres(db) {
return nil
}
statements := []string{
`ALTER TABLE media ALTER COLUMN container TYPE varchar(128)`,
`ALTER TABLE media ALTER COLUMN genres TYPE text`,
`ALTER TABLE media ALTER COLUMN series_id TYPE varchar(128)`,
`ALTER TABLE media ALTER COLUMN duplicate_of TYPE varchar(128)`,
`ALTER TABLE playback_histories ALTER COLUMN media_id TYPE varchar(128)`,
`ALTER TABLE favorites ALTER COLUMN media_id TYPE varchar(128)`,
`ALTER TABLE playlist_items ALTER COLUMN media_id TYPE varchar(128)`,
`ALTER TABLE strm_records ALTER COLUMN media_id TYPE varchar(128)`,
}
for _, stmt := range statements {
if err := db.Exec(stmt).Error; err != nil {
return err
}
}
return nil
}
func ensurePerformanceIndexes(db *gorm.DB) error {
statements := []string{
`CREATE INDEX IF NOT EXISTS idx_media_library_created_active ON media(library_id, created_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_library_episode_active ON media(library_id, season_num, episode_num, created_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_series_active ON media(series_id, season_num, episode_num) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_favorites_user_media_active ON favorites(user_id, media_id) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_playback_histories_user_media_active ON playback_histories(user_id, media_id, watched_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_playback_histories_resume_active ON playback_histories(user_id, completed, watched_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_play_profiles_user_created_active ON play_profiles(user_id, created_at DESC) WHERE deleted_at IS NULL`,
}
if isSQLite(db) {
statements = append(statements,
`CREATE INDEX IF NOT EXISTS idx_media_title_active ON media(title COLLATE NOCASE) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_original_name_active ON media(original_name COLLATE NOCASE) WHERE deleted_at IS NULL`,
)
} else {
statements = append(statements,
`CREATE INDEX IF NOT EXISTS idx_media_title_active ON media(title) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_original_name_active ON media(original_name) WHERE deleted_at IS NULL`,
)
}
for _, stmt := range statements {
if err := db.Exec(stmt).Error; err != nil {
return err
}
}
return nil
}
func isSQLite(db *gorm.DB) bool {
return db != nil && db.Dialector != nil && db.Dialector.Name() == "sqlite"
}
func isPostgres(db *gorm.DB) bool {
return db != nil && db.Dialector != nil && db.Dialector.Name() == "postgres"
}
// mediaSearchIndexSchemaVersion 标识 FTS 索引的物理布局版本。
// v2:FTS 行的 rowid 与 media.rowid 对齐,并由触发器实时维护。
const mediaSearchIndexSchemaVersion = 2
func ensureMediaSearchIndex(db *gorm.DB) error {
if err := db.Exec(`CREATE TABLE IF NOT EXISTS media_search_meta (id INTEGER PRIMARY KEY CHECK (id = 1), version INTEGER NOT NULL)`).Error; err != nil {
return nil
}
var version int
_ = db.Raw(`SELECT version FROM media_search_meta WHERE id = 1`).Scan(&version).Error
if version != mediaSearchIndexSchemaVersion {
// 旧版(v1)FTS 表按 UNINDEXED 的 media_id 寻址。FTS5 的普通列
// 不支持索引查找,按 media_id 的 DELETE / NOT EXISTS 都是整表
// 扫描:十几万行的库每次启动回填要做上百亿次行访问,纯 Go
// sqlite 直接把 CPU 钉满数小时,并隔着全局写锁拖死登录。
// v2 起 FTS 行的 rowid 与 media.rowid 对齐,所有寻址走 rowid
// 点查,索引一致性交给下方触发器维护。
for _, stmt := range []string{
`DROP TRIGGER IF EXISTS media_search_fts_ai`,
`DROP TRIGGER IF EXISTS media_search_fts_au`,
`DROP TRIGGER IF EXISTS media_search_fts_ad`,
`DROP TABLE IF EXISTS media_search_fts`,
} {
_ = db.Exec(stmt).Error
}
}
if err := db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS media_search_fts USING fts5(media_id UNINDEXED, title, original_name, path, genres, tokenize='trigram')`).Error; err != nil {
if fallbackErr := db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS media_search_fts USING fts5(media_id UNINDEXED, title, original_name, path, genres, tokenize='unicode61')`).Error; fallbackErr != nil {
// FTS is an acceleration path. Some embedded SQLite builds may omit
// FTS5; keep startup working and let repository queries fall back to
// LIKE-based Chinese fuzzy search.
return nil
}
}
// 触发器让 FTS 与 media 行保持同步(新增/标题刮削改写/软删/恢复/
// 硬删全覆盖),应用层不再需要按 media_id 手工刷新索引——也顺带
// 修复了刮削直写 Updates() 后新标题搜不到的问题。
for _, stmt := range []string{
`CREATE TRIGGER IF NOT EXISTS media_search_fts_ai AFTER INSERT ON media WHEN new.deleted_at IS NULL BEGIN
DELETE FROM media_search_fts WHERE rowid = new.rowid;
INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
VALUES (new.rowid, new.id, COALESCE(new.title, ''), COALESCE(new.original_name, ''), COALESCE(new.path, ''), COALESCE(new.genres, ''));
END`,
`CREATE TRIGGER IF NOT EXISTS media_search_fts_au AFTER UPDATE OF title, original_name, path, genres, deleted_at ON media BEGIN
DELETE FROM media_search_fts WHERE rowid = old.rowid;
INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
SELECT new.rowid, new.id, COALESCE(new.title, ''), COALESCE(new.original_name, ''), COALESCE(new.path, ''), COALESCE(new.genres, '')
WHERE new.deleted_at IS NULL;
END`,
`CREATE TRIGGER IF NOT EXISTS media_search_fts_ad AFTER DELETE ON media BEGIN
DELETE FROM media_search_fts WHERE rowid = old.rowid;
END`,
} {
if err := db.Exec(stmt).Error; err != nil {
return err
}
}
if version != mediaSearchIndexSchemaVersion {
if err := db.Exec(`INSERT INTO media_search_meta(id, version) VALUES (1, ?) ON CONFLICT(id) DO UPDATE SET version = excluded.version`, mediaSearchIndexSchemaVersion).Error; err != nil {
return err
}
}
return nil
}
func enforceTelegramBindingOneToOne(db *gorm.DB) error {
if !db.Migrator().HasTable(&model.TelegramBinding{}) {
return nil
}
return db.Transaction(func(tx *gorm.DB) error {
if err := tx.Exec(`
DELETE FROM telegram_bindings
WHERE deleted_at IS NULL
AND user_id IN (
SELECT user_id
FROM telegram_bindings
WHERE deleted_at IS NULL
GROUP BY user_id
HAVING COUNT(*) > 1
)
AND id NOT IN (
SELECT id
FROM (
SELECT id,
ROW_NUMBER() OVER (PARTITION BY user_id ORDER BY updated_at DESC, created_at DESC, id DESC) AS rn
FROM telegram_bindings
WHERE deleted_at IS NULL
) AS ranked_bindings
WHERE rn = 1
)
`).Error; err != nil {
return err
}
return tx.Exec(`
CREATE UNIQUE INDEX IF NOT EXISTS idx_telegram_bindings_user_id_active
ON telegram_bindings(user_id)
WHERE deleted_at IS NULL
`).Error
})
}
// zapStdLogger adapts a *zap.Logger to GORM's tiny logger interface.
type zapStdLogger struct{ log *zap.Logger }
func (z zapStdLogger) Printf(format string, args ...interface{}) {
if z.log == nil {
return
}
z.log.Sugar().Infof(format, args...)
}
+94
View File
@@ -2,6 +2,7 @@ package database
import (
"path/filepath"
"strings"
"testing"
"time"
@@ -12,6 +13,47 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestOpenRequiresConfig(t *testing.T) {
db, err := Open(nil, nil)
if err == nil {
t.Fatal("expected nil config to return an error")
}
if db != nil {
t.Fatal("db should be nil when config is missing")
}
if !strings.Contains(err.Error(), "database config") {
t.Fatalf("error = %v, want database config message", err)
}
}
func TestOpenSQLiteWithNilLoggerConfiguresPool(t *testing.T) {
cfg := &config.Config{}
cfg.Database.Type = "sqlite"
cfg.Database.DBPath = filepath.Join(t.TempDir(), "mediastation.db")
cfg.Database.WALMode = true
cfg.Database.BusyTimeout = 5000
cfg.Database.CacheSize = -2000
cfg.Database.MaxOpenConns = 3
cfg.Database.MaxIdleConns = 2
db, err := Open(cfg, nil)
if err != nil {
t.Fatal(err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
}
defer sqlDB.Close()
if err := db.Exec("SELECT 1").Error; err != nil {
t.Fatal(err)
}
stats := sqlDB.Stats()
if stats.MaxOpenConnections != 3 {
t.Fatalf("MaxOpenConnections = %d, want 3", stats.MaxOpenConnections)
}
}
func TestEnforceTelegramBindingOneToOneCleansDuplicatesAndAddsIndex(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
@@ -83,6 +125,58 @@ func TestEnsurePerformanceIndexesCreatesHotPathIndexes(t *testing.T) {
}
}
func TestEnsureMediaSearchIndexCreatesVersionedTriggers(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.Media{}); err != nil {
t.Fatal(err)
}
if err := ensureMediaSearchIndex(db); err != nil {
t.Fatal(err)
}
if !sqliteFTSTableExists(t, db, "media_search_fts") {
t.Skip("SQLite FTS5 is unavailable in this build")
}
var version int
if err := db.Raw(`SELECT version FROM media_search_meta WHERE id = 1`).Scan(&version).Error; err != nil {
t.Fatal(err)
}
if version != mediaSearchIndexSchemaVersion {
t.Fatalf("media search schema version = %d, want %d", version, mediaSearchIndexSchemaVersion)
}
for _, trigger := range []string{"media_search_fts_ai", "media_search_fts_au", "media_search_fts_ad"} {
var count int
if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'trigger' AND name = ?`, trigger).Scan(&count).Error; err != nil {
t.Fatal(err)
}
if count != 1 {
t.Fatalf("trigger %s count = %d, want 1", trigger, count)
}
}
media := model.Media{LibraryID: "lib-1", Title: "中文搜索电影", Path: "/media/movie.mkv", Genres: "动画,冒险"}
if err := db.Create(&media).Error; err != nil {
t.Fatal(err)
}
var indexed int
if err := db.Raw(`SELECT COUNT(1) FROM media_search_fts WHERE media_id = ?`, media.ID).Scan(&indexed).Error; err != nil {
t.Fatal(err)
}
if indexed != 1 {
t.Fatalf("indexed rows = %d, want inserted media indexed", indexed)
}
}
func sqliteFTSTableExists(t *testing.T, db *gorm.DB, table string) bool {
t.Helper()
var count int
if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'table' AND name = ?`, table).Scan(&count).Error; err != nil {
t.Fatal(err)
}
return count == 1
}
func TestCopyModelTablesMigratesExistingSQLiteRows(t *testing.T) {
src, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
+202
View File
@@ -0,0 +1,202 @@
package database
import (
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// AutoMigrate creates tables for every model registered in the model package.
func AutoMigrate(db *gorm.DB) error {
if err := db.AutoMigrate(model.AllModels()...); err != nil {
return err
}
if err := ensurePostgresColumnCompatibility(db); err != nil {
return err
}
if err := enforceTelegramBindingOneToOne(db); err != nil {
return err
}
if err := ensurePerformanceIndexes(db); err != nil {
return err
}
if isSQLite(db) {
return ensureMediaSearchIndex(db)
}
return nil
}
func ensurePostgresColumnCompatibility(db *gorm.DB) error {
if !isPostgres(db) {
return nil
}
statements := []string{
`ALTER TABLE media ALTER COLUMN container TYPE varchar(128)`,
`ALTER TABLE media ALTER COLUMN genres TYPE text`,
`ALTER TABLE media ALTER COLUMN series_id TYPE varchar(128)`,
`ALTER TABLE media ALTER COLUMN duplicate_of TYPE varchar(128)`,
`ALTER TABLE playback_histories ALTER COLUMN media_id TYPE varchar(128)`,
`ALTER TABLE favorites ALTER COLUMN media_id TYPE varchar(128)`,
`ALTER TABLE playlist_items ALTER COLUMN media_id TYPE varchar(128)`,
`ALTER TABLE strm_records ALTER COLUMN media_id TYPE varchar(128)`,
}
for _, stmt := range statements {
if err := db.Exec(stmt).Error; err != nil {
return err
}
}
return nil
}
func ensurePerformanceIndexes(db *gorm.DB) error {
statements := []string{
`CREATE INDEX IF NOT EXISTS idx_media_library_created_active ON media(library_id, created_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_library_episode_active ON media(library_id, season_num, episode_num, created_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_series_active ON media(series_id, season_num, episode_num) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_favorites_user_media_active ON favorites(user_id, media_id) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_playback_histories_user_media_active ON playback_histories(user_id, media_id, watched_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_playback_histories_resume_active ON playback_histories(user_id, completed, watched_at DESC) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_play_profiles_user_created_active ON play_profiles(user_id, created_at DESC) WHERE deleted_at IS NULL`,
}
if isSQLite(db) {
statements = append(statements,
`CREATE INDEX IF NOT EXISTS idx_media_title_active ON media(title COLLATE NOCASE) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_original_name_active ON media(original_name COLLATE NOCASE) WHERE deleted_at IS NULL`,
)
} else {
statements = append(statements,
`CREATE INDEX IF NOT EXISTS idx_media_title_active ON media(title) WHERE deleted_at IS NULL`,
`CREATE INDEX IF NOT EXISTS idx_media_original_name_active ON media(original_name) WHERE deleted_at IS NULL`,
)
}
for _, stmt := range statements {
if err := db.Exec(stmt).Error; err != nil {
return err
}
}
return nil
}
// mediaSearchIndexSchemaVersion identifies the physical FTS index layout.
// v2 aligns FTS rowids with media rowids and keeps the index current with
// triggers.
const mediaSearchIndexSchemaVersion = 2
func ensureMediaSearchIndex(db *gorm.DB) error {
if err := ensureMediaSearchMetaTable(db); err != nil {
return nil
}
version := currentMediaSearchIndexVersion(db)
if version != mediaSearchIndexSchemaVersion {
resetMediaSearchIndex(db)
}
if !createMediaSearchFTSTable(db) {
return nil
}
if err := createMediaSearchTriggers(db); err != nil {
return err
}
if version != mediaSearchIndexSchemaVersion {
return markMediaSearchIndexVersion(db)
}
return nil
}
func ensureMediaSearchMetaTable(db *gorm.DB) error {
return db.Exec(`CREATE TABLE IF NOT EXISTS media_search_meta (id INTEGER PRIMARY KEY CHECK (id = 1), version INTEGER NOT NULL)`).Error
}
func currentMediaSearchIndexVersion(db *gorm.DB) int {
var version int
_ = db.Raw(`SELECT version FROM media_search_meta WHERE id = 1`).Scan(&version).Error
return version
}
func resetMediaSearchIndex(db *gorm.DB) {
for _, stmt := range []string{
`DROP TRIGGER IF EXISTS media_search_fts_ai`,
`DROP TRIGGER IF EXISTS media_search_fts_au`,
`DROP TRIGGER IF EXISTS media_search_fts_ad`,
`DROP TABLE IF EXISTS media_search_fts`,
} {
_ = db.Exec(stmt).Error
}
}
func createMediaSearchFTSTable(db *gorm.DB) bool {
if err := db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS media_search_fts USING fts5(media_id UNINDEXED, title, original_name, path, genres, tokenize='trigram')`).Error; err == nil {
return true
}
if err := db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS media_search_fts USING fts5(media_id UNINDEXED, title, original_name, path, genres, tokenize='unicode61')`).Error; err == nil {
return true
}
// FTS is an acceleration path. Some embedded SQLite builds may omit FTS5;
// keep startup working and let repository queries fall back to LIKE search.
return false
}
func createMediaSearchTriggers(db *gorm.DB) error {
for _, stmt := range mediaSearchTriggerStatements {
if err := db.Exec(stmt).Error; err != nil {
return err
}
}
return nil
}
var mediaSearchTriggerStatements = []string{
`CREATE TRIGGER IF NOT EXISTS media_search_fts_ai AFTER INSERT ON media WHEN new.deleted_at IS NULL BEGIN
DELETE FROM media_search_fts WHERE rowid = new.rowid;
INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
VALUES (new.rowid, new.id, COALESCE(new.title, ''), COALESCE(new.original_name, ''), COALESCE(new.path, ''), COALESCE(new.genres, ''));
END`,
`CREATE TRIGGER IF NOT EXISTS media_search_fts_au AFTER UPDATE OF title, original_name, path, genres, deleted_at ON media BEGIN
DELETE FROM media_search_fts WHERE rowid = old.rowid;
INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
SELECT new.rowid, new.id, COALESCE(new.title, ''), COALESCE(new.original_name, ''), COALESCE(new.path, ''), COALESCE(new.genres, '')
WHERE new.deleted_at IS NULL;
END`,
`CREATE TRIGGER IF NOT EXISTS media_search_fts_ad AFTER DELETE ON media BEGIN
DELETE FROM media_search_fts WHERE rowid = old.rowid;
END`,
}
func markMediaSearchIndexVersion(db *gorm.DB) error {
return db.Exec(`INSERT INTO media_search_meta(id, version) VALUES (1, ?) ON CONFLICT(id) DO UPDATE SET version = excluded.version`, mediaSearchIndexSchemaVersion).Error
}
func enforceTelegramBindingOneToOne(db *gorm.DB) error {
if !db.Migrator().HasTable(&model.TelegramBinding{}) {
return nil
}
return db.Transaction(func(tx *gorm.DB) error {
if err := tx.Exec(`
DELETE FROM telegram_bindings
WHERE deleted_at IS NULL
AND user_id IN (
SELECT user_id
FROM telegram_bindings
WHERE deleted_at IS NULL
GROUP BY user_id
HAVING COUNT(*) > 1
)
AND id NOT IN (
SELECT id
FROM (
SELECT id,
ROW_NUMBER() OVER (PARTITION BY user_id ORDER BY updated_at DESC, created_at DESC, id DESC) AS rn
FROM telegram_bindings
WHERE deleted_at IS NULL
) AS ranked_bindings
WHERE rn = 1
)
`).Error; err != nil {
return err
}
return tx.Exec(`
CREATE UNIQUE INDEX IF NOT EXISTS idx_telegram_bindings_user_id_active
ON telegram_bindings(user_id)
WHERE deleted_at IS NULL
`).Error
})
}
+463
View File
@@ -0,0 +1,463 @@
package database
import (
"errors"
"fmt"
"os"
"path/filepath"
"reflect"
"sort"
"strings"
"time"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"gorm.io/gorm/logger"
"gorm.io/gorm/schema"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// MigrateSQLiteToCurrentIfNeeded copies an existing SQLite database into
// PostgreSQL. Redis is not migrated because it is a rebuildable cache, not a
// source of truth.
const sqliteMigrationCompleteSettingKey = "database.sqlite_migration_complete"
func MigrateSQLiteToCurrentIfNeeded(cfg *config.Config, target *gorm.DB, log *zap.Logger) error {
if cfg == nil || target == nil || target.Dialector == nil || target.Dialector.Name() != "postgres" {
return nil
}
sqlitePath, err := sqliteMigrationSourcePath(cfg, log)
if err != nil {
return err
}
if sqlitePath == "" {
return nil
}
if complete, err := sqliteMigrationMarkedComplete(target); err != nil {
return err
} else if complete {
if log != nil {
log.Info("skip sqlite to postgres migration: already completed")
}
return nil
}
src, err := openSQLiteMigrationSource(cfg, sqlitePath)
if err != nil {
return fmt.Errorf("open sqlite migration source: %w", err)
}
sqlDB, err := src.DB()
if err == nil {
defer sqlDB.Close()
}
started := time.Now()
if err := resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target, log); err != nil {
return err
}
copied, err := copyModelTables(src, target, 500)
if err != nil {
return err
}
if err := markSQLiteMigrationComplete(target); err != nil {
return err
}
if log != nil {
log.Info("sqlite data migrated to postgres",
zap.String("source", sqlitePath),
zap.Int64("rows", copied),
zap.Duration("duration", time.Since(started)))
}
return nil
}
func openSQLiteMigrationSource(cfg *config.Config, sqlitePath string) (*gorm.DB, error) {
srcCfg := *cfg
srcCfg.Database.Type = "sqlite"
srcCfg.Database.DBPath = sqlitePath
return gorm.Open(sqlite.Open(buildSQLiteDSN(&srcCfg)), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
}
func sqliteMigrationSourcePath(cfg *config.Config, log *zap.Logger) (string, error) {
configured := strings.TrimSpace(cfg.Database.DBPath)
if configured != "" {
exists, err := regularFileExists(configured)
if err != nil {
return "", fmt.Errorf("stat sqlite migration source: %w", err)
}
if exists {
return configured, nil
}
}
fallback := filepath.Join(strings.TrimSpace(cfg.App.DataDir), "mediastation.db")
if fallback == "" || sameCleanPath(configured, fallback) {
return "", nil
}
exists, err := regularFileExists(fallback)
if err != nil {
return "", fmt.Errorf("stat default sqlite migration source: %w", err)
}
if !exists {
return "", nil
}
if log != nil && configured != "" {
log.Warn("configured sqlite migration source not found; using data-dir default",
zap.String("configured", configured),
zap.String("fallback", fallback))
}
return fallback, nil
}
func regularFileExists(path string) (bool, error) {
if strings.TrimSpace(path) == "" {
return false, nil
}
info, err := os.Stat(path)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return false, nil
}
return false, err
}
return !info.IsDir(), nil
}
func sameCleanPath(a, b string) bool {
if a == "" || b == "" {
return false
}
return filepath.Clean(a) == filepath.Clean(b)
}
func resetBootstrapTargetBeforeSQLiteMigrationIfSafe(src, target *gorm.DB, log *zap.Logger) error {
hasRows, err := sqliteSourceHasMigratableRows(src)
if err != nil {
return err
}
if !hasRows {
return nil
}
bootstrapOnly, err := targetLooksLikeBootstrapOnly(target)
if err != nil || !bootstrapOnly {
return err
}
for i := len(model.AllModels()) - 1; i >= 0; i-- {
m := model.AllModels()[i]
if !target.Migrator().HasTable(m) {
continue
}
if err := target.Session(&gorm.Session{AllowGlobalUpdate: true}).Unscoped().Delete(m).Error; err != nil {
return fmt.Errorf("clear bootstrap target table %T: %w", m, err)
}
}
if log != nil {
log.Warn("cleared bootstrap postgres rows before sqlite migration")
}
return nil
}
func sqliteSourceHasMigratableRows(src *gorm.DB) (bool, error) {
for _, table := range []string{"users", "libraries", "media", "settings"} {
exists, err := sqliteTableExists(src, table)
if err != nil {
return false, err
}
if !exists {
continue
}
var count int64
if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&count).Error; err != nil {
return false, fmt.Errorf("count sqlite table %s: %w", table, err)
}
if count > 0 {
return true, nil
}
}
return false, nil
}
func targetLooksLikeBootstrapOnly(target *gorm.DB) (bool, error) {
for _, m := range []any{
&model.Library{},
&model.Series{},
&model.Media{},
&model.PlaybackHistory{},
&model.Favorite{},
&model.Playlist{},
&model.PlaylistItem{},
&model.DownloadTask{},
&model.Subscription{},
} {
if !target.Migrator().HasTable(m) {
continue
}
var count int64
if err := target.Unscoped().Model(m).Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return false, nil
}
}
var userCount int64
if !target.Migrator().HasTable(&model.User{}) {
return true, nil
}
if err := target.Model(&model.User{}).Count(&userCount).Error; err != nil {
return false, err
}
if userCount == 0 {
return true, nil
}
if userCount != 1 {
return false, nil
}
var user model.User
if err := target.Unscoped().Where("username = ?", "admin").First(&user).Error; err != nil {
return false, nil
}
return user.Role == "admin", nil
}
func sqliteMigrationMarkedComplete(db *gorm.DB) (bool, error) {
var value string
err := db.Raw("SELECT value FROM "+quoteIdent("settings")+" WHERE "+quoteIdent("key")+" = ?", sqliteMigrationCompleteSettingKey).Scan(&value).Error
if err != nil {
return false, fmt.Errorf("check sqlite migration marker: %w", err)
}
return strings.EqualFold(strings.TrimSpace(value), "true"), nil
}
func markSQLiteMigrationComplete(db *gorm.DB) error {
now := time.Now()
if err := db.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "key"}},
DoUpdates: clause.AssignmentColumns([]string{"value", "updated_at"}),
}).Create(&model.Setting{
Key: sqliteMigrationCompleteSettingKey,
Value: "true",
UpdatedAt: now,
}).Error; err != nil {
return fmt.Errorf("mark sqlite migration complete: %w", err)
}
return nil
}
func copyModelTables(src, target *gorm.DB, batchSize int) (int64, error) {
if batchSize <= 0 {
batchSize = 500
}
var copied int64
for _, m := range model.AllModels() {
table, err := modelTableName(src, m)
if err != nil {
return copied, err
}
primaryColumns, err := modelPrimaryColumns(src, m)
if err != nil {
return copied, fmt.Errorf("inspect model %T primary keys: %w", m, err)
}
exists, err := sqliteTableExists(src, table)
if err != nil {
return copied, err
}
if !exists {
continue
}
var sourceCount int64
if err := src.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&sourceCount).Error; err != nil {
return copied, fmt.Errorf("count sqlite table %s: %w", table, err)
}
if sourceCount == 0 {
continue
}
var targetCount int64
if err := target.Raw("SELECT COUNT(1) FROM " + quoteIdent(table)).Scan(&targetCount).Error; err != nil {
return copied, fmt.Errorf("count target table %s: %w", table, err)
}
modelType := reflect.TypeOf(m)
if modelType.Kind() != reflect.Ptr {
return copied, fmt.Errorf("model %T is not a pointer", m)
}
sliceType := reflect.SliceOf(modelType.Elem())
slicePtr := reflect.New(sliceType)
if err := src.Unscoped().Find(slicePtr.Interface()).Error; err != nil {
return copied, fmt.Errorf("read sqlite table %s: %w", table, err)
}
filtered := slicePtr.Elem()
if targetCount > 0 {
primaryKeySet, err := targetPrimaryKeySet(target, table, primaryColumns)
if err != nil {
return copied, err
}
filtered = filterRowsMissingInTarget(target, table, primaryColumns, filtered, primaryKeySet)
}
if filtered.Len() == 0 {
continue
}
filteredPtr := reflect.New(filtered.Type())
filteredPtr.Elem().Set(filtered)
if err := target.Clauses(clause.OnConflict{DoNothing: true}).CreateInBatches(filteredPtr.Interface(), batchSize).Error; err != nil {
return copied, fmt.Errorf("copy sqlite table %s: %w", table, err)
}
copied += int64(filtered.Len())
}
return copied, nil
}
func modelPrimaryColumns(db *gorm.DB, m any) ([]string, error) {
stmt := &gorm.Statement{DB: db}
if err := stmt.Parse(m); err != nil {
return nil, err
}
var cols []string
for _, field := range stmt.Schema.PrimaryFields {
cols = append(cols, field.DBName)
}
if len(cols) == 0 {
return nil, fmt.Errorf("no primary key columns")
}
return cols, nil
}
func targetPrimaryKeySet(target *gorm.DB, table string, primaryColumns []string) (map[string]struct{}, error) {
if len(primaryColumns) != 1 {
return nil, nil
}
var values []string
if err := target.Raw("SELECT " + quoteIdent(primaryColumns[0]) + " FROM " + quoteIdent(table)).Scan(&values).Error; err != nil {
return nil, fmt.Errorf("read target primary keys for table %s: %w", table, err)
}
set := make(map[string]struct{}, len(values))
for _, value := range values {
set[value] = struct{}{}
}
return set, nil
}
func filterRowsMissingInTarget(target *gorm.DB, table string, primaryColumns []string, rows reflect.Value, primaryKeySet map[string]struct{}) reflect.Value {
if rows.Kind() != reflect.Slice || rows.Len() == 0 || len(primaryColumns) == 0 {
return rows
}
out := reflect.MakeSlice(rows.Type(), 0, rows.Len())
for i := 0; i < rows.Len(); i++ {
row := rows.Index(i)
keys, ok := rowPrimaryKeys(row, primaryColumns)
if !ok {
out = reflect.Append(out, row)
continue
}
if primaryKeySet != nil {
if _, exists := primaryKeySet[fmt.Sprint(keys[primaryColumns[0]])]; !exists {
out = reflect.Append(out, row)
}
continue
}
if !targetHasPrimaryKey(target, table, keys) {
out = reflect.Append(out, row)
}
}
return out
}
func rowPrimaryKeys(row reflect.Value, primaryColumns []string) (map[string]any, bool) {
if row.Kind() == reflect.Pointer {
if row.IsNil() {
return nil, false
}
row = row.Elem()
}
if row.Kind() != reflect.Struct {
return nil, false
}
keys := make(map[string]any, len(primaryColumns))
for _, column := range primaryColumns {
value, ok := fieldByDBName(row, column)
if !ok || value.IsZero() {
return nil, false
}
keys[column] = value.Interface()
}
return keys, true
}
func fieldByDBName(row reflect.Value, column string) (reflect.Value, bool) {
rowType := row.Type()
for i := 0; i < row.NumField(); i++ {
fieldType := rowType.Field(i)
field := row.Field(i)
if fieldType.Anonymous {
if value, ok := fieldByDBName(field, column); ok {
return value, true
}
}
if columnNameForStructField(fieldType) == column {
if field.Kind() == reflect.Pointer && field.IsNil() {
return reflect.Value{}, false
}
return field, field.CanInterface()
}
}
return reflect.Value{}, false
}
func columnNameForStructField(field reflect.StructField) string {
if field.PkgPath != "" && !field.Anonymous {
return ""
}
tag := field.Tag.Get("gorm")
settings := schema.ParseTagSetting(tag, ";")
if column := settings["COLUMN"]; column != "" {
return column
}
return schema.NamingStrategy{}.ColumnName("", field.Name)
}
func targetHasPrimaryKey(target *gorm.DB, table string, keys map[string]any) bool {
where := make([]string, 0, len(keys))
args := make([]any, 0, len(keys))
for _, column := range sortedMapKeys(keys) {
where = append(where, quoteIdent(column)+" = ?")
args = append(args, keys[column])
}
var count int64
err := target.Raw("SELECT COUNT(1) FROM "+quoteIdent(table)+" WHERE "+strings.Join(where, " AND "), args...).Scan(&count).Error
return err == nil && count > 0
}
func sortedMapKeys(m map[string]any) []string {
keys := make([]string, 0, len(m))
for key := range m {
keys = append(keys, key)
}
sort.Strings(keys)
return keys
}
func sqliteTableExists(db *gorm.DB, table string) (bool, error) {
var count int64
if err := db.Raw(`SELECT COUNT(1) FROM sqlite_master WHERE type = 'table' AND name = ?`, table).Scan(&count).Error; err != nil {
return false, fmt.Errorf("inspect sqlite table %s: %w", table, err)
}
return count > 0, nil
}
func modelTableName(db *gorm.DB, m any) (string, error) {
stmt := &gorm.Statement{DB: db}
if err := stmt.Parse(m); err != nil {
return "", err
}
return stmt.Schema.Table, nil
}
func quoteIdent(value string) string {
return `"` + strings.ReplaceAll(value, `"`, `""`) + `"`
}
+104
View File
@@ -0,0 +1,104 @@
package database
import (
"context"
"fmt"
"path/filepath"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
)
func installSQLiteWriteGate(db *gorm.DB) {
if db == nil {
return
}
const lockedKey = "mediastation:sqlite_write_locked"
gate := newSQLiteWriteGate()
lock := func(tx *gorm.DB) {
ctx := context.Background()
if tx.Statement != nil && tx.Statement.Context != nil {
ctx = tx.Statement.Context
}
if err := gate.Lock(ctx); err != nil {
_ = tx.AddError(err)
return
}
tx.InstanceSet(lockedKey, struct{}{})
}
unlock := func(tx *gorm.DB) {
if _, ok := tx.InstanceGet(lockedKey); ok {
gate.Unlock()
}
}
_ = db.Callback().Create().Before("gorm:create").Register("mediastation:sqlite_write_lock", lock)
_ = db.Callback().Create().After("gorm:create").Register("mediastation:sqlite_write_unlock", unlock)
_ = db.Callback().Update().Before("gorm:update").Register("mediastation:sqlite_write_lock", lock)
_ = db.Callback().Update().After("gorm:update").Register("mediastation:sqlite_write_unlock", unlock)
_ = db.Callback().Delete().Before("gorm:delete").Register("mediastation:sqlite_write_lock", lock)
_ = db.Callback().Delete().After("gorm:delete").Register("mediastation:sqlite_write_unlock", unlock)
_ = db.Callback().Raw().Before("gorm:raw").Register("mediastation:sqlite_write_lock", lock)
_ = db.Callback().Raw().After("gorm:raw").Register("mediastation:sqlite_write_unlock", unlock)
}
// sqliteWriteGate serializes in-process SQLite writes while respecting the
// statement context, so request cancellation can break out of a queued write.
type sqliteWriteGate struct {
ch chan struct{}
}
func newSQLiteWriteGate() *sqliteWriteGate {
return &sqliteWriteGate{ch: make(chan struct{}, 1)}
}
func (g *sqliteWriteGate) Lock(ctx context.Context) error {
select {
case g.ch <- struct{}{}:
return nil
default:
}
if ctx == nil {
ctx = context.Background()
}
select {
case g.ch <- struct{}{}:
return nil
case <-ctx.Done():
return ctx.Err()
}
}
func (g *sqliteWriteGate) Unlock() {
select {
case <-g.ch:
default:
}
}
func buildSQLiteDSN(cfg *config.Config) string {
dbPath := cfg.Database.DBPath
if !filepath.IsAbs(dbPath) {
// keep as-is to respect user-provided relative paths.
dbPath = filepath.Clean(dbPath)
}
dsn := dbPath + "?_pragma=foreign_keys(1)"
if cfg.Database.WALMode {
dsn += "&_pragma=journal_mode(WAL)"
}
if cfg.Database.BusyTimeout > 0 {
dsn += fmt.Sprintf("&_pragma=busy_timeout(%d)", cfg.Database.BusyTimeout)
}
if cfg.Database.CacheSize != 0 {
dsn += fmt.Sprintf("&_pragma=cache_size(%d)", cfg.Database.CacheSize)
}
return dsn
}
func isSQLite(db *gorm.DB) bool {
return db != nil && db.Dialector != nil && db.Dialector.Name() == "sqlite"
}
func isPostgres(db *gorm.DB) bool {
return db != nil && db.Dialector != nil && db.Dialector.Name() == "postgres"
}
+7
View File
@@ -24,6 +24,9 @@ func listUsersHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if svc.Sessions != nil {
svc.Sessions.ApplyToUsers(c.Request.Context(), users)
}
c.JSON(http.StatusOK, users)
}
}
@@ -132,6 +135,10 @@ func deleteUserHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusForbidden, gin.H{"error": "default admin cannot be deleted"})
return
}
if svc.Sessions != nil && svc.Sessions.UserRecentlyActive(c.Request.Context(), c.Param("id"), service.RealtimeDeletionGuardWindow()) {
c.JSON(http.StatusConflict, gin.H{"error": "user has a recent realtime session; confirm the user is offline before deletion"})
return
}
if err := svc.Repo.User.Delete(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
+49
View File
@@ -0,0 +1,49 @@
package handler
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func TestDeleteUserRefusesRecentRealtimeSession(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.User{}); err != nil {
t.Fatal(err)
}
repos := repository.New(db)
admin := model.User{Base: model.Base{ID: "admin"}, Username: "admin", PasswordHash: "x", Role: "admin", IsActive: true}
viewer := model.User{Base: model.Base{ID: "viewer"}, Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
if err := repos.DB.Create(&[]model.User{admin, viewer}).Error; err != nil {
t.Fatal(err)
}
tracker := service.NewSessionTrackerService(zap.NewNop())
tracker.RecordLogin(t.Context(), viewer.ID, viewer.Username, "dev-1", "Apple TV", "Yamby", "10.0.0.8")
svc := &service.Container{Repo: repos, Sessions: tracker}
router := gin.New()
router.DELETE("/admin/users/:id", deleteUserHandler(svc))
req := httptest.NewRequest(http.MethodDelete, "/admin/users/viewer", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusConflict {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
if found, _ := repos.User.FindByID(t.Context(), viewer.ID); found == nil {
t.Fatal("recent realtime user should not be deleted")
}
}
+9
View File
@@ -41,6 +41,12 @@ func loginHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if svc.Sessions != nil {
svc.Sessions.RecordLogin(c.Request.Context(), resp.User.ID, resp.User.Username, "", "Web", "Web", c.ClientIP())
}
if resp.Tokens != nil {
setAccessTokenCookie(c, resp.Tokens.AccessToken, int(resp.Tokens.ExpiresIn))
}
c.JSON(http.StatusOK, gin.H{
"user": resp.User,
"tokens": resp.Tokens,
@@ -65,6 +71,9 @@ func registerHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if tokens != nil {
setAccessTokenCookie(c, tokens.AccessToken, int(tokens.ExpiresIn))
}
c.JSON(http.StatusCreated, gin.H{
"user": u,
"tokens": tokens,
+55
View File
@@ -0,0 +1,55 @@
package handler
import (
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func setAccessTokenCookie(c *gin.Context, token string, maxAgeSeconds int) {
token = strings.TrimSpace(token)
if token == "" {
return
}
if maxAgeSeconds <= 0 {
maxAgeSeconds = int(service.AccessTokenDuration.Seconds())
}
writeAccessTokenCookie(c, token, maxAgeSeconds)
}
func clearAccessTokenCookie(c *gin.Context) {
writeAccessTokenCookie(c, "", -1)
}
func writeAccessTokenCookie(c *gin.Context, value string, maxAgeSeconds int) {
cookie := &http.Cookie{
Name: middleware.AccessTokenCookieName,
Value: value,
Path: middleware.AccessTokenCookiePath,
MaxAge: maxAgeSeconds,
HttpOnly: true,
SameSite: http.SameSiteLaxMode,
Secure: requestIsHTTPS(c),
}
if maxAgeSeconds > 0 {
cookie.Expires = time.Now().Add(time.Duration(maxAgeSeconds) * time.Second)
} else if maxAgeSeconds < 0 {
cookie.Expires = time.Unix(0, 0)
}
http.SetCookie(c.Writer, cookie)
}
func requestIsHTTPS(c *gin.Context) bool {
if c == nil || c.Request == nil {
return false
}
if c.Request.TLS != nil {
return true
}
return strings.EqualFold(c.GetHeader("X-Forwarded-Proto"), "https")
}
+71
View File
@@ -0,0 +1,71 @@
package handler
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
)
func TestSetAccessTokenCookie(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "https://media.local/api/auth/login", nil)
setAccessTokenCookie(c, "access-token", 3600)
cookie := findResponseCookie(t, w, middleware.AccessTokenCookieName)
if cookie.Value != "access-token" {
t.Fatalf("cookie value = %q", cookie.Value)
}
if cookie.Path != middleware.AccessTokenCookiePath {
t.Fatalf("cookie path = %q, want %q", cookie.Path, middleware.AccessTokenCookiePath)
}
if cookie.MaxAge != 3600 {
t.Fatalf("cookie max age = %d, want 3600", cookie.MaxAge)
}
if !cookie.HttpOnly {
t.Fatal("cookie should be HttpOnly")
}
if !cookie.Secure {
t.Fatal("https request should set Secure cookie")
}
if cookie.SameSite != http.SameSiteLaxMode {
t.Fatalf("cookie SameSite = %v, want Lax", cookie.SameSite)
}
}
func TestClearAccessTokenCookie(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "http://127.0.0.1:8080/api/me/logout", nil)
clearAccessTokenCookie(c)
cookie := findResponseCookie(t, w, middleware.AccessTokenCookieName)
if cookie.MaxAge >= 0 {
t.Fatalf("clear cookie max age = %d, want negative", cookie.MaxAge)
}
if cookie.Path != middleware.AccessTokenCookiePath {
t.Fatalf("cookie path = %q, want %q", cookie.Path, middleware.AccessTokenCookiePath)
}
if cookie.Secure {
t.Fatal("plain http request should not set Secure cookie")
}
}
func findResponseCookie(t *testing.T, w *httptest.ResponseRecorder, name string) *http.Cookie {
t.Helper()
for _, cookie := range w.Result().Cookies() {
if cookie.Name == name {
return cookie
}
}
t.Fatalf("missing response cookie %q", name)
return nil
}
+1
View File
@@ -39,6 +39,7 @@ func refreshHandler(svc *service.Container) gin.HandlerFunc {
// so the Vue frontend's logout button gets a 200 instead of 404.
func logoutHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
clearAccessTokenCookie(c)
c.Status(http.StatusNoContent)
}
}
+60 -214
View File
@@ -3,18 +3,10 @@
package handler
import (
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"net/url"
"path"
"sort"
"strings"
"time"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/service"
@@ -25,6 +17,10 @@ import (
func cloudListHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
if !service.IsAdminCloudConfigurable(typ) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider", "items": []any{}})
return
}
dir := c.Query("dir")
entries, err := svc.StorageCfg.CloudList(c.Request.Context(), typ, dir)
if err != nil {
@@ -35,10 +31,62 @@ func cloudListHandler(svc *service.Container) gin.HandlerFunc {
}
}
func cloudMkdirHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
if !service.IsAdminCloudConfigurable(typ) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
return
}
var in struct {
Dir string `json:"dir"`
Name string `json:"name" binding:"required"`
}
if err := c.ShouldBindJSON(&in); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
entry, err := svc.StorageCfg.CloudMkdir(c.Request.Context(), typ, in.Dir, in.Name)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"entry": entry})
}
}
func cloudRenameHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
if !service.IsAdminCloudConfigurable(typ) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
return
}
var in struct {
Ref string `json:"ref" binding:"required"`
Name string `json:"name" binding:"required"`
}
if err := c.ShouldBindJSON(&in); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
entry, err := svc.StorageCfg.CloudRename(c.Request.Context(), typ, in.Ref, in.Name)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"entry": entry})
}
}
// cloudImportHandler turns a cloud file into a playable 302-backed media item.
func cloudImportHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
if !service.IsAdminCloudConfigurable(typ) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
return
}
var in struct {
Ref string `json:"ref" binding:"required"`
Name string `json:"name"`
@@ -63,6 +111,10 @@ func cloudImportHandler(svc *service.Container) gin.HandlerFunc {
func cloudMountHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
if !service.IsAdminCloudConfigurable(typ) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
return
}
var in struct {
Dir string `json:"dir"`
DirPath string `json:"dir_path"`
@@ -284,209 +336,3 @@ func cloud115QRPollHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusOK, st)
}
}
// cloudPlayHandler resolves a cloud file to its direct link and either issues a
// 302 redirect (true offload — host does not stream the bytes) or, when the
// provider requires authenticated headers, reverse-proxies the response.
func cloudPlayHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
ref := c.Query("ref")
if ref == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "ref required"})
return
}
if !enforceScopedCloudPlaybackToken(c, svc, typ, ref) {
return
}
serveCloudResolvedLink(svc, c, typ, ref)
}
}
func serveCloudResolvedLink(svc *service.Container, c *gin.Context, typ, ref string) {
if isCloudImageRef(ref) && svc != nil && svc.ImageProxy != nil {
if svc.ImageProxy.ServeCloudCached(c.Writer, c.Request, typ+":"+ref) {
return
}
}
if svc == nil || svc.StorageCfg == nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "cloud storage service unavailable"})
return
}
resolveStart := time.Now()
link, err := svc.StorageCfg.CloudResolve(c.Request.Context(), typ, ref, c.Request.UserAgent())
resolveDur := time.Since(resolveStart)
if err != nil {
logCloudPlayback(svc, "cloud playback resolve failed",
append(cloudPlaybackLogFields(typ, ref, nil, resolveDur), zap.Error(err))...)
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
if isCloudImageRef(ref) && svc.ImageProxy != nil {
if err := svc.ImageProxy.ServeCloudResolved(c.Request.Context(), c.Writer, c.Request, typ+":"+ref, link); err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
}
return
}
if isCloudImageRef(ref) {
c.Header("Cache-Control", "public, max-age=2592000, immutable")
}
if !link.Proxy {
// Pure offload: send the client straight to the cloud CDN.
setRedirectNoStoreHeaders(c)
logCloudPlayback(svc, "cloud playback redirect",
append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
zap.String("mode", "redirect"),
zap.Int("status", http.StatusFound),
zap.String("method", c.Request.Method),
zap.String("range", c.GetHeader("Range")),
)...)
c.Redirect(http.StatusFound, link.URL)
return
}
// Proxy mode: the direct link needs auth headers the browser cannot
// carry. Stream through with Range forwarding.
method := c.Request.Method
if method == "" {
method = http.MethodGet
}
req, err := http.NewRequestWithContext(c.Request.Context(), method, link.URL, nil)
if err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
for k, v := range link.Headers {
req.Header.Set(k, v)
}
if rng := c.GetHeader("Range"); rng != "" {
req.Header.Set("Range", rng)
}
if accept := c.GetHeader("Accept"); accept != "" {
req.Header.Set("Accept", accept)
}
if c.GetHeader("Accept-Encoding") == "" {
req.Header.Set("Accept-Encoding", "identity")
}
upstreamStart := time.Now()
resp, err := http.DefaultClient.Do(req)
upstreamHeaderDur := time.Since(upstreamStart)
if err != nil {
logCloudPlayback(svc, "cloud playback proxy upstream failed",
append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
zap.String("mode", "proxy"),
zap.String("method", method),
zap.String("range", c.GetHeader("Range")),
zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
zap.Error(err),
)...)
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
defer resp.Body.Close()
for _, h := range []string{"Content-Type", "Content-Length", "Content-Range", "Accept-Ranges", "ETag", "Last-Modified"} {
if v := resp.Header.Get(h); v != "" {
c.Header(h, v)
}
}
if c.Writer.Header().Get("Accept-Ranges") == "" {
c.Header("Accept-Ranges", "bytes")
}
if resp.StatusCode >= 400 {
c.Header("Cache-Control", "no-store")
}
c.Status(resp.StatusCode)
var copied int64
var copyErr error
streamStart := time.Now()
if c.Request.Method != http.MethodHead {
copied, copyErr = io.Copy(c.Writer, resp.Body)
}
fields := append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
zap.String("mode", "proxy"),
zap.String("method", method),
zap.String("range", c.GetHeader("Range")),
zap.Int("status", resp.StatusCode),
zap.String("content_range", resp.Header.Get("Content-Range")),
zap.String("content_length", resp.Header.Get("Content-Length")),
zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
zap.Int64("stream_ms", durationMilliseconds(time.Since(streamStart))),
zap.Int64("total_ms", durationMilliseconds(time.Since(resolveStart))),
zap.Int64("bytes", copied),
)
if copyErr != nil {
logCloudPlayback(svc, "cloud playback proxy copy failed", append(fields, zap.Error(copyErr))...)
return
}
logCloudPlayback(svc, "cloud playback proxy finished", fields...)
}
func isCloudImageRef(ref string) bool {
ref = strings.ToLower(strings.TrimSpace(ref))
for _, suffix := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp"} {
if strings.HasSuffix(ref, suffix) {
return true
}
}
return false
}
func logCloudPlayback(svc *service.Container, msg string, fields ...zap.Field) {
if svc == nil || svc.Log == nil {
return
}
svc.Log.Info(msg, fields...)
}
func cloudPlaybackLogFields(typ, ref string, link *cloud.DirectLink, resolveDur time.Duration) []zap.Field {
refHash, refExt := cloudPlaybackRefFingerprint(ref)
fields := []zap.Field{
zap.String("provider", strings.TrimSpace(typ)),
zap.String("ref_hash", refHash),
zap.String("ref_ext", refExt),
zap.Int64("resolve_ms", durationMilliseconds(resolveDur)),
}
if link != nil {
fields = append(fields,
zap.String("target_host", cloudPlaybackLinkHost(link.URL)),
zap.Bool("headers_required", len(link.Headers) > 0),
zap.Strings("header_names", cloudPlaybackHeaderNames(link.Headers)),
)
}
return fields
}
func cloudPlaybackRefFingerprint(ref string) (string, string) {
ref = strings.TrimSpace(ref)
sum := sha256.Sum256([]byte(ref))
ext := strings.ToLower(path.Ext(strings.Trim(strings.ReplaceAll(ref, "\\", "/"), "/")))
return hex.EncodeToString(sum[:])[:12], ext
}
func cloudPlaybackLinkHost(raw string) string {
u, err := url.Parse(strings.TrimSpace(raw))
if err != nil || u.Host == "" {
return ""
}
return u.Host
}
func cloudPlaybackHeaderNames(headers map[string]string) []string {
if len(headers) == 0 {
return nil
}
out := make([]string, 0, len(headers))
for key := range headers {
if key = strings.TrimSpace(key); key != "" {
out = append(out, key)
}
}
sort.Strings(out)
return out
}
func durationMilliseconds(d time.Duration) int64 {
if d <= 0 {
return 0
}
return d.Milliseconds()
}
+289
View File
@@ -0,0 +1,289 @@
package handler
import (
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"net/url"
"path"
"sort"
"strings"
"time"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/service"
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
)
type cloudPlaybackRequest struct {
svc *service.Container
c *gin.Context
typ string
ref string
link *cloud.DirectLink
resolveStart time.Time
resolveDur time.Duration
}
// cloudPlayHandler resolves a cloud file to its direct link and either issues a
// 302 redirect (true offload — host does not stream the bytes) or, when the
// provider requires authenticated headers, reverse-proxies the response.
func cloudPlayHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
ref := c.Query("ref")
if !service.IsAdminCloudConfigurable(typ) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
return
}
if ref == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "ref required"})
return
}
if !enforceScopedCloudPlaybackToken(c, svc, typ, ref) {
return
}
serveCloudResolvedLink(svc, c, typ, ref)
}
}
func serveCloudResolvedLink(svc *service.Container, c *gin.Context, typ, ref string) {
if isCloudImageRef(ref) && svc != nil && svc.ImageProxy != nil {
if svc.ImageProxy.ServeCloudCached(c.Writer, c.Request, typ+":"+ref) {
return
}
}
if svc == nil || svc.StorageCfg == nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "cloud storage service unavailable"})
return
}
resolveStart := time.Now()
link, err := svc.StorageCfg.CloudResolve(c.Request.Context(), typ, ref, c.Request.UserAgent())
resolveDur := time.Since(resolveStart)
if err != nil {
logCloudPlayback(svc, "cloud playback resolve failed",
append(cloudPlaybackLogFields(typ, ref, nil, resolveDur), zap.Error(err))...)
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
if isCloudImageRef(ref) && svc.ImageProxy != nil {
if err := svc.ImageProxy.ServeCloudResolved(c.Request.Context(), c.Writer, c.Request, typ+":"+ref, link); err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
}
return
}
if isCloudImageRef(ref) {
c.Header("Cache-Control", "public, max-age=2592000, immutable")
}
if !link.Proxy {
// Pure offload: send the client straight to the cloud CDN.
setRedirectNoStoreHeaders(c)
logCloudPlayback(svc, "cloud playback redirect",
append(cloudPlaybackLogFields(typ, ref, link, resolveDur),
zap.String("mode", "redirect"),
zap.Int("status", http.StatusFound),
zap.String("method", c.Request.Method),
zap.String("range", c.GetHeader("Range")),
)...)
c.Redirect(http.StatusFound, link.URL)
return
}
proxyCloudResolvedLink(cloudPlaybackRequest{
svc: svc,
c: c,
typ: typ,
ref: ref,
link: link,
resolveStart: resolveStart,
resolveDur: resolveDur,
})
}
func proxyCloudResolvedLink(playback cloudPlaybackRequest) {
c := playback.c
clientMethod := playback.c.Request.Method
if clientMethod == "" {
clientMethod = http.MethodGet
}
upstreamMethod := clientMethod
if upstreamMethod == http.MethodHead {
upstreamMethod = http.MethodGet
}
req, err := http.NewRequestWithContext(c.Request.Context(), upstreamMethod, playback.link.URL, nil)
if err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
for k, v := range playback.link.Headers {
req.Header.Set(k, v)
}
if rng := c.GetHeader("Range"); rng != "" {
req.Header.Set("Range", rng)
} else if clientMethod == http.MethodHead {
req.Header.Set("Range", "bytes=0-0")
}
if accept := c.GetHeader("Accept"); accept != "" {
req.Header.Set("Accept", accept)
}
if c.GetHeader("Accept-Encoding") == "" {
req.Header.Set("Accept-Encoding", "identity")
}
upstreamStart := time.Now()
resp, err := http.DefaultClient.Do(req)
upstreamHeaderDur := time.Since(upstreamStart)
if err != nil {
logCloudPlayback(playback.svc, "cloud playback proxy upstream failed",
append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
zap.String("mode", "proxy"),
zap.String("method", clientMethod),
zap.String("upstream_method", upstreamMethod),
zap.String("range", c.GetHeader("Range")),
zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
zap.Error(err),
)...)
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
defer resp.Body.Close()
for _, h := range []string{"Content-Type", "Content-Length", "Content-Range", "Accept-Ranges", "ETag", "Last-Modified"} {
if v := resp.Header.Get(h); v != "" {
c.Header(h, v)
}
}
if c.Writer.Header().Get("Accept-Ranges") == "" {
c.Header("Accept-Ranges", "bytes")
}
if resp.StatusCode >= 400 {
handleCloudProxyError(playback, req, resp, clientMethod, upstreamMethod, upstreamHeaderDur)
return
}
streamCloudProxyResponse(playback, req, resp, clientMethod, upstreamMethod, upstreamHeaderDur)
}
func handleCloudProxyError(playback cloudPlaybackRequest, req *http.Request, resp *http.Response, clientMethod, upstreamMethod string, upstreamHeaderDur time.Duration) {
c := playback.c
c.Header("Cache-Control", "no-store")
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
fields := append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
zap.String("mode", "proxy"),
zap.String("method", clientMethod),
zap.String("upstream_method", upstreamMethod),
zap.String("range", c.GetHeader("Range")),
zap.String("upstream_range", req.Header.Get("Range")),
zap.Int("status", resp.StatusCode),
zap.String("content_range", resp.Header.Get("Content-Range")),
zap.String("content_length", resp.Header.Get("Content-Length")),
zap.String("upstream_error_body", strings.TrimSpace(string(body))),
zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
zap.Int64("total_ms", durationMilliseconds(time.Since(playback.resolveStart))),
)
logCloudPlayback(playback.svc, "cloud playback proxy upstream returned error", fields...)
c.Status(resp.StatusCode)
if clientMethod != http.MethodHead && len(body) > 0 {
_, _ = c.Writer.Write(body)
}
}
func streamCloudProxyResponse(playback cloudPlaybackRequest, req *http.Request, resp *http.Response, clientMethod, upstreamMethod string, upstreamHeaderDur time.Duration) {
c := playback.c
c.Status(resp.StatusCode)
var copied int64
var copyErr error
streamStart := time.Now()
if c.Request.Method != http.MethodHead {
copied, copyErr = io.Copy(c.Writer, resp.Body)
}
fields := append(cloudPlaybackLogFields(playback.typ, playback.ref, playback.link, playback.resolveDur),
zap.String("mode", "proxy"),
zap.String("method", clientMethod),
zap.String("upstream_method", upstreamMethod),
zap.String("range", c.GetHeader("Range")),
zap.String("upstream_range", req.Header.Get("Range")),
zap.Int("status", resp.StatusCode),
zap.String("content_range", resp.Header.Get("Content-Range")),
zap.String("content_length", resp.Header.Get("Content-Length")),
zap.Int64("upstream_header_ms", durationMilliseconds(upstreamHeaderDur)),
zap.Int64("stream_ms", durationMilliseconds(time.Since(streamStart))),
zap.Int64("total_ms", durationMilliseconds(time.Since(playback.resolveStart))),
zap.Int64("bytes", copied),
)
if copyErr != nil {
logCloudPlayback(playback.svc, "cloud playback proxy copy failed", append(fields, zap.Error(copyErr))...)
return
}
logCloudPlayback(playback.svc, "cloud playback proxy finished", fields...)
}
func isCloudImageRef(ref string) bool {
ref = strings.ToLower(strings.TrimSpace(ref))
for _, suffix := range []string{".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp", ".tbn"} {
if strings.HasSuffix(ref, suffix) {
return true
}
}
return false
}
func logCloudPlayback(svc *service.Container, msg string, fields ...zap.Field) {
if svc == nil || svc.Log == nil {
return
}
svc.Log.Info(msg, fields...)
}
func cloudPlaybackLogFields(typ, ref string, link *cloud.DirectLink, resolveDur time.Duration) []zap.Field {
refHash, refExt := cloudPlaybackRefFingerprint(ref)
fields := []zap.Field{
zap.String("provider", strings.TrimSpace(typ)),
zap.String("ref_hash", refHash),
zap.String("ref_ext", refExt),
zap.Int64("resolve_ms", durationMilliseconds(resolveDur)),
}
if link != nil {
fields = append(fields,
zap.String("target_host", cloudPlaybackLinkHost(link.URL)),
zap.Bool("headers_required", len(link.Headers) > 0),
zap.Strings("header_names", cloudPlaybackHeaderNames(link.Headers)),
)
}
return fields
}
func cloudPlaybackRefFingerprint(ref string) (string, string) {
ref = strings.TrimSpace(ref)
sum := sha256.Sum256([]byte(ref))
ext := strings.ToLower(path.Ext(strings.Trim(strings.ReplaceAll(ref, "\\", "/"), "/")))
return hex.EncodeToString(sum[:])[:12], ext
}
func cloudPlaybackLinkHost(raw string) string {
u, err := url.Parse(strings.TrimSpace(raw))
if err != nil || u.Host == "" {
return ""
}
return u.Host
}
func cloudPlaybackHeaderNames(headers map[string]string) []string {
if len(headers) == 0 {
return nil
}
out := make([]string, 0, len(headers))
for key := range headers {
if key = strings.TrimSpace(key); key != "" {
out = append(out, key)
}
}
sort.Strings(out)
return out
}
func durationMilliseconds(d time.Duration) int64 {
if d <= 0 {
return 0
}
return d.Milliseconds()
}
+104
View File
@@ -1,8 +1,17 @@
package handler
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/service"
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
)
func TestCloudMountLibraryNameDefaultsToDirectoryBaseName(t *testing.T) {
@@ -48,3 +57,98 @@ func TestCloudPlaybackDiagnosticsDoNotExposeRawRefOrURL(t *testing.T) {
t.Fatalf("header names = %q", got)
}
}
func TestAdminCloudHandlersRejectQuarkBrowsing(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.GET("/admin/cloud/:type/list", cloudListHandler(nil))
req := httptest.NewRequest(http.MethodGet, "/admin/cloud/quark/list?dir=0", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
}
if !strings.Contains(w.Body.String(), "unsupported cloud provider") {
t.Fatalf("body = %s, want unsupported cloud provider", w.Body.String())
}
}
func TestCloudPlayRejectsQuarkProvider(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.GET("/api/cloud/play/:type", cloudPlayHandler(nil))
req := httptest.NewRequest(http.MethodGet, "/api/cloud/play/quark?ref=file-1", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
}
if !strings.Contains(w.Body.String(), "unsupported cloud provider") {
t.Fatalf("body = %s, want unsupported cloud provider", w.Body.String())
}
}
func TestCloudArtworkProxyServesCachedImageWithoutCloudResolve(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "image/jpeg")
_, _ = w.Write([]byte("cached-cloud-poster"))
}))
defer upstream.Close()
imageProxy := service.NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, zap.NewNop())
stableKey := "openlist:/Anime/JianLai/poster.jpg"
if err := imageProxy.PrefetchCloudResolved(t.Context(), stableKey, &cloud.DirectLink{URL: upstream.URL + "/poster.jpg"}); err != nil {
t.Fatal(err)
}
router := gin.New()
router.GET("/api/img/cloud/:type", cloudArtworkProxyHandler(&service.Container{ImageProxy: imageProxy}))
req := httptest.NewRequest(http.MethodGet, "/api/img/cloud/openlist?ref=%2FAnime%2FJianLai%2Fposter.jpg", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s, want 200", w.Code, w.Body.String())
}
if got := w.Body.String(); got != "cached-cloud-poster" {
t.Fatalf("body = %q, want cached poster", got)
}
if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "max-age=2592000") {
t.Fatalf("cache-control = %q, want long static cache", got)
}
}
func TestCloudArtworkProxyAcceptsCachedTBNImage(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "image/jpeg")
_, _ = w.Write([]byte("cached-tbn-poster"))
}))
defer upstream.Close()
imageProxy := service.NewImageProxy(&config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}, zap.NewNop())
stableKey := "openlist:/Movies/Movie.tbn"
if err := imageProxy.PrefetchCloudResolved(t.Context(), stableKey, &cloud.DirectLink{URL: upstream.URL + "/Movie.tbn"}); err != nil {
t.Fatal(err)
}
router := gin.New()
router.GET("/api/img/cloud/:type", cloudArtworkProxyHandler(&service.Container{ImageProxy: imageProxy}))
req := httptest.NewRequest(http.MethodGet, "/api/img/cloud/openlist?ref=%2FMovies%2FMovie.tbn", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s, want 200", w.Code, w.Body.String())
}
if got := w.Body.String(); got != "cached-tbn-poster" {
t.Fatalf("body = %q, want cached tbn poster", got)
}
}
+2
View File
@@ -26,6 +26,7 @@ func trendingHandler(svc *service.Container) gin.HandlerFunc {
if items == nil {
items = []service.Match{}
}
svc.Discover.WarmMatchArtwork(items)
c.JSON(http.StatusOK, gin.H{"items": items})
}
}
@@ -41,6 +42,7 @@ func popularHandler(svc *service.Container) gin.HandlerFunc {
if items == nil {
items = []service.Match{}
}
svc.Discover.WarmMatchArtwork(items)
c.JSON(http.StatusOK, gin.H{"items": items})
}
}
+59 -14
View File
@@ -16,24 +16,37 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/service"
)
type discoverSectionDef struct {
Key string
Label string
Provider string
}
var discoverSectionCatalog = []discoverSectionDef{
{Key: "tmdb_trending_day", Label: "TMDb 今日趋势", Provider: "tmdb"},
{Key: "tmdb_trending_week", Label: "TMDb 本周热门", Provider: "tmdb"},
{Key: "tmdb_popular_movie", Label: "TMDb 热门电影", Provider: "tmdb"},
{Key: "tmdb_popular_tv", Label: "TMDb 热门剧集", Provider: "tmdb"},
{Key: "tmdb_top_rated_movie", Label: "TMDb 高分电影", Provider: "tmdb"},
{Key: "douban_hot_movie", Label: "豆瓣热门电影", Provider: "douban"},
{Key: "douban_hot_tv", Label: "豆瓣热门剧集", Provider: "douban"},
{Key: "douban_top_movie", Label: "豆瓣高分电影", Provider: "douban"},
{Key: "bangumi_calendar", Label: "Bangumi 每日放送", Provider: "bangumi"},
}
// discoverSectionsHandler returns the catalog of sections the UI can
// pick from. The names match the upstream Vue UI so existing settings
// keep working.
func discoverSectionsHandler(_ *service.Container) gin.HandlerFunc {
func discoverSectionsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"sections": []gin.H{
{"key": "tmdb_trending_day", "label": "TMDb 今日趋势", "provider": "tmdb"},
{"key": "tmdb_trending_week", "label": "TMDb 本周热门", "provider": "tmdb"},
{"key": "tmdb_popular_movie", "label": "TMDb 热门电影", "provider": "tmdb"},
{"key": "tmdb_popular_tv", "label": "TMDb 热门剧集", "provider": "tmdb"},
{"key": "tmdb_top_rated_movie", "label": "TMDb 高分电影", "provider": "tmdb"},
{"key": "douban_hot_movie", "label": "豆瓣热门电影", "provider": "douban"},
{"key": "douban_hot_tv", "label": "豆瓣热门剧集", "provider": "douban"},
{"key": "douban_top_movie", "label": "豆瓣高分电影", "provider": "douban"},
{"key": "bangumi_calendar", "label": "Bangumi 每日放送", "provider": "bangumi"},
},
})
sections := make([]gin.H, 0, len(discoverSectionCatalog))
for _, section := range discoverSectionCatalog {
if !discoverProviderEnabled(c.Request.Context(), svc, section.Provider) {
continue
}
sections = append(sections, gin.H{"key": section.Key, "label": section.Label, "provider": section.Provider})
}
c.JSON(http.StatusOK, gin.H{"sections": sections})
}
}
@@ -45,19 +58,51 @@ func discoverFeedHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
keys := strings.Split(c.DefaultQuery("sections", "tmdb_trending_day,tmdb_popular_movie,douban_hot_movie,bangumi_calendar"), ",")
out := gin.H{}
artworkItems := []service.ExternalMediaResult{}
for _, raw := range keys {
k := strings.TrimSpace(raw)
if provider := discoverSectionProvider(k); provider != "" && !discoverProviderEnabled(c.Request.Context(), svc, provider) {
out[k] = []service.ExternalMediaResult{}
continue
}
items, err := discoverSectionItems(c.Request.Context(), svc, k)
if err != nil {
svc.Log.Debug("discover fetch failed")
items = nil
}
artworkItems = append(artworkItems, items...)
out[k] = items
}
svc.Discover.WarmExternalArtwork(artworkItems)
c.JSON(http.StatusOK, out)
}
}
func discoverSectionProvider(key string) string {
for _, section := range discoverSectionCatalog {
if section.Key == key {
return section.Provider
}
}
switch key {
case "trending_day", "trending_week", "popular_movie", "popular_tv", "top_rated_movie", "upcoming_movie":
return "tmdb"
default:
return ""
}
}
func discoverProviderEnabled(ctx context.Context, svc *service.Container, provider string) bool {
if svc == nil || svc.APIConfig == nil || strings.TrimSpace(provider) == "" {
return true
}
cfg, err := svc.APIConfig.Get(ctx, provider)
if err != nil || cfg == nil {
return true
}
return cfg.Enabled
}
func discoverSectionItems(ctx context.Context, svc *service.Container, k string) ([]service.ExternalMediaResult, error) {
switch k {
case "tmdb_trending_day", "tmdb_trending_week", "tmdb_popular_movie", "tmdb_popular_tv", "tmdb_top_rated_movie",
+37
View File
@@ -0,0 +1,37 @@
package handler
import (
"testing"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func TestDiscoverProviderEnabledHonorsAPIConfigToggle(t *testing.T) {
db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
t.Fatal(err)
}
repos := repository.New(db)
apiConfig := service.NewAPIConfigService(zap.NewNop(), repos, service.NewCryptoService("", zap.NewNop()))
enabled := false
if _, err := apiConfig.Update(t.Context(), "douban", service.APIConfigPatch{Enabled: &enabled}); err != nil {
t.Fatal(err)
}
svc := &service.Container{APIConfig: apiConfig}
if discoverProviderEnabled(t.Context(), svc, "douban") {
t.Fatal("disabled API config should disable discover provider")
}
if !discoverProviderEnabled(t.Context(), svc, "missing-provider") {
t.Fatal("missing API config should keep discover provider available")
}
}
+2 -1831
View File
File diff suppressed because it is too large Load Diff
+241
View File
@@ -0,0 +1,241 @@
package handler
import (
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
)
// embyError 返回 Emby 风格的错误(顶层 Code/Message)。
func embyError(c *gin.Context, status int, msg string) {
c.JSON(status, gin.H{"Code": status, "Message": msg})
}
// embyUserID 从中间件中获取 user id。Emby auth middleware 写入 CtxUserID。
func embyUserID(c *gin.Context) string {
if uid, ok := c.Get(middleware.CtxUserID); ok {
if s, ok := uid.(string); ok {
return s
}
}
return ""
}
const embyCompatSessionTTL = 30 * time.Minute
type embyCompatSession struct {
token string
expiresAt time.Time
}
var embyCompatSessions = struct {
sync.RWMutex
items map[string]embyCompatSession
}{items: map[string]embyCompatSession{}}
func embyAuthRequiredWithSessionFallback(secret string) gin.HandlerFunc {
required := middleware.EmbyAuthRequired(secret)
return func(c *gin.Context) {
if embyRequestToken(c) == "" {
if token := embyCompatSessionToken(c); token != "" {
c.Request.Header.Set("X-Emby-Token", token)
}
}
required(c)
}
}
func embyRememberCompatSession(c *gin.Context, token string) {
token = strings.TrimSpace(token)
if token == "" {
return
}
keys := embyCompatSessionKeys(c)
if len(keys) == 0 {
return
}
expiresAt := time.Now().Add(embyCompatSessionTTL)
embyCompatSessions.Lock()
defer embyCompatSessions.Unlock()
if len(embyCompatSessions.items) > 1000 {
now := time.Now()
for key, session := range embyCompatSessions.items {
if now.After(session.expiresAt) {
delete(embyCompatSessions.items, key)
}
}
if len(embyCompatSessions.items) > 1000 {
embyCompatSessions.items = map[string]embyCompatSession{}
}
}
for _, key := range keys {
embyCompatSessions.items[key] = embyCompatSession{token: token, expiresAt: expiresAt}
}
}
func embyCompatSessionToken(c *gin.Context) string {
keys := embyCompatSessionKeys(c)
if len(keys) == 0 {
return ""
}
now := time.Now()
embyCompatSessions.RLock()
defer embyCompatSessions.RUnlock()
for _, key := range keys {
session, ok := embyCompatSessions.items[key]
if ok && now.Before(session.expiresAt) {
return session.token
}
}
return ""
}
func embyCompatSessionKeys(c *gin.Context) []string {
if c == nil {
return nil
}
ip := strings.TrimSpace(c.ClientIP())
if ip == "" {
return nil
}
keys := []string{}
add := func(kind, value string) {
value = strings.TrimSpace(value)
if value != "" {
keys = append(keys, ip+"\x00"+kind+"\x00"+value)
}
}
add("device", firstHeaderValue(c, "X-Emby-Device-Id", "X-Emby-DeviceId", "X-MediaBrowser-Device-Id", "X-MediaBrowser-DeviceId"))
add("ua", c.GetHeader("User-Agent"))
return keys
}
func firstHeaderValue(c *gin.Context, names ...string) string {
for _, name := range names {
if value := strings.TrimSpace(c.GetHeader(name)); value != "" {
return value
}
}
return ""
}
type embyClientInfo struct {
DeviceID string
DeviceName string
Client string
}
func embyClientInfoFromRequest(c *gin.Context) embyClientInfo {
auth := parseMediaBrowserAuthorization(firstHeaderValue(c,
"X-Emby-Authorization",
"X-MediaBrowser-Authorization",
"Authorization",
))
info := embyClientInfo{
DeviceID: firstNonEmptyHeaderString(
firstHeaderValue(c, "X-Emby-Device-Id", "X-Emby-DeviceId", "X-MediaBrowser-Device-Id", "X-MediaBrowser-DeviceId"),
auth["DeviceId"],
auth["DeviceID"],
),
DeviceName: firstNonEmptyHeaderString(
firstHeaderValue(c, "X-Emby-Device-Name", "X-Emby-DeviceName", "X-MediaBrowser-Device-Name", "X-MediaBrowser-DeviceName"),
auth["Device"],
),
Client: firstNonEmptyHeaderString(
firstHeaderValue(c, "X-Emby-Client", "X-MediaBrowser-Client"),
auth["Client"],
),
}
ua := strings.TrimSpace(c.GetHeader("User-Agent"))
if info.Client == "" {
info.Client = embyClientFromUserAgent(ua)
}
if info.DeviceName == "" {
info.DeviceName = embyDeviceFromUserAgent(ua)
}
return info
}
func parseMediaBrowserAuthorization(raw string) map[string]string {
out := map[string]string{}
raw = strings.TrimSpace(raw)
if raw == "" {
return out
}
for _, prefix := range []string{"MediaBrowser ", "Emby "} {
if strings.HasPrefix(raw, prefix) {
raw = strings.TrimSpace(strings.TrimPrefix(raw, prefix))
break
}
}
for _, part := range strings.Split(raw, ",") {
key, value, ok := strings.Cut(strings.TrimSpace(part), "=")
if !ok {
continue
}
key = strings.TrimSpace(key)
value = strings.Trim(strings.TrimSpace(value), `"`)
if key != "" && value != "" {
out[key] = value
}
}
return out
}
func firstNonEmptyHeaderString(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
func embyClientFromUserAgent(ua string) string {
ua = strings.TrimSpace(ua)
lower := strings.ToLower(ua)
switch {
case strings.Contains(lower, "infuse"):
return "Infuse"
case strings.Contains(lower, "emby"):
return "Emby"
case strings.Contains(lower, "jellyfin"):
return "Jellyfin"
case strings.Contains(lower, "yamby"):
return "Yamby"
case strings.Contains(lower, "vidhub"):
return "VidHub"
case strings.Contains(lower, "hills"):
return "Hills"
default:
return ua
}
}
func embyDeviceFromUserAgent(ua string) string {
lower := strings.ToLower(strings.TrimSpace(ua))
switch {
case strings.Contains(lower, "android"):
return "Android"
case strings.Contains(lower, "iphone"):
return "iPhone"
case strings.Contains(lower, "ipad"):
return "iPad"
case strings.Contains(lower, "ios"):
return "iOS"
case strings.Contains(lower, "windows"):
return "Windows PC"
case strings.Contains(lower, "macintosh") || strings.Contains(lower, "mac os"):
return "Mac"
case strings.Contains(lower, "linux"):
return "Linux PC"
case strings.Contains(lower, "appletv") || strings.Contains(lower, "apple tv"):
return "Apple TV"
default:
return ""
}
}
+174
View File
@@ -0,0 +1,174 @@
package handler
import (
"bytes"
"encoding/json"
"errors"
"io"
"net/url"
"strings"
"github.com/gin-gonic/gin"
)
type embyAuthByNameReq struct {
Username string `json:"Username"`
Pw string `json:"Pw"`
Password string `json:"Password"`
PasswordMd5 string `json:"PasswordMd5"`
PasswordSha1 string `json:"PasswordSha1"`
}
func parseEmbyAuthByNameReq(c *gin.Context) (embyAuthByNameReq, error) {
req := embyAuthByNameReq{}
if strings.Contains(strings.ToLower(c.GetHeader("Content-Type")), "json") {
var body map[string]any
if err := c.ShouldBindJSON(&body); err != nil && !errors.Is(err, io.EOF) {
return req, err
}
fillEmbyAuthFromMap(&req, body)
}
if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") {
_ = c.Request.ParseForm()
if req.Username == "" {
req.Username = firstFormValue(c, "Username", "username", "Name", "name")
}
if req.Pw == "" {
req.Pw = firstFormValue(c, "Pw", "pw")
}
if req.Password == "" {
req.Password = firstFormValue(c, "Password", "password")
}
if req.PasswordMd5 == "" {
req.PasswordMd5 = firstFormValue(c, "PasswordMd5", "passwordMd5", "password_md5")
}
if req.PasswordSha1 == "" {
req.PasswordSha1 = firstFormValue(c, "PasswordSha1", "passwordSha1", "password_sha1")
}
}
if req.Username == "" {
req.Username = firstQueryValue(c, "Username", "username", "Name", "name")
}
if req.Pw == "" {
req.Pw = firstQueryValue(c, "Pw", "pw")
}
if req.Password == "" {
req.Password = firstQueryValue(c, "Password", "password")
}
if req.PasswordMd5 == "" {
req.PasswordMd5 = firstQueryValue(c, "PasswordMd5", "passwordMd5", "password_md5")
}
if req.PasswordSha1 == "" {
req.PasswordSha1 = firstQueryValue(c, "PasswordSha1", "passwordSha1", "password_sha1")
}
if req.Username == "" || (req.Pw == "" && req.Password == "" && req.PasswordMd5 == "" && req.PasswordSha1 == "") {
fillEmbyAuthFromRawBody(c, &req)
}
return req, nil
}
func fillEmbyAuthFromMap(req *embyAuthByNameReq, body map[string]any) {
if req.Username == "" {
req.Username = firstStringFromMap(body, "Username", "username", "UserName", "userName", "Name", "name", "LoginName", "loginName")
}
if req.Pw == "" {
req.Pw = firstStringFromMap(body, "Pw", "pw", "PW")
}
if req.Password == "" {
req.Password = firstStringFromMap(body, "Password", "password", "Pass", "pass", "Pwd", "pwd")
}
if req.PasswordMd5 == "" {
req.PasswordMd5 = firstStringFromMap(body, "PasswordMd5", "passwordMd5", "password_md5")
}
if req.PasswordSha1 == "" {
req.PasswordSha1 = firstStringFromMap(body, "PasswordSha1", "passwordSha1", "password_sha1")
}
}
func fillEmbyAuthFromRawBody(c *gin.Context, req *embyAuthByNameReq) {
if c.Request == nil || c.Request.Body == nil {
return
}
raw, err := io.ReadAll(io.LimitReader(c.Request.Body, 1<<20))
if err != nil {
return
}
c.Request.Body = io.NopCloser(bytes.NewReader(raw))
raw = bytes.TrimSpace(raw)
if len(raw) == 0 {
return
}
if bytes.HasPrefix(raw, []byte("{")) {
var body map[string]any
if err := json.Unmarshal(raw, &body); err == nil {
fillEmbyAuthFromMap(req, body)
}
return
}
if values, err := url.ParseQuery(string(raw)); err == nil {
fillEmbyAuthFromValues(req, values)
}
}
func fillEmbyAuthFromValues(req *embyAuthByNameReq, values url.Values) {
if req.Username == "" {
req.Username = firstValue(values, "Username", "username", "UserName", "userName", "Name", "name", "LoginName", "loginName")
}
if req.Pw == "" {
req.Pw = firstValue(values, "Pw", "pw", "PW")
}
if req.Password == "" {
req.Password = firstValue(values, "Password", "password", "Pass", "pass", "Pwd", "pwd")
}
if req.PasswordMd5 == "" {
req.PasswordMd5 = firstValue(values, "PasswordMd5", "passwordMd5", "password_md5")
}
if req.PasswordSha1 == "" {
req.PasswordSha1 = firstValue(values, "PasswordSha1", "passwordSha1", "password_sha1")
}
}
func firstValue(values url.Values, keys ...string) string {
for _, key := range keys {
if value := strings.TrimSpace(values.Get(key)); value != "" {
return value
}
}
return ""
}
func firstStringFromMap(body map[string]any, keys ...string) string {
if len(body) == 0 {
return ""
}
for _, key := range keys {
if value, ok := body[key]; ok {
if s, ok := value.(string); ok {
return strings.TrimSpace(s)
}
}
}
return ""
}
func firstFormValue(c *gin.Context, keys ...string) string {
for _, key := range keys {
if values, ok := c.Request.PostForm[key]; ok && len(values) > 0 {
if value := strings.TrimSpace(values[0]); value != "" {
return value
}
}
}
return ""
}
func firstQueryValue(c *gin.Context, keys ...string) string {
for _, key := range keys {
if value := strings.TrimSpace(c.Query(key)); value != "" {
return value
}
}
return ""
}
+168
View File
@@ -0,0 +1,168 @@
package handler
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func TestParseEmbyAuthByNameReqAcceptsLowercaseJSON(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/Users/AuthenticateByName", strings.NewReader(`{"username":"alice","password":"secret"}`))
c.Request.Header.Set("Content-Type", "application/json")
req, err := parseEmbyAuthByNameReq(c)
if err != nil {
t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
}
if req.Username != "alice" || req.Password != "secret" {
t.Fatalf("unexpected request: %#v", req)
}
}
func TestParseEmbyAuthByNameReqAcceptsFormBody(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/Users/AuthenticateByName", strings.NewReader("Username=bob&Pw=secret"))
c.Request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req, err := parseEmbyAuthByNameReq(c)
if err != nil {
t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
}
if req.Username != "bob" || req.Pw != "secret" {
t.Fatalf("unexpected request: %#v", req)
}
}
func TestParseEmbyAuthByNameReqAcceptsJSONWithoutContentType(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/emby/users/authenticatebyname", strings.NewReader(`{"UserName":"carol","PW":"secret"}`))
req, err := parseEmbyAuthByNameReq(c)
if err != nil {
t.Fatalf("parseEmbyAuthByNameReq returned error: %v", err)
}
if req.Username != "carol" || req.Pw != "secret" {
t.Fatalf("unexpected request: %#v", req)
}
}
func TestEmbyAuthenticateByNameAcceptsCaseVariantUsernameAndPath(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(&model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.Setting{}); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
cfg := &config.Config{}
cfg.Secrets.JWTSecret = "test-secret"
log := zap.NewNop()
permissions := service.NewPermissionService(log, repos)
auth := service.NewAuthService(cfg, log, repos, service.NewTokenService(cfg, log, repos), permissions)
if _, _, err := auth.Register(context.Background(), "viewer", "secret-pass"); err != nil {
t.Fatalf("register: %v", err)
}
router := gin.New()
registerEmbyRoutes(router, cfg.Secrets.JWTSecret, &service.Container{
Repo: repos,
Auth: auth,
Emby: service.NewEmbyService(cfg, log, repos),
Audit: service.NewAuditService(log, repos),
})
req := httptest.NewRequest(http.MethodPost, "/emby/users/authenticatebyname", strings.NewReader(`{"Username":"Viewer","Pw":"secret-pass"}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
var payload map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
t.Fatalf("decode response: %v", err)
}
if payload["AccessToken"] == "" {
t.Fatalf("missing AccessToken: %#v", payload)
}
}
func TestEmbyAuthenticateRecordsMediaBrowserClientInfo(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if sqlDB, err := db.DB(); err == nil {
sqlDB.SetMaxOpenConns(1)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
cfg := &config.Config{}
cfg.Secrets.JWTSecret = "test-secret"
log := zap.NewNop()
permissions := service.NewPermissionService(log, repos)
auth := service.NewAuthService(cfg, log, repos, service.NewTokenService(cfg, log, repos), permissions)
if _, _, err := auth.Register(context.Background(), "viewer", "secret-pass"); err != nil {
t.Fatalf("register: %v", err)
}
router := gin.New()
registerEmbyRoutes(router, cfg.Secrets.JWTSecret, &service.Container{
Repo: repos,
Auth: auth,
Emby: service.NewEmbyService(cfg, log, repos),
Device: service.NewDeviceService(log, repos),
Audit: service.NewAuditService(log, repos),
Permissions: permissions,
})
req := httptest.NewRequest(http.MethodPost, "/emby/Users/AuthenticateByName", strings.NewReader(`{"Username":"viewer","Pw":"secret-pass"}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="PC", DeviceId="device-42"`)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
user, err := repos.User.FindByUsername(context.Background(), "viewer")
if err != nil {
t.Fatalf("find user: %v", err)
}
devices, err := repos.UserDevice.ListByUser(context.Background(), user.ID)
if err != nil {
t.Fatalf("list devices: %v", err)
}
if len(devices) != 1 {
t.Fatalf("devices = %#v, want one recorded device", devices)
}
if devices[0].DeviceID != "device-42" || devices[0].DeviceName != "PC" || devices[0].Client != "Infuse" {
t.Fatalf("device info not parsed from MediaBrowser header: %#v", devices[0])
}
}
+72
View File
@@ -0,0 +1,72 @@
package handler
import (
"context"
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
var embyPlaceholderPNG = []byte{
0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52,
0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4,
0x89, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x44, 0x41,
0x54, 0x78, 0x9c, 0x63, 0x50, 0xd1, 0x30, 0xf8,
0x0f, 0x00, 0x02, 0x6c, 0x01, 0x7c, 0x30, 0xed,
0x6e, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x49, 0x45,
0x4e, 0x44, 0xae, 0x42, 0x60, 0x82,
}
// embyItemImageHandler 把 /Items/{id}/Images/Primary 等请求直接输出为图片。
// Emby 客户端缓存图片 URL 时经常不会继续携带 token;如果重定向到受保护的
// /api/img 会变成 401,所以这里复用 ImageProxy 但不再走 /api 路由。
func embyItemImageHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
clearEmbyImageNoStoreHeaders(c)
ctx, cancel := context.WithTimeout(c.Request.Context(), 8*time.Second)
defer cancel()
req := c.Request.WithContext(ctx)
id := c.Param("id")
imgType := strings.ToLower(c.Param("type"))
raw, err := svc.Emby.ImageURL(ctx, id, imgType)
if err != nil || raw == "" {
embyServePlaceholderImage(c)
return
}
if typ, ref, ok := service.ParseCloudArtworkURL(raw); ok {
c.Request = req
serveCloudResolvedLink(svc, c, typ, ref)
return
}
if svc.ImageProxy == nil {
embyServePlaceholderImage(c)
return
}
if err := svc.ImageProxy.Serve(ctx, c.Writer, req, raw); err != nil {
embyServePlaceholderImage(c)
}
}
}
func clearEmbyImageNoStoreHeaders(c *gin.Context) {
c.Writer.Header().Del("Pragma")
c.Writer.Header().Del("Expires")
}
func embyServePlaceholderImage(c *gin.Context) {
c.Header("Content-Type", "image/png")
c.Header("Cache-Control", "public, max-age=3600")
c.Header("Content-Length", strconv.Itoa(len(embyPlaceholderPNG)))
if c.Request.Method == http.MethodHead {
c.Status(http.StatusOK)
return
}
c.Data(http.StatusOK, "image/png", embyPlaceholderPNG)
}
+231
View File
@@ -0,0 +1,231 @@
package handler
import (
"net/http"
"strconv"
"strings"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func parseEmbyItemsParams(c *gin.Context) service.ItemsParams {
limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "50"))
offset, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "StartIndex", "startIndex", "startindex"), "0"))
uid := c.Param("userId")
if uid == "" {
uid = firstQueryValue(c, "UserId", "userId", "userid")
}
if uid == "" {
uid = embyUserID(c)
}
splitOpt := func(s string) []string {
if s == "" {
return nil
}
parts := strings.Split(s, ",")
out := make([]string, 0, len(parts))
for _, p := range parts {
p = strings.TrimSpace(p)
if p != "" {
out = append(out, p)
}
}
return out
}
return service.ItemsParams{
UserID: uid,
ParentID: firstQueryValue(c, "ParentId", "parentId", "parentid"),
IDs: splitOpt(firstQueryValue(c, "Ids", "ids")),
SearchTerm: firstQueryValue(c, "SearchTerm", "searchTerm", "searchterm"),
IncludeItemTypes: splitOpt(firstQueryValue(c, "IncludeItemTypes", "includeItemTypes", "includeitemtypes")),
Filters: splitOpt(firstQueryValue(c, "Filters", "filters")),
Recursive: strings.EqualFold(firstQueryValue(c, "Recursive", "recursive"), "true"),
SortBy: firstQueryValue(c, "SortBy", "sortBy", "sortby"),
SortOrder: firstQueryValue(c, "SortOrder", "sortOrder", "sortorder"),
Limit: limit,
StartIndex: offset,
}
}
func embyFirstNonEmptyString(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
func embyItemsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
out, err := svc.Emby.Items(c.Request.Context(), parseEmbyItemsParams(c))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
embyAttachRequestTokenToMediaSources(c, out)
c.JSON(http.StatusOK, out)
}
}
func embyItemByIDHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
id := c.Param("id")
uid := c.Param("userId")
if uid == "" {
uid = embyUserID(c)
}
out, err := svc.Emby.Item(c.Request.Context(), id, uid)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if out == nil {
embyError(c, http.StatusNotFound, "item not found")
return
}
embyAttachRequestTokenToMediaSources(c, out)
c.JSON(http.StatusOK, out)
}
}
func embyUserItemByIDHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
switch strings.ToLower(c.Param("id")) {
case "latest":
embyLatestItemsHandler(svc)(c)
case "resume":
embyResumeItemsHandler(svc)(c)
default:
embyItemByIDHandler(svc)(c)
}
}
}
func embyLatestItemsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := c.Param("userId")
if uid == "" {
uid = firstQueryValue(c, "UserId", "userId", "userid")
}
if uid == "" {
uid = embyUserID(c)
}
limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "20"))
out, err := svc.Emby.LatestItems(c.Request.Context(), uid, firstQueryValue(c, "ParentId", "parentId", "parentid"), limit)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
embyAttachRequestTokenToMediaSources(c, out)
c.JSON(http.StatusOK, out)
}
}
func embyResumeItemsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := c.Param("userId")
if uid == "" {
uid = firstQueryValue(c, "UserId", "userId", "userid")
}
if uid == "" {
uid = embyUserID(c)
}
limit, _ := strconv.Atoi(embyFirstNonEmptyString(firstQueryValue(c, "Limit", "limit"), "20"))
out, err := svc.Emby.ResumeItems(c.Request.Context(), uid, limit)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
embyAttachRequestTokenToMediaSources(c, out)
c.JSON(http.StatusOK, out)
}
}
func embyItemsCountsHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"MovieCount": 0,
"SeriesCount": 0,
"EpisodeCount": 0,
"ItemCount": 0,
})
}
}
func embyDisplayPreferencesHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"Id": c.Param("id"),
"ViewType": "Poster",
"SortBy": "SortName",
"SortOrder": "Ascending",
"IndexBy": "SortName",
"RememberIndexing": false,
"PrimaryImageHeight": 250,
"PrimaryImageWidth": 250,
"ScrollDirection": "Horizontal",
"ShowSidebar": true,
"CustomPrefs": gin.H{
"homeexploresection": "1",
"homesection0": "smalllibrarytiles",
"homesection1": "resume",
"homesection2": "latestmedia",
"homesection3": "nextup",
"homesection4": "none",
"homesection5": "none",
"homesection6": "none",
"latestItems": "true",
"landing-livetv": "false",
},
})
}
}
func embySaveDisplayPreferencesHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.Status(http.StatusNoContent)
}
}
func embyShowSeasonsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
params := service.ItemsParams{
UserID: firstQueryValue(c, "UserId", "userId"),
ParentID: c.Param("id"),
Limit: 500,
}
out, err := svc.Emby.Items(c.Request.Context(), params)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
embyAttachRequestTokenToMediaSources(c, out)
c.JSON(http.StatusOK, out)
}
}
func embyShowEpisodesHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
parentID := firstQueryValue(c, "SeasonId", "seasonId")
if parentID == "" {
parentID = c.Param("id")
}
params := service.ItemsParams{
UserID: firstQueryValue(c, "UserId", "userId"),
ParentID: parentID,
IncludeItemTypes: []string{"Episode"},
Recursive: true,
Limit: 500,
}
out, err := svc.Emby.Items(c.Request.Context(), params)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
embyAttachRequestTokenToMediaSources(c, out)
c.JSON(http.StatusOK, out)
}
}
+301
View File
@@ -0,0 +1,301 @@
package handler
import (
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
"github.com/ShukeBta/MediaStationGo/internal/service"
"github.com/ShukeBta/MediaStationGo/internal/service/cloud"
)
func TestEmbyItemImageServesWithoutAPIAuth(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(&model.Media{}); err != nil {
t.Fatalf("migrate: %v", err)
}
posterPath := filepath.Join(t.TempDir(), "poster.png")
if err := os.WriteFile(posterPath, []byte{
0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a,
0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52,
0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01,
0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4,
0x89, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x44, 0x41,
0x54, 0x78, 0x9c, 0x63, 0x00, 0x01, 0x00, 0x00,
0x05, 0x00, 0x01, 0x0d, 0x0a, 0x2d, 0xb4, 0x00,
0x00, 0x00, 0x00, 0x49, 0x45, 0x4e, 0x44, 0xae,
0x42, 0x60, 0x82,
}, 0o644); err != nil {
t.Fatalf("write poster: %v", err)
}
repos := repository.New(db)
cfg := &config.Config{
App: config.AppConfig{DataDir: filepath.Dir(posterPath)},
Cache: config.CacheConfig{CacheDir: t.TempDir()},
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "media-1"},
Title: "Poster Test",
Path: "D:\\media\\poster-test.mp4",
PosterURL: posterPath,
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
router := gin.New()
registerEmbyRoutes(router, "test-secret", &service.Container{
Repo: repos,
Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
ImageProxy: service.NewImageProxy(cfg, zap.NewNop()),
})
req := httptest.NewRequest(http.MethodGet, "/Items/media-1/Images/Primary", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
if location := w.Header().Get("Location"); location != "" {
t.Fatalf("expected direct image response, got redirect to %q", location)
}
if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "image/png") {
t.Fatalf("expected png content type, got %q", contentType)
}
if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "max-age=2592000") {
t.Fatalf("image Cache-Control = %q, want long browser cache", got)
}
if got := w.Header().Get("Pragma"); got != "" {
t.Fatalf("image Pragma = %q, want empty", got)
}
if got := w.Header().Get("Expires"); got != "" {
t.Fatalf("image Expires = %q, want empty", got)
}
}
func TestEmbyItemImageServesCachedCloudArtworkWithoutResolve(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "image/jpeg")
_, _ = w.Write([]byte("emby-cached-cloud-poster"))
}))
defer upstream.Close()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(&model.Media{}); err != nil {
t.Fatalf("migrate: %v", err)
}
cfg := &config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}
imageProxy := service.NewImageProxy(cfg, zap.NewNop())
ref := "/Movies/Cloud Movie/poster.jpg"
if err := imageProxy.PrefetchCloudResolved(t.Context(), "openlist:"+ref, &cloud.DirectLink{URL: upstream.URL + "/poster.jpg"}); err != nil {
t.Fatalf("prefetch cloud poster: %v", err)
}
repos := repository.New(db)
if err := db.Create(&model.Media{
Base: model.Base{ID: "cloud-media-1"},
Title: "Cloud Poster Test",
Path: "cloud://openlist/Movies/Cloud Movie/movie.mkv",
PosterURL: service.CloudArtworkURL("openlist", ref),
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
router := gin.New()
registerEmbyRoutes(router, "test-secret", &service.Container{
Repo: repos,
Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
ImageProxy: imageProxy,
})
req := httptest.NewRequest(http.MethodGet, "/Items/cloud-media-1/Images/Primary", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
if got := w.Body.String(); got != "emby-cached-cloud-poster" {
t.Fatalf("body = %q, want cached cloud poster", got)
}
if location := w.Header().Get("Location"); location != "" {
t.Fatalf("expected direct cached image response, got redirect to %q", location)
}
if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "max-age=2592000") {
t.Fatalf("image Cache-Control = %q, want long browser cache", got)
}
}
func TestEmbyMissingItemImageReturnsTransparentPlaceholder(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
cfg := &config.Config{Cache: config.CacheConfig{CacheDir: t.TempDir()}}
router := gin.New()
registerEmbyRoutes(router, "test-secret", &service.Container{
Repo: repos,
Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
ImageProxy: service.NewImageProxy(cfg, zap.NewNop()),
})
req := httptest.NewRequest(http.MethodHead, "/Items/missing/Images/Primary", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected placeholder status 200, got %d body=%s", w.Code, w.Body.String())
}
if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "image/png") {
t.Fatalf("expected png content type, got %q", contentType)
}
if length := w.Header().Get("Content-Length"); length == "" || length == "0" {
t.Fatalf("expected placeholder content length, got %q", length)
}
if got := w.Header().Get("Pragma"); got != "" {
t.Fatalf("placeholder Pragma = %q, want empty", got)
}
if got := w.Header().Get("Expires"); got != "" {
t.Fatalf("placeholder Expires = %q, want empty", got)
}
}
func TestEmbyUserItemByIDRouteReturnsJSON(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.Favorite{}, &model.PlaybackHistory{}); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
lib := model.Library{Name: "剧集", Path: "D:\\media\\tv", Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "episode-1"},
LibraryID: lib.ID,
Title: "Test Show",
Path: "D:\\media\\tv\\Test Show\\Season 01\\Test Show - S01E01.mkv",
SeasonNum: 1,
EpisodeNum: 1,
Container: "mkv",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
})
req := httptest.NewRequest(http.MethodGet, "/Users/user-1/Items/episode-1", nil)
req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
req.Header.Set("If-None-Match", `"stale-client-cache"`)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
if contentType := w.Header().Get("Content-Type"); !strings.Contains(contentType, "application/json") {
t.Fatalf("expected JSON content type, got %q body=%s", contentType, w.Body.String())
}
var item map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &item); err != nil {
t.Fatalf("decode item: %v", err)
}
if item["Id"] != "episode-1" || item["Type"] != "Episode" {
t.Fatalf("unexpected item payload: %#v", item)
}
}
func TestEmbyUserItemByIDRouteReturnsLibraryView(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
lib := model.Library{Base: model.Base{ID: "lib-tv"}, Name: "剧集", Path: "D:\\media\\tv", Type: "tv", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
})
req := httptest.NewRequest(http.MethodGet, "/Users/user-1/Items/lib-tv", nil)
req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
var item map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &item); err != nil {
t.Fatalf("decode item: %v", err)
}
if item["Id"] != "lib-tv" || item["Type"] != "CollectionFolder" || item["CollectionType"] != "tvshows" {
t.Fatalf("unexpected library payload: %#v", item)
}
}
+272
View File
@@ -0,0 +1,272 @@
package handler
import (
"errors"
"net/http"
"net/url"
"strings"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func embyPlaybackInfoHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := c.Param("userId")
if uid == "" {
uid = embyUserID(c)
}
out, err := svc.Emby.PlaybackInfo(c.Request.Context(), c.Param("id"), uid)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if out == nil {
embyError(c, http.StatusNotFound, "not found")
return
}
embyAttachRequestTokenToMediaSources(c, out)
c.JSON(http.StatusOK, out)
}
}
func embyAttachRequestTokenToMediaSources(c *gin.Context, out any) {
token := embyRequestToken(c)
if token == "" || out == nil {
return
}
embyAttachTokenToMediaSourcesValue(out, token)
}
func embyAttachTokenToMediaSourcesValue(value any, token string) {
switch typed := value.(type) {
case map[string]any:
embyAttachTokenToMediaSourcesMap(typed, token)
case gin.H:
embyAttachTokenToMediaSourcesMap(map[string]any(typed), token)
case []map[string]any:
for _, item := range typed {
embyAttachTokenToMediaSourcesMap(item, token)
}
case []any:
for _, item := range typed {
embyAttachTokenToMediaSourcesValue(item, token)
}
}
}
func embyAttachTokenToMediaSourcesMap(out map[string]any, token string) {
if out == nil {
return
}
if sources, ok := out["MediaSources"].([]map[string]any); ok {
embyAttachTokenToMediaSources(sources, token)
} else if sources, ok := out["MediaSources"].([]any); ok {
for _, source := range sources {
if sourceMap, ok := source.(map[string]any); ok {
embyAttachTokenToMediaSources([]map[string]any{sourceMap}, token)
}
}
}
if items, ok := out["Items"]; ok {
embyAttachTokenToMediaSourcesValue(items, token)
}
}
func embyAttachTokenToMediaSources(sources []map[string]any, token string) {
for _, source := range sources {
for _, key := range []string{"DirectStreamUrl", "TranscodingUrl"} {
raw, ok := source[key].(string)
if !ok {
continue
}
source[key] = embyAppendAPIKey(raw, token)
}
}
}
func embyRequestToken(c *gin.Context) string {
if c == nil {
return ""
}
for _, key := range []string{"api_key", "apiKey", "ApiKey", "token", "X-Emby-Token", "X-MediaBrowser-Token"} {
if value := strings.TrimSpace(c.Query(key)); value != "" {
return value
}
}
for _, header := range []string{"X-Emby-Token", "X-MediaBrowser-Token"} {
if value := strings.TrimSpace(c.GetHeader(header)); value != "" {
return value
}
}
for _, header := range []string{"Authorization", "X-Emby-Authorization", "X-MediaBrowser-Authorization"} {
if token := embyTokenFromAuthHeader(c.GetHeader(header)); token != "" {
return token
}
}
return ""
}
func embyTokenFromAuthHeader(value string) string {
value = strings.TrimSpace(value)
if value == "" {
return ""
}
for _, prefix := range []string{"Bearer ", "Emby "} {
if strings.HasPrefix(value, prefix) {
return strings.TrimSpace(strings.TrimPrefix(value, prefix))
}
}
for _, part := range strings.Split(value, ",") {
part = strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(part), "MediaBrowser "))
if !strings.HasPrefix(part, "Token=") {
continue
}
token := strings.TrimSpace(strings.TrimPrefix(part, "Token="))
return strings.Trim(token, `"`)
}
if strings.Contains(value, "Token=") {
return ""
}
return value
}
func embyAppendAPIKey(raw, token string) string {
raw = strings.TrimSpace(raw)
token = strings.TrimSpace(token)
if raw == "" || token == "" {
return raw
}
if strings.HasPrefix(raw, "//") {
return raw
}
u, err := url.Parse(raw)
if err != nil || u.IsAbs() {
return raw
}
q := u.Query()
if q.Get("api_key") == "" && q.Get("apiKey") == "" && q.Get("token") == "" {
q.Set("api_key", token)
u.RawQuery = q.Encode()
}
return u.String()
}
// embyVideoStreamHandler 是 GET /Videos/{id}/stream 的入口,
// 直接代理到我们的 /api/stream/{id}(同一个 ServeFile)。
func embyVideoStreamHandler(svc *service.Container, cloudMode string) gin.HandlerFunc {
return func(c *gin.Context) {
uid := embyUserID(c)
item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if item == nil {
c.Status(http.StatusNotFound)
return
}
if embyShouldRedirectVideoStreamToSTRM(c, svc, c.Param("id"), cloudMode) {
target := "/api/stream/" + url.PathEscape(strings.TrimSpace(c.Param("id")))
if token := embyPlaybackRedirectToken(c, svc); token != "" {
target = embyAppendAPIKey(target, token)
}
setRedirectNoStoreHeaders(c)
c.Redirect(http.StatusFound, absoluteRequestURL(c, target))
return
}
// 直接调用 Stream service 写入 response。
// 此前这里把所有错误一律吞成 404:云盘 Cookie 过期、直链解析失败、
// STRM 播放被关闭……在第三方播放器上全部表现为「404 不存在」,
// 无法排查。现在区分:行不存在→404;云盘播放不可用/上游故障→502+原因。
err = svc.Stream.ServeFileWithCloudMode(c.Writer, c.Request, c.Param("id"), cloudMode)
switch {
case err == nil:
case errors.Is(err, service.ErrMediaNotFound):
c.Status(http.StatusNotFound)
case errors.Is(err, service.ErrCloudPlaybackDisabled):
if !c.Writer.Written() {
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
}
default:
if !c.Writer.Written() {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
}
}
}
}
func embyPlaybackRedirectToken(c *gin.Context, svc *service.Container) string {
if token := embyRequestToken(c); token != "" {
return token
}
if c == nil || svc == nil || svc.Auth == nil || svc.Repo == nil || svc.Repo.User == nil {
return ""
}
uid := embyUserID(c)
if uid == "" {
return ""
}
u, err := svc.Repo.User.FindByID(c.Request.Context(), uid)
if err != nil || u == nil {
return ""
}
token, err := svc.Auth.IssueEmbyToken(u)
if err != nil {
return ""
}
return token
}
func embyShouldRedirectVideoStreamToSTRM(c *gin.Context, svc *service.Container, mediaID, cloudMode string) bool {
if c == nil || svc == nil || svc.Repo == nil || svc.Repo.Media == nil || cloudMode != service.CloudPlaybackModeRedirectProxy {
return false
}
settings := service.CloudPlaybackSettings(c.Request.Context(), svc.Repo)
if settings.PreferredMode != service.CloudPlaybackModeSTRM || !settings.STRMEnabled {
return false
}
m, err := svc.Repo.Media.FindByID(c.Request.Context(), mediaID)
if err != nil || m == nil {
return false
}
return strings.TrimSpace(m.STRMURL) != ""
}
func embyVideoHLSPlaylistHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := embyUserID(c)
item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
if err != nil || item == nil || svc.Stream == nil {
c.Status(http.StatusNotFound)
return
}
err = svc.Stream.ServeHLSPlaylist(c.Writer, c.Request, c.Param("id"))
if errors.Is(err, service.ErrTranscodeDisabled) {
c.JSON(http.StatusConflict, gin.H{"error": "transcode disabled"})
return
}
if errors.Is(err, service.ErrTranscodeBusy) {
c.JSON(http.StatusTooManyRequests, gin.H{"error": "transcode busy"})
return
}
if err != nil {
c.Status(http.StatusNotFound)
}
}
}
func embyVideoHLSSegmentHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := embyUserID(c)
item, err := svc.Emby.Item(c.Request.Context(), c.Param("id"), uid)
if err != nil || item == nil || svc.Stream == nil {
c.Status(http.StatusNotFound)
return
}
if err := svc.Stream.ServeHLSSegment(c.Writer, c.Request, c.Param("id"), c.Param("seg")); err != nil {
c.Status(http.StatusNotFound)
}
}
}
@@ -0,0 +1,690 @@
package handler
import (
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func TestEmbyLowercasePlaybackInfoRouteReturnsJSON(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
lib := model.Library{Name: "电影", Path: t.TempDir(), Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "media-1"},
LibraryID: lib.ID,
Title: "Lowercase Playback",
Path: filepath.Join(lib.Path, "lowercase-playback.mp4"),
Container: "mp4",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
})
req := httptest.NewRequest(http.MethodGet, "/users/user-1/items/media-1/playbackinfo", nil)
req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
var body map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
t.Fatalf("decode playback info: %v", err)
}
if _, ok := body["MediaSources"]; !ok {
t.Fatalf("missing MediaSources: %#v", body)
}
sources, ok := body["MediaSources"].([]any)
if !ok || len(sources) == 0 {
t.Fatalf("unexpected MediaSources: %#v", body["MediaSources"])
}
source, ok := sources[0].(map[string]any)
if !ok {
t.Fatalf("unexpected MediaSource: %#v", sources[0])
}
directURL, _ := source["DirectStreamUrl"].(string)
if !strings.Contains(directURL, "api_key=") {
t.Fatalf("DirectStreamUrl should carry api_key for clients that do not repeat auth headers: %#v", source)
}
transcodeURL, _ := source["TranscodingUrl"].(string)
if transcodeURL != "" && !strings.Contains(transcodeURL, "api_key=") {
t.Fatalf("TranscodingUrl should carry api_key: %#v", source)
}
}
func TestEmbyPlaybackInfoDoesNotExposeTokenInCloudPath(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
t.Fatalf("set cloud playback mode: %v", err)
}
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "cloud-1"},
LibraryID: lib.ID,
Title: "Cloud Movie",
Path: "cloud://openlist/Movies/Movie.mkv",
STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
Container: "mkv",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
})
req := httptest.NewRequest(http.MethodGet, "/users/user-1/items/cloud-1/playbackinfo", nil)
req.Header.Set("X-Emby-Token", signedTestToken(t, secret))
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
var body map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
t.Fatalf("decode playback info: %v", err)
}
source := body["MediaSources"].([]any)[0].(map[string]any)
pathURL, _ := source["Path"].(string)
if pathURL != "/api/stream/cloud-1" {
t.Fatalf("cloud Path should stay as non-tokenized display stream URL, got %#v", source)
}
if strings.Contains(pathURL, "api_key=") || strings.Contains(pathURL, "token=") {
t.Fatalf("cloud Path must not expose auth key/token: %#v", source)
}
if strings.Contains(pathURL, "/api/cloud/play/") {
t.Fatalf("cloud Path should not expose naked cloud play URL: %#v", source)
}
directURL, _ := source["DirectStreamUrl"].(string)
if !strings.HasPrefix(directURL, "/api/stream/cloud-1") || !strings.Contains(directURL, "api_key=") {
t.Fatalf("DirectStreamUrl should stay tokenized: %#v", source)
}
if source["SupportsDirectPlay"] != true {
t.Fatalf("cloud media should advertise DirectPlay when tokenized Path is playable: %#v", source)
}
if source["SupportsTranscoding"] != false {
t.Fatalf("cloud media should not advertise host transcoding: %#v", source)
}
}
func TestEmbyItemsDoNotExposeTokenInEmbeddedCloudPath(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
t.Fatalf("set cloud playback mode: %v", err)
}
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "cloud-1"},
LibraryID: lib.ID,
Title: "Cloud Movie",
Path: "cloud://openlist/Movies/Movie.mkv",
STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
Container: "mkv",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
token := signedTestToken(t, secret)
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
})
req := httptest.NewRequest(http.MethodGet, "/emby/Users/user-1/Items?IncludeItemTypes=Movie&Recursive=true&Limit=5&X-Emby-Token="+token, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
var body map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
t.Fatalf("decode items: %v", err)
}
items := body["Items"].([]any)
if len(items) != 1 {
t.Fatalf("unexpected items: %#v", body["Items"])
}
source := items[0].(map[string]any)["MediaSources"].([]any)[0].(map[string]any)
pathURL, _ := source["Path"].(string)
if pathURL != "/api/stream/cloud-1" {
t.Fatalf("embedded cloud Path should stay as non-tokenized display stream URL, got %#v", source)
}
if strings.Contains(pathURL, "api_key=") || strings.Contains(pathURL, "token=") {
t.Fatalf("embedded cloud Path must not expose auth key/token: %#v", source)
}
}
func TestEmbyVideoStreamUsesSTRMWhenRedirectProxyDisabled(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
t.Fatalf("set cloud playback mode: %v", err)
}
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackSTRMEnabledSettingKey, "true"); err != nil {
t.Fatalf("enable strm playback: %v", err)
}
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackRedirectEnabledSettingKey, "false"); err != nil {
t.Fatalf("disable redirect playback: %v", err)
}
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "cloud-1"},
LibraryID: lib.ID,
Title: "Cloud Movie",
Path: "cloud://openlist/Movies/Movie.mkv",
STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
Container: "mkv",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
cfg := &config.Config{Secrets: config.SecretsConfig{JWTSecret: secret}}
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
Stream: service.NewStreamService(cfg, zap.NewNop(), repos, nil),
})
token := signedTestToken(t, secret)
req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream?api_key="+token, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusFound {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
loc := w.Header().Get("Location")
if !strings.Contains(loc, "/api/stream/cloud-1") || !strings.Contains(loc, "api_key=") {
t.Fatalf("STRM mode should redirect /Videos fallback to tokenized /api/stream, got %q", loc)
}
if got := w.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
t.Fatalf("STRM fallback redirect Cache-Control = %q, want no-store", got)
}
if strings.Contains(loc, "/api/cloud/play/") {
t.Fatalf("STRM mode should not expose cloud play directly from /Videos fallback: %q", loc)
}
}
func TestEmbyVideoStreamIssuesTokenForSessionFallbackSTRMRedirect(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackModeSettingKey, service.CloudPlaybackModeSTRM); err != nil {
t.Fatalf("set cloud playback mode: %v", err)
}
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackSTRMEnabledSettingKey, "true"); err != nil {
t.Fatalf("enable strm playback: %v", err)
}
if err := repos.Setting.Set(t.Context(), service.CloudPlaybackRedirectEnabledSettingKey, "false"); err != nil {
t.Fatalf("disable redirect playback: %v", err)
}
user := model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}
if err := repos.User.Create(t.Context(), &user); err != nil {
t.Fatalf("create user: %v", err)
}
lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "cloud-1"},
LibraryID: lib.ID,
Title: "Cloud Movie",
Path: "cloud://openlist/Movies/Movie.mkv",
STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
Container: "mkv",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
cfg := &config.Config{Secrets: config.SecretsConfig{JWTSecret: secret}}
svc := &service.Container{
Repo: repos,
Auth: service.NewAuthService(cfg, zap.NewNop(), repos, nil, nil),
Emby: service.NewEmbyService(cfg, zap.NewNop(), repos),
Stream: service.NewStreamService(cfg, zap.NewNop(), repos, nil),
}
router := gin.New()
router.GET("/videos/:id/stream", func(c *gin.Context) {
c.Set(middleware.CtxUserID, user.ID)
c.Set(middleware.CtxUserRole, user.Role)
embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy)(c)
})
req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusFound {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
loc := w.Header().Get("Location")
if !strings.Contains(loc, "/api/stream/cloud-1") || !strings.Contains(loc, "api_key=") {
t.Fatalf("session fallback redirect should include api_key for /api/stream, got %q", loc)
}
}
func TestEmbyLowercaseVideoStreamRouteServesMedia(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
dir := t.TempDir()
mediaPath := filepath.Join(dir, "sample.mp4")
if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
t.Fatalf("write media: %v", err)
}
lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "media-1"},
LibraryID: lib.ID,
Title: "Lowercase Stream",
Path: mediaPath,
Container: "mp4",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
})
req := httptest.NewRequest(http.MethodGet, "/videos/media-1/stream?api_key="+signedTestToken(t, secret), nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
if got := w.Body.String(); got != "fake-video-bytes" {
t.Fatalf("unexpected stream body: %q", got)
}
}
func TestEmbyPrefixedAPIStreamRouteServesMedia(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
dir := t.TempDir()
mediaPath := filepath.Join(dir, "sample.mp4")
if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
t.Fatalf("write media: %v", err)
}
lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "media-1"},
LibraryID: lib.ID,
Title: "Prefixed API Stream",
Path: mediaPath,
Container: "mp4",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
})
req := httptest.NewRequest(http.MethodGet, "/emby/api/stream/media-1?api_key="+signedTestToken(t, secret), nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
if got := w.Body.String(); got != "fake-video-bytes" {
t.Fatalf("unexpected stream body: %q", got)
}
}
func TestEmbyVideoStreamRedirectKeepsMediaBrowserAuthorizationToken(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
lib := model.Library{Name: "OpenList", Path: "cloud://openlist/Movies", Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "cloud-1"},
LibraryID: lib.ID,
Title: "Cloud Movie",
Path: "cloud://openlist/Movies/Movie.mkv",
STRMURL: "/api/cloud/play/openlist?ref=%2FMovies%2FMovie.mkv",
Container: "mkv",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
})
token := signedTestToken(t, secret)
req := httptest.NewRequest(http.MethodGet, "/videos/cloud-1/stream", nil)
req.Header.Set("X-MediaBrowser-Authorization", `MediaBrowser Client="Infuse", Device="PC", Token="`+token+`"`)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusFound {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
loc := w.Header().Get("Location")
if !strings.Contains(loc, "/api/cloud/play/openlist?") || !strings.Contains(loc, "token=") {
t.Fatalf("redirect Location should target tokenized cloud play endpoint, got %q", loc)
}
}
func TestEmbyLowercaseOriginalHeadRouteServesHeaders(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
dir := t.TempDir()
mediaPath := filepath.Join(dir, "sample.mp4")
if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
t.Fatalf("write media: %v", err)
}
lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "media-1"},
LibraryID: lib.ID,
Title: "Lowercase Original",
Path: mediaPath,
Container: "mp4",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
})
req := httptest.NewRequest(http.MethodHead, "/videos/media-1/original.mp4?api_key="+signedTestToken(t, secret), nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("unexpected status: %d body=%s", w.Code, w.Body.String())
}
if w.Body.Len() != 0 {
t.Fatalf("HEAD response should not include body, got %q", w.Body.String())
}
}
func TestEmbyLowercaseVideoHLSRouteDoesNot404WhenDirectOnly(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("open db: %v", err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatalf("migrate: %v", err)
}
repos := repository.New(db)
if err := repos.User.Create(t.Context(), &model.User{
Base: model.Base{ID: "user-1"},
Username: "tester",
PasswordHash: "x",
Role: "admin",
Tier: "plus",
IsActive: true,
}); err != nil {
t.Fatalf("create user: %v", err)
}
dir := t.TempDir()
mediaPath := filepath.Join(dir, "sample.mp4")
if err := os.WriteFile(mediaPath, []byte("fake-video-bytes"), 0o644); err != nil {
t.Fatalf("write media: %v", err)
}
lib := model.Library{Name: "电影", Path: dir, Type: "movie", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatalf("create library: %v", err)
}
if err := db.Create(&model.Media{
Base: model.Base{ID: "media-1"},
LibraryID: lib.ID,
Title: "Lowercase HLS",
Path: mediaPath,
Container: "mp4",
}).Error; err != nil {
t.Fatalf("create media: %v", err)
}
if err := repos.Setting.Set(t.Context(), service.PlaybackDirectOnlySettingKey, "true"); err != nil {
t.Fatalf("set direct-only: %v", err)
}
const secret = "test-secret"
router := gin.New()
registerEmbyRoutes(router, secret, &service.Container{
Repo: repos,
Emby: service.NewEmbyService(&config.Config{}, zap.NewNop(), repos),
Stream: service.NewStreamService(&config.Config{}, zap.NewNop(), repos, nil),
})
req := httptest.NewRequest(http.MethodGet, "/videos/media-1/master.m3u8?api_key="+signedTestToken(t, secret), nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code == http.StatusNotFound {
t.Fatalf("lowercase HLS route should be registered, got 404")
}
if w.Code != http.StatusConflict {
t.Fatalf("direct-only HLS should return 409, got %d body=%s", w.Code, w.Body.String())
}
}
+119
View File
@@ -0,0 +1,119 @@
package handler
import (
"net/http"
"strconv"
"strings"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
type embyPlayingReq struct {
ItemId string `json:"ItemId"`
PositionTicks int64 `json:"PositionTicks"`
RunTimeTicks int64 `json:"RunTimeTicks"`
}
func embyPlayingProgressHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := embyUserID(c)
if uid == "" {
c.Status(http.StatusUnauthorized)
return
}
var req embyPlayingReq
_ = c.ShouldBindJSON(&req)
if req.ItemId == "" {
req.ItemId = c.Query("ItemId")
}
if req.PositionTicks == 0 {
req.PositionTicks, _ = strconv.ParseInt(c.Query("PositionTicks"), 10, 64)
}
if req.RunTimeTicks == 0 {
req.RunTimeTicks, _ = strconv.ParseInt(c.Query("RunTimeTicks"), 10, 64)
}
if req.ItemId == "" {
c.Status(http.StatusOK)
return
}
clientInfo := embyClientInfoFromRequest(c)
if svc.Device != nil && svc.Device.IsDeviceKicked(c.Request.Context(), uid, clientInfo.DeviceID) {
c.Status(http.StatusUnauthorized)
return
}
_ = svc.Emby.RecordProgress(c.Request.Context(), uid, req.ItemId, req.PositionTicks, req.RunTimeTicks)
stopped := strings.Contains(strings.ToLower(c.FullPath()+" "+c.Request.URL.Path), "stopped")
if svc.Sessions != nil {
svc.Sessions.RecordPlayback(c.Request.Context(), uid, "",
clientInfo.DeviceID,
clientInfo.DeviceName,
clientInfo.Client,
c.ClientIP(),
req.ItemId,
req.PositionTicks,
req.RunTimeTicks,
stopped)
}
if svc.Device != nil && !stopped {
svc.Device.RecordPlayback(c.Request.Context(), uid,
clientInfo.DeviceID,
clientInfo.DeviceName,
clientInfo.Client)
}
c.Status(http.StatusNoContent)
}
}
func embyFavoriteHandler(svc *service.Container, fav bool) gin.HandlerFunc {
return func(c *gin.Context) {
uid := c.Param("userId")
if uid == "" {
uid = embyUserID(c)
}
mid := c.Param("itemId")
if uid == "" || mid == "" {
c.Status(http.StatusBadRequest)
return
}
if err := svc.Emby.SetFavorite(c.Request.Context(), uid, mid, fav); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
out, _ := svc.Emby.Item(c.Request.Context(), mid, uid)
if out != nil {
c.JSON(http.StatusOK, out["UserData"])
return
}
c.JSON(http.StatusOK, gin.H{"IsFavorite": fav})
}
}
func embyMarkPlayedHandler(svc *service.Container, played bool) gin.HandlerFunc {
return func(c *gin.Context) {
uid := c.Param("userId")
if uid == "" {
uid = embyUserID(c)
}
mid := c.Param("itemId")
if uid == "" || mid == "" {
c.Status(http.StatusBadRequest)
return
}
if err := svc.Emby.MarkPlayed(c.Request.Context(), uid, mid, played); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if played && svc.Device != nil {
clientInfo := embyClientInfoFromRequest(c)
svc.Device.RecordPlayback(c.Request.Context(), uid, clientInfo.DeviceID, clientInfo.DeviceName, clientInfo.Client)
}
out, _ := svc.Emby.Item(c.Request.Context(), mid, uid)
if out != nil {
c.JSON(http.StatusOK, out["UserData"])
return
}
c.JSON(http.StatusOK, gin.H{"Played": played})
}
}
+235
View File
@@ -0,0 +1,235 @@
package handler
import (
"time"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// registerEmbyRoutes 在 r 上挂双前缀("" + "/emby")的 Emby 兼容路由。
func registerEmbyRoutes(r *gin.Engine, jwtSecret string, svc *service.Container) {
for _, prefix := range []string{"/emby", ""} {
grp := r.Group(prefix)
grp.Use(embyNoStoreHeaders())
registerEmbyRootRoutes(grp, prefix, svc)
registerEmbyPublicRoutes(grp, svc)
registerEmbyPublicImageRoutes(grp, svc)
// 鉴权后端点
auth := grp.Group("", embyAuthRequiredWithSessionFallback(jwtSecret), activeEmbyUserRequired(svc))
registerEmbyAuthenticatedRoutes(auth, prefix, svc)
}
}
type embyRouteHandlerFactory func(*service.Container) gin.HandlerFunc
func embyNoStoreHeaders() gin.HandlerFunc {
return func(c *gin.Context) {
c.Header("Cache-Control", "no-store")
c.Header("Pragma", "no-cache")
c.Header("Expires", "0")
c.Next()
}
}
func registerEmbyRootRoutes(grp *gin.RouterGroup, prefix string, svc *service.Container) {
if prefix != "/emby" {
return
}
grp.GET("", embyRootHandler(svc))
grp.HEAD("", embyRootHandler(svc))
grp.GET("/", embyRootHandler(svc))
grp.HEAD("/", embyRootHandler(svc))
}
func registerEmbyPublicRoutes(grp *gin.RouterGroup, svc *service.Container) {
registerEmbyPublicSystemRoutes(grp, svc)
registerEmbyPublicSessionRoutes(grp, svc)
registerEmbyPublicClientRoutes(grp, svc)
}
func registerEmbyPublicSystemRoutes(grp *gin.RouterGroup, svc *service.Container) {
registerEmbyGetHeadRoutes(grp, svc, []string{"/System/Info/Public", "/system/info/public"}, embySystemInfoPublicHandler)
registerEmbyGetHeadRoutes(grp, svc, []string{"/System/Info", "/system/info"}, embySystemInfoHandler)
registerEmbyGetRoutes(grp, svc, []string{"/System/Endpoint", "/system/endpoint"}, embySystemEndpointHandler)
registerEmbyGetHeadRoutes(grp, svc, []string{"/System/Ext/ServerDomains", "/system/ext/serverdomains"}, embyServerDomainsHandler)
registerEmbyGetHeadRoutes(grp, svc, []string{"/System/Configuration/Public", "/system/configuration/public"}, embyPublicServerConfigurationHandler)
registerEmbyGetHeadRoutes(grp, svc, []string{"/Startup/Configuration", "/startup/configuration"}, embyStartupConfigurationHandler)
registerEmbyPostRoutes(grp, svc, []string{"/Startup/Complete", "/startup/complete"}, embyNoContentHandler)
registerEmbyGetHeadRoutes(grp, svc, []string{"/QuickConnect/Enabled", "/quickconnect/enabled"}, embyQuickConnectEnabledHandler)
for _, path := range []string{"/System/Ping", "/system/ping"} {
grp.GET(path, embyPingHandler(svc))
grp.HEAD(path, embyPingHandler(svc))
grp.POST(path, embyPingHandler(svc))
}
}
func registerEmbyPublicSessionRoutes(grp *gin.RouterGroup, svc *service.Container) {
registerEmbyPostRoutes(grp, svc, []string{
"/Sessions/Capabilities", "/Sessions/Capabilities/Full",
"/sessions/capabilities", "/sessions/capabilities/full",
}, embyNoContentHandler)
// 30/min per IP: many Emby clients sit behind a single NAT/reverse-proxy
// IP, so a low limit would throttle legitimate logins into 429s.
embyLoginLimiter := middleware.NewRateLimiter(30, 1*time.Minute)
for _, path := range []string{"/Users/AuthenticateByName", "/Users/authenticatebyname", "/users/AuthenticateByName", "/users/authenticatebyname"} {
grp.POST(path, middleware.RateLimit(embyLoginLimiter), embyAuthByNameHandler(svc))
}
registerEmbyGetRoutes(grp, svc, []string{"/Users/Public", "/users/public"}, embyPublicUsersHandler)
}
func registerEmbyPublicClientRoutes(grp *gin.RouterGroup, svc *service.Container) {
registerEmbyGetRoutes(grp, svc, []string{"/Branding/Configuration", "/branding/configuration"}, embyBrandingConfigHandler)
registerEmbyGetHeadRoutes(grp, svc, []string{"/Branding/Css", "/branding/css"}, embyBrandingCSSHandler)
registerEmbyGetRoutes(grp, svc, []string{"/Localization/Options", "/localization/options"}, embyLocalizationOptionsHandler)
registerEmbyGetRoutes(grp, svc, []string{"/Localization/Cultures", "/Localization/cultures", "/localization/cultures"}, embyLocalizationCulturesHandler)
registerEmbyGetHeadRoutes(grp, svc, []string{"/CustomCssJS/Scripts", "/customcssjs/scripts"}, embyCustomCSSJSScriptsHandler)
for _, path := range []string{"/embywebsocket", "/EmbyWebSocket"} {
grp.GET(path, embyWebSocketHandler(svc))
grp.HEAD(path, embyNoContentHandler(svc))
}
registerEmbyPostRoutes(grp, svc, []string{"/Sessions/Logout", "/sessions/logout"}, embySessionLogoutHandler)
grp.GET("/DisplayPreferences/:id", embyDisplayPreferencesHandler(svc))
grp.POST("/DisplayPreferences/:id", embySaveDisplayPreferencesHandler(svc))
grp.GET("/displaypreferences/:id", embyDisplayPreferencesHandler(svc))
grp.POST("/displaypreferences/:id", embySaveDisplayPreferencesHandler(svc))
}
func registerEmbyPublicImageRoutes(grp *gin.RouterGroup, svc *service.Container) {
// 图片公开(Infuse 缓存 URL 时会丢 token)
grp.GET("/Items/:id/Images/:type", embyItemImageHandler(svc))
grp.GET("/Items/:id/Images/:type/:index", embyItemImageHandler(svc))
grp.HEAD("/Items/:id/Images/:type", embyItemImageHandler(svc))
grp.GET("/items/:id/images/:type", embyItemImageHandler(svc))
grp.GET("/items/:id/images/:type/:index", embyItemImageHandler(svc))
grp.HEAD("/items/:id/images/:type", embyItemImageHandler(svc))
}
func registerEmbyGetRoutes(grp *gin.RouterGroup, svc *service.Container, paths []string, factory embyRouteHandlerFactory) {
for _, path := range paths {
grp.GET(path, factory(svc))
}
}
func registerEmbyGetHeadRoutes(grp *gin.RouterGroup, svc *service.Container, paths []string, factory embyRouteHandlerFactory) {
for _, path := range paths {
grp.GET(path, factory(svc))
grp.HEAD(path, factory(svc))
}
}
func registerEmbyPostRoutes(grp *gin.RouterGroup, svc *service.Container, paths []string, factory embyRouteHandlerFactory) {
for _, path := range paths {
grp.POST(path, factory(svc))
}
}
func registerEmbyAuthenticatedRoutes(auth *gin.RouterGroup, prefix string, svc *service.Container) {
registerEmbyAuthenticatedUserRoutes(auth, svc)
registerEmbyAuthenticatedItemRoutes(auth, svc)
registerEmbyAuthenticatedPlaybackRoutes(auth, prefix, svc)
registerEmbyAuthenticatedProgressRoutes(auth, svc)
registerEmbyAuthenticatedUserDataRoutes(auth, svc)
registerEmbyAuthenticatedSystemRoutes(auth, svc)
registerLowercaseEmbyAuthRoutes(auth, svc)
}
func registerEmbyAuthenticatedUserRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.GET("/Users/Me", embyMeHandler(svc))
auth.GET("/Users", embyListUsersHandler(svc))
auth.GET("/Users/:userId", embyGetUserByIDHandler(svc))
auth.GET("/Users/:userId/Views", embyViewsHandler(svc))
auth.GET("/Library/MediaFolders", embyViewsHandler(svc))
auth.GET("/Library/VirtualFolders", embyVirtualFoldersHandler(svc))
auth.GET("/Library/SelectableMediaFolders", embyVirtualFoldersHandler(svc))
}
func registerEmbyAuthenticatedItemRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.GET("/Items", embyItemsHandler(svc))
auth.GET("/Users/:userId/Items", embyItemsHandler(svc))
auth.GET("/Items/Counts", embyItemsCountsHandler(svc))
auth.GET("/Users/:userId/Items/Counts", embyItemsCountsHandler(svc))
auth.GET("/Items/Latest", embyLatestItemsHandler(svc))
auth.GET("/Items/Resume", embyResumeItemsHandler(svc))
auth.GET("/Items/:id", embyItemByIDHandler(svc))
auth.GET("/Users/:userId/Items/:id", embyUserItemByIDHandler(svc))
auth.GET("/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
auth.GET("/Shows/:id/Episodes", embyShowEpisodesHandler(svc))
auth.GET("/Users/:userId/Shows/:id/Seasons", embyShowSeasonsHandler(svc))
auth.GET("/Users/:userId/Shows/:id/Episodes", embyShowEpisodesHandler(svc))
auth.GET("/Shows/NextUp", embyEmptyItemsHandler(svc))
auth.GET("/Users/:userId/Shows/NextUp", embyEmptyItemsHandler(svc))
auth.GET("/MediaSegments/:id", embyEmptyItemsHandler(svc))
auth.GET("/Artists", embyEmptyItemsHandler(svc))
auth.GET("/Persons", embyEmptyItemsHandler(svc))
auth.GET("/Genres", embyEmptyItemsHandler(svc))
auth.GET("/Shows/Upcoming", embyEmptyItemsHandler(svc))
auth.GET("/Users/:userId/Shows/Upcoming", embyEmptyItemsHandler(svc))
auth.GET("/Items/:id/Similar", embyEmptyItemsHandler(svc))
auth.GET("/Items/:id/ThumbnailSet", embyEmptyItemsHandler(svc))
auth.GET("/Items/:id/ThemeMedia", embyThemeMediaHandler(svc))
auth.GET("/Users/:userId/Items/:id/SpecialFeatures", embyEmptyItemsHandler(svc))
auth.GET("/Users/:userId/Items/:id/Intros", embyEmptyItemsHandler(svc))
auth.GET("/Items/:id/SpecialFeatures", embyEmptyItemsHandler(svc))
auth.GET("/Items/:id/Intros", embyEmptyItemsHandler(svc))
auth.GET("/api/danmu/:id/raw", embyDanmuRawHandler(svc))
}
func registerEmbyAuthenticatedPlaybackRoutes(auth *gin.RouterGroup, prefix string, svc *service.Container) {
auth.GET("/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
auth.POST("/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
auth.GET("/Users/:userId/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
auth.POST("/Users/:userId/Items/:id/PlaybackInfo", embyPlaybackInfoHandler(svc))
registerEmbyVideoStreamRoutes(auth, svc, "/Videos")
if prefix == "/emby" {
auth.GET("/api/stream/:id", embyVideoStreamHandler(svc, service.CloudPlaybackModeSTRM))
auth.HEAD("/api/stream/:id", embyVideoStreamHandler(svc, service.CloudPlaybackModeSTRM))
}
auth.GET("/Videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.HEAD("/Videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.GET("/Videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.HEAD("/Videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.GET("/Videos/:id/:seg", embyVideoHLSSegmentHandler(svc))
}
func registerEmbyVideoStreamRoutes(auth *gin.RouterGroup, svc *service.Container, basePath string) {
streamHandler := func() gin.HandlerFunc {
return embyVideoStreamHandler(svc, service.CloudPlaybackModeRedirectProxy)
}
for _, path := range []string{"/:id/stream", "/:id/stream.:container", "/:id/original", "/:id/original.:container"} {
fullPath := basePath + path
auth.GET(fullPath, streamHandler())
auth.HEAD(fullPath, streamHandler())
}
}
func registerEmbyAuthenticatedProgressRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.POST("/Sessions/Playing", embyPlayingProgressHandler(svc))
auth.POST("/Sessions/Playing/Progress", embyPlayingProgressHandler(svc))
auth.POST("/Sessions/Playing/Stopped", embyPlayingProgressHandler(svc))
}
func registerEmbyAuthenticatedUserDataRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.POST("/Users/:userId/FavoriteItems/:itemId", embyFavoriteHandler(svc, true))
auth.DELETE("/Users/:userId/FavoriteItems/:itemId", embyFavoriteHandler(svc, false))
auth.POST("/Users/:userId/PlayedItems/:itemId", embyMarkPlayedHandler(svc, true))
auth.DELETE("/Users/:userId/PlayedItems/:itemId", embyMarkPlayedHandler(svc, false))
}
func registerEmbyAuthenticatedSystemRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.GET("/Sessions", embySessionsHandler(svc))
auth.GET("/System/Configuration", embyServerConfigurationHandler(svc))
auth.GET("/System/WakeOnLanInfo", embyEmptyArrayHandler(svc))
auth.GET("/ScheduledTasks", embyEmptyArrayHandler(svc))
auth.GET("/LiveTv/Recordings", embyEmptyItemsHandler(svc))
auth.GET("/System/ActivityLog/Entries", embyEmptyItemsHandler(svc))
auth.GET("/Web/ConfigurationPages", embyEmptyArrayHandler(svc))
auth.POST("/Users/:userId/Configuration", embyNoContentHandler(svc))
}
+94
View File
@@ -0,0 +1,94 @@
package handler
import (
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func registerLowercaseEmbyAuthRoutes(auth *gin.RouterGroup, svc *service.Container) {
registerLowercaseEmbyUserRoutes(auth, svc)
registerLowercaseEmbyItemRoutes(auth, svc)
registerLowercaseEmbyPlaybackRoutes(auth, svc)
registerLowercaseEmbyProgressRoutes(auth, svc)
registerLowercaseEmbyUserDataRoutes(auth, svc)
registerLowercaseEmbySystemRoutes(auth, svc)
}
func registerLowercaseEmbyUserRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.GET("/users/me", embyMeHandler(svc))
auth.GET("/users", embyListUsersHandler(svc))
auth.GET("/users/:userId", embyGetUserByIDHandler(svc))
auth.GET("/users/:userId/views", embyViewsHandler(svc))
auth.GET("/library/mediafolders", embyViewsHandler(svc))
auth.GET("/library/virtualfolders", embyVirtualFoldersHandler(svc))
auth.GET("/library/selectablemediafolders", embyVirtualFoldersHandler(svc))
}
func registerLowercaseEmbyItemRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.GET("/items", embyItemsHandler(svc))
auth.GET("/users/:userId/items", embyItemsHandler(svc))
auth.GET("/items/counts", embyItemsCountsHandler(svc))
auth.GET("/users/:userId/items/counts", embyItemsCountsHandler(svc))
auth.GET("/items/latest", embyLatestItemsHandler(svc))
auth.GET("/items/resume", embyResumeItemsHandler(svc))
auth.GET("/items/:id", embyItemByIDHandler(svc))
auth.GET("/users/:userId/items/:id", embyUserItemByIDHandler(svc))
auth.GET("/shows/:id/seasons", embyShowSeasonsHandler(svc))
auth.GET("/shows/:id/episodes", embyShowEpisodesHandler(svc))
auth.GET("/users/:userId/shows/:id/seasons", embyShowSeasonsHandler(svc))
auth.GET("/users/:userId/shows/:id/episodes", embyShowEpisodesHandler(svc))
auth.GET("/shows/nextup", embyEmptyItemsHandler(svc))
auth.GET("/users/:userId/shows/nextup", embyEmptyItemsHandler(svc))
auth.GET("/mediasegments/:id", embyEmptyItemsHandler(svc))
auth.GET("/artists", embyEmptyItemsHandler(svc))
auth.GET("/persons", embyEmptyItemsHandler(svc))
auth.GET("/genres", embyEmptyItemsHandler(svc))
auth.GET("/shows/upcoming", embyEmptyItemsHandler(svc))
auth.GET("/users/:userId/shows/upcoming", embyEmptyItemsHandler(svc))
auth.GET("/items/:id/similar", embyEmptyItemsHandler(svc))
auth.GET("/items/:id/thumbnailset", embyEmptyItemsHandler(svc))
auth.GET("/items/:id/thememedia", embyThemeMediaHandler(svc))
auth.GET("/users/:userId/items/:id/specialfeatures", embyEmptyItemsHandler(svc))
auth.GET("/users/:userId/items/:id/intros", embyEmptyItemsHandler(svc))
auth.GET("/items/:id/specialfeatures", embyEmptyItemsHandler(svc))
auth.GET("/items/:id/intros", embyEmptyItemsHandler(svc))
}
func registerLowercaseEmbyPlaybackRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.GET("/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
auth.POST("/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
auth.GET("/users/:userId/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
auth.POST("/users/:userId/items/:id/playbackinfo", embyPlaybackInfoHandler(svc))
registerEmbyVideoStreamRoutes(auth, svc, "/videos")
auth.GET("/videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.HEAD("/videos/:id/master.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.GET("/videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.HEAD("/videos/:id/main.m3u8", embyVideoHLSPlaylistHandler(svc))
auth.GET("/videos/:id/:seg", embyVideoHLSSegmentHandler(svc))
}
func registerLowercaseEmbyProgressRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.POST("/sessions/playing", embyPlayingProgressHandler(svc))
auth.POST("/sessions/playing/progress", embyPlayingProgressHandler(svc))
auth.POST("/sessions/playing/stopped", embyPlayingProgressHandler(svc))
}
func registerLowercaseEmbyUserDataRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.POST("/users/:userId/favoriteitems/:itemId", embyFavoriteHandler(svc, true))
auth.DELETE("/users/:userId/favoriteitems/:itemId", embyFavoriteHandler(svc, false))
auth.POST("/users/:userId/playeditems/:itemId", embyMarkPlayedHandler(svc, true))
auth.DELETE("/users/:userId/playeditems/:itemId", embyMarkPlayedHandler(svc, false))
}
func registerLowercaseEmbySystemRoutes(auth *gin.RouterGroup, svc *service.Container) {
auth.GET("/sessions", embySessionsHandler(svc))
auth.GET("/system/configuration", embyServerConfigurationHandler(svc))
auth.GET("/system/wakeonlaninfo", embyEmptyArrayHandler(svc))
auth.GET("/scheduledtasks", embyEmptyArrayHandler(svc))
auth.GET("/livetv/recordings", embyEmptyItemsHandler(svc))
auth.GET("/system/activitylog/entries", embyEmptyItemsHandler(svc))
auth.GET("/web/configurationpages", embyEmptyArrayHandler(svc))
auth.POST("/users/:userId/configuration", embyNoContentHandler(svc))
}
+59
View File
@@ -0,0 +1,59 @@
package handler
import (
"net/http"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func embySessionsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if svc.Sessions == nil {
c.JSON(http.StatusOK, []any{})
return
}
out := make([]gin.H, 0)
for _, sess := range svc.Sessions.List(c.Request.Context()) {
last := sess.LastActivityAt
itemID := sess.ItemID
playState := gin.H{
"PositionTicks": sess.PositionTicks,
"IsPaused": sess.IsPaused,
"PlayMethod": "DirectStream",
"CanSeek": true,
}
row := gin.H{
"Id": sess.ID,
"ServerId": "mediastation-go-001",
"Client": sess.Client,
"DeviceId": sess.DeviceID,
"DeviceName": sess.DeviceName,
"UserId": sess.UserID,
"UserName": sess.UserName,
"LastActivityDate": last,
"RemoteEndPoint": sess.RemoteEndPoint,
"PlayState": playState,
"SupportsRemoteControl": true,
}
if itemID != "" && sess.IsPlaying {
row["NowPlayingItem"] = gin.H{"Id": itemID}
}
out = append(out, row)
}
c.Header("Cache-Control", "no-store")
c.JSON(http.StatusOK, out)
}
}
func embySessionLogoutHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if svc.Sessions != nil {
uid := embyUserID(c)
clientInfo := embyClientInfoFromRequest(c)
svc.Sessions.Logout(c.Request.Context(), uid, clientInfo.DeviceID, c.ClientIP())
}
c.Status(http.StatusNoContent)
}
}
+47
View File
@@ -0,0 +1,47 @@
package handler
import (
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func TestEmbySessionsReturnsRealtimeSession(t *testing.T) {
gin.SetMode(gin.TestMode)
tracker := service.NewSessionTrackerService(zap.NewNop())
tracker.RecordPlayback(t.Context(), "user-1", "viewer", "dev-1", "Apple TV", "Yamby", "10.0.0.8", "media-1", 1000, 2000, false)
svc := &service.Container{Sessions: tracker}
router := gin.New()
router.GET("/Sessions", embySessionsHandler(svc))
req := httptest.NewRequest(http.MethodGet, "/Sessions", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
if got := w.Header().Get("Cache-Control"); got != "no-store" {
t.Fatalf("cache-control = %q, want no-store", got)
}
var rows []map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &rows); err != nil {
t.Fatal(err)
}
if len(rows) != 1 {
t.Fatalf("sessions = %d, want 1: %s", len(rows), w.Body.String())
}
if rows[0]["UserId"] != "user-1" || rows[0]["DeviceId"] != "dev-1" || rows[0]["Client"] != "Yamby" {
t.Fatalf("session payload = %#v", rows[0])
}
if _, err := time.Parse(time.RFC3339Nano, rows[0]["LastActivityDate"].(string)); err != nil {
t.Fatalf("LastActivityDate should be RFC3339 time, got %#v", rows[0]["LastActivityDate"])
}
}
+189
View File
@@ -0,0 +1,189 @@
package handler
import (
"net/http"
"time"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func embyNoContentHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.Status(http.StatusNoContent)
}
}
func embyWebSocketHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if !websocket.IsWebSocketUpgrade(c.Request) {
c.Status(http.StatusNoContent)
return
}
conn, err := wsUpgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
return
}
defer conn.Close()
done := make(chan struct{})
go func() {
defer close(done)
for {
if _, _, err := conn.NextReader(); err != nil {
return
}
}
}()
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-done:
return
case <-ticker.C:
_ = conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil {
return
}
}
}
}
}
func embyServerConfigurationHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"EnableFolderView": true,
"EnableGroupingIntoCollections": true,
"EnableExternalContentInSuggestions": false,
"ImageSavingConvention": "Compatible",
})
}
}
func embyPublicServerConfigurationHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"IsStartupWizardCompleted": true,
"EnableRemoteAccess": true,
"EnableUPnP": false,
"EnableHttps": false,
"RequireHttps": false,
"LocalNetworkSubnets": []string{},
"LocalNetworkAddresses": []string{},
"RemoteClientBitrateLimit": 0,
})
}
}
func embyStartupConfigurationHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"IsStartupWizardCompleted": true,
"StartupWizardCompleted": true,
"EnableRemoteAccess": true,
"UICulture": "zh-CN",
"MetadataCountryCode": "CN",
"PreferredMetadataLanguage": "zh-CN",
})
}
}
func embyQuickConnectEnabledHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, false)
}
}
func embyEmptyItemsHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"Items": []any{}, "TotalRecordCount": 0})
}
}
func embyEmptyArrayHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, []any{})
}
}
func embyCustomCSSJSScriptsHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.Data(http.StatusOK, "application/javascript; charset=utf-8", nil)
}
}
func embyLocalizationCulturesHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, []gin.H{
{
"DisplayName": "简体中文",
"Name": "zh-CN",
"ThreeLetterISOLanguageName": "zho",
"TwoLetterISOLanguageName": "zh",
"ThreeLetterISOLanguageNames": []string{"zho", "chi"},
"IsRightToLeft": false,
},
{
"DisplayName": "English",
"Name": "en-US",
"ThreeLetterISOLanguageName": "eng",
"TwoLetterISOLanguageName": "en",
"ThreeLetterISOLanguageNames": []string{"eng"},
"IsRightToLeft": false,
},
})
}
}
func embyThemeMediaHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
empty := gin.H{"Items": []any{}, "TotalRecordCount": 0}
c.JSON(http.StatusOK, gin.H{
"ThemeVideosResult": empty,
"ThemeSongsResult": empty,
"SoundtrackSongsResult": empty,
})
}
}
func embyServerDomainsHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, []any{})
}
}
func embyDanmuRawHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.Data(http.StatusOK, "text/plain; charset=utf-8", nil)
}
}
func embyBrandingConfigHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"LoginDisclaimer": "",
"CustomCss": "",
"SplashscreenEnabled": false,
})
}
}
func embyBrandingCSSHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.Data(http.StatusOK, "text/css; charset=utf-8", []byte(""))
}
}
func embyLocalizationOptionsHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, []map[string]any{
{"Name": "简体中文", "Value": "zh-CN"},
{"Name": "English", "Value": "en-US"},
})
}
}
+98
View File
@@ -0,0 +1,98 @@
package handler
import (
"net/http"
"strings"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func embySystemInfoHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, embyWithRequestAddress(c, svc.Emby.SystemInfo()))
}
}
func embySystemInfoPublicHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, embyWithRequestAddress(c, svc.Emby.SystemInfoPublic()))
}
}
func embyRequestBaseURL(c *gin.Context) string {
proto := strings.TrimSpace(c.GetHeader("X-Forwarded-Proto"))
if proto == "" {
if c.Request != nil && c.Request.TLS != nil {
proto = "https"
} else {
proto = "http"
}
}
if comma := strings.Index(proto, ","); comma >= 0 {
proto = strings.TrimSpace(proto[:comma])
}
host := strings.TrimSpace(c.GetHeader("X-Forwarded-Host"))
if host == "" && c.Request != nil {
host = strings.TrimSpace(c.Request.Host)
}
if host == "" {
return ""
}
return strings.TrimRight(proto+"://"+host, "/")
}
func embyWithRequestAddress(c *gin.Context, payload map[string]any) map[string]any {
out := make(map[string]any, len(payload)+2)
for key, value := range payload {
out[key] = value
}
if address := embyRequestBaseURL(c); address != "" {
out["LocalAddress"] = address
out["WanAddress"] = address
out["PublishedServerUrl"] = address
}
return out
}
func embySystemEndpointHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"IsLocal": true,
"IsInNetwork": true,
})
}
}
func embyPingHandler(_ *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
// Emby/Jellyfin 期望 plain text "Emby Server"
c.String(http.StatusOK, "Emby Server")
}
}
func embyRootHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.JSON(http.StatusOK, embyPublicSystemInfoPayload(c, svc))
}
}
func embyPublicSystemInfoPayload(c *gin.Context, svc *service.Container) map[string]any {
if svc != nil && svc.Emby != nil {
return embyWithRequestAddress(c, svc.Emby.SystemInfoPublic())
}
return embyWithRequestAddress(c, map[string]any{
"Id": "mediastation-go-001",
"ServerId": "mediastation-go-001",
"ServerName": "MediaStationGo",
"Version": "4.8.10.0",
"ServerVersion": "4.8.10.0",
"ProductName": "Emby Server",
"OperatingSystem": "Windows",
"SupportsHttps": false,
"SupportsAutoDiscovery": true,
"StartupWizardCompleted": true,
})
}
File diff suppressed because it is too large Load Diff
+174
View File
@@ -0,0 +1,174 @@
package handler
import (
"net/http"
"strings"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// ─── Users / Auth ────────────────────────────────────────────────────────────
// embyAuthByNameHandler 处理 POST /Users/AuthenticateByName。
//
// 这是 Emby 客户端登录的唯一入口(Infuse / Yamby / Hills 等都走这里)。
// 用户名+密码 → 调用我们已有的 AuthService.Login → 返回 AccessToken + User。
func embyAuthByNameHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
req, err := parseEmbyAuthByNameReq(c)
if err != nil {
embyError(c, http.StatusBadRequest, "invalid body")
return
}
password := req.Pw
if password == "" {
password = req.Password
}
if strings.TrimSpace(req.Username) == "" || password == "" {
if req.PasswordMd5 != "" || req.PasswordSha1 != "" {
embyError(c, http.StatusBadRequest, "plain password required")
return
}
embyError(c, http.StatusBadRequest, "missing username or password")
return
}
resp, err := svc.Auth.Login(c.Request.Context(), req.Username, password)
if err != nil {
embyError(c, http.StatusUnauthorized, err.Error())
return
}
// 记录登录设备会话并执行防共享检测(登录客户端数 / 设备指纹)。
clientInfo := embyClientInfoFromRequest(c)
if svc.Sessions != nil {
svc.Sessions.RecordLogin(c.Request.Context(), resp.User.ID, resp.User.Username,
clientInfo.DeviceID,
clientInfo.DeviceName,
clientInfo.Client,
c.ClientIP())
}
if svc.Device != nil {
svc.Device.RecordLogin(c.Request.Context(), resp.User.ID,
clientInfo.DeviceID,
clientInfo.DeviceName,
clientInfo.Client,
c.ClientIP())
}
userPayload, _ := svc.Emby.FindUser(c.Request.Context(), resp.User.ID)
// Emby/Jellyfin 客户端没有 refresh token 机制:它们把这里返回的
// AccessToken 长期保存并反复使用。若返回 60 分钟的普通 access
// token,客户端每小时就会掉登录、无法播放、媒体库无法刷新。因此
// 签发长期令牌(IssueEmbyToken)匹配 Emby 持久化令牌语义。
accessToken := resp.Tokens.AccessToken
if longLived, err := svc.Auth.IssueEmbyToken(resp.User); err == nil && longLived != "" {
accessToken = longLived
}
embyRememberCompatSession(c, accessToken)
c.JSON(http.StatusOK, gin.H{
"AccessToken": accessToken,
"ServerId": "mediastation-go-001",
"User": userPayload,
"SessionInfo": gin.H{
"Id": resp.User.ID,
"UserId": resp.User.ID,
"UserName": resp.User.Username,
"Client": clientInfo.Client,
"DeviceId": clientInfo.DeviceID,
"DeviceName": clientInfo.DeviceName,
},
})
}
}
func embyPublicUsersHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
// 公开用户列表(Emby Web 客户端登录页拉这个,列出可见用户)。
users, err := svc.Emby.ListUsers(c.Request.Context())
if err != nil {
c.JSON(http.StatusOK, []any{})
return
}
// 公开版本只暴露 Id + Name,不包含 Policy。
out := make([]map[string]any, 0, len(users))
for _, u := range users {
out = append(out, map[string]any{
"Id": u["Id"],
"Name": u["Name"],
"ServerId": u["ServerId"],
"HasPassword": true,
})
}
c.JSON(http.StatusOK, out)
}
}
func embyListUsersHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
users, err := svc.Emby.ListUsers(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, users)
}
}
func embyMeHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := embyUserID(c)
if uid == "" {
embyError(c, http.StatusUnauthorized, "not authenticated")
return
}
u, err := svc.Emby.FindUser(c.Request.Context(), uid)
if err != nil || u == nil {
embyError(c, http.StatusNotFound, "user not found")
return
}
c.JSON(http.StatusOK, u)
}
}
func embyGetUserByIDHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
u, err := svc.Emby.FindUser(c.Request.Context(), c.Param("userId"))
if err == nil && u != nil {
c.JSON(http.StatusOK, u)
return
}
if authUID := embyUserID(c); authUID != "" && authUID != c.Param("userId") {
u, err = svc.Emby.FindUser(c.Request.Context(), authUID)
if err == nil && u != nil {
c.JSON(http.StatusOK, u)
return
}
}
c.JSON(http.StatusOK, embyFallbackUser(c.Param("userId")))
}
}
func embyFallbackUser(id string) gin.H {
if strings.TrimSpace(id) == "" {
id = "mediastation-user"
}
return gin.H{
"Id": id,
"Name": "MediaStationGo",
"ServerId": "mediastation-go-001",
"HasPassword": true,
"HasConfiguredPassword": true,
"HasConfiguredEasyPassword": false,
"EnableAutoLogin": false,
"Policy": gin.H{
"IsAdministrator": true,
"EnableContentDeletion": true,
"EnableRemoteControlOfOtherUsers": true,
"EnableSharedDeviceControl": true,
"EnableRemoteAccess": true,
"EnableAllDevices": true,
"EnableAllChannels": true,
"EnableAllFolders": true,
},
}
}
+63
View File
@@ -0,0 +1,63 @@
package handler
import (
"net/http"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func embyViewsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
uid := c.Param("userId")
if uid == "" {
uid = embyUserID(c)
}
out, err := svc.Emby.Views(c.Request.Context(), uid)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
embyAttachRequestTokenToMediaSources(c, out)
c.JSON(http.StatusOK, out)
}
}
func embyVirtualFoldersHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
c.Header("Cache-Control", "no-store")
libs, err := svc.Repo.Library.List(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
libs = service.FilterDisplayCloudLibraries(c.Request.Context(), svc.Repo, libs)
uid := embyUserID(c)
visibility := service.UserDefaultMediaVisibility(c.Request.Context(), svc.Repo, uid)
out := make([]gin.H, 0, len(libs))
for _, lib := range libs {
if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, lib, visibility) {
continue
}
collectionType := "movies"
switch lib.Type {
case "tv", "anime", "variety":
collectionType = "tvshows"
case "music":
collectionType = "music"
}
out = append(out, gin.H{
"Name": lib.Name,
"Locations": []string{lib.Path},
"CollectionType": collectionType,
"ItemId": lib.ID,
"Id": lib.ID,
"PrimaryImageItemId": lib.ID,
"RefreshStatus": "Idle",
"LibraryOptions": gin.H{},
})
}
c.JSON(http.StatusOK, out)
}
}
+14 -1
View File
@@ -14,6 +14,18 @@ import (
type manualScrapeApplyReq struct {
MediaIDs []string `json:"media_ids"`
Match service.ManualScrapeRequest `json:"match"`
EpisodeArtwork *bool `json:"episode_artwork"`
EpisodeImages *bool `json:"episode_images"`
}
func (r manualScrapeApplyReq) episodeArtworkOption() *bool {
if r.EpisodeImages != nil {
return r.EpisodeImages
}
if r.EpisodeArtwork != nil {
return r.EpisodeArtwork
}
return r.Match.EpisodeArtworkOption()
}
const manualScrapeApplyTimeout = 5 * time.Minute
@@ -72,10 +84,11 @@ func manualScrapeApplyBatchHandler(svc *service.Container) gin.HandlerFunc {
}
applyCtx, cancel := manualScrapeApplyContext(c)
defer cancel()
options := service.ScrapeOptions{EpisodeArtwork: req.episodeArtworkOption()}
applied := 0
errorsOut := make([]string, 0)
for _, id := range ids {
if _, err := svc.Scraper.ApplyManualMatch(applyCtx, id, req.Match); err != nil {
if _, err := svc.Scraper.ApplyManualMatchWithOptions(applyCtx, id, req.Match, options); err != nil {
errorsOut = append(errorsOut, id+": "+err.Error())
continue
}
+3 -2
View File
@@ -27,6 +27,7 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
libs = service.FilterDeprecatedNativeCloudLibraries(libs)
role, _ := c.Get(middleware.CtxUserRole)
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("all") == "1")
if !includeHidden {
@@ -109,7 +110,7 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
})
}
_, _, _ = svc.Scan.StartCloudLibraryScan(id, false)
_, _, _ = svc.Scan.StartCloudLibraryScan(id, true)
finishHTTPTask(task, nil, "queued", "云盘扫描已加入后台队列", map[string]int64{"queued": 1}, nil)
c.JSON(http.StatusAccepted, gin.H{
"library_id": id,
@@ -119,7 +120,7 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
"probed": 0,
"queued": true,
"cloud": true,
"message": "云盘扫描已在后台运行,发现的媒体会自动加入当前媒体库",
"message": "云盘扫描已在后台运行,发现的媒体会自动加入当前媒体库;若已开启自动刮削,会在扫描后补齐元数据",
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
})
return
+8 -2
View File
@@ -80,16 +80,22 @@ func listFavoritesAliasHandler(svc *service.Container) gin.HandlerFunc {
// path; the AI hint comes from svc.AI when configured.
func aiScrapeMediaHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
options, err := scrapeOptionsFromRequest(c, false)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
return
}
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || m == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
return
}
if err := svc.Scraper.EnrichOne(c.Request.Context(), m); err != nil {
if err := svc.Scraper.EnrichOneWithOptions(c.Request.Context(), m, options); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, m)
refreshed, _ := svc.Repo.Media.FindByID(c.Request.Context(), m.ID)
c.JSON(http.StatusOK, refreshed)
}
}
+128
View File
@@ -1,9 +1,13 @@
package handler
import (
"bytes"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
@@ -154,6 +158,82 @@ func TestListMediaGroupsMultipleVersionsByDefault(t *testing.T) {
}
}
func TestListLibrarySeriesDoesNotTruncateLargeEpisodeLibraries(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.User{}, &model.Library{}, &model.Media{}, &model.Setting{}, &model.PlayProfile{}); err != nil {
t.Fatal(err)
}
repos := repository.New(db)
lib := model.Library{Name: "国漫", Path: "cloud://openlist/国漫", Type: "anime", Enabled: true}
if err := repos.Library.Create(t.Context(), &lib); err != nil {
t.Fatal(err)
}
rows := make([]model.Media, 0, 2001)
for i := 1; i <= 2001; i++ {
rows = append(rows, model.Media{
Base: model.Base{ID: fmt.Sprintf("ep-%04d", i), CreatedAt: time.Now().Add(time.Duration(i) * time.Second)},
LibraryID: lib.ID,
Title: "大剧",
Path: fmt.Sprintf("cloud://openlist/国漫/大剧 (2026) {tmdb-123}/Season 1/大剧.S01E%04d.mkv", i),
SeasonNum: 1,
EpisodeNum: i,
})
}
if err := repos.DB.CreateInBatches(rows, 500).Error; err != nil {
t.Fatal(err)
}
svc := &service.Container{
Repo: repos,
Media: service.NewMediaService(&config.Config{}, zap.NewNop(), repos),
}
series := requestLibrarySeries(t, svc, "/api/libraries/"+lib.ID+"/series", lib.ID)
if series.Total != 1 || len(series.Items) != 1 {
t.Fatalf("series response total=%d len=%d body=%#v", series.Total, len(series.Items), series)
}
if series.Items[0].Count != 2001 {
t.Fatalf("series count = %d, want 2001", series.Items[0].Count)
}
if !strings.HasPrefix(series.Items[0].Key, "series:") ||
strings.Contains(series.Items[0].Key, "lib:") ||
strings.Contains(series.Items[0].Key, "show:") {
t.Fatalf("series key = %q, want compact non-raw key", series.Items[0].Key)
}
episodes := requestLibrarySeriesEpisodes(t, svc, "/api/libraries/"+lib.ID+"/series/episodes?key="+url.QueryEscape(series.Items[0].Key), lib.ID)
if episodes.Total != 2001 || len(episodes.Items) != 2001 {
t.Fatalf("episodes total=%d len=%d, want 2001", episodes.Total, len(episodes.Items))
}
if episodes.Items[0].EpisodeNum != 1 || episodes.Items[len(episodes.Items)-1].EpisodeNum != 2001 {
t.Fatalf("episode order first=%d last=%d", episodes.Items[0].EpisodeNum, episodes.Items[len(episodes.Items)-1].EpisodeNum)
}
}
func TestScrapeOptionsFromRequestPreservesEpisodeImagesFalse(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodPost, "/api/media/ep-1/scrape", bytes.NewBufferString(`{"episode_images":false,"refresh_matched":true}`))
c.Request.Header.Set("Content-Type", "application/json")
options, err := scrapeOptionsFromRequest(c, false)
if err != nil {
t.Fatal(err)
}
if options.EpisodeArtwork == nil {
t.Fatal("EpisodeArtwork is nil, want explicit false")
}
if *options.EpisodeArtwork {
t.Fatal("EpisodeArtwork = true, want false")
}
if !options.IncludeMatched {
t.Fatal("IncludeMatched = false, want true from refresh_matched")
}
}
func requestLibraries(t *testing.T, svc *service.Container, userID, role, path string) []model.Library {
t.Helper()
w := httptest.NewRecorder()
@@ -177,6 +257,16 @@ type mediaListResponse struct {
Total int64 `json:"total"`
}
type seriesListResponse struct {
Items []service.SeriesCard `json:"items"`
Total int64 `json:"total"`
}
type seriesEpisodesResponse struct {
Items []model.Media `json:"items"`
Total int64 `json:"total"`
}
func requestMediaList(t *testing.T, svc *service.Container, path, libraryID string) mediaListResponse {
t.Helper()
w := httptest.NewRecorder()
@@ -195,3 +285,41 @@ func requestMediaList(t *testing.T, svc *service.Container, path, libraryID stri
}
return payload
}
func requestLibrarySeries(t *testing.T, svc *service.Container, path, libraryID string) seriesListResponse {
t.Helper()
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Set(middleware.CtxUserID, "user-1")
c.Set(middleware.CtxUserRole, "user")
c.Params = gin.Params{{Key: "id", Value: libraryID}}
c.Request = httptest.NewRequest(http.MethodGet, path, nil)
listLibrarySeriesHandler(svc)(c)
if w.Code != http.StatusOK {
t.Fatalf("GET %s status = %d body=%s", path, w.Code, w.Body.String())
}
var payload seriesListResponse
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
t.Fatalf("decode series list: %v", err)
}
return payload
}
func requestLibrarySeriesEpisodes(t *testing.T, svc *service.Container, path, libraryID string) seriesEpisodesResponse {
t.Helper()
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Set(middleware.CtxUserID, "user-1")
c.Set(middleware.CtxUserRole, "user")
c.Params = gin.Params{{Key: "id", Value: libraryID}}
c.Request = httptest.NewRequest(http.MethodGet, path, nil)
listLibrarySeriesEpisodesHandler(svc)(c)
if w.Code != http.StatusOK {
t.Fatalf("GET %s status = %d body=%s", path, w.Code, w.Body.String())
}
var payload seriesEpisodesResponse
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
t.Fatalf("decode series episodes: %v", err)
}
return payload
}
+67 -2
View File
@@ -74,7 +74,7 @@ func externalPlayersHandler(svc *service.Container) gin.HandlerFunc {
return
}
token := externalPlaybackToken(c, svc, m.ID, m.DurationSec)
streamURL := absoluteRequestURL(c, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c))
streamURL := externalPlaybackURL(c, svc, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c))
escapedStream := url.QueryEscape(streamURL)
c.JSON(http.StatusOK, gin.H{
"url": streamURL,
@@ -100,7 +100,7 @@ func externalURLHandler(svc *service.Container) gin.HandlerFunc {
}
token := externalPlaybackToken(c, svc, m.ID, m.DurationSec)
c.JSON(http.StatusOK, gin.H{
"url": absoluteRequestURL(c, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c)),
"url": externalPlaybackURL(c, svc, "/api/stream/"+m.ID+"?token="+url.QueryEscape(token)+externalProfileQuery(c)),
})
}
}
@@ -137,6 +137,71 @@ func externalPlaybackToken(c *gin.Context, svc *service.Container, mediaID strin
return token
}
func externalPlaybackURL(c *gin.Context, svc *service.Container, path string) string {
if strings.HasPrefix(path, "http://") || strings.HasPrefix(path, "https://") {
return path
}
headerOrigin := sanitizedPublicOrigin(c.GetHeader("X-MediaStation-Public-Origin"))
if headerOrigin != "" && !isLocalPublicOrigin(headerOrigin) {
return joinOriginPath(headerOrigin, path)
}
if svc != nil {
if origin := sanitizedPublicOrigin(service.PublicServerURL(c.Request.Context(), svc.Repo, svc.Cfg)); origin != "" {
return joinOriginPath(origin, path)
}
}
if headerOrigin != "" {
return joinOriginPath(headerOrigin, path)
}
return absoluteRequestURL(c, path)
}
func isLocalPublicOrigin(origin string) bool {
u, err := url.Parse(origin)
if err != nil || u == nil {
return false
}
host := strings.ToLower(strings.Trim(u.Hostname(), "[]"))
switch host {
case "localhost", "127.0.0.1", "::1":
return true
default:
return strings.HasPrefix(host, "127.")
}
}
func sanitizedPublicOrigin(raw string) string {
raw = strings.TrimSpace(strings.Split(raw, ",")[0])
if raw == "" {
return ""
}
u, err := url.Parse(raw)
if err != nil || u == nil {
return ""
}
scheme := strings.ToLower(strings.TrimSpace(u.Scheme))
if scheme != "http" && scheme != "https" {
return ""
}
if strings.TrimSpace(u.Host) == "" {
return ""
}
u.Scheme = scheme
u.User = nil
u.Path = ""
u.RawPath = ""
u.RawQuery = ""
u.Fragment = ""
return strings.TrimRight(u.String(), "/")
}
func joinOriginPath(origin, path string) string {
if !strings.HasPrefix(path, "/") {
path = "/" + path
}
return strings.TrimRight(origin, "/") + path
}
func absoluteRequestURL(c *gin.Context, path string) string {
if strings.HasPrefix(path, "http://") || strings.HasPrefix(path, "https://") {
return path
+132
View File
@@ -84,6 +84,137 @@ func TestExternalURLUsesMediaScopedPlaybackToken(t *testing.T) {
}
}
func TestExternalURLPrefersBrowserPublicOriginOverForwardedSource(t *testing.T) {
router, _, secret := newPlaybackScopeTestRouter(t)
loginToken := signedTestToken(t, secret)
req := httptest.NewRequest(http.MethodGet, "http://origin.internal/api/playback/media-1/external-url", nil)
req.Header.Set("Authorization", "Bearer "+loginToken)
req.Header.Set("X-Forwarded-Proto", "https")
req.Header.Set("X-Forwarded-Host", "media.v6.agonyz.dpdns.org")
req.Header.Set("X-MediaStation-Public-Origin", "https://media.agonyz.dpdns.org")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
var payload struct {
URL string `json:"url"`
}
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
t.Fatalf("decode: %v", err)
}
streamURL, err := url.Parse(payload.URL)
if err != nil {
t.Fatalf("parse stream url: %v", err)
}
if got, want := streamURL.Scheme+"://"+streamURL.Host, "https://media.agonyz.dpdns.org"; got != want {
t.Fatalf("external url origin = %q, want %q; full url=%s", got, want, payload.URL)
}
if strings.Contains(payload.URL, "media.v6.agonyz.dpdns.org") {
t.Fatalf("external url should not use forwarded source host: %s", payload.URL)
}
}
func TestExternalPlayersSanitizeBrowserPublicOrigin(t *testing.T) {
router, _, secret := newPlaybackScopeTestRouter(t)
loginToken := signedTestToken(t, secret)
req := httptest.NewRequest(http.MethodGet, "http://origin.internal/api/playback/media-1/external-players", nil)
req.Header.Set("Authorization", "Bearer "+loginToken)
req.Header.Set("X-MediaStation-Public-Origin", "https://user:pass@media.agonyz.dpdns.org/sneaky/path?x=1#frag")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
var payload struct {
URL string `json:"url"`
Players []struct {
Name string `json:"name"`
URL string `json:"url"`
} `json:"players"`
}
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
t.Fatalf("decode: %v", err)
}
if !strings.HasPrefix(payload.URL, "https://media.agonyz.dpdns.org/api/stream/media-1?") {
t.Fatalf("sanitized stream url = %q", payload.URL)
}
if strings.Contains(payload.URL, "user:pass") || strings.Contains(payload.URL, "sneaky") || strings.Contains(payload.URL, "x=1") || strings.Contains(payload.URL, "#frag") {
t.Fatalf("stream url contains unsafe origin components: %s", payload.URL)
}
for _, player := range payload.Players {
if !strings.Contains(player.URL, "media.agonyz.dpdns.org") {
t.Fatalf("%s player url does not include sanitized public host: %s", player.Name, player.URL)
}
if strings.Contains(player.URL, "user:pass") || strings.Contains(player.URL, "sneaky") {
t.Fatalf("%s player url contains unsafe origin components: %s", player.Name, player.URL)
}
}
}
func TestExternalURLFallsBackToConfiguredPublicServerURL(t *testing.T) {
router, svc, secret := newPlaybackScopeTestRouter(t)
loginToken := signedTestToken(t, secret)
if err := svc.Repo.Setting.Set(t.Context(), "app.server_url", "https://public.example.test"); err != nil {
t.Fatalf("set public url: %v", err)
}
req := httptest.NewRequest(http.MethodGet, "http://origin.internal/api/playback/media-1/external-url", nil)
req.Header.Set("Authorization", "Bearer "+loginToken)
req.Header.Set("X-Forwarded-Proto", "https")
req.Header.Set("X-Forwarded-Host", "source.example.test")
req.Header.Set("X-MediaStation-Public-Origin", "javascript:alert(1)")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
var payload struct {
URL string `json:"url"`
}
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
t.Fatalf("decode: %v", err)
}
if !strings.HasPrefix(payload.URL, "https://public.example.test/api/stream/media-1?") {
t.Fatalf("external url = %q, want configured public origin", payload.URL)
}
}
func TestExternalURLPrefersConfiguredPublicServerURLOverLocalBrowserOrigin(t *testing.T) {
router, svc, secret := newPlaybackScopeTestRouter(t)
loginToken := signedTestToken(t, secret)
if err := svc.Repo.Setting.Set(t.Context(), "app.server_url", "https://media.example.test"); err != nil {
t.Fatalf("set public url: %v", err)
}
req := httptest.NewRequest(http.MethodGet, "http://127.0.0.1:8080/api/playback/media-1/external-url", nil)
req.Header.Set("Authorization", "Bearer "+loginToken)
req.Header.Set("X-MediaStation-Public-Origin", "http://127.0.0.1:8080")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
var payload struct {
URL string `json:"url"`
}
if err := json.Unmarshal(w.Body.Bytes(), &payload); err != nil {
t.Fatalf("decode: %v", err)
}
if !strings.HasPrefix(payload.URL, "https://media.example.test/api/stream/media-1?") {
t.Fatalf("external url = %q, want configured public origin instead of localhost", payload.URL)
}
if strings.Contains(payload.URL, "127.0.0.1:8080") {
t.Fatalf("external url should not keep local browser origin when public url is configured: %s", payload.URL)
}
}
func TestScopedPlaybackTokenCannotStreamAnotherMedia(t *testing.T) {
router, svc, _ := newPlaybackScopeTestRouter(t)
user, err := svc.Repo.User.FindByID(t.Context(), "user-1")
@@ -248,6 +379,7 @@ func newPlaybackScopeTestRouter(t *testing.T) (*gin.Engine, *service.Container,
api := router.Group("/api")
api.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
api.GET("/playback/:id/external-url", externalURLHandler(svc))
api.GET("/playback/:id/external-players", externalPlayersHandler(svc))
api.GET("/stream/:id", streamHandler(svc))
api.GET("/cloud/play/:type", cloudPlayHandler(svc))
return router, svc, cfg.Secrets.JWTSecret
+2
View File
@@ -57,6 +57,7 @@ func (h *RefreshHandler) RefreshToken(c *gin.Context) {
return
}
setAccessTokenCookie(c, tokens.AccessToken, int(tokens.ExpiresIn))
c.JSON(http.StatusOK, gin.H{
"code": 0,
"message": "ok",
@@ -72,6 +73,7 @@ func (h *RefreshHandler) RefreshToken(c *gin.Context) {
// Logout 登出当前用户。
// POST /api/auth/logout
func (h *RefreshHandler) Logout(c *gin.Context) {
clearAccessTokenCookie(c)
userID := c.GetString("ctx_user_id")
if userID == "" {
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": nil})
+22 -6
View File
@@ -17,13 +17,21 @@ import (
// 异步执行, 立即返回 202;通过 WS hub "scrape" topic 推送进度。
func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
options, err := scrapeOptionsFromRequest(c, true)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
return
}
task := startScrapeHTTPTask(svc, "全库修复并重刮", "", "")
go func() {
result, err := svc.RepairAndRescrapeAllLibraries(context.Background())
go func(options service.ScrapeOptions) {
result, err := svc.RepairAndRescrapeAllLibraries(context.Background(), options)
metrics := map[string]int64{
"repaired": int64(result.Repaired),
"reclassified": int64(result.Reclassified),
"libraries": int64(result.Libraries),
"matched": int64(result.Matched),
"processed": int64(result.Processed),
"errors": int64(result.Errors),
"reset": int64(result.Reset),
}
stage := "completed"
@@ -33,7 +41,7 @@ func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
message = "全库修复并重刮失败"
}
finishHTTPTask(task, err, stage, message, metrics, nil)
}()
}(options)
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
}
}
@@ -46,13 +54,21 @@ func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
func repairAndRescrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
libraryID := c.Param("id")
options, err := scrapeOptionsFromRequest(c, true)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
return
}
task := startScrapeHTTPTask(svc, "媒体库修复并重刮", "", "")
go func() {
result, err := svc.RepairAndRescrapeLibrary(context.Background(), libraryID)
go func(options service.ScrapeOptions) {
result, err := svc.RepairAndRescrapeLibrary(context.Background(), libraryID, options)
metrics := map[string]int64{
"repaired": int64(result.Repaired),
"reclassified": int64(result.Reclassified),
"libraries": int64(result.Libraries),
"matched": int64(result.Matched),
"processed": int64(result.Processed),
"errors": int64(result.Errors),
"reset": int64(result.Reset),
}
stage := "completed"
@@ -62,7 +78,7 @@ func repairAndRescrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
message = "媒体库修复并重刮失败"
}
finishHTTPTask(task, err, stage, message, metrics, nil)
}()
}(options)
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
}
}
+42 -20
View File
@@ -10,10 +10,24 @@ import (
)
func registerAdminRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.Container) {
// Admin-only endpoints.
admin := api.Group("/admin")
admin.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret), middleware.AdminRequired())
{
registerAdminUserRoutes(admin, svc)
registerAdminPermissionRoutes(admin, svc)
registerAdminStorageRoutes(admin, svc)
registerAdminCloudRoutes(admin, svc)
registerAdminDownloadClientRoutes(admin, svc)
registerAdminSystemRoutes(admin, svc)
registerAdminBackupRoutes(admin, svc)
registerAdminNotificationRoutes(admin, svc)
registerAdminTelegramRoutes(admin, svc)
registerAdminOrganizerRoutes(admin, svc)
registerAdminRepairRoutes(admin, svc)
registerAdminAPIConfigRoutes(admin, svc)
registerAdminSchedulerRoutes(admin, svc)
}
func registerAdminUserRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/users", listUsersHandler(svc))
admin.POST("/users", createUserHandler(svc))
admin.PATCH("/users/:id", updateUserHandler(svc))
@@ -24,84 +38,92 @@ func registerAdminRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.
admin.GET("/settings", listSettingsHandler(svc))
admin.PUT("/settings", updateSettingHandler(svc))
admin.GET("/logs", recentLogsHandler(svc))
}
// Permissions admin.
func registerAdminPermissionRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/users/:id/permissions", getUserPermissionsHandler(svc))
admin.PUT("/users/:id/permissions", updateUserPermissionsHandler(svc))
admin.POST("/users/:id/permissions/reset", resetUserPermissionsHandler(svc))
}
// Storage configs (Alist / S3 / WebDAV / 网盘).
func registerAdminStorageRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/storage/status", listStorageConfigsHandler(svc))
admin.GET("/storage/:type", getStorageConfigHandler(svc))
admin.PUT("/storage/:type", saveStorageConfigHandler(svc))
admin.POST("/storage/:type/test", testStorageConfigHandler(svc))
admin.POST("/storage/:type/logout", logoutStorageConfigHandler(svc))
admin.POST("/storage/:type/upload-local", storageUploadLocalHandler(svc))
}
// Cloud disk (115 / 夸克) browsing, QR login and 302 import.
func registerAdminCloudRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.POST("/cloud/scan-all", cloudScanAllHandler(svc))
admin.POST("/cloud/scan/cancel", cloudScanCancelHandler(svc))
admin.GET("/cloud/scan/status", cloudScanStatusHandler(svc))
admin.GET("/cloud/:type/list", cloudListHandler(svc))
admin.POST("/cloud/:type/mkdir", cloudMkdirHandler(svc))
admin.PUT("/cloud/:type/rename", cloudRenameHandler(svc))
admin.POST("/cloud/:type/import", cloudImportHandler(svc))
admin.POST("/cloud/:type/mount", cloudMountHandler(svc))
admin.POST("/cloud/:type/qr/start", cloud115QRStartHandler(svc))
admin.POST("/cloud/:type/qr/poll", cloud115QRPollHandler(svc))
}
// Download client CRUD.
func registerAdminDownloadClientRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/download/clients", listDownloadClientsHandler(svc))
admin.POST("/download/clients", createDownloadClientHandler(svc))
admin.PUT("/download/clients/:id", updateDownloadClientHandler(svc))
admin.DELETE("/download/clients/:id", deleteDownloadClientHandler(svc))
admin.POST("/download/clients/:id/test", testDownloadClientHandler(svc))
admin.GET("/download/aria2/stats", aria2StatsHandler(svc))
}
// System scheduler trigger alias.
func registerAdminSystemRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.POST("/system/scheduler/:name/trigger", schedulerTriggerHandler(svc))
}
// Database backup.
func registerAdminBackupRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/backups", listBackupsHandler(svc))
admin.POST("/backups", createBackupHandler(svc))
admin.DELETE("/backups", deleteBackupHandler(svc))
admin.POST("/backups/restore", restoreBackupHandler(svc))
}
// Notifications (test endpoint).
func registerAdminNotificationRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.POST("/notify/test", notifyTestHandler(svc))
// Notify channels CRUD + per-channel test.
admin.GET("/notify/channels", listNotifyChannelsHandler(svc))
admin.POST("/notify/channels", createNotifyChannelHandler(svc))
admin.PUT("/notify/channels/:id", updateNotifyChannelHandler(svc))
admin.DELETE("/notify/channels/:id", deleteNotifyChannelHandler(svc))
admin.POST("/notify/channels/:id/test", testNotifyChannelHandler(svc))
}
// Telegram Bot webhook management.
func registerAdminTelegramRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/telegram/webhook", telegramGetWebhookHandler(svc))
admin.POST("/telegram/webhook", telegramSetWebhookHandler(svc))
admin.POST("/telegram/polling/start", telegramStartPollingHandler(svc))
admin.POST("/telegram/polling/stop", telegramStopPollingHandler(svc))
}
// File organizer.
func registerAdminOrganizerRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.POST("/media/:id/organize", organizeMediaHandler(svc))
admin.POST("/libraries/:id/organize", organizeLibraryHandler(svc))
admin.GET("/organize/sources", organizeSourcesHandler(svc))
admin.POST("/organize/source", organizeDirectoryHandler(svc))
}
// 全库修复+重刮:从路径占位符回填缺失外部 ID,然后批量重刮整库。
func registerAdminRepairRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.POST("/media/repair-rescrape", repairAndRescrapeAllHandler(svc))
// 单库修复+重刮:只对指定媒体库回填占位符外部 ID 并重刮。
admin.POST("/libraries/:id/repair-rescrape", repairAndRescrapeLibraryHandler(svc))
}
// API key management (encrypted at rest).
func registerAdminAPIConfigRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/api-configs", listAPIConfigsHandler(svc))
admin.GET("/api-configs/:provider", getAPIConfigHandler(svc))
admin.PUT("/api-configs/:provider", updateAPIConfigHandler(svc))
admin.DELETE("/api-configs/:provider", deleteAPIConfigHandler(svc))
}
// Scheduled jobs.
func registerAdminSchedulerRoutes(admin *gin.RouterGroup, svc *service.Container) {
admin.GET("/scheduler", schedulerStatusHandler(svc))
admin.POST("/scheduler/:name/run", schedulerRunHandler(svc))
}
}
+45
View File
@@ -0,0 +1,45 @@
package handler
import (
"testing"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func TestAdminRouteSurfacesAreRegistered(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
Register(router, &config.Config{
Secrets: config.SecretsConfig{JWTSecret: "test-secret"},
}, zap.NewNop(), &service.Container{Log: zap.NewNop()})
routes := map[string]bool{}
for _, route := range router.Routes() {
routes[route.Method+" "+route.Path] = true
}
for _, want := range []string{
"GET /api/admin/users",
"GET /api/admin/users/:id/permissions",
"GET /api/admin/storage/status",
"GET /api/admin/cloud/:type/list",
"GET /api/admin/download/clients",
"POST /api/admin/system/scheduler/:name/trigger",
"POST /api/admin/backups",
"GET /api/admin/notify/channels",
"GET /api/admin/telegram/webhook",
"GET /api/admin/organize/sources",
"POST /api/admin/media/repair-rescrape",
"GET /api/admin/api-configs",
"POST /api/admin/scheduler/:name/run",
} {
if !routes[want] {
t.Fatalf("%s route is not registered", want)
}
}
}
+27 -253
View File
@@ -10,261 +10,35 @@ import (
)
func registerAuthenticatedRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.Container) {
// Authenticated endpoints.
authed := api.Group("/")
authed.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
authed.Use(activeUserRequired(svc))
{
authed.GET("/me", meHandler(svc))
authed.PATCH("/me", updateProfileHandler(svc))
authed.POST("/me/password", changePasswordHandler(svc))
authed.POST("/me/logout", logoutHandler(svc))
// Permissions.
authed.GET("/auth/permissions", getMyPermissionsHandler(svc))
// 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))
// Libraries.
authed.GET("/libraries", listLibrariesHandler(svc))
authed.POST("/libraries", middleware.AdminRequired(), createLibraryHandler(svc))
authed.DELETE("/libraries/:id", middleware.AdminRequired(), deleteLibraryHandler(svc))
authed.POST("/libraries/:id/scan", middleware.AdminRequired(), scanLibraryHandler(svc))
authed.POST("/libraries/:id/scrape", middleware.AdminRequired(), scrapeLibraryHandler(svc))
authed.GET("/libraries/:id/media", listMediaHandler(svc))
authed.GET("/libraries/:id/seasons", listSeasonsHandler(svc))
// Media.
authed.GET("/media/:id", getMediaHandler(svc))
authed.GET("/media", searchMediaHandler(svc))
authed.PATCH("/media/:id/metadata", middleware.AdminRequired(), updateMediaMetadataHandler(svc))
authed.POST("/media/:id/scrape", middleware.AdminRequired(), scrapeOneHandler(svc))
authed.GET("/media/:id/scrape/search", middleware.AdminRequired(), manualScrapeSearchHandler(svc))
authed.POST("/media/:id/scrape/apply", middleware.AdminRequired(), manualScrapeApplyOneHandler(svc))
authed.POST("/media/scrape/apply", middleware.AdminRequired(), manualScrapeApplyBatchHandler(svc))
authed.POST("/media/:id/probe", middleware.AdminRequired(), reprobeHandler(svc))
authed.DELETE("/media/:id", middleware.AdminRequired(), deleteMediaHandler(svc))
authed.POST("/media/:id/restore", middleware.AdminRequired(), restoreMediaHandler(svc))
authed.DELETE("/media/:id/purge", middleware.AdminRequired(), purgeMediaHandler(svc))
authed.GET("/media/:id/subtitles", listSubtitlesHandler(svc))
authed.GET("/subtitles/:id", serveSubtitleHandler(svc))
authed.POST("/media/:id/nfo", middleware.AdminRequired(), exportNFOHandler(svc))
authed.POST("/libraries/:id/nfo", middleware.AdminRequired(), exportLibraryNFOHandler(svc))
// Streaming.
authed.GET("/stream/:id", streamHandler(svc))
authed.HEAD("/stream/:id", streamHandler(svc))
authed.GET("/hls/:id/index.m3u8", hlsPlaylistHandler(svc))
authed.GET("/hls/:id/:seg", hlsSegmentHandler(svc))
authed.DELETE("/hls/:id", stopTranscodeHandler(svc))
// Cloud-disk 302 playback redirect (resolves a fresh direct link).
authed.GET("/cloud/play/:type", cloudPlayHandler(svc))
authed.HEAD("/cloud/play/:type", cloudPlayHandler(svc))
// Image proxy (URL passed as ?url=...).
authed.GET("/img", imageProxyHandler(svc))
// History / favourites / playlists.
authed.GET("/history", recentHistoryHandler(svc))
authed.POST("/history", recordProgressHandler(svc))
authed.GET("/favourites", listFavouritesHandler(svc))
authed.POST("/favourites/:id", toggleFavouriteHandler(svc))
// Storage breakdown.
authed.GET("/storage", storageBreakdownHandler(svc))
authed.GET("/playlists", listPlaylistsHandler(svc))
authed.POST("/playlists", createPlaylistHandler(svc))
authed.GET("/playlists/:id", getPlaylistHandler(svc))
authed.POST("/playlists/:id/items", addPlaylistItemHandler(svc))
authed.DELETE("/playlists/:id/items/:media_id", removePlaylistItemHandler(svc))
authed.DELETE("/playlists/:id", deletePlaylistHandler(svc))
// Downloads.
authed.GET("/downloads", requirePermission(svc, "can_manage_downloads"), listDownloadsHandler(svc))
authed.POST("/downloads", requirePermission(svc, "can_manage_downloads"), addDownloadHandler(svc))
authed.DELETE("/downloads/:hash", requirePermission(svc, "can_manage_downloads"), deleteDownloadHandler(svc))
authed.POST("/downloads/relocate", requirePermission(svc, "can_manage_downloads"), relocateDownloadHandler(svc))
authed.POST("/downloads/reload", requirePermission(svc, "can_manage_downloads"), reloadDownloadConfigHandler(svc))
// Subscriptions.
authed.GET("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionsHandler(svc))
authed.GET("/subscriptions/history", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionHistoryHandler(svc))
authed.POST("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), createSubscriptionHandler(svc))
authed.DELETE("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), deleteSubscriptionHandler(svc))
authed.POST("/subscriptions/:id/restore", requirePermission(svc, "can_manage_subscriptions"), restoreSubscriptionHandler(svc))
authed.POST("/subscriptions/:id/run", requirePermission(svc, "can_manage_subscriptions"), runSubscriptionHandler(svc))
// Stats / dashboard.
authed.GET("/stats", statsHandler(svc))
authed.GET("/tasks", middleware.AdminRequired(), tasksHandler(svc))
// Discover (TMDb trending / popular).
authed.GET("/discover/trending", requirePermission(svc, "can_view_discover"), trendingHandler(svc))
authed.GET("/discover/popular", requirePermission(svc, "can_view_discover"), popularHandler(svc))
// AI.
authed.GET("/ai/status", requirePermission(svc, "can_use_ai"), aiStatusHandler(svc))
authed.POST("/ai/search", requirePermission(svc, "can_use_ai"), smartSearchHandler(svc))
authed.GET("/ai/recommend", requirePermission(svc, "can_use_ai"), aiRecommendHandler(svc))
// File browser (used by the library-path picker).
authed.GET("/files", middleware.AdminRequired(), browseFilesHandler(svc))
authed.POST("/files/folders", middleware.AdminRequired(), createFolderHandler(svc))
authed.PUT("/files/rename", middleware.AdminRequired(), renameFileHandler(svc))
authed.DELETE("/files", middleware.AdminRequired(), deleteFileHandler(svc))
authed.POST("/files/transfer", middleware.AdminRequired(), transferFileHandler(svc))
// DLNA discovery + cast.
authed.GET("/dlna/devices", dlnaListHandler(svc))
authed.POST("/dlna/cast", dlnaCastHandler(svc))
// STRM (URL-as-file).
authed.PUT("/media/:id/strm", middleware.AdminRequired(), setSTRMHandler(svc))
authed.DELETE("/media/:id/strm", middleware.AdminRequired(), clearSTRMHandler(svc))
authed.POST("/strm/import", middleware.AdminRequired(), importSTRMHandler(svc))
authed.POST("/strm/generate", middleware.AdminRequired(), generateSTRMHandler(svc))
// Duplicate finder.
authed.GET("/duplicates", middleware.AdminRequired(), listDuplicatesHandler(svc))
authed.POST("/duplicates/scan", middleware.AdminRequired(), detectDuplicatesHandler(svc))
authed.POST("/duplicates/unmark", middleware.AdminRequired(), unmarkDuplicatesHandler(svc))
// Site management + cross-site torrent search (via SiteHandler).
siteHandler := NewSiteHandler(svc)
authed.GET("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.ListSites)
authed.GET("/sites/types", requirePermission(svc, "can_manage_sites"), siteHandler.GetSiteTypes)
authed.GET("/sites/auth-types", requirePermission(svc, "can_manage_sites"), siteHandler.GetAuthTypes)
authed.POST("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.CreateSite)
authed.GET("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.GetSite)
authed.PUT("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.UpdateSite)
authed.DELETE("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.DeleteSite)
authed.POST("/sites/:id/test", requirePermission(svc, "can_manage_sites"), siteHandler.TestSite)
authed.GET("/sites/search", requirePermission(svc, "can_manage_sites"), siteSearchHandler(svc))
// Recycle bin.
authed.GET("/recycle", middleware.AdminRequired(), listRecycleHandler(svc))
authed.POST("/recycle/restore", middleware.AdminRequired(), restoreMediaBatchHandler(svc))
authed.POST("/recycle/purge", middleware.AdminRequired(), purgeMediaBatchHandler(svc))
authed.GET("/ws", wsHandler(svc))
// SSE event stream.
authed.GET("/events", sseHandler(svc))
// Scheduler.
authed.GET("/scheduler/tasks", schedulerListTasksHandler(svc))
authed.POST("/scheduler/tasks/:id/run", middleware.AdminRequired(), schedulerRunTaskHandler(svc))
authed.GET("/scheduler/status", schedulerGetStatusHandler(svc))
// ── Auxiliary endpoints used by the React UI rails ──
authed.GET("/media/recent", recentMediaHandler(svc))
authed.GET("/media/stats", mediaStatsHandler(svc))
// Watch history (extra surface beyond /history).
authed.GET("/watch-history", historyListHandler(svc))
authed.GET("/watch-history/stats", historyStatsHandler(svc))
authed.GET("/watch-history/continue", historyContinueHandler(svc))
authed.DELETE("/watch-history", historyDeleteHandler(svc))
authed.DELETE("/watch-history/:id", historyDeleteOneHandler(svc))
// Multi-section TMDb feed used by DiscoverPage.
authed.GET("/discover/sections", requirePermission(svc, "can_view_discover"), discoverSectionsHandler(svc))
authed.GET("/discover/feed", requirePermission(svc, "can_view_discover"), discoverFeedHandler(svc))
// System metadata + read-only scheduler view.
authed.GET("/system/info", systemInfoHandler(svc))
authed.GET("/system/status", systemStatusHandler(svc))
authed.GET("/system/scheduler", systemSchedulerHandler(svc))
// Richer dashboard rails.
authed.GET("/stats/overview", statsOverviewHandler(svc))
authed.GET("/stats/trend", statsTrendHandler(svc))
authed.GET("/stats/top-content", statsTopContentHandler(svc))
authed.GET("/stats/libraries", statsLibrariesHandler(svc))
authed.GET("/stats/monitor", statsMonitorHandler(svc))
// Multi-persona play profiles (caller-scoped).
authed.GET("/play-profiles", listPlayProfilesHandler(svc))
authed.POST("/play-profiles", createPlayProfileHandler(svc))
authed.PUT("/play-profiles/:id", updatePlayProfileHandler(svc))
authed.POST("/play-profiles/:id/verify-pin", verifyPlayProfilePINHandler(svc))
authed.DELETE("/play-profiles/:id", deletePlayProfileHandler(svc))
// ── Search aliases ──
authed.GET("/search", searchUnifiedHandler(svc))
authed.GET("/search/advanced", searchAdvancedHandler(svc))
authed.GET("/search/tmdb", searchTMDbHandler(svc))
authed.GET("/search/sites", searchSitesHandler(svc))
// ── System extras ──
authed.GET("/system/config", listSystemConfigHandler(svc))
authed.GET("/settings/schema", schemaHandler(svc))
authed.GET("/system/events/ticket", systemEventsTicketHandler(svc))
// ── Per-user stats ──
authed.GET("/stats/user/:id", statsUserHandler(svc))
authed.GET("/stats/top-users", statsTopUsersHandler(svc))
authed.POST("/stats/play", statsPlayHandler(svc))
// ── Sites extras ──
authed.GET("/sites/:id/resource", requirePermission(svc, "can_manage_sites"), siteResourceHandler(svc))
authed.GET("/sites/:id/userdata", requirePermission(svc, "can_manage_sites"), siteUserdataHandler(svc))
// ── Subscription extras ──
authed.PUT("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), updateSubscriptionHandler(svc))
authed.POST("/subscriptions/:id/search", requirePermission(svc, "can_manage_subscriptions"), searchSubscriptionHandler(svc))
// ── Playlist extras ──
authed.POST("/playlists/:id/reorder", reorderPlaylistHandler(svc))
authed.DELETE("/playlists/:id/items/by-id/:item_id", deletePlaylistItemByIDHandler(svc))
// ── DLNA per-renderer control ──
authed.POST("/dlna/:uuid/play", dlnaPlayHandler(svc))
authed.POST("/dlna/:uuid/pause", dlnaPauseHandler(svc))
authed.POST("/dlna/:uuid/stop", dlnaStopHandler(svc))
authed.GET("/dlna/:uuid/status", dlnaStatusHandler(svc))
// ── Media favourite alias surface ──
authed.GET("/favorites", listFavoritesAliasHandler(svc))
authed.POST("/media/:id/favorite", addMediaFavoriteHandler(svc))
authed.DELETE("/media/:id/favorite", removeMediaFavoriteHandler(svc))
authed.GET("/media/:id/favorite/status", getMediaFavoriteStatusHandler(svc))
authed.POST("/media/:id/ai-scrape", requirePermission(svc, "can_rescrape"), aiScrapeMediaHandler(svc))
authed.POST("/media/scrape/test", requirePermission(svc, "can_rescrape"), scrapeTestHandler(svc))
authed.POST("/media/organize", requirePermission(svc, "can_manage_files"), organizeBulkHandler(svc))
// ── Playback metadata + external player handoff ──
authed.GET("/playback/:id/info", playbackInfoHandler(svc))
authed.POST("/playback/:id/progress", playbackProgressHandler(svc))
authed.GET("/playback/:id/external-players", externalPlayersHandler(svc))
authed.GET("/playback/:id/external-url", externalURLHandler(svc))
authed.GET("/playback/transcode/:job_id/status", transcodeStatusHandler(svc))
// ── Download task ops + sync triggers ──
authed.POST("/download/:id/pause", requirePermission(svc, "can_manage_downloads"), downloadPauseHandler(svc))
authed.POST("/download/:id/resume", requirePermission(svc, "can_manage_downloads"), downloadResumeHandler(svc))
authed.POST("/download/:id/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeOneHandler(svc))
authed.POST("/download/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeAllHandler(svc))
authed.POST("/download/sync", requirePermission(svc, "can_manage_downloads"), downloadSyncHandler(svc))
authed.POST("/download/start-auto-sync", requirePermission(svc, "can_manage_downloads"), downloadAutoSyncHandler(svc))
authed.GET("/download/tasks", requirePermission(svc, "can_manage_downloads"), downloadTasksAliasHandler(svc))
// ── Assistant (multi-turn AI chat) ──
authed.GET("/admin/assistant/sessions", listAssistantSessionsHandler(svc))
authed.POST("/admin/assistant/sessions", createAssistantSessionHandler(svc))
authed.GET("/admin/assistant/session/:id", getAssistantSessionHandler(svc))
authed.DELETE("/admin/assistant/session/:id", deleteAssistantSessionHandler(svc))
authed.POST("/admin/assistant/chat", assistantChatHandler(svc))
authed.POST("/admin/assistant/execute", assistantExecuteHandler(svc))
authed.POST("/admin/assistant/undo/:op_id", assistantUndoHandler(svc))
authed.GET("/admin/assistant/history", assistantHistoryHandler(svc))
}
registerAuthedUserAndLicenseRoutes(authed, svc)
registerAuthedLibraryRoutes(authed, svc)
registerAuthedMediaRoutes(authed, svc)
registerAuthedPlaybackAndProxyRoutes(authed, svc)
registerAuthedCollectionRoutes(authed, svc)
registerAuthedDownloadRoutes(authed, svc)
registerAuthedSubscriptionRoutes(authed, svc)
registerAuthedStatsDiscoveryAndAIRoutes(authed, svc)
registerAuthedFileRoutes(authed, svc)
registerAuthedDLNARoutes(authed, svc)
registerAuthedSTRMRoutes(authed, svc)
registerAuthedDuplicateRoutes(authed, svc)
registerAuthedSiteRoutes(authed, svc)
registerAuthedRecycleAndRealtimeRoutes(authed, svc)
registerAuthedSchedulerRoutes(authed, svc)
registerAuthedUISurfaceRoutes(authed, svc)
registerAuthedSearchRoutes(authed, svc)
registerAuthedSystemExtraRoutes(authed, svc)
registerAuthedStatsExtraRoutes(authed, svc)
registerAuthedSitesExtraRoutes(authed, svc)
registerAuthedSubscriptionExtraRoutes(authed, svc)
registerAuthedPlaylistExtraRoutes(authed, svc)
registerAuthedDLNAControlRoutes(authed, svc)
registerAuthedFavoriteAndMediaActionRoutes(authed, svc)
registerAuthedPlaybackExtraRoutes(authed, svc)
registerAuthedDownloadOpsRoutes(authed, svc)
registerAuthedAssistantRoutes(authed, svc)
}
@@ -0,0 +1,84 @@
package handler
import (
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func registerAuthedUserAndLicenseRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/me", meHandler(svc))
authed.PATCH("/me", updateProfileHandler(svc))
authed.POST("/me/password", changePasswordHandler(svc))
authed.POST("/me/logout", logoutHandler(svc))
authed.GET("/auth/permissions", getMyPermissionsHandler(svc))
authed.GET("/license/status", middleware.AdminRequired(), licenseStatusHandler(svc))
authed.POST("/license/activate", middleware.AdminRequired(), licenseActivateHandler(svc))
authed.POST("/license/heartbeat", middleware.AdminRequired(), licenseHeartbeatHandler(svc))
}
func registerAuthedLibraryRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/libraries", listLibrariesHandler(svc))
authed.POST("/libraries", middleware.AdminRequired(), createLibraryHandler(svc))
authed.DELETE("/libraries/:id", middleware.AdminRequired(), deleteLibraryHandler(svc))
authed.POST("/libraries/:id/scan", middleware.AdminRequired(), scanLibraryHandler(svc))
authed.POST("/libraries/:id/scrape", middleware.AdminRequired(), scrapeLibraryHandler(svc))
authed.GET("/libraries/:id/media", listMediaHandler(svc))
authed.GET("/libraries/:id/series", listLibrarySeriesHandler(svc))
authed.GET("/libraries/:id/series/episodes", listLibrarySeriesEpisodesHandler(svc))
authed.GET("/libraries/:id/seasons", listSeasonsHandler(svc))
}
func registerAuthedMediaRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/media/:id", getMediaHandler(svc))
authed.GET("/media", searchMediaHandler(svc))
authed.PATCH("/media/:id/metadata", middleware.AdminRequired(), updateMediaMetadataHandler(svc))
authed.POST("/media/:id/scrape", middleware.AdminRequired(), scrapeOneHandler(svc))
authed.GET("/media/:id/scrape/search", middleware.AdminRequired(), manualScrapeSearchHandler(svc))
authed.POST("/media/:id/scrape/apply", middleware.AdminRequired(), manualScrapeApplyOneHandler(svc))
authed.POST("/media/scrape/apply", middleware.AdminRequired(), manualScrapeApplyBatchHandler(svc))
authed.POST("/media/:id/probe", middleware.AdminRequired(), reprobeHandler(svc))
authed.DELETE("/media/:id", middleware.AdminRequired(), deleteMediaHandler(svc))
authed.POST("/media/:id/restore", middleware.AdminRequired(), restoreMediaHandler(svc))
authed.DELETE("/media/:id/purge", middleware.AdminRequired(), purgeMediaHandler(svc))
authed.GET("/media/:id/subtitles", listSubtitlesHandler(svc))
authed.GET("/subtitles/:id", serveSubtitleHandler(svc))
authed.POST("/media/:id/nfo", middleware.AdminRequired(), exportNFOHandler(svc))
authed.POST("/libraries/:id/nfo", middleware.AdminRequired(), exportLibraryNFOHandler(svc))
}
func registerAuthedPlaybackAndProxyRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/stream/:id", streamHandler(svc))
authed.HEAD("/stream/:id", streamHandler(svc))
authed.GET("/hls/:id/index.m3u8", hlsPlaylistHandler(svc))
authed.GET("/hls/:id/:seg", hlsSegmentHandler(svc))
authed.DELETE("/hls/:id", stopTranscodeHandler(svc))
authed.GET("/cloud/play/:type", cloudPlayHandler(svc))
authed.HEAD("/cloud/play/:type", cloudPlayHandler(svc))
authed.GET("/img/cloud/:type", cloudArtworkProxyHandler(svc))
authed.HEAD("/img/cloud/:type", cloudArtworkProxyHandler(svc))
authed.GET("/img", imageProxyHandler(svc))
}
func registerAuthedCollectionRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/history", recentHistoryHandler(svc))
authed.POST("/history", recordProgressHandler(svc))
authed.GET("/favourites", listFavouritesHandler(svc))
authed.POST("/favourites/:id", toggleFavouriteHandler(svc))
authed.GET("/storage", storageBreakdownHandler(svc))
authed.GET("/playlists", listPlaylistsHandler(svc))
authed.POST("/playlists", createPlaylistHandler(svc))
authed.GET("/playlists/:id", getPlaylistHandler(svc))
authed.POST("/playlists/:id/items", addPlaylistItemHandler(svc))
authed.DELETE("/playlists/:id/items/:media_id", removePlaylistItemHandler(svc))
authed.DELETE("/playlists/:id", deletePlaylistHandler(svc))
}
@@ -0,0 +1,117 @@
package handler
import (
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func registerAuthedUISurfaceRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/media/recent", recentMediaHandler(svc))
authed.GET("/media/stats", mediaStatsHandler(svc))
authed.GET("/watch-history", historyListHandler(svc))
authed.GET("/watch-history/stats", historyStatsHandler(svc))
authed.GET("/watch-history/continue", historyContinueHandler(svc))
authed.DELETE("/watch-history", historyDeleteHandler(svc))
authed.DELETE("/watch-history/:id", historyDeleteOneHandler(svc))
authed.GET("/discover/sections", requirePermission(svc, "can_view_discover"), discoverSectionsHandler(svc))
authed.GET("/discover/feed", requirePermission(svc, "can_view_discover"), discoverFeedHandler(svc))
authed.GET("/system/info", systemInfoHandler(svc))
authed.GET("/system/status", systemStatusHandler(svc))
authed.GET("/system/scheduler", systemSchedulerHandler(svc))
authed.GET("/stats/overview", statsOverviewHandler(svc))
authed.GET("/stats/trend", statsTrendHandler(svc))
authed.GET("/stats/top-content", statsTopContentHandler(svc))
authed.GET("/stats/libraries", statsLibrariesHandler(svc))
authed.GET("/stats/monitor", statsMonitorHandler(svc))
authed.GET("/play-profiles", listPlayProfilesHandler(svc))
authed.POST("/play-profiles", createPlayProfileHandler(svc))
authed.PUT("/play-profiles/:id", updatePlayProfileHandler(svc))
authed.POST("/play-profiles/:id/verify-pin", verifyPlayProfilePINHandler(svc))
authed.DELETE("/play-profiles/:id", deletePlayProfileHandler(svc))
}
func registerAuthedSearchRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/search", searchUnifiedHandler(svc))
authed.GET("/search/advanced", searchAdvancedHandler(svc))
authed.GET("/search/tmdb", searchTMDbHandler(svc))
authed.GET("/search/sites", searchSitesHandler(svc))
}
func registerAuthedSystemExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/system/config", listSystemConfigHandler(svc))
authed.GET("/settings/schema", schemaHandler(svc))
authed.GET("/system/events/ticket", systemEventsTicketHandler(svc))
}
func registerAuthedStatsExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/stats/user/:id", statsUserHandler(svc))
authed.GET("/stats/top-users", statsTopUsersHandler(svc))
authed.POST("/stats/play", statsPlayHandler(svc))
}
func registerAuthedSitesExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/sites/:id/resource", requirePermission(svc, "can_manage_sites"), siteResourceHandler(svc))
authed.GET("/sites/:id/userdata", requirePermission(svc, "can_manage_sites"), siteUserdataHandler(svc))
}
func registerAuthedSubscriptionExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.PUT("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), updateSubscriptionHandler(svc))
authed.POST("/subscriptions/:id/search", requirePermission(svc, "can_manage_subscriptions"), searchSubscriptionHandler(svc))
}
func registerAuthedPlaylistExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.POST("/playlists/:id/reorder", reorderPlaylistHandler(svc))
authed.DELETE("/playlists/:id/items/by-id/:item_id", deletePlaylistItemByIDHandler(svc))
}
func registerAuthedDLNAControlRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.POST("/dlna/:uuid/play", dlnaPlayHandler(svc))
authed.POST("/dlna/:uuid/pause", dlnaPauseHandler(svc))
authed.POST("/dlna/:uuid/stop", dlnaStopHandler(svc))
authed.GET("/dlna/:uuid/status", dlnaStatusHandler(svc))
}
func registerAuthedFavoriteAndMediaActionRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/favorites", listFavoritesAliasHandler(svc))
authed.POST("/media/:id/favorite", addMediaFavoriteHandler(svc))
authed.DELETE("/media/:id/favorite", removeMediaFavoriteHandler(svc))
authed.GET("/media/:id/favorite/status", getMediaFavoriteStatusHandler(svc))
authed.POST("/media/:id/ai-scrape", requirePermission(svc, "can_rescrape"), aiScrapeMediaHandler(svc))
authed.POST("/media/scrape/test", requirePermission(svc, "can_rescrape"), scrapeTestHandler(svc))
authed.POST("/media/organize", requirePermission(svc, "can_manage_files"), organizeBulkHandler(svc))
}
func registerAuthedPlaybackExtraRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/playback/:id/info", playbackInfoHandler(svc))
authed.POST("/playback/:id/progress", playbackProgressHandler(svc))
authed.GET("/playback/:id/external-players", externalPlayersHandler(svc))
authed.GET("/playback/:id/external-url", externalURLHandler(svc))
authed.GET("/playback/transcode/:job_id/status", transcodeStatusHandler(svc))
}
func registerAuthedDownloadOpsRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.POST("/download/:id/pause", requirePermission(svc, "can_manage_downloads"), downloadPauseHandler(svc))
authed.POST("/download/:id/resume", requirePermission(svc, "can_manage_downloads"), downloadResumeHandler(svc))
authed.POST("/download/:id/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeOneHandler(svc))
authed.POST("/download/organize", requirePermission(svc, "can_manage_files"), downloadOrganizeAllHandler(svc))
authed.POST("/download/sync", requirePermission(svc, "can_manage_downloads"), downloadSyncHandler(svc))
authed.POST("/download/start-auto-sync", requirePermission(svc, "can_manage_downloads"), downloadAutoSyncHandler(svc))
authed.GET("/download/tasks", requirePermission(svc, "can_manage_downloads"), downloadTasksAliasHandler(svc))
}
func registerAuthedAssistantRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/admin/assistant/sessions", listAssistantSessionsHandler(svc))
authed.POST("/admin/assistant/sessions", createAssistantSessionHandler(svc))
authed.GET("/admin/assistant/session/:id", getAssistantSessionHandler(svc))
authed.DELETE("/admin/assistant/session/:id", deleteAssistantSessionHandler(svc))
authed.POST("/admin/assistant/chat", assistantChatHandler(svc))
authed.POST("/admin/assistant/execute", assistantExecuteHandler(svc))
authed.POST("/admin/assistant/undo/:op_id", assistantUndoHandler(svc))
authed.GET("/admin/assistant/history", assistantHistoryHandler(svc))
}
@@ -0,0 +1,91 @@
package handler
import (
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func registerAuthedDownloadRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/downloads", requirePermission(svc, "can_manage_downloads"), listDownloadsHandler(svc))
authed.POST("/downloads", requirePermission(svc, "can_manage_downloads"), addDownloadHandler(svc))
authed.DELETE("/downloads/:hash", requirePermission(svc, "can_manage_downloads"), deleteDownloadHandler(svc))
authed.POST("/downloads/relocate", requirePermission(svc, "can_manage_downloads"), relocateDownloadHandler(svc))
authed.POST("/downloads/reload", requirePermission(svc, "can_manage_downloads"), reloadDownloadConfigHandler(svc))
}
func registerAuthedSubscriptionRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionsHandler(svc))
authed.GET("/subscriptions/history", requirePermission(svc, "can_manage_subscriptions"), listSubscriptionHistoryHandler(svc))
authed.POST("/subscriptions", requirePermission(svc, "can_manage_subscriptions"), createSubscriptionHandler(svc))
authed.DELETE("/subscriptions/:id", requirePermission(svc, "can_manage_subscriptions"), deleteSubscriptionHandler(svc))
authed.POST("/subscriptions/:id/restore", requirePermission(svc, "can_manage_subscriptions"), restoreSubscriptionHandler(svc))
authed.POST("/subscriptions/:id/run", requirePermission(svc, "can_manage_subscriptions"), runSubscriptionHandler(svc))
}
func registerAuthedStatsDiscoveryAndAIRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/stats", statsHandler(svc))
authed.GET("/tasks", middleware.AdminRequired(), tasksHandler(svc))
authed.GET("/discover/trending", requirePermission(svc, "can_view_discover"), trendingHandler(svc))
authed.GET("/discover/popular", requirePermission(svc, "can_view_discover"), popularHandler(svc))
authed.GET("/ai/status", requirePermission(svc, "can_use_ai"), aiStatusHandler(svc))
authed.POST("/ai/search", requirePermission(svc, "can_use_ai"), smartSearchHandler(svc))
authed.GET("/ai/recommend", requirePermission(svc, "can_use_ai"), aiRecommendHandler(svc))
}
func registerAuthedFileRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/files", middleware.AdminRequired(), browseFilesHandler(svc))
authed.POST("/files/folders", middleware.AdminRequired(), createFolderHandler(svc))
authed.PUT("/files/rename", middleware.AdminRequired(), renameFileHandler(svc))
authed.DELETE("/files", middleware.AdminRequired(), deleteFileHandler(svc))
authed.POST("/files/transfer", middleware.AdminRequired(), transferFileHandler(svc))
}
func registerAuthedDLNARoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/dlna/devices", dlnaListHandler(svc))
authed.POST("/dlna/cast", dlnaCastHandler(svc))
}
func registerAuthedSTRMRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.PUT("/media/:id/strm", middleware.AdminRequired(), setSTRMHandler(svc))
authed.DELETE("/media/:id/strm", middleware.AdminRequired(), clearSTRMHandler(svc))
authed.POST("/strm/import", middleware.AdminRequired(), importSTRMHandler(svc))
authed.POST("/strm/generate", middleware.AdminRequired(), generateSTRMHandler(svc))
}
func registerAuthedDuplicateRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/duplicates", middleware.AdminRequired(), listDuplicatesHandler(svc))
authed.POST("/duplicates/scan", middleware.AdminRequired(), detectDuplicatesHandler(svc))
authed.POST("/duplicates/unmark", middleware.AdminRequired(), unmarkDuplicatesHandler(svc))
}
func registerAuthedSiteRoutes(authed *gin.RouterGroup, svc *service.Container) {
siteHandler := NewSiteHandler(svc)
authed.GET("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.ListSites)
authed.GET("/sites/types", requirePermission(svc, "can_manage_sites"), siteHandler.GetSiteTypes)
authed.GET("/sites/auth-types", requirePermission(svc, "can_manage_sites"), siteHandler.GetAuthTypes)
authed.POST("/sites", requirePermission(svc, "can_manage_sites"), siteHandler.CreateSite)
authed.GET("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.GetSite)
authed.PUT("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.UpdateSite)
authed.DELETE("/sites/:id", requirePermission(svc, "can_manage_sites"), siteHandler.DeleteSite)
authed.POST("/sites/:id/test", requirePermission(svc, "can_manage_sites"), siteHandler.TestSite)
authed.GET("/sites/search", requirePermission(svc, "can_manage_sites"), siteSearchHandler(svc))
}
func registerAuthedRecycleAndRealtimeRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/recycle", middleware.AdminRequired(), listRecycleHandler(svc))
authed.POST("/recycle/restore", middleware.AdminRequired(), restoreMediaBatchHandler(svc))
authed.POST("/recycle/purge", middleware.AdminRequired(), purgeMediaBatchHandler(svc))
authed.GET("/ws", wsHandler(svc))
authed.GET("/events", sseHandler(svc))
}
func registerAuthedSchedulerRoutes(authed *gin.RouterGroup, svc *service.Container) {
authed.GET("/scheduler/tasks", schedulerListTasksHandler(svc))
authed.POST("/scheduler/tasks/:id/run", middleware.AdminRequired(), schedulerRunTaskHandler(svc))
authed.GET("/scheduler/status", schedulerGetStatusHandler(svc))
}
@@ -0,0 +1,46 @@
package handler
import (
"testing"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
func TestAuthenticatedRouteSurfacesAreRegistered(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
Register(router, &config.Config{
Secrets: config.SecretsConfig{JWTSecret: "test-secret"},
}, zap.NewNop(), &service.Container{Log: zap.NewNop()})
routes := map[string]bool{}
for _, route := range router.Routes() {
routes[route.Method+" "+route.Path] = true
}
for _, want := range []string{
"GET /api/me",
"GET /api/auth/permissions",
"GET /api/libraries",
"GET /api/media",
"GET /api/stream/:id",
"GET /api/storage",
"GET /api/downloads",
"GET /api/subscriptions",
"GET /api/sites/search",
"GET /api/watch-history",
"GET /api/discover/feed",
"GET /api/playback/:id/info",
"GET /api/download/tasks",
"GET /api/admin/assistant/history",
} {
if !routes[want] {
t.Fatalf("%s route is not registered", want)
}
}
}
+72 -4
View File
@@ -8,6 +8,7 @@ package handler
import (
"net/http"
"sort"
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
@@ -31,15 +32,20 @@ func listSeasonsHandler(svc *service.Container) gin.HandlerFunc {
return
}
}
visibility := mediaVisibilityForRequest(c, svc)
var rows []model.Media
err := svc.Repo.DB.Where(&model.Media{LibraryID: libID}).
Order("season_num asc, episode_num asc").
Find(&rows).Error
const pageSize = 2000
for page := 1; ; page++ {
pageRows, total, err := svc.Media.ListMediaVisible(c.Request.Context(), libID, page, pageSize, visibility)
if err != nil && err != gorm.ErrRecordNotFound {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
visibility := mediaVisibilityForRequest(c, svc)
rows = append(rows, pageRows...)
if int64(len(rows)) >= total || len(pageRows) < pageSize {
break
}
}
buckets := make(map[int][]model.Media)
for _, r := range rows {
if !visibility.Allows(&r) {
@@ -55,3 +61,65 @@ func listSeasonsHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusOK, gin.H{"seasons": out})
}
}
func listLibrarySeriesHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
libID := c.Param("id")
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, *lib, mediaVisibilityForRequest(c, svc)) {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
}
items, total, err := svc.Media.ListLibrarySeriesCards(c.Request.Context(), libID, mediaVisibilityForRequest(c, svc))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("page_size", "500"))
if page < 1 {
page = 1
}
if size <= 0 || size > 1000 {
size = 500
}
start := (page - 1) * size
if start > len(items) {
start = len(items)
}
end := start + size
if end > len(items) {
end = len(items)
}
c.JSON(http.StatusOK, gin.H{
"items": items[start:end],
"total": total,
"page": page,
"page_size": size,
})
}
}
func listLibrarySeriesEpisodesHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
libID := c.Param("id")
key := c.Query("key")
if key == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "key is required"})
return
}
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
if !service.LibraryVisibleForUser(c.Request.Context(), svc.Repo, *lib, mediaVisibilityForRequest(c, svc)) {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
}
items, err := svc.Media.ListLibrarySeriesEpisodes(c.Request.Context(), libID, key, mediaVisibilityForRequest(c, svc))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"items": items, "total": len(items)})
}
}
+21 -1
View File
@@ -1,4 +1,4 @@
// Package handler — Alist / S3 / WebDAV storage config endpoints.
// Package handler — external storage config endpoints.
package handler
import (
@@ -25,6 +25,10 @@ func listStorageConfigsHandler(svc *service.Container) gin.HandlerFunc {
// getStorageConfigHandler returns one config (with the decrypted body).
func getStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if !service.IsAdminStorageConfigurable(c.Param("type")) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
return
}
row, err := svc.StorageCfg.Get(c.Request.Context(), c.Param("type"))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -42,6 +46,10 @@ func getStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
// the type via URL and the body as a JSON object.
func saveStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if !service.IsAdminStorageConfigurable(c.Param("type")) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
return
}
var in service.StorageInput
if err := c.ShouldBindJSON(&in); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -63,6 +71,10 @@ func saveStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
// testStorageConfigHandler probes an unsaved config.
func testStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if !service.IsAdminStorageConfigurable(c.Param("type")) {
c.JSON(http.StatusBadRequest, gin.H{"ok": false, "error": "unsupported storage type"})
return
}
var in service.StorageInput
if err := c.ShouldBindJSON(&in); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -80,6 +92,10 @@ func testStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
func logoutStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
if !service.IsAdminStorageConfigurable(typ) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
return
}
row, err := svc.StorageCfg.Logout(c.Request.Context(), typ)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -94,6 +110,10 @@ func logoutStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
func storageUploadLocalHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if !service.IsAdminStorageConfigurable(c.Param("type")) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported storage type"})
return
}
var req service.CloudUploadInput
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
+28
View File
@@ -0,0 +1,28 @@
package handler
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
)
func TestStorageConfigHandlersRejectQuark(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.PUT("/admin/storage/:type", saveStorageConfigHandler(nil))
req := httptest.NewRequest(http.MethodPut, "/admin/storage/quark", strings.NewReader(`{"type":"quark","config":{"cookie":"x"}}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Fatalf("status = %d body=%s, want 400", w.Code, w.Body.String())
}
if !strings.Contains(w.Body.String(), "unsupported storage type") {
t.Fatalf("body = %s, want unsupported storage type", w.Body.String())
}
}
+103 -6
View File
@@ -4,6 +4,7 @@ package handler
import (
"context"
"errors"
"io"
"net/http"
"github.com/gin-gonic/gin"
@@ -79,16 +80,99 @@ func imageProxyHandler(svc *service.Container) gin.HandlerFunc {
}
}
func cloudArtworkProxyHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
typ := c.Param("type")
ref := c.Query("ref")
if !service.IsAdminCloudConfigurable(typ) {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
return
}
if ref == "" || !isCloudImageRef(ref) {
c.JSON(http.StatusBadRequest, gin.H{"error": "image ref required"})
return
}
if svc == nil || svc.ImageProxy == nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "image proxy unavailable"})
return
}
stableKey := typ + ":" + ref
if svc.ImageProxy.ServeCloudCached(c.Writer, c.Request, stableKey) {
return
}
if svc.StorageCfg == nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "cloud storage service unavailable"})
return
}
link, err := svc.StorageCfg.CloudResolve(c.Request.Context(), typ, ref, c.Request.UserAgent())
if err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
if err := svc.ImageProxy.ServeCloudResolved(c.Request.Context(), c.Writer, c.Request, stableKey, link); err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()})
return
}
}
}
type scrapeRequest struct {
EpisodeArtwork *bool `json:"episode_artwork"`
EpisodeImages *bool `json:"episode_images"`
RefreshMatched *bool `json:"refresh_matched"`
IncludeMatched *bool `json:"include_matched"`
}
func (r scrapeRequest) episodeArtworkOption() *bool {
if r.EpisodeImages != nil {
return r.EpisodeImages
}
return r.EpisodeArtwork
}
func (r scrapeRequest) includeMatchedOption() bool {
if r.IncludeMatched != nil {
return *r.IncludeMatched
}
if r.RefreshMatched != nil {
return *r.RefreshMatched
}
return false
}
func scrapeOptionsFromRequest(c *gin.Context, retryNoMatch bool) (service.ScrapeOptions, error) {
options := service.ScrapeOptions{RetryNoMatch: retryNoMatch}
if c.Request.Body == nil || c.Request.ContentLength == 0 {
return options, nil
}
var req scrapeRequest
if err := c.ShouldBindJSON(&req); err != nil {
if errors.Is(err, io.EOF) {
return options, nil
}
return options, err
}
options.EpisodeArtwork = req.episodeArtworkOption()
options.IncludeMatched = req.includeMatchedOption()
return options, nil
}
// scrapeOneHandler enriches a single media via the configured scraper chain.
func scrapeOneHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
options, err := scrapeOptionsFromRequest(c, true)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
return
}
options.IncludeMatched = true
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
if err != nil || m == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
return
}
task := startScrapeHTTPTask(svc, "手动刮削媒体", m.Title, m.Path)
if err := svc.Scraper.EnrichOne(c.Request.Context(), m); err != nil {
if err := svc.Scraper.EnrichOneWithOptions(c.Request.Context(), m, options); err != nil {
finishHTTPTask(task, err, "scrape", "手动刮削媒体失败", nil, nil)
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -103,10 +187,16 @@ func scrapeOneHandler(svc *service.Container) gin.HandlerFunc {
}
}
// scrapeLibraryHandler retries every pending/no_match media in a library.
// scrapeLibraryHandler manually refreshes every scrapeable row in a library.
func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
libID := c.Param("id")
options, err := scrapeOptionsFromRequest(c, true)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid scrape options"})
return
}
options.IncludeMatched = true
var task *service.TaskHandle
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
task = startScrapeHTTPTask(svc, "手动刮削媒体库", lib.Name, lib.Path)
@@ -115,9 +205,16 @@ func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
}
// Run in the background so HTTP returns instantly; the WS hub
// pushes per-item progress on the "scrape" topic.
go func(libID string, task *service.TaskHandle) {
matched, err := svc.Scraper.EnrichLibrary(context.Background(), libID, true)
metrics := map[string]int64{"matched": int64(matched)}
go func(libID string, task *service.TaskHandle, options service.ScrapeOptions) {
result, err := svc.Scraper.EnrichLibraryDetailedWithOptions(context.Background(), libID, options)
metrics := map[string]int64{
"matched": int64(result.Matched),
"processed": int64(result.Processed),
"candidates": int64(result.Candidates),
}
if result.Failed > 0 {
metrics["errors"] = int64(result.Failed)
}
stage := "completed"
message := "手动刮削媒体库结束"
if err != nil {
@@ -125,7 +222,7 @@ func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
message = "手动刮削媒体库失败"
}
finishHTTPTask(task, err, stage, message, metrics, nil)
}(libID, task)
}(libID, task, options)
c.JSON(http.StatusAccepted, gin.H{"status": "scraping"})
}
}
+10 -3
View File
@@ -99,7 +99,7 @@ func importSTRMHandler(svc *service.Container) gin.HandlerFunc {
}
type generateSTRMReq struct {
LibraryID string `json:"library_id" binding:"required"`
LibraryID string `json:"library_id"`
OutputDir string `json:"output_dir"`
BaseURL string `json:"base_url"`
Enabled bool `json:"enabled"`
@@ -122,7 +122,7 @@ func generateSTRMHandler(svc *service.Container) gin.HandlerFunc {
if baseURL == "" {
baseURL = strings.TrimRight(absoluteRequestURL(c, "/"), "/")
}
res, err := strmSvc.GenerateForLibrary(c.Request.Context(), service.GenerateSTRMOptions{
options := service.GenerateSTRMOptions{
LibraryID: req.LibraryID,
OutputDir: req.OutputDir,
BaseURL: baseURL,
@@ -130,7 +130,14 @@ func generateSTRMHandler(svc *service.Container) gin.HandlerFunc {
Overwrite: req.Overwrite,
IncludeLocal: true,
PlaybackToken: strmPlaybackTokenForRequest(c, svc),
})
}
var res *service.GenerateSTRMResult
var err error
if strings.TrimSpace(req.LibraryID) == "*" {
res, err = strmSvc.GenerateForAllLibraries(c.Request.Context(), options)
} else {
res, err = strmSvc.GenerateForLibrary(c.Request.Context(), options)
}
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
+4
View File
@@ -24,6 +24,8 @@ type subscriptionReq struct {
PosterURL string `json:"poster_url"`
BackdropURL string `json:"backdrop_url"`
Overview string `json:"overview"`
OriginalName string `json:"original_name"`
Year int `json:"year"`
Resolution string `json:"resolution"`
Quality string `json:"quality"`
Effects string `json:"effects"`
@@ -62,6 +64,8 @@ func createSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
PosterURL: req.PosterURL,
BackdropURL: req.BackdropURL,
Overview: req.Overview,
OriginalName: req.OriginalName,
Year: req.Year,
Resolution: req.Resolution,
Quality: req.Quality,
Effects: req.Effects,
+2 -3
View File
@@ -85,12 +85,11 @@ func schemaHandler(_ *service.Container) gin.HandlerFunc {
{"key": "cloud.boot_scan_enabled", "type": "toggle", "label": "启动后立即扫描网盘"},
{"key": "cloud.upload_auto_enabled", "type": "toggle", "label": "启用自动转存"},
{"key": "cloud.upload_provider", "type": "select", "label": "转存目标", "options": []gin.H{
{"value": "openlist", "label": "OpenList(推荐,可桥接 115/123/阿里/夸克)"},
{"value": "clouddrive2", "label": "CloudDrive2(推荐,可桥接 115/123/阿里/夸克)"},
{"value": "openlist", "label": "OpenList(推荐,可桥接 115/123/阿里等)"},
{"value": "clouddrive2", "label": "CloudDrive2(推荐,可桥接 115/123/阿里等)"},
{"value": "alist", "label": "Alist(可桥接多网盘)"},
{"value": "webdav", "label": "WebDAV"},
{"value": "cloud115", "label": "115 原生(待接分片上传)"},
{"value": "quark", "label": "夸克原生(待接分片上传)"},
}},
{"key": "cloud.upload_source_dir", "type": "text", "label": "本地源目录"},
{"key": "cloud.upload_dest_path", "type": "text", "label": "网盘目标目录"},
+11 -2
View File
@@ -8,16 +8,25 @@ package handler
import (
"net/http"
"time"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
const tasksLiveTorrentSnapshotMaxAge = 30 * time.Second
func tasksHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
transcodes := svc.Transcoder.Active()
_, torrents, _ := svc.Downloads.List(c.Request.Context())
var transcodes []service.ActiveJob
if svc.Transcoder != nil {
transcodes = svc.Transcoder.Active()
}
var torrents []service.QBitTorrent
if svc.Downloads != nil {
torrents = svc.Downloads.LiveTorrentSnapshot(tasksLiveTorrentSnapshotMaxAge)
}
background := service.TaskSnapshot{}
if svc.Tasks != nil {
background = svc.Tasks.Snapshot()
+51
View File
@@ -21,6 +21,11 @@ const (
CtxUserTier = "ctx_user_tier"
CtxTokenPurpose = "ctx_token_purpose"
CtxTokenMediaID = "ctx_token_media_id"
// AccessTokenCookieName carries the web access token for browser-managed
// resource requests such as <img>, which cannot attach Authorization.
AccessTokenCookieName = "msgo_access_token"
AccessTokenCookiePath = "/api"
)
// RequestLogger logs one structured line per request.
@@ -199,6 +204,7 @@ func AuthRequired(secret string) gin.HandlerFunc {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"code": 40304, "message": "token scope denied"})
return
}
syncAccessTokenCookie(c, raw, claims)
c.Set(CtxUserID, claims.UserID)
c.Set(CtxUserRole, claims.Role)
c.Set(CtxUserTier, claims.Tier)
@@ -208,6 +214,48 @@ func AuthRequired(secret string) gin.HandlerFunc {
}
}
func syncAccessTokenCookie(c *gin.Context, raw string, claims *Claims) {
if c == nil || claims == nil || strings.TrimSpace(raw) == "" || strings.TrimSpace(claims.Purpose) != "" {
return
}
if existing, err := c.Cookie(AccessTokenCookieName); err == nil && existing == raw {
return
}
maxAge := int(time.Hour.Seconds())
expires := time.Now().Add(time.Hour)
if claims.ExpiresAt != nil {
expires = claims.ExpiresAt.Time
ttl := time.Until(expires)
if ttl <= 0 {
return
}
maxAge = int(ttl.Seconds())
if maxAge < 1 {
maxAge = 1
}
}
http.SetCookie(c.Writer, &http.Cookie{
Name: AccessTokenCookieName,
Value: raw,
Path: AccessTokenCookiePath,
MaxAge: maxAge,
Expires: expires,
HttpOnly: true,
SameSite: http.SameSiteLaxMode,
Secure: requestIsHTTPS(c),
})
}
func requestIsHTTPS(c *gin.Context) bool {
if c == nil || c.Request == nil {
return false
}
if c.Request.TLS != nil {
return true
}
return strings.EqualFold(c.GetHeader("X-Forwarded-Proto"), "https")
}
func scopedTokenAllowedForRequest(c *gin.Context, claims *Claims) bool {
if claims == nil || strings.TrimSpace(claims.Purpose) == "" {
return true
@@ -320,5 +368,8 @@ func extractToken(c *gin.Context) string {
return value
}
}
if cookie, err := c.Cookie(AccessTokenCookieName); err == nil {
return strings.TrimSpace(cookie)
}
return ""
}
+160
View File
@@ -4,8 +4,10 @@ import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/golang-jwt/jwt/v5"
)
func TestCORSWildcardOriginAllowsProductionPreflight(t *testing.T) {
@@ -29,3 +31,161 @@ func TestCORSWildcardOriginAllowsProductionPreflight(t *testing.T) {
t.Fatalf("Access-Control-Allow-Origin = %q, want *", got)
}
}
func TestAuthRequiredAcceptsAccessTokenCookie(t *testing.T) {
gin.SetMode(gin.TestMode)
secret := "cookie-secret"
token := signedMiddlewareTestToken(t, secret, Claims{
UserID: "user-1",
Role: "admin",
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
},
})
router := gin.New()
router.Use(AuthRequired(secret))
router.GET("/api/img", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"user_id": c.GetString(CtxUserID)})
})
req := httptest.NewRequest(http.MethodGet, "/api/img?url=https%3A%2F%2Fexample.test%2Fposter.jpg", nil)
req.AddCookie(&http.Cookie{Name: AccessTokenCookieName, Value: token})
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
}
func TestAuthRequiredKeepsExplicitQueryTokenPriority(t *testing.T) {
gin.SetMode(gin.TestMode)
secret := "cookie-secret"
accountToken := signedMiddlewareTestToken(t, secret, Claims{
UserID: "user-1",
Role: "admin",
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
},
})
scopedToken := signedMiddlewareTestToken(t, secret, Claims{
UserID: "user-1",
Role: "admin",
Purpose: "external_play",
MediaID: "media-1",
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
},
})
router := gin.New()
router.Use(AuthRequired(secret))
router.GET("/api/me", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true})
})
req := httptest.NewRequest(http.MethodGet, "/api/me?token="+scopedToken, nil)
req.AddCookie(&http.Cookie{Name: AccessTokenCookieName, Value: accountToken})
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusForbidden {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
}
func TestAuthRequiredSyncsAccessTokenCookieFromBearer(t *testing.T) {
gin.SetMode(gin.TestMode)
secret := "cookie-secret"
token := signedMiddlewareTestToken(t, secret, Claims{
UserID: "user-1",
Role: "admin",
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
},
})
router := gin.New()
router.Use(AuthRequired(secret))
router.GET("/api/discover/feed", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true})
})
req := httptest.NewRequest(http.MethodGet, "/api/discover/feed", nil)
req.Header.Set("Authorization", "Bearer "+token)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
cookie := middlewareTestResponseCookie(t, w, AccessTokenCookieName)
if cookie.Value != token {
t.Fatal("synced cookie should contain the bearer token")
}
if cookie.Path != AccessTokenCookiePath {
t.Fatalf("cookie path = %q, want %q", cookie.Path, AccessTokenCookiePath)
}
if !cookie.HttpOnly || cookie.SameSite != http.SameSiteLaxMode {
t.Fatalf("cookie flags not suitable: httpOnly=%v sameSite=%v", cookie.HttpOnly, cookie.SameSite)
}
}
func TestAuthRequiredDoesNotSyncScopedPlaybackTokenCookie(t *testing.T) {
gin.SetMode(gin.TestMode)
secret := "cookie-secret"
token := signedMiddlewareTestToken(t, secret, Claims{
UserID: "user-1",
Role: "admin",
Purpose: "external_play",
MediaID: "media-1",
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
},
})
router := gin.New()
router.Use(AuthRequired(secret))
router.GET("/api/stream/media-1", func(c *gin.Context) {
c.Status(http.StatusNoContent)
})
req := httptest.NewRequest(http.MethodGet, "/api/stream/media-1?token="+token, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNoContent {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
if cookie := optionalMiddlewareTestResponseCookie(w, AccessTokenCookieName); cookie != nil {
t.Fatalf("scoped playback token should not be synced as web cookie: %#v", cookie)
}
}
func signedMiddlewareTestToken(t *testing.T, secret string, claims Claims) string {
t.Helper()
token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
if err != nil {
t.Fatalf("sign token: %v", err)
}
return token
}
func middlewareTestResponseCookie(t *testing.T, w *httptest.ResponseRecorder, name string) *http.Cookie {
t.Helper()
cookie := optionalMiddlewareTestResponseCookie(w, name)
if cookie == nil {
t.Fatalf("missing response cookie %q", name)
}
return cookie
}
func optionalMiddlewareTestResponseCookie(w *httptest.ResponseRecorder, name string) *http.Cookie {
for _, cookie := range w.Result().Cookies() {
if cookie.Name == name {
return cookie
}
}
return nil
}
+3
View File
@@ -97,6 +97,9 @@ type UserDevice struct {
LastPlayAt *time.Time `gorm:"index" json:"last_play_at,omitempty"`
Warnings int `gorm:"default:0" json:"warnings"` // 指纹不匹配累计告警次数
Kicked bool `gorm:"default:false" json:"kicked"` // 被一键踢下线(强制重新登录)
Realtime bool `gorm:"-" json:"realtime,omitempty"`
Online bool `gorm:"-" json:"online,omitempty"`
Playing bool `gorm:"-" json:"playing,omitempty"`
}
// BeforeCreate 生成 UUID。
+3
View File
@@ -55,6 +55,8 @@ type User struct {
LastShareWarnAt *time.Time `json:"last_share_warn_at,omitempty"`
IsDefaultAdmin bool `gorm:"-" json:"is_default_admin,omitempty"`
IsProtected bool `gorm:"-" json:"is_protected,omitempty"`
RealtimeOnline bool `gorm:"-" json:"realtime_online,omitempty"`
RealtimeDeviceCount int `gorm:"-" json:"realtime_device_count,omitempty"`
}
// Library 表示用户定义的媒体根目录。
@@ -73,6 +75,7 @@ type Media struct {
SeriesID string `gorm:"index;size:128" json:"series_id,omitempty"`
Title string `gorm:"size:255;not null" json:"title"`
OriginalName string `gorm:"size:255" json:"original_name,omitempty"`
EpisodeTitle string `gorm:"size:255" json:"episode_title,omitempty"`
Path string `gorm:"uniqueIndex;size:1024;not null" json:"path"`
SizeBytes int64 `json:"size_bytes"`
DurationSec int `json:"duration_sec"`
@@ -0,0 +1,24 @@
package repository
import (
"context"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// AccessLogRepository persists model.AccessLog records.
type AccessLogRepository struct{ db *gorm.DB }
// Create inserts one structured audit-trail entry.
func (r *AccessLogRepository) Create(ctx context.Context, l *model.AccessLog) error {
return r.db.WithContext(ctx).Create(l).Error
}
// Recent returns the latest access-log entries (admin Activity panel).
func (r *AccessLogRepository) Recent(ctx context.Context, limit int) ([]model.AccessLog, error) {
var rows []model.AccessLog
err := r.db.WithContext(ctx).Order("created_at desc").Limit(limit).Find(&rows).Error
return rows, err
}
@@ -0,0 +1,78 @@
package repository
import (
"context"
"errors"
"time"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// ApiConfigRepository persists model.ApiConfig records.
type ApiConfigRepository struct{ db *gorm.DB }
// Create inserts a new API config record.
func (r *ApiConfigRepository) Create(ctx context.Context, c *model.ApiConfig) error {
return r.db.WithContext(ctx).Create(c).Error
}
// FindByProvider returns the API config for a provider, or (nil, nil).
func (r *ApiConfigRepository) FindByProvider(ctx context.Context, provider string) (*model.ApiConfig, error) {
var c model.ApiConfig
err := r.db.WithContext(ctx).Where("provider = ?", provider).First(&c).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &c, nil
}
// List returns all API configs.
func (r *ApiConfigRepository) List(ctx context.Context) ([]model.ApiConfig, error) {
var rows []model.ApiConfig
err := r.db.WithContext(ctx).Order("provider asc").Find(&rows).Error
return rows, err
}
// Upsert creates or updates an API config.
func (r *ApiConfigRepository) Upsert(ctx context.Context, c *model.ApiConfig) error {
return r.db.WithContext(ctx).Where("provider = ?", c.Provider).
Assign(model.ApiConfig{
Base: model.Base{UpdatedAt: time.Now()},
APIKey: c.APIKey,
BaseURL: c.BaseURL,
Extra: c.Extra,
Enabled: c.Enabled,
}).FirstOrCreate(c).Error
}
// Update updates an API config.
func (r *ApiConfigRepository) Update(ctx context.Context, c *model.ApiConfig) error {
return r.db.WithContext(ctx).Model(&model.ApiConfig{}).
Where("provider = ?", c.Provider).Updates(map[string]any{
"api_key": c.APIKey,
"base_url": c.BaseURL,
"extra": c.Extra,
"enabled": c.Enabled,
"updated_at": time.Now(),
}).Error
}
// Delete removes an API config.
func (r *ApiConfigRepository) Delete(ctx context.Context, provider string) error {
return r.db.WithContext(ctx).Where("provider = ?", provider).Delete(&model.ApiConfig{}).Error
}
// UpdateTestResult 更新测试结果。
func (r *ApiConfigRepository) UpdateTestResult(ctx context.Context, provider, result string) error {
now := time.Now()
return r.db.WithContext(ctx).Model(&model.ApiConfig{}).
Where("provider = ?", provider).Updates(map[string]any{
"test_result": result,
"last_tested_at": &now,
}).Error
}
@@ -0,0 +1,24 @@
package repository
import (
"context"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// DownloadRepository persists model.DownloadTask records.
type DownloadRepository struct{ db *gorm.DB }
// Create inserts a new download task.
func (r *DownloadRepository) Create(ctx context.Context, t *model.DownloadTask) error {
return r.db.WithContext(ctx).Create(t).Error
}
// List returns all download tasks (admin view).
func (r *DownloadRepository) List(ctx context.Context) ([]model.DownloadTask, error) {
var rows []model.DownloadTask
err := r.db.WithContext(ctx).Order("created_at desc").Find(&rows).Error
return rows, err
}
@@ -0,0 +1,34 @@
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// FavoriteRepository persists model.Favorite records.
type FavoriteRepository struct{ db *gorm.DB }
// Toggle flips the favourite flag for (user, media). Returns the new state.
func (r *FavoriteRepository) Toggle(ctx context.Context, userID, mediaID string) (bool, error) {
var f model.Favorite
err := r.db.WithContext(ctx).Where("user_id = ? AND media_id = ?", userID, mediaID).First(&f).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
fav := model.Favorite{UserID: userID, MediaID: mediaID}
return true, r.db.WithContext(ctx).Create(&fav).Error
}
if err != nil {
return false, err
}
return false, r.db.WithContext(ctx).Delete(&f).Error
}
// ListByUser returns all favourite media IDs for a user.
func (r *FavoriteRepository) ListByUser(ctx context.Context, userID string) ([]model.Favorite, error) {
var rows []model.Favorite
err := r.db.WithContext(ctx).Where("user_id = ?", userID).Find(&rows).Error
return rows, err
}
+41
View File
@@ -0,0 +1,41 @@
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// HistoryRepository persists model.PlaybackHistory entries. The application
// upserts on (UserID, MediaID) so resume always reads the latest position.
type HistoryRepository struct{ db *gorm.DB }
// Upsert atomically inserts/updates the resume position.
func (r *HistoryRepository) Upsert(ctx context.Context, h *model.PlaybackHistory) error {
var existing model.PlaybackHistory
err := r.db.WithContext(ctx).
Where("user_id = ? AND media_id = ?", h.UserID, h.MediaID).
First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return r.db.WithContext(ctx).Create(h).Error
}
if err != nil {
return err
}
existing.PositionMs = h.PositionMs
existing.DurationMs = h.DurationMs
existing.WatchedAt = h.WatchedAt
existing.Completed = h.Completed
return r.db.WithContext(ctx).Save(&existing).Error
}
// ListByUser returns the most recent history rows for the user.
func (r *HistoryRepository) ListByUser(ctx context.Context, userID string, limit int) ([]model.PlaybackHistory, error) {
var rows []model.PlaybackHistory
err := r.db.WithContext(ctx).Where("user_id = ?", userID).
Order("watched_at desc").Limit(limit).Find(&rows).Error
return rows, err
}
+44
View File
@@ -0,0 +1,44 @@
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// LibraryRepository persists model.Library records.
type LibraryRepository struct{ db *gorm.DB }
// Create persists a new library row.
func (r *LibraryRepository) Create(ctx context.Context, l *model.Library) error {
return r.db.WithContext(ctx).Create(l).Error
}
// List returns all enabled+disabled libraries.
func (r *LibraryRepository) List(ctx context.Context) ([]model.Library, error) {
var ls []model.Library
err := r.db.WithContext(ctx).Order("created_at asc").Find(&ls).Error
return ls, err
}
// FindByID returns the library, or (nil, nil) when missing.
func (r *LibraryRepository) FindByID(ctx context.Context, id string) (*model.Library, error) {
var l model.Library
err := r.db.WithContext(ctx).Where("id = ?", id).First(&l).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &l, nil
}
// Delete removes a library and (soft) cascades to its media via repository
// callers; we do not run CASCADE here to keep this method narrow.
func (r *LibraryRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.Library{}, "id = ?", id).Error
}
+313
View File
@@ -0,0 +1,313 @@
package repository
import (
"context"
"errors"
"strings"
"sync"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// MediaRepository persists model.Media records.
type MediaRepository struct {
db *gorm.DB
searchIndexOnce sync.Once
searchIndexAvailable bool
searchBackend MediaSearchBackend
}
type MediaSearchBackend interface {
SearchMediaIDs(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]string, int64, error)
}
type MediaSearchSyncBackend interface {
MediaSearchBackend
EnsureIndex(ctx context.Context) error
IndexMedia(ctx context.Context, rows []model.Media) error
}
func (r *MediaRepository) SetSearchBackend(backend MediaSearchBackend) {
if r != nil {
r.searchBackend = backend
}
}
// MediaQueryFilter is applied to user-facing media queries so NSFW items and
// profile-restricted libraries are filtered in SQL instead of only in React.
type MediaQueryFilter struct {
IncludeNSFW bool
AllowedLibraryIDs []string
HiddenLibraryIDs []string
}
func applyMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB {
if !filter.IncludeNSFW {
q = q.Where("nsfw = ?", false)
}
if len(filter.HiddenLibraryIDs) > 0 {
q = q.Where("library_id NOT IN ?", filter.HiddenLibraryIDs)
}
if len(filter.AllowedLibraryIDs) > 0 {
q = q.Where("library_id IN ?", filter.AllowedLibraryIDs)
}
return q
}
// Upsert inserts or updates a media row keyed by Path (unique index).
//
// 重要:当一条行已经存在时,scanner 重扫只应该刷新文件级元数据
// (时长、宽高、编码、容器、大小),不能把刮削器维护的字段(标题改写、
// 海报、TMDb/Bangumi ID、scrape_status 等)覆盖回零值。
//
// 之前用 Assign(*m).FirstOrCreate(m) 会把整张零值结构体写回,导致:
// 1. scrape_status 从 'matched' / 'no_match' 被清空成 ”;
// 2. 新建行使 GORM `default:pending` 也得不到应用(因为 zero value 被
// 显式写入)。这两个问题都让 EnrichLibrary(WHERE scrape_status='pending')
// 永远捞不到数据。
func (r *MediaRepository) Upsert(ctx context.Context, m *model.Media) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.upsert(ctx, m)
})
}
func (r *MediaRepository) upsert(ctx context.Context, m *model.Media) error {
var existing model.Media
err := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
// 新行:保证 scrape_status 走 GORM default:pending(即留空让数据库填)。
if m.ScrapeStatus == "" {
m.ScrapeStatus = "pending"
}
if createErr := r.db.WithContext(ctx).Create(m).Error; createErr == nil {
r.indexMediaBestEffort(ctx, *m)
return nil
} else if retryErr := r.db.WithContext(ctx).Unscoped().Where("path = ?", m.Path).First(&existing).Error; retryErr != nil {
return createErr
}
}
if err != nil {
return err
}
// 已存在:仅刷新文件层面的字段。
updates := map[string]any{}
setIfChanged(updates, "size_bytes", existing.SizeBytes, m.SizeBytes)
setIfChanged(updates, "duration_sec", existing.DurationSec, m.DurationSec)
setIfChanged(updates, "width", existing.Width, m.Width)
setIfChanged(updates, "height", existing.Height, m.Height)
setIfChanged(updates, "video_codec", existing.VideoCodec, m.VideoCodec)
setIfChanged(updates, "audio_codec", existing.AudioCodec, m.AudioCodec)
setIfChanged(updates, "container", existing.Container, m.Container)
if existing.DeletedAt.Valid {
updates["deleted_at"] = nil
}
// 回填硬链接身份标识,便于后续扫描去重(避免重复识别/多倍占用)。
if m.FileID != "" && m.FileID != existing.FileID {
updates["file_id"] = m.FileID
}
if m.Title != "" {
// scanner 给出的标题只是从路径推导,刮削后 title 已被替换为
// 真实剧名。仅在 existing 还停留在 'pending'/'' 时回填扫描标题,
// 避免覆盖刮削结果。
if m.ScrapeStatus == "matched" || existing.ScrapeStatus == "pending" || existing.ScrapeStatus == "" || existing.ScrapeStatus == "no_match" {
setIfChanged(updates, "title", existing.Title, m.Title)
if m.Year > 0 {
setIfChanged(updates, "year", existing.Year, m.Year)
}
}
}
status := strings.TrimSpace(existing.ScrapeStatus)
canRefreshExternalIDs := status == "pending" || status == "" || status == "no_match" ||
m.ScrapeStatus == "matched" || strings.HasPrefix(strings.ToLower(strings.TrimSpace(m.Path)), "cloud://")
if canRefreshExternalIDs {
changedExternalID := false
if m.TMDbID > 0 && existing.TMDbID != m.TMDbID {
updates["tm_db_id"] = m.TMDbID
changedExternalID = true
}
if m.BangumiID > 0 && existing.BangumiID != m.BangumiID {
updates["bangumi_id"] = m.BangumiID
changedExternalID = true
}
if m.DoubanID != "" && strings.TrimSpace(existing.DoubanID) != strings.TrimSpace(m.DoubanID) {
updates["douban_id"] = m.DoubanID
changedExternalID = true
}
if m.TheTVDBID != "" && strings.TrimSpace(existing.TheTVDBID) != strings.TrimSpace(m.TheTVDBID) {
updates["thetvdb_id"] = m.TheTVDBID
changedExternalID = true
}
if m.Year > 0 && existing.Year <= 0 {
updates["year"] = m.Year
}
if changedExternalID && (status == "no_match" || status == "matched") && m.ScrapeStatus != "matched" {
updates["scrape_status"] = "pending"
}
}
if m.ScrapeStatus == "matched" {
setIfChanged(updates, "scrape_status", existing.ScrapeStatus, m.ScrapeStatus)
if m.OriginalName != "" {
setIfChanged(updates, "original_name", existing.OriginalName, m.OriginalName)
}
if m.EpisodeTitle != "" {
setIfChanged(updates, "episode_title", existing.EpisodeTitle, m.EpisodeTitle)
}
if m.PosterURL != "" {
setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL)
}
if m.BackdropURL != "" {
setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL)
}
if m.Overview != "" {
setIfChanged(updates, "overview", existing.Overview, m.Overview)
}
if m.Rating > 0 {
setIfChanged(updates, "rating", existing.Rating, m.Rating)
}
if m.Year > 0 {
setIfChanged(updates, "year", existing.Year, m.Year)
}
if m.TMDbID > 0 {
setIfChanged(updates, "tm_db_id", existing.TMDbID, m.TMDbID)
}
if m.BangumiID > 0 {
setIfChanged(updates, "bangumi_id", existing.BangumiID, m.BangumiID)
}
if m.DoubanID != "" {
setIfChanged(updates, "douban_id", existing.DoubanID, m.DoubanID)
}
if m.TheTVDBID != "" {
setIfChanged(updates, "thetvdb_id", existing.TheTVDBID, m.TheTVDBID)
}
if m.Languages != "" {
setIfChanged(updates, "languages", existing.Languages, m.Languages)
}
if m.Countries != "" {
setIfChanged(updates, "countries", existing.Countries, m.Countries)
}
if m.Genres != "" {
setIfChanged(updates, "genres", existing.Genres, m.Genres)
}
if m.NSFW && !existing.NSFW {
updates["nsfw"] = true
}
}
if m.PosterURL != "" {
setIfChanged(updates, "poster_url", existing.PosterURL, m.PosterURL)
}
if m.BackdropURL != "" {
setIfChanged(updates, "backdrop_url", existing.BackdropURL, m.BackdropURL)
}
// 云盘媒体:同一 cloud:// 文件可能先被父目录库扫描入库,之后用户按二级
// 分类重新挂载/扫描到更精确的分类库。此时让 library_id 迁移到当前扫描库,
// 否则媒体被钉死在旧库、新分类库里看不到(表现为"媒体部分消失")。
// 本地媒体物理位置固定:仅在原 library_id 为空时回填,不迁移。
if isCloudMediaPath := strings.HasPrefix(strings.ToLower(strings.TrimSpace(m.Path)), "cloud://"); m.LibraryID != "" && m.LibraryID != existing.LibraryID {
if isCloudMediaPath || existing.LibraryID == "" {
updates["library_id"] = m.LibraryID
}
}
if (m.SeasonNum > 0 || m.EpisodeNum > 0) && existing.SeasonNum != m.SeasonNum {
updates["season_num"] = m.SeasonNum
}
if m.EpisodeNum > 0 && existing.EpisodeNum != m.EpisodeNum {
updates["episode_num"] = m.EpisodeNum
}
if m.STRMURL != "" {
setIfChanged(updates, "strm_url", existing.STRMURL, m.STRMURL)
}
if len(updates) == 0 {
*m = existing
return nil
}
if err := r.db.WithContext(ctx).Unscoped().Model(&model.Media{}).
Where("id = ?", existing.ID).Updates(updates).Error; err != nil {
return err
}
// 回写 ID / 不可变字段,让 caller 拿到完整的现有行。
*m = existing
if fresh, err := r.FindByID(ctx, existing.ID); err == nil && fresh != nil {
r.indexMediaBestEffort(ctx, *fresh)
}
return nil
}
func (r *MediaRepository) indexMediaBestEffort(ctx context.Context, media model.Media) {
backend, ok := r.searchBackend.(MediaSearchSyncBackend)
if !ok {
return
}
_ = backend.IndexMedia(ctx, []model.Media{media})
}
func setIfChanged[T comparable](updates map[string]any, key string, current, next T) {
if current != next {
updates[key] = next
}
}
// FindByID returns the media row or (nil, nil).
func (r *MediaRepository) FindByID(ctx context.Context, id string) (*model.Media, error) {
var m model.Media
err := r.db.WithContext(ctx).Where("id = ?", id).First(&m).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &m, nil
}
// ListByLibrary returns paginated media items for a library.
func (r *MediaRepository) ListByLibrary(ctx context.Context, libraryID string, offset, limit int) ([]model.Media, int64, error) {
return r.ListByLibraryFiltered(ctx, libraryID, offset, limit, MediaQueryFilter{IncludeNSFW: true})
}
func (r *MediaRepository) ListByLibraryFiltered(ctx context.Context, libraryID string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
return r.ListByLibrariesFiltered(ctx, []string{libraryID}, offset, limit, filter)
}
func (r *MediaRepository) ListByLibrariesFiltered(ctx context.Context, libraryIDs []string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
var items []model.Media
var total int64
if len(libraryIDs) == 0 {
return items, 0, nil
}
q := r.db.WithContext(ctx).Model(&model.Media{})
if len(libraryIDs) == 1 {
q = q.Where("library_id = ?", libraryIDs[0])
} else {
q = q.Where("library_id IN ?", libraryIDs)
}
q = applyMediaQueryFilter(q, filter)
if err := q.Count(&total).Error; err != nil {
return nil, 0, err
}
// 多级排序消除"随机"观感:
// 1. year desc — 上映年份新→旧(用户期望的上映时间维度)
// 2. updated_at desc — 同年按最近更新(刮削/补集会刷新)
// 3. created_at desc — 再按入库时间
// 4. id desc — 稳定 tie-breaker:云盘批量扫描同批 created_at 相同时,
// 没有它 DB 返回顺序不确定,正是"随机排序"的根因。
err := q.Order("year DESC, updated_at DESC, created_at DESC, id DESC").
Offset(offset).Limit(limit).Find(&items).Error
return items, total, err
}
// DeleteByLibrary purges all media tied to a library.
func (r *MediaRepository) DeleteByLibrary(ctx context.Context, libraryID string) error {
// FTS 行由 media 表上的触发器同步清理(软删/硬删都覆盖)。
return r.db.WithContext(ctx).Where("library_id = ?", libraryID).Delete(&model.Media{}).Error
}
// PurgeByLibrary permanently removes media tied to a library. Used for virtual
// cloud mounts where "remove mount" must not populate the recycle bin.
func (r *MediaRepository) PurgeByLibrary(ctx context.Context, libraryID string) error {
return r.db.WithContext(ctx).Unscoped().Where("library_id = ?", libraryID).Delete(&model.Media{}).Error
}
@@ -0,0 +1,268 @@
package repository
import (
"context"
"strings"
"unicode"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// Search runs a LIKE search against the title field. Empty query returns the
// most recently added items.
func (r *MediaRepository) Search(ctx context.Context, query string, limit int) ([]model.Media, error) {
return r.SearchFiltered(ctx, query, limit, MediaQueryFilter{IncludeNSFW: true})
}
func (r *MediaRepository) SearchFiltered(ctx context.Context, query string, limit int, filter MediaQueryFilter) ([]model.Media, error) {
items, _, err := r.SearchFilteredPage(ctx, query, 0, limit, filter)
return items, err
}
func (r *MediaRepository) SearchFilteredPage(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
query = strings.TrimSpace(query)
if limit <= 0 {
limit = 50
}
if query != "" && r.searchBackend != nil {
if items, total, ok := r.searchFilteredBackend(ctx, query, offset, limit, filter); ok {
return items, total, nil
}
}
if query != "" {
if items, total, ok := r.searchFilteredFTS(ctx, query, offset, limit, filter); ok {
if total > 0 {
return items, total, nil
}
}
}
return r.searchFilteredLIKE(ctx, query, offset, limit, filter)
}
func (r *MediaRepository) searchFilteredBackend(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, bool) {
ids, total, err := r.searchBackend.SearchMediaIDs(ctx, query, offset, limit, filter)
if err != nil {
return nil, 0, false
}
if len(ids) == 0 {
return []model.Media{}, total, true
}
var rows []model.Media
q := r.db.WithContext(ctx).Model(&model.Media{}).Where("id IN ?", ids)
q = applyMediaQueryFilter(q, filter)
if err := q.Find(&rows).Error; err != nil {
return nil, 0, false
}
byID := make(map[string]model.Media, len(rows))
for _, row := range rows {
byID[row.ID] = row
}
items := make([]model.Media, 0, len(ids))
for _, id := range ids {
if row, ok := byID[id]; ok {
items = append(items, row)
}
}
if len(items) == 0 && total > 0 {
return nil, 0, false
}
return items, total, true
}
func (r *MediaRepository) searchFilteredFTS(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, bool) {
if !r.searchIndexEnabled(ctx) {
return nil, 0, false
}
ftsQuery := mediaFTSQuery(query)
if ftsQuery == "" {
return nil, 0, false
}
var total int64
var items []model.Media
q := r.db.WithContext(ctx).
Table("media").
Joins("JOIN media_search_fts ON media_search_fts.rowid = media.rowid").
Where("media.deleted_at IS NULL").
Where("media_search_fts MATCH ?", ftsQuery)
q = applyQualifiedMediaQueryFilter(q, filter)
if err := q.Count(&total).Error; err != nil {
return nil, 0, false
}
if total == 0 {
return items, 0, true
}
err := q.Select("media.*").Order("bm25(media_search_fts), media.created_at DESC").Offset(offset).Limit(limit).Find(&items).Error
if err != nil {
return nil, 0, false
}
return items, total, true
}
func (r *MediaRepository) searchFilteredLIKE(ctx context.Context, query string, offset, limit int, filter MediaQueryFilter) ([]model.Media, int64, error) {
var items []model.Media
var total int64
q := r.db.WithContext(ctx).Model(&model.Media{})
q = applyMediaQueryFilter(q, filter)
terms := mediaSearchTerms(query)
for _, term := range terms {
like := "%" + escapeLike(term) + "%"
q = q.Where(
"(title LIKE ? ESCAPE '\\' OR original_name LIKE ? ESCAPE '\\' OR path LIKE ? ESCAPE '\\' OR genres LIKE ? ESCAPE '\\')",
like, like, like, like,
)
}
if err := q.Count(&total).Error; err != nil {
return nil, 0, err
}
if query != "" {
prefix := escapeLike(query) + "%"
exact := query
q = q.Order(gorm.Expr(
"CASE WHEN title = ? THEN 0 WHEN original_name = ? THEN 1 WHEN title LIKE ? ESCAPE '\\' THEN 2 WHEN original_name LIKE ? ESCAPE '\\' THEN 3 ELSE 4 END, created_at desc",
exact, exact, prefix, prefix,
))
} else {
q = q.Order("created_at desc")
}
err := q.Offset(offset).Limit(limit).Find(&items).Error
return items, total, err
}
func applyQualifiedMediaQueryFilter(q *gorm.DB, filter MediaQueryFilter) *gorm.DB {
if !filter.IncludeNSFW {
q = q.Where("media.nsfw = ?", false)
}
if len(filter.HiddenLibraryIDs) > 0 {
q = q.Where("media.library_id NOT IN ?", filter.HiddenLibraryIDs)
}
if len(filter.AllowedLibraryIDs) > 0 {
q = q.Where("media.library_id IN ?", filter.AllowedLibraryIDs)
}
return q
}
func mediaFTSQuery(query string) string {
terms := mediaSearchTerms(query)
if len(terms) == 0 {
return ""
}
quoted := make([]string, 0, len(terms))
for _, term := range terms {
term = strings.ReplaceAll(term, `"`, `""`)
if term != "" {
quoted = append(quoted, `"`+term+`"`)
}
}
return strings.Join(quoted, " AND ")
}
func mediaSearchTerms(query string) []string {
query = strings.TrimSpace(query)
if query == "" {
return nil
}
fields := strings.FieldsFunc(query, func(r rune) bool {
return unicode.IsSpace(r) || unicode.IsPunct(r) || unicode.IsSymbol(r)
})
out := make([]string, 0, len(fields))
seen := map[string]struct{}{}
for _, field := range fields {
field = strings.TrimSpace(field)
if field == "" {
continue
}
lower := strings.ToLower(field)
if _, ok := seen[lower]; ok {
continue
}
seen[lower] = struct{}{}
out = append(out, field)
}
return out
}
func escapeLike(value string) string {
value = strings.ReplaceAll(value, `\`, `\\`)
value = strings.ReplaceAll(value, `%`, `\%`)
value = strings.ReplaceAll(value, `_`, `\_`)
return value
}
func (r *MediaRepository) BackfillSearchIndex(ctx context.Context, batchLimit int) (int64, error) {
if backend, ok := r.searchBackend.(MediaSearchSyncBackend); ok {
return r.backfillExternalSearchIndex(ctx, backend, batchLimit)
}
if batchLimit <= 0 {
batchLimit = 1000
}
if !r.searchIndexEnabled(ctx) {
return 0, nil
}
// 关键性能点:FTS5 普通列(含 UNINDEXED)不支持索引查找,按
// media_id 做 NOT EXISTS 是对 FTS 表的整表扫描,再叠加 ORDER BY
// 后每个批次都要对全部 media 行探测一遍——大库一次启动回填等于
// 上百亿次行访问,曾把 CPU 钉满数小时。v2 布局下 FTS 行 rowid 与
// media.rowid 对齐,NOT EXISTS 走 rowid 点查,且无需排序。
res := r.db.WithContext(ctx).Exec(`
INSERT INTO media_search_fts(rowid, media_id, title, original_name, path, genres)
SELECT m.rowid, m.id, COALESCE(m.title, ''), COALESCE(m.original_name, ''), COALESCE(m.path, ''), COALESCE(m.genres, '')
FROM media AS m
WHERE m.deleted_at IS NULL
AND NOT EXISTS (
SELECT 1 FROM media_search_fts AS f WHERE f.rowid = m.rowid
)
LIMIT ?
`, batchLimit)
return res.RowsAffected, res.Error
}
func (r *MediaRepository) backfillExternalSearchIndex(ctx context.Context, backend MediaSearchSyncBackend, batchLimit int) (int64, error) {
if batchLimit <= 0 {
batchLimit = 1000
}
if err := backend.EnsureIndex(ctx); err != nil {
return 0, err
}
var lastID string
for {
var rows []model.Media
q := r.db.WithContext(ctx).
Model(&model.Media{}).
Where("deleted_at IS NULL")
if lastID != "" {
q = q.Where("id > ?", lastID)
}
if err := q.Order("id ASC").Limit(batchLimit).Find(&rows).Error; err != nil {
return 0, err
}
if len(rows) == 0 {
return 0, nil
}
if err := backend.IndexMedia(ctx, rows); err != nil {
return 0, err
}
lastID = rows[len(rows)-1].ID
if len(rows) < batchLimit {
return 0, nil
}
}
}
func (r *MediaRepository) searchIndexEnabled(ctx context.Context) bool {
if r == nil || r.db == nil {
return false
}
if r.db.Dialector == nil || r.db.Dialector.Name() != "sqlite" {
return false
}
r.searchIndexOnce.Do(func() {
var count int64
err := r.db.WithContext(ctx).
Raw(`SELECT COUNT(*) FROM sqlite_master WHERE name = 'media_search_fts'`).
Scan(&count).Error
r.searchIndexAvailable = err == nil && count > 0
})
return r.searchIndexAvailable
}
@@ -0,0 +1,59 @@
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// PermissionRepository persists model.UserPermission records.
type PermissionRepository struct{ db *gorm.DB }
// Create inserts a new permission record.
func (r *PermissionRepository) Create(ctx context.Context, p *model.UserPermission) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Create(p).Error
})
}
// FindByUserID returns the permission record for a user, or (nil, nil) when absent.
func (r *PermissionRepository) FindByUserID(ctx context.Context, userID string) (*model.UserPermission, error) {
var p model.UserPermission
err := withSQLiteBusyRetry(ctx, func() error {
p = model.UserPermission{}
return r.db.WithContext(ctx).Where("user_id = ?", userID).First(&p).Error
})
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &p, nil
}
// Update updates permission fields for a user.
func (r *PermissionRepository) Update(ctx context.Context, userID string, updates map[string]bool) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.UserPermission{}).
Where("user_id = ?", userID).Updates(updates).Error
})
}
// Upsert creates or updates a permission record.
func (r *PermissionRepository) Upsert(ctx context.Context, p *model.UserPermission) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Where("user_id = ?", p.UserID).
Assign(*p).FirstOrCreate(p).Error
})
}
// Delete removes a permission record.
func (r *PermissionRepository) Delete(ctx context.Context, userID string) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Where("user_id = ?", userID).Delete(&model.UserPermission{}).Error
})
}
@@ -0,0 +1,25 @@
package repository
import (
"context"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// PlaylistRepository persists model.Playlist + PlaylistItem.
type PlaylistRepository struct{ db *gorm.DB }
// Create inserts a new playlist.
func (r *PlaylistRepository) Create(ctx context.Context, p *model.Playlist) error {
return r.db.WithContext(ctx).Create(p).Error
}
// ListByUser returns playlists owned by a user.
func (r *PlaylistRepository) ListByUser(ctx context.Context, userID string) ([]model.Playlist, error) {
var rows []model.Playlist
err := r.db.WithContext(ctx).Where("user_id = ?", userID).
Order("created_at desc").Find(&rows).Error
return rows, err
}
@@ -0,0 +1,94 @@
package repository
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"time"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// RefreshTokenRepository persists model.RefreshToken records.
type RefreshTokenRepository struct{ db *gorm.DB }
// Create inserts a new refresh token record.
func (r *RefreshTokenRepository) Create(ctx context.Context, t *model.RefreshToken) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Create(t).Error
})
}
// FindByHash returns the refresh token matching the hash, or (nil, nil).
func (r *RefreshTokenRepository) FindByHash(ctx context.Context, hash string) (*model.RefreshToken, error) {
var t model.RefreshToken
err := withSQLiteBusyRetry(ctx, func() error {
t = model.RefreshToken{}
return r.db.WithContext(ctx).Where("token_hash = ?", hash).First(&t).Error
})
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &t, nil
}
// RevokeByUserID revokes all refresh tokens for a user.
func (r *RefreshTokenRepository) RevokeByUserID(ctx context.Context, userID string) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
Where("user_id = ?", userID).Update("revoked", true).Error
})
}
// RevokeOldestActiveByUserID keeps at most limit active refresh tokens for a
// user by revoking the oldest non-expired, non-revoked tokens.
func (r *RefreshTokenRepository) RevokeOldestActiveByUserID(ctx context.Context, userID string, limit int) error {
if limit < 1 {
limit = 1
}
return withSQLiteBusyRetry(ctx, func() error {
var tokens []model.RefreshToken
if err := r.db.WithContext(ctx).
Where("user_id = ? AND revoked = ? AND expires_at > ?", userID, false, time.Now()).
Order("created_at desc, id desc").
Find(&tokens).Error; err != nil {
return err
}
if len(tokens) <= limit {
return nil
}
ids := make([]string, 0, len(tokens)-limit)
for _, token := range tokens[limit:] {
ids = append(ids, token.ID)
}
return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
Where("id IN ?", ids).Update("revoked", true).Error
})
}
// DeleteExpired removes all expired refresh tokens.
func (r *RefreshTokenRepository) DeleteExpired(ctx context.Context) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Where("expires_at < ?", time.Now()).Delete(&model.RefreshToken{}).Error
})
}
// Revoke revokes a specific refresh token.
func (r *RefreshTokenRepository) Revoke(ctx context.Context, hash string) error {
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.RefreshToken{}).
Where("token_hash = ?", hash).Update("revoked", true).Error
})
}
// HashToken returns the SHA256 hash of a token.
func HashToken(token string) string {
h := sha256.Sum256([]byte(token))
return hex.EncodeToString(h[:])
}
File diff suppressed because it is too large Load Diff
+33
View File
@@ -0,0 +1,33 @@
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// SeriesRepository persists model.Series records.
type SeriesRepository struct{ db *gorm.DB }
// FindByID returns the series or (nil, nil).
func (r *SeriesRepository) FindByID(ctx context.Context, id string) (*model.Series, error) {
var s model.Series
err := r.db.WithContext(ctx).Where("id = ?", id).First(&s).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &s, nil
}
// List returns all series (ordered by title).
func (r *SeriesRepository) List(ctx context.Context) ([]model.Series, error) {
var s []model.Series
err := r.db.WithContext(ctx).Order("title asc").Find(&s).Error
return s, err
}
+42
View File
@@ -0,0 +1,42 @@
package repository
import (
"context"
"time"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// SettingRepository persists key/value preferences.
type SettingRepository struct{ db *gorm.DB }
// Get returns the value or empty string when absent.
func (r *SettingRepository) Get(ctx context.Context, key string) (string, error) {
var value string
err := r.db.WithContext(ctx).
Model(&model.Setting{}).
Select("value").
Where("key = ?", key).
Scan(&value).Error
return value, err
}
// Set upserts a setting value.
func (r *SettingRepository) Set(ctx context.Context, key, value string) error {
s := model.Setting{Key: key, Value: value, UpdatedAt: time.Now()}
return r.db.WithContext(ctx).Save(&s).Error
}
// Delete removes a setting key.
func (r *SettingRepository) Delete(ctx context.Context, key string) error {
return r.db.WithContext(ctx).Where("key = ?", key).Delete(&model.Setting{}).Error
}
// All returns every key/value pair (used by the admin UI).
func (r *SettingRepository) All(ctx context.Context) ([]model.Setting, error) {
var rows []model.Setting
err := r.db.WithContext(ctx).Find(&rows).Error
return rows, err
}
@@ -0,0 +1,45 @@
package repository
import (
"context"
"time"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// SubscriptionRepository persists model.Subscription records.
type SubscriptionRepository struct{ db *gorm.DB }
// Create inserts a new subscription rule.
func (r *SubscriptionRepository) Create(ctx context.Context, s *model.Subscription) error {
return r.db.WithContext(ctx).Select("*").Omit("DeletedAt").Create(s).Error
}
// List returns active subscription rules. Archived rows live in history and are
// intentionally excluded from scheduler polling and the active management list.
func (r *SubscriptionRepository) List(ctx context.Context) ([]model.Subscription, error) {
var rows []model.Subscription
err := r.db.WithContext(ctx).Where("archived_at IS NULL").Order("created_at desc").Find(&rows).Error
return rows, err
}
// History returns archived subscription rules.
func (r *SubscriptionRepository) History(ctx context.Context) ([]model.Subscription, error) {
var rows []model.Subscription
err := r.db.WithContext(ctx).Where("archived_at IS NOT NULL").Order("archived_at desc, updated_at desc").Find(&rows).Error
return rows, err
}
// Archive moves a completed subscription out of the active list without
// deleting its rule details, so users can audit completed subscriptions later.
func (r *SubscriptionRepository) Archive(ctx context.Context, id, reason string, archivedAt time.Time) error {
return r.db.WithContext(ctx).Model(&model.Subscription{}).
Where("id = ? AND archived_at IS NULL", id).
Updates(map[string]any{
"enabled": false,
"archived_at": &archivedAt,
"archive_reason": reason,
}).Error
}
+161
View File
@@ -0,0 +1,161 @@
package repository
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"time"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// UserRepository persists model.User records.
type UserRepository struct{ db *gorm.DB }
// Create inserts a new user. Caller must pre-hash the password.
func (r *UserRepository) Create(ctx context.Context, u *model.User) error {
return r.db.WithContext(ctx).Create(u).Error
}
// ReleaseDeletedUsername renames soft-deleted rows that still hold a unique
// username so the same account name can be created again.
func (r *UserRepository) ReleaseDeletedUsername(ctx context.Context, username string) error {
if username == "" {
return nil
}
released := username + "__deleted__" + time.Now().Format("20060102150405.000000000")
if len(released) > 64 {
sum := sha256.Sum256([]byte(released))
released = username
if len(released) > 43 {
released = released[:43]
}
released += "__deleted__" + hex.EncodeToString(sum[:])[:10]
}
return r.db.WithContext(ctx).Unscoped().
Model(&model.User{}).
Where("username = ? AND deleted_at IS NOT NULL", username).
Update("username", released).Error
}
// FindByUsername returns the user matching username, or (nil, nil) when absent.
func (r *UserRepository) FindByUsername(ctx context.Context, username string) (*model.User, error) {
var u model.User
err := withSQLiteBusyRetry(ctx, func() error {
u = model.User{}
err := r.db.WithContext(ctx).Where("username = ?", username).First(&u).Error
if errors.Is(err, gorm.ErrRecordNotFound) && username != "" {
err = r.db.WithContext(ctx).Where("LOWER(username) = LOWER(?)", username).First(&u).Error
}
return err
})
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &u, nil
}
// FindByID returns the user with the matching primary key, or (nil, nil).
func (r *UserRepository) FindByID(ctx context.Context, id string) (*model.User, error) {
var u model.User
err := withSQLiteBusyRetry(ctx, func() error {
u = model.User{}
return r.db.WithContext(ctx).Where("id = ?", id).First(&u).Error
})
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &u, nil
}
// Count returns the total number of non-deleted users.
func (r *UserRepository) Count(ctx context.Context) (int64, error) {
var n int64
err := r.db.WithContext(ctx).Model(&model.User{}).Count(&n).Error
return n, err
}
// CountAdmins returns the number of users that hold the admin role.
func (r *UserRepository) CountAdmins(ctx context.Context) (int64, error) {
var n int64
err := r.db.WithContext(ctx).Model(&model.User{}).
Where("role = ?", "admin").Count(&n).Error
return n, err
}
// FirstAdmin returns the earliest admin user. This row represents the protected
// built-in/default administrator even if its username is later changed.
func (r *UserRepository) FirstAdmin(ctx context.Context) (*model.User, error) {
var u model.User
err := r.db.WithContext(ctx).Where("role = ?", "admin").Order("created_at asc").First(&u).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &u, nil
}
// List returns all users ordered by creation time desc.
func (r *UserRepository) List(ctx context.Context) ([]model.User, error) {
var users []model.User
err := r.db.WithContext(ctx).Order("created_at desc").Find(&users).Error
return users, err
}
// UpdateFields applies a narrow set of user field updates.
func (r *UserRepository) UpdateFields(ctx context.Context, id string, updates map[string]any) error {
return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(updates).Error
}
// UpdatePassword sets a new password hash and clears ForcePasswordReset.
func (r *UserRepository) UpdatePassword(ctx context.Context, id, hash string) error {
return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).
Updates(map[string]any{"password_hash": hash, "force_password_reset": false}).Error
}
// TouchLogin updates the last login timestamp.
func (r *UserRepository) TouchLogin(ctx context.Context, id string) error {
now := time.Now()
return withSQLiteBusyRetry(ctx, func() error {
return r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).
Update("last_login_at", &now).Error
})
}
// Delete removes a user (soft-delete via gorm.DeletedAt), releases the unique
// username, and drops Telegram bindings so future re-created users bind cleanly.
func (r *UserRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var user model.User
if err := tx.Where("id = ?", id).First(&user).Error; err != nil {
return err
}
if err := tx.Unscoped().Where("user_id = ?", id).Delete(&model.TelegramBinding{}).Error; err != nil {
return err
}
released := user.Username + "__deleted__" + time.Now().Format("20060102150405.000000000")
if len(released) > 64 {
sum := sha256.Sum256([]byte(user.ID + user.Username))
base := user.Username
if len(base) > 43 {
base = base[:43]
}
released = base + "__deleted__" + hex.EncodeToString(sum[:])[:10]
}
if err := tx.Model(&model.User{}).Where("id = ?", id).Update("username", released).Error; err != nil {
return err
}
return tx.Delete(&model.User{}, "id = ?", id).Error
})
}
+1 -9
View File
@@ -6,9 +6,7 @@ import (
"net/http/httptest"
"testing"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
@@ -88,13 +86,7 @@ func TestAdultProviderUsesConfiguredMultipleSources(t *testing.T) {
}))
defer good.Close()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
t.Fatal(err)
}
db := newServiceTestDB(t, &model.APIConfig{})
apiConfig := NewAPIConfigService(zap.NewNop(), repository.New(db), NewCryptoService("", zap.NewNop()))
baseURL := bad.URL + "\n" + good.URL
if _, err := apiConfig.Update(context.Background(), "adult", APIConfigPatch{BaseURL: &baseURL}); err != nil {
+2 -16
View File
@@ -4,9 +4,7 @@ import (
"context"
"testing"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -14,13 +12,7 @@ import (
)
func TestAIStatusUsesDatabaseOpenAIConfig(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
t.Fatal(err)
}
db := newServiceTestDB(t, &model.APIConfig{})
repo := &repository.Container{DB: db}
crypto := NewCryptoService("test-secret", zap.NewNop())
apiConfig := NewAPIConfigService(zap.NewNop(), repo, crypto)
@@ -52,13 +44,7 @@ func TestAIStatusUsesDatabaseOpenAIConfig(t *testing.T) {
}
func TestAIStatusHonorsDisabledDatabaseOpenAIConfig(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
t.Fatal(err)
}
db := newServiceTestDB(t, &model.APIConfig{})
repo := &repository.Container{DB: db}
apiConfig := NewAPIConfigService(zap.NewNop(), repo, NewCryptoService("test-secret", zap.NewNop()))
key := "sk-test"
+1 -9
View File
@@ -9,11 +9,9 @@ import (
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/golang-jwt/jwt/v5"
"go.uber.org/zap"
"golang.org/x/crypto/bcrypt"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/database"
@@ -23,13 +21,7 @@ import (
func newAuthTestServices(t *testing.T) (*repository.Container, *AuthService, *ProfileService, *PermissionService) {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.TelegramBinding{}, &model.Setting{}); err != nil {
t.Fatal(err)
}
db := newServiceTestDB(t, &model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.TelegramBinding{}, &model.Setting{})
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
+2 -2
View File
@@ -24,7 +24,7 @@ func (c *Container) BootCloudStorageHealthCheck(ctx context.Context) {
cloudConfigs := make([]StorageView, 0)
for _, cfg := range configs {
if cfg.Enabled && (cfg.Type == "quark" || cfg.Type == "cloud115" || cfg.Type == "clouddrive2" || cfg.Type == "openlist") {
if cfg.Enabled && IsAdminCloudConfigurable(cfg.Type) {
cloudConfigs = append(cloudConfigs, cfg)
}
}
@@ -92,7 +92,7 @@ func cloudStorageMissingConfigReason(err error) string {
}
msg := strings.ToLower(strings.TrimSpace(err.Error()))
switch {
case strings.Contains(msg, "missing cookie"):
case strings.Contains(msg, "missing cookie") || (strings.Contains(msg, "missing") && strings.Contains(msg, "cookie")):
return "missing_cookie"
case strings.Contains(msg, "missing webdav url"):
return "missing_webdav_url"
+4 -11
View File
@@ -5,10 +5,8 @@ import (
"errors"
"testing"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
@@ -20,8 +18,9 @@ func TestCloudStorageMissingConfigReason(t *testing.T) {
want string
}{
{errors.New("115: missing cookie"), "missing_cookie"},
{errors.New("openlist: missing cookie"), "missing_cookie"},
{errors.New("clouddrive2: missing WebDAV URL"), "missing_webdav_url"},
{errors.New("quark: token expired"), ""},
{errors.New("openlist: token expired"), ""},
}
for _, tc := range cases {
if got := cloudStorageMissingConfigReason(tc.err); got != tc.want {
@@ -31,19 +30,13 @@ func TestCloudStorageMissingConfigReason(t *testing.T) {
}
func TestWarnMissingCloudStorageConfigOncePersistsMarker(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.Setting{}); err != nil {
t.Fatal(err)
}
db := newServiceTestDB(t, &model.Setting{})
core, observed := observer.New(zap.WarnLevel)
c := &Container{
Log: zap.New(core),
Repo: repository.New(db),
}
err = errors.New("115: missing cookie")
err := errors.New("115: missing cookie")
if !c.warnMissingCloudStorageConfigOnce(context.Background(), "cloud115", err) {
t.Fatal("missing config should be handled")
+246
View File
@@ -0,0 +1,246 @@
package service
import (
"context"
"strings"
"testing"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestBotCleanupRulesDefaultToEmpty(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
cfg := loadBotConfig(ctx, repos)
if len(cfg.AccountCleanupRules) != 0 {
t.Fatalf("default cleanup rules should be empty, got %+v", cfg.AccountCleanupRules)
}
}
func TestBotCleanupRulesCanBeDeletedUntilEmpty(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
if _, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule add watch_hours watch_3_5d_6h 观看3到5天满6小时 3 5 6"); err != nil {
t.Fatal(err)
}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule del watch_3_5d_6h")
if err != nil {
t.Fatal(err)
}
cfg := loadBotConfig(ctx, repos)
if len(cfg.AccountCleanupRules) != 0 {
t.Fatalf("cleanup rules should stay empty after deleting the last rule; reply=%q rules=%+v", reply.Text, cfg.AccountCleanupRules)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup_rule list")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "暂无规则") {
t.Fatalf("expected empty rule list, got %q", reply.Text)
}
}
func TestBotCleanupRunPreviewsBeforeConfirm(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
now := time.Now()
old := now.Add(-30 * 24 * time.Hour)
stale := &model.User{Username: "stale", PasswordHash: "x", Role: "user", IsActive: true}
stale.CreatedAt = old
stale.LastLoginAt = &old
recent := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true}
recent.CreatedAt = old
recent.LastLoginAt = &now
newUser := &model.User{Username: "newbie", PasswordHash: "x", Role: "user", IsActive: true}
newUser.CreatedAt = now
for _, user := range []*model.User{stale, recent, newUser} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupKeepMode, "any"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"最近登录","type":"recent_login","enabled":true,"window_days_max":7},
{"id":"new_7d","name":"新号宽限","type":"account_age_grace","enabled":true,"min_count":7}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "当前只是预览") || !strings.Contains(reply.Text, "stale") || !strings.Contains(reply.Text, "/cleanup run confirm") {
t.Fatalf("cleanup run should preview candidates and confirmation command, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got == nil {
t.Fatal("cleanup preview must not delete the stale user")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/deleted")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "当前只是预览") {
t.Fatalf("/deleted alias should preview only, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got == nil {
t.Fatal("/deleted preview alias must not delete users")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已清理 <b>1</b>") {
t.Fatalf("cleanup confirm should delete exactly one stale user, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got != nil {
t.Fatal("stale user should be deleted after explicit confirmation")
}
for _, user := range []*model.User{recent, newUser, admin} {
if got, _ := repos.User.FindByID(ctx, user.ID); got == nil {
t.Fatalf("%s should be kept by保号 rules/protection", user.Username)
}
}
}
func TestBotCleanupLegacyCountModeStillKeepsSingleMatchedRule(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
now := time.Now()
old := now.Add(-30 * 24 * time.Hour)
recent := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true}
recent.CreatedAt = old
recent.LastLoginAt = &now
stale := &model.User{Username: "stale", PasswordHash: "x", Role: "user", IsActive: true}
stale.CreatedAt = old
stale.LastLoginAt = &old
for _, user := range []*model.User{recent, stale} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupKeepMode, "count"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRequiredCount, "2"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"最近登录","type":"recent_login","enabled":true,"window_days_max":7},
{"id":"new_7d","name":"新号宽限","type":"account_age_grace","enabled":true,"min_count":7}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run")
if err != nil {
t.Fatal(err)
}
if strings.Contains(reply.Text, "recent") {
t.Fatalf("user matching one keep rule must not be a cleanup candidate, got %q", reply.Text)
}
if !strings.Contains(reply.Text, "stale") {
t.Fatalf("user matching no keep rules should be a candidate, got %q", reply.Text)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if got, _ := repos.User.FindByID(ctx, recent.ID); got == nil {
t.Fatal("legacy count mode must not delete a user matching one keep rule")
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got != nil {
t.Fatalf("stale user should be deleted after confirm, reply=%q", reply.Text)
}
}
func TestBotCleanupConfirmRequiresEnabledRules(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
user := &model.User{Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
user.CreatedAt = time.Now().Add(-30 * 24 * time.Hour)
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "没有启用的保号规则") {
t.Fatalf("cleanup confirm without rules should be blocked, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, user.ID); got == nil {
t.Fatal("cleanup confirm without enabled rules must not delete users")
}
}
func TestBotCleanupRuleListInfersDaysAndHidesDuplicateNames(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"login_7d","type":"recent_login","enabled":true,"window_days_min":1,"window_days_max":5,"min_count":1},
{"id":"new_7d","name":"new_7d","type":"account_age_grace","enabled":true,"window_days_min":1,"window_days_max":1,"min_count":1}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule")
if err != nil {
t.Fatal(err)
}
for _, bad := range []string{"login_7d</code> · login_7d", "new_7d</code> · new_7d", "5 天内登录", "新号宽限 1 天", "add watch_hours", "Mgo 保号规则命令"} {
if strings.Contains(reply.Text, bad) {
t.Fatalf("rule list still contains bad fragment %q: %s", bad, reply.Text)
}
}
for _, want := range []string{"login_7d", "7 天内登录", "new_7d", "新号宽限 7 天"} {
if !strings.Contains(reply.Text, want) {
t.Fatalf("rule list missing %q: %s", want, reply.Text)
}
}
}
+1 -496
View File
@@ -2,14 +2,11 @@ package service
import (
"context"
"encoding/json"
"strings"
"testing"
"time"
"github.com/glebarez/sqlite"
"go.uber.org/zap"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
@@ -18,13 +15,7 @@ import (
func newBotTestService(t *testing.T) (*repository.Container, *TelegramBotService) {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(model.AllModels()...); err != nil {
t.Fatal(err)
}
db := newServiceTestDB(t, model.AllModels()...)
repos := repository.New(db)
cfg := &config.Config{}
cfg.Secrets.JWTSecret = "test-secret"
@@ -503,242 +494,6 @@ func TestBotAdminCommandsManageDevicePolicy(t *testing.T) {
}
}
func TestBotCleanupRulesDefaultToEmpty(t *testing.T) {
ctx := context.Background()
repos, _ := newBotTestService(t)
cfg := loadBotConfig(ctx, repos)
if len(cfg.AccountCleanupRules) != 0 {
t.Fatalf("default cleanup rules should be empty, got %+v", cfg.AccountCleanupRules)
}
}
func TestBotCleanupRulesCanBeDeletedUntilEmpty(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
if _, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule add watch_hours watch_3_5d_6h 观看3到5天满6小时 3 5 6"); err != nil {
t.Fatal(err)
}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule del watch_3_5d_6h")
if err != nil {
t.Fatal(err)
}
cfg := loadBotConfig(ctx, repos)
if len(cfg.AccountCleanupRules) != 0 {
t.Fatalf("cleanup rules should stay empty after deleting the last rule; reply=%q rules=%+v", reply.Text, cfg.AccountCleanupRules)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup_rule list")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "暂无规则") {
t.Fatalf("expected empty rule list, got %q", reply.Text)
}
}
func TestBotCleanupRunPreviewsBeforeConfirm(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
now := time.Now()
old := now.Add(-30 * 24 * time.Hour)
stale := &model.User{Username: "stale", PasswordHash: "x", Role: "user", IsActive: true}
stale.CreatedAt = old
stale.LastLoginAt = &old
recent := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true}
recent.CreatedAt = old
recent.LastLoginAt = &now
newUser := &model.User{Username: "newbie", PasswordHash: "x", Role: "user", IsActive: true}
newUser.CreatedAt = now
for _, user := range []*model.User{stale, recent, newUser} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupKeepMode, "any"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"最近登录","type":"recent_login","enabled":true,"window_days_max":7},
{"id":"new_7d","name":"新号宽限","type":"account_age_grace","enabled":true,"min_count":7}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "当前只是预览") || !strings.Contains(reply.Text, "stale") || !strings.Contains(reply.Text, "/cleanup run confirm") {
t.Fatalf("cleanup run should preview candidates and confirmation command, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got == nil {
t.Fatal("cleanup preview must not delete the stale user")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/deleted")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "当前只是预览") {
t.Fatalf("/deleted alias should preview only, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got == nil {
t.Fatal("/deleted preview alias must not delete users")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已清理 <b>1</b>") {
t.Fatalf("cleanup confirm should delete exactly one stale user, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got != nil {
t.Fatal("stale user should be deleted after explicit confirmation")
}
for _, user := range []*model.User{recent, newUser, admin} {
if got, _ := repos.User.FindByID(ctx, user.ID); got == nil {
t.Fatalf("%s should be kept by保号 rules/protection", user.Username)
}
}
}
func TestBotCleanupLegacyCountModeStillKeepsSingleMatchedRule(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
now := time.Now()
old := now.Add(-30 * 24 * time.Hour)
recent := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true}
recent.CreatedAt = old
recent.LastLoginAt = &now
stale := &model.User{Username: "stale", PasswordHash: "x", Role: "user", IsActive: true}
stale.CreatedAt = old
stale.LastLoginAt = &old
for _, user := range []*model.User{recent, stale} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupKeepMode, "count"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRequiredCount, "2"); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"最近登录","type":"recent_login","enabled":true,"window_days_max":7},
{"id":"new_7d","name":"新号宽限","type":"account_age_grace","enabled":true,"min_count":7}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run")
if err != nil {
t.Fatal(err)
}
if strings.Contains(reply.Text, "recent") {
t.Fatalf("user matching one keep rule must not be a cleanup candidate, got %q", reply.Text)
}
if !strings.Contains(reply.Text, "stale") {
t.Fatalf("user matching no keep rules should be a candidate, got %q", reply.Text)
}
reply, err = bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if got, _ := repos.User.FindByID(ctx, recent.ID); got == nil {
t.Fatal("legacy count mode must not delete a user matching one keep rule")
}
if got, _ := repos.User.FindByID(ctx, stale.ID); got != nil {
t.Fatalf("stale user should be deleted after confirm, reply=%q", reply.Text)
}
}
func TestBotCleanupConfirmRequiresEnabledRules(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
user := &model.User{Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
user.CreatedAt = time.Now().Add(-30 * 24 * time.Hour)
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupEnabled, "true"); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup run confirm")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "没有启用的保号规则") {
t.Fatalf("cleanup confirm without rules should be blocked, got %q", reply.Text)
}
if got, _ := repos.User.FindByID(ctx, user.ID); got == nil {
t.Fatal("cleanup confirm without enabled rules must not delete users")
}
}
func TestBotCleanupRuleListInfersDaysAndHidesDuplicateNames(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
if err := repos.User.Create(ctx, admin); err != nil {
t.Fatal(err)
}
if err := repos.Setting.Set(ctx, SettingAccountCleanupRules, `[
{"id":"login_7d","name":"login_7d","type":"recent_login","enabled":true,"window_days_min":1,"window_days_max":5,"min_count":1},
{"id":"new_7d","name":"new_7d","type":"account_age_grace","enabled":true,"window_days_min":1,"window_days_max":1,"min_count":1}
]`); err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9001, Username: "root"}, Chat: TelegramChat{ID: 9001, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/cleanup_rule")
if err != nil {
t.Fatal(err)
}
for _, bad := range []string{"login_7d</code> · login_7d", "new_7d</code> · new_7d", "5 天内登录", "新号宽限 1 天", "add watch_hours", "Mgo 保号规则命令"} {
if strings.Contains(reply.Text, bad) {
t.Fatalf("rule list still contains bad fragment %q: %s", bad, reply.Text)
}
}
for _, want := range []string{"login_7d", "7 天内登录", "new_7d", "新号宽限 7 天"} {
if !strings.Contains(reply.Text, want) {
t.Fatalf("rule list missing %q: %s", want, reply.Text)
}
}
}
func TestBotRegistrationCommandUsesOpenRegQuota(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
@@ -828,137 +583,6 @@ func TestBotUserCommandsAndAdminGate(t *testing.T) {
}
}
func TestBotRedeemRegisterRequiresAllowedTelegramUser(t *testing.T) {
ctx := context.Background()
_, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9201, Username: "outsider"}, Chat: TelegramChat{ID: 9201, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "不在管理员配置") {
t.Fatalf("outsider should not redeem register code, got %q", reply.Text)
}
channel.Config = `{"admin_user_ids":"9201"}`
reply, err = bot.executeCommand(ctx, channel, msg, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("allowed user should redeem register code, got %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9201); binding == nil {
t.Fatal("redeemed account should be bound to telegram user")
}
}
func TestBotRedeemRegisterCodeCreatesOnlyOneAccount(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9201,9202"}`}
first := &TelegramMessage{From: TelegramUser{ID: 9201, Username: "first"}, Chat: TelegramChat{ID: 9201, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, first, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("first redeem should succeed, got %q", reply.Text)
}
second := &TelegramMessage{From: TelegramUser{ID: 9202, Username: "second"}, Chat: TelegramChat{ID: 9202, Type: "private"}}
reply, err = bot.executeCommand(ctx, channel, second, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换码已被使用") && !strings.Contains(reply.Text, "兑换码刚刚被使用") {
t.Fatalf("second redeem should be rejected as used, got %q", reply.Text)
}
var users int64
if err := repos.DB.Model(&model.User{}).Count(&users).Error; err != nil {
t.Fatal(err)
}
if users != 1 {
t.Fatalf("one register code must create exactly one user, got %d", users)
}
if binding := bot.telegramBinding(ctx, 9202); binding != nil {
t.Fatal("second telegram user must not be bound by an already-used register code")
}
}
func TestBotRegisterCommandAcceptsRegistrationCode(t *testing.T) {
ctx := context.Background()
_, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9301"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9301, Username: "codeuser"}, Chat: TelegramChat{ID: 9301, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/register "+strings.ToLower(code.Code[:4])+"-"+strings.ToLower(code.Code[4:]))
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("/register CODE should redeem registration code, got %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9301); binding == nil {
t.Fatal("register code should bind the newly created account")
}
}
func TestBotPlainRegistrationCodeMessageRedeems(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.NotifyChannel{
Name: "Telegram",
Type: "telegram",
Enabled: true,
Config: `{"admin_user_ids":"9302"}`,
}).Error; err != nil {
t.Fatal(err)
}
update, _ := json.Marshal(TelegramUpdate{
UpdateID: 1,
Message: &TelegramMessage{
MessageID: 12,
Text: strings.ToLower(code.Code),
From: TelegramUser{ID: 9302, Username: "plaincode"},
Chat: TelegramChat{ID: 9302, Type: "private"},
},
})
if err := bot.HandleWebhook(ctx, update); err != nil {
t.Fatal(err)
}
if binding := bot.telegramBinding(ctx, 9302); binding == nil {
t.Fatal("plain code private message should redeem and bind account")
}
var used model.RegistrationCode
if err := repos.DB.Where("code = ?", code.Code).First(&used).Error; err != nil {
t.Fatal(err)
}
if used.UsedAt == nil || used.UsedByUserID == "" {
t.Fatal("plain code message should mark registration code as used")
}
}
func TestBotAdminCodeAndUserCommands(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
@@ -1032,122 +656,3 @@ func TestBotGroupMenuShowsAdminActionsOnlyForAdmins(t *testing.T) {
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)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
viewer := &model.User{Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
guest := &model.User{Username: "guest", PasswordHash: "x", Role: "user", IsActive: true}
for _, user := range []*model.User{admin, viewer, guest} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
bindings := []model.TelegramBinding{
{TelegramUserID: 9401, TelegramName: "@root", ChatID: 9401, UserID: admin.ID},
{TelegramUserID: 9402, TelegramName: "@viewer", ChatID: 9402, UserID: viewer.ID},
{TelegramUserID: 9403, TelegramName: "@guest", ChatID: 9403, UserID: guest.ID},
}
for i := range bindings {
if err := repos.DB.Create(&bindings[i]).Error; err != nil {
t.Fatal(err)
}
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9401"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9401, Username: "root"}, Chat: TelegramChat{ID: 9401, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/unbind viewer,guest missing root")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>2</b>") || !strings.Contains(reply.Text, "root(管理员)") || !strings.Contains(reply.Text, "missing") {
t.Fatalf("unexpected unbind reply: %q", reply.Text)
}
for _, user := range []*model.User{viewer, guest} {
var count int64
if err := repos.DB.Model(&model.TelegramBinding{}).Where("user_id = ?", user.ID).Count(&count).Error; err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("%s binding count = %d, want 0", user.Username, count)
}
}
if binding := bot.telegramBinding(ctx, 9401); binding == nil {
t.Fatal("admin binding should be protected from /unbind by username")
}
}
func TestBotAdminUnbindInactiveAndInvalidBindings(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
oldTime := time.Now().Add(-45 * 24 * time.Hour)
recentTime := time.Now().Add(-2 * 24 * time.Hour)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true, LastLoginAt: &oldTime}
oldUser := &model.User{Username: "old", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &oldTime}
recentUser := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &recentTime}
for _, user := range []*model.User{admin, oldUser, recentUser} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
for _, binding := range []model.TelegramBinding{
{TelegramUserID: 9501, TelegramName: "@root", ChatID: 9501, UserID: admin.ID},
{TelegramUserID: 9502, TelegramName: "@old", ChatID: 9502, UserID: oldUser.ID},
{TelegramUserID: 9503, TelegramName: "@recent", ChatID: 9503, UserID: recentUser.ID},
{TelegramUserID: 9504, TelegramName: "@ghost", ChatID: 9504, UserID: "missing-user"},
} {
row := binding
if err := repos.DB.Create(&row).Error; err != nil {
t.Fatal(err)
}
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9501"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9501, Username: "root"}, Chat: TelegramChat{ID: 9501, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/unbind_inactive 30")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>1</b>") || !strings.Contains(reply.Text, "old") {
t.Fatalf("unexpected inactive unbind reply: %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9502); binding != nil {
t.Fatal("old user binding should be removed")
}
if binding := bot.telegramBinding(ctx, 9501); binding == nil {
t.Fatal("admin binding should be skipped by inactive cleanup")
}
if binding := bot.telegramBinding(ctx, 9503); binding == nil {
t.Fatal("recent user binding should remain")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/unbind_duplicates")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>1</b>") || !strings.Contains(reply.Text, "tg:9504") {
t.Fatalf("unexpected duplicate cleanup reply: %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9504); binding != nil {
t.Fatal("invalid binding should be removed")
}
}
func TestTelegramMembershipChatIDsIncludesCommandChatID(t *testing.T) {
_, bot := newBotTestService(t)
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"command_chat_id":"-100123"}`}
got := bot.telegramMembershipChatIDs(channel)
if len(got) != 1 || got[0] != "-100123" {
t.Fatalf("telegramMembershipChatIDs() = %#v, want command_chat_id", got)
}
}
func TestTelegramMembershipChatIDsDedupesGroupChannelAndCommandIDs(t *testing.T) {
_, bot := newBotTestService(t)
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"group_chat_id":"-100123","channel_chat_id":"-100124","command_chat_id":"-100123"}`}
got := bot.telegramMembershipChatIDs(channel)
if len(got) != 2 || got[0] != "-100123" || got[1] != "-100124" {
t.Fatalf("telegramMembershipChatIDs() = %#v, want deduped ids", got)
}
}
@@ -0,0 +1,141 @@
package service
import (
"context"
"encoding/json"
"strings"
"testing"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestBotRedeemRegisterRequiresAllowedTelegramUser(t *testing.T) {
ctx := context.Background()
_, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9001"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9201, Username: "outsider"}, Chat: TelegramChat{ID: 9201, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "不在管理员配置") {
t.Fatalf("outsider should not redeem register code, got %q", reply.Text)
}
channel.Config = `{"admin_user_ids":"9201"}`
reply, err = bot.executeCommand(ctx, channel, msg, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("allowed user should redeem register code, got %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9201); binding == nil {
t.Fatal("redeemed account should be bound to telegram user")
}
}
func TestBotRedeemRegisterCodeCreatesOnlyOneAccount(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9201,9202"}`}
first := &TelegramMessage{From: TelegramUser{ID: 9201, Username: "first"}, Chat: TelegramChat{ID: 9201, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, first, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("first redeem should succeed, got %q", reply.Text)
}
second := &TelegramMessage{From: TelegramUser{ID: 9202, Username: "second"}, Chat: TelegramChat{ID: 9202, Type: "private"}}
reply, err = bot.executeCommand(ctx, channel, second, "/redeem_register "+code.Code)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换码已被使用") && !strings.Contains(reply.Text, "兑换码刚刚被使用") {
t.Fatalf("second redeem should be rejected as used, got %q", reply.Text)
}
var users int64
if err := repos.DB.Model(&model.User{}).Count(&users).Error; err != nil {
t.Fatal(err)
}
if users != 1 {
t.Fatalf("one register code must create exactly one user, got %d", users)
}
if binding := bot.telegramBinding(ctx, 9202); binding != nil {
t.Fatal("second telegram user must not be bound by an already-used register code")
}
}
func TestBotRegisterCommandAcceptsRegistrationCode(t *testing.T) {
ctx := context.Background()
_, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9301"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9301, Username: "codeuser"}, Chat: TelegramChat{ID: 9301, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/register "+strings.ToLower(code.Code[:4])+"-"+strings.ToLower(code.Code[4:]))
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "兑换成功") {
t.Fatalf("/register CODE should redeem registration code, got %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9301); binding == nil {
t.Fatal("register code should bind the newly created account")
}
}
func TestBotPlainRegistrationCodeMessageRedeems(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
code, err := bot.generateCode(ctx, model.RegistrationCodeRegister, 30, 0, "")
if err != nil {
t.Fatal(err)
}
if err := repos.DB.Create(&model.NotifyChannel{
Name: "Telegram",
Type: "telegram",
Enabled: true,
Config: `{"admin_user_ids":"9302"}`,
}).Error; err != nil {
t.Fatal(err)
}
update, _ := json.Marshal(TelegramUpdate{
UpdateID: 1,
Message: &TelegramMessage{
MessageID: 12,
Text: strings.ToLower(code.Code),
From: TelegramUser{ID: 9302, Username: "plaincode"},
Chat: TelegramChat{ID: 9302, Type: "private"},
},
})
if err := bot.HandleWebhook(ctx, update); err != nil {
t.Fatal(err)
}
if binding := bot.telegramBinding(ctx, 9302); binding == nil {
t.Fatal("plain code private message should redeem and bind account")
}
var used model.RegistrationCode
if err := repos.DB.Where("code = ?", code.Code).First(&used).Error; err != nil {
t.Fatal(err)
}
if used.UsedAt == nil || used.UsedByUserID == "" {
t.Fatal("plain code message should mark registration code as used")
}
}
+129
View File
@@ -0,0 +1,129 @@
package service
import (
"context"
"strings"
"testing"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
func TestBotAdminUnbindMultipleUsers(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true}
viewer := &model.User{Username: "viewer", PasswordHash: "x", Role: "user", IsActive: true}
guest := &model.User{Username: "guest", PasswordHash: "x", Role: "user", IsActive: true}
for _, user := range []*model.User{admin, viewer, guest} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
bindings := []model.TelegramBinding{
{TelegramUserID: 9401, TelegramName: "@root", ChatID: 9401, UserID: admin.ID},
{TelegramUserID: 9402, TelegramName: "@viewer", ChatID: 9402, UserID: viewer.ID},
{TelegramUserID: 9403, TelegramName: "@guest", ChatID: 9403, UserID: guest.ID},
}
for i := range bindings {
if err := repos.DB.Create(&bindings[i]).Error; err != nil {
t.Fatal(err)
}
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9401"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9401, Username: "root"}, Chat: TelegramChat{ID: 9401, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/unbind viewer,guest missing root")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>2</b>") || !strings.Contains(reply.Text, "root(管理员)") || !strings.Contains(reply.Text, "missing") {
t.Fatalf("unexpected unbind reply: %q", reply.Text)
}
for _, user := range []*model.User{viewer, guest} {
var count int64
if err := repos.DB.Model(&model.TelegramBinding{}).Where("user_id = ?", user.ID).Count(&count).Error; err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("%s binding count = %d, want 0", user.Username, count)
}
}
if binding := bot.telegramBinding(ctx, 9401); binding == nil {
t.Fatal("admin binding should be protected from /unbind by username")
}
}
func TestBotAdminUnbindInactiveAndInvalidBindings(t *testing.T) {
ctx := context.Background()
repos, bot := newBotTestService(t)
oldTime := time.Now().Add(-45 * 24 * time.Hour)
recentTime := time.Now().Add(-2 * 24 * time.Hour)
admin := &model.User{Username: "root", PasswordHash: "x", Role: "admin", IsActive: true, LastLoginAt: &oldTime}
oldUser := &model.User{Username: "old", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &oldTime}
recentUser := &model.User{Username: "recent", PasswordHash: "x", Role: "user", IsActive: true, LastLoginAt: &recentTime}
for _, user := range []*model.User{admin, oldUser, recentUser} {
if err := repos.User.Create(ctx, user); err != nil {
t.Fatal(err)
}
}
for _, binding := range []model.TelegramBinding{
{TelegramUserID: 9501, TelegramName: "@root", ChatID: 9501, UserID: admin.ID},
{TelegramUserID: 9502, TelegramName: "@old", ChatID: 9502, UserID: oldUser.ID},
{TelegramUserID: 9503, TelegramName: "@recent", ChatID: 9503, UserID: recentUser.ID},
{TelegramUserID: 9504, TelegramName: "@ghost", ChatID: 9504, UserID: "missing-user"},
} {
row := binding
if err := repos.DB.Create(&row).Error; err != nil {
t.Fatal(err)
}
}
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"admin_user_ids":"9501"}`}
msg := &TelegramMessage{From: TelegramUser{ID: 9501, Username: "root"}, Chat: TelegramChat{ID: 9501, Type: "private"}}
reply, err := bot.executeCommand(ctx, channel, msg, "/unbind_inactive 30")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>1</b>") || !strings.Contains(reply.Text, "old") {
t.Fatalf("unexpected inactive unbind reply: %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9502); binding != nil {
t.Fatal("old user binding should be removed")
}
if binding := bot.telegramBinding(ctx, 9501); binding == nil {
t.Fatal("admin binding should be skipped by inactive cleanup")
}
if binding := bot.telegramBinding(ctx, 9503); binding == nil {
t.Fatal("recent user binding should remain")
}
reply, err = bot.executeCommand(ctx, channel, msg, "/unbind_duplicates")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(reply.Text, "已解绑:<b>1</b>") || !strings.Contains(reply.Text, "tg:9504") {
t.Fatalf("unexpected duplicate cleanup reply: %q", reply.Text)
}
if binding := bot.telegramBinding(ctx, 9504); binding != nil {
t.Fatal("invalid binding should be removed")
}
}
func TestTelegramMembershipChatIDsIncludesCommandChatID(t *testing.T) {
_, bot := newBotTestService(t)
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"command_chat_id":"-100123"}`}
got := bot.telegramMembershipChatIDs(channel)
if len(got) != 1 || got[0] != "-100123" {
t.Fatalf("telegramMembershipChatIDs() = %#v, want command_chat_id", got)
}
}
func TestTelegramMembershipChatIDsDedupesGroupChannelAndCommandIDs(t *testing.T) {
_, bot := newBotTestService(t)
channel := &model.NotifyChannel{Name: "Telegram", Type: "telegram", Enabled: true, Config: `{"group_chat_id":"-100123","channel_chat_id":"-100124","command_chat_id":"-100123"}`}
got := bot.telegramMembershipChatIDs(channel)
if len(got) != 2 || got[0] != "-100123" || got[1] != "-100124" {
t.Fatalf("telegramMembershipChatIDs() = %#v, want deduped ids", got)
}
}
+19 -7
View File
@@ -27,7 +27,6 @@ var timeNow = time.Now
// Provider types recognised by the registry.
const (
TypeQuark = "quark" // 夸克网盘
Type115 = "cloud115" // 115 网盘
TypeCloudDrive2 = "clouddrive2" // CloudDrive2 桥接网盘
TypeOpenList = "openlist" // OpenList / AList-compatible bridge
@@ -42,7 +41,7 @@ type FileEntry struct {
Name string `json:"name"`
IsDir bool `json:"is_dir"`
Size int64 `json:"size"`
// PickCode is 115-specific; quark uses ID directly.
// PickCode is 115-specific; other providers use ID directly.
PickCode string `json:"pick_code,omitempty"`
}
@@ -59,7 +58,7 @@ type DirectLink struct {
// Provider is the common cloud-disk interface.
type Provider interface {
// Type returns the provider key (TypeQuark / Type115).
// Type returns the provider key.
Type() string
// Ping validates the stored credentials (cookie). Cheap, used by the
// storage-config Test() probe.
@@ -71,6 +70,21 @@ type Provider interface {
Resolve(ctx context.Context, fileRef string) (*DirectLink, error)
}
// MutableProvider is implemented by cloud bridges that support safe folder
// management through their official API or standard WebDAV methods.
type MutableProvider interface {
Provider
Mkdir(ctx context.Context, parentDir, name string) (*FileEntry, error)
Rename(ctx context.Context, ref, name string) (*FileEntry, error)
}
// MovableProvider is implemented by writable cloud bridges that can move an
// entry across directories, optionally renaming it in the same operation.
type MovableProvider interface {
MutableProvider
Move(ctx context.Context, ref, targetDir, name string) (*FileEntry, error)
}
// New constructs a provider of the given type from a free-form config map
// (as persisted by StorageConfigService). The client is shared so callers can
// inject timeouts / test transports.
@@ -79,8 +93,6 @@ func New(typ string, cfg map[string]any, client *http.Client) (Provider, error)
client = http.DefaultClient
}
switch typ {
case TypeQuark:
return newQuark(cfg, client), nil
case Type115:
return new115(cfg, client), nil
case TypeCloudDrive2:
@@ -94,7 +106,7 @@ func New(typ string, cfg map[string]any, client *http.Client) (Provider, error)
// IsCloudType reports whether typ is a cloud-disk provider.
func IsCloudType(typ string) bool {
return typ == TypeQuark || typ == Type115 || typ == TypeCloudDrive2 || typ == TypeOpenList
return typ == Type115 || typ == TypeCloudDrive2 || typ == TypeOpenList
}
// str coerces a config value to a trimmed string.
@@ -121,5 +133,5 @@ func boolish(v any) bool {
}
}
// defaultUA is a desktop browser UA accepted by both 115 and quark.
// defaultUA is a desktop browser UA accepted by upstream cloud providers.
const defaultUA = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/124.0 Safari/537.36"
+212
View File
@@ -0,0 +1,212 @@
package cloud
import (
"context"
"encoding/base64"
"fmt"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
)
func Test115ListAndResolve(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/files":
if r.URL.Query().Get("cid") != "0" {
t.Errorf("bad cid %q", r.URL.Query().Get("cid"))
}
w.Write([]byte(`{"state":true,"data":[
{"cid":"100","n":"Movies","s":0},
{"fid":"200","n":"Inception.mkv","s":456,"pc":"pick200"}]}`))
default:
t.Errorf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1; CID=2", "base": srv.URL}, srv.Client())
if err != nil {
t.Fatal(err)
}
// The downurl endpoint is m115-encrypted end-to-end (the server side
// requires 115's private key), so stub the decrypted payload via the seam
// and assert the pickcode->URL extraction. The live crypto/transport path is
// exercised by integration testing against the real 115 API.
p115, ok := p.(*pan115Provider)
if !ok {
t.Fatalf("expected *pan115Provider, got %T", p)
}
p115.downURLPayload = func(ctx context.Context, pickcode string) ([]byte, error) {
if pickcode != "pick200" {
t.Errorf("bad pickcode %q", pickcode)
}
return []byte(`{"200":{"file_name":"Inception.mkv","file_size":"456","url":{"url":"https://cdn.115/x.mkv?t=1"}}}`), nil
}
entries, err := p.List(context.Background(), "")
if err != nil {
t.Fatalf("list: %v", err)
}
if len(entries) != 2 {
t.Fatalf("want 2 entries: %#v", entries)
}
if !entries[0].IsDir || entries[0].ID != "100" {
t.Fatalf("dir entry wrong: %#v", entries[0])
}
if entries[1].IsDir || entries[1].PickCode != "pick200" || entries[1].Size != 456 {
t.Fatalf("file entry wrong: %#v", entries[1])
}
link, err := p.Resolve(context.Background(), "pick200")
if err != nil {
t.Fatalf("resolve: %v", err)
}
if link.URL != "https://cdn.115/x.mkv?t=1" {
t.Fatalf("bad url: %s", link.URL)
}
if link.Proxy {
t.Fatalf("115 should default to 302 (no proxy)")
}
}
func Test115ListPaginates(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/files" {
t.Fatalf("unexpected path %s", r.URL.Path)
}
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
count := 100
if offset > 0 {
count = 1
}
items := make([]string, 0, count)
for i := 0; i < count; i++ {
n := offset + i
items = append(items, fmt.Sprintf(`{"fid":"%d","n":"Movie.%03d.mkv","s":%d,"pc":"pick%d"}`, n, n, n, n))
}
w.Write([]byte(`{"state":true,"data":[` + strings.Join(items, ",") + `]}`))
}))
defer srv.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1; CID=2", "base": srv.URL}, srv.Client())
if err != nil {
t.Fatal(err)
}
entries, err := p.List(context.Background(), "0")
if err != nil {
t.Fatalf("list: %v", err)
}
if len(entries) != 101 {
t.Fatalf("entries = %d, want 101", len(entries))
}
if entries[100].ID != "100" || entries[100].PickCode != "pick100" {
t.Fatalf("last entry wrong: %#v", entries[100])
}
}
// Test115DownURLEndpointAndError exercises the live fetchDownURLPayload path:
// it must POST an m115-encrypted `data` body to /app/chrome/downurl?t=... and
// surface 115's error when state=false (no decryption needed for that branch).
func Test115DownURLEndpointAndError(t *testing.T) {
var gotData, gotT string
pro := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/app/chrome/downurl" {
t.Errorf("unexpected path %s", r.URL.Path)
}
gotT = r.URL.Query().Get("t")
_ = r.ParseForm()
gotData = r.PostFormValue("data")
w.Write([]byte(`{"state":false,"error":"not exist"}`))
}))
defer pro.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1", "pro_base": pro.URL}, pro.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.Resolve(context.Background(), "pickX")
if err == nil || !strings.Contains(err.Error(), "not exist") {
t.Fatalf("want upstream error surfaced, got %v", err)
}
if gotT == "" {
t.Errorf("missing t query param")
}
if gotData == "" {
t.Errorf("missing encrypted data body")
}
if _, derr := base64.StdEncoding.DecodeString(gotData); derr != nil {
t.Errorf("data body is not base64: %v", derr)
}
}
func Test115QRFlow(t *testing.T) {
// status sequence: waiting -> scanned -> confirmed
calls := 0
api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/1.0/web/1.0/token/":
w.Write([]byte(`{"state":1,"data":{"uid":"U1","time":1700,"sign":"S1"}}`))
case "/get/status/":
if r.URL.Query().Get("uid") != "U1" {
t.Errorf("bad uid %q", r.URL.Query().Get("uid"))
}
calls++
switch calls {
case 1:
w.Write([]byte(`{"state":1,"data":{"status":0}}`))
case 2:
w.Write([]byte(`{"state":1,"data":{"status":1}}`))
default:
w.Write([]byte(`{"state":1,"data":{"status":2}}`))
}
default:
t.Errorf("unexpected api path %s", r.URL.Path)
}
}))
defer api.Close()
passport := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/app/1.0/web/1.0/login/qrcode/" {
t.Errorf("unexpected passport path %s", r.URL.Path)
}
w.Write([]byte(`{"state":1,"data":{"cookie":{"UID":"u","CID":"c","SEID":"s"}}}`))
}))
defer passport.Close()
oldA, oldP := qr115APIBase, qr115PassportBase
qr115APIBase, qr115PassportBase = api.URL, passport.URL
defer func() { qr115APIBase, qr115PassportBase = oldA, oldP }()
ctx := context.Background()
sess, err := QRStart(ctx, api.Client())
if err != nil {
t.Fatalf("qr start: %v", err)
}
if sess.UID != "U1" || sess.QRImageURL == "" {
t.Fatalf("bad session: %#v", sess)
}
want := []string{"waiting", "scanned", "confirmed"}
for i, exp := range want {
st, err := QRPoll(ctx, api.Client(), sess)
if err != nil {
t.Fatalf("poll %d: %v", i, err)
}
if st.State != exp {
t.Fatalf("poll %d: want %s got %s", i, exp, st.State)
}
if exp == "confirmed" {
if st.Cookie == "" || !containsAll(st.Cookie, "UID=u", "SEID=s") {
t.Fatalf("confirmed must yield cookie: %q", st.Cookie)
}
}
}
}
func containsAll(s string, subs ...string) bool {
for _, sub := range subs {
if !strings.Contains(s, sub) {
return false
}
}
return true
}
+142 -305
View File
@@ -2,115 +2,15 @@ package cloud
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
"time"
)
func TestQuarkListAndResolve(t *testing.T) {
var gotCookie string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotCookie = r.Header.Get("Cookie")
switch {
case r.URL.Path == "/file/sort":
if r.URL.Query().Get("pdir_fid") != "0" {
t.Errorf("unexpected pdir_fid %q", r.URL.Query().Get("pdir_fid"))
}
w.Write([]byte(`{"status":200,"code":0,"data":{"list":[
{"fid":"d1","file_name":"Movies","dir":true,"size":0},
{"fid":"f1","file_name":"Inception.mkv","dir":false,"size":123}]}}`))
case r.URL.Path == "/file/download":
if r.Method != http.MethodPost {
t.Errorf("download must be POST, got %s", r.Method)
}
w.Write([]byte(`{"status":200,"code":0,"data":[{"fid":"f1","download_url":"https://cdn.quark/x.mkv?sign=1"}]}`))
default:
t.Errorf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeQuark, map[string]any{"cookie": "kps=abc", "base": srv.URL}, srv.Client())
if err != nil {
t.Fatal(err)
}
entries, err := p.List(context.Background(), "0")
if err != nil {
t.Fatalf("list: %v", err)
}
if len(entries) != 2 || !entries[0].IsDir || entries[1].Name != "Inception.mkv" || entries[1].Size != 123 {
t.Fatalf("unexpected entries: %#v", entries)
}
if gotCookie != "kps=abc" {
t.Fatalf("cookie not forwarded: %q", gotCookie)
}
link, err := p.Resolve(context.Background(), "f1")
if err != nil {
t.Fatalf("resolve: %v", err)
}
if link.URL != "https://cdn.quark/x.mkv?sign=1" {
t.Fatalf("bad url: %s", link.URL)
}
if !link.Proxy {
t.Fatalf("quark should default to proxy mode")
}
if link.Headers["Cookie"] != "kps=abc" {
t.Fatalf("resolve must carry cookie header: %#v", link.Headers)
}
}
func TestQuarkListPaginates(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/file/sort" {
t.Fatalf("unexpected path %s", r.URL.Path)
}
page, _ := strconv.Atoi(r.URL.Query().Get("_page"))
w.Write([]byte(`{"status":200,"code":0,"data":{"list":[` + quarkPagePayload(page) + `]}}`))
}))
defer srv.Close()
p, err := New(TypeQuark, map[string]any{"cookie": "kps=abc", "base": srv.URL}, srv.Client())
if err != nil {
t.Fatal(err)
}
entries, err := p.List(context.Background(), "0")
if err != nil {
t.Fatalf("list: %v", err)
}
if len(entries) != 101 {
t.Fatalf("entries = %d, want 101", len(entries))
}
if entries[100].ID != "f100" || entries[100].Name != "Movie.100.mkv" {
t.Fatalf("last entry wrong: %#v", entries[100])
}
}
func quarkPagePayload(page int) string {
count := 100
offset := 0
if page > 1 {
count = 1
offset = 100
}
items := make([]string, 0, count)
for i := 0; i < count; i++ {
n := offset + i
items = append(items, fmt.Sprintf(`{"fid":"f%d","file_name":"Movie.%03d.mkv","dir":false,"size":%d}`, n, n, n))
}
return strings.Join(items, ",")
}
func TestDeprecatedProviderPlaybackOverrideKeysAreIgnored(t *testing.T) {
quark := newQuark(map[string]any{"cookie": "c", "force_302": "true"}, http.DefaultClient)
if !quark.proxy {
t.Fatalf("quark should keep safe proxy mode; force_302 is deprecated")
}
pan115 := new115(map[string]any{"cookie": "UID=1; CID=2", "force_proxy": "true"}, http.DefaultClient)
if pan115.proxy {
t.Fatalf("115 should keep safe direct mode; force_proxy is deprecated")
@@ -121,197 +21,6 @@ func TestDeprecatedProviderPlaybackOverrideKeysAreIgnored(t *testing.T) {
}
}
func Test115ListAndResolve(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/files":
if r.URL.Query().Get("cid") != "0" {
t.Errorf("bad cid %q", r.URL.Query().Get("cid"))
}
w.Write([]byte(`{"state":true,"data":[
{"cid":"100","n":"Movies","s":0},
{"fid":"200","n":"Inception.mkv","s":456,"pc":"pick200"}]}`))
default:
t.Errorf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1; CID=2", "base": srv.URL}, srv.Client())
if err != nil {
t.Fatal(err)
}
// The downurl endpoint is m115-encrypted end-to-end (the server side
// requires 115's private key), so stub the decrypted payload via the seam
// and assert the pickcode→URL extraction. The live crypto/transport path is
// exercised by integration testing against the real 115 API.
p115, ok := p.(*pan115Provider)
if !ok {
t.Fatalf("expected *pan115Provider, got %T", p)
}
p115.downURLPayload = func(ctx context.Context, pickcode string) ([]byte, error) {
if pickcode != "pick200" {
t.Errorf("bad pickcode %q", pickcode)
}
return []byte(`{"200":{"file_name":"Inception.mkv","file_size":"456","url":{"url":"https://cdn.115/x.mkv?t=1"}}}`), nil
}
entries, err := p.List(context.Background(), "")
if err != nil {
t.Fatalf("list: %v", err)
}
if len(entries) != 2 {
t.Fatalf("want 2 entries: %#v", entries)
}
if !entries[0].IsDir || entries[0].ID != "100" {
t.Fatalf("dir entry wrong: %#v", entries[0])
}
if entries[1].IsDir || entries[1].PickCode != "pick200" || entries[1].Size != 456 {
t.Fatalf("file entry wrong: %#v", entries[1])
}
link, err := p.Resolve(context.Background(), "pick200")
if err != nil {
t.Fatalf("resolve: %v", err)
}
if link.URL != "https://cdn.115/x.mkv?t=1" {
t.Fatalf("bad url: %s", link.URL)
}
if link.Proxy {
t.Fatalf("115 should default to 302 (no proxy)")
}
}
func Test115ListPaginates(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/files" {
t.Fatalf("unexpected path %s", r.URL.Path)
}
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
count := 100
if offset > 0 {
count = 1
}
items := make([]string, 0, count)
for i := 0; i < count; i++ {
n := offset + i
items = append(items, fmt.Sprintf(`{"fid":"%d","n":"Movie.%03d.mkv","s":%d,"pc":"pick%d"}`, n, n, n, n))
}
w.Write([]byte(`{"state":true,"data":[` + strings.Join(items, ",") + `]}`))
}))
defer srv.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1; CID=2", "base": srv.URL}, srv.Client())
if err != nil {
t.Fatal(err)
}
entries, err := p.List(context.Background(), "0")
if err != nil {
t.Fatalf("list: %v", err)
}
if len(entries) != 101 {
t.Fatalf("entries = %d, want 101", len(entries))
}
if entries[100].ID != "100" || entries[100].PickCode != "pick100" {
t.Fatalf("last entry wrong: %#v", entries[100])
}
}
// Test115DownURLEndpointAndError exercises the live fetchDownURLPayload path:
// it must POST an m115-encrypted `data` body to /app/chrome/downurl?t=... and
// surface 115's error when state=false (no decryption needed for that branch).
func Test115DownURLEndpointAndError(t *testing.T) {
var gotData, gotT string
pro := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/app/chrome/downurl" {
t.Errorf("unexpected path %s", r.URL.Path)
}
gotT = r.URL.Query().Get("t")
_ = r.ParseForm()
gotData = r.PostFormValue("data")
w.Write([]byte(`{"state":false,"error":"not exist"}`))
}))
defer pro.Close()
p, err := New(Type115, map[string]any{"cookie": "UID=1", "pro_base": pro.URL}, pro.Client())
if err != nil {
t.Fatal(err)
}
_, err = p.Resolve(context.Background(), "pickX")
if err == nil || !strings.Contains(err.Error(), "not exist") {
t.Fatalf("want upstream error surfaced, got %v", err)
}
if gotT == "" {
t.Errorf("missing t query param")
}
if gotData == "" {
t.Errorf("missing encrypted data body")
}
if _, derr := base64.StdEncoding.DecodeString(gotData); derr != nil {
t.Errorf("data body is not base64: %v", derr)
}
}
func Test115QRFlow(t *testing.T) {
// status sequence: waiting → scanned → confirmed
calls := 0
api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/1.0/web/1.0/token/":
w.Write([]byte(`{"state":1,"data":{"uid":"U1","time":1700,"sign":"S1"}}`))
case "/get/status/":
if r.URL.Query().Get("uid") != "U1" {
t.Errorf("bad uid %q", r.URL.Query().Get("uid"))
}
calls++
switch calls {
case 1:
w.Write([]byte(`{"state":1,"data":{"status":0}}`))
case 2:
w.Write([]byte(`{"state":1,"data":{"status":1}}`))
default:
w.Write([]byte(`{"state":1,"data":{"status":2}}`))
}
default:
t.Errorf("unexpected api path %s", r.URL.Path)
}
}))
defer api.Close()
passport := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/app/1.0/web/1.0/login/qrcode/" {
t.Errorf("unexpected passport path %s", r.URL.Path)
}
w.Write([]byte(`{"state":1,"data":{"cookie":{"UID":"u","CID":"c","SEID":"s"}}}`))
}))
defer passport.Close()
oldA, oldP := qr115APIBase, qr115PassportBase
qr115APIBase, qr115PassportBase = api.URL, passport.URL
defer func() { qr115APIBase, qr115PassportBase = oldA, oldP }()
ctx := context.Background()
sess, err := QRStart(ctx, api.Client())
if err != nil {
t.Fatalf("qr start: %v", err)
}
if sess.UID != "U1" || sess.QRImageURL == "" {
t.Fatalf("bad session: %#v", sess)
}
want := []string{"waiting", "scanned", "confirmed"}
for i, exp := range want {
st, err := QRPoll(ctx, api.Client(), sess)
if err != nil {
t.Fatalf("poll %d: %v", i, err)
}
if st.State != exp {
t.Fatalf("poll %d: want %s got %s", i, exp, st.State)
}
if exp == "confirmed" {
if st.Cookie == "" || !containsAll(st.Cookie, "UID=u", "SEID=s") {
t.Fatalf("confirmed must yield cookie: %q", st.Cookie)
}
}
}
}
func TestCloudDrive2WebDAVListAndResolve(t *testing.T) {
var gotAuth, gotDepth, gotRange string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@@ -523,6 +232,144 @@ func TestOpenListListUsesAPIUsernamePasswordWithoutWebDAVFallback(t *testing.T)
}
}
func TestOpenListMutableProviderUsesAPI(t *testing.T) {
var mkdirPath, renamePath, renameName, moveSrcDir, moveDstDir string
var moveNames []string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch r.URL.Path {
case "/api/fs/mkdir":
var body map[string]string
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode mkdir body: %v", err)
}
mkdirPath = body["path"]
if r.Header.Get("Authorization") != "alist-token" {
t.Fatalf("mkdir Authorization = %q", r.Header.Get("Authorization"))
}
_, _ = w.Write([]byte(`{"code":200,"message":"success"}`))
case "/api/fs/rename":
var body map[string]string
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode rename body: %v", err)
}
renamePath = body["path"]
renameName = body["name"]
_, _ = w.Write([]byte(`{"code":200,"message":"success"}`))
case "/api/fs/move":
var body struct {
SrcDir string `json:"src_dir"`
DstDir string `json:"dst_dir"`
Names []string `json:"names"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode move body: %v", err)
}
moveSrcDir = body.SrcDir
moveDstDir = body.DstDir
moveNames = body.Names
_, _ = w.Write([]byte(`{"code":200,"message":"success"}`))
default:
t.Fatalf("unexpected path %s", r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeOpenList, map[string]any{"server": srv.URL, "token": "alist-token"}, srv.Client())
if err != nil {
t.Fatal(err)
}
mutable, ok := p.(MutableProvider)
if !ok {
t.Fatal("openlist should support mutable provider")
}
created, err := mutable.Mkdir(context.Background(), "/电视剧", "欧美剧")
if err != nil {
t.Fatalf("mkdir: %v", err)
}
if mkdirPath != "/电视剧/欧美剧" || created.ID != "/电视剧/欧美剧" || !created.IsDir {
t.Fatalf("mkdir path=%q entry=%#v", mkdirPath, created)
}
renamed, err := mutable.Rename(context.Background(), "/电视剧/欧美剧", "美剧")
if err != nil {
t.Fatalf("rename: %v", err)
}
if renamePath != "/电视剧/欧美剧" || renameName != "美剧" || renamed.ID != "/电视剧/美剧" {
t.Fatalf("rename path=%q name=%q entry=%#v", renamePath, renameName, renamed)
}
moved, err := mutable.(MovableProvider).Move(context.Background(), "/待整理/Show.S01E01.mkv", "/动漫/国漫/Show/Season 01", "Show - S01E01.mkv")
if err != nil {
t.Fatalf("move: %v", err)
}
if moveSrcDir != "/待整理" || moveDstDir != "/动漫/国漫/Show/Season 01" || len(moveNames) != 1 || moveNames[0] != "Show.S01E01.mkv" {
t.Fatalf("move src=%q dst=%q names=%#v", moveSrcDir, moveDstDir, moveNames)
}
if renamePath != "/动漫/国漫/Show/Season 01/Show.S01E01.mkv" || renameName != "Show - S01E01.mkv" {
t.Fatalf("post-move rename path=%q name=%q", renamePath, renameName)
}
if moved.ID != "/动漫/国漫/Show/Season 01/Show - S01E01.mkv" {
t.Fatalf("moved entry = %#v", moved)
}
}
func TestCloudDrive2MutableProviderUsesWebDAV(t *testing.T) {
var mkcolSeen bool
var destinations []string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case r.Method == "MKCOL" && r.URL.Path == "/dav/TV":
mkcolSeen = true
w.WriteHeader(http.StatusCreated)
case r.Method == "MOVE" && r.URL.Path == "/dav/TV":
destinations = append(destinations, r.Header.Get("Destination"))
if r.Header.Get("Overwrite") != "F" {
t.Fatalf("Overwrite = %q, want F", r.Header.Get("Overwrite"))
}
w.WriteHeader(http.StatusCreated)
case r.Method == "MOVE" && r.URL.Path == "/dav/Inbox/Movie.mkv":
destinations = append(destinations, r.Header.Get("Destination"))
if r.Header.Get("Overwrite") != "F" {
t.Fatalf("Overwrite = %q, want F", r.Header.Get("Overwrite"))
}
w.WriteHeader(http.StatusCreated)
default:
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
}))
defer srv.Close()
p, err := New(TypeCloudDrive2, map[string]any{"url": srv.URL + "/dav", "username": "u", "password": "p"}, srv.Client())
if err != nil {
t.Fatal(err)
}
mutable, ok := p.(MutableProvider)
if !ok {
t.Fatal("clouddrive2 should support mutable provider")
}
if _, err := mutable.Mkdir(context.Background(), "", "TV"); err != nil {
t.Fatalf("mkdir: %v", err)
}
if _, err := mutable.Rename(context.Background(), "/TV", "电视剧"); err != nil {
t.Fatalf("rename: %v", err)
}
moved, err := mutable.(MovableProvider).Move(context.Background(), "/Inbox/Movie.mkv", "/电影/外语电影/Movie (2026)", "Movie (2026).mkv")
if err != nil {
t.Fatalf("move: %v", err)
}
if !mkcolSeen || len(destinations) != 2 {
t.Fatalf("mkcol=%v destinations=%#v, want mkdir and two MOVE calls", mkcolSeen, destinations)
}
if destinations[0] != srv.URL+"/dav/%E7%94%B5%E8%A7%86%E5%89%A7" {
t.Fatalf("rename Destination = %q", destinations[0])
}
if destinations[1] != srv.URL+"/dav/%E7%94%B5%E5%BD%B1/%E5%A4%96%E8%AF%AD%E7%94%B5%E5%BD%B1/Movie%20%282026%29/Movie%20%282026%29.mkv" {
t.Fatalf("move Destination = %q", destinations[1])
}
if moved.ID != "/电影/外语电影/Movie (2026)/Movie (2026).mkv" {
t.Fatalf("moved entry = %#v", moved)
}
}
func TestOpenListListAPIFailureDoesNotFallbackToWebDAV(t *testing.T) {
var davSeen bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@@ -804,22 +651,12 @@ func TestUnsupportedProvider(t *testing.T) {
if _, err := New("dropbox", nil, nil); err != ErrUnsupported {
t.Fatalf("want ErrUnsupported, got %v", err)
}
if _, err := New("quark", nil, nil); err != ErrUnsupported {
t.Fatalf("quark should be unsupported, got %v", err)
}
func containsAll(s string, subs ...string) bool {
for _, sub := range subs {
found := false
for i := 0; i+len(sub) <= len(s); i++ {
if s[i:i+len(sub)] == sub {
found = true
break
if IsCloudType("quark") {
t.Fatal("quark should not be an active cloud provider")
}
}
if !found {
return false
}
}
return true
}
var _ = time.Second
+2 -624
View File
@@ -1,25 +1,19 @@
package cloud
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"encoding/xml"
"fmt"
"io"
"net/http"
"net/url"
"path"
"sort"
"strconv"
"strings"
)
// cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint.
//
// CloudDrive2 already integrates many cloud disks (115 / 123 / Aliyun / Quark
// and more). Treating it as a WebDAV-backed cloud provider lets MediaStationGo
// CloudDrive2 integrates many cloud disks (115 / 123 / Aliyun and more).
// Treating it as a WebDAV-backed cloud provider lets MediaStationGo
// browse, mount and upload to those disks without carrying every provider's
// private chunk-upload protocol in this project.
type cloudDrive2Provider struct {
@@ -76,133 +70,6 @@ func (p *cloudDrive2Provider) Ping(ctx context.Context) error {
return err
}
func (p *cloudDrive2Provider) List(ctx context.Context, dir string) ([]FileEntry, error) {
if err := p.validate(); err != nil {
return nil, err
}
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
return p.listOpenListAPI(ctx, dir)
}
target := normalizeCloudDAVPath(dir)
req, err := http.NewRequestWithContext(ctx, "PROPFIND", p.urlFor(target), strings.NewReader(cloudDAVPropfindBody))
if err != nil {
return nil, err
}
p.auth(req)
req.Header.Set("Depth", "1")
req.Header.Set("Content-Type", "application/xml; charset=utf-8")
req.Header.Set("Accept", "application/xml,text/xml,*/*")
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.urlFor(target), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, p.decorateDAVStatusError(resp, target)
}
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
var multi cloudDAVMultiStatus
if err := xml.Unmarshal(body, &multi); err != nil {
return nil, fmt.Errorf("%s: decode webdav: %w", p.name, err)
}
basePath := strings.TrimRight(p.base.EscapedPath(), "/")
currentID := normalizeCloudDAVPath(target)
out := make([]FileEntry, 0, len(multi.Responses))
for _, item := range multi.Responses {
entryPath, err := p.entryIDFromHref(item.Href, basePath)
if err != nil || entryPath == "" || sameCloudDAVPath(entryPath, currentID) {
continue
}
name := firstNonEmpty(item.PropStat.Prop.DisplayName, path.Base(strings.TrimRight(entryPath, "/")))
if decoded, err := url.PathUnescape(name); err == nil {
name = decoded
}
if name == "" || name == "." || name == "/" {
continue
}
out = append(out, FileEntry{
ID: entryPath,
Name: name,
IsDir: item.PropStat.Prop.ResourceType.Collection != nil || strings.HasSuffix(item.Href, "/"),
Size: parseDAVSize(item.PropStat.Prop.ContentLength),
})
}
return out, nil
}
func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ([]FileEntry, error) {
token, err := p.openListAPIToken(ctx)
if err != nil {
return nil, err
}
const pageSize = 500
target := normalizeCloudDAVPath(dir)
out := make([]FileEntry, 0, pageSize)
for pageNum := 1; ; pageNum++ {
payload := map[string]any{
"path": target,
"password": "",
"page": pageNum,
"per_page": pageSize,
"refresh": false,
}
body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/fs/list"), bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
if token != "" {
req.Header.Set("Authorization", token)
}
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/list"), err)
}
var decoded openListListResponse
decodeErr := json.NewDecoder(io.LimitReader(resp.Body, 32<<20)).Decode(&decoded)
resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("%s: api list %s returned http %d", p.name, target, resp.StatusCode)
}
if decodeErr != nil {
return nil, fmt.Errorf("%s: decode api list: %w", p.name, decodeErr)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return nil, fmt.Errorf("%s: api list %s failed: %s", p.name, target, msg)
}
for _, item := range decoded.Data.Content {
name := strings.TrimSpace(item.Name)
if name == "" || name == "." || name == "/" {
continue
}
out = append(out, FileEntry{
ID: joinOpenListAPIPath(target, name),
Name: name,
IsDir: item.IsDir,
Size: item.Size,
})
}
total := decoded.Data.Total
if total > 0 {
if len(out) >= total || len(decoded.Data.Content) == 0 {
break
}
continue
}
if len(decoded.Data.Content) == 0 || len(decoded.Data.Content) < pageSize {
break
}
}
return out, nil
}
func (p *cloudDrive2Provider) Resolve(ctx context.Context, fileRef string) (*DirectLink, error) {
if err := p.validate(); err != nil {
return nil, err
@@ -239,323 +106,6 @@ func (p *cloudDrive2Provider) Resolve(ctx context.Context, fileRef string) (*Dir
return &DirectLink{URL: p.urlFor(ref), Headers: headers, Proxy: p.proxy}, nil
}
func (p *cloudDrive2Provider) resolveOpenListAPIDirect(ctx context.Context, fileRef string) (*DirectLink, error) {
token, err := p.openListAPIToken(ctx)
if err != nil {
return nil, err
}
payload, _ := json.Marshal(map[string]string{"path": normalizeCloudDAVPath(fileRef), "password": ""})
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/fs/get"), bytes.NewReader(payload))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
if token != "" {
req.Header.Set("Authorization", token)
}
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/get"), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("%s: api get %s returned http %d", p.name, fileRef, resp.StatusCode)
}
var decoded openListGetResponse
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&decoded); err != nil {
return nil, fmt.Errorf("%s: decode api get: %w", p.name, err)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return nil, fmt.Errorf("%s: api get %s failed: %s", p.name, fileRef, msg)
}
raw := firstNonEmpty(decoded.Data.RawURL, decoded.Data.URL)
if raw == "" {
return nil, fmt.Errorf("%s: api get %s returned empty raw_url", p.name, fileRef)
}
resolved, err := p.resolveOpenListPlaybackURL(raw)
if err != nil {
return nil, err
}
headers := normalizeOpenListPlaybackHeaders(decoded.Data.Header)
if len(headers) > 0 {
return nil, fmt.Errorf("%s: api get %s returned raw_url that requires headers (%s); refusing WebDAV/proxy fallback for pure 302 playback", p.name, fileRef, strings.Join(sortedHeaderNames(headers), ","))
}
resolved, err = p.resolveOpenListCDNRedirect(ctx, fileRef, resolved)
if err != nil {
return nil, err
}
return &DirectLink{URL: resolved, Headers: nil, Proxy: false}, nil
}
func (p *cloudDrive2Provider) resolveOpenListCDNRedirect(ctx context.Context, fileRef, rawURL string) (string, error) {
if p.apiBase == nil || !sameURLHost(rawURL, p.apiBase) {
return rawURL, nil
}
location, status, err := p.firstHTTPRedirectLocation(ctx, rawURL, nil)
if err != nil {
return "", fmt.Errorf("%s: probe raw_url %s failed: %w", p.name, fileRef, err)
}
if location != "" {
return location, nil
}
return "", fmt.Errorf("%s: api get %s returned an OpenList-hosted raw_url with http %d and no CDN Location; refusing OpenList/WebDAV proxy fallback for pure 302 playback", p.name, fileRef, status)
}
func (p *cloudDrive2Provider) resolveCloudDAVRedirectDirect(ctx context.Context, fileRef string) (*DirectLink, error) {
target := p.urlFor(fileRef)
headers := map[string]string{
"User-Agent": p.ua,
}
if p.token != "" {
headers["Authorization"] = p.token
} else if p.username != "" {
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(p.username+":"+p.password))
}
location, status, err := p.firstHTTPRedirectLocation(ctx, target, headers)
if err != nil {
return nil, decorateDAVTransportError(p.name, target, err)
}
if location == "" {
return nil, fmt.Errorf("%s: WebDAV %s returned http %d without CDN Location; refusing WebDAV/proxy fallback for pure 302 playback", p.name, fileRef, status)
}
return &DirectLink{URL: location, Headers: nil, Proxy: false}, nil
}
func (p *cloudDrive2Provider) firstHTTPRedirectLocation(ctx context.Context, target string, headers map[string]string) (string, int, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil)
if err != nil {
return "", 0, err
}
req.Header.Set("Accept", "*/*")
req.Header.Set("Accept-Encoding", "identity")
req.Header.Set("Range", "bytes=0-0")
if strings.TrimSpace(p.ua) != "" {
req.Header.Set("User-Agent", p.ua)
}
for key, value := range headers {
key = strings.TrimSpace(key)
if key != "" && strings.TrimSpace(value) != "" {
req.Header.Set(key, value)
}
}
client := p.client
if client == nil {
client = http.DefaultClient
}
noFollow := *client
noFollow.CheckRedirect = func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
}
resp, err := noFollow.Do(req)
if err != nil {
return "", 0, err
}
defer resp.Body.Close()
status := resp.StatusCode
if status >= 300 && status < 400 {
rawLocation := strings.TrimSpace(resp.Header.Get("Location"))
if rawLocation == "" {
return "", status, fmt.Errorf("%s: upstream returned redirect http %d without Location", p.name, status)
}
location, err := resolveHTTPRedirectLocation(target, rawLocation)
if err != nil {
return "", status, err
}
return location, status, nil
}
return "", status, nil
}
func sortedHeaderNames(headers map[string]string) []string {
if len(headers) == 0 {
return nil
}
out := make([]string, 0, len(headers))
for key := range headers {
key = strings.TrimSpace(key)
if key != "" {
out = append(out, key)
}
}
sort.Strings(out)
return out
}
func (p *cloudDrive2Provider) hasOpenListAPICredentials() bool {
return strings.TrimSpace(p.token) != "" || (strings.TrimSpace(p.username) != "" && p.password != "")
}
func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, error) {
if token := strings.TrimSpace(p.token); token != "" {
return token, nil
}
if strings.TrimSpace(p.username) == "" || p.password == "" {
return "", nil
}
payload, _ := json.Marshal(map[string]string{
"username": p.username,
"password": p.password,
})
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/auth/login"), bytes.NewReader(payload))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
resp, err := p.client.Do(req)
if err != nil {
return "", decorateDAVTransportError(p.name, p.openListAPIURL("/api/auth/login"), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return "", fmt.Errorf("%s: api login returned http %d", p.name, resp.StatusCode)
}
var decoded openListLoginResponse
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&decoded); err != nil {
return "", fmt.Errorf("%s: decode api login: %w", p.name, err)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return "", fmt.Errorf("%s: api login failed: %s", p.name, msg)
}
token := strings.TrimSpace(decoded.Data.Token)
if token == "" {
return "", fmt.Errorf("%s: api login returned empty token", p.name)
}
p.token = token
return token, nil
}
func (p *cloudDrive2Provider) resolveOpenListPlaybackURL(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", fmt.Errorf("%s: empty playback URL", p.name)
}
if strings.HasPrefix(raw, "//") {
if p.apiBase == nil || p.apiBase.Scheme == "" {
return "", fmt.Errorf("%s: protocol-relative playback URL without API base", p.name)
}
raw = p.apiBase.Scheme + ":" + raw
}
u, err := url.Parse(raw)
if err != nil {
return "", fmt.Errorf("%s: invalid playback URL: %w", p.name, err)
}
if u.IsAbs() {
if u.Scheme != "http" && u.Scheme != "https" {
return "", fmt.Errorf("%s: unsupported playback URL scheme %q", p.name, u.Scheme)
}
return u.String(), nil
}
if p.apiBase == nil {
return "", fmt.Errorf("%s: relative playback URL without API base", p.name)
}
base := *p.apiBase
base.RawPath = ""
base.RawQuery = ""
base.Fragment = ""
return base.ResolveReference(u).String(), nil
}
func sameURLHost(raw string, base *url.URL) bool {
if base == nil {
return false
}
u, err := url.Parse(strings.TrimSpace(raw))
if err != nil {
return false
}
if !u.IsAbs() {
return true
}
return strings.EqualFold(u.Host, base.Host)
}
func resolveHTTPRedirectLocation(baseURL, rawLocation string) (string, error) {
rawLocation = strings.TrimSpace(rawLocation)
if rawLocation == "" {
return "", fmt.Errorf("empty redirect Location")
}
if strings.HasPrefix(rawLocation, "//") {
base, err := url.Parse(baseURL)
if err != nil || base.Scheme == "" {
return "", fmt.Errorf("protocol-relative redirect Location without base scheme")
}
rawLocation = base.Scheme + ":" + rawLocation
}
location, err := url.Parse(rawLocation)
if err != nil {
return "", fmt.Errorf("invalid redirect Location: %w", err)
}
if location.IsAbs() {
if location.Scheme != "http" && location.Scheme != "https" {
return "", fmt.Errorf("unsupported redirect Location scheme %q", location.Scheme)
}
return location.String(), nil
}
base, err := url.Parse(baseURL)
if err != nil {
return "", fmt.Errorf("invalid redirect base URL: %w", err)
}
return base.ResolveReference(location).String(), nil
}
func normalizeOpenListPlaybackHeaders(raw json.RawMessage) map[string]string {
if len(raw) == 0 || string(raw) == "null" {
return nil
}
var obj map[string]any
if err := json.Unmarshal(raw, &obj); err != nil {
return nil
}
out := make(map[string]string, len(obj))
for k, v := range obj {
key := strings.TrimSpace(k)
if key == "" {
continue
}
switch value := v.(type) {
case string:
if strings.TrimSpace(value) != "" {
out[key] = strings.TrimSpace(value)
}
case []any:
parts := make([]string, 0, len(value))
for _, item := range value {
if s, ok := item.(string); ok && strings.TrimSpace(s) != "" {
parts = append(parts, strings.TrimSpace(s))
}
}
if len(parts) > 0 {
out[key] = strings.Join(parts, ", ")
}
}
}
if len(out) == 0 {
return nil
}
return out
}
func isCloudVideoPlaybackCandidate(fileRef string) bool {
switch strings.ToLower(path.Ext(strings.TrimSpace(fileRef))) {
case ".mkv", ".mp4", ".m4v", ".avi", ".mov", ".webm", ".ts", ".rmvb", ".rm", ".3gp", ".mpg", ".mpeg":
return true
default:
return false
}
}
func (p *cloudDrive2Provider) validate() error {
if p.base == nil || p.base.Scheme == "" || p.base.Host == "" {
return fmt.Errorf("%s: missing WebDAV URL", p.name)
@@ -563,17 +113,6 @@ func (p *cloudDrive2Provider) validate() error {
return nil
}
func (p *cloudDrive2Provider) auth(req *http.Request) {
req.Header.Set("User-Agent", p.ua)
if p.token != "" {
req.Header.Set("Authorization", p.token)
return
}
if p.username != "" {
req.SetBasicAuth(p.username, p.password)
}
}
func webDAVURLFromConfig(cfg map[string]any, defaultDAVPath string) string {
rawURL := str(cfg["url"])
if rawURL == "" {
@@ -669,153 +208,6 @@ func ensureDefaultDAVPath(rawURL, defaultDAVPath string) string {
return rawURL
}
func (p *cloudDrive2Provider) decorateDAVStatusError(resp *http.Response, target string) error {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
detail := compactDAVErrorBody(string(body))
if detail == "" {
return fmt.Errorf("%s: list %s returned http %d", p.name, target, resp.StatusCode)
}
if resp.StatusCode == http.StatusMethodNotAllowed {
return fmt.Errorf("%s: list %s returned http %d:%s;请确认填写的是 WebDAV 地址(通常以 /dav 结尾),并且桥接网盘已在 OpenList/CloudDrive2 内完成登录或 Cookie 保存", p.name, target, resp.StatusCode, detail)
}
if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden {
return fmt.Errorf("%s: list %s returned http %d:%s;请检查 WebDAV 用户名/密码、Authorization Token,或先在 OpenList/CloudDrive2 中保存对应网盘 Cookie", p.name, target, resp.StatusCode, detail)
}
return fmt.Errorf("%s: list %s returned http %d:%s", p.name, target, resp.StatusCode, detail)
}
func compactDAVErrorBody(raw string) string {
raw = strings.TrimSpace(strings.ReplaceAll(raw, "\x00", ""))
if raw == "" {
return ""
}
raw = strings.Join(strings.Fields(raw), " ")
if len([]rune(raw)) > 180 {
return string([]rune(raw)[:180]) + "…"
}
return raw
}
func decorateDAVTransportError(name, target string, err error) error {
if err == nil {
return nil
}
message := err.Error()
if strings.Contains(message, "server gave HTTP response to HTTPS client") {
return fmt.Errorf("%s: %w;当前地址使用 https://,但服务端返回 HTTP。请改用 http:// 地址,例如 OpenList 默认 WebDAV 通常是 http://host:5244/dav/;如果必须使用 https,请在 OpenList 前配置反向代理和证书", name, err)
}
if strings.Contains(message, "first record does not look like a TLS handshake") {
return fmt.Errorf("%s: %w;疑似把 HTTP 服务配置成了 https://,请检查 %s 的协议头", name, err, target)
}
return err
}
func (p *cloudDrive2Provider) urlFor(remotePath string) string {
u := *p.base
u.RawPath = ""
basePath := strings.TrimRight(u.Path, "/")
remote := strings.Trim(normalizeCloudDAVPath(remotePath), "/")
switch {
case basePath == "" || basePath == "/":
if remote == "" {
u.Path = "/"
} else {
u.Path = "/" + remote
}
case remote == "":
u.Path = basePath
default:
u.Path = basePath + "/" + remote
}
return u.String()
}
func (p *cloudDrive2Provider) entryIDFromHref(href, basePath string) (string, error) {
if href == "" {
return "", nil
}
parsed, err := url.Parse(href)
if err != nil {
return "", err
}
hrefPath := parsed.EscapedPath()
if hrefPath == "" {
hrefPath = href
}
if basePath != "" && basePath != "/" {
hrefPath = strings.TrimPrefix(hrefPath, basePath)
}
if decoded, err := url.PathUnescape(hrefPath); err == nil {
hrefPath = decoded
}
return normalizeCloudDAVPath(hrefPath), nil
}
const cloudDAVPropfindBody = `<?xml version="1.0" encoding="utf-8"?>
<d:propfind xmlns:d="DAV:">
<d:prop>
<d:displayname/>
<d:getcontentlength/>
<d:resourcetype/>
</d:prop>
</d:propfind>`
type cloudDAVMultiStatus struct {
Responses []cloudDAVResponse `xml:"response"`
}
type cloudDAVResponse struct {
Href string `xml:"href"`
PropStat cloudDAVPropStat `xml:"propstat"`
}
type cloudDAVPropStat struct {
Prop cloudDAVProp `xml:"prop"`
}
type cloudDAVProp struct {
DisplayName string `xml:"displayname"`
ContentLength string `xml:"getcontentlength"`
ResourceType cloudDAVResourceType `xml:"resourcetype"`
}
type cloudDAVResourceType struct {
Collection *struct{} `xml:"collection"`
}
type openListListResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
Content []openListListItem `json:"content"`
Total int `json:"total"`
} `json:"data"`
}
type openListListItem struct {
Name string `json:"name"`
Size int64 `json:"size"`
IsDir bool `json:"is_dir"`
}
type openListGetResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
RawURL string `json:"raw_url"`
URL string `json:"url"`
Header json.RawMessage `json:"header"`
} `json:"data"`
}
type openListLoginResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
Token string `json:"token"`
} `json:"data"`
}
func normalizeCloudDAVPath(p string) string {
p = strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
if p == "" || p == "." {
@@ -835,20 +227,6 @@ func sameCloudDAVPath(a, b string) bool {
return strings.TrimRight(normalizeCloudDAVPath(a), "/") == strings.TrimRight(normalizeCloudDAVPath(b), "/")
}
func joinOpenListAPIPath(dir, name string) string {
dir = strings.TrimRight(normalizeCloudDAVPath(dir), "/")
name = strings.Trim(strings.ReplaceAll(name, "\\", "/"), "/")
if dir == "" || dir == "/" {
return normalizeCloudDAVPath(name)
}
return normalizeCloudDAVPath(dir + "/" + name)
}
func parseDAVSize(raw string) int64 {
n, _ := strconv.ParseInt(strings.TrimSpace(raw), 10, 64)
return n
}
func firstNonEmpty(values ...string) string {
for _, v := range values {
if strings.TrimSpace(v) != "" {
+292
View File
@@ -0,0 +1,292 @@
package cloud
import (
"context"
"encoding/base64"
"encoding/xml"
"fmt"
"io"
"net/http"
"net/url"
"path"
"strconv"
"strings"
)
func (p *cloudDrive2Provider) List(ctx context.Context, dir string) ([]FileEntry, error) {
if err := p.validate(); err != nil {
return nil, err
}
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
return p.listOpenListAPI(ctx, dir)
}
target := normalizeCloudDAVPath(dir)
req, err := http.NewRequestWithContext(ctx, "PROPFIND", p.urlFor(target), strings.NewReader(cloudDAVPropfindBody))
if err != nil {
return nil, err
}
p.auth(req)
req.Header.Set("Depth", "1")
req.Header.Set("Content-Type", "application/xml; charset=utf-8")
req.Header.Set("Accept", "application/xml,text/xml,*/*")
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.urlFor(target), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, p.decorateDAVStatusError(resp, target)
}
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
var multi cloudDAVMultiStatus
if err := xml.Unmarshal(body, &multi); err != nil {
return nil, fmt.Errorf("%s: decode webdav: %w", p.name, err)
}
basePath := strings.TrimRight(p.base.EscapedPath(), "/")
currentID := normalizeCloudDAVPath(target)
out := make([]FileEntry, 0, len(multi.Responses))
for _, item := range multi.Responses {
entryPath, err := p.entryIDFromHref(item.Href, basePath)
if err != nil || entryPath == "" || sameCloudDAVPath(entryPath, currentID) {
continue
}
name := firstNonEmpty(item.PropStat.Prop.DisplayName, path.Base(strings.TrimRight(entryPath, "/")))
if decoded, err := url.PathUnescape(name); err == nil {
name = decoded
}
if name == "" || name == "." || name == "/" {
continue
}
out = append(out, FileEntry{
ID: entryPath,
Name: name,
IsDir: item.PropStat.Prop.ResourceType.Collection != nil || strings.HasSuffix(item.Href, "/"),
Size: parseDAVSize(item.PropStat.Prop.ContentLength),
})
}
return out, nil
}
func (p *cloudDrive2Provider) resolveCloudDAVRedirectDirect(ctx context.Context, fileRef string) (*DirectLink, error) {
target := p.urlFor(fileRef)
headers := map[string]string{
"User-Agent": p.ua,
}
if p.token != "" {
headers["Authorization"] = p.token
} else if p.username != "" {
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(p.username+":"+p.password))
}
location, status, err := p.firstHTTPRedirectLocation(ctx, target, headers)
if err != nil {
return nil, decorateDAVTransportError(p.name, target, err)
}
if location == "" {
return nil, fmt.Errorf("%s: WebDAV %s returned http %d without CDN Location; refusing WebDAV/proxy fallback for pure 302 playback", p.name, fileRef, status)
}
return &DirectLink{URL: location, Headers: nil, Proxy: false}, nil
}
func (p *cloudDrive2Provider) firstHTTPRedirectLocation(ctx context.Context, target string, headers map[string]string) (string, int, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil)
if err != nil {
return "", 0, err
}
req.Header.Set("Accept", "*/*")
req.Header.Set("Accept-Encoding", "identity")
req.Header.Set("Range", "bytes=0-0")
if strings.TrimSpace(p.ua) != "" {
req.Header.Set("User-Agent", p.ua)
}
for key, value := range headers {
key = strings.TrimSpace(key)
if key != "" && strings.TrimSpace(value) != "" {
req.Header.Set(key, value)
}
}
client := p.client
if client == nil {
client = http.DefaultClient
}
noFollow := *client
noFollow.CheckRedirect = func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
}
resp, err := noFollow.Do(req)
if err != nil {
return "", 0, err
}
defer resp.Body.Close()
status := resp.StatusCode
if status >= 300 && status < 400 {
rawLocation := strings.TrimSpace(resp.Header.Get("Location"))
if rawLocation == "" {
return "", status, fmt.Errorf("%s: upstream returned redirect http %d without Location", p.name, status)
}
location, err := resolveHTTPRedirectLocation(target, rawLocation)
if err != nil {
return "", status, err
}
return location, status, nil
}
return "", status, nil
}
func resolveHTTPRedirectLocation(baseURL, rawLocation string) (string, error) {
rawLocation = strings.TrimSpace(rawLocation)
if rawLocation == "" {
return "", fmt.Errorf("empty redirect Location")
}
if strings.HasPrefix(rawLocation, "//") {
base, err := url.Parse(baseURL)
if err != nil || base.Scheme == "" {
return "", fmt.Errorf("protocol-relative redirect Location without base scheme")
}
rawLocation = base.Scheme + ":" + rawLocation
}
location, err := url.Parse(rawLocation)
if err != nil {
return "", fmt.Errorf("invalid redirect Location: %w", err)
}
if location.IsAbs() {
if location.Scheme != "http" && location.Scheme != "https" {
return "", fmt.Errorf("unsupported redirect Location scheme %q", location.Scheme)
}
return location.String(), nil
}
base, err := url.Parse(baseURL)
if err != nil {
return "", fmt.Errorf("invalid redirect base URL: %w", err)
}
return base.ResolveReference(location).String(), nil
}
func (p *cloudDrive2Provider) auth(req *http.Request) {
req.Header.Set("User-Agent", p.ua)
if p.token != "" {
req.Header.Set("Authorization", p.token)
return
}
if p.username != "" {
req.SetBasicAuth(p.username, p.password)
}
}
func (p *cloudDrive2Provider) decorateDAVStatusError(resp *http.Response, target string) error {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
detail := compactDAVErrorBody(string(body))
if detail == "" {
return fmt.Errorf("%s: list %s returned http %d", p.name, target, resp.StatusCode)
}
if resp.StatusCode == http.StatusMethodNotAllowed {
return fmt.Errorf("%s: list %s returned http %d:%s;请确认填写的是 WebDAV 地址(通常以 /dav 结尾),并且桥接网盘已在 OpenList/CloudDrive2 内完成登录或 Cookie 保存", p.name, target, resp.StatusCode, detail)
}
if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden {
return fmt.Errorf("%s: list %s returned http %d:%s;请检查 WebDAV 用户名/密码、Authorization Token,或先在 OpenList/CloudDrive2 中保存对应网盘 Cookie", p.name, target, resp.StatusCode, detail)
}
return fmt.Errorf("%s: list %s returned http %d:%s", p.name, target, resp.StatusCode, detail)
}
func compactDAVErrorBody(raw string) string {
raw = strings.TrimSpace(strings.ReplaceAll(raw, "\x00", ""))
if raw == "" {
return ""
}
raw = strings.Join(strings.Fields(raw), " ")
if len([]rune(raw)) > 180 {
return string([]rune(raw)[:180]) + "…"
}
return raw
}
func decorateDAVTransportError(name, target string, err error) error {
if err == nil {
return nil
}
message := err.Error()
if strings.Contains(message, "server gave HTTP response to HTTPS client") {
return fmt.Errorf("%s: %w;当前地址使用 https://,但服务端返回 HTTP。请改用 http:// 地址,例如 OpenList 默认 WebDAV 通常是 http://host:5244/dav/;如果必须使用 https,请在 OpenList 前配置反向代理和证书", name, err)
}
if strings.Contains(message, "first record does not look like a TLS handshake") {
return fmt.Errorf("%s: %w;疑似把 HTTP 服务配置成了 https://,请检查 %s 的协议头", name, err, target)
}
return err
}
func (p *cloudDrive2Provider) urlFor(remotePath string) string {
u := *p.base
u.RawPath = ""
basePath := strings.TrimRight(u.Path, "/")
remote := strings.Trim(normalizeCloudDAVPath(remotePath), "/")
switch {
case basePath == "" || basePath == "/":
if remote == "" {
u.Path = "/"
} else {
u.Path = "/" + remote
}
case remote == "":
u.Path = basePath
default:
u.Path = basePath + "/" + remote
}
return u.String()
}
func (p *cloudDrive2Provider) entryIDFromHref(href, basePath string) (string, error) {
if href == "" {
return "", nil
}
parsed, err := url.Parse(href)
if err != nil {
return "", err
}
hrefPath := parsed.EscapedPath()
if hrefPath == "" {
hrefPath = href
}
if basePath != "" && basePath != "/" {
hrefPath = strings.TrimPrefix(hrefPath, basePath)
}
if decoded, err := url.PathUnescape(hrefPath); err == nil {
hrefPath = decoded
}
return normalizeCloudDAVPath(hrefPath), nil
}
const cloudDAVPropfindBody = `<?xml version="1.0" encoding="utf-8"?>
<d:propfind xmlns:d="DAV:">
<d:prop>
<d:displayname/>
<d:getcontentlength/>
<d:resourcetype/>
</d:prop>
</d:propfind>`
type cloudDAVMultiStatus struct {
Responses []cloudDAVResponse `xml:"response"`
}
type cloudDAVResponse struct {
Href string `xml:"href"`
PropStat cloudDAVPropStat `xml:"propstat"`
}
type cloudDAVPropStat struct {
Prop cloudDAVProp `xml:"prop"`
}
type cloudDAVProp struct {
DisplayName string `xml:"displayname"`
ContentLength string `xml:"getcontentlength"`
ResourceType cloudDAVResourceType `xml:"resourcetype"`
}
type cloudDAVResourceType struct {
Collection *struct{} `xml:"collection"`
}
func parseDAVSize(raw string) int64 {
n, _ := strconv.ParseInt(strings.TrimSpace(raw), 10, 64)
return n
}
@@ -0,0 +1,233 @@
package cloud
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"path"
"strings"
)
func (p *cloudDrive2Provider) Mkdir(ctx context.Context, parentDir, name string) (*FileEntry, error) {
cleanName, err := cleanCloudEntryName(name)
if err != nil {
return nil, err
}
parent := normalizeCloudDAVPath(parentDir)
target := joinOpenListAPIPath(parent, cleanName)
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
if err := p.openListAPIMkdir(ctx, target); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName, IsDir: true}, nil
}
if err := p.webDAVMkdir(ctx, target); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName, IsDir: true}, nil
}
func (p *cloudDrive2Provider) Rename(ctx context.Context, ref, name string) (*FileEntry, error) {
cleanName, err := cleanCloudEntryName(name)
if err != nil {
return nil, err
}
source := normalizeCloudDAVPath(ref)
if source == "/" {
return nil, fmt.Errorf("%s: cannot rename root directory", p.name)
}
target := joinOpenListAPIPath(path.Dir(source), cleanName)
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
if err := p.openListAPIRename(ctx, source, cleanName); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName, IsDir: true}, nil
}
if err := p.webDAVRename(ctx, source, target); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName, IsDir: true}, nil
}
func (p *cloudDrive2Provider) Move(ctx context.Context, ref, targetDir, name string) (*FileEntry, error) {
source := normalizeCloudDAVPath(ref)
if source == "/" {
return nil, fmt.Errorf("%s: cannot move root directory", p.name)
}
cleanName := strings.TrimSpace(name)
if cleanName == "" {
cleanName = path.Base(source)
}
var err error
cleanName, err = cleanCloudEntryName(cleanName)
if err != nil {
return nil, err
}
targetDir = normalizeCloudDAVPath(targetDir)
target := joinOpenListAPIPath(targetDir, cleanName)
if sameCloudDAVPath(source, target) {
return &FileEntry{ID: target, Name: cleanName}, nil
}
if p.typ == TypeOpenList && p.apiBase != nil && p.hasOpenListAPICredentials() {
if err := p.openListAPIMove(ctx, source, targetDir, cleanName); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName}, nil
}
if err := p.webDAVRename(ctx, source, target); err != nil {
return nil, err
}
return &FileEntry{ID: target, Name: cleanName}, nil
}
func cleanCloudEntryName(name string) (string, error) {
name = strings.TrimSpace(name)
if name == "" || name == "." || name == ".." {
return "", fmt.Errorf("entry name is required")
}
if strings.ContainsAny(name, `/\`) {
return "", fmt.Errorf("entry name cannot contain path separators")
}
return name, nil
}
func (p *cloudDrive2Provider) openListAPIMkdir(ctx context.Context, target string) error {
return p.openListAPIPost(ctx, "/api/fs/mkdir", map[string]string{"path": normalizeCloudDAVPath(target)}, "mkdir")
}
func (p *cloudDrive2Provider) openListAPIRename(ctx context.Context, source, name string) error {
return p.openListAPIPost(ctx, "/api/fs/rename", map[string]string{
"path": normalizeCloudDAVPath(source),
"name": name,
}, "rename")
}
func (p *cloudDrive2Provider) openListAPIMove(ctx context.Context, source, targetDir, targetName string) error {
targetDir = normalizeCloudDAVPath(targetDir)
sourceName := path.Base(normalizeCloudDAVPath(source))
if sameCloudDAVPath(path.Dir(source), targetDir) {
if sourceName == targetName {
return nil
}
return p.openListAPIRename(ctx, source, targetName)
}
if err := p.openListAPIPost(ctx, "/api/fs/move", map[string]any{
"src_dir": normalizeCloudDAVPath(path.Dir(source)),
"dst_dir": targetDir,
"names": []string{sourceName},
}, "move"); err != nil {
return err
}
if sourceName != targetName {
moved := joinOpenListAPIPath(targetDir, sourceName)
return p.openListAPIRename(ctx, moved, targetName)
}
return nil
}
func (p *cloudDrive2Provider) openListAPIPost(ctx context.Context, apiPath string, payload any, action string) error {
token, err := p.openListAPIToken(ctx)
if err != nil {
return err
}
body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL(apiPath), bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
if token != "" {
req.Header.Set("Authorization", token)
}
resp, err := p.client.Do(req)
if err != nil {
return decorateDAVTransportError(p.name, p.openListAPIURL(apiPath), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return fmt.Errorf("%s: api %s returned http %d", p.name, action, resp.StatusCode)
}
var decoded struct {
Code int `json:"code"`
Message string `json:"message"`
}
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&decoded); err != nil {
return fmt.Errorf("%s: decode api %s: %w", p.name, action, err)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return fmt.Errorf("%s: api %s failed: %s", p.name, action, msg)
}
return nil
}
func (p *cloudDrive2Provider) webDAVMkdir(ctx context.Context, target string) error {
req, err := http.NewRequestWithContext(ctx, "MKCOL", p.urlFor(target), nil)
if err != nil {
return err
}
p.auth(req)
resp, err := p.client.Do(req)
if err != nil {
return decorateDAVTransportError(p.name, p.urlFor(target), err)
}
defer resp.Body.Close()
switch resp.StatusCode {
case http.StatusCreated, http.StatusOK, http.StatusNoContent:
return nil
case http.StatusMethodNotAllowed:
return fmt.Errorf("%s: mkdir %s returned http %d; the folder may already exist or this WebDAV backend is read-only", p.name, target, resp.StatusCode)
default:
return p.decorateDAVMutationStatusError(resp, "mkdir", target)
}
}
func (p *cloudDrive2Provider) webDAVRename(ctx context.Context, source, target string) error {
req, err := http.NewRequestWithContext(ctx, "MOVE", p.urlFor(source), nil)
if err != nil {
return err
}
p.auth(req)
req.Header.Set("Destination", p.webDAVDestination(target))
req.Header.Set("Overwrite", "F")
resp, err := p.client.Do(req)
if err != nil {
return decorateDAVTransportError(p.name, p.urlFor(source), err)
}
defer resp.Body.Close()
switch resp.StatusCode {
case http.StatusCreated, http.StatusOK, http.StatusNoContent:
return nil
default:
return p.decorateDAVMutationStatusError(resp, "rename", source)
}
}
func (p *cloudDrive2Provider) webDAVDestination(target string) string {
raw := p.urlFor(target)
u, err := url.Parse(raw)
if err != nil {
return raw
}
u.RawQuery = ""
u.Fragment = ""
return u.String()
}
func (p *cloudDrive2Provider) decorateDAVMutationStatusError(resp *http.Response, action, target string) error {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
detail := compactDAVErrorBody(string(body))
if detail == "" {
return fmt.Errorf("%s: %s %s returned http %d", p.name, action, target, resp.StatusCode)
}
return fmt.Errorf("%s: %s %s returned http %d:%s", p.name, action, target, resp.StatusCode, detail)
}
@@ -0,0 +1,352 @@
package cloud
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"path"
"sort"
"strings"
)
func (p *cloudDrive2Provider) listOpenListAPI(ctx context.Context, dir string) ([]FileEntry, error) {
token, err := p.openListAPIToken(ctx)
if err != nil {
return nil, err
}
const pageSize = 500
target := normalizeCloudDAVPath(dir)
out := make([]FileEntry, 0, pageSize)
for pageNum := 1; ; pageNum++ {
payload := map[string]any{
"path": target,
"password": "",
"page": pageNum,
"per_page": pageSize,
"refresh": false,
}
body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/fs/list"), bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
if token != "" {
req.Header.Set("Authorization", token)
}
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/list"), err)
}
var decoded openListListResponse
decodeErr := json.NewDecoder(io.LimitReader(resp.Body, 32<<20)).Decode(&decoded)
resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("%s: api list %s returned http %d", p.name, target, resp.StatusCode)
}
if decodeErr != nil {
return nil, fmt.Errorf("%s: decode api list: %w", p.name, decodeErr)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return nil, fmt.Errorf("%s: api list %s failed: %s", p.name, target, msg)
}
for _, item := range decoded.Data.Content {
name := strings.TrimSpace(item.Name)
if name == "" || name == "." || name == "/" {
continue
}
out = append(out, FileEntry{
ID: joinOpenListAPIPath(target, name),
Name: name,
IsDir: item.IsDir,
Size: item.Size,
})
}
total := decoded.Data.Total
if total > 0 {
if len(out) >= total || len(decoded.Data.Content) == 0 {
break
}
continue
}
if len(decoded.Data.Content) == 0 || len(decoded.Data.Content) < pageSize {
break
}
}
return out, nil
}
func (p *cloudDrive2Provider) resolveOpenListAPIDirect(ctx context.Context, fileRef string) (*DirectLink, error) {
token, err := p.openListAPIToken(ctx)
if err != nil {
return nil, err
}
payload, _ := json.Marshal(map[string]string{"path": normalizeCloudDAVPath(fileRef), "password": ""})
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/fs/get"), bytes.NewReader(payload))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
if token != "" {
req.Header.Set("Authorization", token)
}
resp, err := p.client.Do(req)
if err != nil {
return nil, decorateDAVTransportError(p.name, p.openListAPIURL("/api/fs/get"), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return nil, fmt.Errorf("%s: api get %s returned http %d", p.name, fileRef, resp.StatusCode)
}
var decoded openListGetResponse
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&decoded); err != nil {
return nil, fmt.Errorf("%s: decode api get: %w", p.name, err)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return nil, fmt.Errorf("%s: api get %s failed: %s", p.name, fileRef, msg)
}
raw := firstNonEmpty(decoded.Data.RawURL, decoded.Data.URL)
if raw == "" {
return nil, fmt.Errorf("%s: api get %s returned empty raw_url", p.name, fileRef)
}
resolved, err := p.resolveOpenListPlaybackURL(raw)
if err != nil {
return nil, err
}
headers := normalizeOpenListPlaybackHeaders(decoded.Data.Header)
if len(headers) > 0 {
return nil, fmt.Errorf("%s: api get %s returned raw_url that requires headers (%s); refusing WebDAV/proxy fallback for pure 302 playback", p.name, fileRef, strings.Join(sortedHeaderNames(headers), ","))
}
resolved, err = p.resolveOpenListCDNRedirect(ctx, fileRef, resolved)
if err != nil {
return nil, err
}
return &DirectLink{URL: resolved, Headers: nil, Proxy: false}, nil
}
func (p *cloudDrive2Provider) resolveOpenListCDNRedirect(ctx context.Context, fileRef, rawURL string) (string, error) {
if p.apiBase == nil || !sameURLHost(rawURL, p.apiBase) {
return rawURL, nil
}
location, status, err := p.firstHTTPRedirectLocation(ctx, rawURL, nil)
if err != nil {
return "", fmt.Errorf("%s: probe raw_url %s failed: %w", p.name, fileRef, err)
}
if location != "" {
return location, nil
}
return "", fmt.Errorf("%s: api get %s returned an OpenList-hosted raw_url with http %d and no CDN Location; refusing OpenList/WebDAV proxy fallback for pure 302 playback", p.name, fileRef, status)
}
func sortedHeaderNames(headers map[string]string) []string {
if len(headers) == 0 {
return nil
}
out := make([]string, 0, len(headers))
for key := range headers {
key = strings.TrimSpace(key)
if key != "" {
out = append(out, key)
}
}
sort.Strings(out)
return out
}
func (p *cloudDrive2Provider) hasOpenListAPICredentials() bool {
return strings.TrimSpace(p.token) != "" || (strings.TrimSpace(p.username) != "" && p.password != "")
}
func (p *cloudDrive2Provider) openListAPIToken(ctx context.Context) (string, error) {
if token := strings.TrimSpace(p.token); token != "" {
return token, nil
}
if strings.TrimSpace(p.username) == "" || p.password == "" {
return "", nil
}
payload, _ := json.Marshal(map[string]string{
"username": p.username,
"password": p.password,
})
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.openListAPIURL("/api/auth/login"), bytes.NewReader(payload))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", p.ua)
resp, err := p.client.Do(req)
if err != nil {
return "", decorateDAVTransportError(p.name, p.openListAPIURL("/api/auth/login"), err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return "", fmt.Errorf("%s: api login returned http %d", p.name, resp.StatusCode)
}
var decoded openListLoginResponse
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&decoded); err != nil {
return "", fmt.Errorf("%s: decode api login: %w", p.name, err)
}
if decoded.Code != 0 && decoded.Code != 200 {
msg := strings.TrimSpace(decoded.Message)
if msg == "" {
msg = fmt.Sprintf("code %d", decoded.Code)
}
return "", fmt.Errorf("%s: api login failed: %s", p.name, msg)
}
token := strings.TrimSpace(decoded.Data.Token)
if token == "" {
return "", fmt.Errorf("%s: api login returned empty token", p.name)
}
p.token = token
return token, nil
}
func (p *cloudDrive2Provider) resolveOpenListPlaybackURL(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", fmt.Errorf("%s: empty playback URL", p.name)
}
if strings.HasPrefix(raw, "//") {
if p.apiBase == nil || p.apiBase.Scheme == "" {
return "", fmt.Errorf("%s: protocol-relative playback URL without API base", p.name)
}
raw = p.apiBase.Scheme + ":" + raw
}
u, err := url.Parse(raw)
if err != nil {
return "", fmt.Errorf("%s: invalid playback URL: %w", p.name, err)
}
if u.IsAbs() {
if u.Scheme != "http" && u.Scheme != "https" {
return "", fmt.Errorf("%s: unsupported playback URL scheme %q", p.name, u.Scheme)
}
return u.String(), nil
}
if p.apiBase == nil {
return "", fmt.Errorf("%s: relative playback URL without API base", p.name)
}
base := *p.apiBase
base.RawPath = ""
base.RawQuery = ""
base.Fragment = ""
return base.ResolveReference(u).String(), nil
}
func sameURLHost(raw string, base *url.URL) bool {
if base == nil {
return false
}
u, err := url.Parse(strings.TrimSpace(raw))
if err != nil {
return false
}
if !u.IsAbs() {
return true
}
return strings.EqualFold(u.Host, base.Host)
}
func normalizeOpenListPlaybackHeaders(raw json.RawMessage) map[string]string {
if len(raw) == 0 || string(raw) == "null" {
return nil
}
var obj map[string]any
if err := json.Unmarshal(raw, &obj); err != nil {
return nil
}
out := make(map[string]string, len(obj))
for k, v := range obj {
key := strings.TrimSpace(k)
if key == "" {
continue
}
switch value := v.(type) {
case string:
if strings.TrimSpace(value) != "" {
out[key] = strings.TrimSpace(value)
}
case []any:
parts := make([]string, 0, len(value))
for _, item := range value {
if s, ok := item.(string); ok && strings.TrimSpace(s) != "" {
parts = append(parts, strings.TrimSpace(s))
}
}
if len(parts) > 0 {
out[key] = strings.Join(parts, ", ")
}
}
}
if len(out) == 0 {
return nil
}
return out
}
func isCloudVideoPlaybackCandidate(fileRef string) bool {
switch strings.ToLower(path.Ext(strings.TrimSpace(fileRef))) {
case ".mkv", ".mp4", ".m4v", ".avi", ".mov", ".webm", ".ts", ".rmvb", ".rm", ".3gp", ".mpg", ".mpeg":
return true
default:
return false
}
}
type openListListResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
Content []openListListItem `json:"content"`
Total int `json:"total"`
} `json:"data"`
}
type openListListItem struct {
Name string `json:"name"`
Size int64 `json:"size"`
IsDir bool `json:"is_dir"`
}
type openListGetResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
RawURL string `json:"raw_url"`
URL string `json:"url"`
Header json.RawMessage `json:"header"`
} `json:"data"`
}
type openListLoginResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
Token string `json:"token"`
} `json:"data"`
}
func joinOpenListAPIPath(dir, name string) string {
dir = strings.TrimRight(normalizeCloudDAVPath(dir), "/")
name = strings.Trim(strings.ReplaceAll(name, "\\", "/"), "/")
if dir == "" || dir == "/" {
return normalizeCloudDAVPath(name)
}
return normalizeCloudDAVPath(dir + "/" + name)
}

Some files were not shown because too many files have changed in this diff Show More