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:
2026-07-26 08:43:52 +08:00
commit e6956aa001
397 changed files with 73621 additions and 0 deletions
+565
View File
@@ -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 生成随机 Token64位十六进制)
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
}
+410
View File
@@ -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()
}
+357
View File
@@ -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)
})
}
+530
View File
@@ -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
}
+226
View File
@@ -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
}
+350
View File
@@ -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
}
+140
View File
@@ -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
}
+45
View File
@@ -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
}
+45
View File
@@ -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
}
+45
View File
@@ -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
}
+26
View File
@@ -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
}
+103
View File
@@ -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
}
+62
View File
@@ -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])
}
}
})
}
}
+45
View File
@@ -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
}
+50
View File
@@ -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
}
+46
View File
@@ -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
}
+46
View File
@@ -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
}
+202
View File
@@ -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
}
+45
View File
@@ -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
}
+48
View File
@@ -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
}
+105
View File
@@ -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
}
+87
View File
@@ -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)
}
}
}
+27
View File
@@ -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
}
+46
View File
@@ -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
}
+47
View File
@@ -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
}
+46
View File
@@ -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
}
+46
View File
@@ -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
}
+409
View File
@@ -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{})
}
+66
View File
@@ -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("--------------------------------------------------")
}
+263
View File
@@ -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)
}
+85
View File
@@ -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
}
+128
View File
@@ -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
}
+446
View File
@@ -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)
}
}
}
+405
View File
@@ -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 = &lt
}
}
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
}
+132
View File
@@ -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()
}
+564
View File
@@ -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(&notifyWay)
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(&notifyWays)
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
}
+242
View File
@@ -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
}
}
+27
View File
@@ -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{}
}
}
+71
View File
@@ -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
}
+31
View File
@@ -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 "", ""
}
+251
View File
@@ -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
}
+56
View File
@@ -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
}
+61
View File
@@ -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
}
+204
View File
@@ -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
}
+132
View File
@@ -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
+251
View File
@@ -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
}
+287
View File
@@ -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("")
}
}
}
+386
View File
@@ -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)
}
}
+104
View File
@@ -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))
}
}
+150
View File
@@ -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
}
+181
View File
@@ -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)
}