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:
2026-07-26 08:43:52 +08:00
commit e6956aa001
397 changed files with 73621 additions and 0 deletions
+746
View File
@@ -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, "清理成功")
}
+199
View File
@@ -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)
}
+81
View File
@@ -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,
})
}
+389
View File
@@ -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))
}
+550
View File
@@ -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
}
}
+110
View File
@@ -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
}
+169
View File
@@ -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
}
})
}
+165
View File
@@ -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)
}
+134
View File
@@ -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)
}
+155
View File
@@ -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, "删除成功")
}
+624
View File
@@ -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
}
}
}
}
+872
View File
@@ -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, &param)
}
}
if task == nil {
task = tc.taskService.CreateTask(&param)
}
// 如果是 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, &param)
} else {
savedTask = tc.taskService.CreateTask(&param)
// 如果原始有 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, &param)
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, &param)
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))
}
+476
View File
@@ -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)
}
+93
View File
@@ -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已删除"})
}