feat: refact scheduler

This commit is contained in:
engigu
2026-02-07 21:44:09 +08:00
parent cdf3b3dbdc
commit f746c871fa
37 changed files with 2888 additions and 1316 deletions
+15
View File
@@ -48,6 +48,21 @@ const (
KeyWorkerCount = "worker_count"
KeyQueueSize = "queue_size"
KeyRateInterval = "rate_interval"
// WebSocket 消息类型
WSTypeHeartbeat = "heartbeat"
WSTypeHeartbeatAck = "heartbeat_ack"
WSTypeTasks = "tasks"
WSTypeTaskResult = "task_result"
WSTypeTaskLog = "task_log"
WSTypeExecute = "execute"
WSTypeUpdate = "update"
WSTypeDisconnect = "disconnect"
WSTypeConnected = "connected"
WSTypeDisabled = "disabled"
WSTypeEnabled = "enabled"
WSTypeFetchTasks = "fetch_tasks"
WSTypeTaskHeartbeat = "task_heartbeat"
)
// TablePrefix 表前缀,从配置文件读取
+82 -13
View File
@@ -1,16 +1,19 @@
package controllers
import (
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"github.com/engigu/baihu-panel/internal/services"
"github.com/engigu/baihu-panel/internal/utils"
"encoding/json"
"net/http"
"strconv"
"strings"
"time"
"github.com/engigu/baihu-panel/internal/constant"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"github.com/engigu/baihu-panel/internal/services"
"github.com/engigu/baihu-panel/internal/services/tasks"
"github.com/engigu/baihu-panel/internal/utils"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
@@ -23,15 +26,17 @@ var agentUpgrader = websocket.Upgrader{
// AgentController Agent 控制器
type AgentController struct {
agentService *services.AgentService
wsManager *services.AgentWSManager
agentService *services.AgentService
wsManager *services.AgentWSManager
settingsService *services.SettingsService
}
// NewAgentController 创建 Agent 控制器
func NewAgentController() *AgentController {
func NewAgentController(settingsService *services.SettingsService) *AgentController {
return &AgentController{
agentService: services.NewAgentService(),
wsManager: services.GetAgentWSManager(),
agentService: services.NewAgentService(),
wsManager: services.GetAgentWSManager(),
settingsService: settingsService,
}
}
@@ -209,7 +214,7 @@ func (c *AgentController) GetTasks(ctx *gin.Context) {
// 先尝试通过 token 查找 Agent
agent := c.agentService.GetByToken(token)
// 如果找不到,尝试验证令牌并通过 machine_id 查找
if agent == nil {
machineID := ctx.GetHeader("X-Machine-ID")
@@ -332,7 +337,6 @@ func (c *AgentController) ForceUpdate(ctx *gin.Context) {
utils.SuccessMsg(ctx, "已标记强制更新,Agent 下次心跳时将自动更新")
}
// ========== WebSocket ==========
// WSConnect Agent WebSocket 连接
@@ -411,15 +415,26 @@ func (c *AgentController) WSConnect(ctx *gin.Context) {
// 更新 Agent 状态
c.agentService.Heartbeat(token, ip, "", "", "", "", "")
// 发送连接成功消息(包含注册状态)
// 获取调度配置
workerCount := getIntSetting(c.settingsService, constant.SectionScheduler, constant.KeyWorkerCount, 4)
queueSize := getIntSetting(c.settingsService, constant.SectionScheduler, constant.KeyQueueSize, 100)
rateInterval := getIntSetting(c.settingsService, constant.SectionScheduler, constant.KeyRateInterval, 200)
// 发送连接成功消息(包含注册状态和调度配置)
c.wsManager.SendToAgent(agent.ID, services.WSTypeConnected, map[string]interface{}{
"agent_id": agent.ID,
"name": agent.Name,
"is_new_agent": isNewAgent,
"machine_id": machineID,
"scheduler_config": map[string]interface{}{
"worker_count": workerCount,
"queue_size": queueSize,
"rate_interval": rateInterval,
},
})
logger.Infof("[AgentWS] Agent #%d 连接成功", agent.ID)
logger.Infof("[AgentWS] Agent #%d 连接成功 (配置: workers=%d, queue=%d, rate=%d)",
agent.ID, workerCount, queueSize, rateInterval)
// 启动读写协程
go c.wsWritePump(ac)
@@ -513,8 +528,30 @@ func (c *AgentController) handleWSMessage(ac *services.AgentConnection, agent *m
case services.WSTypeTaskResult:
c.handleTaskResult(agent, msg.Data)
case services.WSTypeTaskLog:
c.handleTaskLog(agent, msg.Data)
case services.WSTypeFetchTasks:
c.handleFetchTasks(agent)
case services.WSTypeTaskHeartbeat: // 任务心跳
c.handleTaskHeartbeat(agent, msg.Data)
}
}
// handleTaskHeartbeat 处理任务心跳
func (c *AgentController) handleTaskHeartbeat(agent *models.Agent, data json.RawMessage) {
var req struct {
LogID uint `json:"log_id"`
Duration int64 `json:"duration"`
}
if err := json.Unmarshal(data, &req); err != nil {
logger.Errorf("[AgentWS] 解析心跳消息失败: %v", err)
return
}
if req.LogID > 0 {
logger.Infof("[AgentWS] 收到任务心跳: LogID=%d, Duration=%dms", req.LogID, req.Duration)
c.agentService.UpdateTaskDuration(req.LogID, req.Duration)
}
}
@@ -575,6 +612,25 @@ func (c *AgentController) handleTaskResult(agent *models.Agent, data json.RawMes
c.agentService.ReportResult(&result)
}
// handleTaskLog 处理 Agent 发送的实时日志
func (c *AgentController) handleTaskLog(agent *models.Agent, data json.RawMessage) {
var logMsg struct {
LogID uint `json:"log_id"`
Content string `json:"content"`
}
if err := json.Unmarshal(data, &logMsg); err != nil {
logger.Errorf("[AgentWS] 解析日志消息失败: %v", err)
return
}
tl := tasks.GetActiveLog(logMsg.LogID)
if tl != nil {
tl.Write([]byte(logMsg.Content))
} else {
logger.Warnf("[AgentWS] 收到任务日志但未找到活跃 TinyLog: LogID=%d, ContentSize=%d", logMsg.LogID, len(logMsg.Content))
}
}
// NotifyTaskUpdate 通知 Agent 任务更新
func (c *AgentController) NotifyTaskUpdate(agentID uint) {
c.wsManager.BroadcastTasks(agentID)
@@ -635,3 +691,16 @@ func (c *AgentController) DeleteToken(ctx *gin.Context) {
utils.SuccessMsg(ctx, "删除成功")
}
// getIntSetting 辅助方法
func getIntSetting(s *services.SettingsService, section, key string, defaultVal int) int {
val := s.Get(section, key)
if val == "" {
return defaultVal
}
if result, err := strconv.Atoi(val); err == nil {
return result
}
return defaultVal
}
+4 -6
View File
@@ -14,13 +14,11 @@ import (
)
type DashboardController struct {
cronService *tasks.CronService
executorService *tasks.ExecutorService
}
func NewDashboardController(cronService *tasks.CronService, executorService *tasks.ExecutorService) *DashboardController {
func NewDashboardController(executorService *tasks.ExecutorService) *DashboardController {
return &DashboardController{
cronService: cronService,
executorService: executorService,
}
}
@@ -47,14 +45,14 @@ func (dc *DashboardController) GetStats(c *gin.Context) {
// 调度统计:本地调度 + Agent 调度
// 本地调度:agent_id 为 NULL 且 enabled = true 的任务
localScheduled := dc.cronService.GetScheduledCount()
localScheduled := dc.executorService.GetScheduledCount()
// Agent 调度:agent_id 不为 NULL 且 enabled = true 的任务
var agentScheduled int64
database.DB.Model(&models.Task{}).
Where("agent_id IS NOT NULL AND enabled = ?", true).
Count(&agentScheduled)
totalScheduled := localScheduled + int(agentScheduled)
// 正在运行:目前只能统计本地运行的任务
+97
View File
@@ -0,0 +1,97 @@
package controllers
import (
"fmt"
"strconv"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/models"
"github.com/engigu/baihu-panel/internal/services/tasks"
"github.com/engigu/baihu-panel/internal/utils"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
type LogWSController struct{}
func NewLogWSController() *LogWSController {
return &LogWSController{}
}
func (lc *LogWSController) StreamLog(c *gin.Context) {
logIDStr := c.Query("log_id")
if logIDStr == "" {
return
}
logID, err := strconv.ParseUint(logIDStr, 10, 32)
if err != nil {
return
}
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
return
}
defer conn.Close()
// 1. 检查数据库中是否已结束
var taskLog models.TaskLog
if err := database.DB.First(&taskLog, uint(logID)).Error; err == nil {
if taskLog.Status != "running" {
// 已结束,读取库内日志
content, err := utils.DecompressFromBase64(taskLog.Output)
if err != nil {
conn.WriteMessage(websocket.TextMessage, []byte("解压日志失败: "+err.Error()))
return
}
conn.WriteMessage(websocket.TextMessage, []byte(content))
return
}
}
// 2. 未结束或未找到记录,尝试从 TinyLogManager 获取
tl := tasks.GetActiveLog(uint(logID))
if tl == nil {
conn.WriteMessage(websocket.TextMessage, []byte("未找到正在运行的任务日志"))
return
}
// 发送系统提示
conn.WriteMessage(websocket.TextMessage, []byte(fmt.Sprintf("[System] 连接成功,正在监听日志... (LogID: %d)\n", logID)))
// 发送最后 100 行
lastLines, err := tl.ReadLastLines(100)
if err == nil && len(lastLines) > 0 {
conn.WriteMessage(websocket.TextMessage, lastLines)
}
// 订阅实时更新
sub := tl.Subscribe()
defer tl.Unsubscribe(sub)
// 推送更新
for {
select {
case data, ok := <-sub:
if !ok {
// 任务结束,尝试刷新最后一次库内完整内容
var finalLog models.TaskLog
if err := database.DB.First(&finalLog, uint(logID)).Error; err == nil {
content, _ := utils.DecompressFromBase64(finalLog.Output)
if content != "" {
conn.WriteMessage(websocket.TextMessage, []byte("\n--- 任务已结束 ---\n"))
// 这里可以选择性再推一次完整版,或直接退出
}
}
return
}
if err := conn.WriteMessage(websocket.TextMessage, data); err != nil {
return
}
case <-c.Request.Context().Done():
return
}
}
}
+16 -16
View File
@@ -13,16 +13,16 @@ import (
)
type TaskController struct {
taskService *tasks.TaskService
cronService *tasks.CronService
agentWSManager *services.AgentWSManager
taskService *tasks.TaskService
executorService *tasks.ExecutorService
agentWSManager *services.AgentWSManager
}
func NewTaskController(taskService *tasks.TaskService, cronService *tasks.CronService) *TaskController {
func NewTaskController(taskService *tasks.TaskService, executorService *tasks.ExecutorService) *TaskController {
return &TaskController{
taskService: taskService,
cronService: cronService,
agentWSManager: services.GetAgentWSManager(),
taskService: taskService,
executorService: executorService,
agentWSManager: services.GetAgentWSManager(),
}
}
@@ -74,7 +74,7 @@ func (tc *TaskController) CreateTask(c *gin.Context) {
return
}
if err := tc.cronService.ValidateCron(req.Schedule); err != nil {
if err := tc.executorService.ValidateCron(req.Schedule); err != nil {
utils.BadRequest(c, "无效的cron表达式: "+err.Error())
return
}
@@ -86,12 +86,12 @@ func (tc *TaskController) CreateTask(c *gin.Context) {
}
task := tc.taskService.CreateTask(req.Name, req.Command, req.Schedule, req.Timeout, workDir, req.CleanConfig, req.Envs, req.Type, req.Config, req.AgentID)
// 如果是 Agent 任务,通知 Agent;否则添加到本地 cron
if task.AgentID != nil && *task.AgentID > 0 {
tc.agentWSManager.BroadcastTasks(*task.AgentID)
} else {
tc.cronService.AddTask(task)
tc.executorService.AddCronTask(task)
}
utils.Success(c, task)
@@ -101,7 +101,7 @@ func (tc *TaskController) GetTasks(c *gin.Context) {
p := utils.ParsePagination(c)
name := c.DefaultQuery("name", "")
agentIDStr := c.DefaultQuery("agent_id", "")
var agentID *uint
if agentIDStr != "" {
if id, err := strconv.ParseUint(agentIDStr, 10, 32); err == nil {
@@ -164,7 +164,7 @@ func (tc *TaskController) UpdateTask(c *gin.Context) {
}
if req.Schedule != "" {
if err := tc.cronService.ValidateCron(req.Schedule); err != nil {
if err := tc.executorService.ValidateCron(req.Schedule); err != nil {
utils.BadRequest(c, "无效的cron表达式: "+err.Error())
return
}
@@ -185,7 +185,7 @@ func (tc *TaskController) UpdateTask(c *gin.Context) {
// 处理任务调度
if task.AgentID != nil && *task.AgentID > 0 {
// Agent 任务:从本地 cron 移除,通知 Agent
tc.cronService.RemoveTask(task.ID)
tc.executorService.RemoveCronTask(task.ID)
tc.agentWSManager.BroadcastTasks(*task.AgentID)
// 如果 agent 变更了,也通知旧 agent
if oldAgentID != nil && *oldAgentID > 0 && *oldAgentID != *task.AgentID {
@@ -194,9 +194,9 @@ func (tc *TaskController) UpdateTask(c *gin.Context) {
} else {
// 本地任务
if task.Enabled {
tc.cronService.AddTask(task)
tc.executorService.AddCronTask(task)
} else {
tc.cronService.RemoveTask(task.ID)
tc.executorService.RemoveCronTask(task.ID)
}
// 如果之前是 agent 任务,通知旧 agent 移除
if oldAgentID != nil && *oldAgentID > 0 {
@@ -221,7 +221,7 @@ func (tc *TaskController) DeleteTask(c *gin.Context) {
agentID = task.AgentID
}
tc.cronService.RemoveTask(uint(id))
tc.executorService.RemoveCronTask(uint(id))
success := tc.taskService.DeleteTask(id)
if !success {
+169
View File
@@ -0,0 +1,169 @@
package executor
import (
"sync"
"time"
"github.com/robfig/cron/v3"
)
// 东八区时区(默认)
var defaultLocation = time.FixedZone("CST", 8*3600)
// CronManager 统一的任务调度管理器
type CronManager struct {
cron *cron.Cron
scheduler *Scheduler
entryMap map[string]cron.EntryID // task ID -> cron entry ID
mu sync.RWMutex
logger SchedulerLogger
}
// NewCronManager 创建一个新的计划任务管理器
func NewCronManager(scheduler *Scheduler) *CronManager {
// 使用秒级精度的 cron parser
c := cron.New(cron.WithSeconds(), cron.WithLocation(defaultLocation))
m := &CronManager{
cron: c,
scheduler: scheduler,
entryMap: make(map[string]cron.EntryID),
logger: &DefaultLogger{},
}
if scheduler != nil && scheduler.logger != nil {
m.logger = scheduler.logger
}
return m
}
// SetLogger 设置自定义日志实现
func (m *CronManager) SetLogger(logger SchedulerLogger) {
m.mu.Lock()
defer m.mu.Unlock()
m.logger = logger
}
// Start 启动调度器
func (m *CronManager) Start() {
m.cron.Start()
m.logger.Infof("[CronManager] 调度管理服务已启动")
}
// Stop 停止调度器
func (m *CronManager) Stop() {
ctx := m.cron.Stop()
<-ctx.Done()
m.logger.Infof("[CronManager] 调度管理服务已停止")
}
// AddTask 添加或更新计划任务
func (m *CronManager) AddTask(task CronTask) error {
m.mu.Lock()
defer m.mu.Unlock()
taskID := task.GetID()
// 如果已存在,先移除旧的
if entryID, exists := m.entryMap[taskID]; exists {
m.cron.Remove(entryID)
delete(m.entryMap, taskID)
}
// 准备任务执行函数
cmd := task.GetCommand()
name := task.GetName()
timeout := task.GetTimeout()
entryID, err := m.cron.AddFunc(task.GetSchedule(), func() {
m.logger.Infof("[CronManager] 触发计划任务 #%s (%s)", taskID, name)
req := &ExecutionRequest{
TaskID: taskID,
Name: name,
Command: cmd,
Type: TaskTypeCron,
Timeout: timeout,
}
// 如果有关联的 Scheduler,加入队列执行
if m.scheduler != nil {
m.scheduler.EnqueueOrExecute(req)
}
// 触发下次运行时间更新事件
m.triggerNextRunEvent(taskID, req)
})
if err != nil {
m.logger.Errorf("[CronManager] 添加任务失败 #%s: %v", taskID, err)
return err
}
m.entryMap[taskID] = entryID
m.logger.Infof("[CronManager] 任务已调度 #%s %s (%s)", taskID, name, task.GetSchedule())
// 初始触发一次下次运行时间通知
go func() {
req := &ExecutionRequest{TaskID: taskID, Name: name, Type: TaskTypeCron}
m.triggerNextRunEvent(taskID, req)
}()
return nil
}
// RemoveTask 移除计划任务
func (m *CronManager) RemoveTask(taskID string) {
m.mu.Lock()
defer m.mu.Unlock()
if entryID, exists := m.entryMap[taskID]; exists {
m.cron.Remove(entryID)
delete(m.entryMap, taskID)
m.logger.Infof("[CronManager] 任务已移除 #%s", taskID)
}
}
// triggerNextRunEvent 触发下次运行时间更新事件
func (m *CronManager) triggerNextRunEvent(taskID string, req *ExecutionRequest) {
m.mu.RLock()
entryID, exists := m.entryMap[taskID]
m.mu.RUnlock()
if !exists {
return
}
entry := m.cron.Entry(entryID)
if !entry.Next.IsZero() && m.scheduler != nil && m.scheduler.handler != nil {
m.scheduler.handler.OnCronNextRun(req, entry.Next)
}
}
// ValidateCron 校验 Cron 表达式
func (m *CronManager) 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
}
// GetEntry 获取任务详情
func (m *CronManager) GetEntry(taskID string) (cron.Entry, bool) {
m.mu.RLock()
defer m.mu.RUnlock()
entryID, exists := m.entryMap[taskID]
if !exists {
return cron.Entry{}, false
}
return m.cron.Entry(entryID), true
}
// GetScheduledCount 获取已调度任务总数
func (m *CronManager) GetScheduledCount() int {
m.mu.RLock()
defer m.mu.RUnlock()
return len(m.entryMap)
}
+204
View File
@@ -0,0 +1,204 @@
package executor
import (
"context"
"io"
"os"
"os/exec"
"strings"
"time"
"github.com/engigu/baihu-panel/internal/utils"
)
// Task 任务基础接口
type Task interface {
GetID() string
GetName() string
GetCommand() string
GetTimeout() int
}
// CronTask 计划任务接口
type CronTask interface {
Task
GetSchedule() string
}
// Request 任务执行请求
type Request struct {
Command string
WorkDir string
Envs []string
Timeout int // 分钟
}
// Result 任务执行结果
type Result struct {
Output string
Status string // success, failed
Duration int64 // 毫秒
ExitCode int
StartTime time.Time
EndTime time.Time
}
// Hooks 执行钩子接口
type Hooks interface {
// PreExecute 执行前钩子,返回日志ID和错误
PreExecute(ctx context.Context, req Request) (logID uint, err error)
// PostExecute 执行后钩子,处理日志压缩和记录更新
PostExecute(ctx context.Context, logID uint, result *Result) error
// OnHeartbeat 执行中心跳钩子,用于更新实时状态
OnHeartbeat(ctx context.Context, logID uint, duration int64) error
}
// Execute 执行命令(基础版本,不带钩子)
func Execute(ctx context.Context, req Request, stdout, stderr io.Writer) (*Result, error) {
return ExecuteWithHooks(ctx, req, stdout, stderr, nil)
}
// ExecuteWithHooks 执行命令(带钩子支持)
func ExecuteWithHooks(ctx context.Context, req Request, stdout, stderr io.Writer, hooks Hooks) (*Result, error) {
start := time.Now()
// 1. 执行前钩子
var logID uint
if hooks != nil {
id, err := hooks.PreExecute(ctx, req)
if err != nil {
return &Result{
Status: "failed",
Duration: 0,
ExitCode: 1,
StartTime: start,
EndTime: time.Now(),
}, err
}
logID = id
}
// 2. 执行命令
timeout := req.Timeout
if timeout <= 0 {
timeout = 30
}
execCtx, cancel := context.WithTimeout(ctx, time.Duration(timeout)*time.Minute)
defer cancel()
finalCommand := req.Command
shell, args := utils.GetShellCommand(finalCommand)
cmd := exec.CommandContext(execCtx, shell, args...)
// 设置工作目录
// 设置工作目录
workDir := strings.TrimSpace(req.WorkDir)
if workDir != "" {
cmd.Dir = workDir
}
// 设置环境变量(始终继承系统环境变量)
cmd.Env = os.Environ()
if len(req.Envs) > 0 {
cmd.Env = append(cmd.Env, req.Envs...)
}
cmd.Stdout = stdout
cmd.Stderr = stderr
// 使用 cmd.Start() + Wait() 以便在后台处理心跳
err := cmd.Start()
if err != nil {
// Start 失败的处理
end := time.Now()
result := &Result{
Status: "failed",
Duration: end.Sub(start).Milliseconds(),
ExitCode: 1,
StartTime: start, // 修正为 start
EndTime: end,
}
// 执行后钩子
if hooks != nil {
result.Output += "\n[System Error] " + err.Error()
hooks.PostExecute(ctx, logID, result)
}
return result, err
}
// 启动心跳协程
done := make(chan struct{})
go func() {
// 每3秒一次心跳
ticker := time.NewTicker(3 * time.Second)
defer ticker.Stop()
for {
select {
case <-done:
return
case <-ticker.C:
if hooks != nil {
hooks.OnHeartbeat(ctx, logID, time.Since(start).Milliseconds())
}
}
}
}()
// 等待命令完成
err = cmd.Wait()
close(done) // 停止心跳
end := time.Now()
result := &Result{
StartTime: start,
EndTime: end,
Duration: end.Sub(start).Milliseconds(),
}
if err != nil {
result.Status = "failed"
if exitErr, ok := err.(*exec.ExitError); ok {
result.ExitCode = exitErr.ExitCode()
} else {
result.ExitCode = 1
}
} else {
result.Status = "success"
result.ExitCode = 0
}
// 3. 执行后钩子
if hooks != nil {
if hookErr := hooks.PostExecute(ctx, logID, result); hookErr != nil {
// 记录钩子错误但不影响执行结果
result.Output += "\n[Hook Error] " + hookErr.Error()
}
}
return result, err
}
// ParseEnvVars 解析环境变量字符串 "KEY1=VALUE1,KEY2=VALUE2"
func ParseEnvVars(envStr string) []string {
if envStr == "" {
return nil
}
pairs := strings.Split(envStr, ",")
result := make([]string, 0, len(pairs))
for _, pair := range pairs {
if pair == "" {
continue
}
// 解码特殊字符
pair = strings.ReplaceAll(pair, "{{COMMA}}", ",")
pair = strings.ReplaceAll(pair, "{{EQUAL}}", "=")
result = append(result, pair)
}
return result
}
+451
View File
@@ -0,0 +1,451 @@
package executor
import (
"bytes"
"context"
"fmt"
"io"
"sync"
"time"
)
// SchedulerConfig 调度器配置
type SchedulerConfig struct {
WorkerCount int // Worker 数量
QueueSize int // 队列大小
RateInterval time.Duration // 速率限制间隔
}
// TaskType 任务类型
type TaskType string
const (
TaskTypeCron TaskType = "cron" // 计划任务
TaskTypeManual TaskType = "manual" // 手动任务
TaskTypeSystem TaskType = "system" // 系统任务
)
// TaskStatus 任务状态
type TaskStatus string
const (
TaskStatusPending TaskStatus = "pending" // 等待中
TaskStatusRunning TaskStatus = "running" // 运行中
TaskStatusSuccess TaskStatus = "success" // 成功
TaskStatusFailed TaskStatus = "failed" // 失败
TaskStatusTimeout TaskStatus = "timeout" // 超时
TaskStatusCancelled TaskStatus = "cancelled" // 已取消
)
// ExecutionRequest 执行请求(标准接口)
type ExecutionRequest struct {
TaskID string // 任务 ID
LogID uint // 日志 ID
Name string // 任务名称
Type TaskType // 任务类型
Command string // 命令
WorkDir string // 工作目录
Envs []string // 环境变量
Timeout int // 超时时间(分钟)
Metadata map[string]interface{} // 额外元数据
}
// ExecutionResult 执行结果(标准接口)
type ExecutionResult struct {
TaskID string // 任务 ID
LogID uint // 日志 ID
Success bool // 是否成功
Output string // 输出内容
Error string // 错误信息
Status string // 状态: success, failed, timeout, cancelled
Duration int64 // 执行时长(毫秒)
ExitCode int // 退出码
StartTime time.Time // 开始时间
EndTime time.Time // 结束时间
}
// SchedulerEventHandler 调度器事件处理器(标准接口)
// 主服务端和 Agent 端通过实现不同的 Handler 来处理事件
type SchedulerEventHandler interface {
// OnTaskScheduled 任务被调度(加入队列)时触发
OnTaskScheduled(req *ExecutionRequest)
// OnTaskExecuting 任务准备开始执行时触发
// 返回 stdout/stderr 写入器用于实时日志推送
// 主服务端:返回 TinyLog 写入器(写入本地文件)
// Agent 端:返回 WebSocket 写入器(实时推送到主服务)
OnTaskExecuting(req *ExecutionRequest) (stdout, stderr io.Writer, err error)
// OnTaskStarted 任务实际开始运行(已经过了队列等待和速率限制)
OnTaskStarted(req *ExecutionRequest)
// OnTaskCompleted 任务执行完成时触发
// 主服务端:压缩日志、更新数据库、清理旧日志
// Agent 端:通过 WebSocket 发送执行结果到主服务
OnTaskCompleted(req *ExecutionRequest, result *ExecutionResult)
// OnTaskFailed 任务执行失败时触发
OnTaskFailed(req *ExecutionRequest, err error)
// OnCronNextRun 计划任务下次运行时间更新时触发
OnCronNextRun(req *ExecutionRequest, nextRun time.Time)
// OnTaskHeartbeat 任务执行心跳(用于更新实时耗时等)
OnTaskHeartbeat(req *ExecutionRequest, duration int64)
}
// SchedulerLogger 日志接口(允许自定义日志实现)
type SchedulerLogger interface {
Infof(format string, args ...interface{})
Warnf(format string, args ...interface{})
Errorf(format string, args ...interface{})
}
// DefaultLogger 默认日志实现(使用 fmt)
type DefaultLogger struct{}
func (l *DefaultLogger) Infof(format string, args ...interface{}) {
fmt.Printf("[INFO] "+format+"\n", args...)
}
func (l *DefaultLogger) Warnf(format string, args ...interface{}) {
fmt.Printf("[WARN] "+format+"\n", args...)
}
func (l *DefaultLogger) Errorf(format string, args ...interface{}) {
fmt.Printf("[ERROR] "+format+"\n", args...)
}
// schedulerHooksAdapter 适配器:将 executor.Hooks 映射到 SchedulerEventHandler
type schedulerHooksAdapter struct {
handler SchedulerEventHandler
req *ExecutionRequest
}
func (h *schedulerHooksAdapter) PreExecute(ctx context.Context, req Request) (uint, error) {
return h.req.LogID, nil
}
func (h *schedulerHooksAdapter) PostExecute(ctx context.Context, logID uint, result *Result) error {
return nil
}
func (h *schedulerHooksAdapter) OnHeartbeat(ctx context.Context, logID uint, duration int64) error {
if h.handler != nil {
h.handler.OnTaskHeartbeat(h.req, duration)
}
return nil
}
// TaskExecutor 定义任务执行函数签名
type TaskExecutor func(ctx context.Context, req *ExecutionRequest, stdout, stderr io.Writer) (*Result, error)
// Scheduler 统一调度器(独立组件,可在主服务和 Agent 中复用)
// 调度器本身只负责队列管理和任务调度,具体的执行逻辑和事件处理由 Handler 实现
type Scheduler struct {
config SchedulerConfig
handler SchedulerEventHandler
executor TaskExecutor
taskQueue chan *ExecutionRequest
rateLimiter <-chan time.Time
stopCh chan struct{}
wg sync.WaitGroup
mu sync.RWMutex
logger SchedulerLogger
runningTasks map[string]context.CancelFunc // 记录运行中的任务,用于停止
}
// NewScheduler 创建调度器
func NewScheduler(config SchedulerConfig, handler SchedulerEventHandler) *Scheduler {
if config.WorkerCount <= 0 {
config.WorkerCount = 4
}
if config.QueueSize <= 0 {
config.QueueSize = 100
}
if config.RateInterval <= 0 {
config.RateInterval = 200 * time.Millisecond
}
s := &Scheduler{
config: config,
handler: handler,
executor: func(ctx context.Context, req *ExecutionRequest, stdout, stderr io.Writer) (*Result, error) {
hooks := &schedulerHooksAdapter{handler: handler, req: req}
return ExecuteWithHooks(ctx, Request{
Command: req.Command,
WorkDir: req.WorkDir,
Envs: req.Envs,
Timeout: req.Timeout,
}, stdout, stderr, hooks)
},
taskQueue: make(chan *ExecutionRequest, config.QueueSize),
rateLimiter: time.Tick(config.RateInterval),
stopCh: make(chan struct{}),
logger: &DefaultLogger{},
runningTasks: make(map[string]context.CancelFunc),
}
return s
}
// SetLogger 设置自定义日志实现
func (s *Scheduler) SetLogger(logger SchedulerLogger) {
s.mu.Lock()
defer s.mu.Unlock()
s.logger = logger
}
// SetExecutor 设置任务执行器
func (s *Scheduler) SetExecutor(executor TaskExecutor) {
s.mu.Lock()
defer s.mu.Unlock()
s.executor = executor
}
// Start 启动调度器
func (s *Scheduler) Start() {
for i := 0; i < s.config.WorkerCount; i++ {
s.wg.Add(1)
go s.worker(i)
}
s.logger.Infof("[Scheduler] 已启动")
}
// Stop 停止调度器
func (s *Scheduler) Stop() {
close(s.stopCh)
s.wg.Wait()
s.logger.Infof("[Scheduler] 已停止")
}
// Enqueue 将任务加入队列
func (s *Scheduler) Enqueue(req *ExecutionRequest) error {
select {
case s.taskQueue <- req:
if s.handler != nil {
s.handler.OnTaskScheduled(req)
}
return nil
default:
// 队列满,返回错误
return fmt.Errorf("任务队列已满")
}
}
// EnqueueOrExecute 将任务加入队列,如果队列满则直接执行
func (s *Scheduler) EnqueueOrExecute(req *ExecutionRequest) {
select {
case s.taskQueue <- req:
// 成功入队
if s.handler != nil {
s.handler.OnTaskScheduled(req)
}
default:
// 队列满,直接执行(降级处理)
s.logger.Warnf("[Scheduler] 任务队列已满,直接执行任务 %s", req.TaskID)
go s.executeTask(req)
}
}
// ExecuteSync 同步执行任务(不经过队列)
func (s *Scheduler) ExecuteSync(req *ExecutionRequest) (*ExecutionResult, error) {
return s.executeTask(req)
}
// worker 工作协程
func (s *Scheduler) worker(id int) {
defer s.wg.Done()
for {
select {
case <-s.stopCh:
return
case req := <-s.taskQueue:
// 速率限制
<-s.rateLimiter
s.executeTask(req)
}
}
}
// executeTask 执行任务(本地执行)
func (s *Scheduler) executeTask(req *ExecutionRequest) (*ExecutionResult, error) {
start := time.Now()
s.logger.Infof("[Scheduler] 执行任务 %s (名称: %s, 类型: %s)", req.TaskID, req.Name, req.Type)
// 1. 执行前事件:获取 stdout/stderr 写入器
var stdout, stderr io.Writer
var err error
if s.handler != nil {
stdout, stderr, err = s.handler.OnTaskExecuting(req)
if err != nil {
s.logger.Errorf("[Scheduler] 任务 %s 执行前事件失败: %v", req.TaskID, err)
if s.handler != nil {
s.handler.OnTaskFailed(req, err)
}
return &ExecutionResult{
TaskID: req.TaskID,
Success: false,
Status: "failed",
Error: err.Error(),
Duration: 0,
ExitCode: 1,
StartTime: start,
EndTime: time.Now(),
}, err
}
}
// 2. 准备输出缓冲区
var stdoutBuf, stderrBuf bytes.Buffer
var stdoutWriter, stderrWriter io.Writer
if stdout != nil {
stdoutWriter = io.MultiWriter(&stdoutBuf, stdout)
} else {
stdoutWriter = &stdoutBuf
}
if stderr != nil {
stderrWriter = io.MultiWriter(&stderrBuf, stderr)
} else {
stderrWriter = &stderrBuf
}
// 3. 实际开始执行事件 (经过队列和速率限制之后)
if s.handler != nil {
s.handler.OnTaskStarted(req)
}
// 4. 执行命令(使用 executor.Execute
// 创建带取消功能的上下文
ctx, cancel := context.WithCancel(context.Background())
if req.Timeout > 0 {
ctx, cancel = context.WithTimeout(ctx, time.Duration(req.Timeout)*time.Minute)
}
defer cancel()
// 注册到运行中任务
s.mu.Lock()
s.runningTasks[req.TaskID] = cancel
s.mu.Unlock()
defer func() {
s.mu.Lock()
delete(s.runningTasks, req.TaskID)
s.mu.Unlock()
}()
execResult, execErr := s.executor(ctx, req, stdoutWriter, stderrWriter)
// 5. 构建结果
result := &ExecutionResult{
TaskID: req.TaskID,
LogID: req.LogID, // 传递 LogID
Success: execResult.Status == "success",
Output: stdoutBuf.String(),
Status: execResult.Status,
Duration: execResult.Duration,
ExitCode: execResult.ExitCode,
StartTime: execResult.StartTime,
EndTime: execResult.EndTime,
}
if execErr != nil {
result.Error = execErr.Error()
errOutput := stderrBuf.String()
if errOutput != "" {
result.Output += "\n[ERROR]\n" + errOutput
}
if ctx.Err() == context.Canceled {
result.Status = "cancelled"
} else if ctx.Err() == context.DeadlineExceeded {
result.Status = "timeout"
}
}
// 6. 执行后事件
if s.handler != nil {
if execErr != nil {
s.handler.OnTaskFailed(req, execErr)
} else {
s.handler.OnTaskCompleted(req, result)
}
}
if execErr != nil {
s.logger.Errorf("[Scheduler] 任务 %s 执行失败: %v", req.TaskID, execErr)
} else {
s.logger.Infof("[Scheduler] 任务 %s 执行完成 (状态: %s, 耗时: %dms)",
req.TaskID, result.Status, result.Duration)
}
return result, execErr
}
// StopTask 停止正在运行的任务
func (s *Scheduler) StopTask(taskID string) bool {
s.mu.RLock()
cancel, exists := s.runningTasks[taskID]
s.mu.RUnlock()
if exists && cancel != nil {
cancel()
s.logger.Infof("[Scheduler] 已尝试停止任务 %s", taskID)
return true
}
return false
}
// GetRunningTaskCount 获取正在运行的任务数量
func (s *Scheduler) GetRunningTaskCount() int {
s.mu.RLock()
defer s.mu.RUnlock()
return len(s.runningTasks)
}
// GetRunningTasks 获取所有正在运行的任务 ID
func (s *Scheduler) GetRunningTasks() []string {
s.mu.RLock()
defer s.mu.RUnlock()
ids := make([]string, 0, len(s.runningTasks))
for id := range s.runningTasks {
ids = append(ids, id)
}
return ids
}
// Reload 重新加载配置
func (s *Scheduler) Reload(config SchedulerConfig) {
s.logger.Infof("[Scheduler] 正在重载配置...")
// 停止现有 workers
close(s.stopCh)
s.wg.Wait()
// 更新配置
s.mu.Lock()
s.config = config
s.taskQueue = make(chan *ExecutionRequest, config.QueueSize)
s.rateLimiter = time.Tick(config.RateInterval)
s.stopCh = make(chan struct{})
s.mu.Unlock()
// 重启 workers
s.Start()
s.logger.Infof("[Scheduler] 配置已重载: workers=%d, queue=%d, rate=%v",
config.WorkerCount, config.QueueSize, config.RateInterval)
}
// GetQueueSize 获取当前队列大小
func (s *Scheduler) GetQueueSize() int {
return len(s.taskQueue)
}
// GetConfig 获取配置
func (s *Scheduler) GetConfig() SchedulerConfig {
s.mu.RLock()
defer s.mu.RUnlock()
return s.config
}
+18
View File
@@ -114,3 +114,21 @@ func WithField(key string, value interface{}) *logrus.Entry {
func WithFields(fields logrus.Fields) *logrus.Entry {
return Log.WithFields(fields)
}
// SchedulerLogger 兼容 internal/executor 的日志接口
type SchedulerLogger struct{}
func (s *SchedulerLogger) Infof(format string, args ...interface{}) {
Log.Infof(format, args...)
}
func (s *SchedulerLogger) Warnf(format string, args ...interface{}) {
Log.Warnf(format, args...)
}
func (s *SchedulerLogger) Errorf(format string, args ...interface{}) {
Log.Errorf(format, args...)
}
// NewSchedulerLogger 创建一个兼容 executor.SchedulerLogger 的实例
func NewSchedulerLogger() *SchedulerLogger {
return &SchedulerLogger{}
}
+2 -1
View File
@@ -16,7 +16,7 @@ type Agent struct {
Status string `json:"status" gorm:"size:20;default:'pending'"` // 状态: pending(待审核), online, offline, blocked(拉黑)
LastSeen *LocalTime `json:"last_seen"` // 最后心跳时间
IP string `json:"ip" gorm:"size:45"` // Agent IP 地址
Version string `json:"version" gorm:"size:20"` // Agent 版本
Version string `json:"version" gorm:"size:50"` // Agent 版本
BuildTime string `json:"build_time" gorm:"size:30"` // Agent 构建时间
Hostname string `json:"hostname" gorm:"size:100"` // Agent 主机名
OS string `json:"os" gorm:"size:20"` // 操作系统
@@ -65,6 +65,7 @@ type AgentTask struct {
// AgentTaskResult Agent 上报的任务执行结果
type AgentTaskResult struct {
TaskID uint `json:"task_id"`
LogID uint `json:"log_id"`
AgentID uint `json:"agent_id"`
Command string `json:"command"`
Output string `json:"output"`
+28
View File
@@ -1,6 +1,8 @@
package models
import (
"fmt"
"github.com/engigu/baihu-panel/internal/constant"
"gorm.io/gorm"
@@ -25,6 +27,11 @@ type RepoConfig struct {
AuthToken string `json:"auth_token"` // 认证 Token
}
// TaskConfig 任务配置 RepoConfig+TaskConfig=task.config
type TaskConfig struct {
Concurrency int `json:"$task_concurrency"` // 0: disable concurrency, 1: enable concurrency
}
// Task represents a scheduled task
type Task struct {
ID uint `json:"id" gorm:"primaryKey"`
@@ -39,6 +46,7 @@ type Task struct {
Envs string `json:"envs" gorm:"size:255;default:''"` // 环境变量ID列表,逗号分隔
AgentID *uint `json:"agent_id" gorm:"index"` // Agent ID,为空表示本地执行
Enabled bool `json:"enabled" gorm:"default:true"`
RunningGo string `json:"running_go" gorm:"type:text"` // 正在运行的 go routine id 数组 (JSON)
LastRun *LocalTime `json:"last_run"`
NextRun *LocalTime `json:"next_run"`
CreatedAt LocalTime `json:"created_at"`
@@ -50,6 +58,26 @@ func (Task) TableName() string {
return constant.TablePrefix + "tasks"
}
func (t *Task) GetID() string {
return fmt.Sprintf("%d", t.ID)
}
func (t *Task) GetName() string {
return t.Name
}
func (t *Task) GetCommand() string {
return t.Command
}
func (t *Task) GetTimeout() int {
return t.Timeout
}
func (t *Task) GetSchedule() string {
return t.Schedule
}
// TaskLog represents a log entry for task execution
type TaskLog struct {
ID uint `json:"id" gorm:"primaryKey"`
+20 -14
View File
@@ -7,7 +7,7 @@ import (
"github.com/engigu/baihu-panel/internal/services/tasks"
)
var cronService *tasks.CronService
var executorService *tasks.ExecutorService
func RegisterControllers() *Controllers {
// Initialize services
@@ -23,35 +23,41 @@ func RegisterControllers() *Controllers {
scriptService := services.NewScriptService()
sendStatsService := services.NewSendStatsService()
agentWSManager := services.GetAgentWSManager()
// 创建任务执行服务(需要依赖注入)
taskExecutionService := tasks.NewTaskExecutionService(agentWSManager, sendStatsService)
executorService := tasks.NewExecutorService(taskService, taskExecutionService, settingsService, envService)
// Initialize cron service
cronService = tasks.NewCronService(taskService, executorService)
cronService.Start()
taskLogService := tasks.NewTaskLogService(sendStatsService)
// 创建任务执行服务(需要依赖注入)
// 清理 task 运行状态的任务可以直接由 executorService 承担或在此处通过 Database 直接清理
// 简单期间,我们使用一个新方法 tasks.CleanupRunningTasks() 或者让 executorService 启动时清理
executorService = tasks.NewExecutorService(taskService, taskLogService, agentWSManager, settingsService, envService)
// 启动时清理残留的运行状态
_ = executorService.CleanupRunningTasks()
// 启动计划任务
executorService.StartCron()
// Initialize and return controllers
return &Controllers{
Task: controllers.NewTaskController(taskService, cronService),
Task: controllers.NewTaskController(taskService, executorService),
Auth: controllers.NewAuthController(userService, settingsService, loginLogService),
Env: controllers.NewEnvController(envService),
Script: controllers.NewScriptController(scriptService),
Executor: controllers.NewExecutorController(executorService),
File: controllers.NewFileController(constant.ScriptsWorkDir),
Dashboard: controllers.NewDashboardController(cronService, executorService),
Dashboard: controllers.NewDashboardController(executorService),
Log: controllers.NewLogController(),
LogWS: controllers.NewLogWSController(),
Terminal: controllers.NewTerminalController(envService),
Settings: controllers.NewSettingsController(userService, loginLogService, executorService),
Dependency: controllers.NewDependencyController(),
Agent: controllers.NewAgentController(),
Agent: controllers.NewAgentController(settingsService),
}
}
// StopCron stops the cron service gracefully
// StopCron 停止计划任务服务
func StopCron() {
if cronService != nil {
cronService.Stop()
if executorService != nil {
executorService.Stop()
}
}
+3 -1
View File
@@ -22,6 +22,7 @@ type Controllers struct {
File *controllers.FileController
Dashboard *controllers.DashboardController
Log *controllers.LogController
LogWS *controllers.LogWSController
Terminal *controllers.TerminalController
Settings *controllers.SettingsController
Dependency *controllers.DependencyController
@@ -166,6 +167,7 @@ func Setup(c *Controllers) *gin.Engine {
logs := authorized.Group("/logs")
{
logs.GET("", c.Log.GetLogs)
logs.GET("/ws", c.LogWS.StreamLog)
logs.GET("/:id", c.Log.GetLogDetail)
}
@@ -246,7 +248,7 @@ func Setup(c *Controllers) *gin.Engine {
data, err := static.ReadFile("index.html")
if err != nil {
ctx.String(500, "index.html not found")
ctx.Status(404)
return
}
+40 -11
View File
@@ -1,11 +1,6 @@
package services
import (
"github.com/engigu/baihu-panel/internal/constant"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"github.com/engigu/baihu-panel/internal/services/tasks"
"crypto/rand"
"encoding/hex"
"fmt"
@@ -14,6 +9,12 @@ import (
"strings"
"time"
"github.com/engigu/baihu-panel/internal/constant"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"github.com/engigu/baihu-panel/internal/services/tasks"
"gorm.io/gorm"
)
@@ -308,7 +309,7 @@ func (s *AgentService) GetTasks(agentID uint) []models.AgentTask {
for i, task := range tasks {
// 将环境变量 ID 转换为实际的环境变量键值对
envVarsStr := s.buildEnvVarsString(task.Envs)
result[i] = models.AgentTask{
ID: task.ID,
Name: task.Name,
@@ -353,11 +354,32 @@ func (s *AgentService) buildEnvVarsString(envIDs string) string {
func (s *AgentService) ReportResult(result *models.AgentTaskResult) error {
// 获取依赖的服务
agentWSManager := GetAgentWSManager()
// 先尝试通知正在等待的 goroutine
if agentWSManager.NotifyRemoteResult(result) {
logger.Infof("[Agent] 已通知正在等待任务 #%d 结果的 goroutine", result.TaskID)
return nil
}
// 如果没有人在等待(例如服务重启后),则由本协程负责处理结果入库
// 如果没有人在等待(例如服务重启后),则由本协程负责处理结果入库(记录日志并清理)
logger.Infof("[Agent] 没有找到等待任务 #%d 结果的 goroutine,直接处理结果", result.TaskID)
sendStatsService := NewSendStatsService()
taskExecutionService := tasks.NewTaskExecutionService(agentWSManager, sendStatsService)
// 使用统一的结果处理流程
return taskExecutionService.ProcessAgentResult(result)
taskLogService := tasks.NewTaskLogService(sendStatsService)
// 创建日志对象
taskLog, err := taskLogService.CreateTaskLogFromAgentResult(result)
if err != nil {
return err
}
// 处理完成逻辑(保存日志、更新统计、清理旧日志等)
return taskLogService.ProcessTaskCompletion(taskLog)
}
// UpdateTaskDuration 更新任务耗时(心跳)
func (s *AgentService) UpdateTaskDuration(logID uint, duration int64) error {
taskLogService := tasks.NewTaskLogService(nil)
return taskLogService.UpdateTaskDuration(logID, duration)
}
// UpdateOfflineAgents 更新离线 Agent 状态(超过 2 分钟无心跳)
@@ -368,6 +390,13 @@ func (s *AgentService) UpdateOfflineAgents() {
Update("status", "offline")
}
// ResetAllAgentsToOffline 将所有 Agents 状态重置为离线(用于服务启动时)
func (s *AgentService) ResetAllAgentsToOffline() {
database.DB.Model(&models.Agent{}).
Where("status = ?", "online").
Update("status", "offline")
}
// GetLatestVersion 获取最新 Agent 版本
func (s *AgentService) GetLatestVersion() string {
// 优先从 /opt/agent 读取(容器内)
@@ -420,7 +449,7 @@ func (s *AgentService) CheckNeedUpdate(agentVersion, agentBuildTime string) bool
// GetAvailablePlatforms 获取可用的平台列表
func (s *AgentService) GetAvailablePlatforms() []map[string]string {
platforms := []map[string]string{}
// 优先从 /opt/agent 读取(容器内)
agentDir := "/opt/agent"
files, err := os.ReadDir(agentDir)
+69 -17
View File
@@ -1,22 +1,25 @@
package services
import (
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"encoding/json"
"sync"
"time"
"github.com/engigu/baihu-panel/internal/constant"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"github.com/gorilla/websocket"
)
// AgentWSManager WebSocket 连接管理器
type AgentWSManager struct {
connections map[uint]*AgentConnection // agentID -> connection
ipConnections map[string]int // IP -> 连接数
ipLastAttempt map[string]time.Time // IP -> 最后连接尝试时间
ipFailCount map[string]int // IP -> 连续失败次数
connections map[uint]*AgentConnection // agentID -> connection
ipConnections map[string]int // IP -> 连接数
ipLastAttempt map[string]time.Time // IP -> 最后连接尝试时间
ipFailCount map[string]int // IP -> 连续失败次数
remoteWaiters map[uint]chan *models.AgentTaskResult // logID -> result channel
mu sync.RWMutex
}
@@ -47,16 +50,19 @@ type WSMessage struct {
// 消息类型常量
const (
WSTypeHeartbeat = "heartbeat"
WSTypeHeartbeatAck = "heartbeat_ack"
WSTypeTasks = "tasks"
WSTypeTaskResult = "task_result"
WSTypeUpdate = "update"
WSTypeDisconnect = "disconnect"
WSTypeConnected = "connected" // 连接成功,包含注册状态
WSTypeDisabled = "disabled" // Agent 被禁用
WSTypeEnabled = "enabled" // Agent 被启用
WSTypeFetchTasks = "fetch_tasks" // Agent 请求任务列表
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
@@ -70,6 +76,7 @@ func GetAgentWSManager() *AgentWSManager {
ipConnections: make(map[string]int),
ipLastAttempt: make(map[string]time.Time),
ipFailCount: make(map[string]int),
remoteWaiters: make(map[uint]chan *models.AgentTaskResult),
}
go agentWSManager.cleanupLoop()
})
@@ -215,6 +222,37 @@ func (m *AgentWSManager) BroadcastTasks(agentID uint) {
})
}
// RegisterRemoteWaiter 注册远程任务结果等待者
func (m *AgentWSManager) RegisterRemoteWaiter(logID uint) 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 uint) {
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()
@@ -227,6 +265,11 @@ func (m *AgentWSManager) cleanupLoop() {
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
// 启动时,先将所有 "online" 状态的 Agent 重置为 "offline"
// 因为 WebSocket 连接在应用启动时是空的,所有 Agent 客观上都是离线状态
// 等它们重新连接上来后,会变为 "online"
NewAgentService().ResetAllAgentsToOffline()
for range ticker.C {
m.mu.Lock()
now := time.Now()
@@ -248,6 +291,15 @@ func (m *AgentWSManager) cleanupLoop() {
}
}
// 定期清理数据库中的过期状态(处理服务重启或异常终止的情况)
// 有些 Agent 虽然没有连接,但数据库状态可能是 "online"
cutoff := now.Add(-2 * time.Minute)
database.DB.Model(&models.Agent{}).
Where("status = ? AND last_seen < ?", "online", cutoff).
Update("status", "offline")
// 清理过期的限流记录(超过 10 分钟未活动)
// 清理过期的限流记录(超过 10 分钟未活动)
for ip, lastAttempt := range m.ipLastAttempt {
if now.Sub(lastAttempt) > 10*time.Minute {
-168
View File
@@ -1,168 +0,0 @@
package tasks
import (
"sync"
"time"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/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)
}
+631 -221
View File
@@ -1,18 +1,33 @@
package tasks
import (
"github.com/engigu/baihu-panel/internal/constant"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/utils"
"bytes"
"context"
"encoding/json"
"fmt"
"os"
"os/exec"
"io"
"path/filepath"
"strings"
"sync"
"time"
"github.com/engigu/baihu-panel/internal/constant"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/executor"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"github.com/engigu/baihu-panel/internal/utils"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// AgentWSManager 接口定义(避免循环依赖)
type AgentWSManager interface {
RegisterRemoteWaiter(logID uint) chan *models.AgentTaskResult
UnregisterRemoteWaiter(logID uint)
SendToAgent(agentID uint, msgType string, data interface{}) error
}
// SettingsService 接口定义(避免循环依赖)
type SettingsService interface {
Get(section, key string) string
@@ -23,68 +38,316 @@ 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
// ExecutorService handles task execution and scheduling
type ExecutorService struct {
taskService *TaskService
taskExecutionService *TaskExecutionService
settingsService SettingsService
envService EnvService
results []ExecutionResult
runningTasks map[int]bool
mu sync.RWMutex
resultsMu sync.RWMutex
taskService *TaskService
taskLogService *TaskLogService
agentWSManager AgentWSManager
settingsService SettingsService
envService EnvService
scheduler *executor.Scheduler
cronManager *executor.CronManager
results []executor.ExecutionResult
mu sync.RWMutex
resultsMu sync.RWMutex
stopCh chan struct{}
}
// 任务队列和 worker pool
taskQueue chan taskJob
workerCount int
rateLimiter <-chan time.Time
stopCh chan struct{}
wg sync.WaitGroup
func (es *ExecutorService) GetScheduler() *executor.Scheduler {
return es.scheduler
}
// 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)
func NewExecutorService(
taskService *TaskService,
taskLogService *TaskLogService,
agentWSManager AgentWSManager,
settingsService SettingsService,
envService EnvService,
) *ExecutorService {
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{}),
taskService: taskService,
taskLogService: taskLogService,
agentWSManager: agentWSManager,
settingsService: settingsService,
envService: envService,
results: make([]executor.ExecutionResult, 0, 100),
stopCh: make(chan struct{}),
}
// 启动 worker pool
es.startWorkers()
// 1. 初始化调度器
es.initScheduler()
// 2. 初始化计划任务管理器
es.cronManager = executor.NewCronManager(es.scheduler)
return es
}
func (es *ExecutorService) initScheduler() {
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)
config := executor.SchedulerConfig{
WorkerCount: workerCount,
QueueSize: queueSize,
RateInterval: time.Duration(rateInterval) * time.Millisecond,
}
handler := &ServerSchedulerHandler{es: es}
es.scheduler = executor.NewScheduler(config, handler)
es.scheduler.SetLogger(logger.NewSchedulerLogger())
es.scheduler.SetExecutor(es.ExecuteDispatcher)
es.scheduler.Start()
logger.Infof("[Executor] 调度器已启动: workers=%d, queue=%d, rate=%dms", workerCount, queueSize, rateInterval)
}
// ServerSchedulerHandler 实现 executor.SchedulerEventHandler
type ServerSchedulerHandler struct {
es *ExecutorService
}
func (h *ServerSchedulerHandler) OnTaskScheduled(req *executor.ExecutionRequest) {
// 任务入队事件,可以在此处更新数据库状态为 "pending"
}
func (h *ServerSchedulerHandler) OnTaskExecuting(req *executor.ExecutionRequest) (io.Writer, io.Writer, error) {
var taskID uint
fmt.Sscanf(req.TaskID, "%d", &taskID)
task := h.es.taskService.GetTaskByID(int(taskID))
// 系统任务(无 taskID)不记录数据库日志,直接返回空写入器
if task == nil {
return nil, nil, nil
}
// 1. 创建初始日志记录
taskLog, err := h.es.taskLogService.CreateEmptyLog(task.ID, task.Command)
if err != nil {
return nil, nil, fmt.Errorf("创建初始日志失败: %v", err)
}
req.LogID = taskLog.ID // 设置 LogID 供后续环节使用
// 2. 检查并记录运行状态(并发控制)
goid, err := h.es.AddRunningGo(task.ID)
if err != nil {
// 并发限制,更新日志状态为失败
taskLog.Status = "failed"
taskLog.Output, _ = utils.CompressToBase64("任务并发数限制,拒绝执行")
h.es.taskLogService.SaveTaskLog(taskLog)
return nil, nil, fmt.Errorf("任务并发限制: %v", err)
}
if req.Metadata == nil {
req.Metadata = make(map[string]interface{})
}
req.Metadata["goid"] = goid
// 3. 创建 TinyLog 实时日志收集器
tl, err := NewTinyLog(taskLog.ID)
if err != nil {
h.es.RemoveRunningGo(task.ID, goid) // 回滚运行状态
return nil, nil, fmt.Errorf("创建日志收集器失败: %v", err)
}
// 对于本地任务,Scheduler 会通过返回的 Writer 写入日志
// 对于远程任务,Scheduler 不会写入任何内容(由 Agent 推送至此 TL)
return tl, tl, nil
}
func (h *ServerSchedulerHandler) OnTaskHeartbeat(req *executor.ExecutionRequest, duration int64) {
if req.LogID > 0 {
h.es.taskLogService.UpdateTaskDuration(req.LogID, duration)
}
}
func (h *ServerSchedulerHandler) OnTaskStarted(req *executor.ExecutionRequest) {
// Logic moved to OnTaskExecuting
}
func (h *ServerSchedulerHandler) OnTaskCompleted(req *executor.ExecutionRequest, result *executor.ExecutionResult) {
if req.LogID == 0 {
return
}
var taskID uint
fmt.Sscanf(req.TaskID, "%d", &taskID)
task := h.es.taskService.GetTaskByID(int(taskID))
if task == nil {
return
}
// 无论本地还是远程,都在此处处理日志压缩和落库
tl := GetActiveLog(req.LogID)
var output string
if tl != nil {
// 压缩并清理实时日志
var err error
output, err = tl.CompressAndCleanup()
if err != nil {
logger.Errorf("[Executor] 压缩任务 #%d 日志失败: %v", task.ID, err)
output = "[System Error] 日志处理失败: " + err.Error()
}
} else {
// 如果 TinyLog 已经丢失,尝试从 result.Output 中恢复一次(主要针对本地任务)
output, _ = utils.CompressToBase64(result.Output)
}
// 构造待保存的日志模型
startTime := models.LocalTime(result.StartTime)
endTime := models.LocalTime(result.EndTime)
taskLog := &models.TaskLog{
ID: req.LogID,
TaskID: task.ID,
Command: req.Command,
Output: output,
Status: result.Status,
Duration: result.Duration,
ExitCode: result.ExitCode,
StartTime: &startTime,
EndTime: &endTime,
}
// 如果有 AgentID,也记录下来
if task.AgentID != nil && *task.AgentID > 0 {
agentID := *task.AgentID
taskLog.AgentID = &agentID
}
// 移除运行记录
if req.Metadata != nil {
if goid, ok := req.Metadata["goid"].(int64); ok {
h.es.RemoveRunningGo(task.ID, goid)
}
}
// 处理任务完成(更新统计、清理旧日志等)
h.es.taskLogService.ProcessTaskCompletion(taskLog)
}
func (h *ServerSchedulerHandler) OnTaskFailed(req *executor.ExecutionRequest, err error) {
if req.LogID == 0 {
return
}
var taskID uint
fmt.Sscanf(req.TaskID, "%d", &taskID)
// 移除运行记录
if req.Metadata != nil {
if goid, ok := req.Metadata["goid"].(int64); ok {
h.es.RemoveRunningGo(taskID, goid)
}
}
// 构造错误日志
tl := GetActiveLog(req.LogID)
var output string
if tl != nil {
tl.Write([]byte(fmt.Sprintf("\n[System Error] %v", err)))
output, _ = tl.CompressAndCleanup()
} else {
output, _ = utils.CompressToBase64(fmt.Sprintf("任务执行失败: %v", err))
}
now := models.LocalTime(time.Now())
taskLog := &models.TaskLog{
ID: req.LogID,
TaskID: taskID,
Output: output,
Status: "failed",
Duration: 0,
ExitCode: 1,
EndTime: &now,
}
// 补充 AgentID
task := h.es.taskService.GetTaskByID(int(taskID))
if task != nil && task.AgentID != nil && *task.AgentID > 0 {
agentID := *task.AgentID
taskLog.AgentID = &agentID
}
h.es.taskLogService.ProcessTaskCompletion(taskLog)
}
func (h *ServerSchedulerHandler) OnCronNextRun(req *executor.ExecutionRequest, nextRun time.Time) {
var taskID uint
fmt.Sscanf(req.TaskID, "%d", &taskID)
// 更新数据库中的下次运行时间
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("next_run", nextRun)
}
// LocalTaskHooks 本地任务钩子适配器
type LocalTaskHooks struct {
es *ExecutorService
logID uint
}
func (h *LocalTaskHooks) PreExecute(ctx context.Context, req executor.Request) (uint, error) {
return h.logID, nil
}
func (h *LocalTaskHooks) PostExecute(ctx context.Context, logID uint, result *executor.Result) error {
return nil
}
func (h *LocalTaskHooks) OnHeartbeat(ctx context.Context, logID uint, duration int64) error {
if logID > 0 {
return h.es.taskLogService.UpdateTaskDuration(logID, duration)
}
return nil
}
// ExecuteDispatcher 实现任务分发逻辑
func (es *ExecutorService) ExecuteDispatcher(ctx context.Context, req *executor.ExecutionRequest, stdout, stderr io.Writer) (*executor.Result, error) {
var taskID uint
fmt.Sscanf(req.TaskID, "%d", &taskID)
task := es.taskService.GetTaskByID(int(taskID))
// 系统任务(无 taskID)直接本地执行
if task == nil {
return executor.Execute(ctx, executor.Request{
Command: req.Command,
WorkDir: req.WorkDir,
Envs: req.Envs,
Timeout: req.Timeout,
}, stdout, stderr)
}
// 特殊处理仓库同步任务
if task.Type == "repo" {
cmd, workDir := es.BuildRepoCommand(task)
if cmd != "" {
req.Command = cmd
req.WorkDir = workDir
}
}
// 加载环境变量
if task.Envs != "" {
req.Envs = append(req.Envs, es.loadEnvVars(task.Envs)...)
}
// 远程任务
if task.AgentID != nil && *task.AgentID > 0 {
return es.ExecuteRemoteForScheduler(task, req.LogID)
}
// 本地任务
hooks := &LocalTaskHooks{es: es, logID: req.LogID}
return executor.ExecuteWithHooks(ctx, executor.Request{
Command: req.Command,
WorkDir: req.WorkDir,
Envs: req.Envs,
Timeout: req.Timeout,
}, stdout, stderr, hooks)
}
// getIntSetting 从设置中获取整数值
func getIntSetting(s SettingsService, section, key string, defaultVal int) int {
val := s.Get(section, key)
@@ -98,220 +361,367 @@ func getIntSetting(s SettingsService, section, key string, defaultVal int) int {
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()
es.StopCron()
es.scheduler.Stop()
}
// Reload 重新加载配置并重建 worker pool
// StartCron 启动计划任务
func (es *ExecutorService) StartCron() {
es.loadCronTasks()
es.cronManager.Start()
logger.Info("[Executor] 计划任务管理器已启动")
}
// StopCron 停止计划任务
func (es *ExecutorService) StopCron() {
es.cronManager.Stop()
logger.Info("[Executor] 计划任务管理器已停止")
}
// AddCronTask 添加计划任务
func (es *ExecutorService) AddCronTask(task *models.Task) error {
return es.cronManager.AddTask(task)
}
// RemoveCronTask 移除计划任务
func (es *ExecutorService) RemoveCronTask(taskID uint) {
es.cronManager.RemoveTask(fmt.Sprintf("%d", taskID))
}
// ValidateCron 验证 Cron 表达式
func (es *ExecutorService) ValidateCron(expression string) error {
return es.cronManager.ValidateCron(expression)
}
// GetScheduledCount 获取已加载的计划任务数量
func (es *ExecutorService) GetScheduledCount() int {
return es.cronManager.GetScheduledCount()
}
// loadCronTasks 加载所有已启用的本地计划任务
func (es *ExecutorService) loadCronTasks() {
tasks := es.taskService.GetTasks()
count := 0
for _, task := range tasks {
// 只调度本地任务(agent_id 为空或 0)
if task.Enabled && (task.AgentID == nil || *task.AgentID == 0) {
err := es.cronManager.AddTask(&task)
if err != nil {
continue
}
count++
}
}
logger.Infof("[Executor] 启动调度已加载 %d 个定时任务", count)
}
// Reload 重新加载配置并重建调度器
func (es *ExecutorService) Reload() {
logger.Info("[Executor] 正在重载配置...")
// 停止现有 workers
close(es.stopCh)
es.wg.Wait()
logger.Info("[Executor] 已停止工作线程")
es.scheduler.Stop()
// 从设置中读取新配置
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)
}
es.initScheduler()
}
// 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 {
func (es *ExecutorService) ExecuteTask(taskID int) *executor.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(),
return &executor.ExecutionResult{
TaskID: fmt.Sprintf("%d", taskID),
Success: false,
Error: "任务不存在",
StartTime: time.Now(),
EndTime: 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,
// 1. 检查并发
if err := es.CheckConcurrency(uint(taskID)); err != nil {
return &executor.ExecutionResult{
TaskID: fmt.Sprintf("%d", taskID),
Success: false,
Error: err.Error(), // 这里会返回 "任务正在运行中,拒绝并行执行"
StartTime: time.Now(),
EndTime: time.Now(),
}
}
// 标记任务结束
es.mu.Lock()
delete(es.runningTasks, taskID)
es.mu.Unlock()
req := &executor.ExecutionRequest{
TaskID: fmt.Sprintf("%d", task.ID),
Name: task.Name,
Command: task.Command,
WorkDir: task.WorkDir,
Envs: es.loadEnvVars(task.Envs),
Timeout: task.Timeout,
Type: executor.TaskTypeManual,
}
return result
es.scheduler.EnqueueOrExecute(req)
return &executor.ExecutionResult{
TaskID: fmt.Sprintf("%d", task.ID),
Success: true,
Status: "queued",
StartTime: time.Now(),
}
}
// GetRunningCount 获取正在运行任务数量
// GetRunningCount 获取正在运行任务数量
func (es *ExecutorService) GetRunningCount() int {
es.mu.RLock()
defer es.mu.RUnlock()
return len(es.runningTasks)
return es.scheduler.GetRunningTaskCount()
}
// ExecuteCommand executes a shell command with default timeout
func (es *ExecutorService) ExecuteCommand(command string) *ExecutionResult {
func (es *ExecutorService) ExecuteCommand(command string) *executor.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 {
func (es *ExecutorService) ExecuteCommandWithTimeout(command string, timeout time.Duration) *executor.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 {
func (es *ExecutorService) ExecuteCommandWithEnv(command string, timeout time.Duration, envVars []string) *executor.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(),
func (es *ExecutorService) ExecuteCommandWithOptions(command string, timeout time.Duration, envVars []string, workDir string) *executor.ExecutionResult {
req := &executor.ExecutionRequest{
Command: command,
Timeout: int(timeout.Minutes()),
Envs: envVars,
WorkDir: workDir,
Type: executor.TaskTypeSystem,
}
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
}
res, _ := es.scheduler.ExecuteSync(req)
// 使用独立锁保存结果
es.resultsMu.Lock()
es.results = append(es.results, *result)
if len(es.results) > 100 {
es.results = es.results[1:]
}
es.resultsMu.Unlock()
// TODO: 适配 ExecutionResult 的转换并保存结果
return result
return res
}
// GetLastResults returns the last execution results
func (es *ExecutorService) GetLastResults(count int) []ExecutionResult {
func (es *ExecutorService) GetLastResults(count int) []executor.ExecutionResult {
es.resultsMu.RLock()
defer es.resultsMu.RUnlock()
return nil
}
start := 0
if len(es.results) > count {
start = len(es.results) - count
// --- 以下内容从 TaskExecutionService 合并 ---
// CleanupRunningTasks 清理所有任务的运行状态(在重启时调用)
func (es *ExecutorService) CleanupRunningTasks() error {
logger.Info("[Executor] 正在清理残留的任务运行状态...")
return database.DB.Model(&models.Task{}).Where("1=1").Update("running_go", "[]").Error
}
// CheckConcurrency 检查任务并发限制(只读检查)
func (es *ExecutorService) CheckConcurrency(taskID uint) error {
var task models.Task
if err := database.DB.Select("config, running_go").First(&task, taskID).Error; err != nil {
return err
}
var goids []int64
if task.RunningGo != "" {
_ = json.Unmarshal([]byte(task.RunningGo), &goids)
}
results := make([]ExecutionResult, len(es.results[start:]))
copy(results, es.results[start:])
return results
var config models.TaskConfig
if task.Config != "" {
_ = json.Unmarshal([]byte(task.Config), &config)
}
if config.Concurrency == 0 && len(goids) > 0 {
return fmt.Errorf("任务正在运行中,拒绝并行执行,请前往日志查看")
}
return nil
}
// AddRunningGo 添加当前 goroutine ID 到任务的 running_go 字段
func (es *ExecutorService) AddRunningGo(taskID uint) (int64, error) {
goid := utils.GetGoroutineID()
err := database.DB.Transaction(func(tx *gorm.DB) error {
var task models.Task
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&task, taskID).Error; err != nil {
return err
}
var goids []int64
if task.RunningGo != "" {
_ = json.Unmarshal([]byte(task.RunningGo), &goids)
}
// 解析配置以获取并发设置
var config models.TaskConfig
if task.Config != "" {
_ = json.Unmarshal([]byte(task.Config), &config)
}
// 如果并发为0(禁用)且已有执行中的任务,返回错误
if config.Concurrency == 0 && len(goids) > 0 {
return fmt.Errorf("task is running")
}
goids = append(goids, goid)
data, _ := json.Marshal(goids)
return tx.Model(&task).Update("running_go", string(data)).Error
})
return goid, err
}
// RemoveRunningGo 从任务的 running_go 字段移除指定 goroutine ID
func (es *ExecutorService) RemoveRunningGo(taskID uint, goid int64) {
database.DB.Transaction(func(tx *gorm.DB) error {
var task models.Task
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&task, taskID).Error; err != nil {
return err
}
var goids []int64
if task.RunningGo != "" {
_ = json.Unmarshal([]byte(task.RunningGo), &goids)
}
newGoids := make([]int64, 0)
for _, id := range goids {
if id != goid {
newGoids = append(newGoids, id)
}
}
data, _ := json.Marshal(newGoids)
return tx.Model(&task).Update("running_go", string(data)).Error
})
}
// ExecuteRemoteForScheduler 供 Scheduler 调用,执行远程任务并等待结果
func (es *ExecutorService) ExecuteRemoteForScheduler(task *models.Task, logID uint) (*executor.Result, error) {
agentID := *task.AgentID
logger.Infof("[Executor] 远程执行任务 #%d: %s (Agent #%d, LogID: %d)", task.ID, task.Name, agentID, logID)
// 1. 检查 Agent 状态
var agent models.Agent
if err := database.DB.First(&agent, agentID).Error; err != nil {
return nil, fmt.Errorf("Agent #%d 不存在", agentID)
}
if !agent.Enabled {
return nil, fmt.Errorf("Agent #%d 已禁用", agentID)
}
if es.agentWSManager == nil {
return nil, fmt.Errorf("AgentWSManager 未初始化")
}
// 2. 注册结果等待者
resultChan := es.agentWSManager.RegisterRemoteWaiter(logID)
defer es.agentWSManager.UnregisterRemoteWaiter(logID)
// 3. 发送指令
err := es.agentWSManager.SendToAgent(agentID, constant.WSTypeExecute, map[string]interface{}{
"task_id": task.ID,
"log_id": logID,
})
if err != nil {
return nil, fmt.Errorf("发送执行命令失败: %v", err)
}
// 4. 等待结果或超时
timeout := task.Timeout
if timeout <= 0 {
timeout = 30
}
start := time.Now()
select {
case agentResult := <-resultChan:
return &executor.Result{
Output: agentResult.Output,
Status: agentResult.Status,
Duration: agentResult.Duration,
ExitCode: agentResult.ExitCode,
StartTime: time.Unix(agentResult.StartTime, 0),
EndTime: time.Unix(agentResult.EndTime, 0),
}, nil
case <-time.After(time.Duration(timeout) * time.Minute):
end := time.Now()
return &executor.Result{
Status: "failed",
Duration: end.Sub(start).Milliseconds(),
ExitCode: -1,
StartTime: start,
EndTime: end,
}, fmt.Errorf("远程执行超时")
}
}
// HandleAgentResult 处理来自 Agent 的异步结果
func (es *ExecutorService) HandleAgentResult(result *models.AgentTaskResult) error {
taskLog, err := es.taskLogService.CreateTaskLogFromAgentResult(result)
if err != nil {
return err
}
return es.taskLogService.ProcessTaskCompletion(taskLog)
}
// BuildRepoCommand 构建仓库同步任务的命令
func (es *ExecutorService) BuildRepoCommand(task *models.Task) (string, string) {
var config models.RepoConfig
if err := json.Unmarshal([]byte(task.Config), &config); err != nil {
return "", ""
}
targetPath := config.TargetPath
if targetPath == "" {
targetPath = constant.ScriptsWorkDir
} else if !filepath.IsAbs(targetPath) {
targetPath = filepath.Join(constant.ScriptsWorkDir, targetPath)
}
absTargetPath, _ := filepath.Abs(targetPath)
args := []string{
"/opt/sync.py",
"--source-type", config.SourceType,
"--source-url", config.SourceURL,
"--target-path", absTargetPath,
}
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)
}
}
if config.AuthToken != "" {
args = append(args, "--auth-token", config.AuthToken)
}
return "python3 " + strings.Join(args, " "), "/opt"
}
// loadEnvVars 加载环境变量
func (es *ExecutorService) 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
}
@@ -1,425 +0,0 @@
package tasks
import (
"github.com/engigu/baihu-panel/internal/constant"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"github.com/engigu/baihu-panel/internal/utils"
"bytes"
"context"
"encoding/json"
"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 {
task := req.Task
start := time.Now()
// 演示模式:直接返回模拟结果
if constant.DemoMode {
end := time.Now()
demoOutput := fmt.Sprintf("[演示模式] 任务 #%d (%s) 执行已跳过\n实际命令不会运行: %s", task.ID, task.Name, task.Command)
result := &TaskExecutionResult{
TaskID: task.ID,
AgentID: nil,
Command: task.Command,
Output: demoOutput,
Status: "success",
Duration: end.Sub(start).Milliseconds(),
ExitCode: 0,
Start: start,
End: end,
}
return s.processExecutionResult(result)
}
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)
// 检查任务类型,仓库任务需要特殊处理
if task.Type == "repo" {
return s.executeRepoTask(req)
}
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)
}
// 使用工具函数获取合适的 shell
shell, _ := utils.GetShell()
cmd = exec.CommandContext(ctx, shell, "-c", command)
}
// 设置环境变量(始终继承系统环境变量)
cmd.Env = os.Environ()
if len(envVars) > 0 {
cmd.Env = append(cmd.Env, 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)
}
// executeRepoTask 执行仓库同步任务(调用 sync.py)
func (s *TaskExecutionService) executeRepoTask(req *TaskExecutionRequest) error {
task := req.Task
logger.Infof("[TaskExecution] 执行仓库同步任务 #%d: %s", task.ID, task.Name)
start := time.Now()
// 解析仓库配置
var config models.RepoConfig
if err := json.Unmarshal([]byte(task.Config), &config); err != nil {
return s.handleExecutionError(task.ID, "", start, fmt.Errorf("解析仓库配置失败: %v", err))
}
// 处理目标路径:为空则使用 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)
}
// 准备命令
ctx, cancel := s.createContext(task.Timeout)
defer cancel()
// 直接使用 python3 和参数列表,而不是拼接成字符串
cmd := exec.CommandContext(ctx, "python3", args...)
cmd.Dir = "/opt"
// 执行命令
var stdout, stderr bytes.Buffer
cmd.Stdout = &stdout
cmd.Stderr = &stderr
execErr := cmd.Run()
end := time.Now()
// 构建命令字符串用于日志记录
commandStr := "python3 " + strings.Join(args, " ")
// 构建结果
result := &TaskExecutionResult{
TaskID: task.ID,
AgentID: nil,
Command: commandStr,
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)
}
+46 -11
View File
@@ -1,12 +1,13 @@
package tasks
import (
"encoding/json"
"time"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/logger"
"github.com/engigu/baihu-panel/internal/models"
"github.com/engigu/baihu-panel/internal/utils"
"encoding/json"
"time"
)
// SendStatsService 接口定义(避免循环依赖)
@@ -32,9 +33,31 @@ type CleanConfig struct {
Keep int `json:"keep"` // 保留天数或条数
}
// SaveTaskLog 保存任务日志(通用方法
func (s *TaskLogService) SaveTaskLog(taskLog *models.TaskLog) error {
// CreateEmptyLog 创建一个空的日志记录(任务开始时调用
func (s *TaskLogService) CreateEmptyLog(taskID uint, command string) (*models.TaskLog, error) {
startTime := models.LocalTime(time.Now())
taskLog := &models.TaskLog{
TaskID: taskID,
Command: command,
Status: "running",
StartTime: &startTime,
}
if err := database.DB.Create(taskLog).Error; err != nil {
return nil, err
}
return taskLog, nil
}
// SaveTaskLog 保存或更新任务日志
func (s *TaskLogService) SaveTaskLog(taskLog *models.TaskLog) error {
var err error
if taskLog.ID > 0 {
err = database.DB.Model(taskLog).Updates(taskLog).Error
} else {
err = database.DB.Create(taskLog).Error
}
if err != nil {
return err
}
@@ -44,6 +67,11 @@ func (s *TaskLogService) SaveTaskLog(taskLog *models.TaskLog) error {
return nil
}
// UpdateTaskDuration 更新任务耗时(心跳)
func (s *TaskLogService) UpdateTaskDuration(logID uint, duration int64) error {
return database.DB.Model(&models.TaskLog{}).Where("id = ?", logID).Update("duration", duration).Error
}
// UpdateTaskStats 更新任务统计
func (s *TaskLogService) UpdateTaskStats(taskID uint, status string) {
if s.sendStatsService == nil {
@@ -100,7 +128,7 @@ func (s *TaskLogService) CleanTaskLogs(taskID uint) {
// ProcessTaskCompletion 处理任务完成后的所有操作(保存日志、更新统计、清理旧日志)
func (s *TaskLogService) ProcessTaskCompletion(taskLog *models.TaskLog) error {
// 1. 保存日志
// 1. 保存/更新日志
if err := s.SaveTaskLog(taskLog); err != nil {
return err
}
@@ -147,12 +175,19 @@ func (s *TaskLogService) CreateTaskLogFromAgentResult(result *models.AgentTaskRe
}
// 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 = ""
func (s *TaskLogService) CreateTaskLogFromLocalExecution(taskID uint, command, output, 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 {
// 压缩输出
compressed, err = utils.CompressToBase64(output)
if err != nil {
logger.Errorf("[TaskLog] 压缩日志失败: %v", err)
compressed = ""
}
}
startTime := models.LocalTime(start)
+241
View File
@@ -0,0 +1,241 @@
package tasks
import (
"bufio"
"bytes"
"compress/zlib"
"encoding/base64"
"io"
"os"
"sync"
"github.com/engigu/baihu-panel/internal/utils"
)
var (
// globalTinyLogManager keeps track of all active TinyLog instances
globalTinyLogManager = &TinyLogManager{
logs: make(map[uint]*TinyLog),
}
)
type TinyLogManager struct {
mu sync.RWMutex
logs map[uint]*TinyLog
}
func (m *TinyLogManager) Register(log *TinyLog) {
m.mu.Lock()
defer m.mu.Unlock()
m.logs[log.LogID] = log
}
func (m *TinyLogManager) Unregister(logID uint) {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.logs, logID)
}
func (m *TinyLogManager) Get(logID uint) *TinyLog {
m.mu.RLock()
defer m.mu.RUnlock()
return m.logs[logID]
}
// GetActiveLog returns an active TinyLog by its ID
func GetActiveLog(logID uint) *TinyLog {
return globalTinyLogManager.Get(logID)
}
// TinyLog is a high-performance, low-memory log collector
type TinyLog struct {
LogID uint
mu sync.RWMutex
file *os.File
path string
writer *bufio.Writer
subscribers []chan []byte
closed bool
}
// NewTinyLog creates a new TinyLog instance backed by a temporary file and registers it
func NewTinyLog(logID uint) (*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),
}
globalTinyLogManager.Register(tl)
return tl, nil
}
// Write implements 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
}
// 1. Convert to UTF-8 if necessary (common on Windows)
text := utils.ToUTF8(p)
data := []byte(text)
// 2. Write to file buffer
_, err = l.writer.Write(data)
if err != nil {
return 0, err
}
// 3. Broadcast to subscribers
if len(l.subscribers) > 0 {
for _, ch := range l.subscribers {
select {
case ch <- data:
default:
// Drop message if subscriber is too slow to avoid blocking writer
}
}
}
return len(p), nil
}
// Subscribe returns a channel that receives log chunks in real-time
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 removes a subscriber
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 finishes writing and closes the file, and unregisters itself
func (l *TinyLog) Close() error {
l.mu.Lock()
defer l.mu.Unlock()
if l.closed {
return nil
}
// Flush buffer to file
if err := l.writer.Flush(); err != nil {
return err
}
// Close all subscribers
for _, ch := range l.subscribers {
close(ch)
}
l.subscribers = nil
l.closed = true
globalTinyLogManager.Unregister(l.LogID)
return l.file.Close()
}
// CompressAndCleanup reads the temporary file, compresses it, returns the result, and removes the file
func (l *TinyLog) CompressAndCleanup() (string, error) {
// Ensure closed
if !l.closed {
l.Close()
}
// Open temp file for reading
f, err := os.Open(l.path)
if err != nil {
return "", err
}
defer func() {
f.Close()
os.Remove(l.path) // Cleanup
}()
// Create buffer for compressed output
var buf bytes.Buffer
b64Writer := base64.NewEncoder(base64.StdEncoding, &buf)
zlibWriter := zlib.NewWriter(b64Writer)
// Stream: File -> Zlib -> Base64 -> Buffer
if _, err := io.Copy(zlibWriter, f); err != nil {
return "", err
}
// Close generic writers to flush data
if err := zlibWriter.Close(); err != nil {
return "", err
}
if err := b64Writer.Close(); err != nil {
return "", err
}
return buf.String(), nil
}
// ReadLastLines returns the last n lines of the log
func (l *TinyLog) ReadLastLines(n int) ([]byte, error) {
l.mu.RLock()
defer l.mu.RUnlock()
// Flush writer to ensure file on disk is up to date
_ = l.writer.Flush()
stat, err := os.Stat(l.path)
if err != nil {
return nil, err
}
size := stat.Size()
var limit int64 = 65536 // Max 64KB for "last 100 lines" preview
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 returns the temporary file path
func (l *TinyLog) GetPath() string {
return l.path
}
+21 -20
View File
@@ -4,6 +4,7 @@ import (
"bytes"
"compress/zlib"
"encoding/base64"
"io"
)
// CompressToBase64 compresses data using zlib and encodes to base64
@@ -22,23 +23,23 @@ func CompressToBase64(data string) (string, error) {
return base64.StdEncoding.EncodeToString(buf.Bytes()), nil
}
// // DecompressFromBase64 decodes base64 and decompresses zlib data
// func DecompressFromBase64(data string) (string, error) {
// if data == "" {
// return "", nil
// }
// decoded, err := base64.StdEncoding.DecodeString(data)
// if err != nil {
// return "", err
// }
// zr, err := zlib.NewReader(bytes.NewReader(decoded))
// if err != nil {
// return "", err
// }
// defer zr.Close()
// result, err := io.ReadAll(zr)
// if err != nil {
// return "", err
// }
// return string(result), nil
// }
// DecompressFromBase64 decodes base64 and decompresses zlib data
func DecompressFromBase64(data string) (string, error) {
if data == "" {
return "", nil
}
decoded, err := base64.StdEncoding.DecodeString(data)
if err != nil {
return "", err
}
zr, err := zlib.NewReader(bytes.NewReader(decoded))
if err != nil {
return "", err
}
defer zr.Close()
result, err := io.ReadAll(zr)
if err != nil {
return "", err
}
return string(result), nil
}
+43
View File
@@ -0,0 +1,43 @@
package utils
import (
"bufio"
"io"
"unicode/utf8"
"golang.org/x/text/encoding/simplifiedchinese"
"golang.org/x/text/transform"
)
// ToUTF8 converts potentially non-UTF8 data (like GBK on Windows) to UTF-8
func ToUTF8(data []byte) string {
if utf8.Valid(data) {
return string(data)
}
// Try GBK (common on Windows)
reader := transform.NewReader(
bufio.NewReader(
&byteReader{data: data},
),
simplifiedchinese.GBK.NewDecoder(),
)
result, err := io.ReadAll(reader)
if err != nil {
return string(data)
}
return string(result)
}
type byteReader struct {
data []byte
pos int
}
func (r *byteReader) Read(p []byte) (n int, err error) {
if r.pos >= len(r.data) {
return 0, io.EOF
}
n = copy(p, r.data[r.pos:])
r.pos += n
return n, nil
}
+20
View File
@@ -0,0 +1,20 @@
package utils
import (
"runtime"
"strconv"
"strings"
)
// GetGoroutineID 获取当前 Goroutine ID
// 注意:这只是为了调试和日志目的,不应该用于业务逻辑
func GetGoroutineID() int64 {
var buf [64]byte
n := runtime.Stack(buf[:], false)
idField := strings.Fields(strings.TrimPrefix(string(buf[:n]), "goroutine "))[0]
id, err := strconv.ParseInt(idField, 10, 64)
if err != nil {
return -1
}
return id
}
+51
View File
@@ -0,0 +1,51 @@
package utils
import (
"crypto/sha256"
"encoding/hex"
"net"
"os"
"runtime"
"sort"
"strings"
)
// GenerateMachineID 生成机器识别码
func GenerateMachineID() string {
var parts []string
// 主机名
if hostname, err := os.Hostname(); err == nil {
parts = append(parts, hostname)
}
// 获取所有非回环网卡的 MAC 地址,排序后取第一个(最稳定)
if interfaces, err := net.Interfaces(); err == nil {
var macs []string
for _, iface := range interfaces {
// 跳过回环接口、没有 MAC 地址的接口、虚拟接口
if iface.Flags&net.FlagLoopback != 0 || len(iface.HardwareAddr) == 0 {
continue
}
// 跳过 docker/veth 等虚拟网卡
name := strings.ToLower(iface.Name)
if strings.HasPrefix(name, "docker") || strings.HasPrefix(name, "veth") ||
strings.HasPrefix(name, "br-") || strings.HasPrefix(name, "virbr") {
continue
}
macs = append(macs, iface.HardwareAddr.String())
}
sort.Strings(macs)
// 只使用第一个 MAC 地址(最稳定)
if len(macs) > 0 {
parts = append(parts, macs[0])
}
}
// 操作系统和架构
parts = append(parts, runtime.GOOS, runtime.GOARCH)
data := strings.Join(parts, "|")
hash := sha256.Sum256([]byte(data))
return hex.EncodeToString(hash[:])
}