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
+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 {