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
|
Dockerfile text eol=lf
|
||||||
|
|
||||||
*.ps1 text eol=crlf
|
*.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 \
|
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
|
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
|
# Tiny entrypoint that lets us run as a NAS host UID/GID via PUID/PGID without
|
||||||
# (handy on NAS deployments where bind-mounted volumes belong to a non-root
|
# rewriting /etc/passwd or /etc/group on every container start.
|
||||||
# user). When PUID == 0 we skip su-exec entirely and run as root.
|
COPY docker-entrypoint.sh /entrypoint.sh
|
||||||
RUN printf '#!/bin/sh\n\
|
RUN chmod +x /entrypoint.sh
|
||||||
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
|
|
||||||
|
|
||||||
CMD ["/entrypoint.sh"]
|
CMD ["/entrypoint.sh"]
|
||||||
|
|||||||
+5
-5
@@ -218,14 +218,14 @@ func serveSPA(r *gin.Engine, webDir string) {
|
|||||||
assets.Static("/", filepath.Join(webDir, "assets"))
|
assets.Static("/", filepath.Join(webDir, "assets"))
|
||||||
brand := r.Group("/brand")
|
brand := r.Group("/brand")
|
||||||
brand.Use(func(c *gin.Context) {
|
brand.Use(func(c *gin.Context) {
|
||||||
c.Header("Cache-Control", "public, max-age=86400")
|
setNoCacheHeaders(c)
|
||||||
c.Next()
|
c.Next()
|
||||||
})
|
})
|
||||||
brand.Static("/", filepath.Join(webDir, "brand"))
|
brand.Static("/", filepath.Join(webDir, "brand"))
|
||||||
for _, icon := range []string{"/favicon.ico", "/favicon.svg"} {
|
for _, rootFile := range []string{"/favicon.ico", "/favicon.svg", "/artwork-cache-sw.js"} {
|
||||||
iconPath := filepath.Join(webDir, strings.TrimPrefix(icon, "/"))
|
filePath := filepath.Join(webDir, strings.TrimPrefix(rootFile, "/"))
|
||||||
r.GET(icon, serveNoCacheFile(iconPath))
|
r.GET(rootFile, serveNoCacheFile(filePath))
|
||||||
r.HEAD(icon, serveNoCacheFile(iconPath))
|
r.HEAD(rootFile, serveNoCacheFile(filePath))
|
||||||
}
|
}
|
||||||
r.NoRoute(func(c *gin.Context) {
|
r.NoRoute(func(c *gin.Context) {
|
||||||
path := c.Request.URL.Path
|
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 {
|
if err := os.WriteFile(filepath.Join(webDir, "brand", "mediastationgo-logo.svg"), []byte("<svg></svg>"), 0o644); err != nil {
|
||||||
t.Fatal(err)
|
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()
|
router := gin.New()
|
||||||
serveSPA(router, webDir)
|
serveSPA(router, webDir)
|
||||||
@@ -84,10 +87,26 @@ func TestServeSPAServesAssetsImmutableAndBypassesAPIRoutes(t *testing.T) {
|
|||||||
if brandResp.Code != http.StatusOK {
|
if brandResp.Code != http.StatusOK {
|
||||||
t.Fatalf("brand asset status = %d, want 200", brandResp.Code)
|
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") {
|
if strings.Contains(brandResp.Body.String(), "index") {
|
||||||
t.Fatalf("brand asset should not serve SPA index: %q", brandResp.Body.String())
|
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{
|
for _, path := range []string{
|
||||||
"/api/missing",
|
"/api/missing",
|
||||||
"/emby",
|
"/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
|
package database
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"reflect"
|
|
||||||
"sort"
|
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
"gorm.io/driver/postgres"
|
"gorm.io/driver/postgres"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
"gorm.io/gorm/clause"
|
|
||||||
"gorm.io/gorm/logger"
|
"gorm.io/gorm/logger"
|
||||||
"gorm.io/gorm/schema"
|
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/config"
|
"github.com/ShukeBta/MediaStationGo/internal/config"
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Open initialises the configured GORM database. database.type=auto chooses
|
// Open initialises the configured GORM database. database.type=auto chooses
|
||||||
// PostgreSQL when database.dsn is present (the Docker Compose default) and
|
// PostgreSQL when database.dsn is present and otherwise falls back to SQLite.
|
||||||
// otherwise falls back to SQLite for old/bare-metal installs.
|
|
||||||
func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
|
func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
|
||||||
gormLogger := logger.New(
|
if cfg == nil {
|
||||||
zapStdLogger{log: log},
|
return nil, errors.New("database config is required")
|
||||||
logger.Config{
|
}
|
||||||
SlowThreshold: 0,
|
|
||||||
LogLevel: logger.Warn,
|
|
||||||
IgnoreRecordNotFoundError: true,
|
|
||||||
Colorful: false,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
dialect := normalizeDatabaseType(cfg.Database.Type)
|
dialect := normalizeDatabaseType(cfg.Database.Type)
|
||||||
if dialect == "auto" {
|
if dialect == "auto" {
|
||||||
dialect = effectiveAutoDatabaseType(cfg)
|
dialect = effectiveAutoDatabaseType(cfg)
|
||||||
@@ -48,7 +31,7 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
db, err := gorm.Open(dialector, &gorm.Config{
|
db, err := gorm.Open(dialector, &gorm.Config{
|
||||||
Logger: gormLogger,
|
Logger: newGormLogger(log),
|
||||||
PrepareStmt: true,
|
PrepareStmt: true,
|
||||||
DisableForeignKeyConstraintWhenMigrating: false,
|
DisableForeignKeyConstraintWhenMigrating: false,
|
||||||
})
|
})
|
||||||
@@ -58,9 +41,31 @@ func Open(cfg *config.Config, log *zap.Logger) (*gorm.DB, error) {
|
|||||||
if dialect == "sqlite" {
|
if dialect == "sqlite" {
|
||||||
installSQLiteWriteGate(db)
|
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()
|
sqlDB, err := db.DB()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("gorm sqldb: %w", err)
|
return fmt.Errorf("gorm sqldb: %w", err)
|
||||||
}
|
}
|
||||||
if cfg.Database.MaxOpenConns > 0 {
|
if cfg.Database.MaxOpenConns > 0 {
|
||||||
sqlDB.SetMaxOpenConns(cfg.Database.MaxOpenConns)
|
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 {
|
if cfg.Database.MaxIdleConns > 0 {
|
||||||
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
|
sqlDB.SetMaxIdleConns(cfg.Database.MaxIdleConns)
|
||||||
}
|
}
|
||||||
return db, nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func normalizeDatabaseType(value string) string {
|
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.
|
// zapStdLogger adapts a *zap.Logger to GORM's tiny logger interface.
|
||||||
type zapStdLogger struct{ log *zap.Logger }
|
type zapStdLogger struct{ log *zap.Logger }
|
||||||
|
|
||||||
func (z zapStdLogger) Printf(format string, args ...interface{}) {
|
func (z zapStdLogger) Printf(format string, args ...interface{}) {
|
||||||
|
if z.log == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
z.log.Sugar().Infof(format, args...)
|
z.log.Sugar().Infof(format, args...)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package database
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -12,6 +13,47 @@ import (
|
|||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
"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) {
|
func TestEnforceTelegramBindingOneToOneCleansDuplicatesAndAddsIndex(t *testing.T) {
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
if err != nil {
|
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) {
|
func TestCopyModelTablesMigratesExistingSQLiteRows(t *testing.T) {
|
||||||
src, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
src, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
if err != nil {
|
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()})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if svc.Sessions != nil {
|
||||||
|
svc.Sessions.ApplyToUsers(c.Request.Context(), users)
|
||||||
|
}
|
||||||
c.JSON(http.StatusOK, 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"})
|
c.JSON(http.StatusForbidden, gin.H{"error": "default admin cannot be deleted"})
|
||||||
return
|
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 {
|
if err := svc.Repo.User.Delete(c.Request.Context(), c.Param("id")); err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||||
return
|
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()})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||||
return
|
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{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"user": resp.User,
|
"user": resp.User,
|
||||||
"tokens": resp.Tokens,
|
"tokens": resp.Tokens,
|
||||||
@@ -65,6 +71,9 @@ func registerHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if tokens != nil {
|
||||||
|
setAccessTokenCookie(c, tokens.AccessToken, int(tokens.ExpiresIn))
|
||||||
|
}
|
||||||
c.JSON(http.StatusCreated, gin.H{
|
c.JSON(http.StatusCreated, gin.H{
|
||||||
"user": u,
|
"user": u,
|
||||||
"tokens": tokens,
|
"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.
|
// so the Vue frontend's logout button gets a 200 instead of 404.
|
||||||
func logoutHandler(_ *service.Container) gin.HandlerFunc {
|
func logoutHandler(_ *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
|
clearAccessTokenCookie(c)
|
||||||
c.Status(http.StatusNoContent)
|
c.Status(http.StatusNoContent)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+60
-214
@@ -3,18 +3,10 @@
|
|||||||
package handler
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/sha256"
|
|
||||||
"encoding/hex"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
|
||||||
"path"
|
|
||||||
"sort"
|
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"go.uber.org/zap"
|
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||||
@@ -25,6 +17,10 @@ import (
|
|||||||
func cloudListHandler(svc *service.Container) gin.HandlerFunc {
|
func cloudListHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
typ := c.Param("type")
|
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")
|
dir := c.Query("dir")
|
||||||
entries, err := svc.StorageCfg.CloudList(c.Request.Context(), typ, dir)
|
entries, err := svc.StorageCfg.CloudList(c.Request.Context(), typ, dir)
|
||||||
if err != nil {
|
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.
|
// cloudImportHandler turns a cloud file into a playable 302-backed media item.
|
||||||
func cloudImportHandler(svc *service.Container) gin.HandlerFunc {
|
func cloudImportHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
typ := c.Param("type")
|
typ := c.Param("type")
|
||||||
|
if !service.IsAdminCloudConfigurable(typ) {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||||
|
return
|
||||||
|
}
|
||||||
var in struct {
|
var in struct {
|
||||||
Ref string `json:"ref" binding:"required"`
|
Ref string `json:"ref" binding:"required"`
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
@@ -63,6 +111,10 @@ func cloudImportHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
func cloudMountHandler(svc *service.Container) gin.HandlerFunc {
|
func cloudMountHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
typ := c.Param("type")
|
typ := c.Param("type")
|
||||||
|
if !service.IsAdminCloudConfigurable(typ) {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported cloud provider"})
|
||||||
|
return
|
||||||
|
}
|
||||||
var in struct {
|
var in struct {
|
||||||
Dir string `json:"dir"`
|
Dir string `json:"dir"`
|
||||||
DirPath string `json:"dir_path"`
|
DirPath string `json:"dir_path"`
|
||||||
@@ -284,209 +336,3 @@ func cloud115QRPollHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
c.JSON(http.StatusOK, st)
|
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
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"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) {
|
func TestCloudMountLibraryNameDefaultsToDirectoryBaseName(t *testing.T) {
|
||||||
@@ -48,3 +57,98 @@ func TestCloudPlaybackDiagnosticsDoNotExposeRawRefOrURL(t *testing.T) {
|
|||||||
t.Fatalf("header names = %q", got)
|
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 {
|
if items == nil {
|
||||||
items = []service.Match{}
|
items = []service.Match{}
|
||||||
}
|
}
|
||||||
|
svc.Discover.WarmMatchArtwork(items)
|
||||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -41,6 +42,7 @@ func popularHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
if items == nil {
|
if items == nil {
|
||||||
items = []service.Match{}
|
items = []service.Match{}
|
||||||
}
|
}
|
||||||
|
svc.Discover.WarmMatchArtwork(items)
|
||||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,24 +16,37 @@ import (
|
|||||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
"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
|
// discoverSectionsHandler returns the catalog of sections the UI can
|
||||||
// pick from. The names match the upstream Vue UI so existing settings
|
// pick from. The names match the upstream Vue UI so existing settings
|
||||||
// keep working.
|
// keep working.
|
||||||
func discoverSectionsHandler(_ *service.Container) gin.HandlerFunc {
|
func discoverSectionsHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
c.JSON(http.StatusOK, gin.H{
|
sections := make([]gin.H, 0, len(discoverSectionCatalog))
|
||||||
"sections": []gin.H{
|
for _, section := range discoverSectionCatalog {
|
||||||
{"key": "tmdb_trending_day", "label": "TMDb 今日趋势", "provider": "tmdb"},
|
if !discoverProviderEnabled(c.Request.Context(), svc, section.Provider) {
|
||||||
{"key": "tmdb_trending_week", "label": "TMDb 本周热门", "provider": "tmdb"},
|
continue
|
||||||
{"key": "tmdb_popular_movie", "label": "TMDb 热门电影", "provider": "tmdb"},
|
}
|
||||||
{"key": "tmdb_popular_tv", "label": "TMDb 热门剧集", "provider": "tmdb"},
|
sections = append(sections, gin.H{"key": section.Key, "label": section.Label, "provider": section.Provider})
|
||||||
{"key": "tmdb_top_rated_movie", "label": "TMDb 高分电影", "provider": "tmdb"},
|
}
|
||||||
{"key": "douban_hot_movie", "label": "豆瓣热门电影", "provider": "douban"},
|
c.JSON(http.StatusOK, gin.H{"sections": sections})
|
||||||
{"key": "douban_hot_tv", "label": "豆瓣热门剧集", "provider": "douban"},
|
|
||||||
{"key": "douban_top_movie", "label": "豆瓣高分电影", "provider": "douban"},
|
|
||||||
{"key": "bangumi_calendar", "label": "Bangumi 每日放送", "provider": "bangumi"},
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -45,19 +58,51 @@ func discoverFeedHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
keys := strings.Split(c.DefaultQuery("sections", "tmdb_trending_day,tmdb_popular_movie,douban_hot_movie,bangumi_calendar"), ",")
|
keys := strings.Split(c.DefaultQuery("sections", "tmdb_trending_day,tmdb_popular_movie,douban_hot_movie,bangumi_calendar"), ",")
|
||||||
out := gin.H{}
|
out := gin.H{}
|
||||||
|
artworkItems := []service.ExternalMediaResult{}
|
||||||
for _, raw := range keys {
|
for _, raw := range keys {
|
||||||
k := strings.TrimSpace(raw)
|
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)
|
items, err := discoverSectionItems(c.Request.Context(), svc, k)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
svc.Log.Debug("discover fetch failed")
|
svc.Log.Debug("discover fetch failed")
|
||||||
items = nil
|
items = nil
|
||||||
}
|
}
|
||||||
|
artworkItems = append(artworkItems, items...)
|
||||||
out[k] = items
|
out[k] = items
|
||||||
}
|
}
|
||||||
|
svc.Discover.WarmExternalArtwork(artworkItems)
|
||||||
c.JSON(http.StatusOK, out)
|
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) {
|
func discoverSectionItems(ctx context.Context, svc *service.Container, k string) ([]service.ExternalMediaResult, error) {
|
||||||
switch k {
|
switch k {
|
||||||
case "tmdb_trending_day", "tmdb_trending_week", "tmdb_popular_movie", "tmdb_popular_tv", "tmdb_top_rated_movie",
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -12,8 +12,20 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type manualScrapeApplyReq struct {
|
type manualScrapeApplyReq struct {
|
||||||
MediaIDs []string `json:"media_ids"`
|
MediaIDs []string `json:"media_ids"`
|
||||||
Match service.ManualScrapeRequest `json:"match"`
|
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
|
const manualScrapeApplyTimeout = 5 * time.Minute
|
||||||
@@ -72,10 +84,11 @@ func manualScrapeApplyBatchHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
}
|
}
|
||||||
applyCtx, cancel := manualScrapeApplyContext(c)
|
applyCtx, cancel := manualScrapeApplyContext(c)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
options := service.ScrapeOptions{EpisodeArtwork: req.episodeArtworkOption()}
|
||||||
applied := 0
|
applied := 0
|
||||||
errorsOut := make([]string, 0)
|
errorsOut := make([]string, 0)
|
||||||
for _, id := range ids {
|
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())
|
errorsOut = append(errorsOut, id+": "+err.Error())
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ func listLibrariesHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
libs = service.FilterDeprecatedNativeCloudLibraries(libs)
|
||||||
role, _ := c.Get(middleware.CtxUserRole)
|
role, _ := c.Get(middleware.CtxUserRole)
|
||||||
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("all") == "1")
|
includeHidden := role == "admin" && (c.Query("include_hidden") == "1" || c.Query("all") == "1")
|
||||||
if !includeHidden {
|
if !includeHidden {
|
||||||
@@ -109,7 +110,7 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
_, _, _ = svc.Scan.StartCloudLibraryScan(id, false)
|
_, _, _ = svc.Scan.StartCloudLibraryScan(id, true)
|
||||||
finishHTTPTask(task, nil, "queued", "云盘扫描已加入后台队列", map[string]int64{"queued": 1}, nil)
|
finishHTTPTask(task, nil, "queued", "云盘扫描已加入后台队列", map[string]int64{"queued": 1}, nil)
|
||||||
c.JSON(http.StatusAccepted, gin.H{
|
c.JSON(http.StatusAccepted, gin.H{
|
||||||
"library_id": id,
|
"library_id": id,
|
||||||
@@ -119,7 +120,7 @@ func scanLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
"probed": 0,
|
"probed": 0,
|
||||||
"queued": true,
|
"queued": true,
|
||||||
"cloud": true,
|
"cloud": true,
|
||||||
"message": "云盘扫描已在后台运行,发现的媒体会自动加入当前媒体库",
|
"message": "云盘扫描已在后台运行,发现的媒体会自动加入当前媒体库;若已开启自动刮削,会在扫描后补齐元数据",
|
||||||
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
"estimate_message": "小目录通常几十秒;几万文件的大目录可能需要数分钟到数小时,取决于网盘接口速度",
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -80,16 +80,22 @@ func listFavoritesAliasHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
// path; the AI hint comes from svc.AI when configured.
|
// path; the AI hint comes from svc.AI when configured.
|
||||||
func aiScrapeMediaHandler(svc *service.Container) gin.HandlerFunc {
|
func aiScrapeMediaHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
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"))
|
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
|
||||||
if err != nil || m == nil {
|
if err != nil || m == nil {
|
||||||
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
|
c.JSON(http.StatusNotFound, gin.H{"error": "media not found"})
|
||||||
return
|
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()})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||||
return
|
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
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"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 {
|
func requestLibraries(t *testing.T, svc *service.Container, userID, role, path string) []model.Library {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
@@ -177,6 +257,16 @@ type mediaListResponse struct {
|
|||||||
Total int64 `json:"total"`
|
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 {
|
func requestMediaList(t *testing.T, svc *service.Container, path, libraryID string) mediaListResponse {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
@@ -195,3 +285,41 @@ func requestMediaList(t *testing.T, svc *service.Container, path, libraryID stri
|
|||||||
}
|
}
|
||||||
return payload
|
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
|
return
|
||||||
}
|
}
|
||||||
token := externalPlaybackToken(c, svc, m.ID, m.DurationSec)
|
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)
|
escapedStream := url.QueryEscape(streamURL)
|
||||||
c.JSON(http.StatusOK, gin.H{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"url": streamURL,
|
"url": streamURL,
|
||||||
@@ -100,7 +100,7 @@ func externalURLHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
}
|
}
|
||||||
token := externalPlaybackToken(c, svc, m.ID, m.DurationSec)
|
token := externalPlaybackToken(c, svc, m.ID, m.DurationSec)
|
||||||
c.JSON(http.StatusOK, gin.H{
|
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
|
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 {
|
func absoluteRequestURL(c *gin.Context, path string) string {
|
||||||
if strings.HasPrefix(path, "http://") || strings.HasPrefix(path, "https://") {
|
if strings.HasPrefix(path, "http://") || strings.HasPrefix(path, "https://") {
|
||||||
return path
|
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) {
|
func TestScopedPlaybackTokenCannotStreamAnotherMedia(t *testing.T) {
|
||||||
router, svc, _ := newPlaybackScopeTestRouter(t)
|
router, svc, _ := newPlaybackScopeTestRouter(t)
|
||||||
user, err := svc.Repo.User.FindByID(t.Context(), "user-1")
|
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 := router.Group("/api")
|
||||||
api.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
|
api.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
|
||||||
api.GET("/playback/:id/external-url", externalURLHandler(svc))
|
api.GET("/playback/:id/external-url", externalURLHandler(svc))
|
||||||
|
api.GET("/playback/:id/external-players", externalPlayersHandler(svc))
|
||||||
api.GET("/stream/:id", streamHandler(svc))
|
api.GET("/stream/:id", streamHandler(svc))
|
||||||
api.GET("/cloud/play/:type", cloudPlayHandler(svc))
|
api.GET("/cloud/play/:type", cloudPlayHandler(svc))
|
||||||
return router, svc, cfg.Secrets.JWTSecret
|
return router, svc, cfg.Secrets.JWTSecret
|
||||||
|
|||||||
@@ -57,6 +57,7 @@ func (h *RefreshHandler) RefreshToken(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
setAccessTokenCookie(c, tokens.AccessToken, int(tokens.ExpiresIn))
|
||||||
c.JSON(http.StatusOK, gin.H{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"code": 0,
|
"code": 0,
|
||||||
"message": "ok",
|
"message": "ok",
|
||||||
@@ -72,6 +73,7 @@ func (h *RefreshHandler) RefreshToken(c *gin.Context) {
|
|||||||
// Logout 登出当前用户。
|
// Logout 登出当前用户。
|
||||||
// POST /api/auth/logout
|
// POST /api/auth/logout
|
||||||
func (h *RefreshHandler) Logout(c *gin.Context) {
|
func (h *RefreshHandler) Logout(c *gin.Context) {
|
||||||
|
clearAccessTokenCookie(c)
|
||||||
userID := c.GetString("ctx_user_id")
|
userID := c.GetString("ctx_user_id")
|
||||||
if userID == "" {
|
if userID == "" {
|
||||||
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": nil})
|
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": nil})
|
||||||
|
|||||||
@@ -17,14 +17,22 @@ import (
|
|||||||
// 异步执行, 立即返回 202;通过 WS hub "scrape" topic 推送进度。
|
// 异步执行, 立即返回 202;通过 WS hub "scrape" topic 推送进度。
|
||||||
func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
|
func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
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, "全库修复并重刮", "", "")
|
task := startScrapeHTTPTask(svc, "全库修复并重刮", "", "")
|
||||||
go func() {
|
go func(options service.ScrapeOptions) {
|
||||||
result, err := svc.RepairAndRescrapeAllLibraries(context.Background())
|
result, err := svc.RepairAndRescrapeAllLibraries(context.Background(), options)
|
||||||
metrics := map[string]int64{
|
metrics := map[string]int64{
|
||||||
"repaired": int64(result.Repaired),
|
"repaired": int64(result.Repaired),
|
||||||
"libraries": int64(result.Libraries),
|
"reclassified": int64(result.Reclassified),
|
||||||
"matched": int64(result.Matched),
|
"libraries": int64(result.Libraries),
|
||||||
"reset": int64(result.Reset),
|
"matched": int64(result.Matched),
|
||||||
|
"processed": int64(result.Processed),
|
||||||
|
"errors": int64(result.Errors),
|
||||||
|
"reset": int64(result.Reset),
|
||||||
}
|
}
|
||||||
stage := "completed"
|
stage := "completed"
|
||||||
message := "全库修复并重刮完成"
|
message := "全库修复并重刮完成"
|
||||||
@@ -33,7 +41,7 @@ func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
message = "全库修复并重刮失败"
|
message = "全库修复并重刮失败"
|
||||||
}
|
}
|
||||||
finishHTTPTask(task, err, stage, message, metrics, nil)
|
finishHTTPTask(task, err, stage, message, metrics, nil)
|
||||||
}()
|
}(options)
|
||||||
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
|
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -46,14 +54,22 @@ func repairAndRescrapeAllHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
func repairAndRescrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
func repairAndRescrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
libraryID := c.Param("id")
|
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, "媒体库修复并重刮", "", "")
|
task := startScrapeHTTPTask(svc, "媒体库修复并重刮", "", "")
|
||||||
go func() {
|
go func(options service.ScrapeOptions) {
|
||||||
result, err := svc.RepairAndRescrapeLibrary(context.Background(), libraryID)
|
result, err := svc.RepairAndRescrapeLibrary(context.Background(), libraryID, options)
|
||||||
metrics := map[string]int64{
|
metrics := map[string]int64{
|
||||||
"repaired": int64(result.Repaired),
|
"repaired": int64(result.Repaired),
|
||||||
"libraries": int64(result.Libraries),
|
"reclassified": int64(result.Reclassified),
|
||||||
"matched": int64(result.Matched),
|
"libraries": int64(result.Libraries),
|
||||||
"reset": int64(result.Reset),
|
"matched": int64(result.Matched),
|
||||||
|
"processed": int64(result.Processed),
|
||||||
|
"errors": int64(result.Errors),
|
||||||
|
"reset": int64(result.Reset),
|
||||||
}
|
}
|
||||||
stage := "completed"
|
stage := "completed"
|
||||||
message := "媒体库修复并重刮完成"
|
message := "媒体库修复并重刮完成"
|
||||||
@@ -62,7 +78,7 @@ func repairAndRescrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
message = "媒体库修复并重刮失败"
|
message = "媒体库修复并重刮失败"
|
||||||
}
|
}
|
||||||
finishHTTPTask(task, err, stage, message, metrics, nil)
|
finishHTTPTask(task, err, stage, message, metrics, nil)
|
||||||
}()
|
}(options)
|
||||||
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
|
c.JSON(http.StatusAccepted, gin.H{"status": "started"})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,98 +10,120 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func registerAdminRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.Container) {
|
func registerAdminRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.Container) {
|
||||||
// Admin-only endpoints.
|
|
||||||
admin := api.Group("/admin")
|
admin := api.Group("/admin")
|
||||||
admin.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret), middleware.AdminRequired())
|
admin.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret), middleware.AdminRequired())
|
||||||
{
|
registerAdminUserRoutes(admin, svc)
|
||||||
admin.GET("/users", listUsersHandler(svc))
|
registerAdminPermissionRoutes(admin, svc)
|
||||||
admin.POST("/users", createUserHandler(svc))
|
registerAdminStorageRoutes(admin, svc)
|
||||||
admin.PATCH("/users/:id", updateUserHandler(svc))
|
registerAdminCloudRoutes(admin, svc)
|
||||||
admin.PATCH("/users/:id/password", resetUserPasswordHandler(svc))
|
registerAdminDownloadClientRoutes(admin, svc)
|
||||||
admin.PATCH("/users/:id/status", updateUserStatusHandler(svc))
|
registerAdminSystemRoutes(admin, svc)
|
||||||
admin.PATCH("/users/:id/role", adminUpdateRoleHandler(svc))
|
registerAdminBackupRoutes(admin, svc)
|
||||||
admin.DELETE("/users/:id", deleteUserHandler(svc))
|
registerAdminNotificationRoutes(admin, svc)
|
||||||
admin.GET("/settings", listSettingsHandler(svc))
|
registerAdminTelegramRoutes(admin, svc)
|
||||||
admin.PUT("/settings", updateSettingHandler(svc))
|
registerAdminOrganizerRoutes(admin, svc)
|
||||||
admin.GET("/logs", recentLogsHandler(svc))
|
registerAdminRepairRoutes(admin, svc)
|
||||||
|
registerAdminAPIConfigRoutes(admin, svc)
|
||||||
// Permissions admin.
|
registerAdminSchedulerRoutes(admin, svc)
|
||||||
admin.GET("/users/:id/permissions", getUserPermissionsHandler(svc))
|
}
|
||||||
admin.PUT("/users/:id/permissions", updateUserPermissionsHandler(svc))
|
|
||||||
admin.POST("/users/:id/permissions/reset", resetUserPermissionsHandler(svc))
|
func registerAdminUserRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
|
admin.GET("/users", listUsersHandler(svc))
|
||||||
// Storage configs (Alist / S3 / WebDAV / 网盘).
|
admin.POST("/users", createUserHandler(svc))
|
||||||
admin.GET("/storage/status", listStorageConfigsHandler(svc))
|
admin.PATCH("/users/:id", updateUserHandler(svc))
|
||||||
admin.GET("/storage/:type", getStorageConfigHandler(svc))
|
admin.PATCH("/users/:id/password", resetUserPasswordHandler(svc))
|
||||||
admin.PUT("/storage/:type", saveStorageConfigHandler(svc))
|
admin.PATCH("/users/:id/status", updateUserStatusHandler(svc))
|
||||||
admin.POST("/storage/:type/test", testStorageConfigHandler(svc))
|
admin.PATCH("/users/:id/role", adminUpdateRoleHandler(svc))
|
||||||
admin.POST("/storage/:type/logout", logoutStorageConfigHandler(svc))
|
admin.DELETE("/users/:id", deleteUserHandler(svc))
|
||||||
admin.POST("/storage/:type/upload-local", storageUploadLocalHandler(svc))
|
admin.GET("/settings", listSettingsHandler(svc))
|
||||||
|
admin.PUT("/settings", updateSettingHandler(svc))
|
||||||
// Cloud disk (115 / 夸克) browsing, QR login and 302 import.
|
admin.GET("/logs", recentLogsHandler(svc))
|
||||||
admin.POST("/cloud/scan-all", cloudScanAllHandler(svc))
|
}
|
||||||
admin.POST("/cloud/scan/cancel", cloudScanCancelHandler(svc))
|
|
||||||
admin.GET("/cloud/scan/status", cloudScanStatusHandler(svc))
|
func registerAdminPermissionRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
admin.GET("/cloud/:type/list", cloudListHandler(svc))
|
admin.GET("/users/:id/permissions", getUserPermissionsHandler(svc))
|
||||||
admin.POST("/cloud/:type/import", cloudImportHandler(svc))
|
admin.PUT("/users/:id/permissions", updateUserPermissionsHandler(svc))
|
||||||
admin.POST("/cloud/:type/mount", cloudMountHandler(svc))
|
admin.POST("/users/:id/permissions/reset", resetUserPermissionsHandler(svc))
|
||||||
admin.POST("/cloud/:type/qr/start", cloud115QRStartHandler(svc))
|
}
|
||||||
admin.POST("/cloud/:type/qr/poll", cloud115QRPollHandler(svc))
|
|
||||||
|
func registerAdminStorageRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
// Download client CRUD.
|
admin.GET("/storage/status", listStorageConfigsHandler(svc))
|
||||||
admin.GET("/download/clients", listDownloadClientsHandler(svc))
|
admin.GET("/storage/:type", getStorageConfigHandler(svc))
|
||||||
admin.POST("/download/clients", createDownloadClientHandler(svc))
|
admin.PUT("/storage/:type", saveStorageConfigHandler(svc))
|
||||||
admin.PUT("/download/clients/:id", updateDownloadClientHandler(svc))
|
admin.POST("/storage/:type/test", testStorageConfigHandler(svc))
|
||||||
admin.DELETE("/download/clients/:id", deleteDownloadClientHandler(svc))
|
admin.POST("/storage/:type/logout", logoutStorageConfigHandler(svc))
|
||||||
admin.POST("/download/clients/:id/test", testDownloadClientHandler(svc))
|
admin.POST("/storage/:type/upload-local", storageUploadLocalHandler(svc))
|
||||||
admin.GET("/download/aria2/stats", aria2StatsHandler(svc))
|
}
|
||||||
|
|
||||||
// System scheduler trigger alias.
|
func registerAdminCloudRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
admin.POST("/system/scheduler/:name/trigger", schedulerTriggerHandler(svc))
|
admin.POST("/cloud/scan-all", cloudScanAllHandler(svc))
|
||||||
|
admin.POST("/cloud/scan/cancel", cloudScanCancelHandler(svc))
|
||||||
// Database backup.
|
admin.GET("/cloud/scan/status", cloudScanStatusHandler(svc))
|
||||||
admin.GET("/backups", listBackupsHandler(svc))
|
admin.GET("/cloud/:type/list", cloudListHandler(svc))
|
||||||
admin.POST("/backups", createBackupHandler(svc))
|
admin.POST("/cloud/:type/mkdir", cloudMkdirHandler(svc))
|
||||||
admin.DELETE("/backups", deleteBackupHandler(svc))
|
admin.PUT("/cloud/:type/rename", cloudRenameHandler(svc))
|
||||||
admin.POST("/backups/restore", restoreBackupHandler(svc))
|
admin.POST("/cloud/:type/import", cloudImportHandler(svc))
|
||||||
|
admin.POST("/cloud/:type/mount", cloudMountHandler(svc))
|
||||||
// Notifications (test endpoint).
|
admin.POST("/cloud/:type/qr/start", cloud115QRStartHandler(svc))
|
||||||
admin.POST("/notify/test", notifyTestHandler(svc))
|
admin.POST("/cloud/:type/qr/poll", cloud115QRPollHandler(svc))
|
||||||
|
}
|
||||||
// Notify channels CRUD + per-channel test.
|
|
||||||
admin.GET("/notify/channels", listNotifyChannelsHandler(svc))
|
func registerAdminDownloadClientRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
admin.POST("/notify/channels", createNotifyChannelHandler(svc))
|
admin.GET("/download/clients", listDownloadClientsHandler(svc))
|
||||||
admin.PUT("/notify/channels/:id", updateNotifyChannelHandler(svc))
|
admin.POST("/download/clients", createDownloadClientHandler(svc))
|
||||||
admin.DELETE("/notify/channels/:id", deleteNotifyChannelHandler(svc))
|
admin.PUT("/download/clients/:id", updateDownloadClientHandler(svc))
|
||||||
admin.POST("/notify/channels/:id/test", testNotifyChannelHandler(svc))
|
admin.DELETE("/download/clients/:id", deleteDownloadClientHandler(svc))
|
||||||
|
admin.POST("/download/clients/:id/test", testDownloadClientHandler(svc))
|
||||||
// Telegram Bot webhook management.
|
admin.GET("/download/aria2/stats", aria2StatsHandler(svc))
|
||||||
admin.GET("/telegram/webhook", telegramGetWebhookHandler(svc))
|
}
|
||||||
admin.POST("/telegram/webhook", telegramSetWebhookHandler(svc))
|
|
||||||
admin.POST("/telegram/polling/start", telegramStartPollingHandler(svc))
|
func registerAdminSystemRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
admin.POST("/telegram/polling/stop", telegramStopPollingHandler(svc))
|
admin.POST("/system/scheduler/:name/trigger", schedulerTriggerHandler(svc))
|
||||||
|
}
|
||||||
// File organizer.
|
|
||||||
admin.POST("/media/:id/organize", organizeMediaHandler(svc))
|
func registerAdminBackupRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
admin.POST("/libraries/:id/organize", organizeLibraryHandler(svc))
|
admin.GET("/backups", listBackupsHandler(svc))
|
||||||
admin.GET("/organize/sources", organizeSourcesHandler(svc))
|
admin.POST("/backups", createBackupHandler(svc))
|
||||||
admin.POST("/organize/source", organizeDirectoryHandler(svc))
|
admin.DELETE("/backups", deleteBackupHandler(svc))
|
||||||
|
admin.POST("/backups/restore", restoreBackupHandler(svc))
|
||||||
// 全库修复+重刮:从路径占位符回填缺失外部 ID,然后批量重刮整库。
|
}
|
||||||
admin.POST("/media/repair-rescrape", repairAndRescrapeAllHandler(svc))
|
|
||||||
// 单库修复+重刮:只对指定媒体库回填占位符外部 ID 并重刮。
|
func registerAdminNotificationRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
admin.POST("/libraries/:id/repair-rescrape", repairAndRescrapeLibraryHandler(svc))
|
admin.POST("/notify/test", notifyTestHandler(svc))
|
||||||
|
admin.GET("/notify/channels", listNotifyChannelsHandler(svc))
|
||||||
// API key management (encrypted at rest).
|
admin.POST("/notify/channels", createNotifyChannelHandler(svc))
|
||||||
admin.GET("/api-configs", listAPIConfigsHandler(svc))
|
admin.PUT("/notify/channels/:id", updateNotifyChannelHandler(svc))
|
||||||
admin.GET("/api-configs/:provider", getAPIConfigHandler(svc))
|
admin.DELETE("/notify/channels/:id", deleteNotifyChannelHandler(svc))
|
||||||
admin.PUT("/api-configs/:provider", updateAPIConfigHandler(svc))
|
admin.POST("/notify/channels/:id/test", testNotifyChannelHandler(svc))
|
||||||
admin.DELETE("/api-configs/:provider", deleteAPIConfigHandler(svc))
|
}
|
||||||
|
|
||||||
// Scheduled jobs.
|
func registerAdminTelegramRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
admin.GET("/scheduler", schedulerStatusHandler(svc))
|
admin.GET("/telegram/webhook", telegramGetWebhookHandler(svc))
|
||||||
admin.POST("/scheduler/:name/run", schedulerRunHandler(svc))
|
admin.POST("/telegram/webhook", telegramSetWebhookHandler(svc))
|
||||||
|
admin.POST("/telegram/polling/start", telegramStartPollingHandler(svc))
|
||||||
}
|
admin.POST("/telegram/polling/stop", telegramStopPollingHandler(svc))
|
||||||
|
}
|
||||||
|
|
||||||
|
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))
|
||||||
|
}
|
||||||
|
|
||||||
|
func registerAdminRepairRoutes(admin *gin.RouterGroup, svc *service.Container) {
|
||||||
|
admin.POST("/media/repair-rescrape", repairAndRescrapeAllHandler(svc))
|
||||||
|
admin.POST("/libraries/:id/repair-rescrape", repairAndRescrapeLibraryHandler(svc))
|
||||||
|
}
|
||||||
|
|
||||||
|
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))
|
||||||
|
}
|
||||||
|
|
||||||
|
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) {
|
func registerAuthenticatedRoutes(api *gin.RouterGroup, cfg *config.Config, svc *service.Container) {
|
||||||
// Authenticated endpoints.
|
|
||||||
authed := api.Group("/")
|
authed := api.Group("/")
|
||||||
authed.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
|
authed.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret))
|
||||||
authed.Use(activeUserRequired(svc))
|
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 (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
@@ -31,15 +32,20 @@ func listSeasonsHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
var rows []model.Media
|
|
||||||
err := svc.Repo.DB.Where(&model.Media{LibraryID: libID}).
|
|
||||||
Order("season_num asc, episode_num asc").
|
|
||||||
Find(&rows).Error
|
|
||||||
if err != nil && err != gorm.ErrRecordNotFound {
|
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
visibility := mediaVisibilityForRequest(c, svc)
|
visibility := mediaVisibilityForRequest(c, svc)
|
||||||
|
var rows []model.Media
|
||||||
|
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
|
||||||
|
}
|
||||||
|
rows = append(rows, pageRows...)
|
||||||
|
if int64(len(rows)) >= total || len(pageRows) < pageSize {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
buckets := make(map[int][]model.Media)
|
buckets := make(map[int][]model.Media)
|
||||||
for _, r := range rows {
|
for _, r := range rows {
|
||||||
if !visibility.Allows(&r) {
|
if !visibility.Allows(&r) {
|
||||||
@@ -55,3 +61,65 @@ func listSeasonsHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
c.JSON(http.StatusOK, gin.H{"seasons": out})
|
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
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -25,6 +25,10 @@ func listStorageConfigsHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
// getStorageConfigHandler returns one config (with the decrypted body).
|
// getStorageConfigHandler returns one config (with the decrypted body).
|
||||||
func getStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
func getStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
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"))
|
row, err := svc.StorageCfg.Get(c.Request.Context(), c.Param("type"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
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.
|
// the type via URL and the body as a JSON object.
|
||||||
func saveStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
func saveStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
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
|
var in service.StorageInput
|
||||||
if err := c.ShouldBindJSON(&in); err != nil {
|
if err := c.ShouldBindJSON(&in); err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
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.
|
// testStorageConfigHandler probes an unsaved config.
|
||||||
func testStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
func testStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
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
|
var in service.StorageInput
|
||||||
if err := c.ShouldBindJSON(&in); err != nil {
|
if err := c.ShouldBindJSON(&in); err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
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 {
|
func logoutStorageConfigHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
typ := c.Param("type")
|
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)
|
row, err := svc.StorageCfg.Logout(c.Request.Context(), typ)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
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 {
|
func storageUploadLocalHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
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
|
var req service.CloudUploadInput
|
||||||
if err := c.ShouldBindJSON(&req); err != nil {
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
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 (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"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.
|
// scrapeOneHandler enriches a single media via the configured scraper chain.
|
||||||
func scrapeOneHandler(svc *service.Container) gin.HandlerFunc {
|
func scrapeOneHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
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"))
|
m, err := svc.Repo.Media.FindByID(c.Request.Context(), c.Param("id"))
|
||||||
if err != nil || m == nil {
|
if err != nil || m == nil {
|
||||||
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
c.JSON(http.StatusNotFound, gin.H{"error": "not found"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
task := startScrapeHTTPTask(svc, "手动刮削媒体", m.Title, m.Path)
|
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)
|
finishHTTPTask(task, err, "scrape", "手动刮削媒体失败", nil, nil)
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||||
return
|
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 {
|
func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
libID := c.Param("id")
|
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
|
var task *service.TaskHandle
|
||||||
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
|
if lib, err := svc.Repo.Library.FindByID(c.Request.Context(), libID); err == nil && lib != nil {
|
||||||
task = startScrapeHTTPTask(svc, "手动刮削媒体库", lib.Name, lib.Path)
|
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
|
// Run in the background so HTTP returns instantly; the WS hub
|
||||||
// pushes per-item progress on the "scrape" topic.
|
// pushes per-item progress on the "scrape" topic.
|
||||||
go func(libID string, task *service.TaskHandle) {
|
go func(libID string, task *service.TaskHandle, options service.ScrapeOptions) {
|
||||||
matched, err := svc.Scraper.EnrichLibrary(context.Background(), libID, true)
|
result, err := svc.Scraper.EnrichLibraryDetailedWithOptions(context.Background(), libID, options)
|
||||||
metrics := map[string]int64{"matched": int64(matched)}
|
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"
|
stage := "completed"
|
||||||
message := "手动刮削媒体库结束"
|
message := "手动刮削媒体库结束"
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -125,7 +222,7 @@ func scrapeLibraryHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
message = "手动刮削媒体库失败"
|
message = "手动刮削媒体库失败"
|
||||||
}
|
}
|
||||||
finishHTTPTask(task, err, stage, message, metrics, nil)
|
finishHTTPTask(task, err, stage, message, metrics, nil)
|
||||||
}(libID, task)
|
}(libID, task, options)
|
||||||
c.JSON(http.StatusAccepted, gin.H{"status": "scraping"})
|
c.JSON(http.StatusAccepted, gin.H{"status": "scraping"})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -99,7 +99,7 @@ func importSTRMHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type generateSTRMReq struct {
|
type generateSTRMReq struct {
|
||||||
LibraryID string `json:"library_id" binding:"required"`
|
LibraryID string `json:"library_id"`
|
||||||
OutputDir string `json:"output_dir"`
|
OutputDir string `json:"output_dir"`
|
||||||
BaseURL string `json:"base_url"`
|
BaseURL string `json:"base_url"`
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
@@ -122,7 +122,7 @@ func generateSTRMHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
if baseURL == "" {
|
if baseURL == "" {
|
||||||
baseURL = strings.TrimRight(absoluteRequestURL(c, "/"), "/")
|
baseURL = strings.TrimRight(absoluteRequestURL(c, "/"), "/")
|
||||||
}
|
}
|
||||||
res, err := strmSvc.GenerateForLibrary(c.Request.Context(), service.GenerateSTRMOptions{
|
options := service.GenerateSTRMOptions{
|
||||||
LibraryID: req.LibraryID,
|
LibraryID: req.LibraryID,
|
||||||
OutputDir: req.OutputDir,
|
OutputDir: req.OutputDir,
|
||||||
BaseURL: baseURL,
|
BaseURL: baseURL,
|
||||||
@@ -130,7 +130,14 @@ func generateSTRMHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
Overwrite: req.Overwrite,
|
Overwrite: req.Overwrite,
|
||||||
IncludeLocal: true,
|
IncludeLocal: true,
|
||||||
PlaybackToken: strmPlaybackTokenForRequest(c, svc),
|
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 {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -24,6 +24,8 @@ type subscriptionReq struct {
|
|||||||
PosterURL string `json:"poster_url"`
|
PosterURL string `json:"poster_url"`
|
||||||
BackdropURL string `json:"backdrop_url"`
|
BackdropURL string `json:"backdrop_url"`
|
||||||
Overview string `json:"overview"`
|
Overview string `json:"overview"`
|
||||||
|
OriginalName string `json:"original_name"`
|
||||||
|
Year int `json:"year"`
|
||||||
Resolution string `json:"resolution"`
|
Resolution string `json:"resolution"`
|
||||||
Quality string `json:"quality"`
|
Quality string `json:"quality"`
|
||||||
Effects string `json:"effects"`
|
Effects string `json:"effects"`
|
||||||
@@ -62,6 +64,8 @@ func createSubscriptionHandler(svc *service.Container) gin.HandlerFunc {
|
|||||||
PosterURL: req.PosterURL,
|
PosterURL: req.PosterURL,
|
||||||
BackdropURL: req.BackdropURL,
|
BackdropURL: req.BackdropURL,
|
||||||
Overview: req.Overview,
|
Overview: req.Overview,
|
||||||
|
OriginalName: req.OriginalName,
|
||||||
|
Year: req.Year,
|
||||||
Resolution: req.Resolution,
|
Resolution: req.Resolution,
|
||||||
Quality: req.Quality,
|
Quality: req.Quality,
|
||||||
Effects: req.Effects,
|
Effects: req.Effects,
|
||||||
|
|||||||
@@ -85,12 +85,11 @@ func schemaHandler(_ *service.Container) gin.HandlerFunc {
|
|||||||
{"key": "cloud.boot_scan_enabled", "type": "toggle", "label": "启动后立即扫描网盘"},
|
{"key": "cloud.boot_scan_enabled", "type": "toggle", "label": "启动后立即扫描网盘"},
|
||||||
{"key": "cloud.upload_auto_enabled", "type": "toggle", "label": "启用自动转存"},
|
{"key": "cloud.upload_auto_enabled", "type": "toggle", "label": "启用自动转存"},
|
||||||
{"key": "cloud.upload_provider", "type": "select", "label": "转存目标", "options": []gin.H{
|
{"key": "cloud.upload_provider", "type": "select", "label": "转存目标", "options": []gin.H{
|
||||||
{"value": "openlist", "label": "OpenList(推荐,可桥接 115/123/阿里/夸克)"},
|
{"value": "openlist", "label": "OpenList(推荐,可桥接 115/123/阿里等)"},
|
||||||
{"value": "clouddrive2", "label": "CloudDrive2(推荐,可桥接 115/123/阿里/夸克)"},
|
{"value": "clouddrive2", "label": "CloudDrive2(推荐,可桥接 115/123/阿里等)"},
|
||||||
{"value": "alist", "label": "Alist(可桥接多网盘)"},
|
{"value": "alist", "label": "Alist(可桥接多网盘)"},
|
||||||
{"value": "webdav", "label": "WebDAV"},
|
{"value": "webdav", "label": "WebDAV"},
|
||||||
{"value": "cloud115", "label": "115 原生(待接分片上传)"},
|
{"value": "cloud115", "label": "115 原生(待接分片上传)"},
|
||||||
{"value": "quark", "label": "夸克原生(待接分片上传)"},
|
|
||||||
}},
|
}},
|
||||||
{"key": "cloud.upload_source_dir", "type": "text", "label": "本地源目录"},
|
{"key": "cloud.upload_source_dir", "type": "text", "label": "本地源目录"},
|
||||||
{"key": "cloud.upload_dest_path", "type": "text", "label": "网盘目标目录"},
|
{"key": "cloud.upload_dest_path", "type": "text", "label": "网盘目标目录"},
|
||||||
|
|||||||
@@ -8,16 +8,25 @@ package handler
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/service"
|
"github.com/ShukeBta/MediaStationGo/internal/service"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const tasksLiveTorrentSnapshotMaxAge = 30 * time.Second
|
||||||
|
|
||||||
func tasksHandler(svc *service.Container) gin.HandlerFunc {
|
func tasksHandler(svc *service.Container) gin.HandlerFunc {
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
transcodes := svc.Transcoder.Active()
|
var transcodes []service.ActiveJob
|
||||||
_, torrents, _ := svc.Downloads.List(c.Request.Context())
|
if svc.Transcoder != nil {
|
||||||
|
transcodes = svc.Transcoder.Active()
|
||||||
|
}
|
||||||
|
var torrents []service.QBitTorrent
|
||||||
|
if svc.Downloads != nil {
|
||||||
|
torrents = svc.Downloads.LiveTorrentSnapshot(tasksLiveTorrentSnapshotMaxAge)
|
||||||
|
}
|
||||||
background := service.TaskSnapshot{}
|
background := service.TaskSnapshot{}
|
||||||
if svc.Tasks != nil {
|
if svc.Tasks != nil {
|
||||||
background = svc.Tasks.Snapshot()
|
background = svc.Tasks.Snapshot()
|
||||||
|
|||||||
@@ -21,6 +21,11 @@ const (
|
|||||||
CtxUserTier = "ctx_user_tier"
|
CtxUserTier = "ctx_user_tier"
|
||||||
CtxTokenPurpose = "ctx_token_purpose"
|
CtxTokenPurpose = "ctx_token_purpose"
|
||||||
CtxTokenMediaID = "ctx_token_media_id"
|
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.
|
// 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"})
|
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"code": 40304, "message": "token scope denied"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
syncAccessTokenCookie(c, raw, claims)
|
||||||
c.Set(CtxUserID, claims.UserID)
|
c.Set(CtxUserID, claims.UserID)
|
||||||
c.Set(CtxUserRole, claims.Role)
|
c.Set(CtxUserRole, claims.Role)
|
||||||
c.Set(CtxUserTier, claims.Tier)
|
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 {
|
func scopedTokenAllowedForRequest(c *gin.Context, claims *Claims) bool {
|
||||||
if claims == nil || strings.TrimSpace(claims.Purpose) == "" {
|
if claims == nil || strings.TrimSpace(claims.Purpose) == "" {
|
||||||
return true
|
return true
|
||||||
@@ -320,5 +368,8 @@ func extractToken(c *gin.Context) string {
|
|||||||
return value
|
return value
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if cookie, err := c.Cookie(AccessTokenCookieName); err == nil {
|
||||||
|
return strings.TrimSpace(cookie)
|
||||||
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,8 +4,10 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/golang-jwt/jwt/v5"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestCORSWildcardOriginAllowsProductionPreflight(t *testing.T) {
|
func TestCORSWildcardOriginAllowsProductionPreflight(t *testing.T) {
|
||||||
@@ -29,3 +31,161 @@ func TestCORSWildcardOriginAllowsProductionPreflight(t *testing.T) {
|
|||||||
t.Fatalf("Access-Control-Allow-Origin = %q, want *", got)
|
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"`
|
LastPlayAt *time.Time `gorm:"index" json:"last_play_at,omitempty"`
|
||||||
Warnings int `gorm:"default:0" json:"warnings"` // 指纹不匹配累计告警次数
|
Warnings int `gorm:"default:0" json:"warnings"` // 指纹不匹配累计告警次数
|
||||||
Kicked bool `gorm:"default:false" json:"kicked"` // 被一键踢下线(强制重新登录)
|
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。
|
// BeforeCreate 生成 UUID。
|
||||||
|
|||||||
+49
-46
@@ -51,10 +51,12 @@ type User struct {
|
|||||||
// ShareWarnings counts anti-account-sharing warnings, mainly device
|
// ShareWarnings counts anti-account-sharing warnings, mainly device
|
||||||
// fingerprint mismatches. Once it exceeds the configured threshold a
|
// fingerprint mismatches. Once it exceeds the configured threshold a
|
||||||
// re-offence disables the account until an admin re-enables it.
|
// re-offence disables the account until an admin re-enables it.
|
||||||
ShareWarnings int `gorm:"default:0" json:"share_warnings"`
|
ShareWarnings int `gorm:"default:0" json:"share_warnings"`
|
||||||
LastShareWarnAt *time.Time `json:"last_share_warn_at,omitempty"`
|
LastShareWarnAt *time.Time `json:"last_share_warn_at,omitempty"`
|
||||||
IsDefaultAdmin bool `gorm:"-" json:"is_default_admin,omitempty"`
|
IsDefaultAdmin bool `gorm:"-" json:"is_default_admin,omitempty"`
|
||||||
IsProtected bool `gorm:"-" json:"is_protected,omitempty"`
|
IsProtected bool `gorm:"-" json:"is_protected,omitempty"`
|
||||||
|
RealtimeOnline bool `gorm:"-" json:"realtime_online,omitempty"`
|
||||||
|
RealtimeDeviceCount int `gorm:"-" json:"realtime_device_count,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Library 表示用户定义的媒体根目录。
|
// Library 表示用户定义的媒体根目录。
|
||||||
@@ -73,6 +75,7 @@ type Media struct {
|
|||||||
SeriesID string `gorm:"index;size:128" json:"series_id,omitempty"`
|
SeriesID string `gorm:"index;size:128" json:"series_id,omitempty"`
|
||||||
Title string `gorm:"size:255;not null" json:"title"`
|
Title string `gorm:"size:255;not null" json:"title"`
|
||||||
OriginalName string `gorm:"size:255" json:"original_name,omitempty"`
|
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"`
|
Path string `gorm:"uniqueIndex;size:1024;not null" json:"path"`
|
||||||
SizeBytes int64 `json:"size_bytes"`
|
SizeBytes int64 `json:"size_bytes"`
|
||||||
DurationSec int `json:"duration_sec"`
|
DurationSec int `json:"duration_sec"`
|
||||||
@@ -200,17 +203,17 @@ type PlaylistItem struct {
|
|||||||
// DownloadTask 是待处理(或已完成)的 torrent / HTTP 下载。
|
// DownloadTask 是待处理(或已完成)的 torrent / HTTP 下载。
|
||||||
type DownloadTask struct {
|
type DownloadTask struct {
|
||||||
Base
|
Base
|
||||||
UserID string `gorm:"index;size:36" json:"user_id"`
|
UserID string `gorm:"index;size:36" json:"user_id"`
|
||||||
SubscriptionID string `gorm:"index;size:36" json:"subscription_id,omitempty"`
|
SubscriptionID string `gorm:"index;size:36" json:"subscription_id,omitempty"`
|
||||||
Source string `gorm:"size:32;not null" json:"source"` // qbittorrent / transmission / http
|
Source string `gorm:"size:32;not null" json:"source"` // qbittorrent / transmission / http
|
||||||
URL string `gorm:"size:2048;not null" json:"-"`
|
URL string `gorm:"size:2048;not null" json:"-"`
|
||||||
Title string `gorm:"size:512" json:"title,omitempty"`
|
Title string `gorm:"size:512" json:"title,omitempty"`
|
||||||
PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
|
PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
|
||||||
BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
|
BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
|
||||||
Overview string `gorm:"type:text" json:"overview,omitempty"`
|
Overview string `gorm:"type:text" json:"overview,omitempty"`
|
||||||
SavePath string `gorm:"size:1024" json:"save_path"`
|
SavePath string `gorm:"size:1024" json:"save_path"`
|
||||||
MediaType string `gorm:"size:16" json:"media_type,omitempty"`
|
MediaType string `gorm:"size:16" json:"media_type,omitempty"`
|
||||||
MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
|
MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
|
||||||
// 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。
|
// 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。
|
||||||
OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
|
OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
|
||||||
OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
|
OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
|
||||||
@@ -228,38 +231,38 @@ type DownloadTask struct {
|
|||||||
// Subscription 是自动化规则,轮询 RSS 源并将匹配种子排队到配置的下载客户端。
|
// Subscription 是自动化规则,轮询 RSS 源并将匹配种子排队到配置的下载客户端。
|
||||||
type Subscription struct {
|
type Subscription struct {
|
||||||
Base
|
Base
|
||||||
UserID string `gorm:"index;size:36" json:"user_id"`
|
UserID string `gorm:"index;size:36" json:"user_id"`
|
||||||
Name string `gorm:"size:128;not null" json:"name"`
|
Name string `gorm:"size:128;not null" json:"name"`
|
||||||
FeedURL string `gorm:"size:2048;not null" json:"feed_url"`
|
FeedURL string `gorm:"size:2048;not null" json:"feed_url"`
|
||||||
Filter string `gorm:"size:512" json:"filter"`
|
Filter string `gorm:"size:512" json:"filter"`
|
||||||
MediaType string `gorm:"size:16" json:"media_type,omitempty"`
|
MediaType string `gorm:"size:16" json:"media_type,omitempty"`
|
||||||
MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
|
MediaCategory string `gorm:"size:128" json:"media_category,omitempty"`
|
||||||
SavePath string `gorm:"size:1024" json:"save_path,omitempty"`
|
SavePath string `gorm:"size:1024" json:"save_path,omitempty"`
|
||||||
SearchMode string `gorm:"size:16;default:keyword" json:"search_mode,omitempty"` // keyword / imdb
|
SearchMode string `gorm:"size:16;default:keyword" json:"search_mode,omitempty"` // keyword / imdb
|
||||||
IMDBID string `gorm:"size:32" json:"imdb_id,omitempty"`
|
IMDBID string `gorm:"size:32" json:"imdb_id,omitempty"`
|
||||||
Source string `gorm:"size:32" json:"source,omitempty"`
|
Source string `gorm:"size:32" json:"source,omitempty"`
|
||||||
PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
|
PosterURL string `gorm:"size:2048" json:"poster_url,omitempty"`
|
||||||
BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
|
BackdropURL string `gorm:"size:2048" json:"backdrop_url,omitempty"`
|
||||||
Overview string `gorm:"type:text" json:"overview,omitempty"`
|
Overview string `gorm:"type:text" json:"overview,omitempty"`
|
||||||
// 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。
|
// 媒体展示元数据(用于 Telegram 富通知模板等):原始片名/语言/年份/评分/类型。
|
||||||
OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
|
OriginalName string `gorm:"size:512" json:"original_name,omitempty"`
|
||||||
OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
|
OriginalLanguage string `gorm:"size:32" json:"original_language,omitempty"`
|
||||||
Year int `json:"year,omitempty"`
|
Year int `json:"year,omitempty"`
|
||||||
Rating float32 `json:"rating,omitempty"`
|
Rating float32 `json:"rating,omitempty"`
|
||||||
Genres string `gorm:"size:255" json:"genres,omitempty"` // comma separated
|
Genres string `gorm:"size:255" json:"genres,omitempty"` // comma separated
|
||||||
Resolution string `gorm:"size:32" json:"resolution,omitempty"` // 2160p / 1080p / 720p / best
|
Resolution string `gorm:"size:32" json:"resolution,omitempty"` // 2160p / 1080p / 720p / best
|
||||||
Quality string `gorm:"size:64" json:"quality,omitempty"` // remux / bluray / web-dl / hdtv
|
Quality string `gorm:"size:64" json:"quality,omitempty"` // remux / bluray / web-dl / hdtv
|
||||||
Effects string `gorm:"size:128" json:"effects,omitempty"` // hdr,dolby-vision,atmos
|
Effects string `gorm:"size:128" json:"effects,omitempty"` // hdr,dolby-vision,atmos
|
||||||
ReleaseGroups string `gorm:"size:255" json:"release_groups,omitempty"` // comma separated
|
ReleaseGroups string `gorm:"size:255" json:"release_groups,omitempty"` // comma separated
|
||||||
ExcludeWords string `gorm:"size:255" json:"exclude_words,omitempty"` // comma separated
|
ExcludeWords string `gorm:"size:255" json:"exclude_words,omitempty"` // comma separated
|
||||||
WashEnabled bool `gorm:"default:false" json:"wash_enabled"`
|
WashEnabled bool `gorm:"default:false" json:"wash_enabled"`
|
||||||
WashPriority string `gorm:"size:32" json:"wash_priority,omitempty"` // balanced / resolution / quality / effects / seeders
|
WashPriority string `gorm:"size:32" json:"wash_priority,omitempty"` // balanced / resolution / quality / effects / seeders
|
||||||
TotalEpisodes int `gorm:"default:0" json:"total_episodes,omitempty"`
|
TotalEpisodes int `gorm:"default:0" json:"total_episodes,omitempty"`
|
||||||
Priority int `gorm:"default:50" json:"priority,omitempty"` // lower is earlier when schedulers sort later
|
Priority int `gorm:"default:50" json:"priority,omitempty"` // lower is earlier when schedulers sort later
|
||||||
Enabled bool `gorm:"default:true" json:"enabled"`
|
Enabled bool `gorm:"default:true" json:"enabled"`
|
||||||
LastRunAt *time.Time `json:"last_run_at,omitempty"`
|
LastRunAt *time.Time `json:"last_run_at,omitempty"`
|
||||||
ArchivedAt *time.Time `gorm:"index" json:"archived_at,omitempty"`
|
ArchivedAt *time.Time `gorm:"index" json:"archived_at,omitempty"`
|
||||||
ArchiveReason string `gorm:"size:255" json:"archive_reason,omitempty"`
|
ArchiveReason string `gorm:"size:255" json:"archive_reason,omitempty"`
|
||||||
|
|
||||||
DownloadedEpisodes int `gorm:"-" json:"downloaded_episodes,omitempty"`
|
DownloadedEpisodes int `gorm:"-" json:"downloaded_episodes,omitempty"`
|
||||||
LocalMediaCount int `gorm:"-" json:"local_media_count,omitempty"`
|
LocalMediaCount int `gorm:"-" json:"local_media_count,omitempty"`
|
||||||
|
|||||||
@@ -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"
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
"gorm.io/gorm"
|
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||||
@@ -88,13 +86,7 @@ func TestAdultProviderUsesConfiguredMultipleSources(t *testing.T) {
|
|||||||
}))
|
}))
|
||||||
defer good.Close()
|
defer good.Close()
|
||||||
|
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db := newServiceTestDB(t, &model.APIConfig{})
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
apiConfig := NewAPIConfigService(zap.NewNop(), repository.New(db), NewCryptoService("", zap.NewNop()))
|
apiConfig := NewAPIConfigService(zap.NewNop(), repository.New(db), NewCryptoService("", zap.NewNop()))
|
||||||
baseURL := bad.URL + "\n" + good.URL
|
baseURL := bad.URL + "\n" + good.URL
|
||||||
if _, err := apiConfig.Update(context.Background(), "adult", APIConfigPatch{BaseURL: &baseURL}); err != nil {
|
if _, err := apiConfig.Update(context.Background(), "adult", APIConfigPatch{BaseURL: &baseURL}); err != nil {
|
||||||
|
|||||||
@@ -4,9 +4,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
"gorm.io/gorm"
|
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/config"
|
"github.com/ShukeBta/MediaStationGo/internal/config"
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||||
@@ -14,13 +12,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestAIStatusUsesDatabaseOpenAIConfig(t *testing.T) {
|
func TestAIStatusUsesDatabaseOpenAIConfig(t *testing.T) {
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db := newServiceTestDB(t, &model.APIConfig{})
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
repo := &repository.Container{DB: db}
|
repo := &repository.Container{DB: db}
|
||||||
crypto := NewCryptoService("test-secret", zap.NewNop())
|
crypto := NewCryptoService("test-secret", zap.NewNop())
|
||||||
apiConfig := NewAPIConfigService(zap.NewNop(), repo, crypto)
|
apiConfig := NewAPIConfigService(zap.NewNop(), repo, crypto)
|
||||||
@@ -52,13 +44,7 @@ func TestAIStatusUsesDatabaseOpenAIConfig(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestAIStatusHonorsDisabledDatabaseOpenAIConfig(t *testing.T) {
|
func TestAIStatusHonorsDisabledDatabaseOpenAIConfig(t *testing.T) {
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db := newServiceTestDB(t, &model.APIConfig{})
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := db.AutoMigrate(&model.APIConfig{}); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
repo := &repository.Container{DB: db}
|
repo := &repository.Container{DB: db}
|
||||||
apiConfig := NewAPIConfigService(zap.NewNop(), repo, NewCryptoService("test-secret", zap.NewNop()))
|
apiConfig := NewAPIConfigService(zap.NewNop(), repo, NewCryptoService("test-secret", zap.NewNop()))
|
||||||
key := "sk-test"
|
key := "sk-test"
|
||||||
|
|||||||
@@ -9,11 +9,9 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
|
||||||
"github.com/golang-jwt/jwt/v5"
|
"github.com/golang-jwt/jwt/v5"
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
"golang.org/x/crypto/bcrypt"
|
"golang.org/x/crypto/bcrypt"
|
||||||
"gorm.io/gorm"
|
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/config"
|
"github.com/ShukeBta/MediaStationGo/internal/config"
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/database"
|
"github.com/ShukeBta/MediaStationGo/internal/database"
|
||||||
@@ -23,13 +21,7 @@ import (
|
|||||||
|
|
||||||
func newAuthTestServices(t *testing.T) (*repository.Container, *AuthService, *ProfileService, *PermissionService) {
|
func newAuthTestServices(t *testing.T) (*repository.Container, *AuthService, *ProfileService, *PermissionService) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db := newServiceTestDB(t, &model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.TelegramBinding{}, &model.Setting{})
|
||||||
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)
|
|
||||||
}
|
|
||||||
sqlDB, err := db.DB()
|
sqlDB, err := db.DB()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ func (c *Container) BootCloudStorageHealthCheck(ctx context.Context) {
|
|||||||
|
|
||||||
cloudConfigs := make([]StorageView, 0)
|
cloudConfigs := make([]StorageView, 0)
|
||||||
for _, cfg := range configs {
|
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)
|
cloudConfigs = append(cloudConfigs, cfg)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -92,7 +92,7 @@ func cloudStorageMissingConfigReason(err error) string {
|
|||||||
}
|
}
|
||||||
msg := strings.ToLower(strings.TrimSpace(err.Error()))
|
msg := strings.ToLower(strings.TrimSpace(err.Error()))
|
||||||
switch {
|
switch {
|
||||||
case strings.Contains(msg, "missing cookie"):
|
case strings.Contains(msg, "missing cookie") || (strings.Contains(msg, "missing") && strings.Contains(msg, "cookie")):
|
||||||
return "missing_cookie"
|
return "missing_cookie"
|
||||||
case strings.Contains(msg, "missing webdav url"):
|
case strings.Contains(msg, "missing webdav url"):
|
||||||
return "missing_webdav_url"
|
return "missing_webdav_url"
|
||||||
|
|||||||
@@ -5,10 +5,8 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
"go.uber.org/zap/zaptest/observer"
|
"go.uber.org/zap/zaptest/observer"
|
||||||
"gorm.io/gorm"
|
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
"github.com/ShukeBta/MediaStationGo/internal/repository"
|
||||||
@@ -20,8 +18,9 @@ func TestCloudStorageMissingConfigReason(t *testing.T) {
|
|||||||
want string
|
want string
|
||||||
}{
|
}{
|
||||||
{errors.New("115: missing cookie"), "missing_cookie"},
|
{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("clouddrive2: missing WebDAV URL"), "missing_webdav_url"},
|
||||||
{errors.New("quark: token expired"), ""},
|
{errors.New("openlist: token expired"), ""},
|
||||||
}
|
}
|
||||||
for _, tc := range cases {
|
for _, tc := range cases {
|
||||||
if got := cloudStorageMissingConfigReason(tc.err); got != tc.want {
|
if got := cloudStorageMissingConfigReason(tc.err); got != tc.want {
|
||||||
@@ -31,19 +30,13 @@ func TestCloudStorageMissingConfigReason(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWarnMissingCloudStorageConfigOncePersistsMarker(t *testing.T) {
|
func TestWarnMissingCloudStorageConfigOncePersistsMarker(t *testing.T) {
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db := newServiceTestDB(t, &model.Setting{})
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := db.AutoMigrate(&model.Setting{}); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
core, observed := observer.New(zap.WarnLevel)
|
core, observed := observer.New(zap.WarnLevel)
|
||||||
c := &Container{
|
c := &Container{
|
||||||
Log: zap.New(core),
|
Log: zap.New(core),
|
||||||
Repo: repository.New(db),
|
Repo: repository.New(db),
|
||||||
}
|
}
|
||||||
err = errors.New("115: missing cookie")
|
err := errors.New("115: missing cookie")
|
||||||
|
|
||||||
if !c.warnMissingCloudStorageConfigOnce(context.Background(), "cloud115", err) {
|
if !c.warnMissingCloudStorageConfigOnce(context.Background(), "cloud115", err) {
|
||||||
t.Fatal("missing config should be handled")
|
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 (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
"gorm.io/gorm"
|
|
||||||
|
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/config"
|
"github.com/ShukeBta/MediaStationGo/internal/config"
|
||||||
"github.com/ShukeBta/MediaStationGo/internal/model"
|
"github.com/ShukeBta/MediaStationGo/internal/model"
|
||||||
@@ -18,13 +15,7 @@ import (
|
|||||||
|
|
||||||
func newBotTestService(t *testing.T) (*repository.Container, *TelegramBotService) {
|
func newBotTestService(t *testing.T) (*repository.Container, *TelegramBotService) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db := newServiceTestDB(t, model.AllModels()...)
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := db.AutoMigrate(model.AllModels()...); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
repos := repository.New(db)
|
repos := repository.New(db)
|
||||||
cfg := &config.Config{}
|
cfg := &config.Config{}
|
||||||
cfg.Secrets.JWTSecret = "test-secret"
|
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) {
|
func TestBotRegistrationCommandUsesOpenRegQuota(t *testing.T) {
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
repos, bot := newBotTestService(t)
|
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) {
|
func TestBotAdminCodeAndUserCommands(t *testing.T) {
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
repos, bot := newBotTestService(t)
|
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)
|
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.
|
// Provider types recognised by the registry.
|
||||||
const (
|
const (
|
||||||
TypeQuark = "quark" // 夸克网盘
|
|
||||||
Type115 = "cloud115" // 115 网盘
|
Type115 = "cloud115" // 115 网盘
|
||||||
TypeCloudDrive2 = "clouddrive2" // CloudDrive2 桥接网盘
|
TypeCloudDrive2 = "clouddrive2" // CloudDrive2 桥接网盘
|
||||||
TypeOpenList = "openlist" // OpenList / AList-compatible bridge
|
TypeOpenList = "openlist" // OpenList / AList-compatible bridge
|
||||||
@@ -42,7 +41,7 @@ type FileEntry struct {
|
|||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
IsDir bool `json:"is_dir"`
|
IsDir bool `json:"is_dir"`
|
||||||
Size int64 `json:"size"`
|
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"`
|
PickCode string `json:"pick_code,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -59,7 +58,7 @@ type DirectLink struct {
|
|||||||
|
|
||||||
// Provider is the common cloud-disk interface.
|
// Provider is the common cloud-disk interface.
|
||||||
type Provider interface {
|
type Provider interface {
|
||||||
// Type returns the provider key (TypeQuark / Type115).
|
// Type returns the provider key.
|
||||||
Type() string
|
Type() string
|
||||||
// Ping validates the stored credentials (cookie). Cheap, used by the
|
// Ping validates the stored credentials (cookie). Cheap, used by the
|
||||||
// storage-config Test() probe.
|
// storage-config Test() probe.
|
||||||
@@ -71,6 +70,21 @@ type Provider interface {
|
|||||||
Resolve(ctx context.Context, fileRef string) (*DirectLink, error)
|
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
|
// 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
|
// (as persisted by StorageConfigService). The client is shared so callers can
|
||||||
// inject timeouts / test transports.
|
// inject timeouts / test transports.
|
||||||
@@ -79,8 +93,6 @@ func New(typ string, cfg map[string]any, client *http.Client) (Provider, error)
|
|||||||
client = http.DefaultClient
|
client = http.DefaultClient
|
||||||
}
|
}
|
||||||
switch typ {
|
switch typ {
|
||||||
case TypeQuark:
|
|
||||||
return newQuark(cfg, client), nil
|
|
||||||
case Type115:
|
case Type115:
|
||||||
return new115(cfg, client), nil
|
return new115(cfg, client), nil
|
||||||
case TypeCloudDrive2:
|
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.
|
// IsCloudType reports whether typ is a cloud-disk provider.
|
||||||
func IsCloudType(typ string) bool {
|
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.
|
// 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"
|
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 (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/base64"
|
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"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) {
|
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)
|
pan115 := new115(map[string]any{"cookie": "UID=1; CID=2", "force_proxy": "true"}, http.DefaultClient)
|
||||||
if pan115.proxy {
|
if pan115.proxy {
|
||||||
t.Fatalf("115 should keep safe direct mode; force_proxy is deprecated")
|
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) {
|
func TestCloudDrive2WebDAVListAndResolve(t *testing.T) {
|
||||||
var gotAuth, gotDepth, gotRange string
|
var gotAuth, gotDepth, gotRange string
|
||||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
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) {
|
func TestOpenListListAPIFailureDoesNotFallbackToWebDAV(t *testing.T) {
|
||||||
var davSeen bool
|
var davSeen bool
|
||||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
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 {
|
if _, err := New("dropbox", nil, nil); err != ErrUnsupported {
|
||||||
t.Fatalf("want ErrUnsupported, got %v", err)
|
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 {
|
if IsCloudType("quark") {
|
||||||
found := false
|
t.Fatal("quark should not be an active cloud provider")
|
||||||
for i := 0; i+len(sub) <= len(s); i++ {
|
|
||||||
if s[i:i+len(sub)] == sub {
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var _ = time.Second
|
var _ = time.Second
|
||||||
|
|||||||
@@ -1,25 +1,19 @@
|
|||||||
package cloud
|
package cloud
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"context"
|
"context"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"encoding/json"
|
|
||||||
"encoding/xml"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
"path"
|
"path"
|
||||||
"sort"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
// cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint.
|
// cloudDrive2Provider bridges CloudDrive2 through its WebDAV endpoint.
|
||||||
//
|
//
|
||||||
// CloudDrive2 already integrates many cloud disks (115 / 123 / Aliyun / Quark
|
// CloudDrive2 integrates many cloud disks (115 / 123 / Aliyun and more).
|
||||||
// and more). Treating it as a WebDAV-backed cloud provider lets MediaStationGo
|
// Treating it as a WebDAV-backed cloud provider lets MediaStationGo
|
||||||
// browse, mount and upload to those disks without carrying every provider's
|
// browse, mount and upload to those disks without carrying every provider's
|
||||||
// private chunk-upload protocol in this project.
|
// private chunk-upload protocol in this project.
|
||||||
type cloudDrive2Provider struct {
|
type cloudDrive2Provider struct {
|
||||||
@@ -76,133 +70,6 @@ func (p *cloudDrive2Provider) Ping(ctx context.Context) error {
|
|||||||
return err
|
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) {
|
func (p *cloudDrive2Provider) Resolve(ctx context.Context, fileRef string) (*DirectLink, error) {
|
||||||
if err := p.validate(); err != nil {
|
if err := p.validate(); err != nil {
|
||||||
return nil, err
|
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
|
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 {
|
func (p *cloudDrive2Provider) validate() error {
|
||||||
if p.base == nil || p.base.Scheme == "" || p.base.Host == "" {
|
if p.base == nil || p.base.Scheme == "" || p.base.Host == "" {
|
||||||
return fmt.Errorf("%s: missing WebDAV URL", p.name)
|
return fmt.Errorf("%s: missing WebDAV URL", p.name)
|
||||||
@@ -563,17 +113,6 @@ func (p *cloudDrive2Provider) validate() error {
|
|||||||
return nil
|
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 {
|
func webDAVURLFromConfig(cfg map[string]any, defaultDAVPath string) string {
|
||||||
rawURL := str(cfg["url"])
|
rawURL := str(cfg["url"])
|
||||||
if rawURL == "" {
|
if rawURL == "" {
|
||||||
@@ -669,153 +208,6 @@ func ensureDefaultDAVPath(rawURL, defaultDAVPath string) string {
|
|||||||
return rawURL
|
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 {
|
func normalizeCloudDAVPath(p string) string {
|
||||||
p = strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
|
p = strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
|
||||||
if p == "" || p == "." {
|
if p == "" || p == "." {
|
||||||
@@ -835,20 +227,6 @@ func sameCloudDAVPath(a, b string) bool {
|
|||||||
return strings.TrimRight(normalizeCloudDAVPath(a), "/") == strings.TrimRight(normalizeCloudDAVPath(b), "/")
|
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 {
|
func firstNonEmpty(values ...string) string {
|
||||||
for _, v := range values {
|
for _, v := range values {
|
||||||
if strings.TrimSpace(v) != "" {
|
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