feat: refact id column define

This commit is contained in:
engigu
2026-03-03 17:25:38 +08:00
parent 117ae126f7
commit 4589c64923
62 changed files with 881 additions and 454 deletions
+41 -41
View File
@@ -49,8 +49,8 @@ func (c *AgentController) List(ctx *gin.Context) {
// Update 更新 Agent
func (c *AgentController) Update(ctx *gin.Context) {
id, err := strconv.ParseUint(ctx.Param("id"), 10, 32)
if err != nil {
id := ctx.Param("id")
if id == "" {
utils.BadRequest(ctx, "无效的 ID")
return
}
@@ -67,14 +67,14 @@ func (c *AgentController) Update(ctx *gin.Context) {
}
// 获取旧状态
oldAgent := c.agentService.GetByID(uint(id))
oldAgent := c.agentService.GetByID(id)
if oldAgent == nil {
utils.NotFound(ctx, "Agent 不存在")
return
}
wasEnabled := oldAgent.Enabled
if err := c.agentService.Update(uint(id), req.Name, req.Description, req.Enabled); err != nil {
if err := c.agentService.Update(id, req.Name, req.Description, req.Enabled); err != nil {
utils.ServerError(ctx, err.Error())
return
}
@@ -83,14 +83,14 @@ func (c *AgentController) Update(ctx *gin.Context) {
if wasEnabled != req.Enabled {
if req.Enabled {
// 启用:发送任务列表
c.wsManager.SendToAgent(uint(id), services.WSTypeEnabled, map[string]interface{}{
c.wsManager.SendToAgent(id, services.WSTypeEnabled, map[string]interface{}{
"message": "Agent 已启用",
})
// 发送任务列表
c.wsManager.BroadcastTasks(uint(id))
c.wsManager.BroadcastTasks(id)
} else {
// 禁用:发送禁用消息,Agent 收到后清空任务
c.wsManager.SendToAgent(uint(id), services.WSTypeDisabled, map[string]interface{}{
c.wsManager.SendToAgent(id, services.WSTypeDisabled, map[string]interface{}{
"message": "Agent 已禁用",
})
}
@@ -101,13 +101,13 @@ func (c *AgentController) Update(ctx *gin.Context) {
// Delete 删除 Agent
func (c *AgentController) Delete(ctx *gin.Context) {
id, err := strconv.ParseUint(ctx.Param("id"), 10, 32)
if err != nil {
id := ctx.Param("id")
if id == "" {
utils.BadRequest(ctx, "无效的 ID")
return
}
if err := c.agentService.Delete(uint(id)); err != nil {
if err := c.agentService.Delete(id); err != nil {
utils.BadRequest(ctx, err.Error())
return
}
@@ -117,13 +117,13 @@ func (c *AgentController) Delete(ctx *gin.Context) {
// RegenerateToken 重新生成 Token
func (c *AgentController) RegenerateToken(ctx *gin.Context) {
id, err := strconv.ParseUint(ctx.Param("id"), 10, 32)
if err != nil {
id := ctx.Param("id")
if id == "" {
utils.BadRequest(ctx, "无效的 ID")
return
}
token, err := c.agentService.RegenerateToken(uint(id))
token, err := c.agentService.RegenerateToken(id)
if err != nil {
utils.ServerError(ctx, err.Error())
return
@@ -324,13 +324,13 @@ func (c *AgentController) GetVersion(ctx *gin.Context) {
// ForceUpdate 强制更新指定 Agent
func (c *AgentController) ForceUpdate(ctx *gin.Context) {
id, err := strconv.ParseUint(ctx.Param("id"), 10, 32)
if err != nil {
id := ctx.Param("id")
if id == "" {
utils.BadRequest(ctx, "无效的 ID")
return
}
if err := c.agentService.SetForceUpdate(uint(id)); err != nil {
if err := c.agentService.SetForceUpdate(id); err != nil {
utils.ServerError(ctx, err.Error())
return
}
@@ -390,20 +390,20 @@ func (c *AgentController) WSConnect(ctx *gin.Context) {
ctx.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
return
}
logger.Infof("[AgentWS] 注册成功: Agent #%d, isNew=%v", agent.ID, isNewAgent)
logger.Infof("[AgentWS] 注册成功: Agent #%s, isNew=%v", agent.ID, isNewAgent)
}
if !agent.Enabled {
c.wsManager.RecordConnectFail(ip)
logger.Warnf("[AgentWS] Agent #%d 已禁用, IP=%s", agent.ID, ip)
logger.Warnf("[AgentWS] Agent #%s 已禁用, IP=%s", agent.ID, ip)
ctx.JSON(http.StatusForbidden, gin.H{"error": "Agent 已禁用"})
return
}
logger.Infof("[AgentWS] 准备升级连接: Agent #%d, IP=%s", agent.ID, ip)
logger.Infof("[AgentWS] 准备升级连接: Agent #%s, IP=%s", agent.ID, ip)
conn, err := agentUpgrader.Upgrade(ctx.Writer, ctx.Request, nil)
if err != nil {
logger.Errorf("[AgentWS] 升级连接失败: %v, Agent #%d, IP=%s", err, agent.ID, ip)
logger.Errorf("[AgentWS] 升级连接失败: %v, Agent #%s, IP=%s", err, agent.ID, ip)
return
}
@@ -434,7 +434,7 @@ func (c *AgentController) WSConnect(ctx *gin.Context) {
},
})
logger.Infof("[AgentWS] Agent #%d 连接成功 (配置: workers=%d, queue=%d, rate=%d)",
logger.Infof("[AgentWS] Agent #%s 连接成功 (配置: workers=%d, queue=%d, rate=%d)",
agent.ID, workerCount, queueSize, rateInterval)
// 启动读写协程
@@ -449,9 +449,9 @@ func (c *AgentController) WSConnect(ctx *gin.Context) {
func (c *AgentController) wsReadPump(ac *services.AgentConnection, agent *models.Agent) {
defer func() {
if r := recover(); r != nil {
logger.Errorf("[AgentWS] Agent #%d wsReadPump panic: %v", agent.ID, r)
logger.Errorf("[AgentWS] Agent #%s wsReadPump panic: %v", agent.ID, r)
}
logger.Infof("[AgentWS] Agent #%d wsReadPump 退出", agent.ID)
logger.Infof("[AgentWS] Agent #%s wsReadPump 退出", agent.ID)
c.wsManager.Unregister(agent.ID, ac)
}()
@@ -471,7 +471,7 @@ func (c *AgentController) wsReadPump(ac *services.AgentConnection, agent *models
for {
_, message, err := ac.ReadMessage()
if err != nil {
logger.Warnf("[AgentWS] Agent #%d 读取错误: %v", agent.ID, err)
logger.Warnf("[AgentWS] Agent #%s 读取错误: %v", agent.ID, err)
break
}
@@ -488,9 +488,9 @@ func (c *AgentController) wsReadPump(ac *services.AgentConnection, agent *models
func (c *AgentController) wsWritePump(ac *services.AgentConnection) {
defer func() {
if r := recover(); r != nil {
logger.Errorf("[AgentWS] Agent #%d wsWritePump panic: %v", ac.AgentID, r)
logger.Errorf("[AgentWS] Agent #%s wsWritePump panic: %v", ac.AgentID, r)
}
logger.Infof("[AgentWS] Agent #%d wsWritePump 退出", ac.AgentID)
logger.Infof("[AgentWS] Agent #%s wsWritePump 退出", ac.AgentID)
}()
ticker := time.NewTicker(30 * time.Second)
@@ -500,15 +500,15 @@ func (c *AgentController) wsWritePump(ac *services.AgentConnection) {
select {
case message, ok := <-ac.Send:
if !ok {
logger.Warnf("[AgentWS] Agent #%d Send channel 已关闭", ac.AgentID)
logger.Warnf("[AgentWS] Agent #%s Send channel 已关闭", ac.AgentID)
return
}
if ac.IsClosed() {
logger.Warnf("[AgentWS] Agent #%d 连接已关闭(write)", ac.AgentID)
logger.Warnf("[AgentWS] Agent #%s 连接已关闭(write)", ac.AgentID)
return
}
if err := ac.WriteMessage(message); err != nil {
logger.Warnf("[AgentWS] Agent #%d 写入消息失败: %v", ac.AgentID, err)
logger.Warnf("[AgentWS] Agent #%s 写入消息失败: %v", ac.AgentID, err)
return
}
case <-ticker.C:
@@ -516,7 +516,7 @@ func (c *AgentController) wsWritePump(ac *services.AgentConnection) {
return
}
if err := ac.WritePing(); err != nil {
logger.Warnf("[AgentWS] Agent #%d 发送 Ping 失败: %v", ac.AgentID, err)
logger.Warnf("[AgentWS] Agent #%s 发送 Ping 失败: %v", ac.AgentID, err)
return
}
}
@@ -546,15 +546,15 @@ func (c *AgentController) handleWSMessage(ac *services.AgentConnection, agent *m
// handleTaskHeartbeat 处理任务心跳
func (c *AgentController) handleTaskHeartbeat(agent *models.Agent, data json.RawMessage) {
var req struct {
LogID uint `json:"log_id"`
Duration int64 `json:"duration"`
LogID string `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)
if req.LogID != "" {
logger.Infof("[AgentWS] 收到任务心跳: LogID=%s, Duration=%dms", req.LogID, req.Duration)
c.agentService.UpdateTaskDuration(req.LogID, req.Duration)
}
}
@@ -565,7 +565,7 @@ func (c *AgentController) handleFetchTasks(agent *models.Agent) {
c.wsManager.SendToAgent(agent.ID, services.WSTypeTasks, map[string]interface{}{
"tasks": tasks,
})
logger.Infof("[AgentWS] Agent #%d 请求任务列表,返回 %d 个任务", agent.ID, len(tasks))
logger.Infof("[AgentWS] Agent #%s 请求任务列表,返回 %d 个任务", agent.ID, len(tasks))
}
// handleHeartbeat 处理心跳
@@ -619,7 +619,7 @@ func (c *AgentController) handleTaskResult(agent *models.Agent, data json.RawMes
// handleTaskLog 处理 Agent 发送的实时日志
func (c *AgentController) handleTaskLog(agent *models.Agent, data json.RawMessage) {
var logMsg struct {
LogID uint `json:"log_id"`
LogID string `json:"log_id"`
Content string `json:"content"`
}
if err := json.Unmarshal(data, &logMsg); err != nil {
@@ -631,12 +631,12 @@ func (c *AgentController) handleTaskLog(agent *models.Agent, data json.RawMessag
if tl != nil {
tl.Write([]byte(logMsg.Content))
} else {
logger.Warnf("[AgentWS] 收到任务日志但未找到活跃 TinyLog: LogID=%d, ContentSize=%d", logMsg.LogID, len(logMsg.Content))
logger.Warnf("[AgentWS] 收到任务日志 but could not find active TinyLog: LogID=%s, ContentSize=%d", logMsg.LogID, len(logMsg.Content))
}
}
// NotifyTaskUpdate 通知 Agent 任务更新
func (c *AgentController) NotifyTaskUpdate(agentID uint) {
func (c *AgentController) NotifyTaskUpdate(agentID string) {
c.wsManager.BroadcastTasks(agentID)
}
@@ -682,13 +682,13 @@ func (c *AgentController) CreateToken(ctx *gin.Context) {
// DeleteToken 删除令牌
func (c *AgentController) DeleteToken(ctx *gin.Context) {
id, err := strconv.ParseUint(ctx.Param("id"), 10, 32)
if err != nil {
id := ctx.Param("id")
if id == "" {
utils.BadRequest(ctx, "无效的 ID")
return
}
if err := c.agentService.DeleteToken(uint(id)); err != nil {
if err := c.agentService.DeleteToken(id); err != nil {
utils.ServerError(ctx, err.Error())
return
}
+4 -4
View File
@@ -139,7 +139,7 @@ func (dc *DashboardController) GetSendStats(c *gin.Context) {
// TaskStats 任务执行统计
type TaskStats struct {
TaskID uint `json:"task_id"`
TaskID string `json:"task_id"`
TaskName string `json:"task_name"`
Count int `json:"count"`
}
@@ -159,7 +159,7 @@ func (dc *DashboardController) GetTaskStats(c *gin.Context) {
// 按 task_id 聚合统计
var results []struct {
TaskID uint
TaskID string
Total int
}
database.DB.Model(&models.SendStats{}).
@@ -170,7 +170,7 @@ func (dc *DashboardController) GetTaskStats(c *gin.Context) {
Find(&results)
// 获取任务名称
taskIDs := make([]uint, 0, len(results))
taskIDs := make([]string, 0, len(results))
for _, r := range results {
taskIDs = append(taskIDs, r.TaskID)
}
@@ -179,7 +179,7 @@ func (dc *DashboardController) GetTaskStats(c *gin.Context) {
if len(taskIDs) > 0 {
database.DB.Where("id IN ?", taskIDs).Find(&tasks)
}
taskNameMap := make(map[uint]string)
taskNameMap := make(map[string]string)
for _, t := range tasks {
taskNameMap[t.ID] = t.Name
}
@@ -1,7 +1,6 @@
package controllers
import (
"strconv"
"strings"
"github.com/engigu/baihu-panel/internal/models"
@@ -71,8 +70,8 @@ func (c *DependencyController) Create(ctx *gin.Context) {
// Delete 删除依赖
func (c *DependencyController) Delete(ctx *gin.Context) {
id, err := strconv.Atoi(ctx.Param("id"))
if err != nil {
id := ctx.Param("id")
if id == "" {
utils.BadRequest(ctx, "无效的 ID")
return
}
@@ -186,8 +185,8 @@ func (c *DependencyController) GetReinstallAllCommand(ctx *gin.Context) {
// Uninstall 卸载依赖
func (c *DependencyController) Uninstall(ctx *gin.Context) {
id, err := strconv.Atoi(ctx.Param("id"))
if err != nil {
id := ctx.Param("id")
if id == "" {
utils.BadRequest(ctx, "无效的 ID")
return
}
@@ -220,8 +219,8 @@ func (c *DependencyController) Uninstall(ctx *gin.Context) {
// Reinstall 重新安装依赖
func (c *DependencyController) Reinstall(ctx *gin.Context) {
id, err := strconv.Atoi(ctx.Param("id"))
if err != nil {
id := ctx.Param("id")
if id == "" {
utils.BadRequest(ctx, "无效的 ID")
return
}
+9 -10
View File
@@ -1,7 +1,6 @@
package controllers
import (
"strconv"
"github.com/engigu/baihu-panel/internal/models/vo"
"github.com/engigu/baihu-panel/internal/services"
@@ -19,7 +18,7 @@ func NewEnvController(envService *services.EnvService) *EnvController {
}
func (ec *EnvController) CreateEnvVar(c *gin.Context) {
userID := 1
userID := c.GetString("userID")
var req struct {
Name string `json:"name" binding:"required"`
@@ -43,7 +42,7 @@ func (ec *EnvController) CreateEnvVar(c *gin.Context) {
}
func (ec *EnvController) GetEnvVars(c *gin.Context) {
userID := 1
userID := c.GetString("userID")
p := utils.ParsePagination(c)
name := c.DefaultQuery("name", "")
envVars, total := ec.envService.GetEnvVarsWithPagination(userID, name, p.Page, p.PageSize)
@@ -51,14 +50,14 @@ func (ec *EnvController) GetEnvVars(c *gin.Context) {
}
func (ec *EnvController) GetAllEnvVars(c *gin.Context) {
userID := 1
userID := c.GetString("userID")
envVars := ec.envService.GetEnvVarsByUserID(userID)
utils.Success(c, vo.ToEnvVOListFromModels(envVars))
}
func (ec *EnvController) GetEnvVar(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的环境变量ID")
return
}
@@ -73,8 +72,8 @@ func (ec *EnvController) GetEnvVar(c *gin.Context) {
}
func (ec *EnvController) UpdateEnvVar(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的环境变量ID")
return
}
@@ -112,8 +111,8 @@ func (ec *EnvController) UpdateEnvVar(c *gin.Context) {
}
func (ec *EnvController) DeleteEnvVar(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的环境变量ID")
return
}
+2 -2
View File
@@ -19,8 +19,8 @@ func NewExecutorController(executorService *tasks.ExecutorService) *ExecutorCont
}
func (ec *ExecutorController) ExecuteTask(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的任务ID")
return
}
+13 -14
View File
@@ -1,7 +1,6 @@
package controllers
import (
"strconv"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/models"
@@ -19,7 +18,7 @@ func NewLogController() *LogController {
func (lc *LogController) GetLogs(c *gin.Context) {
p := utils.ParsePagination(c)
taskID, _ := strconv.Atoi(c.DefaultQuery("task_id", "0"))
taskID := c.DefaultQuery("task_id", "")
taskName := c.DefaultQuery("task_name", "")
status := c.DefaultQuery("status", "")
@@ -27,7 +26,7 @@ func (lc *LogController) GetLogs(c *gin.Context) {
var total int64
query := database.DB.Model(&models.TaskLog{})
if taskID > 0 {
if taskID != "" {
query = query.Where("task_id = ?", taskID)
}
if status != "" {
@@ -36,7 +35,7 @@ func (lc *LogController) GetLogs(c *gin.Context) {
// 按任务名称过滤
if taskName != "" {
var taskIDs []uint
var taskIDs []string
database.DB.Model(&models.Task{}).Where("name LIKE ?", "%"+taskName+"%").Pluck("id", &taskIDs)
if len(taskIDs) > 0 {
query = query.Where("task_id IN ?", taskIDs)
@@ -49,14 +48,14 @@ func (lc *LogController) GetLogs(c *gin.Context) {
query.Count(&total)
query.Order("id DESC").Offset(p.Offset()).Limit(p.PageSize).Find(&logs)
taskIDList := make([]uint, 0)
taskIDList := make([]string, 0)
for _, log := range logs {
taskIDList = append(taskIDList, log.TaskID)
}
var tasks []models.Task
database.DB.Where("id IN ?", taskIDList).Find(&tasks)
taskMap := make(map[uint]models.Task)
taskMap := make(map[string]models.Task)
for _, t := range tasks {
taskMap[t.ID] = t
}
@@ -87,14 +86,14 @@ func (lc *LogController) GetLogs(c *gin.Context) {
}
func (lc *LogController) GetLogDetail(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的日志ID")
return
}
var log models.TaskLog
if err := database.DB.First(&log, id).Error; err != nil {
if err := database.DB.Where("id = ?", id).First(&log).Error; err != nil {
utils.NotFound(c, "日志不存在")
return
}
@@ -104,7 +103,7 @@ func (lc *LogController) GetLogDetail(c *gin.Context) {
func (lc *LogController) ClearLogs(c *gin.Context) {
var req struct {
TaskID *int `json:"task_id"`
TaskID *string `json:"task_id"`
}
if err := c.ShouldBindJSON(&req); err != nil {
@@ -113,7 +112,7 @@ func (lc *LogController) ClearLogs(c *gin.Context) {
}
query := database.DB.Model(&models.TaskLog{})
if req.TaskID != nil && *req.TaskID > 0 {
if req.TaskID != nil && *req.TaskID != "" {
query = query.Where("task_id = ?", *req.TaskID)
} else {
query = query.Where("1 = 1") // Allow delete all without GORM safety block
@@ -128,13 +127,13 @@ func (lc *LogController) ClearLogs(c *gin.Context) {
}
func (lc *LogController) DeleteLog(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的日志ID")
return
}
if err := database.DB.Delete(&models.TaskLog{}, id).Error; err != nil {
if err := database.DB.Where("id = ?", id).Delete(&models.TaskLog{}).Error; err != nil {
utils.ServerError(c, "删除日志失败")
return
}
+5 -9
View File
@@ -2,7 +2,6 @@ package controllers
import (
"fmt"
"strconv"
"github.com/engigu/baihu-panel/internal/database"
"github.com/engigu/baihu-panel/internal/models"
@@ -25,10 +24,7 @@ func (lc *LogWSController) StreamLog(c *gin.Context) {
return
}
logID, err := strconv.ParseUint(logIDStr, 10, 32)
if err != nil {
return
}
logID := logIDStr
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
@@ -38,7 +34,7 @@ func (lc *LogWSController) StreamLog(c *gin.Context) {
// 1. 检查数据库中是否已结束
var taskLog models.TaskLog
if err := database.DB.First(&taskLog, uint(logID)).Error; err == nil {
if err := database.DB.Where("id = ?", logID).First(&taskLog).Error; err == nil {
if taskLog.Status != "running" {
// 已结束,读取库内日志
content, err := utils.DecompressFromBase64(taskLog.Output)
@@ -52,14 +48,14 @@ func (lc *LogWSController) StreamLog(c *gin.Context) {
}
// 2. 未结束或未找到记录,尝试从 TinyLogManager 获取
tl := tasks.GetActiveLog(uint(logID))
tl := tasks.GetActiveLog(logID)
if tl == nil {
conn.WriteMessage(websocket.TextMessage, []byte("未找到正在运行的任务日志"))
return
}
// 发送系统提示
conn.WriteMessage(websocket.TextMessage, []byte(fmt.Sprintf("[System] 连接成功,正在监听日志... (LogID: %d)\n", logID)))
conn.WriteMessage(websocket.TextMessage, []byte(fmt.Sprintf("[System] 连接成功,正在监听日志... (LogID: %s)\n", logID)))
// 发送最后 100 行
lastLines, err := tl.ReadLastLines(100)
@@ -78,7 +74,7 @@ func (lc *LogWSController) StreamLog(c *gin.Context) {
if !ok {
// 任务结束,尝试刷新最后一次库内完整内容
var finalLog models.TaskLog
if err := database.DB.First(&finalLog, uint(logID)).Error; err == nil {
if err := database.DB.Where("id = ?", logID).First(&finalLog).Error; err == nil {
content, _ := utils.DecompressFromBase64(finalLog.Output)
if content != "" {
conn.WriteMessage(websocket.TextMessage, []byte("\n--- 任务已结束 ---\n"))
+8 -9
View File
@@ -1,7 +1,6 @@
package controllers
import (
"strconv"
"github.com/engigu/baihu-panel/internal/models/vo"
"github.com/engigu/baihu-panel/internal/services"
@@ -19,7 +18,7 @@ func NewScriptController(scriptService *services.ScriptService) *ScriptControlle
}
func (sc *ScriptController) CreateScript(c *gin.Context) {
userID := 1
userID := c.GetString("userID")
var req struct {
Name string `json:"name" binding:"required"`
@@ -36,7 +35,7 @@ func (sc *ScriptController) CreateScript(c *gin.Context) {
}
func (sc *ScriptController) GetScripts(c *gin.Context) {
userID := 1
userID := c.GetString("userID")
scripts := sc.scriptService.GetScriptsByUserID(userID)
vos := vo.ToScriptVOListFromModels(scripts)
for i := range vos {
@@ -46,8 +45,8 @@ func (sc *ScriptController) GetScripts(c *gin.Context) {
}
func (sc *ScriptController) GetScript(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的脚本ID")
return
}
@@ -62,8 +61,8 @@ func (sc *ScriptController) GetScript(c *gin.Context) {
}
func (sc *ScriptController) UpdateScript(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的脚本ID")
return
}
@@ -88,8 +87,8 @@ func (sc *ScriptController) UpdateScript(c *gin.Context) {
}
func (sc *ScriptController) DeleteScript(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的脚本ID")
return
}
+3 -3
View File
@@ -59,9 +59,9 @@ func (sc *SettingsController) ChangePassword(c *gin.Context) {
return
}
// 暂时使用固定用户名 admin
user := sc.userService.GetUserByUsername("admin")
if user == nil {
userID := c.GetString("userID")
var user *models.User
if err := database.DB.Where("id = ?", userID).First(&user).Error; err != nil {
utils.NotFound(c, "用户不存在")
return
}
+23 -27
View File
@@ -2,7 +2,6 @@ package controllers
import (
"path/filepath"
"strconv"
"github.com/engigu/baihu-panel/internal/constant"
"github.com/engigu/baihu-panel/internal/models/vo"
@@ -63,7 +62,7 @@ func (tc *TaskController) CreateTask(c *gin.Context) {
CleanConfig string `json:"clean_config"`
Envs string `json:"envs"`
Languages []map[string]string `json:"languages"`
AgentID *uint `json:"agent_id"`
AgentID *string `json:"agent_id"`
TriggerType string `json:"trigger_type"`
RetryCount int `json:"retry_count"`
RetryInterval int `json:"retry_interval"`
@@ -90,14 +89,14 @@ func (tc *TaskController) CreateTask(c *gin.Context) {
// 转换为绝对路径(Agent 任务保持原样)
workDir := req.WorkDir
if req.AgentID == nil || *req.AgentID == 0 {
if req.AgentID == nil || *req.AgentID == "" {
workDir = resolveWorkDir(req.WorkDir)
}
task := tc.taskService.CreateTask(req.Name, req.Command, req.Schedule, req.Timeout, workDir, req.CleanConfig, req.Envs, req.Type, req.Config, req.AgentID, req.Languages, req.TriggerType, req.Tags, req.RetryCount, req.RetryInterval, req.RandomRange)
// 如果是 Agent 任务,通知 Agent;否则添加到本地 cron
if task.AgentID != nil && *task.AgentID > 0 {
if task.AgentID != nil && *task.AgentID != "" {
tc.agentWSManager.BroadcastTasks(*task.AgentID)
} else {
tc.executorService.AddCronTask(task)
@@ -114,12 +113,9 @@ func (tc *TaskController) GetTasks(c *gin.Context) {
tags := c.DefaultQuery("tags", "")
taskType := c.DefaultQuery("type", "")
var agentID *uint
var agentID *string
if agentIDStr != "" {
if id, err := strconv.ParseUint(agentIDStr, 10, 32); err == nil {
uid := uint(id)
agentID = &uid
}
agentID = &agentIDStr
}
tasks, total := tc.taskService.GetTasksWithPagination(p.Page, p.PageSize, name, agentID, tags, taskType)
@@ -127,8 +123,8 @@ func (tc *TaskController) GetTasks(c *gin.Context) {
}
func (tc *TaskController) GetTask(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的任务ID")
return
}
@@ -143,15 +139,15 @@ func (tc *TaskController) GetTask(c *gin.Context) {
}
func (tc *TaskController) UpdateTask(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的任务ID")
return
}
// 获取旧任务信息(用于判断 agent 变更)
oldTask := tc.taskService.GetTaskByID(id)
var oldAgentID *uint
var oldAgentID *string
if oldTask != nil {
oldAgentID = oldTask.AgentID
}
@@ -169,7 +165,7 @@ func (tc *TaskController) UpdateTask(c *gin.Context) {
Envs string `json:"envs"`
Enabled bool `json:"enabled"`
Languages []map[string]string `json:"languages"`
AgentID *uint `json:"agent_id"`
AgentID *string `json:"agent_id"`
TriggerType string `json:"trigger_type"`
RetryCount int `json:"retry_count"`
RetryInterval int `json:"retry_interval"`
@@ -190,7 +186,7 @@ func (tc *TaskController) UpdateTask(c *gin.Context) {
// 转换为绝对路径(Agent 任务保持原样)
workDir := req.WorkDir
if req.AgentID == nil || *req.AgentID == 0 {
if req.AgentID == nil || *req.AgentID == "" {
workDir = resolveWorkDir(req.WorkDir)
}
@@ -201,12 +197,12 @@ func (tc *TaskController) UpdateTask(c *gin.Context) {
}
// 处理任务调度
if task.AgentID != nil && *task.AgentID > 0 {
if task.AgentID != nil && *task.AgentID != "" {
// Agent 任务:从本地 cron 移除,通知 Agent
tc.executorService.RemoveCronTask(task.ID)
tc.agentWSManager.BroadcastTasks(*task.AgentID)
// 如果 agent 变更了,也通知旧 agent
if oldAgentID != nil && *oldAgentID > 0 && *oldAgentID != *task.AgentID {
if oldAgentID != nil && *oldAgentID != "" && *oldAgentID != *task.AgentID {
tc.agentWSManager.BroadcastTasks(*oldAgentID)
}
} else {
@@ -217,7 +213,7 @@ func (tc *TaskController) UpdateTask(c *gin.Context) {
tc.executorService.RemoveCronTask(task.ID)
}
// 如果之前是 agent 任务,通知旧 agent 移除
if oldAgentID != nil && *oldAgentID > 0 {
if oldAgentID != nil && *oldAgentID != "" {
tc.agentWSManager.BroadcastTasks(*oldAgentID)
}
}
@@ -226,20 +222,20 @@ func (tc *TaskController) UpdateTask(c *gin.Context) {
}
func (tc *TaskController) DeleteTask(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
id := c.Param("id")
if id == "" {
utils.BadRequest(c, "无效的任务ID")
return
}
// 获取任务信息(用于通知 agent)
task := tc.taskService.GetTaskByID(id)
var agentID *uint
var agentID *string
if task != nil {
agentID = task.AgentID
}
tc.executorService.RemoveCronTask(uint(id))
tc.executorService.RemoveCronTask(id)
success := tc.taskService.DeleteTask(id)
if !success {
@@ -248,7 +244,7 @@ func (tc *TaskController) DeleteTask(c *gin.Context) {
}
// 如果是 agent 任务,通知 agent
if agentID != nil && *agentID > 0 {
if agentID != nil && *agentID != "" {
tc.agentWSManager.BroadcastTasks(*agentID)
}
@@ -256,13 +252,13 @@ func (tc *TaskController) DeleteTask(c *gin.Context) {
}
func (tc *TaskController) StopTask(c *gin.Context) {
logID, err := strconv.ParseUint(c.Param("logID"), 10, 32)
if err != nil {
logID := c.Param("logID")
if logID == "" {
utils.BadRequest(c, "无效的日志ID")
return
}
err = tc.executorService.StopTaskExecution(uint(logID))
err := tc.executorService.StopTaskExecution(logID)
if err != nil {
utils.BadRequest(c, err.Error())
return
+5 -7
View File
@@ -85,11 +85,9 @@ func (tc *TerminalController) HandleWebSocket(c *gin.Context) {
}
// Windows 使用 pipe 模式,Unix 使用 PTY 模式
userID := 1
if v, exists := c.Get("userID"); exists {
if id, ok := v.(uint); ok {
userID = int(id)
}
userID := c.GetString("userID")
if userID == "" {
userID = "1" // 兜底
}
if runtime.GOOS == "windows" {
@@ -100,7 +98,7 @@ func (tc *TerminalController) HandleWebSocket(c *gin.Context) {
}
// handlePtyMode 使用 PTY 处理终端(Unix/macOS
func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID int) {
func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) {
// 发送 PTY 模式标识
conn.WriteMessage(websocket.TextMessage, []byte("__PTY_MODE__"))
@@ -174,7 +172,7 @@ func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID int) {
}
// handlePipeMode 使用 pipe 处理终端(Windows
func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID int) {
func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string) {
// 发送 pipe 模式标识
conn.WriteMessage(websocket.TextMessage, []byte("__PIPE_MODE__"))