feat: merge conflict resolution, site management, UI fixes

This commit is contained in:
ShukeBta
2026-05-16 17:57:34 +08:00
parent 0be8c7952e
commit cbb4b806be
84 changed files with 13826 additions and 202 deletions
+37 -31
View File
@@ -1,17 +1,10 @@
// Package config loads layered configuration from defaults, config files and
// environment variables, mirroring the conventions used by nowen-video.
// Package config 加载分层配置:默认值、配置文件和环境变量。
//
// Priority (low -> high):
// 1. Built-in defaults
// 2. config.yaml in the working directory (nested format)
// 3. config/*.yaml shard files (per-module)
// 4. Environment variables prefixed with MEDIASTATION_
//
// Environment variable example:
//
// MEDIASTATION_APP_PORT=8080
// MEDIASTATION_SECRETS_JWT_SECRET=please-change-me
// MEDIASTATION_DATABASE_DB_PATH=/data/mediastation.db
// 优先级(低 -> 高):
// 1. 内置默认值
// 2. 工作目录中的 config.yaml(嵌套格式)
// 3. config/*.yaml 分片文件(按模块)
// 4. 以 MEDIASTATION_ 为前缀的环境变量
package config
import (
@@ -25,10 +18,10 @@ import (
"github.com/spf13/viper"
)
// EnvPrefix is the prefix used for all env-var-driven overrides.
// EnvPrefix 是所有环境变量驱动的覆盖使用的前缀。
const EnvPrefix = "MEDIASTATION"
// Config is the root config aggregate.
// Config 是根配置聚合。
type Config struct {
App AppConfig `mapstructure:"app"`
Database DatabaseConfig `mapstructure:"database"`
@@ -38,9 +31,18 @@ type Config struct {
Media MediaConfig `mapstructure:"media"`
Transcoder TranscoderConfig `mapstructure:"transcoder"`
AI AIConfig `mapstructure:"ai"`
ApiConfig ApiConfigConfig `mapstructure:"api_config"`
}
// TranscoderConfig controls the HLS / ffmpeg backend.
// ApiConfigConfig API 配置相关设置。
type ApiConfigConfig struct {
// AutoEncrypt 是否自动加密敏感字段
AutoEncrypt bool `mapstructure:"auto_encrypt"`
// DefaultTimeout 默认请求超时(秒)
DefaultTimeout int `mapstructure:"default_timeout"`
}
// TranscoderConfig 控制 HLS / ffmpeg 后端。
type TranscoderConfig struct {
Encoder string `mapstructure:"encoder"` // "" / nvenc / qsv / vaapi
Preset string `mapstructure:"preset"`
@@ -51,7 +53,7 @@ type TranscoderConfig struct {
SegmentSeconds int `mapstructure:"segment_seconds"`
}
// AppConfig holds runtime app parameters.
// AppConfig 保存运行时应用参数。
type AppConfig struct {
Port int `mapstructure:"port"`
Debug bool `mapstructure:"debug"`
@@ -65,7 +67,7 @@ type AppConfig struct {
ServerURL string `mapstructure:"server_url"`
}
// DatabaseConfig configures GORM + SQLite.
// DatabaseConfig 配置 GORM + SQLite。
type DatabaseConfig struct {
DBPath string `mapstructure:"db_path"`
WALMode bool `mapstructure:"wal_mode"`
@@ -75,7 +77,7 @@ type DatabaseConfig struct {
MaxIdleConns int `mapstructure:"max_idle_conns"`
}
// SecretsConfig holds JWT / 3rd-party API keys (do NOT commit values).
// SecretsConfig 保存 JWT / 第三方 API 密钥(不要提交值)。
type SecretsConfig struct {
JWTSecret string `mapstructure:"jwt_secret"`
TMDbAPIKey string `mapstructure:"tmdb_api_key"`
@@ -85,9 +87,11 @@ type SecretsConfig struct {
TheTVDBAPIKey string `mapstructure:"thetvdb_api_key"`
FanartAPIKey string `mapstructure:"fanart_tv_api_key"`
DoubanCookie string `mapstructure:"douban_cookie"`
// 用于加密的密钥,如果为空则使用 JWTSecret
EncryptionKey string `mapstructure:"encryption_key"`
}
// LoggingConfig configures Zap.
// LoggingConfig 配置 Zap。
type LoggingConfig struct {
Level string `mapstructure:"level"`
Format string `mapstructure:"format"`
@@ -98,7 +102,7 @@ type LoggingConfig struct {
MaxBackups int `mapstructure:"max_backups"`
}
// CacheConfig controls the on-disk transcode/scrape cache.
// CacheConfig 控制磁盘转码/刮削缓存。
type CacheConfig struct {
CacheDir string `mapstructure:"cache_dir"`
MaxDiskUsageMB int `mapstructure:"max_disk_usage_mb"`
@@ -107,14 +111,14 @@ type CacheConfig struct {
CleanupIntervalMin int `mapstructure:"cleanup_interval_min"`
}
// MediaConfig holds default library locations (used by the bootstrap library).
// MediaConfig 保存默认库位置(用于引导库)。
type MediaConfig struct {
MoviesDir string `mapstructure:"movies_dir"`
TVDir string `mapstructure:"tv_dir"`
AnimeDir string `mapstructure:"anime_dir"`
}
// AIConfig configures the optional LLM provider.
// AIConfig 配置可选的 LLM 提供者。
type AIConfig struct {
Enabled bool `mapstructure:"enabled"`
Provider string `mapstructure:"provider"`
@@ -125,9 +129,9 @@ type AIConfig struct {
MaxConcurrent int `mapstructure:"max_concurrent"`
}
// Load reads configuration from defaults / files / environment.
// Load 从默认值 / 文件 / 环境读取配置。
//
// It always returns a usable Config, even if no files are present.
// 即使没有文件也始终返回可用的 Config。
func Load() (*Config, error) {
v := viper.New()
setDefaults(v)
@@ -143,7 +147,7 @@ func Load() (*Config, error) {
}
}
// Merge sharded files under ./config/*.yaml.
// 合并 ./config/*.yaml 下的分片文件。
if entries, err := os.ReadDir("config"); err == nil {
for _, e := range entries {
if e.IsDir() || !strings.HasSuffix(e.Name(), ".yaml") {
@@ -215,9 +219,13 @@ func setDefaults(v *viper.Viper) {
v.SetDefault("transcoder.buf_size", "3000k")
v.SetDefault("transcoder.max_height", 720)
v.SetDefault("transcoder.segment_seconds", 4)
// API Config 默认设置
v.SetDefault("api_config.auto_encrypt", true)
v.SetDefault("api_config.default_timeout", 30)
}
// normalize fills derived defaults and self-heals empty critical fields.
// normalize 填充派生默认值并自愈空的关键字段。
func (c *Config) normalize() error {
if c.App.DataDir == "" {
c.App.DataDir = "./data"
@@ -229,8 +237,7 @@ func (c *Config) normalize() error {
c.Cache.CacheDir = filepath.Join(c.App.DataDir, "cache")
}
if c.Secrets.JWTSecret == "" {
// Persist an auto-generated secret to keep sessions stable across
// restarts even when the operator forgot to configure one.
// 持久化自动生成的密钥以在操作员忘记配置时保持会话稳定。
path := filepath.Join(c.App.DataDir, ".jwt_secret")
if data, err := os.ReadFile(path); err == nil && len(data) > 0 {
c.Secrets.JWTSecret = strings.TrimSpace(string(data))
@@ -247,8 +254,7 @@ func (c *Config) normalize() error {
return nil
}
// asConfigFileNotFound is a small helper around errors.As that avoids importing
// errors in this short file.
// asConfigFileNotFound 是 errors.As 的小辅助函数,避免在这个短文件中导入 errors。
func asConfigFileNotFound(err error, target *viper.ConfigFileNotFoundError) bool {
if err == nil {
return false
+186
View File
@@ -0,0 +1,186 @@
// Package handler — API 配置 HTTP Handler。
package handler
import (
"net/http"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// ApiConfigHandler API 配置 HTTP 处理。
type ApiConfigHandler struct {
svc *service.Container
log *zap.Logger
}
// NewApiConfigHandler 创建 API 配置处理器。
func NewApiConfigHandler(svc *service.Container, log *zap.Logger) *ApiConfigHandler {
return &ApiConfigHandler{svc: svc, log: log}
}
// ListApiConfigs 获取所有 API 配置。
// GET /api/api-config
func (h *ApiConfigHandler) ListApiConfigs(c *gin.Context) {
configs, err := h.svc.ApiConfig.List(c.Request.Context())
if err != nil {
h.log.Error("list api configs failed", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
// 遮蔽 API Key
for i := range configs {
if configs[i].APIKey != "" {
configs[i].APIKey = h.svc.ApiConfig.MaskAPIKey(configs[i].APIKey)
}
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": configs})
}
// ListProviders 获取预定义的提供者列表。
// GET /api/api-config/providers/list
func (h *ApiConfigHandler) ListProviders(c *gin.Context) {
providers := h.svc.ApiConfig.GetProviders()
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": providers})
}
// GetApiConfig 获取指定提供者的配置。
// GET /api/api-config/:provider
func (h *ApiConfigHandler) GetApiConfig(c *gin.Context) {
provider := c.Param("provider")
if provider == "" {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "provider required", "data": nil})
return
}
cfg, err := h.svc.ApiConfig.GetByProvider(c.Request.Context(), provider)
if err != nil {
if err == service.ErrApiConfigNotFound {
c.JSON(http.StatusNotFound, gin.H{"code": 40401, "message": "api config not found", "data": nil})
return
}
h.log.Error("get api config failed", zap.Error(err), zap.String("provider", provider))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
// 遮蔽 API Key
if cfg.APIKey != "" {
cfg.APIKey = h.svc.ApiConfig.MaskAPIKey(cfg.APIKey)
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": cfg})
}
// GetEffectiveConfig 获取生效的配置(数据库配置优先于配置文件)。
// GET /api/api-config/:provider/effective
func (h *ApiConfigHandler) GetEffectiveConfig(c *gin.Context) {
provider := c.Param("provider")
if provider == "" {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "provider required", "data": nil})
return
}
cfg, err := h.svc.ApiConfig.GetEffectiveConfig(c.Request.Context(), provider)
if err != nil {
if err == service.ErrApiConfigNotFound {
c.JSON(http.StatusNotFound, gin.H{"code": 40401, "message": "api config not found", "data": nil})
return
}
h.log.Error("get effective config failed", zap.Error(err), zap.String("provider", provider))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
// 遮蔽 API Key
if cfg.APIKey != "" {
cfg.APIKey = h.svc.ApiConfig.MaskAPIKey(cfg.APIKey)
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": cfg})
}
// UpsertApiConfig 创建或更新 API 配置。
// POST /api/api-config/:provider
func (h *ApiConfigHandler) UpsertApiConfig(c *gin.Context) {
provider := c.Param("provider")
if provider == "" {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "provider required", "data": nil})
return
}
var req struct {
APIKey string `json:"api_key"`
BaseURL string `json:"base_url"`
Extra string `json:"extra"`
Enabled bool `json:"enabled"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "invalid request", "data": nil})
return
}
cfg, err := h.svc.ApiConfig.Upsert(c.Request.Context(), provider, req.APIKey, req.BaseURL, req.Extra, req.Enabled)
if err != nil {
if err == service.ErrInvalidProvider {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "invalid provider", "data": nil})
return
}
h.log.Error("upsert api config failed", zap.Error(err), zap.String("provider", provider))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
// 返回遮蔽后的配置
cfg.APIKey = h.svc.ApiConfig.MaskAPIKey(cfg.APIKey)
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": cfg})
}
// DeleteApiConfig 删除 API 配置。
// DELETE /api/api-config/:provider
func (h *ApiConfigHandler) DeleteApiConfig(c *gin.Context) {
provider := c.Param("provider")
if provider == "" {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "provider required", "data": nil})
return
}
if err := h.svc.ApiConfig.Delete(c.Request.Context(), provider); err != nil {
h.log.Error("delete api config failed", zap.Error(err), zap.String("provider", provider))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": nil})
}
// TestApiConfig 测试 API 连接。
// POST /api/api-config/:provider/test
func (h *ApiConfigHandler) TestApiConfig(c *gin.Context) {
provider := c.Param("provider")
if provider == "" {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "provider required", "data": nil})
return
}
result, err := h.svc.ApiConfig.TestConnection(c.Request.Context(), provider)
if err != nil {
h.log.Debug("test api config failed", zap.Error(err), zap.String("provider", provider))
// 不返回错误,只返回测试结果
}
// 更新测试结果
_ = h.svc.ApiConfig.UpdateTestResult(c.Request.Context(), provider, result)
c.JSON(http.StatusOK, gin.H{
"code": 0,
"message": "ok",
"data": gin.H{
"result": result,
},
})
}
+13 -6
View File
@@ -28,20 +28,24 @@ func loginHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
u, token, err := svc.Auth.Login(c.Request.Context(), req.Username, req.Password)
resp, err := svc.Auth.Login(c.Request.Context(), req.Username, req.Password)
if err != nil {
if errors.Is(err, service.ErrInvalidCredentials) {
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid credentials"})
return
}
if errors.Is(err, service.ErrUserInactive) {
c.JSON(http.StatusForbidden, gin.H{"error": "user account is inactive"})
return
}
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{
"token": token,
"user": u,
"user": resp.User,
"tokens": resp.Tokens,
})
svc.Audit.Record(c.Request.Context(), u.ID, "auth.login", u.Username, c.ClientIP(), "")
svc.Audit.Record(c.Request.Context(), resp.User.ID, "auth.login", resp.User.Username, c.ClientIP(), "")
}
}
@@ -52,7 +56,7 @@ func registerHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
u, err := svc.Auth.Register(c.Request.Context(), req.Username, req.Password)
u, tokens, err := svc.Auth.Register(c.Request.Context(), req.Username, req.Password)
if err != nil {
if errors.Is(err, service.ErrUsernameTaken) {
c.JSON(http.StatusConflict, gin.H{"error": "username taken"})
@@ -61,7 +65,10 @@ func registerHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, u)
c.JSON(http.StatusOK, gin.H{
"user": u,
"tokens": tokens,
})
}
}
+235
View File
@@ -0,0 +1,235 @@
// Package handler — 下载客户端管理 HTTP 端点。
package handler
import (
"encoding/json"
"net/http"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// DownloadClientHandler 处理下载客户端的 CRUD 操作。
type DownloadClientHandler struct {
svc *service.Container
log *zap.Logger
}
// NewDownloadClientHandler 创建下载客户端处理器。
func NewDownloadClientHandler(svc *service.Container, log *zap.Logger) *DownloadClientHandler {
return &DownloadClientHandler{svc: svc, log: log}
}
// downloadClientCreateRequest 创建下载客户端请求体。
type downloadClientCreateRequest struct {
Name string `json:"name" binding:"required"`
Type string `json:"type" binding:"required,oneof=qbittorrent transmission aria2"`
Host string `json:"host" binding:"required"`
Username string `json:"username"`
Password string `json:"password"`
IsDefault bool `json:"is_default"`
Extra map[string]string `json:"extra,omitempty"`
}
// downloadClientUpdateRequest 更新下载客户端请求体。
type downloadClientUpdateRequest struct {
Name string `json:"name"`
Type string `json:"type" binding:"omitempty,oneof=qbittorrent transmission aria2"`
Host string `json:"host"`
Username string `json:"username"`
Password string `json:"password"`
IsDefault *bool `json:"is_default"`
Enabled *bool `json:"enabled"`
Extra map[string]string `json:"extra,omitempty"`
}
// Create 创建新的下载客户端。
func (h *DownloadClientHandler) Create(c *gin.Context) {
var req downloadClientCreateRequest
if err := c.ShouldBindJSON(&req); err != nil {
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
return
}
ctx := c.Request.Context()
// 加密密码
password := req.Password
if password != "" && h.svc.Crypto != nil {
password = h.svc.Crypto.Encrypt(password)
}
// 加密 Extra 配置
extraStr := ""
if len(req.Extra) > 0 {
extraJSON, _ := json.Marshal(req.Extra)
extraStr = string(extraJSON)
if h.svc.Crypto != nil {
extraStr = h.svc.Crypto.Encrypt(extraStr)
}
}
// 如果设为默认,先清除其他默认
if req.IsDefault {
_ = h.svc.Repo.DownloadClient.ClearDefault(ctx)
}
client := &model.DownloadClient{
Name: req.Name,
Type: req.Type,
Host: req.Host,
Username: req.Username,
Password: password,
IsDefault: req.IsDefault,
Enabled: true,
Extra: extraStr,
}
if err := h.svc.Repo.DownloadClient.Create(ctx, client); err != nil {
Error(c, http.StatusInternalServerError, ErrInternal, "创建失败: "+err.Error())
return
}
// 热插拔:加载新客户端
go func() {
if initErr := h.svc.DownloadMgr.AddClient(ctx, client); initErr != nil {
h.log.Warn("failed to hot-add download client", zap.Error(initErr))
}
}()
Success(c, client)
}
// List 返回所有下载客户端。
func (h *DownloadClientHandler) List(c *gin.Context) {
ctx := c.Request.Context()
clients, err := h.svc.Repo.DownloadClient.List(ctx)
if err != nil {
Error(c, http.StatusInternalServerError, ErrInternal, "查询失败")
return
}
Success(c, clients)
}
// Get 返回指定下载客户端详情。
func (h *DownloadClientHandler) Get(c *gin.Context) {
id := c.Param("id")
ctx := c.Request.Context()
client, err := h.svc.Repo.DownloadClient.FindByID(ctx, id)
if err != nil || client == nil {
Error(c, http.StatusNotFound, ErrNotFound, "客户端不存在")
return
}
Success(c, client)
}
// Update 更新下载客户端。
func (h *DownloadClientHandler) Update(c *gin.Context) {
id := c.Param("id")
var req downloadClientUpdateRequest
if err := c.ShouldBindJSON(&req); err != nil {
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
return
}
ctx := c.Request.Context()
client, err := h.svc.Repo.DownloadClient.FindByID(ctx, id)
if err != nil || client == nil {
Error(c, http.StatusNotFound, ErrNotFound, "客户端不存在")
return
}
if req.Name != "" {
client.Name = req.Name
}
if req.Type != "" {
client.Type = req.Type
}
if req.Host != "" {
client.Host = req.Host
}
if req.Username != "" {
client.Username = req.Username
}
if req.Password != "" {
if h.svc.Crypto != nil {
client.Password = h.svc.Crypto.Encrypt(req.Password)
} else {
client.Password = req.Password
}
}
if req.IsDefault != nil && *req.IsDefault {
_ = h.svc.Repo.DownloadClient.ClearDefault(ctx)
client.IsDefault = *req.IsDefault
}
if req.Enabled != nil {
client.Enabled = *req.Enabled
}
if len(req.Extra) > 0 {
extraJSON, _ := json.Marshal(req.Extra)
extraStr := string(extraJSON)
if h.svc.Crypto != nil {
client.Extra = h.svc.Crypto.Encrypt(extraStr)
} else {
client.Extra = extraStr
}
}
if err := h.svc.Repo.DownloadClient.Update(ctx, client); err != nil {
Error(c, http.StatusInternalServerError, ErrInternal, "更新失败")
return
}
// 热更新适配器
go func() {
if updateErr := h.svc.DownloadMgr.UpdateClient(ctx, client); updateErr != nil {
h.log.Warn("failed to hot-update download client", zap.Error(updateErr))
}
}()
Success(c, client)
}
// Delete 删除下载客户端。
func (h *DownloadClientHandler) Delete(c *gin.Context) {
id := c.Param("id")
ctx := c.Request.Context()
_, err := h.svc.Repo.DownloadClient.FindByID(ctx, id)
if err != nil {
Error(c, http.StatusNotFound, ErrNotFound, "客户端不存在")
return
}
if delErr := h.svc.Repo.DownloadClient.Delete(ctx, id); delErr != nil {
Error(c, http.StatusInternalServerError, ErrInternal, "删除失败")
return
}
// 热移除
h.svc.DownloadMgr.RemoveClient(id)
SuccessWithMessage(c, "已删除", nil)
}
// Test 测试下载客户端连接。
func (h *DownloadClientHandler) Test(c *gin.Context) {
id := c.Param("id")
ctx := c.Request.Context()
client, err := h.svc.Repo.DownloadClient.FindByID(ctx, id)
if err != nil || client == nil {
Error(c, http.StatusNotFound, ErrNotFound, "客户端不存在")
return
}
if err := h.svc.DownloadMgr.TestConnection(ctx, client); err != nil {
Error(c, http.StatusBadRequest, ErrExternal, "连接测试失败: "+err.Error())
return
}
SuccessWithMessage(c, "连接成功", nil)
}
+254
View File
@@ -25,6 +25,7 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
{
auth.POST("/login", loginHandler(svc))
auth.POST("/register", registerHandler(svc))
auth.POST("/refresh", refreshHandler(svc))
}
// Authenticated endpoints.
@@ -34,6 +35,10 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
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))
// Libraries.
authed.GET("/libraries", listLibrariesHandler(svc))
@@ -129,6 +134,42 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
authed.GET("/recycle", middleware.AdminRequired(), listRecycleHandler(svc))
authed.GET("/ws", wsHandler(svc))
// SSE event stream.
authed.GET("/events", sseHandler(svc))
// Download clients.
authed.GET("/download-clients", listDownloadClientsHandler(svc))
authed.POST("/download-clients", middleware.AdminRequired(), createDownloadClientHandler(svc))
authed.GET("/download-clients/:id", getDownloadClientHandler(svc))
authed.PUT("/download-clients/:id", middleware.AdminRequired(), updateDownloadClientHandler(svc))
authed.DELETE("/download-clients/:id", middleware.AdminRequired(), deleteDownloadClientHandler(svc))
authed.POST("/download-clients/:id/test", middleware.AdminRequired(), testDownloadClientHandler(svc))
// Notify channels.
authed.GET("/notify-channels", listNotifyChannelsHandler(svc))
authed.GET("/notify-channels/types", getNotifyChannelTypesHandler(svc))
authed.POST("/notify-channels", middleware.AdminRequired(), createNotifyChannelHandler(svc))
authed.GET("/notify-channels/:id", getNotifyChannelHandler(svc))
authed.PUT("/notify-channels/:id", middleware.AdminRequired(), updateNotifyChannelHandler(svc))
authed.DELETE("/notify-channels/:id", middleware.AdminRequired(), deleteNotifyChannelHandler(svc))
authed.POST("/notify-channels/:id/test", middleware.AdminRequired(), testNotifyChannelHandler(svc))
// Scheduler.
authed.GET("/scheduler/tasks", schedulerListTasksHandler(svc))
authed.POST("/scheduler/tasks/:id/run", middleware.AdminRequired(), schedulerRunTaskHandler(svc))
authed.GET("/scheduler/status", schedulerGetStatusHandler(svc))
// Sites (PT 站点管理).
siteHandler := NewSiteHandler(svc)
authed.GET("/sites", siteHandler.ListSites)
authed.GET("/sites/types", siteHandler.GetSiteTypes)
authed.GET("/sites/auth-types", siteHandler.GetAuthTypes)
authed.POST("/sites", middleware.AdminRequired(), siteHandler.CreateSite)
authed.GET("/sites/:id", siteHandler.GetSite)
authed.PUT("/sites/:id", middleware.AdminRequired(), siteHandler.UpdateSite)
authed.DELETE("/sites/:id", middleware.AdminRequired(), siteHandler.DeleteSite)
authed.POST("/sites/:id/test", middleware.AdminRequired(), siteHandler.TestSite)
}
// Admin-only endpoints.
@@ -164,6 +205,24 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
// Scheduled jobs.
admin.GET("/scheduler", schedulerStatusHandler(svc))
admin.POST("/scheduler/:name/run", schedulerRunHandler(svc))
// User permissions management.
admin.GET("/users/:id/permissions", getUserPermissionsHandler(svc))
admin.PUT("/users/:id/permissions", updateUserPermissionsHandler(svc))
admin.POST("/users/:id/permissions/reset", resetUserPermissionsHandler(svc))
}
// API Config management (admin only).
apiConfig := api.Group("/api-config")
apiConfig.Use(middleware.AuthRequired(cfg.Secrets.JWTSecret), middleware.AdminRequired())
{
apiConfig.GET("", listApiConfigsHandler(svc))
apiConfig.GET("/providers/list", listProvidersHandler(svc))
apiConfig.GET("/:provider", getApiConfigHandler(svc))
apiConfig.GET("/:provider/effective", getEffectiveConfigHandler(svc))
apiConfig.POST("/:provider", upsertApiConfigHandler(svc))
apiConfig.DELETE("/:provider", deleteApiConfigHandler(svc))
apiConfig.POST("/:provider/test", testApiConfigHandler(svc))
}
// Emby/Jellyfin compatibility shim (read-only).
@@ -188,3 +247,198 @@ func healthCheck(c *gin.Context) {
func versionInfo(c *gin.Context) {
c.JSON(200, gin.H{"name": "MediaStationGo", "version": "0.1.0"})
}
// ─── 权限 Handler 包装 ────────────────────────────────────────────────────────
func getUserPermissionsHandler(svc *service.Container) gin.HandlerFunc {
h := NewPermissionHandler(svc, svc.Log)
return h.GetUserPermissions
}
func updateUserPermissionsHandler(svc *service.Container) gin.HandlerFunc {
h := NewPermissionHandler(svc, svc.Log)
return h.UpdateUserPermissions
}
func resetUserPermissionsHandler(svc *service.Container) gin.HandlerFunc {
h := NewPermissionHandler(svc, svc.Log)
return h.ResetUserPermissions
}
func getMyPermissionsHandler(svc *service.Container) gin.HandlerFunc {
h := NewPermissionHandler(svc, svc.Log)
return h.GetMyPermissions
}
// ─── 刷新 Handler 包装 ────────────────────────────────────────────────────────
func refreshHandler(svc *service.Container) gin.HandlerFunc {
h := NewRefreshHandler(svc, svc.Log)
return h.RefreshToken
}
func logoutHandler(svc *service.Container) gin.HandlerFunc {
h := NewRefreshHandler(svc, svc.Log)
return h.Logout
}
// ─── API Config Handler 包装 ───────────────────────────────────────────────────
func listApiConfigsHandler(svc *service.Container) gin.HandlerFunc {
h := NewApiConfigHandler(svc, svc.Log)
return h.ListApiConfigs
}
func listProvidersHandler(svc *service.Container) gin.HandlerFunc {
h := NewApiConfigHandler(svc, svc.Log)
return h.ListProviders
}
func getApiConfigHandler(svc *service.Container) gin.HandlerFunc {
h := NewApiConfigHandler(svc, svc.Log)
return h.GetApiConfig
}
func getEffectiveConfigHandler(svc *service.Container) gin.HandlerFunc {
h := NewApiConfigHandler(svc, svc.Log)
return h.GetEffectiveConfig
}
func upsertApiConfigHandler(svc *service.Container) gin.HandlerFunc {
h := NewApiConfigHandler(svc, svc.Log)
return h.UpsertApiConfig
}
func deleteApiConfigHandler(svc *service.Container) gin.HandlerFunc {
h := NewApiConfigHandler(svc, svc.Log)
return h.DeleteApiConfig
}
func testApiConfigHandler(svc *service.Container) gin.HandlerFunc {
h := NewApiConfigHandler(svc, svc.Log)
return h.TestApiConfig
}
// ─── Download Client Handler 包装 ─────────────────────────────────────────────
func listDownloadClientsHandler(svc *service.Container) gin.HandlerFunc {
h := NewDownloadClientHandler(svc, svc.Log)
return h.List
}
func createDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
h := NewDownloadClientHandler(svc, svc.Log)
return h.Create
}
func getDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
h := NewDownloadClientHandler(svc, svc.Log)
return h.Get
}
func updateDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
h := NewDownloadClientHandler(svc, svc.Log)
return h.Update
}
func deleteDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
h := NewDownloadClientHandler(svc, svc.Log)
return h.Delete
}
func testDownloadClientHandler(svc *service.Container) gin.HandlerFunc {
h := NewDownloadClientHandler(svc, svc.Log)
return h.Test
}
// ─── Notify Channel Handler 包装 ──────────────────────────────────────────────
func listNotifyChannelsHandler(svc *service.Container) gin.HandlerFunc {
h := NewNotifyHandler(svc, svc.Log)
return h.List
}
func getNotifyChannelTypesHandler(svc *service.Container) gin.HandlerFunc {
h := NewNotifyHandler(svc, svc.Log)
return h.GetTypes
}
func createNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
h := NewNotifyHandler(svc, svc.Log)
return h.Create
}
func getNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
h := NewNotifyHandler(svc, svc.Log)
return h.Get
}
func updateNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
h := NewNotifyHandler(svc, svc.Log)
return h.Update
}
func deleteNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
h := NewNotifyHandler(svc, svc.Log)
return h.Delete
}
func testNotifyChannelHandler(svc *service.Container) gin.HandlerFunc {
h := NewNotifyHandler(svc, svc.Log)
return h.Test
}
// ─── Scheduler Handler 包装 ──────────────────────────────────────────────────
func schedulerListTasksHandler(svc *service.Container) gin.HandlerFunc {
h := NewSchedulerHandler(svc, svc.Log)
return h.ListTasks
}
func schedulerRunTaskHandler(svc *service.Container) gin.HandlerFunc {
h := NewSchedulerHandler(svc, svc.Log)
return h.RunTask
}
func schedulerGetStatusHandler(svc *service.Container) gin.HandlerFunc {
h := NewSchedulerHandler(svc, svc.Log)
return h.GetStatus
}
// ─── SSE Handler ──────────────────────────────────────────────────────────────
func sseHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
// 获取 SSE Hub
hub := svc.SSEHub
// 设置 SSE 响应头
c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive")
c.Header("X-Accel-Buffering", "no")
// 订阅事件流
client := hub.Subscribe()
defer hub.Unsubscribe(client)
// 发送初始连接成功事件
c.SSEvent("connected", gin.H{"status": "ok"})
c.Writer.Flush()
// 持续发送事件直到客户端断开连接
clientGone := c.Request.Context().Done()
for {
select {
case <-clientGone:
return
case event, ok := <-client.Ch:
if !ok {
return
}
c.SSEvent(event.Type, event.Payload)
c.Writer.Flush()
}
}
}
}
+196
View File
@@ -0,0 +1,196 @@
// Package handler — 通知渠道管理 HTTP 端点。
package handler
import (
"encoding/json"
"net/http"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// NotifyHandler 处理通知渠道的 CRUD 操作。
type NotifyHandler struct {
svc *service.Container
log *zap.Logger
}
// NewNotifyHandler 创建通知渠道处理器。
func NewNotifyHandler(svc *service.Container, log *zap.Logger) *NotifyHandler {
return &NotifyHandler{svc: svc, log: log}
}
// notifyCreateRequest 创建通知渠道请求体。
type notifyCreateRequest struct {
Name string `json:"name" binding:"required"`
Type string `json:"type" binding:"required,oneof=telegram wechat bark webhook email"`
Enabled bool `json:"enabled"`
Config map[string]string `json:"config" binding:"required"`
Events []string `json:"events"`
}
// notifyUpdateRequest 更新通知渠道请求体。
type notifyUpdateRequest struct {
Name string `json:"name"`
Enabled *bool `json:"enabled"`
Config map[string]string `json:"config"`
Events []string `json:"events"`
}
// Create 创建新的通知渠道。
func (h *NotifyHandler) Create(c *gin.Context) {
var req notifyCreateRequest
if err := c.ShouldBindJSON(&req); err != nil {
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
return
}
ctx := c.Request.Context()
// 验证配置
if err := h.svc.Notify.ValidateChannelConfig(req.Type, req.Config); err != nil {
Error(c, http.StatusBadRequest, ErrInvalidParams, "配置验证失败: "+err.Error())
return
}
// 加密配置
configJSON, _ := json.Marshal(req.Config)
configStr := string(configJSON)
if h.svc.Crypto != nil {
configStr = h.svc.Crypto.Encrypt(configStr)
}
// 序列化事件列表
eventsJSON, _ := json.Marshal(req.Events)
eventsStr := string(eventsJSON)
channel := &model.NotifyChannel{
Name: req.Name,
Type: req.Type,
Enabled: req.Enabled,
Config: configStr,
Events: eventsStr,
}
if err := h.svc.Repo.NotifyChannel.Create(ctx, channel); err != nil {
Error(c, http.StatusInternalServerError, ErrInternal, "创建失败: "+err.Error())
return
}
Success(c, channel)
}
// List 返回所有通知渠道。
func (h *NotifyHandler) List(c *gin.Context) {
ctx := c.Request.Context()
channels, err := h.svc.Repo.NotifyChannel.List(ctx)
if err != nil {
Error(c, http.StatusInternalServerError, ErrInternal, "查询失败")
return
}
Success(c, channels)
}
// Get 返回指定通知渠道详情。
func (h *NotifyHandler) Get(c *gin.Context) {
id := c.Param("id")
ctx := c.Request.Context()
channel, err := h.svc.Repo.NotifyChannel.FindByID(ctx, id)
if err != nil || channel == nil {
Error(c, http.StatusNotFound, ErrNotFound, "通知渠道不存在")
return
}
Success(c, channel)
}
// Update 更新通知渠道。
func (h *NotifyHandler) Update(c *gin.Context) {
id := c.Param("id")
var req notifyUpdateRequest
if err := c.ShouldBindJSON(&req); err != nil {
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
return
}
ctx := c.Request.Context()
channel, err := h.svc.Repo.NotifyChannel.FindByID(ctx, id)
if err != nil || channel == nil {
Error(c, http.StatusNotFound, ErrNotFound, "通知渠道不存在")
return
}
if req.Name != "" {
channel.Name = req.Name
}
if req.Enabled != nil {
channel.Enabled = *req.Enabled
}
// 更新配置
if len(req.Config) > 0 {
if err := h.svc.Notify.ValidateChannelConfig(channel.Type, req.Config); err != nil {
Error(c, http.StatusBadRequest, ErrInvalidParams, "配置验证失败: "+err.Error())
return
}
configJSON, _ := json.Marshal(req.Config)
configStr := string(configJSON)
if h.svc.Crypto != nil {
configStr = h.svc.Crypto.Encrypt(configStr)
}
channel.Config = configStr
}
// 更新事件列表
if req.Events != nil {
eventsJSON, _ := json.Marshal(req.Events)
channel.Events = string(eventsJSON)
}
if err := h.svc.Repo.NotifyChannel.Update(ctx, channel); err != nil {
Error(c, http.StatusInternalServerError, ErrInternal, "更新失败")
return
}
Success(c, channel)
}
// Delete 删除通知渠道。
func (h *NotifyHandler) Delete(c *gin.Context) {
id := c.Param("id")
ctx := c.Request.Context()
_, err := h.svc.Repo.NotifyChannel.FindByID(ctx, id)
if err != nil {
Error(c, http.StatusNotFound, ErrNotFound, "通知渠道不存在")
return
}
if delErr := h.svc.Repo.NotifyChannel.Delete(ctx, id); delErr != nil {
Error(c, http.StatusInternalServerError, ErrInternal, "删除失败")
return
}
SuccessWithMessage(c, "已删除", nil)
}
// Test 发送测试通知。
func (h *NotifyHandler) Test(c *gin.Context) {
id := c.Param("id")
ctx := c.Request.Context()
if err := h.svc.Notify.SendTest(ctx, id); err != nil {
Error(c, http.StatusBadRequest, ErrExternal, "测试通知发送失败: "+err.Error())
return
}
SuccessWithMessage(c, "测试通知已发送", nil)
}
// GetTypes 返回支持的通知渠道类型列表。
func (h *NotifyHandler) GetTypes(c *gin.Context) {
types := h.svc.Notify.GetProviderTypes()
Success(c, types)
}
+139
View File
@@ -0,0 +1,139 @@
// Package handler — 权限相关 HTTP Handler。
package handler
import (
"net/http"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// PermissionHandler 权限相关 HTTP 处理。
type PermissionHandler struct {
svc *service.Container
log *zap.Logger
}
// NewPermissionHandler 创建权限处理器。
func NewPermissionHandler(svc *service.Container, log *zap.Logger) *PermissionHandler {
return &PermissionHandler{svc: svc, log: log}
}
// GetUserPermissions 获取指定用户的权限(管理员)。
// GET /api/users/:id/permissions
func (h *PermissionHandler) GetUserPermissions(c *gin.Context) {
userID := c.Param("id")
if userID == "" {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "user id required", "data": nil})
return
}
// 权限检查:需要管理员权限
role := middleware.GetUserRole(c)
if role != "admin" {
c.JSON(http.StatusForbidden, gin.H{"code": 40301, "message": "admin only", "data": nil})
return
}
perms, err := h.svc.Permission.GetByUserID(c.Request.Context(), userID)
if err != nil {
h.log.Error("get user permissions failed", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "data": perms, "message": "ok"})
}
// UpdateUserPermissions 更新指定用户的权限(管理员)。
// PUT /api/users/:id/permissions
func (h *PermissionHandler) UpdateUserPermissions(c *gin.Context) {
userID := c.Param("id")
if userID == "" {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "user id required", "data": nil})
return
}
// 权限检查:需要管理员权限
role := middleware.GetUserRole(c)
if role != "admin" {
c.JSON(http.StatusForbidden, gin.H{"code": 40301, "message": "admin only", "data": nil})
return
}
var req struct {
Permissions map[string]bool `json:"permissions"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "invalid request", "data": nil})
return
}
if err := h.svc.Permission.Update(c.Request.Context(), userID, req.Permissions); err != nil {
h.log.Error("update user permissions failed", zap.Error(err), zap.String("user_id", userID))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "permissions updated", "data": nil})
}
// ResetUserPermissions 重置指定用户的权限为默认值(管理员)。
// POST /api/users/:id/permissions/reset
func (h *PermissionHandler) ResetUserPermissions(c *gin.Context) {
userID := c.Param("id")
if userID == "" {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "user id required", "data": nil})
return
}
// 权限检查:需要管理员权限
role := middleware.GetUserRole(c)
if role != "admin" {
c.JSON(http.StatusForbidden, gin.H{"code": 40301, "message": "admin only", "data": nil})
return
}
if err := h.svc.Permission.ResetToDefault(c.Request.Context(), userID); err != nil {
h.log.Error("reset user permissions failed", zap.Error(err), zap.String("user_id", userID))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "permissions reset to default", "data": nil})
}
// GetMyPermissions 获取当前用户的权限。
// GET /api/auth/permissions
func (h *PermissionHandler) GetMyPermissions(c *gin.Context) {
currentUserID := middleware.GetUserID(c)
if currentUserID == "" {
c.JSON(http.StatusUnauthorized, gin.H{"code": 40101, "message": "authentication required", "data": nil})
return
}
perms, err := h.svc.Permission.GetPermissionMap(c.Request.Context(), currentUserID)
if err != nil {
h.log.Error("get my permissions failed", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
return
}
// 同时返回角色和等级信息
role := middleware.GetUserRole(c)
tier := middleware.GetUserTier(c)
c.JSON(http.StatusOK, gin.H{
"code": 0,
"message": "ok",
"data": gin.H{
"permissions": perms,
"role": role,
"tier": tier,
"is_super": role == "admin" || tier == "plus",
},
})
}
+82
View File
@@ -0,0 +1,82 @@
// Package handler — 令牌刷新 HTTP Handler。
package handler
import (
"net/http"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// RefreshHandler 令牌刷新 HTTP 处理。
type RefreshHandler struct {
svc *service.Container
log *zap.Logger
}
// NewRefreshHandler 创建刷新令牌处理器。
func NewRefreshHandler(svc *service.Container, log *zap.Logger) *RefreshHandler {
return &RefreshHandler{svc: svc, log: log}
}
// RefreshTokenRequest 刷新令牌请求结构。
type RefreshTokenRequest struct {
RefreshToken string `json:"refresh_token" binding:"required"`
}
// RefreshToken 刷新访问令牌。
// POST /api/auth/refresh
func (h *RefreshHandler) RefreshToken(c *gin.Context) {
var req RefreshTokenRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 40001, "message": "refresh_token required", "data": nil})
return
}
tokens, err := h.svc.Auth.RefreshTokens(c.Request.Context(), req.RefreshToken)
if err != nil {
h.log.Debug("token refresh failed", zap.Error(err))
// 根据错误类型返回不同状态码
switch err {
case service.ErrInvalidRefreshToken:
c.JSON(http.StatusUnauthorized, gin.H{"code": 40101, "message": "invalid refresh token", "data": nil})
case service.ErrTokenExpired:
c.JSON(http.StatusUnauthorized, gin.H{"code": 40102, "message": "refresh token expired", "data": nil})
case service.ErrTokenRevoked:
c.JSON(http.StatusUnauthorized, gin.H{"code": 40103, "message": "refresh token revoked", "data": nil})
default:
c.JSON(http.StatusInternalServerError, gin.H{"code": 50001, "message": "internal error", "data": nil})
}
return
}
c.JSON(http.StatusOK, gin.H{
"code": 0,
"message": "ok",
"data": gin.H{
"token": tokens.AccessToken,
"refresh_token": tokens.RefreshToken,
"expires_in": tokens.ExpiresIn,
"token_type": tokens.TokenType,
},
})
}
// Logout 登出当前用户。
// POST /api/auth/logout
func (h *RefreshHandler) Logout(c *gin.Context) {
userID := c.GetString("ctx_user_id")
if userID == "" {
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": nil})
return
}
if err := h.svc.Auth.Logout(c.Request.Context(), userID); err != nil {
h.log.Warn("logout failed", zap.Error(err), zap.String("user_id", userID))
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": nil})
}
+69
View File
@@ -0,0 +1,69 @@
// Package handler — 统一响应格式和错误码。
package handler
import (
"net/http"
"github.com/gin-gonic/gin"
)
// ─── 错误码定义 ───────────────────────────────────────────────────────────────
// 统一错误码
const (
ErrOK = 0
ErrInvalidParams = 40001
ErrUnauthorized = 40101
ErrForbidden = 40301
ErrNotFound = 40401
ErrConflict = 40901
ErrInternal = 50001
ErrExternal = 50201
ErrEncryptFailed = 50801
)
// ─── Response Helpers ─────────────────────────────────────────────────────────
// APIResponse 统一 API 响应格式。
type APIResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data interface{} `json:"data"`
}
// PaginatedResponse 分页响应格式。
type PaginatedResponse struct {
Code int `json:"code"`
Message string `json:"message"`
Data interface{} `json:"data"`
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
}
// Success 返回成功响应。
func Success(c *gin.Context, data interface{}) {
c.JSON(http.StatusOK, APIResponse{Code: 0, Message: "ok", Data: data})
}
// SuccessWithMessage 返回带消息的成功响应。
func SuccessWithMessage(c *gin.Context, message string, data interface{}) {
c.JSON(http.StatusOK, APIResponse{Code: 0, Message: message, Data: data})
}
// Error 返回错误响应。
func Error(c *gin.Context, httpStatus int, code int, message string) {
c.JSON(httpStatus, APIResponse{Code: code, Message: message, Data: nil})
}
// Paginated 返回分页响应。
func Paginated(c *gin.Context, items interface{}, total int64, page, pageSize int) {
c.JSON(http.StatusOK, PaginatedResponse{
Code: 0,
Message: "ok",
Data: items,
Total: total,
Page: page,
PageSize: pageSize,
})
}
+47
View File
@@ -0,0 +1,47 @@
// Package handler — 定时任务管理 HTTP 端点。
package handler
import (
"net/http"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// SchedulerHandler 处理定时任务的查询和管理操作。
type SchedulerHandler struct {
svc *service.Container
log *zap.Logger
}
// NewSchedulerHandler 创建定时任务处理器。
func NewSchedulerHandler(svc *service.Container, log *zap.Logger) *SchedulerHandler {
return &SchedulerHandler{svc: svc, log: log}
}
// ListTasks 返回所有定时任务列表。
func (h *SchedulerHandler) ListTasks(c *gin.Context) {
tasks := h.svc.Scheduler.Status()
Success(c, tasks)
}
// RunTask 手动触发指定任务。
func (h *SchedulerHandler) RunTask(c *gin.Context) {
name := c.Param("id")
ctx := c.Request.Context()
if err := h.svc.Scheduler.RunNow(ctx, name); err != nil {
Error(c, http.StatusBadRequest, ErrInternal, "任务执行失败: "+err.Error())
return
}
SuccessWithMessage(c, "任务已触发执行", nil)
}
// GetStatus 返回调度器运行状态。
func (h *SchedulerHandler) GetStatus(c *gin.Context) {
status := h.svc.Scheduler.Status()
Success(c, status)
}
+115
View File
@@ -0,0 +1,115 @@
// Package handler — PT 站点管理 HTTP 处理。
package handler
import (
"net/http"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
// SiteHandler 站点管理 CRUD。
type SiteHandler struct {
svc *service.Container
}
// NewSiteHandler 创建站点管理 Handler。
func NewSiteHandler(svc *service.Container) *SiteHandler {
return &SiteHandler{svc: svc}
}
// ListSites 列出所有站点。
func (h *SiteHandler) ListSites(c *gin.Context) {
sites, err := h.svc.Site.List(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"code": 1, "message": err.Error()})
return
}
if sites == nil {
sites = []model.Site{}
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": sites})
}
// GetSite 获取单个站点详情(解密敏感字段)。
func (h *SiteHandler) GetSite(c *gin.Context) {
site, err := h.svc.Site.GetByID(c.Request.Context(), c.Param("id"))
if err != nil {
if err == service.ErrSiteNotFound {
c.JSON(http.StatusNotFound, gin.H{"code": 1, "message": "site not found"})
return
}
c.JSON(http.StatusInternalServerError, gin.H{"code": 1, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": site})
}
// CreateSite 创建站点。
func (h *SiteHandler) CreateSite(c *gin.Context) {
var site model.Site
if err := c.ShouldBindJSON(&site); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
return
}
created, err := h.svc.Site.Create(c.Request.Context(), &site)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
return
}
c.JSON(http.StatusCreated, gin.H{"code": 0, "message": "ok", "data": created})
}
// UpdateSite 更新站点。
func (h *SiteHandler) UpdateSite(c *gin.Context) {
var site model.Site
if err := c.ShouldBindJSON(&site); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
return
}
site.ID = c.Param("id")
updated, err := h.svc.Site.Update(c.Request.Context(), &site)
if err != nil {
if err == service.ErrSiteNotFound {
c.JSON(http.StatusNotFound, gin.H{"code": 1, "message": "site not found"})
return
}
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": updated})
}
// DeleteSite 删除站点。
func (h *SiteHandler) DeleteSite(c *gin.Context) {
if err := h.svc.Site.Delete(c.Request.Context(), c.Param("id")); err != nil {
if err == service.ErrSiteNotFound {
c.JSON(http.StatusNotFound, gin.H{"code": 1, "message": "site not found"})
return
}
c.JSON(http.StatusInternalServerError, gin.H{"code": 1, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok"})
}
// TestSite 测试站点连通性。
func (h *SiteHandler) TestSite(c *gin.Context) {
if err := h.svc.Site.Authenticate(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok"})
}
// GetSiteTypes 返回支持的站点类型列表。
func (h *SiteHandler) GetSiteTypes(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": model.SiteTypes()})
}
// GetAuthTypes 返回支持的认证方式列表。
func (h *SiteHandler) GetAuthTypes(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": model.AuthTypes()})
}
+95
View File
@@ -0,0 +1,95 @@
// Package middleware — Emby API 兼容层认证中间件。
// 支持 X-Emby-Token / Bearer / URL token / Username+Password 四种认证方式。
package middleware
import (
"errors"
"net/http"
"github.com/gin-gonic/gin"
"github.com/golang-jwt/jwt/v5"
)
// EmbyCtxUserID 是 Emby 认证中间件设置的用户 ID 上下文键。
const EmbyCtxUserID = "emby_user_id"
// EmbyAuthRequired Emby 认证中间件。
// 按优先级尝试以下认证方式:
// 1. X-Emby-Token 请求头
// 2. Authorization: Bearer <token> 请求头
// 3. ?token=<token> URL 参数
// 4. (仅 AuthenticateByName 端点)POST body 中的 Username+Password
func EmbyAuthRequired(secret string) gin.HandlerFunc {
return func(c *gin.Context) {
token := ""
// 1. X-Emby-Token 头
if t := c.GetHeader("X-Emby-Token"); t != "" {
token = t
}
// 2. Authorization: Bearer <token> 或 Emby <token>
if token == "" {
if authHeader := c.GetHeader("Authorization"); authHeader != "" {
// Strip "Bearer " or "Emby " prefix
for _, prefix := range []string{"Bearer ", "Emby "} {
if len(authHeader) > len(prefix) && authHeader[:len(prefix)] == prefix {
token = authHeader[len(prefix):]
break
}
}
if token == "" {
token = authHeader
}
}
}
// 3. URL 参数 token
if token == "" {
if t := c.Query("token"); t != "" {
token = t
}
}
if token == "" {
c.JSON(http.StatusUnauthorized, gin.H{
"Code": 40101,
"Message": "Unauthorized",
})
c.Abort()
return
}
// 解析 JWT
claims := &Claims{}
parsed, err := jwt.ParseWithClaims(token, claims, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, errors.New("unexpected signing method")
}
return []byte(secret), nil
})
if err != nil || !parsed.Valid || claims.UserID == "" {
c.JSON(http.StatusUnauthorized, gin.H{
"Code": 40101,
"Message": "Invalid token",
})
c.Abort()
return
}
c.Set(EmbyCtxUserID, claims.UserID)
c.Set(CtxUserID, claims.UserID)
c.Set(CtxUserRole, claims.Role)
c.Set(CtxUserTier, claims.Tier)
c.Next()
}
}
// GetEmbyUserID 从上下文中获取 Emby 用户 ID。
func GetEmbyUserID(c *gin.Context) string {
if uid, exists := c.Get(EmbyCtxUserID); exists {
return uid.(string)
}
return ""
}
+60 -5
View File
@@ -1,5 +1,5 @@
// Package middleware exposes Gin middlewares used by the HTTP server:
// request logging, CORS, JWT authentication and admin guard.
// Package middleware 暴露 Gin 中间件,用于 HTTP 服务器:
// 请求日志、CORS、JWT 认证、管理员守卫和权限检查。
package middleware
import (
@@ -17,6 +17,7 @@ import (
const (
CtxUserID = "ctx_user_id"
CtxUserRole = "ctx_user_role"
CtxUserTier = "ctx_user_tier"
)
// RequestLogger logs one structured line per request.
@@ -65,6 +66,7 @@ func CORS(origins []string) gin.HandlerFunc {
type Claims struct {
UserID string `json:"uid"`
Role string `json:"role"`
Tier string `json:"tier,omitempty"`
jwt.RegisteredClaims
}
@@ -74,7 +76,7 @@ func AuthRequired(secret string) gin.HandlerFunc {
return func(c *gin.Context) {
raw := extractToken(c)
if raw == "" {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "missing token"})
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"code": 40101, "message": "missing token"})
return
}
claims := &Claims{}
@@ -85,11 +87,12 @@ func AuthRequired(secret string) gin.HandlerFunc {
return []byte(secret), nil
})
if err != nil || claims.UserID == "" {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid token"})
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"code": 40101, "message": "invalid token"})
return
}
c.Set(CtxUserID, claims.UserID)
c.Set(CtxUserRole, claims.Role)
c.Set(CtxUserTier, claims.Tier)
c.Next()
}
}
@@ -99,13 +102,65 @@ func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
role, _ := c.Get(CtxUserRole)
if role != "admin" {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "admin only"})
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"code": 40301, "message": "admin only"})
return
}
c.Next()
}
}
// PlusOrAdminRequired enforces role == "admin" or tier == "plus".
func PlusOrAdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
role, _ := c.Get(CtxUserRole)
tier, _ := c.Get(CtxUserTier)
if role != "admin" && tier != "plus" {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"code": 40301, "message": "plus or admin only"})
return
}
c.Next()
}
}
// GetUserID extracts the user ID from the Gin context.
func GetUserID(c *gin.Context) string {
if uid, exists := c.Get(CtxUserID); exists {
return uid.(string)
}
return ""
}
// GetUserRole extracts the user role from the Gin context.
func GetUserRole(c *gin.Context) string {
if role, exists := c.Get(CtxUserRole); exists {
return role.(string)
}
return ""
}
// GetUserTier extracts the user tier from the Gin context.
func GetUserTier(c *gin.Context) string {
if tier, exists := c.Get(CtxUserTier); exists {
return tier.(string)
}
return ""
}
// IsAdmin checks if the current user is an admin.
func IsAdmin(c *gin.Context) bool {
return GetUserRole(c) == "admin"
}
// IsPlus checks if the current user is a plus subscriber.
func IsPlus(c *gin.Context) bool {
return GetUserTier(c) == "plus" || GetUserRole(c) == "admin"
}
// IsSuperUser checks if the current user is a super user (admin or plus).
func IsSuperUser(c *gin.Context) bool {
return IsAdmin(c) || IsPlus(c)
}
func extractToken(c *gin.Context) string {
if h := c.GetHeader("Authorization"); strings.HasPrefix(h, "Bearer ") {
return strings.TrimSpace(strings.TrimPrefix(h, "Bearer "))
+111
View File
@@ -0,0 +1,111 @@
// Package middleware — 权限检查中间件。
// 注意:实际的权限检查在 handler 层通过 PermissionService 实现。
// 此中间件主要用于设置上下文和基本的角色/等级检查。
package middleware
import (
"net/http"
"github.com/gin-gonic/gin"
)
// RequirePermission 创建权限检查中间件标记。
// 实际的权限检查由 handler 中的 PermissionService 执行。
// 此中间件确保请求已经过身份验证。
func RequirePermission(permissionKey string) gin.HandlerFunc {
return func(c *gin.Context) {
userID := GetUserID(c)
role := GetUserRole(c)
tier := GetUserTier(c)
if userID == "" {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
"code": 40101,
"message": "authentication required",
"data": nil,
})
return
}
// admin 和 plus 用户拥有所有权限
if role == "admin" || tier == "plus" {
c.Next()
return
}
// 将权限键存储到上下文中供 handler 使用
c.Set("permission_key", permissionKey)
c.Next()
}
}
// RequireAnyPermission 创建需要任意一个权限的中间件标记。
func RequireAnyPermission(permissionKeys ...string) gin.HandlerFunc {
return func(c *gin.Context) {
userID := GetUserID(c)
role := GetUserRole(c)
tier := GetUserTier(c)
if userID == "" {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
"code": 40101,
"message": "authentication required",
"data": nil,
})
return
}
// admin 和 plus 用户拥有所有权限
if role == "admin" || tier == "plus" {
c.Next()
return
}
// 将权限键数组存储到上下文中供 handler 使用
c.Set("permission_keys", permissionKeys)
c.Next()
}
}
// RequireAllPermissions 创建需要所有权限的中间件标记。
func RequireAllPermissions(permissionKeys ...string) gin.HandlerFunc {
return func(c *gin.Context) {
userID := GetUserID(c)
role := GetUserRole(c)
tier := GetUserTier(c)
if userID == "" {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{
"code": 40101,
"message": "authentication required",
"data": nil,
})
return
}
// admin 和 plus 用户拥有所有权限
if role == "admin" || tier == "plus" {
c.Next()
return
}
c.Set("permission_keys", permissionKeys)
c.Next()
}
}
// GetPermissionKey 从上下文中获取存储的权限键。
func GetPermissionKey(c *gin.Context) string {
if key, exists := c.Get("permission_key"); exists {
return key.(string)
}
return ""
}
// GetPermissionKeys 从上下文中获取存储的权限键数组。
func GetPermissionKeys(c *gin.Context) []string {
if keys, exists := c.Get("permission_keys"); exists {
return keys.([]string)
}
return nil
}
+62
View File
@@ -0,0 +1,62 @@
// Package model 定义第三方 API 配置数据模型。
package model
import (
"time"
"github.com/google/uuid"
"gorm.io/gorm"
)
// ApiConfig 存储第三方 API 密钥和配置信息。
// APIKey 字段在 JSON 序列化时隐藏(json:"-"),通过加密存储。
type ApiConfig struct {
ID string `gorm:"primaryKey;size:36" json:"id"`
Provider string `gorm:"size:64;uniqueIndex;not null" json:"provider"`
APIKey string `gorm:"size:512" json:"-"`
BaseURL string `gorm:"size:512" json:"base_url,omitempty"`
Extra string `gorm:"type:text" json:"extra,omitempty"`
Enabled bool `gorm:"default:true" json:"enabled"`
Description string `gorm:"size:255" json:"description,omitempty"`
LastTestedAt *time.Time `json:"last_tested_at,omitempty"`
TestResult string `gorm:"size:32" json:"test_result,omitempty"`
UpdatedAt time.Time `json:"updated_at"`
}
// BeforeCreate 生成 UUID。
func (c *ApiConfig) BeforeCreate(_ *gorm.DB) error {
if c.ID == "" {
c.ID = uuid.NewString()
}
return nil
}
// BeforeUpdate 更新时间戳。
func (c *ApiConfig) BeforeUpdate(_ *gorm.DB) error {
c.UpdatedAt = time.Now()
return nil
}
// ApiProvider 定义支持的 API 提供者列表。
type ApiProvider struct {
ID string `json:"id"`
Name string `json:"name"`
Description string `json:"description"`
HasAPIKey bool `json:"has_api_key"`
HasBaseURL bool `json:"has_base_url"`
}
// PredefinedProviders 返回预定义的 API 提供者列表。
func PredefinedProviders() []ApiProvider {
return []ApiProvider{
{ID: "tmdb", Name: "TMDb", Description: "The Movie Database - 电影/剧集元数据", HasAPIKey: true, HasBaseURL: true},
{ID: "douban", Name: "豆瓣", Description: "豆瓣电影/音乐/书籍数据", HasAPIKey: true, HasBaseURL: false},
{ID: "bangumi", Name: "Bangumi", Description: "番剧/动漫数据库", HasAPIKey: true, HasBaseURL: false},
{ID: "thetvdb", Name: "TheTVDB", Description: "TV Series Database", HasAPIKey: true, HasBaseURL: false},
{ID: "fanart", Name: "Fanart.tv", Description: "影视海报/背景图", HasAPIKey: true, HasBaseURL: false},
{ID: "openai", Name: "OpenAI", Description: "GPT 系列模型", HasAPIKey: true, HasBaseURL: true},
{ID: "deepseek", Name: "DeepSeek", Description: "DeepSeek 大模型", HasAPIKey: true, HasBaseURL: true},
{ID: "siliconflow", Name: "SiliconFlow", Description: "AI 模型聚合 API", HasAPIKey: true, HasBaseURL: true},
{ID: "adult", Name: "Adult API", Description: "Adult 内容元数据(需额外权限)", HasAPIKey: true, HasBaseURL: false},
}
}
+16
View File
@@ -0,0 +1,16 @@
// Package model 定义下载客户端配置数据模型。
package model
// DownloadClient 下载客户端配置(qBittorrent / Transmission / Aria2)。
// 支持多客户端并行运行,一个标记为默认客户端。
type DownloadClient struct {
Base
Name string `gorm:"size:128;not null" json:"name"`
Type string `gorm:"size:32;not null" json:"type"` // qbittorrent / transmission / aria2
Host string `gorm:"size:512;not null" json:"host"` // http://host:port
Username string `gorm:"size:256" json:"username"`
Password string `gorm:"size:1024" json:"-"` // AES加密存储
IsDefault bool `gorm:"default:false" json:"is_default"`
Enabled bool `gorm:"default:true" json:"enabled"`
Extra string `gorm:"type:text" json:"-"` // JSON配置, AES加密
}
+530
View File
@@ -0,0 +1,530 @@
// Package model — Emby API 兼容层请求/响应类型。
package model
import "time"
// ─── Emby 认证 ────────────────────────────────────────────────────────────────
// EmbyAuthRequest Emby 认证请求(用户名+密码)。
type EmbyAuthRequest struct {
Username string `json:"Username"`
Password string `json:"Password"`
}
// EmbyAuthResponse Emby 认证响应。
type EmbyAuthResponse struct {
User EmbyUser `json:"User"`
AccessToken string `json:"AccessToken"`
ServerID string `json:"ServerId"`
}
// EmbyApiKeyAuthRequest API Key 认证请求。
type EmbyApiKeyAuthRequest struct {
ApiKey string `json:"ApiKey"`
}
// ─── Emby 用户 ────────────────────────────────────────────────────────────────
// EmbyUser Emby 用户信息。
type EmbyUser struct {
Id string `json:"Id"`
Name string `json:"Name"`
ServerId string `json:"ServerId"`
HasPassword bool `json:"HasPassword"`
PrimaryImageTag string `json:"PrimaryImageTag,omitempty"`
Configuration *EmbyUserConfiguration `json:"Configuration,omitempty"`
LastActivityDate *time.Time `json:"LastActivityDate,omitempty"`
LastLoginDate *time.Time `json:"LastLoginDate,omitempty"`
}
// EmbyUserConfiguration 用户配置。
type EmbyUserConfiguration struct {
AudioLanguagePreference string `json:"AudioLanguagePreference"`
SubtitleLanguagePreference string `json:"SubtitleLanguagePreference"`
EnableAutoPlay bool `json:"EnableAutoPlay"`
EnableNextEpisodeAutoPlay bool `json:"EnableNextEpisodeAutoPlay"`
}
// ─── Emby 系统信息 ────────────────────────────────────────────────────────────
// EmbySystemInfo 系统信息。
type EmbySystemInfo struct {
Id string `json:"Id"`
ServerName string `json:"ServerName"`
Version string `json:"Version"`
ProductName string `json:"ProductName"`
OperatingSystem string `json:"OperatingSystem"`
Architecture string `json:"Architecture"`
LocalAddress string `json:"LocalAddress"`
WanAddress string `json:"WanAddress,omitempty"`
HasPendingRestart bool `json:"HasPendingRestart"`
IsShuttingDown bool `json:"IsShuttingDown"`
SupportsLibraryScan bool `json:"SupportsLibraryScan"`
SupportsHttps bool `json:"SupportsHttps"`
SupportsAutoDiscovery bool `json:"SupportsAutoDiscovery"`
WebSocketPortNumber int `json:"WebSocketPortNumber"`
TranscodingTempPath string `json:"TranscodingTempPath,omitempty"`
CanSelfUpdate bool `json:"CanSelfUpdate"`
CanLaunchWebBrowser bool `json:"CanLaunchWebBrowser"`
CanRestart bool `json:"CanRestart"`
CodecCount int `json:"CodecCount"`
}
// EmbyLogEntry 日志条目。
type EmbyLogEntry struct {
Id string `json:"Id"`
DateCreated time.Time `json:"DateCreated"`
Level string `json:"Level"`
Message string `json:"Message"`
}
// EmbyServerConfiguration 服务器配置。
type EmbyServerConfiguration struct {
Name string `json:"Name"`
ServerName string `json:"ServerName"`
EnableUPnP bool `json:"EnableUPnP"`
PublicPort int `json:"PublicPort"`
EnableHttps bool `json:"EnableHttps"`
HttpServerPortNumber int `json:"HttpServerPortNumber"`
HttpsPortNumber int `json:"HttpsPortNumber"`
EnableRemoteAccess bool `json:"EnableRemoteAccess"`
}
// EmbySession 会话信息。
type EmbySession struct {
Id string `json:"Id"`
Client string `json:"Client"`
ClientVersion string `json:"ClientVersion"`
DeviceId string `json:"DeviceId"`
DeviceName string `json:"DeviceName"`
UserName string `json:"UserName,omitempty"`
UserId string `json:"UserId,omitempty"`
LastActivityDate *time.Time `json:"LastActivityDate,omitempty"`
RemoteEndPoint string `json:"RemoteEndPoint,omitempty"`
NowPlayingItem *EmbyItem `json:"NowPlayingItem,omitempty"`
PlayState *EmbyPlaybackState `json:"PlayState,omitempty"`
Capabilities EmbyClientCapabilities `json:"Capabilities,omitempty"`
SupportsRemoteControl bool `json:"SupportsRemoteControl"`
AdditionalUsers []EmbySessionUserInfo `json:"AdditionalUsers,omitempty"`
}
// EmbySessionUserInfo 会话中的附加用户。
type EmbySessionUserInfo struct {
UserId string `json:"UserId"`
UserName string `json:"UserName"`
}
// EmbyPlaybackState 播放状态。
type EmbyPlaybackState struct {
PositionTicks int64 `json:"PositionTicks"`
VolumeLevel int `json:"VolumeLevel"`
IsMuted bool `json:"IsMuted"`
IsPaused bool `json:"IsPaused"`
PlayMethod string `json:"PlayMethod,omitempty"`
CanSeek bool `json:"CanSeek"`
}
// EmbyClientCapabilities 客户端能力描述。
type EmbyClientCapabilities struct {
PlayableMediaTypes []string `json:"PlayableMediaTypes"`
SupportedCommands []string `json:"SupportedCommands"`
SupportsMediaControl bool `json:"SupportsMediaControl"`
SupportsSync bool `json:"SupportsSync"`
}
// ─── Emby 虚拟文件夹 / 媒体库 ────────────────────────────────────────────────
// EmbyVirtualFolder 虚拟文件夹(媒体库)。
type EmbyVirtualFolder struct {
Name string `json:"Name"`
Locations []string `json:"Locations"`
CollectionType string `json:"CollectionType"`
LibraryOptions EmbyLibraryOptions `json:"LibraryOptions,omitempty"`
RefreshStatus *EmbyRefreshStatus `json:"RefreshStatus,omitempty"`
ItemId string `json:"ItemId"`
}
// EmbyLibraryOptions 媒体库选项。
type EmbyLibraryOptions struct {
PreferredMetadataLanguage string `json:"PreferredMetadataLanguage"`
MetadataCountryCode string `json:"MetadataCountryCode"`
EnableRealtimeMonitor bool `json:"EnableRealtimeMonitor"`
EnableAutomaticSeriesGrouping bool `json:"EnableAutomaticSeriesGrouping"`
}
// EmbyRefreshStatus 刷新状态。
type EmbyRefreshStatus struct {
LastRefreshResult string `json:"LastRefreshResult"`
LastRefreshedAt time.Time `json:"LastRefreshedAt"`
IsActive bool `json:"IsActive"`
}
// EmbyItemsCounts 项目计数。
type EmbyItemsCounts struct {
MovieCount int `json:"MovieCount"`
SeriesCount int `json:"SeriesCount"`
EpisodeCount int `json:"EpisodeCount"`
ArtistCount int `json:"ArtistCount"`
AlbumCount int `json:"AlbumCount"`
SongCount int `json:"SongCount"`
MusicVideoCount int `json:"MusicVideoCount"`
BookCount int `json:"BookCount"`
BoxSetCount int `json:"BoxSetCount"`
}
// ─── Emby Items ───────────────────────────────────────────────────────────────
// EmbyItemsResponse Emby 标准分页响应包装。
type EmbyItemsResponse struct {
Items []EmbyItem `json:"Items"`
TotalRecordCount int `json:"TotalRecordCount"`
StartIndex int `json:"StartIndex"`
}
// EmbyItem Emby 媒体项。
type EmbyItem struct {
Id string `json:"Id"`
Name string `json:"Name"`
Type string `json:"Type"` // Movie / Series / Episode / Season / BoxSet / Folder
Overview string `json:"Overview,omitempty"`
ProductionYear int `json:"ProductionYear,omitempty"`
PremiereDate *time.Time `json:"PremiereDate,omitempty"`
CommunityRating float64 `json:"CommunityRating,omitempty"`
OfficialRating string `json:"OfficialRating,omitempty"`
RunTimeTicks int64 `json:"RunTimeTicks,omitempty"`
ParentId string `json:"ParentId,omitempty"`
SeriesId string `json:"SeriesId,omitempty"`
SeasonId string `json:"SeasonId,omitempty"`
IndexNumber int `json:"IndexNumber,omitempty"`
ParentIndexNumber int `json:"ParentIndexNumber,omitempty"`
UserData *EmbyUserData `json:"UserData,omitempty"`
ImageTags map[string]string `json:"ImageTags,omitempty"`
BackdropImageTags []string `json:"BackdropImageTags,omitempty"`
Genres []string `json:"Genres,omitempty"`
Studios []EmbyNameId `json:"Studios,omitempty"`
People []EmbyPerson `json:"People,omitempty"`
MediaSources []EmbyMediaSource `json:"MediaSources,omitempty"`
RecursiveItemCount int `json:"RecursiveItemCount,omitempty"`
ChildCount int `json:"ChildCount,omitempty"`
Status string `json:"Status,omitempty"`
AirDays []string `json:"AirDays,omitempty"`
EndDate *time.Time `json:"EndDate,omitempty"`
ProviderIds map[string]string `json:"ProviderIds,omitempty"`
Taglines []string `json:"Taglines,omitempty"`
GenreItems []EmbyNameId `json:"GenreItems,omitempty"`
DateCreated *time.Time `json:"DateCreated,omitempty"`
Path string `json:"Path,omitempty"`
SortName string `json:"SortName,omitempty"`
ForcedSortName string `json:"ForcedSortName,omitempty"`
Width int `json:"Width,omitempty"`
Height int `json:"Height,omitempty"`
Container string `json:"Container,omitempty"`
}
// EmbyUserData 用户播放数据。
type EmbyUserData struct {
PlaybackPositionTicks int64 `json:"PlaybackPositionTicks"`
PlayCount int `json:"PlayCount"`
IsFavorite bool `json:"IsFavorite"`
Played bool `json:"Played"`
UnplayedItemCount int `json:"UnplayedItemCount"`
PercentagePlayed float64 `json:"PercentagePlayed"`
Rating float64 `json:"Rating,omitempty"`
PlayedPercentage float64 `json:"PlayedPercentage,omitempty"`
}
// EmbyMediaSource 媒体源。
type EmbyMediaSource struct {
Id string `json:"Id"`
Name string `json:"Name"`
Path string `json:"Path"`
Size int64 `json:"Size"`
Container string `json:"Container,omitempty"`
Bitrate int64 `json:"Bitrate,omitempty"`
MediaStreams []EmbyMediaStream `json:"MediaStreams"`
SupportsTranscoding bool `json:"SupportsTranscoding"`
SupportsDirectStream bool `json:"SupportsDirectStream"`
SupportsDirectPlay bool `json:"SupportsDirectPlay"`
TranscodingUrl string `json:"TranscodingUrl,omitempty"`
Protocol string `json:"Protocol,omitempty"`
Type string `json:"Type,omitempty"`
IsRemote bool `json:"IsRemote,omitempty"`
RunTimeTicks int64 `json:"RunTimeTicks,omitempty"`
ETag string `json:"ETag,omitempty"`
SupportsProbing bool `json:"SupportsProbing,omitempty"`
}
// EmbyMediaStream 媒体流(视频/音频/字幕)。
type EmbyMediaStream struct {
Codec string `json:"Codec"`
Type string `json:"Type"` // Video / Audio / Subtitle
Language string `json:"Language,omitempty"`
DisplayTitle string `json:"DisplayTitle,omitempty"`
Index int `json:"Index"`
IsDefault bool `json:"IsDefault"`
IsForced bool `json:"IsForced"`
IsExternal bool `json:"IsExternal,omitempty"`
Height int `json:"Height,omitempty"`
Width int `json:"Width,omitempty"`
BitRate int64 `json:"BitRate,omitempty"`
Channels int `json:"Channels,omitempty"`
SampleRate int `json:"SampleRate,omitempty"`
AspectRatio string `json:"AspectRatio,omitempty"`
VideoRange string `json:"VideoRange,omitempty"`
DeliveryUrl string `json:"DeliveryUrl,omitempty"`
DeliveryMethod string `json:"DeliveryMethod,omitempty"`
ExternalUrl string `json:"ExternalUrl,omitempty"`
ExternalSubtitleId string `json:"ExternalSubtitleId,omitempty"`
SubtitleFileName string `json:"SubtitleFileName,omitempty"`
Title string `json:"Title,omitempty"`
Comment string `json:"Comment,omitempty"`
Path string `json:"Path,omitempty"`
}
// EmbyPerson 人员信息。
type EmbyPerson struct {
Id string `json:"Id"`
Name string `json:"Name"`
Role string `json:"Role,omitempty"`
Type string `json:"Type,omitempty"`
PrimaryImageTag string `json:"PrimaryImageTag,omitempty"`
}
// EmbyNameId 名称+ID 对。
type EmbyNameId struct {
Name string `json:"Name"`
Id string `json:"Id"`
}
// ─── Emby PlaybackInfo ────────────────────────────────────────────────────────
// EmbyPlaybackInfoRequest 播放信息请求。
type EmbyPlaybackInfoRequest struct {
UserId string `json:"UserId,omitempty"`
MaxStreamingBitrate int64 `json:"MaxStreamingBitrate,omitempty"`
StartTimeTicks int64 `json:"StartTimeTicks,omitempty"`
AudioStreamIndex int `json:"AudioStreamIndex,omitempty"`
SubtitleStreamIndex int `json:"SubtitleStreamIndex,omitempty"`
MaxAudioChannels int `json:"MaxAudioChannels,omitempty"`
ItemId string `json:"ItemId,omitempty"`
DeviceProfile *EmbyDeviceProfile `json:"DeviceProfile,omitempty"`
EnableDirectStream bool `json:"EnableDirectStream,omitempty"`
EnableDirectPlay bool `json:"EnableDirectPlay,omitempty"`
AutoOpenLiveStream bool `json:"AutoOpenLiveStream,omitempty"`
}
// EmbyPlaybackInfoResponse 播放信息响应。
type EmbyPlaybackInfoResponse struct {
MediaSources []EmbyMediaSource `json:"MediaSources"`
PlaySessionId string `json:"PlaySessionId"`
}
// EmbyDeviceProfile 设备配置文件。
type EmbyDeviceProfile struct {
Name string `json:"Name,omitempty"`
MaxStaticBitrate int `json:"MaxStaticBitrate,omitempty"`
MaxStreamingBitrate int `json:"MaxStreamingBitrate,omitempty"`
MusicStreamingTranscodingBitrate int `json:"MusicStreamingTranscodingBitrate,omitempty"`
DirectPlayProfiles []EmbyDirectPlayProfile `json:"DirectPlayProfiles,omitempty"`
TranscodingProfiles []EmbyTranscodingProfile `json:"TranscodingProfiles,omitempty"`
ContainerProfiles []EmbyContainerProfile `json:"ContainerProfiles,omitempty"`
CodecProfiles []EmbyCodecProfile `json:"CodecProfiles,omitempty"`
SubtitleProfiles []EmbySubtitleProfile `json:"SubtitleProfiles,omitempty"`
}
// EmbyDirectPlayProfile 直接播放配置。
type EmbyDirectPlayProfile struct {
Container string `json:"Container,omitempty"`
AudioCodec string `json:"AudioCodec,omitempty"`
VideoCodec string `json:"VideoCodec,omitempty"`
Type string `json:"Type,omitempty"`
}
// EmbyTranscodingProfile 转码配置。
type EmbyTranscodingProfile struct {
Container string `json:"Container,omitempty"`
Type string `json:"Type,omitempty"`
VideoCodec string `json:"VideoCodec,omitempty"`
AudioCodec string `json:"AudioCodec,omitempty"`
Protocol string `json:"Protocol,omitempty"`
EstimateContentLength bool `json:"EstimateContentLength,omitempty"`
EnableMpegtsM2TsMode bool `json:"EnableMpegtsM2TsMode,omitempty"`
TranscodeSeekInfo string `json:"TranscodeSeekInfo,omitempty"`
Context string `json:"Context,omitempty"`
EnableSubtitlesInManifest bool `json:"EnableSubtitlesInManifest,omitempty"`
MaxAudioChannels int `json:"MaxAudioChannels,omitempty"`
MinSegments int `json:"MinSegments,omitempty"`
SegmentLength int `json:"SegmentLength,omitempty"`
BreakOnNonKeyFrames bool `json:"BreakOnNonKeyFrames,omitempty"`
}
// EmbyContainerProfile 容器配置。
type EmbyContainerProfile struct {
Type string `json:"Type,omitempty"`
Conditions []string `json:"Conditions,omitempty"`
Container string `json:"Container,omitempty"`
}
// EmbyCodecProfile 编解码器配置。
type EmbyCodecProfile struct {
Type string `json:"Type,omitempty"`
Conditions []EmbyProfileCondition `json:"Conditions,omitempty"`
Codec string `json:"Codec,omitempty"`
Container string `json:"Container,omitempty"`
}
// EmbyProfileCondition 配置条件。
type EmbyProfileCondition struct {
Condition string `json:"Condition,omitempty"`
Property string `json:"Property,omitempty"`
Value string `json:"Value,omitempty"`
IsRequired bool `json:"IsRequired,omitempty"`
}
// EmbySubtitleProfile 字幕配置。
type EmbySubtitleProfile struct {
Format string `json:"Format,omitempty"`
Method string `json:"Method,omitempty"`
DidlMode string `json:"DidlMode,omitempty"`
Language string `json:"Language,omitempty"`
Container string `json:"Container,omitempty"`
}
// ─── Emby 播放进度上报 ────────────────────────────────────────────────────────
// EmbyPlaybackProgressRequest 播放进度上报。
type EmbyPlaybackProgressRequest struct {
CanSeek bool `json:"CanSeek"`
ItemId string `json:"ItemId"`
MediaSourceId string `json:"MediaSourceId,omitempty"`
PositionTicks int64 `json:"PositionTicks"`
RunTimeTicks int64 `json:"RunTimeTicks,omitempty"`
IsPaused bool `json:"IsPaused"`
IsMuted bool `json:"IsMuted"`
VolumeLevel int `json:"VolumeLevel,omitempty"`
PlayMethod string `json:"PlayMethod,omitempty"`
PlaySessionId string `json:"PlaySessionId,omitempty"`
LiveStreamId string `json:"LiveStreamId,omitempty"`
QueueableMediaTypes []string `json:"QueueableMediaTypes,omitempty"`
}
// EmbyStopPlaybackRequest 停止播放上报。
type EmbyStopPlaybackRequest struct {
ItemId string `json:"ItemId"`
MediaSourceId string `json:"MediaSourceId,omitempty"`
PositionTicks int64 `json:"PositionTicks"`
RunTimeTicks int64 `json:"RunTimeTicks,omitempty"`
PlaySessionId string `json:"PlaySessionId,omitempty"`
LiveStreamId string `json:"LiveStreamId,omitempty"`
}
// EmbyUserDataRequest 用户数据更新。
type EmbyUserDataRequest struct {
PlaybackPositionTicks int64 `json:"PlaybackPositionTicks,omitempty"`
PlayCount int `json:"PlayCount,omitempty"`
IsFavorite bool `json:"IsFavorite,omitempty"`
Played bool `json:"Played,omitempty"`
PlayedPercentage float64 `json:"PlayedPercentage,omitempty"`
}
// ─── Emby Hubs ────────────────────────────────────────────────────────────────
// EmbyHubResponse Hub 响应。
type EmbyHubResponse struct {
Items []EmbyHubItem `json:"Items"`
}
// EmbyHubItem Hub 条目。
type EmbyHubItem struct {
Id string `json:"Id"`
Name string `json:"Name"`
Type string `json:"Type"`
Items []EmbyItem `json:"Items"`
TotalCount int `json:"TotalCount,omitempty"`
ImageUrl string `json:"ImageUrl,omitempty"`
}
// ─── Emby 字幕 ────────────────────────────────────────────────────────────────
// EmbyRemoteSubtitleInfo 远程字幕信息。
type EmbyRemoteSubtitleInfo struct {
ThreeLetterISOLanguageName string `json:"ThreeLetterISOLanguageName"`
Id string `json:"Id"`
ProviderName string `json:"ProviderName"`
Name string `json:"Name"`
Format string `json:"Format"`
Author string `json:"Author"`
Comment string `json:"Comment"`
DateCreated *time.Time `json:"DateCreated,omitempty"`
CommunityRating float64 `json:"CommunityRating,omitempty"`
DownloadCount int `json:"DownloadCount"`
IsHashMatch bool `json:"IsHashMatch,omitempty"`
IsForced bool `json:"IsForced,omitempty"`
IsHearingImpaired bool `json:"IsHearingImpaired,omitempty"`
}
// EmbySubtitleSearchRequest 字幕搜索请求。
type EmbySubtitleSearchRequest struct {
ItemId string `json:"ItemId"`
Language string `json:"Language"`
IsPerfectMatch bool `json:"IsPerfectMatch,omitempty"`
}
// EmbyImageRemoteInfo 远程图片信息。
type EmbyImageRemoteInfo struct {
Providers []EmbyImageProviderInfo `json:"Providers"`
TotalRecordCount int `json:"TotalRecordCount"`
}
// EmbyImageProviderInfo 图片提供者信息。
type EmbyImageProviderInfo struct {
Name string `json:"Name"`
RemoteImages []EmbyRemoteImageInfo `json:"RemoteImages,omitempty"`
SupportedImages []string `json:"SupportedImages"`
}
// EmbyRemoteImageInfo 远程图片信息。
type EmbyRemoteImageInfo struct {
Url string `json:"Url"`
ThumbnailUrl string `json:"ThumbnailUrl,omitempty"`
Height int `json:"Height"`
Width int `json:"Width"`
CommunityRating float64 `json:"CommunityRating,omitempty"`
VoteCount int `json:"VoteCount,omitempty"`
Language string `json:"Language,omitempty"`
Type string `json:"Type"`
RatingType string `json:"RatingType,omitempty"`
ProviderName string `json:"ProviderName"`
}
// ─── Emby Genre / Person ─────────────────────────────────────────────────────
// EmbyGenre 类型。
type EmbyGenre struct {
Name string `json:"Name"`
Id string `json:"Id,omitempty"`
}
// EmbyPersonInfo 人物信息。
type EmbyPersonInfo struct {
Id string `json:"Id"`
Name string `json:"Name"`
Type string `json:"Type,omitempty"`
PrimaryImageTag string `json:"PrimaryImageTag,omitempty"`
Overview string `json:"Overview,omitempty"`
BirthDate string `json:"BirthDate,omitempty"`
ProductionYear int `json:"ProductionYear,omitempty"`
EndDate string `json:"EndDate,omitempty"`
PremiereDate *time.Time `json:"PremiereDate,omitempty"`
}
// ─── Emby Active Encoding ────────────────────────────────────────────────────
// EmbyActiveEncodingRequest 活跃编码请求(客户端报告转码进度)。
type EmbyActiveEncodingRequest struct {
PlaySessionId string `json:"PlaySessionId"`
When string `json:"When"`
PositionTicks int64 `json:"PositionTicks,omitempty"`
IsPaused bool `json:"IsPaused,omitempty"`
IsUserPaused bool `json:"IsUserPaused,omitempty"`
}
+36 -30
View File
@@ -1,6 +1,5 @@
// Package model defines GORM data models and the registry used by
// auto-migration. Each subsystem in MediaStationGo owns a slice of tables
// here; AllModels returns the union for db.AutoMigrate.
// Package model 定义 GORM 数据模型和自动迁移使用的注册表。
// 每个子系统在 MediaStationGo 中拥有一个表切片;AllModels 返回联合以供 db.AutoMigrate 使用。
package model
import (
@@ -10,11 +9,11 @@ import (
"gorm.io/gorm"
)
// Base captures the fields embedded in every domain entity:
// Base 嵌入每个域实体共享的字段:
//
// - ID: UUID v4 string primary key.
// - CreatedAt / UpdatedAt: managed by GORM.
// - DeletedAt: soft-delete (queries auto-filter on it).
// - ID: UUID v4 字符串主键
// - CreatedAt / UpdatedAt: 由 GORM 管理
// - DeletedAt: 软删除(查询自动过滤)
type Base struct {
ID string `gorm:"primaryKey;type:varchar(36)" json:"id"`
CreatedAt time.Time `json:"created_at"`
@@ -22,7 +21,7 @@ type Base struct {
DeletedAt gorm.DeletedAt `gorm:"index" json:"-"`
}
// BeforeCreate generates a UUID if the caller did not supply one.
// BeforeCreate 如果调用者未提供则生成 UUID。
func (b *Base) BeforeCreate(_ *gorm.DB) error {
if b.ID == "" {
b.ID = uuid.NewString()
@@ -30,20 +29,23 @@ func (b *Base) BeforeCreate(_ *gorm.DB) error {
return nil
}
// User is a local account. The first registered admin (or seeded admin)
// gains the "admin" role; everyone else defaults to "user".
// User 是本地账户。第一个注册的管理员(或种子管理员)获得 "admin" 角色;
// 其他所有用户默认为 "user"。
type User struct {
Base
Username string `gorm:"uniqueIndex;size:64;not null" json:"username"`
PasswordHash string `gorm:"size:128;not null" json:"-"`
Role string `gorm:"size:16;not null;default:user" json:"role"`
Tier string `gorm:"size:16;default:free" json:"tier"` // free / plus
Nickname string `gorm:"size:128" json:"nickname,omitempty"`
Email string `gorm:"size:128" json:"email,omitempty"`
AvatarURL string `gorm:"size:255" json:"avatar_url,omitempty"`
ForcePasswordReset bool `gorm:"default:false" json:"force_password_reset"`
IsActive bool `gorm:"default:true" json:"is_active"`
LastLoginAt *time.Time `json:"last_login_at,omitempty"`
}
// Library represents a user-defined media root directory.
// Library 表示用户定义的媒体根目录。
type Library struct {
Base
Name string `gorm:"size:128;not null" json:"name"`
@@ -52,8 +54,7 @@ type Library struct {
Enabled bool `gorm:"default:true" json:"enabled"`
}
// Media is a single playable item. Series episodes link to a SeriesID; movies
// have SeriesID == "".
// Media 是单个可播放项。剧集链接到 SeriesID;电影 SeriesID == ""。
type Media struct {
Base
LibraryID string `gorm:"index;size:36" json:"library_id"`
@@ -116,7 +117,7 @@ type APIConfig struct {
Description string `gorm:"size:255" json:"description,omitempty"`
}
// Series groups episodes that belong to the same show.
// Series 将属于同一节目的剧集分组。
type Series struct {
Base
LibraryID string `gorm:"index;size:36" json:"library_id"`
@@ -130,7 +131,7 @@ type Series struct {
BangumiID int `json:"bangumi_id"`
}
// PlaybackHistory records the current playback position for resume support.
// PlaybackHistory 记录当前播放位置以支持续播。
type PlaybackHistory struct {
Base
UserID string `gorm:"index;size:36;not null" json:"user_id"`
@@ -141,14 +142,14 @@ type PlaybackHistory struct {
Completed bool `json:"completed"`
}
// Favorite marks a media item as favourite for a given user.
// Favorite 将媒体项标记为给定用户的收藏。
type Favorite struct {
Base
UserID string `gorm:"index;size:36;not null;uniqueIndex:uniq_user_media" json:"user_id"`
MediaID string `gorm:"index;size:36;not null;uniqueIndex:uniq_user_media" json:"media_id"`
}
// Playlist is a user-curated, ordered list of media items.
// Playlist 是用户策划的、有序的媒体列表。
type Playlist struct {
Base
UserID string `gorm:"index;size:36;not null" json:"user_id"`
@@ -156,7 +157,7 @@ type Playlist struct {
IsPublic bool `gorm:"default:false" json:"is_public"`
}
// PlaylistItem is the join table between Playlists and Media with ordering.
// PlaylistItem 是 Playlist 和 Media 的连接表,带有排序。
type PlaylistItem struct {
Base
PlaylistID string `gorm:"index;size:36;not null" json:"playlist_id"`
@@ -164,7 +165,7 @@ type PlaylistItem struct {
Position int `json:"position"`
}
// DownloadTask is an outstanding (or completed) torrent / HTTP download.
// DownloadTask 是待处理(或已完成)的 torrent / HTTP 下载。
type DownloadTask struct {
Base
UserID string `gorm:"index;size:36" json:"user_id"`
@@ -175,27 +176,25 @@ type DownloadTask struct {
Progress float32 `json:"progress"`
}
// Subscription is an automation rule that polls an RSS feed and queues
// matching torrents into the configured download client.
// Subscription 是自动化规则,轮询 RSS 源并将匹配种子排队到配置的下载客户端。
type Subscription struct {
Base
UserID string `gorm:"index;size:36" json:"user_id"`
Name string `gorm:"size:128;not null" json:"name"`
FeedURL string `gorm:"size:2048;not null" json:"feed_url"`
Filter string `gorm:"size:512" json:"filter"`
Enabled bool `gorm:"default:true" json:"enabled"`
UserID string `gorm:"index;size:36" json:"user_id"`
Name string `gorm:"size:128;not null" json:"name"`
FeedURL string `gorm:"size:2048;not null" json:"feed_url"`
Filter string `gorm:"size:512" json:"filter"`
Enabled bool `gorm:"default:true" json:"enabled"`
LastRunAt *time.Time `json:"last_run_at,omitempty"`
}
// Setting is a single key/value system-wide preference (used by the admin UI).
// Setting 是单个键/值系统级偏好(供管理 UI 使用)。
type Setting struct {
Key string `gorm:"primaryKey;size:128" json:"key"`
Value string `gorm:"type:text" json:"value"`
UpdatedAt time.Time `json:"updated_at"`
}
// AccessLog is a structured audit-trail entry. Stored in SQLite for the
// admin Activity panel.
// AccessLog 是结构化审计跟踪条目。存储在 SQLite 中供管理活动面板使用。
type AccessLog struct {
Base
UserID string `gorm:"index;size:36" json:"user_id"`
@@ -205,7 +204,7 @@ type AccessLog struct {
Detail string `gorm:"type:text" json:"detail"`
}
// AllModels returns the slice consumed by gorm.AutoMigrate.
// AllModels 返回 gorm.AutoMigrate 使用的切片。
func AllModels() []interface{} {
return []interface{}{
&User{},
@@ -221,5 +220,12 @@ func AllModels() []interface{} {
&Setting{},
&AccessLog{},
&APIConfig{},
&UserPermission{},
&RefreshToken{},
&ApiConfig{},
&DownloadClient{},
&NotifyChannel{},
&Site{},
&STRMRecord{},
}
}
+14
View File
@@ -0,0 +1,14 @@
// Package model 定义通知渠道配置数据模型。
package model
// NotifyChannel 通知渠道配置。
// 支持多种通知渠道:telegram / wechat / bark / webhook / email。
// Events 字段存储 JSON array,表示该渠道订阅的事件类型。
type NotifyChannel struct {
Base
Name string `gorm:"size:128;not null" json:"name"`
Type string `gorm:"size:32;not null" json:"type"` // telegram / wechat / bark / webhook / email
Enabled bool `gorm:"default:true" json:"enabled"`
Config string `gorm:"type:text" json:"-"` // JSON配置, AES加密
Events string `gorm:"type:text" json:"events"` // 订阅的事件列表, JSON array
}
+103
View File
@@ -0,0 +1,103 @@
// Package model 定义权限相关的数据模型。
package model
import (
"time"
"github.com/google/uuid"
"gorm.io/gorm"
)
// UserPermission 定义用户细粒度权限(19项)。
// 默认开启(6项):CanViewDashboard, CanPlayMedia, CanCast, CanExternalPlayer, CanFavorite, CanViewHistory
// 默认关闭(13项):其他权限需要管理员分配
type UserPermission struct {
ID string `gorm:"primaryKey;size:36" json:"id"`
UserID string `gorm:"uniqueIndex;size:36;not null" json:"user_id"`
// 默认开启(6项)- Basic
CanViewDashboard bool `gorm:"default:true" json:"can_view_dashboard"`
CanPlayMedia bool `gorm:"default:true" json:"can_play_media"`
CanCast bool `gorm:"default:true" json:"can_cast"`
CanExternalPlayer bool `gorm:"default:true" json:"can_external_player"`
CanFavorite bool `gorm:"default:true" json:"can_favorite"`
CanViewHistory bool `gorm:"default:true" json:"can_view_history"`
// 默认关闭(13项)- Advanced
CanEditMedia bool `gorm:"default:false" json:"can_edit_media"`
CanRescrape bool `gorm:"default:false" json:"can_rescrape"`
CanUseAI bool `gorm:"default:false" json:"can_use_ai"`
CanCaptureFrames bool `gorm:"default:false" json:"can_capture_frames"`
CanManageDownloads bool `gorm:"default:false" json:"can_manage_downloads"`
CanViewDiscover bool `gorm:"default:false" json:"can_view_discover"`
CanManageSubscriptions bool `gorm:"default:false" json:"can_manage_subscriptions"`
CanManageSites bool `gorm:"default:false" json:"can_manage_sites"`
CanUseAIAssistant bool `gorm:"default:false" json:"can_use_ai_assistant"`
CanManageUsers bool `gorm:"default:false" json:"can_manage_users"`
CanManageFiles bool `gorm:"default:false" json:"can_manage_files"`
CanManageStrm bool `gorm:"default:false" json:"can_manage_strm"`
CanAccessSettings bool `gorm:"default:false" json:"can_access_settings"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// BeforeCreate 生成 UUID。
func (p *UserPermission) BeforeCreate(_ *gorm.DB) error {
if p.ID == "" {
p.ID = uuid.NewString()
}
return nil
}
// NewDefaultPermission 创建带有默认权限的 UserPermission。
func NewDefaultPermission(userID string) *UserPermission {
return &UserPermission{
ID: uuid.NewString(),
UserID: userID,
CanViewDashboard: true,
CanPlayMedia: true,
CanCast: true,
CanExternalPlayer: true,
CanFavorite: true,
CanViewHistory: true,
CanEditMedia: false,
CanRescrape: false,
CanUseAI: false,
CanCaptureFrames: false,
CanManageDownloads: false,
CanViewDiscover: false,
CanManageSubscriptions: false,
CanManageSites: false,
CanUseAIAssistant: false,
CanManageUsers: false,
CanManageFiles: false,
CanManageStrm: false,
CanAccessSettings: false,
}
}
// PermissionMap 将权限结构转换为 map[string]bool 便于检查。
func (p *UserPermission) PermissionMap() map[string]bool {
return map[string]bool{
"can_view_dashboard": p.CanViewDashboard,
"can_play_media": p.CanPlayMedia,
"can_cast": p.CanCast,
"can_external_player": p.CanExternalPlayer,
"can_favorite": p.CanFavorite,
"can_view_history": p.CanViewHistory,
"can_edit_media": p.CanEditMedia,
"can_rescrape": p.CanRescrape,
"can_use_ai": p.CanUseAI,
"can_capture_frames": p.CanCaptureFrames,
"can_manage_downloads": p.CanManageDownloads,
"can_view_discover": p.CanViewDiscover,
"can_manage_subscriptions": p.CanManageSubscriptions,
"can_manage_sites": p.CanManageSites,
"can_use_ai_assistant": p.CanUseAIAssistant,
"can_manage_users": p.CanManageUsers,
"can_manage_files": p.CanManageFiles,
"can_manage_strm": p.CanManageStrm,
"can_access_settings": p.CanAccessSettings,
}
}
+38
View File
@@ -0,0 +1,38 @@
// Package model 定义刷新令牌数据模型。
package model
import (
"time"
"github.com/google/uuid"
"gorm.io/gorm"
)
// RefreshToken 用于双令牌认证机制中的刷新令牌。
// 存储时使用 SHA256 哈希,原始令牌不存储。
type RefreshToken struct {
ID string `gorm:"primaryKey;size:36" json:"id"`
UserID string `gorm:"index;size:36;not null" json:"user_id"`
TokenHash string `gorm:"uniqueIndex;size:128;not null" json:"-"`
ExpiresAt time.Time `gorm:"index" json:"expires_at"`
CreatedAt time.Time `json:"created_at"`
Revoked bool `gorm:"default:false" json:"revoked"`
}
// BeforeCreate 生成 UUID。
func (t *RefreshToken) BeforeCreate(_ *gorm.DB) error {
if t.ID == "" {
t.ID = uuid.NewString()
}
return nil
}
// IsExpired 检查刷新令牌是否已过期。
func (t *RefreshToken) IsExpired() bool {
return time.Now().After(t.ExpiresAt)
}
// IsValid 检查刷新令牌是否有效(未撤销且未过期)。
func (t *RefreshToken) IsValid() bool {
return !t.Revoked && !t.IsExpired()
}
+33
View File
@@ -0,0 +1,33 @@
// Package model — PT 站点配置数据模型。
package model
import (
"time"
)
// Site PT 站点配置。
type Site struct {
Base
Name string `gorm:"size:128;not null" json:"name"`
Type string `gorm:"size:32;not null" json:"type"` // nexusphp / gazelle / unit3d / mteam / discuz / custom_rss
URL string `gorm:"size:512;not null" json:"url"`
AuthType string `gorm:"size:32;not null" json:"auth_type"` // cookie / api_key / auth_header
Cookie string `gorm:"type:text" json:"-"` // AES 加密
APIKey string `gorm:"type:text" json:"-"` // AES 加密
AuthHeader string `gorm:"type:text" json:"-"` // AES 加密
Enabled bool `gorm:"default:true" json:"enabled"`
IsDefault bool `gorm:"default:false" json:"is_default"`
Extra string `gorm:"type:text" json:"-"` // JSON 扩展配置, AES 加密
LastError string `gorm:"size:1024" json:"last_error"`
LastCheckAt *time.Time `json:"last_check_at"`
}
// SiteType 返回支持的站点类型列表。
func SiteTypes() []string {
return []string{"nexusphp", "gazelle", "unit3d", "mteam", "discuz", "custom_rss"}
}
// AuthTypes 返回支持的认证方式列表。
func AuthTypes() []string {
return []string{"cookie", "api_key", "auth_header"}
}
+35
View File
@@ -0,0 +1,35 @@
// Package model — STRM 文件记录数据模型。
package model
// STRMRecord STRM 文件记录。
// 外部存储以"文件"形式入库,URL 指向实际资源。
type STRMRecord struct {
Base
Title string `gorm:"size:512;not null;index" json:"title"`
URL string `gorm:"size:2048;not null" json:"url"` // STRM 文件指向的 URL
FilePath string `gorm:"size:1024;not null" json:"file_path"` // 本地 STRM 文件路径
Protocol string `gorm:"size:32;not null" json:"protocol"` // webdav / alist / s3 / http / https
FileSize int64 `json:"file_size"`
MediaID string `gorm:"size:36;index" json:"media_id"` // 关联媒体 ID
MediaType string `gorm:"size:16" json:"media_type"` // movie / series
SeasonNum int `json:"season_num"`
EpisodeNum int `json:"episode_num"`
}
// AllowedSTRMProtocols 协议白名单。
var AllowedSTRMProtocols = []string{
"webdav", "davs",
"alist", "alists",
"s3",
"http", "https",
}
// IsAllowedProtocol 检查协议是否在白名单中。
func IsAllowedProtocol(protocol string) bool {
for _, p := range AllowedSTRMProtocols {
if p == protocol {
return true
}
}
return false
}
@@ -0,0 +1,92 @@
// Package repository 实现下载客户端配置的数据访问层。
package repository
import (
"context"
"errors"
"time"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// DownloadClientRepository persists model.DownloadClient records.
type DownloadClientRepository struct{ db *gorm.DB }
// Create inserts a new download client.
func (r *DownloadClientRepository) Create(ctx context.Context, c *model.DownloadClient) error {
return r.db.WithContext(ctx).Create(c).Error
}
// FindByID returns the download client by ID, or (nil, nil) when absent.
func (r *DownloadClientRepository) FindByID(ctx context.Context, id string) (*model.DownloadClient, error) {
var c model.DownloadClient
err := r.db.WithContext(ctx).Where("id = ?", id).First(&c).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &c, nil
}
// FindDefault returns the default download client, or (nil, nil).
func (r *DownloadClientRepository) FindDefault(ctx context.Context) (*model.DownloadClient, error) {
var c model.DownloadClient
err := r.db.WithContext(ctx).Where("is_default = ? AND enabled = ?", true, true).First(&c).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &c, nil
}
// List returns all download clients ordered by creation time.
func (r *DownloadClientRepository) List(ctx context.Context) ([]model.DownloadClient, error) {
var rows []model.DownloadClient
err := r.db.WithContext(ctx).Order("created_at asc").Find(&rows).Error
return rows, err
}
// ListEnabled returns all enabled download clients.
func (r *DownloadClientRepository) ListEnabled(ctx context.Context) ([]model.DownloadClient, error) {
var rows []model.DownloadClient
err := r.db.WithContext(ctx).Where("enabled = ?", true).Order("created_at asc").Find(&rows).Error
return rows, err
}
// Update persists changes to a download client.
func (r *DownloadClientRepository) Update(ctx context.Context, c *model.DownloadClient) error {
return r.db.WithContext(ctx).Save(c).Error
}
// Delete removes a download client (soft-delete).
func (r *DownloadClientRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.DownloadClient{}, "id = ?", id).Error
}
// ClearDefault unsets the default flag for all clients.
func (r *DownloadClientRepository) ClearDefault(ctx context.Context) error {
return r.db.WithContext(ctx).Model(&model.DownloadClient{}).
Where("is_default = ?", true).Update("is_default", false).Error
}
// SetDefault sets a specific client as default and clears others.
func (r *DownloadClientRepository) SetDefault(ctx context.Context, id string) error {
now := time.Now()
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&model.DownloadClient{}).
Where("is_default = ?", true).Update("is_default", false).Error; err != nil {
return err
}
return tx.Model(&model.DownloadClient{}).
Where("id = ?", id).Updates(map[string]any{
"is_default": true,
"updated_at": now,
}).Error
})
}
@@ -0,0 +1,67 @@
// Package repository 实现通知渠道配置的数据访问层。
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// NotifyChannelRepository persists model.NotifyChannel records.
type NotifyChannelRepository struct{ db *gorm.DB }
// Create inserts a new notification channel.
func (r *NotifyChannelRepository) Create(ctx context.Context, c *model.NotifyChannel) error {
return r.db.WithContext(ctx).Create(c).Error
}
// FindByID returns the notification channel by ID, or (nil, nil) when absent.
func (r *NotifyChannelRepository) FindByID(ctx context.Context, id string) (*model.NotifyChannel, error) {
var c model.NotifyChannel
err := r.db.WithContext(ctx).Where("id = ?", id).First(&c).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &c, nil
}
// List returns all notification channels ordered by creation time.
func (r *NotifyChannelRepository) List(ctx context.Context) ([]model.NotifyChannel, error) {
var rows []model.NotifyChannel
err := r.db.WithContext(ctx).Order("created_at asc").Find(&rows).Error
return rows, err
}
// ListEnabled returns all enabled notification channels.
func (r *NotifyChannelRepository) ListEnabled(ctx context.Context) ([]model.NotifyChannel, error) {
var rows []model.NotifyChannel
err := r.db.WithContext(ctx).Where("enabled = ?", true).Order("created_at asc").Find(&rows).Error
return rows, err
}
// ListByEvent returns all enabled channels that subscribe to the given event type.
func (r *NotifyChannelRepository) ListByEvent(ctx context.Context, eventType string) ([]model.NotifyChannel, error) {
var rows []model.NotifyChannel
// Events is a JSON array stored as text; use LIKE for simple matching.
// This works for exact event type matches within the JSON array.
err := r.db.WithContext(ctx).
Where("enabled = ? AND events LIKE ?", true, "%\""+eventType+"\"%").
Find(&rows).Error
return rows, err
}
// Update persists changes to a notification channel.
func (r *NotifyChannelRepository) Update(ctx context.Context, c *model.NotifyChannel) error {
return r.db.WithContext(ctx).Save(c).Error
}
// Delete removes a notification channel (soft-delete).
func (r *NotifyChannelRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.NotifyChannel{}, "id = ?", id).Error
}
+180 -10
View File
@@ -1,13 +1,14 @@
// Package repository implements a thin GORM-based data-access layer over the
// types declared in internal/model. Each method takes a context.Context so we
// can plug in cancellation / tracing later.
// Package repository 实现基于 GORM 的数据访问层。
// 每个方法接受 context.Context 以便后续插入取消/追踪。
//
// Repositories are intentionally narrow: they only know how to persist data,
// not how to interpret it. Domain logic lives in internal/service.
// Repository 故意保持精简:它们只负责持久化数据,不处理业务逻辑。
// 业务逻辑位于 internal/service。
package repository
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"time"
@@ -16,7 +17,7 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// Container is the registry of all repositories injected into services.
// Container 是所有 repositories 的注册表,注入到 services 中。
type Container struct {
DB *gorm.DB
User *UserRepository
@@ -30,9 +31,16 @@ type Container struct {
Subscription *SubscriptionRepository
Setting *SettingRepository
Log *AccessLogRepository
Permission *PermissionRepository
RefreshToken *RefreshTokenRepository
ApiConfig *ApiConfigRepository
DownloadClient *DownloadClientRepository
NotifyChannel *NotifyChannelRepository
Site *SiteRepository
STRM *STRMRepository
}
// New wires every repository to a single *gorm.DB.
// New 将每个 repository 连接到单个 *gorm.DB。
func New(db *gorm.DB) *Container {
return &Container{
DB: db,
@@ -47,6 +55,13 @@ func New(db *gorm.DB) *Container {
Subscription: &SubscriptionRepository{db: db},
Setting: &SettingRepository{db: db},
Log: &AccessLogRepository{db: db},
Permission: &PermissionRepository{db: db},
RefreshToken: &RefreshTokenRepository{db: db},
ApiConfig: &ApiConfigRepository{db: db},
DownloadClient: &DownloadClientRepository{db: db},
NotifyChannel: &NotifyChannelRepository{db: db},
Site: &SiteRepository{db: db},
STRM: &STRMRepository{db: db},
}
}
@@ -267,7 +282,7 @@ func (r *HistoryRepository) ListByUser(ctx context.Context, userID string, limit
return rows, err
}
// ─── Favorite ────────────────────────────────────────────────────────────────
// ─── Favorite ───────────────────────────────────────────────────────────────
// FavoriteRepository persists model.Favorite records.
type FavoriteRepository struct{ db *gorm.DB }
@@ -311,7 +326,7 @@ func (r *PlaylistRepository) ListByUser(ctx context.Context, userID string) ([]m
return rows, err
}
// ─── Download ────────────────────────────────────────────────────────────────
// ─── Download ───────────────────────────────────────────────────────────────
// DownloadRepository persists model.DownloadTask records.
type DownloadRepository struct{ db *gorm.DB }
@@ -328,7 +343,7 @@ func (r *DownloadRepository) List(ctx context.Context) ([]model.DownloadTask, er
return rows, err
}
// ─── Subscription ────────────────────────────────────────────────────────────
// ─── Subscription ───────────────────────────────────────────────────────────
// SubscriptionRepository persists model.Subscription records.
type SubscriptionRepository struct{ db *gorm.DB }
@@ -389,3 +404,158 @@ func (r *AccessLogRepository) Recent(ctx context.Context, limit int) ([]model.Ac
err := r.db.WithContext(ctx).Order("created_at desc").Limit(limit).Find(&rows).Error
return rows, err
}
// ─── Permission ──────────────────────────────────────────────────────────────
// 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 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 := 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 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 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 r.db.WithContext(ctx).Where("user_id = ?", userID).Delete(&model.UserPermission{}).Error
}
// ─── Refresh Token ───────────────────────────────────────────────────────────
// 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 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 := 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 r.db.WithContext(ctx).Model(&model.RefreshToken{}).
Where("user_id = ?", userID).Update("revoked", true).Error
}
// DeleteExpired removes all expired refresh tokens.
func (r *RefreshTokenRepository) DeleteExpired(ctx context.Context) 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 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[:])
}
// ─── API Config ──────────────────────────────────────────────────────────────
// 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{
APIKey: c.APIKey,
BaseURL: c.BaseURL,
Extra: c.Extra,
Enabled: c.Enabled,
UpdatedAt: time.Now(),
}).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
}
+56
View File
@@ -0,0 +1,56 @@
// Package repository — PT 站点数据访问层。
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// SiteRepository persists model.Site records.
type SiteRepository struct{ db *gorm.DB }
// Create inserts a new site.
func (r *SiteRepository) Create(ctx context.Context, s *model.Site) error {
return r.db.WithContext(ctx).Create(s).Error
}
// FindByID returns the site by ID, or (nil, nil) when absent.
func (r *SiteRepository) FindByID(ctx context.Context, id string) (*model.Site, error) {
var s model.Site
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 sites ordered by name.
func (r *SiteRepository) List(ctx context.Context) ([]model.Site, error) {
var rows []model.Site
err := r.db.WithContext(ctx).Order("name asc").Find(&rows).Error
return rows, err
}
// ListEnabled returns all enabled sites.
func (r *SiteRepository) ListEnabled(ctx context.Context) ([]model.Site, error) {
var rows []model.Site
err := r.db.WithContext(ctx).Where("enabled = ?", true).Order("name asc").Find(&rows).Error
return rows, err
}
// Update updates site fields.
func (r *SiteRepository) Update(ctx context.Context, s *model.Site) error {
return r.db.WithContext(ctx).Save(s).Error
}
// Delete removes a site (soft-delete).
func (r *SiteRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.Site{}, "id = ?", id).Error
}
+82
View File
@@ -0,0 +1,82 @@
// Package repository — STRM 文件记录数据访问层。
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// STRMRepository persists model.STRMRecord records.
type STRMRepository struct{ db *gorm.DB }
// Create inserts a new STRM record.
func (r *STRMRepository) Create(ctx context.Context, s *model.STRMRecord) error {
return r.db.WithContext(ctx).Create(s).Error
}
// CreateBatch inserts multiple STRM records.
func (r *STRMRepository) CreateBatch(ctx context.Context, records []model.STRMRecord) error {
if len(records) == 0 {
return nil
}
return r.db.WithContext(ctx).CreateInBatches(records, 100).Error
}
// FindByID returns the STRM record by ID, or (nil, nil) when absent.
func (r *STRMRepository) FindByID(ctx context.Context, id string) (*model.STRMRecord, error) {
var s model.STRMRecord
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 STRM records with optional filters. Supports pagination.
// Filters: media_id, media_type, protocol
func (r *STRMRepository) List(ctx context.Context, filters map[string]string, offset, limit int) ([]model.STRMRecord, int64, error) {
q := r.db.WithContext(ctx).Model(&model.STRMRecord{})
if mediaID, ok := filters["media_id"]; ok && mediaID != "" {
q = q.Where("media_id = ?", mediaID)
}
if mediaType, ok := filters["media_type"]; ok && mediaType != "" {
q = q.Where("media_type = ?", mediaType)
}
if protocol, ok := filters["protocol"]; ok && protocol != "" {
q = q.Where("protocol = ?", protocol)
}
var total int64
if err := q.Count(&total).Error; err != nil {
return nil, 0, err
}
var rows []model.STRMRecord
err := q.Order("created_at desc").Offset(offset).Limit(limit).Find(&rows).Error
return rows, total, err
}
// Update updates a STRM record.
func (r *STRMRepository) Update(ctx context.Context, s *model.STRMRecord) error {
return r.db.WithContext(ctx).Save(s).Error
}
// Delete removes a STRM record (soft-delete).
func (r *STRMRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.STRMRecord{}, "id = ?", id).Error
}
// FindByMediaID returns STRM records for a given media ID.
func (r *STRMRepository) FindByMediaID(ctx context.Context, mediaID string) ([]model.STRMRecord, error) {
var rows []model.STRMRecord
err := r.db.WithContext(ctx).Where("media_id = ?", mediaID).Find(&rows).Error
return rows, err
}
+383
View File
@@ -0,0 +1,383 @@
// Package service — API 配置管理服务。
package service
import (
"context"
"errors"
"fmt"
"net/http"
"net/url"
"strings"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// ApiConfigService 负责第三方 API 配置的 CRUD 和加密管理。
type ApiConfigService struct {
cfg *config.Config
log *zap.Logger
repo *repository.Container
crypto *CryptoService
}
// NewApiConfigService 创建 API 配置服务实例。
func NewApiConfigService(cfg *config.Config, log *zap.Logger, repo *repository.Container, crypto *CryptoService) *ApiConfigService {
return &ApiConfigService{cfg: cfg, log: log, repo: repo, crypto: crypto}
}
// ApiConfigService 错误定义。
var (
ErrApiConfigNotFound = errors.New("API configuration not found")
ErrInvalidProvider = errors.New("invalid provider")
ErrTestFailed = errors.New("connection test failed")
)
// GetByProvider 获取指定提供者的 API 配置。
func (s *ApiConfigService) GetByProvider(ctx context.Context, provider string) (*model.ApiConfig, error) {
cfg, err := s.repo.ApiConfig.FindByProvider(ctx, provider)
if err != nil {
return nil, err
}
if cfg == nil {
return nil, ErrApiConfigNotFound
}
// 解密敏感字段
if cfg.APIKey != "" && s.crypto.IsEncrypted(cfg.APIKey) {
cfg.APIKey = s.crypto.Decrypt(cfg.APIKey)
}
return cfg, nil
}
// List 返回所有 API 配置。
func (s *ApiConfigService) List(ctx context.Context) ([]model.ApiConfig, error) {
configs, err := s.repo.ApiConfig.List(ctx)
if err != nil {
return nil, err
}
// 解密敏感字段
for i := range configs {
if configs[i].APIKey != "" && s.crypto.IsEncrypted(configs[i].APIKey) {
configs[i].APIKey = s.crypto.Decrypt(configs[i].APIKey)
}
}
return configs, nil
}
// GetProviders 返回预定义的提供者列表。
func (s *ApiConfigService) GetProviders() []model.ApiProvider {
return model.PredefinedProviders()
}
// Upsert 创建或更新 API 配置,自动加密敏感字段。
func (s *ApiConfigService) Upsert(ctx context.Context, provider string, apiKey, baseURL, extra string, enabled bool) (*model.ApiConfig, error) {
// 验证提供者是否有效
if !s.isValidProvider(provider) {
return nil, ErrInvalidProvider
}
// 加密 API Key
encryptedKey := apiKey
if apiKey != "" && !s.crypto.IsEncrypted(apiKey) {
encryptedKey = s.crypto.Encrypt(apiKey)
}
cfg := &model.ApiConfig{
Provider: provider,
APIKey: encryptedKey,
BaseURL: baseURL,
Extra: extra,
Enabled: enabled,
Description: s.getProviderDescription(provider),
}
if err := s.repo.ApiConfig.Upsert(ctx, cfg); err != nil {
return nil, err
}
// 返回解密后的配置
cfg.APIKey = apiKey
return cfg, nil
}
// Delete 删除 API 配置。
func (s *ApiConfigService) Delete(ctx context.Context, provider string) error {
return s.repo.ApiConfig.Delete(ctx, provider)
}
// Update 更新 API 配置。
func (s *ApiConfigService) Update(ctx context.Context, provider string, apiKey, baseURL, extra string, enabled bool) error {
// 加密 API Key
encryptedKey := apiKey
if apiKey != "" && !s.crypto.IsEncrypted(apiKey) {
encryptedKey = s.crypto.Encrypt(apiKey)
}
cfg := &model.ApiConfig{
Provider: provider,
APIKey: encryptedKey,
BaseURL: baseURL,
Extra: extra,
Enabled: enabled,
}
return s.repo.ApiConfig.Update(ctx, cfg)
}
// TestConnection 测试 API 连接。
func (s *ApiConfigService) TestConnection(ctx context.Context, provider string) (string, error) {
cfg, err := s.GetByProvider(ctx, provider)
if err != nil {
return "error", err
}
// 根据不同提供者执行不同的测试逻辑
switch provider {
case "tmdb":
return s.testTMDb(cfg)
case "openai":
return s.testOpenAI(cfg)
case "deepseek":
return s.testDeepSeek(cfg)
case "siliconflow":
return s.testSiliconFlow(cfg)
default:
return "unknown", fmt.Errorf("no test implemented for provider: %s", provider)
}
}
// testTMDb 测试 TMDb API 连接。
func (s *ApiConfigService) testTMDb(cfg *model.ApiConfig) (string, error) {
if cfg.APIKey == "" {
return "error", errors.New("API key is required")
}
testURL := "https://api.themoviedb.org/3/configuration?api_key=" + cfg.APIKey
resp, err := http.Get(testURL)
if err != nil {
// 如果配置了代理,使用代理
if s.cfg.Secrets.TMDbAPIProxy != "" {
proxyURL := s.cfg.Secrets.TMDbAPIProxy + "?api_key=" + cfg.APIKey
resp, err = http.Get(proxyURL)
if err != nil {
return "error", fmt.Errorf("TMDb connection failed: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == 200 {
return "success", nil
}
return "error", fmt.Errorf("TMDb API returned status %d", resp.StatusCode)
}
return "error", fmt.Errorf("TMDb connection failed: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == 200 {
return "success", nil
}
if resp.StatusCode == 401 {
return "invalid", errors.New("invalid API key")
}
return "error", fmt.Errorf("TMDb API returned status %d", resp.StatusCode)
}
// testOpenAI 测试 OpenAI API 连接。
func (s *ApiConfigService) testOpenAI(cfg *model.ApiConfig) (string, error) {
if cfg.APIKey == "" {
return "error", errors.New("API key is required")
}
baseURL := cfg.BaseURL
if baseURL == "" {
baseURL = "https://api.openai.com/v1"
}
testURL := baseURL + "/models"
req, err := http.NewRequest("GET", testURL, nil)
if err != nil {
return "error", err
}
req.Header.Set("Authorization", "Bearer "+cfg.APIKey)
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Do(req.WithContext(context.Background()))
if err != nil {
return "error", fmt.Errorf("OpenAI connection failed: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == 200 {
return "success", nil
}
if resp.StatusCode == 401 {
return "invalid", errors.New("invalid API key")
}
return "error", fmt.Errorf("OpenAI API returned status %d", resp.StatusCode)
}
// testDeepSeek 测试 DeepSeek API 连接。
func (s *ApiConfigService) testDeepSeek(cfg *model.ApiConfig) (string, error) {
if cfg.APIKey == "" {
return "error", errors.New("API key is required")
}
baseURL := cfg.BaseURL
if baseURL == "" {
baseURL = "https://api.deepseek.com"
}
testURL := baseURL + "/models"
req, err := http.NewRequest("GET", testURL, nil)
if err != nil {
return "error", err
}
req.Header.Set("Authorization", "Bearer "+cfg.APIKey)
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Do(req.WithContext(context.Background()))
if err != nil {
return "error", fmt.Errorf("DeepSeek connection failed: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == 200 {
return "success", nil
}
if resp.StatusCode == 401 {
return "invalid", errors.New("invalid API key")
}
return "error", fmt.Errorf("DeepSeek API returned status %d", resp.StatusCode)
}
// testSiliconFlow 测试 SiliconFlow API 连接。
func (s *ApiConfigService) testSiliconFlow(cfg *model.ApiConfig) (string, error) {
if cfg.APIKey == "" {
return "error", errors.New("API key is required")
}
baseURL := cfg.BaseURL
if baseURL == "" {
baseURL = "https://api.siliconflow.cn/v1"
}
testURL := baseURL + "/models"
req, err := http.NewRequest("GET", testURL, nil)
if err != nil {
return "error", err
}
req.Header.Set("Authorization", "Bearer "+cfg.APIKey)
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Do(req.WithContext(context.Background()))
if err != nil {
return "error", fmt.Errorf("SiliconFlow connection failed: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == 200 {
return "success", nil
}
if resp.StatusCode == 401 {
return "invalid", errors.New("invalid API key")
}
return "error", fmt.Errorf("SiliconFlow API returned status %d", resp.StatusCode)
}
// GetEffectiveConfig 获取生效的 API 配置(数据库配置优先于配置文件)。
func (s *ApiConfigService) GetEffectiveConfig(ctx context.Context, provider string) (*model.ApiConfig, error) {
// 首先尝试从数据库获取
cfg, err := s.GetByProvider(ctx, provider)
if err == nil && cfg != nil {
return cfg, nil
}
// 如果数据库没有,尝试从配置文件获取
return s.getConfigFromFile(provider)
}
// getConfigFromFile 从配置文件获取 API 配置。
func (s *ApiConfigService) getConfigFromFile(provider string) (*model.ApiConfig, error) {
var apiKey string
var hasKey bool
switch provider {
case "tmdb":
apiKey = s.cfg.Secrets.TMDbAPIKey
hasKey = apiKey != ""
case "bangumi":
apiKey = s.cfg.Secrets.BangumiToken
hasKey = apiKey != ""
case "thetvdb":
apiKey = s.cfg.Secrets.TheTVDBAPIKey
hasKey = apiKey != ""
case "fanart":
apiKey = s.cfg.Secrets.FanartAPIKey
hasKey = apiKey != ""
}
if !hasKey {
return nil, ErrApiConfigNotFound
}
return &model.ApiConfig{
Provider: provider,
APIKey: apiKey,
Enabled: true,
}, nil
}
// isValidProvider 检查提供者是否有效。
func (s *ApiConfigService) isValidProvider(provider string) bool {
providers := model.PredefinedProviders()
for _, p := range providers {
if p.ID == provider {
return true
}
}
return false
}
// getProviderDescription 获取提供者描述。
func (s *ApiConfigService) getProviderDescription(provider string) string {
providers := model.PredefinedProviders()
for _, p := range providers {
if p.ID == provider {
return p.Description
}
}
return ""
}
// UpdateTestResult 更新测试结果。
func (s *ApiConfigService) UpdateTestResult(ctx context.Context, provider, result string) error {
return s.repo.ApiConfig.UpdateTestResult(ctx, provider, result)
}
// MaskAPIKey 遮蔽 API Key 的中间部分。
func (s *ApiConfigService) MaskAPIKey(apiKey string) string {
if len(apiKey) <= 8 {
return "***"
}
return apiKey[:4] + "..." + apiKey[len(apiKey)-4:]
}
// ExtractBaseURL 从 URL 中提取域名。
func ExtractBaseURL(rawURL string) string {
if rawURL == "" {
return ""
}
u, err := url.Parse(rawURL)
if err != nil {
return rawURL
}
return u.Scheme + "://" + u.Host
}
// ProviderMatches 检查请求的提供者是否与配置的提供者匹配。
func ProviderMatches(requested, configured string) bool {
return strings.EqualFold(requested, configured)
}
+424
View File
@@ -0,0 +1,424 @@
// Package service — Aria2 下载适配器。
//
// Aria2Adapter 实现了 DownloadAdapter 接口,通过 Aria2 JSON-RPC API
// 管理下载任务。
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"sync"
"time"
)
// aria2Request 是 Aria2 JSON-RPC 请求结构。
type aria2Request struct {
JSONRPC string `json:"jsonrpc"`
Method string `json:"method"`
ID string `json:"id"`
Params []interface{} `json:"params"`
}
// aria2Response 是 Aria2 JSON-RPC 响应结构。
type aria2Response struct {
JSONRPC string `json:"jsonrpc"`
ID string `json:"id"`
Result json.RawMessage `json:"result"`
Error *aria2Error `json:"error"`
}
// aria2Error 是 Aria2 JSON-RPC 错误结构。
type aria2Error struct {
Code int `json:"code"`
Message string `json:"message"`
}
// Aria2Adapter 是 Aria2 的 DownloadAdapter 实现。
type Aria2Adapter struct {
mu sync.Mutex
cfg DownloadClientConfig
client *http.Client
idSeq int
}
// NewAria2Adapter 创建新的 Aria2 适配器。
func NewAria2Adapter() *Aria2Adapter {
return &Aria2Adapter{
client: &http.Client{Timeout: 20 * time.Second},
}
}
// Initialize 配置并初始化 Aria2 RPC 连接。
func (a *Aria2Adapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
a.mu.Lock()
defer a.mu.Unlock()
a.cfg = cfg
a.idSeq = 0
return a.getVersionLocked(ctx)
}
// Ping 测试连接。
func (a *Aria2Adapter) Ping(ctx context.Context) error {
a.mu.Lock()
defer a.mu.Unlock()
return a.getVersionLocked(ctx)
}
// getVersionLocked 内部版本检查(调用者必须持有锁)。
func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error {
rpcURL := a.cfg.Host
if !strings.HasSuffix(rpcURL, "/jsonrpc") {
rpcURL = strings.TrimRight(rpcURL, "/") + "/jsonrpc"
}
req := &aria2Request{
JSONRPC: "2.0",
Method: "aria2.getVersion",
ID: a.nextID(),
Params: []interface{}{"token:" + a.cfg.Password},
}
body, err := json.Marshal(req)
if err != nil {
return err
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
if err != nil {
return err
}
httpReq.Header.Set("Content-Type", "application/json")
if a.cfg.Username != "" {
httpReq.SetBasicAuth(a.cfg.Username, a.cfg.Password)
}
resp, err := a.client.Do(httpReq)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return fmt.Errorf("aria2 rpc: %d", resp.StatusCode)
}
return nil
}
// rpcLocked 发送 JSON-RPC 请求(调用者必须持有锁)。
func (a *Aria2Adapter) rpcLocked(ctx context.Context, method string, params []interface{}) (json.RawMessage, error) {
rpcURL := a.cfg.Host
if !strings.HasSuffix(rpcURL, "/jsonrpc") {
rpcURL = strings.TrimRight(rpcURL, "/") + "/jsonrpc"
}
if params == nil {
params = []interface{}{}
}
// 如果 secret 不在 params 中,添加到第一位
if len(params) > 0 {
if secret, ok := params[0].(string); ok && strings.HasPrefix(secret, "token:") {
// 已经有 secret
} else {
newParams := make([]interface{}, 0, len(params)+1)
newParams = append(newParams, "token:"+a.cfg.Password)
newParams = append(newParams, params...)
params = newParams
}
} else {
params = []interface{}{"token:" + a.cfg.Password}
}
req := &aria2Request{
JSONRPC: "2.0",
Method: method,
ID: a.nextID(),
Params: params,
}
body, err := json.Marshal(req)
if err != nil {
return nil, err
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
if err != nil {
return nil, err
}
httpReq.Header.Set("Content-Type", "application/json")
if a.cfg.Username != "" {
httpReq.SetBasicAuth(a.cfg.Username, a.cfg.Password)
}
resp, err := a.client.Do(httpReq)
if err != nil {
return nil, err
}
defer resp.Body.Close()
respBody, err := io.ReadAll(resp.Body)
if err != nil {
return nil, err
}
var rpcResp aria2Response
if err := json.Unmarshal(respBody, &rpcResp); err != nil {
return nil, err
}
if rpcResp.Error != nil {
return nil, fmt.Errorf("aria2 rpc error [%d]: %s", rpcResp.Error.Code, rpcResp.Error.Message)
}
return rpcResp.Result, nil
}
// AddTorrent 通过 URL 添加种子或磁力链接。
func (a *Aria2Adapter) AddTorrent(ctx context.Context, torrentURL, savePath string) (string, error) {
a.mu.Lock()
defer a.mu.Unlock()
// Aria2 addUri 的参数: [secret, [uris], options]
uris := []string{torrentURL}
options := map[string]string{}
if savePath != "" {
options["dir"] = savePath
}
result, err := a.rpcLocked(ctx, "aria2.addUri", []interface{}{uris, options})
if err != nil {
return "", err
}
var gid string
if err := json.Unmarshal(result, &gid); err != nil {
return "", err
}
return gid, nil
}
// AddMagnet 通过磁力链接添加下载。
func (a *Aria2Adapter) AddMagnet(ctx context.Context, magnet, savePath string) (string, error) {
return a.AddTorrent(ctx, magnet, savePath)
}
// Pause 暂停下载任务(通过 GID)。
func (a *Aria2Adapter) Pause(ctx context.Context, hash string) error {
a.mu.Lock()
defer a.mu.Unlock()
_, err := a.rpcLocked(ctx, "aria2.pause", []interface{}{hash})
return err
}
// Resume 恢复下载任务(通过 GID)。
func (a *Aria2Adapter) Resume(ctx context.Context, hash string) error {
a.mu.Lock()
defer a.mu.Unlock()
_, err := a.rpcLocked(ctx, "aria2.unpause", []interface{}{hash})
return err
}
// Remove 移除下载任务。
func (a *Aria2Adapter) Remove(ctx context.Context, hash string, deleteFiles bool) error {
a.mu.Lock()
defer a.mu.Unlock()
if deleteFiles {
_, err := a.rpcLocked(ctx, "aria2.removeDownloadResult", []interface{}{hash})
return err
}
_, err := a.rpcLocked(ctx, "aria2.remove", []interface{}{hash})
return err
}
// List 列出所有活动/等待/已停止的任务。
func (a *Aria2Adapter) List(ctx context.Context, filter string) ([]TorrentInfo, error) {
a.mu.Lock()
defer a.mu.Unlock()
var allResults []TorrentInfo
// 获取活动任务
active, err := a.rpcLocked(ctx, "aria2.tellActive", []interface{}{
[]string{"gid", "bittorrent", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"},
})
if err == nil && active != nil {
items := a.parseAria2Items(active)
allResults = append(allResults, items...)
}
// 获取等待中的任务
waiting, err := a.rpcLocked(ctx, "aria2.tellWaiting", []interface{}{
0, 100,
[]string{"gid", "bittorrent", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"},
})
if err == nil && waiting != nil {
items := a.parseAria2Items(waiting)
allResults = append(allResults, items...)
}
// 获取已停止的任务
stopped, err := a.rpcLocked(ctx, "aria2.tellStopped", []interface{}{
0, 100,
[]string{"gid", "bittorrent", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections", "errorCode"},
})
if err == nil && stopped != nil {
items := a.parseAria2Items(stopped)
allResults = append(allResults, items...)
}
if filter != "" {
filtered := make([]TorrentInfo, 0, len(allResults))
for _, item := range allResults {
if strings.EqualFold(item.State, filter) {
filtered = append(filtered, item)
}
}
return filtered, nil
}
return allResults, nil
}
// GetInfo 获取单个任务信息。
func (a *Aria2Adapter) GetInfo(ctx context.Context, hash string) (*TorrentInfo, error) {
a.mu.Lock()
defer a.mu.Unlock()
result, err := a.rpcLocked(ctx, "aria2.tellStatus", []interface{}{
hash,
[]string{"gid", "bittorrent", "totalLength", "completedLength", "downloadSpeed", "uploadSpeed", "status", "dir", "numSeeders", "connections"},
})
if err != nil {
return nil, err
}
var item map[string]interface{}
if err := json.Unmarshal(result, &item); err != nil {
return nil, err
}
info := a.parseSingleItem(item)
if info == nil {
return nil, fmt.Errorf("task %s not found", hash)
}
return info, nil
}
// parseAria2Items 解析 Aria2 返回的任务列表。
func (a *Aria2Adapter) parseAria2Items(raw json.RawMessage) []TorrentInfo {
var items []map[string]interface{}
if err := json.Unmarshal(raw, &items); err != nil {
return nil
}
result := make([]TorrentInfo, 0, len(items))
for _, item := range items {
info := a.parseSingleItem(item)
if info != nil {
result = append(result, *info)
}
}
return result
}
// parseSingleItem 解析单个 Aria2 任务项。
func (a *Aria2Adapter) parseSingleItem(item map[string]interface{}) *TorrentInfo {
gid := strVal(item["gid"])
totalLength := toInt64(item["totalLength"])
completedLength := toInt64(item["completedLength"])
dlSpeed := toInt64(item["downloadSpeed"])
upSpeed := toInt64(item["uploadSpeed"])
status := strVal(item["status"])
dir := strVal(item["dir"])
numSeeders := int(toInt64(item["numSeeders"]))
connections := int(toInt64(item["connections"]))
var name string
var hash string
// 尝试从 bittorrent info 获取名称和 hash
if bt, ok := item["bittorrent"].(map[string]interface{}); ok {
if info, ok := bt["info"].(map[string]interface{}); ok {
name = strVal(info["name"])
}
hash = strVal(bt["infoHash"])
}
// 如果没有 bittorrent 信息,使用 GID 作为 hash
if hash == "" {
hash = gid
}
if name == "" {
// 尝试从 files 获取文件名
if files, ok := item["files"].([]interface{}); ok && len(files) > 0 {
if f, ok := files[0].(map[string]interface{}); ok {
paths, ok := f["path"].([]interface{})
if ok && len(paths) > 0 {
name = strVal(paths[len(paths)-1])
}
if name == "" {
name = strVal(f["uris"])
}
}
}
}
if name == "" {
name = gid
}
var progress float64
if totalLength > 0 {
progress = float64(completedLength) / float64(totalLength) * 100
}
// Aria2 状态映射
state := aria2StatusStr(status)
return &TorrentInfo{
Hash: hash,
Name: name,
Size: totalLength,
Progress: progress,
DLSpeed: dlSpeed,
UPSpeed: upSpeed,
State: state,
SavePath: dir,
NumSeeds: numSeeders,
NumLeechs: max(connections-numSeeders, 0),
AddedOn: time.Now(),
}
}
// aria2StatusStr 将 Aria2 状态转为可读字符串。
func aria2StatusStr(status string) string {
switch status {
case "active":
return "downloading"
case "waiting":
return "queued"
case "paused":
return "paused"
case "error":
return "error"
case "complete":
return "seeding"
case "removed":
return "removed"
default:
return status
}
}
// nextID 生成递增的请求 ID。
func (a *Aria2Adapter) nextID() string {
a.idSeq++
return fmt.Sprintf("msg-%d", a.idSeq)
}
func max(a, b int) int {
if a > b {
return a
}
return b
}
+65 -26
View File
@@ -14,27 +14,29 @@ import (
"golang.org/x/crypto/bcrypt"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// AuthService handles registration, login, and JWT issuance.
type AuthService struct {
cfg *config.Config
log *zap.Logger
repo *repository.Container
cfg *config.Config
log *zap.Logger
repo *repository.Container
tokenSvc *TokenService
permissionSvc *PermissionService
}
// NewAuthService is the constructor.
func NewAuthService(cfg *config.Config, log *zap.Logger, repo *repository.Container) *AuthService {
return &AuthService{cfg: cfg, log: log, repo: repo}
func NewAuthService(cfg *config.Config, log *zap.Logger, repo *repository.Container, tokenSvc *TokenService, permissionSvc *PermissionService) *AuthService {
return &AuthService{cfg: cfg, log: log, repo: repo, tokenSvc: tokenSvc, permissionSvc: permissionSvc}
}
// Common service-level errors.
var (
ErrInvalidCredentials = errors.New("invalid username or password")
ErrUsernameTaken = errors.New("username already taken")
ErrUserInactive = errors.New("user account is inactive")
)
// SeedAdmin makes sure at least one admin user exists. It mirrors the
@@ -60,11 +62,14 @@ func (s *AuthService) SeedAdmin(ctx context.Context) error {
Username: "admin",
PasswordHash: hash,
Role: "admin",
Tier: "plus",
ForcePasswordReset: pwd == "admin123",
}
if err := s.repo.User.Create(ctx, user); err != nil {
return err
}
// 确保管理员有权限记录
_, _ = s.permissionSvc.EnsureForUser(ctx, user.ID)
s.log.Warn("default admin created — change the password after first login",
zap.String("username", "admin"),
zap.String("password_source", "ADMIN_INITIAL_PASSWORD or admin123"),
@@ -74,49 +79,72 @@ func (s *AuthService) SeedAdmin(ctx context.Context) error {
// Register creates a new user. The first registered user is auto-promoted to
// admin to support fresh installs that did not run SeedAdmin.
func (s *AuthService) Register(ctx context.Context, username, password string) (*model.User, error) {
func (s *AuthService) Register(ctx context.Context, username, password string) (*model.User, *TokenPair, error) {
username = strings.TrimSpace(username)
if username == "" || password == "" {
return nil, fmt.Errorf("username and password required")
return nil, nil, fmt.Errorf("username and password required")
}
if existing, err := s.repo.User.FindByUsername(ctx, username); err != nil {
return nil, err
return nil, nil, err
} else if existing != nil {
return nil, ErrUsernameTaken
return nil, nil, ErrUsernameTaken
}
hash, err := hashPassword(password)
if err != nil {
return nil, err
return nil, nil, err
}
role := "user"
if n, err := s.repo.User.CountAdmins(ctx); err == nil && n == 0 {
role = "admin"
}
u := &model.User{Username: username, PasswordHash: hash, Role: role}
if err := s.repo.User.Create(ctx, u); err != nil {
return nil, err
u := &model.User{
Username: username,
PasswordHash: hash,
Role: role,
Tier: "free",
}
return u, nil
if err := s.repo.User.Create(ctx, u); err != nil {
return nil, nil, err
}
// 自动为新用户创建默认权限
_, _ = s.permissionSvc.EnsureForUser(ctx, u.ID)
// 签发令牌对
tokens, err := s.tokenSvc.IssuePair(ctx, u.ID, u.Role, u.Tier)
if err != nil {
return u, nil, nil // 用户已创建,令牌签发失败不影响注册成功
}
return u, tokens, nil
}
// Login validates credentials and returns the user + a fresh JWT.
func (s *AuthService) Login(ctx context.Context, username, password string) (*model.User, string, error) {
// LoginResponse 登录响应结构。
type LoginResponse struct {
User *model.User `json:"user"`
Tokens *TokenPair `json:"tokens"`
}
// Login validates credentials and returns the user + a fresh JWT token pair.
func (s *AuthService) Login(ctx context.Context, username, password string) (*LoginResponse, error) {
u, err := s.repo.User.FindByUsername(ctx, username)
if err != nil {
return nil, "", err
return nil, err
}
if u == nil {
return nil, "", ErrInvalidCredentials
return nil, ErrInvalidCredentials
}
// 检查用户是否激活
if !u.IsActive {
return nil, ErrUserInactive
}
if err := bcrypt.CompareHashAndPassword([]byte(u.PasswordHash), []byte(password)); err != nil {
return nil, "", ErrInvalidCredentials
return nil, ErrInvalidCredentials
}
token, err := s.IssueToken(u)
// 签发令牌对
tokens, err := s.tokenSvc.IssuePair(ctx, u.ID, u.Role, u.Tier)
if err != nil {
return nil, "", err
return nil, err
}
_ = s.repo.User.TouchLogin(ctx, u.ID)
return u, token, nil
return &LoginResponse{User: u, Tokens: tokens}, nil
}
// ChangePassword updates the user password if the old one matches.
@@ -138,14 +166,15 @@ func (s *AuthService) ChangePassword(ctx context.Context, userID, oldPwd, newPwd
return s.repo.User.UpdatePassword(ctx, userID, hash)
}
// IssueToken signs a JWT for the given user (24h validity).
// IssueToken signs a JWT for the given user (60min validity, includes tier).
func (s *AuthService) IssueToken(u *model.User) (string, error) {
claims := middleware.Claims{
claims := Claims{
UserID: u.ID,
Role: u.Role,
Tier: u.Tier,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(time.Now()),
ExpiresAt: jwt.NewNumericDate(time.Now().Add(24 * time.Hour)),
ExpiresAt: jwt.NewNumericDate(time.Now().Add(60 * time.Minute)),
Issuer: "mediastationgo",
Subject: u.ID,
},
@@ -154,6 +183,16 @@ func (s *AuthService) IssueToken(u *model.User) (string, error) {
return t.SignedString([]byte(s.cfg.Secrets.JWTSecret))
}
// RefreshTokens 使用刷新令牌获取新的令牌对。
func (s *AuthService) RefreshTokens(ctx context.Context, refreshToken string) (*TokenPair, error) {
return s.tokenSvc.Refresh(ctx, refreshToken)
}
// Logout 撤销用户的所有刷新令牌。
func (s *AuthService) Logout(ctx context.Context, userID string) error {
return s.tokenSvc.RevokeAll(ctx, userID)
}
func hashPassword(p string) (string, error) {
h, err := bcrypt.GenerateFromPassword([]byte(p), bcrypt.DefaultCost)
if err != nil {
+5
View File
@@ -99,6 +99,11 @@ func (c *CryptoService) Decrypt(value string) string {
return string(plain)
}
// IsEncrypted returns true if value carries the encrypted prefix.
func (c *CryptoService) IsEncrypted(value string) bool {
return strings.HasPrefix(value, encPrefix)
}
// MaskAPIKey returns "abcd****wxyz" so the key can be displayed in the
// admin UI without leaking it. Inputs shorter than 8 chars become "****".
func MaskAPIKey(plain string) string {
+69
View File
@@ -0,0 +1,69 @@
// Package service 定义下载适配器接口和通用数据结构。
package service
import (
"context"
"time"
)
// DownloadAdapter 定义下载客户端的统一接口。
// 所有下载客户端(qBittorrent / Transmission / Aria2)必须实现此接口。
type DownloadAdapter interface {
// Initialize 使用配置初始化客户端连接。
Initialize(ctx context.Context, cfg DownloadClientConfig) error
// Ping 测试客户端连接是否可用。
Ping(ctx context.Context) error
// AddTorrent 通过 URL(磁力链接或种子 URL)添加下载任务。
AddTorrent(ctx context.Context, url, savePath string) (string, error)
// AddMagnet 通过磁力链接添加下载任务。
AddMagnet(ctx context.Context, magnet, savePath string) (string, error)
// Pause 暂停指定下载任务。
Pause(ctx context.Context, hash string) error
// Resume 恢复指定下载任务。
Resume(ctx context.Context, hash string) error
// Remove 移除指定下载任务,deleteFiles 控制是否同时删除文件。
Remove(ctx context.Context, hash string, deleteFiles bool) error
// List 列出所有或过滤后的种子任务。filter 可为空字符串表示全部。
List(ctx context.Context, filter string) ([]TorrentInfo, error)
// GetInfo 获取指定种子的详细信息。
GetInfo(ctx context.Context, hash string) (*TorrentInfo, error)
}
// TorrentInfo 是各种下载客户端的种子信息的统一表示。
type TorrentInfo struct {
Hash string `json:"hash"`
Name string `json:"name"`
Size int64 `json:"size"`
Progress float64 `json:"progress"`
DLSpeed int64 `json:"dl_speed"`
UPSpeed int64 `json:"up_speed"`
State string `json:"state"`
SavePath string `json:"save_path"`
NumSeeds int `json:"num_seeds"`
NumLeechs int `json:"num_leechs"`
AddedOn time.Time `json:"added_on"`
Category string `json:"category"`
Tags string `json:"tags"`
}
// DownloadClientConfig 是下载客户端的连接配置。
type DownloadClientConfig struct {
Host string `json:"host"`
Username string `json:"username"`
Password string `json:"password"`
Extra map[string]string `json:"extra,omitempty"`
}
// AdapterFactory 根据客户端类型创建适配器实例。
func AdapterFactory(clientType string) DownloadAdapter {
switch clientType {
case "qbittorrent":
return NewQBitAdapter()
case "transmission":
return NewTransmissionAdapter()
case "aria2":
return NewAria2Adapter()
default:
return nil
}
}
+253
View File
@@ -0,0 +1,253 @@
// Package service — 下载管理器,管理多个下载客户端适配器。
//
// DownloadManager 提供多客户端分发能力,支持运行时热插拔。
// 调用方通过 GetDefault() 或 GetClient(id) 获取适配器来执行下载操作。
package service
import (
"context"
"encoding/json"
"errors"
"sync"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// DownloadManager 管理多个下载客户端适配器实例。
type DownloadManager struct {
log *zap.Logger
repo *repository.Container
crypto *CryptoService
mu sync.RWMutex
clients map[string]DownloadAdapter // clientID -> adapter
configs map[string]DownloadClientConfig
}
// NewDownloadManager 创建新的下载管理器。
func NewDownloadManager(log *zap.Logger, repo *repository.Container, crypto *CryptoService) *DownloadManager {
return &DownloadManager{
log: log,
repo: repo,
crypto: crypto,
clients: make(map[string]DownloadAdapter),
configs: make(map[string]DownloadClientConfig),
}
}
// LoadAll 从数据库加载所有已启用的客户端并初始化适配器。
func (m *DownloadManager) LoadAll(ctx context.Context) error {
dbClients, err := m.repo.DownloadClient.ListEnabled(ctx)
if err != nil {
return err
}
m.mu.Lock()
defer m.mu.Unlock()
// 清空现有
m.clients = make(map[string]DownloadAdapter, len(dbClients))
m.configs = make(map[string]DownloadClientConfig, len(dbClients))
for _, dc := range dbClients {
cfg, err := m.buildConfig(&dc)
if err != nil {
m.log.Warn("failed to build config for download client",
zap.String("id", dc.ID),
zap.String("name", dc.Name),
zap.Error(err),
)
continue
}
adapter := AdapterFactory(dc.Type)
if adapter == nil {
m.log.Warn("unknown download client type",
zap.String("type", dc.Type),
zap.String("id", dc.ID),
)
continue
}
if initErr := adapter.Initialize(ctx, cfg); initErr != nil {
m.log.Warn("failed to initialize download client",
zap.String("id", dc.ID),
zap.String("name", dc.Name),
zap.Error(initErr),
)
continue
}
m.clients[dc.ID] = adapter
m.configs[dc.ID] = cfg
m.log.Info("download client initialized",
zap.String("id", dc.ID),
zap.String("name", dc.Name),
zap.String("type", dc.Type),
)
}
return nil
}
// GetDefault 返回默认下载客户端适配器。
// 如果没有设置默认客户端,返回第一个可用的客户端。
func (m *DownloadManager) GetDefault() (string, DownloadAdapter, error) {
m.mu.RLock()
defer m.mu.RUnlock()
// 首先找默认的
defaultClient, err := m.repo.DownloadClient.FindDefault(context.Background())
if err != nil {
return "", nil, err
}
if defaultClient != nil {
if adapter, ok := m.clients[defaultClient.ID]; ok {
return defaultClient.ID, adapter, nil
}
}
// 返回第一个可用的
for id, adapter := range m.clients {
return id, adapter, nil
}
return "", nil, errors.New("no download client available")
}
// GetClient 返回指定 ID 的下载客户端适配器。
func (m *DownloadManager) GetClient(id string) (DownloadAdapter, error) {
m.mu.RLock()
defer m.mu.RUnlock()
adapter, ok := m.clients[id]
if !ok {
return nil, errors.New("download client not found or not initialized")
}
return adapter, nil
}
// AddClient 动态添加并初始化一个下载客户端。
func (m *DownloadManager) AddClient(ctx context.Context, dc *model.DownloadClient) error {
cfg, err := m.buildConfig(dc)
if err != nil {
return err
}
adapter := AdapterFactory(dc.Type)
if adapter == nil {
return errors.New("unknown download client type: " + dc.Type)
}
if err := adapter.Initialize(ctx, cfg); err != nil {
return err
}
m.mu.Lock()
defer m.mu.Unlock()
m.clients[dc.ID] = adapter
m.configs[dc.ID] = cfg
return nil
}
// RemoveClient 移除一个下载客户端(停止适配器,不删除数据库记录)。
func (m *DownloadManager) RemoveClient(id string) {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.clients, id)
delete(m.configs, id)
}
// UpdateClient 更新已有客户端的配置并重新初始化。
func (m *DownloadManager) UpdateClient(ctx context.Context, dc *model.DownloadClient) error {
m.RemoveClient(dc.ID)
return m.AddClient(ctx, dc)
}
// TestConnection 测试客户端连接。
func (m *DownloadManager) TestConnection(ctx context.Context, dc *model.DownloadClient) error {
cfg, err := m.buildConfig(dc)
if err != nil {
return err
}
adapter := AdapterFactory(dc.Type)
if adapter == nil {
return errors.New("unknown download client type: " + dc.Type)
}
return adapter.Initialize(ctx, cfg)
}
// ListAll 获取所有已加载客户端的种子列表。
func (m *DownloadManager) ListAll(ctx context.Context, filter string) (map[string][]TorrentInfo, error) {
m.mu.RLock()
ids := make([]string, 0, len(m.clients))
for id := range m.clients {
ids = append(ids, id)
}
adapters := make([]DownloadAdapter, 0, len(m.clients))
for _, id := range ids {
adapters = append(adapters, m.clients[id])
}
m.mu.RUnlock()
result := make(map[string][]TorrentInfo)
for i, id := range ids {
list, err := adapters[i].List(ctx, filter)
if err != nil {
m.log.Warn("failed to list torrents from client",
zap.String("id", id),
zap.Error(err),
)
continue
}
result[id] = list
}
return result, nil
}
// GetAdapterTypes 返回支持的下载客户端类型列表。
func (m *DownloadManager) GetAdapterTypes() []AdapterTypeInfo {
return []AdapterTypeInfo{
{Type: "qbittorrent", Name: "qBittorrent", Description: "qBittorrent WebUI API (v2)"},
{Type: "transmission", Name: "Transmission", Description: "Transmission RPC API"},
{Type: "aria2", Name: "Aria2", Description: "Aria2 JSON-RPC API"},
}
}
// AdapterTypeInfo 描述下载客户端类型信息。
type AdapterTypeInfo struct {
Type string `json:"type"`
Name string `json:"name"`
Description string `json:"description"`
}
// buildConfig 从数据库模型构建适配器配置。
func (m *DownloadManager) buildConfig(dc *model.DownloadClient) (DownloadClientConfig, error) {
password := dc.Password
if m.crypto != nil && password != "" {
password = m.crypto.Decrypt(password)
}
cfg := DownloadClientConfig{
Host: dc.Host,
Username: dc.Username,
Password: password,
}
// 解析 Extra JSON 配置
if dc.Extra != "" {
extraStr := dc.Extra
if m.crypto != nil {
extraStr = m.crypto.Decrypt(extraStr)
}
var extra map[string]string
if err := json.Unmarshal([]byte(extraStr), &extra); err == nil {
cfg.Extra = extra
}
}
return cfg, nil
}
+76
View File
@@ -141,3 +141,79 @@ func (p *ImageProxy) Serve(ctx context.Context, w http.ResponseWriter, raw strin
http.ServeContent(w, &http.Request{}, key, stat.ModTime(), f)
return nil
}
// Fetch 拉取远程图片并返回字节和 Content-Type(带缓存)。
func (p *ImageProxy) Fetch(ctx context.Context, raw string) ([]byte, string, error) {
if raw == "" {
return nil, "", errors.New("missing url")
}
u, err := url.Parse(raw)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, "", errors.New("invalid url")
}
if _, ok := p.allowHost[strings.ToLower(u.Host)]; !ok {
return nil, "", errors.New("host not allowed")
}
// Cache lookup
sum := sha1.Sum([]byte(raw))
key := hex.EncodeToString(sum[:])
cachePath := filepath.Join(p.cacheDir, key)
if data, err := os.ReadFile(cachePath); err == nil {
// Content-Type from file extension or upstream headers — use a simple detect
ctype := detectContentType(data)
return data, ctype, nil
}
// Fetch upstream
if err := os.MkdirAll(p.cacheDir, 0o755); err != nil {
return nil, "", err
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, raw, nil)
if err != nil {
return nil, "", err
}
req.Header.Set("User-Agent", "MediaStationGo/0.1")
resp, err := p.client.Do(req)
if err != nil {
return nil, "", err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return nil, "", errors.New("upstream returned " + resp.Status)
}
data, err := io.ReadAll(resp.Body)
if err != nil {
return nil, "", err
}
// Write to cache
tmp, err := os.CreateTemp(p.cacheDir, "img-*.tmp")
if err == nil {
if _, err := tmp.Write(data); err == nil {
tmp.Close()
os.Rename(tmp.Name(), cachePath)
} else {
tmp.Close()
os.Remove(tmp.Name())
}
}
ctype := resp.Header.Get("Content-Type")
if ctype == "" {
ctype = detectContentType(data)
}
return data, ctype, nil
}
// detectContentType 通过前 512 字节检测 MIME 类型。
func detectContentType(data []byte) string {
if len(data) > 512 {
return http.DetectContentType(data[:512])
}
return http.DetectContentType(data)
}
+77
View File
@@ -0,0 +1,77 @@
// Package service — Bark 通知 Provider。
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
)
// BarkProvider 通过 Bark 推送通知到 iOS 设备。
// Bark API 文档: https://github.com/Finb/bark-server
type BarkProvider struct{}
// Send 发送 Bark 推送通知。
func (p *BarkProvider) Send(ctx context.Context, cfg map[string]string, event NotifyEvent) error {
serverURL := cfg["server_url"]
deviceKey := cfg["device_key"]
if serverURL == "" {
serverURL = "https://api.day.app"
}
serverURL = strings.TrimRight(serverURL, "/")
if deviceKey == "" {
return fmt.Errorf("bark: device_key is required")
}
payload := map[string]interface{}{
"title": event.Title,
"body": event.Message,
"group": "MediaStationGo",
}
if len(event.Data) > 0 {
var extra string
for k, v := range event.Data {
extra += fmt.Sprintf("%s: %v\n", k, v)
}
payload["body"] = event.Message + "\n\n" + extra
}
body, err := json.Marshal(payload)
if err != nil {
return err
}
apiURL := fmt.Sprintf("%s/%s", serverURL, deviceKey)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: 15 * time.Second}
resp, err := client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
respBody, _ := io.ReadAll(resp.Body)
if resp.StatusCode >= 400 {
return fmt.Errorf("bark api error %d: %s", resp.StatusCode, string(respBody))
}
return nil
}
// ValidateConfig 验证 Bark 配置。
func (p *BarkProvider) ValidateConfig(cfg map[string]string) error {
if cfg["device_key"] == "" {
return fmt.Errorf("bark: device_key is required")
}
return nil
}
+147
View File
@@ -0,0 +1,147 @@
// Package service — Email(SMTP) 通知 Provider。
package service
import (
"context"
"crypto/tls"
"fmt"
"net/mail"
"net/smtp"
"strconv"
"strings"
)
// EmailProvider 通过 SMTP 发送邮件通知。
type EmailProvider struct{}
// Send 通过 SMTP 发送邮件。
func (p *EmailProvider) Send(ctx context.Context, cfg map[string]string, event NotifyEvent) error {
smtpHost := cfg["smtp_host"]
smtpPortStr := cfg["smtp_port"]
username := cfg["username"]
password := cfg["password"]
from := cfg["from"]
to := cfg["to"]
tlsStr := cfg["tls"]
if smtpHost == "" || smtpPortStr == "" || username == "" || from == "" || to == "" {
return fmt.Errorf("email: smtp_host, smtp_port, username, from, and to are required")
}
smtpPort, err := strconv.Atoi(smtpPortStr)
if err != nil {
return fmt.Errorf("email: invalid smtp_port: %s", smtpPortStr)
}
useTLS := true
if tlsStr == "false" || tlsStr == "0" || tlsStr == "no" {
useTLS = false
}
// 构建邮件内容
subject := fmt.Sprintf("[MediaStationGo] %s", event.Title)
body := event.Message
if len(event.Data) > 0 {
body += "\n\n---\n详细信息:\n"
for k, v := range event.Data {
body += fmt.Sprintf(" %s: %v\n", k, v)
}
}
recipients := strings.Split(to, ",")
for i, r := range recipients {
recipients[i] = strings.TrimSpace(r)
}
// 构建邮件
fromAddr := mail.Address{Name: "MediaStationGo", Address: from}
toAddrs := make([]mail.Address, 0, len(recipients))
for _, r := range recipients {
toAddrs = append(toAddrs, mail.Address{Address: r})
}
msg := "From: " + fromAddr.String() + "\r\n"
msg += "To: "
for i, addr := range toAddrs {
if i > 0 {
msg += ", "
}
msg += addr.String()
}
msg += "\r\n"
msg += "Subject: " + subject + "\r\n"
msg += "MIME-Version: 1.0\r\n"
msg += "Content-Type: text/plain; charset=\"utf-8\"\r\n"
msg += "Content-Transfer-Encoding: base64\r\n"
msg += "\r\n"
msg += body
addr := fmt.Sprintf("%s:%d", smtpHost, smtpPort)
auth := smtp.PlainAuth("", username, password, smtpHost)
if useTLS {
// 使用 TLS 连接
tlsConfig := &tls.Config{
ServerName: smtpHost,
MinVersion: tls.VersionTLS12,
}
conn, err := tls.Dial("tcp", addr, tlsConfig)
if err != nil {
return fmt.Errorf("email tls dial: %w", err)
}
client, err := smtp.NewClient(conn, smtpHost)
if err != nil {
return fmt.Errorf("email smtp client: %w", err)
}
defer client.Close()
if err = client.Auth(auth); err != nil {
return fmt.Errorf("email auth: %w", err)
}
if err = client.Mail(from); err != nil {
return fmt.Errorf("email mail from: %w", err)
}
for _, r := range recipients {
if err = client.Rcpt(r); err != nil {
return fmt.Errorf("email rcpt to: %w", err)
}
}
w, err := client.Data()
if err != nil {
return fmt.Errorf("email data: %w", err)
}
if _, err = w.Write([]byte(msg)); err != nil {
return fmt.Errorf("email write: %w", err)
}
if err = w.Close(); err != nil {
return fmt.Errorf("email close: %w", err)
}
return client.Quit()
}
// 不使用 TLS(STARTTLS 或明文)
return smtp.SendMail(addr, auth, from, recipients, []byte(msg))
}
// ValidateConfig 验证 Email 配置。
func (p *EmailProvider) ValidateConfig(cfg map[string]string) error {
if cfg["smtp_host"] == "" {
return fmt.Errorf("email: smtp_host is required")
}
if cfg["smtp_port"] == "" {
return fmt.Errorf("email: smtp_port is required")
}
if cfg["username"] == "" {
return fmt.Errorf("email: username is required")
}
if cfg["from"] == "" {
return fmt.Errorf("email: from is required")
}
if cfg["to"] == "" {
return fmt.Errorf("email: to is required")
}
if _, err := strconv.Atoi(cfg["smtp_port"]); err != nil {
return fmt.Errorf("email: invalid smtp_port")
}
return nil
}
+185
View File
@@ -0,0 +1,185 @@
// Package service — 通知服务事件分发引擎。
//
// NotifyService 管理所有通知渠道,根据事件类型将通知分发给
// 订阅了该事件的渠道。支持 4 种内置事件类型和 5 种通知渠道。
package service
import (
"context"
"encoding/json"
"sync"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// 通知事件类型常量。
const (
EventSubscriptionHit = "subscription_hit"
EventDownloadComplete = "download_complete"
EventScrapeFailed = "scrape_failed"
EventSystemAlert = "system_alert"
)
// NotifyEvent 是通知事件的数据结构。
type NotifyEvent struct {
Type string `json:"type"`
Title string `json:"title"`
Message string `json:"message"`
Data map[string]interface{} `json:"data,omitempty"`
}
// NotifyProvider 定义通知渠道的发送接口。
type NotifyProvider interface {
Send(ctx context.Context, cfg map[string]string, event NotifyEvent) error
ValidateConfig(cfg map[string]string) error
}
// NotifyService 是事件驱动的通知分发引擎。
type NotifyService struct {
log *zap.Logger
repo *repository.Container
crypto *CryptoService
mu sync.RWMutex
providers map[string]NotifyProvider // type -> provider
}
// NewNotifyService 创建通知服务。
func NewNotifyService(log *zap.Logger, repo *repository.Container, crypto *CryptoService) *NotifyService {
ns := &NotifyService{
log: log,
repo: repo,
crypto: crypto,
providers: make(map[string]NotifyProvider),
}
// 注册内置 Provider
ns.registerProviders()
return ns
}
// registerProviders 注册所有内置通知 Provider。
func (s *NotifyService) registerProviders() {
s.mu.Lock()
defer s.mu.Unlock()
s.providers["telegram"] = &TelegramProvider{}
s.providers["wechat"] = &WechatProvider{}
s.providers["bark"] = &BarkProvider{}
s.providers["webhook"] = &WebhookProvider{}
s.providers["email"] = &EmailProvider{}
}
// Dispatch 将事件分发给所有订阅了该事件类型的已启用渠道。
func (s *NotifyService) Dispatch(ctx context.Context, event NotifyEvent) {
channels, err := s.repo.NotifyChannel.ListByEvent(ctx, event.Type)
if err != nil {
s.log.Error("failed to list channels for event",
zap.String("event", event.Type),
zap.Error(err),
)
return
}
for _, ch := range channels {
go func(channel model.NotifyChannel) {
if sendErr := s.sendToChannel(ctx, channel, event); sendErr != nil {
s.log.Error("failed to send notification",
zap.String("channel", channel.Name),
zap.String("type", channel.Type),
zap.String("event", event.Type),
zap.Error(sendErr),
)
}
}(ch)
}
}
// SendTest 向指定渠道发送测试通知。
func (s *NotifyService) SendTest(ctx context.Context, channelID string) error {
ch, err := s.repo.NotifyChannel.FindByID(ctx, channelID)
if err != nil {
return err
}
if ch == nil {
return ErrNotifyChannelNotFound
}
testEvent := NotifyEvent{
Type: "test",
Title: "MediaStationGo 测试通知",
Message: "这是一条测试通知,如果您看到此消息,说明通知渠道配置正确。",
}
return s.sendToChannel(ctx, *ch, testEvent)
}
// ValidateChannelConfig 验证渠道配置是否合法。
func (s *NotifyService) ValidateChannelConfig(channelType string, config map[string]string) error {
s.mu.RLock()
provider, ok := s.providers[channelType]
s.mu.RUnlock()
if !ok {
return ErrUnknownNotifyType
}
return provider.ValidateConfig(config)
}
// GetProviderTypes 返回支持的通知渠道类型列表。
func (s *NotifyService) GetProviderTypes() []NotifyProviderInfo {
return []NotifyProviderInfo{
{Type: "telegram", Name: "Telegram", Description: "通过 Telegram Bot 发送消息"},
{Type: "wechat", Name: "Server酱", Description: "通过 Server酱 推送到微信"},
{Type: "bark", Name: "Bark", Description: "通过 Bark 推送到 iOS"},
{Type: "webhook", Name: "Webhook", Description: "通过自定义 HTTP Webhook 发送"},
{Type: "email", Name: "Email", Description: "通过 SMTP 发送邮件"},
}
}
// NotifyProviderInfo 描述通知渠道类型信息。
type NotifyProviderInfo struct {
Type string `json:"type"`
Name string `json:"name"`
Description string `json:"description"`
}
// sendToChannel 解密渠道配置并通过对应的 Provider 发送通知。
func (s *NotifyService) sendToChannel(ctx context.Context, channel model.NotifyChannel, event NotifyEvent) error {
// 解密配置
configStr := channel.Config
if s.crypto != nil && configStr != "" {
configStr = s.crypto.Decrypt(configStr)
}
var cfg map[string]string
if err := json.Unmarshal([]byte(configStr), &cfg); err != nil {
return err
}
s.mu.RLock()
provider, ok := s.providers[channel.Type]
s.mu.RUnlock()
if !ok {
return ErrUnknownNotifyType
}
return provider.Send(ctx, cfg, event)
}
// 通知服务错误定义。
var (
ErrNotifyChannelNotFound = &NotifyError{Code: "CHANNEL_NOT_FOUND", Message: "notification channel not found"}
ErrUnknownNotifyType = &NotifyError{Code: "UNKNOWN_TYPE", Message: "unknown notification type"}
)
// NotifyError 是通知服务专用错误类型。
type NotifyError struct {
Code string `json:"code"`
Message string `json:"message"`
}
// Error 实现 error 接口。
func (e *NotifyError) Error() string {
return e.Message
}
+100
View File
@@ -0,0 +1,100 @@
// Package service — Telegram 通知 Provider。
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
)
// TelegramProvider 通过 Telegram Bot API 发送通知。
type TelegramProvider struct{}
// Send 发送 Telegram 消息。
func (p *TelegramProvider) Send(ctx context.Context, cfg map[string]string, event NotifyEvent) error {
botToken := cfg["bot_token"]
chatID := cfg["chat_id"]
parseMode := cfg["parse_mode"]
if parseMode == "" {
parseMode = "HTML"
}
if botToken == "" || chatID == "" {
return fmt.Errorf("telegram: bot_token and chat_id are required")
}
text := formatTelegramMessage(event, parseMode)
payload := map[string]string{
"chat_id": chatID,
"text": text,
"parse_mode": parseMode,
}
body, err := json.Marshal(payload)
if err != nil {
return err
}
apiURL := fmt.Sprintf("https://api.telegram.org/bot%s/sendMessage", botToken)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: 15 * time.Second}
resp, err := client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
respBody, _ := io.ReadAll(resp.Body)
if resp.StatusCode >= 400 {
return fmt.Errorf("telegram api error %d: %s", resp.StatusCode, string(respBody))
}
return nil
}
// ValidateConfig 验证 Telegram 配置。
func (p *TelegramProvider) ValidateConfig(cfg map[string]string) error {
if cfg["bot_token"] == "" {
return fmt.Errorf("telegram: bot_token is required")
}
if cfg["chat_id"] == "" {
return fmt.Errorf("telegram: chat_id is required")
}
return nil
}
// formatTelegramMessage 格式化消息内容。
func formatTelegramMessage(event NotifyEvent, parseMode string) string {
var sb strings.Builder
sb.WriteString(fmt.Sprintf("<b>%s</b>\n\n", escapeHTML(event.Title)))
sb.WriteString(escapeHTML(event.Message))
if len(event.Data) > 0 {
sb.WriteString("\n\n")
for k, v := range event.Data {
sb.WriteString(fmt.Sprintf("• <b>%s</b>: %v\n", escapeHTML(k), v))
}
}
if parseMode != "HTML" {
// Markdown 模式
result := sb.String()
result = strings.ReplaceAll(result, "<b>", "**")
result = strings.ReplaceAll(result, "</b>", "**")
result = strings.ReplaceAll(result, "&lt;", "<")
result = strings.ReplaceAll(result, "&gt;", ">")
result = strings.ReplaceAll(result, "&amp;", "&")
return result
}
return sb.String()
}
+119
View File
@@ -0,0 +1,119 @@
// Package service — Webhook 通知 Provider。
package service
import (
"context"
"fmt"
"io"
"net/http"
"strings"
"time"
)
// WebhookProvider 通过自定义 HTTP Webhook 发送通知。
// 支持自定义 HTTP 方法和请求头。
type WebhookProvider struct{}
// Send 发送 Webhook 通知。
func (p *WebhookProvider) Send(ctx context.Context, cfg map[string]string, event NotifyEvent) error {
webhookURL := cfg["url"]
if webhookURL == "" {
return fmt.Errorf("webhook: url is required")
}
method := cfg["method"]
if method == "" {
method = "POST"
}
method = strings.ToUpper(method)
// 构建请求体
bodyTemplate := cfg["body_template"]
var bodyStr string
if bodyTemplate != "" {
bodyStr = renderTemplate(bodyTemplate, event)
} else {
// 默认 JSON 格式
bodyStr = fmt.Sprintf(`{"type":"%s","title":"%s","message":"%s","data":{}}`,
event.Type, event.Title, event.Message)
if len(event.Data) > 0 {
var dataParts []string
for k, v := range event.Data {
dataParts = append(dataParts, fmt.Sprintf(`"%s":%v`, k, v))
}
bodyStr = fmt.Sprintf(`{"type":"%s","title":"%s","message":"%s","data":{%s}}`,
event.Type, event.Title, event.Message, strings.Join(dataParts, ","))
}
}
req, err := http.NewRequestWithContext(ctx, method, webhookURL, strings.NewReader(bodyStr))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
// 自定义请求头
headersJSON := cfg["headers_json"]
if headersJSON != "" {
headers := parseHeadersJSON(headersJSON)
for k, v := range headers {
req.Header.Set(k, v)
}
}
client := &http.Client{Timeout: 15 * time.Second}
resp, err := client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
respBody, _ := io.ReadAll(resp.Body)
return fmt.Errorf("webhook error %d: %s", resp.StatusCode, string(respBody))
}
return nil
}
// ValidateConfig 验证 Webhook 配置。
func (p *WebhookProvider) ValidateConfig(cfg map[string]string) error {
if cfg["url"] == "" {
return fmt.Errorf("webhook: url is required")
}
return nil
}
// renderTemplate 简单模板渲染,支持 {{title}}, {{message}}, {{type}} 占位符。
func renderTemplate(template string, event NotifyEvent) string {
result := template
result = strings.ReplaceAll(result, "{{title}}", event.Title)
result = strings.ReplaceAll(result, "{{message}}", event.Message)
result = strings.ReplaceAll(result, "{{type}}", event.Type)
return result
}
// parseHeadersJSON 简单解析 headers JSON(格式: {"key":"value",...})。
func parseHeadersJSON(jsonStr string) map[string]string {
result := make(map[string]string)
jsonStr = strings.TrimSpace(jsonStr)
if jsonStr == "" || (jsonStr[0] != '{' && jsonStr[len(jsonStr)-1] != '}') {
return result
}
// 简单 key:value 解析
inner := jsonStr[1 : len(jsonStr)-1]
parts := strings.Split(inner, ",")
for _, part := range parts {
kv := strings.SplitN(part, ":", 2)
if len(kv) != 2 {
continue
}
key := strings.Trim(strings.TrimSpace(kv[0]), `"`)
value := strings.Trim(strings.TrimSpace(kv[1]), `"`)
if key != "" {
result[key] = value
}
}
return result
}
+77
View File
@@ -0,0 +1,77 @@
// Package service — Server酱(WeChat) 通知 Provider。
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"time"
)
// WechatProvider 通过 Server酱 API 推送消息到微信。
// Server酱 API 文档: https://sct.ftqq.com/
type WechatProvider struct{}
// Send 发送 Server酱 推送消息。
func (p *WechatProvider) Send(ctx context.Context, cfg map[string]string, event NotifyEvent) error {
sendkey := cfg["sendkey"]
if sendkey == "" {
return fmt.Errorf("wechat: sendkey is required")
}
payload := map[string]string{
"title": event.Title,
"desp": event.Message,
}
if len(event.Data) > 0 {
payload["desp"] += "\n\n---\n\n"
for k, v := range event.Data {
payload["desp"] += fmt.Sprintf("- **%s**: %v\n", k, v)
}
}
body, err := json.Marshal(payload)
if err != nil {
return err
}
apiURL := fmt.Sprintf("https://sctapi.ftqq.com/%s.send", sendkey)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: 15 * time.Second}
resp, err := client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
respBody, _ := io.ReadAll(resp.Body)
if resp.StatusCode >= 400 {
return fmt.Errorf("wechat server酱 api error %d: %s", resp.StatusCode, string(respBody))
}
// 检查 Server酱 响应
var result map[string]interface{}
if err := json.Unmarshal(respBody, &result); err == nil {
if code, ok := result["code"].(float64); ok && code != 0 {
msg, _ := result["message"].(string)
return fmt.Errorf("wechat server酱 error: %s", msg)
}
}
return nil
}
// ValidateConfig 验证 Server酱 配置。
func (p *WechatProvider) ValidateConfig(cfg map[string]string) error {
if cfg["sendkey"] == "" {
return fmt.Errorf("wechat: sendkey is required")
}
return nil
}
+117
View File
@@ -0,0 +1,117 @@
// Package service — 权限管理服务。
package service
import (
"context"
"errors"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// PermissionService 负责用户细粒度权限管理。
type PermissionService struct {
cfg *config.Config
log *zap.Logger
repo *repository.Container
}
// NewPermissionService 创建权限服务实例。
func NewPermissionService(cfg *config.Config, log *zap.Logger, repo *repository.Container) *PermissionService {
return &PermissionService{cfg: cfg, log: log, repo: repo}
}
// 权限服务错误定义。
var (
ErrPermissionDenied = errors.New("permission denied")
ErrPermissionNotFound = errors.New("permission not found")
)
// GetByUserID 获取用户的权限记录,不存在则返回默认权限。
func (s *PermissionService) GetByUserID(ctx context.Context, userID string) (*model.UserPermission, error) {
perm, err := s.repo.Permission.FindByUserID(ctx, userID)
if err != nil {
return nil, err
}
if perm == nil {
// 返回默认权限但不持久化
return model.NewDefaultPermission(userID), nil
}
return perm, nil
}
// EnsureForUser 确保用户拥有权限记录,不存在则创建默认权限。
func (s *PermissionService) EnsureForUser(ctx context.Context, userID string) (*model.UserPermission, error) {
perm, err := s.repo.Permission.FindByUserID(ctx, userID)
if err != nil {
return nil, err
}
if perm != nil {
return perm, nil
}
// 创建默认权限
defaultPerm := model.NewDefaultPermission(userID)
if err := s.repo.Permission.Upsert(ctx, defaultPerm); err != nil {
return nil, err
}
return defaultPerm, nil
}
// Check 检查用户是否拥有特定权限。
// 权限检查优先级:admin → 全权限 > plus → 全权限 > user → 查表
func (s *PermissionService) Check(ctx context.Context, userID, role, tier, permissionKey string) bool {
// admin 拥有所有权限
if role == "admin" {
return true
}
// plus 用户拥有所有权限
if tier == "plus" {
return true
}
// free 用户查表
perm, err := s.GetByUserID(ctx, userID)
if err != nil || perm == nil {
return false
}
permMap := perm.PermissionMap()
hasPermission, ok := permMap[permissionKey]
if !ok {
return false
}
return hasPermission
}
// Update 更新用户的权限。
func (s *PermissionService) Update(ctx context.Context, userID string, updates map[string]bool) error {
// 确保权限记录存在
if _, err := s.EnsureForUser(ctx, userID); err != nil {
return err
}
return s.repo.Permission.Update(ctx, userID, updates)
}
// ResetToDefault 将用户权限重置为默认值。
func (s *PermissionService) ResetToDefault(ctx context.Context, userID string) error {
defaultPerm := model.NewDefaultPermission(userID)
return s.repo.Permission.Upsert(ctx, defaultPerm)
}
// GetPermissionMap 获取用户权限的 map 表示。
func (s *PermissionService) GetPermissionMap(ctx context.Context, userID string) (map[string]bool, error) {
perm, err := s.GetByUserID(ctx, userID)
if err != nil {
return nil, err
}
return perm.PermissionMap(), nil
}
// IsSuperUser 检查用户是否为超级用户(admin 或 plus)。
func (s *PermissionService) IsSuperUser(role, tier string) bool {
return role == "admin" || tier == "plus"
}
+344
View File
@@ -0,0 +1,344 @@
// Package service — qBittorrent 下载适配器。
//
// QBitAdapter 实现了 DownloadAdapter 接口,通过 qBittorrent WebUI API
// 管理下载任务。底层使用与 QBitClient 相同的 HTTP API 调用逻辑。
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"mime/multipart"
"net/http"
"net/http/cookiejar"
"net/url"
"strconv"
"strings"
"sync"
"time"
)
// QBitAdapter 是 qBittorrent 的 DownloadAdapter 实现。
type QBitAdapter struct {
mu sync.Mutex
cfg DownloadClientConfig
client *http.Client
LoggedIn bool
}
// NewQBitAdapter 创建新的 qBittorrent 适配器。
func NewQBitAdapter() *QBitAdapter {
jar, _ := cookiejar.New(nil)
return &QBitAdapter{
client: &http.Client{Jar: jar, Timeout: 20 * time.Second},
}
}
// Initialize 配置并初始化 qBittorrent 连接。
func (a *QBitAdapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
a.mu.Lock()
defer a.mu.Unlock()
a.cfg = cfg
a.LoggedIn = false
jar, _ := cookiejar.New(nil)
a.client.Jar = jar
return a.loginLocked(ctx)
}
// Ping 测试连接。
func (a *QBitAdapter) Ping(ctx context.Context) error {
a.mu.Lock()
defer a.mu.Unlock()
return a.loginLocked(ctx)
}
// AddTorrent 通过 URL 添加种子。
func (a *QBitAdapter) AddTorrent(ctx context.Context, torrentURL, savePath string) (string, error) {
a.mu.Lock()
defer a.mu.Unlock()
if err := a.ensureAuthLocked(ctx); err != nil {
return "", err
}
body := &bytes.Buffer{}
w := multipart.NewWriter(body)
_ = w.WriteField("urls", torrentURL)
if savePath != "" {
_ = w.WriteField("savepath", savePath)
}
_ = w.Close()
baseURL := strings.TrimRight(a.cfg.Host, "/")
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
baseURL+"/api/v2/torrents/add", body)
if err != nil {
return "", err
}
req.Header.Set("Content-Type", w.FormDataContentType())
req.Header.Set("Referer", baseURL)
resp, err := a.client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
raw, _ := io.ReadAll(resp.Body)
return "", fmt.Errorf("qbittorrent add torrent: %d: %s", resp.StatusCode, strings.TrimSpace(string(raw)))
}
return "", nil
}
// AddMagnet 通过磁力链接添加种子。
func (a *QBitAdapter) AddMagnet(ctx context.Context, magnet, savePath string) (string, error) {
return a.AddTorrent(ctx, magnet, savePath)
}
// Pause 暂停种子。
func (a *QBitAdapter) Pause(ctx context.Context, hash string) error {
a.mu.Lock()
defer a.mu.Unlock()
if err := a.ensureAuthLocked(ctx); err != nil {
return err
}
baseURL := strings.TrimRight(a.cfg.Host, "/")
form := url.Values{}
form.Set("hashes", hash)
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
baseURL+"/api/v2/torrents/pause", strings.NewReader(form.Encode()))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("Referer", baseURL)
resp, err := a.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return fmt.Errorf("qbittorrent pause: %d", resp.StatusCode)
}
return nil
}
// Resume 恢复种子。
func (a *QBitAdapter) Resume(ctx context.Context, hash string) error {
a.mu.Lock()
defer a.mu.Unlock()
if err := a.ensureAuthLocked(ctx); err != nil {
return err
}
baseURL := strings.TrimRight(a.cfg.Host, "/")
form := url.Values{}
form.Set("hashes", hash)
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
baseURL+"/api/v2/torrents/resume", strings.NewReader(form.Encode()))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("Referer", baseURL)
resp, err := a.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return fmt.Errorf("qbittorrent resume: %d", resp.StatusCode)
}
return nil
}
// Remove 删除种子。
func (a *QBitAdapter) Remove(ctx context.Context, hash string, deleteFiles bool) error {
a.mu.Lock()
defer a.mu.Unlock()
if err := a.ensureAuthLocked(ctx); err != nil {
return err
}
baseURL := strings.TrimRight(a.cfg.Host, "/")
form := url.Values{}
form.Set("hashes", hash)
if deleteFiles {
form.Set("deleteFiles", "true")
} else {
form.Set("deleteFiles", "false")
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
baseURL+"/api/v2/torrents/delete", strings.NewReader(form.Encode()))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("Referer", baseURL)
resp, err := a.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return fmt.Errorf("qbittorrent delete: %d", resp.StatusCode)
}
return nil
}
// List 列出种子。
func (a *QBitAdapter) List(ctx context.Context, filter string) ([]TorrentInfo, error) {
a.mu.Lock()
defer a.mu.Unlock()
if err := a.ensureAuthLocked(ctx); err != nil {
return nil, err
}
baseURL := strings.TrimRight(a.cfg.Host, "/")
u := baseURL + "/api/v2/torrents/info"
if filter != "" {
u += "?filter=" + url.QueryEscape(filter)
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
if err != nil {
return nil, err
}
req.Header.Set("Referer", baseURL)
resp, err := a.client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode >= 400 {
return nil, fmt.Errorf("qbittorrent list: %d", resp.StatusCode)
}
// qBittorrent 返回的字段名与 TorrentInfo 不同,需要转换
type qbTorrent struct {
Hash string `json:"hash"`
Name string `json:"name"`
State string `json:"state"`
Progress float32 `json:"progress"`
DLSpeed int64 `json:"dlspeed"`
UPSpeed int64 `json:"upspeed"`
NumSeeds int `json:"num_seeds"`
NumLeechs int `json:"num_leechs"`
Size int64 `json:"size"`
SavePath string `json:"save_path"`
AddedOn int64 `json:"added_on"`
Category string `json:"category"`
Tags string `json:"tags"`
}
var qbList []qbTorrent
if err := json.NewDecoder(resp.Body).Decode(&qbList); err != nil {
return nil, err
}
result := make([]TorrentInfo, 0, len(qbList))
for _, t := range qbList {
result = append(result, TorrentInfo{
Hash: t.Hash,
Name: t.Name,
Size: t.Size,
Progress: float64(t.Progress),
DLSpeed: t.DLSpeed,
UPSpeed: t.UPSpeed,
State: t.State,
SavePath: t.SavePath,
NumSeeds: t.NumSeeds,
NumLeechs: t.NumLeechs,
AddedOn: time.Unix(t.AddedOn, 0),
Category: t.Category,
Tags: t.Tags,
})
}
return result, nil
}
// GetInfo 获取单个种子信息。
func (a *QBitAdapter) GetInfo(ctx context.Context, hash string) (*TorrentInfo, error) {
list, err := a.List(ctx, "")
if err != nil {
return nil, err
}
for _, t := range list {
if t.Hash == hash {
return &t, nil
}
}
return nil, fmt.Errorf("torrent %s not found", hash)
}
// loginLocked 执行登录(调用者必须持有锁)。
func (a *QBitAdapter) loginLocked(ctx context.Context) error {
if a.cfg.Host == "" {
return fmt.Errorf("qbittorrent host not configured")
}
form := url.Values{}
form.Set("username", a.cfg.Username)
form.Set("password", a.cfg.Password)
baseURL := strings.TrimRight(a.cfg.Host, "/")
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
baseURL+"/api/v2/auth/login", strings.NewReader(form.Encode()))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("Referer", baseURL)
resp, err := a.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode >= 400 || strings.TrimSpace(string(body)) != "Ok." {
return fmt.Errorf("qbittorrent login failed: %s", strings.TrimSpace(string(body)))
}
a.LoggedIn = true
return nil
}
// ensureAuthLocked 确保已认证(调用者必须持有锁)。
func (a *QBitAdapter) ensureAuthLocked(ctx context.Context) error {
if a.LoggedIn {
return nil
}
return a.loginLocked(ctx)
}
// --- 为了与现有的 QBitClient 兼容,添加转换辅助函数 ---
// QBitTorrentToInfo 将旧的 QBitTorrent 转换为新的 TorrentInfo。
func QBitTorrentToInfo(q QBitTorrent) TorrentInfo {
return TorrentInfo{
Hash: q.Hash,
Name: q.Name,
Size: q.Size,
Progress: float64(q.Progress),
DLSpeed: q.DLSpeed,
UPSpeed: q.UpSpeed,
State: q.State,
SavePath: q.SavePath,
NumSeeds: q.NumSeeds,
NumLeechs: q.NumLeech,
}
}
// TorrentInfoToQBit 将 TorrentInfo 转换回旧的 QBitTorrent 格式(兼容性)。
func TorrentInfoToQBit(t TorrentInfo) QBitTorrent {
return QBitTorrent{
Hash: t.Hash,
Name: t.Name,
State: t.State,
Progress: float32(t.Progress),
DLSpeed: t.DLSpeed,
UpSpeed: t.UPSpeed,
NumSeeds: t.NumSeeds,
NumLeech: t.NumLeechs,
Size: t.Size,
SavePath: t.SavePath,
}
}
// unused import guard
var _ = strconv.Itoa
+52 -12
View File
@@ -1,11 +1,11 @@
// Package service contains the business logic of MediaStationGo. Handlers
// deserialize the HTTP request, call into a Service method, then serialize
// the response. Services own all cross-cutting policy (auth, scanning,
// transcoding, etc.) and never deal with HTTP types directly.
// Package service 包含 MediaStationGo 的业务逻辑。
// Handler 反序列化 HTTP 请求,调用 Service 方法,然后序列化响应。
// Services 拥有所有横切策略(认证、扫描、转码等)且不直接处理 HTTP 类型。
package service
import (
"context"
"time"
"go.uber.org/zap"
@@ -13,13 +13,13 @@ import (
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// Container holds every service initialized at startup. Handlers receive a
// pointer to it and pick the relevant fields.
// Container 持有在启动时初始化的每个服务。Handler 接收指向它的指针并选择相关字段。
type Container struct {
Cfg *config.Config
Log *zap.Logger
Repo *repository.Container
WSHub *Hub
SSEHub *SSEHub
Auth *AuthService
Media *MediaService
Scan *ScannerService
@@ -55,16 +55,26 @@ type Container struct {
Notifier *NotifierService
Organizer *OrganizerService
Douban *DoubanProvider
Permission *PermissionService
Token *TokenService
ApiConfig *ApiConfigService
DownloadMgr *DownloadManager
Notify *NotifyService
Site *SiteService
stopCtx context.Context
stopCancel context.CancelFunc
}
// New builds the service container.
// New 构建服务容器。
func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Container {
hub := NewHub(log)
go hub.Run()
// 初始化 SSE Hub
sseHub := NewSSEHub(log)
go sseHub.Run()
probe := NewFFprobeService(cfg, log)
tmdb := NewTMDbProvider(cfg, log)
bangumi := NewBangumiProvider(cfg, log)
@@ -92,6 +102,14 @@ func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Cont
douban := NewDoubanProvider(cfg, log)
scheduler := NewSchedulerService(log, repos, scanner, transcoder, hub, cfg.Cache.CacheDir)
// 初始化认证相关服务
tokenSvc := NewTokenService(cfg, log, repos)
permissionSvc := NewPermissionService(cfg, log, repos)
apiConfigSvc := NewApiConfigService(cfg, log, repos, crypto)
downloadMgr := NewDownloadManager(log, repos, crypto)
notifySvc := NewNotifyService(log, repos, crypto)
siteSvc := NewSiteService(log, repos, crypto)
ctx, cancel := context.WithCancel(context.Background())
return &Container{
@@ -99,7 +117,8 @@ func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Cont
Log: log,
Repo: repos,
WSHub: hub,
Auth: NewAuthService(cfg, log, repos),
SSEHub: sseHub,
Auth: NewAuthService(cfg, log, repos, tokenSvc, permissionSvc),
Media: NewMediaService(cfg, log, repos),
Scan: scanner,
Stream: NewStreamService(cfg, log, repos, transcoder),
@@ -134,13 +153,19 @@ func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Cont
Notifier: notifier,
Organizer: organizer,
Douban: douban,
Permission: permissionSvc,
Token: tokenSvc,
ApiConfig: apiConfigSvc,
DownloadMgr: downloadMgr,
Notify: notifySvc,
Site: siteSvc,
stopCtx: ctx,
stopCancel: cancel,
}
}
// Boot kicks off background workers (watcher, downloads poller,
// subscription scheduler). Called once after AutoMigrate.
// Boot 启动后台工作进程(watcher, downloads poller, subscription scheduler)。
// 在 AutoMigrate 后调用一次。
func (c *Container) Boot() {
if err := c.Watcher.Start(c.stopCtx); err != nil {
c.Log.Warn("watcher start failed", zap.Error(err))
@@ -150,11 +175,17 @@ func (c *Container) Boot() {
if err := c.APIConfig.SeedDefaults(c.stopCtx); err != nil {
c.Log.Warn("api config seed failed", zap.Error(err))
}
// 加载所有已配置的下载客户端
if err := c.DownloadMgr.LoadAll(c.stopCtx); err != nil {
c.Log.Warn("failed to load download clients", zap.Error(err))
}
// 启动调度器定时任务
c.Scheduler.Start(c.stopCtx)
}
// Close releases any resources held by services (websocket hub, ffmpeg
// transcodes, fsnotify, background pollers).
// Close 释放 services 持有的任何资源(websocket hub, ffmpeg 转码, fsnotify, 后台轮询器)。
func (c *Container) Close() {
if c.stopCancel != nil {
c.stopCancel()
@@ -177,4 +208,13 @@ func (c *Container) Close() {
if c.WSHub != nil {
c.WSHub.Stop()
}
if c.SSEHub != nil {
c.SSEHub.Stop()
}
if c.Scheduler != nil {
c.Scheduler.Stop()
}
}
// unused guard
var _ = time.Now
File diff suppressed because it is too large Load Diff
+214
View File
@@ -0,0 +1,214 @@
// Package service — 跨站聚合搜索服务.
package service
import (
"context"
"fmt"
"sort"
"strings"
"sync"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// SiteSearchService 跨站聚合搜索服务.
type SiteSearchService struct {
log *zap.Logger
repo *repository.Container
site *SiteService
}
// NewSiteSearchService 创建跨站搜索服务.
func NewSiteSearchService(log *zap.Logger, repo *repository.Container, siteSvc *SiteService) *SiteSearchService {
return &SiteSearchService{log: log, repo: repo, site: siteSvc}
}
// SearchAll 在所有启用的站点中搜索关键字.
func (s *SiteSearchService) SearchAll(ctx context.Context, keyword string, page, pageSize int) (*AggregatedResult, error) {
sites, err := s.repo.Site.ListEnabled(ctx)
if err != nil {
return nil, fmt.Errorf("list enabled sites: %w", err)
}
if len(sites) == 0 {
return &AggregatedResult{
Keyword: keyword,
Items: []TorrentItem{},
Total: 0,
Page: page,
PageSize: pageSize,
}, nil
}
return s.SearchSites(ctx, keyword, sites, page, pageSize)
}
// SearchSites 在指定站点中搜索关键字.
func (s *SiteSearchService) SearchSites(ctx context.Context, keyword string, sites []model.Site, page, pageSize int) (*AggregatedResult, error) {
var mu sync.Mutex
var wg sync.WaitGroup
var allItems []TorrentItem
for _, site := range sites {
wg.Add(1)
go func(siteModel model.Site) {
defer wg.Done()
cfg, err := s.site.GetSiteConfig(ctx, siteModel.ID)
if err != nil {
s.log.Warn("get site config failed", zap.String("site_id", siteModel.ID), zap.Error(err))
return
}
adapter := GetAdapterForType(siteModel.Type)
result, err := adapter.Search(ctx, *cfg, keyword, page)
if err != nil {
s.log.Warn("site search failed",
zap.String("site_id", siteModel.ID),
zap.String("site_name", siteModel.Name),
zap.Error(err),
)
return
}
mu.Lock()
allItems = append(allItems, result.Items...)
mu.Unlock()
}(site)
}
wg.Wait()
sort.Slice(allItems, func(i, j int) bool {
if allItems[i].Seeders != allItems[j].Seeders {
return allItems[i].Seeders > allItems[j].Seeders
}
return allItems[i].UploadTime.After(allItems[j].UploadTime)
})
allItems = deduplicateItems(allItems)
total := len(allItems)
start := (page - 1) * pageSize
end := start + pageSize
if start > total {
start = total
}
if end > total {
end = total
}
return &AggregatedResult{
Keyword: keyword,
Items: allItems[start:end],
Total: total,
Page: page,
PageSize: pageSize,
}, nil
}
// SearchSite 在单个站点中搜索.
func (s *SiteSearchService) SearchSite(ctx context.Context, siteID, keyword string, page int) (*SearchResult, error) {
cfg, err := s.site.GetSiteConfig(ctx, siteID)
if err != nil {
return nil, fmt.Errorf("get site config: %w", err)
}
siteModel, err := s.repo.Site.FindByID(ctx, siteID)
if err != nil || siteModel == nil {
return nil, fmt.Errorf("find site: %w", err)
}
adapter := GetAdapterForType(siteModel.Type)
result, err := adapter.Search(ctx, *cfg, keyword, page)
if err != nil {
return nil, fmt.Errorf("search site %s: %w", siteModel.Name, err)
}
return result, nil
}
// BrowseSite 浏览站点资源.
func (s *SiteSearchService) BrowseSite(ctx context.Context, siteID, category string, page int) (*SearchResult, error) {
cfg, err := s.site.GetSiteConfig(ctx, siteID)
if err != nil {
return nil, fmt.Errorf("get site config: %w", err)
}
siteModel, err := s.repo.Site.FindByID(ctx, siteID)
if err != nil || siteModel == nil {
return nil, fmt.Errorf("find site: %w", err)
}
adapter := GetAdapterForType(siteModel.Type)
result, err := adapter.Browse(ctx, *cfg, category, page)
if err != nil {
return nil, fmt.Errorf("browse site %s: %w", siteModel.Name, err)
}
return result, nil
}
// GetTorrentDetail 获取种子详情.
func (s *SiteSearchService) GetTorrentDetail(ctx context.Context, siteID, torrentID string) (*TorrentDetail, error) {
cfg, err := s.site.GetSiteConfig(ctx, siteID)
if err != nil {
return nil, fmt.Errorf("get site config: %w", err)
}
siteModel, err := s.repo.Site.FindByID(ctx, siteID)
if err != nil || siteModel == nil {
return nil, fmt.Errorf("find site: %w", err)
}
adapter := GetAdapterForType(siteModel.Type)
detail, err := adapter.GetDetail(ctx, *cfg, torrentID)
if err != nil {
return nil, fmt.Errorf("get detail from %s: %w", siteModel.Name, err)
}
return detail, nil
}
// AggregatedResult 聚合搜索结果.
type AggregatedResult struct {
Keyword string `json:"keyword"`
Items []TorrentItem `json:"items"`
Total int `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
}
// deduplicateItems 通过标题相似性去重.
func deduplicateItems(items []TorrentItem) []TorrentItem {
seen := make(map[string]bool)
result := make([]TorrentItem, 0, len(items))
for _, item := range items {
key := normalizeTitle(item.Title)
if key == "" {
continue
}
if !seen[key] {
seen[key] = true
result = append(result, item)
}
}
return result
}
// normalizeTitle 标题标准化.
func normalizeTitle(title string) string {
title = strings.ToLower(strings.TrimSpace(title))
title = strings.ReplaceAll(title, ".", " ")
title = strings.ReplaceAll(title, "_", " ")
title = strings.ReplaceAll(title, "-", " ")
for strings.Contains(title, " ") {
title = strings.ReplaceAll(title, " ", " ")
}
return strings.TrimSpace(title)
}
+246
View File
@@ -0,0 +1,246 @@
// Package service — PT 站点管理服务。
package service
import (
"context"
"encoding/json"
"errors"
"strings"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// 站点管理错误码。
var (
ErrSiteNotFound = errors.New("site not found")
ErrSiteAuthFailed = errors.New("site authentication failed")
ErrSiteTypeInvalid = errors.New("invalid site type")
ErrSiteAuthInvalid = errors.New("invalid auth type")
)
// SiteService 站点管理服务。
type SiteService struct {
log *zap.Logger
repo *repository.Container
crypto *CryptoService
}
// NewSiteService 创建站点管理服务。
func NewSiteService(log *zap.Logger, repo *repository.Container, crypto *CryptoService) *SiteService {
return &SiteService{log: log, repo: repo, crypto: crypto}
}
// Create 创建站点,加密敏感字段。
func (s *SiteService) Create(ctx context.Context, site *model.Site) (*model.Site, error) {
if !isValidSiteType(site.Type) {
return nil, ErrSiteTypeInvalid
}
if !isValidAuthType(site.AuthType) {
return nil, ErrSiteAuthInvalid
}
// 加密敏感字段
s.encryptSite(site)
if err := s.repo.Site.Create(ctx, site); err != nil {
s.log.Error("create site failed", zap.Error(err))
return nil, err
}
return site, nil
}
// GetByID 获取站点(敏感字段解密)。
func (s *SiteService) GetByID(ctx context.Context, id string) (*model.Site, error) {
site, err := s.repo.Site.FindByID(ctx, id)
if err != nil {
return nil, err
}
if site == nil {
return nil, ErrSiteNotFound
}
s.decryptSite(site)
return site, nil
}
// List 获取所有站点(不含敏感字段)。
func (s *SiteService) List(ctx context.Context) ([]model.Site, error) {
sites, err := s.repo.Site.List(ctx)
if err != nil {
return nil, err
}
return sites, nil
}
// Update 更新站点。
func (s *SiteService) Update(ctx context.Context, site *model.Site) (*model.Site, error) {
existing, err := s.repo.Site.FindByID(ctx, site.ID)
if err != nil {
return nil, err
}
if existing == nil {
return nil, ErrSiteNotFound
}
if !isValidSiteType(site.Type) {
return nil, ErrSiteTypeInvalid
}
if !isValidAuthType(site.AuthType) {
return nil, ErrSiteAuthInvalid
}
s.encryptSite(site)
if err := s.repo.Site.Update(ctx, site); err != nil {
s.log.Error("update site failed", zap.Error(err))
return nil, err
}
return site, nil
}
// Delete 删除站点。
func (s *SiteService) Delete(ctx context.Context, id string) error {
existing, err := s.repo.Site.FindByID(ctx, id)
if err != nil {
return err
}
if existing == nil {
return ErrSiteNotFound
}
return s.repo.Site.Delete(ctx, id)
}
// Authenticate 测试站点认证。
func (s *SiteService) Authenticate(ctx context.Context, id string) error {
site, err := s.repo.Site.FindByID(ctx, id)
if err != nil {
return err
}
if site == nil {
return ErrSiteNotFound
}
cfg, err := s.toSiteConfig(site)
if err != nil {
return err
}
adapter := GetAdapterForType(site.Type)
if err := adapter.Authenticate(ctx, *cfg); err != nil {
// 更新错误状态
now := time.Now()
site.LastError = err.Error()
site.LastCheckAt = &now
_ = s.repo.Site.Update(ctx, site)
return ErrSiteAuthFailed
}
// 清除错误状态
now := time.Now()
site.LastError = ""
site.LastCheckAt = &now
_ = s.repo.Site.Update(ctx, site)
return nil
}
// GetSiteConfig 获取解密后的站点配置(供内部使用)。
func (s *SiteService) GetSiteConfig(ctx context.Context, id string) (*SiteConfig, error) {
site, err := s.repo.Site.FindByID(ctx, id)
if err != nil {
return nil, err
}
if site == nil {
return nil, ErrSiteNotFound
}
return s.toSiteConfig(site)
}
// encryptSite 加密站点敏感字段。
func (s *SiteService) encryptSite(site *model.Site) {
if site.Cookie != "" {
site.Cookie = s.crypto.Encrypt(site.Cookie)
}
if site.APIKey != "" {
site.APIKey = s.crypto.Encrypt(site.APIKey)
}
if site.AuthHeader != "" {
site.AuthHeader = s.crypto.Encrypt(site.AuthHeader)
}
if site.Extra != "" {
site.Extra = s.crypto.Encrypt(site.Extra)
}
}
// decryptSite 解密站点敏感字段。
func (s *SiteService) decryptSite(site *model.Site) {
if site.Cookie != "" {
site.Cookie = s.crypto.Decrypt(site.Cookie)
}
if site.APIKey != "" {
site.APIKey = s.crypto.Decrypt(site.APIKey)
}
if site.AuthHeader != "" {
site.AuthHeader = s.crypto.Decrypt(site.AuthHeader)
}
if site.Extra != "" {
site.Extra = s.crypto.Decrypt(site.Extra)
}
}
// toSiteConfig 将 model.Site 转换为 SiteConfig(解密后)。
func (s *SiteService) toSiteConfig(site *model.Site) (*SiteConfig, error) {
cfg := &SiteConfig{
Name: site.Name,
Type: site.Type,
URL: strings.TrimRight(site.URL, "/"),
AuthType: site.AuthType,
Extra: map[string]string{},
}
// 解密
if site.Cookie != "" {
cfg.Cookie = s.crypto.Decrypt(site.Cookie)
}
if site.APIKey != "" {
cfg.APIKey = s.crypto.Decrypt(site.APIKey)
}
if site.AuthHeader != "" {
cfg.AuthHeader = s.crypto.Decrypt(site.AuthHeader)
}
if site.Extra != "" {
dec := s.crypto.Decrypt(site.Extra)
if dec != "" {
if err := json.Unmarshal([]byte(dec), &cfg.Extra); err != nil {
s.log.Warn("parse site extra config failed", zap.Error(err))
}
}
}
return cfg, nil
}
// isValidSiteType 检查站点类型是否有效。
func isValidSiteType(siteType string) bool {
for _, t := range model.SiteTypes() {
if t == siteType {
return true
}
}
return false
}
// isValidAuthType 检查认证方式是否有效。
func isValidAuthType(authType string) bool {
for _, t := range model.AuthTypes() {
if t == authType {
return true
}
}
return false
}
+228
View File
@@ -0,0 +1,228 @@
// Package service — SSE (Server-Sent Events) 事件流服务。
package service
import (
"crypto/rand"
"encoding/hex"
"encoding/json"
"sync"
"time"
"go.uber.org/zap"
)
// SSEHub 管理 SSE 客户端连接和事件广播。
type SSEHub struct {
clients map[chan SSEEvent]bool
broadcast chan SSEEvent
register chan chan SSEEvent
unregister chan chan SSEEvent
log *zap.Logger
tickets map[string]*sseTicket
ticketMu sync.RWMutex
stopCh chan struct{}
}
type SSEEvent struct {
Type string `json:"type"`
Payload interface{} `json:"payload"`
}
// SSEEvent 事件类型常量。
const (
EventTypeScan = "scan"
EventTypeDownload = "download"
EventTypeSubscribe = "subscribe"
EventTypeTask = "task"
EventTypeSystem = "system"
EventTypeAuth = "auth"
)
// sseTicket 是一次性 OTP 票据。
type sseTicket struct {
UserID string
ExpiresAt time.Time
}
// NewSSEHub 创建 SSE Hub 实例。
func NewSSEHub(log *zap.Logger) *SSEHub {
return &SSEHub{
clients: make(map[chan SSEEvent]bool),
broadcast: make(chan SSEEvent, 256),
register: make(chan chan SSEEvent),
unregister: make(chan chan SSEEvent),
tickets: make(map[string]*sseTicket),
log: log,
stopCh: make(chan struct{}),
}
}
// Run 启动 SSE Hub 的事件循环。
func (h *SSEHub) Run() {
for {
select {
case client := <-h.register:
h.clients[client] = true
h.log.Debug("SSE client connected", zap.Int("total", len(h.clients)))
case client := <-h.unregister:
if _, ok := h.clients[client]; ok {
delete(h.clients, client)
close(client)
h.log.Debug("SSE client disconnected", zap.Int("total", len(h.clients)))
}
case event := <-h.broadcast:
h.distribute(event)
case <-h.stopCh:
h.closeAll()
return
}
}
}
// Stop 停止 SSE Hub。
func (h *SSEHub) Stop() {
close(h.stopCh)
}
// ClientChannel SSE 客户端通道包装器。
type ClientChannel struct {
Ch chan SSEEvent
}
// Subscribe 注册一个新的 SSE 客户端,返回 SSE 客户端包装器。
func (h *SSEHub) Subscribe() *ClientChannel {
ch := make(chan SSEEvent, 100)
h.register <- ch
return &ClientChannel{Ch: ch}
}
// Unsubscribe 取消注册 SSE 客户端。
func (h *SSEHub) Unsubscribe(client *ClientChannel) {
if client != nil && client.Ch != nil {
h.unregister <- client.Ch
}
}
// Broadcast 向所有连接的客户端广播事件。
func (h *SSEHub) Broadcast(eventType string, payload interface{}) {
event := SSEEvent{
Type: eventType,
Payload: payload,
}
select {
case h.broadcast <- event:
default:
h.log.Warn("SSE broadcast queue full, dropping event", zap.String("type", eventType))
}
}
// SendToUser 向指定用户发送事件(通过 UserID 匹配)。
// 注意:此方法需要在客户端连接时关联 UserID。
func (h *SSEHub) SendToUser(userID string, eventType string, payload interface{}) {
// 目前通过广播实现,未来可扩展为按用户分组
h.Broadcast(eventType, payload)
}
// distribute 将事件分发给所有客户端。
func (h *SSEHub) distribute(event SSEEvent) {
data, err := json.Marshal(event)
if err != nil {
h.log.Error("failed to marshal SSE event", zap.Error(err))
return
}
for client := range h.clients {
select {
case client <- event:
default:
// 客户端通道已满,跳过
h.log.Warn("SSE client buffer full", zap.String("event", string(data)))
}
}
}
// closeAll 关闭所有客户端连接。
func (h *SSEHub) closeAll() {
for client := range h.clients {
close(client)
}
h.clients = make(map[chan SSEEvent]bool)
}
// GenerateTicket 生成一次性 SSE 连接票据(用于无 JWT 场景下的安全连接)。
func (h *SSEHub) GenerateTicket(userID string) (string, error) {
buf := make([]byte, 16)
if _, err := rand.Read(buf); err != nil {
return "", err
}
ticket := hex.EncodeToString(buf)
h.ticketMu.Lock()
defer h.ticketMu.Unlock()
h.tickets[ticket] = &sseTicket{
UserID: userID,
ExpiresAt: time.Now().Add(10 * time.Second),
}
return ticket, nil
}
// ValidateTicket 验证 SSE 连接票据,返回关联的用户 ID。
func (h *SSEHub) ValidateTicket(ticket string) (string, error) {
h.ticketMu.Lock()
defer h.ticketMu.Unlock()
t, ok := h.tickets[ticket]
if !ok {
return "", ErrInvalidTicket
}
if time.Now().After(t.ExpiresAt) {
delete(h.tickets, ticket)
return "", ErrTicketExpired
}
userID := t.UserID
delete(h.tickets, ticket)
return userID, nil
}
// CleanupTickets 清理过期的票据。
func (h *SSEHub) CleanupTickets() {
h.ticketMu.Lock()
defer h.ticketMu.Unlock()
now := time.Now()
for ticket, t := range h.tickets {
if now.After(t.ExpiresAt) {
delete(h.tickets, ticket)
}
}
}
// SSE Hub 错误定义。
var (
ErrInvalidTicket = &SSEError{Message: "invalid ticket"}
ErrTicketExpired = &SSEError{Message: "ticket expired"}
)
// SSEError SSE 相关错误。
type SSEError struct {
Message string
}
func (e *SSEError) Error() string {
return e.Message
}
// ClientCount 返回当前连接的客户端数量。
func (h *SSEHub) ClientCount() int {
h.ticketMu.RLock()
defer h.ticketMu.RUnlock()
return len(h.clients)
}
+235
View File
@@ -0,0 +1,235 @@
// Package service — STRM 文件管理服务。
package service
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// STRM 错误定义。
var (
ErrSTRMNotFound = errors.New("strm record not found")
ErrSTRMProtocolInvalid = errors.New("invalid strm protocol")
ErrSTRMURLInvalid = errors.New("invalid strm url")
)
// STRMService STRM 文件管理服务。
type STRMService struct {
log *zap.Logger
repo *repository.Container
cfg *config.Config
}
// NewSTRMService 创建 STRM 服务。
func NewSTRMService(log *zap.Logger, repo *repository.Container, cfg *config.Config) *STRMService {
return &STRMService{log: log, repo: repo, cfg: cfg}
}
// Create 创建 STRM 记录。
func (s *STRMService) Create(ctx context.Context, record *model.STRMRecord) (*model.STRMRecord, error) {
if err := s.validateSTRM(record); err != nil {
return nil, err
}
if err := s.repo.STRM.Create(ctx, record); err != nil {
s.log.Error("create strm failed", zap.Error(err))
return nil, err
}
return record, nil
}
// CreateBatch 批量创建 STRM 记录。
func (s *STRMService) CreateBatch(ctx context.Context, records []model.STRMRecord) (int, error) {
created := 0
for i := range records {
if err := s.validateSTRM(&records[i]); err != nil {
s.log.Warn("skip invalid strm record",
zap.String("title", records[i].Title),
zap.Error(err),
)
continue
}
created++
}
validRecords := make([]model.STRMRecord, 0, created)
for _, r := range records {
if model.IsAllowedProtocol(r.Protocol) && r.URL != "" {
validRecords = append(validRecords, r)
}
}
if len(validRecords) == 0 {
return 0, nil
}
if err := s.repo.STRM.CreateBatch(ctx, validRecords); err != nil {
s.log.Error("batch create strm failed", zap.Error(err))
return 0, err
}
return len(validRecords), nil
}
// GetByID 获取 STRM 记录。
func (s *STRMService) GetByID(ctx context.Context, id string) (*model.STRMRecord, error) {
record, err := s.repo.STRM.FindByID(ctx, id)
if err != nil {
return nil, err
}
if record == nil {
return nil, ErrSTRMNotFound
}
return record, nil
}
// List 列出 STRM 记录(支持筛选和分页)。
func (s *STRMService) List(ctx context.Context, filters map[string]string, page, pageSize int) ([]model.STRMRecord, int64, error) {
offset := (page - 1) * pageSize
if offset < 0 {
offset = 0
}
records, total, err := s.repo.STRM.List(ctx, filters, offset, pageSize)
if err != nil {
return nil, 0, err
}
return records, total, nil
}
// Update 更新 STRM 记录。
func (s *STRMService) Update(ctx context.Context, record *model.STRMRecord) (*model.STRMRecord, error) {
existing, err := s.repo.STRM.FindByID(ctx, record.ID)
if err != nil {
return nil, err
}
if existing == nil {
return nil, ErrSTRMNotFound
}
if record.Protocol != "" {
if !model.IsAllowedProtocol(record.Protocol) {
return nil, ErrSTRMProtocolInvalid
}
}
if err := s.repo.STRM.Update(ctx, record); err != nil {
s.log.Error("update strm failed", zap.Error(err))
return nil, err
}
return record, nil
}
// Delete 删除 STRM 记录。
func (s *STRMService) Delete(ctx context.Context, id string) error {
existing, err := s.repo.STRM.FindByID(ctx, id)
if err != nil {
return err
}
if existing == nil {
return ErrSTRMNotFound
}
return s.repo.STRM.Delete(ctx, id)
}
// GetProtocols 获取支持的协议列表。
func (s *STRMService) GetProtocols() []string {
return model.AllowedSTRMProtocols
}
// ProxySTRM 代理访问 STRM 资源。
// 支持 Range 请求(206 Partial Content)。
func (s *STRMService) ProxySTRM(ctx context.Context, id string, req *http.Request, w http.ResponseWriter) error {
record, err := s.repo.STRM.FindByID(ctx, id)
if err != nil {
return err
}
if record == nil {
return ErrSTRMNotFound
}
if !model.IsAllowedProtocol(record.Protocol) {
return ErrSTRMProtocolInvalid
}
// 创建代理请求
proxyReq, err := http.NewRequestWithContext(ctx, req.Method, record.URL, nil)
if err != nil {
return fmt.Errorf("create proxy request: %w", err)
}
// 复制 Range 等关键请求头
for _, header := range []string{
"Range", "If-Range", "If-Match", "If-None-Match",
"If-Modified-Since", "If-Unmodified-Since",
"Accept", "Accept-Encoding", "Accept-Language",
} {
if v := req.Header.Get(header); v != "" {
proxyReq.Header.Set(header, v)
}
}
// 对 alist/webdav 协议可能需要特殊处理认证
if record.Protocol == "alist" || record.Protocol == "alists" {
// alist 协议可以直接访问,无需额外认证
}
client := &http.Client{Timeout: 60 * time.Second}
resp, err := client.Do(proxyReq)
if err != nil {
return fmt.Errorf("proxy request failed: %w", err)
}
defer resp.Body.Close()
// 复制响应头
for _, header := range []string{
"Content-Type", "Content-Length", "Content-Range",
"Accept-Ranges", "Last-Modified", "ETag",
"Cache-Control", "Content-Disposition",
} {
if v := resp.Header.Get(header); v != "" {
w.Header().Set(header, v)
}
}
w.WriteHeader(resp.StatusCode)
_, err = io.Copy(w, resp.Body)
return err
}
// validateSTRM 验证 STRM 记录。
func (s *STRMService) validateSTRM(record *model.STRMRecord) error {
if record.Title == "" {
return errors.New("title is required")
}
if record.URL == "" {
return ErrSTRMURLInvalid
}
if !model.IsAllowedProtocol(record.Protocol) {
return ErrSTRMProtocolInvalid
}
// 标准化协议名
record.Protocol = strings.ToLower(record.Protocol)
return nil
}
// ListByMediaID 获取关联到指定媒体的 STRM 记录。
func (s *STRMService) ListByMediaID(ctx context.Context, mediaID string) ([]model.STRMRecord, error) {
return s.repo.STRM.FindByMediaID(ctx, mediaID)
}
+186
View File
@@ -0,0 +1,186 @@
// Package service — 双令牌认证服务。
package service
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"time"
"github.com/golang-jwt/jwt/v5"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/config"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
const (
// AccessTokenDuration Access Token 有效期(60分钟)
AccessTokenDuration = 60 * time.Minute
// RefreshTokenDuration Refresh Token 有效期(30天)
RefreshTokenDuration = 30 * 24 * time.Hour
// RefreshTokenLength Refresh Token 随机字节长度
RefreshTokenLength = 32
)
// Claims 是 JWT 载荷(复制自 middleware 以避免循环导入)。
type Claims struct {
UserID string `json:"uid"`
Role string `json:"role"`
Tier string `json:"tier,omitempty"`
jwt.RegisteredClaims
}
// TokenService 处理双令牌认证(Access Token + Refresh Token)。
type TokenService struct {
cfg *config.Config
log *zap.Logger
repo *repository.Container
}
// NewTokenService 创建令牌服务实例。
func NewTokenService(cfg *config.Config, log *zap.Logger, repo *repository.Container) *TokenService {
return &TokenService{cfg: cfg, log: log, repo: repo}
}
// TokenPair 包含访问令牌和刷新令牌。
type TokenPair struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
ExpiresIn int64 `json:"expires_in"` // 秒
TokenType string `json:"token_type"`
}
// TokenService 错误定义。
var (
ErrInvalidRefreshToken = errors.New("invalid refresh token")
ErrTokenExpired = errors.New("token expired")
ErrTokenRevoked = errors.New("token revoked")
)
// IssuePair 为用户签发新的令牌对。
func (s *TokenService) IssuePair(ctx context.Context, userID, role, tier string) (*TokenPair, error) {
// 生成 Access Token
accessToken, err := s.issueAccessToken(userID, role, tier)
if err != nil {
return nil, err
}
// 生成 Refresh Token
refreshToken, err := s.generateRefreshToken()
if err != nil {
return nil, err
}
// 存储 Refresh Token 哈希
tokenHash := repository.HashToken(refreshToken)
rt := &model.RefreshToken{
UserID: userID,
TokenHash: tokenHash,
ExpiresAt: time.Now().Add(RefreshTokenDuration),
}
if err := s.repo.RefreshToken.Create(ctx, rt); err != nil {
return nil, err
}
return &TokenPair{
AccessToken: accessToken,
RefreshToken: refreshToken,
ExpiresIn: int64(AccessTokenDuration.Seconds()),
TokenType: "Bearer",
}, nil
}
// issueAccessToken 签发 JWT Access Token(HS256,60分钟有效期)。
func (s *TokenService) issueAccessToken(userID, role, tier string) (string, error) {
claims := Claims{
UserID: userID,
Role: role,
Tier: tier,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(time.Now()),
ExpiresAt: jwt.NewNumericDate(time.Now().Add(AccessTokenDuration)),
Issuer: "mediastationgo",
Subject: userID,
},
}
t := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return t.SignedString([]byte(s.cfg.Secrets.JWTSecret))
}
// generateRefreshToken 生成安全的随机 Refresh Token。
func (s *TokenService) generateRefreshToken() (string, error) {
buf := make([]byte, RefreshTokenLength)
if _, err := rand.Read(buf); err != nil {
return "", err
}
return hex.EncodeToString(buf), nil
}
// Refresh 使用 Refresh Token 轮换获取新的令牌对。
func (s *TokenService) Refresh(ctx context.Context, refreshToken string) (*TokenPair, error) {
tokenHash := repository.HashToken(refreshToken)
// 查找 Refresh Token 记录
rt, err := s.repo.RefreshToken.FindByHash(ctx, tokenHash)
if err != nil {
return nil, err
}
if rt == nil {
return nil, ErrInvalidRefreshToken
}
// 检查是否已撤销
if rt.Revoked {
return nil, ErrTokenRevoked
}
// 检查是否过期
if rt.IsExpired() {
return nil, ErrTokenExpired
}
// 获取用户信息
user, err := s.repo.User.FindByID(ctx, rt.UserID)
if err != nil {
return nil, err
}
if user == nil {
return nil, ErrInvalidRefreshToken
}
// 撤销旧的 Refresh Token
if err := s.repo.RefreshToken.Revoke(ctx, tokenHash); err != nil {
s.log.Warn("failed to revoke old refresh token", zap.Error(err))
}
// 签发新的令牌对
return s.IssuePair(ctx, user.ID, user.Role, user.Tier)
}
// RevokeAll 撤销用户的所有 Refresh Token(用于登出)。
func (s *TokenService) RevokeAll(ctx context.Context, userID string) error {
return s.repo.RefreshToken.RevokeByUserID(ctx, userID)
}
// ValidateAccessToken 验证 Access Token 并返回 Claims。
func (s *TokenService) ValidateAccessToken(tokenString string) (*Claims, error) {
claims := &Claims{}
_, err := jwt.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, errors.New("unexpected signing method")
}
return []byte(s.cfg.Secrets.JWTSecret), nil
})
if err != nil {
return nil, err
}
return claims, nil
}
// CleanupExpired 清理过期的 Refresh Token。
func (s *TokenService) CleanupExpired(ctx context.Context) error {
return s.repo.RefreshToken.DeleteExpired(ctx)
}
+417
View File
@@ -0,0 +1,417 @@
// Package service — Transmission 下载适配器。
//
// TransmissionAdapter 实现了 DownloadAdapter 接口,通过 Transmission RPC API
// 管理下载任务。
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"sync"
"time"
)
// transmissionRPCRequest 是 Transmission RPC 请求的通用结构。
type transmissionRPCRequest struct {
Method string `json:"method"`
Arguments map[string]interface{} `json:"arguments"`
Tag int `json:"tag,omitempty"`
}
// transmissionRPCResponse 是 Transmission RPC 响应的通用结构。
type transmissionRPCResponse struct {
Result string `json:"result"`
Arguments map[string]interface{} `json:"arguments"`
Tag int `json:"tag"`
}
// TransmissionAdapter 是 Transmission 的 DownloadAdapter 实现。
type TransmissionAdapter struct {
mu sync.Mutex
cfg DownloadClientConfig
client *http.Client
tag int
sessionID string
}
// NewTransmissionAdapter 创建新的 Transmission 适配器。
func NewTransmissionAdapter() *TransmissionAdapter {
return &TransmissionAdapter{
client: &http.Client{Timeout: 20 * time.Second},
}
}
// Initialize 配置并初始化 Transmission RPC 连接。
func (a *TransmissionAdapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
a.mu.Lock()
defer a.mu.Unlock()
a.cfg = cfg
a.sessionID = ""
a.tag = 0
return a.pingLocked(ctx)
}
// Ping 测试连接。
func (a *TransmissionAdapter) Ping(ctx context.Context) error {
a.mu.Lock()
defer a.mu.Unlock()
return a.pingLocked(ctx)
}
// pingLocked 内部 ping 实现(调用者必须持有锁)。
func (a *TransmissionAdapter) pingLocked(ctx context.Context) error {
rpcURL := a.cfg.Host
if !strings.HasSuffix(rpcURL, "/rpc") && !strings.HasSuffix(rpcURL, "/transmission/rpc") {
if !strings.Contains(rpcURL, "/rpc") {
rpcURL = strings.TrimRight(rpcURL, "/") + "/transmission/rpc"
}
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, rpcURL, nil)
if err != nil {
return err
}
if a.cfg.Username != "" {
req.SetBasicAuth(a.cfg.Username, a.cfg.Password)
}
resp, err := a.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
io.Copy(io.Discard, resp.Body)
if resp.StatusCode == 409 {
// 正常:需要 CSRF token
a.sessionID = resp.Header.Get("X-Transmission-Session-Id")
return nil
}
if resp.StatusCode >= 400 {
return fmt.Errorf("transmission rpc: %d", resp.StatusCode)
}
return nil
}
// rpcLocked 发送 RPC 请求(调用者必须持有锁)。
func (a *TransmissionAdapter) rpcLocked(ctx context.Context, method string, args map[string]interface{}) (*transmissionRPCResponse, error) {
rpcURL := a.cfg.Host
if !strings.Contains(rpcURL, "/rpc") {
rpcURL = strings.TrimRight(rpcURL, "/") + "/transmission/rpc"
}
a.tag++
body, err := json.Marshal(transmissionRPCRequest{
Method: method,
Arguments: args,
Tag: a.tag,
})
if err != nil {
return nil, err
}
for attempt := 0; attempt < 2; attempt++ {
req, err := http.NewRequestWithContext(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
if a.sessionID != "" {
req.Header.Set("X-Transmission-Session-Id", a.sessionID)
}
if a.cfg.Username != "" {
req.SetBasicAuth(a.cfg.Username, a.cfg.Password)
}
resp, err := a.client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode == 409 {
a.sessionID = resp.Header.Get("X-Transmission-Session-Id")
continue
}
if resp.StatusCode >= 400 {
raw, _ := io.ReadAll(resp.Body)
return nil, fmt.Errorf("transmission rpc error: %d: %s", resp.StatusCode, string(raw))
}
var result transmissionRPCResponse
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
return nil, err
}
if result.Result != "success" {
return nil, fmt.Errorf("transmission rpc result: %s", result.Result)
}
return &result, nil
}
return nil, fmt.Errorf("transmission: failed after CSRF retry")
}
// AddTorrent 通过 URL 添加种子。
func (a *TransmissionAdapter) AddTorrent(ctx context.Context, torrentURL, savePath string) (string, error) {
a.mu.Lock()
defer a.mu.Unlock()
args := map[string]interface{}{"filename": torrentURL}
if savePath != "" {
args["download-dir"] = savePath
}
resp, err := a.rpcLocked(ctx, "torrent-add", args)
if err != nil {
return "", err
}
if added, ok := resp.Arguments["torrent-added"].(map[string]interface{}); ok {
if hashStr, ok := added["hashString"].(string); ok {
return hashStr, nil
}
}
if dup, ok := resp.Arguments["torrent-duplicate"].(map[string]interface{}); ok {
if hashStr, ok := dup["hashString"].(string); ok {
return hashStr, nil
}
}
return "", nil
}
// AddMagnet 通过磁力链接添加种子。
func (a *TransmissionAdapter) AddMagnet(ctx context.Context, magnet, savePath string) (string, error) {
return a.AddTorrent(ctx, magnet, savePath)
}
// Pause 暂停种子。
func (a *TransmissionAdapter) Pause(ctx context.Context, hash string) error {
a.mu.Lock()
defer a.mu.Unlock()
_, err := a.rpcLocked(ctx, "torrent-stop", map[string]interface{}{
"ids": []string{hash},
})
return err
}
// Resume 恢复种子。
func (a *TransmissionAdapter) Resume(ctx context.Context, hash string) error {
a.mu.Lock()
defer a.mu.Unlock()
_, err := a.rpcLocked(ctx, "torrent-start", map[string]interface{}{
"ids": []string{hash},
})
return err
}
// Remove 删除种子。
func (a *TransmissionAdapter) Remove(ctx context.Context, hash string, deleteFiles bool) error {
a.mu.Lock()
defer a.mu.Unlock()
_, err := a.rpcLocked(ctx, "torrent-remove", map[string]interface{}{
"ids": []string{hash},
"delete-local-data": deleteFiles,
})
return err
}
// List 列出种子。
func (a *TransmissionAdapter) List(ctx context.Context, filter string) ([]TorrentInfo, error) {
a.mu.Lock()
defer a.mu.Unlock()
args := map[string]interface{}{
"fields": []string{
"hashString", "name", "totalSize", "percentDone",
"rateDownload", "rateUpload", "status", "downloadDir",
"peersSendingToUs", "peersGettingFromUs", "addedDate",
"labels", "isStalled",
},
}
resp, err := a.rpcLocked(ctx, "torrent-get", args)
if err != nil {
return nil, err
}
torrentsRaw, ok := resp.Arguments["torrents"].([]interface{})
if !ok {
return nil, nil
}
result := make([]TorrentInfo, 0, len(torrentsRaw))
for _, tr := range torrentsRaw {
t, ok := tr.(map[string]interface{})
if !ok {
continue
}
hash, _ := t["hashString"].(string)
name, _ := t["name"].(string)
size := toInt64(t["totalSize"])
progress := toFloat64(t["percentDone"])
dlSpeed := toInt64(t["rateDownload"])
upSpeed := toInt64(t["rateUpload"])
savePath, _ := t["downloadDir"].(string)
numSeeds := int(toInt64(t["peersSendingToUs"]))
numLeechs := int(toInt64(t["peersGettingFromUs"]))
addedOn := int64(toFloat64(t["addedDate"]))
// Transmission 状态码转字符串
status := int(toFloat64(t["status"]))
state := transmissionStateStr(status)
// 过滤
if filter != "" && !strings.EqualFold(state, filter) {
continue
}
result = append(result, TorrentInfo{
Hash: hash,
Name: name,
Size: size,
Progress: progress * 100,
DLSpeed: dlSpeed,
UPSpeed: upSpeed,
State: state,
SavePath: savePath,
NumSeeds: numSeeds,
NumLeechs: numLeechs,
AddedOn: time.Unix(addedOn, 0),
Tags: toJSONLabels(t["labels"]),
})
}
return result, nil
}
// GetInfo 获取单个种子信息。
func (a *TransmissionAdapter) GetInfo(ctx context.Context, hash string) (*TorrentInfo, error) {
a.mu.Lock()
defer a.mu.Unlock()
args := map[string]interface{}{
"ids": []string{hash},
"fields": []string{
"hashString", "name", "totalSize", "percentDone",
"rateDownload", "rateUpload", "status", "downloadDir",
"peersSendingToUs", "peersGettingFromUs", "addedDate", "labels",
},
}
resp, err := a.rpcLocked(ctx, "torrent-get", args)
if err != nil {
return nil, err
}
torrentsRaw, ok := resp.Arguments["torrents"].([]interface{})
if !ok || len(torrentsRaw) == 0 {
return nil, fmt.Errorf("torrent %s not found", hash)
}
t, ok := torrentsRaw[0].(map[string]interface{})
if !ok {
return nil, fmt.Errorf("torrent %s: invalid response", hash)
}
status := int(toFloat64(t["status"]))
info := &TorrentInfo{
Hash: hash,
Name: strVal(t["name"]),
Size: toInt64(t["totalSize"]),
Progress: toFloat64(t["percentDone"]) * 100,
DLSpeed: toInt64(t["rateDownload"]),
UPSpeed: toInt64(t["rateUpload"]),
State: transmissionStateStr(status),
SavePath: strVal(t["downloadDir"]),
NumSeeds: int(toInt64(t["peersSendingToUs"])),
NumLeechs: int(toInt64(t["peersGettingFromUs"])),
AddedOn: time.Unix(int64(toFloat64(t["addedDate"])), 0),
Tags: toJSONLabels(t["labels"]),
}
return info, nil
}
// transmissionStateStr 将 Transmission 状态码转为可读字符串。
func transmissionStateStr(status int) string {
switch status {
case 0:
return "stopped"
case 1:
return "check_pending"
case 2:
return "checking"
case 3:
return "download_pending"
case 4:
return "downloading"
case 5:
return "seed_pending"
case 6:
return "seeding"
default:
return "unknown"
}
}
// toInt64 安全地将 interface{} 转为 int64。
func toInt64(v interface{}) int64 {
switch val := v.(type) {
case float64:
return int64(val)
case int:
return int64(val)
case int64:
return val
case json.Number:
n, _ := val.Int64()
return n
case string:
n, _ := strconv.ParseInt(val, 10, 64)
return n
default:
return 0
}
}
// toFloat64 安全地将 interface{} 转为 float64。
func toFloat64(v interface{}) float64 {
switch val := v.(type) {
case float64:
return val
case int:
return float64(val)
case int64:
return float64(val)
case json.Number:
n, _ := val.Float64()
return n
case string:
n, _ := strconv.ParseFloat(val, 64)
return n
default:
return 0
}
}
// strVal 安全地提取字符串。
func strVal(v interface{}) string {
if v == nil {
return ""
}
s, ok := v.(string)
if ok {
return s
}
return fmt.Sprintf("%v", v)
}
// toJSONLabels 将 Transmission labels 转为逗号分隔字符串。
func toJSONLabels(v interface{}) string {
if v == nil {
return ""
}
arr, ok := v.([]interface{})
if !ok {
return ""
}
labels := make([]string, 0, len(arr))
for _, item := range arr {
if s, ok := item.(string); ok {
labels = append(labels, s)
}
}
return strings.Join(labels, ",")
}