mirror of
https://github.com/truewhile/MeBox.git
synced 2026-09-28 03:06:38 +08:00
refactor: split modules and harden scraping workflows
This commit is contained in:
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
Vendored
+28
@@ -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
@@ -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...)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -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, `"`, `""`) + `"`
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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 ""
|
||||
}
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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})
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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"])
|
||||
}
|
||||
}
|
||||
@@ -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"},
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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,
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -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,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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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"})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()})
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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"})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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": "网盘目标目录"},
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 ""
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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。
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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) != "" {
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user