mirror of
https://github.com/truewhile/MeBox.git
synced 2026-10-04 04:26:38 +08:00
feat: merge conflict resolution, site management, UI fixes
This commit is contained in:
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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, "<", "<")
|
||||
result = strings.ReplaceAll(result, ">", ">")
|
||||
result = strings.ReplaceAll(result, "&", "&")
|
||||
return result
|
||||
}
|
||||
|
||||
return sb.String()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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, ",")
|
||||
}
|
||||
Reference in New Issue
Block a user