feat: refact scheduler
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
// 正在运行:目前只能统计本地运行的任务
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user