更新前后端代码
This commit is contained in:
@@ -0,0 +1,214 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"filestoragesystem/internal/config"
|
||||
"filestoragesystem/internal/model"
|
||||
"filestoragesystem/internal/repository"
|
||||
"filestoragesystem/internal/utils"
|
||||
"filestoragesystem/pkg/apperr"
|
||||
)
|
||||
|
||||
// AdminService 系统管理服务(设置/日志/备份)
|
||||
type AdminService struct {
|
||||
cfg *config.Config
|
||||
db *gorm.DB
|
||||
settingRepo *repository.SettingRepo
|
||||
opLogRepo *repository.OpLogRepo
|
||||
sysLogRepo *repository.SysLogRepo
|
||||
}
|
||||
|
||||
// ---- 系统设置 ----
|
||||
|
||||
// GetSettings 获取全部设置(按分组组织)
|
||||
func (s *AdminService) GetSettings() (map[string][]model.SystemSetting, error) {
|
||||
settings, err := s.settingRepo.GetAll()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
groups := make(map[string][]model.SystemSetting)
|
||||
for _, st := range settings {
|
||||
groups[st.GroupName] = append(groups[st.GroupName], st)
|
||||
}
|
||||
return groups, nil
|
||||
}
|
||||
|
||||
// GetSettingsByGroup 按分组获取设置
|
||||
func (s *AdminService) GetSettingsByGroup(group string) ([]model.SystemSetting, error) {
|
||||
return s.settingRepo.GetByGroup(group)
|
||||
}
|
||||
|
||||
// UpdateSettings 批量更新设置(仅更新已存在的键)
|
||||
func (s *AdminService) UpdateSettings(values map[string]string) error {
|
||||
if len(values) == 0 {
|
||||
return nil
|
||||
}
|
||||
existing := make(map[string]bool)
|
||||
all, err := s.settingRepo.GetAll()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, st := range all {
|
||||
existing[st.Key] = true
|
||||
}
|
||||
valid := make(map[string]string)
|
||||
for k, v := range values {
|
||||
if existing[k] {
|
||||
valid[k] = v
|
||||
}
|
||||
}
|
||||
if len(valid) == 0 {
|
||||
return fmt.Errorf("没有可更新的有效设置项")
|
||||
}
|
||||
return s.settingRepo.UpdateValues(valid)
|
||||
}
|
||||
|
||||
// ---- 日志 ----
|
||||
|
||||
// ListOperationLogs 操作日志列表
|
||||
func (s *AdminService) ListOperationLogs(f repository.OpLogFilter) ([]model.OperationLog, int64, error) {
|
||||
if f.Page <= 0 {
|
||||
f.Page = 1
|
||||
}
|
||||
if f.PageSize <= 0 {
|
||||
f.PageSize = 20
|
||||
}
|
||||
return s.opLogRepo.List(f)
|
||||
}
|
||||
|
||||
// GetOperationLog 操作日志详情
|
||||
func (s *AdminService) GetOperationLog(id uint) (*model.OperationLog, error) {
|
||||
l, err := s.opLogRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
return l, nil
|
||||
}
|
||||
|
||||
// ListSystemLogs 系统日志列表
|
||||
func (s *AdminService) ListSystemLogs(f repository.SysLogFilter) ([]model.SystemLog, int64, error) {
|
||||
if f.Page <= 0 {
|
||||
f.Page = 1
|
||||
}
|
||||
if f.PageSize <= 0 {
|
||||
f.PageSize = 20
|
||||
}
|
||||
return s.sysLogRepo.List(f)
|
||||
}
|
||||
|
||||
// GetSystemLog 系统日志详情
|
||||
func (s *AdminService) GetSystemLog(id uint) (*model.SystemLog, error) {
|
||||
l, err := s.sysLogRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
return l, nil
|
||||
}
|
||||
|
||||
// CleanupLogs 清理days天前的日志
|
||||
func (s *AdminService) CleanupLogs(days int) (opDeleted, sysDeleted int64, err error) {
|
||||
if days <= 0 {
|
||||
days = 30
|
||||
}
|
||||
opDeleted, err = s.opLogRepo.DeleteBefore(days)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
sysDeleted, err = s.sysLogRepo.DeleteBefore(days)
|
||||
utils.Info("system", "清理%d天前日志:操作日志%d条,系统日志%d条", days, opDeleted, sysDeleted)
|
||||
return
|
||||
}
|
||||
|
||||
// ---- 备份 ----
|
||||
|
||||
// BackupDatabase 使用VACUUM INTO备份数据库
|
||||
func (s *AdminService) BackupDatabase() (string, error) {
|
||||
if err := os.MkdirAll(s.cfg.Storage.Backups, 0o755); err != nil {
|
||||
return "", err
|
||||
}
|
||||
name := fmt.Sprintf("backup-%s.db", time.Now().Format("20060102-150405"))
|
||||
target := filepath.Join(s.cfg.Storage.Backups, name)
|
||||
// SQLite安全备份:VACUUM INTO
|
||||
quoted := strings.ReplaceAll(target, "'", "''")
|
||||
if err := s.db.Exec(fmt.Sprintf("VACUUM INTO '%s'", quoted)).Error; err != nil {
|
||||
utils.Error("system", "数据库备份失败: %v", err)
|
||||
return "", fmt.Errorf("备份失败: %w", err)
|
||||
}
|
||||
utils.Info("system", "数据库备份完成: %s", name)
|
||||
return name, nil
|
||||
}
|
||||
|
||||
// ListBackups 备份列表
|
||||
func (s *AdminService) ListBackups() ([]map[string]interface{}, error) {
|
||||
entries, err := os.ReadDir(s.cfg.Storage.Backups)
|
||||
if err != nil {
|
||||
return []map[string]interface{}{}, nil
|
||||
}
|
||||
var list []map[string]interface{}
|
||||
for _, e := range entries {
|
||||
if e.IsDir() || !strings.HasSuffix(e.Name(), ".db") {
|
||||
continue
|
||||
}
|
||||
info, _ := e.Info()
|
||||
list = append(list, map[string]interface{}{
|
||||
"name": e.Name(),
|
||||
"size": info.Size(),
|
||||
"time": info.ModTime().Format(time.RFC3339),
|
||||
})
|
||||
}
|
||||
sort.Slice(list, func(i, j int) bool {
|
||||
return list[i]["name"].(string) > list[j]["name"].(string)
|
||||
})
|
||||
return list, nil
|
||||
}
|
||||
|
||||
// BackupPath 校验备份文件名并返回其物理路径(供下载)
|
||||
func (s *AdminService) BackupPath(name string) (string, error) {
|
||||
// 只允许 backup-*.db 格式,防止路径穿越
|
||||
if strings.ContainsAny(name, "/\\") || strings.Contains(name, "..") || !strings.HasPrefix(name, "backup-") || !strings.HasSuffix(name, ".db") {
|
||||
return "", fmt.Errorf("非法的备份文件名")
|
||||
}
|
||||
p := filepath.Join(s.cfg.Storage.Backups, name)
|
||||
if _, err := os.Stat(p); err != nil {
|
||||
return "", fmt.Errorf("备份文件不存在")
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
// ---- 站点状态 ----
|
||||
|
||||
// SiteStatus 公开的站点状态信息
|
||||
func (s *AdminService) SiteStatus() map[string]interface{} {
|
||||
get := func(k, def string) string {
|
||||
if v, err := s.settingRepo.GetValue(k); err == nil {
|
||||
return v
|
||||
}
|
||||
return def
|
||||
}
|
||||
return map[string]interface{}{
|
||||
"site_name": get("site_name", "云泽文件存储云平台"),
|
||||
"site_description": get("site_description", ""),
|
||||
"site_logo": get("site_logo", ""),
|
||||
"site_icp": get("site_icp", ""),
|
||||
"site_copyright": get("site_copyright", ""),
|
||||
"site_enabled": get("site_enabled", "true") == "true",
|
||||
"register_enabled": get("register_enabled", "true") == "true",
|
||||
"maintenance_msg": get("site_maintenance_msg", ""),
|
||||
}
|
||||
}
|
||||
|
||||
// IsSiteEnabled 站点是否开启
|
||||
func (s *AdminService) IsSiteEnabled() bool {
|
||||
if v, err := s.settingRepo.GetValue("site_enabled"); err == nil {
|
||||
return v == "true" || v == "1"
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,346 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
"unicode"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"filestoragesystem/internal/config"
|
||||
"filestoragesystem/internal/model"
|
||||
"filestoragesystem/internal/repository"
|
||||
"filestoragesystem/internal/utils"
|
||||
"filestoragesystem/pkg/apperr"
|
||||
)
|
||||
|
||||
const (
|
||||
maxLoginFails = 5 // 最大登录失败次数
|
||||
lockDuration = 15 * time.Minute // 锁定时长
|
||||
signTimeWindow = 5 * time.Minute // 签名时间戳有效窗口
|
||||
)
|
||||
|
||||
// AuthService 认证服务
|
||||
type AuthService struct {
|
||||
cfg *config.Config
|
||||
userRepo *repository.UserRepo
|
||||
apiKeyRepo *repository.APIKeyRepo
|
||||
roleRepo *repository.RoleRepo
|
||||
settingRepo *repository.SettingRepo
|
||||
opLogRepo *repository.OpLogRepo
|
||||
|
||||
mu sync.Mutex
|
||||
loginFails map[string]*loginState // 登录失败状态
|
||||
nonces map[string]time.Time // 防重放nonce
|
||||
}
|
||||
|
||||
type loginState struct {
|
||||
fails int
|
||||
lockedUntil time.Time
|
||||
}
|
||||
|
||||
// ensureState 惰性初始化登录失败与nonce缓存
|
||||
func (s *AuthService) ensureState() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.loginFails == nil {
|
||||
s.loginFails = make(map[string]*loginState)
|
||||
}
|
||||
if s.nonces == nil {
|
||||
s.nonces = make(map[string]time.Time)
|
||||
}
|
||||
}
|
||||
|
||||
// Register 用户注册
|
||||
func (s *AuthService) Register(username, email, password string) (*model.User, error) {
|
||||
// 注册开关
|
||||
if v, err := s.settingRepo.GetValue("register_enabled"); err == nil && v != "true" {
|
||||
return nil, apperr.ErrRegisterClosed
|
||||
}
|
||||
username = strings.TrimSpace(username)
|
||||
email = strings.TrimSpace(strings.ToLower(email))
|
||||
|
||||
if len(username) < 3 || len(username) > 50 {
|
||||
return nil, fmt.Errorf("用户名长度需在3-50之间")
|
||||
}
|
||||
for _, ch := range username {
|
||||
if !unicode.IsLetter(ch) && !unicode.IsDigit(ch) && ch != '_' && ch != '-' && ch < 0x80 {
|
||||
return nil, fmt.Errorf("用户名仅支持字母、数字、下划线和横线")
|
||||
}
|
||||
}
|
||||
if !utils.IsEmail(email) {
|
||||
return nil, fmt.Errorf("邮箱格式不正确")
|
||||
}
|
||||
if !utils.IsValidPassword(password) {
|
||||
return nil, fmt.Errorf("密码至少8位且需包含字母和数字")
|
||||
}
|
||||
if _, err := s.userRepo.FindByUsername(username); err == nil {
|
||||
return nil, apperr.ErrUserExists
|
||||
}
|
||||
if _, err := s.userRepo.FindByEmail(email); err == nil {
|
||||
return nil, apperr.ErrUserExists
|
||||
}
|
||||
|
||||
// 默认角色与配额
|
||||
role, err := s.roleRepo.FindByCode("user")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("默认角色不存在: %w", err)
|
||||
}
|
||||
limit := int64(10737418240)
|
||||
if v, err := s.settingRepo.GetValue("default_storage_limit"); err == nil {
|
||||
var n int64
|
||||
if _, err := fmt.Sscanf(v, "%d", &n); err == nil && n > 0 {
|
||||
limit = n
|
||||
}
|
||||
}
|
||||
|
||||
hash, err := utils.HashPassword(password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u := &model.User{
|
||||
Username: username, Email: email, PasswordHash: hash,
|
||||
RoleID: role.ID, Status: 1, StorageLimit: limit,
|
||||
}
|
||||
if err := s.userRepo.Create(u); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
utils.Info("auth", "新用户注册: %s", username)
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// Login 用户登录,成功返回用户与JWT token
|
||||
func (s *AuthService) Login(account, password, ip string) (*model.User, string, error) {
|
||||
s.ensureState()
|
||||
account = strings.TrimSpace(account)
|
||||
|
||||
// 锁定检查
|
||||
s.mu.Lock()
|
||||
if st, ok := s.loginFails[account]; ok {
|
||||
if time.Now().Before(st.lockedUntil) {
|
||||
s.mu.Unlock()
|
||||
return nil, "", apperr.ErrAccountLocked
|
||||
}
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
u, err := s.userRepo.FindByUsernameOrEmail(account)
|
||||
if err != nil {
|
||||
s.recordLoginFail(account)
|
||||
return nil, "", apperr.ErrPasswordWrong
|
||||
}
|
||||
if u.Status != 1 {
|
||||
return nil, "", apperr.ErrUserDisabled
|
||||
}
|
||||
if !utils.CheckPassword(u.PasswordHash, password) {
|
||||
s.recordLoginFail(account)
|
||||
return nil, "", apperr.ErrPasswordWrong
|
||||
}
|
||||
|
||||
// 登录成功:清除失败计数
|
||||
s.mu.Lock()
|
||||
delete(s.loginFails, account)
|
||||
s.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
_ = s.userRepo.UpdateFields(u.ID, map[string]interface{}{
|
||||
"last_login_at": now, "last_login_ip": ip,
|
||||
})
|
||||
|
||||
token, err := utils.GenerateToken(u.ID, u.Username, roleCodeOf(u))
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
return u, token, nil
|
||||
}
|
||||
|
||||
func (s *AuthService) recordLoginFail(account string) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.loginFails == nil {
|
||||
s.loginFails = make(map[string]*loginState)
|
||||
}
|
||||
st, ok := s.loginFails[account]
|
||||
if !ok {
|
||||
st = &loginState{}
|
||||
s.loginFails[account] = st
|
||||
}
|
||||
st.fails++
|
||||
if st.fails >= maxLoginFails {
|
||||
st.lockedUntil = time.Now().Add(lockDuration)
|
||||
st.fails = 0
|
||||
utils.Warn("auth", "账户 %s 因连续登录失败被锁定%d分钟", account, int(lockDuration.Minutes()))
|
||||
}
|
||||
}
|
||||
|
||||
// roleCodeOf 用户角色编码(内存缓存于User.Role)
|
||||
func roleCodeOf(u *model.User) string {
|
||||
if u.Role != nil {
|
||||
return u.Role.Code
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// CreateAPIKey 为用户生成API密钥,secret仅此一次返回
|
||||
func (s *AuthService) CreateAPIKey(userID uint, name string) (*model.APIKey, string, error) {
|
||||
u, err := s.userRepo.FindByID(userID)
|
||||
if err != nil {
|
||||
return nil, "", apperr.ErrNotFound
|
||||
}
|
||||
key := &model.APIKey{
|
||||
UserID: userID,
|
||||
AccessKey: utils.RandomKey(32),
|
||||
SecretKey: utils.RandomKey(64),
|
||||
Name: name,
|
||||
Status: 1,
|
||||
}
|
||||
if key.Name == "" {
|
||||
key.Name = u.Username + "的密钥"
|
||||
}
|
||||
if err := s.apiKeyRepo.Create(key); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
utils.Info("auth", "用户 %s 创建API密钥: %s", u.Username, key.AccessKey)
|
||||
return key, key.SecretKey, nil
|
||||
}
|
||||
|
||||
// ListAPIKeys 用户密钥列表
|
||||
func (s *AuthService) ListAPIKeys(userID uint) ([]model.APIKey, error) {
|
||||
return s.apiKeyRepo.ListByUser(userID)
|
||||
}
|
||||
|
||||
// DeleteAPIKey 删除密钥
|
||||
func (s *AuthService) DeleteAPIKey(userID, id uint) error {
|
||||
if err := s.apiKeyRepo.Delete(id, userID); err != nil {
|
||||
return apperr.ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// VerifyAPIKeySignature 校验API Key HMAC-SHA256签名,返回所属用户
|
||||
// 签名串: accessKey + "\n" + timestamp + "\n" + nonce + "\n" + method + "\n" + path
|
||||
func (s *AuthService) VerifyAPIKeySignature(accessKey, timestamp, nonce, signature, method, path string) (*model.User, error) {
|
||||
s.ensureState()
|
||||
|
||||
// 时间窗口
|
||||
var ts int64
|
||||
if _, err := fmt.Sscanf(timestamp, "%d", &ts); err != nil {
|
||||
return nil, apperr.ErrSignInvalid
|
||||
}
|
||||
t := time.Unix(ts, 0)
|
||||
if diff := time.Since(t); diff > signTimeWindow || diff < -signTimeWindow {
|
||||
return nil, apperr.ErrSignExpired
|
||||
}
|
||||
|
||||
// 防重放
|
||||
s.mu.Lock()
|
||||
if s.nonces == nil {
|
||||
s.nonces = make(map[string]time.Time)
|
||||
}
|
||||
if _, used := s.nonces[nonce]; used {
|
||||
s.mu.Unlock()
|
||||
return nil, apperr.ErrSignInvalid
|
||||
}
|
||||
s.nonces[nonce] = time.Now()
|
||||
// 顺带清理过期nonce
|
||||
if len(s.nonces) > 10000 {
|
||||
cutoff := time.Now().Add(-signTimeWindow * 2)
|
||||
for k, v := range s.nonces {
|
||||
if v.Before(cutoff) {
|
||||
delete(s.nonces, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
key, err := s.apiKeyRepo.FindByAccessKey(accessKey)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrSignInvalid
|
||||
}
|
||||
if key.Status != 1 {
|
||||
return nil, apperr.ErrForbidden
|
||||
}
|
||||
|
||||
message := strings.Join([]string{accessKey, timestamp, nonce, method, path}, "\n")
|
||||
expect := utils.HMACSHA256(key.SecretKey, message)
|
||||
if !strings.EqualFold(expect, signature) {
|
||||
return nil, apperr.ErrSignInvalid
|
||||
}
|
||||
|
||||
_ = s.apiKeyRepo.Touch(key.ID)
|
||||
|
||||
u, err := s.userRepo.FindByID(key.UserID)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
if u.Status != 1 {
|
||||
return nil, apperr.ErrUserDisabled
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// LoadUserByID 加载用户(含角色权限)
|
||||
func (s *AuthService) LoadUserByID(id uint) (*model.User, error) {
|
||||
u, err := s.userRepo.FindByID(id)
|
||||
if err != nil {
|
||||
if err == gorm.ErrRecordNotFound {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if u.Status != 1 {
|
||||
return nil, apperr.ErrUserDisabled
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// PermissionCodesOf 用户拥有的权限编码集合
|
||||
func PermissionCodesOf(u *model.User) map[string]bool {
|
||||
codes := make(map[string]bool)
|
||||
if u == nil || u.Role == nil {
|
||||
return codes
|
||||
}
|
||||
if u.Role.Code == "super_admin" {
|
||||
codes["*"] = true
|
||||
return codes
|
||||
}
|
||||
for _, p := range u.Role.Permissions {
|
||||
codes[p.Code] = true
|
||||
}
|
||||
return codes
|
||||
}
|
||||
|
||||
// BootstrapAdmin 按配置确保管理员账户存在
|
||||
func (s *AuthService) BootstrapAdmin() error {
|
||||
ac := s.cfg.Admin
|
||||
if ac.Username == "" {
|
||||
return nil
|
||||
}
|
||||
u, err := s.userRepo.FindByUsername(ac.Username)
|
||||
if err == nil {
|
||||
// 已存在:确保其为配置的角色
|
||||
if role, rerr := s.roleRepo.FindByCode(ac.Role); rerr == nil && u.RoleID != role.ID {
|
||||
_ = s.userRepo.UpdateFields(u.ID, map[string]interface{}{"role_id": role.ID, "status": 1})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
role, err := s.roleRepo.FindByCode(ac.Role)
|
||||
if err != nil {
|
||||
return fmt.Errorf("配置的管理员角色 %s 不存在: %w", ac.Role, err)
|
||||
}
|
||||
hash, err := utils.HashPassword(ac.Password)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
admin := &model.User{
|
||||
Username: ac.Username, Email: ac.Email, PasswordHash: hash,
|
||||
RoleID: role.ID, Status: 1, StorageLimit: 1 << 40,
|
||||
}
|
||||
if err := s.userRepo.Create(admin); err != nil {
|
||||
return err
|
||||
}
|
||||
utils.Info("auth", "已创建管理员账户: %s", ac.Username)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,423 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/md5"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"mime/multipart"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"filestoragesystem/internal/config"
|
||||
"filestoragesystem/internal/model"
|
||||
"filestoragesystem/internal/repository"
|
||||
"filestoragesystem/internal/utils"
|
||||
"filestoragesystem/pkg/apperr"
|
||||
)
|
||||
|
||||
// FileService 文件服务
|
||||
type FileService struct {
|
||||
cfg *config.Config
|
||||
fileRepo *repository.FileRepo
|
||||
projectRepo *repository.ProjectRepo
|
||||
userRepo *repository.UserRepo
|
||||
tempLinkRepo *repository.TempLinkRepo
|
||||
trafficRepo *repository.TrafficRepo
|
||||
settingRepo *repository.SettingRepo
|
||||
opLogRepo *repository.OpLogRepo
|
||||
webhookSvc *WebhookService
|
||||
}
|
||||
|
||||
// settingInt 读取数字型设置
|
||||
func (s *FileService) settingInt(key string, def int64) int64 {
|
||||
if v, err := s.settingRepo.GetValue(key); err == nil {
|
||||
var n int64
|
||||
if _, e := fmt.Sscanf(v, "%d", &n); e == nil && n > 0 {
|
||||
return n
|
||||
}
|
||||
}
|
||||
return def
|
||||
}
|
||||
|
||||
// settingBool 读取布尔型设置
|
||||
func (s *FileService) settingBool(key string, def bool) bool {
|
||||
if v, err := s.settingRepo.GetValue(key); err == nil {
|
||||
return v == "true" || v == "1"
|
||||
}
|
||||
return def
|
||||
}
|
||||
|
||||
// Upload 上传文件:校验配额/类型/大小,MD5去重,落盘,记录统计
|
||||
func (s *FileService) Upload(c *gin.Context, userID uint, fh *multipart.FileHeader) (*model.File, error) {
|
||||
projectID, _ := strconvUint(c.PostForm("project_id"))
|
||||
if projectID == 0 {
|
||||
return nil, fmt.Errorf("缺少项目ID")
|
||||
}
|
||||
|
||||
project, err := s.projectRepo.FindByID(projectID)
|
||||
if err != nil || project.UserID != userID {
|
||||
return nil, apperr.ErrForbidden
|
||||
}
|
||||
if project.Status != 1 {
|
||||
return nil, fmt.Errorf("项目已被禁用")
|
||||
}
|
||||
|
||||
src, err := fh.Open()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer src.Close()
|
||||
|
||||
size := fh.Size
|
||||
// 大小限制
|
||||
maxSize := s.settingInt("max_file_size", 104857600)
|
||||
if size > maxSize {
|
||||
return nil, apperr.ErrFileTooLarge
|
||||
}
|
||||
// 类型限制
|
||||
allowed := "*"
|
||||
if v, err := s.settingRepo.GetValue("allowed_file_types"); err == nil {
|
||||
allowed = v
|
||||
}
|
||||
if !s.checkTypeAllowed(fh.Filename, allowed) {
|
||||
return nil, apperr.ErrFileTypeDenied
|
||||
}
|
||||
// 配额
|
||||
if project.StorageUsed+size > project.StorageLimit {
|
||||
return nil, apperr.ErrQuotaExceeded
|
||||
}
|
||||
user, err := s.userRepo.FindByID(userID)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
if user.StorageUsed+size > user.StorageLimit {
|
||||
return nil, apperr.ErrQuotaExceeded
|
||||
}
|
||||
|
||||
// 虚拟路径 + 文件名
|
||||
virtualPath := utils.SanitizePath(c.PostForm("path"))
|
||||
baseName := filepath.Base(utils.SanitizePath(fh.Filename))
|
||||
if baseName == "" || baseName == "." || baseName == "/" {
|
||||
baseName = "unnamed"
|
||||
}
|
||||
filename := baseName
|
||||
if virtualPath != "" {
|
||||
filename = virtualPath + "/" + baseName
|
||||
}
|
||||
visibility := int8(1)
|
||||
if c.PostForm("visibility") == "2" || c.PostForm("visibility") == "public" {
|
||||
visibility = 2
|
||||
}
|
||||
|
||||
// 写入临时文件并计算MD5
|
||||
tmpFile, err := os.CreateTemp("", "fss-upload-*")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tmpPath := tmpFile.Name()
|
||||
defer os.Remove(tmpPath)
|
||||
|
||||
hasher := md5.New()
|
||||
if _, err := io.Copy(io.MultiWriter(tmpFile, hasher), src); err != nil {
|
||||
tmpFile.Close()
|
||||
return nil, err
|
||||
}
|
||||
tmpFile.Close()
|
||||
md5sum := hex.EncodeToString(hasher.Sum(nil))
|
||||
|
||||
// 去重:相同MD5复用已有物理文件
|
||||
storedRel := ""
|
||||
if exist, err := s.fileRepo.FindByMD5AndPath(md5sum); err == nil && exist.StoredPath != "" {
|
||||
if _, serr := os.Stat(filepath.Join(s.cfg.Storage.Root, exist.StoredPath)); serr == nil {
|
||||
storedRel = exist.StoredPath
|
||||
}
|
||||
}
|
||||
|
||||
// 无可复用文件则保存到 blob/<md5前2位>/<md5>
|
||||
if storedRel == "" {
|
||||
blobRel := filepath.ToSlash(filepath.Join("blobs", md5sum[:2], md5sum))
|
||||
blobAbs := filepath.Join(s.cfg.Storage.Root, blobRel)
|
||||
if err := os.MkdirAll(filepath.Dir(blobAbs), 0o755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.Rename(tmpPath, blobAbs); err != nil {
|
||||
// 跨盘时回退为复制
|
||||
if cerr := copyFile(tmpPath, blobAbs); cerr != nil {
|
||||
return nil, cerr
|
||||
}
|
||||
}
|
||||
storedRel = blobRel
|
||||
}
|
||||
|
||||
mimeType := fh.Header.Get("Content-Type")
|
||||
if mimeType == "" || mimeType == "application/octet-stream" {
|
||||
if t := mime.TypeByExtension(strings.ToLower(filepath.Ext(baseName))); t != "" {
|
||||
mimeType = t
|
||||
}
|
||||
}
|
||||
|
||||
f := &model.File{
|
||||
ProjectID: projectID,
|
||||
UserID: userID,
|
||||
Filename: filename,
|
||||
StoredPath: storedRel,
|
||||
Size: size,
|
||||
MimeType: mimeType,
|
||||
MD5: md5sum,
|
||||
Visibility: visibility,
|
||||
Status: 1,
|
||||
}
|
||||
if err := s.fileRepo.Create(f); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 更新配额、流量、webhook
|
||||
_ = s.projectRepo.UpdateStorageUsed(projectID, size)
|
||||
_ = s.userRepo.UpdateStorageUsed(userID, size)
|
||||
_ = s.trafficRepo.Create(&model.TrafficLog{
|
||||
UserID: userID, ProjectID: projectID, FileID: f.ID,
|
||||
Type: 1, Size: size, IP: utils.ClientIP(c), CreatedAt: time.Now(),
|
||||
})
|
||||
s.webhookSvc.Dispatch(userID, projectID, "file.upload", map[string]interface{}{
|
||||
"file_id": f.ID, "filename": f.Filename, "size": f.Size, "md5": f.MD5, "project_id": projectID,
|
||||
})
|
||||
utils.Info("storage", "用户#%d 上传文件 %s (%d字节)", userID, f.Filename, f.Size)
|
||||
return f, nil
|
||||
}
|
||||
|
||||
func (s *FileService) checkTypeAllowed(filename, allowed string) bool {
|
||||
allowed = strings.TrimSpace(allowed)
|
||||
if allowed == "" || allowed == "*" {
|
||||
return true
|
||||
}
|
||||
ext := strings.ToLower(strings.TrimPrefix(filepath.Ext(filename), "."))
|
||||
if ext == "" {
|
||||
return true // 无扩展名不限制
|
||||
}
|
||||
for _, t := range strings.Split(allowed, ",") {
|
||||
t = strings.ToLower(strings.TrimSpace(strings.TrimPrefix(t, ".")))
|
||||
if t == "*" || t == ext {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// GetInfo 获取文件信息(所有者或管理员)
|
||||
func (s *FileService) GetInfo(userID uint, id uint, isAdmin bool) (*model.File, error) {
|
||||
f, err := s.fileRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
if !isAdmin && f.UserID != userID {
|
||||
return nil, apperr.ErrForbidden
|
||||
}
|
||||
return f, nil
|
||||
}
|
||||
|
||||
// List 文件列表
|
||||
func (s *FileService) List(userID uint, isAdmin bool, filter repository.FileListFilter) ([]model.File, int64, error) {
|
||||
if !isAdmin {
|
||||
filter.UserID = userID
|
||||
}
|
||||
if filter.Page <= 0 {
|
||||
filter.Page = 1
|
||||
}
|
||||
if filter.PageSize <= 0 {
|
||||
filter.PageSize = 20
|
||||
}
|
||||
return s.fileRepo.List(filter)
|
||||
}
|
||||
|
||||
// Open 打开文件物理句柄并校验访问权,记录流量
|
||||
func (s *FileService) Open(c *gin.Context, userID uint, id uint, isAdmin bool) (*model.File, *os.File, error) {
|
||||
f, err := s.GetInfo(userID, id, isAdmin)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if f.Status != 1 {
|
||||
return nil, nil, apperr.ErrFileDeleted
|
||||
}
|
||||
abs := filepath.Join(s.cfg.Storage.Root, f.StoredPath)
|
||||
fp, err := os.Open(abs)
|
||||
if err != nil {
|
||||
utils.Error("storage", "打开文件失败 path=%s err=%v", abs, err)
|
||||
return nil, nil, fmt.Errorf("文件读取失败")
|
||||
}
|
||||
return f, fp, nil
|
||||
}
|
||||
|
||||
// RecordDownload 记录下载行为
|
||||
func (s *FileService) RecordDownload(c *gin.Context, f *model.File) {
|
||||
_ = s.fileRepo.IncrDownloadCount(f.ID)
|
||||
_ = s.trafficRepo.Create(&model.TrafficLog{
|
||||
UserID: f.UserID, ProjectID: f.ProjectID, FileID: f.ID,
|
||||
Type: 2, Size: f.Size, IP: utils.ClientIP(c), CreatedAt: time.Now(),
|
||||
})
|
||||
}
|
||||
|
||||
// Delete 删除文件;permanent=true时硬删除
|
||||
func (s *FileService) Delete(userID uint, id uint, isAdmin bool, permanent bool) error {
|
||||
f, err := s.GetInfo(userID, id, isAdmin)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if !permanent {
|
||||
if f.Status == 0 {
|
||||
return apperr.ErrFileDeleted
|
||||
}
|
||||
if err := s.fileRepo.SoftDelete(id); err != nil {
|
||||
return err
|
||||
}
|
||||
// 回收站不占配额
|
||||
_ = s.projectRepo.UpdateStorageUsed(f.ProjectID, -f.Size)
|
||||
_ = s.userRepo.UpdateStorageUsed(f.UserID, -f.Size)
|
||||
return nil
|
||||
}
|
||||
|
||||
// 硬删除
|
||||
if err := s.fileRepo.Delete(id); err != nil {
|
||||
return err
|
||||
}
|
||||
_ = s.tempLinkRepo.DeleteByFile(id)
|
||||
// 若处于回收站,配额已扣减过
|
||||
if f.Status == 1 {
|
||||
_ = s.projectRepo.UpdateStorageUsed(f.ProjectID, -f.Size)
|
||||
_ = s.userRepo.UpdateStorageUsed(f.UserID, -f.Size)
|
||||
}
|
||||
// 物理文件无其他引用时删除
|
||||
if n, err := s.fileRepo.CountByStoredPath(f.StoredPath); err == nil && n == 0 {
|
||||
_ = os.Remove(filepath.Join(s.cfg.Storage.Root, f.StoredPath))
|
||||
}
|
||||
s.webhookSvc.Dispatch(f.UserID, f.ProjectID, "file.delete", map[string]interface{}{
|
||||
"file_id": f.ID, "filename": f.Filename, "permanent": permanent,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
// Restore 从回收站恢复文件
|
||||
func (s *FileService) Restore(userID uint, id uint, isAdmin bool) error {
|
||||
f, err := s.GetInfo(userID, id, isAdmin)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if f.Status != 0 {
|
||||
return fmt.Errorf("文件不在回收站中")
|
||||
}
|
||||
// 恢复前检查配额
|
||||
project, err := s.projectRepo.FindByID(f.ProjectID)
|
||||
if err == nil {
|
||||
if project.StorageUsed+f.Size > project.StorageLimit {
|
||||
return apperr.ErrQuotaExceeded
|
||||
}
|
||||
}
|
||||
user, err := s.userRepo.FindByID(f.UserID)
|
||||
if err == nil && user.StorageUsed+f.Size > user.StorageLimit {
|
||||
return apperr.ErrQuotaExceeded
|
||||
}
|
||||
if err := s.fileRepo.Restore(id); err != nil {
|
||||
return err
|
||||
}
|
||||
_ = s.projectRepo.UpdateStorageUsed(f.ProjectID, f.Size)
|
||||
_ = s.userRepo.UpdateStorageUsed(f.UserID, f.Size)
|
||||
return nil
|
||||
}
|
||||
|
||||
// CreateTempLink 生成临时访问链接
|
||||
func (s *FileService) CreateTempLink(userID uint, id uint, isAdmin bool, password string, maxCount int, expiresInSeconds int64) (*model.TempLink, error) {
|
||||
f, err := s.GetInfo(userID, id, isAdmin)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if f.Status != 1 {
|
||||
return nil, apperr.ErrFileDeleted
|
||||
}
|
||||
if expiresInSeconds <= 0 {
|
||||
expiresInSeconds = s.settingInt("temp_link_max_age", 3600)
|
||||
}
|
||||
if expiresInSeconds > 7*24*3600 {
|
||||
expiresInSeconds = 7 * 24 * 3600
|
||||
}
|
||||
link := &model.TempLink{
|
||||
FileID: f.ID,
|
||||
UserID: userID,
|
||||
Token: utils.RandomKey(48),
|
||||
MaxCount: maxCount,
|
||||
ExpiresAt: time.Now().Add(time.Duration(expiresInSeconds) * time.Second),
|
||||
}
|
||||
if password != "" {
|
||||
link.Password = password
|
||||
}
|
||||
if err := s.tempLinkRepo.Create(link); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
link.HasPwd = link.Password != ""
|
||||
return link, nil
|
||||
}
|
||||
|
||||
// ResolveTempLink 通过token解析临时链接(校验有效期/次数/密码),返回文件与物理句柄
|
||||
func (s *FileService) ResolveTempLink(token, password string) (*model.File, *os.File, error) {
|
||||
link, err := s.tempLinkRepo.FindByToken(token)
|
||||
if err != nil {
|
||||
return nil, nil, apperr.ErrLinkExpired
|
||||
}
|
||||
if time.Now().After(link.ExpiresAt) {
|
||||
return nil, nil, apperr.ErrLinkExpired
|
||||
}
|
||||
if link.MaxCount > 0 && link.UsedCount >= link.MaxCount {
|
||||
return nil, nil, apperr.ErrLinkExpired
|
||||
}
|
||||
if link.Password != "" && link.Password != password {
|
||||
return nil, nil, apperr.ErrLinkPassword
|
||||
}
|
||||
f, err := s.fileRepo.FindByID(link.FileID)
|
||||
if err != nil || f.Status != 1 {
|
||||
return nil, nil, apperr.ErrNotFound
|
||||
}
|
||||
fp, err := os.Open(filepath.Join(s.cfg.Storage.Root, f.StoredPath))
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("文件读取失败")
|
||||
}
|
||||
_ = s.tempLinkRepo.IncrUsed(link.ID)
|
||||
return f, fp, nil
|
||||
}
|
||||
|
||||
// ---- 辅助函数 ----
|
||||
|
||||
func copyFile(src, dst string) error {
|
||||
in, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer in.Close()
|
||||
out, err := os.Create(dst)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer out.Close()
|
||||
if _, err := io.Copy(out, in); err != nil {
|
||||
return err
|
||||
}
|
||||
return out.Sync()
|
||||
}
|
||||
|
||||
func strconvUint(s string) (uint, error) {
|
||||
var n uint
|
||||
if s == "" {
|
||||
return 0, fmt.Errorf("empty")
|
||||
}
|
||||
for _, ch := range s {
|
||||
if ch < '0' || ch > '9' {
|
||||
return 0, fmt.Errorf("invalid")
|
||||
}
|
||||
n = n*10 + uint(ch-'0')
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"filestoragesystem/internal/config"
|
||||
"filestoragesystem/internal/model"
|
||||
"filestoragesystem/internal/repository"
|
||||
"filestoragesystem/pkg/apperr"
|
||||
)
|
||||
|
||||
// ProjectService 项目服务
|
||||
type ProjectService struct {
|
||||
cfg *config.Config
|
||||
projectRepo *repository.ProjectRepo
|
||||
fileRepo *repository.FileRepo
|
||||
settingRepo *repository.SettingRepo
|
||||
opLogRepo *repository.OpLogRepo
|
||||
webhookSvc *WebhookService
|
||||
}
|
||||
|
||||
// Create 创建项目
|
||||
func (s *ProjectService) Create(userID uint, name, description string, storageLimit int64) (*model.Project, error) {
|
||||
if name == "" || len(name) > 100 {
|
||||
return nil, fmt.Errorf("项目名称不能为空且不超过100字符")
|
||||
}
|
||||
if _, err := s.projectRepo.FindByName(userID, name); err == nil {
|
||||
return nil, apperr.ErrProjectExists
|
||||
}
|
||||
if storageLimit <= 0 {
|
||||
storageLimit = 5368709120 // 5GB
|
||||
if v, err := s.settingRepo.GetValue("default_project_limit"); err == nil {
|
||||
var n int64
|
||||
if _, e := fmt.Sscanf(v, "%d", &n); e == nil && n > 0 {
|
||||
storageLimit = n
|
||||
}
|
||||
}
|
||||
}
|
||||
p := &model.Project{
|
||||
UserID: userID, Name: name, Description: description,
|
||||
StorageLimit: storageLimit, Status: 1,
|
||||
}
|
||||
if err := s.projectRepo.Create(p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
// List 项目列表,all=true时管理员查看全部
|
||||
func (s *ProjectService) List(userID uint, keyword string, all bool, page, pageSize int) ([]model.Project, int64, error) {
|
||||
return s.projectRepo.List(userID, keyword, all, page, pageSize)
|
||||
}
|
||||
|
||||
// Get 项目详情(所有者或管理员)
|
||||
func (s *ProjectService) Get(userID uint, id uint, isAdmin bool) (*model.Project, error) {
|
||||
p, err := s.projectRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
if !isAdmin && p.UserID != userID {
|
||||
return nil, apperr.ErrForbidden
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
// Update 更新项目
|
||||
func (s *ProjectService) Update(userID uint, id uint, isAdmin bool, name, description string, storageLimit int64) (*model.Project, error) {
|
||||
p, err := s.Get(userID, id, isAdmin)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if name != "" && name != p.Name {
|
||||
if _, err := s.projectRepo.FindByName(userID, name); err == nil {
|
||||
return nil, apperr.ErrProjectExists
|
||||
}
|
||||
p.Name = name
|
||||
}
|
||||
if description != "" {
|
||||
p.Description = description
|
||||
}
|
||||
if storageLimit > 0 {
|
||||
p.StorageLimit = storageLimit
|
||||
}
|
||||
if err := s.projectRepo.Update(p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
// Delete 删除项目及其全部文件(硬删除)
|
||||
func (s *ProjectService) Delete(userID uint, id uint, isAdmin bool) error {
|
||||
if _, err := s.Get(userID, id, isAdmin); err != nil {
|
||||
return err
|
||||
}
|
||||
// 删除物理文件
|
||||
files, _ := s.fileRepo.ListByProject(id)
|
||||
seen := make(map[string]bool)
|
||||
for _, f := range files {
|
||||
if f.StoredPath != "" && !seen[f.StoredPath] {
|
||||
seen[f.StoredPath] = true
|
||||
_ = os.Remove(filepath.Join(s.cfg.Storage.Root, f.StoredPath))
|
||||
}
|
||||
}
|
||||
_ = s.fileRepo.DB.Exec("DELETE FROM files WHERE project_id = ?", id).Error
|
||||
_ = s.fileRepo.DB.Exec("DELETE FROM temp_links WHERE file_id IN (SELECT id FROM files WHERE project_id = ?)", id).Error
|
||||
|
||||
return s.projectRepo.Delete(id)
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"filestoragesystem/internal/model"
|
||||
"filestoragesystem/internal/repository"
|
||||
"filestoragesystem/pkg/apperr"
|
||||
)
|
||||
|
||||
// RoleService 角色权限服务
|
||||
type RoleService struct {
|
||||
roleRepo *repository.RoleRepo
|
||||
permRepo *repository.PermissionRepo
|
||||
}
|
||||
|
||||
// ListRoles 角色列表
|
||||
func (s *RoleService) ListRoles() ([]model.Role, error) {
|
||||
return s.roleRepo.List()
|
||||
}
|
||||
|
||||
// GetRole 角色详情(含权限)
|
||||
func (s *RoleService) GetRole(id uint) (*model.Role, error) {
|
||||
r, err := s.roleRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
// CreateRole 创建角色
|
||||
func (s *RoleService) CreateRole(name, code, description string, permIDs []uint) (*model.Role, error) {
|
||||
if name == "" || code == "" {
|
||||
return nil, fmt.Errorf("角色名称和编码不能为空")
|
||||
}
|
||||
if _, err := s.roleRepo.FindByCode(code); err == nil {
|
||||
return nil, fmt.Errorf("角色编码已存在")
|
||||
}
|
||||
role := &model.Role{
|
||||
Name: name, Code: code, Description: description,
|
||||
IsSystem: false, Status: 1,
|
||||
Permissions: s.mustPerms(permIDs),
|
||||
}
|
||||
if err := s.roleRepo.Create(role); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return role, nil
|
||||
}
|
||||
|
||||
// UpdateRole 编辑角色(系统角色仅允许改描述与权限)
|
||||
func (s *RoleService) UpdateRole(id uint, name, description string, permIDs []uint) (*model.Role, error) {
|
||||
role, err := s.roleRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
if role.Code == "super_admin" {
|
||||
return nil, apperr.ErrSystemRole
|
||||
}
|
||||
if name != "" && !role.IsSystem {
|
||||
role.Name = name
|
||||
}
|
||||
if description != "" {
|
||||
role.Description = description
|
||||
}
|
||||
if err := s.roleRepo.Update(role); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if permIDs != nil {
|
||||
if err := s.roleRepo.ReplacePermissions(id, permIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return s.roleRepo.FindByID(id)
|
||||
}
|
||||
|
||||
// DeleteRole 删除角色(系统角色不可删,有用户的角色不可删)
|
||||
func (s *RoleService) DeleteRole(id uint) error {
|
||||
role, err := s.roleRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return apperr.ErrNotFound
|
||||
}
|
||||
if role.IsSystem {
|
||||
return apperr.ErrSystemRole
|
||||
}
|
||||
if n, _ := s.roleRepo.CountUsersByRole(id); n > 0 {
|
||||
return fmt.Errorf("该角色下仍有 %d 个用户,无法删除", n)
|
||||
}
|
||||
return s.roleRepo.Delete(id)
|
||||
}
|
||||
|
||||
// SetPermissions 配置角色权限
|
||||
func (s *RoleService) SetPermissions(id uint, permIDs []uint) error {
|
||||
role, err := s.roleRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return apperr.ErrNotFound
|
||||
}
|
||||
if role.Code == "super_admin" {
|
||||
return apperr.ErrSystemRole
|
||||
}
|
||||
return s.roleRepo.ReplacePermissions(id, permIDs)
|
||||
}
|
||||
|
||||
// ListPermissions 全部权限
|
||||
func (s *RoleService) ListPermissions() ([]model.Permission, error) {
|
||||
return s.permRepo.List()
|
||||
}
|
||||
|
||||
// PermissionsGrouped 按模块分组权限
|
||||
func (s *RoleService) PermissionsGrouped() (map[string][]model.Permission, []string, error) {
|
||||
perms, err := s.permRepo.List()
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
groups := make(map[string][]model.Permission)
|
||||
var order []string
|
||||
seen := make(map[string]bool)
|
||||
for _, p := range perms {
|
||||
if !seen[p.Module] {
|
||||
seen[p.Module] = true
|
||||
order = append(order, p.Module)
|
||||
}
|
||||
groups[p.Module] = append(groups[p.Module], p)
|
||||
}
|
||||
return groups, order, nil
|
||||
}
|
||||
|
||||
func (s *RoleService) mustPerms(ids []uint) []model.Permission {
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
perms, err := s.permRepo.ListByIDs(ids)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return perms
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"filestoragesystem/internal/config"
|
||||
"filestoragesystem/internal/model"
|
||||
"filestoragesystem/internal/repository"
|
||||
)
|
||||
|
||||
// OpLogContext 操作日志上下文(由controller从请求中提取)
|
||||
type OpLogContext struct {
|
||||
UserID *uint
|
||||
Username string
|
||||
IP string
|
||||
UserAgent string
|
||||
Path string
|
||||
Method string
|
||||
}
|
||||
|
||||
// Services 服务容器
|
||||
type Services struct {
|
||||
cfg *config.Config
|
||||
db *gorm.DB
|
||||
Auth *AuthService
|
||||
User *UserService
|
||||
Project *ProjectService
|
||||
File *FileService
|
||||
Role *RoleService
|
||||
Admin *AdminService
|
||||
Stats *StatsService
|
||||
Webhook *WebhookService
|
||||
opLogRepo *repository.OpLogRepo
|
||||
}
|
||||
|
||||
// NewServices 创建服务容器并初始化各服务
|
||||
func NewServices(cfg *config.Config, db *gorm.DB) *Services {
|
||||
userRepo := &repository.UserRepo{DB: db}
|
||||
roleRepo := &repository.RoleRepo{DB: db}
|
||||
permRepo := &repository.PermissionRepo{DB: db}
|
||||
projectRepo := &repository.ProjectRepo{DB: db}
|
||||
fileRepo := &repository.FileRepo{DB: db}
|
||||
tempLinkRepo := &repository.TempLinkRepo{DB: db}
|
||||
trafficRepo := &repository.TrafficRepo{DB: db}
|
||||
settingRepo := &repository.SettingRepo{DB: db}
|
||||
opLogRepo := &repository.OpLogRepo{DB: db}
|
||||
sysLogRepo := &repository.SysLogRepo{DB: db}
|
||||
webhookRepo := &repository.WebhookRepo{DB: db}
|
||||
apiKeyRepo := &repository.APIKeyRepo{DB: db}
|
||||
|
||||
s := &Services{cfg: cfg, db: db, opLogRepo: opLogRepo}
|
||||
|
||||
s.Auth = &AuthService{
|
||||
cfg: cfg, userRepo: userRepo, apiKeyRepo: apiKeyRepo,
|
||||
roleRepo: roleRepo, settingRepo: settingRepo, opLogRepo: opLogRepo,
|
||||
}
|
||||
s.User = &UserService{
|
||||
cfg: cfg, userRepo: userRepo, roleRepo: roleRepo,
|
||||
settingRepo: settingRepo, opLogRepo: opLogRepo, fileRepo: fileRepo, projectRepo: projectRepo,
|
||||
}
|
||||
s.Project = &ProjectService{
|
||||
cfg: cfg, projectRepo: projectRepo, fileRepo: fileRepo,
|
||||
settingRepo: settingRepo, opLogRepo: opLogRepo, webhookSvc: nil,
|
||||
}
|
||||
s.Webhook = &WebhookService{cfg: cfg, webhookRepo: webhookRepo, settingRepo: settingRepo}
|
||||
s.File = &FileService{
|
||||
cfg: cfg, fileRepo: fileRepo, projectRepo: projectRepo, userRepo: userRepo,
|
||||
tempLinkRepo: tempLinkRepo, trafficRepo: trafficRepo,
|
||||
settingRepo: settingRepo, opLogRepo: opLogRepo, webhookSvc: s.Webhook,
|
||||
}
|
||||
s.Role = &RoleService{roleRepo: roleRepo, permRepo: permRepo}
|
||||
s.Admin = &AdminService{
|
||||
cfg: cfg, db: db, settingRepo: settingRepo,
|
||||
opLogRepo: opLogRepo, sysLogRepo: sysLogRepo,
|
||||
}
|
||||
s.Stats = &StatsService{
|
||||
cfg: cfg, userRepo: userRepo, projectRepo: projectRepo,
|
||||
fileRepo: fileRepo, trafficRepo: trafficRepo,
|
||||
}
|
||||
s.Project.webhookSvc = s.Webhook
|
||||
return s
|
||||
}
|
||||
|
||||
// RecordOpLog 记录操作日志(异步容忍失败)
|
||||
func (s *Services) RecordOpLog(ctx *OpLogContext, action, resourceType string, resourceID uint, resourceName string, success bool, errMsg string) {
|
||||
log := model.OperationLog{
|
||||
Action: action,
|
||||
ResourceType: resourceType,
|
||||
ResourceID: resourceID,
|
||||
ResourceName: resourceName,
|
||||
IP: ctx.IP,
|
||||
UserAgent: truncate(ctx.UserAgent, 500),
|
||||
RequestPath: truncate(ctx.Path, 500),
|
||||
RequestMethod: ctx.Method,
|
||||
Status: 1,
|
||||
ErrorMessage: errMsg,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if ctx.UserID != nil {
|
||||
log.UserID = ctx.UserID
|
||||
}
|
||||
log.Username = ctx.Username
|
||||
if !success {
|
||||
log.Status = 0
|
||||
}
|
||||
_ = s.opLogRepo.Create(&log)
|
||||
}
|
||||
|
||||
func truncate(s string, n int) string {
|
||||
if len(s) <= n {
|
||||
return s
|
||||
}
|
||||
return s[:n]
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"filestoragesystem/internal/config"
|
||||
"filestoragesystem/internal/repository"
|
||||
"filestoragesystem/pkg/apperr"
|
||||
)
|
||||
|
||||
// StatsService 存储与流量统计服务
|
||||
type StatsService struct {
|
||||
cfg *config.Config
|
||||
userRepo *repository.UserRepo
|
||||
projectRepo *repository.ProjectRepo
|
||||
fileRepo *repository.FileRepo
|
||||
trafficRepo *repository.TrafficRepo
|
||||
}
|
||||
|
||||
// UserStats 用户存储统计:总体 + 项目明细 + 流量汇总
|
||||
func (s *StatsService) UserStats(userID uint) (map[string]interface{}, error) {
|
||||
user, err := s.userRepo.FindByID(userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
fileCount, totalSize, err := s.fileRepo.StatsByUser(userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
projects, err := s.projectRepo.ListByUser(userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
projectStats := make([]map[string]interface{}, 0, len(projects))
|
||||
for _, p := range projects {
|
||||
count, size, err := s.fileRepo.StatsByProject(p.ID)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
projectStats = append(projectStats, map[string]interface{}{
|
||||
"project_id": p.ID,
|
||||
"name": p.Name,
|
||||
"file_count": count,
|
||||
"storage_used": size,
|
||||
"storage_limit": p.StorageLimit,
|
||||
})
|
||||
}
|
||||
|
||||
upload30, download30, _ := s.trafficRepo.SumByUser(userID, 30)
|
||||
uploadAll, downloadAll, _ := s.trafficRepo.SumByUser(userID, 0)
|
||||
|
||||
return map[string]interface{}{
|
||||
"storage_used": user.StorageUsed,
|
||||
"storage_limit": user.StorageLimit,
|
||||
"file_count": fileCount,
|
||||
"total_size": totalSize,
|
||||
"project_count": len(projects),
|
||||
"projects": projectStats,
|
||||
"traffic": map[string]interface{}{
|
||||
"upload_30d": upload30,
|
||||
"download_30d": download30,
|
||||
"upload_total": uploadAll,
|
||||
"download_total": downloadAll,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ProjectStats 项目存储统计
|
||||
func (s *StatsService) ProjectStats(userID uint, projectID uint, isAdmin bool) (map[string]interface{}, error) {
|
||||
p, err := s.projectRepo.FindByID(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !isAdmin && p.UserID != userID {
|
||||
return nil, apperr.ErrForbidden
|
||||
}
|
||||
count, size, err := s.fileRepo.StatsByProject(projectID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var trashed int64
|
||||
s.fileRepo.DB.Table("files").Where("project_id = ? AND status = 0", projectID).Count(&trashed)
|
||||
|
||||
return map[string]interface{}{
|
||||
"project_id": p.ID,
|
||||
"name": p.Name,
|
||||
"file_count": count,
|
||||
"storage_used": size,
|
||||
"storage_limit": p.StorageLimit,
|
||||
"trash_count": trashed,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// TrafficStats 流量统计:按天聚合 + 汇总
|
||||
func (s *StatsService) TrafficStats(userID uint, days int) (map[string]interface{}, error) {
|
||||
if days <= 0 {
|
||||
days = 30
|
||||
}
|
||||
daily, err := s.trafficRepo.DailyStats(userID, days)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upload, download, _ := s.trafficRepo.SumByUser(userID, days)
|
||||
return map[string]interface{}{
|
||||
"days": days,
|
||||
"daily": daily,
|
||||
"upload": upload,
|
||||
"download": download,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,258 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"filestoragesystem/internal/config"
|
||||
"filestoragesystem/internal/model"
|
||||
"filestoragesystem/internal/repository"
|
||||
"filestoragesystem/internal/utils"
|
||||
"filestoragesystem/pkg/apperr"
|
||||
)
|
||||
|
||||
// UserService 用户管理服务(个人资料 + 管理员用户管理)
|
||||
type UserService struct {
|
||||
cfg *config.Config
|
||||
userRepo *repository.UserRepo
|
||||
roleRepo *repository.RoleRepo
|
||||
settingRepo *repository.SettingRepo
|
||||
opLogRepo *repository.OpLogRepo
|
||||
fileRepo *repository.FileRepo
|
||||
projectRepo *repository.ProjectRepo
|
||||
}
|
||||
|
||||
// GetProfile 获取个人资料
|
||||
func (s *UserService) GetProfile(userID uint) (*model.User, error) {
|
||||
u, err := s.userRepo.FindByID(userID)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// UpdateProfile 更新个人资料(邮箱)
|
||||
func (s *UserService) UpdateProfile(userID uint, email string) error {
|
||||
email = strings.TrimSpace(strings.ToLower(email))
|
||||
if email != "" && !utils.IsEmail(email) {
|
||||
return fmt.Errorf("邮箱格式不正确")
|
||||
}
|
||||
if email != "" {
|
||||
if other, err := s.userRepo.FindByEmail(email); err == nil && other.ID != userID {
|
||||
return apperr.ErrUserExists
|
||||
}
|
||||
return s.userRepo.UpdateFields(userID, map[string]interface{}{"email": email})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ChangePassword 修改自己的密码
|
||||
func (s *UserService) ChangePassword(userID uint, oldPwd, newPwd string) error {
|
||||
if !utils.IsValidPassword(newPwd) {
|
||||
return fmt.Errorf("新密码至少8位且需包含字母和数字")
|
||||
}
|
||||
u, err := s.userRepo.FindByID(userID)
|
||||
if err != nil {
|
||||
return apperr.ErrNotFound
|
||||
}
|
||||
if !utils.CheckPassword(u.PasswordHash, oldPwd) {
|
||||
return fmt.Errorf("原密码错误")
|
||||
}
|
||||
hash, err := utils.HashPassword(newPwd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.userRepo.UpdateFields(userID, map[string]interface{}{"password_hash": hash})
|
||||
}
|
||||
|
||||
// ---- 管理员操作 ----
|
||||
|
||||
// AdminListUsers 用户列表
|
||||
func (s *UserService) AdminListUsers(keyword, roleID, status string, page, pageSize int) ([]model.User, int64, error) {
|
||||
return s.userRepo.List(keyword, roleID, status, page, pageSize)
|
||||
}
|
||||
|
||||
// AdminGetUser 用户详情
|
||||
func (s *UserService) AdminGetUser(id uint) (*model.User, error) {
|
||||
u, err := s.userRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// AdminCreateUser 管理员创建用户
|
||||
func (s *UserService) AdminCreateUser(username, email, password string, roleID uint, storageLimit int64) (*model.User, error) {
|
||||
username = strings.TrimSpace(username)
|
||||
email = strings.TrimSpace(strings.ToLower(email))
|
||||
if len(username) < 3 {
|
||||
return nil, fmt.Errorf("用户名至少3个字符")
|
||||
}
|
||||
if !utils.IsEmail(email) {
|
||||
return nil, fmt.Errorf("邮箱格式不正确")
|
||||
}
|
||||
if !utils.IsValidPassword(password) {
|
||||
return nil, fmt.Errorf("密码至少8位且需包含字母和数字")
|
||||
}
|
||||
if _, err := s.userRepo.FindByUsername(username); err == nil {
|
||||
return nil, apperr.ErrUserExists
|
||||
}
|
||||
if _, err := s.userRepo.FindByEmail(email); err == nil {
|
||||
return nil, apperr.ErrUserExists
|
||||
}
|
||||
role, err := s.roleRepo.FindByID(roleID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("角色不存在")
|
||||
}
|
||||
if storageLimit <= 0 {
|
||||
storageLimit = 10737418240
|
||||
if v, err := s.settingRepo.GetValue("default_storage_limit"); err == nil {
|
||||
var n int64
|
||||
if _, e := fmt.Sscanf(v, "%d", &n); e == nil && n > 0 {
|
||||
storageLimit = n
|
||||
}
|
||||
}
|
||||
}
|
||||
hash, err := utils.HashPassword(password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u := &model.User{
|
||||
Username: username, Email: email, PasswordHash: hash,
|
||||
RoleID: role.ID, Status: 1, StorageLimit: storageLimit,
|
||||
}
|
||||
if err := s.userRepo.Create(u); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
utils.Info("user", "管理员创建用户: %s", username)
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// AdminUpdateUser 编辑用户(邮箱/角色/配额)
|
||||
func (s *UserService) AdminUpdateUser(id uint, email string, roleID uint, storageLimit int64) (*model.User, error) {
|
||||
u, err := s.userRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
updates := map[string]interface{}{}
|
||||
if email = strings.TrimSpace(strings.ToLower(email)); email != "" {
|
||||
if !utils.IsEmail(email) {
|
||||
return nil, fmt.Errorf("邮箱格式不正确")
|
||||
}
|
||||
if other, err := s.userRepo.FindByEmail(email); err == nil && other.ID != id {
|
||||
return nil, apperr.ErrUserExists
|
||||
}
|
||||
updates["email"] = email
|
||||
}
|
||||
if roleID > 0 {
|
||||
role, err := s.roleRepo.FindByID(roleID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("角色不存在")
|
||||
}
|
||||
// 防止移除最后一个超管
|
||||
if u.Role != nil && u.Role.Code == "super_admin" && role.Code != "super_admin" {
|
||||
if n, _ := s.roleRepo.CountUsersByRole(u.RoleID); n <= 1 {
|
||||
return nil, fmt.Errorf("系统至少需要保留一个超级管理员")
|
||||
}
|
||||
}
|
||||
updates["role_id"] = role.ID
|
||||
}
|
||||
if storageLimit > 0 {
|
||||
updates["storage_limit"] = storageLimit
|
||||
}
|
||||
if len(updates) == 0 {
|
||||
return u, nil
|
||||
}
|
||||
if err := s.userRepo.UpdateFields(id, updates); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.userRepo.FindByID(id)
|
||||
}
|
||||
|
||||
// AdminSetUserStatus 启用/禁用用户
|
||||
func (s *UserService) AdminSetUserStatus(id uint, status int8) error {
|
||||
u, err := s.userRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return apperr.ErrNotFound
|
||||
}
|
||||
if u.Role != nil && u.Role.Code == "super_admin" && status == 0 {
|
||||
if n, _ := s.roleRepo.CountUsersByRole(u.RoleID); n <= 1 {
|
||||
return fmt.Errorf("不能禁用最后一个超级管理员")
|
||||
}
|
||||
}
|
||||
return s.userRepo.UpdateFields(id, map[string]interface{}{"status": status})
|
||||
}
|
||||
|
||||
// AdminAssignRole 修改用户角色
|
||||
func (s *UserService) AdminAssignRole(id, roleID uint) error {
|
||||
role, err := s.roleRepo.FindByID(roleID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("角色不存在")
|
||||
}
|
||||
u, err := s.userRepo.FindByID(id)
|
||||
if err != nil {
|
||||
return apperr.ErrNotFound
|
||||
}
|
||||
if u.Role != nil && u.Role.Code == "super_admin" && role.Code != "super_admin" {
|
||||
if n, _ := s.roleRepo.CountUsersByRole(u.RoleID); n <= 1 {
|
||||
return fmt.Errorf("系统至少需要保留一个超级管理员")
|
||||
}
|
||||
}
|
||||
return s.userRepo.UpdateFields(id, map[string]interface{}{"role_id": role.ID})
|
||||
}
|
||||
|
||||
// AdminResetPassword 重置用户密码
|
||||
func (s *UserService) AdminResetPassword(id uint, newPassword string) error {
|
||||
if !utils.IsValidPassword(newPassword) {
|
||||
return fmt.Errorf("密码至少8位且需包含字母和数字")
|
||||
}
|
||||
hash, err := utils.HashPassword(newPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.userRepo.UpdateFields(id, map[string]interface{}{"password_hash": hash})
|
||||
}
|
||||
|
||||
// AdminDeleteUser 删除用户及其全部数据
|
||||
func (s *UserService) AdminDeleteUser(id uint) error {
|
||||
u, err := s.userRepo.FindByID(id)
|
||||
if err != nil {
|
||||
if err == gorm.ErrRecordNotFound {
|
||||
return apperr.ErrNotFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
if u.Role != nil && u.Role.Code == "super_admin" {
|
||||
if n, _ := s.roleRepo.CountUsersByRole(u.RoleID); n <= 1 {
|
||||
return fmt.Errorf("不能删除最后一个超级管理员")
|
||||
}
|
||||
}
|
||||
|
||||
// 删除物理文件
|
||||
files, _ := s.fileRepo.ListByUser(id)
|
||||
seen := make(map[string]bool)
|
||||
for _, f := range files {
|
||||
if f.StoredPath != "" && !seen[f.StoredPath] {
|
||||
seen[f.StoredPath] = true
|
||||
_ = os.Remove(filepath.Join(s.cfg.Storage.Root, f.StoredPath))
|
||||
}
|
||||
}
|
||||
|
||||
// 级联清理(文件/项目/密钥由外键无级联,手动删除)
|
||||
_ = s.fileRepo.DB.Exec("DELETE FROM files WHERE user_id = ?", id).Error
|
||||
_ = s.fileRepo.DB.Exec("DELETE FROM temp_links WHERE user_id = ?", id).Error
|
||||
_ = s.fileRepo.DB.Exec("DELETE FROM projects WHERE user_id = ?", id).Error
|
||||
_ = s.fileRepo.DB.Exec("DELETE FROM api_keys WHERE user_id = ?", id).Error
|
||||
_ = s.fileRepo.DB.Exec("DELETE FROM webhooks WHERE user_id = ?", id).Error
|
||||
_ = s.fileRepo.DB.Exec("DELETE FROM traffic_logs WHERE user_id = ?", id).Error
|
||||
|
||||
if err := s.userRepo.Delete(id); err != nil {
|
||||
return err
|
||||
}
|
||||
utils.Info("user", "管理员删除用户: %s", u.Username)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,170 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"filestoragesystem/internal/config"
|
||||
"filestoragesystem/internal/model"
|
||||
"filestoragesystem/internal/repository"
|
||||
"filestoragesystem/internal/utils"
|
||||
"filestoragesystem/pkg/apperr"
|
||||
)
|
||||
|
||||
// WebhookService Webhook服务
|
||||
type WebhookService struct {
|
||||
cfg *config.Config
|
||||
webhookRepo *repository.WebhookRepo
|
||||
settingRepo *repository.SettingRepo
|
||||
}
|
||||
|
||||
// Create 创建webhook
|
||||
func (s *WebhookService) Create(userID uint, projectID *uint, url, secret, events string, status int8) (*model.Webhook, error) {
|
||||
if !strings.HasPrefix(url, "http://") && !strings.HasPrefix(url, "https://") {
|
||||
return nil, fmt.Errorf("回调地址必须以http://或https://开头")
|
||||
}
|
||||
events = normalizeEvents(events)
|
||||
if events == "" {
|
||||
return nil, fmt.Errorf("至少订阅一个事件")
|
||||
}
|
||||
if projectID != nil && *projectID > 0 {
|
||||
// 校验项目归属
|
||||
var count int64
|
||||
s.webhookRepo.DB.Model(&model.Project{}).Where("id = ? AND user_id = ?", *projectID, userID).Count(&count)
|
||||
if count == 0 {
|
||||
return nil, apperr.ErrForbidden
|
||||
}
|
||||
}
|
||||
if secret == "" {
|
||||
secret = utils.RandomKey(32)
|
||||
}
|
||||
if status == 0 {
|
||||
status = 1
|
||||
}
|
||||
w := &model.Webhook{
|
||||
UserID: userID, ProjectID: projectID, URL: url,
|
||||
Secret: secret, Events: events, Status: status,
|
||||
}
|
||||
if err := s.webhookRepo.Create(w); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return w, nil
|
||||
}
|
||||
|
||||
// List 用户webhook列表
|
||||
func (s *WebhookService) List(userID uint) ([]model.Webhook, error) {
|
||||
return s.webhookRepo.ListByUser(userID)
|
||||
}
|
||||
|
||||
// Update 更新webhook
|
||||
func (s *WebhookService) Update(userID, id uint, url, events string, status int8) (*model.Webhook, error) {
|
||||
w, err := s.webhookRepo.FindByID(id)
|
||||
if err != nil || w.UserID != userID {
|
||||
return nil, apperr.ErrNotFound
|
||||
}
|
||||
if url != "" {
|
||||
if !strings.HasPrefix(url, "http://") && !strings.HasPrefix(url, "https://") {
|
||||
return nil, fmt.Errorf("回调地址必须以http://或https://开头")
|
||||
}
|
||||
w.URL = url
|
||||
}
|
||||
if events != "" {
|
||||
events = normalizeEvents(events)
|
||||
if events == "" {
|
||||
return nil, fmt.Errorf("至少订阅一个事件")
|
||||
}
|
||||
w.Events = events
|
||||
}
|
||||
if status == 0 || status == 1 {
|
||||
w.Status = status
|
||||
}
|
||||
if err := s.webhookRepo.Update(w); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return w, nil
|
||||
}
|
||||
|
||||
// Delete 删除webhook
|
||||
func (s *WebhookService) Delete(userID, id uint) error {
|
||||
if err := s.webhookRepo.Delete(id, userID); err != nil {
|
||||
return apperr.ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeEvents(events string) string {
|
||||
seen := make(map[string]bool)
|
||||
parts := make([]string, 0, 4)
|
||||
for _, e := range strings.Split(events, ",") {
|
||||
e = strings.TrimSpace(e)
|
||||
if e == "" || seen[e] {
|
||||
continue
|
||||
}
|
||||
seen[e] = true
|
||||
parts = append(parts, e)
|
||||
}
|
||||
return strings.Join(parts, ",")
|
||||
}
|
||||
|
||||
func matchEvent(events, event string) bool {
|
||||
for _, e := range strings.Split(events, ",") {
|
||||
if strings.TrimSpace(e) == event {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Dispatch 异步派发事件回调
|
||||
func (s *WebhookService) Dispatch(userID, projectID uint, event string, payload map[string]interface{}) {
|
||||
if v, err := s.settingRepo.GetValue("webhook_enabled"); err == nil && v != "true" {
|
||||
return
|
||||
}
|
||||
hooks, err := s.webhookRepo.ListActive(userID, projectID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if len(hooks) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
body, _ := json.Marshal(map[string]interface{}{
|
||||
"event": event,
|
||||
"timestamp": time.Now().Unix(),
|
||||
"data": payload,
|
||||
})
|
||||
|
||||
for _, hook := range hooks {
|
||||
if !matchEvent(hook.Events, event) {
|
||||
continue
|
||||
}
|
||||
go s.deliver(hook, body)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *WebhookService) deliver(hook model.Webhook, body []byte) {
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
req, err := http.NewRequest(http.MethodPost, hook.URL, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
utils.Warn("webhook", "构建回调请求失败 url=%s err=%v", hook.URL, err)
|
||||
return
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-Webhook-Event", "filestoragesystem")
|
||||
req.Header.Set("X-Webhook-Signature", utils.HMACSHA256(hook.Secret, string(body)))
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
utils.Warn("webhook", "回调发送失败 url=%s err=%v", hook.URL, err)
|
||||
} else {
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode >= 300 {
|
||||
utils.Warn("webhook", "回调响应异常 url=%s status=%d", hook.URL, resp.StatusCode)
|
||||
}
|
||||
}
|
||||
_ = s.webhookRepo.Touch(hook.ID)
|
||||
}
|
||||
Reference in New Issue
Block a user