chore: remove license code, fix compilation, update README bilingual

- Remove all license-related code (handler/service/repository/model)
  License authorization is managed by separate server:
  https://github.com/ShukeBta/MediaStationLicenseServer
- Fix compilation errors: model field alignment, method name fixes,
  route conflicts, struct literal corrections (7 files)
- Add Chinese README.md as primary, English README_EN.md
- Update .gitignore: exclude .workbuddy/, editor backups
- Add new repository files: assistant, play_profile, storage_config
This commit is contained in:
ShukeBta
2026-05-17 00:02:41 +08:00
parent 37dd83c56d
commit 8e06b7e149
30 changed files with 756 additions and 1625 deletions
-94
View File
@@ -265,12 +265,6 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
authed.GET("/download/tasks", downloadTasksAliasHandler(svc))
authed.POST("/download/add", addDownloadHandler(svc))
// ── License (anyone authenticated can activate / heartbeat) ──
authed.POST("/license/activate", licenseActivateHandler(svc))
authed.POST("/license/heartbeat", licenseHeartbeatHandler(svc))
authed.GET("/license/status", licenseStatusHandler(svc))
authed.GET("/license/heartbeat-status", licenseStatusHandler(svc))
// ── Assistant (multi-turn AI chat) ──
authed.GET("/admin/assistant/sessions", listAssistantSessionsHandler(svc))
authed.POST("/admin/assistant/sessions", createAssistantSessionHandler(svc))
@@ -312,13 +306,6 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
admin.POST("/download/clients/:id/test", testDownloadClientHandler(svc))
admin.GET("/download/aria2/stats", aria2StatsHandler(svc))
// License generation / revocation.
admin.POST("/license/generate", licenseGenerateHandler(svc))
admin.GET("/license/list", licenseListHandler(svc))
admin.GET("/license/:id/activations", licenseListActivationsHandler(svc))
admin.POST("/license/activation/:id/unbind", licenseUnbindHandler(svc))
admin.POST("/license/:id/revoke", licenseRevokeHandler(svc))
// System scheduler trigger alias.
admin.POST("/system/scheduler/:name/trigger", schedulerTriggerHandler(svc))
@@ -352,10 +339,6 @@ func Register(r *gin.Engine, cfg *config.Config, log *zap.Logger, svc *service.C
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).
@@ -396,38 +379,11 @@ func versionInfo(c *gin.Context) {
// ─── 权限 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 {
@@ -467,73 +423,23 @@ func testApiConfigHandler(svc *service.Container) gin.HandlerFunc {
// ─── 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 {
-145
View File
@@ -1,145 +0,0 @@
// Package handler — license key endpoints.
package handler
import (
"net/http"
"time"
"github.com/gin-gonic/gin"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
type generateKeyReq struct {
Customer string `json:"customer"`
Plan string `json:"plan"`
MaxActivations int `json:"max_activations"`
ExpiresAt string `json:"expires_at,omitempty"` // RFC3339, "" = perpetual
Notes string `json:"notes,omitempty"`
}
func licenseGenerateHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req generateKeyReq
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
var expires *time.Time
if req.ExpiresAt != "" {
t, err := time.Parse(time.RFC3339, req.ExpiresAt)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "expires_at must be RFC3339"})
return
}
expires = &t
}
k, err := svc.License.Generate(
c.Request.Context(),
req.Customer, req.Plan, req.Notes,
req.MaxActivations, expires,
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, k)
}
}
func licenseListHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
rows, err := svc.License.List(c.Request.Context())
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, rows)
}
}
type activateReq struct {
Key string `json:"key" binding:"required"`
DeviceID string `json:"device_id" binding:"required"`
DeviceName string `json:"device_name"`
}
func licenseActivateHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
var req activateReq
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
a, err := svc.License.Activate(
c.Request.Context(), req.Key, req.DeviceID, req.DeviceName, c.ClientIP(),
)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, a)
}
}
func licenseListActivationsHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
rows, err := svc.License.ListActivations(c.Request.Context(), c.Param("id"))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, rows)
}
}
func licenseUnbindHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if err := svc.License.Unbind(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.Status(http.StatusNoContent)
}
}
func licenseRevokeHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
if err := svc.License.Revoke(c.Request.Context(), c.Param("id")); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.Status(http.StatusNoContent)
}
}
func licenseHeartbeatHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
actID := c.Query("activation_id")
if actID == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "activation_id required"})
return
}
if err := svc.License.Heartbeat(c.Request.Context(), actID); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
}
func licenseStatusHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
keyID := c.Query("key_id")
if keyID == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "key_id required"})
return
}
out, err := svc.License.Status(c.Request.Context(), keyID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, out)
}
}
+7 -8
View File
@@ -8,6 +8,7 @@ import (
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/middleware"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/service"
)
@@ -38,7 +39,7 @@ func (h *PermissionHandler) GetUserPermissions(c *gin.Context) {
return
}
perms, err := h.svc.Permission.GetByUserID(c.Request.Context(), userID)
perms, err := h.svc.Permissions.Effective(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})
@@ -64,15 +65,13 @@ func (h *PermissionHandler) UpdateUserPermissions(c *gin.Context) {
return
}
var req struct {
Permissions map[string]bool `json:"permissions"`
}
var req model.UserPermission
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 {
if err := h.svc.Permissions.Save(c.Request.Context(), userID, &req); 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
@@ -97,7 +96,7 @@ func (h *PermissionHandler) ResetUserPermissions(c *gin.Context) {
return
}
if err := h.svc.Permission.ResetToDefault(c.Request.Context(), userID); err != nil {
if _, err := h.svc.Permissions.Reset(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
@@ -115,7 +114,7 @@ func (h *PermissionHandler) GetMyPermissions(c *gin.Context) {
return
}
perms, err := h.svc.Permission.GetPermissionMap(c.Request.Context(), currentUserID)
perms, err := h.svc.Permissions.Effective(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})
@@ -127,7 +126,7 @@ func (h *PermissionHandler) GetMyPermissions(c *gin.Context) {
tier := middleware.GetUserTier(c)
c.JSON(http.StatusOK, gin.H{
"code": 0,
"code": 0,
"message": "ok",
"data": gin.H{
"permissions": perms,
+20 -26
View File
@@ -33,17 +33,17 @@ func (h *SiteHandler) ListSites(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": sites})
}
// GetSite 获取单个站点详情(解密敏感字段)。
// GetSite 获取单个站点详情。
func (h *SiteHandler) GetSite(c *gin.Context) {
site, err := h.svc.Site.GetByID(c.Request.Context(), c.Param("id"))
site, err := h.svc.Site.FindByID(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
}
if site == nil {
c.JSON(http.StatusNotFound, gin.H{"code": 1, "message": "site not found"})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": site})
}
@@ -54,41 +54,30 @@ func (h *SiteHandler) CreateSite(c *gin.Context) {
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 {
if err := h.svc.Site.Create(c.Request.Context(), &site); 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})
c.JSON(http.StatusCreated, gin.H{"code": 0, "message": "ok", "data": site})
}
// UpdateSite 更新站点。
func (h *SiteHandler) UpdateSite(c *gin.Context) {
var site model.Site
if err := c.ShouldBindJSON(&site); err != nil {
patch := make(map[string]any)
if err := c.ShouldBindJSON(&patch); 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
}
if err := h.svc.Site.Update(c.Request.Context(), c.Param("id"), patch); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": updated})
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok"})
}
// 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
}
@@ -97,11 +86,16 @@ func (h *SiteHandler) DeleteSite(c *gin.Context) {
// 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()})
ok, msg, err := h.svc.Site.TestConnection(c.Request.Context(), c.Param("id"))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"code": 1, "message": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok"})
if !ok {
c.JSON(http.StatusBadRequest, gin.H{"code": 1, "message": msg})
return
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": msg})
}
// GetSiteTypes 返回支持的站点类型列表。
+24 -8
View File
@@ -2,6 +2,7 @@
package handler
import (
"encoding/json"
"net/http"
"github.com/gin-gonic/gin"
@@ -66,21 +67,36 @@ func createSiteHandler(svc *service.Container) gin.HandlerFunc {
if req.Enabled != nil {
enabled = *req.Enabled
}
// Pack fields not in the core model into Extra JSON.
extraMap := map[string]any{}
if req.UserAgent != "" {
extraMap["user_agent"] = req.UserAgent
}
if req.RSSURL != "" {
extraMap["rss_url"] = req.RSSURL
}
if req.Timeout > 0 {
extraMap["timeout"] = req.Timeout
}
if req.Priority > 0 {
extraMap["priority"] = req.Priority
}
extraMap["use_proxy"] = req.UseProxy
if req.Downloader != "" {
extraMap["downloader"] = req.Downloader
}
extraJSON, _ := json.Marshal(extraMap)
site := &model.Site{
Name: req.Name,
BaseURL: req.BaseURL,
SiteType: req.SiteType,
URL: req.BaseURL,
Type: req.SiteType,
AuthType: req.AuthType,
Cookie: req.Cookie,
APIKey: req.APIKey,
AuthHeader: req.AuthHeader,
UserAgent: req.UserAgent,
RSSURL: req.RSSURL,
Timeout: req.Timeout,
Priority: req.Priority,
UseProxy: req.UseProxy,
Extra: string(extraJSON),
Enabled: enabled,
Downloader: req.Downloader,
}
if err := svc.Site.Create(c.Request.Context(), site); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
+7 -1
View File
@@ -50,11 +50,17 @@ func siteUserdataHandler(svc *service.Container) gin.HandlerFunc {
c.JSON(http.StatusNotFound, gin.H{"error": "site not found"})
return
}
loginStatus := "unknown"
if s.LastError == "ok" {
loginStatus = "ok"
} else if s.LastError != "" {
loginStatus = "fail"
}
c.JSON(http.StatusOK, gin.H{
"site_id": s.ID,
"name": s.Name,
"cookie_set": len(s.Cookie) > 0,
"login_status": s.LoginStatus,
"login_status": loginStatus,
"note": "userdata parsing not implemented; stub",
})
}
-110
View File
@@ -204,51 +204,6 @@ type AccessLog struct {
Detail string `gorm:"type:text" json:"detail"`
}
// Site stores a PT/BT tracker site configuration used by the subscription
// and cross-site search system. Mirrors the original MediaStation sites table.
//
// Supported site types: nexusphp / gazelle / unit3d / mteam / custom_rss
// Supported auth types: cookie / api_key / authorization
type Site struct {
Base
Name string `gorm:"size:128;not null" json:"name"`
BaseURL string `gorm:"size:512;not null" json:"base_url"`
SiteType string `gorm:"size:32;default:nexusphp" json:"site_type"`
AuthType string `gorm:"size:32;default:cookie" json:"auth_type"`
Cookie string `gorm:"type:text" json:"cookie,omitempty"`
APIKey string `gorm:"size:512" json:"api_key,omitempty"`
AuthHeader string `gorm:"size:512" json:"auth_header,omitempty"`
UserAgent string `gorm:"size:512" json:"user_agent,omitempty"`
RSSURL string `gorm:"size:1024" json:"rss_url,omitempty"`
Timeout int `gorm:"default:15" json:"timeout"`
Priority int `gorm:"default:50" json:"priority"`
UseProxy bool `gorm:"default:false" json:"use_proxy"`
Enabled bool `gorm:"default:true" json:"enabled"`
LoginStatus string `gorm:"size:20;default:unknown" json:"login_status"`
Downloader string `gorm:"size:50" json:"downloader,omitempty"`
}
// NotifyChannel is one named outbound notification destination.
//
// The Config column holds a JSON blob whose schema depends on the
// ChannelType (telegram/wechat/bark/webhook):
//
// telegram → {bot_token, chat_id}
// wechat → {sendkey}
// bark → {device_key, server?}
// webhook → {url, method, headers (JSON string), body_template}
//
// The Events column is a JSON array of event-type strings the channel
// subscribes to; an empty array means "all events".
type NotifyChannel struct {
Base
Name string `gorm:"size:128;not null" json:"name"`
ChannelType string `gorm:"size:32;not null" json:"channel_type"`
Config string `gorm:"type:text;not null" json:"config"`
Enabled bool `gorm:"default:true" json:"enabled"`
Events string `gorm:"type:text;default:'[]'" json:"events"`
}
// PlayProfile lets one user define multiple "viewing personas" with
// different content-rating limits, library access, and player defaults.
// The original Vue project sketched this out as a forward-looking
@@ -274,28 +229,6 @@ type PlayProfile struct {
LastActiveAt *time.Time `json:"last_active_at,omitempty"`
}
// UserPermission stores per-user feature toggles for the React UI's
// menu visibility + route guards. The original Python project surfaces
// 11 boolean flags; we mirror the same set so the existing frontend
// can swap to the Go API without code changes.
type UserPermission struct {
UserID string `gorm:"primaryKey;size:36" json:"user_id"`
CanPlayMedia bool `gorm:"default:true" json:"can_play_media"`
CanFavorite bool `gorm:"default:true" json:"can_favorite"`
CanViewHistory bool `gorm:"default:true" json:"can_view_history"`
CanViewDashboard bool `gorm:"default:true" json:"can_view_dashboard"`
CanViewDiscover bool `gorm:"default:true" json:"can_view_discover"`
CanManageDownloads bool `gorm:"default:false" json:"can_manage_downloads"`
CanManageSubscriptions bool `gorm:"default:false" json:"can_manage_subscriptions"`
CanManageSites bool `gorm:"default:false" json:"can_manage_sites"`
CanManageFiles bool `gorm:"default:false" json:"can_manage_files"`
CanManageSTRM bool `gorm:"default:false" json:"can_manage_strm"`
CanCast bool `gorm:"default:true" json:"can_cast"`
CanUseAIAssistant bool `gorm:"default:false" json:"can_use_ai_assistant"`
CanAccessSettings bool `gorm:"default:false" json:"can_access_settings"`
UpdatedAt time.Time `json:"updated_at"`
}
// StorageConfig holds the connection settings for one external storage
// backend (Alist / S3 / WebDAV). Type column makes the row poly-typed
// — Config is a JSON blob whose shape is determined by Type.
@@ -311,47 +244,6 @@ type StorageConfig struct {
LastError string `gorm:"size:512" json:"last_error,omitempty"`
}
// LicenseKey is one issued license for a customer. Activations live in
// a child table so a single key can bind to multiple devices when its
// MaxActivations > 1.
type LicenseKey struct {
Base
Key string `gorm:"uniqueIndex;size:64;not null" json:"key"`
Customer string `gorm:"size:128" json:"customer,omitempty"`
Plan string `gorm:"size:32;default:basic" json:"plan"`
MaxActivations int `gorm:"default:1" json:"max_activations"`
IssuedAt time.Time `json:"issued_at"`
ExpiresAt *time.Time `json:"expires_at,omitempty"`
Revoked bool `gorm:"default:false" json:"revoked"`
Notes string `gorm:"type:text" json:"notes,omitempty"`
}
// LicenseActivation is one (key, device) binding.
type LicenseActivation struct {
Base
KeyID string `gorm:"index;size:36;not null" json:"key_id"`
DeviceID string `gorm:"size:128;not null" json:"device_id"`
DeviceName string `gorm:"size:128" json:"device_name,omitempty"`
IP string `gorm:"size:64" json:"ip,omitempty"`
UnboundAt *time.Time `json:"unbound_at,omitempty"`
HeartbeatAt *time.Time `json:"heartbeat_at,omitempty"`
}
// DownloadClient is one configured downloader (qBittorrent / Aria2 /
// Transmission). We keep the password column out of JSON so list calls
// don't leak secrets to the React UI.
type DownloadClient struct {
Base
Name string `gorm:"size:128;not null" json:"name"`
Type string `gorm:"size:16;not null" json:"type"` // qbittorrent / transmission / aria2
URL string `gorm:"size:512;not null" json:"url"`
Username string `gorm:"size:128" json:"username,omitempty"`
Password string `gorm:"size:512" json:"-"`
SavePath string `gorm:"size:1024" json:"save_path,omitempty"`
IsDefault bool `gorm:"default:false" json:"is_default"`
Enabled bool `gorm:"default:true" json:"enabled"`
}
// AssistantSession groups a multi-turn chat with the AI assistant.
type AssistantSession struct {
Base
@@ -397,8 +289,6 @@ func AllModels() []interface{} {
&STRMRecord{},
&PlayProfile{},
&StorageConfig{},
&LicenseKey{},
&LicenseActivation{},
&AssistantSession{},
&AssistantMessage{},
}
+64
View File
@@ -0,0 +1,64 @@
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// AssistantRepository persists model.AssistantSession + AssistantMessage records.
type AssistantRepository struct{ db *gorm.DB }
// ─── Session ────────────────────────────────────────────────────────────
// CreateSession inserts a new chat session.
func (r *AssistantRepository) CreateSession(ctx context.Context, s *model.AssistantSession) error {
return r.db.WithContext(ctx).Create(s).Error
}
// FindSession returns a session by ID, or (nil, nil).
func (r *AssistantRepository) FindSession(ctx context.Context, id string) (*model.AssistantSession, error) {
var s model.AssistantSession
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
}
// ListSessions returns sessions for a user, or all when userID is empty.
func (r *AssistantRepository) ListSessions(ctx context.Context, userID string) ([]model.AssistantSession, error) {
q := r.db.WithContext(ctx).Model(&model.AssistantSession{})
if userID != "" {
q = q.Where("user_id = ?", userID)
}
var rows []model.AssistantSession
err := q.Order("created_at desc").Find(&rows).Error
return rows, err
}
// DeleteSession soft-deletes a session (cascade handled by GORM hooks if set).
func (r *AssistantRepository) DeleteSession(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.AssistantSession{}, "id = ?", id).Error
}
// ─── Message ────────────────────────────────────────────────────────────
// AppendMessage inserts a new message into a session.
func (r *AssistantRepository) AppendMessage(ctx context.Context, m *model.AssistantMessage) error {
return r.db.WithContext(ctx).Create(m).Error
}
// ListMessages returns all messages for a session in chronological order.
func (r *AssistantRepository) ListMessages(ctx context.Context, sessionID string) ([]model.AssistantMessage, error) {
var rows []model.AssistantMessage
err := r.db.WithContext(ctx).Where("session_id = ?", sessionID).
Order("created_at asc").Find(&rows).Error
return rows, err
}
+63
View File
@@ -0,0 +1,63 @@
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// PlayProfileRepository persists model.PlayProfile records.
type PlayProfileRepository struct{ db *gorm.DB }
// Create inserts a new play profile.
func (r *PlayProfileRepository) Create(ctx context.Context, p *model.PlayProfile) error {
return r.db.WithContext(ctx).Create(p).Error
}
// FindByID returns the profile or (nil, nil).
func (r *PlayProfileRepository) FindByID(ctx context.Context, id string) (*model.PlayProfile, error) {
var p model.PlayProfile
err := r.db.WithContext(ctx).Where("id = ?", id).First(&p).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &p, nil
}
// List returns every profile (admin view).
func (r *PlayProfileRepository) List(ctx context.Context) ([]model.PlayProfile, error) {
var rows []model.PlayProfile
err := r.db.WithContext(ctx).Order("created_at desc").Find(&rows).Error
return rows, err
}
// ListByUser returns profiles owned by a user.
func (r *PlayProfileRepository) ListByUser(ctx context.Context, userID string) ([]model.PlayProfile, error) {
var rows []model.PlayProfile
err := r.db.WithContext(ctx).Where("user_id = ?", userID).
Order("created_at desc").Find(&rows).Error
return rows, err
}
// Update applies a partial update to a profile row.
func (r *PlayProfileRepository) Update(ctx context.Context, id string, patch map[string]any) error {
return r.db.WithContext(ctx).Model(&model.PlayProfile{}).
Where("id = ?", id).Updates(patch).Error
}
// Delete soft-deletes a profile.
func (r *PlayProfileRepository) Delete(ctx context.Context, id string) error {
return r.db.WithContext(ctx).Delete(&model.PlayProfile{}, "id = ?", id).Error
}
// ClearDefaultsFor resets is_default for all of a user's profiles.
func (r *PlayProfileRepository) ClearDefaultsFor(ctx context.Context, userID string) error {
return r.db.WithContext(ctx).Model(&model.PlayProfile{}).
Where("user_id = ?", userID).Update("is_default", false).Error
}
+6
View File
@@ -38,6 +38,9 @@ type Container struct {
NotifyChannel *NotifyChannelRepository
Site *SiteRepository
STRM *STRMRepository
PlayProfile *PlayProfileRepository
StorageConfig *StorageConfigRepository
Assistant *AssistantRepository
}
// New 将每个 repository 连接到单个 *gorm.DB。
@@ -62,6 +65,9 @@ func New(db *gorm.DB) *Container {
NotifyChannel: &NotifyChannelRepository{db: db},
Site: &SiteRepository{db: db},
STRM: &STRMRepository{db: db},
PlayProfile: &PlayProfileRepository{db: db},
StorageConfig: &StorageConfigRepository{db: db},
Assistant: &AssistantRepository{db: db},
}
}
@@ -0,0 +1,39 @@
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// StorageConfigRepository persists model.StorageConfig records.
type StorageConfigRepository struct{ db *gorm.DB }
// Get returns the config row by type, or (nil, nil).
func (r *StorageConfigRepository) Get(ctx context.Context, kind string) (*model.StorageConfig, error) {
var c model.StorageConfig
err := r.db.WithContext(ctx).Where("type = ?", kind).First(&c).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &c, nil
}
// List returns all storage configs.
func (r *StorageConfigRepository) List(ctx context.Context) ([]model.StorageConfig, error) {
var rows []model.StorageConfig
err := r.db.WithContext(ctx).Order("type asc").Find(&rows).Error
return rows, err
}
// Upsert creates or replaces a storage config keyed by Type.
func (r *StorageConfigRepository) Upsert(ctx context.Context, c *model.StorageConfig) error {
return r.db.WithContext(ctx).Where("type = ?", c.Type).
Assign(*c).FirstOrCreate(c).Error
}
+26 -12
View File
@@ -39,10 +39,9 @@ func NewDownloadClientService(log *zap.Logger, repo *repository.Container) *Down
type DownloadClientInput struct {
Name string `json:"name" binding:"required"`
Type string `json:"type" binding:"required"`
URL string `json:"url" binding:"required"`
Host string `json:"host" binding:"required"`
Username string `json:"username,omitempty"`
Password string `json:"password,omitempty"`
SavePath string `json:"save_path,omitempty"`
IsDefault bool `json:"is_default"`
Enabled bool `json:"enabled"`
}
@@ -60,10 +59,9 @@ func (s *DownloadClientService) Create(ctx context.Context, in DownloadClientInp
c := &model.DownloadClient{
Name: strings.TrimSpace(in.Name),
Type: in.Type,
URL: strings.TrimSpace(in.URL),
Host: strings.TrimSpace(in.Host),
Username: in.Username,
Password: in.Password,
SavePath: in.SavePath,
IsDefault: in.IsDefault,
Enabled: in.Enabled,
}
@@ -81,9 +79,8 @@ func (s *DownloadClientService) Update(ctx context.Context, id string, in Downlo
patch := map[string]any{
"name": strings.TrimSpace(in.Name),
"type": in.Type,
"url": strings.TrimSpace(in.URL),
"host": strings.TrimSpace(in.Host),
"username": in.Username,
"save_path": in.SavePath,
"is_default": in.IsDefault,
"enabled": in.Enabled,
}
@@ -91,7 +88,24 @@ func (s *DownloadClientService) Update(ctx context.Context, id string, in Downlo
if in.Password != "" {
patch["password"] = in.Password
}
if err := s.repo.DownloadClient.Update(ctx, id, patch); err != nil {
// Fetch existing row, apply patch via Save
existing, err := s.repo.DownloadClient.FindByID(ctx, id)
if err != nil {
return nil, err
}
if existing == nil {
return nil, errors.New("client not found")
}
existing.Name = patch["name"].(string)
existing.Type = patch["type"].(string)
existing.Host = patch["host"].(string)
existing.Username = patch["username"].(string)
existing.IsDefault = patch["is_default"].(bool)
existing.Enabled = patch["enabled"].(bool)
if pw, ok := patch["password"]; ok {
existing.Password = pw.(string)
}
if err := s.repo.DownloadClient.Update(ctx, existing); err != nil {
return nil, err
}
return s.repo.DownloadClient.FindByID(ctx, id)
@@ -120,7 +134,7 @@ func (s *DownloadClientService) Test(ctx context.Context, id string) error {
body.Set("password", c.Password)
req, _ := http.NewRequestWithContext(
ctx, http.MethodPost,
strings.TrimRight(c.URL, "/")+"/api/v2/auth/login",
strings.TrimRight(c.Host, "/")+"/api/v2/auth/login",
strings.NewReader(body.Encode()),
)
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
@@ -134,7 +148,7 @@ func (s *DownloadClientService) Test(ctx context.Context, id string) error {
}
return nil
case "aria2", "transmission":
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, c.URL, nil)
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, c.Host, nil)
resp, err := s.client.Do(req)
if err != nil {
return err
@@ -163,7 +177,7 @@ func (s *DownloadClientService) Aria2GlobalStats(ctx context.Context, clientID s
`{"jsonrpc":"2.0","id":"x","method":"aria2.getGlobalStat","params":["token:%s"]}`,
c.Password,
)
req, _ := http.NewRequestWithContext(ctx, http.MethodPost, c.URL,
req, _ := http.NewRequestWithContext(ctx, http.MethodPost, c.Host,
strings.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
resp, err := s.client.Do(req)
@@ -183,8 +197,8 @@ func validateClient(in DownloadClientInput) error {
if strings.TrimSpace(in.Name) == "" {
return errors.New("name required")
}
if strings.TrimSpace(in.URL) == "" {
return errors.New("url required")
if strings.TrimSpace(in.Host) == "" {
return errors.New("host required")
}
switch in.Type {
case "qbittorrent", "aria2", "transmission":
-160
View File
@@ -1,160 +0,0 @@
// Package service — license key management.
//
// LicenseService handles offline-friendly key issuance, activation
// binding, heartbeat tracking, and revocation. Keys are 24 random
// uppercase chars in groups of four (e.g. ABCD-1234-EFGH-5678-IJKL-90MN)
// — the same shape the Vue admin UI expects.
package service
import (
"context"
"crypto/rand"
"errors"
"strings"
"time"
"go.uber.org/zap"
"github.com/ShukeBta/MediaStationGo/internal/model"
"github.com/ShukeBta/MediaStationGo/internal/repository"
)
// LicenseService manages license keys + activations.
type LicenseService struct {
log *zap.Logger
repo *repository.Container
}
// NewLicenseService is the constructor.
func NewLicenseService(log *zap.Logger, repo *repository.Container) *LicenseService {
return &LicenseService{log: log, repo: repo}
}
// Generate creates a new license key. ExpiresAt nil means "perpetual".
func (s *LicenseService) Generate(
ctx context.Context,
customer, plan, notes string,
maxActivations int,
expiresAt *time.Time,
) (*model.LicenseKey, error) {
if maxActivations <= 0 {
maxActivations = 1
}
k := &model.LicenseKey{
Key: randomLicenseKey(),
Customer: strings.TrimSpace(customer),
Plan: strings.TrimSpace(plan),
MaxActivations: maxActivations,
Notes: strings.TrimSpace(notes),
IssuedAt: time.Now(),
ExpiresAt: expiresAt,
}
if err := s.repo.License.Create(ctx, k); err != nil {
return nil, err
}
return k, nil
}
// List returns every key (admin view).
func (s *LicenseService) List(ctx context.Context) ([]model.LicenseKey, error) {
return s.repo.License.List(ctx)
}
// Activate binds a key to a device. Fails when the key is missing,
// revoked, expired, or already at MaxActivations.
func (s *LicenseService) Activate(
ctx context.Context,
key, deviceID, deviceName, ip string,
) (*model.LicenseActivation, error) {
k, err := s.repo.License.FindByKey(ctx, key)
if err != nil {
return nil, err
}
if k == nil {
return nil, errors.New("invalid key")
}
if k.Revoked {
return nil, errors.New("key revoked")
}
if k.ExpiresAt != nil && k.ExpiresAt.Before(time.Now()) {
return nil, errors.New("key expired")
}
count, err := s.repo.License.CountActiveActivations(ctx, k.ID)
if err != nil {
return nil, err
}
if int(count) >= k.MaxActivations {
return nil, errors.New("activation limit reached")
}
a := &model.LicenseActivation{
KeyID: k.ID,
DeviceID: strings.TrimSpace(deviceID),
DeviceName: strings.TrimSpace(deviceName),
IP: ip,
}
if err := s.repo.License.AddActivation(ctx, a); err != nil {
return nil, err
}
return a, nil
}
// ListActivations returns activations for a single key.
func (s *LicenseService) ListActivations(ctx context.Context, keyID string) ([]model.LicenseActivation, error) {
return s.repo.License.ListActivations(ctx, keyID)
}
// Unbind marks one activation as released.
func (s *LicenseService) Unbind(ctx context.Context, activationID string) error {
return s.repo.License.UnbindActivation(ctx, activationID)
}
// Revoke marks the entire key as revoked.
func (s *LicenseService) Revoke(ctx context.Context, keyID string) error {
return s.repo.License.Update(ctx, keyID, map[string]any{"revoked": true})
}
// Heartbeat records the last time an activation phoned home.
func (s *LicenseService) Heartbeat(ctx context.Context, activationID string) error {
return s.repo.License.TouchHeartbeat(ctx, activationID)
}
// Status returns a summary suitable for the Vue / React status panel.
func (s *LicenseService) Status(ctx context.Context, keyID string) (map[string]any, error) {
k, err := s.repo.License.FindByID(ctx, keyID)
if err != nil {
return nil, err
}
if k == nil {
return nil, errors.New("key not found")
}
count, _ := s.repo.License.CountActiveActivations(ctx, keyID)
valid := !k.Revoked
if k.ExpiresAt != nil && k.ExpiresAt.Before(time.Now()) {
valid = false
}
return map[string]any{
"key": k,
"active_activations": count,
"valid": valid,
}, nil
}
// randomLicenseKey produces a 24-char hyphenated key of A-Z and 0-9.
func randomLicenseKey() string {
const alphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" // omit confusables
out := make([]byte, 24)
buf := make([]byte, 24)
_, _ = rand.Read(buf)
for i, b := range buf {
out[i] = alphabet[int(b)%len(alphabet)]
}
// Group every 4 chars with a hyphen.
var sb strings.Builder
for i, c := range out {
if i > 0 && i%4 == 0 {
sb.WriteByte('-')
}
sb.WriteByte(byte(c))
}
return sb.String()
}
+30 -15
View File
@@ -41,11 +41,11 @@ func NewNotifyChannelService(log *zap.Logger, repo *repository.Container) *Notif
// ChannelInput is the shape accepted by Create / Update. Config is a
// generic map; it gets serialised to JSON before being persisted.
type ChannelInput struct {
Name string `json:"name" binding:"required"`
ChannelType string `json:"channel_type" binding:"required"`
Config map[string]any `json:"config"`
Events []string `json:"events"`
Enabled *bool `json:"enabled,omitempty"`
Name string `json:"name" binding:"required"`
Type string `json:"type" binding:"required"`
Config map[string]any `json:"config"`
Events []string `json:"events"`
Enabled *bool `json:"enabled,omitempty"`
}
// channelView is the public shape — Config is decoded back to a map so
@@ -96,7 +96,7 @@ func (s *NotifyChannelService) Create(ctx context.Context, in ChannelInput) (*ch
evBlob, _ := json.Marshal(in.Events)
n := &model.NotifyChannel{
Name: strings.TrimSpace(in.Name),
ChannelType: in.ChannelType,
Type: in.Type,
Config: string(cfgBlob),
Events: string(evBlob),
Enabled: true,
@@ -119,15 +119,30 @@ func (s *NotifyChannelService) Update(ctx context.Context, id string, in Channel
cfgBlob, _ := json.Marshal(in.Config)
evBlob, _ := json.Marshal(in.Events)
patch := map[string]any{
"name": strings.TrimSpace(in.Name),
"channel_type": in.ChannelType,
"config": string(cfgBlob),
"events": string(evBlob),
"name": strings.TrimSpace(in.Name),
"type": in.Type,
"config": string(cfgBlob),
"events": string(evBlob),
}
if in.Enabled != nil {
patch["enabled"] = *in.Enabled
}
if err := s.repo.NotifyChannel.Update(ctx, id, patch); err != nil {
// Fetch existing row, apply patch via repo Update
existing, err := s.repo.NotifyChannel.FindByID(ctx, id)
if err != nil {
return nil, err
}
if existing == nil {
return nil, errors.New("channel not found")
}
existing.Name = patch["name"].(string)
existing.Type = patch["type"].(string)
existing.Config = patch["config"].(string)
existing.Events = patch["events"].(string)
if en, ok := patch["enabled"]; ok {
existing.Enabled = en.(bool)
}
if err := s.repo.NotifyChannel.Update(ctx, existing); err != nil {
return nil, err
}
row, err := s.repo.NotifyChannel.FindByID(ctx, id)
@@ -201,7 +216,7 @@ func (s *NotifyChannelService) dispatchOne(ctx context.Context, n model.NotifyCh
cfg := map[string]any{}
_ = json.Unmarshal([]byte(n.Config), &cfg)
switch n.ChannelType {
switch n.Type {
case "telegram":
token := str(cfg["bot_token"])
chat := str(cfg["chat_id"])
@@ -279,7 +294,7 @@ func (s *NotifyChannelService) dispatchOne(ctx context.Context, n model.NotifyCh
}
return s.do(req)
}
return fmt.Errorf("unknown channel type %q", n.ChannelType)
return fmt.Errorf("unknown channel type %q", n.Type)
}
func (s *NotifyChannelService) do(req *http.Request) error {
@@ -300,10 +315,10 @@ func validateChannel(in ChannelInput) error {
if strings.TrimSpace(in.Name) == "" {
return errors.New("name required")
}
switch in.ChannelType {
switch in.Type {
case "telegram", "wechat", "bark", "webhook":
default:
return fmt.Errorf("unsupported channel type %q", in.ChannelType)
return fmt.Errorf("unsupported channel type %q", in.Type)
}
return nil
}
+23 -6
View File
@@ -40,7 +40,7 @@ func DefaultPermissions(userID string) *model.UserPermission {
CanManageSubscriptions: false,
CanManageSites: false,
CanManageFiles: false,
CanManageSTRM: false,
CanManageStrm: false,
CanUseAIAssistant: false,
CanAccessSettings: false,
}
@@ -59,7 +59,7 @@ func adminGrant(userID string) *model.UserPermission {
CanManageSubscriptions: true,
CanManageSites: true,
CanManageFiles: true,
CanManageSTRM: true,
CanManageStrm: true,
CanCast: true,
CanUseAIAssistant: true,
CanAccessSettings: true,
@@ -79,7 +79,7 @@ func (s *PermissionService) Effective(ctx context.Context, userID string) (*mode
if u.Role == "admin" {
return adminGrant(userID), nil
}
row, err := s.repo.Permission.Get(ctx, userID)
row, err := s.repo.Permission.FindByUserID(ctx, userID)
if err != nil {
return nil, err
}
@@ -89,7 +89,7 @@ func (s *PermissionService) Effective(ctx context.Context, userID string) (*mode
// Seed defaults on first read so subsequent updates have a row to
// patch.
def := DefaultPermissions(userID)
if err := s.repo.Permission.Save(ctx, def); err != nil {
if err := s.repo.Permission.Upsert(ctx, def); err != nil {
return nil, err
}
return def, nil
@@ -98,13 +98,30 @@ func (s *PermissionService) Effective(ctx context.Context, userID string) (*mode
// Save persists the user permission patch (admin only — caller checks).
func (s *PermissionService) Save(ctx context.Context, userID string, in *model.UserPermission) error {
in.UserID = userID
return s.repo.Permission.Save(ctx, in)
return s.repo.Permission.Upsert(ctx, in)
}
// EnsureForUser guarantees a permission row exists for the given user.
// If one already exists it is a no-op; otherwise a default row is created.
func (s *PermissionService) EnsureForUser(ctx context.Context, userID string) (*model.UserPermission, error) {
row, err := s.repo.Permission.FindByUserID(ctx, userID)
if err != nil {
return nil, err
}
if row != nil {
return row, nil
}
def := DefaultPermissions(userID)
if err := s.repo.Permission.Upsert(ctx, def); err != nil {
return nil, err
}
return def, nil
}
// Reset reverts to the non-admin defaults.
func (s *PermissionService) Reset(ctx context.Context, userID string) (*model.UserPermission, error) {
def := DefaultPermissions(userID)
if err := s.repo.Permission.Save(ctx, def); err != nil {
if err := s.repo.Permission.Upsert(ctx, def); err != nil {
return nil, err
}
return def, nil
-117
View File
@@ -1,117 +0,0 @@
// 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"
}
+2 -7
View File
@@ -57,7 +57,6 @@ type Container struct {
PlayProfiles *PlayProfileService
Permissions *PermissionService
StorageCfg *StorageConfigService
License *LicenseService
DownloadClients *DownloadClientService
Assistant *AssistantService
Organizer *OrganizerService
@@ -108,21 +107,18 @@ func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Cont
playProfiles := NewPlayProfileService(log, repos)
permissions := NewPermissionService(log, repos)
storageCfg := NewStorageConfigService(log, repos, crypto)
licenseSvc := NewLicenseService(log, repos)
downloadClients := NewDownloadClientService(log, repos)
assistant := NewAssistantService(log, repos, ai)
organizer := NewOrganizerService(cfg, log, repos)
douban := NewDoubanProvider(cfg, log)
siteService := NewSiteService(log, repos)
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)
siteSvc := NewSiteService(log, repos)
ctx, cancel := context.WithCancel(context.Background())
@@ -132,7 +128,7 @@ func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Cont
Repo: repos,
WSHub: hub,
SSEHub: sseHub,
Auth: NewAuthService(cfg, log, repos, tokenSvc, permissionSvc),
Auth: NewAuthService(cfg, log, repos, tokenSvc, permissions),
Media: NewMediaService(cfg, log, repos),
Scan: scanner,
Stream: NewStreamService(cfg, log, repos, transcoder),
@@ -169,7 +165,6 @@ func New(cfg *config.Config, log *zap.Logger, repos *repository.Container) *Cont
PlayProfiles: playProfiles,
Permissions: permissions,
StorageCfg: storageCfg,
License: licenseSvc,
DownloadClients: downloadClients,
Assistant: assistant,
Organizer: organizer,
+27 -22
View File
@@ -32,26 +32,23 @@ func NewSiteService(log *zap.Logger, repo *repository.Container) *SiteService {
// Create persists a new site.
func (s *SiteService) Create(ctx context.Context, site *model.Site) error {
if strings.TrimSpace(site.Name) == "" || strings.TrimSpace(site.BaseURL) == "" {
return errors.New("name and base_url required")
if strings.TrimSpace(site.Name) == "" || strings.TrimSpace(site.URL) == "" {
return errors.New("name and url required")
}
site.BaseURL = strings.TrimRight(site.BaseURL, "/")
if site.SiteType == "" {
site.SiteType = "nexusphp"
site.URL = strings.TrimRight(site.URL, "/")
if site.Type == "" {
site.Type = "nexusphp"
}
if site.AuthType == "" {
site.AuthType = "cookie"
}
if site.Timeout <= 0 {
site.Timeout = 15
}
return s.repo.DB.WithContext(ctx).Create(site).Error
}
// List returns every site ordered by priority (lower = higher priority).
func (s *SiteService) List(ctx context.Context) ([]model.Site, error) {
var sites []model.Site
err := s.repo.DB.WithContext(ctx).Order("priority asc, created_at asc").Find(&sites).Error
err := s.repo.DB.WithContext(ctx).Order("created_at asc").Find(&sites).Error
return sites, err
}
@@ -83,14 +80,14 @@ func (s *SiteService) TestConnection(ctx context.Context, id string) (bool, stri
return false, "site not found", err
}
client := &http.Client{Timeout: time.Duration(site.Timeout) * time.Second}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, site.BaseURL, nil)
client := &http.Client{Timeout: 15 * time.Second}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, site.URL, nil)
if err != nil {
return false, err.Error(), nil
}
// Apply auth headers.
req.Header.Set("User-Agent", effectiveUA(site))
req.Header.Set("User-Agent", "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/122.0.0.0 Safari/537.36")
switch site.AuthType {
case "cookie":
if site.Cookie != "" {
@@ -108,9 +105,9 @@ func (s *SiteService) TestConnection(ctx context.Context, id string) (bool, stri
resp, err := client.Do(req)
if err != nil {
status := "fail"
now := time.Now()
_ = s.repo.DB.WithContext(ctx).Model(&model.Site{}).Where("id = ?", id).
Update("login_status", status).Error
Updates(map[string]any{"last_error": err.Error(), "last_check_at": &now}).Error
return false, err.Error(), nil
}
defer resp.Body.Close()
@@ -132,8 +129,9 @@ func (s *SiteService) TestConnection(ctx context.Context, id string) (bool, stri
if !ok {
loginStatus = "fail"
}
now := time.Now()
_ = s.repo.DB.WithContext(ctx).Model(&model.Site{}).Where("id = ?", id).
Updates(map[string]any{"login_status": loginStatus, "last_check": time.Now()}).Error
Updates(map[string]any{"last_error": loginStatus, "last_check_at": &now}).Error
return ok, msg, nil
}
@@ -170,18 +168,19 @@ func (s *SiteService) Search(ctx context.Context, keyword string) ([]SearchResul
if adapter == nil {
continue
}
items, err := adapter.Search(ctx, keyword)
cfg := siteModelToConfig(&sites[i])
result, err := adapter.Search(ctx, cfg, keyword, 1)
if err != nil {
s.log.Debug("site search failed",
zap.String("site", sites[i].Name), zap.Error(err))
continue
}
for _, item := range items {
for _, item := range result.Items {
results = append(results, SearchResult{
SiteName: sites[i].Name,
SiteID: sites[i].ID,
Title: item.Title,
TorrentURL: item.TorrentURL,
TorrentURL: item.DetailURL,
DownloadURL: item.DownloadURL,
Size: item.Size,
Seeders: item.Seeders,
@@ -202,9 +201,15 @@ func (s *SiteService) Search(ctx context.Context, keyword string) ([]SearchResul
return results, nil
}
func effectiveUA(site *model.Site) string {
if site.UserAgent != "" {
return site.UserAgent
// siteModelToConfig 将 model.Site 转换为适配器使用的 SiteConfig。
func siteModelToConfig(s *model.Site) SiteConfig {
return SiteConfig{
Name: s.Name,
Type: s.Type,
URL: s.URL,
AuthType: s.AuthType,
Cookie: s.Cookie,
APIKey: s.APIKey,
AuthHeader: s.AuthHeader,
}
return "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/122.0.0.0 Safari/537.36"
}
+35 -28
View File
@@ -12,6 +12,8 @@ import (
"strconv"
"strings"
"time"
"github.com/ShukeBta/MediaStationGo/internal/model"
)
// SiteConfig 站点配置(从 model.Site 解密后的纯文本)。
@@ -26,8 +28,8 @@ type SiteConfig struct {
Extra map[string]string // JSON 扩展配置
}
// SearchResult 站点搜索结果。
type SearchResult struct {
// SiteSearchResult 站点搜索结果(按站点分组的批量搜索结果)。
type SiteSearchResult struct {
SiteName string `json:"site_name"`
Items []TorrentItem `json:"items"`
Total int `json:"total"`
@@ -78,10 +80,10 @@ type SiteAdapter interface {
Authenticate(ctx context.Context, cfg SiteConfig) error
// Search 搜索种子。
Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SearchResult, error)
Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SiteSearchResult, error)
// Browse 浏览种子列表。
Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SearchResult, error)
Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SiteSearchResult, error)
// GetDetail 获取种子详情。
GetDetail(ctx context.Context, cfg SiteConfig, id string) (*TorrentDetail, error)
@@ -194,7 +196,7 @@ func (a *NexusPHPAdapter) Authenticate(ctx context.Context, cfg SiteConfig) erro
return nil
}
func (a *NexusPHPAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SearchResult, error) {
func (a *NexusPHPAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SiteSearchResult, error) {
params := url.Values{}
params.Set("search", keyword)
params.Set("page", strconv.Itoa(page))
@@ -213,7 +215,7 @@ func (a *NexusPHPAdapter) Search(ctx context.Context, cfg SiteConfig, keyword st
return parseNexusPHPHTML(string(data), cfg.Name, cfg.URL)
}
func (a *NexusPHPAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SearchResult, error) {
func (a *NexusPHPAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SiteSearchResult, error) {
params := url.Values{}
if category != "" {
params.Set("cat", category)
@@ -250,8 +252,8 @@ func (a *NexusPHPAdapter) GetDownloadURL(ctx context.Context, cfg SiteConfig, id
}
// parseNexusPHPHTML 解析 NexusPHP 种子列表 HTML。
func parseNexusPHPHTML(html, siteName, baseURL string) (*SearchResult, error) {
result := &SearchResult{
func parseNexusPHPHTML(html, siteName, baseURL string) (*SiteSearchResult, error) {
result := &SiteSearchResult{
SiteName: siteName,
Items: []TorrentItem{},
Page: 1,
@@ -427,7 +429,7 @@ func (a *GazelleAdapter) Authenticate(ctx context.Context, cfg SiteConfig) error
return nil
}
func (a *GazelleAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SearchResult, error) {
func (a *GazelleAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SiteSearchResult, error) {
params := url.Values{}
params.Set("action", "browse")
params.Set("searchstr", keyword)
@@ -445,7 +447,7 @@ func (a *GazelleAdapter) Search(ctx context.Context, cfg SiteConfig, keyword str
return parseGazelleJSON(data, cfg.Name, cfg.URL)
}
func (a *GazelleAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SearchResult, error) {
func (a *GazelleAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SiteSearchResult, error) {
params := url.Values{}
params.Set("action", "browse")
if category != "" {
@@ -534,13 +536,13 @@ func (a *GazelleAdapter) GetDownloadURL(ctx context.Context, cfg SiteConfig, id
}
// parseGazelleJSON 解析 Gazelle JSON 响应。
func parseGazelleJSON(data []byte, siteName, baseURL string) (*SearchResult, error) {
func parseGazelleJSON(data []byte, siteName, baseURL string) (*SiteSearchResult, error) {
var resp map[string]interface{}
if err := json.Unmarshal(data, &resp); err != nil {
return nil, fmt.Errorf("parse JSON: %w", err)
}
result := &SearchResult{
result := &SiteSearchResult{
SiteName: siteName,
Items: []TorrentItem{},
}
@@ -644,7 +646,7 @@ func (a *UNIT3DAdapter) Authenticate(ctx context.Context, cfg SiteConfig) error
return nil
}
func (a *UNIT3DAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SearchResult, error) {
func (a *UNIT3DAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SiteSearchResult, error) {
params := url.Values{}
params.Set("search", keyword)
params.Set("page", strconv.Itoa(page))
@@ -661,7 +663,7 @@ func (a *UNIT3DAdapter) Search(ctx context.Context, cfg SiteConfig, keyword stri
return parseUNIT3DJSON(data, cfg.Name, cfg.URL)
}
func (a *UNIT3DAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SearchResult, error) {
func (a *UNIT3DAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SiteSearchResult, error) {
params := url.Values{}
if category != "" {
params.Set("category", category)
@@ -734,7 +736,7 @@ func (a *UNIT3DAdapter) GetDownloadURL(ctx context.Context, cfg SiteConfig, id s
}
// parseUNIT3DJSON 解析 UNIT3D JSON 响应。
func parseUNIT3DJSON(data []byte, siteName, baseURL string) (*SearchResult, error) {
func parseUNIT3DJSON(data []byte, siteName, baseURL string) (*SiteSearchResult, error) {
var resp struct {
Data []map[string]interface{} `json:"data"`
Meta struct {
@@ -746,7 +748,7 @@ func parseUNIT3DJSON(data []byte, siteName, baseURL string) (*SearchResult, erro
return nil, fmt.Errorf("parse JSON: %w", err)
}
result := &SearchResult{
result := &SiteSearchResult{
SiteName: siteName,
Items: []TorrentItem{},
Page: resp.Meta.CurrentPage,
@@ -831,7 +833,7 @@ func (a *MTeamAdapter) Authenticate(ctx context.Context, cfg SiteConfig) error {
return nil
}
func (a *MTeamAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SearchResult, error) {
func (a *MTeamAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SiteSearchResult, error) {
payload := map[string]interface{}{
"mode": "search",
"keyword": keyword,
@@ -852,7 +854,7 @@ func (a *MTeamAdapter) Search(ctx context.Context, cfg SiteConfig, keyword strin
return parseMTeamJSON(data, cfg.Name, cfg.URL)
}
func (a *MTeamAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SearchResult, error) {
func (a *MTeamAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SiteSearchResult, error) {
payload := map[string]interface{}{
"mode": "browse",
"category": category,
@@ -936,7 +938,7 @@ func (a *MTeamAdapter) GetDownloadURL(ctx context.Context, cfg SiteConfig, id st
}
// parseMTeamJSON 解析 MTeam JSON 响应。
func parseMTeamJSON(data []byte, siteName, baseURL string) (*SearchResult, error) {
func parseMTeamJSON(data []byte, siteName, baseURL string) (*SiteSearchResult, error) {
var resp struct {
Code int `json:"code"`
Data struct {
@@ -948,7 +950,7 @@ func parseMTeamJSON(data []byte, siteName, baseURL string) (*SearchResult, error
return nil, fmt.Errorf("parse JSON: %w", err)
}
result := &SearchResult{
result := &SiteSearchResult{
SiteName: siteName,
Items: []TorrentItem{},
Total: resp.Data.Total,
@@ -1033,7 +1035,7 @@ func (a *DiscuzAdapter) Authenticate(ctx context.Context, cfg SiteConfig) error
return nil
}
func (a *DiscuzAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SearchResult, error) {
func (a *DiscuzAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SiteSearchResult, error) {
params := url.Values{}
params.Set("mod", "forum")
params.Set("srchtxt", keyword)
@@ -1052,7 +1054,7 @@ func (a *DiscuzAdapter) Search(ctx context.Context, cfg SiteConfig, keyword stri
return parseDiscuzHTML(string(data), cfg.Name, cfg.URL)
}
func (a *DiscuzAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SearchResult, error) {
func (a *DiscuzAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SiteSearchResult, error) {
params := url.Values{}
if category != "" {
params.Set("fid", category)
@@ -1117,8 +1119,8 @@ func (a *DiscuzAdapter) GetDownloadURL(ctx context.Context, cfg SiteConfig, id s
}
// parseDiscuzHTML 解析 Discuz HTML 响应。
func parseDiscuzHTML(html, siteName, baseURL string) (*SearchResult, error) {
result := &SearchResult{
func parseDiscuzHTML(html, siteName, baseURL string) (*SiteSearchResult, error) {
result := &SiteSearchResult{
SiteName: siteName,
Items: []TorrentItem{},
Page: 1,
@@ -1183,7 +1185,7 @@ func (a *CustomRSSAdapter) Authenticate(ctx context.Context, cfg SiteConfig) err
return nil
}
func (a *CustomRSSAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SearchResult, error) {
func (a *CustomRSSAdapter) Search(ctx context.Context, cfg SiteConfig, keyword string, page int) (*SiteSearchResult, error) {
searchURL := cfg.URL
// If extra has search URL template, use it
if searchTpl, ok := cfg.Extra["search_url"]; ok && searchTpl != "" {
@@ -1218,7 +1220,7 @@ func (a *CustomRSSAdapter) Search(ctx context.Context, cfg SiteConfig, keyword s
return result, nil
}
func (a *CustomRSSAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SearchResult, error) {
func (a *CustomRSSAdapter) Browse(ctx context.Context, cfg SiteConfig, category string, page int) (*SiteSearchResult, error) {
// RSS browse is essentially the same as search with empty keyword
return a.Search(ctx, cfg, "", page)
}
@@ -1236,8 +1238,8 @@ func (a *CustomRSSAdapter) GetDownloadURL(ctx context.Context, cfg SiteConfig, i
}
// parseRSSXML 解析 RSS XML 内容。
func parseRSSXML(data []byte, siteName, keyword string) (*SearchResult, error) {
result := &SearchResult{
func parseRSSXML(data []byte, siteName, keyword string) (*SiteSearchResult, error) {
result := &SiteSearchResult{
SiteName: siteName,
Items: []TorrentItem{},
}
@@ -1385,3 +1387,8 @@ func GetAdapterForType(siteType string) SiteAdapter {
return NewNexusPHPAdapter()
}
}
// NewSiteAdapter 根据站点模型创建对应的适配器。
func NewSiteAdapter(site *model.Site) SiteAdapter {
return GetAdapterForType(site.Type)
}
-214
View File
@@ -1,214 +0,0 @@
// 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
@@ -1,246 +0,0 @@
// 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
}