Initial commit: TaskPool React panel
- React frontend with route-level code splitting - Backend rebranded from Baihu to TaskPool - DB brand migration script and local compatibility
This commit is contained in:
@@ -0,0 +1,746 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
var agentUpgrader = websocket.Upgrader{
|
||||
CheckOrigin: utils.CheckWSOrigin,
|
||||
}
|
||||
|
||||
// AgentController Agent 控制器
|
||||
type AgentController struct {
|
||||
agentService *services.AgentService
|
||||
wsManager *services.AgentWSManager
|
||||
settingsService *services.SettingsService
|
||||
}
|
||||
|
||||
// NewAgentController 创建 Agent 控制器
|
||||
func NewAgentController(settingsService *services.SettingsService) *AgentController {
|
||||
return &AgentController{
|
||||
agentService: services.NewAgentService(),
|
||||
wsManager: services.GetAgentWSManager(),
|
||||
settingsService: settingsService,
|
||||
}
|
||||
}
|
||||
|
||||
// List 获取 Agent 列表
|
||||
func (c *AgentController) List(ctx *gin.Context) {
|
||||
agents := c.agentService.List()
|
||||
utils.Success(ctx, vo.ToAgentVOListFromModels(agents))
|
||||
}
|
||||
|
||||
// getActiveSchedulerConfig 获取 Agent 的实际调度配置(若为空或零值,则使用系统默认的 settings)
|
||||
func (c *AgentController) getActiveSchedulerConfig(agent *models.Agent) map[string]interface{} {
|
||||
workerCount := agent.SchedulerConfig.WorkerCount
|
||||
queueSize := agent.SchedulerConfig.QueueSize
|
||||
rateInterval := int(agent.SchedulerConfig.RateInterval / time.Millisecond)
|
||||
strictQueue := agent.SchedulerConfig.StrictQueue
|
||||
|
||||
// 如果未配置(WorkerCount <= 0),则使用全局系统设置
|
||||
if workerCount <= 0 {
|
||||
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)
|
||||
strictQueue = false
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"worker_count": workerCount,
|
||||
"queue_size": queueSize,
|
||||
"rate_interval": rateInterval,
|
||||
"strict_queue": strictQueue,
|
||||
}
|
||||
}
|
||||
|
||||
// Update 更新 Agent
|
||||
func (c *AgentController) Update(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Description string `json:"description"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SchedulerConfig *vo.AgentSchedulerConfigVO `json:"scheduler_config"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取旧状态
|
||||
oldAgent := c.agentService.GetByID(id)
|
||||
if oldAgent == nil {
|
||||
utils.NotFound(ctx, "Agent 不存在")
|
||||
return
|
||||
}
|
||||
wasEnabled := utils.DerefBool(oldAgent.Enabled, true)
|
||||
|
||||
var schedulerConfig models.AgentSchedulerConfig
|
||||
if req.SchedulerConfig != nil {
|
||||
schedulerConfig.WorkerCount = req.SchedulerConfig.WorkerCount
|
||||
schedulerConfig.QueueSize = req.SchedulerConfig.QueueSize
|
||||
schedulerConfig.RateInterval = time.Duration(req.SchedulerConfig.RateInterval) * time.Millisecond
|
||||
schedulerConfig.Verbose = req.SchedulerConfig.Verbose
|
||||
schedulerConfig.StrictQueue = req.SchedulerConfig.StrictQueue
|
||||
}
|
||||
|
||||
if err := c.agentService.Update(id, req.Name, req.Description, req.Enabled, schedulerConfig); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 如果启用状态发生变化,通知 Agent
|
||||
if wasEnabled != req.Enabled {
|
||||
if req.Enabled {
|
||||
// 启用:发送任务列表
|
||||
c.wsManager.SendToAgent(id, services.WSTypeEnabled, map[string]interface{}{
|
||||
"message": "Agent 已启用",
|
||||
})
|
||||
// 发送任务列表
|
||||
c.wsManager.BroadcastTasks(id)
|
||||
} else {
|
||||
// 禁用:发送禁用消息,Agent 收到后清空任务
|
||||
c.wsManager.SendToAgent(id, services.WSTypeDisabled, map[string]interface{}{
|
||||
"message": "Agent 已禁用",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 推送最新的调度配置给 Agent (如果 Agent 在线)
|
||||
if req.Enabled {
|
||||
// 重新加载已更新的 Agent 信息以获取正确的 SchedulerConfig
|
||||
updatedAgent := c.agentService.GetByID(id)
|
||||
if updatedAgent != nil {
|
||||
c.wsManager.SendToAgent(id, services.WSTypeConnected, map[string]interface{}{
|
||||
"agent_id": id,
|
||||
"name": req.Name,
|
||||
"scheduler_config": c.getActiveSchedulerConfig(updatedAgent),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "更新成功")
|
||||
}
|
||||
|
||||
// Delete 删除 Agent
|
||||
func (c *AgentController) Delete(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.agentService.Delete(id); err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "删除成功")
|
||||
}
|
||||
|
||||
// RegenerateToken 重新生成 Token
|
||||
func (c *AgentController) RegenerateToken(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
token, err := c.agentService.RegenerateToken(id)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"token": token})
|
||||
}
|
||||
|
||||
// ========== Agent API(供 Agent 调用)==========
|
||||
|
||||
// Register Agent 注册(无需认证)
|
||||
func (c *AgentController) Register(ctx *gin.Context) {
|
||||
var req models.AgentRegisterRequest
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" {
|
||||
utils.BadRequest(ctx, "名称不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
ip := ctx.ClientIP()
|
||||
agent, token, err := c.agentService.Register(&req, ip)
|
||||
if err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{
|
||||
"agent_id": agent.ID,
|
||||
"token": token,
|
||||
"message": "注册成功",
|
||||
})
|
||||
}
|
||||
|
||||
// Heartbeat Agent 心跳
|
||||
func (c *AgentController) Heartbeat(ctx *gin.Context) {
|
||||
token := c.getAgentToken(ctx)
|
||||
if token == "" {
|
||||
utils.Unauthorized(ctx, "缺少认证 Token")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Version string `json:"version"`
|
||||
BuildTime string `json:"build_time"`
|
||||
Hostname string `json:"hostname"`
|
||||
OS string `json:"os"`
|
||||
Arch string `json:"arch"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
}
|
||||
ctx.ShouldBindJSON(&req)
|
||||
|
||||
ip := ctx.ClientIP()
|
||||
agent, err := c.agentService.Heartbeat(token, ip, req.Version, req.BuildTime, req.Hostname, req.OS, req.Arch)
|
||||
if err != nil {
|
||||
utils.Unauthorized(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 检查是否需要更新
|
||||
latestVersion := c.agentService.GetLatestVersion()
|
||||
needUpdate := c.agentService.CheckNeedUpdate(req.Version, req.BuildTime)
|
||||
forceUpdate := agent.ForceUpdate
|
||||
|
||||
// 如果强制更新已触发,重置标志
|
||||
if forceUpdate && needUpdate {
|
||||
c.agentService.ClearForceUpdate(agent.ID)
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{
|
||||
"agent_id": agent.ID,
|
||||
"name": agent.Name,
|
||||
"need_update": needUpdate,
|
||||
"force_update": forceUpdate,
|
||||
"latest_version": latestVersion,
|
||||
})
|
||||
}
|
||||
|
||||
// GetTasks Agent 获取任务列表
|
||||
func (c *AgentController) GetTasks(ctx *gin.Context) {
|
||||
token := c.getAgentToken(ctx)
|
||||
if token == "" {
|
||||
utils.Unauthorized(ctx, "缺少认证 Token")
|
||||
return
|
||||
}
|
||||
|
||||
// 先尝试通过 token 查找 Agent
|
||||
agent := c.agentService.GetByToken(token)
|
||||
|
||||
// 如果找不到,尝试验证令牌并通过 machine_id 查找
|
||||
if agent == nil {
|
||||
machineID := ctx.GetHeader("X-Machine-ID")
|
||||
if machineID != "" {
|
||||
// 验证令牌是否有效
|
||||
if _, err := c.agentService.ValidateToken(token); err == nil {
|
||||
// 令牌有效,尝试通过 machine_id 查找 Agent
|
||||
agent = c.agentService.GetByMachineID(machineID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if agent == nil {
|
||||
utils.Unauthorized(ctx, "无效的 Token")
|
||||
return
|
||||
}
|
||||
|
||||
if !utils.DerefBool(agent.Enabled, true) {
|
||||
utils.Forbidden(ctx, "Agent 已禁用")
|
||||
return
|
||||
}
|
||||
|
||||
tasks := c.agentService.GetTasks(agent.ID)
|
||||
utils.Success(ctx, gin.H{
|
||||
"agent_id": agent.ID,
|
||||
"tasks": tasks,
|
||||
})
|
||||
}
|
||||
|
||||
// ReportResult Agent 上报执行结果
|
||||
func (c *AgentController) ReportResult(ctx *gin.Context) {
|
||||
token := c.getAgentToken(ctx)
|
||||
if token == "" {
|
||||
utils.Unauthorized(ctx, "缺少认证 Token")
|
||||
return
|
||||
}
|
||||
|
||||
agent := c.agentService.GetByToken(token)
|
||||
if agent == nil {
|
||||
utils.Unauthorized(ctx, "无效的 Token")
|
||||
return
|
||||
}
|
||||
|
||||
if !utils.DerefBool(agent.Enabled, true) {
|
||||
utils.Forbidden(ctx, "Agent 已禁用")
|
||||
return
|
||||
}
|
||||
|
||||
var result models.AgentTaskResult
|
||||
if err := ctx.ShouldBindJSON(&result); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
result.AgentID = agent.ID
|
||||
|
||||
if err := c.agentService.ReportResult(&result); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "上报成功")
|
||||
}
|
||||
|
||||
// getAgentToken 从请求头获取 Agent Token
|
||||
func (c *AgentController) getAgentToken(ctx *gin.Context) string {
|
||||
auth := ctx.GetHeader("Authorization")
|
||||
if auth == "" {
|
||||
return ""
|
||||
}
|
||||
// Bearer <token>
|
||||
parts := strings.SplitN(auth, " ", 2)
|
||||
if len(parts) == 2 && parts[0] == "Bearer" {
|
||||
return parts[1]
|
||||
}
|
||||
return auth
|
||||
}
|
||||
|
||||
// Download 下载 Agent 程序
|
||||
func (c *AgentController) Download(ctx *gin.Context) {
|
||||
osType := ctx.DefaultQuery("os", "linux")
|
||||
arch := ctx.DefaultQuery("arch", "amd64")
|
||||
|
||||
data, filename, err := c.agentService.GetAgentBinary(osType, arch)
|
||||
if err != nil {
|
||||
utils.NotFound(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx.Header("Content-Disposition", "attachment; filename="+filename)
|
||||
ctx.Header("Content-Type", "application/gzip")
|
||||
ctx.Header("Content-Length", strconv.Itoa(len(data)))
|
||||
ctx.Data(200, "application/gzip", data)
|
||||
}
|
||||
|
||||
// GetVersion 获取 Agent 最新版本信息
|
||||
func (c *AgentController) GetVersion(ctx *gin.Context) {
|
||||
version := c.agentService.GetLatestVersion()
|
||||
platforms := c.agentService.GetAvailablePlatforms()
|
||||
|
||||
utils.Success(ctx, gin.H{
|
||||
"version": version,
|
||||
"platforms": platforms,
|
||||
})
|
||||
}
|
||||
|
||||
// ForceUpdate 强制更新指定 Agent
|
||||
func (c *AgentController) ForceUpdate(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.agentService.SetForceUpdate(id); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "已标记强制更新,Agent 下次心跳时将自动更新")
|
||||
}
|
||||
|
||||
// ========== WebSocket ==========
|
||||
|
||||
// WSConnect Agent WebSocket 连接
|
||||
func (c *AgentController) WSConnect(ctx *gin.Context) {
|
||||
// 添加 panic 恢复
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf("[AgentWS] WSConnect panic: %v", r)
|
||||
ctx.JSON(http.StatusInternalServerError, gin.H{"error": "服务器内部错误"})
|
||||
}
|
||||
}()
|
||||
|
||||
ip := ctx.ClientIP()
|
||||
|
||||
// 打印请求信息用于调试
|
||||
logger.Infof("[AgentWS] 收到连接请求: IP=%s, URL=%s", ip, ctx.Request.URL.String())
|
||||
|
||||
// 检查 IP 限流
|
||||
if allowed, reason := c.wsManager.CheckRateLimit(ip); !allowed {
|
||||
logger.Warnf("[AgentWS] IP %s 被限流: %s", ip, reason)
|
||||
ctx.JSON(http.StatusTooManyRequests, gin.H{"error": reason})
|
||||
return
|
||||
}
|
||||
|
||||
token := ctx.Query("token")
|
||||
if token == "" {
|
||||
c.wsManager.RecordConnectFail(ip)
|
||||
logger.Warnf("[AgentWS] 连接失败: 缺少 token, IP=%s", ip)
|
||||
ctx.JSON(http.StatusUnauthorized, gin.H{"error": "缺少 token"})
|
||||
return
|
||||
}
|
||||
|
||||
machineID := ctx.Query("machine_id")
|
||||
logger.Infof("[AgentWS] Token: %s..., MachineID: %s...", token[:8], machineID[:16])
|
||||
|
||||
isNewAgent := false
|
||||
|
||||
// 先尝试用 token 查找已有 Agent
|
||||
agent := c.agentService.GetByToken(token)
|
||||
logger.Infof("[AgentWS] GetByToken 结果: agent=%v", agent != nil)
|
||||
|
||||
// 如果没找到,尝试用令牌注册(会检查 machine_id 是否已存在)
|
||||
if agent == nil {
|
||||
logger.Infof("[AgentWS] 尝试注册新 Agent")
|
||||
var err error
|
||||
agent, isNewAgent, err = c.agentService.RegisterByToken(token, machineID, ip)
|
||||
if err != nil {
|
||||
c.wsManager.RecordConnectFail(ip)
|
||||
logger.Warnf("[AgentWS] 注册失败: %v, IP=%s, token=%s", err, ip, token[:8]+"...")
|
||||
ctx.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
logger.Infof("[AgentWS] 注册成功: Agent #%s, isNew=%v", agent.ID, isNewAgent)
|
||||
}
|
||||
|
||||
if !utils.DerefBool(agent.Enabled, true) {
|
||||
c.wsManager.RecordConnectFail(ip)
|
||||
logger.Warnf("[AgentWS] Agent #%s 已禁用, IP=%s", agent.ID, ip)
|
||||
ctx.JSON(http.StatusForbidden, gin.H{"error": "Agent 已禁用"})
|
||||
return
|
||||
}
|
||||
|
||||
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 #%s, IP=%s", err, agent.ID, ip)
|
||||
return
|
||||
}
|
||||
conn.SetReadLimit(constant.MaxMessageSize)
|
||||
|
||||
// 连接成功,重置失败计数
|
||||
c.wsManager.RecordConnectSuccess(ip)
|
||||
|
||||
// 注册连接
|
||||
ac := c.wsManager.Register(agent.ID, conn, ip)
|
||||
|
||||
// 更新 Agent 状态
|
||||
c.agentService.Heartbeat(token, ip, "", "", "", "", "")
|
||||
|
||||
// 获取调度配置并发送连接成功消息(包含注册状态和调度配置)
|
||||
schedCfg := c.getActiveSchedulerConfig(agent)
|
||||
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": schedCfg,
|
||||
})
|
||||
|
||||
logger.Infof("[AgentWS] Agent #%s 连接成功 (配置: %v)", agent.ID, schedCfg)
|
||||
|
||||
// 启动读写协程
|
||||
go c.wsWritePump(ac)
|
||||
go c.wsReadPump(ac, agent)
|
||||
|
||||
// 主动推送任务列表
|
||||
go c.wsManager.BroadcastTasks(agent.ID)
|
||||
}
|
||||
|
||||
// wsReadPump 读取消息
|
||||
func (c *AgentController) wsReadPump(ac *services.AgentConnection, agent *models.Agent) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf("[AgentWS] Agent #%s wsReadPump panic: %v", agent.ID, r)
|
||||
}
|
||||
logger.Infof("[AgentWS] Agent #%s wsReadPump 退出", agent.ID)
|
||||
c.wsManager.Unregister(agent.ID, ac)
|
||||
}()
|
||||
|
||||
// 检查连接是否有效(可能是旧连接被新连接替换)
|
||||
if ac == nil || ac.IsClosed() {
|
||||
return
|
||||
}
|
||||
|
||||
ac.SetReadDeadline(time.Now().Add(90 * time.Second))
|
||||
// 注意:SetPongHandler 需要直接访问 Conn,但这里我们在连接建立后立即设置
|
||||
// 所以是安全的,因为此时连接还没有被其他 goroutine 关闭
|
||||
ac.Conn.SetPongHandler(func(string) error {
|
||||
ac.SetReadDeadline(time.Now().Add(90 * time.Second))
|
||||
return nil
|
||||
})
|
||||
|
||||
for {
|
||||
_, message, err := ac.ReadMessage()
|
||||
if err != nil {
|
||||
logger.Warnf("[AgentWS] Agent #%s 读取错误: %v", agent.ID, err)
|
||||
break
|
||||
}
|
||||
|
||||
var msg services.WSMessage
|
||||
if err := json.Unmarshal(message, &msg); err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
c.handleWSMessage(ac, agent, &msg)
|
||||
}
|
||||
}
|
||||
|
||||
// wsWritePump 写入消息
|
||||
func (c *AgentController) wsWritePump(ac *services.AgentConnection) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf("[AgentWS] Agent #%s wsWritePump panic: %v", ac.AgentID, r)
|
||||
}
|
||||
logger.Infof("[AgentWS] Agent #%s wsWritePump 退出", ac.AgentID)
|
||||
}()
|
||||
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case message, ok := <-ac.Send:
|
||||
if !ok {
|
||||
logger.Warnf("[AgentWS] Agent #%s Send channel 已关闭", ac.AgentID)
|
||||
return
|
||||
}
|
||||
if ac.IsClosed() {
|
||||
logger.Warnf("[AgentWS] Agent #%s 连接已关闭(write)", ac.AgentID)
|
||||
return
|
||||
}
|
||||
if err := ac.WriteMessage(message); err != nil {
|
||||
logger.Warnf("[AgentWS] Agent #%s 写入消息失败: %v", ac.AgentID, err)
|
||||
return
|
||||
}
|
||||
case <-ticker.C:
|
||||
if ac.IsClosed() {
|
||||
return
|
||||
}
|
||||
if err := ac.WritePing(); err != nil {
|
||||
logger.Warnf("[AgentWS] Agent #%s 发送 Ping 失败: %v", ac.AgentID, err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// handleWSMessage 处理 WebSocket 消息
|
||||
func (c *AgentController) handleWSMessage(ac *services.AgentConnection, agent *models.Agent, msg *services.WSMessage) {
|
||||
switch msg.Type {
|
||||
case services.WSTypeHeartbeat:
|
||||
c.handleHeartbeat(ac, agent, msg.Data)
|
||||
|
||||
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(_ *models.Agent, data json.RawMessage) {
|
||||
var req struct {
|
||||
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 != "" {
|
||||
logger.Infof("[AgentWS] 收到任务心跳: LogID=%s, Duration=%dms", req.LogID, req.Duration)
|
||||
c.agentService.UpdateTaskDuration(req.LogID, req.Duration)
|
||||
}
|
||||
}
|
||||
|
||||
// handleFetchTasks 处理 Agent 请求任务列表
|
||||
func (c *AgentController) handleFetchTasks(agent *models.Agent) {
|
||||
tasks := c.agentService.GetTasks(agent.ID)
|
||||
c.wsManager.SendToAgent(agent.ID, services.WSTypeTasks, map[string]interface{}{
|
||||
"tasks": tasks,
|
||||
})
|
||||
logger.Infof("[AgentWS] Agent #%s 请求任务列表,返回 %d 个任务", agent.ID, len(tasks))
|
||||
}
|
||||
|
||||
// handleHeartbeat 处理心跳
|
||||
func (c *AgentController) handleHeartbeat(ac *services.AgentConnection, agent *models.Agent, data json.RawMessage) {
|
||||
var req struct {
|
||||
Version string `json:"version"`
|
||||
BuildTime string `json:"build_time"`
|
||||
Hostname string `json:"hostname"`
|
||||
OS string `json:"os"`
|
||||
Arch string `json:"arch"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
}
|
||||
json.Unmarshal(data, &req)
|
||||
|
||||
ac.UpdatePing()
|
||||
|
||||
// 更新 Agent 信息(使用连接时保存的 IP)
|
||||
c.agentService.Heartbeat(agent.Token, ac.IP, req.Version, req.BuildTime, req.Hostname, req.OS, req.Arch)
|
||||
|
||||
// 检查是否需要更新
|
||||
latestVersion := c.agentService.GetLatestVersion()
|
||||
needUpdate := c.agentService.CheckNeedUpdate(req.Version, req.BuildTime)
|
||||
forceUpdate := agent.ForceUpdate
|
||||
|
||||
if forceUpdate && needUpdate {
|
||||
c.agentService.ClearForceUpdate(agent.ID)
|
||||
}
|
||||
|
||||
// 发送心跳响应
|
||||
response := map[string]interface{}{
|
||||
"agent_id": agent.ID,
|
||||
"name": agent.Name,
|
||||
"need_update": needUpdate,
|
||||
"force_update": forceUpdate,
|
||||
"latest_version": latestVersion,
|
||||
}
|
||||
c.wsManager.SendToAgent(agent.ID, services.WSTypeHeartbeatAck, response)
|
||||
}
|
||||
|
||||
// handleTaskResult 处理任务结果
|
||||
func (c *AgentController) handleTaskResult(agent *models.Agent, data json.RawMessage) {
|
||||
var result models.AgentTaskResult
|
||||
if err := json.Unmarshal(data, &result); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
result.AgentID = agent.ID
|
||||
c.agentService.ReportResult(&result)
|
||||
}
|
||||
|
||||
// handleTaskLog 处理 Agent 发送的实时日志
|
||||
func (c *AgentController) handleTaskLog(_ *models.Agent, data json.RawMessage) {
|
||||
var logMsg struct {
|
||||
LogID string `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] 收到任务日志 but could not find active TinyLog: LogID=%s, ContentSize=%d", logMsg.LogID, len(logMsg.Content))
|
||||
}
|
||||
}
|
||||
|
||||
// NotifyTaskUpdate 通知 Agent 任务更新
|
||||
func (c *AgentController) NotifyTaskUpdate(agentID string) {
|
||||
c.wsManager.BroadcastTasks(agentID)
|
||||
}
|
||||
|
||||
// ========== 令牌管理 ==========
|
||||
|
||||
// ListTokens 获取令牌列表
|
||||
func (c *AgentController) ListTokens(ctx *gin.Context) {
|
||||
tokens := c.agentService.ListTokens()
|
||||
utils.Success(ctx, vo.ToAgentTokenVOListFromModels(tokens))
|
||||
}
|
||||
|
||||
// CreateToken 创建令牌
|
||||
func (c *AgentController) CreateToken(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Remark string `json:"remark"`
|
||||
MaxUses int `json:"max_uses"`
|
||||
ExpiresAt string `json:"expires_at"` // 格式: 2006-01-02 15:04:05
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
var expiresAt *time.Time
|
||||
if req.ExpiresAt != "" {
|
||||
t, err := time.ParseInLocation("2006-01-02 15:04:05", req.ExpiresAt, time.Local)
|
||||
if err != nil {
|
||||
utils.BadRequest(ctx, "过期时间格式错误")
|
||||
return
|
||||
}
|
||||
expiresAt = &t
|
||||
}
|
||||
|
||||
token, err := c.agentService.CreateToken(req.Remark, req.MaxUses, expiresAt)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, vo.ToAgentTokenVO(token))
|
||||
}
|
||||
|
||||
// DeleteToken 删除令牌
|
||||
func (c *AgentController) DeleteToken(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.agentService.DeleteToken(id); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type AppLogController struct {
|
||||
appLogService *services.AppLogService
|
||||
}
|
||||
|
||||
func NewAppLogController() *AppLogController {
|
||||
return &AppLogController{
|
||||
appLogService: services.NewAppLogService(),
|
||||
}
|
||||
}
|
||||
|
||||
// GetLogs 获取应用日志列表
|
||||
func (ac *AppLogController) GetLogs(c *gin.Context) {
|
||||
p := utils.ParsePagination(c)
|
||||
category := c.Query("category")
|
||||
status := c.Query("status")
|
||||
level := c.Query("level")
|
||||
keyword := c.Query("keyword")
|
||||
|
||||
logs, total, err := ac.appLogService.List(category, status, level, p.Page, p.PageSize, keyword)
|
||||
if err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
utils.PaginatedResponse(c, logs, total, p)
|
||||
}
|
||||
|
||||
// MarkAsRead 标记已读
|
||||
func (ac *AppLogController) MarkAsRead(c *gin.Context) {
|
||||
var req struct {
|
||||
ID string `json:"id"`
|
||||
Category string `json:"category"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.ID != "" {
|
||||
if err := ac.appLogService.MarkAsRead(req.ID); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
} else if req.Category != "" {
|
||||
if err := ac.appLogService.MarkAllAsRead(req.Category); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
} else {
|
||||
utils.BadRequest(c, "id 或 category 必须提供")
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(c, "标记成功")
|
||||
}
|
||||
|
||||
// ClearLogs 清理日志
|
||||
func (ac *AppLogController) ClearLogs(c *gin.Context) {
|
||||
var req struct {
|
||||
Category string `json:"category"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := ac.appLogService.Clear(req.Category); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(c, "清理成功")
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
"github.com/engigu/taskpool/internal/middleware"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type AuthController struct {
|
||||
userService *services.UserService
|
||||
settingsService *services.SettingsService
|
||||
loginLogService *services.LoginLogService
|
||||
}
|
||||
|
||||
type loginAttempt struct {
|
||||
Count int
|
||||
LastAttempt time.Time
|
||||
}
|
||||
|
||||
var loginAttempts sync.Map
|
||||
|
||||
func init() {
|
||||
// 定期清理过期的登录尝试统计,防止内存溢出
|
||||
go func() {
|
||||
ticker := time.NewTicker(30 * time.Minute)
|
||||
for range ticker.C {
|
||||
loginAttempts.Range(func(key, value any) bool {
|
||||
attempt := value.(*loginAttempt)
|
||||
if time.Since(attempt.LastAttempt) > 10*time.Minute {
|
||||
loginAttempts.Delete(key)
|
||||
}
|
||||
return true
|
||||
})
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func NewAuthController(userService *services.UserService, settingsService *services.SettingsService, loginLogService *services.LoginLogService) *AuthController {
|
||||
return &AuthController{
|
||||
userService: userService,
|
||||
settingsService: settingsService,
|
||||
loginLogService: loginLogService,
|
||||
}
|
||||
}
|
||||
|
||||
func (ac *AuthController) Login(c *gin.Context) {
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
|
||||
ip := c.ClientIP()
|
||||
userAgent := c.GetHeader("User-Agent")
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 暴力破解防御
|
||||
if val, ok := loginAttempts.Load(ip); ok {
|
||||
attempt := val.(*loginAttempt)
|
||||
if attempt.Count >= 5 && time.Since(attempt.LastAttempt) < time.Minute {
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventBruteForceLogin,
|
||||
Payload: map[string]interface{}{
|
||||
"ip": ip,
|
||||
"username": req.Username,
|
||||
"userAgent": userAgent,
|
||||
},
|
||||
})
|
||||
utils.TooManyRequests(c, "尝试次数过多,请一分钟后再试")
|
||||
return
|
||||
}
|
||||
// 如果距离上次尝试已超过一分钟,重置计数
|
||||
if time.Since(attempt.LastAttempt) >= time.Minute {
|
||||
loginAttempts.Delete(ip)
|
||||
}
|
||||
}
|
||||
|
||||
user := ac.userService.GetUserByUsername(req.Username)
|
||||
if user == nil || !ac.userService.ValidatePassword(user, req.Password) {
|
||||
// 记录失败尝试
|
||||
val, _ := loginAttempts.LoadOrStore(ip, &loginAttempt{Count: 0, LastAttempt: time.Now()})
|
||||
attempt := val.(*loginAttempt)
|
||||
attempt.Count++
|
||||
attempt.LastAttempt = time.Now()
|
||||
|
||||
// 记录登录失败日志
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventUserLogin,
|
||||
Payload: map[string]interface{}{
|
||||
"ip": ip,
|
||||
"username": req.Username,
|
||||
"userAgent": userAgent,
|
||||
"status": "failed",
|
||||
"message": "用户名或密码错误",
|
||||
},
|
||||
})
|
||||
utils.Unauthorized(c, "用户名或密码错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 登录成功,清除尝试记录
|
||||
loginAttempts.Delete(ip)
|
||||
|
||||
// 获取 cookie 过期天数
|
||||
expireDays := 7
|
||||
if days := ac.settingsService.Get(constant.SectionSite, constant.KeyCookieDays); days != "" {
|
||||
if d, err := strconv.Atoi(days); err == nil && d > 0 {
|
||||
expireDays = d
|
||||
}
|
||||
}
|
||||
|
||||
// 生成 token
|
||||
token, err := utils.GenerateToken(user.ID, user.Username, user.TokenVersion, expireDays, constant.Secret)
|
||||
if err != nil {
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventUserLogin,
|
||||
Payload: map[string]interface{}{
|
||||
"ip": ip,
|
||||
"username": req.Username,
|
||||
"userAgent": userAgent,
|
||||
"status": "failed",
|
||||
"message": "Token生成失败",
|
||||
},
|
||||
})
|
||||
utils.ServerError(c, "登录失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 设置 Cookie
|
||||
middleware.SetAuthCookie(c, token, expireDays)
|
||||
|
||||
// 记录登录成功日志
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventUserLogin,
|
||||
Payload: map[string]interface{}{
|
||||
"ip": ip,
|
||||
"username": req.Username,
|
||||
"userAgent": userAgent,
|
||||
"status": "success",
|
||||
"message": "登录成功",
|
||||
},
|
||||
})
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"user": user.Username,
|
||||
})
|
||||
}
|
||||
|
||||
func (ac *AuthController) Logout(c *gin.Context) {
|
||||
if userID, exists := c.Get("userID"); exists {
|
||||
ac.userService.InvalidateUserTokens(userID.(string))
|
||||
}
|
||||
middleware.ClearAuthCookie(c)
|
||||
utils.SuccessMsg(c, "退出成功")
|
||||
}
|
||||
|
||||
func (ac *AuthController) GetCurrentUser(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
user, err := ac.userService.GetUserByID(userID)
|
||||
if err != nil {
|
||||
utils.Unauthorized(c, "会话无效")
|
||||
return
|
||||
}
|
||||
utils.Success(c, gin.H{
|
||||
"username": user.Username,
|
||||
"role": user.Role,
|
||||
})
|
||||
}
|
||||
|
||||
func (ac *AuthController) Register(c *gin.Context) {
|
||||
/*
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Email string `json:"email" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 安全性:强制设定角色为 user,防止注册时篡改角色为 admin
|
||||
user := ac.userService.CreateUser(req.Username, req.Password, req.Email, constant.DefaultRole)
|
||||
utils.Success(c, vo.ToUserVO(user))
|
||||
*/
|
||||
utils.BadRequest(c, "注册功能已关闭")
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type DashboardController struct {
|
||||
executorService *tasks.ExecutorService
|
||||
}
|
||||
|
||||
func NewDashboardController(executorService *tasks.ExecutorService) *DashboardController {
|
||||
return &DashboardController{
|
||||
executorService: executorService,
|
||||
}
|
||||
}
|
||||
|
||||
type StatsResponse struct {
|
||||
Tasks int64 `json:"tasks"`
|
||||
TodayExecs int64 `json:"today_execs"`
|
||||
Envs int64 `json:"envs"`
|
||||
Logs int64 `json:"logs"`
|
||||
Scheduled int `json:"scheduled"`
|
||||
Running int `json:"running"`
|
||||
}
|
||||
|
||||
func (dc *DashboardController) GetStats(c *gin.Context) {
|
||||
var taskCount, envCount, logCount, todayExecs int64
|
||||
|
||||
database.DB.Model(&models.Task{}).Count(&taskCount)
|
||||
database.DB.Model(&models.EnvironmentVariable{}).Count(&envCount)
|
||||
database.DB.Model(&models.TaskLog{}).Count(&logCount)
|
||||
|
||||
// 今日执行总数
|
||||
today := time.Now().Format("2006-01-02")
|
||||
database.DB.Model(&models.SendStats{}).Where("day = ?", today).Select("COALESCE(SUM(num), 0)").Scan(&todayExecs)
|
||||
|
||||
// 调度统计:本地调度 + Agent 调度
|
||||
// 本地调度:agent_id 为 NULL 且 enabled = true 的任务
|
||||
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)
|
||||
|
||||
// 正在运行:目前只能统计本地运行的任务
|
||||
// Agent 端的运行状态需要通过心跳上报(未来优化)
|
||||
running := dc.executorService.GetRunningCount()
|
||||
|
||||
stats := StatsResponse{
|
||||
Tasks: taskCount,
|
||||
TodayExecs: todayExecs,
|
||||
Envs: envCount,
|
||||
Logs: logCount,
|
||||
Scheduled: totalScheduled,
|
||||
Running: running,
|
||||
}
|
||||
|
||||
utils.Success(c, stats)
|
||||
}
|
||||
|
||||
// GetSentence 获取随机古诗词
|
||||
func (dc *DashboardController) GetSentence(c *gin.Context) {
|
||||
utils.Success(c, gin.H{
|
||||
"sentence": constant.GetRandomSentence(),
|
||||
})
|
||||
}
|
||||
|
||||
// DailyStats 每日统计数据
|
||||
type DailyStats struct {
|
||||
Day string `json:"day"`
|
||||
Total int `json:"total"`
|
||||
Success int `json:"success"`
|
||||
Failed int `json:"failed"`
|
||||
}
|
||||
|
||||
// GetSendStats 获取发送统计
|
||||
func (dc *DashboardController) GetSendStats(c *gin.Context) {
|
||||
// 获取天数参数,默认30天
|
||||
days := 30
|
||||
if d := c.Query("days"); d != "" {
|
||||
if parsed, err := utils.ParseInt(d); err == nil && parsed > 0 && parsed <= 90 {
|
||||
days = parsed
|
||||
}
|
||||
}
|
||||
|
||||
// 获取日期范围
|
||||
now := time.Now()
|
||||
startDay := now.AddDate(0, 0, -(days - 1)).Format("2006-01-02")
|
||||
|
||||
var stats []models.SendStats
|
||||
database.DB.Where("day >= ?", startDay).Find(&stats)
|
||||
|
||||
// 按日期聚合
|
||||
dayMap := make(map[string]*DailyStats)
|
||||
for _, s := range stats {
|
||||
if _, ok := dayMap[s.Day]; !ok {
|
||||
dayMap[s.Day] = &DailyStats{Day: s.Day}
|
||||
}
|
||||
ds := dayMap[s.Day]
|
||||
ds.Total += s.Num
|
||||
if s.Status == constant.TaskStatusSuccess {
|
||||
ds.Success += s.Num
|
||||
} else {
|
||||
ds.Failed += s.Num
|
||||
}
|
||||
}
|
||||
|
||||
// 填充缺失的日期
|
||||
result := make([]DailyStats, 0, days)
|
||||
for i := days - 1; i >= 0; i-- {
|
||||
day := now.AddDate(0, 0, -i).Format("2006-01-02")
|
||||
if ds, ok := dayMap[day]; ok {
|
||||
result = append(result, *ds)
|
||||
} else {
|
||||
result = append(result, DailyStats{Day: day})
|
||||
}
|
||||
}
|
||||
|
||||
// 按日期排序
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
return result[i].Day < result[j].Day
|
||||
})
|
||||
|
||||
utils.Success(c, result)
|
||||
}
|
||||
|
||||
// TaskStats 任务执行统计
|
||||
type TaskStats struct {
|
||||
TaskID string `json:"task_id"`
|
||||
TaskName string `json:"task_name"`
|
||||
Count int `json:"count"`
|
||||
}
|
||||
|
||||
// GetTaskStats 获取任务执行占比
|
||||
func (dc *DashboardController) GetTaskStats(c *gin.Context) {
|
||||
// 获取天数参数,默认30天
|
||||
days := 30
|
||||
if d := c.Query("days"); d != "" {
|
||||
if parsed, err := utils.ParseInt(d); err == nil && parsed > 0 && parsed <= 90 {
|
||||
days = parsed
|
||||
}
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
startDay := now.AddDate(0, 0, -(days - 1)).Format("2006-01-02")
|
||||
|
||||
// 按 task_id 聚合统计
|
||||
var results []struct {
|
||||
TaskID string
|
||||
Total int
|
||||
}
|
||||
database.DB.Model(&models.SendStats{}).
|
||||
Select("task_id, SUM(num) as total").
|
||||
Where("day >= ?", startDay).
|
||||
Group("task_id").
|
||||
Order("total DESC").
|
||||
Find(&results)
|
||||
|
||||
// 获取任务名称
|
||||
taskIDs := make([]string, 0, len(results))
|
||||
for _, r := range results {
|
||||
taskIDs = append(taskIDs, r.TaskID)
|
||||
}
|
||||
|
||||
var tasks []models.Task
|
||||
if len(taskIDs) > 0 {
|
||||
database.DB.Where("id IN ?", taskIDs).Find(&tasks)
|
||||
}
|
||||
taskNameMap := make(map[string]string)
|
||||
for _, t := range tasks {
|
||||
taskNameMap[t.ID] = t.Name
|
||||
}
|
||||
|
||||
// 构建结果
|
||||
stats := make([]TaskStats, 0, len(results))
|
||||
for _, r := range results {
|
||||
name := taskNameMap[r.TaskID]
|
||||
if name == "" {
|
||||
name = "未知任务"
|
||||
}
|
||||
stats = append(stats, TaskStats{
|
||||
TaskID: r.TaskID,
|
||||
TaskName: name,
|
||||
Count: r.Total,
|
||||
})
|
||||
}
|
||||
|
||||
utils.Success(c, stats)
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type DataController struct {
|
||||
dataService *services.DataService
|
||||
taskController *TaskController
|
||||
envController *EnvController
|
||||
}
|
||||
|
||||
func NewDataController(tc *TaskController, ec *EnvController) *DataController {
|
||||
return &DataController{
|
||||
dataService: services.NewDataService(),
|
||||
taskController: tc,
|
||||
envController: ec,
|
||||
}
|
||||
}
|
||||
|
||||
// ExportBusinessData 导出业务数据
|
||||
func (dc *DataController) ExportBusinessData(c *gin.Context) {
|
||||
var req struct {
|
||||
TaskIDs []string `json:"task_ids"`
|
||||
EnvIDs []string `json:"env_ids"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
exportData := dc.dataService.ExportBusinessData(req.TaskIDs, req.EnvIDs)
|
||||
utils.Success(c, exportData)
|
||||
}
|
||||
|
||||
// ImportBusinessData 导入业务数据
|
||||
func (dc *DataController) ImportBusinessData(c *gin.Context) {
|
||||
var req models.ExportData
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.Version == "" {
|
||||
utils.BadRequest(c, "无效的导入数据格式")
|
||||
return
|
||||
}
|
||||
|
||||
// 停止相关的定时任务
|
||||
if len(req.Tasks) > 0 {
|
||||
for _, task := range req.Tasks {
|
||||
dc.taskController.executorService.RemoveCronTask(task.ID)
|
||||
dc.taskController.executorService.GetScheduler().StopTask(task.ID)
|
||||
}
|
||||
}
|
||||
|
||||
// 导入数据
|
||||
if err := dc.dataService.ImportBusinessData(&req); err != nil {
|
||||
utils.ServerError(c, "导入失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 重新启动任务和通知相关的代理
|
||||
if len(req.Tasks) > 0 {
|
||||
for i := range req.Tasks {
|
||||
task := &req.Tasks[i]
|
||||
if utils.DerefBool(task.Enabled, true) && (task.AgentID == nil || *task.AgentID == "") {
|
||||
dc.taskController.executorService.AddCronTask(task)
|
||||
}
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
dc.taskController.agentWSManager.BroadcastTasks(*task.AgentID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "导入成功")
|
||||
}
|
||||
@@ -0,0 +1,421 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/deps"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type DependencyController struct {
|
||||
service *services.DependencyService
|
||||
}
|
||||
|
||||
func NewDependencyController() *DependencyController {
|
||||
return &DependencyController{
|
||||
service: services.NewDependencyService(),
|
||||
}
|
||||
}
|
||||
|
||||
// List 获取依赖列表
|
||||
func (c *DependencyController) List(ctx *gin.Context) {
|
||||
language := ctx.Query("language")
|
||||
langVersion := ctx.Query("lang_version")
|
||||
deps, err := c.service.List(language, langVersion)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取依赖列表失败")
|
||||
return
|
||||
}
|
||||
vos := vo.ToDependencyVOListFromModels(deps)
|
||||
utils.Success(ctx, vos)
|
||||
}
|
||||
|
||||
// Create 添加依赖
|
||||
func (c *DependencyController) Create(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language" binding:"required"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
dep := &models.Dependency{
|
||||
Name: req.Name,
|
||||
Version: req.Version,
|
||||
Language: req.Language,
|
||||
LangVersion: req.LangVersion,
|
||||
Remark: req.Remark,
|
||||
}
|
||||
|
||||
if err := c.service.Create(dep); err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, vo.ToDependencyVO(dep))
|
||||
}
|
||||
|
||||
// Delete 删除依赖
|
||||
func (c *DependencyController) Delete(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.service.Delete(id); err != nil {
|
||||
utils.ServerError(ctx, "删除失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "删除成功")
|
||||
}
|
||||
|
||||
func (c *DependencyController) Install(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
language := req.Language
|
||||
if language == "" {
|
||||
language = ctx.Query("language")
|
||||
}
|
||||
langVersion := req.LangVersion
|
||||
if langVersion == "" {
|
||||
langVersion = ctx.Query("lang_version")
|
||||
}
|
||||
|
||||
dep := &models.Dependency{
|
||||
Name: req.Name,
|
||||
Version: req.Version,
|
||||
Language: language,
|
||||
LangVersion: langVersion,
|
||||
Remark: req.Remark,
|
||||
}
|
||||
|
||||
err := c.service.Install(dep)
|
||||
// 无论成功失败,都同步记录日志
|
||||
c.service.Create(dep)
|
||||
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(ctx, "安装成功")
|
||||
}
|
||||
|
||||
// GetInstallCommand 获取安装命令
|
||||
func (c *DependencyController) GetInstallCommand(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
language := req.Language
|
||||
if language == "" {
|
||||
language = ctx.Query("language")
|
||||
}
|
||||
langVersion := req.LangVersion
|
||||
if langVersion == "" {
|
||||
langVersion = ctx.Query("lang_version")
|
||||
}
|
||||
|
||||
dep := &models.Dependency{
|
||||
Name: req.Name,
|
||||
Version: req.Version,
|
||||
Language: language,
|
||||
LangVersion: langVersion,
|
||||
}
|
||||
|
||||
cmd, err := c.service.GetInstallCommand(dep)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"command": cmd})
|
||||
}
|
||||
|
||||
// GetReinstallAllCommand 获取全部重装命令
|
||||
func (c *DependencyController) GetReinstallAllCommand(ctx *gin.Context) {
|
||||
language := ctx.Query("language")
|
||||
langVersion := ctx.Query("lang_version")
|
||||
if language == "" {
|
||||
utils.BadRequest(ctx, "缺少 language 参数")
|
||||
return
|
||||
}
|
||||
|
||||
cmd, err := c.service.GetReinstallAllCommand(language, langVersion)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"command": cmd})
|
||||
}
|
||||
|
||||
// Uninstall 卸载依赖
|
||||
func (c *DependencyController) Uninstall(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
force := ctx.Query("force") == "true"
|
||||
|
||||
// 获取依赖信息
|
||||
deps, _ := c.service.List("", "")
|
||||
var dep *models.Dependency
|
||||
for i := range deps {
|
||||
if deps[i].ID == id {
|
||||
dep = &deps[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if dep == nil {
|
||||
utils.NotFound(ctx, "依赖不存在")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.service.Uninstall(dep); err != nil {
|
||||
if !force {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 卸载成功(或强制删除)后从数据库删除
|
||||
c.service.Delete(id)
|
||||
|
||||
utils.SuccessMsg(ctx, "卸载成功")
|
||||
}
|
||||
|
||||
// Reinstall 重新安装依赖
|
||||
func (c *DependencyController) Reinstall(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取依赖信息
|
||||
deps, _ := c.service.List("", "")
|
||||
var dep *models.Dependency
|
||||
for i := range deps {
|
||||
if deps[i].ID == id {
|
||||
dep = &deps[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if dep == nil {
|
||||
utils.NotFound(ctx, "依赖不存在")
|
||||
return
|
||||
}
|
||||
|
||||
err := c.service.Install(dep)
|
||||
// 无论成功失败,都同步记录日志
|
||||
c.service.Create(dep)
|
||||
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(ctx, "重新安装成功")
|
||||
}
|
||||
|
||||
// ReinstallAll 重新安装所有依赖
|
||||
func (c *DependencyController) ReinstallAll(ctx *gin.Context) {
|
||||
language := ctx.Query("language")
|
||||
langVersion := ctx.Query("lang_version")
|
||||
if language == "" {
|
||||
utils.BadRequest(ctx, "缺少 language 参数")
|
||||
return
|
||||
}
|
||||
|
||||
deps, err := c.service.List(language, langVersion)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取依赖列表失败")
|
||||
return
|
||||
}
|
||||
|
||||
var failed []string
|
||||
for i := range deps {
|
||||
d := &deps[i]
|
||||
err := c.service.Install(d)
|
||||
if err != nil {
|
||||
failed = append(failed, d.Name)
|
||||
}
|
||||
// 无论成功失败,都同步记录日志到数据库
|
||||
c.service.Create(d)
|
||||
}
|
||||
|
||||
if len(failed) > 0 {
|
||||
utils.ServerError(ctx, "部分包安装失败: "+strings.Join(failed, ", "))
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "全部重新安装成功")
|
||||
}
|
||||
|
||||
// GetInstalled 获取已安装的包
|
||||
func (c *DependencyController) GetInstalled(ctx *gin.Context) {
|
||||
language := ctx.Query("language")
|
||||
langVersion := ctx.Query("lang_version")
|
||||
if language == "" {
|
||||
utils.BadRequest(ctx, "缺少 language 参数")
|
||||
return
|
||||
}
|
||||
|
||||
packages, err := c.service.GetInstalledPackages(language, langVersion)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取已安装包失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, packages)
|
||||
}
|
||||
|
||||
// GetBatchInstallCommand 获取批量安装依赖包的命令
|
||||
func (c *DependencyController) GetBatchInstallCommand(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Items []struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language" binding:"required"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
} `json:"items" binding:"required,gt=0"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: items 不能为空且必须包含 name 和 language")
|
||||
return
|
||||
}
|
||||
|
||||
var depsList []models.Dependency
|
||||
for _, item := range req.Items {
|
||||
depsList = append(depsList, models.Dependency{
|
||||
Name: item.Name,
|
||||
Version: item.Version,
|
||||
Language: item.Language,
|
||||
LangVersion: item.LangVersion,
|
||||
})
|
||||
}
|
||||
|
||||
cmd, err := c.service.GetBatchInstallCommand(depsList)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"command": cmd})
|
||||
}
|
||||
|
||||
// ParseAndImport 解析上传/粘贴的清单文件内容并批量导入至数据库
|
||||
func (c *DependencyController) ParseAndImport(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Language string `json:"language" binding:"required"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
ImportDB bool `json:"import_db"` // 是否持久化到数据库做可视化管理
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: language 和 content 必填")
|
||||
return
|
||||
}
|
||||
|
||||
// 1. 解析文本清单内容
|
||||
parsedDeps, err := deps.ParseManifest(req.Language, req.Content)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "清单文件解析失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if len(parsedDeps) == 0 {
|
||||
utils.BadRequest(ctx, "未解析到任何有效依赖包")
|
||||
return
|
||||
}
|
||||
|
||||
// 2. 补全语言和版本属性
|
||||
for i := range parsedDeps {
|
||||
parsedDeps[i].Language = req.Language
|
||||
parsedDeps[i].LangVersion = req.LangVersion
|
||||
}
|
||||
|
||||
// 3. 根据需求决定是否导入数据库
|
||||
var finalDeps []models.Dependency
|
||||
if req.ImportDB {
|
||||
imported, err := c.service.ImportDependencies(parsedDeps)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "导入依赖记录至数据库失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
finalDeps = imported
|
||||
} else {
|
||||
finalDeps = parsedDeps
|
||||
}
|
||||
|
||||
// 4. 为这一批包生成合并批量安装命令
|
||||
cmd, err := c.service.GetBatchInstallCommand(finalDeps)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "生成安装命令失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{
|
||||
"dependencies": vo.ToDependencyVOListFromModels(finalDeps),
|
||||
"command": cmd,
|
||||
})
|
||||
}
|
||||
|
||||
// GetDepInstallCommand 获取自动补全的命令,返回给前端执行
|
||||
func (c *DependencyController) GetDepInstallCommand(ctx *gin.Context) {
|
||||
logID := ctx.Query("log_id")
|
||||
if logID == "" {
|
||||
utils.BadRequest(ctx, "参数错误: log_id 不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
execPath, err := os.Executable()
|
||||
if err != nil {
|
||||
execPath = "taskpool" // 兜底
|
||||
}
|
||||
|
||||
// 构造命令,比如: "F:\workspace\taskpool\taskpool.exe" depinstall <log_id>
|
||||
cmdStr := fmt.Sprintf("%q depinstall %s", execPath, logID)
|
||||
utils.Success(ctx, gin.H{
|
||||
"command": cmdStr,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,389 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/relation"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type EnvController struct {
|
||||
envService *services.EnvService
|
||||
}
|
||||
|
||||
func NewEnvController(envService *services.EnvService) *EnvController {
|
||||
return &EnvController{envService: envService}
|
||||
}
|
||||
|
||||
// GetSecretStatus 获取加密秘钥状态
|
||||
// @Summary 获取加密秘钥状态
|
||||
// @Description 返回系统是否已配置加密秘钥
|
||||
// @Tags Env
|
||||
// @Produce json
|
||||
// @Success 200 {object} utils.Response{data=bool} "成功"
|
||||
// @Router /env/secret-status [get]
|
||||
// @Security BearerAuth
|
||||
func (ec *EnvController) GetSecretStatus(c *gin.Context) {
|
||||
utils.Success(c, utils.IsSecretKeySet())
|
||||
}
|
||||
|
||||
// CreateEnvVar 创建环境变量
|
||||
// @Summary 创建环境变量
|
||||
// @Description 创建一个新的环境变量
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param body body object true "环境变量信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.EnvVO}
|
||||
// @Router /env [post]
|
||||
func (ec *EnvController) CreateEnvVar(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Value string `json:"value" binding:"required"`
|
||||
Remark string `json:"remark"`
|
||||
Type string `json:"type"`
|
||||
Hidden *bool `json:"hidden"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Tags string `json:"tags"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.Type == "" {
|
||||
req.Type = constant.EnvTypeNormal
|
||||
}
|
||||
|
||||
hidden := true
|
||||
if req.Hidden != nil {
|
||||
hidden = *req.Hidden
|
||||
}
|
||||
|
||||
enabled := true
|
||||
if req.Enabled != nil {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
|
||||
envVar := ec.envService.CreateEnvVar(req.Name, req.Value, req.Remark, req.Type, hidden, enabled, userID)
|
||||
if envVar != nil {
|
||||
relation.DataRelation.SaveTags(envVar.ID, constant.RelationTypeEnvTag, req.Tags)
|
||||
envVar.Tags = req.Tags
|
||||
}
|
||||
|
||||
// Broadcast tasks to all agents because global envs changed
|
||||
services.GetAgentWSManager().BroadcastTasksToAll()
|
||||
|
||||
utils.Success(c, vo.ToEnvVO(envVar))
|
||||
}
|
||||
|
||||
// GetEnvVars 获取环境变量列表
|
||||
// @Summary 获取环境变量列表
|
||||
// @Description 分页获取环境变量列表,支持按名称筛选
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param name query string false "按名称模糊查询"
|
||||
// @Param page query int false "页码"
|
||||
// @Param page_size query int false "每页数量"
|
||||
// @Param type query string false "按类型筛选"
|
||||
// @Param tags query string false "按标签筛选"
|
||||
// @Success 200 {object} utils.Response{data=utils.PaginationData{data=[]vo.EnvVO}}
|
||||
// @Router /env [get]
|
||||
func (ec *EnvController) GetEnvVars(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
p := utils.ParsePagination(c)
|
||||
name := c.DefaultQuery("name", "")
|
||||
envType := c.DefaultQuery("type", "")
|
||||
tags := c.DefaultQuery("tags", "")
|
||||
envVars, total := ec.envService.GetEnvVarsWithPagination(userID, name, envType, tags, p.Page, p.PageSize)
|
||||
utils.PaginatedResponse(c, vo.ToEnvVOListFromModels(envVars), total, p)
|
||||
}
|
||||
|
||||
// GetAllEnvVars 获取所有环境变量
|
||||
// @Summary 获取所有环境变量
|
||||
// @Description 获取当前用户的所有环境变量(不分页)
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Success 200 {object} utils.Response{data=[]vo.EnvVO}
|
||||
// @Router /env/all [get]
|
||||
func (ec *EnvController) GetAllEnvVars(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
envVars := ec.envService.GetEnvVarsByUserID(userID)
|
||||
utils.Success(c, vo.ToEnvVOListFromModels(envVars))
|
||||
}
|
||||
|
||||
// GetEnvVar 获取环境变量详情
|
||||
// @Summary 获取环境变量详情
|
||||
// @Description 根据 ID 获取环境变量详情
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "环境变量ID"
|
||||
// @Success 200 {object} utils.Response{data=vo.EnvVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /env/{id} [get]
|
||||
func (ec *EnvController) GetEnvVar(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的环境变量ID")
|
||||
return
|
||||
}
|
||||
|
||||
envVar := ec.envService.GetEnvVarByID(id)
|
||||
if envVar == nil {
|
||||
utils.NotFound(c, "环境变量不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToEnvVO(envVar))
|
||||
}
|
||||
|
||||
// UpdateEnvVar 更新环境变量
|
||||
// @Summary 更新环境变量
|
||||
// @Description 根据 ID 更新环境变量信息
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "环境变量ID"
|
||||
// @Param body body object true "环境变量更新信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.EnvVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /env/{id} [put]
|
||||
func (ec *EnvController) UpdateEnvVar(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的环境变量ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name"`
|
||||
Value string `json:"value"`
|
||||
Remark string `json:"remark"`
|
||||
Type string `json:"type"`
|
||||
Hidden *bool `json:"hidden"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Tags string `json:"tags"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.Type == "" {
|
||||
req.Type = constant.EnvTypeNormal
|
||||
}
|
||||
|
||||
// 对于更新,获取现有数据
|
||||
existing := ec.envService.GetEnvVarByID(id)
|
||||
if existing == nil {
|
||||
utils.NotFound(c, "环境变量不存在")
|
||||
return
|
||||
}
|
||||
|
||||
hidden := existing.Hidden
|
||||
if req.Hidden != nil {
|
||||
hidden = req.Hidden
|
||||
}
|
||||
|
||||
enabled := existing.Enabled
|
||||
if req.Enabled != nil {
|
||||
enabled = req.Enabled
|
||||
}
|
||||
|
||||
envVar := ec.envService.UpdateEnvVar(id, req.Name, req.Value, req.Remark, req.Type, utils.DerefBool(hidden, true), utils.DerefBool(enabled, true))
|
||||
if envVar == nil {
|
||||
utils.NotFound(c, "环境变量不存在")
|
||||
return
|
||||
}
|
||||
|
||||
relation.DataRelation.SaveTags(envVar.ID, constant.RelationTypeEnvTag, req.Tags)
|
||||
envVar.Tags = req.Tags
|
||||
|
||||
// Broadcast tasks to all agents because global envs changed
|
||||
services.GetAgentWSManager().BroadcastTasksToAll()
|
||||
|
||||
utils.Success(c, vo.ToEnvVO(envVar))
|
||||
}
|
||||
|
||||
// DeleteEnvVar 删除环境变量
|
||||
// @Summary 删除环境变量
|
||||
// @Description 根据 ID 删除环境变量
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "环境变量ID"
|
||||
// @Param force query boolean false "强制删除(忽略任务关联)"
|
||||
// @Success 200 {object} utils.Response
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Failure 409 {object} utils.Response{data=[]vo.TaskVO}
|
||||
// @Router /env/{id} [delete]
|
||||
func (ec *EnvController) DeleteEnvVar(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的环境变量ID")
|
||||
return
|
||||
}
|
||||
|
||||
force := c.Query("force") == "true"
|
||||
success, associatedTasks := ec.envService.DeleteEnvVar(id, force)
|
||||
|
||||
if len(associatedTasks) > 0 {
|
||||
c.JSON(200, utils.Response{
|
||||
Code: 409,
|
||||
Msg: "该环境变量已被任务引用,请先在任务中删除引用或选择强制删除",
|
||||
Data: vo.ToTaskVOListFromModels(associatedTasks),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if !success {
|
||||
utils.NotFound(c, "环境变量不存在或删除失败")
|
||||
return
|
||||
}
|
||||
|
||||
// Broadcast tasks to all agents because global envs changed
|
||||
services.GetAgentWSManager().BroadcastTasksToAll()
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
// GetAssociatedTasks 获取关联任务
|
||||
// @Summary 获取关联任务
|
||||
// @Description 获取引用了该环境变量的任务列表
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "环境变量ID"
|
||||
// @Success 200 {object} utils.Response{data=[]vo.TaskVO}
|
||||
// @Router /env/{id}/tasks [get]
|
||||
func (ec *EnvController) GetAssociatedTasks(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的环境变量ID")
|
||||
return
|
||||
}
|
||||
tasks := ec.envService.GetAssociatedTasks(id)
|
||||
utils.Success(c, vo.ToTaskVOListFromModels(tasks))
|
||||
}
|
||||
|
||||
// GetTags 获取所有环境变量标签
|
||||
// @Summary 获取所有环境变量标签
|
||||
// @Description 获取所有环境变量中使用的标签列表
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Success 200 {object} utils.Response{data=[]string}
|
||||
// @Router /env/tags [get]
|
||||
func (ec *EnvController) GetTags(c *gin.Context) {
|
||||
tags, err := ec.envService.GetAllEnvTags()
|
||||
if err != nil {
|
||||
utils.ServerError(c, "获取标签失败")
|
||||
return
|
||||
}
|
||||
utils.Success(c, tags)
|
||||
}
|
||||
|
||||
// BulkSaveEnv 批量保存环境变量
|
||||
func (ec *EnvController) BulkSaveEnv(c *gin.Context) {
|
||||
var reqs []struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name" binding:"required"`
|
||||
Value string `json:"value" binding:"required"`
|
||||
Remark string `json:"remark"`
|
||||
Type string `json:"type"`
|
||||
Hidden *bool `json:"hidden"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&reqs); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
userID := c.GetString("userID")
|
||||
|
||||
for _, req := range reqs {
|
||||
if req.Type == constant.EnvTypeSecret {
|
||||
continue // 二次严苛拦截,机密变量不应下发/保存
|
||||
}
|
||||
|
||||
hidden := true
|
||||
if req.Hidden != nil {
|
||||
hidden = *req.Hidden
|
||||
}
|
||||
enabled := true
|
||||
if req.Enabled != nil {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
|
||||
var existingEnv *models.EnvironmentVariable
|
||||
// 优先按 ID 匹配
|
||||
if req.ID != "" {
|
||||
var e models.EnvironmentVariable
|
||||
if err := database.DB.Where("id = ?", req.ID).First(&e).Error; err == nil {
|
||||
existingEnv = &e
|
||||
}
|
||||
}
|
||||
// 如果 ID 没找到,按 Name 匹配
|
||||
if existingEnv == nil {
|
||||
var e models.EnvironmentVariable
|
||||
if err := database.DB.Where("name = ?", req.Name).First(&e).Error; err == nil {
|
||||
existingEnv = &e
|
||||
}
|
||||
}
|
||||
|
||||
if existingEnv != nil {
|
||||
existingEnv.Name = req.Name
|
||||
existingEnv.Value = models.BigText(req.Value)
|
||||
existingEnv.Remark = req.Remark
|
||||
existingEnv.Type = req.Type
|
||||
existingEnv.Hidden = &hidden
|
||||
existingEnv.Enabled = &enabled
|
||||
database.DB.Save(existingEnv)
|
||||
|
||||
if req.ID != "" && existingEnv.ID != req.ID {
|
||||
database.DB.Model(existingEnv).Update("id", req.ID)
|
||||
}
|
||||
} else {
|
||||
envVar := &models.EnvironmentVariable{
|
||||
ID: req.ID,
|
||||
Name: req.Name,
|
||||
Value: models.BigText(req.Value),
|
||||
Remark: req.Remark,
|
||||
Type: req.Type,
|
||||
Hidden: &hidden,
|
||||
Enabled: &enabled,
|
||||
UserID: userID,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
if envVar.ID == "" {
|
||||
envVar.ID = utils.GenerateID()
|
||||
}
|
||||
database.DB.Create(envVar)
|
||||
}
|
||||
}
|
||||
|
||||
services.GetAgentWSManager().BroadcastTasksToAll()
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type ExecutorController struct {
|
||||
executorService *tasks.ExecutorService
|
||||
}
|
||||
|
||||
func NewExecutorController(executorService *tasks.ExecutorService) *ExecutorController {
|
||||
return &ExecutorController{executorService: executorService}
|
||||
}
|
||||
|
||||
// ExecuteTask 运行任务
|
||||
// @Summary 运行任务
|
||||
// @Description 立即执行指定的任务
|
||||
// @Tags 任务执行
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "任务ID"
|
||||
// @Param body body object false "执行参数 (envs: 环境变量字典)"
|
||||
// @Success 200 {object} utils.Response{data=vo.ExecutionResultVO}
|
||||
// @Failure 400 {object} utils.Response
|
||||
// @Router /execute/task/{id} [post]
|
||||
func (ec *ExecutorController) ExecuteTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Envs map[string]string `json:"envs"`
|
||||
}
|
||||
// 尝试绑定 JSON 体,但不强制要求
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
|
||||
var extraEnvs []string
|
||||
if req.Envs != nil {
|
||||
for k, v := range req.Envs {
|
||||
extraEnvs = append(extraEnvs, k+"="+v)
|
||||
}
|
||||
}
|
||||
|
||||
result := ec.executorService.ExecuteTask(id, extraEnvs)
|
||||
utils.Success(c, vo.ToExecutionResultVO(result))
|
||||
}
|
||||
|
||||
// ExecuteCommand 执行命令
|
||||
func (ec *ExecutorController) ExecuteCommand(c *gin.Context) {
|
||||
var req struct {
|
||||
Command string `json:"command" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
result := ec.executorService.ExecuteCommand(req.Command)
|
||||
utils.Success(c, vo.ToExecutionResultVO(result))
|
||||
}
|
||||
|
||||
// GetLastResults 获取最新执行结果
|
||||
// @Summary 获取最新执行结果
|
||||
// @Description 获取最新任务或命令执行的结果列表
|
||||
// @Tags 任务执行
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param count query int false "数量 (默认 10)"
|
||||
// @Success 200 {object} utils.Response{data=[]vo.ExecutionResultVO}
|
||||
// @Router /execute/results [get]
|
||||
func (ec *ExecutorController) GetLastResults(c *gin.Context) {
|
||||
count := 10
|
||||
if c.Query("count") != "" {
|
||||
if parsedCount, err := strconv.Atoi(c.Query("count")); err == nil && parsedCount > 0 {
|
||||
count = parsedCount
|
||||
}
|
||||
}
|
||||
|
||||
results := ec.executorService.GetLastResults(count)
|
||||
utils.Success(c, vo.ToExecutionResultVOList(results))
|
||||
}
|
||||
@@ -0,0 +1,550 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
var (
|
||||
extractZip = utils.ExtractZip
|
||||
extractTar = utils.ExtractTar
|
||||
extractTarGz = utils.ExtractTarGz
|
||||
)
|
||||
|
||||
type FileController struct {
|
||||
workDir string
|
||||
}
|
||||
|
||||
func NewFileController(workDir string) *FileController {
|
||||
os.MkdirAll(workDir, 0755)
|
||||
absPath, err := filepath.Abs(workDir)
|
||||
if err != nil {
|
||||
absPath = workDir
|
||||
}
|
||||
return &FileController{workDir: absPath}
|
||||
}
|
||||
|
||||
type FileNode struct {
|
||||
Name string `json:"name"`
|
||||
Path string `json:"path"`
|
||||
IsDir bool `json:"isDir"`
|
||||
ModTime int64 `json:"modTime"`
|
||||
Children []*FileNode `json:"children,omitempty"`
|
||||
}
|
||||
|
||||
// checkPath 校验路径是否在工作目录内且安全。
|
||||
// 它返回完整的绝对路径以及一个表示路径是否安全的布尔值。
|
||||
func (fc *FileController) checkPath(path string, allowRoot bool) (string, bool) {
|
||||
fullPath := filepath.Join(fc.workDir, filepath.Clean(path))
|
||||
rel, err := filepath.Rel(fc.workDir, fullPath)
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
|
||||
// 基础的目录穿越检查
|
||||
if rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
|
||||
return "", false
|
||||
}
|
||||
|
||||
// 根目录检查
|
||||
if !allowRoot && rel == "." {
|
||||
return "", false
|
||||
}
|
||||
|
||||
return fullPath, true
|
||||
}
|
||||
|
||||
func (fc *FileController) GetFileTree(c *gin.Context) {
|
||||
root := &FileNode{
|
||||
Name: filepath.Base(fc.workDir),
|
||||
Path: "",
|
||||
IsDir: true,
|
||||
Children: []*FileNode{},
|
||||
}
|
||||
|
||||
err := filepath.WalkDir(fc.workDir, func(path string, d fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
if path == fc.workDir {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 过滤 __pycache__ 文件夹
|
||||
if d.IsDir() && d.Name() == "__pycache__" {
|
||||
return filepath.SkipDir
|
||||
}
|
||||
|
||||
relPath, _ := filepath.Rel(fc.workDir, path)
|
||||
parts := strings.Split(relPath, string(filepath.Separator))
|
||||
|
||||
info, err := d.Info()
|
||||
var modTime int64
|
||||
if err == nil {
|
||||
modTime = info.ModTime().UnixMilli()
|
||||
}
|
||||
|
||||
current := root
|
||||
for i, part := range parts {
|
||||
found := false
|
||||
for _, child := range current.Children {
|
||||
if child.Name == part {
|
||||
current = child
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
isLast := i == len(parts)-1
|
||||
isDir := !isLast || d.IsDir()
|
||||
node := &FileNode{
|
||||
Name: part,
|
||||
Path: strings.Join(parts[:i+1], "/"),
|
||||
IsDir: isDir,
|
||||
ModTime: modTime,
|
||||
}
|
||||
if isDir {
|
||||
node.Children = []*FileNode{}
|
||||
}
|
||||
current.Children = append(current.Children, node)
|
||||
current = node
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, root.Children)
|
||||
}
|
||||
|
||||
func (fc *FileController) GetFileContent(c *gin.Context) {
|
||||
filePath := c.Query("path")
|
||||
if filePath == "" {
|
||||
utils.BadRequest(c, "path参数必填")
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(filePath, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(fullPath)
|
||||
if err != nil {
|
||||
utils.NotFound(c, "文件不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"path": filePath,
|
||||
"content": string(content),
|
||||
})
|
||||
}
|
||||
|
||||
func (fc *FileController) SaveFileContent(c *gin.Context) {
|
||||
var req struct {
|
||||
Path string `json:"path" binding:"required"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(req.Path, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
os.MkdirAll(filepath.Dir(fullPath), 0755)
|
||||
|
||||
if err := os.WriteFile(fullPath, []byte(req.Content), 0644); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
func (fc *FileController) CreateFile(c *gin.Context) {
|
||||
var req struct {
|
||||
Path string `json:"path" binding:"required"`
|
||||
IsDir bool `json:"isDir"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(req.Path, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if req.IsDir {
|
||||
if err := os.MkdirAll(fullPath, 0755); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
} else {
|
||||
os.MkdirAll(filepath.Dir(fullPath), 0755)
|
||||
if err := os.WriteFile(fullPath, []byte(""), 0644); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "创建成功")
|
||||
}
|
||||
|
||||
func (fc *FileController) DeleteFile(c *gin.Context) {
|
||||
var req struct {
|
||||
Path string `json:"path" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(req.Path, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if err := os.RemoveAll(fullPath); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
func (fc *FileController) MoveFile(c *gin.Context) {
|
||||
var req struct {
|
||||
OldPath string `json:"oldPath" binding:"required"`
|
||||
NewPath string `json:"newPath" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
oldFull, oldSafe := fc.checkPath(req.OldPath, false)
|
||||
newFull, newSafe := fc.checkPath(req.NewPath, false)
|
||||
|
||||
if !oldSafe || !newSafe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if oldFull == newFull {
|
||||
utils.Success(c, nil)
|
||||
return
|
||||
}
|
||||
|
||||
// 检查目标是否存在
|
||||
if _, err := os.Stat(newFull); err == nil {
|
||||
utils.BadRequest(c, "目标已存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 确保目标目录存在
|
||||
os.MkdirAll(filepath.Dir(newFull), 0755)
|
||||
|
||||
if err := os.Rename(oldFull, newFull); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
func (fc *FileController) CopyFile(c *gin.Context) {
|
||||
var req struct {
|
||||
SourcePath string `json:"sourcePath" binding:"required"`
|
||||
TargetPath string `json:"targetPath" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
sourceFull, sourceSafe := fc.checkPath(req.SourcePath, false)
|
||||
targetFull, targetSafe := fc.checkPath(req.TargetPath, false)
|
||||
|
||||
if !sourceSafe || !targetSafe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if sourceFull == targetFull {
|
||||
utils.Success(c, nil)
|
||||
return
|
||||
}
|
||||
|
||||
// Read content
|
||||
content, err := os.ReadFile(sourceFull)
|
||||
if err != nil {
|
||||
utils.NotFound(c, "源文件不存在或无法读取")
|
||||
return
|
||||
}
|
||||
|
||||
// 确保目标目录存在
|
||||
os.MkdirAll(filepath.Dir(targetFull), 0755)
|
||||
|
||||
// 检查目标是否存在
|
||||
if _, err := os.Stat(targetFull); err == nil {
|
||||
utils.BadRequest(c, "目标已存在")
|
||||
return
|
||||
}
|
||||
|
||||
if err := os.WriteFile(targetFull, content, 0644); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
func (fc *FileController) RenameFile(c *gin.Context) {
|
||||
var req struct {
|
||||
OldPath string `json:"oldPath" binding:"required"`
|
||||
NewPath string `json:"newPath" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 校验:重命名禁止跨目录
|
||||
if filepath.Dir(filepath.Clean(req.OldPath)) != filepath.Dir(filepath.Clean(req.NewPath)) {
|
||||
utils.BadRequest(c, "禁止跨目录重命名")
|
||||
return
|
||||
}
|
||||
|
||||
oldFull, oldSafe := fc.checkPath(req.OldPath, false)
|
||||
newFull, newSafe := fc.checkPath(req.NewPath, false)
|
||||
|
||||
if !oldSafe || !newSafe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if oldFull == newFull {
|
||||
utils.Success(c, nil)
|
||||
return
|
||||
}
|
||||
|
||||
// 检查目标是否存在
|
||||
if _, err := os.Stat(newFull); err == nil {
|
||||
utils.BadRequest(c, "文件已存在")
|
||||
return
|
||||
}
|
||||
|
||||
if err := os.Rename(oldFull, newFull); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
// UploadArchive 处理归档文件的上传和解压
|
||||
func (fc *FileController) UploadArchive(c *gin.Context) {
|
||||
targetDir := c.PostForm("path")
|
||||
|
||||
file, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
utils.BadRequest(c, "请选择文件")
|
||||
return
|
||||
}
|
||||
|
||||
// 检查文件类型
|
||||
ext := strings.ToLower(filepath.Ext(file.Filename))
|
||||
if ext != ".zip" && ext != ".tar" && ext != ".gz" && ext != ".tgz" {
|
||||
utils.BadRequest(c, "仅支持 zip、tar、gz、tgz 格式")
|
||||
return
|
||||
}
|
||||
|
||||
// 确定解压目标目录
|
||||
extractDir, safe := fc.checkPath(targetDir, true)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
os.MkdirAll(extractDir, 0755)
|
||||
|
||||
// 保存临时文件
|
||||
// 安全修复:使用 filepath.Base 提取纯文件名,防止路径穿越攻击
|
||||
tempFile := filepath.Join(os.TempDir(), filepath.Base(file.Filename))
|
||||
if err := c.SaveUploadedFile(file, tempFile); err != nil {
|
||||
utils.ServerError(c, "保存文件失败")
|
||||
return
|
||||
}
|
||||
defer os.Remove(tempFile)
|
||||
|
||||
// 解压文件
|
||||
var extractErr error
|
||||
switch {
|
||||
case ext == ".zip":
|
||||
extractErr = extractZip(tempFile, extractDir)
|
||||
case ext == ".tar":
|
||||
extractErr = extractTar(tempFile, extractDir)
|
||||
case ext == ".gz" || ext == ".tgz":
|
||||
extractErr = extractTarGz(tempFile, extractDir)
|
||||
}
|
||||
|
||||
if extractErr != nil {
|
||||
utils.ServerError(c, "解压失败: "+extractErr.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "导入成功")
|
||||
}
|
||||
|
||||
// UploadFiles 处理多个文件的上传
|
||||
func (fc *FileController) UploadFiles(c *gin.Context) {
|
||||
targetDir := c.PostForm("path")
|
||||
|
||||
// 确定目标目录
|
||||
destDir, safe := fc.checkPath(targetDir, true)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
os.MkdirAll(destDir, 0755)
|
||||
|
||||
form, err := c.MultipartForm()
|
||||
if err != nil {
|
||||
utils.BadRequest(c, "请选择文件")
|
||||
return
|
||||
}
|
||||
|
||||
files := form.File["files"]
|
||||
paths := form.Value["paths"] // 相对路径数组,用于保持文件夹结构
|
||||
|
||||
if len(files) == 0 {
|
||||
utils.BadRequest(c, "请选择文件")
|
||||
return
|
||||
}
|
||||
|
||||
for i, file := range files {
|
||||
// 获取相对路径(如果有)
|
||||
// 安全修复:清理文件名
|
||||
relPath := filepath.Base(file.Filename)
|
||||
if i < len(paths) && paths[i] != "" {
|
||||
relPath = paths[i]
|
||||
}
|
||||
|
||||
// 构建完整路径
|
||||
fullPath, safe := fc.checkPath(filepath.Join(targetDir, relPath), false)
|
||||
if !safe {
|
||||
continue
|
||||
}
|
||||
|
||||
// 确保父目录存在
|
||||
os.MkdirAll(filepath.Dir(fullPath), 0755)
|
||||
|
||||
// 保存文件
|
||||
if err := c.SaveUploadedFile(file, fullPath); err != nil {
|
||||
utils.ServerError(c, "保存文件失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "上传成功")
|
||||
}
|
||||
|
||||
func (fc *FileController) DownloadFile(c *gin.Context) {
|
||||
filePath := c.Query("path")
|
||||
if filePath == "" {
|
||||
utils.BadRequest(c, "path参数必填")
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(filePath, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
info, err := os.Stat(fullPath)
|
||||
if err != nil || info.IsDir() {
|
||||
utils.NotFound(c, "文件不存在")
|
||||
return
|
||||
}
|
||||
|
||||
c.Header("Content-Description", "File Transfer")
|
||||
c.Header("Content-Transfer-Encoding", "binary")
|
||||
c.Header("Content-Disposition", "attachment; filename="+filepath.Base(fullPath))
|
||||
c.Header("Content-Type", "application/octet-stream")
|
||||
c.File(fullPath)
|
||||
}
|
||||
|
||||
func (fc *FileController) DownloadZip(c *gin.Context) {
|
||||
paths := c.QueryArray("path")
|
||||
if len(paths) == 0 || c.ContentType() == "application/json" {
|
||||
var req struct {
|
||||
Paths []string `json:"paths"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err == nil && len(paths) == 0 {
|
||||
paths = req.Paths
|
||||
}
|
||||
}
|
||||
|
||||
if len(paths) == 0 {
|
||||
utils.BadRequest(c, "path参数必填")
|
||||
return
|
||||
}
|
||||
|
||||
validatedAbsPaths := make([]string, 0, len(paths))
|
||||
for _, path := range paths {
|
||||
fullPath, safe := fc.checkPath(path, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if _, err := os.Stat(fullPath); err != nil {
|
||||
utils.NotFound(c, "文件不存在")
|
||||
return
|
||||
}
|
||||
|
||||
validatedAbsPaths = append(validatedAbsPaths, fullPath)
|
||||
}
|
||||
|
||||
fileName := "taskpool-export-" + time.Now().Format("20060102-150405") + ".zip"
|
||||
if len(validatedAbsPaths) == 1 {
|
||||
fileName = filepath.Base(validatedAbsPaths[0]) + ".zip"
|
||||
}
|
||||
|
||||
c.Header("Content-Description", "File Transfer")
|
||||
c.Header("Content-Transfer-Encoding", "binary")
|
||||
c.Header("Content-Disposition", "attachment; filename="+fileName)
|
||||
c.Header("Content-Type", "application/zip")
|
||||
|
||||
if err := utils.CreateZip(c.Writer, validatedAbsPaths); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type InstallController struct {
|
||||
installService *services.InstallService
|
||||
}
|
||||
|
||||
func NewInstallController() *InstallController {
|
||||
return &InstallController{
|
||||
installService: services.NewInstallService(),
|
||||
}
|
||||
}
|
||||
|
||||
// GetInstallStatus 获取安装状态
|
||||
// @Summary 获取安装状态
|
||||
// @Description 检查系统是否已完成安装
|
||||
// @Tags 安装
|
||||
// @Produce json
|
||||
// @Success 200 {object} services.InstallStatus
|
||||
// @Router /api/v1/install/status [get]
|
||||
func (c *InstallController) GetInstallStatus(ctx *gin.Context) {
|
||||
status, err := c.installService.CheckInstallStatus()
|
||||
if err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
ctx.JSON(http.StatusOK, status)
|
||||
}
|
||||
|
||||
// Install 执行安装
|
||||
// @Summary 执行安装
|
||||
// @Description 初始化系统配置和管理员账号
|
||||
// @Tags 安装
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body services.InstallRequest true "安装请求"
|
||||
// @Success 200 {object} map[string]interface{}
|
||||
// @Failure 400 {object} map[string]string
|
||||
// @Failure 500 {object} map[string]string
|
||||
// @Router /api/v1/install [post]
|
||||
func (c *InstallController) Install(ctx *gin.Context) {
|
||||
// 先检查是否已安装
|
||||
status, err := c.installService.CheckInstallStatus()
|
||||
if err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, gin.H{"error": "检查安装状态失败"})
|
||||
return
|
||||
}
|
||||
|
||||
if status.Installed {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "系统已安装,无法重复安装"})
|
||||
return
|
||||
}
|
||||
|
||||
var req services.InstallRequest
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "参数错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
// 验证必填字段
|
||||
if req.AdminUsername == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "管理员用户名不能为空"})
|
||||
return
|
||||
}
|
||||
if req.AdminPassword == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "管理员密码不能为空"})
|
||||
return
|
||||
}
|
||||
if len(req.AdminPassword) < 6 {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "管理员密码至少6位"})
|
||||
return
|
||||
}
|
||||
|
||||
// MySQL 必填验证
|
||||
if req.DBType == "mysql" {
|
||||
if req.DBHost == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "MySQL 主机不能为空"})
|
||||
return
|
||||
}
|
||||
if req.DBName == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "MySQL 数据库名不能为空"})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Redis 启用时的验证
|
||||
if req.RedisEnabled {
|
||||
if req.RedisHost == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "Redis 主机不能为空"})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 执行安装
|
||||
if err := c.installService.Install(&req); err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, gin.H{"error": "安装失败: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
ctx.JSON(http.StatusOK, gin.H{
|
||||
"message": "安装成功",
|
||||
"admin_username": req.AdminUsername,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,532 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/tunnel"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type InterconnectController struct {
|
||||
interconnectService *services.InterconnectService
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewInterconnectController(interconnectService *services.InterconnectService) *InterconnectController {
|
||||
return &InterconnectController{
|
||||
interconnectService: interconnectService,
|
||||
httpClient: &http.Client{
|
||||
Timeout: 10 * time.Second,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// GetNodes 获取互联节点列表
|
||||
func (ic *InterconnectController) GetNodes(c *gin.Context) {
|
||||
nodes, err := ic.interconnectService.GetNodes()
|
||||
if err != nil {
|
||||
utils.ServerError(c, "获取互联节点失败")
|
||||
return
|
||||
}
|
||||
utils.Success(c, nodes)
|
||||
}
|
||||
|
||||
// CreateNode 创建互联节点
|
||||
func (ic *InterconnectController) CreateNode(c *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
URL string `json:"url"`
|
||||
Token string `json:"token" binding:"required"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
node, err := ic.interconnectService.CreateNode(req.Name, req.URL, req.Token, req.Remark)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "创建互联节点失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, node)
|
||||
}
|
||||
|
||||
// UpdateNode 更新互联节点
|
||||
func (ic *InterconnectController) UpdateNode(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的节点ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
URL string `json:"url"`
|
||||
Token string `json:"token" binding:"required"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
node, err := ic.interconnectService.UpdateNode(id, req.Name, req.URL, req.Token, req.Remark)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "更新互联节点失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, node)
|
||||
}
|
||||
|
||||
// DeleteNode 删除互联节点
|
||||
func (ic *InterconnectController) DeleteNode(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的节点ID")
|
||||
return
|
||||
}
|
||||
|
||||
err := ic.interconnectService.DeleteNode(id)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "删除互联节点失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
// GetNodeStatus 获取单个子节点的状态
|
||||
func (ic *InterconnectController) GetNodeStatus(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
node, err := ic.interconnectService.GetNodeByID(id)
|
||||
if err != nil {
|
||||
utils.NotFound(c, "节点不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 针对反向隧道节点状态检测的特判
|
||||
if strings.HasPrefix(node.URL, "tunnel://") {
|
||||
sess := tunnel.GetSession(node.ID)
|
||||
if sess == nil {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "节点离线或反向隧道未建立", "data": nil})
|
||||
return
|
||||
}
|
||||
|
||||
// 使用当前 Yamux Session 的虚拟底层连接进行拨号
|
||||
transport := &http.Transport{
|
||||
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
return sess.Session.Open()
|
||||
},
|
||||
}
|
||||
client := &http.Client{
|
||||
Transport: transport,
|
||||
Timeout: 5 * time.Second,
|
||||
}
|
||||
|
||||
req, err := http.NewRequest("GET", "http://tunnel.local/api/v1/monitor", nil)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "构建检测请求失败")
|
||||
return
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "与子节点逆向连接通讯失败", "data": nil})
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode != 200 {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "子节点检测异常", "data": string(body)})
|
||||
return
|
||||
}
|
||||
|
||||
var jsonResp map[string]interface{}
|
||||
if err := json.Unmarshal(body, &jsonResp); err != nil {
|
||||
utils.ServerError(c, "解析节点检测数据失败")
|
||||
return
|
||||
}
|
||||
|
||||
if dataMap, ok := jsonResp["data"].(map[string]interface{}); ok {
|
||||
dataMap["tunnel_connected"] = true
|
||||
dataMap["tunnel_url"] = node.URL
|
||||
if hostMap, ok := dataMap["host"].(map[string]interface{}); ok {
|
||||
hostMap["tx_bytes"] = node.Metrics.TxBytes
|
||||
hostMap["rx_bytes"] = node.Metrics.RxBytes
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, jsonResp["data"])
|
||||
return
|
||||
}
|
||||
|
||||
apiURL := strings.TrimRight(node.URL, "/") + "/api/v1/monitor"
|
||||
req, err := http.NewRequest("GET", apiURL, nil)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "构建请求失败")
|
||||
return
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
|
||||
resp, err := ic.httpClient.Do(req)
|
||||
if err != nil {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "节点离线或网络不可达", "data": nil})
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode != 200 {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "节点返回异常", "data": string(body)})
|
||||
return
|
||||
}
|
||||
|
||||
var jsonResp map[string]interface{}
|
||||
if err := json.Unmarshal(body, &jsonResp); err != nil {
|
||||
utils.ServerError(c, "解析节点响应失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, jsonResp["data"])
|
||||
}
|
||||
|
||||
// SyncScript 将脚本同步到指定的节点列表
|
||||
func (ic *InterconnectController) SyncScript(c *gin.Context) {
|
||||
var req struct {
|
||||
NodeIDs []string `json:"node_ids" binding:"required"`
|
||||
Filename string `json:"filename" binding:"required"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
results := make([]map[string]interface{}, 0)
|
||||
|
||||
for _, nodeID := range req.NodeIDs {
|
||||
node, err := ic.interconnectService.GetNodeByID(nodeID)
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "节点不存在"})
|
||||
continue
|
||||
}
|
||||
|
||||
client, apiURL, err := ic.getClientAndURL(node, "/api/v1/scripts/save")
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "反向隧道未连接"})
|
||||
continue
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"filename": req.Filename,
|
||||
"content": req.Content,
|
||||
}
|
||||
payloadBytes, _ := json.Marshal(payload)
|
||||
|
||||
httpReq, err := http.NewRequest("POST", apiURL, bytes.NewBuffer(payloadBytes))
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "构建请求失败"})
|
||||
continue
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil || resp.StatusCode != 200 {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "同步请求失败或超时"})
|
||||
} else {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": true, "msg": "同步成功"})
|
||||
}
|
||||
if resp != nil {
|
||||
resp.Body.Close()
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, results)
|
||||
}
|
||||
|
||||
// SyncEnv 将环境变量同步到指定的节点列表
|
||||
func (ic *InterconnectController) SyncEnv(c *gin.Context) {
|
||||
var req struct {
|
||||
NodeIDs []string `json:"node_ids" binding:"required"`
|
||||
Envs []struct{ ID string `json:"id"` } `json:"envs" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var envIDs []string
|
||||
for _, e := range req.Envs {
|
||||
envIDs = append(envIDs, e.ID)
|
||||
}
|
||||
|
||||
dataService := services.NewDataService()
|
||||
exportData := dataService.ExportBusinessData(nil, envIDs)
|
||||
|
||||
results := make([]map[string]interface{}, 0)
|
||||
|
||||
for _, nodeID := range req.NodeIDs {
|
||||
node, err := ic.interconnectService.GetNodeByID(nodeID)
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "节点不存在"})
|
||||
continue
|
||||
}
|
||||
|
||||
client, apiURL, err := ic.getClientAndURL(node, "/api/v1/system/import")
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "反向隧道未连接"})
|
||||
continue
|
||||
}
|
||||
|
||||
payloadBytes, _ := json.Marshal(exportData)
|
||||
|
||||
httpReq, err := http.NewRequest("POST", apiURL, bytes.NewBuffer(payloadBytes))
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "构建请求失败"})
|
||||
continue
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil || resp.StatusCode != 200 {
|
||||
msg := "同步失败"
|
||||
if err != nil {
|
||||
msg = err.Error()
|
||||
}
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": msg})
|
||||
} else {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": true, "msg": "同步成功"})
|
||||
}
|
||||
if resp != nil {
|
||||
resp.Body.Close()
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, results)
|
||||
}
|
||||
|
||||
// SyncTask 将任务同步到指定的节点列表
|
||||
func (ic *InterconnectController) SyncTask(c *gin.Context) {
|
||||
var req struct {
|
||||
NodeIDs []string `json:"node_ids" binding:"required"`
|
||||
Tasks []struct{ ID string `json:"id"` } `json:"tasks" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var taskIDs []string
|
||||
for _, t := range req.Tasks {
|
||||
taskIDs = append(taskIDs, t.ID)
|
||||
}
|
||||
|
||||
dataService := services.NewDataService()
|
||||
exportData := dataService.ExportBusinessData(taskIDs, nil)
|
||||
|
||||
results := make([]map[string]interface{}, 0)
|
||||
|
||||
for _, nodeID := range req.NodeIDs {
|
||||
node, err := ic.interconnectService.GetNodeByID(nodeID)
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "节点不存在"})
|
||||
continue
|
||||
}
|
||||
|
||||
client, apiURL, err := ic.getClientAndURL(node, "/api/v1/system/import")
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "反向隧道未连接"})
|
||||
continue
|
||||
}
|
||||
|
||||
payloadBytes, _ := json.Marshal(exportData)
|
||||
|
||||
httpReq, err := http.NewRequest("POST", apiURL, bytes.NewBuffer(payloadBytes))
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "构建请求失败"})
|
||||
continue
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil || resp.StatusCode != 200 {
|
||||
msg := "同步失败"
|
||||
if err != nil {
|
||||
msg = err.Error()
|
||||
}
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": msg})
|
||||
} else {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": true, "msg": "同步成功"})
|
||||
}
|
||||
if resp != nil {
|
||||
resp.Body.Close()
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, results)
|
||||
}
|
||||
|
||||
// HandleTunnel 接受子节点 WebSocket 连接请求
|
||||
func (ic *InterconnectController) HandleTunnel(c *gin.Context) {
|
||||
tunnel.HandleTunnel(c)
|
||||
}
|
||||
|
||||
// ProxyRequest 代理转发请求至目标节点
|
||||
func (ic *InterconnectController) ProxyRequest(c *gin.Context) {
|
||||
nodeID := c.Param("node_id")
|
||||
path := c.Param("path")
|
||||
if nodeID == "" {
|
||||
utils.BadRequest(c, "Node ID required")
|
||||
return
|
||||
}
|
||||
|
||||
node, err := ic.interconnectService.GetNodeByID(nodeID)
|
||||
if err != nil {
|
||||
utils.NotFound(c, "Node not found")
|
||||
return
|
||||
}
|
||||
|
||||
if strings.HasPrefix(node.URL, "tunnel://") {
|
||||
// 走 WebSocket 逆向隧道 (基于 Yamux 流式多路复用)
|
||||
err := tunnel.ProxyHTTP(nodeID, c, path)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "Tunnel request failed: "+err.Error())
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// 走普通 HTTP 直连
|
||||
// Construct the target URL
|
||||
targetURL := strings.TrimRight(node.URL, "/") + path
|
||||
if c.Request.URL.RawQuery != "" {
|
||||
targetURL += "?" + c.Request.URL.RawQuery
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(c.Request.Method, targetURL, c.Request.Body)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "Failed to create proxy request")
|
||||
return
|
||||
}
|
||||
|
||||
// Copy headers
|
||||
req.Header = c.Request.Header.Clone()
|
||||
|
||||
// If the node token exists, append it as Bearer Auth
|
||||
if node.Token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
}
|
||||
|
||||
resp, err := ic.httpClient.Do(req)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "Failed to connect to target node: "+err.Error())
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
for k, v := range resp.Header {
|
||||
for _, vv := range v {
|
||||
c.Writer.Header().Add(k, vv)
|
||||
}
|
||||
}
|
||||
c.Status(resp.StatusCode)
|
||||
io.Copy(c.Writer, resp.Body)
|
||||
}
|
||||
|
||||
// ReportMonitorData 接收子节点上报的监控数据
|
||||
func (ic *InterconnectController) ReportMonitorData(c *gin.Context) {
|
||||
var req models.NodeMetrics
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if authHeader == "" {
|
||||
c.JSON(401, gin.H{"error": "missing authorization"})
|
||||
return
|
||||
}
|
||||
tokenStr := strings.TrimSpace(strings.TrimPrefix(authHeader, "Bearer "))
|
||||
|
||||
node, err := ic.interconnectService.GetNodeByToken(tokenStr)
|
||||
if err != nil {
|
||||
c.JSON(401, gin.H{"error": "invalid token"})
|
||||
return
|
||||
}
|
||||
|
||||
err = ic.interconnectService.UpdateNodeMonitorData(node.ID, req)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "更新节点数据失败")
|
||||
return
|
||||
}
|
||||
utils.Success(c, gin.H{
|
||||
"tunnel_url": node.URL,
|
||||
})
|
||||
}
|
||||
|
||||
// GetChildStatus 获取本机作为子节点的连接状态
|
||||
func (ic *InterconnectController) GetChildStatus(c *gin.Context) {
|
||||
settingsSvc := services.NewSettingsService()
|
||||
parentURL := settingsSvc.Get(constant.SectionInterconnect, constant.KeyInterconnectParentURL)
|
||||
parentToken := settingsSvc.Get(constant.SectionInterconnect, constant.KeyInterconnectParentToken)
|
||||
|
||||
connected := tunnel.IsTunnelConnected()
|
||||
tunnelURL := tunnel.GetLocalTunnelURL()
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"parent_url": parentURL,
|
||||
"parent_token": parentToken,
|
||||
"connected": connected,
|
||||
"tunnel_url": tunnelURL,
|
||||
"tx_bytes": tunnel.GetTxBytes(),
|
||||
"rx_bytes": tunnel.GetRxBytes(),
|
||||
})
|
||||
}
|
||||
|
||||
// getClientAndURL 辅助方法:根据节点类型决定走直连还是隧道,并返回对应的 Client 和完整 URL
|
||||
func (ic *InterconnectController) getClientAndURL(node *models.InterconnectNode, path string) (*http.Client, string, error) {
|
||||
if strings.HasPrefix(node.URL, "tunnel://") {
|
||||
sess := tunnel.GetSession(node.ID)
|
||||
if sess == nil {
|
||||
return nil, "", net.ErrClosed
|
||||
}
|
||||
transport := &http.Transport{
|
||||
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
return sess.Session.Open()
|
||||
},
|
||||
}
|
||||
client := &http.Client{
|
||||
Transport: transport,
|
||||
Timeout: 10 * time.Second,
|
||||
}
|
||||
return client, "http://tunnel.local" + path, nil
|
||||
}
|
||||
|
||||
targetURL := strings.TrimRight(node.URL, "/") + path
|
||||
return ic.httpClient, targetURL, nil
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type LogController struct{}
|
||||
|
||||
func NewLogController() *LogController {
|
||||
return &LogController{}
|
||||
}
|
||||
|
||||
// GetLogs 获取任务日志列表
|
||||
// @Summary 获取任务日志列表
|
||||
// @Description 分页获取任务日志列表,支持按任务 ID、任务名称、状态筛选
|
||||
// @Tags 日志管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param task_id query string false "任务 ID"
|
||||
// @Param task_name query string false "任务名称"
|
||||
// @Param status query string false "状态"
|
||||
// @Param page query int false "页码"
|
||||
// @Param page_size query int false "每页数量"
|
||||
// @Success 200 {object} utils.Response{data=utils.PaginationData{data=[]vo.TaskLogVO}}
|
||||
// @Router /logs [get]
|
||||
func (lc *LogController) GetLogs(c *gin.Context) {
|
||||
p := utils.ParsePagination(c)
|
||||
taskID := c.DefaultQuery("task_id", "")
|
||||
taskName := c.DefaultQuery("task_name", "")
|
||||
status := c.DefaultQuery("status", "")
|
||||
|
||||
var logs []models.TaskLog
|
||||
var total int64
|
||||
|
||||
query := database.DB.Model(&models.TaskLog{})
|
||||
if taskID != "" {
|
||||
query = query.Where("task_id = ?", taskID)
|
||||
}
|
||||
if status != "" {
|
||||
query = query.Where("status = ?", status)
|
||||
}
|
||||
|
||||
// 按任务名称过滤
|
||||
if taskName != "" {
|
||||
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)
|
||||
} else {
|
||||
utils.PaginatedResponse(c, []vo.TaskLogVO{}, 0, p)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
query.Count(&total)
|
||||
query.Order("id DESC").Offset(p.Offset()).Limit(p.PageSize).Find(&logs)
|
||||
|
||||
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[string]models.Task)
|
||||
for _, t := range tasks {
|
||||
taskMap[t.ID] = t
|
||||
}
|
||||
|
||||
result := make([]vo.TaskLogVO, len(logs))
|
||||
for i, log := range logs {
|
||||
task := taskMap[log.TaskID]
|
||||
taskType := task.Type
|
||||
if taskType == "" {
|
||||
taskType = "task"
|
||||
}
|
||||
result[i] = vo.TaskLogVO{
|
||||
ID: log.ID,
|
||||
TaskID: log.TaskID,
|
||||
TaskName: task.Name,
|
||||
TaskType: taskType,
|
||||
AgentID: log.AgentID,
|
||||
Command: string(log.Command),
|
||||
Status: log.Status,
|
||||
Duration: log.Duration,
|
||||
StartTime: log.StartTime,
|
||||
EndTime: log.EndTime,
|
||||
CreatedAt: log.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
utils.PaginatedResponse(c, result, total, p)
|
||||
}
|
||||
|
||||
// GetLogDetail 获取日志详情
|
||||
// @Summary 获取日志详情
|
||||
// @Description 根据 ID 获取任务日志详细内容(包含输出)
|
||||
// @Tags 日志管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "日志ID"
|
||||
// @Success 200 {object} utils.Response{data=vo.TaskLogVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /logs/{id} [get]
|
||||
func (lc *LogController) GetLogDetail(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的日志ID")
|
||||
return
|
||||
}
|
||||
|
||||
var log models.TaskLog
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&log)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
utils.NotFound(c, "日志不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskLogVO(&log))
|
||||
}
|
||||
|
||||
// ClearLogs 清空日志
|
||||
func (lc *LogController) ClearLogs(c *gin.Context) {
|
||||
var req struct {
|
||||
TaskID *string `json:"task_id"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
query := database.DB.Model(&models.TaskLog{})
|
||||
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
|
||||
}
|
||||
|
||||
if err := query.Delete(&models.TaskLog{}).Error; err != nil {
|
||||
utils.ServerError(c, "清空日志失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "日志清空成功")
|
||||
}
|
||||
|
||||
// DeleteLog 删除日志
|
||||
func (lc *LogController) DeleteLog(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的日志ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := database.DB.Where("id = ?", id).Delete(&models.TaskLog{}).Error; err != nil {
|
||||
utils.ServerError(c, "删除日志失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "日志已删除")
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type LogSSEController struct{}
|
||||
|
||||
func NewLogSSEController() *LogSSEController {
|
||||
return &LogSSEController{}
|
||||
}
|
||||
|
||||
func (lc *LogSSEController) StreamLog(c *gin.Context) {
|
||||
logIDStr := c.Query("log_id")
|
||||
if logIDStr == "" {
|
||||
c.JSON(400, gin.H{"error": "log_id is required"})
|
||||
return
|
||||
}
|
||||
|
||||
logID := logIDStr
|
||||
|
||||
c.Header("Content-Type", "text/event-stream")
|
||||
c.Header("Cache-Control", "no-cache")
|
||||
c.Header("Connection", "keep-alive")
|
||||
c.Header("Transfer-Encoding", "chunked")
|
||||
// c.Header("Access-Control-Allow-Origin", "*")
|
||||
|
||||
// 1. 检查数据库中是否已结束
|
||||
var taskLog models.TaskLog
|
||||
res := database.DB.Where("id = ?", logID).Limit(1).Find(&taskLog)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
if taskLog.Status != "running" {
|
||||
// 已结束,读取库内日志
|
||||
content, err := utils.DecompressFromBase64(string(taskLog.Output))
|
||||
if err != nil {
|
||||
c.SSEvent("message", gin.H{"text": "解压日志失败: " + err.Error()})
|
||||
c.Writer.Flush()
|
||||
return
|
||||
}
|
||||
c.SSEvent("message", gin.H{"text": content})
|
||||
c.Writer.Flush()
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 未结束或未找到记录,尝试从 TinyLogManager 获取
|
||||
tl := tasks.GetActiveLog(logID)
|
||||
if tl == nil {
|
||||
c.SSEvent("message", gin.H{"text": "未找到正在运行的任务日志"})
|
||||
c.Writer.Flush()
|
||||
return
|
||||
}
|
||||
|
||||
// 发送系统提示
|
||||
c.SSEvent("message", gin.H{"text": fmt.Sprintf("[System] 连接成功,正在监听日志... (LogID: %s)\n", logID)})
|
||||
c.Writer.Flush()
|
||||
|
||||
// 发送最后 100 行
|
||||
lastLines, err := tl.ReadLastLines(100)
|
||||
if err == nil && len(lastLines) > 0 {
|
||||
c.SSEvent("message", gin.H{"text": string(lastLines)})
|
||||
c.Writer.Flush()
|
||||
}
|
||||
|
||||
// 订阅实时更新
|
||||
sub := tl.Subscribe()
|
||||
defer tl.Unsubscribe(sub)
|
||||
|
||||
// 推送更新
|
||||
c.Stream(func(w io.Writer) bool {
|
||||
select {
|
||||
case data, ok := <-sub:
|
||||
if !ok {
|
||||
// 任务结束,尝试刷新最后一次库内完整内容
|
||||
var finalLog models.TaskLog
|
||||
res := database.DB.Where("id = ?", logID).Limit(1).Find(&finalLog)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
content, _ := utils.DecompressFromBase64(string(finalLog.Output))
|
||||
if content != "" {
|
||||
c.SSEvent("message", gin.H{"text": "\n--- 任务已结束 ---\n"})
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
c.SSEvent("message", gin.H{"text": string(data)})
|
||||
return true
|
||||
case <-c.Request.Context().Done():
|
||||
return false
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,165 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type MiseController struct {
|
||||
service *services.MiseService
|
||||
}
|
||||
|
||||
func NewMiseController(service *services.MiseService) *MiseController {
|
||||
return &MiseController{
|
||||
service: service,
|
||||
}
|
||||
}
|
||||
|
||||
// List 获取语言列表
|
||||
func (c *MiseController) List(ctx *gin.Context) {
|
||||
langs, err := c.service.List()
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取语言列表失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, langs)
|
||||
}
|
||||
|
||||
// Sync 同步本地环境到数据库
|
||||
func (c *MiseController) Sync(ctx *gin.Context) {
|
||||
if err := c.service.Sync(); err != nil {
|
||||
utils.ServerError(ctx, "同步本地环境失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
|
||||
// Plugins 获取可用插件列表
|
||||
func (c *MiseController) Plugins(ctx *gin.Context) {
|
||||
plugins, err := c.service.Plugins()
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取插件列表失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, plugins)
|
||||
}
|
||||
|
||||
// Versions 获取指定插件的可用版本列表
|
||||
func (c *MiseController) Versions(ctx *gin.Context) {
|
||||
plugin := ctx.Query("plugin")
|
||||
if plugin == "" {
|
||||
utils.BadRequest(ctx, "参数 plugin 不能为空")
|
||||
return
|
||||
}
|
||||
versions, err := c.service.Versions(plugin)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取版本列表失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, versions)
|
||||
}
|
||||
|
||||
// VerifyCommand 获取验证命令
|
||||
func (c *MiseController) VerifyCommand(ctx *gin.Context) {
|
||||
plugin := ctx.Query("plugin")
|
||||
version := ctx.Query("version")
|
||||
if plugin == "" {
|
||||
utils.BadRequest(ctx, "参数 plugin 不能为空")
|
||||
return
|
||||
}
|
||||
cmd, err := c.service.GetVerifyCommand(plugin, version)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取验证命令失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, gin.H{"command": cmd})
|
||||
}
|
||||
|
||||
// UseGlobal 设置全局默认版本
|
||||
func (c *MiseController) UseGlobal(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Plugin string `json:"plugin"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
||||
return
|
||||
}
|
||||
if req.Plugin == "" || req.Version == "" {
|
||||
utils.BadRequest(ctx, "参数 plugin 和 version 不能为空")
|
||||
return
|
||||
}
|
||||
if err := c.service.UseGlobal(req.Plugin, req.Version); err != nil {
|
||||
utils.ServerError(ctx, "设置全局版本失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
|
||||
// UnsetGlobal 取消全局默认版本
|
||||
func (c *MiseController) UnsetGlobal(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Plugin string `json:"plugin"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
||||
return
|
||||
}
|
||||
if req.Plugin == "" {
|
||||
utils.BadRequest(ctx, "参数 plugin 不能为空")
|
||||
return
|
||||
}
|
||||
if err := c.service.UnsetGlobal(req.Plugin, req.Version); err != nil {
|
||||
utils.ServerError(ctx, "取消全局版本失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
|
||||
// Envs 获取全局环境变量
|
||||
func (c *MiseController) Envs(ctx *gin.Context) {
|
||||
envs, err := c.service.Envs()
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取全局环境变量失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, envs)
|
||||
}
|
||||
|
||||
// SetEnv 设置全局环境变量
|
||||
func (c *MiseController) SetEnv(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"`
|
||||
}
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
||||
return
|
||||
}
|
||||
if req.Key == "" {
|
||||
utils.BadRequest(ctx, "参数 key 不能为空")
|
||||
return
|
||||
}
|
||||
if err := c.service.SetEnv(req.Key, req.Value); err != nil {
|
||||
utils.ServerError(ctx, "设置环境变量失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
|
||||
// UnsetEnv 取消全局环境变量
|
||||
func (c *MiseController) UnsetEnv(ctx *gin.Context) {
|
||||
key := ctx.Query("key")
|
||||
if key == "" {
|
||||
utils.BadRequest(ctx, "参数 key 不能为空")
|
||||
return
|
||||
}
|
||||
if err := c.service.UnsetEnv(key); err != nil {
|
||||
utils.ServerError(ctx, "取消环境变量失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"runtime"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type MonitorController struct {
|
||||
executorService *tasks.ExecutorService
|
||||
}
|
||||
|
||||
func NewMonitorController(executorService *tasks.ExecutorService) *MonitorController {
|
||||
return &MonitorController{
|
||||
executorService: executorService,
|
||||
}
|
||||
}
|
||||
|
||||
var monitorUpgrader = websocket.Upgrader{
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true // 开发环境允许所有跨域,生产环境可根据配置限制
|
||||
},
|
||||
}
|
||||
|
||||
// GetSystemMonitor 获取系统和内存监控信息 (HTTP)
|
||||
func (mc *MonitorController) GetSystemMonitor(c *gin.Context) {
|
||||
data := mc.getMonitorData()
|
||||
utils.Success(c, data)
|
||||
}
|
||||
|
||||
// MonitorSSE Server-Sent Events 获取系统监控数据
|
||||
func (mc *MonitorController) MonitorSSE(c *gin.Context) {
|
||||
// 设置 SSE 响应头
|
||||
c.Writer.Header().Set("Content-Type", "text/event-stream")
|
||||
c.Writer.Header().Set("Cache-Control", "no-cache")
|
||||
c.Writer.Header().Set("Connection", "keep-alive")
|
||||
c.Writer.Header().Set("Transfer-Encoding", "chunked")
|
||||
|
||||
// 初始发送一次数据
|
||||
if err := mc.sendMonitorDataSSE(c); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(5 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
if err := mc.sendMonitorDataSSE(c); err != nil {
|
||||
return // 客户端断开连接或发送失败
|
||||
}
|
||||
case <-c.Request.Context().Done():
|
||||
return // 连接已断开,立即退出
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (mc *MonitorController) sendMonitorDataSSE(c *gin.Context) error {
|
||||
data := mc.getMonitorData()
|
||||
// 使用 Gin 提供的 SSE 方法
|
||||
c.SSEvent("message", gin.H{
|
||||
"code": 200,
|
||||
"data": data,
|
||||
"msg": "success",
|
||||
})
|
||||
c.Writer.Flush()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (mc *MonitorController) getMonitorData() gin.H {
|
||||
rt := services.GetMonitorService().GetRuntimeMetrics()
|
||||
m := rt.MemStats
|
||||
|
||||
// 调用统一的监控服务获取物理机指标
|
||||
metrics := services.GetMonitorService().GetHostMetrics()
|
||||
|
||||
return gin.H{
|
||||
"env": gin.H{
|
||||
"os": runtime.GOOS,
|
||||
"arch": runtime.GOARCH,
|
||||
"go_version": runtime.Version(),
|
||||
"num_cpu": runtime.NumCPU(),
|
||||
"goroutines": rt.NumGoroutine,
|
||||
},
|
||||
"host": gin.H{
|
||||
"cpu_percent": metrics.CPUPercent,
|
||||
"mem_total": metrics.VMem.Total,
|
||||
"mem_used": metrics.VMem.Used,
|
||||
"mem_percent": metrics.VMem.UsedPercent,
|
||||
"disk_total": metrics.DiskUsage.Total,
|
||||
"disk_used": metrics.DiskUsage.Used,
|
||||
"disk_percent": metrics.DiskUsage.UsedPercent,
|
||||
"uptime": metrics.HostInfo.Uptime,
|
||||
"platform": metrics.HostInfo.Platform + " " + metrics.HostInfo.PlatformVersion,
|
||||
},
|
||||
"mem": gin.H{
|
||||
"alloc": m.Alloc,
|
||||
"total_alloc": m.TotalAlloc,
|
||||
"sys": m.Sys,
|
||||
"lookups": m.Lookups,
|
||||
"mallocs": m.Mallocs,
|
||||
"frees": m.Frees,
|
||||
},
|
||||
"heap": gin.H{
|
||||
"heap_alloc": m.HeapAlloc,
|
||||
"heap_sys": m.HeapSys,
|
||||
"heap_idle": m.HeapIdle,
|
||||
"heap_inuse": m.HeapInuse,
|
||||
"heap_released": m.HeapReleased,
|
||||
"heap_objects": m.HeapObjects,
|
||||
},
|
||||
"gc": gin.H{
|
||||
"next_gc": m.NextGC,
|
||||
"last_gc": m.LastGC,
|
||||
"pause_total_ns": m.PauseTotalNs,
|
||||
"num_gc": m.NumGC,
|
||||
},
|
||||
"scheduler": gin.H{
|
||||
"scheduled": mc.executorService.GetScheduledCount(),
|
||||
"running": mc.executorService.GetRunningCount(),
|
||||
"queue_size": mc.executorService.GetScheduler().GetQueueSize(),
|
||||
"worker_count": mc.executorService.GetScheduler().GetConfig().WorkerCount,
|
||||
"workers": mc.executorService.GetScheduler().GetWorkerStatuses(),
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type NotificationController struct {
|
||||
notifyService *services.NotificationService
|
||||
}
|
||||
|
||||
func NewNotificationController() *NotificationController {
|
||||
return &NotificationController{
|
||||
notifyService: services.NewNotificationService(),
|
||||
}
|
||||
}
|
||||
|
||||
// GetChannelTypes 获取支持的渠道类型
|
||||
func (nc *NotificationController) GetChannelTypes(c *gin.Context) {
|
||||
utils.Success(c, gin.H{
|
||||
"channel_types": services.SupportedChannelTypes,
|
||||
"event_types": services.SupportedEvents,
|
||||
})
|
||||
}
|
||||
|
||||
// GetChannels 获取所有渠道
|
||||
func (nc *NotificationController) GetChannels(c *gin.Context) {
|
||||
channels := nc.notifyService.GetChannels()
|
||||
utils.Success(c, channels)
|
||||
}
|
||||
|
||||
// SaveChannel 保存/更新渠道
|
||||
func (nc *NotificationController) SaveChannel(c *gin.Context) {
|
||||
var req services.NotifyChannel
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" || req.Type == "" {
|
||||
utils.BadRequest(c, "渠道名称和类型不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
if err := nc.notifyService.SaveChannel(req); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// DeleteChannel 删除渠道
|
||||
func (nc *NotificationController) DeleteChannel(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "缺少渠道ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := nc.notifyService.DeleteChannel(id); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
// TestChannel 测试渠道
|
||||
func (nc *NotificationController) TestChannel(c *gin.Context) {
|
||||
var req services.NotifyChannel
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
result := nc.notifyService.SendToChannel(req, &services.NotifyMessage{
|
||||
Title: "🔔 任务池测试通知",
|
||||
Text: "如果你看到这条消息,说明通知渠道配置正确!",
|
||||
})
|
||||
|
||||
utils.Success(c, result)
|
||||
}
|
||||
|
||||
// GetBindings 获取事件绑定列表
|
||||
func (nc *NotificationController) GetBindings(c *gin.Context) {
|
||||
bindings := nc.notifyService.GetBindings()
|
||||
utils.Success(c, bindings)
|
||||
}
|
||||
|
||||
// SaveBinding 保存事件绑定
|
||||
func (nc *NotificationController) SaveBinding(c *gin.Context) {
|
||||
var req struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
Event string `json:"event"`
|
||||
WayID string `json:"way_id"`
|
||||
DataID string `json:"data_id"`
|
||||
Extra models.BigText `json:"extra"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Type == "" || req.Event == "" || req.WayID == "" {
|
||||
utils.BadRequest(c, "类型、事件和渠道ID不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
binding := &models.NotifyBinding{
|
||||
ID: req.ID,
|
||||
Type: req.Type,
|
||||
Event: req.Event,
|
||||
WayID: req.WayID,
|
||||
DataID: req.DataID,
|
||||
Extra: req.Extra,
|
||||
}
|
||||
|
||||
if err := nc.notifyService.SaveBinding(binding); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, binding)
|
||||
}
|
||||
|
||||
// DeleteBinding 删除事件绑定
|
||||
func (nc *NotificationController) DeleteBinding(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "缺少绑定ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := nc.notifyService.DeleteBinding(id); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
// BatchSaveBindings 批量保存事件绑定
|
||||
func (nc *NotificationController) BatchSaveBindings(c *gin.Context) {
|
||||
var req struct {
|
||||
Type string `json:"type"`
|
||||
DataID string `json:"data_id"`
|
||||
Bindings []models.NotifyBinding `json:"bindings"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Type == "" {
|
||||
utils.BadRequest(c, "类型不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
if err := nc.notifyService.BatchSaveBindings(req.Type, req.DataID, req.Bindings); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// SendNotification API 发送通知(供脚本调用)
|
||||
func (nc *NotificationController) SendNotification(c *gin.Context) {
|
||||
var req struct {
|
||||
ChannelID string `json:"channel_id"`
|
||||
Title string `json:"title"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.ChannelID == "" || req.Title == "" {
|
||||
utils.BadRequest(c, "channel_id 和 title 不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
result := nc.notifyService.SendByChannelID(req.ChannelID, &services.NotifyMessage{
|
||||
Title: req.Title,
|
||||
Text: req.Text,
|
||||
})
|
||||
|
||||
utils.Success(c, result)
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type ScriptController struct {
|
||||
scriptService *services.ScriptService
|
||||
}
|
||||
|
||||
func NewScriptController(scriptService *services.ScriptService) *ScriptController {
|
||||
return &ScriptController{scriptService: scriptService}
|
||||
}
|
||||
|
||||
// CreateScript 创建脚本
|
||||
// @Summary 创建脚本
|
||||
// @Description 创建一个新的脚本
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param body body object true "脚本信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.ScriptVO}
|
||||
// @Router /scripts [post]
|
||||
func (sc *ScriptController) CreateScript(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
script := sc.scriptService.CreateScript(req.Name, req.Content, userID)
|
||||
utils.Success(c, vo.ToScriptVO(script))
|
||||
}
|
||||
|
||||
// GetScripts 获取脚本列表
|
||||
// @Summary 获取脚本列表
|
||||
// @Description 获取当前用户的所有脚本(内容字段为空)
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Success 200 {object} utils.Response{data=[]vo.ScriptVO}
|
||||
// @Router /scripts [get]
|
||||
func (sc *ScriptController) GetScripts(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
scripts := sc.scriptService.GetScriptsByUserID(userID)
|
||||
vos := vo.ToScriptVOListFromModels(scripts)
|
||||
for i := range vos {
|
||||
vos[i].Content = "" // 列表不返回内容
|
||||
}
|
||||
utils.Success(c, vos)
|
||||
}
|
||||
|
||||
// GetScript 获取脚本详情
|
||||
// @Summary 获取脚本详情
|
||||
// @Description 根据 ID 获取脚本详情
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "脚本ID"
|
||||
// @Success 200 {object} utils.Response{data=vo.ScriptVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /scripts/{id} [get]
|
||||
func (sc *ScriptController) GetScript(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的脚本ID")
|
||||
return
|
||||
}
|
||||
|
||||
script := sc.scriptService.GetScriptByID(id)
|
||||
if script == nil {
|
||||
utils.NotFound(c, "脚本不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToScriptVO(script))
|
||||
}
|
||||
|
||||
// UpdateScript 更新脚本
|
||||
// @Summary 更新脚本
|
||||
// @Description 根据 ID 更新脚本信息
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "脚本ID"
|
||||
// @Param body body object true "脚本更新信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.ScriptVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /scripts/{id} [put]
|
||||
func (sc *ScriptController) UpdateScript(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的脚本ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
script := sc.scriptService.UpdateScript(id, req.Name, req.Content)
|
||||
if script == nil {
|
||||
utils.NotFound(c, "脚本不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToScriptVO(script))
|
||||
}
|
||||
|
||||
// DeleteScript 删除脚本
|
||||
// @Summary 删除脚本
|
||||
// @Description 根据 ID 删除脚本
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "脚本ID"
|
||||
// @Success 200 {object} utils.Response
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /scripts/{id} [delete]
|
||||
func (sc *ScriptController) DeleteScript(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的脚本ID")
|
||||
return
|
||||
}
|
||||
|
||||
success := sc.scriptService.DeleteScript(id)
|
||||
if !success {
|
||||
utils.NotFound(c, "脚本不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
@@ -0,0 +1,624 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/tunnel"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/shirou/gopsutil/v3/process"
|
||||
)
|
||||
|
||||
type SettingsController struct {
|
||||
userService *services.UserService
|
||||
settingsService *services.SettingsService
|
||||
loginLogService *services.LoginLogService
|
||||
backupService *services.BackupService
|
||||
executorService *tasks.ExecutorService
|
||||
}
|
||||
|
||||
func NewSettingsController(userService *services.UserService, loginLogService *services.LoginLogService, executorService *tasks.ExecutorService) *SettingsController {
|
||||
return &SettingsController{
|
||||
userService: userService,
|
||||
settingsService: services.NewSettingsService(),
|
||||
loginLogService: loginLogService,
|
||||
backupService: services.NewBackupService(),
|
||||
executorService: executorService,
|
||||
}
|
||||
}
|
||||
|
||||
// ChangePassword 修改密码及账号信息
|
||||
func (sc *SettingsController) ChangePassword(c *gin.Context) {
|
||||
// 演示模式下禁止修改
|
||||
if constant.DemoMode {
|
||||
utils.BadRequest(c, "演示模式下不能修改账号或密码")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
OldUsername string `json:"old_username"`
|
||||
Username string `json:"username"`
|
||||
OldPassword string `json:"old_password" binding:"required"`
|
||||
NewPassword string `json:"new_password"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
userID := c.GetString("userID")
|
||||
var user *models.User
|
||||
res := database.DB.Where("id = ?", userID).Limit(1).Find(&user)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
utils.NotFound(c, "用户不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 统一校验原账密
|
||||
if req.OldUsername != "" && req.OldUsername != user.Username {
|
||||
utils.BadRequest(c, "原账号不正确")
|
||||
return
|
||||
}
|
||||
|
||||
if !sc.userService.AuthenticateUser(user.Username, req.OldPassword) {
|
||||
utils.BadRequest(c, "原密码错误")
|
||||
return
|
||||
}
|
||||
|
||||
var updated bool
|
||||
var logoutRequired bool
|
||||
|
||||
// 1. 处理用户名修改
|
||||
if req.Username != "" && req.Username != user.Username {
|
||||
if err := sc.userService.UpdateAccount(user.ID, req.Username); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
updated = true
|
||||
logoutRequired = true
|
||||
}
|
||||
|
||||
// 2. 处理密码修改
|
||||
if req.NewPassword != "" {
|
||||
if len(req.NewPassword) < 6 {
|
||||
utils.BadRequest(c, "新密码至少6位")
|
||||
return
|
||||
}
|
||||
if err := sc.userService.UpdatePassword(user.ID, req.NewPassword); err != nil {
|
||||
utils.ServerError(c, "修改密码失败")
|
||||
return
|
||||
}
|
||||
updated = true
|
||||
logoutRequired = true
|
||||
}
|
||||
|
||||
if !updated {
|
||||
utils.SuccessMsg(c, "未检测到变更内容")
|
||||
return
|
||||
}
|
||||
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventPasswordChanged,
|
||||
Payload: map[string]interface{}{
|
||||
"username": user.Username,
|
||||
},
|
||||
})
|
||||
|
||||
msg := "保存成功"
|
||||
if logoutRequired {
|
||||
msg += ",请重新登录"
|
||||
}
|
||||
utils.SuccessMsg(c, msg)
|
||||
}
|
||||
|
||||
// CleanLogs 清理日志 - 已移除,改为任务级别的日志清理配置
|
||||
|
||||
// GetSiteSettings 获取站点设置
|
||||
func (sc *SettingsController) GetSiteSettings(c *gin.Context) {
|
||||
settings := sc.settingsService.GetSection(constant.SectionSite)
|
||||
|
||||
// 纠正数据库中的空值,防止因配置冲突被意外置空
|
||||
if settings[constant.KeyTitle] == "" {
|
||||
settings[constant.KeyTitle] = "任务池"
|
||||
sc.settingsService.Set(constant.SectionSite, constant.KeyTitle, "任务池")
|
||||
}
|
||||
if settings[constant.KeySubtitle] == "" {
|
||||
settings[constant.KeySubtitle] = "极致轻量、高性能的自动化任务调度平台"
|
||||
sc.settingsService.Set(constant.SectionSite, constant.KeySubtitle, "极致轻量、高性能的自动化任务调度平台")
|
||||
}
|
||||
if settings[constant.KeyIcon] == "" {
|
||||
settings[constant.KeyIcon] = constant.DefaultIcon
|
||||
sc.settingsService.Set(constant.SectionSite, constant.KeyIcon, constant.DefaultIcon)
|
||||
}
|
||||
|
||||
// 解析 JSON 格式的 OpenAPI Token
|
||||
if tokenJson, ok := settings[constant.KeyOpenapiToken]; ok && tokenJson != "" {
|
||||
var tokenConfig vo.TokenConfig
|
||||
if err := json.Unmarshal([]byte(tokenJson), &tokenConfig); err == nil {
|
||||
settings["openapi_token"] = tokenConfig.Token
|
||||
settings["openapi_token_expire"] = tokenConfig.ExpireAt
|
||||
if tokenConfig.Enabled {
|
||||
settings["openapi_enabled"] = "true"
|
||||
} else {
|
||||
settings["openapi_enabled"] = "false"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 获取日志清理配置
|
||||
settings["system_notice_days"] = sc.settingsService.Get(constant.SectionSystem, constant.KeySystemNoticeDays)
|
||||
settings["system_notice_max_count"] = sc.settingsService.Get(constant.SectionSystem, constant.KeySystemNoticeMaxCount)
|
||||
settings["push_log_days"] = sc.settingsService.Get(constant.SectionSystem, constant.KeyPushLogDays)
|
||||
settings["push_log_max_count"] = sc.settingsService.Get(constant.SectionSystem, constant.KeyPushLogMaxCount)
|
||||
settings["login_log_days"] = sc.settingsService.Get(constant.SectionSystem, constant.KeyLoginLogDays)
|
||||
settings["login_log_max_count"] = sc.settingsService.Get(constant.SectionSystem, constant.KeyLoginLogMaxCount)
|
||||
settings["scheduler_log_days"] = sc.settingsService.Get(constant.SectionSystem, constant.KeySchedulerLogDays)
|
||||
settings["scheduler_log_max_count"] = sc.settingsService.Get(constant.SectionSystem, constant.KeySchedulerLogMaxCount)
|
||||
|
||||
utils.Success(c, settings)
|
||||
}
|
||||
|
||||
// GetPublicSiteSettings 获取公开的站点设置(无需认证)
|
||||
func (sc *SettingsController) GetPublicSiteSettings(c *gin.Context) {
|
||||
settings := sc.settingsService.GetSection(constant.SectionSite)
|
||||
|
||||
title := settings[constant.KeyTitle]
|
||||
if title == "" {
|
||||
title = "任务池"
|
||||
}
|
||||
subtitle := settings[constant.KeySubtitle]
|
||||
if subtitle == "" {
|
||||
subtitle = "极致轻量、高性能的自动化任务调度平台"
|
||||
}
|
||||
icon := settings[constant.KeyIcon]
|
||||
if icon == "" {
|
||||
icon = constant.DefaultIcon
|
||||
}
|
||||
|
||||
// 只返回公开信息
|
||||
utils.Success(c, gin.H{
|
||||
constant.KeyTitle: title,
|
||||
constant.KeySubtitle: subtitle,
|
||||
constant.KeyIcon: icon,
|
||||
"demo_mode": constant.DemoMode,
|
||||
})
|
||||
}
|
||||
|
||||
// UpdateSiteSettings 更新站点设置
|
||||
func (sc *SettingsController) UpdateSiteSettings(c *gin.Context) {
|
||||
var req struct {
|
||||
Title string `json:"title"`
|
||||
Subtitle string `json:"subtitle"`
|
||||
Icon string `json:"icon"`
|
||||
PageSize string `json:"page_size"`
|
||||
CookieDays string `json:"cookie_days"`
|
||||
OpenapiEnabled bool `json:"openapi_enabled"`
|
||||
OpenapiToken string `json:"openapi_token"`
|
||||
OpenapiTokenExpire string `json:"openapi_token_expire"`
|
||||
SystemNoticeDays string `json:"system_notice_days"`
|
||||
SystemNoticeMaxCount string `json:"system_notice_max_count"`
|
||||
PushLogDays string `json:"push_log_days"`
|
||||
PushLogMaxCount string `json:"push_log_max_count"`
|
||||
LoginLogDays string `json:"login_log_days"`
|
||||
LoginLogMaxCount string `json:"login_log_max_count"`
|
||||
SchedulerLogDays string `json:"scheduler_log_days"`
|
||||
SchedulerLogMaxCount string `json:"scheduler_log_max_count"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
openapiTokenJson := ""
|
||||
if req.OpenapiToken != "" || req.OpenapiTokenExpire != "" || req.OpenapiEnabled {
|
||||
tokenConfig := vo.TokenConfig{
|
||||
Enabled: req.OpenapiEnabled,
|
||||
Token: req.OpenapiToken,
|
||||
ExpireAt: req.OpenapiTokenExpire,
|
||||
}
|
||||
if b, err := json.Marshal(tokenConfig); err == nil {
|
||||
openapiTokenJson = string(b)
|
||||
}
|
||||
}
|
||||
|
||||
values := map[string]string{
|
||||
constant.KeyTitle: req.Title,
|
||||
constant.KeySubtitle: req.Subtitle,
|
||||
constant.KeyIcon: req.Icon,
|
||||
constant.KeyPageSize: req.PageSize,
|
||||
constant.KeyCookieDays: req.CookieDays,
|
||||
constant.KeyOpenapiToken: openapiTokenJson,
|
||||
}
|
||||
|
||||
if err := sc.settingsService.SetSection(constant.SectionSite, values); err != nil {
|
||||
utils.ServerError(c, "保存失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 保存日志清理配置
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeySystemNoticeDays, req.SystemNoticeDays)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeySystemNoticeMaxCount, req.SystemNoticeMaxCount)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeyPushLogDays, req.PushLogDays)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeyPushLogMaxCount, req.PushLogMaxCount)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeyLoginLogDays, req.LoginLogDays)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeyLoginLogMaxCount, req.LoginLogMaxCount)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeySchedulerLogDays, req.SchedulerLogDays)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeySchedulerLogMaxCount, req.SchedulerLogMaxCount)
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// GenerateOpenapiToken 随机生成OpenAPI Token
|
||||
func (sc *SettingsController) GenerateOpenapiToken(c *gin.Context) {
|
||||
utils.Success(c, gin.H{
|
||||
"token": strings.ToLower(utils.RandomString(32)),
|
||||
})
|
||||
}
|
||||
|
||||
// GetSchedulerSettings 获取调度设置
|
||||
func (sc *SettingsController) GetSchedulerSettings(c *gin.Context) {
|
||||
settings := sc.settingsService.GetSection(constant.SectionScheduler)
|
||||
utils.Success(c, settings)
|
||||
}
|
||||
|
||||
// UpdateSchedulerSettings 更新调度设置
|
||||
func (sc *SettingsController) UpdateSchedulerSettings(c *gin.Context) {
|
||||
var req struct {
|
||||
WorkerCount string `json:"worker_count"`
|
||||
QueueSize string `json:"queue_size"`
|
||||
RateInterval string `json:"rate_interval"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
var workerCount int
|
||||
if _, err := fmt.Sscanf(req.WorkerCount, "%d", &workerCount); err != nil || workerCount < 1 || workerCount > 1000 {
|
||||
utils.BadRequest(c, "工作线程数必须在 1 至 1000 之间")
|
||||
return
|
||||
}
|
||||
|
||||
var queueSize int
|
||||
if _, err := fmt.Sscanf(req.QueueSize, "%d", &queueSize); err != nil || queueSize < 1 || queueSize > 50000 {
|
||||
utils.BadRequest(c, "等待队列容量必须在 1 至 50000 之间")
|
||||
return
|
||||
}
|
||||
|
||||
var rateInterval int
|
||||
if _, err := fmt.Sscanf(req.RateInterval, "%d", &rateInterval); err != nil || rateInterval < 1 {
|
||||
utils.BadRequest(c, "限频间隔必须为正整数")
|
||||
return
|
||||
}
|
||||
|
||||
values := map[string]string{
|
||||
constant.KeyWorkerCount: req.WorkerCount,
|
||||
constant.KeyQueueSize: req.QueueSize,
|
||||
constant.KeyRateInterval: req.RateInterval,
|
||||
}
|
||||
|
||||
if err := sc.settingsService.SetSection(constant.SectionScheduler, values); err != nil {
|
||||
utils.ServerError(c, "保存失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 重新加载 executor service
|
||||
if sc.executorService != nil {
|
||||
sc.executorService.Reload()
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// GetPaths 获取系统路径信息
|
||||
func (sc *SettingsController) GetPaths(c *gin.Context) {
|
||||
absScriptsDir, _ := filepath.Abs(constant.ScriptsWorkDir)
|
||||
utils.Success(c, gin.H{
|
||||
"scripts_dir": absScriptsDir,
|
||||
})
|
||||
}
|
||||
|
||||
// GetAbout 获取关于信息
|
||||
func (sc *SettingsController) GetAbout(c *gin.Context) {
|
||||
var taskCount, logCount, envCount int64
|
||||
database.DB.Model(&models.Task{}).Count(&taskCount)
|
||||
database.DB.Model(&models.TaskLog{}).Count(&logCount)
|
||||
database.DB.Model(&models.EnvironmentVariable{}).Count(&envCount)
|
||||
|
||||
// 内存使用
|
||||
memUsage := "N/A"
|
||||
if p, err := process.NewProcess(int32(os.Getpid())); err == nil {
|
||||
if memInfo, err := p.MemoryInfo(); err == nil {
|
||||
memUsage = formatBytes(memInfo.RSS)
|
||||
}
|
||||
}
|
||||
|
||||
// 运行时间
|
||||
uptime := formatDuration(time.Since(constant.StartTime))
|
||||
|
||||
// 获取远程最新版本
|
||||
remoteVersion := ""
|
||||
client := &http.Client{Timeout: 2 * time.Second}
|
||||
req, err := http.NewRequest("GET", "https://api.github.com/repos/engigu/taskpool/releases/latest", nil)
|
||||
if err == nil {
|
||||
req.Header.Set("User-Agent", "taskpool")
|
||||
if resp, err := client.Do(req); err == nil {
|
||||
defer resp.Body.Close()
|
||||
var release struct {
|
||||
TagName string `json:"tag_name"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&release); err == nil {
|
||||
remoteVersion = release.TagName
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"version": constant.Version,
|
||||
"remote_version": remoteVersion,
|
||||
"build_time": constant.BuildTime,
|
||||
"mem_usage": memUsage,
|
||||
"goroutines": runtime.NumGoroutine(),
|
||||
"uptime": uptime,
|
||||
"task_count": taskCount,
|
||||
"log_count": logCount,
|
||||
"env_count": envCount,
|
||||
})
|
||||
}
|
||||
|
||||
// GetChangelog 获取更新日志
|
||||
func (sc *SettingsController) GetChangelog(c *gin.Context) {
|
||||
content, err := os.ReadFile("docs/guide/changelog.md")
|
||||
if err != nil {
|
||||
utils.Success(c, "暂无更新日志")
|
||||
return
|
||||
}
|
||||
utils.Success(c, string(content))
|
||||
}
|
||||
|
||||
// formatBytes 格式化字节数
|
||||
func formatBytes(bytes uint64) string {
|
||||
const unit = 1024
|
||||
if bytes < unit {
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
div, exp := uint64(unit), 0
|
||||
for n := bytes / unit; n >= unit; n /= unit {
|
||||
div *= unit
|
||||
exp++
|
||||
}
|
||||
return fmt.Sprintf("%.1f %cB", float64(bytes)/float64(div), "KMGTPE"[exp])
|
||||
}
|
||||
|
||||
// formatDuration 格式化时间间隔
|
||||
func formatDuration(d time.Duration) string {
|
||||
days := int(d.Hours()) / 24
|
||||
hours := int(d.Hours()) % 24
|
||||
minutes := int(d.Minutes()) % 60
|
||||
seconds := int(d.Seconds()) % 60
|
||||
|
||||
if days > 0 {
|
||||
return fmt.Sprintf("%d天%d小时%d分钟%d秒", days, hours, minutes, seconds)
|
||||
}
|
||||
if hours > 0 {
|
||||
return fmt.Sprintf("%d小时%d分钟%d秒", hours, minutes, seconds)
|
||||
}
|
||||
if minutes > 0 {
|
||||
return fmt.Sprintf("%d分钟%d秒", minutes, seconds)
|
||||
}
|
||||
return fmt.Sprintf("%d秒", seconds)
|
||||
}
|
||||
|
||||
// GetLoginLogs 获取登录日志
|
||||
func (sc *SettingsController) GetLoginLogs(c *gin.Context) {
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "10"))
|
||||
username := c.Query("username")
|
||||
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 || pageSize > 100 {
|
||||
pageSize = 10
|
||||
}
|
||||
|
||||
logs, total, err := sc.loginLogService.List(page, pageSize, username)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "获取登录日志失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 将 AppLog 转换为 LoginLogVO 返回,保持前端兼容性
|
||||
vos := make([]*vo.LoginLogVO, len(logs))
|
||||
for i, log := range logs {
|
||||
vos[i] = &vo.LoginLogVO{
|
||||
ID: log.ID,
|
||||
Username: log.Title,
|
||||
IP: log.RefID,
|
||||
UserAgent: string(log.Content),
|
||||
Status: log.Status,
|
||||
Message: string(log.ErrorMsg),
|
||||
CreatedAt: log.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, utils.PaginationData{
|
||||
Data: vos,
|
||||
Total: total,
|
||||
Page: page,
|
||||
PageSize: pageSize,
|
||||
})
|
||||
}
|
||||
|
||||
// CreateBackup 创建备份
|
||||
func (sc *SettingsController) CreateBackup(c *gin.Context) {
|
||||
_, err := sc.backupService.CreateBackup()
|
||||
if err != nil {
|
||||
utils.ServerError(c, "创建备份失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(c, "备份创建成功")
|
||||
}
|
||||
|
||||
// GetBackupStatus 获取备份状态
|
||||
func (sc *SettingsController) GetBackupStatus(c *gin.Context) {
|
||||
filePath := sc.backupService.GetBackupFile()
|
||||
var backupTime string
|
||||
if filePath != "" {
|
||||
if info, err := os.Stat(filePath); err == nil {
|
||||
backupTime = info.ModTime().Format("2006-01-02 15:04:05")
|
||||
}
|
||||
}
|
||||
utils.Success(c, gin.H{
|
||||
"has_backup": filePath != "",
|
||||
"backup_time": backupTime,
|
||||
})
|
||||
}
|
||||
|
||||
// DownloadBackup 下载备份文件
|
||||
func (sc *SettingsController) DownloadBackup(c *gin.Context) {
|
||||
filePath := sc.backupService.GetBackupFile()
|
||||
if filePath == "" {
|
||||
utils.NotFound(c, "没有可下载的备份")
|
||||
return
|
||||
}
|
||||
|
||||
// 检查文件是否存在
|
||||
if _, err := os.Stat(filePath); os.IsNotExist(err) {
|
||||
sc.backupService.ClearBackup()
|
||||
utils.NotFound(c, "备份文件不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 设置响应头
|
||||
c.Header("Content-Disposition", "attachment; filename="+filepath.Base(filePath))
|
||||
c.Header("Content-Type", "application/zip")
|
||||
c.File(filePath)
|
||||
|
||||
// 下载后清除备份记录和文件
|
||||
go func() {
|
||||
time.Sleep(time.Minute * 5) // 等待下载完成
|
||||
sc.backupService.ClearBackup()
|
||||
}()
|
||||
}
|
||||
|
||||
// RestoreBackup 恢复备份
|
||||
func (sc *SettingsController) RestoreBackup(c *gin.Context) {
|
||||
file, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
utils.BadRequest(c, "请上传备份文件")
|
||||
return
|
||||
}
|
||||
|
||||
// 保存上传的文件
|
||||
tempPath := filepath.Join(os.TempDir(), file.Filename)
|
||||
if err := c.SaveUploadedFile(file, tempPath); err != nil {
|
||||
utils.ServerError(c, "保存文件失败")
|
||||
return
|
||||
}
|
||||
defer os.Remove(tempPath)
|
||||
|
||||
// 恢复备份
|
||||
if err := sc.backupService.Restore(tempPath); err != nil {
|
||||
utils.ServerError(c, "恢复失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "恢复成功")
|
||||
}
|
||||
|
||||
// GetSectionSettings 获取指定 section 的所有设置
|
||||
func (sc *SettingsController) GetSectionSettings(c *gin.Context) {
|
||||
section := c.Param("section")
|
||||
if section == "" {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
settings := sc.settingsService.GetSection(section)
|
||||
utils.Success(c, settings)
|
||||
}
|
||||
|
||||
// UpdateSectionSettings 批量更新指定 section 的设置
|
||||
func (sc *SettingsController) UpdateSectionSettings(c *gin.Context) {
|
||||
section := c.Param("section")
|
||||
if section == "" {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
var values map[string]string
|
||||
if err := c.ShouldBindJSON(&values); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if err := sc.settingsService.SetSection(section, values); err != nil {
|
||||
utils.ServerError(c, "更新失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 当互联配置发生改变时,通知 tunnel 模块立刻应用新角色,启动或停止相关的后台协程
|
||||
if section == constant.SectionInterconnect {
|
||||
if role, ok := values[constant.KeyInterconnectRole]; ok {
|
||||
tunnel.ApplyRole(role)
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// GetSetting 获取单个设置值
|
||||
func (sc *SettingsController) GetSetting(c *gin.Context) {
|
||||
section := c.Param("section")
|
||||
key := c.Param("key")
|
||||
|
||||
if section == "" || key == "" {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
value := sc.settingsService.Get(section, key)
|
||||
utils.Success(c, value)
|
||||
}
|
||||
|
||||
// GenerateSettingToken 为指定设置生成随机token
|
||||
func (sc *SettingsController) GenerateSettingToken(c *gin.Context) {
|
||||
section := c.Param("section")
|
||||
key := c.Param("key")
|
||||
|
||||
if section == "" || key == "" {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 生成32位随机token
|
||||
token := strings.ToLower(utils.RandomString(32))
|
||||
|
||||
// 保存到数据库
|
||||
if err := sc.settingsService.Set(section, key, token); err != nil {
|
||||
utils.ServerError(c, "保存失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, token)
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type SystemWSController struct {
|
||||
manager *services.SystemWSManager
|
||||
}
|
||||
|
||||
func NewSystemWSController() *SystemWSController {
|
||||
return &SystemWSController{
|
||||
manager: services.GetSystemWSManager(),
|
||||
}
|
||||
}
|
||||
|
||||
func (sc *SystemWSController) HandleEvents(c *gin.Context) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
logger.Errorf("[SystemWS] 升级 WebSocket 失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
client := sc.manager.Register(conn)
|
||||
defer sc.manager.Unregister(client)
|
||||
|
||||
// 启动写循环
|
||||
go sc.writeLoop(client)
|
||||
|
||||
// 启动读循环 (主要用于检测连接断开和维持心跳)
|
||||
sc.readLoop(client)
|
||||
}
|
||||
|
||||
func (sc *SystemWSController) readLoop(client *services.ClientConnection) {
|
||||
defer client.Close()
|
||||
|
||||
client.Conn.SetReadLimit(constant.MaxMessageSize)
|
||||
client.Conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
client.Conn.SetPongHandler(func(string) error {
|
||||
client.Conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
return nil
|
||||
})
|
||||
|
||||
for {
|
||||
_, _, err := client.Conn.ReadMessage()
|
||||
if err != nil {
|
||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) {
|
||||
logger.Warnf("[SystemWS] 客户端异常断开: %v", err)
|
||||
}
|
||||
break
|
||||
}
|
||||
// 暂时不处理来自前端的消息,前端仅作为接收方
|
||||
}
|
||||
}
|
||||
|
||||
func (sc *SystemWSController) writeLoop(client *services.ClientConnection) {
|
||||
ticker := time.NewTicker(constant.PingPeriod)
|
||||
defer func() {
|
||||
ticker.Stop()
|
||||
client.Close()
|
||||
}()
|
||||
|
||||
for {
|
||||
select {
|
||||
case message, ok := <-client.Send:
|
||||
client.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
if !ok {
|
||||
// 通道关闭
|
||||
client.Conn.WriteMessage(websocket.CloseMessage, []byte{})
|
||||
return
|
||||
}
|
||||
|
||||
if err := client.Conn.WriteMessage(websocket.TextMessage, message); err != nil {
|
||||
return
|
||||
}
|
||||
case <-ticker.C:
|
||||
client.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
if err := client.Conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,872 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"os"
|
||||
)
|
||||
|
||||
type TaskController struct {
|
||||
taskService *tasks.TaskService
|
||||
executorService *tasks.ExecutorService
|
||||
agentWSManager *services.AgentWSManager
|
||||
}
|
||||
|
||||
func NewTaskController(taskService *tasks.TaskService, executorService *tasks.ExecutorService) *TaskController {
|
||||
return &TaskController{
|
||||
taskService: taskService,
|
||||
executorService: executorService,
|
||||
agentWSManager: services.GetAgentWSManager(),
|
||||
}
|
||||
}
|
||||
|
||||
// resolveWorkDir 将相对路径转换为绝对路径
|
||||
func resolveWorkDir(workDir string) string {
|
||||
if workDir == "" {
|
||||
// 空则使用默认 scripts 目录
|
||||
absPath, err := filepath.Abs(constant.ScriptsWorkDir)
|
||||
if err != nil {
|
||||
return constant.ScriptsWorkDir
|
||||
}
|
||||
return absPath
|
||||
}
|
||||
// 如果已经是绝对路径,直接返回
|
||||
if strings.HasPrefix(workDir, constant.ScriptsDirPlaceholder) {
|
||||
return workDir
|
||||
}
|
||||
if filepath.IsAbs(workDir) {
|
||||
return workDir
|
||||
}
|
||||
// 相对路径,基于 scripts 目录
|
||||
fullPath := filepath.Join(constant.ScriptsWorkDir, workDir)
|
||||
absPath, err := filepath.Abs(fullPath)
|
||||
if err != nil {
|
||||
return fullPath
|
||||
}
|
||||
return absPath
|
||||
}
|
||||
|
||||
// isValidDirName 校验目录名是否合法
|
||||
func isValidDirName(dirName string) bool {
|
||||
if dirName == "." || strings.Contains(dirName, "/") || strings.Contains(dirName, "\\") || strings.Contains(dirName, "..") {
|
||||
return false
|
||||
}
|
||||
for _, ch := range dirName {
|
||||
if !((ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || (ch >= '0' && ch <= '9') || ch == '_' || ch == '-' || ch == '.') {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// getRepoPhysicalPath 计算仓库任务的最终物理绝对路径
|
||||
func getRepoPhysicalPath(targetPath, dirName, sourceURL, branch string) string {
|
||||
if dirName == "." {
|
||||
return "" // 如果不追加目录,此逻辑不负责判断其根目录(共享的 scripts 目录)
|
||||
}
|
||||
finalDirName := dirName
|
||||
if finalDirName == "" {
|
||||
finalDirName = utils.GetRepoIdentifier(sourceURL, branch)
|
||||
}
|
||||
if finalDirName == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
basePath := targetPath
|
||||
if basePath == "" || basePath == constant.ScriptsDirPlaceholder {
|
||||
basePath = constant.ScriptsWorkDir
|
||||
} else if strings.HasPrefix(basePath, constant.ScriptsDirPlaceholder) {
|
||||
basePath = filepath.Join(constant.ScriptsWorkDir, strings.TrimPrefix(basePath, constant.ScriptsDirPlaceholder))
|
||||
} else if !filepath.IsAbs(basePath) {
|
||||
basePath = filepath.Join(constant.ScriptsWorkDir, basePath)
|
||||
}
|
||||
|
||||
fullPath := filepath.Join(basePath, finalDirName)
|
||||
absPath, err := filepath.Abs(fullPath)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return absPath
|
||||
}
|
||||
// CreateTask 创建任务
|
||||
// @Summary 创建任务
|
||||
// @Description 创建一个新的任务
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param body body vo.TaskCreateReq true "任务创建信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.TaskVO}
|
||||
// @Failure 400 {object} utils.Response
|
||||
// @Router /tasks [post]
|
||||
func (tc *TaskController) CreateTask(c *gin.Context) {
|
||||
var req vo.TaskCreateReq
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 普通任务需要命令
|
||||
if req.Type != constant.TaskTypeRepo && req.Command == "" {
|
||||
utils.BadRequest(c, "命令不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Schedule != "" {
|
||||
if err := tc.executorService.ValidateCron(req.Schedule); err != nil {
|
||||
utils.BadRequest(c, "无效的cron表达式: "+err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 转换为绝对路径(Agent 任务保持原样)
|
||||
workDir := req.WorkDir
|
||||
if req.AgentID == nil || *req.AgentID == "" {
|
||||
workDir = resolveWorkDir(req.WorkDir)
|
||||
}
|
||||
|
||||
var sourceID string
|
||||
// 如果是仓库同步任务,根据 URL 生成 SourceID 用于去重
|
||||
if req.Type == constant.TaskTypeRepo && req.Config != "" {
|
||||
var repoCfg struct {
|
||||
SourceURL string `json:"source_url"`
|
||||
Branch string `json:"branch"`
|
||||
RepoDirName string `json:"repo_dir_name"`
|
||||
TargetPath string `json:"target_path"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(req.Config), &repoCfg); err == nil && repoCfg.SourceURL != "" {
|
||||
if repoCfg.RepoDirName != "" {
|
||||
if !isValidDirName(repoCfg.RepoDirName) {
|
||||
utils.BadRequest(c, "自定义目录名只能包含字母、数字、下划线、短划线和点,不能只有点,且不能包含路径逻辑")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 如果配置了自定义名字,使用配置的名字。没有配置的话,使用以前的username_reponame
|
||||
if repoCfg.RepoDirName != "" {
|
||||
sourceID = "repo_" + repoCfg.RepoDirName
|
||||
} else {
|
||||
sourceID = "repo_" + utils.GetRepoIdentifier(repoCfg.SourceURL, repoCfg.Branch)
|
||||
}
|
||||
|
||||
// 校验 SourceID 是否已存在(任务唯一性)
|
||||
existingTask := tc.taskService.GetTaskBySourceID(sourceID)
|
||||
if existingTask != nil {
|
||||
utils.BadRequest(c, "当前任务已存在,请检查或更换仓库目录名称")
|
||||
return
|
||||
}
|
||||
|
||||
// 校验物理目录是否存在
|
||||
newAbsPath := getRepoPhysicalPath(repoCfg.TargetPath, repoCfg.RepoDirName, repoCfg.SourceURL, repoCfg.Branch)
|
||||
if newAbsPath != "" {
|
||||
if info, err := os.Stat(newAbsPath); err == nil && info.IsDir() {
|
||||
utils.BadRequest(c, "本地已存在同名仓库文件夹,请更换自定义目录名或清理残留文件")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
param := tasks.TaskParam{
|
||||
Name: req.Name,
|
||||
Remark: req.Remark,
|
||||
Command: req.Command,
|
||||
PreCommand: req.PreCommand,
|
||||
PostCommand: req.PostCommand,
|
||||
Tags: req.Tags,
|
||||
Type: req.Type,
|
||||
Config: req.Config,
|
||||
Schedule: req.Schedule,
|
||||
Timeout: req.Timeout,
|
||||
WorkDir: workDir,
|
||||
CleanConfig: req.CleanConfig,
|
||||
Envs: req.Envs,
|
||||
Languages: req.Languages,
|
||||
AgentID: req.AgentID,
|
||||
TriggerType: req.TriggerType,
|
||||
RetryCount: req.RetryCount,
|
||||
RetryInterval: req.RetryInterval,
|
||||
RandomRange: req.RandomRange,
|
||||
SourceID: sourceID,
|
||||
PinType: req.PinType,
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
var task *models.Task
|
||||
// 去重逻辑:如果已存在相同 SourceID 的仓库任务,则改为更新
|
||||
if sourceID != "" {
|
||||
task = tc.taskService.GetTaskBySourceID(sourceID)
|
||||
if task != nil {
|
||||
task = tc.taskService.UpdateTask(task.ID, ¶m)
|
||||
}
|
||||
}
|
||||
|
||||
if task == nil {
|
||||
task = tc.taskService.CreateTask(¶m)
|
||||
}
|
||||
|
||||
// 如果是 Agent 任务,通知 Agent;否则添加到本地 cron
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*task.AgentID)
|
||||
} else {
|
||||
tc.executorService.AddCronTask(task)
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskVO(task))
|
||||
}
|
||||
|
||||
// BulkSaveTask 批量保存/导入任务配置(用于主节点下发同步)
|
||||
// @Summary 批量保存任务
|
||||
// @Description 批量导入任务配置,如果ID或同名存在则更新,不存在则创建
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Router /tasks/bulk_save [post]
|
||||
func (tc *TaskController) BulkSaveTask(c *gin.Context) {
|
||||
var reqs []vo.TaskVO
|
||||
|
||||
if err := c.ShouldBindJSON(&reqs); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
for _, req := range reqs {
|
||||
param := tasks.TaskParam{
|
||||
Name: req.Name,
|
||||
Remark: req.Remark,
|
||||
Command: req.Command,
|
||||
PreCommand: req.PreCommand,
|
||||
PostCommand: req.PostCommand,
|
||||
Tags: req.Tags,
|
||||
Type: req.Type,
|
||||
Config: req.Config,
|
||||
Schedule: req.Schedule,
|
||||
Timeout: req.Timeout,
|
||||
WorkDir: req.WorkDir,
|
||||
CleanConfig: req.CleanConfig,
|
||||
Envs: req.Envs,
|
||||
Languages: req.Languages,
|
||||
AgentID: req.AgentID,
|
||||
TriggerType: req.TriggerType,
|
||||
RetryCount: req.RetryCount,
|
||||
RetryInterval: req.RetryInterval,
|
||||
RandomRange: req.RandomRange,
|
||||
PinType: req.PinType,
|
||||
Enabled: req.Enabled,
|
||||
SourceID: "", // 不直接覆盖
|
||||
}
|
||||
|
||||
var existingTask *models.Task
|
||||
// 优先按 ID 匹配
|
||||
if req.ID != "" {
|
||||
existingTask = tc.taskService.GetTaskByID(req.ID)
|
||||
}
|
||||
// 如果 ID 没找到,尝试按 Name 匹配
|
||||
if existingTask == nil {
|
||||
var t models.Task
|
||||
res := database.DB.Where("name = ?", req.Name).First(&t)
|
||||
if res.Error == nil {
|
||||
existingTask = &t
|
||||
}
|
||||
}
|
||||
|
||||
var savedTask *models.Task
|
||||
if existingTask != nil {
|
||||
savedTask = tc.taskService.UpdateTask(existingTask.ID, ¶m)
|
||||
} else {
|
||||
savedTask = tc.taskService.CreateTask(¶m)
|
||||
// 如果原始有 ID,强制覆盖更新 ID 保持强同步一致性
|
||||
if req.ID != "" && savedTask != nil {
|
||||
database.DB.Model(savedTask).Update("id", req.ID)
|
||||
savedTask.ID = req.ID
|
||||
}
|
||||
}
|
||||
|
||||
// 如果是 Agent 任务,通知 Agent;否则添加到本地 cron
|
||||
if savedTask != nil {
|
||||
if savedTask.AgentID != nil && *savedTask.AgentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*savedTask.AgentID)
|
||||
} else {
|
||||
tc.executorService.AddCronTask(savedTask)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
// GetTasks 获取任务列表
|
||||
// @Summary 获取任务列表
|
||||
// @Description 分页获取任务列表,支持按名称、Agent ID、标签、类型筛选
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param name query string false "任务名称"
|
||||
// @Param agent_id query string false "Agent ID"
|
||||
// @Param tags query string false "标签"
|
||||
// @Param type query string false "任务类型"
|
||||
// @Param page query int false "页码"
|
||||
// @Param page_size query int false "每页数量"
|
||||
// @Success 200 {object} utils.Response{data=utils.PaginationData{data=[]vo.TaskVO}}
|
||||
// @Router /tasks [get]
|
||||
func (tc *TaskController) GetTasks(c *gin.Context) {
|
||||
p := utils.ParsePagination(c)
|
||||
name := c.DefaultQuery("name", "")
|
||||
agentIDStr := c.DefaultQuery("agent_id", "")
|
||||
|
||||
tags := c.DefaultQuery("tags", "")
|
||||
taskType := c.DefaultQuery("type", "")
|
||||
|
||||
var agentID *string
|
||||
if agentIDStr != "" {
|
||||
agentID = &agentIDStr
|
||||
}
|
||||
|
||||
sortBy := c.DefaultQuery("sort_by", "")
|
||||
order := c.DefaultQuery("order", "")
|
||||
|
||||
tasks, total := tc.taskService.GetTasksWithPagination(p.Page, p.PageSize, name, agentID, tags, taskType, sortBy, order)
|
||||
utils.PaginatedResponse(c, vo.ToTaskVOListFromModels(tasks), total, p)
|
||||
}
|
||||
|
||||
// GetTask 获取任务详情
|
||||
// @Summary 获取任务详情
|
||||
// @Description 根据 ID 获取任务详情
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "任务ID"
|
||||
// @Success 200 {object} utils.Response{data=vo.TaskVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /tasks/{id} [get]
|
||||
func (tc *TaskController) GetTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
task := tc.taskService.GetTaskByID(id)
|
||||
if task == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskVO(task))
|
||||
}
|
||||
|
||||
// UpdateTask 更新任务
|
||||
// @Summary 更新任务
|
||||
// @Description 根据 ID 更新任务信息
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "任务ID"
|
||||
// @Param body body vo.TaskUpdateReq true "任务更新信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.TaskVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /tasks/{id} [put]
|
||||
func (tc *TaskController) UpdateTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取旧任务信息(用于判断 agent 变更)
|
||||
oldTask := tc.taskService.GetTaskByID(id)
|
||||
var oldAgentID *string
|
||||
if oldTask != nil {
|
||||
oldAgentID = oldTask.AgentID
|
||||
}
|
||||
|
||||
var req vo.TaskUpdateReq
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.Schedule != "" {
|
||||
if err := tc.executorService.ValidateCron(req.Schedule); err != nil {
|
||||
utils.BadRequest(c, "无效的cron表达式: "+err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 转换为绝对路径(Agent 任务保持原样)
|
||||
workDir := req.WorkDir
|
||||
if req.AgentID == nil || *req.AgentID == "" {
|
||||
workDir = resolveWorkDir(req.WorkDir)
|
||||
}
|
||||
|
||||
var sourceID string
|
||||
if req.Type == constant.TaskTypeRepo && req.Config != "" {
|
||||
var repoCfg struct {
|
||||
SourceURL string `json:"source_url"`
|
||||
Branch string `json:"branch"`
|
||||
RepoDirName string `json:"repo_dir_name"`
|
||||
TargetPath string `json:"target_path"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(req.Config), &repoCfg); err == nil && repoCfg.SourceURL != "" {
|
||||
if repoCfg.RepoDirName != "" {
|
||||
if !isValidDirName(repoCfg.RepoDirName) {
|
||||
utils.BadRequest(c, "自定义目录名只能包含字母、数字、下划线、短划线和点,不能只有点,且不能包含路径逻辑")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 如果配置了自定义名字,使用配置的名字。没有配置的话,使用以前的username_reponame
|
||||
if repoCfg.RepoDirName != "" {
|
||||
sourceID = "repo_" + repoCfg.RepoDirName
|
||||
} else {
|
||||
sourceID = "repo_" + utils.GetRepoIdentifier(repoCfg.SourceURL, repoCfg.Branch)
|
||||
}
|
||||
|
||||
// 验证更新后的 SourceID 是否和别的任务冲突
|
||||
if sourceID != oldTask.SourceID {
|
||||
existingTask := tc.taskService.GetTaskBySourceID(sourceID)
|
||||
if existingTask != nil && existingTask.ID != oldTask.ID {
|
||||
utils.BadRequest(c, "当前任务已存在,请检查或更换仓库目录名称")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 计算新的物理路径
|
||||
newAbsPath := getRepoPhysicalPath(repoCfg.TargetPath, repoCfg.RepoDirName, repoCfg.SourceURL, repoCfg.Branch)
|
||||
|
||||
var oldAbsPath string
|
||||
if oldTask != nil && oldTask.Type == constant.TaskTypeRepo && oldTask.Config != "" {
|
||||
var oldCfg struct {
|
||||
SourceURL string `json:"source_url"`
|
||||
Branch string `json:"branch"`
|
||||
RepoDirName string `json:"repo_dir_name"`
|
||||
TargetPath string `json:"target_path"`
|
||||
}
|
||||
if json.Unmarshal([]byte(oldTask.Config), &oldCfg) == nil {
|
||||
oldAbsPath = getRepoPhysicalPath(oldCfg.TargetPath, oldCfg.RepoDirName, oldCfg.SourceURL, oldCfg.Branch)
|
||||
}
|
||||
}
|
||||
|
||||
// 如果路径发生了改变(或者是个全新计算的路径),并且新路径已存在,则报错拦截
|
||||
if newAbsPath != "" && newAbsPath != oldAbsPath {
|
||||
if info, err := os.Stat(newAbsPath); err == nil && info.IsDir() {
|
||||
utils.BadRequest(c, "目标目录在本地已存在同名文件夹,请更换目录名或清理残留文件")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if oldTask != nil {
|
||||
sourceID = oldTask.SourceID
|
||||
}
|
||||
|
||||
param := tasks.TaskParam{
|
||||
Name: req.Name,
|
||||
Remark: req.Remark,
|
||||
Command: req.Command,
|
||||
PreCommand: req.PreCommand,
|
||||
PostCommand: req.PostCommand,
|
||||
Tags: req.Tags,
|
||||
Type: req.Type,
|
||||
Config: req.Config,
|
||||
Schedule: req.Schedule,
|
||||
Timeout: req.Timeout,
|
||||
WorkDir: workDir,
|
||||
CleanConfig: req.CleanConfig,
|
||||
Envs: req.Envs,
|
||||
Languages: req.Languages,
|
||||
AgentID: req.AgentID,
|
||||
TriggerType: req.TriggerType,
|
||||
RetryCount: req.RetryCount,
|
||||
RetryInterval: req.RetryInterval,
|
||||
RandomRange: req.RandomRange,
|
||||
SourceID: sourceID,
|
||||
PinType: req.PinType,
|
||||
Enabled: req.Enabled,
|
||||
}
|
||||
|
||||
task := tc.taskService.UpdateTask(id, ¶m)
|
||||
if task == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 处理任务调度
|
||||
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 != "" && *oldAgentID != *task.AgentID {
|
||||
tc.agentWSManager.BroadcastTasks(*oldAgentID)
|
||||
}
|
||||
} else {
|
||||
// 本地任务
|
||||
if utils.DerefBool(task.Enabled, true) {
|
||||
tc.executorService.AddCronTask(task)
|
||||
} else {
|
||||
tc.executorService.RemoveCronTask(task.ID)
|
||||
}
|
||||
// 如果之前是 agent 任务,通知旧 agent 移除
|
||||
if oldAgentID != nil && *oldAgentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*oldAgentID)
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskVO(task))
|
||||
}
|
||||
|
||||
// DeleteTask 删除任务
|
||||
// @Summary 删除任务
|
||||
// @Description 根据 ID 删除任务
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "任务ID"
|
||||
// @Success 200 {object} utils.Response
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /tasks/{id} [delete]
|
||||
func (tc *TaskController) DeleteTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取任务信息(用于通知 agent 和物理删除校验)
|
||||
task := tc.taskService.GetTaskByID(id)
|
||||
if task == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
agentID := task.AgentID
|
||||
deleteFiles := c.Query("delete_files") == "true"
|
||||
|
||||
// 如果需要删除物理文件且是仓库任务
|
||||
if deleteFiles && task.Type == constant.TaskTypeRepo {
|
||||
tc.deleteRepoPhysicalFiles(task)
|
||||
}
|
||||
|
||||
tc.executorService.RemoveCronTask(id)
|
||||
tc.executorService.GetScheduler().StopTask(id)
|
||||
|
||||
success := tc.taskService.DeleteTask(id)
|
||||
if !success {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 如果是 agent 任务,通知 agent
|
||||
if agentID != nil && *agentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*agentID)
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
// deleteRepoPhysicalFiles 删除仓库关联的物理文件
|
||||
func (tc *TaskController) deleteRepoPhysicalFiles(task *models.Task) {
|
||||
if task.Type != constant.TaskTypeRepo {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof("[Controller] 开始尝试物理删除任务关联文件: %s", task.Name)
|
||||
var repoCfg models.RepoConfig
|
||||
if err := json.Unmarshal([]byte(task.Config), &repoCfg); err != nil {
|
||||
logger.Errorf("[Controller] 解析任务配置失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
targetPath := repoCfg.TargetPath
|
||||
if targetPath == "" {
|
||||
// 如果 TargetPath 为空,调用系统的计算函数获取默认目录名
|
||||
repoId := utils.GetRepoIdentifier(repoCfg.SourceURL, repoCfg.Branch)
|
||||
if repoId != "" {
|
||||
targetPath = repoId
|
||||
logger.Infof("[Controller] TargetPath 为空,使用计算出的标识符: %s", targetPath)
|
||||
}
|
||||
}
|
||||
|
||||
if targetPath == "" || targetPath == constant.ScriptsDirPlaceholder {
|
||||
logger.Warnf("[Controller] 任务 %s 无法确定有效的物理删除路径,跳过", task.Name)
|
||||
return
|
||||
}
|
||||
|
||||
// 确定绝对路径
|
||||
scriptsDir, _ := filepath.Abs(constant.ScriptsWorkDir)
|
||||
fullPath := targetPath
|
||||
if strings.HasPrefix(targetPath, constant.ScriptsDirPlaceholder) {
|
||||
fullPath = filepath.Join(scriptsDir, strings.TrimPrefix(targetPath, constant.ScriptsDirPlaceholder))
|
||||
} else if !filepath.IsAbs(targetPath) {
|
||||
fullPath = filepath.Join(scriptsDir, targetPath)
|
||||
}
|
||||
|
||||
absTargetPath, _ := filepath.Abs(fullPath)
|
||||
logger.Infof("[Controller] 最终计算的绝对路径: %s, Scripts目录: %s", absTargetPath, scriptsDir)
|
||||
scriptsDir, _ = filepath.Abs(constant.ScriptsWorkDir)
|
||||
|
||||
// 安全检查:使用 Rel 判断路径关系
|
||||
rel, err := filepath.Rel(scriptsDir, absTargetPath)
|
||||
if err != nil {
|
||||
logger.Errorf("[Controller] 计算相对路径失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
// 必须是在 scripts 目录下(不以 .. 开头)且不能是 scripts 目录本身 (.)
|
||||
if rel != "." && !strings.HasPrefix(rel, "..") {
|
||||
err := os.RemoveAll(absTargetPath)
|
||||
if err != nil {
|
||||
logger.Errorf("[Controller] 物理删除文件夹失败: %s, 路径: %s, 错误: %v", task.Name, absTargetPath, err)
|
||||
} else {
|
||||
logger.Infof("[Controller] 已成功物理删除文件夹: %s, 路径: %s", task.Name, absTargetPath)
|
||||
}
|
||||
} else {
|
||||
logger.Warnf("[Controller] 拒绝物理删除安全目录之外的路径: %s", absTargetPath)
|
||||
}
|
||||
}
|
||||
|
||||
func (tc *TaskController) BatchDeleteTasks(c *gin.Context) {
|
||||
var req struct {
|
||||
IDs []string `json:"ids" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 收集涉及到的 AgentID
|
||||
agentIDs := make(map[string]struct{})
|
||||
for _, id := range req.IDs {
|
||||
// 获取任务信息
|
||||
task := tc.taskService.GetTaskByID(id)
|
||||
if task != nil {
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
agentIDs[*task.AgentID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
// 移除 cron 调度
|
||||
tc.executorService.RemoveCronTask(id)
|
||||
tc.executorService.GetScheduler().StopTask(id)
|
||||
}
|
||||
|
||||
// 执行批量删除
|
||||
count := tc.taskService.BatchDeleteTasks(req.IDs)
|
||||
|
||||
// 通知受影响的 Agent
|
||||
for agentID := range agentIDs {
|
||||
tc.agentWSManager.BroadcastTasks(agentID)
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{"count": count})
|
||||
}
|
||||
|
||||
// BatchDeleteByQuery 根据查询条件批量删除任务
|
||||
// @Summary 根据查询条件批量删除任务
|
||||
// @Description 根据查询条件批量删除匹配的所有任务
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param name query string false "任务名称关键词"
|
||||
// @Param tags query string false "标签关键词"
|
||||
// @Param type query string false "任务类型"
|
||||
// @Param agent_id query string false "执行位置(节点ID)"
|
||||
// @Success 200 {object} utils.Response{data=map[string]int}
|
||||
// @Failure 401 {object} utils.Response "未授权"
|
||||
// @Router /tasks/batch-by-query [delete]
|
||||
func (tc *TaskController) BatchDeleteByQuery(c *gin.Context) {
|
||||
name := c.Query("name")
|
||||
agentIDStr := c.Query("agent_id")
|
||||
tags := c.Query("tags")
|
||||
taskType := c.Query("type")
|
||||
|
||||
var agentID *string
|
||||
if agentIDStr != "" {
|
||||
agentID = &agentIDStr
|
||||
}
|
||||
|
||||
tasks, _ := tc.taskService.GetTasksWithPagination(1, 999999, name, agentID, tags, taskType, "", "")
|
||||
if len(tasks) == 0 {
|
||||
utils.Success(c, gin.H{"count": 0})
|
||||
return
|
||||
}
|
||||
|
||||
var ids []string
|
||||
agentIDs := make(map[string]struct{})
|
||||
for _, task := range tasks {
|
||||
ids = append(ids, task.ID)
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
agentIDs[*task.AgentID] = struct{}{}
|
||||
}
|
||||
// 移除 cron 调度
|
||||
tc.executorService.RemoveCronTask(task.ID)
|
||||
tc.executorService.GetScheduler().StopTask(task.ID)
|
||||
}
|
||||
|
||||
// 执行批量删除
|
||||
count := tc.taskService.BatchDeleteTasks(ids)
|
||||
|
||||
// 通知受影响的 Agent
|
||||
for aID := range agentIDs {
|
||||
tc.agentWSManager.BroadcastTasks(aID)
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{"count": count})
|
||||
}
|
||||
|
||||
// StopTask 停止任务
|
||||
// @Summary 停止任务
|
||||
// @Description 根据运行日志 ID 停止正在执行的任务
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param logID path string true "运行日志ID"
|
||||
// @Success 200 {object} utils.Response
|
||||
// @Failure 400 {object} utils.Response
|
||||
// @Router /tasks/stop/{logID} [post]
|
||||
func (tc *TaskController) StopTask(c *gin.Context) {
|
||||
logID := c.Param("logID")
|
||||
if logID == "" {
|
||||
utils.BadRequest(c, "无效的日志ID")
|
||||
return
|
||||
}
|
||||
|
||||
err := tc.executorService.StopTaskExecution(logID)
|
||||
if err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "停止请求已发送")
|
||||
}
|
||||
|
||||
// GetTags 获取所有任务标签
|
||||
// @Summary 获取所有任务标签
|
||||
// @Description 获取系统中所有任务已使用的唯一标签列表
|
||||
// @Tags 任务管理
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Success 200 {object} utils.Response{data=[]string}
|
||||
// @Router /tasks/tags [get]
|
||||
func (tc *TaskController) GetTags(c *gin.Context) {
|
||||
tags, err := tc.taskService.GetAllTags()
|
||||
if err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(c, tags)
|
||||
}
|
||||
|
||||
// SyncRepoTasks 增量同步仓库任务状态(供本地 reposync 进程调用)
|
||||
func (tc *TaskController) SyncRepoTasks(c *gin.Context) {
|
||||
var req struct {
|
||||
RepoID string `json:"repo_id"`
|
||||
UpsertedIDs []string `json:"upserted_ids"`
|
||||
DeletedIDs []string `json:"deleted_ids"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
tc.executorService.SyncRepoTasks(req.UpsertedIDs, req.DeletedIDs)
|
||||
utils.SuccessMsg(c, "增量同步成功")
|
||||
}
|
||||
|
||||
// ToggleTask 切换任务启用/禁用状态
|
||||
func (tc *TaskController) ToggleTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
task := tc.taskService.GetTaskByID(id)
|
||||
if task == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取旧 AgentID
|
||||
var oldAgentID *string
|
||||
oldAgentID = task.AgentID
|
||||
|
||||
// 构造更新参数,仅修改 Enabled
|
||||
param := tasks.TaskParam{
|
||||
Name: task.Name,
|
||||
Remark: task.Remark,
|
||||
Command: string(task.Command),
|
||||
PreCommand: string(task.PreCommand),
|
||||
PostCommand: string(task.PostCommand),
|
||||
Tags: task.Tags,
|
||||
Type: task.Type,
|
||||
Config: string(task.Config),
|
||||
Schedule: task.Schedule,
|
||||
Timeout: task.Timeout,
|
||||
WorkDir: task.WorkDir,
|
||||
CleanConfig: task.CleanConfig,
|
||||
Envs: string(task.Envs),
|
||||
Languages: task.Languages,
|
||||
AgentID: task.AgentID,
|
||||
TriggerType: task.TriggerType,
|
||||
RetryCount: task.RetryCount,
|
||||
RetryInterval: task.RetryInterval,
|
||||
RandomRange: task.RandomRange,
|
||||
SourceID: task.SourceID,
|
||||
PinType: task.PinType,
|
||||
Enabled: req.Enabled,
|
||||
}
|
||||
|
||||
updatedTask := tc.taskService.UpdateTask(id, ¶m)
|
||||
if updatedTask == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 处理调度器更新
|
||||
if updatedTask.AgentID != nil && *updatedTask.AgentID != "" {
|
||||
tc.executorService.RemoveCronTask(updatedTask.ID)
|
||||
tc.agentWSManager.BroadcastTasks(*updatedTask.AgentID)
|
||||
} else {
|
||||
if req.Enabled {
|
||||
tc.executorService.AddCronTask(updatedTask)
|
||||
} else {
|
||||
tc.executorService.RemoveCronTask(updatedTask.ID)
|
||||
}
|
||||
if oldAgentID != nil && *oldAgentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*oldAgentID)
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskVO(updatedTask))
|
||||
}
|
||||
@@ -0,0 +1,476 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/creack/pty"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
"golang.org/x/text/encoding/simplifiedchinese"
|
||||
"golang.org/x/text/transform"
|
||||
)
|
||||
|
||||
type TerminalController struct {
|
||||
envService *services.EnvService
|
||||
}
|
||||
|
||||
func NewTerminalController(envService *services.EnvService) *TerminalController {
|
||||
return &TerminalController{
|
||||
envService: envService,
|
||||
}
|
||||
}
|
||||
|
||||
var upgrader = websocket.Upgrader{
|
||||
CheckOrigin: utils.CheckWSOrigin,
|
||||
}
|
||||
|
||||
// toUTF8 将可能是 GBK 编码的字节转换为 UTF-8
|
||||
func toUTF8(data []byte) string {
|
||||
if utf8.Valid(data) {
|
||||
return string(data)
|
||||
}
|
||||
// 尝试从 GBK 转换
|
||||
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
|
||||
}
|
||||
|
||||
func (tc *TerminalController) HandleWebSocket(c *gin.Context) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// 演示模式下禁用终端
|
||||
if constant.DemoMode {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("\r\n\033[1;33m[演示模式] 终端功能已禁用\033[0m\r\n"))
|
||||
return
|
||||
}
|
||||
|
||||
// Windows 使用 pipe 模式,Unix 使用 PTY 模式
|
||||
userID := c.GetString("userID")
|
||||
if userID == "" {
|
||||
userID = "1" // 兜底
|
||||
}
|
||||
if runtime.GOOS == "windows" {
|
||||
tc.handlePipeMode(conn, userID)
|
||||
} else {
|
||||
tc.handlePtyMode(conn, userID)
|
||||
}
|
||||
}
|
||||
|
||||
// handlePtyMode 使用 PTY 处理终端(Unix/macOS)
|
||||
func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) {
|
||||
conn.SetReadLimit(constant.MaxMessageSize)
|
||||
conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
return nil
|
||||
})
|
||||
|
||||
// 发送 PTY 模式标识
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("__PTY_MODE__"))
|
||||
|
||||
cmd := utils.NewShellCmd()
|
||||
|
||||
if absDir, err := filepath.Abs(constant.ScriptsWorkDir); err == nil {
|
||||
cmd.Dir = absDir
|
||||
}
|
||||
|
||||
cmd.Env = tc.buildTerminalEnv(userID, "TERM=xterm-256color")
|
||||
|
||||
ptmx, err := pty.Start(cmd)
|
||||
if err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error starting shell: "+err.Error()))
|
||||
return
|
||||
}
|
||||
defer ptmx.Close()
|
||||
|
||||
pty.Setsize(ptmx, &pty.Winsize{Rows: 24, Cols: 80})
|
||||
|
||||
var wg sync.WaitGroup
|
||||
var connMu sync.Mutex
|
||||
|
||||
writeMessage := func(data []byte) {
|
||||
connMu.Lock()
|
||||
defer connMu.Unlock()
|
||||
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
conn.WriteMessage(websocket.TextMessage, data)
|
||||
}
|
||||
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
buf := make([]byte, 4096)
|
||||
var remainder []byte
|
||||
for {
|
||||
n, err := ptmx.Read(buf)
|
||||
if n > 0 {
|
||||
chunk := append(remainder, buf[:n]...)
|
||||
lastSafe := len(chunk)
|
||||
for i := len(chunk); i > 0 && i > len(chunk)-4; i-- {
|
||||
if utf8.RuneStart(chunk[i-1]) {
|
||||
if !utf8.FullRune(chunk[i-1 : len(chunk)]) {
|
||||
lastSafe = i - 1
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
safe := chunk[:lastSafe]
|
||||
|
||||
if len(safe) == 0 && len(chunk) >= 4 {
|
||||
safe = chunk
|
||||
remainder = nil
|
||||
} else {
|
||||
remainder = make([]byte, len(chunk[lastSafe:]))
|
||||
copy(remainder, chunk[lastSafe:])
|
||||
}
|
||||
|
||||
if len(safe) > 0 {
|
||||
text := toUTF8(safe)
|
||||
writeMessage([]byte(text))
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
if len(remainder) > 0 {
|
||||
writeMessage([]byte(toUTF8(remainder)))
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// 启动 ping 协程
|
||||
pingDone := make(chan struct{})
|
||||
go func() {
|
||||
ticker := time.NewTicker(constant.PingPeriod)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
connMu.Lock()
|
||||
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||||
connMu.Unlock()
|
||||
return
|
||||
}
|
||||
connMu.Unlock()
|
||||
case <-pingDone:
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
_, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
|
||||
// 处理调整窗口大小的消息
|
||||
if len(message) > 0 && message[0] == '{' {
|
||||
var resizeMsg struct {
|
||||
Type string `json:"type"`
|
||||
Rows uint16 `json:"rows"`
|
||||
Cols uint16 `json:"cols"`
|
||||
}
|
||||
if err := json.Unmarshal(message, &resizeMsg); err == nil && resizeMsg.Type == "resize" {
|
||||
pty.Setsize(ptmx, &pty.Winsize{Rows: resizeMsg.Rows, Cols: resizeMsg.Cols})
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := ptmx.Write(message); err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
close(pingDone)
|
||||
cmd.Process.Kill()
|
||||
cmd.Wait()
|
||||
ptmx.Close() // Force close PTY to interrupt the blocking ptmx.Read() in the goroutine
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// handlePipeMode 使用 pipe 处理终端(Windows)
|
||||
func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string) {
|
||||
conn.SetReadLimit(constant.MaxMessageSize)
|
||||
conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
return nil
|
||||
})
|
||||
|
||||
// 发送 pipe 模式标识
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("__PIPE_MODE__"))
|
||||
|
||||
cmd := utils.NewShellCmd()
|
||||
|
||||
if absDir, err := filepath.Abs(constant.ScriptsWorkDir); err == nil {
|
||||
cmd.Dir = absDir
|
||||
}
|
||||
|
||||
// 注入环境变量
|
||||
cmd.Env = tc.buildTerminalEnv(userID)
|
||||
|
||||
stdin, err := cmd.StdinPipe()
|
||||
if err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
stdout, err := cmd.StdoutPipe()
|
||||
if err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
stderr, err := cmd.StderrPipe()
|
||||
if err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
var connMu sync.Mutex
|
||||
|
||||
writeMessage := func(data []byte) {
|
||||
connMu.Lock()
|
||||
defer connMu.Unlock()
|
||||
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
conn.WriteMessage(websocket.TextMessage, data)
|
||||
}
|
||||
|
||||
readOutput := func(reader io.Reader) {
|
||||
defer wg.Done()
|
||||
defer func() { recover() }()
|
||||
buf := make([]byte, 4096)
|
||||
var remainder []byte
|
||||
for {
|
||||
n, err := reader.Read(buf)
|
||||
if n > 0 {
|
||||
chunk := append(remainder, buf[:n]...)
|
||||
lastSafe := len(chunk)
|
||||
for i := len(chunk); i > 0 && i > len(chunk)-4; i-- {
|
||||
if utf8.RuneStart(chunk[i-1]) {
|
||||
if !utf8.FullRune(chunk[i-1 : len(chunk)]) {
|
||||
lastSafe = i - 1
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
safe := chunk[:lastSafe]
|
||||
|
||||
if len(safe) == 0 && len(chunk) >= 4 {
|
||||
safe = chunk
|
||||
remainder = nil
|
||||
} else {
|
||||
remainder = make([]byte, len(chunk[lastSafe:]))
|
||||
copy(remainder, chunk[lastSafe:])
|
||||
}
|
||||
|
||||
if len(safe) > 0 {
|
||||
text := toUTF8(safe)
|
||||
writeMessage([]byte(text))
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
if len(remainder) > 0 {
|
||||
writeMessage([]byte(toUTF8(remainder)))
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
wg.Add(2)
|
||||
go readOutput(stdout)
|
||||
go readOutput(stderr)
|
||||
|
||||
// 启动 ping 协程
|
||||
pingDone := make(chan struct{})
|
||||
go func() {
|
||||
ticker := time.NewTicker(constant.PingPeriod)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
connMu.Lock()
|
||||
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||||
connMu.Unlock()
|
||||
return
|
||||
}
|
||||
connMu.Unlock()
|
||||
case <-pingDone:
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
_, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
// 过滤调整窗口大小的消息(Windows Pipe 模式不支持调整尺寸,需过滤掉避免写入 stdin)
|
||||
if len(message) > 0 && message[0] == '{' {
|
||||
var resizeMsg struct {
|
||||
Type string `json:"type"`
|
||||
}
|
||||
if err := json.Unmarshal(message, &resizeMsg); err == nil && resizeMsg.Type == "resize" {
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := stdin.Write(message); err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
close(pingDone)
|
||||
cmd.Process.Kill()
|
||||
cmd.Wait()
|
||||
|
||||
if stdinCloser, ok := stdin.(io.Closer); ok {
|
||||
stdinCloser.Close()
|
||||
}
|
||||
if stdoutCloser, ok := stdout.(io.Closer); ok {
|
||||
stdoutCloser.Close()
|
||||
}
|
||||
if stderrCloser, ok := stderr.(io.Closer); ok {
|
||||
stderrCloser.Close()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// ExecuteShellCommand 执行单个命令并返回结果
|
||||
func (tc *TerminalController) ExecuteShellCommand(c *gin.Context) {
|
||||
// 演示模式下禁止执行命令
|
||||
if constant.DemoMode {
|
||||
utils.BadRequest(c, "演示模式下不能执行命令")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Command string `json:"command" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
cmd := utils.NewShellCommandCmd(req.Command)
|
||||
userID := c.GetString("userID")
|
||||
if userID == "" {
|
||||
userID = "1" // 与 WebSocket 终端保持一致,保留原有兜底行为
|
||||
}
|
||||
cmd.Env = tc.buildTerminalEnv(userID)
|
||||
output, err := cmd.CombinedOutput()
|
||||
|
||||
if err != nil {
|
||||
utils.Success(c, gin.H{
|
||||
"output": string(output),
|
||||
"error": err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"output": string(output),
|
||||
})
|
||||
}
|
||||
|
||||
func (tc *TerminalController) buildTerminalEnv(userID string, extraEnvs ...string) []string {
|
||||
env := os.Environ()
|
||||
env = append(env, extraEnvs...)
|
||||
|
||||
// 注入 taskpool 命令运行时路径与配置,保证在终端里手动执行 taskpool 子命令时
|
||||
// 仍然能连接到与主服务一致的数据库,而不是回退到默认 sqlite。
|
||||
if absBinDir, err := filepath.Abs(filepath.Join(constant.DataDir, "bin")); err == nil {
|
||||
pathStr := absBinDir + string(os.PathListSeparator) + os.Getenv("PATH")
|
||||
env = append(env, "PATH="+pathStr)
|
||||
}
|
||||
env = append(env, utils.BuildRuntimeProcessEnv()...)
|
||||
|
||||
// 注入环境变量(支持同名合并)
|
||||
if tc.envService != nil {
|
||||
envVars := tc.envService.GetFormattedEnvVarsByUserID(userID)
|
||||
env = append(env, envVars...)
|
||||
}
|
||||
|
||||
// 为 Docker 环境或二进制版本注入所有 mise 已安装 Node 的全局依赖路径到 NODE_PATH (Issue-90)
|
||||
if !utils.IsInDocker() || (!strings.Contains(os.Args[0], "go-build") && !strings.Contains(os.Args[0], "tmp")) {
|
||||
versions, _ := utils.ListMiseInstalledVersions("node")
|
||||
var nodePaths []string
|
||||
for _, v := range versions {
|
||||
if p := utils.GetMiseNodePath(v); p != "" {
|
||||
nodePaths = append(nodePaths, p)
|
||||
}
|
||||
}
|
||||
if len(nodePaths) > 0 {
|
||||
sep := ":"
|
||||
if runtime.GOOS == "windows" {
|
||||
sep = ";"
|
||||
}
|
||||
env = append(env, "NODE_PATH="+strings.Join(nodePaths, sep))
|
||||
}
|
||||
}
|
||||
|
||||
return env
|
||||
}
|
||||
|
||||
// GetCommands 获取所有可用的 cmd 列表及说明
|
||||
func (tc *TerminalController) GetCommands(c *gin.Context) {
|
||||
var cmds []map[string]string
|
||||
for _, cmdInfo := range constant.Commands {
|
||||
cmds = append(cmds, map[string]string{
|
||||
"name": cmdInfo.Name,
|
||||
"description": cmdInfo.Description,
|
||||
})
|
||||
}
|
||||
utils.Success(c, cmds)
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type WebUIController struct {
|
||||
webuiService *services.WebUIService
|
||||
}
|
||||
|
||||
func NewWebUIController(webuiService *services.WebUIService) *WebUIController {
|
||||
return &WebUIController{
|
||||
webuiService: webuiService,
|
||||
}
|
||||
}
|
||||
|
||||
// GetWebUIs 获取所有WebUI
|
||||
func (c *WebUIController) GetWebUIs(ctx *gin.Context) {
|
||||
webuis, err := c.webuiService.GetWebUIs()
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, webuis)
|
||||
}
|
||||
|
||||
// UploadWebUI 上传新WebUI
|
||||
func (c *WebUIController) UploadWebUI(ctx *gin.Context) {
|
||||
file, err := ctx.FormFile("file")
|
||||
if err != nil {
|
||||
utils.BadRequest(ctx, "获取上传文件失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 临时保存上传的文件到挂载目录,避免 /tmp 跨分区移动或权限问题
|
||||
tmpDir := filepath.Join(constant.DataDir, "tmp")
|
||||
os.MkdirAll(tmpDir, 0755)
|
||||
tmpFile := filepath.Join(tmpDir, file.Filename)
|
||||
|
||||
if err := ctx.SaveUploadedFile(file, tmpFile); err != nil {
|
||||
utils.ServerError(ctx, "保存临时文件失败")
|
||||
return
|
||||
}
|
||||
defer os.Remove(tmpFile) // 自动清理临时文件
|
||||
|
||||
webuiName, err := c.webuiService.ExtractWebUI(tmpFile)
|
||||
if err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"message": "WebUI上传成功", "webui": webuiName})
|
||||
}
|
||||
|
||||
// SetActiveWebUI 切换活动WebUI
|
||||
func (c *WebUIController) SetActiveWebUI(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
}
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "无效的请求参数")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.webuiService.SetActiveWebUI(req.Name); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"message": "WebUI已切换成功,部分页面可能需要刷新"})
|
||||
}
|
||||
|
||||
// DeleteWebUI 删除自定义WebUI
|
||||
func (c *WebUIController) DeleteWebUI(ctx *gin.Context) {
|
||||
name := ctx.Param("name")
|
||||
if name == "" {
|
||||
utils.BadRequest(ctx, "未提供WebUI名称")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.webuiService.DeleteWebUI(name); err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"message": "WebUI已删除"})
|
||||
}
|
||||
Reference in New Issue
Block a user