fix: agent exec log error
This commit is contained in:
@@ -0,0 +1,168 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/logger"
|
||||
"baihu/internal/models"
|
||||
|
||||
"github.com/robfig/cron/v3"
|
||||
)
|
||||
|
||||
// 东八区时区
|
||||
var cstZone = time.FixedZone("CST", 8*3600)
|
||||
|
||||
// CronService manages scheduled tasks using robfig/cron
|
||||
type CronService struct {
|
||||
cron *cron.Cron
|
||||
taskService *TaskService
|
||||
executorService *ExecutorService
|
||||
entryMap map[uint]cron.EntryID // task ID -> cron entry ID
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewCronService creates a new cron service
|
||||
func NewCronService(taskService *TaskService, executorService *ExecutorService) *CronService {
|
||||
// 使用秒级精度的 cron parser,支持 6 位表达式(秒 分 时 日 月 周),使用东八区时区
|
||||
c := cron.New(cron.WithSeconds(), cron.WithLocation(cstZone))
|
||||
|
||||
return &CronService{
|
||||
cron: c,
|
||||
taskService: taskService,
|
||||
executorService: executorService,
|
||||
entryMap: make(map[uint]cron.EntryID),
|
||||
}
|
||||
}
|
||||
|
||||
// Start starts the cron service and loads all enabled tasks
|
||||
func (cs *CronService) Start() {
|
||||
cs.loadTasks()
|
||||
cs.cron.Start()
|
||||
logger.Info("[Cron] 调度服务已启动")
|
||||
}
|
||||
|
||||
// Stop stops the cron service
|
||||
func (cs *CronService) Stop() {
|
||||
ctx := cs.cron.Stop()
|
||||
<-ctx.Done()
|
||||
logger.Info("[Cron] 调度服务已停止")
|
||||
}
|
||||
|
||||
// loadTasks loads all enabled tasks from database
|
||||
func (cs *CronService) loadTasks() {
|
||||
tasks := cs.taskService.GetTasks()
|
||||
count := 0
|
||||
for _, task := range tasks {
|
||||
// 只调度本地任务(agent_id 为空)
|
||||
if task.Enabled && task.AgentID == nil {
|
||||
err := cs.addTask(&task, false)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
count++
|
||||
}
|
||||
}
|
||||
logger.Infof("[Cron] 启动调度已加载 %d 个定时任务", count)
|
||||
}
|
||||
|
||||
// addTask 内部添加任务方法,silent 控制是否打印日志
|
||||
func (cs *CronService) addTask(task *models.Task, logEnabled bool) error {
|
||||
cs.mu.Lock()
|
||||
|
||||
// 如果已存在,先移除
|
||||
if entryID, exists := cs.entryMap[task.ID]; exists {
|
||||
cs.cron.Remove(entryID)
|
||||
delete(cs.entryMap, task.ID)
|
||||
}
|
||||
|
||||
taskID := task.ID
|
||||
entryID, err := cs.cron.AddFunc(task.Schedule, func() {
|
||||
cs.runTask(taskID)
|
||||
})
|
||||
if err != nil {
|
||||
cs.mu.Unlock()
|
||||
logger.Errorf("[Cron] 添加任务失败 #%d: %v", task.ID, err)
|
||||
return err
|
||||
}
|
||||
|
||||
cs.entryMap[task.ID] = entryID
|
||||
cs.mu.Unlock()
|
||||
|
||||
if logEnabled {
|
||||
logger.Infof("[Cron] 任务已调度 #%d %s (%s)", task.ID, task.Name, task.Schedule)
|
||||
}
|
||||
|
||||
// 更新下次运行时间
|
||||
cs.updateNextRun(task.ID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddTask adds a task to the cron scheduler
|
||||
func (cs *CronService) AddTask(task *models.Task) error {
|
||||
return cs.addTask(task, true)
|
||||
}
|
||||
|
||||
// RemoveTask removes a task from the cron scheduler
|
||||
func (cs *CronService) RemoveTask(taskID uint) {
|
||||
cs.mu.Lock()
|
||||
defer cs.mu.Unlock()
|
||||
|
||||
if entryID, exists := cs.entryMap[taskID]; exists {
|
||||
cs.cron.Remove(entryID)
|
||||
delete(cs.entryMap, taskID)
|
||||
logger.Infof("[Cron] 任务已移除 #%d", taskID)
|
||||
}
|
||||
}
|
||||
|
||||
// runTask executes a task and updates its status
|
||||
func (cs *CronService) runTask(taskID uint) {
|
||||
// 获取任务信息用于日志
|
||||
task := cs.taskService.GetTaskByID(int(taskID))
|
||||
if task != nil {
|
||||
logger.Infof("[Cron] 执行任务 #%d %s", taskID, task.Name)
|
||||
} else {
|
||||
logger.Infof("[Cron] 执行任务 #%d", taskID)
|
||||
}
|
||||
|
||||
// 更新 last_run
|
||||
now := time.Now()
|
||||
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("last_run", now)
|
||||
|
||||
// 将任务加入队列执行(通过 worker pool 控制并发)
|
||||
cs.executorService.EnqueueTask(int(taskID))
|
||||
|
||||
// 更新 next_run
|
||||
cs.updateNextRun(taskID)
|
||||
}
|
||||
|
||||
// updateNextRun updates the next run time for a task
|
||||
func (cs *CronService) updateNextRun(taskID uint) {
|
||||
cs.mu.RLock()
|
||||
entryID, exists := cs.entryMap[taskID]
|
||||
cs.mu.RUnlock()
|
||||
|
||||
if !exists {
|
||||
return
|
||||
}
|
||||
|
||||
entry := cs.cron.Entry(entryID)
|
||||
if !entry.Next.IsZero() {
|
||||
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("next_run", entry.Next)
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateCron validates a cron expression (6 fields: second minute hour day month weekday)
|
||||
func (cs *CronService) ValidateCron(expression string) error {
|
||||
parser := cron.NewParser(cron.Second | cron.Minute | cron.Hour | cron.Dom | cron.Month | cron.Dow | cron.Descriptor)
|
||||
_, err := parser.Parse(expression)
|
||||
return err
|
||||
}
|
||||
|
||||
// GetScheduledCount returns the number of scheduled tasks
|
||||
func (cs *CronService) GetScheduledCount() int {
|
||||
cs.mu.RLock()
|
||||
defer cs.mu.RUnlock()
|
||||
return len(cs.entryMap)
|
||||
}
|
||||
@@ -0,0 +1,434 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"baihu/internal/constant"
|
||||
"baihu/internal/logger"
|
||||
"baihu/internal/models"
|
||||
"baihu/internal/utils"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// SettingsService 接口定义(避免循环依赖)
|
||||
type SettingsService interface {
|
||||
Get(section, key string) string
|
||||
}
|
||||
|
||||
// EnvService 接口定义(避免循环依赖)
|
||||
type EnvService interface {
|
||||
GetEnvVarsByIDs(ids string) []string
|
||||
}
|
||||
|
||||
// ExecutionResult represents the result of a task execution
|
||||
type ExecutionResult struct {
|
||||
TaskID int
|
||||
Success bool
|
||||
Output string
|
||||
Error string
|
||||
Start time.Time
|
||||
End time.Time
|
||||
}
|
||||
|
||||
// taskJob 任务队列项
|
||||
type taskJob struct {
|
||||
taskID int
|
||||
}
|
||||
|
||||
// ExecutorService handles task execution
|
||||
type ExecutorService struct {
|
||||
taskService *TaskService
|
||||
taskExecutionService *TaskExecutionService
|
||||
settingsService SettingsService
|
||||
envService EnvService
|
||||
results []ExecutionResult
|
||||
runningTasks map[int]bool
|
||||
mu sync.RWMutex
|
||||
resultsMu sync.RWMutex
|
||||
|
||||
// 任务队列和 worker pool
|
||||
taskQueue chan taskJob
|
||||
workerCount int
|
||||
rateLimiter <-chan time.Time
|
||||
stopCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
// NewExecutorService creates a new executor service
|
||||
func NewExecutorService(taskService *TaskService, taskExecutionService *TaskExecutionService, settingsService SettingsService, envService EnvService) *ExecutorService {
|
||||
// 从设置中读取调度配置
|
||||
workerCount := getIntSetting(settingsService, constant.SectionScheduler, constant.KeyWorkerCount, 4)
|
||||
queueSize := getIntSetting(settingsService, constant.SectionScheduler, constant.KeyQueueSize, 100)
|
||||
rateInterval := getIntSetting(settingsService, constant.SectionScheduler, constant.KeyRateInterval, 200)
|
||||
|
||||
logger.Infof("[Executor] 配置: workers=%d, queue=%d, rate=%dms", workerCount, queueSize, rateInterval)
|
||||
|
||||
es := &ExecutorService{
|
||||
taskService: taskService,
|
||||
taskExecutionService: taskExecutionService,
|
||||
settingsService: settingsService,
|
||||
envService: envService,
|
||||
results: make([]ExecutionResult, 0, 100),
|
||||
runningTasks: make(map[int]bool),
|
||||
taskQueue: make(chan taskJob, queueSize),
|
||||
workerCount: workerCount,
|
||||
rateLimiter: time.Tick(time.Duration(rateInterval) * time.Millisecond),
|
||||
stopCh: make(chan struct{}),
|
||||
}
|
||||
|
||||
// 启动 worker pool
|
||||
es.startWorkers()
|
||||
|
||||
return es
|
||||
}
|
||||
|
||||
// getIntSetting 从设置中获取整数值
|
||||
func getIntSetting(s SettingsService, section, key string, defaultVal int) int {
|
||||
val := s.Get(section, key)
|
||||
if val == "" {
|
||||
return defaultVal
|
||||
}
|
||||
var result int
|
||||
if _, err := fmt.Sscanf(val, "%d", &result); err != nil {
|
||||
return defaultVal
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// startWorkers 启动 worker pool
|
||||
func (es *ExecutorService) startWorkers() {
|
||||
for i := 0; i < es.workerCount; i++ {
|
||||
es.wg.Add(1)
|
||||
go es.worker(i)
|
||||
}
|
||||
}
|
||||
|
||||
// worker 从队列中取任务执行
|
||||
func (es *ExecutorService) worker(id int) {
|
||||
defer es.wg.Done()
|
||||
for {
|
||||
select {
|
||||
case <-es.stopCh:
|
||||
return
|
||||
case job := <-es.taskQueue:
|
||||
// 速率限制
|
||||
<-es.rateLimiter
|
||||
es.executeTaskInternal(job.taskID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Stop 停止 executor service
|
||||
func (es *ExecutorService) Stop() {
|
||||
close(es.stopCh)
|
||||
es.wg.Wait()
|
||||
}
|
||||
|
||||
// Reload 重新加载配置并重建 worker pool
|
||||
func (es *ExecutorService) Reload() {
|
||||
logger.Info("[Executor] 正在重载配置...")
|
||||
|
||||
// 停止现有 workers
|
||||
close(es.stopCh)
|
||||
es.wg.Wait()
|
||||
logger.Info("[Executor] 已停止工作线程")
|
||||
|
||||
// 从设置中读取新配置
|
||||
workerCount := getIntSetting(es.settingsService, constant.SectionScheduler, constant.KeyWorkerCount, 4)
|
||||
queueSize := getIntSetting(es.settingsService, constant.SectionScheduler, constant.KeyQueueSize, 100)
|
||||
rateInterval := getIntSetting(es.settingsService, constant.SectionScheduler, constant.KeyRateInterval, 200)
|
||||
|
||||
// 重建 channel 和配置
|
||||
es.mu.Lock()
|
||||
es.taskQueue = make(chan taskJob, queueSize)
|
||||
es.workerCount = workerCount
|
||||
es.rateLimiter = time.Tick(time.Duration(rateInterval) * time.Millisecond)
|
||||
es.stopCh = make(chan struct{})
|
||||
es.mu.Unlock()
|
||||
|
||||
// 启动新的 workers
|
||||
es.startWorkers()
|
||||
|
||||
logger.Infof("[Executor] 配置已重载: workers=%d, queue=%d, rate=%dms", workerCount, queueSize, rateInterval)
|
||||
}
|
||||
|
||||
// EnqueueTask 将任务加入队列(供 cron 调度器调用)
|
||||
func (es *ExecutorService) EnqueueTask(taskID int) {
|
||||
select {
|
||||
case es.taskQueue <- taskJob{taskID: taskID}:
|
||||
// 成功入队
|
||||
default:
|
||||
// 队列满,直接执行(降级处理)
|
||||
logger.Warnf("[Executor] 任务队列已满,直接执行任务 #%d", taskID)
|
||||
go es.executeTaskInternal(taskID)
|
||||
}
|
||||
}
|
||||
|
||||
// ExecuteTask executes a task by ID(同步执行,供 API 调用)
|
||||
func (es *ExecutorService) ExecuteTask(taskID int) *ExecutionResult {
|
||||
return es.executeTaskInternal(taskID)
|
||||
}
|
||||
|
||||
// executeTaskInternal 内部执行任务逻辑
|
||||
func (es *ExecutorService) executeTaskInternal(taskID int) *ExecutionResult {
|
||||
task := es.taskService.GetTaskByID(taskID)
|
||||
if task == nil {
|
||||
return &ExecutionResult{
|
||||
TaskID: taskID,
|
||||
Success: false,
|
||||
Error: "Task not found",
|
||||
Start: time.Now(),
|
||||
End: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
// 标记任务开始运行
|
||||
es.mu.Lock()
|
||||
es.runningTasks[taskID] = true
|
||||
es.mu.Unlock()
|
||||
|
||||
var result *ExecutionResult
|
||||
|
||||
// 使用统一的任务执行服务
|
||||
req := &TaskExecutionRequest{
|
||||
TaskID: uint(taskID),
|
||||
Task: task,
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
err := es.taskExecutionService.ExecuteTask(req)
|
||||
end := time.Now()
|
||||
|
||||
if err != nil {
|
||||
result = &ExecutionResult{
|
||||
TaskID: taskID,
|
||||
Success: false,
|
||||
Error: err.Error(),
|
||||
Start: start,
|
||||
End: end,
|
||||
}
|
||||
} else {
|
||||
result = &ExecutionResult{
|
||||
TaskID: taskID,
|
||||
Success: true,
|
||||
Output: "任务已提交执行",
|
||||
Start: start,
|
||||
End: end,
|
||||
}
|
||||
}
|
||||
|
||||
// 标记任务结束
|
||||
es.mu.Lock()
|
||||
delete(es.runningTasks, taskID)
|
||||
es.mu.Unlock()
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// executeNormalTask 执行普通任务
|
||||
func (es *ExecutorService) executeNormalTask(task *models.Task) *ExecutionResult {
|
||||
// 演示模式下使用 echo 替换实际命令
|
||||
if constant.DemoMode {
|
||||
return es.ExecuteCommandWithOptions("echo '[演示模式] 任务执行已跳过,实际命令不会运行'", time.Minute, nil, "")
|
||||
}
|
||||
|
||||
// 加载环境变量
|
||||
envVars := es.envService.GetEnvVarsByIDs(task.Envs)
|
||||
|
||||
// 确定工作目录
|
||||
workDir := task.WorkDir
|
||||
if workDir == "" {
|
||||
workDir = constant.ScriptsWorkDir
|
||||
}
|
||||
|
||||
// 使用任务配置的超时时间
|
||||
timeout := task.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = constant.DefaultTaskTimeout
|
||||
}
|
||||
return es.ExecuteCommandWithOptions(task.Command, time.Duration(timeout)*time.Minute, envVars, workDir)
|
||||
}
|
||||
|
||||
// executeRepoTask 执行仓库同步任务(调用 sync.py)
|
||||
func (es *ExecutorService) executeRepoTask(task *models.Task) *ExecutionResult {
|
||||
// 演示模式下使用 echo 替换实际命令
|
||||
if constant.DemoMode {
|
||||
return es.ExecuteCommandWithOptions("echo '[演示模式] 仓库同步已跳过,实际命令不会运行'", time.Minute, nil, "")
|
||||
}
|
||||
|
||||
result := &ExecutionResult{
|
||||
Success: false,
|
||||
Start: time.Now(),
|
||||
}
|
||||
|
||||
// 解析仓库配置
|
||||
var config models.RepoConfig
|
||||
if err := json.Unmarshal([]byte(task.Config), &config); err != nil {
|
||||
result.End = time.Now()
|
||||
result.Error = "解析仓库配置失败: " + err.Error()
|
||||
return result
|
||||
}
|
||||
|
||||
// 处理目标路径:为空则使用 scripts 目录,相对路径则基于 scripts 目录
|
||||
targetPath := config.TargetPath
|
||||
if targetPath == "" {
|
||||
targetPath = constant.ScriptsWorkDir
|
||||
} else if !filepath.IsAbs(targetPath) {
|
||||
targetPath = filepath.Join(constant.ScriptsWorkDir, targetPath)
|
||||
}
|
||||
// 转换为绝对路径
|
||||
absTargetPath, err := filepath.Abs(targetPath)
|
||||
if err != nil {
|
||||
absTargetPath = targetPath
|
||||
}
|
||||
|
||||
// 构建 sync.py 命令参数
|
||||
args := []string{
|
||||
"/opt/sync.py",
|
||||
"--source-type", config.SourceType,
|
||||
"--source-url", config.SourceURL,
|
||||
"--target-path", absTargetPath,
|
||||
}
|
||||
|
||||
// Git 分支
|
||||
if config.Branch != "" {
|
||||
args = append(args, "--branch", config.Branch)
|
||||
}
|
||||
|
||||
// 稀疏路径
|
||||
if config.SparsePath != "" {
|
||||
args = append(args, "--path", config.SparsePath)
|
||||
}
|
||||
|
||||
// 单文件模式
|
||||
if config.SingleFile {
|
||||
args = append(args, "--single-file")
|
||||
}
|
||||
|
||||
// 代理设置
|
||||
if config.Proxy != "" && config.Proxy != "none" {
|
||||
args = append(args, "--proxy", config.Proxy)
|
||||
if config.Proxy == "custom" && config.ProxyURL != "" {
|
||||
args = append(args, "--proxy-url", config.ProxyURL)
|
||||
}
|
||||
}
|
||||
|
||||
// 认证 Token
|
||||
if config.AuthToken != "" {
|
||||
args = append(args, "--auth-token", config.AuthToken)
|
||||
}
|
||||
|
||||
// 构建命令
|
||||
command := "python3 " + strings.Join(args, " ")
|
||||
|
||||
// 使用任务配置的超时时间
|
||||
timeout := task.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = constant.DefaultTaskTimeout
|
||||
}
|
||||
|
||||
// 执行命令
|
||||
execResult := es.ExecuteCommandWithOptions(command, time.Duration(timeout)*time.Minute, nil, "/opt")
|
||||
|
||||
result.End = time.Now()
|
||||
result.Output = execResult.Output
|
||||
result.Success = execResult.Success
|
||||
result.Error = execResult.Error
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetRunningCount 获取正在运行的任务数量
|
||||
func (es *ExecutorService) GetRunningCount() int {
|
||||
es.mu.RLock()
|
||||
defer es.mu.RUnlock()
|
||||
return len(es.runningTasks)
|
||||
}
|
||||
|
||||
// ExecuteCommand executes a shell command with default timeout
|
||||
func (es *ExecutorService) ExecuteCommand(command string) *ExecutionResult {
|
||||
return es.ExecuteCommandWithTimeout(command, time.Duration(constant.DefaultTaskTimeout)*time.Minute)
|
||||
}
|
||||
|
||||
// ExecuteCommandWithTimeout executes a shell command with specified timeout
|
||||
func (es *ExecutorService) ExecuteCommandWithTimeout(command string, timeout time.Duration) *ExecutionResult {
|
||||
return es.ExecuteCommandWithEnv(command, timeout, nil)
|
||||
}
|
||||
|
||||
// ExecuteCommandWithEnv executes a shell command with specified timeout and environment variables
|
||||
func (es *ExecutorService) ExecuteCommandWithEnv(command string, timeout time.Duration, envVars []string) *ExecutionResult {
|
||||
return es.ExecuteCommandWithOptions(command, timeout, envVars, "")
|
||||
}
|
||||
|
||||
// ExecuteCommandWithOptions executes a shell command with specified timeout, environment variables and working directory
|
||||
func (es *ExecutorService) ExecuteCommandWithOptions(command string, timeout time.Duration, envVars []string, workDir string) *ExecutionResult {
|
||||
result := &ExecutionResult{
|
||||
Success: false,
|
||||
Start: time.Now(),
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
|
||||
shell, args := utils.GetShellCommand(command)
|
||||
cmd := exec.CommandContext(ctx, shell, args...)
|
||||
var stdout, stderr bytes.Buffer
|
||||
cmd.Stdout = &stdout
|
||||
cmd.Stderr = &stderr
|
||||
|
||||
// 设置工作目录
|
||||
if workDir != "" {
|
||||
cmd.Dir = workDir
|
||||
}
|
||||
|
||||
// 设置环境变量:继承系统环境变量 + 自定义环境变量
|
||||
if len(envVars) > 0 {
|
||||
cmd.Env = append(os.Environ(), envVars...)
|
||||
}
|
||||
|
||||
err := cmd.Run()
|
||||
result.End = time.Now()
|
||||
|
||||
result.Output = stdout.String()
|
||||
if err != nil {
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
result.Error = "执行超时\n" + stderr.String()
|
||||
} else {
|
||||
result.Error = err.Error() + "\n" + stderr.String()
|
||||
}
|
||||
} else {
|
||||
result.Success = true
|
||||
}
|
||||
|
||||
// 使用独立锁保存结果
|
||||
es.resultsMu.Lock()
|
||||
es.results = append(es.results, *result)
|
||||
if len(es.results) > 100 {
|
||||
es.results = es.results[1:]
|
||||
}
|
||||
es.resultsMu.Unlock()
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetLastResults returns the last execution results
|
||||
func (es *ExecutorService) GetLastResults(count int) []ExecutionResult {
|
||||
es.resultsMu.RLock()
|
||||
defer es.resultsMu.RUnlock()
|
||||
|
||||
start := 0
|
||||
if len(es.results) > count {
|
||||
start = len(es.results) - count
|
||||
}
|
||||
|
||||
results := make([]ExecutionResult, len(es.results[start:]))
|
||||
copy(results, es.results[start:])
|
||||
return results
|
||||
}
|
||||
@@ -0,0 +1,284 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/logger"
|
||||
"baihu/internal/models"
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// AgentWSManager 接口定义(避免循环依赖)
|
||||
type AgentWSManager interface {
|
||||
SendToAgent(agentID uint, msgType string, data interface{}) error
|
||||
}
|
||||
|
||||
// TaskExecutionService 统一的任务执行服务
|
||||
type TaskExecutionService struct {
|
||||
taskLogService *TaskLogService
|
||||
agentWSManager AgentWSManager
|
||||
}
|
||||
|
||||
// NewTaskExecutionService 创建任务执行服务
|
||||
func NewTaskExecutionService(agentWSManager AgentWSManager, sendStatsService SendStatsService) *TaskExecutionService {
|
||||
return &TaskExecutionService{
|
||||
taskLogService: NewTaskLogService(sendStatsService),
|
||||
agentWSManager: agentWSManager,
|
||||
}
|
||||
}
|
||||
|
||||
// TaskExecutionRequest 任务执行请求
|
||||
type TaskExecutionRequest struct {
|
||||
TaskID uint
|
||||
Task *models.Task
|
||||
AgentID *uint // nil 表示本地执行
|
||||
}
|
||||
|
||||
// TaskExecutionResult 任务执行结果
|
||||
type TaskExecutionResult struct {
|
||||
TaskID uint
|
||||
AgentID *uint
|
||||
Command string
|
||||
Output string
|
||||
Status string // success, failed
|
||||
Duration int64 // milliseconds
|
||||
ExitCode int
|
||||
Start time.Time
|
||||
End time.Time
|
||||
}
|
||||
|
||||
// ExecuteTask 执行任务(统一入口)
|
||||
func (s *TaskExecutionService) ExecuteTask(req *TaskExecutionRequest) error {
|
||||
if req.Task.AgentID != nil && *req.Task.AgentID > 0 {
|
||||
// 远程执行:通过 Agent
|
||||
return s.executeRemote(req)
|
||||
}
|
||||
// 本地执行
|
||||
return s.executeLocal(req)
|
||||
}
|
||||
|
||||
// executeLocal 本地执行任务
|
||||
func (s *TaskExecutionService) executeLocal(req *TaskExecutionRequest) error {
|
||||
task := req.Task
|
||||
logger.Infof("[TaskExecution] 本地执行任务 #%d: %s", task.ID, task.Name)
|
||||
|
||||
start := time.Now()
|
||||
|
||||
// 准备命令
|
||||
ctx, cancel := s.createContext(task.Timeout)
|
||||
defer cancel()
|
||||
|
||||
cmd, err := s.prepareCommand(ctx, task)
|
||||
if err != nil {
|
||||
return s.handleExecutionError(task.ID, task.Command, start, err)
|
||||
}
|
||||
|
||||
// 执行命令
|
||||
var stdout, stderr bytes.Buffer
|
||||
cmd.Stdout = &stdout
|
||||
cmd.Stderr = &stderr
|
||||
|
||||
execErr := cmd.Run()
|
||||
end := time.Now()
|
||||
|
||||
// 构建结果
|
||||
result := &TaskExecutionResult{
|
||||
TaskID: task.ID,
|
||||
AgentID: nil,
|
||||
Command: task.Command,
|
||||
Output: stdout.String(),
|
||||
Start: start,
|
||||
End: end,
|
||||
Duration: end.Sub(start).Milliseconds(),
|
||||
}
|
||||
|
||||
if execErr != nil {
|
||||
result.Status = "failed"
|
||||
result.Output += "\n[ERROR]\n" + stderr.String() + "\n" + execErr.Error()
|
||||
if exitErr, ok := execErr.(*exec.ExitError); ok {
|
||||
result.ExitCode = exitErr.ExitCode()
|
||||
} else {
|
||||
result.ExitCode = 1
|
||||
}
|
||||
} else {
|
||||
result.Status = "success"
|
||||
result.ExitCode = 0
|
||||
}
|
||||
|
||||
// 处理执行结果
|
||||
return s.processExecutionResult(result)
|
||||
}
|
||||
|
||||
// executeRemote 远程执行任务(通过 Agent)
|
||||
func (s *TaskExecutionService) executeRemote(req *TaskExecutionRequest) error {
|
||||
task := req.Task
|
||||
agentID := *task.AgentID
|
||||
|
||||
logger.Infof("[TaskExecution] 远程执行任务 #%d: %s (Agent #%d)", task.ID, task.Name, agentID)
|
||||
|
||||
// 检查 Agent 是否在线
|
||||
var agent models.Agent
|
||||
if err := database.DB.First(&agent, agentID).Error; err != nil {
|
||||
return fmt.Errorf("Agent #%d 不存在", agentID)
|
||||
}
|
||||
|
||||
if !agent.Enabled {
|
||||
return fmt.Errorf("Agent #%d 已禁用", agentID)
|
||||
}
|
||||
|
||||
// 通过 WebSocket 发送立即执行命令给 Agent
|
||||
if s.agentWSManager == nil {
|
||||
return fmt.Errorf("AgentWSManager 未初始化")
|
||||
}
|
||||
err := s.agentWSManager.SendToAgent(agentID, "execute", map[string]interface{}{
|
||||
"task_id": task.ID,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("发送执行命令失败: %v", err)
|
||||
}
|
||||
|
||||
logger.Infof("[TaskExecution] 已发送立即执行命令给 Agent #%d,任务 #%d", agentID, task.ID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// prepareCommand 准备执行命令
|
||||
func (s *TaskExecutionService) prepareCommand(ctx context.Context, task *models.Task) (*exec.Cmd, error) {
|
||||
command := task.Command
|
||||
|
||||
// 处理工作目录
|
||||
if task.WorkDir != "" {
|
||||
// 验证工作目录
|
||||
if _, err := os.Stat(task.WorkDir); err != nil {
|
||||
return nil, fmt.Errorf("工作目录不存在或无法访问: %s", task.WorkDir)
|
||||
}
|
||||
}
|
||||
|
||||
// 处理环境变量
|
||||
envVars := s.loadEnvVars(task.Envs)
|
||||
|
||||
// 根据操作系统创建命令
|
||||
var cmd *exec.Cmd
|
||||
if runtime.GOOS == "windows" {
|
||||
cmd = exec.CommandContext(ctx, "cmd", "/c", command)
|
||||
} else {
|
||||
// 如果有工作目录,在命令前加 cd
|
||||
if task.WorkDir != "" {
|
||||
command = fmt.Sprintf("cd %s && %s", task.WorkDir, command)
|
||||
}
|
||||
cmd = exec.CommandContext(ctx, "sh", "-c", command)
|
||||
}
|
||||
|
||||
// 设置环境变量
|
||||
if len(envVars) > 0 {
|
||||
cmd.Env = append(os.Environ(), envVars...)
|
||||
}
|
||||
|
||||
return cmd, nil
|
||||
}
|
||||
|
||||
// createContext 创建带超时的上下文
|
||||
func (s *TaskExecutionService) createContext(timeout int) (context.Context, context.CancelFunc) {
|
||||
if timeout <= 0 {
|
||||
timeout = 30 // 默认 30 分钟
|
||||
}
|
||||
return context.WithTimeout(context.Background(), time.Duration(timeout)*time.Minute)
|
||||
}
|
||||
|
||||
// loadEnvVars 加载环境变量
|
||||
func (s *TaskExecutionService) loadEnvVars(envIDs string) []string {
|
||||
if envIDs == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
var envVars []models.EnvironmentVariable
|
||||
ids := strings.Split(envIDs, ",")
|
||||
database.DB.Where("id IN ?", ids).Find(&envVars)
|
||||
|
||||
result := make([]string, 0, len(envVars))
|
||||
for _, env := range envVars {
|
||||
result = append(result, fmt.Sprintf("%s=%s", env.Name, env.Value))
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// handleExecutionError 处理执行错误
|
||||
func (s *TaskExecutionService) handleExecutionError(taskID uint, command string, start time.Time, err error) error {
|
||||
end := time.Now()
|
||||
result := &TaskExecutionResult{
|
||||
TaskID: taskID,
|
||||
Command: command,
|
||||
Output: fmt.Sprintf("[ERROR] 任务执行失败: %v", err),
|
||||
Status: "failed",
|
||||
Duration: end.Sub(start).Milliseconds(),
|
||||
ExitCode: 1,
|
||||
Start: start,
|
||||
End: end,
|
||||
}
|
||||
return s.processExecutionResult(result)
|
||||
}
|
||||
|
||||
// processExecutionResult 处理执行结果(统一的结果处理)
|
||||
func (s *TaskExecutionService) processExecutionResult(result *TaskExecutionResult) error {
|
||||
// 创建任务日志
|
||||
taskLog, err := s.taskLogService.CreateTaskLogFromLocalExecution(
|
||||
result.TaskID,
|
||||
result.Command,
|
||||
result.Output,
|
||||
result.Status,
|
||||
result.Duration,
|
||||
result.ExitCode,
|
||||
result.Start,
|
||||
result.End,
|
||||
)
|
||||
if err != nil {
|
||||
logger.Errorf("[TaskExecution] 创建任务日志失败: %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
// 如果是 Agent 执行的,设置 AgentID
|
||||
if result.AgentID != nil {
|
||||
taskLog.AgentID = result.AgentID
|
||||
}
|
||||
|
||||
// 处理任务完成(保存日志、更新统计、清理旧日志)
|
||||
if err := s.taskLogService.ProcessTaskCompletion(taskLog); err != nil {
|
||||
logger.Errorf("[TaskExecution] 处理任务完成失败: %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
logger.Infof("[TaskExecution] 任务 #%d 执行完成 (%s)", result.TaskID, result.Status)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ProcessAgentResult 处理 Agent 上报的结果(统一入口)
|
||||
func (s *TaskExecutionService) ProcessAgentResult(agentResult *models.AgentTaskResult) error {
|
||||
logger.Infof("[TaskExecution] 处理 Agent #%d 上报的任务 #%d 结果", agentResult.AgentID, agentResult.TaskID)
|
||||
|
||||
// 转换为统一的执行结果
|
||||
result := &TaskExecutionResult{
|
||||
TaskID: agentResult.TaskID,
|
||||
AgentID: &agentResult.AgentID,
|
||||
Command: agentResult.Command,
|
||||
Output: agentResult.Output,
|
||||
Status: agentResult.Status,
|
||||
Duration: agentResult.Duration,
|
||||
ExitCode: agentResult.ExitCode,
|
||||
Start: time.Unix(agentResult.StartTime, 0),
|
||||
End: time.Unix(agentResult.EndTime, 0),
|
||||
}
|
||||
|
||||
// 使用统一的结果处理流程
|
||||
return s.processExecutionResult(result)
|
||||
}
|
||||
|
||||
// GetScriptPath 获取脚本路径
|
||||
func (s *TaskExecutionService) GetScriptPath(scriptName string) string {
|
||||
return filepath.Join("data", "scripts", scriptName)
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/logger"
|
||||
"baihu/internal/models"
|
||||
"baihu/internal/utils"
|
||||
"encoding/json"
|
||||
"time"
|
||||
)
|
||||
|
||||
// SendStatsService 接口定义(避免循环依赖)
|
||||
type SendStatsService interface {
|
||||
IncrementStats(taskID uint, 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"` // 保留天数或条数
|
||||
}
|
||||
|
||||
// SaveTaskLog 保存任务日志(通用方法)
|
||||
func (s *TaskLogService) SaveTaskLog(taskLog *models.TaskLog) error {
|
||||
if err := database.DB.Create(taskLog).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 更新任务的 last_run
|
||||
database.DB.Model(&models.Task{}).Where("id = ?", taskLog.TaskID).Update("last_run", time.Now())
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateTaskStats 更新任务统计
|
||||
func (s *TaskLogService) UpdateTaskStats(taskID uint, 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 uint) {
|
||||
var task models.Task
|
||||
if err := database.DB.First(&task, taskID).Error; err != nil {
|
||||
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 := 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
|
||||
err := database.DB.Where("task_id = ?", taskID).Order("id DESC").Offset(config.Keep - 1).Limit(1).First(&boundaryLog).Error
|
||||
if err == nil {
|
||||
result := database.DB.Where("task_id = ? AND id < ?", taskID, boundaryLog.ID).Delete(&models.TaskLog{})
|
||||
deleted = result.RowsAffected
|
||||
}
|
||||
}
|
||||
|
||||
if deleted > 0 {
|
||||
logger.Infof("[TaskLog] 清理任务 #%d 的 %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) {
|
||||
// 压缩输出
|
||||
compressed, err := utils.CompressToBase64(result.Output)
|
||||
if err != nil {
|
||||
logger.Errorf("[TaskLog] 压缩日志失败: %v", err)
|
||||
compressed = ""
|
||||
}
|
||||
|
||||
taskLog := &models.TaskLog{
|
||||
TaskID: result.TaskID,
|
||||
AgentID: &result.AgentID,
|
||||
Command: result.Command,
|
||||
Output: compressed,
|
||||
Status: result.Status,
|
||||
Duration: result.Duration,
|
||||
ExitCode: result.ExitCode,
|
||||
}
|
||||
|
||||
// 处理开始和结束时间
|
||||
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 uint, command, output, status string, duration int64, exitCode int, start, end time.Time) (*models.TaskLog, error) {
|
||||
// 压缩输出
|
||||
compressed, err := utils.CompressToBase64(output)
|
||||
if err != nil {
|
||||
logger.Errorf("[TaskLog] 压缩日志失败: %v", err)
|
||||
compressed = ""
|
||||
}
|
||||
|
||||
startTime := models.LocalTime(start)
|
||||
endTime := models.LocalTime(end)
|
||||
|
||||
taskLog := &models.TaskLog{
|
||||
TaskID: taskID,
|
||||
Command: command,
|
||||
Output: compressed,
|
||||
Status: status,
|
||||
Duration: duration,
|
||||
ExitCode: exitCode,
|
||||
StartTime: &startTime,
|
||||
EndTime: &endTime,
|
||||
}
|
||||
|
||||
return taskLog, nil
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/models"
|
||||
)
|
||||
|
||||
type TaskService struct{}
|
||||
|
||||
func NewTaskService() *TaskService {
|
||||
return &TaskService{}
|
||||
}
|
||||
|
||||
func (ts *TaskService) CreateTask(name, command, schedule string, timeout int, workDir, cleanConfig, envs, taskType, config string, agentID *uint) *models.Task {
|
||||
if taskType == "" {
|
||||
taskType = "task"
|
||||
}
|
||||
task := &models.Task{
|
||||
Name: name,
|
||||
Command: command,
|
||||
Type: taskType,
|
||||
Config: config,
|
||||
Schedule: schedule,
|
||||
Timeout: timeout,
|
||||
WorkDir: workDir,
|
||||
CleanConfig: cleanConfig,
|
||||
Envs: envs,
|
||||
AgentID: agentID,
|
||||
Enabled: true,
|
||||
}
|
||||
database.DB.Create(task)
|
||||
return task
|
||||
}
|
||||
|
||||
func (ts *TaskService) GetTasks() []models.Task {
|
||||
var tasks []models.Task
|
||||
database.DB.Find(&tasks)
|
||||
return tasks
|
||||
}
|
||||
|
||||
// GetTasksWithPagination 分页获取任务列表
|
||||
func (ts *TaskService) GetTasksWithPagination(page, pageSize int, name string, agentID *uint) ([]models.Task, int64) {
|
||||
var tasks []models.Task
|
||||
var total int64
|
||||
|
||||
query := database.DB.Model(&models.Task{})
|
||||
if name != "" {
|
||||
query = query.Where("name LIKE ?", "%"+name+"%")
|
||||
}
|
||||
if agentID != nil {
|
||||
query = query.Where("agent_id = ?", *agentID)
|
||||
}
|
||||
|
||||
query.Count(&total)
|
||||
query.Order("id DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks)
|
||||
|
||||
return tasks, total
|
||||
}
|
||||
|
||||
func (ts *TaskService) GetTaskByID(id int) *models.Task {
|
||||
var task models.Task
|
||||
if err := database.DB.First(&task, id).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &task
|
||||
}
|
||||
|
||||
func (ts *TaskService) UpdateTask(id int, name, command, schedule string, timeout int, workDir, cleanConfig, envs string, enabled bool, taskType, config string, agentID *uint) *models.Task {
|
||||
var task models.Task
|
||||
if err := database.DB.First(&task, id).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
task.Name = name
|
||||
task.Command = command
|
||||
task.Schedule = schedule
|
||||
task.Timeout = timeout
|
||||
task.WorkDir = workDir
|
||||
task.CleanConfig = cleanConfig
|
||||
task.Envs = envs
|
||||
task.Enabled = enabled
|
||||
task.AgentID = agentID
|
||||
if taskType != "" {
|
||||
task.Type = taskType
|
||||
}
|
||||
task.Config = config
|
||||
database.DB.Save(&task)
|
||||
return &task
|
||||
}
|
||||
|
||||
func (ts *TaskService) DeleteTask(id int) bool {
|
||||
result := database.DB.Delete(&models.Task{}, id)
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
Reference in New Issue
Block a user