Initial commit: TaskPool React panel
- React frontend with route-level code splitting - Backend rebranded from Baihu to TaskPool - DB brand migration script and local compatibility
This commit is contained in:
@@ -0,0 +1,565 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/executor"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/relation"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// AgentService Agent 服务
|
||||
type AgentService struct{}
|
||||
|
||||
// NewAgentService 创建 Agent 服务
|
||||
func NewAgentService() *AgentService {
|
||||
return &AgentService{}
|
||||
}
|
||||
|
||||
// generateToken 生成随机 Token(64位十六进制)
|
||||
func generateToken() string {
|
||||
bytes := make([]byte, 32)
|
||||
rand.Read(bytes)
|
||||
return hex.EncodeToString(bytes)
|
||||
}
|
||||
|
||||
// ========== 令牌管理 ==========
|
||||
|
||||
// CreateToken 创建令牌
|
||||
func (s *AgentService) CreateToken(remark string, maxUses int, expiresAt *time.Time) (*models.AgentToken, error) {
|
||||
var expires *models.LocalTime
|
||||
if expiresAt != nil {
|
||||
t := models.LocalTime(*expiresAt)
|
||||
expires = &t
|
||||
}
|
||||
|
||||
token := generateToken()
|
||||
|
||||
agentToken := &models.AgentToken{
|
||||
ID: utils.GenerateID(),
|
||||
Token: token,
|
||||
Remark: remark,
|
||||
MaxUses: maxUses,
|
||||
ExpiresAt: expires,
|
||||
Enabled: utils.BoolPtr(true),
|
||||
}
|
||||
|
||||
if err := database.DB.Create(agentToken).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
logger.Infof("[Agent] 创建令牌: %s (max_uses=%d)", token[:8]+"...", maxUses)
|
||||
return agentToken, nil
|
||||
}
|
||||
|
||||
// ListTokens 获取令牌列表
|
||||
func (s *AgentService) ListTokens() []models.AgentToken {
|
||||
var tokens []models.AgentToken
|
||||
database.DB.Order("id DESC").Find(&tokens)
|
||||
return tokens
|
||||
}
|
||||
|
||||
// DeleteToken 删除令牌
|
||||
func (s *AgentService) DeleteToken(id string) error {
|
||||
return database.DB.Where("id = ?", id).Delete(&models.AgentToken{}).Error
|
||||
}
|
||||
|
||||
// ValidateToken 验证令牌
|
||||
func (s *AgentService) ValidateToken(token string) (*models.AgentToken, error) {
|
||||
var agentToken models.AgentToken
|
||||
res := database.DB.Where("token = ?", token).Limit(1).Find(&agentToken)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil, &ServiceError{Message: "无效的令牌"}
|
||||
}
|
||||
|
||||
if !utils.DerefBool(agentToken.Enabled, true) {
|
||||
return nil, &ServiceError{Message: "令牌已禁用"}
|
||||
}
|
||||
|
||||
// 检查使用次数
|
||||
if agentToken.MaxUses > 0 && agentToken.UsedCount >= agentToken.MaxUses {
|
||||
return nil, &ServiceError{Message: "令牌已达到使用上限"}
|
||||
}
|
||||
|
||||
// 检查过期时间
|
||||
if agentToken.ExpiresAt != nil && time.Time(*agentToken.ExpiresAt).Before(time.Now()) {
|
||||
return nil, &ServiceError{Message: "令牌已过期"}
|
||||
}
|
||||
|
||||
return &agentToken, nil
|
||||
}
|
||||
|
||||
// UseToken 使用令牌(增加使用计数)
|
||||
func (s *AgentService) UseToken(id string) {
|
||||
database.DB.Model(&models.AgentToken{}).Where("id = ?", id).UpdateColumn("used_count", gorm.Expr("used_count + 1"))
|
||||
}
|
||||
|
||||
// ========== Agent 注册 ==========
|
||||
|
||||
// RegisterByToken 通过令牌注册 Agent(首次 WebSocket 连接时调用)
|
||||
// 返回: agent, isNewAgent, error
|
||||
func (s *AgentService) RegisterByToken(token string, machineID string, ip string) (*models.Agent, bool, error) {
|
||||
// 验证令牌
|
||||
agentToken, err := s.ValidateToken(token)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
// 如果提供了 machine_id,先检查是否已存在
|
||||
if machineID != "" {
|
||||
var existing models.Agent
|
||||
res := database.DB.Where("machine_id = ?", machineID).Limit(1).Find(&existing)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
// 已存在,更新 token 和状态,复用已有 Agent
|
||||
now := models.LocalTime(time.Now())
|
||||
database.DB.Model(&existing).Updates(map[string]interface{}{
|
||||
"token": token,
|
||||
"ip": ip,
|
||||
"status": constant.AgentStatusOnline,
|
||||
"last_seen": now,
|
||||
})
|
||||
s.UseToken(agentToken.ID)
|
||||
logger.Infof("[Agent] Agent #%s 通过 machine_id 复用 (%s)", existing.ID, machineID[:8]+"...")
|
||||
return &existing, false, nil
|
||||
}
|
||||
}
|
||||
|
||||
// 创建 Agent,使用令牌作为认证 Token
|
||||
now := models.LocalTime(time.Now())
|
||||
agent := &models.Agent{
|
||||
ID: utils.GenerateID(),
|
||||
Name: fmt.Sprintf("agent-%d", time.Now().Unix()),
|
||||
Token: token,
|
||||
MachineID: machineID,
|
||||
IP: ip,
|
||||
Status: constant.AgentStatusOnline,
|
||||
LastSeen: &now,
|
||||
Enabled: utils.BoolPtr(true),
|
||||
}
|
||||
|
||||
if err := database.DB.Create(agent).Error; err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
s.UseToken(agentToken.ID)
|
||||
logger.Infof("[Agent] Agent 通过令牌注册: #%s (%s)", agent.ID, ip)
|
||||
return agent, true, nil
|
||||
}
|
||||
|
||||
// Register Agent 注册(必须使用令牌)- 保留兼容旧版本
|
||||
func (s *AgentService) Register(req *models.AgentRegisterRequest, ip string) (*models.Agent, string, error) {
|
||||
// 必须提供令牌
|
||||
if req.Token == "" {
|
||||
return nil, "", &ServiceError{Message: "缺少令牌"}
|
||||
}
|
||||
|
||||
agentToken, err := s.ValidateToken(req.Token)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
|
||||
// 检查是否已存在同名 Agent
|
||||
var existing models.Agent
|
||||
res := database.DB.Where("name = ?", req.Name).Limit(1).Find(&existing)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
return nil, "", &ServiceError{Message: "Agent 名称已存在"}
|
||||
}
|
||||
|
||||
// 创建新 Agent,使用令牌作为认证 Token
|
||||
now := models.LocalTime(time.Now())
|
||||
agent := &models.Agent{
|
||||
ID: utils.GenerateID(),
|
||||
Name: req.Name,
|
||||
Token: req.Token,
|
||||
Hostname: req.Hostname,
|
||||
Version: req.Version,
|
||||
BuildTime: req.BuildTime,
|
||||
IP: ip,
|
||||
Status: constant.AgentStatusOnline,
|
||||
LastSeen: &now,
|
||||
Enabled: utils.BoolPtr(true),
|
||||
}
|
||||
|
||||
if err := database.DB.Create(agent).Error; err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
|
||||
s.UseToken(agentToken.ID)
|
||||
logger.Infof("[Agent] Agent 注册成功: %s (%s)", req.Name, ip)
|
||||
return agent, req.Token, nil
|
||||
}
|
||||
|
||||
// Update 更新 Agent
|
||||
func (s *AgentService) Update(id string, name, description string, enabled bool, schedulerConfig models.AgentSchedulerConfig) error {
|
||||
return database.DB.Model(&models.Agent{}).Where("id = ?", id).Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"description": description,
|
||||
"enabled": &enabled,
|
||||
"scheduler_config": schedulerConfig,
|
||||
}).Error
|
||||
}
|
||||
|
||||
// Delete 删除 Agent(物理删除)
|
||||
func (s *AgentService) Delete(id string) error {
|
||||
// 检查是否有关联任务
|
||||
var count int64
|
||||
database.DB.Model(&models.Task{}).Where("agent_id = ?", id).Count(&count)
|
||||
if count > 0 {
|
||||
return &ServiceError{Message: "该 Agent 下还有关联任务,无法删除"}
|
||||
}
|
||||
|
||||
return database.DB.Where("id = ?", id).Delete(&models.Agent{}).Error
|
||||
}
|
||||
|
||||
// GetByID 根据 ID 获取 Agent
|
||||
func (s *AgentService) GetByID(id string) *models.Agent {
|
||||
var agent models.Agent
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&agent)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
return &agent
|
||||
}
|
||||
|
||||
// GetByToken 根据 Token 获取 Agent
|
||||
func (s *AgentService) GetByToken(token string) *models.Agent {
|
||||
var agent models.Agent
|
||||
res := database.DB.Where("token = ?", token).Limit(1).Find(&agent)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
return &agent
|
||||
}
|
||||
|
||||
// GetByMachineID 根据 MachineID 获取 Agent
|
||||
func (s *AgentService) GetByMachineID(machineID string) *models.Agent {
|
||||
var agent models.Agent
|
||||
res := database.DB.Where("machine_id = ?", machineID).Limit(1).Find(&agent)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
return &agent
|
||||
}
|
||||
|
||||
// List 获取 Agent 列表
|
||||
func (s *AgentService) List() []models.Agent {
|
||||
var agents []models.Agent
|
||||
database.DB.Order("id DESC").Find(&agents)
|
||||
return agents
|
||||
}
|
||||
|
||||
// RegenerateToken 重新生成 Token - 已废弃,保留空实现避免路由错误
|
||||
func (s *AgentService) RegenerateToken(id string) (string, error) {
|
||||
return "", &ServiceError{Message: "此功能已禁用"}
|
||||
}
|
||||
|
||||
// Heartbeat Agent 心跳
|
||||
func (s *AgentService) Heartbeat(token, ip, version, buildTime, hostname, osType, arch string) (*models.Agent, error) {
|
||||
agent := s.GetByToken(token)
|
||||
if agent == nil {
|
||||
return nil, &ServiceError{Message: "无效的 Token"}
|
||||
}
|
||||
|
||||
if !utils.DerefBool(agent.Enabled, true) {
|
||||
return nil, &ServiceError{Message: "Agent 已禁用"}
|
||||
}
|
||||
|
||||
now := models.LocalTime(time.Now())
|
||||
updates := map[string]interface{}{
|
||||
"status": "online",
|
||||
"last_seen": now,
|
||||
"ip": ip,
|
||||
}
|
||||
if version != "" {
|
||||
updates["version"] = version
|
||||
}
|
||||
if buildTime != "" {
|
||||
updates["build_time"] = buildTime
|
||||
}
|
||||
if hostname != "" {
|
||||
updates["hostname"] = hostname
|
||||
}
|
||||
if osType != "" {
|
||||
updates["os"] = osType
|
||||
}
|
||||
if arch != "" {
|
||||
updates["arch"] = arch
|
||||
}
|
||||
|
||||
database.DB.Model(&models.Agent{}).Where("id = ?", agent.ID).Updates(updates)
|
||||
|
||||
agent.Status = constant.AgentStatusOnline
|
||||
agent.LastSeen = &now
|
||||
agent.IP = ip
|
||||
agent.Version = version
|
||||
agent.BuildTime = buildTime
|
||||
agent.Hostname = hostname
|
||||
agent.OS = osType
|
||||
agent.Arch = arch
|
||||
|
||||
return agent, nil
|
||||
}
|
||||
|
||||
// GetTasks 获取 Agent 的任务列表
|
||||
func (s *AgentService) GetTasks(agentID string) []models.AgentTask {
|
||||
var tasksList []models.Task
|
||||
database.DB.Where("agent_id = ? AND enabled = ?", agentID, true).Find(&tasksList)
|
||||
|
||||
// 装载关联的变量信息
|
||||
if len(tasksList) > 0 {
|
||||
taskIDs := make([]string, len(tasksList))
|
||||
for i, t := range tasksList {
|
||||
taskIDs[i] = t.ID
|
||||
}
|
||||
envsMap := relation.DataRelation.LoadRelations(taskIDs, constant.RelationTypeTaskEnv)
|
||||
for i, t := range tasksList {
|
||||
if envs, ok := envsMap[t.ID]; ok {
|
||||
tasksList[i].Envs = models.BigText(strings.Join(envs, ","))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result := make([]models.AgentTask, len(tasksList))
|
||||
envService := NewEnvService()
|
||||
|
||||
for i, task := range tasksList {
|
||||
// 加载环境配置
|
||||
var envVars []string
|
||||
|
||||
// 检查全量注入模式
|
||||
allEnvs := false
|
||||
if task.Config != "" {
|
||||
var config models.TaskConfig
|
||||
if err := json.Unmarshal([]byte(task.Config), &config); err == nil {
|
||||
if config.AllEnvs {
|
||||
allEnvs = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var secrets []string
|
||||
if allEnvs {
|
||||
envVars, secrets = envService.GetAllEnvVarsAndSecrets()
|
||||
} else if string(task.Envs) != "" {
|
||||
envVars, secrets = envService.GetEnvVarsAndSecretsByIDs(string(task.Envs))
|
||||
}
|
||||
|
||||
envVarsStr := executor.FormatEnvVars(envVars)
|
||||
|
||||
command := string(task.Command)
|
||||
preCommand := string(task.PreCommand)
|
||||
postCommand := string(task.PostCommand)
|
||||
workDir := task.WorkDir
|
||||
|
||||
// 仓库同步任务特殊处理:将配置转换为 reposync 命令行
|
||||
if task.Type == constant.TaskTypeRepo {
|
||||
command, workDir = tasks.BuildRepoCommand(&task)
|
||||
// 仓库任务的前置/后置命令已作为参数传给 reposync 内部处理,此处清空防止重复执行
|
||||
preCommand = ""
|
||||
postCommand = ""
|
||||
}
|
||||
|
||||
result[i] = models.AgentTask{
|
||||
ID: task.ID,
|
||||
Name: task.Name,
|
||||
Command: command,
|
||||
PreCommand: preCommand,
|
||||
PostCommand: postCommand,
|
||||
Schedule: task.Schedule,
|
||||
Timeout: task.Timeout,
|
||||
WorkDir: workDir,
|
||||
Envs: envVarsStr,
|
||||
Languages: []map[string]string(task.Languages),
|
||||
RandomRange: task.RandomRange,
|
||||
Secrets: secrets,
|
||||
Enabled: utils.DerefBool(task.Enabled, true),
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// ReportResult Agent 上报执行结果
|
||||
func (s *AgentService) ReportResult(result *models.AgentTaskResult) error {
|
||||
// 获取依赖的服务
|
||||
agentWSManager := GetAgentWSManager()
|
||||
|
||||
// 先尝试通知正在等待的 goroutine
|
||||
if agentWSManager.NotifyRemoteResult(result) {
|
||||
logger.Infof("[Agent] 已通知正在等待任务 #%s 结果的 goroutine", result.TaskID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// 如果没有人在等待(例如服务重启后),则由本协程负责处理结果入库
|
||||
// 如果没有人在等待(例如服务重启后),则由本协程负责处理结果入库(记录日志并清理)
|
||||
logger.Infof("[Agent] 没有找到等待任务 #%s 结果的 goroutine,直接处理结果", result.TaskID)
|
||||
sendStatsService := NewSendStatsService()
|
||||
taskLogService := tasks.NewTaskLogService(sendStatsService)
|
||||
|
||||
// 创建日志对象前进行指令脱敏
|
||||
result.Command = utils.MaskSecrets(result.Command, utils.GetSystemSecrets())
|
||||
taskLog, err := taskLogService.CreateTaskLogFromAgentResult(result)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// 处理完成逻辑(保存日志、更新统计、清理旧日志等)
|
||||
return taskLogService.ProcessTaskCompletion(taskLog)
|
||||
}
|
||||
|
||||
// UpdateTaskDuration 更新任务耗时(心跳)
|
||||
func (s *AgentService) UpdateTaskDuration(logID string, duration int64) error {
|
||||
taskLogService := tasks.NewTaskLogService(nil)
|
||||
return taskLogService.UpdateTaskDuration(logID, duration)
|
||||
}
|
||||
|
||||
// UpdateOfflineAgents 更新离线 Agent 状态(超过 2 分钟无心跳)
|
||||
func (s *AgentService) UpdateOfflineAgents() {
|
||||
cutoff := time.Now().Add(-2 * time.Minute)
|
||||
database.DB.Model(&models.Agent{}).
|
||||
Where("status = ? AND last_seen < ?", constant.AgentStatusOnline, cutoff).
|
||||
Update("status", constant.AgentStatusOffline)
|
||||
}
|
||||
|
||||
// ResetAllAgentsToOffline 将所有 Agents 状态重置为离线(用于服务启动时)
|
||||
func (s *AgentService) ResetAllAgentsToOffline() {
|
||||
database.DB.Model(&models.Agent{}).
|
||||
Where("status = ?", constant.AgentStatusOnline).
|
||||
Update("status", constant.AgentStatusOffline)
|
||||
}
|
||||
|
||||
// GetLatestVersion 获取最新 Agent 版本
|
||||
func (s *AgentService) GetLatestVersion() string {
|
||||
// 优先从 /opt/agent 读取(容器内)
|
||||
versionFile := "/opt/agent/version.txt"
|
||||
data, err := os.ReadFile(versionFile)
|
||||
if err != nil {
|
||||
// 回退到 data/agent(本地开发)
|
||||
data, err = os.ReadFile("data/agent/version.txt")
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
return strings.TrimSpace(string(data))
|
||||
}
|
||||
|
||||
// GetLatestBuildTime 获取最新 Agent 构建时间
|
||||
func (s *AgentService) GetLatestBuildTime() string {
|
||||
return constant.BuildTime
|
||||
}
|
||||
|
||||
// CheckNeedUpdate 检查 Agent 是否需要更新
|
||||
// 比较 version 和 build_time,任一不同则需要更新
|
||||
func (s *AgentService) CheckNeedUpdate(agentVersion, agentBuildTime string) bool {
|
||||
latestVersion := s.GetLatestVersion()
|
||||
latestBuildTime := s.GetLatestBuildTime()
|
||||
|
||||
// 如果没有最新版本信息,不需要更新
|
||||
if latestVersion == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
// 如果 Agent 没有版本信息,需要更新
|
||||
if agentVersion == "" {
|
||||
return true
|
||||
}
|
||||
|
||||
// 版本不同,需要更新
|
||||
if agentVersion != latestVersion {
|
||||
return true
|
||||
}
|
||||
|
||||
// 版本相同但构建时间不同,也需要更新
|
||||
if latestBuildTime != "" && latestBuildTime != "unknown" && agentBuildTime != "" && agentBuildTime != latestBuildTime {
|
||||
return true
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// GetAvailablePlatforms 获取可用的平台列表
|
||||
func (s *AgentService) GetAvailablePlatforms() []map[string]string {
|
||||
platforms := []map[string]string{}
|
||||
|
||||
// 优先从 /opt/agent 读取(容器内)
|
||||
agentDir := "/opt/agent"
|
||||
files, err := os.ReadDir(agentDir)
|
||||
if err != nil {
|
||||
// 回退到 data/agent(本地开发)
|
||||
agentDir = "data/agent"
|
||||
files, err = os.ReadDir(agentDir)
|
||||
if err != nil {
|
||||
return platforms
|
||||
}
|
||||
}
|
||||
|
||||
for _, f := range files {
|
||||
name := f.Name()
|
||||
// taskpool-agent-linux-amd64.tar.gz
|
||||
if strings.HasPrefix(name, "taskpool-agent-") && strings.HasSuffix(name, ".tar.gz") {
|
||||
// 去掉 .tar.gz 后缀
|
||||
baseName := strings.TrimSuffix(name, ".tar.gz")
|
||||
parts := strings.Split(baseName, "-")
|
||||
if len(parts) >= 4 {
|
||||
platforms = append(platforms, map[string]string{
|
||||
"os": parts[2],
|
||||
"arch": parts[3],
|
||||
"filename": name,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return platforms
|
||||
}
|
||||
|
||||
// GetAgentBinary 获取 Agent 压缩包
|
||||
func (s *AgentService) GetAgentBinary(osType, arch string) ([]byte, string, error) {
|
||||
filename := fmt.Sprintf("taskpool-agent-%s-%s.tar.gz", osType, arch)
|
||||
|
||||
// 优先从 /opt/agent 读取(容器内)
|
||||
filePath := filepath.Join("/opt/agent", filename)
|
||||
data, err := os.ReadFile(filePath)
|
||||
if err != nil {
|
||||
// 回退到 data/agent(本地开发)
|
||||
filePath = filepath.Join("data/agent", filename)
|
||||
data, err = os.ReadFile(filePath)
|
||||
if err != nil {
|
||||
return nil, "", &ServiceError{Message: "未找到对应平台的 Agent 程序"}
|
||||
}
|
||||
}
|
||||
|
||||
return data, filename, nil
|
||||
}
|
||||
|
||||
// SetForceUpdate 设置强制更新标志
|
||||
func (s *AgentService) SetForceUpdate(id string) error {
|
||||
return database.DB.Model(&models.Agent{}).Where("id = ?", id).Update("force_update", true).Error
|
||||
}
|
||||
|
||||
// ClearForceUpdate 清除强制更新标志
|
||||
func (s *AgentService) ClearForceUpdate(id string) error {
|
||||
return database.DB.Model(&models.Agent{}).Where("id = ?", id).Update("force_update", false).Error
|
||||
}
|
||||
|
||||
// ServiceError 服务错误
|
||||
type ServiceError struct {
|
||||
Message string
|
||||
}
|
||||
|
||||
func (e *ServiceError) Error() string {
|
||||
return e.Message
|
||||
}
|
||||
@@ -0,0 +1,410 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/executor"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// AgentWSManager WebSocket 连接管理器
|
||||
type AgentWSManager struct {
|
||||
connections map[string]*AgentConnection // Agent ID -> 连接对象
|
||||
ipConnections map[string]int // IP -> 连接数
|
||||
ipLastAttempt map[string]time.Time // IP -> 最后连接尝试时间
|
||||
ipFailCount map[string]int // IP -> 连续失败次数
|
||||
remoteWaiters map[string]chan *models.AgentTaskResult // 日志 ID -> 结果通道
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// 限流配置
|
||||
const (
|
||||
maxConnectionsPerIP = 10 // 每个 IP 最大连接数
|
||||
minConnectInterval = 5 * time.Second // 同一 IP 最小连接间隔
|
||||
maxFailCount = 5 // 最大连续失败次数
|
||||
failBlockDuration = 5 * time.Minute // 失败后封禁时长
|
||||
)
|
||||
|
||||
// AgentConnection Agent WebSocket 连接
|
||||
type AgentConnection struct {
|
||||
AgentID string
|
||||
IP string
|
||||
Conn *websocket.Conn
|
||||
Send chan []byte
|
||||
LastPing time.Time
|
||||
closed bool
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// WSMessage WebSocket 消息结构
|
||||
type WSMessage struct {
|
||||
Type string `json:"type"`
|
||||
Data json.RawMessage `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
// 消息类型常量
|
||||
const (
|
||||
WSTypeHeartbeat = constant.WSTypeHeartbeat
|
||||
WSTypeHeartbeatAck = constant.WSTypeHeartbeatAck
|
||||
WSTypeTasks = constant.WSTypeTasks
|
||||
WSTypeTaskResult = constant.WSTypeTaskResult
|
||||
WSTypeUpdate = constant.WSTypeUpdate
|
||||
WSTypeDisconnect = constant.WSTypeDisconnect
|
||||
WSTypeConnected = constant.WSTypeConnected
|
||||
WSTypeDisabled = constant.WSTypeDisabled
|
||||
WSTypeEnabled = constant.WSTypeEnabled
|
||||
WSTypeFetchTasks = constant.WSTypeFetchTasks
|
||||
WSTypeTaskLog = constant.WSTypeTaskLog
|
||||
WSTypeExecute = constant.WSTypeExecute
|
||||
WSTypeTaskHeartbeat = constant.WSTypeTaskHeartbeat
|
||||
)
|
||||
|
||||
var agentWSManager *AgentWSManager
|
||||
var agentWSOnce sync.Once
|
||||
|
||||
// GetAgentWSManager 获取单例
|
||||
func GetAgentWSManager() *AgentWSManager {
|
||||
agentWSOnce.Do(func() {
|
||||
agentWSManager = &AgentWSManager{
|
||||
connections: make(map[string]*AgentConnection),
|
||||
ipConnections: make(map[string]int),
|
||||
ipLastAttempt: make(map[string]time.Time),
|
||||
ipFailCount: make(map[string]int),
|
||||
remoteWaiters: make(map[string]chan *models.AgentTaskResult),
|
||||
}
|
||||
// 启动时,先将所有 "online" 状态的 Agent 重置为 "offline"
|
||||
NewAgentService().ResetAllAgentsToOffline()
|
||||
|
||||
// 将清理任务注册到系统内部 Cron,每 30 秒执行一次
|
||||
executor.GetSysCron().AddJob("@every 30s", agentWSManager.cleanupLoop)
|
||||
})
|
||||
return agentWSManager
|
||||
}
|
||||
|
||||
// CheckRateLimit 检查 IP 限流,返回是否允许连接
|
||||
func (m *AgentWSManager) CheckRateLimit(ip string) (bool, string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
|
||||
// 检查是否被封禁(连续失败过多)
|
||||
if failCount, exists := m.ipFailCount[ip]; exists && failCount >= maxFailCount {
|
||||
if lastAttempt, ok := m.ipLastAttempt[ip]; ok {
|
||||
if now.Sub(lastAttempt) < failBlockDuration {
|
||||
remaining := failBlockDuration - now.Sub(lastAttempt)
|
||||
return false, "连接失败次数过多,请 " + remaining.Round(time.Second).String() + " 后重试"
|
||||
}
|
||||
// 封禁时间已过,重置计数
|
||||
delete(m.ipFailCount, ip)
|
||||
}
|
||||
}
|
||||
|
||||
// 检查连接频率
|
||||
if lastAttempt, exists := m.ipLastAttempt[ip]; exists {
|
||||
if now.Sub(lastAttempt) < minConnectInterval {
|
||||
return false, "连接过于频繁,请稍后重试"
|
||||
}
|
||||
}
|
||||
|
||||
// 检查 IP 连接数
|
||||
if count, exists := m.ipConnections[ip]; exists && count >= maxConnectionsPerIP {
|
||||
return false, "该 IP 连接数已达上限"
|
||||
}
|
||||
|
||||
m.ipLastAttempt[ip] = now
|
||||
return true, ""
|
||||
}
|
||||
|
||||
// RecordConnectFail 记录连接失败
|
||||
func (m *AgentWSManager) RecordConnectFail(ip string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.ipFailCount[ip]++
|
||||
m.ipLastAttempt[ip] = time.Now()
|
||||
if m.ipFailCount[ip] >= maxFailCount {
|
||||
logger.Warnf("[AgentWS] IP %s 连续失败 %d 次,已封禁 %v", ip, m.ipFailCount[ip], failBlockDuration)
|
||||
}
|
||||
}
|
||||
|
||||
// RecordConnectSuccess 记录连接成功,重置失败计数
|
||||
func (m *AgentWSManager) RecordConnectSuccess(ip string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.ipFailCount, ip)
|
||||
}
|
||||
|
||||
// Register 注册连接
|
||||
func (m *AgentWSManager) Register(agentID string, conn *websocket.Conn, ip string) *AgentConnection {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// 关闭旧连接
|
||||
if old, exists := m.connections[agentID]; exists {
|
||||
// 减少旧 IP 的连接计数
|
||||
if old.IP != "" {
|
||||
if count, ok := m.ipConnections[old.IP]; ok && count > 0 {
|
||||
m.ipConnections[old.IP] = count - 1
|
||||
}
|
||||
}
|
||||
old.Close()
|
||||
}
|
||||
|
||||
ac := &AgentConnection{
|
||||
AgentID: agentID,
|
||||
IP: ip,
|
||||
Conn: conn,
|
||||
Send: make(chan []byte, 256),
|
||||
LastPing: time.Now(),
|
||||
}
|
||||
m.connections[agentID] = ac
|
||||
|
||||
// 增加 IP 连接计数
|
||||
m.ipConnections[ip]++
|
||||
|
||||
logger.Infof("[AgentWS] Agent #%s 已连接 (%s)", agentID, ip)
|
||||
return ac
|
||||
}
|
||||
|
||||
// Unregister 注销连接(只注销指定的连接实例)
|
||||
func (m *AgentWSManager) Unregister(agentID string, ac *AgentConnection) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// 只有当前连接和 map 中的连接是同一个实例时才删除
|
||||
if conn, exists := m.connections[agentID]; exists && conn == ac {
|
||||
// 减少 IP 连接计数
|
||||
if conn.IP != "" {
|
||||
if count, ok := m.ipConnections[conn.IP]; ok && count > 0 {
|
||||
m.ipConnections[conn.IP] = count - 1
|
||||
}
|
||||
}
|
||||
conn.Close()
|
||||
delete(m.connections, agentID)
|
||||
logger.Infof("[AgentWS] Agent #%s 已断开", agentID)
|
||||
}
|
||||
}
|
||||
|
||||
// GetConnection 获取连接
|
||||
func (m *AgentWSManager) GetConnection(agentID string) *AgentConnection {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.connections[agentID]
|
||||
}
|
||||
|
||||
// IsAgentOnline 检查指定 Agent 是否在线
|
||||
func (m *AgentWSManager) IsAgentOnline(agentID string) bool {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
_, exists := m.connections[agentID]
|
||||
return exists
|
||||
}
|
||||
|
||||
// SendToAgent 发送消息给指定 Agent
|
||||
func (m *AgentWSManager) SendToAgent(agentID string, msgType string, data interface{}) error {
|
||||
conn := m.GetConnection(agentID)
|
||||
if conn == nil {
|
||||
return nil // Agent 不在线
|
||||
}
|
||||
|
||||
dataBytes, _ := json.Marshal(data)
|
||||
msg := WSMessage{Type: msgType, Data: dataBytes}
|
||||
msgBytes, _ := json.Marshal(msg)
|
||||
|
||||
select {
|
||||
case conn.Send <- msgBytes:
|
||||
return nil
|
||||
default:
|
||||
return nil // 缓冲区满,丢弃
|
||||
}
|
||||
}
|
||||
|
||||
// BroadcastTasks 广播任务更新给指定 Agent
|
||||
func (m *AgentWSManager) BroadcastTasks(agentID string) {
|
||||
agentService := NewAgentService()
|
||||
tasks := agentService.GetTasks(agentID)
|
||||
m.SendToAgent(agentID, WSTypeTasks, map[string]interface{}{
|
||||
"tasks": tasks,
|
||||
})
|
||||
}
|
||||
|
||||
// BroadcastTasksToAll 广播任务更新给所有在线 Agent
|
||||
func (m *AgentWSManager) BroadcastTasksToAll() {
|
||||
m.mu.RLock()
|
||||
var agentIDs []string
|
||||
for agentID := range m.connections {
|
||||
agentIDs = append(agentIDs, agentID)
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
for _, agentID := range agentIDs {
|
||||
m.BroadcastTasks(agentID)
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRemoteWaiter 注册远程任务结果等待者
|
||||
func (m *AgentWSManager) RegisterRemoteWaiter(logID string) chan *models.AgentTaskResult {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
ch := make(chan *models.AgentTaskResult, 1)
|
||||
m.remoteWaiters[logID] = ch
|
||||
return ch
|
||||
}
|
||||
|
||||
// UnregisterRemoteWaiter 注销远程任务结果等待者
|
||||
func (m *AgentWSManager) UnregisterRemoteWaiter(logID string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.remoteWaiters, logID)
|
||||
}
|
||||
|
||||
// NotifyRemoteResult 通知远程任务结果
|
||||
func (m *AgentWSManager) NotifyRemoteResult(result *models.AgentTaskResult) bool {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
if ch, ok := m.remoteWaiters[result.LogID]; ok {
|
||||
select {
|
||||
case ch <- result:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// OnlineCount 在线 Agent 数量
|
||||
func (m *AgentWSManager) OnlineCount() int {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return len(m.connections)
|
||||
}
|
||||
|
||||
// cleanupLoop 清理超时连接 (由 SysCron 每 30 秒调用一次)
|
||||
func (m *AgentWSManager) cleanupLoop() {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf("[AgentWS] cleanupLoop panic: %v", r)
|
||||
}
|
||||
}()
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
now := time.Now()
|
||||
|
||||
// 清理超时连接
|
||||
for agentID, conn := range m.connections {
|
||||
if now.Sub(conn.LastPing) > 2*time.Minute {
|
||||
// 减少 IP 连接计数
|
||||
if conn.IP != "" {
|
||||
if count, ok := m.ipConnections[conn.IP]; ok && count > 0 {
|
||||
m.ipConnections[conn.IP] = count - 1
|
||||
}
|
||||
}
|
||||
conn.Close()
|
||||
delete(m.connections, agentID)
|
||||
// 更新数据库状态
|
||||
database.DB.Model(&models.Agent{}).Where("id = ?", agentID).Update("status", constant.AgentStatusOffline)
|
||||
logger.Infof("[AgentWS] Agent #%s 心跳超时,已断开", agentID)
|
||||
}
|
||||
}
|
||||
|
||||
// 定期清理数据库中的过期状态(处理服务重启或异常终止的情况)
|
||||
cutoff := now.Add(-2 * time.Minute)
|
||||
database.DB.Model(&models.Agent{}).
|
||||
Where("status = ? AND last_seen < ?", constant.AgentStatusOnline, cutoff).
|
||||
Update("status", constant.AgentStatusOffline)
|
||||
|
||||
// 清理过期的限流记录(超过 10 分钟未活动)
|
||||
for ip, lastAttempt := range m.ipLastAttempt {
|
||||
if now.Sub(lastAttempt) > 10*time.Minute {
|
||||
delete(m.ipLastAttempt, ip)
|
||||
delete(m.ipFailCount, ip)
|
||||
// 只清理没有活跃连接的 IP 计数
|
||||
if m.ipConnections[ip] == 0 {
|
||||
delete(m.ipConnections, ip)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Close 关闭连接
|
||||
func (c *AgentConnection) Close() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed {
|
||||
return
|
||||
}
|
||||
c.closed = true
|
||||
if c.Conn != nil {
|
||||
c.Conn.Close()
|
||||
}
|
||||
if c.Send != nil {
|
||||
close(c.Send)
|
||||
}
|
||||
}
|
||||
|
||||
// IsClosed 检查连接是否已关闭
|
||||
func (c *AgentConnection) IsClosed() bool {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.closed
|
||||
}
|
||||
|
||||
// WriteMessage 写入消息
|
||||
func (c *AgentConnection) WriteMessage(data []byte) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed || c.Conn == nil {
|
||||
return nil
|
||||
}
|
||||
c.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
return c.Conn.WriteMessage(websocket.TextMessage, data)
|
||||
}
|
||||
|
||||
// SetReadDeadline 设置读取超时
|
||||
func (c *AgentConnection) SetReadDeadline(t time.Time) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed || c.Conn == nil {
|
||||
return nil
|
||||
}
|
||||
return c.Conn.SetReadDeadline(t)
|
||||
}
|
||||
|
||||
// ReadMessage 读取消息
|
||||
func (c *AgentConnection) ReadMessage() (int, []byte, error) {
|
||||
// 不加锁,因为 ReadMessage 是阻塞的
|
||||
// 但需要先检查连接状态
|
||||
c.mu.Lock()
|
||||
if c.closed || c.Conn == nil {
|
||||
c.mu.Unlock()
|
||||
return 0, nil, websocket.ErrCloseSent
|
||||
}
|
||||
conn := c.Conn
|
||||
c.mu.Unlock()
|
||||
return conn.ReadMessage()
|
||||
}
|
||||
|
||||
// WritePing 发送 ping 消息
|
||||
func (c *AgentConnection) WritePing() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed || c.Conn == nil {
|
||||
return nil
|
||||
}
|
||||
c.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
return c.Conn.WriteMessage(websocket.PingMessage, nil)
|
||||
}
|
||||
|
||||
// UpdatePing 更新心跳时间
|
||||
func (c *AgentConnection) UpdatePing() {
|
||||
c.LastPing = time.Now()
|
||||
}
|
||||
@@ -0,0 +1,357 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type LogRetentionConfig struct {
|
||||
Days int `json:"days"`
|
||||
MaxCount int `json:"max_count"`
|
||||
}
|
||||
|
||||
type AppLogService struct {
|
||||
settingsService *SettingsService
|
||||
}
|
||||
|
||||
func NewAppLogService() *AppLogService {
|
||||
return &AppLogService{
|
||||
settingsService: NewSettingsService(),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *AppLogService) Add(log *models.AppLog) error {
|
||||
if log.ID == "" {
|
||||
log.ID = utils.GenerateID()
|
||||
}
|
||||
err := database.DB.Create(log).Error
|
||||
if err == nil {
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventAppLogAdded,
|
||||
Payload: log,
|
||||
})
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *AppLogService) List(category string, status string, level string, page, pageSize int, keyword string) ([]models.AppLog, int64, error) {
|
||||
var logs []models.AppLog
|
||||
var total int64
|
||||
db := database.DB
|
||||
|
||||
// 如果是推送日志,尝试关联查询渠道名称
|
||||
if category == constant.LogCategoryPushLog {
|
||||
db = db.Table(models.AppLog{}.TableName() + " AS al").
|
||||
Select("al.*, nw.name as channel_name").
|
||||
Joins(fmt.Sprintf("LEFT JOIN %s AS nw ON al.ref_id = nw.id", models.NotifyWay{}.TableName()))
|
||||
} else {
|
||||
db = db.Model(&models.AppLog{})
|
||||
}
|
||||
|
||||
if category != "" {
|
||||
if category == constant.LogCategoryPushLog {
|
||||
db = db.Where("al.category = ?", category)
|
||||
} else {
|
||||
db = db.Where("category = ?", category)
|
||||
}
|
||||
}
|
||||
if status != "" {
|
||||
field := "status"
|
||||
if category == constant.LogCategoryPushLog {
|
||||
field = "al.status"
|
||||
}
|
||||
db = db.Where(field+" = ?", status)
|
||||
}
|
||||
if level != "" {
|
||||
field := "level"
|
||||
if category == constant.LogCategoryPushLog {
|
||||
field = "al.level"
|
||||
}
|
||||
db = db.Where(field+" = ?", level)
|
||||
}
|
||||
if keyword != "" {
|
||||
if category == constant.LogCategoryPushLog {
|
||||
db = db.Where("(al.title LIKE ? OR al.content LIKE ?)", "%"+keyword+"%", "%"+keyword+"%")
|
||||
} else {
|
||||
db = db.Where("(title LIKE ? OR content LIKE ?)", "%"+keyword+"%", "%"+keyword+"%")
|
||||
}
|
||||
}
|
||||
|
||||
db.Count(&total)
|
||||
offset := (page - 1) * pageSize
|
||||
order := "created_at desc"
|
||||
if category == constant.LogCategoryPushLog {
|
||||
order = "al.created_at desc"
|
||||
}
|
||||
err := db.Order(order).Offset(offset).Limit(pageSize).Scan(&logs).Error
|
||||
return logs, total, err
|
||||
}
|
||||
|
||||
func (s *AppLogService) MarkAsRead(id string) error {
|
||||
now := models.LocalTime(time.Now())
|
||||
return database.DB.Model(&models.AppLog{}).Where("id = ?", id).Updates(map[string]interface{}{
|
||||
"status": constant.LogStatusRead,
|
||||
"read_at": &now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (s *AppLogService) MarkAllAsRead(category string) error {
|
||||
now := models.LocalTime(time.Now())
|
||||
return database.DB.Model(&models.AppLog{}).Where("category = ? AND status = ?", category, constant.LogStatusUnread).Updates(map[string]interface{}{
|
||||
"status": constant.LogStatusRead,
|
||||
"read_at": &now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (s *AppLogService) Clear(category string) error {
|
||||
query := database.DB.Model(&models.AppLog{})
|
||||
if category != "" {
|
||||
query = query.Where("category = ?", category)
|
||||
}
|
||||
return query.Delete(&models.AppLog{}).Error
|
||||
}
|
||||
|
||||
func (s *AppLogService) GetRetentionConfigs() map[string]LogRetentionConfig {
|
||||
configs := map[string]LogRetentionConfig{
|
||||
constant.LogCategorySystemNotice: {
|
||||
Days: utils.ToInt(s.settingsService.Get(constant.SectionSystem, constant.KeySystemNoticeDays), 30),
|
||||
MaxCount: utils.ToInt(s.settingsService.Get(constant.SectionSystem, constant.KeySystemNoticeMaxCount), 500),
|
||||
},
|
||||
constant.LogCategoryPushLog: {
|
||||
Days: utils.ToInt(s.settingsService.Get(constant.SectionSystem, constant.KeyPushLogDays), 15),
|
||||
MaxCount: utils.ToInt(s.settingsService.Get(constant.SectionSystem, constant.KeyPushLogMaxCount), 5000),
|
||||
},
|
||||
constant.LogCategoryLoginLog: {
|
||||
Days: utils.ToInt(s.settingsService.Get(constant.SectionSystem, constant.KeyLoginLogDays), 30),
|
||||
MaxCount: utils.ToInt(s.settingsService.Get(constant.SectionSystem, constant.KeyLoginLogMaxCount), 1000),
|
||||
},
|
||||
constant.LogCategorySchedulerLog: {
|
||||
Days: utils.ToInt(s.settingsService.Get(constant.SectionSystem, constant.KeySchedulerLogDays), 30),
|
||||
MaxCount: utils.ToInt(s.settingsService.Get(constant.SectionSystem, constant.KeySchedulerLogMaxCount), 10000),
|
||||
},
|
||||
constant.LogCategoryDefault: {
|
||||
Days: 30,
|
||||
MaxCount: 10000,
|
||||
},
|
||||
}
|
||||
return configs
|
||||
}
|
||||
|
||||
func (s *AppLogService) AddSchedulerLog(title, content, level string) error {
|
||||
if level == "" {
|
||||
level = constant.LogLevelInfo
|
||||
}
|
||||
return s.Add(&models.AppLog{
|
||||
Category: constant.LogCategorySchedulerLog,
|
||||
Title: title,
|
||||
Content: models.BigText(content),
|
||||
Level: level,
|
||||
Status: constant.LogStatusRead,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *AppLogService) CleanUp() {
|
||||
configs := s.GetRetentionConfigs()
|
||||
categories := []string{constant.LogCategorySystemNotice, constant.LogCategoryPushLog, constant.LogCategoryLoginLog, constant.LogCategorySchedulerLog}
|
||||
|
||||
var totalDeleted int64
|
||||
var summaryBuilder strings.Builder
|
||||
|
||||
for _, cat := range categories {
|
||||
cfg, ok := configs[cat]
|
||||
if !ok {
|
||||
cfg = configs[constant.LogCategoryDefault]
|
||||
}
|
||||
|
||||
var daysDeleted int64
|
||||
if cfg.Days > 0 {
|
||||
deadline := time.Now().AddDate(0, 0, -cfg.Days)
|
||||
res := database.DB.Where("category = ? AND created_at < ?", cat, deadline).Delete(&models.AppLog{})
|
||||
if res.Error == nil {
|
||||
daysDeleted = res.RowsAffected
|
||||
}
|
||||
}
|
||||
|
||||
var countDeleted int64
|
||||
if cfg.MaxCount > 0 {
|
||||
var total int64
|
||||
database.DB.Model(&models.AppLog{}).Where("category = ?", cat).Count(&total)
|
||||
if total > int64(cfg.MaxCount) {
|
||||
deleteCount := total - int64(cfg.MaxCount)
|
||||
var ids []string
|
||||
database.DB.Model(&models.AppLog{}).Where("category = ?", cat).Order("created_at asc").Limit(int(deleteCount)).Pluck("id", &ids)
|
||||
if len(ids) > 0 {
|
||||
res := database.DB.Where("id IN ?", ids).Delete(&models.AppLog{})
|
||||
if res.Error == nil {
|
||||
countDeleted = res.RowsAffected
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
catTotal := daysDeleted + countDeleted
|
||||
if catTotal > 0 {
|
||||
totalDeleted += catTotal
|
||||
var catLabel string
|
||||
switch cat {
|
||||
case constant.LogCategorySystemNotice:
|
||||
catLabel = "系统通知"
|
||||
case constant.LogCategoryPushLog:
|
||||
catLabel = "推送日志"
|
||||
case constant.LogCategoryLoginLog:
|
||||
catLabel = "登录日志"
|
||||
case constant.LogCategorySchedulerLog:
|
||||
catLabel = "调度日志"
|
||||
default:
|
||||
catLabel = cat
|
||||
}
|
||||
summaryBuilder.WriteString(fmt.Sprintf("分类 [%s]: 过期清理 %d 条,溢出限制清理 %d 条;\n", catLabel, daysDeleted, countDeleted))
|
||||
}
|
||||
}
|
||||
|
||||
if totalDeleted > 0 {
|
||||
logger.Infof("[AppLog] 周期巡检完成清理,共删除 %d 条陈旧日志", totalDeleted)
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventSystemNotice,
|
||||
Payload: map[string]interface{}{
|
||||
"title": "系统日志定时容量收敛",
|
||||
"content": fmt.Sprintf("后台巡检已自动执行日志容量清理,共计清除陈旧或溢出记录 %d 条。\n详情明细:\n%s", totalDeleted, summaryBuilder.String()),
|
||||
"level": constant.LogLevelInfo,
|
||||
},
|
||||
})
|
||||
} else {
|
||||
logger.Debugf("[AppLog] 完成应用日志清理策略,未检测到需要剔除的陈旧记录")
|
||||
}
|
||||
}
|
||||
|
||||
func (s *AppLogService) SubscribeEvents(bus *eventbus.EventBus) {
|
||||
// 1. [订阅] 系统通知 -> 存储到数据库表现为红点消息
|
||||
bus.Subscribe(constant.EventSystemNotice, func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
title, _ := payload["title"].(string)
|
||||
content, _ := payload["content"].(string)
|
||||
level, _ := payload["level"].(string)
|
||||
if level == "" {
|
||||
level = constant.LogLevelInfo
|
||||
}
|
||||
|
||||
s.Add(&models.AppLog{
|
||||
Category: constant.LogCategorySystemNotice,
|
||||
Title: title,
|
||||
Content: models.BigText(content),
|
||||
Level: level,
|
||||
Status: constant.LogStatusUnread,
|
||||
})
|
||||
})
|
||||
|
||||
// 2. [订阅] 推送结果 -> 存储到数据库供推送日志查看
|
||||
bus.Subscribe(constant.EventNotifySent, func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
title, _ := payload["title"].(string)
|
||||
content, _ := payload["content"].(string)
|
||||
success, _ := payload["success"].(bool)
|
||||
errorMsg, _ := payload["error_msg"].(string)
|
||||
channelID, _ := payload["channel_id"].(string)
|
||||
|
||||
status := constant.LogStatusSuccess
|
||||
level := constant.LogLevelInfo
|
||||
if !success {
|
||||
status = constant.LogStatusFailed
|
||||
level = constant.LogLevelError
|
||||
}
|
||||
|
||||
s.Add(&models.AppLog{
|
||||
Category: constant.LogCategoryPushLog,
|
||||
Title: title,
|
||||
Content: models.BigText(content),
|
||||
Level: level,
|
||||
Status: status,
|
||||
RefID: channelID,
|
||||
ErrorMsg: models.BigText(errorMsg),
|
||||
})
|
||||
})
|
||||
|
||||
// 3. 将某些业务事件转化为系统内部通知 (自动出现在小铃铛)
|
||||
/*
|
||||
bus.Subscribe(constant.EventTaskFailed, func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
taskName, _ := payload["task_name"].(string)
|
||||
errMsg, _ := payload["error"].(string)
|
||||
|
||||
bus.Publish(eventbus.Event{
|
||||
Type: constant.EventSystemNotice,
|
||||
Payload: map[string]interface{}{
|
||||
"title": fmt.Sprintf("任务 [%s] 执行失败", taskName),
|
||||
"content": fmt.Sprintf("错误详情: %s", errMsg),
|
||||
"level": constant.LogLevelError,
|
||||
},
|
||||
})
|
||||
})
|
||||
|
||||
bus.Subscribe(constant.EventTaskTimeout, func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
taskName, _ := payload["task_name"].(string)
|
||||
|
||||
bus.Publish(eventbus.Event{
|
||||
Type: constant.EventSystemNotice,
|
||||
Payload: map[string]interface{}{
|
||||
"title": fmt.Sprintf("任务 [%s] 执行超时", taskName),
|
||||
"content": "任务已经超过预设的运行时间并被系统强制中止。",
|
||||
"level": constant.LogLevelWarning,
|
||||
},
|
||||
})
|
||||
})
|
||||
*/
|
||||
|
||||
bus.Subscribe(constant.EventPasswordChanged, func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
username, _ := payload["username"].(string)
|
||||
|
||||
bus.Publish(eventbus.Event{
|
||||
Type: constant.EventSystemNotice,
|
||||
Payload: map[string]interface{}{
|
||||
"title": "账户安全提醒",
|
||||
"content": fmt.Sprintf("用户 %s 的账号密码已被修改。", username),
|
||||
"level": constant.LogLevelInfo,
|
||||
},
|
||||
})
|
||||
})
|
||||
|
||||
// 4. [订阅] 调度日志写入
|
||||
bus.Subscribe(constant.EventSchedulerLog, func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
title, _ := payload["title"].(string)
|
||||
content, _ := payload["content"].(string)
|
||||
level, _ := payload["level"].(string)
|
||||
s.AddSchedulerLog(title, content, level)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,530 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/cache"
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/systime"
|
||||
"github.com/rs/xid"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type BackupService struct {
|
||||
settingsService *SettingsService
|
||||
}
|
||||
|
||||
func NewBackupService() *BackupService {
|
||||
return &BackupService{
|
||||
settingsService: NewSettingsService(),
|
||||
}
|
||||
}
|
||||
|
||||
const (
|
||||
BackupSection = "backup"
|
||||
BackupFileKey = "backup_file"
|
||||
BackupDir = "./data/backups"
|
||||
)
|
||||
|
||||
// tableConfig 表备份配置
|
||||
type tableConfig struct {
|
||||
filename string
|
||||
export func(io.Writer) error
|
||||
restore func([]byte) error
|
||||
}
|
||||
|
||||
func (s *BackupService) getTableConfigs() []tableConfig {
|
||||
return []tableConfig{
|
||||
{"users.json", s.exportTable(&[]models.User{}), s.restoreTable(&[]models.User{})},
|
||||
{"tasks.json", s.exportTable(&[]models.Task{}), s.restoreTable(&[]models.Task{})},
|
||||
{"task_logs.json", s.exportTable(&[]models.TaskLog{}), s.restoreTable(&[]models.TaskLog{})},
|
||||
{"envs.json", s.exportTable(&[]models.EnvironmentVariable{}), s.restoreTable(&[]models.EnvironmentVariable{})},
|
||||
{"scripts.json", s.exportTable(&[]models.Script{}), s.restoreTable(&[]models.Script{})},
|
||||
{"settings.json", s.exportSettings, s.restoreSettings},
|
||||
{"send_stats.json", s.exportTable(&[]models.SendStats{}), s.restoreTable(&[]models.SendStats{})},
|
||||
|
||||
{"agents.json", s.exportTable(&[]models.Agent{}), s.restoreTable(&[]models.Agent{})},
|
||||
{"tokens.json", s.exportTable(&[]models.AgentToken{}), s.restoreTable(&[]models.AgentToken{})},
|
||||
{"languages.json", s.exportTable(&[]models.Language{}), s.restoreTable(&[]models.Language{})},
|
||||
{"deps.json", s.exportTable(&[]models.Dependency{}), s.restoreTable(&[]models.Dependency{})},
|
||||
{"notify_ways.json", s.exportTable(&[]models.NotifyWay{}), s.restoreTable(&[]models.NotifyWay{})},
|
||||
{"notify_bindings.json", s.exportTable(&[]models.NotifyBinding{}), s.restoreTable(&[]models.NotifyBinding{})},
|
||||
{"app_logs.json", s.exportTable(&[]models.AppLog{}), s.restoreTable(&[]models.AppLog{})},
|
||||
{"data_storage.json", s.exportTable(&[]models.DataStorage{}), s.restoreTable(&[]models.DataStorage{})},
|
||||
{"data_relations.json", s.exportTable(&[]models.DataRelation{}), s.restoreTable(&[]models.DataRelation{})},
|
||||
}
|
||||
}
|
||||
|
||||
func (s *BackupService) exportTable(modelPtr any) func(io.Writer) error {
|
||||
return func(w io.Writer) error {
|
||||
db := database.DB
|
||||
|
||||
if _, err := w.Write([]byte("[\n")); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
first := true
|
||||
err := db.FindInBatches(modelPtr, 1000, func(tx *gorm.DB, batch int) error {
|
||||
val := reflect.ValueOf(modelPtr).Elem()
|
||||
count := val.Len()
|
||||
for i := 0; i < count; i++ {
|
||||
if !first {
|
||||
if _, err := w.Write([]byte(",\n")); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
item := val.Index(i).Interface()
|
||||
jsonData, err := json.MarshalIndent(item, " ", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := w.Write(jsonData); err != nil {
|
||||
return err
|
||||
}
|
||||
first = false
|
||||
}
|
||||
return nil
|
||||
}).Error
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, err = w.Write([]byte("\n]"))
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
func (s *BackupService) restoreTable(dest any) func([]byte) error {
|
||||
return func(data []byte) error {
|
||||
if err := json.Unmarshal(data, dest); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (s *BackupService) exportSettings(w io.Writer) error {
|
||||
var data []models.Setting
|
||||
if err := database.DB.Where("section != ?", BackupSection).Find(&data).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
jsonData, err := json.MarshalIndent(data, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = w.Write(jsonData)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *BackupService) restoreSettings(data []byte) error {
|
||||
var settings []models.Setting
|
||||
return json.Unmarshal(data, &settings)
|
||||
}
|
||||
|
||||
// CreateBackup 创建备份
|
||||
func (s *BackupService) CreateBackup() (string, error) {
|
||||
if err := os.MkdirAll(BackupDir, 0755); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
timestamp := systime.FormatDatetime(time.Now())
|
||||
zipPath := filepath.Join(BackupDir, fmt.Sprintf("backup_%s.zip", timestamp))
|
||||
|
||||
zipFile, err := os.Create(zipPath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer zipFile.Close()
|
||||
|
||||
zipWriter := zip.NewWriter(zipFile)
|
||||
defer zipWriter.Close()
|
||||
|
||||
// 导出各表
|
||||
for _, cfg := range s.getTableConfigs() {
|
||||
w, err := zipWriter.Create(cfg.filename)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := cfg.export(w); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
// 写入元数据信息
|
||||
sysInfo := map[string]interface{}{
|
||||
"version": "v3",
|
||||
"ts": time.Now().Format("2006-01-02 15:04:05"),
|
||||
}
|
||||
sysFile, err := zipWriter.Create("__sys__.json")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
sysData, _ := json.MarshalIndent(sysInfo, "", " ")
|
||||
if _, err := sysFile.Write(sysData); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 打包 scripts 文件夹
|
||||
scriptsDir := constant.ScriptsWorkDir
|
||||
if _, err := os.Stat(scriptsDir); err == nil {
|
||||
if err := s.addDirToZip(zipWriter, scriptsDir, "scripts"); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
s.settingsService.Set(BackupSection, BackupFileKey, zipPath)
|
||||
return zipPath, nil
|
||||
}
|
||||
|
||||
// Restore 恢复备份
|
||||
func (s *BackupService) Restore(zipPath string) error {
|
||||
r, err := zip.OpenReader(zipPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
// 构建文件名到配置的映射
|
||||
configs := s.getTableConfigs()
|
||||
fileMap := make(map[string]*zip.File)
|
||||
for _, f := range r.File {
|
||||
fileMap[f.Name] = f
|
||||
}
|
||||
|
||||
// 校验版本
|
||||
if f, ok := fileMap["__sys__.json"]; ok {
|
||||
rc, err := f.Open()
|
||||
if err == nil {
|
||||
var sysInfo map[string]interface{}
|
||||
json.NewDecoder(rc).Decode(&sysInfo)
|
||||
rc.Close()
|
||||
if v, ok := sysInfo["version"]; ok {
|
||||
vs, _ := v.(string)
|
||||
if vs < "v3" {
|
||||
return fmt.Errorf("只能数据随版本升级上来,当前备份版本为 %s,限制 v3 以下的不能导入", vs)
|
||||
}
|
||||
}
|
||||
}
|
||||
// } else {
|
||||
// return fmt.Errorf("非法备份包:缺失版本标记")
|
||||
}
|
||||
|
||||
// 开启全局事务
|
||||
err = database.DB.Transaction(func(tx *gorm.DB) error {
|
||||
// 1. 清空现有数据(物理删除)
|
||||
tx.Where("1=1").Delete(&models.User{})
|
||||
tx.Where("1=1").Delete(&models.Task{})
|
||||
tx.Where("1=1").Delete(&models.TaskLog{})
|
||||
tx.Where("1=1").Delete(&models.EnvironmentVariable{})
|
||||
tx.Where("1=1").Delete(&models.Script{})
|
||||
tx.Where("section != ?", BackupSection).Delete(&models.Setting{})
|
||||
tx.Where("1=1").Delete(&models.SendStats{})
|
||||
|
||||
tx.Where("1=1").Delete(&models.Agent{})
|
||||
tx.Where("1=1").Delete(&models.AgentToken{})
|
||||
tx.Where("1=1").Delete(&models.Language{})
|
||||
tx.Where("1=1").Delete(&models.Dependency{})
|
||||
tx.Where("1=1").Delete(&models.NotifyWay{})
|
||||
tx.Where("1=1").Delete(&models.NotifyBinding{})
|
||||
tx.Where("1=1").Delete(&models.AppLog{})
|
||||
tx.Where("1=1").Delete(&models.DataStorage{})
|
||||
tx.Where("1=1").Delete(&models.DataRelation{})
|
||||
|
||||
// 2. 依次恢复每个表
|
||||
for _, cfg := range configs {
|
||||
if f, ok := fileMap[cfg.filename]; ok {
|
||||
if err := s.restoreFromZipFile(tx, f, cfg.filename); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 恢复 scripts 文件夹
|
||||
s.restoreScriptsDir(r)
|
||||
|
||||
return nil
|
||||
})
|
||||
|
||||
if err == nil {
|
||||
// 备份恢复成功后,需要同时刷新内存中的配置缓存以免数据不一致导致异常
|
||||
constant.Secret = s.settingsService.Get(constant.SectionSecurity, constant.KeySecret)
|
||||
cache.LoadSiteCache()
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func restoreStreamBatch[T any](tx *gorm.DB, decoder *json.Decoder) error {
|
||||
batchSize := 1000
|
||||
var batch []*T
|
||||
|
||||
for decoder.More() {
|
||||
var m T
|
||||
if err := decoder.Decode(&m); err != nil {
|
||||
return err
|
||||
}
|
||||
batch = append(batch, &m)
|
||||
|
||||
if len(batch) >= batchSize {
|
||||
if err := tx.Select("*").CreateInBatches(batch, batchSize).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
batch = nil // reset batch
|
||||
}
|
||||
}
|
||||
|
||||
if len(batch) > 0 {
|
||||
return tx.Select("*").CreateInBatches(batch, len(batch)).Error
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *BackupService) restoreFromZipFile(tx *gorm.DB, f *zip.File, filename string) error {
|
||||
rc, err := f.Open()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer rc.Close()
|
||||
|
||||
// 特殊处理设置表(设置表通常很小,直接反序列化)
|
||||
if filename == "settings.json" {
|
||||
data, _ := io.ReadAll(rc)
|
||||
var settings []models.Setting
|
||||
if err := json.Unmarshal(data, &settings); err == nil {
|
||||
if len(settings) > 0 {
|
||||
return tx.Select("*").Create(&settings).Error
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// 流式解析 JSON 数组
|
||||
decoder := json.NewDecoder(rc)
|
||||
|
||||
// 找到数组开始 [
|
||||
if t, err := decoder.Token(); err != nil || t != json.Delim('[') {
|
||||
return fmt.Errorf("invalid json format: expected %s", filename)
|
||||
}
|
||||
|
||||
switch filename {
|
||||
case "users.json":
|
||||
return restoreStreamBatch[models.User](tx, decoder)
|
||||
case "tasks.json":
|
||||
return s.restoreTasks(tx, decoder)
|
||||
case "task_logs.json":
|
||||
return restoreStreamBatch[models.TaskLog](tx, decoder)
|
||||
case "envs.json":
|
||||
return restoreStreamBatch[models.EnvironmentVariable](tx, decoder)
|
||||
case "scripts.json":
|
||||
return restoreStreamBatch[models.Script](tx, decoder)
|
||||
case "send_stats.json":
|
||||
return restoreStreamBatch[models.SendStats](tx, decoder)
|
||||
|
||||
case "agents.json":
|
||||
return restoreStreamBatch[models.Agent](tx, decoder)
|
||||
case "tokens.json":
|
||||
return restoreStreamBatch[models.AgentToken](tx, decoder)
|
||||
case "languages.json":
|
||||
return restoreStreamBatch[models.Language](tx, decoder)
|
||||
case "deps.json":
|
||||
return restoreStreamBatch[models.Dependency](tx, decoder)
|
||||
case "notify_ways.json":
|
||||
return restoreStreamBatch[models.NotifyWay](tx, decoder)
|
||||
case "notify_bindings.json":
|
||||
return restoreStreamBatch[models.NotifyBinding](tx, decoder)
|
||||
case "app_logs.json":
|
||||
return restoreStreamBatch[models.AppLog](tx, decoder)
|
||||
case "data_storage.json":
|
||||
return restoreStreamBatch[models.DataStorage](tx, decoder)
|
||||
case "data_relations.json":
|
||||
return restoreStreamBatch[models.DataRelation](tx, decoder)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (s *BackupService) restoreTasks(tx *gorm.DB, decoder *json.Decoder) error {
|
||||
batchSize := 1000
|
||||
var batch []*models.Task
|
||||
|
||||
for decoder.More() {
|
||||
var m models.Task
|
||||
if err := decoder.Decode(&m); err != nil {
|
||||
return err
|
||||
}
|
||||
batch = append(batch, &m)
|
||||
|
||||
if len(batch) >= batchSize {
|
||||
if err := s.insertTasksBatch(tx, batch); err != nil {
|
||||
return err
|
||||
}
|
||||
batch = nil // reset
|
||||
}
|
||||
}
|
||||
|
||||
if len(batch) > 0 {
|
||||
return s.insertTasksBatch(tx, batch)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *BackupService) insertTasksBatch(tx *gorm.DB, batch []*models.Task) error {
|
||||
if err := tx.Select("*").CreateInBatches(batch, len(batch)).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// [COMPATIBILITY CODE - TEMPORARY]
|
||||
// 为旧版备份数据做向下兼容(在旧版本中 tags/envs 还是直接存储在 tasks.json 中)。
|
||||
// 此段代码将旧数据迁移到新的 data_relations 和 data_storages 中,确保旧备份包能正常恢复。
|
||||
// 由于新版本的备份中 tasks.json 已不再包含这些字段,未来某次大版本更新后,这段兼容代码将被移除。
|
||||
for _, t := range batch {
|
||||
if t.Tags != "" {
|
||||
tags := strings.Split(t.Tags, ",")
|
||||
for _, tag := range tags {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
var storage models.DataStorage
|
||||
res := tx.Where("type = ? AND name = ?", constant.RelationTypeTaskTag, tag).Limit(1).Find(&storage)
|
||||
if res.RowsAffected == 0 {
|
||||
storage = models.DataStorage{
|
||||
ID: xid.New().String(),
|
||||
Type: constant.RelationTypeTaskTag,
|
||||
Name: tag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
tx.Create(&storage)
|
||||
}
|
||||
|
||||
var count int64
|
||||
tx.Model(&models.DataRelation{}).Where("data_id = ? AND relate_id = ? AND type = ?", t.ID, storage.ID, constant.RelationTypeTaskTag).Count(&count)
|
||||
if count == 0 {
|
||||
relation := models.DataRelation{
|
||||
ID: xid.New().String(),
|
||||
DataID: t.ID,
|
||||
RelateID: storage.ID,
|
||||
Type: constant.RelationTypeTaskTag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
tx.Create(&relation)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if string(t.Envs) != "" {
|
||||
envs := strings.Split(string(t.Envs), ",")
|
||||
for _, envID := range envs {
|
||||
envID = strings.TrimSpace(envID)
|
||||
if envID == "" {
|
||||
continue
|
||||
}
|
||||
var count int64
|
||||
tx.Model(&models.DataRelation{}).Where("data_id = ? AND relate_id = ? AND type = ?", t.ID, envID, constant.RelationTypeTaskEnv).Count(&count)
|
||||
if count == 0 {
|
||||
relation := models.DataRelation{
|
||||
ID: xid.New().String(),
|
||||
DataID: t.ID,
|
||||
RelateID: envID,
|
||||
Type: constant.RelationTypeTaskEnv,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
tx.Create(&relation)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// insertRecords, restoreFromData 方法已合并入 restoreFromZipFile,此处删除冗余方法
|
||||
|
||||
func (s *BackupService) restoreScriptsDir(r *zip.ReadCloser) {
|
||||
scriptsDir := constant.ScriptsWorkDir
|
||||
for _, f := range r.File {
|
||||
if len(f.Name) > 8 && f.Name[:8] == "scripts/" {
|
||||
relPath := f.Name[8:]
|
||||
if relPath == "" {
|
||||
continue
|
||||
}
|
||||
fpath := filepath.Join(scriptsDir, relPath)
|
||||
if f.FileInfo().IsDir() {
|
||||
os.MkdirAll(fpath, 0755)
|
||||
continue
|
||||
}
|
||||
os.MkdirAll(filepath.Dir(fpath), 0755)
|
||||
if outFile, err := os.Create(fpath); err == nil {
|
||||
if rc, err := f.Open(); err == nil {
|
||||
io.Copy(outFile, rc)
|
||||
rc.Close()
|
||||
}
|
||||
outFile.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *BackupService) addDirToZip(zipWriter *zip.Writer, srcDir, prefix string) error {
|
||||
return filepath.Walk(srcDir, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
relPath, err := filepath.Rel(srcDir, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
zipPath := filepath.ToSlash(filepath.Join(prefix, relPath))
|
||||
if info.IsDir() {
|
||||
if relPath != "." {
|
||||
_, err := zipWriter.Create(zipPath + "/")
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
w, err := zipWriter.Create(zipPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
file, err := os.Open(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer file.Close()
|
||||
_, err = io.Copy(w, file)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func (s *BackupService) GetBackupFile() string {
|
||||
var setting models.Setting
|
||||
res := database.DB.Where(&models.Setting{Section: BackupSection, Key: BackupFileKey}).Limit(1).Find(&setting)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return ""
|
||||
}
|
||||
return string(setting.Value)
|
||||
}
|
||||
|
||||
func (s *BackupService) ClearBackup() error {
|
||||
filePath := s.GetBackupFile()
|
||||
if filePath != "" {
|
||||
os.Remove(filePath)
|
||||
database.DB.Where(&models.Setting{Section: BackupSection, Key: BackupFileKey}).Delete(&models.Setting{})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"gopkg.in/ini.v1"
|
||||
)
|
||||
|
||||
type ServerConfig struct {
|
||||
Port int `ini:"port"`
|
||||
Host string `ini:"host"`
|
||||
URLPrefix string `ini:"url_prefix"`
|
||||
PprofEnabled bool `ini:"pprof_enabled"`
|
||||
CookieName string `ini:"cookie_name"`
|
||||
}
|
||||
|
||||
type DatabaseConfig struct {
|
||||
Type string `ini:"type"`
|
||||
Host string `ini:"host"`
|
||||
Port int `ini:"port"`
|
||||
User string `ini:"user"`
|
||||
Password string `ini:"password"`
|
||||
DBName string `ini:"dbname"`
|
||||
Path string `ini:"path"`
|
||||
DSN string `ini:"dsn"`
|
||||
TablePrefix string `ini:"table_prefix"`
|
||||
SSLMode string `ini:"ssl_mode"`
|
||||
}
|
||||
|
||||
type SecurityConfig struct {
|
||||
Secret string `ini:"secret"`
|
||||
}
|
||||
|
||||
type RedisConfig struct {
|
||||
Enabled bool `ini:"enabled"`
|
||||
Host string `ini:"host"`
|
||||
Port int `ini:"port"`
|
||||
Password string `ini:"password"`
|
||||
DB int `ini:"db"`
|
||||
}
|
||||
|
||||
type AppConfig struct {
|
||||
Server ServerConfig `ini:"server"`
|
||||
Database DatabaseConfig `ini:"database"`
|
||||
Security SecurityConfig `ini:"security"`
|
||||
Redis RedisConfig `ini:"redis"`
|
||||
}
|
||||
|
||||
var Config *AppConfig
|
||||
|
||||
// getEnvStr 获取环境变量字符串
|
||||
func getEnvStr(key string, target *string) {
|
||||
if v := os.Getenv(key); v != "" {
|
||||
*target = v
|
||||
_ = os.Unsetenv(key)
|
||||
}
|
||||
}
|
||||
|
||||
// getEnvBool 获取环境变量布尔值
|
||||
func getEnvBool(key string, target *bool) {
|
||||
if v := os.Getenv(key); v != "" {
|
||||
if b, err := strconv.ParseBool(v); err == nil {
|
||||
*target = b
|
||||
}
|
||||
_ = os.Unsetenv(key)
|
||||
}
|
||||
}
|
||||
|
||||
// getEnvInt 获取环境变量整数
|
||||
func getEnvInt(key string, target *int) {
|
||||
if v := os.Getenv(key); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil {
|
||||
*target = n
|
||||
}
|
||||
_ = os.Unsetenv(key)
|
||||
}
|
||||
}
|
||||
|
||||
func LoadConfig(path string) (*AppConfig, error) {
|
||||
// 路径发现逻辑:参数优先 -> 环境变量优先 -> 默认常量
|
||||
if path == "" {
|
||||
if envPath := os.Getenv("BH_CONFIG_PATH"); envPath != "" {
|
||||
path = envPath
|
||||
} else {
|
||||
path = constant.ConfigPath
|
||||
}
|
||||
}
|
||||
|
||||
// 初始化默认配置
|
||||
Config = &AppConfig{
|
||||
Server: ServerConfig{
|
||||
Port: 8052,
|
||||
Host: "0.0.0.0",
|
||||
PprofEnabled: false,
|
||||
CookieName: "BHToken",
|
||||
},
|
||||
Database: DatabaseConfig{
|
||||
Type: "sqlite",
|
||||
Host: "localhost",
|
||||
Port: 3306,
|
||||
User: "root",
|
||||
Password: "",
|
||||
DBName: "taskpool",
|
||||
Path: constant.DefaultDBPath,
|
||||
TablePrefix: "taskpool_",
|
||||
},
|
||||
Security: SecurityConfig{
|
||||
Secret: "",
|
||||
},
|
||||
}
|
||||
|
||||
// 检查配置文件是否存在
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
// 配置文件存在,从文件加载
|
||||
logger.Infof("[Config] 从文件加载配置: %s", path)
|
||||
cfg, err := ini.Load(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := cfg.MapTo(Config); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
} else {
|
||||
// 配置文件不存在,使用环境变量
|
||||
logger.Info("[Config] 配置文件不存在,从环境变量加载")
|
||||
applyEnvOverrides()
|
||||
}
|
||||
|
||||
// 设置默认数据库路径
|
||||
if Config.Database.Path == "" {
|
||||
Config.Database.Path = constant.DefaultDBPath
|
||||
}
|
||||
// sqlite:新路径不存在时,回退旧版 data/baihu.db,避免升级后丢库
|
||||
if Config.Database.Type == "sqlite" || Config.Database.Type == "" {
|
||||
if _, err := os.Stat(Config.Database.Path); err != nil {
|
||||
legacy := strings.ReplaceAll(Config.Database.Path, "taskpool.db", "baihu.db")
|
||||
if legacy == Config.Database.Path {
|
||||
legacy = filepath.Join(constant.DataDir, "baihu.db")
|
||||
}
|
||||
if _, err2 := os.Stat(legacy); err2 == nil {
|
||||
logger.Infof("[Config] 使用旧版数据库文件: %s", legacy)
|
||||
Config.Database.Path = legacy
|
||||
}
|
||||
}
|
||||
}
|
||||
// 表前缀:若仍是默认 taskpool_ 且配置文件未显式迁移,保留兼容旧库 baihu_
|
||||
if Config.Database.TablePrefix == "" {
|
||||
Config.Database.TablePrefix = "taskpool_"
|
||||
}
|
||||
|
||||
// 设置配置到 constant 包
|
||||
if Config.Server.CookieName != "" {
|
||||
constant.CookieName = Config.Server.CookieName
|
||||
}
|
||||
constant.TablePrefix = Config.Database.TablePrefix
|
||||
constant.RuntimeDBType = Config.Database.Type
|
||||
constant.RuntimeDBHost = Config.Database.Host
|
||||
constant.RuntimeDBPort = Config.Database.Port
|
||||
constant.RuntimeDBUser = Config.Database.User
|
||||
constant.RuntimeDBPassword = Config.Database.Password
|
||||
constant.RuntimeDBName = Config.Database.DBName
|
||||
constant.RuntimeDBPath = Config.Database.Path
|
||||
constant.RuntimeDBDSN = Config.Database.DSN
|
||||
constant.RuntimeDBTablePrefix = Config.Database.TablePrefix
|
||||
constant.RuntimeDBSSLMode = Config.Database.SSLMode
|
||||
|
||||
// 暂存旧的 Secret,不再直接给 constant 赋值(改为到 settings 初始化时判断)
|
||||
// constant.Secret = Config.Security.Secret
|
||||
|
||||
// 设置演示模式
|
||||
if v := os.Getenv("BH_DEMO_MODE"); v == "true" || v == "1" {
|
||||
constant.DemoMode = true
|
||||
logger.Info("[Config] 演示模式已启用")
|
||||
_ = os.Unsetenv("BH_DEMO_MODE")
|
||||
}
|
||||
|
||||
// 输出配置信息(隐藏敏感信息)
|
||||
logger.Infof("[Config] 服务地址: %s:%d", Config.Server.Host, Config.Server.Port)
|
||||
if Config.Server.URLPrefix != "" {
|
||||
logger.Infof("[Config] URL前缀: %s", Config.Server.URLPrefix)
|
||||
}
|
||||
|
||||
maskedHost := utils.MaskString(Config.Database.Host)
|
||||
maskedDBName := utils.MaskString(Config.Database.DBName)
|
||||
logger.Infof("[Config] 数据库: type=%s, host=%s, port=%d, dbname=%s, dsn=%v",
|
||||
Config.Database.Type, maskedHost, Config.Database.Port, maskedDBName, Config.Database.DSN != "")
|
||||
|
||||
return Config, nil
|
||||
}
|
||||
|
||||
// applyEnvOverrides 从环境变量加载配置
|
||||
func applyEnvOverrides() {
|
||||
// Server
|
||||
getEnvInt("BH_SERVER_PORT", &Config.Server.Port)
|
||||
getEnvStr("BH_SERVER_HOST", &Config.Server.Host)
|
||||
getEnvStr("BH_SERVER_URL_PREFIX", &Config.Server.URLPrefix)
|
||||
getEnvBool("BH_SERVER_PPROF", &Config.Server.PprofEnabled)
|
||||
getEnvStr("BH_COOKIE_NAME", &Config.Server.CookieName)
|
||||
|
||||
// Database
|
||||
getEnvStr("BH_DB_TYPE", &Config.Database.Type)
|
||||
getEnvStr("BH_DB_HOST", &Config.Database.Host)
|
||||
getEnvInt("BH_DB_PORT", &Config.Database.Port)
|
||||
getEnvStr("BH_DB_USER", &Config.Database.User)
|
||||
getEnvStr("BH_DB_PASSWORD", &Config.Database.Password)
|
||||
getEnvStr("BH_DB_NAME", &Config.Database.DBName)
|
||||
getEnvStr("BH_DB_PATH", &Config.Database.Path)
|
||||
getEnvStr("BH_DB_DSN", &Config.Database.DSN)
|
||||
getEnvStr("BH_DB_TABLE_PREFIX", &Config.Database.TablePrefix)
|
||||
getEnvStr("BH_DB_SSL_MODE", &Config.Database.SSLMode)
|
||||
|
||||
// Security
|
||||
getEnvStr("BH_SECRET", &Config.Security.Secret)
|
||||
|
||||
}
|
||||
|
||||
func GetConfig() *AppConfig {
|
||||
return Config
|
||||
}
|
||||
@@ -0,0 +1,350 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/relation"
|
||||
"github.com/rs/xid"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type DataService struct{}
|
||||
|
||||
func NewDataService() *DataService {
|
||||
return &DataService{}
|
||||
}
|
||||
|
||||
// ExportBusinessData 智能解析依赖并导出业务数据
|
||||
func (s *DataService) ExportBusinessData(taskIDs []string, envIDs []string) *models.ExportData {
|
||||
export := models.NewExportData()
|
||||
|
||||
export.Tasks = s.collectTasksAndRelations(taskIDs)
|
||||
export.Envs = s.collectEnvironmentVariables(export.Tasks, envIDs)
|
||||
export.Bindings = s.collectNotifyBindings(export.Tasks)
|
||||
export.Tags = s.collectTagStorages(export.Tasks)
|
||||
|
||||
return export
|
||||
}
|
||||
|
||||
// collectTasksAndRelations 收集任务及其子任务,并加载关联关系 (Envs, Tags)
|
||||
func (s *DataService) collectTasksAndRelations(taskIDs []string) []models.Task {
|
||||
if len(taskIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var tasks []models.Task
|
||||
targetTaskIDs := make(map[string]bool)
|
||||
for _, id := range taskIDs {
|
||||
targetTaskIDs[id] = true
|
||||
}
|
||||
|
||||
var initialTasks []models.Task
|
||||
database.DB.Where("id IN ?", taskIDs).Find(&initialTasks)
|
||||
tasks = append(tasks, initialTasks...)
|
||||
|
||||
// 检查哪些是仓库任务,并递归查询其子任务
|
||||
var parentIDs []string
|
||||
for _, t := range initialTasks {
|
||||
if t.Type == "repo" {
|
||||
parentIDs = append(parentIDs, t.ID)
|
||||
}
|
||||
}
|
||||
|
||||
if len(parentIDs) > 0 {
|
||||
var childTasks []models.Task
|
||||
database.DB.Where("repo_task_id IN ?", parentIDs).Find(&childTasks)
|
||||
for _, ct := range childTasks {
|
||||
if !targetTaskIDs[ct.ID] {
|
||||
tasks = append(tasks, ct)
|
||||
targetTaskIDs[ct.ID] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 填充任务的 Envs 与 Tags 关系数据,以保证依赖解析和后续导入成功
|
||||
if len(tasks) > 0 {
|
||||
var exportTaskIDs []string
|
||||
for _, t := range tasks {
|
||||
exportTaskIDs = append(exportTaskIDs, t.ID)
|
||||
}
|
||||
|
||||
envsMap := relation.DataRelation.LoadRelations(exportTaskIDs, constant.RelationTypeTaskEnv)
|
||||
tagsMap := relation.DataRelation.LoadTags(exportTaskIDs, constant.RelationTypeTaskTag)
|
||||
|
||||
for i, t := range tasks {
|
||||
if envs, ok := envsMap[t.ID]; ok {
|
||||
tasks[i].Envs = models.BigText(strings.Join(envs, ","))
|
||||
}
|
||||
if tags, ok := tagsMap[t.ID]; ok {
|
||||
tasks[i].Tags = strings.Join(tags, ",")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return tasks
|
||||
}
|
||||
|
||||
// collectEnvironmentVariables 收集所需和任务所依赖的环境变量
|
||||
func (s *DataService) collectEnvironmentVariables(tasks []models.Task, envIDs []string) []models.EnvironmentVariable {
|
||||
targetEnvIDs := make(map[string]bool)
|
||||
for _, id := range envIDs {
|
||||
targetEnvIDs[id] = true
|
||||
}
|
||||
|
||||
// 从任务中解析依赖的环境变量
|
||||
for _, t := range tasks {
|
||||
if t.Envs != "" {
|
||||
envArray := strings.Split(string(t.Envs), ",")
|
||||
for _, eID := range envArray {
|
||||
eID = strings.TrimSpace(eID)
|
||||
if eID != "" {
|
||||
targetEnvIDs[eID] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(targetEnvIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var finalEnvIDs []string
|
||||
for id := range targetEnvIDs {
|
||||
finalEnvIDs = append(finalEnvIDs, id)
|
||||
}
|
||||
|
||||
var envs []models.EnvironmentVariable
|
||||
database.DB.Where("id IN ?", finalEnvIDs).Find(&envs)
|
||||
return envs
|
||||
}
|
||||
|
||||
// collectNotifyBindings 收集相关的通知规则 (NotifyBindings)
|
||||
func (s *DataService) collectNotifyBindings(tasks []models.Task) []models.NotifyBinding {
|
||||
if len(tasks) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var taskIDList []string
|
||||
for _, t := range tasks {
|
||||
taskIDList = append(taskIDList, t.ID)
|
||||
}
|
||||
|
||||
var bindings []models.NotifyBinding
|
||||
database.DB.Where("type = ? AND data_id IN ?", "task", taskIDList).Find(&bindings)
|
||||
return bindings
|
||||
}
|
||||
|
||||
// collectTagStorages 收集标签定义 (DataStorage)
|
||||
func (s *DataService) collectTagStorages(tasks []models.Task) []models.DataStorage {
|
||||
targetTagNames := make(map[string]bool)
|
||||
for _, t := range tasks {
|
||||
if t.Tags != "" {
|
||||
tagArray := strings.Split(t.Tags, ",")
|
||||
for _, tagName := range tagArray {
|
||||
tagName = strings.TrimSpace(tagName)
|
||||
if tagName != "" {
|
||||
targetTagNames[tagName] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(targetTagNames) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var finalTagNames []string
|
||||
for name := range targetTagNames {
|
||||
finalTagNames = append(finalTagNames, name)
|
||||
}
|
||||
|
||||
var tagStorages []models.DataStorage
|
||||
database.DB.Where("type = ? AND name IN ?", constant.RelationTypeTaskTag, finalTagNames).Find(&tagStorages)
|
||||
return tagStorages
|
||||
}
|
||||
|
||||
// ImportBusinessData 导入业务数据
|
||||
func (s *DataService) ImportBusinessData(data *models.ExportData) error {
|
||||
tx := database.DB.Begin()
|
||||
if tx.Error != nil {
|
||||
return tx.Error
|
||||
}
|
||||
|
||||
var adminUser models.User
|
||||
if err := tx.Where("role = ?", constant.AdminRole).First(&adminUser).Error; err != nil {
|
||||
tx.Rollback()
|
||||
return err
|
||||
}
|
||||
|
||||
importer := &businessImporter{
|
||||
tx: tx,
|
||||
adminID: adminUser.ID,
|
||||
}
|
||||
|
||||
// 1. 导入环境变量
|
||||
if err := importer.importEnvs(data.Envs); err != nil {
|
||||
tx.Rollback()
|
||||
return err
|
||||
}
|
||||
|
||||
// 2. 导入标签定义 (DataStorage)
|
||||
if err := importer.importTags(data.Tags); err != nil {
|
||||
tx.Rollback()
|
||||
return err
|
||||
}
|
||||
|
||||
// 3. 导入任务及关联映射关系
|
||||
if err := importer.importTasks(data.Tasks); err != nil {
|
||||
tx.Rollback()
|
||||
return err
|
||||
}
|
||||
|
||||
// 4. 导入通知规则
|
||||
if err := importer.importBindings(data.Bindings); err != nil {
|
||||
tx.Rollback()
|
||||
return err
|
||||
}
|
||||
|
||||
return tx.Commit().Error
|
||||
}
|
||||
|
||||
type businessImporter struct {
|
||||
tx *gorm.DB
|
||||
adminID string
|
||||
}
|
||||
|
||||
func (importer *businessImporter) importEnvs(envs []models.EnvironmentVariable) error {
|
||||
if len(envs) == 0 {
|
||||
return nil
|
||||
}
|
||||
for _, env := range envs {
|
||||
env.UserID = importer.adminID
|
||||
if err := importer.tx.Save(&env).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (importer *businessImporter) importTags(tags []models.DataStorage) error {
|
||||
if len(tags) == 0 {
|
||||
return nil
|
||||
}
|
||||
for _, tagStorage := range tags {
|
||||
if err := importer.tx.Save(&tagStorage).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (importer *businessImporter) importTasks(tasks []models.Task) error {
|
||||
if len(tasks) == 0 {
|
||||
return nil
|
||||
}
|
||||
for _, task := range tasks {
|
||||
if err := importer.tx.Save(&task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := importer.importTaskRelations(task); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (importer *businessImporter) importTaskRelations(task models.Task) error {
|
||||
if err := importer.importTaskTagRelations(task); err != nil {
|
||||
return err
|
||||
}
|
||||
return importer.importTaskEnvRelations(task)
|
||||
}
|
||||
|
||||
func (importer *businessImporter) importTaskTagRelations(task models.Task) error {
|
||||
if task.Tags == "" {
|
||||
return nil
|
||||
}
|
||||
if err := importer.tx.Where("data_id = ? AND type = ?", task.ID, constant.RelationTypeTaskTag).Delete(&models.DataRelation{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
tags := strings.Split(task.Tags, ",")
|
||||
for _, tag := range tags {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
var storage models.DataStorage
|
||||
res := importer.tx.Where("type = ? AND name = ?", constant.RelationTypeTaskTag, tag).Limit(1).Find(&storage)
|
||||
if res.Error != nil {
|
||||
return res.Error
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
storage = models.DataStorage{
|
||||
ID: xid.New().String(),
|
||||
Type: constant.RelationTypeTaskTag,
|
||||
Name: tag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
if err := importer.tx.Create(&storage).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
rel := models.DataRelation{
|
||||
ID: xid.New().String(),
|
||||
DataID: task.ID,
|
||||
RelateID: storage.ID,
|
||||
Type: constant.RelationTypeTaskTag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
if err := importer.tx.Create(&rel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (importer *businessImporter) importTaskEnvRelations(task models.Task) error {
|
||||
if string(task.Envs) == "" {
|
||||
return nil
|
||||
}
|
||||
if err := importer.tx.Where("data_id = ? AND type = ?", task.ID, constant.RelationTypeTaskEnv).Delete(&models.DataRelation{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
ids := strings.Split(string(task.Envs), ",")
|
||||
for _, relateID := range ids {
|
||||
relateID = strings.TrimSpace(relateID)
|
||||
if relateID == "" {
|
||||
continue
|
||||
}
|
||||
rel := models.DataRelation{
|
||||
ID: xid.New().String(),
|
||||
DataID: task.ID,
|
||||
RelateID: relateID,
|
||||
Type: constant.RelationTypeTaskEnv,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
if err := importer.tx.Create(&rel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (importer *businessImporter) importBindings(bindings []models.NotifyBinding) error {
|
||||
if len(bindings) == 0 {
|
||||
return nil
|
||||
}
|
||||
for _, binding := range bindings {
|
||||
if err := importer.tx.Save(&binding).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/deps"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type DependencyService struct{}
|
||||
|
||||
func NewDependencyService() *DependencyService {
|
||||
return &DependencyService{}
|
||||
}
|
||||
|
||||
// List 获取依赖列表
|
||||
func (s *DependencyService) List(language, langVersion string) ([]models.Dependency, error) {
|
||||
var results []models.Dependency
|
||||
query := database.DB
|
||||
if language != "" {
|
||||
query = query.Where("language = ?", language)
|
||||
}
|
||||
if langVersion != "" {
|
||||
query = query.Where("lang_version = ?", langVersion)
|
||||
}
|
||||
err := query.Order("id desc").Find(&results).Error
|
||||
return results, err
|
||||
}
|
||||
|
||||
// Create 创建依赖记录
|
||||
func (s *DependencyService) Create(dep *models.Dependency) error {
|
||||
// 检查是否已存在(名称、版本、语言及版本必须完全匹配)
|
||||
var existing models.Dependency
|
||||
res := database.DB.Where("name = ? AND version = ? AND language = ? AND lang_version = ?", dep.Name, dep.Version, dep.Language, dep.LangVersion).Limit(1).Find(&existing)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
// 如果已存在,更新 ID 并执行更新
|
||||
dep.ID = existing.ID
|
||||
return database.DB.Model(&existing).Updates(dep).Error
|
||||
}
|
||||
|
||||
// 不存在则新建
|
||||
if dep.ID == "" {
|
||||
dep.ID = utils.GenerateID()
|
||||
}
|
||||
return database.DB.Create(dep).Error
|
||||
}
|
||||
|
||||
// Delete 删除依赖记录
|
||||
func (s *DependencyService) Delete(id string) error {
|
||||
return database.DB.Where("id = ?", id).Delete(&models.Dependency{}).Error
|
||||
}
|
||||
|
||||
// Install 安装依赖
|
||||
func (s *DependencyService) Install(dep *models.Dependency) error {
|
||||
m := deps.GetManager(dep.Language)
|
||||
if m == nil {
|
||||
return errors.New("不支持的依赖类型: " + dep.Language)
|
||||
}
|
||||
return m.Install(dep)
|
||||
}
|
||||
|
||||
// Uninstall 卸载依赖
|
||||
func (s *DependencyService) Uninstall(dep *models.Dependency) error {
|
||||
m := deps.GetManager(dep.Language)
|
||||
if m == nil {
|
||||
return errors.New("不支持的依赖类型: " + dep.Language)
|
||||
}
|
||||
return m.Uninstall(dep)
|
||||
}
|
||||
|
||||
// GetInstalledPackages 获取已安装的包列表
|
||||
func (s *DependencyService) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
m := deps.GetManager(language)
|
||||
if m == nil {
|
||||
return nil, errors.New("不支持的依赖类型: " + language)
|
||||
}
|
||||
return m.GetInstalledPackages(language, langVersion)
|
||||
}
|
||||
|
||||
// GetInstallCommand 获取安装命令
|
||||
func (s *DependencyService) GetInstallCommand(dep *models.Dependency) (string, error) {
|
||||
m := deps.GetManager(dep.Language)
|
||||
if m == nil {
|
||||
return "", errors.New("不支持的依赖类型: " + dep.Language)
|
||||
}
|
||||
return m.GetInstallCommand(dep)
|
||||
}
|
||||
|
||||
// GetReinstallAllCommand 获取全部重装命令
|
||||
func (s *DependencyService) GetReinstallAllCommand(language, langVersion string) (string, error) {
|
||||
m := deps.GetManager(language)
|
||||
if m == nil {
|
||||
return "", errors.New("不支持的依赖类型: " + language)
|
||||
}
|
||||
|
||||
deps_list, err := s.List(language, langVersion)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return m.GetReinstallAllCommand(deps_list)
|
||||
}
|
||||
|
||||
// GetVerifyCommand 获取环境验证命令
|
||||
func (s *DependencyService) GetVerifyCommand(language, langVersion string) (string, error) {
|
||||
m := deps.GetManager(language)
|
||||
if m == nil {
|
||||
return "", errors.New("不支持的依赖类型: " + language)
|
||||
}
|
||||
return m.GetVerifyCommand(langVersion)
|
||||
}
|
||||
|
||||
// GetBatchInstallCommand 获取批量安装命令
|
||||
func (s *DependencyService) GetBatchInstallCommand(depsList []models.Dependency) (string, error) {
|
||||
if len(depsList) == 0 {
|
||||
return "", errors.New("依赖包列表不能为空")
|
||||
}
|
||||
|
||||
firstDep := depsList[0]
|
||||
m := deps.GetManager(firstDep.Language)
|
||||
if m == nil {
|
||||
return "", errors.New("不支持的依赖类型: " + firstDep.Language)
|
||||
}
|
||||
|
||||
return m.GetBatchInstallCommand(depsList)
|
||||
}
|
||||
|
||||
// ImportDependencies 批量导入依赖并自动入库去重
|
||||
func (s *DependencyService) ImportDependencies(depsList []models.Dependency) ([]models.Dependency, error) {
|
||||
var imported []models.Dependency
|
||||
for i := range depsList {
|
||||
dep := &depsList[i]
|
||||
if err := s.Create(dep); err == nil {
|
||||
imported = append(imported, *dep)
|
||||
}
|
||||
}
|
||||
return imported, nil
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type BunManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewBunManager(language string) *BunManager {
|
||||
return &BunManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"bun", "add", "-g"},
|
||||
UninstallCmd: []string{"bun", "remove", "-g"},
|
||||
ListCmd: []string{"bun", "pm", "ls", "-g"},
|
||||
Separator: "@",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *BunManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "Tool") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type CrystalManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewCrystalManager(language string) *CrystalManager {
|
||||
return &CrystalManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"shards", "install"}, // 主要是针对当前目录,但 mise exec 下可以安装 tools
|
||||
UninstallCmd: []string{"rm", "-rf"}, // 这里的逻辑可能不完全正确
|
||||
ListCmd: []string{"shards", "list"},
|
||||
Separator: " ",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *CrystalManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "Shards") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 1 {
|
||||
packages = append(packages, models.Dependency{Name: fields[1], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type DartManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewDartManager(language string) *DartManager {
|
||||
return &DartManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"dart", "pub", "global", "activate"},
|
||||
UninstallCmd: []string{"dart", "pub", "global", "deactivate"},
|
||||
ListCmd: []string{"dart", "pub", "global", "list"},
|
||||
Separator: " ",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *DartManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type DenoManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewDenoManager(language string) *DenoManager {
|
||||
return &DenoManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"deno", "install", "--global", "-A"},
|
||||
UninstallCmd: []string{"deno", "uninstall"},
|
||||
ListCmd: []string{"ls", "-1", "/root/.deno/bin"}, // 这是一个简单的 fallback 尝试
|
||||
Separator: "@",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *DenoManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
// Deno 没有官方的 list 命令,暂时留空或通过文件系统猜测
|
||||
return []models.Dependency{}, nil
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Detector 依赖检测器接口
|
||||
type Detector interface {
|
||||
Detect(logContent string) []string
|
||||
}
|
||||
|
||||
// PythonDetector Python 依赖检测器
|
||||
type PythonDetector struct{}
|
||||
|
||||
func (d *PythonDetector) Detect(logContent string) []string {
|
||||
var pkgs []string
|
||||
seen := make(map[string]bool)
|
||||
pythonRegex1 := regexp.MustCompile(`ModuleNotFoundError: No module named '([^']+)'`)
|
||||
pythonRegex2 := regexp.MustCompile(`No module named ([a-zA-Z0-9_\-]+)`)
|
||||
|
||||
matches := pythonRegex1.FindAllStringSubmatch(logContent, -1)
|
||||
for _, m := range matches {
|
||||
if len(m) > 1 {
|
||||
name := strings.TrimSpace(m[1])
|
||||
if name != "" && !seen[name] {
|
||||
seen[name] = true
|
||||
pkgs = append(pkgs, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
matches2 := pythonRegex2.FindAllStringSubmatch(logContent, -1)
|
||||
for _, m := range matches2 {
|
||||
if len(m) > 1 {
|
||||
name := strings.TrimSpace(m[1])
|
||||
if name != "" && !seen[name] {
|
||||
seen[name] = true
|
||||
pkgs = append(pkgs, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
return pkgs
|
||||
}
|
||||
|
||||
// NodeDetector Node.js 依赖检测器
|
||||
type NodeDetector struct{}
|
||||
|
||||
func (d *NodeDetector) Detect(logContent string) []string {
|
||||
var pkgs []string
|
||||
seen := make(map[string]bool)
|
||||
nodeRegex1 := regexp.MustCompile(`Error: Cannot find module '([^']+)'`)
|
||||
nodeRegex2 := regexp.MustCompile(`Cannot find module '([^']+)'`)
|
||||
|
||||
matches := nodeRegex1.FindAllStringSubmatch(logContent, -1)
|
||||
for _, m := range matches {
|
||||
if len(m) > 1 {
|
||||
name := strings.TrimSpace(m[1])
|
||||
if name != "" && !seen[name] {
|
||||
seen[name] = true
|
||||
pkgs = append(pkgs, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
matches2 := nodeRegex2.FindAllStringSubmatch(logContent, -1)
|
||||
for _, m := range matches2 {
|
||||
if len(m) > 1 {
|
||||
name := strings.TrimSpace(m[1])
|
||||
if name != "" && !seen[name] {
|
||||
seen[name] = true
|
||||
pkgs = append(pkgs, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
return pkgs
|
||||
}
|
||||
|
||||
var languageDetectors = map[string]Detector{
|
||||
"python": &PythonDetector{},
|
||||
"python3": &PythonDetector{},
|
||||
"node": &NodeDetector{},
|
||||
"js": &NodeDetector{},
|
||||
"ts": &NodeDetector{},
|
||||
"bun": &NodeDetector{},
|
||||
}
|
||||
|
||||
// DetectMissingDependencies 从日志内容中检测缺失的依赖包名
|
||||
func DetectMissingDependencies(language, logContent string) ([]string, bool) {
|
||||
lang := strings.ToLower(language)
|
||||
var det Detector
|
||||
for key, d := range languageDetectors {
|
||||
if strings.Contains(lang, key) {
|
||||
det = d
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if det == nil {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
pkgs := det.Detect(logContent)
|
||||
return pkgs, len(pkgs) > 0
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDetectMissingDependencies(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
language string
|
||||
logContent string
|
||||
wantPkgs []string
|
||||
wantFound bool
|
||||
}{
|
||||
{
|
||||
name: "Python ModuleNotFoundError",
|
||||
language: "python3",
|
||||
logContent: "Traceback (most recent call last):\n File \"main.py\", line 1, in <module>\n import requests\nModuleNotFoundError: No module named 'requests'",
|
||||
wantPkgs: []string{"requests"},
|
||||
wantFound: true,
|
||||
},
|
||||
{
|
||||
name: "Python No module named",
|
||||
language: "python",
|
||||
logContent: "ImportError: No module named yaml",
|
||||
wantPkgs: []string{"yaml"},
|
||||
wantFound: true,
|
||||
},
|
||||
{
|
||||
name: "Node Error Cannot find module",
|
||||
language: "node",
|
||||
logContent: "Error: Cannot find module 'axios'\nRequire stack:\n- /app/index.js",
|
||||
wantPkgs: []string{"axios"},
|
||||
wantFound: true,
|
||||
},
|
||||
{
|
||||
name: "No match",
|
||||
language: "python",
|
||||
logContent: "Success running script",
|
||||
wantPkgs: nil,
|
||||
wantFound: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
gotPkgs, gotFound := DetectMissingDependencies(tt.language, tt.logContent)
|
||||
if gotFound != tt.wantFound {
|
||||
t.Errorf("DetectMissingDependencies() gotFound = %v, want %v", gotFound, tt.wantFound)
|
||||
}
|
||||
if len(gotPkgs) != len(tt.wantPkgs) {
|
||||
t.Errorf("DetectMissingDependencies() gotPkgs = %v, want %v", gotPkgs, tt.wantPkgs)
|
||||
return
|
||||
}
|
||||
for i, p := range gotPkgs {
|
||||
if p != tt.wantPkgs[i] {
|
||||
t.Errorf("DetectMissingDependencies() gotPkgs[%d] = %v, want %v", i, p, tt.wantPkgs[i])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type DotnetManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewDotnetManager(language string) *DotnetManager {
|
||||
return &DotnetManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"dotnet", "tool", "install", "-g"},
|
||||
UninstallCmd: []string{"dotnet", "tool", "uninstall", "-g"},
|
||||
ListCmd: []string{"dotnet", "tool", "list", "-g"},
|
||||
Separator: " ",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *DotnetManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for i, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if i < 2 || line == "" { // 跳过表头
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language, Version: fields[1]})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type ElixirManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewElixirManager(language string) *ElixirManager {
|
||||
verifyCmd := []string{"elixir", "-v"}
|
||||
if strings.Contains(strings.ToLower(language), "erlang") {
|
||||
verifyCmd = []string{"erl", "+V"}
|
||||
}
|
||||
return &ElixirManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"mix", "archive.install", "hex", "--force"},
|
||||
UninstallCmd: []string{"mix", "archive.uninstall"},
|
||||
ListCmd: []string{"mix", "archive"},
|
||||
VerifyCmd: verifyCmd,
|
||||
Separator: " ",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *ElixirManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "*") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type GoManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewGoManager(language string) *GoManager {
|
||||
return &GoManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"go", "install"},
|
||||
UninstallCmd: []string{"go", "clean", "-i"},
|
||||
ListCmd: []string{"go", "list", "..."},
|
||||
VerifyCmd: []string{"go", "version"},
|
||||
Separator: "@",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *GoManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "Tool") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type LuaManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewLuaManager(language string) *LuaManager {
|
||||
return &LuaManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"luarocks", "install"},
|
||||
UninstallCmd: []string{"luarocks", "remove"},
|
||||
ListCmd: []string{"luarocks", "list"},
|
||||
VerifyCmd: []string{"lua", "-v"},
|
||||
Separator: " ",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *LuaManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "Rock") || strings.HasPrefix(line, "--") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os/exec"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// Manager 依赖管理器接口
|
||||
type Manager interface {
|
||||
Install(dep *models.Dependency) error
|
||||
Uninstall(dep *models.Dependency) error
|
||||
GetInstalledPackages(language, langVersion string) ([]models.Dependency, error)
|
||||
GetInstallCommand(dep *models.Dependency) (string, error)
|
||||
GetBatchInstallCommand(deps []models.Dependency) (string, error)
|
||||
GetReinstallAllCommand(deps []models.Dependency) (string, error)
|
||||
GetVerifyCommand(langVersion string) (string, error)
|
||||
}
|
||||
|
||||
// BaseManager 基础管理器,提供通用方法
|
||||
type BaseManager struct {
|
||||
Language string
|
||||
InstallCmd []string
|
||||
UninstallCmd []string
|
||||
ListCmd []string
|
||||
VerifyCmd []string
|
||||
Separator string
|
||||
}
|
||||
|
||||
func (m *BaseManager) runMiseCommand(langVersion string, cmdArgs []string) ([]byte, error) {
|
||||
args := utils.BuildMiseCommandArgsSimple(cmdArgs, m.Language, langVersion)
|
||||
cmd := exec.Command(args[0], args[1:]...)
|
||||
return cmd.CombinedOutput()
|
||||
}
|
||||
|
||||
func (m *BaseManager) Install(dep *models.Dependency) error {
|
||||
var packageSpec string
|
||||
if dep.Version != "" {
|
||||
packageSpec = dep.Name + m.Separator + dep.Version
|
||||
} else {
|
||||
packageSpec = dep.Name
|
||||
}
|
||||
|
||||
args := append([]string{}, m.InstallCmd...)
|
||||
args = append(args, packageSpec)
|
||||
|
||||
logger.Infof("Installing %s package: %s", m.Language, packageSpec)
|
||||
output, err := m.runMiseCommand(dep.LangVersion, args)
|
||||
dep.Log = models.BigText(output)
|
||||
|
||||
if err != nil {
|
||||
logger.Errorf("Install failed: %v, output: %s", err, string(output))
|
||||
return errors.New("安装失败: " + string(output))
|
||||
}
|
||||
logger.Infof("Install success: %s", packageSpec)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *BaseManager) GetInstallCommand(dep *models.Dependency) (string, error) {
|
||||
var packageSpec string
|
||||
if dep.Version != "" {
|
||||
packageSpec = dep.Name + m.Separator + dep.Version
|
||||
} else {
|
||||
packageSpec = dep.Name
|
||||
}
|
||||
|
||||
args := append([]string{}, m.InstallCmd...)
|
||||
args = append(args, packageSpec)
|
||||
|
||||
fullCmd := utils.BuildMiseCommandSimple(strings.Join(args, " "), m.Language, dep.LangVersion)
|
||||
return fullCmd + " && echo \"__INSTALL_SUCCESS__\" || echo \"__INSTALL_FAILED__\"", nil
|
||||
}
|
||||
|
||||
func (m *BaseManager) GetBatchInstallCommand(deps []models.Dependency) (string, error) {
|
||||
if len(deps) == 0 {
|
||||
return "echo \"没有需要安装的依赖\"", nil
|
||||
}
|
||||
|
||||
var packageSpecs []string
|
||||
var langVersion string
|
||||
var language string
|
||||
for _, dep := range deps {
|
||||
language = dep.Language
|
||||
if dep.LangVersion != "" {
|
||||
langVersion = dep.LangVersion
|
||||
}
|
||||
if dep.Version != "" {
|
||||
packageSpecs = append(packageSpecs, dep.Name+m.Separator+dep.Version)
|
||||
} else {
|
||||
packageSpecs = append(packageSpecs, dep.Name)
|
||||
}
|
||||
}
|
||||
|
||||
args := append([]string{}, m.InstallCmd...)
|
||||
args = append(args, packageSpecs...)
|
||||
|
||||
fullCmd := utils.BuildMiseCommandSimple(strings.Join(args, " "), language, langVersion)
|
||||
return fullCmd + " && echo \"__INSTALL_SUCCESS__\" || echo \"__INSTALL_FAILED__\"", nil
|
||||
}
|
||||
|
||||
func (m *BaseManager) GetReinstallAllCommand(deps []models.Dependency) (string, error) {
|
||||
if len(deps) == 0 {
|
||||
return "echo \"没有需要安装的依赖\"", nil
|
||||
}
|
||||
|
||||
var packageSpecs []string
|
||||
var langVersion string
|
||||
for _, dep := range deps {
|
||||
if dep.LangVersion != "" {
|
||||
langVersion = dep.LangVersion
|
||||
}
|
||||
if dep.Version != "" {
|
||||
packageSpecs = append(packageSpecs, dep.Name+m.Separator+dep.Version)
|
||||
} else {
|
||||
packageSpecs = append(packageSpecs, dep.Name)
|
||||
}
|
||||
}
|
||||
|
||||
args := append([]string{}, m.InstallCmd...)
|
||||
args = append(args, packageSpecs...)
|
||||
|
||||
fullCmd := utils.BuildMiseCommandSimple(strings.Join(args, " "), m.Language, langVersion)
|
||||
return fullCmd + " && echo \"__INSTALL_SUCCESS__\" || echo \"__INSTALL_FAILED__\"", nil
|
||||
}
|
||||
|
||||
func (m *BaseManager) GetVerifyCommand(langVersion string) (string, error) {
|
||||
var cmd string
|
||||
if len(m.VerifyCmd) > 0 {
|
||||
cmd = strings.Join(m.VerifyCmd, " ")
|
||||
} else {
|
||||
cmd = m.Language + " --version"
|
||||
}
|
||||
return utils.BuildMiseCommandSimple(cmd, m.Language, langVersion), nil
|
||||
}
|
||||
|
||||
func (m *BaseManager) Uninstall(dep *models.Dependency) error {
|
||||
args := append([]string{}, m.UninstallCmd...)
|
||||
args = append(args, dep.Name)
|
||||
|
||||
logger.Infof("Uninstalling %s package: %s", m.Language, dep.Name)
|
||||
output, err := m.runMiseCommand(dep.LangVersion, args)
|
||||
if err != nil {
|
||||
logger.Errorf("Uninstall failed: %v, output: %s", err, string(output))
|
||||
return errors.New("卸载失败: " + string(output))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetManager 根据语言获取对应的管理器
|
||||
func GetManager(language string) Manager {
|
||||
lang := strings.ToLower(language)
|
||||
if strings.Contains(lang, "python") {
|
||||
return NewPythonManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "node") {
|
||||
return NewNodeManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "ruby") {
|
||||
return NewRubyManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "go") {
|
||||
return NewGoManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "rust") {
|
||||
return NewRustManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "bun") {
|
||||
return NewBunManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "php") {
|
||||
return NewPhpManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "deno") {
|
||||
return NewDenoManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "dotnet") {
|
||||
return NewDotnetManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "elixir") || strings.Contains(lang, "erlang") {
|
||||
return NewElixirManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "lua") {
|
||||
return NewLuaManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "nim") {
|
||||
return NewNimManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "dart") || strings.Contains(lang, "flutter") {
|
||||
return NewDartManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "perl") {
|
||||
return NewPerlManager(language)
|
||||
}
|
||||
if strings.Contains(lang, "crystal") {
|
||||
return NewCrystalManager(language)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type NimManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewNimManager(language string) *NimManager {
|
||||
return &NimManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"nimble", "install", "-y"},
|
||||
UninstallCmd: []string{"nimble", "uninstall", "-y"},
|
||||
ListCmd: []string{"nimble", "list", "-i"},
|
||||
Separator: "@",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *NimManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || !strings.Contains(line, "[") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type NodeManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewNodeManager(language string) *NodeManager {
|
||||
return &NodeManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"npm", "install", "-g"},
|
||||
UninstallCmd: []string{"npm", "uninstall", "-g"},
|
||||
ListCmd: []string{"npm", "list", "-g", "--depth=0", "--json"},
|
||||
VerifyCmd: []string{"node", "-v"},
|
||||
Separator: "@",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *NodeManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if strings.Contains(line, `"version"`) || strings.HasPrefix(line, "{") || strings.HasPrefix(line, "}") || strings.HasPrefix(line, "]") {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(line, `"`) && strings.Contains(line, ":") {
|
||||
name := strings.Trim(strings.Split(line, ":")[0], `" `)
|
||||
if name != "" && name != "dependencies" && name != "name" {
|
||||
packages = append(packages, models.Dependency{Name: name, Language: language})
|
||||
}
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
// ParseManifest 根据语言解析依赖清单文件内容
|
||||
func ParseManifest(language, content string) ([]models.Dependency, error) {
|
||||
lang := strings.ToLower(language)
|
||||
if strings.Contains(lang, "python") {
|
||||
return ParseRequirements(content), nil
|
||||
}
|
||||
if strings.Contains(lang, "node") {
|
||||
return ParsePackageJson(content)
|
||||
}
|
||||
return []models.Dependency{}, nil
|
||||
}
|
||||
|
||||
// ParseRequirements 解析 Python requirements.txt
|
||||
func ParseRequirements(content string) []models.Dependency {
|
||||
var deps []models.Dependency
|
||||
// 使用正则表达式按行分割,兼容 Windows 和 Linux 的换行符
|
||||
lines := regexp.MustCompile(`\r?\n`).Split(content, -1)
|
||||
|
||||
// 用于分割包名和版本的正则 (支持 ==, >=, <=, ~=, >, <, @)
|
||||
versionRegex := regexp.MustCompile(`[=><~@]+`)
|
||||
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
// 忽略空行、注释行以及参数行 (以 - 开头的行如 -i, -r)
|
||||
if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, "-") {
|
||||
continue
|
||||
}
|
||||
|
||||
// 分割名称与版本
|
||||
parts := versionRegex.Split(line, 2)
|
||||
name := strings.TrimSpace(parts[0])
|
||||
version := ""
|
||||
if len(parts) > 1 {
|
||||
// 清除可能存在的后续参数,比如 requests==2.31.0 --hash=sha256:...
|
||||
versionPart := strings.TrimSpace(parts[1])
|
||||
|
||||
// 如果有逗号分隔的多个范围限制,比如 >=1.20,<2.0,只取第一个范围作为参考版本号
|
||||
if idx := strings.Index(versionPart, ","); idx != -1 {
|
||||
versionPart = versionPart[:idx]
|
||||
}
|
||||
|
||||
versionFields := strings.Fields(versionPart)
|
||||
if len(versionFields) > 0 {
|
||||
version = strings.TrimSpace(versionFields[0])
|
||||
// 清除可能残留的首部版本符号
|
||||
version = strings.TrimLeft(version, "=><~@ ")
|
||||
}
|
||||
}
|
||||
|
||||
if name != "" {
|
||||
deps = append(deps, models.Dependency{
|
||||
Name: name,
|
||||
Version: version,
|
||||
Language: "python3",
|
||||
})
|
||||
}
|
||||
}
|
||||
return deps
|
||||
}
|
||||
|
||||
// PackageJson 代表 package.json 的结构定义
|
||||
type PackageJson struct {
|
||||
Dependencies map[string]string `json:"dependencies"`
|
||||
DevDependencies map[string]string `json:"devDependencies"`
|
||||
}
|
||||
|
||||
// ParsePackageJson 解析 Node.js package.json
|
||||
func ParsePackageJson(content string) ([]models.Dependency, error) {
|
||||
var pkg PackageJson
|
||||
if err := json.Unmarshal([]byte(content), &pkg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var deps []models.Dependency
|
||||
collect := func(m map[string]string, isDev bool) {
|
||||
for name, versionRange := range m {
|
||||
// 移除 npm 常见版本范围修饰符(如 ^1.2.3 或 ~2.3.0,保留底线版本号)
|
||||
version := strings.TrimLeft(versionRange, "^~>=<* ")
|
||||
remark := ""
|
||||
if isDev {
|
||||
remark = "devDependencies"
|
||||
}
|
||||
deps = append(deps, models.Dependency{
|
||||
Name: name,
|
||||
Version: version,
|
||||
Language: "node",
|
||||
Remark: remark,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
collect(pkg.Dependencies, false)
|
||||
collect(pkg.DevDependencies, true)
|
||||
return deps, nil
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseRequirements(t *testing.T) {
|
||||
content := `
|
||||
# This is a comment
|
||||
requests==2.31.0
|
||||
numpy>=1.20,<2.0
|
||||
gunicorn
|
||||
-r other-requirements.txt
|
||||
pandas ~= 1.3.0
|
||||
`
|
||||
deps := ParseRequirements(content)
|
||||
if len(deps) != 4 {
|
||||
t.Fatalf("expected 4 dependencies, got %d", len(deps))
|
||||
}
|
||||
|
||||
expected := []struct {
|
||||
name string
|
||||
version string
|
||||
}{
|
||||
{"requests", "2.31.0"},
|
||||
{"numpy", "1.20"},
|
||||
{"gunicorn", ""},
|
||||
{"pandas", "1.3.0"},
|
||||
}
|
||||
|
||||
for i, exp := range expected {
|
||||
if deps[i].Name != exp.name {
|
||||
t.Errorf("expected name %s, got %s", exp.name, deps[i].Name)
|
||||
}
|
||||
if deps[i].Version != exp.version {
|
||||
t.Errorf("expected version %s, got %s", exp.version, deps[i].Version)
|
||||
}
|
||||
if deps[i].Language != "python3" {
|
||||
t.Errorf("expected language python3, got %s", deps[i].Language)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParsePackageJson(t *testing.T) {
|
||||
content := `{
|
||||
"dependencies": {
|
||||
"lodash": "^4.17.21",
|
||||
"express": "~4.18.2"
|
||||
},
|
||||
"devDependencies": {
|
||||
"typescript": "^5.0.4"
|
||||
}
|
||||
}`
|
||||
deps, err := ParsePackageJson(content)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to parse package.json: %v", err)
|
||||
}
|
||||
|
||||
if len(deps) != 3 {
|
||||
t.Fatalf("expected 3 dependencies, got %d", len(deps))
|
||||
}
|
||||
|
||||
expected := []struct {
|
||||
name string
|
||||
version string
|
||||
remark string
|
||||
}{
|
||||
{"lodash", "4.17.21", ""},
|
||||
{"express", "4.18.2", ""},
|
||||
{"typescript", "5.0.4", "devDependencies"},
|
||||
}
|
||||
|
||||
for i, exp := range expected {
|
||||
if deps[i].Name != exp.name {
|
||||
t.Errorf("expected name %s, got %s", exp.name, deps[i].Name)
|
||||
}
|
||||
if deps[i].Version != exp.version {
|
||||
t.Errorf("expected version %s, got %s", exp.version, deps[i].Version)
|
||||
}
|
||||
if deps[i].Remark != exp.remark {
|
||||
t.Errorf("expected remark %s, got %s", exp.remark, deps[i].Remark)
|
||||
}
|
||||
if deps[i].Language != "node" {
|
||||
t.Errorf("expected language node, got %s", deps[i].Language)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type PerlManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewPerlManager(language string) *PerlManager {
|
||||
return &PerlManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"cpanm"}, // 需要系统预装 cpanm
|
||||
UninstallCmd: []string{"cpanm", "--uninstall"}, // 部分 cpanm 版本不支持,这是一个占位
|
||||
ListCmd: []string{"perldoc", "-l"}, // 很难列出所有,暂时简单处理
|
||||
VerifyCmd: []string{"perl", "-v"},
|
||||
Separator: " ",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *PerlManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
// Perl 依赖列出比较复杂,暂时返回空
|
||||
return []models.Dependency{}, nil
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type PhpManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewPhpManager(language string) *PhpManager {
|
||||
return &PhpManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"composer", "global", "require"},
|
||||
UninstallCmd: []string{"composer", "global", "remove"},
|
||||
ListCmd: []string{"composer", "global", "show", "--name-only"},
|
||||
VerifyCmd: []string{"php", "-v"},
|
||||
Separator: ":",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *PhpManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "Tool") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type PythonManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewPythonManager(language string) *PythonManager {
|
||||
return &PythonManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"pip", "install"},
|
||||
UninstallCmd: []string{"pip", "uninstall", "-y"},
|
||||
ListCmd: []string{"pip", "list", "--format=freeze"},
|
||||
Separator: "==",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *PythonManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
parts := strings.SplitN(line, "==", 2)
|
||||
pkg := models.Dependency{Name: parts[0], Language: language}
|
||||
if len(parts) > 1 {
|
||||
pkg.Version = parts[1]
|
||||
}
|
||||
packages = append(packages, pkg)
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type RubyManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewRubyManager(language string) *RubyManager {
|
||||
return &RubyManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"gem", "install"},
|
||||
UninstallCmd: []string{"gem", "uninstall", "-a", "-x"},
|
||||
ListCmd: []string{"gem", "list", "--local"},
|
||||
VerifyCmd: []string{"ruby", "-v"},
|
||||
Separator: " ",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *RubyManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "Tool") || strings.HasPrefix(line, "(") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package deps
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
type RustManager struct {
|
||||
BaseManager
|
||||
}
|
||||
|
||||
func NewRustManager(language string) *RustManager {
|
||||
return &RustManager{
|
||||
BaseManager: BaseManager{
|
||||
Language: language,
|
||||
InstallCmd: []string{"cargo", "install"},
|
||||
UninstallCmd: []string{"cargo", "uninstall"},
|
||||
ListCmd: []string{"cargo", "install", "--list"},
|
||||
VerifyCmd: []string{"rustc", "--version"},
|
||||
Separator: " v",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (m *RustManager) GetInstalledPackages(language, langVersion string) ([]models.Dependency, error) {
|
||||
output, err := m.runMiseCommand(langVersion, m.ListCmd)
|
||||
if err != nil {
|
||||
logger.Warnf("GetInstalledPackages for %s failed: %v", language, err)
|
||||
}
|
||||
|
||||
var packages []models.Dependency
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || !strings.Contains(line, " v") {
|
||||
continue
|
||||
}
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) > 0 {
|
||||
packages = append(packages, models.Dependency{Name: fields[0], Language: language})
|
||||
}
|
||||
}
|
||||
return packages, nil
|
||||
}
|
||||
@@ -0,0 +1,409 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/relation"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type EnvService struct{}
|
||||
|
||||
func NewEnvService() *EnvService {
|
||||
return &EnvService{}
|
||||
}
|
||||
|
||||
func (es *EnvService) CreateEnvVar(name, value, remark, envType string, hidden, enabled bool, userID string) *models.EnvironmentVariable {
|
||||
if envType == constant.EnvTypeSecret {
|
||||
if encValue, err := utils.Encrypt(value); err == nil {
|
||||
value = encValue
|
||||
}
|
||||
}
|
||||
|
||||
env := &models.EnvironmentVariable{
|
||||
ID: utils.GenerateID(),
|
||||
Name: name,
|
||||
Value: models.BigText(value),
|
||||
Remark: remark,
|
||||
Type: envType,
|
||||
Hidden: &hidden,
|
||||
Enabled: &enabled,
|
||||
UserID: userID,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
database.DB.Select("*").Create(env)
|
||||
return env
|
||||
}
|
||||
|
||||
func (es *EnvService) GetEnvVarsByUserID(userID string) []models.EnvironmentVariable {
|
||||
var envs []models.EnvironmentVariable
|
||||
database.DB.Where("user_id = ?", userID).Find(&envs)
|
||||
es.LoadEnvTags(envs)
|
||||
return envs
|
||||
}
|
||||
|
||||
// GetFormattedEnvVarsByUserID 获取用户环境变量并格式化为 NAME=VALUE 格式(支持重名合并)
|
||||
func (es *EnvService) GetFormattedEnvVarsByUserID(userID string) []string {
|
||||
envs := es.GetEnvVarsByUserID(userID)
|
||||
return es.formatEnvVars(envs)
|
||||
}
|
||||
|
||||
func (es *EnvService) GetEnvVarsWithPagination(userID string, name string, envType string, tags string, page, pageSize int) ([]models.EnvironmentVariable, int64) {
|
||||
var envs []models.EnvironmentVariable
|
||||
var total int64
|
||||
|
||||
query := database.DB.Model(&models.EnvironmentVariable{}).Where("user_id = ?", userID)
|
||||
if name != "" {
|
||||
query = query.Where("name LIKE ? OR remark LIKE ?", "%"+name+"%", "%"+name+"%")
|
||||
}
|
||||
if envType != "" {
|
||||
query = query.Where("type = ?", envType)
|
||||
}
|
||||
|
||||
if tags != "" {
|
||||
tagList := strings.Split(tags, ",")
|
||||
var validTags []string
|
||||
for _, t := range tagList {
|
||||
t = strings.TrimSpace(t)
|
||||
if t != "" {
|
||||
validTags = append(validTags, t)
|
||||
}
|
||||
}
|
||||
if len(validTags) > 0 {
|
||||
var storageIDs []string
|
||||
database.DB.Model(&models.DataStorage{}).Where("type = ? AND name IN ?", constant.RelationTypeEnvTag, validTags).Pluck("id", &storageIDs)
|
||||
|
||||
var envIDs []string
|
||||
if len(storageIDs) > 0 {
|
||||
database.DB.Model(&models.DataRelation{}).Where("type = ? AND relate_id IN ?", constant.RelationTypeEnvTag, storageIDs).Pluck("data_id", &envIDs)
|
||||
}
|
||||
|
||||
if len(envIDs) > 0 {
|
||||
query = query.Where("id IN ?", envIDs)
|
||||
} else {
|
||||
query = query.Where("1 = 0")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
query.Count(&total)
|
||||
query.Order("id DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&envs)
|
||||
es.LoadEnvTags(envs)
|
||||
return envs, total
|
||||
}
|
||||
|
||||
func (es *EnvService) GetEnvVarByID(id string) *models.EnvironmentVariable {
|
||||
var env models.EnvironmentVariable
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&env)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
envs := []models.EnvironmentVariable{env}
|
||||
es.LoadEnvTags(envs)
|
||||
return &envs[0]
|
||||
}
|
||||
|
||||
func (es *EnvService) UpdateEnvVar(id string, name, value, remark, envType string, hidden, enabled bool) *models.EnvironmentVariable {
|
||||
var env models.EnvironmentVariable
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&env)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
if envType == constant.EnvTypeSecret && value != "********" && value != "" {
|
||||
if encValue, err := utils.Encrypt(value); err == nil {
|
||||
value = encValue
|
||||
}
|
||||
} else if envType == constant.EnvTypeSecret && (value == "********" || value == "") {
|
||||
// Keep the original encrypted value
|
||||
value = string(env.Value)
|
||||
}
|
||||
|
||||
updates := map[string]interface{}{
|
||||
"name": name,
|
||||
"value": models.BigText(value),
|
||||
"remark": remark,
|
||||
"type": envType,
|
||||
"hidden": &hidden,
|
||||
"enabled": &enabled,
|
||||
}
|
||||
database.DB.Model(&env).Updates(updates)
|
||||
return &env
|
||||
}
|
||||
|
||||
func (es *EnvService) GetAssociatedTasks(id string) []models.Task {
|
||||
var associatedTasks []models.Task
|
||||
var taskIDs []string
|
||||
database.DB.Model(&models.DataRelation{}).Where("type = ? AND relate_id = ?", constant.RelationTypeTaskEnv, id).Pluck("data_id", &taskIDs)
|
||||
if len(taskIDs) > 0 {
|
||||
database.DB.Where("id IN ?", taskIDs).Find(&associatedTasks)
|
||||
}
|
||||
return associatedTasks
|
||||
}
|
||||
|
||||
func (es *EnvService) DeleteEnvVar(id string, force bool) (bool, []models.Task) {
|
||||
associatedTasks := es.GetAssociatedTasks(id)
|
||||
|
||||
if len(associatedTasks) > 0 && !force {
|
||||
return false, associatedTasks
|
||||
}
|
||||
|
||||
if force {
|
||||
err := database.DB.Transaction(func(tx *gorm.DB) error {
|
||||
// Delete the relations mapping this env to any tasks
|
||||
if err := tx.Where("type = ? AND relate_id = ?", constant.RelationTypeTaskEnv, id).Delete(&models.DataRelation{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
// Delete the env var
|
||||
if err := tx.Where("id = ?", id).Delete(&models.EnvironmentVariable{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err == nil {
|
||||
relation.DataRelation.CleanRelations(id, constant.RelationTypeEnvTag)
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
result := database.DB.Where("id = ?", id).Delete(&models.EnvironmentVariable{})
|
||||
if result.RowsAffected > 0 {
|
||||
relation.DataRelation.CleanRelations(id, constant.RelationTypeEnvTag)
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// GetEnvVarsByIDs 根据逗号分隔的ID字符串获取环境变量列表,返回 NAME=VALUE 格式
|
||||
// 如果存在重名变量,会类似青龙面板一样使用 & 拼接
|
||||
func (es *EnvService) GetEnvVarsByIDs(envIDs string) []string {
|
||||
if envIDs == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
ids := splitEnvIDs(envIDs)
|
||||
var envs []models.EnvironmentVariable
|
||||
for _, id := range ids {
|
||||
env := es.GetEnvVarByID(id)
|
||||
if env != nil {
|
||||
envs = append(envs, *env)
|
||||
}
|
||||
}
|
||||
|
||||
return es.formatEnvVars(envs)
|
||||
}
|
||||
|
||||
// GetEnvVarsAndSecretsByIDs 根据逗号分隔的ID字符串获取环境变量列表和安全机密值列表
|
||||
func (es *EnvService) GetEnvVarsAndSecretsByIDs(envIDs string) ([]string, []string) {
|
||||
if envIDs == "" {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
ids := splitEnvIDs(envIDs)
|
||||
var envs []models.EnvironmentVariable
|
||||
for _, id := range ids {
|
||||
env := es.GetEnvVarByID(id)
|
||||
if env != nil {
|
||||
envs = append(envs, *env)
|
||||
}
|
||||
}
|
||||
|
||||
return es.formatEnvVarsAndSecrets(envs)
|
||||
}
|
||||
|
||||
// GetAllEnvVars获取系统中所有的环境变量,并按 NAME=VALUE 格式返回
|
||||
func (es *EnvService) GetAllEnvVars() []string {
|
||||
var envs []models.EnvironmentVariable
|
||||
if err := database.DB.Find(&envs).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return es.formatEnvVars(envs)
|
||||
}
|
||||
|
||||
// GetAllEnvVarsAndSecrets 获取系统中所有的环境变量和安全机密值列表
|
||||
func (es *EnvService) GetAllEnvVarsAndSecrets() ([]string, []string) {
|
||||
var envs []models.EnvironmentVariable
|
||||
if err := database.DB.Find(&envs).Error; err != nil {
|
||||
return nil, nil
|
||||
}
|
||||
return es.formatEnvVarsAndSecrets(envs)
|
||||
}
|
||||
|
||||
// formatEnvVars 将环境变量列表格式化为 NAME=VALUE 数组,并处理重名合并 (过滤掉所有的 Secret)
|
||||
func (es *EnvService) formatEnvVars(envs []models.EnvironmentVariable) []string {
|
||||
if len(envs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
type mergedEnv struct {
|
||||
name string
|
||||
values []string
|
||||
}
|
||||
var mergedList []mergedEnv
|
||||
nameToIndex := make(map[string]int)
|
||||
|
||||
for _, env := range envs {
|
||||
// 非调度器入口,直接当做没有(跳过 Secret)
|
||||
if env.Type == constant.EnvTypeSecret {
|
||||
continue
|
||||
}
|
||||
|
||||
value := string(env.Value)
|
||||
if !utils.DerefBool(env.Enabled, true) {
|
||||
value = ""
|
||||
}
|
||||
|
||||
if idx, ok := nameToIndex[env.Name]; ok {
|
||||
mergedList[idx].values = append(mergedList[idx].values, value)
|
||||
} else {
|
||||
nameToIndex[env.Name] = len(mergedList)
|
||||
mergedList = append(mergedList, mergedEnv{
|
||||
name: env.Name,
|
||||
values: []string{value},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
var result []string
|
||||
for _, item := range mergedList {
|
||||
val := strings.Join(item.values, "&")
|
||||
result = append(result, item.name+"="+val)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// formatEnvVarsAndSecrets 将环境变量列表格式化为 NAME=VALUE 数组,并提取明文安全机密列表
|
||||
func (es *EnvService) formatEnvVarsAndSecrets(envs []models.EnvironmentVariable) ([]string, []string) {
|
||||
if len(envs) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
type mergedEnv struct {
|
||||
name string
|
||||
values []string
|
||||
}
|
||||
var mergedList []mergedEnv
|
||||
var secrets []string
|
||||
nameToIndex := make(map[string]int)
|
||||
|
||||
for _, env := range envs {
|
||||
value := string(env.Value)
|
||||
if env.Type == constant.EnvTypeSecret {
|
||||
if decValue, err := utils.Decrypt(value); err == nil {
|
||||
value = decValue
|
||||
if utils.DerefBool(env.Enabled, true) && value != "" {
|
||||
secrets = append(secrets, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if !utils.DerefBool(env.Enabled, true) {
|
||||
value = ""
|
||||
}
|
||||
|
||||
if idx, ok := nameToIndex[env.Name]; ok {
|
||||
mergedList[idx].values = append(mergedList[idx].values, value)
|
||||
} else {
|
||||
nameToIndex[env.Name] = len(mergedList)
|
||||
mergedList = append(mergedList, mergedEnv{
|
||||
name: env.Name,
|
||||
values: []string{value},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
var result []string
|
||||
for _, item := range mergedList {
|
||||
// 多个值使用 & 拼接
|
||||
val := strings.Join(item.values, "&")
|
||||
result = append(result, item.name+"="+val)
|
||||
}
|
||||
return result, secrets
|
||||
}
|
||||
|
||||
// splitEnvIDs 解析逗号分隔的ID字符串
|
||||
func splitEnvIDs(envIDs string) []string {
|
||||
var ids []string
|
||||
for _, s := range strings.Split(envIDs, ",") {
|
||||
s = strings.TrimSpace(s)
|
||||
if s != "" {
|
||||
ids = append(ids, s)
|
||||
}
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// SaveEnvTags 保存环境变量标签
|
||||
func (es *EnvService) SaveEnvTags(envID string, tagsStr string) {
|
||||
database.DB.Where("data_id = ? AND type = ?", envID, constant.RelationTypeEnvTag).Delete(&models.DataRelation{})
|
||||
if tagsStr == "" {
|
||||
return
|
||||
}
|
||||
tags := strings.Split(tagsStr, ",")
|
||||
for _, tag := range tags {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
var storage models.DataStorage
|
||||
res := database.DB.Where("type = ? AND name = ?", constant.RelationTypeEnvTag, tag).Limit(1).Find(&storage)
|
||||
if res.RowsAffected == 0 {
|
||||
storage = models.DataStorage{
|
||||
ID: utils.GenerateID(),
|
||||
Type: constant.RelationTypeEnvTag,
|
||||
Name: tag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
database.DB.Create(&storage)
|
||||
}
|
||||
relation := models.DataRelation{
|
||||
ID: utils.GenerateID(),
|
||||
DataID: envID,
|
||||
RelateID: storage.ID,
|
||||
Type: constant.RelationTypeEnvTag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
database.DB.Create(&relation)
|
||||
}
|
||||
}
|
||||
|
||||
// LoadEnvTags 为环境变量列表加载标签
|
||||
func (es *EnvService) LoadEnvTags(envs []models.EnvironmentVariable) {
|
||||
if len(envs) == 0 {
|
||||
return
|
||||
}
|
||||
envIDs := make([]string, len(envs))
|
||||
for i, e := range envs {
|
||||
envIDs[i] = e.ID
|
||||
}
|
||||
|
||||
tagsMap := relation.DataRelation.LoadTags(envIDs, constant.RelationTypeEnvTag)
|
||||
|
||||
for i, e := range envs {
|
||||
if tags, ok := tagsMap[e.ID]; ok {
|
||||
envs[i].Tags = strings.Join(tags, ",")
|
||||
} else {
|
||||
envs[i].Tags = ""
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GetAllEnvTags 获取全局环境变量标签
|
||||
func (es *EnvService) GetAllEnvTags() ([]string, error) {
|
||||
return relation.DataRelation.GetAllTags(constant.RelationTypeEnvTag)
|
||||
}
|
||||
|
||||
// CleanEnvTags 删除环境变量时清理关联标签记录
|
||||
func (es *EnvService) CleanEnvTags(id string) {
|
||||
database.DB.Where("data_id = ? AND type = ?", id, constant.RelationTypeEnvTag).Delete(&models.DataRelation{})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type InitService struct {
|
||||
settingsService *SettingsService
|
||||
}
|
||||
|
||||
func NewInitService(settingsService *SettingsService) *InitService {
|
||||
return &InitService{
|
||||
settingsService: settingsService,
|
||||
}
|
||||
}
|
||||
|
||||
// Initialize 执行系统初始化,返回 UserService
|
||||
func (s *InitService) Initialize() *UserService {
|
||||
logger.Info("[Initialize] 开始初始化系统...")
|
||||
|
||||
// 初始化默认设置
|
||||
if err := s.settingsService.InitSettings(); err != nil {
|
||||
logger.Warnf("[Initialize] 初始化设置失败: %v", err)
|
||||
}
|
||||
|
||||
// 创建 UserService
|
||||
userService := NewUserService()
|
||||
|
||||
// 创建管理员账号
|
||||
s.initializeAdmin(userService)
|
||||
|
||||
// 初始化语言环境
|
||||
s.initializeLanguages()
|
||||
|
||||
return userService
|
||||
}
|
||||
|
||||
// initializeLanguages 初始化同步语言环境
|
||||
func (s *InitService) initializeLanguages() {
|
||||
logger.Info("[Languages] 开始初始化编程语言环境...")
|
||||
miseService := NewMiseService()
|
||||
if err := miseService.Sync(); err != nil {
|
||||
logger.Errorf("[Languages] 初始化同步语言环境失败: %v", err)
|
||||
} else {
|
||||
logger.Info("[Languages] 初始化语言环境同步完成")
|
||||
}
|
||||
}
|
||||
|
||||
// initializeAdmin 创建管理员账号
|
||||
func (s *InitService) initializeAdmin(userService *UserService) {
|
||||
existingUser := userService.GetUserByUsername("admin")
|
||||
if existingUser != nil {
|
||||
logger.Info("[Init] 管理员账号已存在,跳过创建")
|
||||
return
|
||||
}
|
||||
|
||||
password := utils.RandomString(12)
|
||||
userService.CreateUser("admin", password, "admin@local", "admin")
|
||||
logger.Infof("--------------------------------------------------")
|
||||
logger.Infof("[Init] 管理员账号创建成功:")
|
||||
logger.Infof("[Init] 用户名: admin")
|
||||
logger.Infof("[Init] 密 码: %s", password)
|
||||
logger.Infof("[Init] 请妥善保管您的密码,并登录后及时修改。")
|
||||
logger.Infof("--------------------------------------------------")
|
||||
}
|
||||
@@ -0,0 +1,263 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"gopkg.in/ini.v1"
|
||||
)
|
||||
|
||||
// InstallRequest 安装请求
|
||||
type InstallRequest struct {
|
||||
// 数据库配置
|
||||
DBType string `json:"db_type" binding:"required"` // sqlite 或 mysql
|
||||
DBHost string `json:"db_host"` // MySQL 主机
|
||||
DBPort int `json:"db_port"` // MySQL 端口
|
||||
DBUser string `json:"db_user"` // MySQL 用户名
|
||||
DBPassword string `json:"db_password"` // MySQL 密码
|
||||
DBName string `json:"db_name"` // MySQL 数据库名
|
||||
DBPath string `json:"db_path"` // SQLite 数据库路径
|
||||
DBSSLMode string `json:"db_ssl_mode"` // SSL 模式
|
||||
|
||||
// Redis 配置(可选)
|
||||
RedisEnabled bool `json:"redis_enabled"` // 是否启用 Redis
|
||||
RedisHost string `json:"redis_host"` // Redis 主机
|
||||
RedisPort int `json:"redis_port"` // Redis 端口
|
||||
RedisPassword string `json:"redis_password"` // Redis 密码
|
||||
RedisDB int `json:"redis_db"` // Redis 数据库索引
|
||||
|
||||
// 管理员账号
|
||||
AdminUsername string `json:"admin_username" binding:"required"` // 管理员用户名
|
||||
AdminPassword string `json:"admin_password" binding:"required"` // 管理员密码
|
||||
AdminEmail string `json:"admin_email"` // 管理员邮箱
|
||||
|
||||
// 站点设置
|
||||
SiteTitle string `json:"site_title"` // 站点标题
|
||||
SiteSubtitle string `json:"site_subtitle"` // 站点副标题
|
||||
}
|
||||
|
||||
// InstallStatus 安装状态
|
||||
type InstallStatus struct {
|
||||
Installed bool `json:"installed"`
|
||||
ConfigPath string `json:"config_path"`
|
||||
DBType string `json:"db_type"`
|
||||
}
|
||||
|
||||
// InstallService 安装服务
|
||||
type InstallService struct {
|
||||
settingsService *SettingsService
|
||||
userService *UserService
|
||||
}
|
||||
|
||||
// NewInstallService 创建安装服务
|
||||
func NewInstallService() *InstallService {
|
||||
return &InstallService{
|
||||
settingsService: NewSettingsService(),
|
||||
userService: NewUserService(),
|
||||
}
|
||||
}
|
||||
|
||||
// CheckInstallStatus 检查安装状态
|
||||
func (s *InstallService) CheckInstallStatus() (*InstallStatus, error) {
|
||||
status := &InstallStatus{
|
||||
ConfigPath: constant.ConfigPath,
|
||||
}
|
||||
|
||||
// 检查配置文件是否存在
|
||||
configExists := false
|
||||
if _, err := os.Stat(constant.ConfigPath); err == nil {
|
||||
configExists = true
|
||||
}
|
||||
|
||||
// 检查是否已初始化(数据库中有管理员用户)
|
||||
initialized := false
|
||||
if configExists {
|
||||
// 尝试检查是否有管理员用户
|
||||
admin := s.userService.GetUserByUsername("admin")
|
||||
if admin != nil {
|
||||
initialized = true
|
||||
}
|
||||
}
|
||||
|
||||
// 同时检查数据库中的初始化标志
|
||||
dbInitialized := s.settingsService.Get(constant.SectionSystem, constant.KeyInitialized)
|
||||
if dbInitialized == "true" {
|
||||
initialized = true
|
||||
}
|
||||
|
||||
status.Installed = initialized
|
||||
status.DBType = Config.Database.Type
|
||||
|
||||
return status, nil
|
||||
}
|
||||
|
||||
// Install 执行安装
|
||||
func (s *InstallService) Install(req *InstallRequest) error {
|
||||
// 1. 创建配置文件
|
||||
if err := s.createConfigFile(req); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 2. 重新加载数据库配置
|
||||
if err := s.reloadDatabase(req); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 3. 初始化数据库
|
||||
if err := database.Migrate(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 4. 初始化设置
|
||||
if err := s.settingsService.InitSettings(); err != nil {
|
||||
logger.Warnf("[Install] 初始化设置失败: %v", err)
|
||||
}
|
||||
|
||||
// 5. 创建管理员账号
|
||||
if err := s.createAdmin(req); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 6. 保存站点设置
|
||||
if req.SiteTitle != "" || req.SiteSubtitle != "" {
|
||||
siteSettings := make(map[string]string)
|
||||
if req.SiteTitle != "" {
|
||||
siteSettings[constant.KeyTitle] = req.SiteTitle
|
||||
}
|
||||
if req.SiteSubtitle != "" {
|
||||
siteSettings[constant.KeySubtitle] = req.SiteSubtitle
|
||||
}
|
||||
if err := s.settingsService.SetSection(constant.SectionSite, siteSettings); err != nil {
|
||||
logger.Warnf("[Install] 保存站点设置失败: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 7. 标记已初始化
|
||||
if err := s.settingsService.Set(constant.SectionSystem, constant.KeyInitialized, "true"); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
logger.Info("[Install] 安装完成")
|
||||
return nil
|
||||
}
|
||||
|
||||
// createConfigFile 创建配置文件
|
||||
func (s *InstallService) createConfigFile(req *InstallRequest) error {
|
||||
// 确保配置目录存在
|
||||
configDir := filepath.Dir(constant.ConfigPath)
|
||||
if err := os.MkdirAll(configDir, 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 创建配置文件
|
||||
cfg := ini.Empty()
|
||||
|
||||
// [server] 配置
|
||||
serverSection, _ := cfg.NewSection("server")
|
||||
serverSection.NewKey("port", "8052")
|
||||
serverSection.NewKey("host", "0.0.0.0")
|
||||
serverSection.NewKey("cookie_name", "BHToken")
|
||||
|
||||
// [database] 配置
|
||||
dbSection, _ := cfg.NewSection("database")
|
||||
dbSection.NewKey("type", req.DBType)
|
||||
|
||||
if req.DBType == "sqlite" {
|
||||
dbPath := req.DBPath
|
||||
if dbPath == "" {
|
||||
dbPath = constant.DefaultDBPath
|
||||
}
|
||||
dbSection.NewKey("path", dbPath)
|
||||
} else if req.DBType == "mysql" {
|
||||
dbSection.NewKey("host", req.DBHost)
|
||||
dbSection.NewKey("port", intToStr(req.DBPort))
|
||||
dbSection.NewKey("user", req.DBUser)
|
||||
dbSection.NewKey("password", req.DBPassword)
|
||||
dbSection.NewKey("dbname", req.DBName)
|
||||
if req.DBSSLMode != "" {
|
||||
dbSection.NewKey("ssl_mode", req.DBSSLMode)
|
||||
}
|
||||
}
|
||||
dbSection.NewKey("table_prefix", "taskpool_")
|
||||
|
||||
// [redis] 配置(如果启用)
|
||||
if req.RedisEnabled {
|
||||
redisSection, _ := cfg.NewSection("redis")
|
||||
redisSection.NewKey("enabled", "true")
|
||||
redisSection.NewKey("host", req.RedisHost)
|
||||
redisSection.NewKey("port", intToStr(req.RedisPort))
|
||||
if req.RedisPassword != "" {
|
||||
redisSection.NewKey("password", req.RedisPassword)
|
||||
}
|
||||
redisSection.NewKey("db", intToStr(req.RedisDB))
|
||||
}
|
||||
|
||||
// [security] 配置
|
||||
securitySection, _ := cfg.NewSection("security")
|
||||
secret := utils.RandomString(32)
|
||||
securitySection.NewKey("secret", secret)
|
||||
|
||||
// 保存配置文件
|
||||
if err := cfg.SaveTo(constant.ConfigPath); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
logger.Infof("[Install] 配置文件已保存到: %s", constant.ConfigPath)
|
||||
return nil
|
||||
}
|
||||
|
||||
// reloadDatabase 重新加载数据库
|
||||
func (s *InstallService) reloadDatabase(req *InstallRequest) error {
|
||||
dbCfg := &database.Config{
|
||||
Type: req.DBType,
|
||||
Host: req.DBHost,
|
||||
Port: req.DBPort,
|
||||
User: req.DBUser,
|
||||
Password: req.DBPassword,
|
||||
DBName: req.DBName,
|
||||
Path: req.DBPath,
|
||||
SSLMode: req.DBSSLMode,
|
||||
}
|
||||
|
||||
if req.DBType == "sqlite" && req.DBPath == "" {
|
||||
dbCfg.Path = constant.DefaultDBPath
|
||||
}
|
||||
|
||||
if err := database.Init(dbCfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// createAdmin 创建管理员账号
|
||||
func (s *InstallService) createAdmin(req *InstallRequest) error {
|
||||
// 检查用户是否已存在
|
||||
existingUser := s.userService.GetUserByUsername(req.AdminUsername)
|
||||
if existingUser != nil {
|
||||
logger.Info("[Install] 管理员账号已存在,跳过创建")
|
||||
return nil
|
||||
}
|
||||
|
||||
email := req.AdminEmail
|
||||
if email == "" {
|
||||
email = "admin@local"
|
||||
}
|
||||
|
||||
s.userService.CreateUser(req.AdminUsername, req.AdminPassword, email, "admin")
|
||||
logger.Infof("[Install] 管理员账号创建成功: %s", req.AdminUsername)
|
||||
return nil
|
||||
}
|
||||
|
||||
// intToStr 整数转字符串
|
||||
func intToStr(n int) string {
|
||||
if n == 0 {
|
||||
return ""
|
||||
}
|
||||
return utils.IntToStr(n)
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type InterconnectService struct{}
|
||||
|
||||
func NewInterconnectService() *InterconnectService {
|
||||
return &InterconnectService{}
|
||||
}
|
||||
|
||||
func (s *InterconnectService) GetNodes() ([]*models.InterconnectNode, error) {
|
||||
var nodes []*models.InterconnectNode
|
||||
err := database.DB.Find(&nodes).Error
|
||||
return nodes, err
|
||||
}
|
||||
|
||||
func (s *InterconnectService) GetNodeByID(id string) (*models.InterconnectNode, error) {
|
||||
var node models.InterconnectNode
|
||||
err := database.DB.Where("id = ?", id).First(&node).Error
|
||||
return &node, err
|
||||
}
|
||||
|
||||
func (s *InterconnectService) CreateNode(name, url, token, remark string) (*models.InterconnectNode, error) {
|
||||
nodeID := utils.GenerateID()
|
||||
if url == "" {
|
||||
url = "tunnel://" + nodeID
|
||||
}
|
||||
node := &models.InterconnectNode{
|
||||
ID: nodeID,
|
||||
Name: name,
|
||||
URL: url,
|
||||
Token: strings.ToLower(token),
|
||||
Remark: remark,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
err := database.DB.Create(node).Error
|
||||
return node, err
|
||||
}
|
||||
|
||||
func (s *InterconnectService) UpdateNode(id, name, url, token, remark string) (*models.InterconnectNode, error) {
|
||||
node, err := s.GetNodeByID(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
node.Name = name
|
||||
if url != "" {
|
||||
node.URL = url
|
||||
}
|
||||
node.Token = token
|
||||
node.Remark = remark
|
||||
node.UpdatedAt = models.Now()
|
||||
|
||||
err = database.DB.Save(node).Error
|
||||
return node, err
|
||||
}
|
||||
|
||||
func (s *InterconnectService) DeleteNode(id string) error {
|
||||
return database.DB.Where("id = ?", id).Delete(&models.InterconnectNode{}).Error
|
||||
}
|
||||
|
||||
func (s *InterconnectService) GetNodeByToken(token string) (*models.InterconnectNode, error) {
|
||||
var node models.InterconnectNode
|
||||
err := database.DB.Where("token = ?", token).First(&node).Error
|
||||
return &node, err
|
||||
}
|
||||
|
||||
func (s *InterconnectService) UpdateNodeMonitorData(id string, metrics models.NodeMetrics) error {
|
||||
now := models.Now()
|
||||
return database.DB.Model(&models.InterconnectNode{}).
|
||||
Where("id = ?", id).
|
||||
Select("status", "metrics", "last_heartbeat_at", "updated_at").
|
||||
Updates(models.InterconnectNode{
|
||||
Status: "online",
|
||||
Metrics: metrics,
|
||||
LastHeartbeatAt: &now,
|
||||
UpdatedAt: now,
|
||||
}).Error
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type LoginLogService struct{}
|
||||
|
||||
func NewLoginLogService() *LoginLogService {
|
||||
return &LoginLogService{}
|
||||
}
|
||||
|
||||
// Create 创建登录日志
|
||||
func (s *LoginLogService) Create(username, ip, userAgent, status, message string) error {
|
||||
level := constant.LogLevelInfo
|
||||
if status != "success" {
|
||||
level = constant.LogLevelWarning
|
||||
}
|
||||
log := &models.AppLog{
|
||||
ID: utils.GenerateID(),
|
||||
Category: constant.LogCategoryLoginLog,
|
||||
Title: username,
|
||||
Content: models.BigText(userAgent),
|
||||
Level: level,
|
||||
Status: status,
|
||||
RefID: ip,
|
||||
ErrorMsg: models.BigText(message),
|
||||
}
|
||||
err := database.DB.Create(log).Error
|
||||
if err == nil {
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventAppLogAdded,
|
||||
Payload: log,
|
||||
})
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// SubscribeEvents 注册订阅事件
|
||||
func (s *LoginLogService) SubscribeEvents(bus *eventbus.EventBus) {
|
||||
// 用户登录事件
|
||||
bus.Subscribe(constant.EventUserLogin, func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
username, _ := payload["username"].(string)
|
||||
ip, _ := payload["ip"].(string)
|
||||
userAgent, _ := payload["userAgent"].(string)
|
||||
status, _ := payload["status"].(string)
|
||||
message, _ := payload["message"].(string)
|
||||
|
||||
s.Create(username, ip, userAgent, status, message)
|
||||
|
||||
// 如果登录成功,触发系统通知
|
||||
if status == "success" {
|
||||
bus.Publish(eventbus.Event{
|
||||
Type: constant.EventSystemNotice,
|
||||
Payload: map[string]interface{}{
|
||||
"title": "登录提醒",
|
||||
"content": fmt.Sprintf("用户 %s 已从 IP %s 登录系统", username, ip),
|
||||
"level": constant.LogLevelWarning,
|
||||
},
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
// 暴力破解防御触发事件
|
||||
bus.Subscribe(constant.EventBruteForceLogin, func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
username, _ := payload["username"].(string)
|
||||
ip, _ := payload["ip"].(string)
|
||||
userAgent, _ := payload["userAgent"].(string)
|
||||
|
||||
s.Create(username, ip, userAgent, "failed", "尝试次数过多,由于暴力破解防御机制已锁定")
|
||||
|
||||
// 触发系统通知
|
||||
bus.Publish(eventbus.Event{
|
||||
Type: constant.EventSystemNotice,
|
||||
Payload: map[string]interface{}{
|
||||
"title": "系统安全警告",
|
||||
"content": fmt.Sprintf("检测到 IP %s 正在尝试暴力破解用户 %s", ip, username),
|
||||
"level": constant.LogLevelError,
|
||||
},
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
// List 获取登录日志列表
|
||||
func (s *LoginLogService) List(page, pageSize int, username string) ([]models.AppLog, int64, error) {
|
||||
var logs []models.AppLog
|
||||
var total int64
|
||||
|
||||
query := database.DB.Model(&models.AppLog{}).Where("category = ?", constant.LogCategoryLoginLog)
|
||||
if username != "" {
|
||||
query = query.Where("title LIKE ?", "%"+username+"%")
|
||||
}
|
||||
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
offset := (page - 1) * pageSize
|
||||
if err := query.Order("created_at DESC").Offset(offset).Limit(pageSize).Find(&logs).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
return logs, total, nil
|
||||
}
|
||||
|
||||
// CleanOldLogs 清理指定天数前的日志
|
||||
func (s *LoginLogService) CleanOldLogs(days int) (int64, error) {
|
||||
deadline := time.Now().AddDate(0, 0, -days)
|
||||
result := database.DB.Unscoped().Where("category = ? AND created_at < ?", constant.LogCategoryLoginLog, deadline).Delete(&models.AppLog{})
|
||||
return result.RowsAffected, result.Error
|
||||
}
|
||||
@@ -0,0 +1,446 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// MigrationTable 定义迁移配置
|
||||
type MigrationTable struct {
|
||||
Model any
|
||||
EntityName string
|
||||
FKs map[string]string // 单个外键映射: 字段名 -> 实体名
|
||||
MultiFKs map[string]string // 复合外键映射 (如逗号分隔): 字段名 -> 实体名
|
||||
}
|
||||
|
||||
func getMigrationTables() []MigrationTable {
|
||||
return []MigrationTable{
|
||||
{&models.User{}, "users", nil, nil},
|
||||
{&models.Agent{}, "agents", nil, nil},
|
||||
{&models.AgentToken{}, "tokens", nil, nil},
|
||||
{&models.EnvironmentVariable{}, "envs", map[string]string{"UserID": "users"}, nil},
|
||||
{&models.Task{}, "tasks", map[string]string{"AgentID": "agents"}, map[string]string{"Envs": "envs"}},
|
||||
{&models.TaskLog{}, "task_logs", map[string]string{"TaskID": "tasks", "AgentID": "agents"}, nil},
|
||||
{&models.Script{}, "scripts", map[string]string{"UserID": "users"}, nil},
|
||||
{&models.Setting{}, "settings", nil, nil},
|
||||
{&models.SendStats{}, "send_stats", map[string]string{"TaskID": "tasks"}, nil},
|
||||
{&models.Language{}, "languages", nil, nil},
|
||||
{&models.Dependency{}, "deps", nil, nil},
|
||||
}
|
||||
}
|
||||
|
||||
func getTableName(db *gorm.DB, model any) string {
|
||||
stmt := &gorm.Statement{DB: db}
|
||||
if err := stmt.Parse(model); err != nil {
|
||||
return ""
|
||||
}
|
||||
return stmt.Schema.Table
|
||||
}
|
||||
|
||||
func isTableStringID(db *gorm.DB, model any) bool {
|
||||
if !db.Migrator().HasTable(model) {
|
||||
return true
|
||||
}
|
||||
columnTypes, err := db.Migrator().ColumnTypes(model)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
for _, ct := range columnTypes {
|
||||
if strings.ToLower(ct.Name()) == "id" {
|
||||
typeName := strings.ToLower(ct.DatabaseTypeName())
|
||||
return strings.Contains(typeName, "char") || strings.Contains(typeName, "text") || strings.Contains(typeName, "string")
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func getValFromMap(m map[string]interface{}, key string) (interface{}, bool) {
|
||||
lowerKey := strings.ToLower(key)
|
||||
for k, v := range m {
|
||||
if strings.ToLower(k) == lowerKey {
|
||||
return v, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func RunMigrationV3() error {
|
||||
db := database.DB
|
||||
if db == nil {
|
||||
return fmt.Errorf("数据库未初始化")
|
||||
}
|
||||
|
||||
// 0. 检查迁移标记,防止重复迁移逻辑被误判触发
|
||||
if db.Migrator().HasTable(&models.Setting{}) {
|
||||
var migrationFlag models.Setting
|
||||
res := db.Where(&models.Setting{Section: "system", Key: "migration_v3_success"}).Limit(1).Find(&migrationFlag)
|
||||
if res.Error == nil && res.RowsAffected > 0 && migrationFlag.Value == "true" {
|
||||
// 如果已经是字符串 ID 模式,双重确认
|
||||
logger.Info("[MigrationV3] 系统已处于 V3 模式,跳过检查")
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
tables := getMigrationTables()
|
||||
needMigration := false
|
||||
for _, t := range tables {
|
||||
if !isTableStringID(db, t.Model) {
|
||||
needMigration = true
|
||||
logger.Infof("[MigrationV3] 表 [%s] ID 为数字,需迁移", getTableName(db, t.Model))
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !needMigration {
|
||||
// 没有数字 ID 表了,但也可能还没打标(比如之前手动修过表),此时补一个标
|
||||
return markMigrationSuccess(db)
|
||||
}
|
||||
|
||||
// 1. 备份过程
|
||||
backupDir := "./data"
|
||||
os.MkdirAll(backupDir, 0755)
|
||||
// 如果已经有还原过的记录,为了防止反复循环,我们可以检查还原标记
|
||||
backups, _ := filepath.Glob(filepath.Join(backupDir, "migration_v3_backup_*.zip"))
|
||||
if len(backups) == 0 {
|
||||
logger.Infof("[MigrationV3] 执行关键备份...")
|
||||
backupService := NewBackupService()
|
||||
zipPath, err := backupService.CreateBackup()
|
||||
if err != nil {
|
||||
return fmt.Errorf("自动备份失败,流程终止: %v", err)
|
||||
}
|
||||
baseName := strings.TrimSuffix(filepath.Base(zipPath), ".zip")
|
||||
newPath := filepath.Join(backupDir, fmt.Sprintf("migration_v3_backup_%s.zip", baseName))
|
||||
os.Rename(zipPath, newPath)
|
||||
logger.Infof("[MigrationV3] 备份成功: %s", newPath)
|
||||
}
|
||||
|
||||
mappings := make(map[string]map[uint]string)
|
||||
// 根据数据库类型决定是否使用事务:
|
||||
// - PostgreSQL: DDL 完全支持事务,使用事务保证原子性
|
||||
// - MySQL: DDL 会隐式提交事务,包裹事务无意义
|
||||
// - SQLite: DDL+DML 混合在 GORM 事务中会导致数据丢失
|
||||
dbType := db.Dialector.Name()
|
||||
if dbType == "postgres" {
|
||||
err := db.Transaction(func(tx *gorm.DB) error {
|
||||
return performHardMigration(tx, mappings)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
// SQLite / MySQL: 不使用事务包裹
|
||||
if err := performHardMigration(db, mappings); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 标记成功
|
||||
return markMigrationSuccess(db)
|
||||
}
|
||||
|
||||
func markMigrationSuccess(db *gorm.DB) error {
|
||||
if !db.Migrator().HasTable(&models.Setting{}) {
|
||||
return nil
|
||||
}
|
||||
var flag models.Setting
|
||||
res := db.Where(&models.Setting{Section: "system", Key: "migration_v3_success"}).Limit(1).Find(&flag)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
// 创建或更新
|
||||
flag = models.Setting{
|
||||
ID: utils.GenerateID(),
|
||||
Section: "system",
|
||||
Key: "migration_v3_success",
|
||||
Value: "true",
|
||||
}
|
||||
return db.Create(&flag).Error
|
||||
}
|
||||
return db.Model(&flag).Update("value", "true").Error
|
||||
}
|
||||
|
||||
func performHardMigration(db *gorm.DB, mappings map[string]map[uint]string) error {
|
||||
allTables := getMigrationTables()
|
||||
|
||||
// ---------------------------------------------------------
|
||||
// 第一阶段:全量构建 ID 映射映射表 (Pass 1)
|
||||
// ---------------------------------------------------------
|
||||
for _, t := range allTables {
|
||||
actualName := getTableName(db, t.Model)
|
||||
if actualName == "" || !db.Migrator().HasTable(actualName) {
|
||||
continue
|
||||
}
|
||||
mappings[t.EntityName] = make(map[uint]string)
|
||||
oldTableName := actualName + "_v2_bak"
|
||||
|
||||
// 如果还没有备份表,说明这是第一次处理该表,先重命名
|
||||
if !db.Migrator().HasTable(oldTableName) {
|
||||
if isTableStringID(db, t.Model) {
|
||||
continue // 已经是字符串 ID 且无备份,跳过
|
||||
}
|
||||
if err := db.Migrator().RenameTable(actualName, oldTableName); err != nil {
|
||||
return fmt.Errorf("重命名表 %s 失败: %v", actualName, err)
|
||||
}
|
||||
// SQLite 重命名表后,索引名称不变,会导致 AutoMigrate 创建新表时索引冲突
|
||||
// 需要先删除备份表上的旧索引
|
||||
dropOldIndexes(db, oldTableName)
|
||||
}
|
||||
|
||||
// 预先为该表所有记录生成新的 xid
|
||||
var rows []map[string]interface{}
|
||||
db.Table(oldTableName).Select("id").Find(&rows)
|
||||
for _, row := range rows {
|
||||
if val, ok := getValFromMap(row, "id"); ok {
|
||||
uid := parseUint(val)
|
||||
if uid > 0 {
|
||||
mappings[t.EntityName][uid] = utils.GenerateID()
|
||||
}
|
||||
}
|
||||
}
|
||||
logger.Infof("[MigrationV3] Pass 1: 构建关键表 %s 的 ID 映射, 共 %d 条", actualName, len(mappings[t.EntityName]))
|
||||
}
|
||||
|
||||
// 构建 fallback ID:FK 映射找不到时使用第一条记录的新 ID
|
||||
fallbackIDs := make(map[string]string)
|
||||
for _, t := range allTables {
|
||||
// 优先从 mappings 中取第一个新 xid(本次迁移生成的)
|
||||
if m, ok := mappings[t.EntityName]; ok && len(m) > 0 {
|
||||
for _, newID := range m {
|
||||
fallbackIDs[t.EntityName] = newID
|
||||
break
|
||||
}
|
||||
continue
|
||||
}
|
||||
// 如果该表已经是 string ID(被跳过了),从实际表查第一条
|
||||
actualName := getTableName(db, t.Model)
|
||||
if actualName != "" && db.Migrator().HasTable(actualName) {
|
||||
var row map[string]interface{}
|
||||
if err := db.Table(actualName).Select("id").Order("id").Limit(1).Find(&row).Error; err == nil {
|
||||
if id, ok := getValFromMap(row, "id"); ok && id != nil {
|
||||
fallbackIDs[t.EntityName] = fmt.Sprintf("%v", id)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------
|
||||
// 第二、三阶段:正式转换数据并处理关联字段 (Pass 2 & 3)
|
||||
// ---------------------------------------------------------
|
||||
for _, t := range allTables {
|
||||
actualName := getTableName(db, t.Model)
|
||||
oldTableName := actualName + "_v2_bak"
|
||||
|
||||
if !db.Migrator().HasTable(oldTableName) {
|
||||
// 虽然可能已经改过格式,但为了安全还是 AutoMigrate 一下
|
||||
db.AutoMigrate(t.Model)
|
||||
continue
|
||||
}
|
||||
|
||||
logger.Infof("[MigrationV3] Pass 2&3: 正在转换数据并修复关联: %s", actualName)
|
||||
db.AutoMigrate(t.Model)
|
||||
|
||||
// 获取新表的有效列名(小写)
|
||||
columnTypes, _ := db.Migrator().ColumnTypes(t.Model)
|
||||
validColumns := make(map[string]bool)
|
||||
for _, ct := range columnTypes {
|
||||
validColumns[strings.ToLower(ct.Name())] = true
|
||||
}
|
||||
|
||||
var oldData []map[string]interface{}
|
||||
if err := db.Table(oldTableName).Find(&oldData).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, row := range oldData {
|
||||
// 1. 处理主键 ID
|
||||
if val, ok := getValFromMap(row, "id"); ok {
|
||||
uid := parseUint(val)
|
||||
if nid, exists := mappings[t.EntityName][uid]; exists {
|
||||
row["id"] = nid
|
||||
} else {
|
||||
row["id"] = utils.GenerateID()
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 处理单外键关联 (Phase 3)
|
||||
for field, parentEntity := range t.FKs {
|
||||
columnName := getColumnName(field)
|
||||
if val, ok := getValFromMap(row, columnName); ok && val != nil {
|
||||
strVal := fmt.Sprintf("%v", val)
|
||||
if strVal == "" || strVal == "0" || strVal == "<nil>" {
|
||||
row[columnName] = nil
|
||||
continue
|
||||
}
|
||||
// 如果值已经是非数字字符串(例如已迁移过的 xid),直接保留
|
||||
if !utils.IsNumeric(strVal) {
|
||||
continue
|
||||
}
|
||||
// 数字型外键,查映射表转换
|
||||
ufk := parseUint(val)
|
||||
if ufk > 0 {
|
||||
if nid, exists := mappings[parentEntity][ufk]; exists {
|
||||
row[columnName] = nid
|
||||
} else if fbID, hasFB := fallbackIDs[parentEntity]; hasFB {
|
||||
logger.Warnf("[MigrationV3] 表 %s 的 %s=%d 在 %s 映射中未找到,使用首条记录 ID: %s", getTableName(db, t.Model), columnName, ufk, parentEntity, fbID)
|
||||
row[columnName] = fbID
|
||||
} else {
|
||||
logger.Warnf("[MigrationV3] 表 %s 的 %s=%d 在 %s 映射中未找到且无 fallback,置空", getTableName(db, t.Model), columnName, ufk, parentEntity)
|
||||
row[columnName] = nil
|
||||
}
|
||||
} else {
|
||||
row[columnName] = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 处理复合多外键字段 (envs)
|
||||
for field, parentEntity := range t.MultiFKs {
|
||||
columnName := getColumnName(field)
|
||||
if val, ok := getValFromMap(row, columnName); ok && val != nil {
|
||||
if strVal, ok := val.(string); ok && strVal != "" {
|
||||
row[columnName] = transformMultiIDs(strVal, parentEntity, mappings)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 4. 过滤不存在的列并插入
|
||||
filteredRow := make(map[string]interface{})
|
||||
for k, v := range row {
|
||||
if validColumns[strings.ToLower(k)] {
|
||||
filteredRow[k] = v
|
||||
}
|
||||
}
|
||||
if err := db.Table(actualName).Create(filteredRow).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 迁移完成,清理备份表
|
||||
db.Migrator().DropTable(oldTableName)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// 辅助函数:解析各种数字 ID
|
||||
func parseUint(val interface{}) uint {
|
||||
if val == nil {
|
||||
return 0
|
||||
}
|
||||
switch v := val.(type) {
|
||||
case uint:
|
||||
return v
|
||||
case int64:
|
||||
return uint(v)
|
||||
case int:
|
||||
return uint(v)
|
||||
case uint64:
|
||||
return uint(v)
|
||||
case float64:
|
||||
return uint(v)
|
||||
case string:
|
||||
var u uint
|
||||
fmt.Sscanf(v, "%d", &u)
|
||||
return u
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// 辅助函数:字段名转列名
|
||||
func getColumnName(field string) string {
|
||||
switch field {
|
||||
case "AgentID":
|
||||
return "agent_id"
|
||||
case "TaskID":
|
||||
return "task_id"
|
||||
case "UserID":
|
||||
return "user_id"
|
||||
case "LogID":
|
||||
return "log_id"
|
||||
case "Envs":
|
||||
return "envs"
|
||||
default:
|
||||
return strings.ToLower(field)
|
||||
}
|
||||
}
|
||||
|
||||
// 辅助函数:处理逗号分隔的 ID 列表
|
||||
func transformMultiIDs(oldStr string, parentEntity string, mappings map[string]map[uint]string) string {
|
||||
parts := strings.Split(oldStr, ",")
|
||||
var result []string
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p == "" {
|
||||
continue
|
||||
}
|
||||
if len(p) == 20 && !utils.IsNumeric(p) {
|
||||
result = append(result, p) // 已经是 xid,保留
|
||||
continue
|
||||
}
|
||||
uid := parseUint(p)
|
||||
if uid > 0 {
|
||||
if nid, exists := mappings[parentEntity][uid]; exists {
|
||||
result = append(result, nid)
|
||||
}
|
||||
}
|
||||
}
|
||||
return strings.Join(result, ",")
|
||||
}
|
||||
|
||||
// dropOldIndexes 删除备份表上的旧索引,防止 AutoMigrate 创建新表时索引名冲突
|
||||
// SQLite 重命名表后索引名不变,MySQL/PostgreSQL 也可能存在类似问题
|
||||
func dropOldIndexes(db *gorm.DB, tableName string) {
|
||||
dbType := db.Dialector.Name()
|
||||
var indexNames []string
|
||||
|
||||
switch dbType {
|
||||
case "sqlite":
|
||||
var indexes []struct{ Name string }
|
||||
db.Raw("SELECT name FROM sqlite_master WHERE type='index' AND tbl_name=? AND name NOT LIKE 'sqlite_%'", tableName).Scan(&indexes)
|
||||
for _, idx := range indexes {
|
||||
indexNames = append(indexNames, idx.Name)
|
||||
}
|
||||
case "mysql":
|
||||
var indexes []struct {
|
||||
KeyName string `gorm:"column:Key_name"`
|
||||
}
|
||||
db.Raw("SHOW INDEX FROM `" + tableName + "`").Scan(&indexes)
|
||||
seen := make(map[string]bool)
|
||||
for _, idx := range indexes {
|
||||
if idx.KeyName != "PRIMARY" && !seen[idx.KeyName] {
|
||||
indexNames = append(indexNames, idx.KeyName)
|
||||
seen[idx.KeyName] = true
|
||||
}
|
||||
}
|
||||
case "postgres":
|
||||
var indexes []struct {
|
||||
IndexName string `gorm:"column:indexname"`
|
||||
}
|
||||
db.Raw("SELECT indexname FROM pg_indexes WHERE tablename=?", tableName).Scan(&indexes)
|
||||
for _, idx := range indexes {
|
||||
if !strings.HasSuffix(idx.IndexName, "_pkey") {
|
||||
indexNames = append(indexNames, idx.IndexName)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for _, name := range indexNames {
|
||||
var dropSQL string
|
||||
switch dbType {
|
||||
case "mysql":
|
||||
dropSQL = fmt.Sprintf("DROP INDEX `%s` ON `%s`", name, tableName)
|
||||
default:
|
||||
dropSQL = fmt.Sprintf("DROP INDEX IF EXISTS \"%s\"", name)
|
||||
}
|
||||
if err := db.Exec(dropSQL).Error; err != nil {
|
||||
logger.Warnf("[MigrationV3] 删除旧索引 %s 失败 (可忽略): %v", name, err)
|
||||
} else {
|
||||
logger.Infof("[MigrationV3] 已删除备份表旧索引: %s", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,405 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/deps"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type MiseService struct{}
|
||||
|
||||
func NewMiseService() *MiseService {
|
||||
return &MiseService{}
|
||||
}
|
||||
|
||||
type MiseLanguage struct {
|
||||
Plugin string `json:"plugin"`
|
||||
Version string `json:"version"`
|
||||
Source MiseSource `json:"source"`
|
||||
IsGlobal bool `json:"is_global"`
|
||||
InstallPath string `json:"install_path,omitempty"`
|
||||
InstalledAt string `json:"installed_at,omitempty"` // 安装日期
|
||||
}
|
||||
|
||||
type MiseSource struct {
|
||||
Type string `json:"type"`
|
||||
Path string `json:"path"`
|
||||
}
|
||||
|
||||
// List 实时从系统检测 mise 环境并同步到数据库
|
||||
func (s *MiseService) List() ([]MiseLanguage, error) {
|
||||
langs, err := s.fetchLiveLanguages()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 异步同步到数据库,确保列表响应速度
|
||||
go s.syncToDB(langs)
|
||||
|
||||
return langs, nil
|
||||
}
|
||||
|
||||
// Sync 实时检测本地 mise 环境并同步到数据库
|
||||
func (s *MiseService) Sync() error {
|
||||
// 获取实时数据
|
||||
langs, err := s.fetchLiveLanguages()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 同步到数据库
|
||||
s.syncToDB(langs)
|
||||
return nil
|
||||
}
|
||||
|
||||
// fetchLiveLanguages 实时从系统检测 mise 语言列表
|
||||
func (s *MiseService) fetchLiveLanguages() ([]MiseLanguage, error) {
|
||||
// 使用 --json 获取格式化数据
|
||||
cmd := exec.Command("mise", "ls", "--json")
|
||||
// 继承父进程环境变量
|
||||
cmd.Env = os.Environ()
|
||||
cmd.Env = append(cmd.Env, "MISE_NO_COLOR=1", "TERM=dumb")
|
||||
|
||||
output, err := cmd.Output()
|
||||
if err != nil {
|
||||
logger.Warnf("[Mise] mise ls --json failed: %v", err)
|
||||
return s.listFallback()
|
||||
}
|
||||
|
||||
// 1. 尝试解析为数组格式 [{}, {}]
|
||||
var languages []MiseLanguage
|
||||
if err := json.Unmarshal(output, &languages); err == nil {
|
||||
s.enrichInstallDates(languages)
|
||||
s.enrichSourceInfo(languages)
|
||||
s.sortByInstallDate(languages)
|
||||
return languages, nil
|
||||
}
|
||||
|
||||
// 2. 如果失败,尝试解析为对象格式 {"node": [{}], "python": [{}]}
|
||||
var langMap map[string][]MiseLanguage
|
||||
if err := json.Unmarshal(output, &langMap); err == nil {
|
||||
var result []MiseLanguage
|
||||
for plugin, items := range langMap {
|
||||
for _, item := range items {
|
||||
// 填充插件名称(对象格式中插件名通常是 key)
|
||||
if item.Plugin == "" {
|
||||
item.Plugin = plugin
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
}
|
||||
s.enrichInstallDates(result)
|
||||
s.enrichSourceInfo(result)
|
||||
s.sortByInstallDate(result)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
return s.listFallback()
|
||||
}
|
||||
|
||||
func (s *MiseService) listFallback() ([]MiseLanguage, error) {
|
||||
cmd := exec.Command("mise", "ls")
|
||||
cmd.Env = os.Environ()
|
||||
cmd.Env = append(cmd.Env, "MISE_NO_COLOR=1", "TERM=dumb")
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("mise ls failed: %v, output: %s", err, string(output))
|
||||
}
|
||||
|
||||
lines := strings.Split(string(output), "\n")
|
||||
languages := []MiseLanguage{}
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(strings.ToLower(line), "tool") {
|
||||
continue
|
||||
}
|
||||
|
||||
parts := strings.Fields(line)
|
||||
if len(parts) < 2 {
|
||||
continue
|
||||
}
|
||||
|
||||
lang := MiseLanguage{
|
||||
Plugin: parts[0],
|
||||
Version: parts[1],
|
||||
}
|
||||
if len(parts) >= 3 {
|
||||
lang.Source = MiseSource{Path: parts[2]}
|
||||
}
|
||||
languages = append(languages, lang)
|
||||
}
|
||||
|
||||
return languages, nil
|
||||
}
|
||||
|
||||
// Plugins 获取主流的 mise 插件列表 (固定返回以确保速度和稳定性)
|
||||
func (s *MiseService) Plugins() ([]string, error) {
|
||||
mainstream := constant.MainstreamMisePlugins
|
||||
logger.Infof("[Mise] Returning fixed list of %d mainstream plugins", len(mainstream))
|
||||
return mainstream, nil
|
||||
}
|
||||
|
||||
// Versions 获取指定插件的所有可用版本
|
||||
func (s *MiseService) Versions(plugin string) ([]string, error) {
|
||||
if plugin == "" {
|
||||
return []string{}, nil
|
||||
}
|
||||
|
||||
// 只获取最新版本列表
|
||||
cmd := exec.Command("mise", "ls-remote", plugin)
|
||||
cmd.Env = append(cmd.Env, "MISE_NO_COLOR=1", "TERM=dumb")
|
||||
output, err := cmd.CombinedOutput()
|
||||
|
||||
if err != nil && len(output) == 0 {
|
||||
logger.Errorf("[Mise] Fetch versions for %s failed: %v", plugin, err)
|
||||
return []string{}, nil
|
||||
}
|
||||
|
||||
lines := strings.Split(string(output), "\n")
|
||||
var versions []string
|
||||
// 倒序排列,优先显示新版本
|
||||
for i := len(lines) - 1; i >= 0; i-- {
|
||||
line := strings.TrimSpace(lines[i])
|
||||
if line == "" || strings.Contains(line, " ") || strings.Contains(line, "Usage:") || line == "latest" {
|
||||
continue
|
||||
}
|
||||
versions = append(versions, line)
|
||||
if len(versions) >= 300 {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return versions, nil
|
||||
}
|
||||
|
||||
// enrichInstallDates 为语言列表添加安装日期信息
|
||||
func (s *MiseService) enrichInstallDates(languages []MiseLanguage) {
|
||||
for i := range languages {
|
||||
// 判断是否是 global
|
||||
path := strings.ToLower(languages[i].Source.Path)
|
||||
// 归一化路径分隔符
|
||||
normPath := strings.ReplaceAll(path, "\\", "/")
|
||||
if languages[i].Source.Type == "global" {
|
||||
languages[i].IsGlobal = true
|
||||
} else if strings.Contains(normPath, "mise/config.toml") {
|
||||
languages[i].IsGlobal = true
|
||||
}
|
||||
|
||||
if languages[i].InstallPath != "" {
|
||||
if installDate := s.getInstallDate(languages[i].InstallPath); installDate != "" {
|
||||
languages[i].InstalledAt = installDate
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// enrichSourceInfo 为语言列表添加来源信息
|
||||
func (s *MiseService) enrichSourceInfo(languages []MiseLanguage) {
|
||||
for i := range languages {
|
||||
// 如果source为空,使用install_path作为来源
|
||||
if languages[i].Source.Type == "" && languages[i].Source.Path == "" {
|
||||
if languages[i].InstallPath != "" {
|
||||
languages[i].Source.Path = languages[i].InstallPath
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// getInstallDate 获取安装路径的创建时间
|
||||
func (s *MiseService) getInstallDate(installPath string) string {
|
||||
if installPath == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
fileInfo, err := os.Stat(installPath)
|
||||
if err != nil {
|
||||
logger.Debugf("[Mise] Failed to stat install path %s: %v", installPath, err)
|
||||
return ""
|
||||
}
|
||||
|
||||
// 获取修改时间作为安装时间的近似值
|
||||
modTime := fileInfo.ModTime()
|
||||
return modTime.Format("2006-01-02 15:04:05")
|
||||
}
|
||||
|
||||
// sortByInstallDate 按安装时间降序排序(最新的在前面)
|
||||
func (s *MiseService) sortByInstallDate(languages []MiseLanguage) {
|
||||
sort.Slice(languages, func(i, j int) bool {
|
||||
// 如果都有安装时间,按时间降序排序
|
||||
if languages[i].InstalledAt != "" && languages[j].InstalledAt != "" {
|
||||
return languages[i].InstalledAt > languages[j].InstalledAt
|
||||
}
|
||||
// 有安装时间的排在前面
|
||||
if languages[i].InstalledAt != "" {
|
||||
return true
|
||||
}
|
||||
if languages[j].InstalledAt != "" {
|
||||
return false
|
||||
}
|
||||
// 都没有安装时间,按插件名排序
|
||||
return languages[i].Plugin < languages[j].Plugin
|
||||
})
|
||||
}
|
||||
|
||||
// syncToDB 将实时检测到的语言信息同步到数据库表中
|
||||
func (s *MiseService) syncToDB(languages []MiseLanguage) {
|
||||
db := database.GetDB()
|
||||
if db == nil {
|
||||
return
|
||||
}
|
||||
|
||||
var currentIds []string
|
||||
for _, lang := range languages {
|
||||
var model models.Language
|
||||
// 以 plugin 和 version 作为联合唯一标识(业务逻辑上)
|
||||
res := db.Where("plugin = ? AND version = ?", lang.Plugin, lang.Version).Limit(1).Find(&model)
|
||||
queryErr := res.Error
|
||||
rowsAffected := res.RowsAffected
|
||||
|
||||
sourceStr := ""
|
||||
if lang.Source.Path != "" {
|
||||
sourceStr = lang.Source.Path
|
||||
} else if lang.Source.Type != "" {
|
||||
sourceStr = lang.Source.Type
|
||||
}
|
||||
|
||||
var installTime *models.LocalTime
|
||||
if lang.InstalledAt != "" {
|
||||
t, err := time.Parse("2006-01-02 15:04:05", lang.InstalledAt)
|
||||
if err == nil {
|
||||
lt := models.LocalTime(t)
|
||||
installTime = <
|
||||
}
|
||||
}
|
||||
|
||||
if queryErr == nil && rowsAffected == 0 {
|
||||
// 如果不存在,则创建
|
||||
newLang := models.Language{
|
||||
ID: utils.GenerateID(),
|
||||
Plugin: lang.Plugin,
|
||||
Version: lang.Version,
|
||||
InstallPath: lang.InstallPath,
|
||||
Source: sourceStr,
|
||||
InstalledAt: installTime,
|
||||
}
|
||||
if err := db.Create(&newLang).Error; err == nil {
|
||||
currentIds = append(currentIds, newLang.ID)
|
||||
}
|
||||
} else {
|
||||
// 如果已存在,更新可能变动的信息
|
||||
updates := map[string]interface{}{
|
||||
"install_path": lang.InstallPath,
|
||||
"source": sourceStr,
|
||||
"installed_at": installTime,
|
||||
}
|
||||
db.Model(&model).Updates(updates)
|
||||
currentIds = append(currentIds, model.ID)
|
||||
}
|
||||
}
|
||||
|
||||
// 清理数据库中存在但实际 mise 已经卸载的记录
|
||||
if len(currentIds) > 0 {
|
||||
db.Where("id NOT IN ?", currentIds).Delete(&models.Language{})
|
||||
} else {
|
||||
// 如果本地一个都没有,清空表
|
||||
db.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&models.Language{})
|
||||
}
|
||||
}
|
||||
|
||||
// GetVerifyCommand 获取环境验证命令
|
||||
func (s *MiseService) GetVerifyCommand(plugin, version string) (string, error) {
|
||||
m := deps.GetManager(plugin)
|
||||
if m == nil {
|
||||
// 如果没找到对应的包管理器(比如 java 等不支持依赖管理的),
|
||||
// 则提供一个基础的通用验证命令
|
||||
return utils.BuildMiseCommandSimple(plugin+" --version", plugin, version), nil
|
||||
}
|
||||
return m.GetVerifyCommand(version)
|
||||
}
|
||||
|
||||
// UseGlobal 设置全局默认版本
|
||||
func (s *MiseService) UseGlobal(plugin, version string) error {
|
||||
cmd := exec.Command("mise", "use", "-g", fmt.Sprintf("%s@%s", plugin, version))
|
||||
cmd.Env = os.Environ()
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("mise use -g failed: %v, output: %s", err, string(output))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UnsetGlobal 取消全局默认版本
|
||||
func (s *MiseService) UnsetGlobal(plugin, version string) error {
|
||||
// 使用 mise unuse -g <tool>@<version> 移出全局配置
|
||||
target := plugin
|
||||
if version != "" {
|
||||
target = fmt.Sprintf("%s@%s", plugin, version)
|
||||
}
|
||||
cmd := exec.Command("mise", "unuse", "-g", target)
|
||||
cmd.Env = os.Environ()
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("mise unuse global failed: %v, output: %s", err, string(output))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Envs 获取全局环境变量
|
||||
func (s *MiseService) Envs() (map[string]string, error) {
|
||||
cmd := exec.Command("mise", "set")
|
||||
cmd.Env = os.Environ()
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("mise set failed: %v, output: %s", err, string(output))
|
||||
}
|
||||
|
||||
envs := make(map[string]string)
|
||||
lines := strings.Split(string(output), "\n")
|
||||
for _, line := range lines {
|
||||
fields := strings.Fields(line)
|
||||
// 期望格式: KEY VALUE SOURCE
|
||||
// 跳过标题行或不合规行
|
||||
if len(fields) < 3 || strings.ToLower(fields[0]) == "key" {
|
||||
continue
|
||||
}
|
||||
// 仅显示来自全局配置文件的环境变量
|
||||
if strings.Contains(strings.ToLower(fields[len(fields)-1]), "config.toml") {
|
||||
envs[fields[0]] = fields[1]
|
||||
}
|
||||
}
|
||||
return envs, nil
|
||||
}
|
||||
|
||||
// SetEnv 设置全局环境变量
|
||||
func (s *MiseService) SetEnv(key, value string) error {
|
||||
cmd := exec.Command("mise", "set", "-g", fmt.Sprintf("%s=%s", key, value))
|
||||
cmd.Env = os.Environ()
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("mise set -g failed: %v, output: %s", err, string(output))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UnsetEnv 取消全局环境变量
|
||||
func (s *MiseService) UnsetEnv(key string) error {
|
||||
cmd := exec.Command("mise", "unset", "-g", key)
|
||||
cmd.Env = os.Environ()
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("mise unset -g failed: %v, output: %s", err, string(output))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"math/rand"
|
||||
"runtime"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/shirou/gopsutil/v3/cpu"
|
||||
"github.com/shirou/gopsutil/v3/disk"
|
||||
"github.com/shirou/gopsutil/v3/host"
|
||||
"github.com/shirou/gopsutil/v3/mem"
|
||||
)
|
||||
|
||||
type HostMetrics struct {
|
||||
CPUPercent float64
|
||||
VMem *mem.VirtualMemoryStat
|
||||
DiskUsage *disk.UsageStat
|
||||
HostInfo *host.InfoStat
|
||||
}
|
||||
|
||||
type MonitorService struct {
|
||||
hostMu sync.RWMutex
|
||||
lastUpdate time.Time
|
||||
metrics HostMetrics
|
||||
}
|
||||
|
||||
var (
|
||||
monitorServiceInstance *MonitorService
|
||||
monitorServiceOnce sync.Once
|
||||
)
|
||||
|
||||
// GetMonitorService 获取系统监控服务单例
|
||||
func GetMonitorService() *MonitorService {
|
||||
monitorServiceOnce.Do(func() {
|
||||
monitorServiceInstance = &MonitorService{}
|
||||
})
|
||||
return monitorServiceInstance
|
||||
}
|
||||
|
||||
// GetHostMetrics 获取并返回物理机状态(带有缓存和演示模式伪装)
|
||||
func (ms *MonitorService) GetHostMetrics() HostMetrics {
|
||||
ms.hostMu.Lock()
|
||||
defer ms.hostMu.Unlock()
|
||||
|
||||
// 缓存 2 秒
|
||||
if time.Since(ms.lastUpdate) < 2*time.Second && ms.metrics.VMem != nil {
|
||||
return ms.metrics
|
||||
}
|
||||
|
||||
if constant.DemoMode {
|
||||
ms.updateDemoMetrics()
|
||||
return ms.metrics
|
||||
}
|
||||
|
||||
cpuPercents, _ := cpu.Percent(0, false)
|
||||
if len(cpuPercents) > 0 {
|
||||
ms.metrics.CPUPercent = cpuPercents[0]
|
||||
}
|
||||
ms.metrics.VMem, _ = mem.VirtualMemory()
|
||||
ms.metrics.DiskUsage, _ = disk.Usage("/")
|
||||
ms.metrics.HostInfo, _ = host.Info()
|
||||
ms.lastUpdate = time.Now()
|
||||
|
||||
// 提供默认值防空指针
|
||||
if ms.metrics.VMem == nil {
|
||||
ms.metrics.VMem = &mem.VirtualMemoryStat{}
|
||||
}
|
||||
if ms.metrics.DiskUsage == nil {
|
||||
ms.metrics.DiskUsage = &disk.UsageStat{}
|
||||
}
|
||||
if ms.metrics.HostInfo == nil {
|
||||
ms.metrics.HostInfo = &host.InfoStat{}
|
||||
}
|
||||
|
||||
return ms.metrics
|
||||
}
|
||||
|
||||
type RuntimeMetrics struct {
|
||||
NumGoroutine int
|
||||
MemStats runtime.MemStats
|
||||
}
|
||||
|
||||
var (
|
||||
runtimeMu sync.RWMutex
|
||||
lastRuntime time.Time
|
||||
cachedRuntime RuntimeMetrics
|
||||
)
|
||||
|
||||
// GetRuntimeMetrics 获取 Go 运行时指标(缓存 2 秒,防止高并发下频繁触发 STW)
|
||||
func (ms *MonitorService) GetRuntimeMetrics() RuntimeMetrics {
|
||||
runtimeMu.Lock()
|
||||
defer runtimeMu.Unlock()
|
||||
|
||||
if time.Since(lastRuntime) < 2*time.Second && cachedRuntime.NumGoroutine > 0 {
|
||||
return cachedRuntime
|
||||
}
|
||||
|
||||
cachedRuntime.NumGoroutine = runtime.NumGoroutine()
|
||||
runtime.ReadMemStats(&cachedRuntime.MemStats)
|
||||
lastRuntime = time.Now()
|
||||
|
||||
return cachedRuntime
|
||||
}
|
||||
|
||||
func (ms *MonitorService) updateDemoMetrics() {
|
||||
ms.metrics.CPUPercent = 10 + rand.Float64()*40 // 10% - 50% 的随机 CPU 波动
|
||||
|
||||
totalMem := uint64(8 * 1024 * 1024 * 1024) // 8GB 内存
|
||||
usedMem := uint64(float64(totalMem) * (0.3 + rand.Float64()*0.3)) // 30% - 60% 随机使用率
|
||||
ms.metrics.VMem = &mem.VirtualMemoryStat{
|
||||
Total: totalMem,
|
||||
Used: usedMem,
|
||||
UsedPercent: float64(usedMem) / float64(totalMem) * 100,
|
||||
}
|
||||
|
||||
totalDisk := uint64(500 * 1024 * 1024 * 1024) // 500GB 硬盘
|
||||
usedDisk := uint64(float64(totalDisk) * 0.45) // 固定 45% 使用率
|
||||
ms.metrics.DiskUsage = &disk.UsageStat{
|
||||
Total: totalDisk,
|
||||
Used: usedDisk,
|
||||
UsedPercent: float64(usedDisk) / float64(totalDisk) * 100,
|
||||
}
|
||||
|
||||
ms.metrics.HostInfo = &host.InfoStat{
|
||||
Platform: "Demo Environment",
|
||||
OS: "linux",
|
||||
Uptime: uint64(time.Now().Unix() - 1700000000), // 生成一个较长且持续增加的运行时间
|
||||
}
|
||||
ms.lastUpdate = time.Now()
|
||||
}
|
||||
@@ -0,0 +1,564 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/engigu/taskpool/internal/sdk/messenger"
|
||||
"gorm.io/gorm"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// NotifyChannel 通知渠道配置
|
||||
type NotifyChannel struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Enabled bool `json:"enabled"`
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
Config map[string]string `json:"config"`
|
||||
}
|
||||
|
||||
// NotifyMessage 通知消息
|
||||
type NotifyMessage struct {
|
||||
Title string `json:"title"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
// NotifyResult 发送结果
|
||||
type NotifyResult struct {
|
||||
Success bool `json:"success"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// SupportedChannelTypes 支持的渠道类型
|
||||
var SupportedChannelTypes = []map[string]string{
|
||||
{"type": messenger.ChannelTelegram, "label": "Telegram"},
|
||||
{"type": messenger.ChannelBark, "label": "Bark"},
|
||||
{"type": messenger.ChannelDtalk, "label": "钉钉"},
|
||||
{"type": messenger.ChannelQyWeiXin, "label": "企业微信"},
|
||||
{"type": messenger.ChannelFeishu, "label": "飞书"},
|
||||
{"type": messenger.ChannelEmail, "label": "邮件"},
|
||||
{"type": messenger.ChannelCustom, "label": "自定义Webhook"},
|
||||
{"type": messenger.ChannelNtfy, "label": "Ntfy"},
|
||||
{"type": messenger.ChannelGotify, "label": "Gotify"},
|
||||
{"type": messenger.ChannelPushMe, "label": "PushMe"},
|
||||
// {"type": messenger.ChannelWeChatOFAccount, "label": "微信公众号"},
|
||||
{"type": messenger.ChannelAliyunSMS, "label": "阿里云短信"},
|
||||
{"type": messenger.ChannelPushPlus, "label": "PushPlus"},
|
||||
{"type": messenger.ChannelVoceChat, "label": "VoceChat"},
|
||||
{"type": messenger.ChannelWxPusher, "label": "WxPusher"},
|
||||
}
|
||||
|
||||
// SupportedEvents 支持的事件类型
|
||||
var SupportedEvents = []map[string]string{
|
||||
{"type": constant.EventUserLogin, "label": "用户登录", "binding_type": constant.BindingTypeSystem},
|
||||
{"type": constant.EventBruteForceLogin, "label": "密码多次错误", "binding_type": constant.BindingTypeSystem},
|
||||
{"type": constant.EventPasswordChanged, "label": "密码修改", "binding_type": constant.BindingTypeSystem},
|
||||
{"type": constant.EventTaskSuccess, "label": "任务成功", "binding_type": constant.BindingTypeTask},
|
||||
{"type": constant.EventTaskFailed, "label": "任务失败", "binding_type": constant.BindingTypeTask},
|
||||
{"type": constant.EventTaskTimeout, "label": "任务超时", "binding_type": constant.BindingTypeTask},
|
||||
}
|
||||
|
||||
type NotificationService struct {
|
||||
settingsService *SettingsService
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
func NewNotificationService() *NotificationService {
|
||||
return &NotificationService{
|
||||
settingsService: NewSettingsService(),
|
||||
}
|
||||
}
|
||||
|
||||
// GetChannels 获取所有渠道
|
||||
func (s *NotificationService) GetChannels() []NotifyChannel {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.getChannelsInternal()
|
||||
}
|
||||
|
||||
// SaveChannel 保存/更新渠道
|
||||
func (s *NotificationService) SaveChannel(channel NotifyChannel) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
configJSON, err := json.Marshal(channel.Config)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if channel.ID == "" {
|
||||
// 新建
|
||||
channel.ID = utils.GenerateID()
|
||||
notifyWay := &models.NotifyWay{
|
||||
ID: channel.ID,
|
||||
Name: channel.Name,
|
||||
Type: channel.Type,
|
||||
Config: models.BigText(configJSON),
|
||||
Enabled: utils.BoolPtr(channel.Enabled),
|
||||
}
|
||||
return database.DB.Create(notifyWay).Error
|
||||
}
|
||||
|
||||
// 更新
|
||||
updates := map[string]interface{}{
|
||||
"name": channel.Name,
|
||||
"type": channel.Type,
|
||||
"config": models.BigText(configJSON),
|
||||
"enabled": &channel.Enabled,
|
||||
}
|
||||
return database.DB.Model(&models.NotifyWay{}).Where("id = ?", channel.ID).Updates(updates).Error
|
||||
}
|
||||
|
||||
// DeleteChannel 删除渠道
|
||||
func (s *NotificationService) DeleteChannel(id string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
// 检查渠道是否存在
|
||||
var count int64
|
||||
database.DB.Model(&models.NotifyWay{}).Where("id = ?", id).Count(&count)
|
||||
if count == 0 {
|
||||
return fmt.Errorf("渠道 %s 不存在", id)
|
||||
}
|
||||
|
||||
// 删除渠道
|
||||
if err := database.DB.Where("id = ?", id).Delete(&models.NotifyWay{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 同时清理事件绑定中引用此渠道的配置
|
||||
if err := database.DB.Where("way_id = ?", id).Delete(&models.NotifyBinding{}).Error; err != nil {
|
||||
logger.Errorf("[Notify] 清理事件绑定失败: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetBindings 获取事件绑定列表(新接口,用于前端展示)
|
||||
func (s *NotificationService) GetBindings() []models.NotifyBinding {
|
||||
var bindings []models.NotifyBinding
|
||||
database.DB.Find(&bindings)
|
||||
return bindings
|
||||
}
|
||||
|
||||
// SaveBinding 保存事件绑定
|
||||
func (s *NotificationService) SaveBinding(binding *models.NotifyBinding) error {
|
||||
if binding.ID == "" {
|
||||
// 检查是否已经存在相同的绑定(避免重复点击导致多个记录)
|
||||
var existing models.NotifyBinding
|
||||
res := database.DB.Where("type = ? AND event = ? AND way_id = ? AND data_id = ?",
|
||||
binding.Type, binding.Event, binding.WayID, binding.DataID).Limit(1).Find(&existing)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
// 如果已存在且未删除,更新现有记录(特别是 Extra 字段)
|
||||
existing.Extra = binding.Extra
|
||||
err := database.DB.Save(&existing).Error
|
||||
if err == nil {
|
||||
*binding = existing
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
binding.ID = utils.GenerateID()
|
||||
return database.DB.Create(binding).Error
|
||||
}
|
||||
return database.DB.Save(binding).Error
|
||||
}
|
||||
|
||||
// BatchSaveBindings 批量保存事件绑定
|
||||
func (s *NotificationService) BatchSaveBindings(bindingType, dataID string, bindings []models.NotifyBinding) error {
|
||||
return database.DB.Transaction(func(tx *gorm.DB) error {
|
||||
// 如果指定了 dataID,先清理该对象的所有现有绑定
|
||||
if dataID != "" {
|
||||
if err := tx.Where("type = ? AND data_id = ?", bindingType, dataID).Delete(&models.NotifyBinding{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 批量插入新绑定
|
||||
for i := range bindings {
|
||||
bindings[i].ID = utils.GenerateID()
|
||||
bindings[i].Type = bindingType
|
||||
bindings[i].DataID = dataID
|
||||
if err := tx.Create(&bindings[i]).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// DeleteBinding 删除事件绑定
|
||||
func (s *NotificationService) DeleteBinding(id string) error {
|
||||
return database.DB.Where("id = ?", id).Delete(&models.NotifyBinding{}).Error
|
||||
}
|
||||
|
||||
// GetBindingsByEvent 根据事件类型和数据ID获取绑定
|
||||
func (s *NotificationService) GetBindingsByEvent(bindingType, event, dataID string) []models.NotifyBinding {
|
||||
var bindings []models.NotifyBinding
|
||||
|
||||
// 如果是任务事件且带有 dataID,只获取特定任务的绑定(禁用全局任务配置)
|
||||
if bindingType == constant.BindingTypeTask && dataID != "" {
|
||||
database.DB.Where("type = ? AND event = ? AND data_id = ?", constant.BindingTypeTask, event, dataID).Find(&bindings)
|
||||
return bindings
|
||||
}
|
||||
|
||||
// 对于系统事件或其他情况
|
||||
query := database.DB.Where("event = ?", event)
|
||||
if bindingType != "" {
|
||||
query = query.Where("type = ?", bindingType)
|
||||
}
|
||||
|
||||
if dataID != "" {
|
||||
query = query.Where("data_id = ?", dataID)
|
||||
} else {
|
||||
query = query.Where("data_id = ? OR data_id IS NULL", "")
|
||||
}
|
||||
|
||||
query.Find(&bindings)
|
||||
return bindings
|
||||
}
|
||||
|
||||
// SendToChannel 使用 messenger SDK 发送通知到指定渠道
|
||||
func (s *NotificationService) SendToChannel(channel NotifyChannel, msg *NotifyMessage) *NotifyResult {
|
||||
result, err := messenger.Send(channel.Type, messenger.ChannelConfig(channel.Config), &messenger.Message{
|
||||
Title: msg.Title,
|
||||
Text: msg.Text,
|
||||
})
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"title": msg.Title,
|
||||
"content": msg.Text,
|
||||
"channel_id": channel.ID,
|
||||
"channel_name": channel.Name,
|
||||
"success": false,
|
||||
"error_msg": "",
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
payload["error_msg"] = err.Error()
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventNotifySent,
|
||||
Payload: payload,
|
||||
})
|
||||
return &NotifyResult{Success: false, Error: err.Error()}
|
||||
}
|
||||
|
||||
if !result.Success {
|
||||
payload["error_msg"] = result.Error
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventNotifySent,
|
||||
Payload: payload,
|
||||
})
|
||||
return &NotifyResult{Success: false, Error: result.Error}
|
||||
}
|
||||
|
||||
payload["success"] = true
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventNotifySent,
|
||||
Payload: payload,
|
||||
})
|
||||
|
||||
return &NotifyResult{Success: true}
|
||||
}
|
||||
|
||||
// SendByChannelID 根据渠道ID发送通知
|
||||
func (s *NotificationService) SendByChannelID(channelID string, msg *NotifyMessage) *NotifyResult {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
var notifyWay models.NotifyWay
|
||||
res := database.DB.Where("id = ?", channelID).Limit(1).Find(¬ifyWay)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return &NotifyResult{Success: false, Error: "渠道不存在"}
|
||||
}
|
||||
|
||||
if !utils.DerefBool(notifyWay.Enabled, true) {
|
||||
return &NotifyResult{Success: false, Error: "渠道已禁用"}
|
||||
}
|
||||
|
||||
var config map[string]string
|
||||
if err := json.Unmarshal([]byte(notifyWay.Config), &config); err != nil {
|
||||
return &NotifyResult{Success: false, Error: "渠道配置解析失败"}
|
||||
}
|
||||
|
||||
ch := NotifyChannel{
|
||||
ID: notifyWay.ID,
|
||||
Name: notifyWay.Name,
|
||||
Type: notifyWay.Type,
|
||||
Enabled: utils.DerefBool(notifyWay.Enabled, true),
|
||||
Config: config,
|
||||
}
|
||||
return s.SendToChannel(ch, msg)
|
||||
}
|
||||
|
||||
// SubscribeEvents 注册通知服务自身为事件流的订阅者
|
||||
func (s *NotificationService) SubscribeEvents(bus *eventbus.EventBus) {
|
||||
// 系统事件
|
||||
systemEvents := []string{constant.EventUserLogin, constant.EventBruteForceLogin, constant.EventPasswordChanged}
|
||||
for _, evt := range systemEvents {
|
||||
bus.Subscribe(evt, s.handleEvent(constant.BindingTypeSystem))
|
||||
}
|
||||
|
||||
// 任务事件
|
||||
taskEvents := []string{constant.EventTaskSuccess, constant.EventTaskFailed, constant.EventTaskTimeout}
|
||||
for _, evt := range taskEvents {
|
||||
bus.Subscribe(evt, s.handleEvent(constant.BindingTypeTask))
|
||||
}
|
||||
|
||||
// 通用系统通知
|
||||
bus.Subscribe(constant.EventSystemNotice, s.handleEvent(constant.BindingTypeSystem))
|
||||
}
|
||||
|
||||
var ansiRegexp = regexp.MustCompile(`[\x1b\x9b][\[()#;?]*([0-9]{1,4}(;[0-9]{0,4})*)?[0-9A-ORZcf-nqry=><]`)
|
||||
|
||||
// stripAnsi 移除字符串中的 ANSI 转义码(如颜色代码)
|
||||
func stripAnsi(str string) string {
|
||||
return ansiRegexp.ReplaceAllString(str, "")
|
||||
}
|
||||
|
||||
// parseTemplate 简单的 {{key}} 模板替换
|
||||
func (s *NotificationService) parseTemplate(tmpl string, payload map[string]interface{}) string {
|
||||
result := tmpl
|
||||
for k, v := range payload {
|
||||
placeholder := fmt.Sprintf("{{%s}}", k)
|
||||
valStr := fmt.Sprintf("%v", v)
|
||||
result = strings.ReplaceAll(result, placeholder, valStr)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// getDefaultMessage 兜底默认消息内容
|
||||
func (s *NotificationService) getDefaultMessage(eventType string, payload map[string]interface{}) (string, string) {
|
||||
var title, text string
|
||||
switch eventType {
|
||||
case constant.EventUserLogin:
|
||||
status, _ := payload["status"].(string)
|
||||
if status == "success" {
|
||||
title = "用户登录成功"
|
||||
text = fmt.Sprintf("用户 %v 在 IP %v 登录成功", payload["username"], payload["ip"])
|
||||
} else {
|
||||
title = "用户登录失败"
|
||||
reason, _ := payload["message"].(string)
|
||||
text = fmt.Sprintf("用户 %v 在 IP %v 登录失败\n原因: %v", payload["username"], payload["ip"], reason)
|
||||
}
|
||||
case constant.EventBruteForceLogin:
|
||||
title = "系统安全警告"
|
||||
text = fmt.Sprintf("检测到 IP %v 正在尝试暴力破解用户 %v", payload["ip"], payload["username"])
|
||||
case constant.EventPasswordChanged:
|
||||
title = "账户安全通知"
|
||||
text = fmt.Sprintf("用户 %v 刚刚修改了密码", payload["username"])
|
||||
case constant.EventTaskSuccess:
|
||||
title = fmt.Sprintf("任务[%v] 成功", payload["task_name"])
|
||||
text = fmt.Sprintf("任务 #%v %v\n状态: 成功\n执行时间: %v\n耗时: %vms", payload["task_id"], payload["task_name"], payload["start_time"], payload["duration"])
|
||||
case constant.EventTaskFailed:
|
||||
title = fmt.Sprintf("任务[%v] 失败", payload["task_name"])
|
||||
if errStr, ok := payload["error"]; ok {
|
||||
text = fmt.Sprintf("任务 #%v %v\n执行失败\n执行时间: %v\n错误: %v", payload["task_id"], payload["task_name"], payload["start_time"], errStr)
|
||||
} else {
|
||||
text = fmt.Sprintf("任务 #%v %v\n执行失败\n状态: %v\n执行时间: %v\n耗时: %vms", payload["task_id"], payload["task_name"], payload["status"], payload["start_time"], payload["duration"])
|
||||
}
|
||||
case constant.EventTaskTimeout:
|
||||
title = fmt.Sprintf("任务[%v] 超时", payload["task_name"])
|
||||
text = fmt.Sprintf("任务 #%v %v\n执行超时\n执行时间: %v\n耗时: %vms", payload["task_id"], payload["task_name"], payload["start_time"], payload["duration"])
|
||||
}
|
||||
return title, text
|
||||
}
|
||||
|
||||
// resolveEvent 解析不同事件类型,返回对应的模板Key、静态内容(非模板事件)和原始任务输出
|
||||
func (s *NotificationService) resolveEvent(eventType string, payload map[string]interface{}) (tmplTitleKey, tmplTextKey, title, text, rawOutput string, ok bool) {
|
||||
switch eventType {
|
||||
case constant.EventUserLogin:
|
||||
tmplTitleKey = constant.KeyNotifyTemplateUserLoginTitle
|
||||
tmplTextKey = constant.KeyNotifyTemplateUserLoginText
|
||||
// 特殊处理登录状态
|
||||
status, _ := payload["status"].(string)
|
||||
if status == "success" {
|
||||
payload["status_label"] = "成功"
|
||||
} else {
|
||||
payload["status_label"] = "失败"
|
||||
}
|
||||
|
||||
case constant.EventBruteForceLogin:
|
||||
tmplTitleKey = constant.KeyNotifyTemplateBruteForceLoginTitle
|
||||
tmplTextKey = constant.KeyNotifyTemplateBruteForceLoginText
|
||||
|
||||
case constant.EventPasswordChanged:
|
||||
tmplTitleKey = constant.KeyNotifyTemplatePasswordChangedTitle
|
||||
tmplTextKey = constant.KeyNotifyTemplatePasswordChangedText
|
||||
|
||||
case constant.EventTaskSuccess, constant.EventTaskFailed, constant.EventTaskTimeout:
|
||||
switch eventType {
|
||||
case constant.EventTaskSuccess:
|
||||
tmplTitleKey = constant.KeyNotifyTemplateTaskSuccessTitle
|
||||
tmplTextKey = constant.KeyNotifyTemplateTaskSuccessText
|
||||
case constant.EventTaskFailed:
|
||||
tmplTitleKey = constant.KeyNotifyTemplateTaskFailedTitle
|
||||
tmplTextKey = constant.KeyNotifyTemplateTaskFailedText
|
||||
case constant.EventTaskTimeout:
|
||||
tmplTitleKey = constant.KeyNotifyTemplateTaskTimeoutTitle
|
||||
tmplTextKey = constant.KeyNotifyTemplateTaskTimeoutText
|
||||
}
|
||||
|
||||
// 处理输出内容,避免过长
|
||||
if output, ok := payload["output"].(string); ok {
|
||||
rawOutput = output
|
||||
// trimmed := utils.TrimLastRunes(output, 1000)
|
||||
// if len(trimmed) < len(output) {
|
||||
// payload["output"] = trimmed + "\n...(截断)"
|
||||
// }
|
||||
}
|
||||
|
||||
case constant.EventSystemNotice:
|
||||
title, _ = payload["title"].(string)
|
||||
text, _ = payload["content"].(string)
|
||||
default:
|
||||
return "", "", "", "", "", false
|
||||
}
|
||||
return tmplTitleKey, tmplTextKey, title, text, rawOutput, true
|
||||
}
|
||||
|
||||
// buildMessage 匹配并解析模板内容,提供兜底消息并拼接全局前缀
|
||||
func (s *NotificationService) buildMessage(eventType string, tmplTitleKey, tmplTextKey, defaultTitle, defaultText, prefix string, payload map[string]interface{}) (title, text string) {
|
||||
title = defaultTitle
|
||||
text = defaultText
|
||||
|
||||
if tmplTitleKey != "" {
|
||||
tmplTitle := s.settingsService.Get(constant.SectionNotify, tmplTitleKey)
|
||||
tmplText := s.settingsService.Get(constant.SectionNotify, tmplTextKey)
|
||||
|
||||
if tmplTitle != "" {
|
||||
title = s.parseTemplate(tmplTitle, payload)
|
||||
}
|
||||
if tmplText != "" {
|
||||
text = s.parseTemplate(tmplText, payload)
|
||||
}
|
||||
|
||||
// 如果模板为空,使用兜底默认逻辑(保持向上兼容)
|
||||
if title == "" || text == "" {
|
||||
title, text = s.getDefaultMessage(eventType, payload)
|
||||
}
|
||||
}
|
||||
|
||||
// 添加全局前缀
|
||||
if prefix != "" && title != "" {
|
||||
title = fmt.Sprintf("%s %s", prefix, title)
|
||||
}
|
||||
return title, text
|
||||
}
|
||||
|
||||
// handleEvent 处理事件订阅并发送通知
|
||||
func (s *NotificationService) handleEvent(bindingType string) eventbus.Handler {
|
||||
return func(e eventbus.Event) {
|
||||
payload, ok := e.Payload.(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
var dataID string
|
||||
if id, ok := payload["task_id"].(string); ok {
|
||||
dataID = id
|
||||
}
|
||||
|
||||
// 获取全局前缀并解析事件数据
|
||||
prefix := s.settingsService.Get(constant.SectionNotify, constant.KeyNotifyPrefix)
|
||||
tmplTitleKey, tmplTextKey, title, text, rawOutput, ok := s.resolveEvent(e.Type, payload)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
// 构建最终的通知标题和正文文本
|
||||
title, text = s.buildMessage(e.Type, tmplTitleKey, tmplTextKey, title, text, prefix, payload)
|
||||
|
||||
bindings := s.GetBindingsByEvent(bindingType, e.Type, dataID)
|
||||
if len(bindings) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
var cleanLog string
|
||||
if rawOutput != "" {
|
||||
// cleanLog = stripAnsi(rawOutput)
|
||||
cleanLog = rawOutput
|
||||
}
|
||||
|
||||
channels := s.GetChannels()
|
||||
channelMap := make(map[string]NotifyChannel)
|
||||
for _, ch := range channels {
|
||||
channelMap[ch.ID] = ch
|
||||
}
|
||||
|
||||
for _, binding := range bindings {
|
||||
ch, ok := channelMap[binding.WayID]
|
||||
if !ok || !ch.Enabled {
|
||||
continue
|
||||
}
|
||||
|
||||
// 克隆文本以便修改
|
||||
currentText := text
|
||||
|
||||
// 解析额外配置
|
||||
var extra models.BindingExtra
|
||||
if binding.Extra != "" {
|
||||
_ = json.Unmarshal([]byte(binding.Extra), &extra)
|
||||
}
|
||||
// 默认日志限制为 1000
|
||||
if extra.LogLimit <= 0 {
|
||||
extra.LogLimit = 1000
|
||||
}
|
||||
|
||||
// 如果开启了日志推送
|
||||
if extra.EnableLog {
|
||||
if cleanLog != "" {
|
||||
trimmed := utils.TrimLastRunes(cleanLog, extra.LogLimit)
|
||||
if len(trimmed) < len(cleanLog) {
|
||||
trimmed = "...\n" + trimmed
|
||||
}
|
||||
currentText += "\n\n[执行日志]\n" + trimmed
|
||||
}
|
||||
}
|
||||
|
||||
go func(channel NotifyChannel, msgTitle, msgText string) {
|
||||
result := s.SendToChannel(channel, &NotifyMessage{Title: msgTitle, Text: msgText})
|
||||
if !result.Success {
|
||||
logger.Warnf("[Notify] 发送事件 %s 到渠道 %s(%s) 失败: %s", e.Type, channel.Name, channel.Type, result.Error)
|
||||
}
|
||||
}(ch, title, currentText)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- 内部方法 ---
|
||||
|
||||
// getChannelsInternal 从 notify_ways 表中读取所有渠道配置
|
||||
func (s *NotificationService) getChannelsInternal() []NotifyChannel {
|
||||
var notifyWays []models.NotifyWay
|
||||
database.DB.Find(¬ifyWays)
|
||||
|
||||
channels := make([]NotifyChannel, 0, len(notifyWays))
|
||||
for _, nw := range notifyWays {
|
||||
var config map[string]string
|
||||
if err := json.Unmarshal([]byte(nw.Config), &config); err != nil {
|
||||
logger.Warnf("[Notify] 解析渠道 %s 配置失败: %v", nw.ID, err)
|
||||
continue
|
||||
}
|
||||
channels = append(channels, NotifyChannel{
|
||||
ID: nw.ID,
|
||||
Name: nw.Name,
|
||||
Type: nw.Type,
|
||||
Enabled: utils.DerefBool(nw.Enabled, true),
|
||||
CreatedAt: nw.CreatedAt,
|
||||
Config: config,
|
||||
})
|
||||
}
|
||||
return channels
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
package relation
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/rs/xid"
|
||||
)
|
||||
|
||||
type DataRelationService struct{}
|
||||
|
||||
var DataRelation = &DataRelationService{}
|
||||
|
||||
// SaveTags 保存带有 Storage (文本标签) 的关系映射
|
||||
func (s *DataRelationService) SaveTags(dataID string, relType string, tagsStr string) {
|
||||
database.DB.Where("data_id = ? AND type = ?", dataID, relType).Delete(&models.DataRelation{})
|
||||
if tagsStr == "" {
|
||||
return
|
||||
}
|
||||
tags := strings.Split(tagsStr, ",")
|
||||
for _, tag := range tags {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
var storage models.DataStorage
|
||||
res := database.DB.Where("type = ? AND name = ?", relType, tag).Limit(1).Find(&storage)
|
||||
if res.RowsAffected == 0 {
|
||||
storage = models.DataStorage{
|
||||
ID: xid.New().String(),
|
||||
Type: relType,
|
||||
Name: tag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
database.DB.Create(&storage)
|
||||
}
|
||||
relation := models.DataRelation{
|
||||
ID: xid.New().String(),
|
||||
DataID: dataID,
|
||||
RelateID: storage.ID,
|
||||
Type: relType,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
database.DB.Create(&relation)
|
||||
}
|
||||
}
|
||||
|
||||
// LoadTags 加载带有 Storage (文本标签) 的映射,返回 map[DataID][]TagName
|
||||
func (s *DataRelationService) LoadTags(dataIDs []string, relType string) map[string][]string {
|
||||
if len(dataIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
var relations []models.DataRelation
|
||||
database.DB.Where("data_id IN ? AND type = ?", dataIDs, relType).Find(&relations)
|
||||
|
||||
if len(relations) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var relateIDs []string
|
||||
for _, r := range relations {
|
||||
relateIDs = append(relateIDs, r.RelateID)
|
||||
}
|
||||
|
||||
var storages []models.DataStorage
|
||||
database.DB.Where("id IN ?", relateIDs).Find(&storages)
|
||||
|
||||
storageMap := make(map[string]string)
|
||||
for _, storage := range storages {
|
||||
storageMap[storage.ID] = storage.Name
|
||||
}
|
||||
|
||||
resultMap := make(map[string][]string)
|
||||
for _, r := range relations {
|
||||
if name, ok := storageMap[r.RelateID]; ok {
|
||||
resultMap[r.DataID] = append(resultMap[r.DataID], name)
|
||||
}
|
||||
}
|
||||
return resultMap
|
||||
}
|
||||
|
||||
// SaveRelations 保存单纯的关系映射 (例如 ID关联)
|
||||
func (s *DataRelationService) SaveRelations(dataID string, relType string, relateIDsStr string) {
|
||||
database.DB.Where("data_id = ? AND type = ?", dataID, relType).Delete(&models.DataRelation{})
|
||||
if relateIDsStr == "" {
|
||||
return
|
||||
}
|
||||
ids := strings.Split(relateIDsStr, ",")
|
||||
for _, relateID := range ids {
|
||||
relateID = strings.TrimSpace(relateID)
|
||||
if relateID == "" {
|
||||
continue
|
||||
}
|
||||
relation := models.DataRelation{
|
||||
ID: xid.New().String(),
|
||||
DataID: dataID,
|
||||
RelateID: relateID,
|
||||
Type: relType,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
database.DB.Create(&relation)
|
||||
}
|
||||
}
|
||||
|
||||
// LoadRelations 加载单纯的关系映射,返回 map[DataID][]RelateID
|
||||
func (s *DataRelationService) LoadRelations(dataIDs []string, relType string) map[string][]string {
|
||||
if len(dataIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
var relations []models.DataRelation
|
||||
database.DB.Where("data_id IN ? AND type = ?", dataIDs, relType).Find(&relations)
|
||||
|
||||
resultMap := make(map[string][]string)
|
||||
for _, r := range relations {
|
||||
resultMap[r.DataID] = append(resultMap[r.DataID], r.RelateID)
|
||||
}
|
||||
return resultMap
|
||||
}
|
||||
|
||||
// CleanRelations 删除某种类型的所有关联映射
|
||||
func (s *DataRelationService) CleanRelations(dataID string, relType string) {
|
||||
database.DB.Where("data_id = ? AND type = ?", dataID, relType).Delete(&models.DataRelation{})
|
||||
}
|
||||
|
||||
// GetAllTags 获取全局范围内某种类型的所有的 Tag Name
|
||||
func (s *DataRelationService) GetAllTags(relType string) ([]string, error) {
|
||||
var storages []models.DataStorage
|
||||
err := database.DB.Where("type = ?", relType).Find(&storages).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var tags []string
|
||||
for _, s := range storages {
|
||||
tags = append(tags, s.Name)
|
||||
}
|
||||
return tags, nil
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/relation"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// ParseRepoScriptsAndAddCron 扫描仓库目录中的脚本,解析 cron 和环境注释,并注册任务
|
||||
func ParseRepoScriptsAndAddCron(taskID string, logWriter io.Writer, forceCommentToTask bool) ([]string, []string) {
|
||||
// 帮助函数:如果提供了 logWriter,则将日志输出到该处
|
||||
log := func(format string, a ...interface{}) {
|
||||
msg := fmt.Sprintf(format, a...)
|
||||
if !strings.HasSuffix(msg, "\n") {
|
||||
msg += "\n"
|
||||
}
|
||||
if logWriter != nil {
|
||||
logWriter.Write([]byte(msg))
|
||||
}
|
||||
}
|
||||
|
||||
var repoTask models.Task
|
||||
res := database.DB.Where("id = ?", taskID).Limit(1).Find(&repoTask)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if repoTask.Type != constant.TaskTypeRepo {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
var repoCfg models.RepoConfig
|
||||
if err := json.Unmarshal([]byte(repoTask.Config), &repoCfg); err != nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// 如果命令行强制开启,则覆盖配置
|
||||
if forceCommentToTask {
|
||||
repoCfg.CommentToTask = "true"
|
||||
}
|
||||
|
||||
// 1. 确定解析策略
|
||||
strategy := GetParserStrategy(repoCfg.RepoSource)
|
||||
|
||||
// 目标路径
|
||||
targetPath := repoCfg.TargetPath
|
||||
if targetPath == "" {
|
||||
targetPath = repoTask.WorkDir
|
||||
} else if !filepath.IsAbs(targetPath) {
|
||||
targetPath = filepath.Join(utils.ResolveAbsScriptsDir(), targetPath)
|
||||
}
|
||||
if targetPath == "" {
|
||||
return nil, nil
|
||||
}
|
||||
targetPath = filepath.Clean(targetPath)
|
||||
|
||||
// 获取仓库标识符
|
||||
repoId := utils.GetRepoIdentifier(repoCfg.SourceURL, repoCfg.Branch)
|
||||
|
||||
gitDir := filepath.Join(targetPath, ".git")
|
||||
if !isDir(targetPath) || !pathExists(gitDir) {
|
||||
repoPath := filepath.Join(targetPath, repoId)
|
||||
if pathExists(repoPath) {
|
||||
targetPath = repoPath
|
||||
}
|
||||
}
|
||||
|
||||
if !pathExists(targetPath) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// 同步过程中使用的标签
|
||||
tag := fmt.Sprintf("%s", repoId)
|
||||
|
||||
exts := getValidExtensions(repoCfg.Extensions)
|
||||
|
||||
log("\n----------------------------------------")
|
||||
log(" 开始扫描脚本并自动注册定时任务 ")
|
||||
log("----------------------------------------")
|
||||
|
||||
foundSourceIDs := make(map[string]bool)
|
||||
var upsertedIDs []string
|
||||
var deletedIDs []string
|
||||
newTaskCount := 0
|
||||
updateTaskCount := 0
|
||||
|
||||
filepath.WalkDir(targetPath, func(path string, d fs.DirEntry, err error) error {
|
||||
if err != nil || d.IsDir() {
|
||||
return nil
|
||||
}
|
||||
|
||||
if strings.Contains(path, ".git") {
|
||||
return nil
|
||||
}
|
||||
|
||||
ext := filepath.Ext(path)
|
||||
if !strategy.SupportExtension(ext, exts) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 2. 元数据提取
|
||||
taskName, taskCron := strategy.ExtractMeta(path, ext, repoCfg)
|
||||
|
||||
// 3. 过滤处理
|
||||
relRepoPath, _ := filepath.Rel(targetPath, path)
|
||||
filename := filepath.Base(path)
|
||||
if !strategy.ShouldProcess(relRepoPath, filename, repoCfg) {
|
||||
return nil
|
||||
}
|
||||
|
||||
if taskName != "" && taskCron != "" && repoCfg.AutoAddCron {
|
||||
// 获取脚本相对于数据目录的路径
|
||||
absScriptsDir := utils.ResolveAbsScriptsDir()
|
||||
// absTargetPath, _ := filepath.Abs(targetPath)
|
||||
absPath, _ := filepath.Abs(path)
|
||||
|
||||
// 计算 SourceID: 相对于脚本目录的完整路径,并清洗特殊符号
|
||||
relPath, _ := filepath.Rel(absScriptsDir, absPath)
|
||||
sourceID := sanitizeIdentifier(relPath)
|
||||
|
||||
// 替换绝对路径为代号 $SCRIPTS_DIR$
|
||||
displayPath := path
|
||||
displayWorkDir := targetPath
|
||||
if strings.HasPrefix(absPath, absScriptsDir) {
|
||||
// 命令仅使用文件名,工作目录设置为脚本所在目录
|
||||
displayPath = filepath.Base(path)
|
||||
relDir, _ := filepath.Rel(absScriptsDir, filepath.Dir(absPath))
|
||||
if relDir == "." {
|
||||
displayWorkDir = constant.ScriptsDirPlaceholder
|
||||
} else {
|
||||
displayWorkDir = filepath.ToSlash(filepath.Join(constant.ScriptsDirPlaceholder, relDir))
|
||||
}
|
||||
}
|
||||
|
||||
// 找到任务,进行保存
|
||||
command := getCommandByExt(ext, displayPath)
|
||||
taskID, isNew := upsertRepoTask(&repoTask, sourceID, taskName, command, taskCron, displayWorkDir, tag)
|
||||
|
||||
if isNew {
|
||||
log("[新增] 任务: %s (%s)", taskName, filename)
|
||||
newTaskCount++
|
||||
} else {
|
||||
log("[更新] 任务: %s (%s)", taskName, filename)
|
||||
updateTaskCount++
|
||||
}
|
||||
foundSourceIDs[sourceID] = true
|
||||
upsertedIDs = append(upsertedIDs, taskID)
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
|
||||
// 清理该仓库下不再存在的旧脚本任务
|
||||
deletedTaskCount := 0
|
||||
var oldTasks []models.Task
|
||||
if err := database.DB.Where("repo_task_id = ?", repoTask.ID).Find(&oldTasks).Error; err == nil {
|
||||
for _, ot := range oldTasks {
|
||||
if !foundSourceIDs[ot.SourceID] {
|
||||
log("[移除] 脚本已不存在,删除对应任务: %s", ot.Name)
|
||||
deletedTaskCount++
|
||||
deletedIDs = append(deletedIDs, ot.ID)
|
||||
|
||||
// 删除关联
|
||||
relation.DataRelation.CleanRelations(ot.ID, constant.RelationTypeTaskTag)
|
||||
relation.DataRelation.CleanRelations(ot.ID, constant.RelationTypeTaskEnv)
|
||||
|
||||
database.DB.Unscoped().Where("id = ?", ot.ID).Delete(&models.Task{})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
log("\n扫描完成: [新增 %d] [更新 %d] [移除 %d]", newTaskCount, updateTaskCount, deletedTaskCount)
|
||||
log("----------------------------------------")
|
||||
return upsertedIDs, deletedIDs
|
||||
}
|
||||
|
||||
// upsertRepoTask 处理来自仓库的任务的创建或更新
|
||||
func upsertRepoTask(parentTask *models.Task, sourceID, name, command, cron, workDir, tag string) (string, bool) {
|
||||
defaultTaskConfig := `{"$task_all_envs":true,"$task_concurrency":0}`
|
||||
var existing models.Task
|
||||
tx := database.DB.Where("source_id = ? AND repo_task_id = ?", sourceID, parentTask.ID).Limit(1).Find(&existing)
|
||||
|
||||
if tx.RowsAffected > 0 {
|
||||
// 更新操作
|
||||
existing.Name = name
|
||||
existing.Command = models.BigText(command)
|
||||
existing.Schedule = normalizeCron(cron)
|
||||
existing.Languages = parentTask.Languages
|
||||
existing.SourceID = sourceID
|
||||
existing.RepoTaskID = parentTask.ID
|
||||
existing.WorkDir = workDir
|
||||
// 如果原配置为空或者是 {},则应用默认配置
|
||||
if string(existing.Config) == "" || string(existing.Config) == "{}" {
|
||||
existing.Config = models.BigText(defaultTaskConfig)
|
||||
}
|
||||
// 默认开启按条数清理30条
|
||||
if existing.CleanConfig == "" {
|
||||
existing.CleanConfig = `{"type":"count","keep":30}`
|
||||
}
|
||||
// 显式白名单模式:只更新脚本核心相关的字段,其他所有字段(如 Enabled, Pin, Remark 等)均不触碰
|
||||
database.DB.Model(&existing).
|
||||
Select("Name", "Command", "Schedule", "WorkDir", "Languages").
|
||||
Updates(&existing)
|
||||
return existing.ID, false
|
||||
} else {
|
||||
// 创建新任务
|
||||
newTask := &models.Task{
|
||||
Name: name,
|
||||
Command: models.BigText(command),
|
||||
Schedule: normalizeCron(cron),
|
||||
Type: "task",
|
||||
TriggerType: constant.TriggerTypeCron,
|
||||
Tags: tag,
|
||||
Languages: parentTask.Languages,
|
||||
Timeout: parentTask.Timeout,
|
||||
Config: models.BigText(defaultTaskConfig),
|
||||
Enabled: utils.BoolPtr(true),
|
||||
WorkDir: workDir,
|
||||
SourceID: sourceID,
|
||||
RepoTaskID: parentTask.ID,
|
||||
CleanConfig: `{"type":"count","keep":30}`,
|
||||
}
|
||||
newTask.ID = utils.GenerateID()
|
||||
database.DB.Create(newTask)
|
||||
|
||||
// 确保将标签同步写入到新的 DataRelation 中
|
||||
if tag != "" {
|
||||
relation.DataRelation.SaveTags(newTask.ID, constant.RelationTypeTaskTag, tag)
|
||||
}
|
||||
|
||||
return newTask.ID, true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
// RepoParserStrategy 定义不同仓库解析策略的接口
|
||||
type RepoParserStrategy interface {
|
||||
// SupportExtension 判断给定后缀的文件是否应该被处理
|
||||
SupportExtension(ext string, exts []string) bool
|
||||
|
||||
// ShouldProcess 应用白名单/黑名单过滤,决定是否处理该文件
|
||||
ShouldProcess(relRepoPath, filename string, cfg models.RepoConfig) bool
|
||||
|
||||
// ExtractMeta 从脚本文件中提取任务元数据(名称和 cron 表达式)
|
||||
ExtractMeta(path string, ext string, cfg models.RepoConfig) (taskName string, taskCron string)
|
||||
}
|
||||
|
||||
// GetParserStrategy 根据来源类型返回相应的策略实现
|
||||
func GetParserStrategy(sourceType string) RepoParserStrategy {
|
||||
switch sourceType {
|
||||
case "ql":
|
||||
return &QinglongStrategy{}
|
||||
default:
|
||||
return &StandardStrategy{}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// QinglongStrategy 实现与青龙兼容的解析逻辑
|
||||
type QinglongStrategy struct{}
|
||||
|
||||
func (s *QinglongStrategy) SupportExtension(ext string, exts []string) bool {
|
||||
for _, e := range exts {
|
||||
if ext == e {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *QinglongStrategy) ShouldProcess(relRepoPath, filename string, cfg models.RepoConfig) bool {
|
||||
// 只有在显式设置了白名单时才进行白名单校验 (青龙行为)
|
||||
if cfg.WhitelistPaths != "" {
|
||||
if !matchesQLPattern(relRepoPath, filename, cfg.WhitelistPaths) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// 校验黑名单
|
||||
if cfg.Blacklist != "" {
|
||||
if matchesQLPattern(relRepoPath, filename, cfg.Blacklist) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *QinglongStrategy) ExtractMeta(path string, ext string, cfg models.RepoConfig) (taskName string, taskCron string) {
|
||||
return ExtractScriptMeta(path, ext)
|
||||
}
|
||||
|
||||
// matchesQLPattern 应用关键字过滤逻辑(正则或包含匹配)
|
||||
func matchesQLPattern(rel, filename string, keywordsStr string) bool {
|
||||
if keywordsStr == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
keywords := splitKeywords(keywordsStr)
|
||||
for _, k := range keywords {
|
||||
// 1. 尝试作为正则整体进行匹配,默认不区分大小写 (?i)
|
||||
pattern := k
|
||||
if !strings.HasPrefix(pattern, "(?i)") {
|
||||
pattern = "(?i)" + pattern
|
||||
}
|
||||
|
||||
reg, err := regexp.Compile(pattern)
|
||||
if err == nil {
|
||||
// 优先匹配文件名(解决 ^jd[^_] 这种锚点在相对路径下失效的问题)
|
||||
if reg.MatchString(filename) || reg.MatchString(rel) {
|
||||
return true
|
||||
}
|
||||
} else {
|
||||
// 回退逻辑:全小写包含判断
|
||||
kLower := strings.ToLower(k)
|
||||
if strings.Contains(strings.ToLower(rel), kLower) || strings.Contains(strings.ToLower(filename), kLower) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
// StandardStrategy 实现默认的解析逻辑
|
||||
type StandardStrategy struct{}
|
||||
|
||||
func (s *StandardStrategy) SupportExtension(ext string, exts []string) bool {
|
||||
for _, e := range exts {
|
||||
if ext == e {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *StandardStrategy) ShouldProcess(relRepoPath, filename string, cfg models.RepoConfig) bool {
|
||||
// 标准策略未来可能具有不同的过滤规则
|
||||
// 目前如果提供了白名单/黑名单,则遵循相同的逻辑,但使用更简单的匹配
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *StandardStrategy) ExtractMeta(path string, ext string, cfg models.RepoConfig) (taskName string, taskCron string) {
|
||||
// 仅在开启兼容 QL 配置时才解析脚本注释
|
||||
if cfg.CommentToTask == "true" {
|
||||
return ExtractScriptMeta(path, ext)
|
||||
}
|
||||
return "", ""
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"fmt"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"github.com/robfig/cron/v3"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var (
|
||||
// envRegex 匹配脚本中的环境名称设置,如 Env("名称")
|
||||
envRegex = regexp.MustCompile(`(?i)(?:new[ \t]+)?Env\(['"]?([^'"]+)['"]?\)`)
|
||||
// cronRegex 匹配脚本中的 cron 表达式设置
|
||||
cronRegex = regexp.MustCompile(`(?i)(?:cron[ \t]*[:=][ \t]*['"]?([^'"\r\n]+))|(?:(?:^|[ \t\*\/])(([0-9\*\/\-,L?#]+[ \t]+){4,5}[0-9\*\/\-,L?#]+))`)
|
||||
// cronFormatRegex 用于校验提取出的字符串是否符合 Cron 表达式格式 (5位或6位)
|
||||
cronFormatRegex = regexp.MustCompile(`^(([0-9\*\/\-,L?#]+)[ \t]+){4,5}([0-9\*\/\-,L?#]+)$`)
|
||||
)
|
||||
|
||||
// ExtractScriptMeta 读取文件以提取任务名称和 cron 表达式
|
||||
func ExtractScriptMeta(path string, ext string) (taskName string, taskCron string) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return "", ""
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
scanner := bufio.NewScanner(f)
|
||||
var firstCommentLine string
|
||||
inBlockComment := false
|
||||
|
||||
// 特殊处理:针对当前文件名的 Cron 关联正则表达式 (对标青龙 perl 逻辑)
|
||||
// 寻找类似 "// 0 0 * * * jd_task.js" 的行
|
||||
fileNameEscaped := regexp.QuoteMeta(filepath.Base(path))
|
||||
associatedCronRegex := regexp.MustCompile(fmt.Sprintf(`(?i)(?:^|[ \t\*\//])(([0-9\*\/\-,L?#]+[ \t]+){4,5}[0-9\*\/\-,L?#]+)[ \t,"]+.*%s`, fileNameEscaped))
|
||||
|
||||
for i := 0; i < 15 && scanner.Scan(); i++ { // 限制扫描范围,避免误匹配代码深处的属性
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// 处理块注释开始/结束
|
||||
if strings.HasPrefix(line, "/*") {
|
||||
inBlockComment = true
|
||||
line = strings.TrimPrefix(line, "/*")
|
||||
line = strings.TrimPrefix(line, "*")
|
||||
line = strings.TrimSpace(line)
|
||||
}
|
||||
if strings.HasSuffix(line, "*/") {
|
||||
inBlockComment = false
|
||||
line = strings.TrimSuffix(line, "*/")
|
||||
line = strings.TrimSpace(line)
|
||||
}
|
||||
|
||||
// 1. 尝试提取任务名称 (优先使用 Env)
|
||||
if taskName == "" {
|
||||
if envMatch := envRegex.FindStringSubmatch(line); len(envMatch) > 1 {
|
||||
taskName = strings.TrimSpace(envMatch[1])
|
||||
} else if strings.Contains(line, "name:") {
|
||||
// 兼容 name: "xxx" 格式
|
||||
nameRegex := regexp.MustCompile(`(?i)name:[ \t]*['"]([^'"]+)['"]`)
|
||||
if nameMatch := nameRegex.FindStringSubmatch(line); len(nameMatch) > 1 {
|
||||
taskName = strings.TrimSpace(nameMatch[1])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 如果还没找到名称,且在注释中,记录第一行非空注释作为备选名称
|
||||
if taskName == "" && (inBlockComment || strings.HasPrefix(line, "//") || strings.HasPrefix(line, "*") || strings.HasPrefix(line, "#")) {
|
||||
cleanLine := line
|
||||
if strings.HasPrefix(line, "//") {
|
||||
cleanLine = strings.TrimPrefix(line, "//")
|
||||
} else if strings.HasPrefix(line, "#") {
|
||||
cleanLine = strings.TrimPrefix(line, "#")
|
||||
} else if strings.HasPrefix(line, "*") {
|
||||
cleanLine = strings.TrimPrefix(line, "*")
|
||||
}
|
||||
cleanLine = strings.TrimSpace(cleanLine)
|
||||
|
||||
// 排除掉包含 "Env" 或 "cron" 的行, 且排除掉可能是路径或URL的行
|
||||
if cleanLine != "" && !strings.Contains(strings.ToLower(cleanLine), "env") &&
|
||||
!strings.Contains(strings.ToLower(cleanLine), "cron") &&
|
||||
!strings.Contains(cleanLine, "http") &&
|
||||
!strings.Contains(cleanLine, "/") &&
|
||||
firstCommentLine == "" {
|
||||
// 且排除掉纯 cron 表达式
|
||||
if !cronRegex.MatchString(cleanLine) {
|
||||
firstCommentLine = cleanLine
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 提取 Cron
|
||||
if taskCron == "" {
|
||||
// A. 优先查找关联了当前文件名的 Cron (对标 QL)
|
||||
if assocMatch := associatedCronRegex.FindStringSubmatch(line); len(assocMatch) > 1 {
|
||||
tempCron := strings.Trim(strings.TrimSpace(assocMatch[1]), "\"' \t")
|
||||
if isLikelyCron(tempCron) {
|
||||
taskCron = tempCron
|
||||
}
|
||||
}
|
||||
|
||||
// B. 如果没找到,尝试普通的 cron: "..." 或 cron 表达式
|
||||
if taskCron == "" {
|
||||
if cronMatch := cronRegex.FindStringSubmatch(line); len(cronMatch) > 0 {
|
||||
for _, m := range cronMatch[1:] {
|
||||
if m != "" {
|
||||
tempCron := strings.Trim(strings.TrimSpace(m), "\"' \t")
|
||||
if isLikelyCron(tempCron) {
|
||||
taskCron = tempCron
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if taskName != "" && taskCron != "" {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// 如果最后还是没找到 taskName,使用备选名称或文件名
|
||||
if taskName == "" {
|
||||
if firstCommentLine != "" {
|
||||
taskName = firstCommentLine
|
||||
} else {
|
||||
taskName = strings.TrimSuffix(filepath.Base(path), ext)
|
||||
}
|
||||
}
|
||||
|
||||
return taskName, taskCron
|
||||
}
|
||||
|
||||
// isLikelyCron 校验字符串是否符合 Cron 表达式的格式特征且数值合法
|
||||
func isLikelyCron(s string) bool {
|
||||
if !cronFormatRegex.MatchString(s) {
|
||||
return false
|
||||
}
|
||||
|
||||
// 语义校验:确保数值在合法范围内(解决如 JS 数组数字超出 Cron 范围的问题)
|
||||
fields := strings.Fields(s)
|
||||
testCron := s
|
||||
if len(fields) == 5 {
|
||||
testCron = "0 " + s
|
||||
}
|
||||
|
||||
// 使用与面板执行器一致的 Parser 进行校验
|
||||
parser := cron.NewParser(cron.Second | cron.Minute | cron.Hour | cron.Dom | cron.Month | cron.Dow | cron.Descriptor)
|
||||
_, err := parser.Parse(testCron)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// splitKeywords 按竖线或逗号分割字符串
|
||||
func splitKeywords(s string) []string {
|
||||
if s == "" {
|
||||
return nil
|
||||
}
|
||||
var parts []string
|
||||
if strings.Contains(s, "|") {
|
||||
parts = strings.Split(s, "|")
|
||||
} else if strings.Contains(s, ",") {
|
||||
parts = strings.Split(s, ",")
|
||||
} else {
|
||||
parts = []string{s}
|
||||
}
|
||||
|
||||
var res []string
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p != "" {
|
||||
res = append(res, p)
|
||||
}
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
// getValidExtensions 返回支持的文件扩展名列表,优先考虑自定义扩展名
|
||||
func getValidExtensions(customExtensions string) []string {
|
||||
exts := []string{".js", ".py", ".ts", ".sh"}
|
||||
if customExtensions != "" {
|
||||
customExts := splitKeywords(customExtensions)
|
||||
if len(customExts) > 0 {
|
||||
exts = nil
|
||||
for _, e := range customExts {
|
||||
e = strings.TrimSpace(e)
|
||||
if e != "" {
|
||||
if !strings.HasPrefix(e, ".") {
|
||||
e = "." + e
|
||||
}
|
||||
exts = append(exts, e)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return exts
|
||||
}
|
||||
|
||||
// sanitizeIdentifier 将非字母数字字符替换为下划线
|
||||
func sanitizeIdentifier(s string) string {
|
||||
reg := regexp.MustCompile(`[^a-zA-Z0-9]+`)
|
||||
res := reg.ReplaceAllString(s, "_")
|
||||
return strings.ToLower(strings.Trim(res, "_"))
|
||||
}
|
||||
|
||||
// normalizeCron 确保 cron 表达式具有 6 个字段
|
||||
func normalizeCron(cron string) string {
|
||||
fields := strings.Fields(cron)
|
||||
if len(fields) == 5 {
|
||||
return "0 " + cron
|
||||
}
|
||||
return cron
|
||||
}
|
||||
|
||||
// pathExists 检查路径是否存在
|
||||
func pathExists(path string) bool {
|
||||
_, err := os.Stat(path)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// isDir 检查路径是否为目录
|
||||
func isDir(path string) bool {
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return info.IsDir()
|
||||
}
|
||||
|
||||
// getCommandByExt 根据文件扩展名返回默认执行命令
|
||||
func getCommandByExt(ext, path string) string {
|
||||
quotedPath := utils.QuotePath(path)
|
||||
switch ext {
|
||||
case ".js", ".ts":
|
||||
return fmt.Sprintf("node %s", quotedPath)
|
||||
case ".py":
|
||||
return fmt.Sprintf("python %s", quotedPath)
|
||||
case ".sh":
|
||||
return fmt.Sprintf("bash %s", quotedPath)
|
||||
case ".php":
|
||||
return fmt.Sprintf("php %s", quotedPath)
|
||||
case ".cs":
|
||||
return fmt.Sprintf("dotnet run %s", quotedPath)
|
||||
}
|
||||
return quotedPath
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type ScriptService struct{}
|
||||
|
||||
func NewScriptService() *ScriptService {
|
||||
return &ScriptService{}
|
||||
}
|
||||
|
||||
func (ss *ScriptService) CreateScript(name, content string, userID string) *models.Script {
|
||||
script := &models.Script{
|
||||
ID: utils.GenerateID(),
|
||||
Name: name,
|
||||
Content: models.BigText(content),
|
||||
UserID: userID,
|
||||
}
|
||||
database.DB.Create(script)
|
||||
return script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) GetScriptsByUserID(userID string) []models.Script {
|
||||
var scripts []models.Script
|
||||
database.DB.Where("user_id = ?", userID).Find(&scripts)
|
||||
return scripts
|
||||
}
|
||||
|
||||
func (ss *ScriptService) GetScriptByID(id string) *models.Script {
|
||||
var script models.Script
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&script)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
return &script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) UpdateScript(id string, name, content string) *models.Script {
|
||||
var script models.Script
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&script)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
script.Name = name
|
||||
script.Content = models.BigText(content)
|
||||
database.DB.Save(&script)
|
||||
return &script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) DeleteScript(id string) bool {
|
||||
result := database.DB.Where("id = ?", id).Delete(&models.Script{})
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/systime"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type SendStatsService struct{}
|
||||
|
||||
func NewSendStatsService() *SendStatsService {
|
||||
return &SendStatsService{}
|
||||
}
|
||||
|
||||
// IncrementStats 增加任务执行统计
|
||||
func (s *SendStatsService) IncrementStats(taskID string, status string) error {
|
||||
day := systime.FormatDate(time.Now())
|
||||
|
||||
var stats models.SendStats
|
||||
res := database.DB.Where("task_id = ? AND day = ? AND status = ?", taskID, day, status).Limit(1).Find(&stats)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
// 不存在则创建
|
||||
stats = models.SendStats{
|
||||
ID: utils.GenerateID(),
|
||||
TaskID: taskID,
|
||||
Day: day,
|
||||
Status: status,
|
||||
Num: 1,
|
||||
}
|
||||
return database.DB.Create(&stats).Error
|
||||
}
|
||||
|
||||
// 存在则增加计数
|
||||
return database.DB.Model(&stats).Update("num", stats.Num+1).Error
|
||||
}
|
||||
|
||||
// GetStatsByTaskID 获取任务的统计数据
|
||||
func (s *SendStatsService) GetStatsByTaskID(taskID string) []models.SendStats {
|
||||
var stats []models.SendStats
|
||||
database.DB.Where("task_id = ?", taskID).Order("day DESC").Find(&stats)
|
||||
return stats
|
||||
}
|
||||
|
||||
// GetTodayStats 获取今日统计
|
||||
func (s *SendStatsService) GetTodayStats() []models.SendStats {
|
||||
day := systime.FormatDate(time.Now())
|
||||
var stats []models.SendStats
|
||||
database.DB.Where("day = ?", day).Find(&stats)
|
||||
return stats
|
||||
}
|
||||
|
||||
// GetRecentStats 获取最近N天的统计
|
||||
func (s *SendStatsService) GetRecentStats(days int) []models.SendStats {
|
||||
startDay := systime.FormatDate(time.Now().AddDate(0, 0, -days))
|
||||
var stats []models.SendStats
|
||||
database.DB.Where("day >= ?", startDay).Order("day DESC").Find(&stats)
|
||||
return stats
|
||||
}
|
||||
@@ -0,0 +1,204 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/engigu/taskpool/internal/cache"
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type SettingsService struct{}
|
||||
|
||||
func NewSettingsService() *SettingsService {
|
||||
return &SettingsService{}
|
||||
}
|
||||
|
||||
// InitSettings 初始化默认设置
|
||||
func (s *SettingsService) InitSettings() error {
|
||||
for section, keys := range constant.DefaultSettings {
|
||||
for key, value := range keys {
|
||||
var count int64
|
||||
database.DB.Model(&models.Setting{}).Where(&models.Setting{Section: section, Key: key}).Count(&count)
|
||||
if count == 0 {
|
||||
if err := database.DB.Create(&models.Setting{
|
||||
ID: utils.GenerateID(),
|
||||
Section: section,
|
||||
Key: key,
|
||||
Value: models.BigText(value),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 初始化日志清理配置
|
||||
// 检查是否需要从旧的 JSON 迁移
|
||||
oldVal := s.Get(constant.SectionSystem, "log_retention")
|
||||
if oldVal != "" {
|
||||
var oldConfigs map[string]struct {
|
||||
Days int `json:"days"`
|
||||
MaxCount int `json:"max_count"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(oldVal), &oldConfigs); err == nil {
|
||||
migrationMap := map[string]string{}
|
||||
if cfg, ok := oldConfigs[constant.LogCategorySystemNotice]; ok {
|
||||
migrationMap[constant.KeySystemNoticeDays] = fmt.Sprintf("%d", cfg.Days)
|
||||
migrationMap[constant.KeySystemNoticeMaxCount] = fmt.Sprintf("%d", cfg.MaxCount)
|
||||
}
|
||||
if cfg, ok := oldConfigs[constant.LogCategoryPushLog]; ok {
|
||||
migrationMap[constant.KeyPushLogDays] = fmt.Sprintf("%d", cfg.Days)
|
||||
migrationMap[constant.KeyPushLogMaxCount] = fmt.Sprintf("%d", cfg.MaxCount)
|
||||
}
|
||||
if cfg, ok := oldConfigs[constant.LogCategoryLoginLog]; ok {
|
||||
migrationMap[constant.KeyLoginLogDays] = fmt.Sprintf("%d", cfg.Days)
|
||||
migrationMap[constant.KeyLoginLogMaxCount] = fmt.Sprintf("%d", cfg.MaxCount)
|
||||
}
|
||||
|
||||
if len(migrationMap) > 0 {
|
||||
for k, v := range migrationMap {
|
||||
s.Set(constant.SectionSystem, k, v)
|
||||
}
|
||||
// 迁移完成后删除旧键
|
||||
s.Delete(constant.SectionSystem, "log_retention")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 默认值初始化
|
||||
defaultRetention := map[string]string{
|
||||
constant.KeySystemNoticeDays: "30",
|
||||
constant.KeySystemNoticeMaxCount: "500",
|
||||
constant.KeyPushLogDays: "15",
|
||||
constant.KeyPushLogMaxCount: "5000",
|
||||
constant.KeyLoginLogDays: "30",
|
||||
constant.KeyLoginLogMaxCount: "1000",
|
||||
constant.KeySchedulerLogDays: "30",
|
||||
constant.KeySchedulerLogMaxCount: "10000",
|
||||
}
|
||||
|
||||
for k, v := range defaultRetention {
|
||||
var count int64
|
||||
database.DB.Model(&models.Setting{}).Where(&models.Setting{Section: constant.SectionSystem, Key: k}).Count(&count)
|
||||
if count == 0 {
|
||||
s.Set(constant.SectionSystem, k, v)
|
||||
}
|
||||
}
|
||||
|
||||
// 从 constant.DefaultSettings 初始化所有缺少的通知模板
|
||||
if notifyDefaults, ok := constant.DefaultSettings[constant.SectionNotify]; ok {
|
||||
for k, v := range notifyDefaults {
|
||||
var count int64
|
||||
database.DB.Model(&models.Setting{}).Where(&models.Setting{Section: constant.SectionNotify, Key: k}).Count(&count)
|
||||
if count == 0 {
|
||||
s.Set(constant.SectionNotify, k, v)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 初始化或获取 JWT Secret 密码
|
||||
var secCount int64
|
||||
database.DB.Model(&models.Setting{}).Where(&models.Setting{Section: constant.SectionSecurity, Key: constant.KeySecret}).Count(&secCount)
|
||||
var secretValue string
|
||||
if secCount == 0 {
|
||||
// 先尝试从配置文件读取遗留下来的旧设
|
||||
if Config != nil && Config.Security.Secret != "" {
|
||||
secretValue = Config.Security.Secret
|
||||
} else {
|
||||
secretValue = utils.RandomString(32)
|
||||
}
|
||||
if err := database.DB.Create(&models.Setting{
|
||||
ID: utils.GenerateID(),
|
||||
Section: constant.SectionSecurity,
|
||||
Key: constant.KeySecret,
|
||||
Value: models.BigText(secretValue),
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
secretValue = s.Get(constant.SectionSecurity, constant.KeySecret)
|
||||
}
|
||||
constant.Secret = secretValue
|
||||
|
||||
cache.LoadSiteCache()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get 获取单个设置
|
||||
func (s *SettingsService) Get(section, key string) string {
|
||||
if section == constant.SectionSite {
|
||||
return cache.GetSiteCache(key)
|
||||
}
|
||||
var setting models.Setting
|
||||
res := database.DB.Where(&models.Setting{Section: section, Key: key}).Limit(1).Find(&setting)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
if def, ok := constant.DefaultSettings[section][key]; ok {
|
||||
return def
|
||||
}
|
||||
return ""
|
||||
}
|
||||
return string(setting.Value)
|
||||
}
|
||||
|
||||
// Set 设置单个值
|
||||
func (s *SettingsService) Set(section, key, value string) error {
|
||||
var setting models.Setting
|
||||
res := database.DB.Where(&models.Setting{Section: section, Key: key}).Limit(1).Find(&setting)
|
||||
var err error
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
err = database.DB.Create(&models.Setting{
|
||||
ID: utils.GenerateID(),
|
||||
Section: section,
|
||||
Key: key,
|
||||
Value: models.BigText(value),
|
||||
}).Error
|
||||
} else {
|
||||
err = database.DB.Model(&setting).Update("value", models.BigText(value)).Error
|
||||
}
|
||||
|
||||
if err == nil && section == constant.SectionSite {
|
||||
cache.SetSiteCache(key, value)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Delete 删除单个设置
|
||||
func (s *SettingsService) Delete(section, key string) error {
|
||||
return database.DB.Where(&models.Setting{Section: section, Key: key}).Delete(&models.Setting{}).Error
|
||||
}
|
||||
|
||||
// GetSection 获取整个 section 的设置
|
||||
func (s *SettingsService) GetSection(section string) map[string]string {
|
||||
if section == constant.SectionSite {
|
||||
return cache.GetSiteCacheAll()
|
||||
}
|
||||
result := make(map[string]string)
|
||||
if defaults, ok := constant.DefaultSettings[section]; ok {
|
||||
for k, v := range defaults {
|
||||
result[k] = v
|
||||
}
|
||||
}
|
||||
var settings []models.Setting
|
||||
database.DB.Where("section = ?", section).Find(&settings)
|
||||
for _, setting := range settings {
|
||||
result[setting.Key] = string(setting.Value)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// SetSection 批量设置
|
||||
func (s *SettingsService) SetSection(section string, values map[string]string) error {
|
||||
for key, value := range values {
|
||||
if err := s.Set(section, key, value); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if section == constant.SectionSite {
|
||||
cache.SetSiteCacheBatch(values)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// SystemWSManager 前端系统事件 WebSocket 管理器 (单例)
|
||||
type SystemWSManager struct {
|
||||
clients map[*ClientConnection]bool
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// ClientConnection 代表一个前端页面的 WebSocket 连接
|
||||
type ClientConnection struct {
|
||||
Conn *websocket.Conn
|
||||
Send chan []byte
|
||||
closed bool
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
var systemWSManager *SystemWSManager
|
||||
var systemWSOnce sync.Once
|
||||
|
||||
// GetSystemWSManager 获取系统 WebSocket 管理器单例
|
||||
func GetSystemWSManager() *SystemWSManager {
|
||||
systemWSOnce.Do(func() {
|
||||
systemWSManager = &SystemWSManager{
|
||||
clients: make(map[*ClientConnection]bool),
|
||||
}
|
||||
})
|
||||
return systemWSManager
|
||||
}
|
||||
|
||||
// Register 注册一个新的前端连接
|
||||
func (m *SystemWSManager) Register(conn *websocket.Conn) *ClientConnection {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
client := &ClientConnection{
|
||||
Conn: conn,
|
||||
Send: make(chan []byte, 256),
|
||||
}
|
||||
m.clients[client] = true
|
||||
return client
|
||||
}
|
||||
|
||||
// Unregister 注销一个前端连接
|
||||
func (m *SystemWSManager) Unregister(client *ClientConnection) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if _, ok := m.clients[client]; ok {
|
||||
delete(m.clients, client)
|
||||
client.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// Broadcast 广播消息给所有在线前端
|
||||
func (m *SystemWSManager) Broadcast(msgType string, payload interface{}) {
|
||||
msg := vo.WSMessage{
|
||||
Type: msgType,
|
||||
Timestamp: time.Now().UnixMilli(),
|
||||
Payload: payload,
|
||||
}
|
||||
|
||||
data, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
logger.Errorf("[SystemWS] 序列化消息失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
for client := range m.clients {
|
||||
select {
|
||||
case client.Send <- data:
|
||||
default:
|
||||
// 缓冲区满,可能该客户端连接已死
|
||||
go m.Unregister(client)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// SubscribeEvents 订阅系统事件总线并分发给 WebSocket
|
||||
func (m *SystemWSManager) SubscribeEvents(bus *eventbus.EventBus) {
|
||||
// 任务相关事件
|
||||
taskEvents := []string{
|
||||
constant.EventTaskSuccess,
|
||||
constant.EventTaskFailed,
|
||||
constant.EventTaskTimeout,
|
||||
constant.EventTaskRunning,
|
||||
constant.EventTaskQueued,
|
||||
constant.EventTaskCancelled,
|
||||
}
|
||||
|
||||
for _, evt := range taskEvents {
|
||||
bus.Subscribe(evt, func(e eventbus.Event) {
|
||||
m.Broadcast(e.Type, e.Payload)
|
||||
})
|
||||
}
|
||||
|
||||
// 系统通知事件
|
||||
bus.Subscribe(constant.EventSystemNotice, func(e eventbus.Event) {
|
||||
m.Broadcast("notice", e.Payload)
|
||||
})
|
||||
|
||||
// 应用日志新增事件(驱动运行日志下属4大标签页实时流式刷新列表)
|
||||
bus.Subscribe(constant.EventAppLogAdded, func(e eventbus.Event) {
|
||||
m.Broadcast(e.Type, e.Payload)
|
||||
})
|
||||
}
|
||||
|
||||
func (c *ClientConnection) Close() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed {
|
||||
return
|
||||
}
|
||||
c.closed = true
|
||||
c.Conn.Close()
|
||||
close(c.Send)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,251 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/systime"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// SendStatsService 接口定义(避免循环依赖)
|
||||
type SendStatsService interface {
|
||||
IncrementStats(taskID string, status string) error
|
||||
}
|
||||
|
||||
// TaskLogService 任务日志服务
|
||||
type TaskLogService struct {
|
||||
sendStatsService SendStatsService
|
||||
}
|
||||
|
||||
// NewTaskLogService 创建任务日志服务
|
||||
func NewTaskLogService(sendStatsService SendStatsService) *TaskLogService {
|
||||
return &TaskLogService{
|
||||
sendStatsService: sendStatsService,
|
||||
}
|
||||
}
|
||||
|
||||
// CleanConfig 清理配置
|
||||
type CleanConfig struct {
|
||||
Type string `json:"type"` // day 或 count
|
||||
Keep int `json:"keep"` // 保留天数或条数
|
||||
}
|
||||
|
||||
// CreateEmptyLog 创建一个空的日志记录(任务开始时调用)
|
||||
func (s *TaskLogService) CreateEmptyLog(taskID string, command string) (*models.TaskLog, error) {
|
||||
startTime := models.Now()
|
||||
taskLog := &models.TaskLog{
|
||||
ID: utils.GenerateID(),
|
||||
TaskID: taskID,
|
||||
Command: models.BigText(command),
|
||||
Status: "running",
|
||||
StartTime: &startTime,
|
||||
CreatedAt: models.Now(),
|
||||
}
|
||||
if err := database.DB.Create(taskLog).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 任务开始时即更新任务的 last_run 为启动时间
|
||||
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("last_run", startTime)
|
||||
|
||||
return taskLog, nil
|
||||
}
|
||||
|
||||
// SaveTaskLog 保存或更新任务日志
|
||||
func (s *TaskLogService) SaveTaskLog(taskLog *models.TaskLog) error {
|
||||
var err error
|
||||
if taskLog.ID != "" {
|
||||
// 先检查记录是否存在,如果不存在则创建,存在则更新
|
||||
var count int64
|
||||
database.DB.Model(&models.TaskLog{}).Where("id = ?", taskLog.ID).Count(&count)
|
||||
if count > 0 {
|
||||
err = database.DB.Model(taskLog).Where("id = ?", taskLog.ID).Updates(taskLog).Error
|
||||
} else {
|
||||
err = database.DB.Create(taskLog).Error
|
||||
}
|
||||
} else {
|
||||
taskLog.ID = utils.GenerateID()
|
||||
if taskLog.CreatedAt.Time().IsZero() {
|
||||
taskLog.CreatedAt = models.Now()
|
||||
}
|
||||
err = database.DB.Create(taskLog).Error
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 更新任务的 last_run
|
||||
// 更新任务的 last_run,优先使用日志记录的启动时间
|
||||
lastRun := models.Now()
|
||||
if taskLog.StartTime != nil {
|
||||
lastRun = *taskLog.StartTime
|
||||
}
|
||||
database.DB.Model(&models.Task{}).Where("id = ?", taskLog.TaskID).Update("last_run", lastRun)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateTaskDuration 更新任务耗时(心跳)
|
||||
func (s *TaskLogService) UpdateTaskDuration(logID string, duration int64) error {
|
||||
return database.DB.Model(&models.TaskLog{}).Where("id = ?", logID).Update("duration", duration).Error
|
||||
}
|
||||
|
||||
// UpdateLogCommand 更新日志中的命令内容(用于动态生成的命令脱敏)
|
||||
func (s *TaskLogService) UpdateLogCommand(logID string, command string) error {
|
||||
return database.DB.Model(&models.TaskLog{}).Where("id = ?", logID).Update("command", models.BigText(command)).Error
|
||||
}
|
||||
|
||||
// UpdateTaskStats 更新任务统计
|
||||
func (s *TaskLogService) UpdateTaskStats(taskID string, status string) {
|
||||
if s.sendStatsService == nil {
|
||||
logger.Error("[TaskLog] SendStatsService 未初始化")
|
||||
return
|
||||
}
|
||||
err := s.sendStatsService.IncrementStats(taskID, status)
|
||||
if err != nil {
|
||||
logger.Errorf("UpdateTaskStats err: %v", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// CleanTaskLogs 清理任务日志
|
||||
func (s *TaskLogService) CleanTaskLogs(taskID string) {
|
||||
var task models.Task
|
||||
res := database.DB.Where("id = ?", taskID).Limit(1).Find(&task)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
if task.CleanConfig == "" {
|
||||
return
|
||||
}
|
||||
|
||||
var config CleanConfig
|
||||
if err := json.Unmarshal([]byte(task.CleanConfig), &config); err != nil {
|
||||
logger.Errorf("[TaskLog] 解析清理配置失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if config.Keep <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
var deleted int64
|
||||
switch config.Type {
|
||||
case "day":
|
||||
cutoff := systime.InCST(time.Now()).AddDate(0, 0, -config.Keep)
|
||||
result := database.DB.Where("task_id = ? AND created_at < ?", taskID, cutoff).Delete(&models.TaskLog{})
|
||||
deleted = result.RowsAffected
|
||||
case "count":
|
||||
var boundaryLog models.TaskLog
|
||||
res := database.DB.Where("task_id = ?", taskID).Order("id DESC").Offset(config.Keep - 1).Limit(1).Find(&boundaryLog)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
result := database.DB.Where("task_id = ? AND id < ?", taskID, boundaryLog.ID).Delete(&models.TaskLog{})
|
||||
deleted = result.RowsAffected
|
||||
}
|
||||
}
|
||||
|
||||
if deleted > 0 {
|
||||
logger.Infof("[TaskLog] 清理旧日志: #%s 共 %d 条", taskID, deleted)
|
||||
}
|
||||
}
|
||||
|
||||
// ProcessTaskCompletion 处理任务完成后的所有操作(保存日志、更新统计、清理旧日志)
|
||||
func (s *TaskLogService) ProcessTaskCompletion(taskLog *models.TaskLog) error {
|
||||
// 1. 保存/更新日志
|
||||
if err := s.SaveTaskLog(taskLog); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 2. 更新统计
|
||||
s.UpdateTaskStats(taskLog.TaskID, taskLog.Status)
|
||||
|
||||
// 3. 异步清理旧日志
|
||||
go s.CleanTaskLogs(taskLog.TaskID)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// CreateTaskLogFromAgentResult 从 Agent 结果创建任务日志
|
||||
func (s *TaskLogService) CreateTaskLogFromAgentResult(result *models.AgentTaskResult) (*models.TaskLog, error) {
|
||||
// 裁剪并压缩输出
|
||||
trimmedOutput := utils.TrimLog(result.Output, constant.MaxLogSize)
|
||||
compressed, err := utils.CompressToBase64(trimmedOutput)
|
||||
if err != nil {
|
||||
logger.Errorf("[TaskLog] 压缩日志失败: %v", err)
|
||||
compressed = ""
|
||||
}
|
||||
|
||||
logID := result.LogID
|
||||
if logID == "" {
|
||||
logID = utils.GenerateID()
|
||||
}
|
||||
|
||||
taskLog := &models.TaskLog{
|
||||
ID: logID,
|
||||
TaskID: result.TaskID,
|
||||
AgentID: &result.AgentID,
|
||||
Command: models.BigText(result.Command),
|
||||
Output: models.BigText(compressed),
|
||||
Error: models.BigText(result.Error),
|
||||
Status: result.Status,
|
||||
Duration: result.Duration,
|
||||
ExitCode: result.ExitCode,
|
||||
CreatedAt: models.Now(),
|
||||
}
|
||||
|
||||
// 处理开始和结束时间
|
||||
if result.StartTime > 0 {
|
||||
startTime := models.LocalTime(time.Unix(result.StartTime, 0))
|
||||
taskLog.StartTime = &startTime
|
||||
}
|
||||
if result.EndTime > 0 {
|
||||
endTime := models.LocalTime(time.Unix(result.EndTime, 0))
|
||||
taskLog.EndTime = &endTime
|
||||
}
|
||||
|
||||
return taskLog, nil
|
||||
}
|
||||
|
||||
// CreateTaskLogFromLocalExecution 从本地执行结果创建任务日志
|
||||
func (s *TaskLogService) CreateTaskLogFromLocalExecution(taskID string, command, output, systemErr, status string, duration int64, exitCode int, start, end time.Time, isCompressed bool) (*models.TaskLog, error) {
|
||||
var compressed string
|
||||
var err error
|
||||
|
||||
if isCompressed {
|
||||
compressed = output
|
||||
} else {
|
||||
// 裁剪并压缩输出
|
||||
trimmedOutput := utils.TrimLog(output, constant.MaxLogSize)
|
||||
compressed, err = utils.CompressToBase64(trimmedOutput)
|
||||
if err != nil {
|
||||
logger.Errorf("[TaskLog] 压缩日志失败: %v", err)
|
||||
compressed = ""
|
||||
}
|
||||
}
|
||||
|
||||
startTime := models.LocalTime(start)
|
||||
endTime := models.LocalTime(end)
|
||||
|
||||
taskLog := &models.TaskLog{
|
||||
ID: utils.GenerateID(),
|
||||
TaskID: taskID,
|
||||
Command: models.BigText(command),
|
||||
Output: models.BigText(compressed),
|
||||
Error: models.BigText(systemErr),
|
||||
Status: status,
|
||||
Duration: duration,
|
||||
ExitCode: exitCode,
|
||||
StartTime: &startTime,
|
||||
EndTime: &endTime,
|
||||
CreatedAt: models.Now(),
|
||||
}
|
||||
|
||||
return taskLog, nil
|
||||
}
|
||||
@@ -0,0 +1,287 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/relation"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// TaskParam 任务创建与更新参数传输对象
|
||||
type TaskParam struct {
|
||||
Name string
|
||||
Remark string
|
||||
Command string
|
||||
PreCommand string
|
||||
PostCommand string
|
||||
Tags string
|
||||
Type string
|
||||
Config string
|
||||
Schedule string
|
||||
Timeout int
|
||||
WorkDir string
|
||||
CleanConfig string
|
||||
Envs string
|
||||
Languages models.TaskLanguages
|
||||
AgentID *string
|
||||
TriggerType string
|
||||
RetryCount int
|
||||
RetryInterval int
|
||||
RandomRange int
|
||||
SourceID string
|
||||
PinType string
|
||||
Enabled bool
|
||||
}
|
||||
|
||||
type TaskService struct {
|
||||
}
|
||||
|
||||
func NewTaskService() *TaskService {
|
||||
return &TaskService{}
|
||||
}
|
||||
|
||||
func (ts *TaskService) GetTaskBySourceID(sourceID string) *models.Task {
|
||||
var task models.Task
|
||||
res := database.DB.Where("source_id = ?", sourceID).Limit(1).Find(&task)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
ts.loadTagsAndEnvs([]models.Task{task})
|
||||
|
||||
return &task
|
||||
}
|
||||
|
||||
func (ts *TaskService) CreateTask(p *TaskParam) *models.Task {
|
||||
if p.Type == "" {
|
||||
p.Type = "task"
|
||||
}
|
||||
if p.TriggerType == "" {
|
||||
p.TriggerType = constant.TriggerTypeCron
|
||||
}
|
||||
if p.PinType == "" {
|
||||
p.PinType = constant.PinTypeNone
|
||||
}
|
||||
task := &models.Task{
|
||||
ID: utils.GenerateID(),
|
||||
Name: p.Name,
|
||||
Remark: p.Remark,
|
||||
Command: models.BigText(p.Command),
|
||||
PreCommand: models.BigText(p.PreCommand),
|
||||
PostCommand: models.BigText(p.PostCommand),
|
||||
PinType: p.PinType,
|
||||
Tags: p.Tags,
|
||||
Type: p.Type,
|
||||
TriggerType: p.TriggerType,
|
||||
Config: models.BigText(p.Config),
|
||||
Schedule: p.Schedule,
|
||||
Timeout: p.Timeout,
|
||||
WorkDir: p.WorkDir,
|
||||
CleanConfig: p.CleanConfig,
|
||||
Envs: models.BigText(p.Envs),
|
||||
Languages: p.Languages,
|
||||
AgentID: p.AgentID,
|
||||
Enabled: utils.BoolPtr(true),
|
||||
RetryCount: p.RetryCount,
|
||||
RetryInterval: p.RetryInterval,
|
||||
RandomRange: p.RandomRange,
|
||||
SourceID: p.SourceID,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
if p.TriggerType != constant.TriggerTypeCron {
|
||||
task.NextRun = nil
|
||||
}
|
||||
database.DB.Select("*").Create(task)
|
||||
relation.DataRelation.SaveTags(task.ID, constant.RelationTypeTaskTag, p.Tags)
|
||||
task.Tags = p.Tags
|
||||
relation.DataRelation.SaveRelations(task.ID, constant.RelationTypeTaskEnv, p.Envs)
|
||||
task.Envs = models.BigText(p.Envs)
|
||||
|
||||
return task
|
||||
}
|
||||
|
||||
func (ts *TaskService) GetTasks() []models.Task {
|
||||
var tasks []models.Task
|
||||
database.DB.Find(&tasks)
|
||||
ts.loadTagsAndEnvs(tasks)
|
||||
return tasks
|
||||
}
|
||||
|
||||
// GetTasksWithPagination 分页获取任务列表
|
||||
func (ts *TaskService) GetTasksWithPagination(page, pageSize int, name string, agentID *string, tags string, taskType string, sortBy string, order string) ([]models.Task, int64) {
|
||||
var tasks []models.Task
|
||||
var total int64
|
||||
|
||||
query := database.DB.Model(&models.Task{})
|
||||
if name != "" {
|
||||
query = query.Where("name LIKE ? OR remark LIKE ?", "%"+name+"%", "%"+name+"%")
|
||||
}
|
||||
|
||||
// 标签筛选 (交集或并集均可,这里保留原本的逻辑为并集,但是利用数据关联表)
|
||||
if tags != "" {
|
||||
tagList := strings.Split(tags, ",")
|
||||
var validTags []string
|
||||
for _, tag := range tagList {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag != "" {
|
||||
validTags = append(validTags, tag)
|
||||
}
|
||||
}
|
||||
if len(validTags) > 0 {
|
||||
var storageIDs []string
|
||||
database.DB.Model(&models.DataStorage{}).Where("type = ? AND name IN ?", constant.RelationTypeTaskTag, validTags).Pluck("id", &storageIDs)
|
||||
|
||||
var taskIDs []string
|
||||
if len(storageIDs) > 0 {
|
||||
database.DB.Model(&models.DataRelation{}).Where("type = ? AND relate_id IN ?", constant.RelationTypeTaskTag, storageIDs).Pluck("data_id", &taskIDs)
|
||||
}
|
||||
|
||||
if len(taskIDs) > 0 {
|
||||
query = query.Where("id IN ?", taskIDs)
|
||||
} else {
|
||||
query = query.Where("1 = 0")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if taskType != "" && taskType != "all" {
|
||||
query = query.Where("type = ?", taskType)
|
||||
}
|
||||
if agentID != nil {
|
||||
query = query.Where("agent_id = ?", *agentID)
|
||||
}
|
||||
|
||||
sortColumn := "created_at"
|
||||
if sortBy != "" {
|
||||
switch sortBy {
|
||||
case "name", "next_run", "last_run", "created_at", "enabled":
|
||||
sortColumn = sortBy
|
||||
}
|
||||
}
|
||||
sortOrder := "DESC"
|
||||
if strings.ToUpper(order) == "ASC" {
|
||||
sortOrder = "ASC"
|
||||
}
|
||||
|
||||
query.Count(&total)
|
||||
query.Order("pin_type DESC, " + sortColumn + " " + sortOrder).Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks)
|
||||
ts.loadTagsAndEnvs(tasks)
|
||||
|
||||
return tasks, total
|
||||
}
|
||||
|
||||
func (ts *TaskService) GetTaskByID(id string) *models.Task {
|
||||
var task models.Task
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&task)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
tasks := []models.Task{task}
|
||||
ts.loadTagsAndEnvs(tasks)
|
||||
|
||||
return &tasks[0]
|
||||
}
|
||||
|
||||
func (ts *TaskService) UpdateTask(id string, p *TaskParam) *models.Task {
|
||||
var task models.Task
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&task)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
task.Name = p.Name
|
||||
task.Remark = p.Remark
|
||||
task.Command = models.BigText(p.Command)
|
||||
task.PreCommand = models.BigText(p.PreCommand)
|
||||
task.PostCommand = models.BigText(p.PostCommand)
|
||||
task.PinType = p.PinType
|
||||
task.Schedule = p.Schedule
|
||||
task.Timeout = p.Timeout
|
||||
task.WorkDir = p.WorkDir
|
||||
task.CleanConfig = p.CleanConfig
|
||||
task.Enabled = &p.Enabled
|
||||
task.AgentID = p.AgentID
|
||||
task.Languages = p.Languages
|
||||
task.Config = models.BigText(p.Config)
|
||||
task.RetryCount = p.RetryCount
|
||||
task.RetryInterval = p.RetryInterval
|
||||
task.RandomRange = p.RandomRange
|
||||
if p.Type != "" {
|
||||
task.Type = p.Type
|
||||
}
|
||||
if p.TriggerType != "" {
|
||||
task.TriggerType = p.TriggerType
|
||||
}
|
||||
if p.SourceID != "" {
|
||||
task.SourceID = p.SourceID
|
||||
}
|
||||
|
||||
database.DB.Model(&task).Select(
|
||||
"Name", "Remark", "Command", "Tags", "Schedule", "Timeout", "WorkDir",
|
||||
"CleanConfig", "Enabled", "AgentID", "Languages",
|
||||
"RetryCount", "RetryInterval", "RandomRange", "Type",
|
||||
"TriggerType", "Config", "SourceID", "PinType",
|
||||
"PreCommand", "PostCommand",
|
||||
).Updates(&task)
|
||||
|
||||
relation.DataRelation.SaveTags(task.ID, constant.RelationTypeTaskTag, p.Tags)
|
||||
task.Tags = p.Tags
|
||||
relation.DataRelation.SaveRelations(task.ID, constant.RelationTypeTaskEnv, p.Envs)
|
||||
task.Envs = models.BigText(p.Envs)
|
||||
|
||||
return &task
|
||||
}
|
||||
|
||||
func (ts *TaskService) DeleteTask(id string) bool {
|
||||
// 同时删除关联的通知推送设置
|
||||
database.DB.Where("type = ? AND data_id = ?", constant.BindingTypeTask, id).Delete(&models.NotifyBinding{})
|
||||
relation.DataRelation.CleanRelations(id, constant.RelationTypeTaskTag)
|
||||
relation.DataRelation.CleanRelations(id, constant.RelationTypeTaskEnv)
|
||||
|
||||
result := database.DB.Where("id = ?", id).Delete(&models.Task{})
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
|
||||
func (ts *TaskService) BatchDeleteTasks(ids []string) int64 {
|
||||
// 同时删除关联的通知推送设置
|
||||
database.DB.Where("type = ? AND data_id IN ?", constant.BindingTypeTask, ids).Delete(&models.NotifyBinding{})
|
||||
database.DB.Where("type = ? AND data_id IN ?", constant.RelationTypeTaskTag, ids).Delete(&models.DataRelation{})
|
||||
database.DB.Where("type = ? AND data_id IN ?", constant.RelationTypeTaskEnv, ids).Delete(&models.DataRelation{})
|
||||
|
||||
result := database.DB.Where("id IN ?", ids).Delete(&models.Task{})
|
||||
return result.RowsAffected
|
||||
}
|
||||
|
||||
// GetAllTags 获取所有任务标签
|
||||
func (ts *TaskService) GetAllTags() ([]string, error) {
|
||||
return relation.DataRelation.GetAllTags(constant.RelationTypeTaskTag)
|
||||
}
|
||||
|
||||
func (ts *TaskService) loadTagsAndEnvs(tasks []models.Task) {
|
||||
if len(tasks) == 0 {
|
||||
return
|
||||
}
|
||||
taskIDs := make([]string, len(tasks))
|
||||
for i, t := range tasks {
|
||||
taskIDs[i] = t.ID
|
||||
}
|
||||
|
||||
tagsMap := relation.DataRelation.LoadTags(taskIDs, constant.RelationTypeTaskTag)
|
||||
envsMap := relation.DataRelation.LoadRelations(taskIDs, constant.RelationTypeTaskEnv)
|
||||
|
||||
for i, t := range tasks {
|
||||
if tags, ok := tagsMap[t.ID]; ok {
|
||||
tasks[i].Tags = strings.Join(tags, ",")
|
||||
} else {
|
||||
tasks[i].Tags = ""
|
||||
}
|
||||
|
||||
if envs, ok := envsMap[t.ID]; ok {
|
||||
tasks[i].Envs = models.BigText(strings.Join(envs, ","))
|
||||
} else {
|
||||
tasks[i].Envs = models.BigText("")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,386 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
const (
|
||||
// maxLogBufferLen 定义了没有换行符时的最大缓冲长度 (4KB)
|
||||
maxLogBufferLen = 4096
|
||||
)
|
||||
|
||||
var (
|
||||
// globalTinyLogManager 跟踪所有活跃的 TinyLog 实例
|
||||
globalTinyLogManager = &TinyLogManager{
|
||||
logs: make(map[string]*TinyLog),
|
||||
}
|
||||
)
|
||||
|
||||
type TinyLogManager struct {
|
||||
mu sync.RWMutex
|
||||
logs map[string]*TinyLog
|
||||
}
|
||||
|
||||
func (m *TinyLogManager) Register(log *TinyLog) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.logs[log.LogID] = log
|
||||
}
|
||||
|
||||
func (m *TinyLogManager) Unregister(logID string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.logs, logID)
|
||||
}
|
||||
|
||||
func (m *TinyLogManager) Get(logID string) *TinyLog {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.logs[logID]
|
||||
}
|
||||
|
||||
// GetActiveLog 通过 ID 获取活跃的 TinyLog 实例
|
||||
func GetActiveLog(logID string) *TinyLog {
|
||||
return globalTinyLogManager.Get(logID)
|
||||
}
|
||||
|
||||
// TinyLog 是一个高性能、低内存占用的日志收集器
|
||||
type TinyLog struct {
|
||||
LogID string
|
||||
mu sync.RWMutex
|
||||
file *os.File
|
||||
path string
|
||||
writer *bufio.Writer
|
||||
subscribers []chan []byte
|
||||
remainder []byte // Leftover bytes from previous write (partial lines)
|
||||
masks []string // Secrets to mask
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewTinyLog 创建一个新的 TinyLog 实例(基于临时文件存储)并注册它,支持将配置的 masks 替换为 ********
|
||||
func NewTinyLog(logID string, masks []string) (*TinyLog, error) {
|
||||
f, err := os.CreateTemp("", "task_log_*.log")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
tl := &TinyLog{
|
||||
LogID: logID,
|
||||
file: f,
|
||||
path: f.Name(),
|
||||
writer: bufio.NewWriter(f),
|
||||
subscribers: make([]chan []byte, 0),
|
||||
masks: masks,
|
||||
}
|
||||
globalTinyLogManager.Register(tl)
|
||||
return tl, nil
|
||||
}
|
||||
|
||||
// Write 实现 io.Writer 接口
|
||||
func (l *TinyLog) Write(p []byte) (n int, err error) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
if l.closed {
|
||||
return 0, os.ErrClosed
|
||||
}
|
||||
|
||||
originalInputLen := len(p)
|
||||
var payload []byte
|
||||
if len(l.remainder) > 0 {
|
||||
// 为了防止 p 和 l.remainder 底层数组有重叠或不可预期的修改,这里分配新内存
|
||||
payload = make([]byte, len(l.remainder)+len(p))
|
||||
copy(payload, l.remainder)
|
||||
copy(payload[len(l.remainder):], p)
|
||||
l.remainder = nil
|
||||
} else {
|
||||
payload = p
|
||||
}
|
||||
|
||||
// 1. 寻找最后一个换行符 (\n 或 \r)
|
||||
lastLineBreak := bytes.LastIndexAny(payload, "\n\r")
|
||||
var completeBytes []byte
|
||||
var remainder []byte
|
||||
|
||||
if lastLineBreak != -1 {
|
||||
// 2. 提取出完整的行
|
||||
completeBytes = payload[:lastLineBreak+1]
|
||||
remainder = payload[lastLineBreak+1:]
|
||||
} else {
|
||||
// 3. 没有换行符,且如果长度超过最大缓冲,强制截断并输出,防止内存无限制增长
|
||||
if len(payload) > maxLogBufferLen {
|
||||
// 寻找最后一个完整的 UTF-8 字符边界,避免乱码
|
||||
lastSafe := maxLogBufferLen
|
||||
for i := maxLogBufferLen; i > 0 && i > maxLogBufferLen-4; i-- {
|
||||
if utf8.RuneStart(payload[i-1]) {
|
||||
if !utf8.FullRune(payload[i-1 : maxLogBufferLen]) {
|
||||
lastSafe = i - 1
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
completeBytes = payload[:lastSafe]
|
||||
remainder = payload[lastSafe:]
|
||||
} else {
|
||||
// 保留当前所有内容到下一轮 (必须 copy,因为 payload 底层可能是 io.Copy 的复用 buf)
|
||||
l.remainder = make([]byte, len(payload))
|
||||
copy(l.remainder, payload)
|
||||
return originalInputLen, nil
|
||||
}
|
||||
}
|
||||
|
||||
// 4. 将剩余部分保存 (必须 copy,防止后续 Read 覆盖底层数组)
|
||||
if len(remainder) > 0 {
|
||||
l.remainder = make([]byte, len(remainder))
|
||||
copy(l.remainder, remainder)
|
||||
} else {
|
||||
l.remainder = nil
|
||||
}
|
||||
|
||||
// 5. 将完整行转换为 UTF-8 并脱敏
|
||||
text := utils.MaskSecrets(utils.ToUTF8(completeBytes), l.masks)
|
||||
outData := []byte(text)
|
||||
|
||||
// 6. 输出安全部分
|
||||
_, err = l.writer.Write(outData)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
// 6. 广播给所有订阅者
|
||||
if len(l.subscribers) > 0 {
|
||||
for _, ch := range l.subscribers {
|
||||
select {
|
||||
case ch <- outData:
|
||||
default:
|
||||
// 如果订阅者处理太慢,丢弃消息以避免阻塞写入
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return originalInputLen, nil
|
||||
}
|
||||
|
||||
// WriteString 方便地写入字符串
|
||||
func (l *TinyLog) WriteString(s string) (n int, err error) {
|
||||
return l.Write([]byte(s))
|
||||
}
|
||||
|
||||
// Subscribe 返回一个实时接收日志块的通道
|
||||
func (l *TinyLog) Subscribe() chan []byte {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
ch := make(chan []byte, 100) // Buffer to handle bursts
|
||||
l.subscribers = append(l.subscribers, ch)
|
||||
return ch
|
||||
}
|
||||
|
||||
// Unsubscribe 移除订阅者
|
||||
func (l *TinyLog) Unsubscribe(ch chan []byte) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
for i, sub := range l.subscribers {
|
||||
if sub == ch {
|
||||
l.subscribers = append(l.subscribers[:i], l.subscribers[i+1:]...)
|
||||
close(ch)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Close 完成写入,关闭文件并注销实例
|
||||
func (l *TinyLog) Close() error {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
if l.closed {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 处理剩余的字节
|
||||
if len(l.remainder) > 0 {
|
||||
text := utils.MaskSecrets(utils.ToUTF8(l.remainder), l.masks)
|
||||
data := []byte(text)
|
||||
_, _ = l.writer.Write(data)
|
||||
|
||||
// 通知订阅者最后一部分内容
|
||||
for _, ch := range l.subscribers {
|
||||
select {
|
||||
case ch <- data:
|
||||
default:
|
||||
}
|
||||
}
|
||||
l.remainder = nil
|
||||
}
|
||||
|
||||
// 将缓冲区刷新到文件
|
||||
if err := l.writer.Flush(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 关闭所有订阅者通道
|
||||
for _, ch := range l.subscribers {
|
||||
close(ch)
|
||||
}
|
||||
l.subscribers = nil
|
||||
|
||||
l.closed = true
|
||||
globalTinyLogManager.Unregister(l.LogID)
|
||||
return l.file.Close()
|
||||
}
|
||||
|
||||
// CompressAndCleanup 读取临时文件,进行压缩处理,返回结果并删除临时文件
|
||||
func (l *TinyLog) CompressAndCleanup() (string, error) {
|
||||
// Ensure closed
|
||||
if !l.closed {
|
||||
l.Close()
|
||||
}
|
||||
|
||||
// 打开临时文件进行读取
|
||||
f, err := os.Open(l.path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() {
|
||||
f.Close()
|
||||
os.Remove(l.path) // Cleanup
|
||||
}()
|
||||
|
||||
// 获取文件大小
|
||||
stat, err := f.Stat()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
size := stat.Size()
|
||||
|
||||
// 如果日志极短,免去压缩和 Base64 编码,直接以 raw: 明文形式返回
|
||||
if size <= int64(utils.MinCompressSize) {
|
||||
content, err := io.ReadAll(f)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "raw:" + string(content), nil
|
||||
}
|
||||
|
||||
// 创建压缩输出缓冲区
|
||||
var buf bytes.Buffer
|
||||
b64Writer := base64.NewEncoder(base64.StdEncoding, &buf)
|
||||
|
||||
// 使用 Pool 优化压缩
|
||||
zw := utils.GetZstdWriter(b64Writer)
|
||||
defer utils.PutZstdWriter(zw)
|
||||
|
||||
maxSize := int64(constant.MaxLogSize)
|
||||
if maxSize < 1024*1024 {
|
||||
maxSize = 1024 * 1024
|
||||
}
|
||||
|
||||
var readStart int64 = 0
|
||||
if size > maxSize {
|
||||
readStart = size - maxSize
|
||||
// 写入一条截断提示
|
||||
truncatedMsg := fmt.Sprintf("\n\n[System] 日志过长,已自动截断,仅保留末尾 %d MB...\n\n", maxSize/1024/1024)
|
||||
if _, err := zw.Write([]byte(truncatedMsg)); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
if readStart > 0 {
|
||||
if _, err := f.Seek(readStart, io.SeekStart); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
// 流处理: 文件 -> Zstd -> Base64 -> 缓冲区
|
||||
if _, err := io.Copy(zw, f); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 关闭写入器以刷新数据
|
||||
if err := zw.Close(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := b64Writer.Close(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return "zstd:" + buf.String(), nil
|
||||
}
|
||||
|
||||
// ReadLastLines 返回日志的最后 n 行
|
||||
func (l *TinyLog) ReadLastLines(n int) ([]byte, error) {
|
||||
l.mu.RLock()
|
||||
defer l.mu.RUnlock()
|
||||
|
||||
// 刷新写入器以确保磁盘上的文件是最新的
|
||||
_ = l.writer.Flush()
|
||||
|
||||
stat, err := os.Stat(l.path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
size := stat.Size()
|
||||
var limit int64 = 65536 // 预览限制:最大 64KB
|
||||
if size < limit {
|
||||
limit = size
|
||||
}
|
||||
offset := size - limit
|
||||
|
||||
data := make([]byte, limit)
|
||||
f, err := os.Open(l.path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
_, err = f.ReadAt(data, offset)
|
||||
if err != nil && err != io.EOF {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
lines := bytes.Split(data, []byte{'\n'})
|
||||
if len(lines) > n+1 {
|
||||
return bytes.Join(lines[len(lines)-n-1:], []byte{'\n'}), nil
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// GetPath 返回临时文件路径
|
||||
func (l *TinyLog) GetPath() string {
|
||||
return l.path
|
||||
}
|
||||
|
||||
// CleanupOrphanedTinyLogs 启动时清理残留的临时日志文件
|
||||
func CleanupOrphanedTinyLogs() {
|
||||
tmpDir := os.TempDir()
|
||||
files, err := os.ReadDir(tmpDir)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
count := 0
|
||||
for _, file := range files {
|
||||
if !file.IsDir() && len(file.Name()) > 9 && file.Name()[:9] == "task_log_" && filepath.Ext(file.Name()) == ".log" {
|
||||
os.Remove(filepath.Join(tmpDir, file.Name()))
|
||||
count++
|
||||
}
|
||||
}
|
||||
if count > 0 {
|
||||
logger.Infof("[System] 清理了 %d 个残留的任务日志临时文件", count)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestTinyLog_UTF8Splitting(t *testing.T) {
|
||||
tl, err := NewTinyLog("test-utf8", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create TinyLog: %v", err)
|
||||
}
|
||||
defer tl.Close()
|
||||
|
||||
// "你好" in UTF-8: E4 BD A0, E5 a5 bd
|
||||
part1 := []byte{0xE4, 0xBD} // Partial "你"
|
||||
part2 := []byte{0xA0, 0xE5, 0xA5} // Rest of "你", partial "好"
|
||||
part3 := []byte{0xBD, '\n'} // Rest of "好", newline
|
||||
|
||||
_, _ = tl.Write(part1)
|
||||
if len(tl.remainder) != 2 {
|
||||
t.Errorf("Expected remainder len 2, got %d", len(tl.remainder))
|
||||
}
|
||||
|
||||
_, _ = tl.Write(part2)
|
||||
// Currently it should collect both parts but still no newline,
|
||||
// so remainder should be 5 bytes.
|
||||
if len(tl.remainder) != 5 {
|
||||
t.Errorf("Expected remainder len 5, got %d", len(tl.remainder))
|
||||
}
|
||||
|
||||
_, _ = tl.Write(part3)
|
||||
if len(tl.remainder) != 0 {
|
||||
t.Errorf("Expected remainder len 0 after newline, got %d", len(tl.remainder))
|
||||
}
|
||||
|
||||
// Read and verify
|
||||
data, err := tl.ReadLastLines(1)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadLastLines failed: %v", err)
|
||||
}
|
||||
if !bytes.Contains(data, []byte("你好")) {
|
||||
t.Errorf("Expected output to contain '你好', got %q", data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTinyLog_CarriageReturn(t *testing.T) {
|
||||
tl, err := NewTinyLog("test-cr", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create TinyLog: %v", err)
|
||||
}
|
||||
defer tl.Close()
|
||||
|
||||
input := []byte("progress: 50%\rprogress: 100%\r\n")
|
||||
_, _ = tl.Write(input)
|
||||
|
||||
data, err := tl.ReadLastLines(10)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadLastLines failed: %v", err)
|
||||
}
|
||||
|
||||
// Should contain both progress lines (or at least be split correctly)
|
||||
if !bytes.Contains(data, []byte("progress: 50%")) {
|
||||
t.Errorf("Expected output to contain 'progress: 50%%', got %q", data)
|
||||
}
|
||||
if !bytes.Contains(data, []byte("progress: 100%")) {
|
||||
t.Errorf("Expected output to contain 'progress: 100%%', got %q", data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTinyLog_LongLineCut(t *testing.T) {
|
||||
tl, err := NewTinyLog("test-long", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to create TinyLog: %v", err)
|
||||
}
|
||||
defer tl.Close()
|
||||
|
||||
// Create a buffer of maxLogBufferLen-1 bytes with a multi-byte character at the maxLogBufferLen boundary
|
||||
// We want to ensure it doesn't cut in the middle of a 3-byte char.
|
||||
longData := make([]byte, maxLogBufferLen-1)
|
||||
for i := range longData {
|
||||
longData[i] = 'A'
|
||||
}
|
||||
// "你" is E4 BD A0
|
||||
longData = append(longData, 0xE4, 0xBD, 0xA0) // This starts at index maxLogBufferLen-1.
|
||||
// Index maxLogBufferLen-1: E4
|
||||
// Index maxLogBufferLen: BD
|
||||
// Index maxLogBufferLen+1: A0
|
||||
// If we cut at maxLogBufferLen, we split E4 and BD.
|
||||
|
||||
_, _ = tl.Write(longData)
|
||||
|
||||
// Since it's > maxLogBufferLen and no newline, it should trigger the cut.
|
||||
// Our logic finds the last safe boundary before maxLogBufferLen.
|
||||
// RuneStart(E4) at maxLogBufferLen-1 is true. FullRune(E4 at maxLogBufferLen-1 in payload[:maxLogBufferLen]) is false.
|
||||
// So lastSafe should be maxLogBufferLen-1.
|
||||
|
||||
// The first maxLogBufferLen-1 bytes (all 'A') should be processed.
|
||||
// The "你" should be in remainder.
|
||||
|
||||
if len(tl.remainder) != 3 {
|
||||
t.Errorf("Expected remainder len 3 (the char '你'), got %d", len(tl.remainder))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type UserService struct{}
|
||||
|
||||
func NewUserService() *UserService {
|
||||
return &UserService{}
|
||||
}
|
||||
|
||||
func (us *UserService) hashPassword(password string) (string, error) {
|
||||
bytes, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(bytes), nil
|
||||
}
|
||||
|
||||
func (us *UserService) legacyHashPassword(password string) string {
|
||||
hash := sha256.Sum256([]byte(password + constant.Secret))
|
||||
return hex.EncodeToString(hash[:])
|
||||
}
|
||||
|
||||
func (us *UserService) CreateUser(username, password, email, role string) *models.User {
|
||||
hashedPassword, _ := us.hashPassword(password)
|
||||
user := &models.User{
|
||||
ID: utils.GenerateID(),
|
||||
Username: username,
|
||||
Password: hashedPassword,
|
||||
Email: email,
|
||||
Role: role,
|
||||
TokenVersion: 1,
|
||||
}
|
||||
database.DB.Create(user)
|
||||
return user
|
||||
}
|
||||
|
||||
func (us *UserService) GetUserByUsername(username string) *models.User {
|
||||
var user models.User
|
||||
res := database.DB.Where("username = ?", username).Limit(1).Find(&user)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil
|
||||
}
|
||||
return &user
|
||||
}
|
||||
|
||||
func (us *UserService) GetUserByID(id string) (*models.User, error) {
|
||||
var user models.User
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&user)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return nil, res.Error
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func (us *UserService) ValidatePassword(user *models.User, password string) bool {
|
||||
// 尝试 bcrypt 校验
|
||||
err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(password))
|
||||
if err == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
// 如果 bcrypt 失败,检查是否为旧的 SHA256 格式
|
||||
// 旧格式是 64 位十六进制字符串
|
||||
if len(user.Password) == 64 && !strings.HasPrefix(user.Password, "$2") {
|
||||
if user.Password == us.legacyHashPassword(password) {
|
||||
// 校验成功,迁移到 bcrypt
|
||||
newHash, err := us.hashPassword(password)
|
||||
if err == nil {
|
||||
database.DB.Model(user).Update("password", newHash)
|
||||
}
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func (us *UserService) EnsureAdminExists() {
|
||||
var count int64
|
||||
database.DB.Model(&models.User{}).Where("role = ?", "admin").Count(&count)
|
||||
if count == 0 {
|
||||
us.CreateUser("admin", "admin123", "admin@local", "admin")
|
||||
}
|
||||
}
|
||||
|
||||
func (us *UserService) AuthenticateUser(username, password string) bool {
|
||||
user := us.GetUserByUsername(username)
|
||||
if user == nil {
|
||||
return false
|
||||
}
|
||||
return us.ValidatePassword(user, password)
|
||||
}
|
||||
|
||||
func (us *UserService) UpdatePassword(userID string, newPassword string) error {
|
||||
hashedPassword, err := us.hashPassword(newPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// 修改密码时同时失效旧 Token
|
||||
return database.DB.Model(&models.User{}).Where("id = ?", userID).Updates(map[string]interface{}{
|
||||
"password": hashedPassword,
|
||||
"token_version": gorm.Expr("token_version + 1"),
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (us *UserService) InvalidateUserTokens(userID string) error {
|
||||
return database.DB.Model(&models.User{}).Where("id = ?", userID).Update("token_version", gorm.Expr("token_version + 1")).Error
|
||||
}
|
||||
|
||||
func (us *UserService) UpdateAccount(userID string, newUsername string) error {
|
||||
var user models.User
|
||||
res := database.DB.Where("id = ?", userID).Limit(1).Find(&user)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
return fmt.Errorf("未找到对应的用户")
|
||||
}
|
||||
|
||||
updates := make(map[string]interface{})
|
||||
if newUsername != "" && newUsername != user.Username {
|
||||
// 检查用户名是否已存在
|
||||
var count int64
|
||||
database.DB.Model(&models.User{}).Where("username = ? AND id <> ?", newUsername, userID).Count(&count)
|
||||
if count > 0 {
|
||||
return fmt.Errorf("用户名 [%s] 已被占用", newUsername)
|
||||
}
|
||||
updates["username"] = newUsername
|
||||
// 用户名变更,必须失效所有 Token,因为 Token 中包含 Username 且中间件会校验
|
||||
updates["token_version"] = gorm.Expr("token_version + 1")
|
||||
}
|
||||
|
||||
if len(updates) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
return database.DB.Model(&user).Updates(updates).Error
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
type WebUIService struct {
|
||||
settingsService *SettingsService
|
||||
}
|
||||
|
||||
func NewWebUIService(settingsService *SettingsService) *WebUIService {
|
||||
return &WebUIService{
|
||||
settingsService: settingsService,
|
||||
}
|
||||
}
|
||||
|
||||
// GetActiveWebUIFS 返回当前激活WebUI的 fs.FS 接口。
|
||||
// 如果激活的WebUI是 "default" 或者不存在,则返回 nil。
|
||||
func (s *WebUIService) GetActiveWebUIFS() fs.FS {
|
||||
activeWebUI := s.settingsService.Get(constant.SectionSite, constant.KeyActiveWebUI)
|
||||
if activeWebUI == "" || activeWebUI == "default" {
|
||||
return nil
|
||||
}
|
||||
|
||||
webuiDir := filepath.Join(constant.DataDir, "webuis", activeWebUI)
|
||||
|
||||
// 检查是否存在 uimanifest.json 以确认这是一个有效的WebUI目录
|
||||
if _, err := os.Stat(filepath.Join(webuiDir, "uimanifest.json")); os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
|
||||
return os.DirFS(webuiDir)
|
||||
}
|
||||
|
||||
// WebUIManifest 代表 uimanifest.json 中的元数据
|
||||
type WebUIManifest struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
Author string `json:"author"`
|
||||
Description string `json:"description"`
|
||||
MinPanelVersion string `json:"min_panel_version"`
|
||||
}
|
||||
|
||||
// GetWebUIs 获取所有可用的WebUI列表
|
||||
func (s *WebUIService) GetWebUIs() ([]WebUIManifest, error) {
|
||||
// 默认WebUI总是可用的
|
||||
webuis := []WebUIManifest{
|
||||
{
|
||||
Name: "default",
|
||||
Version: "builtin",
|
||||
Author: "TaskPool",
|
||||
Description: "内置默认WebUI",
|
||||
},
|
||||
}
|
||||
|
||||
records := s.settingsService.GetSection("webui")
|
||||
for name, val := range records {
|
||||
var manifest WebUIManifest
|
||||
if err := json.Unmarshal([]byte(val), &manifest); err == nil {
|
||||
manifest.Name = name // 强制名称匹配
|
||||
webuis = append(webuis, manifest)
|
||||
}
|
||||
}
|
||||
|
||||
return webuis, nil
|
||||
}
|
||||
|
||||
// ExtractWebUI 将 zip 或 tar.gz 压缩包解压到WebUI目录
|
||||
func (s *WebUIService) ExtractWebUI(zipPath string) (string, error) {
|
||||
// 1. 创建临时解压目录(放在 DataDir 下避免跨分区移动失败)
|
||||
baseWebUIDir := filepath.Join(constant.DataDir, "webuis")
|
||||
if err := os.MkdirAll(baseWebUIDir, 0755); err != nil {
|
||||
return "", fmt.Errorf("无法创建 WebUI 基础目录: %v", err)
|
||||
}
|
||||
tmpDir, err := os.MkdirTemp(baseWebUIDir, "tmp-webui-*")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("无法创建临时解压目录: %v", err)
|
||||
}
|
||||
// 确保在出错时清理临时目录
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
// 2. 根据后缀名选择解压方法
|
||||
var extractErr error
|
||||
if strings.HasSuffix(strings.ToLower(zipPath), ".tar.gz") || strings.HasSuffix(strings.ToLower(zipPath), ".tgz") {
|
||||
extractErr = utils.ExtractTarGz(zipPath, tmpDir)
|
||||
} else {
|
||||
extractErr = utils.ExtractZip(zipPath, tmpDir)
|
||||
}
|
||||
if extractErr != nil {
|
||||
return "", fmt.Errorf("解压WebUI包失败: %v", extractErr)
|
||||
}
|
||||
|
||||
// 3. 读取并解析 uimanifest.json
|
||||
manifestPath := filepath.Join(tmpDir, "uimanifest.json")
|
||||
manifestData, err := os.ReadFile(manifestPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", fmt.Errorf("webui package must contain a uimanifest.json file")
|
||||
}
|
||||
return "", fmt.Errorf("无法读取 uimanifest.json: %v", err)
|
||||
}
|
||||
|
||||
var webuiManifest WebUIManifest
|
||||
if err := json.Unmarshal(manifestData, &webuiManifest); err != nil {
|
||||
return "", fmt.Errorf("invalid uimanifest.json format")
|
||||
}
|
||||
|
||||
webuiName := webuiManifest.Name
|
||||
if webuiName == "" || webuiName == "default" {
|
||||
return "", fmt.Errorf("invalid webui name in uimanifest.json")
|
||||
}
|
||||
|
||||
// 确保WebUI名称安全,防止目录穿越
|
||||
webuiName = filepath.Base(filepath.Clean(webuiName))
|
||||
|
||||
// 4. 确保压缩包中包含 index.html 入口文件
|
||||
if _, err := os.Stat(filepath.Join(tmpDir, "index.html")); os.IsNotExist(err) {
|
||||
return "", fmt.Errorf("webui package must contain an index.html file")
|
||||
}
|
||||
|
||||
// 5. 移动临时目录到最终的目标目录
|
||||
targetDir := filepath.Join(constant.DataDir, "webuis", webuiName)
|
||||
// 如果目标目录已存在,先删除旧版本
|
||||
os.RemoveAll(targetDir)
|
||||
if err := os.Rename(tmpDir, targetDir); err != nil {
|
||||
return "", fmt.Errorf("覆盖安装WebUI失败: %v", err)
|
||||
}
|
||||
|
||||
// 6. 将记录保存到 settings 表中
|
||||
manifestJSON, _ := json.Marshal(webuiManifest)
|
||||
if err := s.settingsService.Set("webui", webuiName, string(manifestJSON)); err != nil {
|
||||
// 回滚
|
||||
os.RemoveAll(targetDir)
|
||||
return "", fmt.Errorf("保存WebUI记录失败: %v", err)
|
||||
}
|
||||
|
||||
return webuiName, nil
|
||||
}
|
||||
|
||||
// DeleteWebUI 删除自定义WebUI
|
||||
func (s *WebUIService) DeleteWebUI(name string) error {
|
||||
if name == "" || name == "default" {
|
||||
return fmt.Errorf("cannot delete default webui")
|
||||
}
|
||||
|
||||
name = filepath.Base(filepath.Clean(name))
|
||||
targetDir := filepath.Join(constant.DataDir, "webuis", name)
|
||||
|
||||
activeWebUI := s.settingsService.Get(constant.SectionSite, constant.KeyActiveWebUI)
|
||||
if activeWebUI == name {
|
||||
return fmt.Errorf("cannot delete currently active webui")
|
||||
}
|
||||
|
||||
if err := os.RemoveAll(targetDir); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 从 settings 表中移除记录
|
||||
return s.settingsService.Delete("webui", name)
|
||||
}
|
||||
|
||||
// SetActiveWebUI 设置当前的活动WebUI
|
||||
func (s *WebUIService) SetActiveWebUI(name string) error {
|
||||
if name != "default" {
|
||||
name = filepath.Base(filepath.Clean(name))
|
||||
targetDir := filepath.Join(constant.DataDir, "webuis", name)
|
||||
if _, err := os.Stat(filepath.Join(targetDir, "uimanifest.json")); os.IsNotExist(err) {
|
||||
return fmt.Errorf("webui %s not found", name)
|
||||
}
|
||||
}
|
||||
|
||||
return s.settingsService.Set(constant.SectionSite, constant.KeyActiveWebUI, name)
|
||||
}
|
||||
Reference in New Issue
Block a user